Clean up on Sarvam STT and TTS classes
This commit is contained in:
1
changelog/3671.added.2.md
Normal file
1
changelog/3671.added.2.md
Normal 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
1
changelog/3671.added.md
Normal 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
1
changelog/3671.fixed.md
Normal 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.
|
||||||
@@ -1,3 +1,9 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024–2026, 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
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
Reference in New Issue
Block a user