add goodbye option
This commit is contained in:
parent
385d1386af
commit
9a0677e88d
3 changed files with 74 additions and 28 deletions
|
|
@ -4,8 +4,20 @@ import signal
|
||||||
|
|
||||||
from vocode.conversation import Conversation
|
from vocode.conversation import Conversation
|
||||||
from vocode.helpers import create_microphone_input_and_speaker_output
|
from vocode.helpers import create_microphone_input_and_speaker_output
|
||||||
from vocode.models.transcriber import DeepgramTranscriberConfig, PunctuationEndpointingConfig, GoogleTranscriberConfig
|
from vocode.models.transcriber import (
|
||||||
from vocode.models.agent import ChatGPTAgentConfig, RESTfulUserImplementedAgentConfig, WebSocketUserImplementedAgentConfig, EchoAgentConfig, ChatGPTAlphaAgentConfig, ChatGPTAgentConfig
|
DeepgramTranscriberConfig,
|
||||||
|
PunctuationEndpointingConfig,
|
||||||
|
GoogleTranscriberConfig,
|
||||||
|
)
|
||||||
|
from vocode.models.agent import (
|
||||||
|
ChatGPTAgentConfig,
|
||||||
|
RESTfulUserImplementedAgentConfig,
|
||||||
|
WebSocketUserImplementedAgentConfig,
|
||||||
|
EchoAgentConfig,
|
||||||
|
ChatGPTAlphaAgentConfig,
|
||||||
|
LLMAgentConfig,
|
||||||
|
ChatGPTAgentConfig,
|
||||||
|
)
|
||||||
from vocode.models.synthesizer import AzureSynthesizerConfig
|
from vocode.models.synthesizer import AzureSynthesizerConfig
|
||||||
from vocode.user_implemented_agent.restful_agent import RESTfulAgent
|
from vocode.user_implemented_agent.restful_agent import RESTfulAgent
|
||||||
|
|
||||||
|
|
@ -14,25 +26,30 @@ logging.root.setLevel(logging.INFO)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
microphone_input, speaker_output = create_microphone_input_and_speaker_output(use_default_devices=False)
|
microphone_input, speaker_output = create_microphone_input_and_speaker_output(
|
||||||
|
use_default_devices=False
|
||||||
|
)
|
||||||
|
|
||||||
conversation = Conversation(
|
conversation = Conversation(
|
||||||
input_device=microphone_input,
|
input_device=microphone_input,
|
||||||
output_device=speaker_output,
|
output_device=speaker_output,
|
||||||
transcriber_config=DeepgramTranscriberConfig.from_input_device(
|
transcriber_config=DeepgramTranscriberConfig.from_input_device(
|
||||||
microphone_input,
|
microphone_input, endpointing_config=PunctuationEndpointingConfig()
|
||||||
endpointing_config=PunctuationEndpointingConfig()
|
|
||||||
),
|
),
|
||||||
agent_config=WebSocketUserImplementedAgentConfig(
|
# agent_config=WebSocketUserImplementedAgentConfig(
|
||||||
initial_message="Hello!",
|
# initial_message="Hello!",
|
||||||
respond=WebSocketUserImplementedAgentConfig.RouteConfig(
|
# respond=WebSocketUserImplementedAgentConfig.RouteConfig(
|
||||||
url="wss://8b7425d5b2ab.ngrok.io/respond",
|
# url="wss://8b7425d5b2ab.ngrok.io/respond",
|
||||||
)
|
# )
|
||||||
|
# ),
|
||||||
|
# id="ajay",
|
||||||
|
agent_config=ChatGPTAgentConfig(
|
||||||
|
initial_message="goodbye",
|
||||||
|
prompt_preamble="you are an expert on the NBA",
|
||||||
|
generate_responses=True,
|
||||||
|
end_conversation_on_goodbye=True,
|
||||||
),
|
),
|
||||||
id="ajay",
|
synthesizer_config=AzureSynthesizerConfig.from_output_device(speaker_output),
|
||||||
# agent_config=ChatGPTAgentConfig(initial_message="hello", prompt_preamble="you are an expert on the NBA"),
|
|
||||||
synthesizer_config=AzureSynthesizerConfig.from_output_device(speaker_output)
|
|
||||||
)
|
)
|
||||||
signal.signal(signal.SIGINT, lambda _0, _1: conversation.deactivate())
|
signal.signal(signal.SIGINT, lambda _0, _1: conversation.deactivate())
|
||||||
asyncio.run(conversation.start())
|
asyncio.run(conversation.start())
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -18,20 +18,25 @@ class AgentConfig(TypedModel, type=AgentType.BASE):
|
||||||
initial_message: Optional[str] = None
|
initial_message: Optional[str] = None
|
||||||
generate_responses: bool = True
|
generate_responses: bool = True
|
||||||
allowed_idle_time_seconds: Optional[float] = None
|
allowed_idle_time_seconds: Optional[float] = None
|
||||||
|
end_conversation_on_goodbye: bool = False
|
||||||
|
|
||||||
|
|
||||||
class LLMAgentConfig(AgentConfig, type=AgentType.LLM):
|
class LLMAgentConfig(AgentConfig, type=AgentType.LLM):
|
||||||
prompt_preamble: str
|
prompt_preamble: str
|
||||||
expected_first_prompt: Optional[str] = None
|
expected_first_prompt: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
class ChatGPTAlphaAgentConfig(AgentConfig, type=AgentType.CHAT_GPT_ALPHA):
|
class ChatGPTAlphaAgentConfig(AgentConfig, type=AgentType.CHAT_GPT_ALPHA):
|
||||||
prompt_preamble: str
|
prompt_preamble: str
|
||||||
expected_first_prompt: Optional[str] = None
|
expected_first_prompt: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
class ChatGPTAgentConfig(AgentConfig, type=AgentType.CHAT_GPT):
|
class ChatGPTAgentConfig(AgentConfig, type=AgentType.CHAT_GPT):
|
||||||
prompt_preamble: str
|
prompt_preamble: str
|
||||||
expected_first_prompt: Optional[str] = None
|
expected_first_prompt: Optional[str] = None
|
||||||
generate_responses: bool = False
|
generate_responses: bool = False
|
||||||
|
|
||||||
|
|
||||||
class InformationRetrievalAgentConfig(
|
class InformationRetrievalAgentConfig(
|
||||||
AgentConfig, type=AgentType.INFORMATION_RETRIEVAL
|
AgentConfig, type=AgentType.INFORMATION_RETRIEVAL
|
||||||
):
|
):
|
||||||
|
|
@ -45,7 +50,10 @@ class InformationRetrievalAgentConfig(
|
||||||
class EchoAgentConfig(AgentConfig, type=AgentType.ECHO):
|
class EchoAgentConfig(AgentConfig, type=AgentType.ECHO):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class RESTfulUserImplementedAgentConfig(AgentConfig, type=AgentType.RESTFUL_USER_IMPLEMENTED):
|
|
||||||
|
class RESTfulUserImplementedAgentConfig(
|
||||||
|
AgentConfig, type=AgentType.RESTFUL_USER_IMPLEMENTED
|
||||||
|
):
|
||||||
class EndpointConfig(BaseModel):
|
class EndpointConfig(BaseModel):
|
||||||
url: str
|
url: str
|
||||||
method: str = "POST"
|
method: str = "POST"
|
||||||
|
|
@ -55,25 +63,33 @@ class RESTfulUserImplementedAgentConfig(AgentConfig, type=AgentType.RESTFUL_USER
|
||||||
# generate_response: Optional[EndpointConfig]
|
# generate_response: Optional[EndpointConfig]
|
||||||
# update_last_bot_message_on_cut_off: Optional[EndpointConfig]
|
# update_last_bot_message_on_cut_off: Optional[EndpointConfig]
|
||||||
|
|
||||||
|
|
||||||
class RESTfulAgentInput(BaseModel):
|
class RESTfulAgentInput(BaseModel):
|
||||||
conversation_id: str
|
conversation_id: str
|
||||||
human_input: str
|
human_input: str
|
||||||
|
|
||||||
|
|
||||||
class RESTfulAgentOutputType(str, Enum):
|
class RESTfulAgentOutputType(str, Enum):
|
||||||
BASE = "restful_agent_base"
|
BASE = "restful_agent_base"
|
||||||
TEXT = "restful_agent_text"
|
TEXT = "restful_agent_text"
|
||||||
END = "restful_agent_end"
|
END = "restful_agent_end"
|
||||||
|
|
||||||
|
|
||||||
class RESTfulAgentOutput(TypedModel, type=RESTfulAgentOutputType.BASE):
|
class RESTfulAgentOutput(TypedModel, type=RESTfulAgentOutputType.BASE):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
class RESTfulAgentText(RESTfulAgentOutput, type=RESTfulAgentOutputType.TEXT):
|
class RESTfulAgentText(RESTfulAgentOutput, type=RESTfulAgentOutputType.TEXT):
|
||||||
response: str
|
response: str
|
||||||
|
|
||||||
|
|
||||||
class RESTfulAgentEnd(RESTfulAgentOutput, type=RESTfulAgentOutputType.END):
|
class RESTfulAgentEnd(RESTfulAgentOutput, type=RESTfulAgentOutputType.END):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class WebSocketUserImplementedAgentConfig(AgentConfig, type=AgentType.WEBSOCKET_USER_IMPLEMENTED):
|
|
||||||
|
class WebSocketUserImplementedAgentConfig(
|
||||||
|
AgentConfig, type=AgentType.WEBSOCKET_USER_IMPLEMENTED
|
||||||
|
):
|
||||||
class RouteConfig(BaseModel):
|
class RouteConfig(BaseModel):
|
||||||
url: str
|
url: str
|
||||||
|
|
||||||
|
|
@ -82,17 +98,22 @@ class WebSocketUserImplementedAgentConfig(AgentConfig, type=AgentType.WEBSOCKET_
|
||||||
# generate_response: Optional[RouteConfig]
|
# generate_response: Optional[RouteConfig]
|
||||||
# send_message_on_cut_off: bool = False
|
# send_message_on_cut_off: bool = False
|
||||||
|
|
||||||
|
|
||||||
class WebSocketAgentMessageType(str, Enum):
|
class WebSocketAgentMessageType(str, Enum):
|
||||||
BASE = 'websocket_agent_base'
|
BASE = "websocket_agent_base"
|
||||||
START = 'websocket_agent_start'
|
START = "websocket_agent_start"
|
||||||
TEXT = 'websocket_agent_text'
|
TEXT = "websocket_agent_text"
|
||||||
READY = 'websocket_agent_ready'
|
READY = "websocket_agent_ready"
|
||||||
STOP = 'websocket_agent_stop'
|
STOP = "websocket_agent_stop"
|
||||||
|
|
||||||
|
|
||||||
class WebSocketAgentMessage(TypedModel, type=WebSocketAgentMessageType.BASE):
|
class WebSocketAgentMessage(TypedModel, type=WebSocketAgentMessageType.BASE):
|
||||||
conversation_id: Optional[str] = None
|
conversation_id: Optional[str] = None
|
||||||
|
|
||||||
class WebSocketAgentTextMessage(WebSocketAgentMessage, type=WebSocketAgentMessageType.TEXT):
|
|
||||||
|
class WebSocketAgentTextMessage(
|
||||||
|
WebSocketAgentMessage, type=WebSocketAgentMessageType.TEXT
|
||||||
|
):
|
||||||
class Payload(BaseModel):
|
class Payload(BaseModel):
|
||||||
text: str
|
text: str
|
||||||
|
|
||||||
|
|
@ -103,11 +124,19 @@ class WebSocketAgentTextMessage(WebSocketAgentMessage, type=WebSocketAgentMessag
|
||||||
return cls(data=cls.Payload(text=text), conversation_id=conversation_id)
|
return cls(data=cls.Payload(text=text), conversation_id=conversation_id)
|
||||||
|
|
||||||
|
|
||||||
class WebSocketAgentStartMessage(WebSocketAgentMessage, type=WebSocketAgentMessageType.START):
|
class WebSocketAgentStartMessage(
|
||||||
|
WebSocketAgentMessage, type=WebSocketAgentMessageType.START
|
||||||
|
):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class WebSocketAgentReadyMessage(WebSocketAgentMessage, type=WebSocketAgentMessageType.READY):
|
|
||||||
|
class WebSocketAgentReadyMessage(
|
||||||
|
WebSocketAgentMessage, type=WebSocketAgentMessageType.READY
|
||||||
|
):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class WebSocketAgentStopMessage(WebSocketAgentMessage, type=WebSocketAgentMessageType.STOP):
|
|
||||||
|
class WebSocketAgentStopMessage(
|
||||||
|
WebSocketAgentMessage, type=WebSocketAgentMessageType.STOP
|
||||||
|
):
|
||||||
pass
|
pass
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,10 @@
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
import uvicorn
|
import uvicorn
|
||||||
|
|
||||||
class BaseAgent():
|
|
||||||
|
|
||||||
|
class BaseAgent:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.app = FastAPI()
|
self.app = FastAPI()
|
||||||
|
|
||||||
def run(self, host="localhost", port=3000):
|
def run(self, host="localhost", port=3000):
|
||||||
uvicorn.run(self.app, host=host, port=port)
|
uvicorn.run(self.app, host=host, port=port)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue