Update GladiaSTTService to use WebsocketSTTService

This commit is contained in:
Mark Backman
2025-12-16 09:20:53 -05:00
parent 88d909d468
commit d3f918eb58

View File

@@ -23,7 +23,6 @@ from pipecat import __version__ as pipecat_version
from pipecat.frames.frames import ( from pipecat.frames.frames import (
CancelFrame, CancelFrame,
EndFrame, EndFrame,
ErrorFrame,
Frame, Frame,
InterimTranscriptionFrame, InterimTranscriptionFrame,
StartFrame, StartFrame,
@@ -31,7 +30,7 @@ from pipecat.frames.frames import (
TranslationFrame, TranslationFrame,
) )
from pipecat.services.gladia.config import GladiaInputParams from pipecat.services.gladia.config import GladiaInputParams
from pipecat.services.stt_service import STTService from pipecat.services.stt_service import WebsocketSTTService
from pipecat.transcriptions.language import Language, resolve_language from pipecat.transcriptions.language import Language, resolve_language
from pipecat.utils.time import time_now_iso8601 from pipecat.utils.time import time_now_iso8601
from pipecat.utils.tracing.service_decorators import traced_stt from pipecat.utils.tracing.service_decorators import traced_stt
@@ -176,7 +175,7 @@ class _InputParamsDescriptor:
return GladiaInputParams return GladiaInputParams
class GladiaSTTService(STTService): class GladiaSTTService(WebsocketSTTService):
"""Speech-to-Text service using Gladia's API. """Speech-to-Text service using Gladia's API.
This service connects to Gladia's WebSocket API for real-time transcription This service connects to Gladia's WebSocket API for real-time transcription
@@ -202,8 +201,6 @@ class GladiaSTTService(STTService):
sample_rate: Optional[int] = None, sample_rate: Optional[int] = None,
model: str = "solaria-1", model: str = "solaria-1",
params: Optional[GladiaInputParams] = None, params: Optional[GladiaInputParams] = None,
max_reconnection_attempts: int = 5,
reconnection_delay: float = 1.0,
max_buffer_size: int = 1024 * 1024 * 20, # 20MB default buffer max_buffer_size: int = 1024 * 1024 * 20, # 20MB default buffer
**kwargs, **kwargs,
): ):
@@ -222,8 +219,6 @@ class GladiaSTTService(STTService):
sample_rate: Audio sample rate in Hz. If None, uses service default. sample_rate: Audio sample rate in Hz. If None, uses service default.
model: Model to use for transcription. Defaults to "solaria-1". model: Model to use for transcription. Defaults to "solaria-1".
params: Additional configuration parameters for Gladia service. params: Additional configuration parameters for Gladia service.
max_reconnection_attempts: Maximum number of reconnection attempts. Defaults to 5.
reconnection_delay: Initial delay between reconnection attempts in seconds.
max_buffer_size: Maximum size of audio buffer in bytes. Defaults to 20MB. max_buffer_size: Maximum size of audio buffer in bytes. Defaults to 20MB.
**kwargs: Additional arguments passed to the STTService parent class. **kwargs: Additional arguments passed to the STTService parent class.
""" """
@@ -256,15 +251,11 @@ class GladiaSTTService(STTService):
self._url = url self._url = url
self.set_model_name(model) self.set_model_name(model)
self._params = params self._params = params
self._websocket = None
self._receive_task = None self._receive_task = None
self._keepalive_task = None self._keepalive_task = None
self._settings = {} self._settings = {}
# Reconnection settings # Session management
self._max_reconnection_attempts = max_reconnection_attempts
self._reconnection_delay = reconnection_delay
self._reconnection_attempts = 0
self._session_url = None self._session_url = None
self._connection_active = False self._connection_active = False
@@ -274,10 +265,6 @@ class GladiaSTTService(STTService):
self._max_buffer_size = max_buffer_size self._max_buffer_size = max_buffer_size
self._buffer_lock = asyncio.Lock() self._buffer_lock = asyncio.Lock()
# Connection management
self._connection_task = None
self._should_reconnect = True
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
"""Check if the service can generate performance metrics. """Check if the service can generate performance metrics.
@@ -355,11 +342,7 @@ class GladiaSTTService(STTService):
frame: The start frame triggering service startup. frame: The start frame triggering service startup.
""" """
await super().start(frame) await super().start(frame)
if self._connection_task: await self._connect()
return
self._should_reconnect = True
self._connection_task = self.create_task(self._connection_handler())
async def stop(self, frame: EndFrame): async def stop(self, frame: EndFrame):
"""Stop the Gladia STT websocket connection. """Stop the Gladia STT websocket connection.
@@ -368,14 +351,8 @@ class GladiaSTTService(STTService):
frame: The end frame triggering service shutdown. frame: The end frame triggering service shutdown.
""" """
await super().stop(frame) await super().stop(frame)
self._should_reconnect = False
await self._send_stop_recording() await self._send_stop_recording()
await self._disconnect()
if self._connection_task:
await self.cancel_task(self._connection_task)
self._connection_task = None
await self._cleanup_connection()
async def cancel(self, frame: CancelFrame): async def cancel(self, frame: CancelFrame):
"""Cancel the Gladia STT websocket connection. """Cancel the Gladia STT websocket connection.
@@ -384,13 +361,7 @@ class GladiaSTTService(STTService):
frame: The cancel frame triggering service cancellation. frame: The cancel frame triggering service cancellation.
""" """
await super().cancel(frame) await super().cancel(frame)
self._should_reconnect = False await self._disconnect()
if self._connection_task:
await self.cancel_task(self._connection_task)
self._connection_task = None
await self._cleanup_connection()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
"""Run speech-to-text on audio data. """Run speech-to-text on audio data.
@@ -424,76 +395,75 @@ class GladiaSTTService(STTService):
yield None yield None
async def _connection_handler(self): async def _connect(self):
"""Handle WebSocket connection with automatic reconnection.""" """Connect to the Gladia service.
while self._should_reconnect:
try:
# Initialize session if needed
if not self._session_url:
settings = self._prepare_settings()
response = await self._setup_gladia(settings)
self._session_url = response["url"]
self._reconnection_attempts = 0
logger.info(f"Session URL : {self._session_url}")
# Connect with automatic reconnection Initializes the session if needed and establishes websocket connection.
async with websocket_connect(self._session_url) as websocket: """
try: # Initialize session if needed
self._websocket = websocket if not self._session_url:
self._connection_active = True settings = self._prepare_settings()
logger.debug(f"{self} Connected to Gladia WebSocket") response = await self._setup_gladia(settings)
self._session_url = response["url"]
logger.info(f"Session URL : {self._session_url}")
# Send buffered audio if any await self._connect_websocket()
await self._send_buffered_audio()
# Start tasks if self._websocket and not self._receive_task:
self._receive_task = self.create_task(self._receive_task_handler()) self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
self._keepalive_task = self.create_task(self._keepalive_task_handler())
# Wait for tasks to complete if self._websocket and not self._keepalive_task:
await asyncio.gather(self._receive_task, self._keepalive_task) self._keepalive_task = self.create_task(self._keepalive_task_handler())
except websockets.exceptions.ConnectionClosed as e: async def _disconnect(self):
logger.warning(f"WebSocket connection closed: {e}") """Disconnect from the Gladia service.
self._connection_active = False
# Clean up tasks Cleans up tasks and closes websocket connection.
if self._receive_task: """
await self.cancel_task(self._receive_task)
if self._keepalive_task:
await self.cancel_task(self._keepalive_task)
# Attempt reconnect using helper
if not await self._maybe_reconnect():
break
except Exception as e:
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)
self._connection_active = False
if not self._should_reconnect:
break
# Reset session URL to get a new one
self._session_url = None
await asyncio.sleep(self._reconnection_delay)
async def _cleanup_connection(self):
"""Clean up connection resources."""
self._connection_active = False self._connection_active = False
if self._keepalive_task: if self._keepalive_task:
await self.cancel_task(self._keepalive_task) await self.cancel_task(self._keepalive_task)
self._keepalive_task = None self._keepalive_task = None
if self._websocket:
await self._websocket.close()
self._websocket = None
if self._receive_task: if self._receive_task:
await self.cancel_task(self._receive_task) await self.cancel_task(self._receive_task)
self._receive_task = None self._receive_task = None
await self._disconnect_websocket()
async def _connect_websocket(self):
"""Establish the websocket connection to Gladia."""
try:
if self._websocket and self._websocket.state is State.OPEN:
return
logger.debug("Connecting to Gladia WebSocket")
self._websocket = await websocket_connect(self._session_url)
self._connection_active = True
await self._call_event_handler("on_connected")
# Send buffered audio if any
await self._send_buffered_audio()
logger.debug(f"{self} Connected to Gladia WebSocket")
except Exception as e:
await self.push_error(error_msg=f"Unable to connect to Gladia: {e}", exception=e)
raise
async def _disconnect_websocket(self):
"""Close the websocket connection to Gladia."""
try:
if self._websocket and self._websocket.state is State.OPEN:
logger.debug("Disconnecting from Gladia WebSocket")
await self._websocket.close()
except Exception as e:
await self.push_error(error_msg=f"Error closing websocket: {e}", exception=e)
finally:
self._websocket = None
await self._call_event_handler("on_disconnected")
async def _setup_gladia(self, settings: Dict[str, Any]): async def _setup_gladia(self, settings: Dict[str, Any]):
async with aiohttp.ClientSession() as session: async with aiohttp.ClientSession() as session:
params = {} params = {}
@@ -541,28 +511,26 @@ class GladiaSTTService(STTService):
if self._websocket and self._websocket.state is State.OPEN: if self._websocket and self._websocket.state is State.OPEN:
await self._websocket.send(json.dumps({"type": "stop_recording"})) await self._websocket.send(json.dumps({"type": "stop_recording"}))
async def _keepalive_task_handler(self): def _get_websocket(self):
"""Send periodic empty audio chunks to keep the connection alive.""" """Get the current WebSocket connection.
try:
KEEPALIVE_SLEEP = 20
while self._connection_active:
# Send keepalive (Gladia times out after 30 seconds)
await asyncio.sleep(KEEPALIVE_SLEEP)
if self._websocket and self._websocket.state is State.OPEN:
# Send an empty audio chunk as keepalive
empty_audio = b""
await self._send_audio(empty_audio)
else:
logger.debug("Websocket closed, stopping keepalive")
break
except websockets.exceptions.ConnectionClosed:
logger.debug("Connection closed during keepalive")
except Exception as e:
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)
async def _receive_task_handler(self): Returns:
try: The WebSocket connection.
async for message in self._websocket:
Raises:
Exception: If WebSocket is not connected.
"""
if self._websocket:
return self._websocket
raise Exception("Websocket not connected")
async def _receive_messages(self):
"""Receive and process websocket messages.
Continuously processes messages from the websocket connection.
"""
async for message in self._get_websocket():
try:
content = json.loads(message) content = json.loads(message)
# Handle audio chunk acknowledgments # Handle audio chunk acknowledgments
@@ -617,26 +585,24 @@ class GladiaSTTService(STTService):
translation, "", time_now_iso8601(), translated_language translation, "", time_now_iso8601(), translated_language
) )
) )
except json.JSONDecodeError:
logger.warning(f"Received non-JSON message: {message}")
async def _keepalive_task_handler(self):
"""Send periodic empty audio chunks to keep the connection alive."""
try:
KEEPALIVE_SLEEP = 20
while self._connection_active:
# Send keepalive (Gladia times out after 30 seconds)
await asyncio.sleep(KEEPALIVE_SLEEP)
if self._websocket and self._websocket.state is State.OPEN:
# Send an empty audio chunk as keepalive
empty_audio = b""
await self._send_audio(empty_audio)
else:
logger.debug("Websocket closed, stopping keepalive")
break
except websockets.exceptions.ConnectionClosed: except websockets.exceptions.ConnectionClosed:
# Expected when closing the connection logger.debug("Connection closed during keepalive")
pass
except Exception as e: except Exception as e:
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e) await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)
async def _maybe_reconnect(self) -> bool:
"""Handle exponential backoff reconnection logic."""
if not self._should_reconnect:
return False
self._reconnection_attempts += 1
if self._reconnection_attempts > self._max_reconnection_attempts:
await self.push_error(
error_msg=f"Max reconnection attempts ({self._max_reconnection_attempts}) reached",
)
self._should_reconnect = False
return False
delay = self._reconnection_delay * (2 ** (self._reconnection_attempts - 1))
logger.debug(
f"{self} Reconnecting in {delay} seconds (attempt {self._reconnection_attempts}/{self._max_reconnection_attempts})"
)
await asyncio.sleep(delay)
return True