Merge pull request #3038 from pipecat-ai/filipi/flux_improvements
Deepgram Flux improvements
This commit is contained in:
12
CHANGELOG.md
12
CHANGELOG.md
@@ -9,6 +9,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
||||||
|
- Added a watchdog to `DeepgramFluxSTTService` to prevent dangling tasks in case the
|
||||||
|
user was speaking and we stop receiving audio.
|
||||||
|
|
||||||
|
- Introduced a minimum confidence parameter in `DeepgramFluxSTTService` to avoid
|
||||||
|
generating transcriptions below a defined threshold.
|
||||||
|
|
||||||
- Added `ElevenLabsRealtimeSTTService` which implements the Realtime STT
|
- Added `ElevenLabsRealtimeSTTService` which implements the Realtime STT
|
||||||
service from ElevenLabs.
|
service from ElevenLabs.
|
||||||
|
|
||||||
@@ -18,6 +24,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Changed
|
### Changed
|
||||||
|
|
||||||
|
- Extracted the logic for retrying connections, and create a new `send_with_retry`
|
||||||
|
method inside `WebSocketService`.
|
||||||
|
|
||||||
|
- Refactored `DeepgramFluxSTTService` to automatically reconnect if sending a
|
||||||
|
message fails.
|
||||||
|
|
||||||
- Updated all STT and TTS services to use consistent error handling pattern with
|
- Updated all STT and TTS services to use consistent error handling pattern with
|
||||||
`push_error()` method for better pipeline error event integration.
|
`push_error()` method for better pipeline error event integration.
|
||||||
|
|
||||||
|
|||||||
@@ -52,7 +52,10 @@ transport_params = {
|
|||||||
async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
||||||
logger.info(f"Starting bot")
|
logger.info(f"Starting bot")
|
||||||
|
|
||||||
stt = DeepgramFluxSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
stt = DeepgramFluxSTTService(
|
||||||
|
api_key=os.getenv("DEEPGRAM_API_KEY"),
|
||||||
|
params=DeepgramFluxSTTService.InputParams(min_confidence=0.3),
|
||||||
|
)
|
||||||
|
|
||||||
tts = DeepgramTTSService(api_key=os.getenv("DEEPGRAM_API_KEY"), voice="aura-2-andromeda-en")
|
tts = DeepgramTTSService(api_key=os.getenv("DEEPGRAM_API_KEY"), voice="aura-2-andromeda-en")
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ class WebsocketService(ABC):
|
|||||||
"""
|
"""
|
||||||
self._websocket: Optional[websockets.WebSocketClientProtocol] = None
|
self._websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||||||
self._reconnect_on_error = reconnect_on_error
|
self._reconnect_on_error = reconnect_on_error
|
||||||
|
self._reconnect_in_progress: bool = False # Add this flag
|
||||||
|
|
||||||
async def _verify_connection(self) -> bool:
|
async def _verify_connection(self) -> bool:
|
||||||
"""Verify the websocket connection is active and responsive.
|
"""Verify the websocket connection is active and responsive.
|
||||||
@@ -66,6 +67,59 @@ class WebsocketService(ABC):
|
|||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
return await self._verify_connection()
|
return await self._verify_connection()
|
||||||
|
|
||||||
|
async def _try_reconnect(
|
||||||
|
self,
|
||||||
|
max_retries: int = 3,
|
||||||
|
report_error: Optional[Callable[[ErrorFrame], Awaitable[None]]] = None,
|
||||||
|
) -> bool:
|
||||||
|
# Prevent concurrent reconnection attempts
|
||||||
|
if self._reconnect_in_progress:
|
||||||
|
logger.warning(f"{self} reconnect attempt aborted: already in progress")
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._reconnect_in_progress = True
|
||||||
|
last_exception: Optional[Exception] = None
|
||||||
|
try:
|
||||||
|
for attempt in range(1, max_retries + 1):
|
||||||
|
try:
|
||||||
|
logger.warning(f"{self} reconnecting, attempt {attempt}")
|
||||||
|
if await self._reconnect_websocket(attempt):
|
||||||
|
logger.info(f"{self} reconnected successfully on attempt {attempt}")
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
last_exception = e
|
||||||
|
logger.error(f"{self} reconnection attempt {attempt} failed: {e}")
|
||||||
|
if report_error:
|
||||||
|
await report_error(
|
||||||
|
ErrorFrame(f"{self} reconnection attempt {attempt} failed: {e}")
|
||||||
|
)
|
||||||
|
wait_time = exponential_backoff_time(attempt)
|
||||||
|
await asyncio.sleep(wait_time)
|
||||||
|
fatal_msg = f"{self} failed to reconnect after {max_retries} attempts"
|
||||||
|
if last_exception:
|
||||||
|
fatal_msg += f": {last_exception}"
|
||||||
|
logger.error(fatal_msg)
|
||||||
|
if report_error:
|
||||||
|
await report_error(ErrorFrame(fatal_msg, fatal=True))
|
||||||
|
return False
|
||||||
|
finally:
|
||||||
|
self._reconnect_in_progress = False
|
||||||
|
|
||||||
|
async def send_with_retry(self, message, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
||||||
|
"""Attempt to send a message, retrying after reconnect if necessary."""
|
||||||
|
try:
|
||||||
|
await self._websocket.send(message)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"{self} send failed: {e}, will try to reconnect")
|
||||||
|
# Try to reconnect before retrying
|
||||||
|
success = await self._try_reconnect(report_error=report_error)
|
||||||
|
if success:
|
||||||
|
logger.info(f"{self} reconnected successfully, will retry send the message")
|
||||||
|
# trying to send the message one more time
|
||||||
|
await self._websocket.send(message)
|
||||||
|
else:
|
||||||
|
logger.error(f"{self} send failed; unable to reconnect")
|
||||||
|
|
||||||
async def _receive_task_handler(self, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
async def _receive_task_handler(self, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
||||||
"""Handle websocket message receiving with automatic retry logic.
|
"""Handle websocket message receiving with automatic retry logic.
|
||||||
|
|
||||||
@@ -76,13 +130,9 @@ class WebsocketService(ABC):
|
|||||||
Args:
|
Args:
|
||||||
report_error: Callback function to report connection errors.
|
report_error: Callback function to report connection errors.
|
||||||
"""
|
"""
|
||||||
retry_count = 0
|
|
||||||
MAX_RETRIES = 3
|
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
await self._receive_messages()
|
await self._receive_messages()
|
||||||
retry_count = 0 # Reset counter on successful message receive
|
|
||||||
except ConnectionClosedOK as e:
|
except ConnectionClosedOK as e:
|
||||||
# Normal closure, don't retry
|
# Normal closure, don't retry
|
||||||
logger.debug(f"{self} connection closed normally: {e}")
|
logger.debug(f"{self} connection closed normally: {e}")
|
||||||
@@ -92,21 +142,9 @@ class WebsocketService(ABC):
|
|||||||
logger.error(message)
|
logger.error(message)
|
||||||
|
|
||||||
if self._reconnect_on_error:
|
if self._reconnect_on_error:
|
||||||
retry_count += 1
|
success = await self._try_reconnect(report_error=report_error)
|
||||||
if retry_count >= MAX_RETRIES:
|
if not success:
|
||||||
await report_error(ErrorFrame(message))
|
|
||||||
break
|
break
|
||||||
|
|
||||||
logger.warning(f"{self} connection error, will retry: {e}")
|
|
||||||
await report_error(ErrorFrame(message))
|
|
||||||
|
|
||||||
try:
|
|
||||||
if await self._reconnect_websocket(retry_count):
|
|
||||||
retry_count = 0 # Reset counter on successful reconnection
|
|
||||||
wait_time = exponential_backoff_time(retry_count)
|
|
||||||
await asyncio.sleep(wait_time)
|
|
||||||
except Exception as reconnect_error:
|
|
||||||
logger.error(f"{self} reconnection failed: {reconnect_error}")
|
|
||||||
else:
|
else:
|
||||||
await report_error(ErrorFrame(message))
|
await report_error(ErrorFrame(message))
|
||||||
break
|
break
|
||||||
|
|||||||
Reference in New Issue
Block a user