processors(rtvi): messages now require an id
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user