Refactored DeepgramFluxSTTService to automatically reconnect if sending a message fails.

This commit is contained in:
Filipi Fuchter
2025-11-17 09:54:20 -03:00
parent 77cd106795
commit 19cc0177b8

View File

@@ -6,7 +6,9 @@
"""Deepgram Flux speech-to-text service implementation.""" """Deepgram Flux speech-to-text service implementation."""
import asyncio
import json import json
import time
from enum import Enum from enum import Enum
from typing import Any, AsyncGenerator, Dict, Optional from typing import Any, AsyncGenerator, Dict, Optional
from urllib.parse import urlencode from urllib.parse import urlencode
@@ -94,6 +96,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
mip_opt_out: Optional. Opts out requests from the Deepgram Model Improvement Program mip_opt_out: Optional. Opts out requests from the Deepgram Model Improvement Program
(default False). (default False).
tag: List of tags to label requests for identification during usage reporting. tag: List of tags to label requests for identification during usage reporting.
min_confidence: Optional. Minimum confidence required confidence to create a TranscriptionFrame
""" """
eager_eot_threshold: Optional[float] = None eager_eot_threshold: Optional[float] = None
@@ -102,6 +105,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
keyterm: list = [] keyterm: list = []
mip_opt_out: Optional[bool] = None mip_opt_out: Optional[bool] = None
tag: list = [] tag: list = []
min_confidence: Optional[float] = None # New parameter
def __init__( def __init__(
self, self,
@@ -163,6 +167,13 @@ class DeepgramFluxSTTService(WebsocketSTTService):
self._register_event_handler("on_end_of_turn") self._register_event_handler("on_end_of_turn")
self._register_event_handler("on_eager_end_of_turn") self._register_event_handler("on_eager_end_of_turn")
self._register_event_handler("on_update") self._register_event_handler("on_update")
self._connection_established_event = asyncio.Event()
# Watchdog task to prevent dangling tasks
# If we stop sending audio to Flux after we have received that the User has started speaking
# we never receive the user stopped speaking event unless we resume sending audio to it.
self._last_stt_time = None
self._watchdog_task = None
self._user_is_speaking = False
async def _connect(self): async def _connect(self):
"""Connect to WebSocket and start background tasks. """Connect to WebSocket and start background tasks.
@@ -172,9 +183,6 @@ class DeepgramFluxSTTService(WebsocketSTTService):
""" """
await self._connect_websocket() await self._connect_websocket()
if self._websocket and not self._receive_task:
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
async def _disconnect(self): async def _disconnect(self):
"""Disconnect from WebSocket and clean up tasks. """Disconnect from WebSocket and clean up tasks.
@@ -182,14 +190,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
and cleans up resources to prevent memory leaks. and cleans up resources to prevent memory leaks.
""" """
try: try:
# Cancel background tasks BEFORE closing websocket
if self._receive_task:
await self.cancel_task(self._receive_task, timeout=2.0)
self._receive_task = None
# Now close the websocket
await self._disconnect_websocket() await self._disconnect_websocket()
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") logger.error(f"{self} exception: {e}")
await self.push_error(ErrorFrame(error=f"{self} error: {e}")) await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
@@ -197,6 +198,25 @@ class DeepgramFluxSTTService(WebsocketSTTService):
# Reset state only after everything is cleaned up # Reset state only after everything is cleaned up
self._websocket = None self._websocket = None
async def _send_silence(self, duration_secs: float = 0.5):
"""Send a block of silence of the specified duration (default 500 ms)."""
sample_width = 2 # bytes per sample for 16-bit PCM
num_channels = 1 # mono
num_samples = int(self.sample_rate * duration_secs)
silence = b"\x00" * (num_samples * sample_width * num_channels)
await self._websocket.send(silence)
async def _watchdog_task_handler(self):
while self._websocket and self._websocket.state is State.OPEN:
now = time.monotonic()
# More than 500 ms without sending new audio to Flux
if self._user_is_speaking and self._last_stt_time and now - self._last_stt_time > 0.5:
logger.warning("Sending silence to Flux to prevent dangling task")
await self._send_silence()
self._last_stt_time = time.monotonic()
# check every 100ms
await asyncio.sleep(0.1)
async def _connect_websocket(self): async def _connect_websocket(self):
"""Establish WebSocket connection to API. """Establish WebSocket connection to API.
@@ -208,10 +228,26 @@ class DeepgramFluxSTTService(WebsocketSTTService):
if self._websocket and self._websocket.state is State.OPEN: if self._websocket and self._websocket.state is State.OPEN:
return return
self._connection_established_event.clear()
self._user_is_speaking = False
self._websocket = await websocket_connect( self._websocket = await websocket_connect(
self._websocket_url, self._websocket_url,
additional_headers={"Authorization": f"Token {self._api_key}"}, additional_headers={"Authorization": f"Token {self._api_key}"},
) )
# Creating the receiver task
if not self._receive_task:
self._receive_task = self.create_task(
self._receive_task_handler(self._report_error)
)
# Creating the watchdog task
if not self._watchdog_task:
self._watchdog_task = self.create_task(self._watchdog_task_handler())
# Now wait for the connection established event
logger.debug("WebSocket connected, waiting for server confirmation...")
await self._connection_established_event.wait()
logger.debug("Connected to Deepgram Flux Websocket") logger.debug("Connected to Deepgram Flux Websocket")
await self._call_event_handler("on_connected") await self._call_event_handler("on_connected")
except Exception as e: except Exception as e:
@@ -227,6 +263,16 @@ class DeepgramFluxSTTService(WebsocketSTTService):
metrics collection. Handles disconnection errors gracefully. metrics collection. Handles disconnection errors gracefully.
""" """
try: try:
# Cancel background tasks BEFORE closing websocket
if self._receive_task:
await self.cancel_task(self._receive_task, timeout=2.0)
self._receive_task = None
if self._watchdog_task:
await self.cancel_task(self._watchdog_task, timeout=2.0)
self._watchdog_task = None
self._last_stt_time = None
self._connection_established_event.clear()
await self.stop_all_metrics() await self.stop_all_metrics()
if self._websocket: if self._websocket:
@@ -340,7 +386,8 @@ class DeepgramFluxSTTService(WebsocketSTTService):
return return
try: try:
await self._websocket.send(audio) self._last_stt_time = time.monotonic()
await self.send_with_retry(audio, self._report_error)
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") logger.error(f"{self} exception: {e}")
yield ErrorFrame(error=f"{self} error: {e}") yield ErrorFrame(error=f"{self} error: {e}")
@@ -463,6 +510,8 @@ class DeepgramFluxSTTService(WebsocketSTTService):
transcription processing. transcription processing.
""" """
logger.info("Connected to Flux - ready to stream audio") logger.info("Connected to Flux - ready to stream audio")
# Notify connection is established
self._connection_established_event.set()
async def _handle_fatal_error(self, data: Dict[str, Any]): async def _handle_fatal_error(self, data: Dict[str, Any]):
"""Handle fatal error messages from Deepgram Flux. """Handle fatal error messages from Deepgram Flux.
@@ -530,6 +579,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
transcript: maybe the first few words of the turn. transcript: maybe the first few words of the turn.
""" """
logger.debug("User started speaking") logger.debug("User started speaking")
self._user_is_speaking = True
await self.push_interruption_task_frame_and_wait() await self.push_interruption_task_frame_and_wait()
await self.broadcast_frame(UserStartedSpeakingFrame) await self.broadcast_frame(UserStartedSpeakingFrame)
await self.start_metrics() await self.start_metrics()
@@ -550,6 +600,22 @@ class DeepgramFluxSTTService(WebsocketSTTService):
logger.trace(f"Received event TurnResumed: {event}") logger.trace(f"Received event TurnResumed: {event}")
await self._call_event_handler("on_turn_resumed") await self._call_event_handler("on_turn_resumed")
def _calculate_average_confidence(self, transcript_data) -> Optional[float]:
"""Calculate the average confidence from transcript data.
Return None if the data is missing or invalid.
"""
# Example: Assume transcript_data has a list of words with confidence
words = transcript_data.get("words")
if not words or not isinstance(words, list):
return None
confidences = [
w.get("confidence") for w in words if isinstance(w.get("confidence"), (float, int))
]
if not confidences:
return None
return sum(confidences) / len(confidences)
async def _handle_end_of_turn(self, transcript: str, data: Dict[str, Any]): async def _handle_end_of_turn(self, transcript: str, data: Dict[str, Any]):
"""Handle EndOfTurn events from Deepgram Flux. """Handle EndOfTurn events from Deepgram Flux.
@@ -569,16 +635,26 @@ class DeepgramFluxSTTService(WebsocketSTTService):
data: The TurnInfo message data containing event type, transcript and some extra metadata. data: The TurnInfo message data containing event type, transcript and some extra metadata.
""" """
logger.debug("User stopped speaking") logger.debug("User stopped speaking")
self._user_is_speaking = False
await self.push_frame( # Compute the average confidence
TranscriptionFrame( average_confidence = self._calculate_average_confidence(data)
transcript,
self._user_id, if not self._params.min_confidence or average_confidence > self._params.min_confidence:
time_now_iso8601(), await self.push_frame(
self._language, TranscriptionFrame(
result=data, transcript,
self._user_id,
time_now_iso8601(),
self._language,
result=data,
)
) )
) else:
logger.warning(
f"Transcription confidence below min_confidence threshold: {average_confidence}"
)
await self._handle_transcription(transcript, True, self._language) await self._handle_transcription(transcript, True, self._language)
await self.stop_processing_metrics() await self.stop_processing_metrics()
await self.push_frame(UserStoppedSpeakingFrame(), FrameDirection.DOWNSTREAM) await self.push_frame(UserStoppedSpeakingFrame(), FrameDirection.DOWNSTREAM)