conversation ID tracking
This commit is contained in:
parent
2b3dc08f99
commit
3ad8a109e7
9 changed files with 24 additions and 13 deletions
|
|
@ -10,10 +10,10 @@ class RESTfulAgent(BaseAgent):
|
|||
super().__init__()
|
||||
self.app.post("/respond")(self.respond_rest)
|
||||
|
||||
async def respond(self, human_input) -> RESTfulAgentOutput:
|
||||
async def respond(self, human_input, conversation_id) -> RESTfulAgentOutput:
|
||||
raise NotImplementedError
|
||||
|
||||
async def respond_rest(self, request: RESTfulAgentInput) -> Union[RESTfulAgentText, RESTfulAgentEnd]:
|
||||
response = await self.respond(request.human_input)
|
||||
response = await self.respond(request.human_input, request.conversation_id)
|
||||
return response
|
||||
|
||||
|
|
|
|||
|
|
@ -23,7 +23,6 @@ class WebSocketAgent(BaseAgent):
|
|||
|
||||
async def respond_websocket(self, websocket: WebSocket):
|
||||
await websocket.accept()
|
||||
conversation_id = str(uuid.uuid4())
|
||||
WebSocketAgentStartMessage.parse_obj(await websocket.receive_json())
|
||||
await websocket.send_text(WebSocketAgentReadyMessage().json())
|
||||
while True:
|
||||
|
|
@ -31,7 +30,8 @@ class WebSocketAgent(BaseAgent):
|
|||
if input_message.type == WebSocketAgentMessageType.STOP:
|
||||
break
|
||||
text_message = typing.cast(WebSocketAgentTextMessage, input_message)
|
||||
output_response = await self.respond(text_message.data.text, conversation_id=conversation_id)
|
||||
print(text_message)
|
||||
output_response = await self.respond(text_message.data.text, text_message.conversation_id)
|
||||
await websocket.send_text(output_response.json())
|
||||
await websocket.close()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue