first pass at turn based conversation
This commit is contained in:
parent
d1118d375e
commit
518a0f2b53
40 changed files with 503 additions and 99 deletions
10
vocode/streaming/user_implemented_agent/base_agent.py
Normal file
10
vocode/streaming/user_implemented_agent/base_agent.py
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
from fastapi import FastAPI
|
||||
import uvicorn
|
||||
|
||||
|
||||
class BaseAgent:
|
||||
def __init__(self):
|
||||
self.app = FastAPI()
|
||||
|
||||
def run(self, host="localhost", port=3000):
|
||||
uvicorn.run(self.app, host=host, port=port)
|
||||
19
vocode/streaming/user_implemented_agent/restful_agent.py
Normal file
19
vocode/streaming/user_implemented_agent/restful_agent.py
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
from .base_agent import BaseAgent
|
||||
from ..models.agent import RESTfulAgentInput, RESTfulAgentOutput, RESTfulAgentText, RESTfulAgentEnd
|
||||
from pydantic import BaseModel
|
||||
from typing import Union
|
||||
from fastapi import APIRouter
|
||||
|
||||
class RESTfulAgent(BaseAgent):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.app.post("/respond")(self.respond_rest)
|
||||
|
||||
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, request.conversation_id)
|
||||
return response
|
||||
|
||||
56
vocode/streaming/user_implemented_agent/websocket_agent.py
Normal file
56
vocode/streaming/user_implemented_agent/websocket_agent.py
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
from .base_agent import BaseAgent
|
||||
import uuid
|
||||
import typing
|
||||
from typing import AsyncGenerator, Union, Optional
|
||||
from fastapi import WebSocket
|
||||
from ..models.agent import (
|
||||
WebSocketAgentStartMessage,
|
||||
WebSocketAgentReadyMessage,
|
||||
WebSocketAgentTextEndMessage,
|
||||
WebSocketAgentTextMessage,
|
||||
WebSocketAgentStopMessage,
|
||||
WebSocketAgentMessage,
|
||||
WebSocketAgentMessageType,
|
||||
)
|
||||
|
||||
|
||||
class WebSocketAgent(BaseAgent):
|
||||
def __init__(self, generate_responses: bool = False):
|
||||
super().__init__()
|
||||
self.generate_responses = generate_responses
|
||||
self.app.websocket("/respond")(self.respond_websocket)
|
||||
|
||||
async def respond(
|
||||
self, human_input: str, conversation_id: Optional[str] = None
|
||||
) -> Union[WebSocketAgentTextMessage, WebSocketAgentStopMessage]:
|
||||
raise NotImplementedError
|
||||
|
||||
async def generate_response(
|
||||
self, human_input: str, conversation_id: Optional[str] = None
|
||||
) -> AsyncGenerator[
|
||||
Union[WebSocketAgentTextMessage, WebSocketAgentTextEndMessage], None
|
||||
]:
|
||||
raise NotImplementedError
|
||||
|
||||
async def respond_websocket(self, websocket: WebSocket):
|
||||
await websocket.accept()
|
||||
WebSocketAgentStartMessage.parse_obj(await websocket.receive_json())
|
||||
await websocket.send_text(WebSocketAgentReadyMessage().json())
|
||||
while True:
|
||||
input_message: WebSocketAgentMessage = WebSocketAgentMessage.parse_obj(
|
||||
await websocket.receive_json()
|
||||
)
|
||||
if input_message.type == WebSocketAgentMessageType.STOP:
|
||||
break
|
||||
text_message = typing.cast(WebSocketAgentTextMessage, input_message)
|
||||
if self.generate_responses:
|
||||
async for output_response in self.generate_response(
|
||||
text_message.data.text, text_message.conversation_id
|
||||
):
|
||||
await websocket.send_text(output_response.json())
|
||||
else:
|
||||
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