Merge pull request #3827 from pipecat-ai/pk/gemini-tts-service-remove-model-ivar
Remove unnecessary `_model` ivar from `GeminiTTSService`, using `_set…
This commit is contained in:
@@ -279,7 +279,6 @@ class ElevenLabsSTTService(SegmentedSTTService):
|
|||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._base_url = base_url
|
self._base_url = base_url
|
||||||
self._session = aiohttp_session
|
self._session = aiohttp_session
|
||||||
self._model_id = model
|
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
"""Check if the service can generate processing metrics.
|
"""Check if the service can generate processing metrics.
|
||||||
@@ -300,25 +299,6 @@ class ElevenLabsSTTService(SegmentedSTTService):
|
|||||||
"""
|
"""
|
||||||
return language_to_elevenlabs_language(language)
|
return language_to_elevenlabs_language(language)
|
||||||
|
|
||||||
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
|
||||||
"""Apply a settings delta.
|
|
||||||
|
|
||||||
Converts language to ElevenLabs format before applying and keeps
|
|
||||||
``_model_id`` in sync with the model setting.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
delta: A :class:`STTSettings` (or ``ElevenLabsSTTSettings``) delta.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dict mapping changed field names to their previous values.
|
|
||||||
"""
|
|
||||||
changed = await super()._update_settings(delta)
|
|
||||||
|
|
||||||
if "model" in changed:
|
|
||||||
self._model_id = self._settings.model
|
|
||||||
|
|
||||||
return changed
|
|
||||||
|
|
||||||
async def _transcribe_audio(self, audio_data: bytes) -> dict:
|
async def _transcribe_audio(self, audio_data: bytes) -> dict:
|
||||||
"""Upload audio data to ElevenLabs and get transcription result.
|
"""Upload audio data to ElevenLabs and get transcription result.
|
||||||
|
|
||||||
@@ -344,7 +324,7 @@ class ElevenLabsSTTService(SegmentedSTTService):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Add required model_id, language_code, and tag_audio_events
|
# Add required model_id, language_code, and tag_audio_events
|
||||||
data.add_field("model_id", self._model_id)
|
data.add_field("model_id", self._settings.model)
|
||||||
data.add_field("language_code", self._settings.language)
|
data.add_field("language_code", self._settings.language)
|
||||||
data.add_field("tag_audio_events", str(self._settings.tag_audio_events).lower())
|
data.add_field("tag_audio_events", str(self._settings.tag_audio_events).lower())
|
||||||
|
|
||||||
@@ -522,7 +502,6 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._base_url = base_url
|
self._base_url = base_url
|
||||||
self._model_id = model
|
|
||||||
self._audio_format = "" # initialized in start()
|
self._audio_format = "" # initialized in start()
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
@@ -540,9 +519,6 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
||||||
"""Apply a settings delta and reconnect if anything changed.
|
"""Apply a settings delta and reconnect if anything changed.
|
||||||
|
|
||||||
Converts language to ElevenLabs format before applying and keeps
|
|
||||||
``_model_id`` in sync.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
delta: A :class:`STTSettings` (or ``ElevenLabsRealtimeSTTSettings``) delta.
|
delta: A :class:`STTSettings` (or ``ElevenLabsRealtimeSTTSettings``) delta.
|
||||||
|
|
||||||
@@ -554,11 +530,9 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
if not changed:
|
if not changed:
|
||||||
return changed
|
return changed
|
||||||
|
|
||||||
if "model" in changed:
|
|
||||||
self._model_id = self._settings.model
|
|
||||||
|
|
||||||
await self._disconnect()
|
await self._disconnect()
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
return changed
|
return changed
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
@@ -704,7 +678,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
logger.debug("Connecting to ElevenLabs Realtime STT")
|
logger.debug("Connecting to ElevenLabs Realtime STT")
|
||||||
|
|
||||||
# Build query parameters
|
# Build query parameters
|
||||||
params = [f"model_id={self._model_id}"]
|
params = [f"model_id={self._settings.model}"]
|
||||||
|
|
||||||
if self._settings.language:
|
if self._settings.language:
|
||||||
params.append(f"language_code={self._settings.language}")
|
params.append(f"language_code={self._settings.language}")
|
||||||
|
|||||||
@@ -1236,7 +1236,7 @@ class GeminiTTSService(GoogleBaseTTSService):
|
|||||||
super().__init__(
|
super().__init__(
|
||||||
sample_rate=sample_rate,
|
sample_rate=sample_rate,
|
||||||
settings=GeminiTTSSettings(
|
settings=GeminiTTSSettings(
|
||||||
model=None,
|
model=model,
|
||||||
language=self.language_to_service_language(params.language)
|
language=self.language_to_service_language(params.language)
|
||||||
if params.language
|
if params.language
|
||||||
else "en-US",
|
else "en-US",
|
||||||
@@ -1249,7 +1249,6 @@ class GeminiTTSService(GoogleBaseTTSService):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self._location = location
|
self._location = location
|
||||||
self._model = model
|
|
||||||
self._client: texttospeech_v1.TextToSpeechAsyncClient = self._create_client(
|
self._client: texttospeech_v1.TextToSpeechAsyncClient = self._create_client(
|
||||||
credentials, credentials_path
|
credentials, credentials_path
|
||||||
)
|
)
|
||||||
@@ -1327,7 +1326,7 @@ class GeminiTTSService(GoogleBaseTTSService):
|
|||||||
|
|
||||||
voice = texttospeech_v1.VoiceSelectionParams(
|
voice = texttospeech_v1.VoiceSelectionParams(
|
||||||
language_code=self._settings.language,
|
language_code=self._settings.language,
|
||||||
model_name=self._model,
|
model_name=self._settings.model,
|
||||||
multi_speaker_voice_config=multi_speaker_voice_config,
|
multi_speaker_voice_config=multi_speaker_voice_config,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -1335,7 +1334,7 @@ class GeminiTTSService(GoogleBaseTTSService):
|
|||||||
voice = texttospeech_v1.VoiceSelectionParams(
|
voice = texttospeech_v1.VoiceSelectionParams(
|
||||||
language_code=self._settings.language,
|
language_code=self._settings.language,
|
||||||
name=self._settings.voice,
|
name=self._settings.voice,
|
||||||
model_name=self._model,
|
model_name=self._settings.model,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Create streaming config
|
# Create streaming config
|
||||||
|
|||||||
@@ -101,8 +101,6 @@ class HathoraSTTService(SegmentedSTTService):
|
|||||||
),
|
),
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._model = model
|
|
||||||
self._api_key = api_key or os.getenv("HATHORA_API_KEY")
|
self._api_key = api_key or os.getenv("HATHORA_API_KEY")
|
||||||
self._base_url = base_url
|
self._base_url = base_url
|
||||||
|
|
||||||
@@ -136,7 +134,7 @@ class HathoraSTTService(SegmentedSTTService):
|
|||||||
url = f"{self._base_url}"
|
url = f"{self._base_url}"
|
||||||
|
|
||||||
payload = {
|
payload = {
|
||||||
"model": self._model,
|
"model": self._settings.model,
|
||||||
}
|
}
|
||||||
|
|
||||||
if self._settings.language is not None:
|
if self._settings.language is not None:
|
||||||
|
|||||||
@@ -120,7 +120,6 @@ class HathoraTTSService(TTSService):
|
|||||||
),
|
),
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
self._model = model
|
|
||||||
self._api_key = api_key or os.getenv("HATHORA_API_KEY")
|
self._api_key = api_key or os.getenv("HATHORA_API_KEY")
|
||||||
self._base_url = base_url
|
self._base_url = base_url
|
||||||
|
|
||||||
@@ -149,7 +148,7 @@ class HathoraTTSService(TTSService):
|
|||||||
|
|
||||||
url = f"{self._base_url}"
|
url = f"{self._base_url}"
|
||||||
|
|
||||||
payload = {"model": self._model, "text": text}
|
payload = {"model": self._settings.model, "text": text}
|
||||||
|
|
||||||
if self._settings.voice is not None:
|
if self._settings.voice is not None:
|
||||||
payload["voice"] = self._settings.voice
|
payload["voice"] = self._settings.voice
|
||||||
|
|||||||
Reference in New Issue
Block a user