Update GladiaSTTService to use WebsocketSTTService
This commit is contained in:
@@ -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
|
|
||||||
|
|||||||
Reference in New Issue
Block a user