conversation ID tracking

This commit is contained in:
Ajay Raj 2023-03-08 15:33:59 -08:00
commit 3ad8a109e7
9 changed files with 24 additions and 13 deletions

View file

@ -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

View file

@ -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()