Clean up on Sarvam STT and TTS classes

This commit is contained in:
Mark Backman
2026-02-07 10:53:43 -05:00
parent 3ff9b7b5ad
commit 6305e04569
5 changed files with 50 additions and 44 deletions

View File

@@ -0,0 +1 @@
- Added `bulbul:v3-beta` TTS model support for Sarvam AI with temperature control and 25 new speaker voices.

1
changelog/3671.added.md Normal file
View File

@@ -0,0 +1 @@
- Added `saaras:v3` STT model support for Sarvam AI with new `mode` parameter (transcribe, translate, verbatim, translit, codemix) and prompt support.

1
changelog/3671.fixed.md Normal file
View File

@@ -0,0 +1 @@
- Fixed issues in Sarvam STT and TTS services: missing event handler registration for VAD signals, `Optional[bool]` type annotations, WebSocket state cleanup on API errors, and TTS disconnect/reconnection state management.

View File

@@ -1,3 +1,9 @@
#
# Copyright (c) 20242026, Daily
#
# SPDX-License-Identifier: BSD 2-Clause License
#
"""Sarvam AI Speech-to-Text service implementation. """Sarvam AI Speech-to-Text service implementation.
This module provides a streaming Speech-to-Text service using Sarvam AI's WebSocket-based This module provides a streaming Speech-to-Text service using Sarvam AI's WebSocket-based
@@ -7,7 +13,7 @@ can handle multiple audio formats for Indian language speech recognition.
import base64 import base64
from dataclasses import dataclass from dataclasses import dataclass
from typing import Dict, Literal, Optional from typing import AsyncGenerator, Dict, Literal, Optional
from loguru import logger from loguru import logger
from pydantic import BaseModel from pydantic import BaseModel
@@ -149,8 +155,8 @@ class SarvamSTTService(STTService):
language: Optional[Language] = None language: Optional[Language] = None
prompt: Optional[str] = None prompt: Optional[str] = None
mode: Optional[Literal["transcribe", "translate", "verbatim", "translit", "codemix"]] = None mode: Optional[Literal["transcribe", "translate", "verbatim", "translit", "codemix"]] = None
vad_signals: bool = None vad_signals: Optional[bool] = None
high_vad_sensitivity: bool = None high_vad_sensitivity: Optional[bool] = None
def __init__( def __init__(
self, self,
@@ -233,6 +239,12 @@ class SarvamSTTService(STTService):
self._websocket_context = None self._websocket_context = None
self._socket_client = None self._socket_client = None
self._receive_task = None self._receive_task = None
if self._vad_signals:
self._register_event_handler("on_speech_started")
self._register_event_handler("on_speech_stopped")
self._register_event_handler("on_utterance_end")
logger.info(f"Sarvam STT initialized with SDK headers: {self._sdk_headers}") logger.info(f"Sarvam STT initialized with SDK headers: {self._sdk_headers}")
def language_to_service_language(self, language: Language) -> str: def language_to_service_language(self, language: Language) -> str:
@@ -337,7 +349,7 @@ class SarvamSTTService(STTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes): async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
"""Send audio data to Sarvam for transcription. """Send audio data to Sarvam for transcription.
Args: Args:
@@ -385,18 +397,25 @@ class SarvamSTTService(STTService):
logger.debug("Connecting to Sarvam") logger.debug("Connecting to Sarvam")
try: try:
# Convert boolean parameters to string for SDK
vad_signals_str = "true" if self._vad_signals else "false"
high_vad_sensitivity_str = "true" if self._high_vad_sensitivity else "false"
# Build common connection parameters # Build common connection parameters
connect_kwargs = { connect_kwargs = {
"model": self.model_name, "model": self.model_name,
"vad_signals": vad_signals_str,
"high_vad_sensitivity": high_vad_sensitivity_str,
"sample_rate": str(self.sample_rate), "sample_rate": str(self.sample_rate),
} }
# Enable flush signal when using Pipecat's VAD (not Sarvam's) so that
# the flush() call on user-stopped-speaking is honored by the server.
if not self._vad_signals:
connect_kwargs["flush_signal"] = "true"
# Only send vad parameters when explicitly set (avoid overriding server defaults)
if self._vad_signals is not None:
connect_kwargs["vad_signals"] = "true" if self._vad_signals else "false"
if self._high_vad_sensitivity is not None:
connect_kwargs["high_vad_sensitivity"] = (
"true" if self._high_vad_sensitivity else "false"
)
# Add language_code for models that support it # Add language_code for models that support it
if self._language_string is not None: if self._language_string is not None:
connect_kwargs["language_code"] = self._language_string connect_kwargs["language_code"] = self._language_string
@@ -447,6 +466,8 @@ class SarvamSTTService(STTService):
logger.info("Connected to Sarvam successfully") logger.info("Connected to Sarvam successfully")
except ApiError as e: except ApiError as e:
self._socket_client = None
self._websocket_context = None
await self.push_error(error_msg=f"Sarvam API error: {e}", exception=e) await self.push_error(error_msg=f"Sarvam API error: {e}", exception=e)
except Exception as e: except Exception as e:
self._socket_client = None self._socket_client = None

View File

@@ -525,7 +525,7 @@ class SarvamHttpTTSService(TTSService):
audio_data = base64.b64decode(base64_audio) audio_data = base64.b64decode(base64_audio)
# Strip WAV header (first 44 bytes) if present # Strip WAV header (first 44 bytes) if present
if audio_data.startswith(b"RIFF"): if len(audio_data) > 44 and audio_data.startswith(b"RIFF"):
logger.debug("Stripping WAV header from Sarvam audio data") logger.debug("Stripping WAV header from Sarvam audio data")
audio_data = audio_data[44:] audio_data = audio_data[44:]
@@ -792,7 +792,6 @@ class SarvamTTSService(InterruptibleTTSService):
self._started = False self._started = False
self._receive_task = None self._receive_task = None
self._keepalive_task = None self._keepalive_task = None
self._disconnecting = False
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
"""Check if this service can generate processing metrics. """Check if this service can generate processing metrics.
@@ -844,10 +843,13 @@ class SarvamTTSService(InterruptibleTTSService):
await self._disconnect() await self._disconnect()
async def flush_audio(self): async def flush_audio(self):
"""Flush any pending audio synthesis by sending stop command.""" """Flush any pending audio synthesis by sending flush command."""
if self._websocket: try:
msg = {"type": "flush"} if self._websocket:
await self._websocket.send(json.dumps(msg)) msg = {"type": "flush"}
await self._websocket.send(json.dumps(msg))
except Exception as e:
await self.push_error(error_msg=f"Error sending flush to Sarvam: {e}", exception=e)
async def push_frame(self, frame: Frame, direction: FrameDirection = FrameDirection.DOWNSTREAM): async def push_frame(self, frame: Frame, direction: FrameDirection = FrameDirection.DOWNSTREAM):
"""Push a frame downstream with special handling for stop conditions. """Push a frame downstream with special handling for stop conditions.
@@ -894,29 +896,15 @@ class SarvamTTSService(InterruptibleTTSService):
"""Disconnect from Sarvam WebSocket and clean up tasks.""" """Disconnect from Sarvam WebSocket and clean up tasks."""
await super()._disconnect() await super()._disconnect()
try: if self._receive_task:
# First, set a flag to prevent new operations await self.cancel_task(self._receive_task)
self._disconnecting = True self._receive_task = None
# Cancel background tasks BEFORE closing websocket if self._keepalive_task:
if self._receive_task: await self.cancel_task(self._keepalive_task)
await self.cancel_task(self._receive_task, timeout=2.0) self._keepalive_task = None
self._receive_task = None
if self._keepalive_task: await self._disconnect_websocket()
await self.cancel_task(self._keepalive_task, timeout=2.0)
self._keepalive_task = None
# Now close the websocket
await self._disconnect_websocket()
except Exception as e:
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)
finally:
# Reset state only after everything is cleaned up
self._started = False
self._websocket = None
self._disconnecting = False
async def _connect_websocket(self): async def _connect_websocket(self):
"""Establish WebSocket connection to Sarvam API.""" """Establish WebSocket connection to Sarvam API."""
@@ -996,6 +984,7 @@ class SarvamTTSService(InterruptibleTTSService):
if "too long" in error_msg.lower() or "timeout" in error_msg.lower(): if "too long" in error_msg.lower() or "timeout" in error_msg.lower():
logger.warning("Connection timeout detected, service may need restart") logger.warning("Connection timeout detected, service may need restart")
self._started = False
await self.push_frame(ErrorFrame(error=f"TTS Error: {error_msg}")) await self.push_frame(ErrorFrame(error=f"TTS Error: {error_msg}"))
async def _keepalive_task_handler(self): async def _keepalive_task_handler(self):
@@ -1007,19 +996,12 @@ class SarvamTTSService(InterruptibleTTSService):
async def _send_keepalive(self): async def _send_keepalive(self):
"""Send keepalive message to maintain connection.""" """Send keepalive message to maintain connection."""
if self._disconnecting:
return
if self._websocket and self._websocket.state == State.OPEN: if self._websocket and self._websocket.state == State.OPEN:
msg = {"type": "ping"} msg = {"type": "ping"}
await self._websocket.send(json.dumps(msg)) await self._websocket.send(json.dumps(msg))
async def _send_text(self, text: str): async def _send_text(self, text: str):
"""Send text to Sarvam WebSocket for synthesis.""" """Send text to Sarvam WebSocket for synthesis."""
if self._disconnecting:
logger.warning("Service is disconnecting, ignoring text send")
return
if self._websocket and self._websocket.state == State.OPEN: if self._websocket and self._websocket.state == State.OPEN:
msg = {"type": "text", "data": {"text": text}} msg = {"type": "text", "data": {"text": text}}
await self._websocket.send(json.dumps(msg)) await self._websocket.send(json.dumps(msg))