livekit--agents
161 行
5.3 KiB
Python
161 行
5.3 KiB
Python
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()
|