from __future__ import annotations import asyncio from collections.abc import AsyncIterable from typing import TYPE_CHECKING, Any from .. import utils from ..types import DEFAULT_API_CONNECT_OPTIONS, NOT_GIVEN, APIConnectOptions, NotGivenOr from ..vad import VAD, VADEventType from .stt import STT, RecognizeStream, SpeechEvent, SpeechEventType, STTCapabilities if TYPE_CHECKING: from ..voice.events import ConversationItemAddedEvent # already a retry mechanism in STT.recognize, don't retry in stream adapter DEFAULT_STREAM_ADAPTER_API_CONNECT_OPTIONS = APIConnectOptions( max_retry=0, timeout=DEFAULT_API_CONNECT_OPTIONS.timeout ) class StreamAdapter(STT): def __init__(self, *, stt: STT, vad: VAD) -> None: super().__init__( capabilities=STTCapabilities( streaming=True, interim_results=False, diarization=False, # diarization requires streaming STT keyterms=stt.capabilities.keyterms, chat_context=stt.capabilities.chat_context, ) ) self._vad = vad self._stt = stt # TODO(theomonnom): The segment_id needs to be populated! self._stt.on("metrics_collected", self._on_metrics_collected) @property def wrapped_stt(self) -> STT: return self._stt @property def model(self) -> str: return self._stt.model @property def provider(self) -> str: return self._stt.provider def _update_session_keyterms(self, keyterms: list[str]) -> None: self._stt._update_session_keyterms(keyterms) def _push_conversation_item(self, item: ConversationItemAddedEvent) -> None: self._stt._push_conversation_item(item) async def _recognize_impl( self, buffer: utils.AudioBuffer, *, language: NotGivenOr[str] = NOT_GIVEN, conn_options: APIConnectOptions = DEFAULT_API_CONNECT_OPTIONS, ) -> SpeechEvent: return await self._stt.recognize( buffer=buffer, language=language, conn_options=conn_options ) def stream( self, *, language: NotGivenOr[str] = NOT_GIVEN, conn_options: APIConnectOptions = DEFAULT_API_CONNECT_OPTIONS, ) -> RecognizeStream: return StreamAdapterWrapper( self, vad=self._vad, wrapped_stt=self._stt, language=language, conn_options=conn_options, ) def _on_metrics_collected(self, *args: Any, **kwargs: Any) -> None: self.emit("metrics_collected", *args, **kwargs) async def aclose(self) -> None: self._stt.off("metrics_collected", self._on_metrics_collected) class StreamAdapterWrapper(RecognizeStream): def __init__( self, stt: STT, *, vad: VAD, wrapped_stt: STT, language: NotGivenOr[str], conn_options: APIConnectOptions, ) -> None: super().__init__(stt=stt, conn_options=DEFAULT_STREAM_ADAPTER_API_CONNECT_OPTIONS) self._vad = vad self._wrapped_stt = wrapped_stt self._wrapped_stt_conn_options = conn_options self._language = language async def _metrics_monitor_task(self, event_aiter: AsyncIterable[SpeechEvent]) -> None: async for _ in event_aiter: pass async def _run(self) -> None: vad_stream = self._vad.stream() async def _forward_input() -> None: """forward input to vad""" async for input in self._input_ch: if isinstance(input, self._FlushSentinel): vad_stream.flush() continue vad_stream.push_frame(input) vad_stream.end_input() async def _recognize() -> None: """recognize speech from vad""" async for event in vad_stream: if event.type == VADEventType.START_OF_SPEECH: self._event_ch.send_nowait(SpeechEvent(SpeechEventType.START_OF_SPEECH)) elif event.type == VADEventType.END_OF_SPEECH: self._event_ch.send_nowait( SpeechEvent( type=SpeechEventType.END_OF_SPEECH, ) ) merged_frames = utils.merge_frames(event.frames) t_event = await self._wrapped_stt.recognize( buffer=merged_frames, language=self._language, conn_options=self._wrapped_stt_conn_options, ) if len(t_event.alternatives) == 0: continue elif not t_event.alternatives[0].text: continue self._event_ch.send_nowait( SpeechEvent( type=SpeechEventType.FINAL_TRANSCRIPT, alternatives=[t_event.alternatives[0]], ) ) tasks = [ asyncio.create_task(_forward_input(), name="forward_input"), asyncio.create_task(_recognize(), name="recognize"), ] try: await asyncio.gather(*tasks) finally: await utils.aio.cancel_and_wait(*tasks) await vad_stream.aclose()