Add GoogleSTTService
This commit is contained in:
@@ -20,6 +20,7 @@ from pipecat.pipeline.task import PipelineParams, PipelineTask
|
|||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
|
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
|
||||||
from pipecat.services.deepgram import DeepgramSTTService
|
from pipecat.services.deepgram import DeepgramSTTService
|
||||||
from pipecat.services.google import GoogleTTSService
|
from pipecat.services.google import GoogleTTSService
|
||||||
|
from pipecat.services.google.google import GoogleSTTService
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
@@ -46,11 +47,14 @@ async def main():
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
stt = GoogleSTTService(
|
||||||
|
credentials=os.getenv("GOOGLE_TEST_CREDENTIALS"),
|
||||||
|
)
|
||||||
|
|
||||||
tts = GoogleTTSService(
|
tts = GoogleTTSService(
|
||||||
voice_id="en-US-Journey-F",
|
voice_id="en-US-Journey-F",
|
||||||
params=GoogleTTSService.InputParams(language=Language.EN_US),
|
params=GoogleTTSService.InputParams(language=Language.EN_US),
|
||||||
|
credentials=os.getenv("GOOGLE_TEST_CREDENTIALS"),
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"), model="gpt-4o")
|
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"), model="gpt-4o")
|
||||||
|
|||||||
@@ -8,6 +8,11 @@ import asyncio
|
|||||||
import base64
|
import base64
|
||||||
import io
|
import io
|
||||||
import json
|
import json
|
||||||
|
import os
|
||||||
|
|
||||||
|
# Suppress gRPC fork warnings
|
||||||
|
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
|
||||||
|
|
||||||
@@ -17,15 +22,20 @@ from pydantic import BaseModel, Field
|
|||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
AudioRawFrame,
|
||||||
|
CancelFrame,
|
||||||
|
EndFrame,
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
FunctionCallResultProperties,
|
FunctionCallResultProperties,
|
||||||
|
InterimTranscriptionFrame,
|
||||||
LLMFullResponseEndFrame,
|
LLMFullResponseEndFrame,
|
||||||
LLMFullResponseStartFrame,
|
LLMFullResponseStartFrame,
|
||||||
LLMMessagesFrame,
|
LLMMessagesFrame,
|
||||||
LLMTextFrame,
|
LLMTextFrame,
|
||||||
LLMUpdateSettingsFrame,
|
LLMUpdateSettingsFrame,
|
||||||
OpenAILLMContextAssistantTimestampFrame,
|
OpenAILLMContextAssistantTimestampFrame,
|
||||||
|
StartFrame,
|
||||||
|
TranscriptionFrame,
|
||||||
TTSAudioRawFrame,
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
@@ -38,7 +48,7 @@ from pipecat.processors.aggregators.openai_llm_context import (
|
|||||||
OpenAILLMContextFrame,
|
OpenAILLMContextFrame,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.ai_services import ImageGenService, LLMService, TTSService
|
from pipecat.services.ai_services import ImageGenService, LLMService, STTService, TTSService
|
||||||
from pipecat.services.google.frames import LLMSearchResponseFrame
|
from pipecat.services.google.frames import LLMSearchResponseFrame
|
||||||
from pipecat.services.openai import (
|
from pipecat.services.openai import (
|
||||||
OpenAIAssistantContextAggregator,
|
OpenAIAssistantContextAggregator,
|
||||||
@@ -51,10 +61,12 @@ try:
|
|||||||
import google.ai.generativelanguage as glm
|
import google.ai.generativelanguage as glm
|
||||||
import google.generativeai as gai
|
import google.generativeai as gai
|
||||||
from google import genai
|
from google import genai
|
||||||
from google.cloud import texttospeech_v1
|
from google.cloud import speech_v2, texttospeech_v1
|
||||||
|
from google.cloud.speech_v2.types import cloud_speech
|
||||||
from google.genai import types
|
from google.genai import types
|
||||||
from google.generativeai.types import GenerationConfig
|
from google.generativeai.types import GenerationConfig
|
||||||
from google.oauth2 import service_account
|
from google.oauth2 import service_account
|
||||||
|
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -1097,3 +1109,230 @@ class GoogleImageGenService(ImageGenService):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} error generating image: {e}")
|
logger.error(f"{self} error generating image: {e}")
|
||||||
yield ErrorFrame(f"Image generation error: {str(e)}")
|
yield ErrorFrame(f"Image generation error: {str(e)}")
|
||||||
|
|
||||||
|
|
||||||
|
class GoogleSTTService(STTService):
|
||||||
|
class InputParams(BaseModel):
|
||||||
|
language: Optional[Language] = Language.EN_US
|
||||||
|
model: Optional[str] = "latest_long"
|
||||||
|
use_separate_recognition_per_channel: Optional[bool] = False
|
||||||
|
enable_automatic_punctuation: Optional[bool] = True
|
||||||
|
enable_spoken_punctuation: Optional[bool] = False
|
||||||
|
enable_spoken_emojis: Optional[bool] = False
|
||||||
|
profanity_filter: Optional[bool] = False
|
||||||
|
enable_word_time_offsets: Optional[bool] = False
|
||||||
|
enable_word_confidence: Optional[bool] = False
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
credentials: Optional[str] = None,
|
||||||
|
credentials_path: Optional[str] = None,
|
||||||
|
location: str = "global",
|
||||||
|
recognition_config: Optional[dict] = None,
|
||||||
|
sample_rate: Optional[int] = None,
|
||||||
|
params: InputParams = InputParams(),
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
super().__init__(sample_rate=sample_rate, **kwargs)
|
||||||
|
|
||||||
|
self._location = location
|
||||||
|
self._stream = None
|
||||||
|
self._config = None
|
||||||
|
self._request_queue = asyncio.Queue()
|
||||||
|
self._streaming_task = None
|
||||||
|
|
||||||
|
# Extract project ID and create client
|
||||||
|
if credentials:
|
||||||
|
json_account_info = json.loads(credentials)
|
||||||
|
self._project_id = json_account_info.get("project_id")
|
||||||
|
creds = service_account.Credentials.from_service_account_info(json_account_info)
|
||||||
|
elif credentials_path:
|
||||||
|
with open(credentials_path) as f:
|
||||||
|
json_account_info = json.load(f)
|
||||||
|
self._project_id = json_account_info.get("project_id")
|
||||||
|
creds = service_account.Credentials.from_service_account_file(credentials_path)
|
||||||
|
else:
|
||||||
|
raise ValueError("Either credentials or credentials_path must be provided")
|
||||||
|
|
||||||
|
if not self._project_id:
|
||||||
|
raise ValueError("Project ID not found in credentials")
|
||||||
|
|
||||||
|
logger.debug(f"Using project ID from credentials: {self._project_id}")
|
||||||
|
|
||||||
|
self._client = speech_v2.SpeechAsyncClient(credentials=creds)
|
||||||
|
|
||||||
|
self._settings = {
|
||||||
|
"language_code": self.language_to_service_language(params.language or Language.EN_US),
|
||||||
|
"model": params.model,
|
||||||
|
"use_separate_recognition_per_channel": params.use_separate_recognition_per_channel,
|
||||||
|
"enable_automatic_punctuation": params.enable_automatic_punctuation,
|
||||||
|
"enable_spoken_punctuation": params.enable_spoken_punctuation,
|
||||||
|
"enable_spoken_emojis": params.enable_spoken_emojis,
|
||||||
|
"profanity_filter": params.profanity_filter,
|
||||||
|
"enable_word_time_offsets": params.enable_word_time_offsets,
|
||||||
|
"enable_word_confidence": params.enable_word_confidence,
|
||||||
|
}
|
||||||
|
|
||||||
|
if recognition_config:
|
||||||
|
self._settings.update(recognition_config)
|
||||||
|
|
||||||
|
def language_to_service_language(self, language: Language) -> str:
|
||||||
|
return str(language.value)
|
||||||
|
|
||||||
|
async def set_language(self, language: Language):
|
||||||
|
logger.info(f"Switching STT language to: [{language}]")
|
||||||
|
self._settings["language_code"] = self.language_to_service_language(language)
|
||||||
|
# Recreate stream with new language
|
||||||
|
if self._streaming_task:
|
||||||
|
await self._disconnect()
|
||||||
|
await self._connect()
|
||||||
|
|
||||||
|
async def set_model(self, model: str):
|
||||||
|
await super().set_model(model)
|
||||||
|
self._settings["model"] = model
|
||||||
|
# Recreate stream with new model
|
||||||
|
if self._streaming_task:
|
||||||
|
await self._disconnect()
|
||||||
|
await self._connect()
|
||||||
|
|
||||||
|
async def start(self, frame: StartFrame):
|
||||||
|
await super().start(frame)
|
||||||
|
await self._connect()
|
||||||
|
|
||||||
|
async def stop(self, frame: EndFrame):
|
||||||
|
await super().stop(frame)
|
||||||
|
await self._disconnect()
|
||||||
|
|
||||||
|
async def cancel(self, frame: CancelFrame):
|
||||||
|
await super().cancel(frame)
|
||||||
|
await self._disconnect()
|
||||||
|
|
||||||
|
async def _connect(self):
|
||||||
|
"""Initialize streaming recognition config and stream"""
|
||||||
|
logger.debug("Connecting to Google Speech-to-Text")
|
||||||
|
|
||||||
|
# Create recognition config with explicit audio format
|
||||||
|
self._config = cloud_speech.StreamingRecognitionConfig(
|
||||||
|
config=cloud_speech.RecognitionConfig(
|
||||||
|
explicit_decoding_config=cloud_speech.ExplicitDecodingConfig(
|
||||||
|
encoding=cloud_speech.ExplicitDecodingConfig.AudioEncoding.LINEAR16,
|
||||||
|
sample_rate_hertz=self.sample_rate,
|
||||||
|
audio_channel_count=1,
|
||||||
|
),
|
||||||
|
language_codes=[self._settings["language_code"]],
|
||||||
|
model=self._settings["model"],
|
||||||
|
features=cloud_speech.RecognitionFeatures(
|
||||||
|
enable_automatic_punctuation=self._settings["enable_automatic_punctuation"],
|
||||||
|
enable_spoken_punctuation=self._settings["enable_spoken_punctuation"],
|
||||||
|
enable_spoken_emojis=self._settings["enable_spoken_emojis"],
|
||||||
|
profanity_filter=self._settings["profanity_filter"],
|
||||||
|
enable_word_time_offsets=self._settings["enable_word_time_offsets"],
|
||||||
|
enable_word_confidence=self._settings["enable_word_confidence"],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Start the streaming task using task manager
|
||||||
|
self._streaming_task = self.create_task(self._stream_audio())
|
||||||
|
|
||||||
|
async def _disconnect(self):
|
||||||
|
"""Clean up streaming recognition resources"""
|
||||||
|
if self._streaming_task:
|
||||||
|
logger.debug("Disconnecting from Google Speech-to-Text")
|
||||||
|
# Send sentinel value to stop request generator
|
||||||
|
await self._request_queue.put(None)
|
||||||
|
await self.cancel_task(self._streaming_task)
|
||||||
|
self._streaming_task = None
|
||||||
|
# Clear any remaining items in the queue
|
||||||
|
while not self._request_queue.empty():
|
||||||
|
try:
|
||||||
|
self._request_queue.get_nowait()
|
||||||
|
self._request_queue.task_done()
|
||||||
|
except asyncio.QueueEmpty:
|
||||||
|
break
|
||||||
|
|
||||||
|
async def _request_generator(self):
|
||||||
|
"""Generates requests for the streaming recognize method."""
|
||||||
|
recognizer_path = f"projects/{self._project_id}/locations/{self._location}/recognizers/_"
|
||||||
|
logger.debug(f"Using recognizer path: {recognizer_path}")
|
||||||
|
|
||||||
|
try:
|
||||||
|
# First, send the recognition config
|
||||||
|
config_request = cloud_speech.StreamingRecognizeRequest(
|
||||||
|
recognizer=recognizer_path,
|
||||||
|
streaming_config=self._config,
|
||||||
|
)
|
||||||
|
yield config_request
|
||||||
|
|
||||||
|
# Then send all audio data requests
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
audio_data = await self._request_queue.get()
|
||||||
|
if audio_data is None: # Sentinel value to stop
|
||||||
|
break
|
||||||
|
yield cloud_speech.StreamingRecognizeRequest(audio=audio_data)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
break
|
||||||
|
finally:
|
||||||
|
self._request_queue.task_done()
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error in request generator: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def _stream_audio(self):
|
||||||
|
"""Handle bi-directional streaming with Google STT"""
|
||||||
|
try:
|
||||||
|
# Start bi-directional streaming
|
||||||
|
streaming_recognize = await self._client.streaming_recognize(
|
||||||
|
requests=self._request_generator()
|
||||||
|
)
|
||||||
|
|
||||||
|
# Process responses using task manager
|
||||||
|
response_task = self.create_task(self._process_responses(streaming_recognize))
|
||||||
|
|
||||||
|
# Wait for the response processing to complete
|
||||||
|
await self.wait_for_task(response_task)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error in streaming task: {e}")
|
||||||
|
await self.push_frame(ErrorFrame(str(e)))
|
||||||
|
|
||||||
|
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
||||||
|
"""Process an audio chunk for STT transcription"""
|
||||||
|
if self._streaming_task:
|
||||||
|
# Queue the audio data
|
||||||
|
await self._request_queue.put(audio)
|
||||||
|
yield None
|
||||||
|
|
||||||
|
async def _process_responses(self, streaming_recognize):
|
||||||
|
"""Process streaming recognition responses"""
|
||||||
|
try:
|
||||||
|
async for response in streaming_recognize:
|
||||||
|
if not response.results:
|
||||||
|
continue
|
||||||
|
|
||||||
|
for result in response.results:
|
||||||
|
if not result.alternatives:
|
||||||
|
continue
|
||||||
|
|
||||||
|
transcript = result.alternatives[0].transcript
|
||||||
|
if not transcript:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if result.is_final:
|
||||||
|
await self.push_frame(
|
||||||
|
TranscriptionFrame(
|
||||||
|
transcript, "", time_now_iso8601(), self._settings["language_code"]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await self.push_frame(
|
||||||
|
InterimTranscriptionFrame(
|
||||||
|
transcript, "", time_now_iso8601(), self._settings["language_code"]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error processing Google STT responses: {e}")
|
||||||
|
await self.push_frame(ErrorFrame(str(e)))
|
||||||
|
|||||||
Reference in New Issue
Block a user