Update InputParams to languages: support str or List of Languages

This commit is contained in:
Mark Backman
2025-02-11 16:21:58 -05:00
parent 8c2071f248
commit 32d8f6153f
2 changed files with 52 additions and 28 deletions

View File

@@ -45,7 +45,7 @@ async def main():
) )
stt = GoogleSTTService( stt = GoogleSTTService(
params=GoogleSTTService.InputParams(language=Language.EN_US), params=GoogleSTTService.InputParams(languages=Language.EN_US),
) )
tts = GoogleTTSService( tts = GoogleTTSService(

View File

@@ -14,11 +14,11 @@ import os
os.environ["GRPC_ENABLE_FORK_SUPPORT"] = "false" os.environ["GRPC_ENABLE_FORK_SUPPORT"] = "false"
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, AsyncGenerator, Dict, List, Literal, Optional from typing import Any, AsyncGenerator, Dict, List, Literal, Optional, Union
from loguru import logger from loguru import logger
from PIL import Image from PIL import Image
from pydantic import BaseModel, Field from pydantic import BaseModel, Field, field_validator
from pipecat.frames.frames import ( from pipecat.frames.frames import (
AudioRawFrame, AudioRawFrame,
@@ -1425,7 +1425,7 @@ class GoogleSTTService(STTService):
"""Configuration parameters for Google Speech-to-Text. """Configuration parameters for Google Speech-to-Text.
Attributes: Attributes:
language: Recognition language (defaults to US English). languages: Single language or list of recognition languages. First language is primary.
model: Speech recognition model to use. model: Speech recognition model to use.
use_separate_recognition_per_channel: Process each audio channel separately. use_separate_recognition_per_channel: Process each audio channel separately.
enable_automatic_punctuation: Add punctuation to transcripts. enable_automatic_punctuation: Add punctuation to transcripts.
@@ -1438,7 +1438,7 @@ class GoogleSTTService(STTService):
enable_voice_activity_events: Detect voice activity in audio. enable_voice_activity_events: Detect voice activity in audio.
""" """
language: Optional[Language] = Language.EN_US languages: Union[Language, List[Language]] = Field(default_factory=lambda: [Language.EN_US])
model: Optional[str] = "latest_long" model: Optional[str] = "latest_long"
use_separate_recognition_per_channel: Optional[bool] = False use_separate_recognition_per_channel: Optional[bool] = False
enable_automatic_punctuation: Optional[bool] = True enable_automatic_punctuation: Optional[bool] = True
@@ -1450,6 +1450,19 @@ class GoogleSTTService(STTService):
enable_interim_results: Optional[bool] = True enable_interim_results: Optional[bool] = True
enable_voice_activity_events: Optional[bool] = False enable_voice_activity_events: Optional[bool] = False
@field_validator("languages", mode="before")
@classmethod
def validate_languages(cls, v) -> List[Language]:
if isinstance(v, Language):
return [v]
return v
@property
def language_list(self) -> List[Language]:
"""Get languages as a guaranteed list."""
assert isinstance(self.languages, list)
return self.languages
def __init__( def __init__(
self, self,
*, *,
@@ -1506,9 +1519,9 @@ class GoogleSTTService(STTService):
self._client = speech_v2.SpeechAsyncClient(credentials=creds, client_options=client_options) self._client = speech_v2.SpeechAsyncClient(credentials=creds, client_options=client_options)
self._settings = { self._settings = {
"language_code": self.language_to_service_language(params.language) "language_codes": [
if params.language self.language_to_service_language(lang) for lang in params.language_list
else "en-US", ],
"model": params.model, "model": params.model,
"use_separate_recognition_per_channel": params.use_separate_recognition_per_channel, "use_separate_recognition_per_channel": params.use_separate_recognition_per_channel,
"enable_automatic_punctuation": params.enable_automatic_punctuation, "enable_automatic_punctuation": params.enable_automatic_punctuation,
@@ -1521,22 +1534,30 @@ class GoogleSTTService(STTService):
"enable_voice_activity_events": params.enable_voice_activity_events, "enable_voice_activity_events": params.enable_voice_activity_events,
} }
def language_to_service_language(self, language: Language) -> Optional[str]: def language_to_service_language(self, language: Language | List[Language]) -> str | List[str]:
"""Convert Language enum to Google STT language code. """Convert Language enum(s) to Google STT language code(s).
Args: Args:
language: Language enum value. language: Single Language enum or list of Language enums.
Returns: Returns:
str: Google STT language code. str | List[str]: Google STT language code(s).
""" """
return language_to_google_stt_language(language) if isinstance(language, list):
return [language_to_google_stt_language(lang) or "en-US" for lang in language]
return language_to_google_stt_language(language) or "en-US"
async def set_language(self, language: Language): async def set_languages(self, languages: List[Language]):
"""Update the service's recognition language.""" """Update the service's recognition languages.
logger.info(f"Switching STT language to: [{language}]")
self._settings["language_code"] = self.language_to_service_language(language) Args:
# Recreate stream with new language languages: List of languages for recognition. First language is primary.
"""
logger.info(f"Switching STT languages to: {languages}")
self._settings["language_codes"] = [
self.language_to_service_language(lang) for lang in languages
]
# Recreate stream with new languages
if self._streaming_task: if self._streaming_task:
await self._disconnect() await self._disconnect()
await self._connect() await self._connect()
@@ -1565,7 +1586,7 @@ class GoogleSTTService(STTService):
async def update_options( async def update_options(
self, self,
*, *,
language: Optional[Language] = None, languages: Optional[List[Language]] = None,
model: Optional[str] = None, model: Optional[str] = None,
enable_automatic_punctuation: Optional[bool] = None, enable_automatic_punctuation: Optional[bool] = None,
enable_spoken_punctuation: Optional[bool] = None, enable_spoken_punctuation: Optional[bool] = None,
@@ -1580,7 +1601,7 @@ class GoogleSTTService(STTService):
"""Update service options dynamically. """Update service options dynamically.
Args: Args:
language: New recognition language. languages: New list of recongition languages.
model: New recognition model. model: New recognition model.
enable_automatic_punctuation: Enable/disable automatic punctuation. enable_automatic_punctuation: Enable/disable automatic punctuation.
enable_spoken_punctuation: Enable/disable spoken punctuation. enable_spoken_punctuation: Enable/disable spoken punctuation.
@@ -1599,9 +1620,11 @@ class GoogleSTTService(STTService):
needs_reconnect = False needs_reconnect = False
# Update settings with new values # Update settings with new values
if language is not None: if languages is not None:
logger.debug(f"Updating language to: {language}") logger.debug(f"Updating language to: {languages}")
self._settings["language_code"] = self.language_to_service_language(language) self._settings["language_codes"] = [
self.language_to_service_language(lang) for lang in languages
]
needs_reconnect = True needs_reconnect = True
if model is not None: if model is not None:
@@ -1672,7 +1695,7 @@ class GoogleSTTService(STTService):
sample_rate_hertz=self.sample_rate, sample_rate_hertz=self.sample_rate,
audio_channel_count=1, audio_channel_count=1,
), ),
language_codes=[self._settings["language_code"]], language_codes=self._settings["language_codes"],
model=self._settings["model"], model=self._settings["model"],
features=cloud_speech.RecognitionFeatures( features=cloud_speech.RecognitionFeatures(
enable_automatic_punctuation=self._settings["enable_automatic_punctuation"], enable_automatic_punctuation=self._settings["enable_automatic_punctuation"],
@@ -1775,16 +1798,17 @@ class GoogleSTTService(STTService):
if not transcript: if not transcript:
continue continue
# Use the primary language (first in the list)
primary_language = self._settings["language_codes"][0]
if result.is_final: if result.is_final:
await self.push_frame( await self.push_frame(
TranscriptionFrame( TranscriptionFrame(transcript, "", time_now_iso8601(), primary_language)
transcript, "", time_now_iso8601(), self._settings["language_code"]
)
) )
else: else:
await self.push_frame( await self.push_frame(
InterimTranscriptionFrame( InterimTranscriptionFrame(
transcript, "", time_now_iso8601(), self._settings["language_code"] transcript, "", time_now_iso8601(), primary_language
) )
) )