processors(rtvi): messages now require an id

This commit is contained in:
Aleix Conchillo Flaqué
2024-07-21 16:54:40 -07:00
parent 37b04ed283
commit 302ea90dce

View File

@@ -83,35 +83,47 @@ class RTVIMessageData(BaseModel):
class RTVIMessage(BaseModel): class RTVIMessage(BaseModel):
label: Literal["realtime-ai"] = "realtime-ai" label: Literal["realtime-ai"] = "realtime-ai"
type: str type: str
id: str
data: Optional[RTVIMessageData] = None data: Optional[RTVIMessageData] = None
class RTVIResponseData(BaseModel): class RTVIResponseData(BaseModel):
type: str
success: bool success: bool
error: Optional[str] = None error: Optional[str] = None
class RTVIResponse(BaseModel): class RTVIResponse(BaseModel):
label: Literal["realtime-ai"] = "realtime-ai" label: Literal["realtime-ai"] = "realtime-ai"
type: Literal["response"] = "response" type: Literal["response"] = "response"
id: str
data: RTVIResponseData data: RTVIResponseData
class RTVIErrorData(BaseModel):
message: str
class RTVIError(BaseModel):
label: Literal["realtime-ai"] = "realtime-ai"
type: Literal["error"] = "error"
data: RTVIErrorData
class RTVILLMContextMessageData(BaseModel): class RTVILLMContextMessageData(BaseModel):
messages: List[dict] messages: List[dict]
class RTVIBotReady(BaseModel):
label: Literal["realtime-ai"] = "realtime-ai"
type: Literal["bot-ready"] = "bot-ready"
class RTVILLMContextMessage(BaseModel): class RTVILLMContextMessage(BaseModel):
label: Literal["realtime-ai"] = "realtime-ai" label: Literal["realtime-ai"] = "realtime-ai"
type: Literal["llm-context"] = "llm-context" type: Literal["llm-context"] = "llm-context"
data: RTVILLMContextMessageData data: RTVILLMContextMessageData
class RTVIBotReady(BaseModel):
label: Literal["realtime-ai"] = "realtime-ai"
type: Literal["bot-ready"] = "bot-ready"
class RTVITranscriptionMessageData(BaseModel): class RTVITranscriptionMessageData(BaseModel):
text: str text: str
user_id: str user_id: str
@@ -178,7 +190,10 @@ class RTVIProcessor(FrameProcessor):
if isinstance(frame, StartFrame): if isinstance(frame, StartFrame):
self._start_frame = frame self._start_frame = frame
try:
await self._handle_setup(self._setup) await self._handle_setup(self._setup)
except Exception as e:
await self._send_error(f"unable to setup RTVI: {e}")
async def cleanup(self): async def cleanup(self):
self._frame_handler_task.cancel() self._frame_handler_task.cancel()
@@ -239,16 +254,18 @@ class RTVIProcessor(FrameProcessor):
try: try:
message = RTVIMessage.model_validate(frame.message) message = RTVIMessage.model_validate(frame.message)
except ValidationError as e: except ValidationError as e:
await self._send_response("setup", False, f"invalid message: {e}") await self._send_error(f"invalid message: {e}")
return return
try: try:
success = True
error = None
match message.type: match message.type:
case "setup": case "setup":
setup = None setup = None
if message.data: if message.data:
setup = message.data.setup setup = message.data.setup
await self._handle_setup(setup) await self._handle_setup(message.id, setup)
case "config-update": case "config-update":
await self._handle_config_update(message.data.config) await self._handle_config_update(message.data.config)
case "llm-get-context": case "llm-get-context":
@@ -261,14 +278,17 @@ class RTVIProcessor(FrameProcessor):
await self._handle_tts_speak(message.data.tts) await self._handle_tts_speak(message.data.tts)
case "tts-interrupt": case "tts-interrupt":
await self._handle_tts_interrupt() await self._handle_tts_interrupt()
case _:
success = False
error = f"unsupported type {message.type}"
# Send a message to indicate we successfully executed the command. await self._send_response(message.id, success, error)
await self._send_response(message.type, True)
except ValidationError as e: except ValidationError as e:
await self._send_response(message.type, False, f"invalid message: {e}") await self._send_response(message.id, False, f"invalid message: {e}")
except Exception as e:
await self._send_response(message.id, False, f"{e}")
async def _handle_setup(self, setup: RTVISetup | None): async def _handle_setup(self, setup: RTVISetup | None):
try:
model = DEFAULT_MODEL model = DEFAULT_MODEL
if setup and setup.config and setup.config.llm and setup.config.llm.model: if setup and setup.config and setup.config.llm and setup.config.llm.model:
model = setup.config.llm.model model = setup.config.llm.model
@@ -312,8 +332,6 @@ class RTVIProcessor(FrameProcessor):
message = RTVIBotReady() message = RTVIBotReady()
frame = TransportMessageFrame(message=message.model_dump(exclude_none=True)) frame = TransportMessageFrame(message=message.model_dump(exclude_none=True))
await self.push_frame(frame) await self.push_frame(frame)
except Exception as e:
await self._send_response("setup", False, f"unable to create pipeline: {e}")
async def _handle_config_update(self, config: RTVIConfig): async def _handle_config_update(self, config: RTVIConfig):
if config.llm and config.llm.model: if config.llm and config.llm.model:
@@ -352,7 +370,12 @@ class RTVIProcessor(FrameProcessor):
async def _handle_tts_interrupt(self): async def _handle_tts_interrupt(self):
await self.push_frame(BotInterruptionFrame(), FrameDirection.UPSTREAM) await self.push_frame(BotInterruptionFrame(), FrameDirection.UPSTREAM)
async def _send_response(self, type: str, success: bool, error: str | None = None): async def _send_error(self, error: str):
message = RTVIError(data=RTVIErrorData(message=error))
frame = TransportMessageFrame(message=message.model_dump(exclude_none=True))
await self.push_frame(frame)
async def _send_response(self, id: str, success: bool, error: str | None = None):
# TODO(aleix): This is a bit hacky, but we might get invalid # TODO(aleix): This is a bit hacky, but we might get invalid
# configuration or something might going wrong during setup and we would # configuration or something might going wrong during setup and we would
# like to send the error to the client. However, if the pipeline is not # like to send the error to the client. However, if the pipeline is not
@@ -367,6 +390,6 @@ class RTVIProcessor(FrameProcessor):
if parent and self._start_frame: if parent and self._start_frame:
parent.link(pipeline) parent.link(pipeline)
message = RTVIResponse(data=RTVIResponseData(type=type, success=success, error=error)) message = RTVIResponse(id=id, data=RTVIResponseData(success=success, error=error))
frame = TransportMessageFrame(message=message.model_dump(exclude_none=True)) frame = TransportMessageFrame(message=message.model_dump(exclude_none=True))
await self.push_frame(frame) await self.push_frame(frame)