Merge pull request #3376 from dhruvladia-sarvam/update/sarvam-plugins
headers update
This commit is contained in:
20
src/pipecat/services/sarvam/_sdk.py
Normal file
20
src/pipecat/services/sarvam/_sdk.py
Normal file
@@ -0,0 +1,20 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024–2026, Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import platform
|
||||||
|
from typing import Dict
|
||||||
|
|
||||||
|
from pipecat import version as pipecat_version
|
||||||
|
|
||||||
|
|
||||||
|
def sdk_headers() -> Dict[str, str]:
|
||||||
|
"""SDK identification headers for upstream providers."""
|
||||||
|
return {
|
||||||
|
"X-SDK-Source": "Pipecat",
|
||||||
|
"X-SDK-Version": pipecat_version(),
|
||||||
|
"SDK-Language": "Python",
|
||||||
|
"sdk-language-version": platform.python_version(),
|
||||||
|
}
|
||||||
@@ -18,6 +18,7 @@ from pipecat.frames.frames import (
|
|||||||
StartFrame,
|
StartFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
)
|
)
|
||||||
|
from pipecat.services.sarvam._sdk import sdk_headers
|
||||||
from pipecat.services.stt_service import STTService
|
from pipecat.services.stt_service import STTService
|
||||||
from pipecat.transcriptions.language import Language, resolve_language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
@@ -125,7 +126,7 @@ class SarvamSTTService(STTService):
|
|||||||
|
|
||||||
self.set_model_name(model)
|
self.set_model_name(model)
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._language_code = params.language
|
self._language_code: Optional[Language] = params.language
|
||||||
# For saarika models, default to "unknown" if language is not provided
|
# For saarika models, default to "unknown" if language is not provided
|
||||||
if params.language:
|
if params.language:
|
||||||
self._language_string = language_to_sarvam_language(params.language)
|
self._language_string = language_to_sarvam_language(params.language)
|
||||||
@@ -141,7 +142,16 @@ class SarvamSTTService(STTService):
|
|||||||
self._input_audio_codec = input_audio_codec
|
self._input_audio_codec = input_audio_codec
|
||||||
|
|
||||||
# Initialize Sarvam SDK client
|
# Initialize Sarvam SDK client
|
||||||
|
self._sdk_headers = sdk_headers()
|
||||||
|
# NOTE: We avoid passing non-standard kwargs here because different sarvamai
|
||||||
|
# versions expose different constructor signatures (static type checkers
|
||||||
|
# complain otherwise). We instead inject headers best-effort below.
|
||||||
self._sarvam_client = AsyncSarvamAI(api_subscription_key=api_key)
|
self._sarvam_client = AsyncSarvamAI(api_subscription_key=api_key)
|
||||||
|
for attr in ("default_headers", "_default_headers", "headers", "_headers"):
|
||||||
|
d = getattr(self._sarvam_client, attr, None)
|
||||||
|
if isinstance(d, dict):
|
||||||
|
d.update(self._sdk_headers)
|
||||||
|
break
|
||||||
self._websocket_context = None
|
self._websocket_context = None
|
||||||
self._socket_client = None
|
self._socket_client = None
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
@@ -297,17 +307,28 @@ class SarvamSTTService(STTService):
|
|||||||
"sample_rate": str(self.sample_rate),
|
"sample_rate": str(self.sample_rate),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def _connect_with_sdk_headers(connect_fn, **kwargs):
|
||||||
|
# Different SDK versions may use different kwarg names.
|
||||||
|
for header_kw in ("headers", "additional_headers", "extra_headers"):
|
||||||
|
try:
|
||||||
|
return connect_fn(**kwargs, **{header_kw: self._sdk_headers})
|
||||||
|
except TypeError:
|
||||||
|
pass
|
||||||
|
return connect_fn(**kwargs)
|
||||||
|
|
||||||
# Choose the appropriate service based on model
|
# Choose the appropriate service based on model
|
||||||
if "saarika" in self.model_name.lower():
|
if "saarika" in self.model_name.lower():
|
||||||
# STT service - requires language_code
|
# STT service - requires language_code
|
||||||
connect_kwargs["language_code"] = self._language_string
|
connect_kwargs["language_code"] = self._language_string
|
||||||
self._websocket_context = self._sarvam_client.speech_to_text_streaming.connect(
|
self._websocket_context = _connect_with_sdk_headers(
|
||||||
**connect_kwargs
|
self._sarvam_client.speech_to_text_streaming.connect,
|
||||||
|
**connect_kwargs,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# STT-Translate service - auto-detects input language and returns translated text
|
# STT-Translate service - auto-detects input language and returns translated text
|
||||||
self._websocket_context = (
|
self._websocket_context = _connect_with_sdk_headers(
|
||||||
self._sarvam_client.speech_to_text_translate_streaming.connect(**connect_kwargs)
|
self._sarvam_client.speech_to_text_translate_streaming.connect,
|
||||||
|
**connect_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Enter the async context manager
|
# Enter the async context manager
|
||||||
@@ -394,6 +415,10 @@ class SarvamSTTService(STTService):
|
|||||||
logger.debug("User started speaking")
|
logger.debug("User started speaking")
|
||||||
await self._call_event_handler("on_speech_started")
|
await self._call_event_handler("on_speech_started")
|
||||||
|
|
||||||
|
elif signal == "END_SPEECH":
|
||||||
|
logger.debug("User stopped speaking")
|
||||||
|
await self._call_event_handler("on_speech_stopped")
|
||||||
|
|
||||||
elif message.type == "data":
|
elif message.type == "data":
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
transcript = message.data.transcript
|
transcript = message.data.transcript
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
#
|
#
|
||||||
# Copyright (c) 2024-2026, Daily
|
# Copyright (c) 2024–2026, Daily
|
||||||
#
|
#
|
||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
@@ -28,6 +28,7 @@ from pipecat.frames.frames import (
|
|||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
|
from pipecat.services.sarvam._sdk import sdk_headers
|
||||||
from pipecat.services.tts_service import InterruptibleTTSService, TTSService
|
from pipecat.services.tts_service import InterruptibleTTSService, TTSService
|
||||||
from pipecat.transcriptions.language import Language, resolve_language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
from pipecat.utils.tracing.service_decorators import traced_tts
|
from pipecat.utils.tracing.service_decorators import traced_tts
|
||||||
@@ -245,6 +246,7 @@ class SarvamHttpTTSService(TTSService):
|
|||||||
headers = {
|
headers = {
|
||||||
"api-subscription-key": self._api_key,
|
"api-subscription-key": self._api_key,
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
|
**sdk_headers(),
|
||||||
}
|
}
|
||||||
|
|
||||||
url = f"{self._base_url}/text-to-speech"
|
url = f"{self._base_url}/text-to-speech"
|
||||||
@@ -576,6 +578,7 @@ class SarvamTTSService(InterruptibleTTSService):
|
|||||||
self._websocket_url,
|
self._websocket_url,
|
||||||
additional_headers={
|
additional_headers={
|
||||||
"api-subscription-key": self._api_key,
|
"api-subscription-key": self._api_key,
|
||||||
|
**sdk_headers(),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
logger.debug("Connected to Sarvam TTS Websocket")
|
logger.debug("Connected to Sarvam TTS Websocket")
|
||||||
|
|||||||
Reference in New Issue
Block a user