services(deepgram): inherit from STTService instead of AsyncAIService

This commit is contained in:
Aleix Conchillo Flaqué
2024-08-26 11:10:42 -07:00
parent 3931cb3235
commit 6629b853c5

View File

@@ -16,12 +16,10 @@ from pipecat.frames.frames import (
Frame, Frame,
InterimTranscriptionFrame, InterimTranscriptionFrame,
StartFrame, StartFrame,
SystemFrame,
TTSStartedFrame, TTSStartedFrame,
TTSStoppedFrame, TTSStoppedFrame,
TranscriptionFrame) TranscriptionFrame)
from pipecat.processors.frame_processor import FrameDirection from pipecat.services.ai_services import STTService, TTSService
from pipecat.services.ai_services import AsyncAIService, TTSService
from pipecat.transcriptions.language import Language from pipecat.transcriptions.language import Language
from pipecat.utils.time import time_now_iso8601 from pipecat.utils.time import time_now_iso8601
@@ -45,15 +43,6 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def deepgram_language_to_language(language: str) -> Language | None:
match language:
case "en":
return Language.EN
case "es":
return Language.ES
return None
class DeepgramTTSService(TTSService): class DeepgramTTSService(TTSService):
def __init__( def __init__(
@@ -119,7 +108,7 @@ class DeepgramTTSService(TTSService):
logger.exception(f"{self} exception: {e}") logger.exception(f"{self} exception: {e}")
class DeepgramSTTService(AsyncAIService): class DeepgramSTTService(STTService):
def __init__(self, def __init__(self,
*, *,
api_key: str, api_key: str,
@@ -132,6 +121,8 @@ class DeepgramSTTService(AsyncAIService):
channels=1, channels=1,
interim_results=True, interim_results=True,
smart_format=True, smart_format=True,
punctuate=True,
profanity_filter=True,
), ),
**kwargs): **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
@@ -143,30 +134,46 @@ class DeepgramSTTService(AsyncAIService):
self._connection: AsyncListenWebSocketClient = self._client.listen.asyncwebsocket.v("1") self._connection: AsyncListenWebSocketClient = self._client.listen.asyncwebsocket.v("1")
self._connection.on(LiveTranscriptionEvents.Transcript, self._on_message) self._connection.on(LiveTranscriptionEvents.Transcript, self._on_message)
async def process_frame(self, frame: Frame, direction: FrameDirection): async def set_model(self, model: str):
await super().process_frame(frame, direction) logger.debug(f"Switching STT model to: [{model}]")
self._live_options.model = model
await self._disconnect()
await self._connect()
if isinstance(frame, SystemFrame): async def set_language(self, language: Language):
await self.push_frame(frame, direction) logger.debug(f"Switching STT language to: [{language}]")
elif isinstance(frame, AudioRawFrame): self._live_options.language = language.value
await self._connection.send(frame.audio) await self._disconnect()
else: await self._connect()
await self.queue_frame(frame, direction)
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) 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 run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
await self.start_processing_metrics()
await self._connection.send(audio)
yield None
await self.stop_processing_metrics()
async def _connect(self):
if await self._connection.start(self._live_options): if await self._connection.start(self._live_options):
logger.debug(f"{self}: Connected to Deepgram") logger.debug(f"{self}: Connected to Deepgram")
else: else:
logger.error(f"{self}: Unable to connect to Deepgram") logger.error(f"{self}: Unable to connect to Deepgram")
async def stop(self, frame: EndFrame): async def _disconnect(self):
await super().stop(frame) if self._connection.is_connected:
await self._connection.finish() await self._connection.finish()
logger.debug(f"{self}: Disconnected from Deepgram")
async def cancel(self, frame: CancelFrame):
await super().cancel(frame)
await self._connection.finish()
async def _on_message(self, *args, **kwargs): async def _on_message(self, *args, **kwargs):
result: LiveResultResponse = kwargs["result"] result: LiveResultResponse = kwargs["result"]
@@ -177,9 +184,9 @@ class DeepgramSTTService(AsyncAIService):
language = None language = None
if result.channel.alternatives[0].languages: if result.channel.alternatives[0].languages:
language = result.channel.alternatives[0].languages[0] language = result.channel.alternatives[0].languages[0]
language = deepgram_language_to_language(language) language = Language(language)
if len(transcript) > 0: if len(transcript) > 0:
if is_final: if is_final:
await self.queue_frame(TranscriptionFrame(transcript, "", time_now_iso8601(), language)) await self.push_frame(TranscriptionFrame(transcript, "", time_now_iso8601(), language))
else: else:
await self.queue_frame(InterimTranscriptionFrame(transcript, "", time_now_iso8601(), language)) await self.push_frame(InterimTranscriptionFrame(transcript, "", time_now_iso8601(), language))