Refactoring the services to use push_error and push_error_frame

This commit is contained in:
Filipi Fuchter
2025-11-18 18:22:45 -03:00
parent 79f43ece74
commit 50bef86d33
14 changed files with 28 additions and 58 deletions

View File

@@ -126,6 +126,4 @@ class WakeCheckFilter(FrameProcessor):
else: else:
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
except Exception as e: except Exception as e:
error_msg = f"Error in wake word filter: {e}" await self.push_error(error_msg=f"Error in wake word filter: {e}", exception=e)
logger.exception(error_msg)
await self.push_error(ErrorFrame(error_msg))

View File

@@ -166,6 +166,6 @@ class AIService(FrameProcessor):
async for f in generator: async for f in generator:
if f: if f:
if isinstance(f, ErrorFrame): if isinstance(f, ErrorFrame):
await self.push_error(f) await self.push_error_frame(f)
else: else:
await self.push_frame(f) await self.push_frame(f)

View File

@@ -458,8 +458,7 @@ class AnthropicLLMService(LLMService):
except httpx.TimeoutException: except httpx.TimeoutException:
await self._call_event_handler("on_completion_timeout") await self._call_event_handler("on_completion_timeout")
except Exception as e: except Exception as e:
logger.exception(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(f"{e}"))
finally: finally:
await self.stop_processing_metrics() await self.stop_processing_metrics()
await self.push_frame(LLMFullResponseEndFrame()) await self.push_frame(LLMFullResponseEndFrame())

View File

@@ -206,9 +206,8 @@ class AssemblyAISTTService(STTService):
await self._call_event_handler("on_connected") await self._call_event_handler("on_connected")
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}")
self._connected = False self._connected = False
await self.push_error(ErrorFrame(error=f"{self} error: {e}")) await self.push_error(exception=e)
raise raise
async def _disconnect(self): async def _disconnect(self):
@@ -233,8 +232,7 @@ class AssemblyAISTTService(STTService):
logger.warning("Timed out waiting for termination message from server") logger.warning("Timed out waiting for termination message from server")
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
if self._receive_task: if self._receive_task:
await self.cancel_task(self._receive_task) await self.cancel_task(self._receive_task)
@@ -262,13 +260,11 @@ class AssemblyAISTTService(STTService):
except websockets.exceptions.ConnectionClosedOK: except websockets.exceptions.ConnectionClosedOK:
break break
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
break break
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
def _parse_message(self, message: Dict[str, Any]) -> BaseMessage: def _parse_message(self, message: Dict[str, Any]) -> BaseMessage:
"""Parse a raw message into the appropriate message type.""" """Parse a raw message into the appropriate message type."""

View File

@@ -228,8 +228,7 @@ class AsyncAITTSService(InterruptibleTTSService):
await self._call_event_handler("on_connected") await self._call_event_handler("on_connected")
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
self._websocket = None self._websocket = None
await self._call_event_handler("on_connection_error", f"{e}") await self._call_event_handler("on_connection_error", f"{e}")
@@ -290,7 +289,7 @@ class AsyncAITTSService(InterruptibleTTSService):
logger.error(f"{self} error: {msg}") logger.error(f"{self} error: {msg}")
await self.push_frame(TTSStoppedFrame()) await self.push_frame(TTSStoppedFrame())
await self.stop_all_metrics() await self.stop_all_metrics()
await self.push_error(ErrorFrame(error=f"{self} error: {msg['message']}")) await self.push_error(error_msg=f"{self} error: {msg['message']}", exception=e)
else: else:
logger.error(f"{self} error, unknown message type: {msg}") logger.error(f"{self} error, unknown message type: {msg}")
@@ -494,8 +493,7 @@ class AsyncAIHttpTTSService(TTSService):
yield frame yield frame
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
finally: finally:
await self.stop_ttfb_metrics() await self.stop_ttfb_metrics()
yield TTSStoppedFrame() yield TTSStoppedFrame()

View File

@@ -140,8 +140,7 @@ class AWSTranscribeSTTService(STTService):
return return
logger.warning("WebSocket connection not established after connect") logger.warning("WebSocket connection not established after connect")
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
retry_count += 1 retry_count += 1
if retry_count < max_retries: if retry_count < max_retries:
await asyncio.sleep(1) # Wait before retrying await asyncio.sleep(1) # Wait before retrying
@@ -310,8 +309,7 @@ class AWSTranscribeSTTService(STTService):
await self._ws_client.send(json.dumps(end_stream)) await self._ws_client.send(json.dumps(end_stream))
await self._ws_client.close() await self._ws_client.close()
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
finally: finally:
self._ws_client = None self._ws_client = None
await self._call_event_handler("on_disconnected") await self._call_event_handler("on_disconnected")

View File

@@ -151,8 +151,7 @@ class AzureSTTService(STTService):
self._speech_recognizer.recognized.connect(self._on_handle_recognized) self._speech_recognizer.recognized.connect(self._on_handle_recognized)
self._speech_recognizer.start_continuous_recognition_async() self._speech_recognizer.start_continuous_recognition_async()
except Exception as e: except Exception as e:
logger.error(f"{self} exception during initialization: {e}") await self.push_error(error_msg=f"{self} exception during initialization: {e}", exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
async def stop(self, frame: EndFrame): async def stop(self, frame: EndFrame):
"""Stop the speech recognition service. """Stop the speech recognition service.

View File

@@ -276,8 +276,7 @@ class CartesiaSTTService(WebsocketSTTService):
self._websocket = await websocket_connect(ws_url, additional_headers=headers) self._websocket = await websocket_connect(ws_url, additional_headers=headers)
await self._call_event_handler("on_connected") await self._call_event_handler("on_connected")
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
async def _disconnect_websocket(self): async def _disconnect_websocket(self):
try: try:
@@ -319,8 +318,7 @@ class CartesiaSTTService(WebsocketSTTService):
elif data["type"] == "error": elif data["type"] == "error":
error_msg = data.get("message", "Unknown error") error_msg = data.get("message", "Unknown error")
logger.error(f"Cartesia error: {error_msg}") await self.push_error(error_msg=error_msg, exception=e)
await self.push_error(ErrorFrame(error=error_msg))
@traced_stt @traced_stt
async def _handle_transcription( async def _handle_transcription(

View File

@@ -397,8 +397,7 @@ class CartesiaTTSService(AudioContextWordTTSService):
) )
await self._call_event_handler("on_connected") await self._call_event_handler("on_connected")
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
self._websocket = None self._websocket = None
await self._call_event_handler("on_connection_error", f"{e}") await self._call_event_handler("on_connection_error", f"{e}")
@@ -464,10 +463,9 @@ class CartesiaTTSService(AudioContextWordTTSService):
) )
await self.append_to_audio_context(msg["context_id"], frame) await self.append_to_audio_context(msg["context_id"], frame)
elif msg["type"] == "error": elif msg["type"] == "error":
logger.error(f"{self} error: {msg}")
await self.push_frame(TTSStoppedFrame()) await self.push_frame(TTSStoppedFrame())
await self.stop_all_metrics() await self.stop_all_metrics()
await self.push_error(ErrorFrame(error=f"{self} error: {msg['error']}")) await self.push_error(error_msg=f"{self} error: {msg}", exception=e)
self._context_id = None self._context_id = None
else: else:
logger.error(f"{self} error, unknown message type: {msg}") logger.error(f"{self} error, unknown message type: {msg}")
@@ -708,8 +706,7 @@ class CartesiaHttpTTSService(TTSService):
async with session.post(url, json=payload, headers=headers) as response: async with session.post(url, json=payload, headers=headers) as response:
if response.status != 200: if response.status != 200:
error_text = await response.text() error_text = await response.text()
logger.error(f"Cartesia API error: {error_text}") await self.push_error(error_msg=f"Cartesia API error: {error_text}")
await self.push_error(ErrorFrame(error=f"Cartesia API error: {error_text}"))
raise Exception(f"Cartesia API returned status {response.status}: {error_text}") raise Exception(f"Cartesia API returned status {response.status}: {error_text}")
audio_data = await response.read() audio_data = await response.read()
@@ -725,8 +722,7 @@ class CartesiaHttpTTSService(TTSService):
yield frame yield frame
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
finally: finally:
await self.stop_ttfb_metrics() await self.stop_ttfb_metrics()
yield TTSStoppedFrame() yield TTSStoppedFrame()

View File

@@ -192,8 +192,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
try: try:
await self._disconnect_websocket() await self._disconnect_websocket()
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
finally: finally:
# Reset state only after everything is cleaned up # Reset state only after everything is cleaned up
self._websocket = None self._websocket = None
@@ -280,8 +279,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
logger.debug("Disconnecting from Deepgram Flux Websocket") logger.debug("Disconnecting from Deepgram Flux Websocket")
await self._websocket.close() await self._websocket.close()
except Exception as e: except Exception as e:
logger.error(f"{self} error closing websocket: {e}") await self.push_error(error_msg=f"{self} error closing websocket: {e}", exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
finally: finally:
self._websocket = None self._websocket = None
await self._call_event_handler("on_disconnected") await self._call_event_handler("on_disconnected")
@@ -467,8 +465,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
# Skip malformed messages # Skip malformed messages
continue continue
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
# Error will be handled inside WebsocketService->_receive_task_handler # Error will be handled inside WebsocketService->_receive_task_handler
raise raise
else: else:

View File

@@ -256,7 +256,7 @@ class DeepgramSTTService(STTService):
async def _on_error(self, *args, **kwargs): async def _on_error(self, *args, **kwargs):
error: ErrorResponse = kwargs["error"] error: ErrorResponse = kwargs["error"]
logger.warning(f"{self} connection error, will retry: {error}") logger.warning(f"{self} connection error, will retry: {error}")
await self.push_error(ErrorFrame(error=f"{error}")) await self.push_error(error_msg=f"{error}")
await self.stop_all_metrics() await self.stop_all_metrics()
# NOTE(aleix): we don't disconnect (i.e. call finish on the connection) # NOTE(aleix): we don't disconnect (i.e. call finish on the connection)
# because this triggers more errors internally in the Deepgram SDK. So, # because this triggers more errors internally in the Deepgram SDK. So,

View File

@@ -645,8 +645,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
await self._call_event_handler("on_connected") await self._call_event_handler("on_connected")
logger.debug("Connected to ElevenLabs Realtime STT") logger.debug("Connected to ElevenLabs Realtime STT")
except Exception as e: except Exception as e:
logger.error(f"{self}: unable to connect to ElevenLabs Realtime STT: {e}") await self.push_error(error_msg=f"{self}: unable to connect to ElevenLabs Realtime STT: {e}", exception=e)
await self.push_error(ErrorFrame(f"Connection error: {str(e)}"))
async def _disconnect_websocket(self): async def _disconnect_websocket(self):
"""Disconnect from ElevenLabs Realtime STT WebSocket.""" """Disconnect from ElevenLabs Realtime STT WebSocket."""
@@ -714,13 +713,11 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
elif message_type == "input_error": elif message_type == "input_error":
error_msg = data.get("error", "Unknown input error") error_msg = data.get("error", "Unknown input error")
logger.error(f"ElevenLabs input error: {error_msg}") await self.push_error(error_msg=f"ElevenLabs input error: {error_msg}")
await self.push_error(ErrorFrame(f"Input error: {error_msg}"))
elif message_type in ["auth_error", "quota_exceeded", "transcriber_error", "error"]: elif message_type in ["auth_error", "quota_exceeded", "transcriber_error", "error"]:
error_msg = data.get("error", data.get("message", "Unknown error")) error_msg = data.get("error", data.get("message", "Unknown error"))
logger.error(f"ElevenLabs error ({message_type}): {error_msg}") await self.push_error(error_msg=f"ElevenLabs error ({message_type}): {error_msg}")
await self.push_error(ErrorFrame(f"{message_type}: {error_msg}"))
else: else:
logger.debug(f"Unknown message type: {message_type}") logger.debug(f"Unknown message type: {message_type}")

View File

@@ -242,8 +242,7 @@ class FishAudioTTSService(InterruptibleTTSService):
await self._websocket.send(ormsgpack.packb(stop_message)) await self._websocket.send(ormsgpack.packb(stop_message))
await self._websocket.close() await self._websocket.close()
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") await self.push_error(exception=e)
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
finally: finally:
self._request_id = None self._request_id = None
self._started = False self._started = False

View File

@@ -141,13 +141,8 @@ class XTTSService(TTSService):
async with self._aiohttp_session.get(self._settings["base_url"] + "/studio_speakers") as r: async with self._aiohttp_session.get(self._settings["base_url"] + "/studio_speakers") as r:
if r.status != 200: if r.status != 200:
text = await r.text() text = await r.text()
logger.error(
f"{self} error getting studio speakers (status: {r.status}, error: {text})"
)
await self.push_error( await self.push_error(
ErrorFrame( error_msg=f"{self} error getting studio speakers (status: {r.status}, error: {text})"
error=f"Error getting studio speakers (status: {r.status}, error: {text})"
)
) )
return return
self._studio_speakers = await r.json() self._studio_speakers = await r.json()