fix: Fix language param and include suggested way of handling STT response

This commit is contained in:
shreyas-sarvam
2025-10-31 13:23:08 +05:30
parent 8d0e7e5e16
commit 1433df4de2
3 changed files with 42 additions and 32 deletions

View File

@@ -66,7 +66,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
stt = SarvamSTTService( stt = SarvamSTTService(
api_key=os.getenv("SARVAM_API_KEY"), api_key=os.getenv("SARVAM_API_KEY"),
model="saarika:v2.5", model="saarika:v2.5",
params=SarvamSTTService.InputParams(language=None),
) )
tts = SarvamHttpTTSService( tts = SarvamHttpTTSService(

View File

@@ -27,6 +27,7 @@ from pipecat.runner.utils import create_transport
from pipecat.services.openai.llm import OpenAILLMService from pipecat.services.openai.llm import OpenAILLMService
from pipecat.services.sarvam.stt import SarvamSTTService from pipecat.services.sarvam.stt import SarvamSTTService
from pipecat.services.sarvam.tts import SarvamTTSService from pipecat.services.sarvam.tts import SarvamTTSService
from pipecat.transcriptions.language import Language
from pipecat.transports.base_transport import BaseTransport, TransportParams from pipecat.transports.base_transport import BaseTransport, TransportParams
from pipecat.transports.daily.transport import DailyParams from pipecat.transports.daily.transport import DailyParams
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
@@ -65,8 +66,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
stt = SarvamSTTService( stt = SarvamSTTService(
api_key=os.getenv("SARVAM_API_KEY"), api_key=os.getenv("SARVAM_API_KEY"),
model="saarika:v2.5", model="saarika:v2.5",
# Example: set Hindi; omit or change via set_language at runtime
params=SarvamSTTService.InputParams(language=None),
) )
tts = SarvamTTSService( tts = SarvamTTSService(

View File

@@ -5,7 +5,6 @@ API. It supports real-time transcription with Voice Activity Detection (VAD) and
can handle multiple audio formats for Indian language speech recognition. can handle multiple audio formats for Indian language speech recognition.
""" """
import asyncio
import base64 import base64
from typing import Optional from typing import Optional
@@ -55,7 +54,6 @@ def language_to_sarvam_language(language: Language) -> str:
Language.TE_IN: "te-IN", Language.TE_IN: "te-IN",
Language.PA_IN: "pa-IN", Language.PA_IN: "pa-IN",
Language.OR_IN: "od-IN", Language.OR_IN: "od-IN",
Language.EN_US: "en-US",
Language.EN_IN: "en-IN", Language.EN_IN: "en-IN",
Language.AS_IN: "as-IN", Language.AS_IN: "as-IN",
} }
@@ -119,7 +117,7 @@ class SarvamSTTService(STTService):
self._sarvam_client = AsyncSarvamAI(api_subscription_key=api_key) self._sarvam_client = AsyncSarvamAI(api_subscription_key=api_key)
self._websocket_context = None self._websocket_context = None
self._socket_client = None self._socket_client = None
self._listening_task = None self._receive_task = None
def language_to_service_language(self, language: Language) -> str: def language_to_service_language(self, language: Language) -> str:
"""Convert pipecat Language enum to Sarvam's language code. """Convert pipecat Language enum to Sarvam's language code.
@@ -256,17 +254,13 @@ class SarvamSTTService(STTService):
# Register event handler for incoming messages # Register event handler for incoming messages
def _message_handler(message): def _message_handler(message):
"""Wrapper to handle async response handler.""" """Wrapper to handle async response handler."""
try: # Use Pipecat's built-in task management
loop = asyncio.get_running_loop() self.create_task(self._handle_message(message))
loop.create_task(self._handle_response(message))
except RuntimeError:
# Fallback if no running loop
asyncio.create_task(self._handle_response(message))
self._socket_client.on(EventType.MESSAGE, _message_handler) self._socket_client.on(EventType.MESSAGE, _message_handler)
# Start listening for messages # Start receive task using Pipecat's task management
self._listening_task = asyncio.create_task(self._socket_client.start_listening()) self._receive_task = self.create_task(self._receive_task_handler())
logger.info("Connected to Sarvam successfully") logger.info("Connected to Sarvam successfully")
@@ -281,13 +275,9 @@ class SarvamSTTService(STTService):
async def _disconnect(self): async def _disconnect(self):
"""Disconnect from Sarvam WebSocket API using SDK.""" """Disconnect from Sarvam WebSocket API using SDK."""
if self._listening_task: if self._receive_task:
self._listening_task.cancel() await self.cancel_task(self._receive_task)
try: self._receive_task = None
await self._listening_task
except asyncio.CancelledError:
pass
self._listening_task = None
if self._websocket_context and self._socket_client: if self._websocket_context and self._socket_client:
try: try:
@@ -300,8 +290,27 @@ class SarvamSTTService(STTService):
self._socket_client = None self._socket_client = None
self._websocket_context = None self._websocket_context = None
async def _handle_response(self, message): async def _receive_task_handler(self):
"""Handle transcription response from Sarvam SDK. """Handle incoming messages from Sarvam WebSocket.
This task wraps the SDK's start_listening() method which processes
messages via the registered event handler callback.
"""
if not self._socket_client:
return
try:
# Start listening for messages from the Sarvam SDK
# Messages will be handled via the _message_handler callback
await self._socket_client.start_listening()
except Exception as e:
logger.error(f"Error in Sarvam receive task: {e}")
await self.push_error(ErrorFrame(f"Sarvam receive task error: {e}"))
async def _handle_message(self, message):
"""Handle incoming WebSocket message from Sarvam SDK.
Processes transcription data and VAD events from the Sarvam service.
Args: Args:
message: The parsed response object from Sarvam WebSocket. message: The parsed response object from Sarvam WebSocket.
@@ -351,10 +360,20 @@ class SarvamSTTService(STTService):
await self.stop_processing_metrics() await self.stop_processing_metrics()
except Exception as e: except Exception as e:
logger.error(f"Error handling Sarvam response: {e}") logger.error(f"Error handling Sarvam message: {e}")
await self.push_error(ErrorFrame(f"Failed to handle response: {e}")) await self.push_error(ErrorFrame(f"Failed to handle message: {e}"))
await self.stop_all_metrics() await self.stop_all_metrics()
@traced_stt
async def _handle_transcription(
self, transcript: str, is_final: bool, language: Optional[Language] = None
):
"""Handle a transcription result with tracing.
This method is decorated with @traced_stt for observability.
"""
pass
def _map_language_code_to_enum(self, language_code: str) -> Language: def _map_language_code_to_enum(self, language_code: str) -> Language:
"""Map Sarvam language code to pipecat Language enum.""" """Map Sarvam language code to pipecat Language enum."""
mapping = { mapping = {
@@ -378,10 +397,3 @@ class SarvamSTTService(STTService):
"""Start TTFB and processing metrics collection.""" """Start TTFB and processing metrics collection."""
await self.start_ttfb_metrics() await self.start_ttfb_metrics()
await self.start_processing_metrics() await self.start_processing_metrics()
@traced_stt
async def _handle_transcription(
self, transcript: str, is_final: bool, language: Optional[Language] = None
):
"""Handle a transcription result with tracing."""
pass