Files
2026-07-13 13:39:38 +08:00

476 lines
17 KiB
Python

from __future__ import annotations
import asyncio
import contextlib
import dataclasses
import time
from collections.abc import AsyncIterable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Literal
from livekit import rtc
from .. import utils
from .._exceptions import APIConnectionError, APIError
from ..log import logger
from ..types import DEFAULT_API_CONNECT_OPTIONS, NOT_GIVEN, APIConnectOptions, NotGivenOr
from ..utils import aio
from ..utils.audio import AudioBuffer
from ..vad import VAD
from .stt import STT, RecognizeStream, SpeechEvent, SpeechEventType, STTCapabilities
if TYPE_CHECKING:
from ..voice.events import ConversationItemAddedEvent
# don't retry when using the fallback adapter
DEFAULT_FALLBACK_API_CONNECT_OPTIONS = APIConnectOptions(
max_retry=0, timeout=DEFAULT_API_CONNECT_OPTIONS.timeout
)
@dataclass
class AvailabilityChangedEvent:
stt: STT
available: bool
@dataclass
class _STTStatus:
available: bool
recovering_recognize_task: asyncio.Task[None] | None
recovering_stream_task: asyncio.Task[None] | None
class FallbackAdapter(
STT[Literal["stt_availability_changed"]],
):
"""Agent Fallback Adapter for STT. Manages multiple STT instances with automatic fallback
when the primary provider fails.
"""
def __init__(
self,
stt: list[STT],
*,
vad: VAD | None = None,
attempt_timeout: float = 10.0,
max_retry_per_stt: int = 1,
retry_interval: float = 5,
) -> None:
if len(stt) < 1:
raise ValueError("At least one STT instance must be provided.")
non_streaming_stt = [t for t in stt if not t.capabilities.streaming]
if non_streaming_stt:
if vad is None:
labels = ", ".join(t.label for t in non_streaming_stt)
raise ValueError(
f"STTs do not support streaming: {labels}. "
"Provide a VAD to enable stt.StreamAdapter automatically "
"or wrap them with stt.StreamAdapter before using this adapter."
)
from ..stt import StreamAdapter
stt = [
StreamAdapter(stt=t, vad=vad) if not t.capabilities.streaming else t for t in stt
]
# Use the primary STT's aligned_transcript if all providers support it, since
# the SDK only checks truthiness, not the specific granularity.
aligned_transcript: Literal["word", "chunk", False] = False
if all(t.capabilities.aligned_transcript for t in stt):
aligned_transcript = stt[0].capabilities.aligned_transcript
super().__init__(
capabilities=STTCapabilities(
streaming=True,
interim_results=all(t.capabilities.interim_results for t in stt),
diarization=all(t.capabilities.diarization for t in stt),
aligned_transcript=aligned_transcript,
keyterms=any(t.capabilities.keyterms for t in stt),
chat_context=any(t.capabilities.chat_context for t in stt),
)
)
self._stt_instances = stt
self._attempt_timeout = attempt_timeout
self._max_retry_per_stt = max_retry_per_stt
self._retry_interval = retry_interval
self._status: list[_STTStatus] = [
_STTStatus(
available=True,
recovering_recognize_task=None,
recovering_stream_task=None,
)
for _ in self._stt_instances
]
for stt_instance in self._stt_instances:
stt_instance.on("metrics_collected", self._on_metrics_collected)
self._recognize_metrics_needed = False # don't emit metrics via fallback adapter
@property
def model(self) -> str:
return "FallbackAdapter"
@property
def provider(self) -> str:
return "livekit"
def _update_session_keyterms(self, keyterms: list[str]) -> None:
# forward to every underlying STT; unsupported ones warn-and-skip internally
for stt_instance in self._stt_instances:
stt_instance._update_session_keyterms(keyterms)
def _push_conversation_item(self, ev: ConversationItemAddedEvent) -> None:
# forward to every underlying STT; unsupported ones warn-and-skip internally
for stt_instance in self._stt_instances:
stt_instance._push_conversation_item(ev)
async def _try_recognize(
self,
*,
stt: STT,
buffer: utils.AudioBuffer,
language: NotGivenOr[str] = NOT_GIVEN,
conn_options: APIConnectOptions,
recovering: bool = False,
) -> SpeechEvent:
try:
return await stt.recognize(
buffer,
language=language,
conn_options=dataclasses.replace(
conn_options,
max_retry=self._max_retry_per_stt,
timeout=self._attempt_timeout,
retry_interval=self._retry_interval,
),
)
except asyncio.TimeoutError:
if recovering:
logger.warning(f"{stt.label} recovery timed out", extra={"streamed": False})
raise
logger.warning(
f"{stt.label} timed out, switching to next STT",
extra={"streamed": False},
)
raise
except APIError as e:
if recovering:
logger.warning(
"%s recovery failed: %s",
stt.label,
e,
extra={"streamed": False},
)
raise
logger.warning(
"%s failed, switching to next STT: %s",
stt.label,
e,
extra={"streamed": False},
)
raise
except Exception:
if recovering:
logger.exception(
f"{stt.label} recovery unexpected error", extra={"streamed": False}
)
raise
logger.exception(
f"{stt.label} unexpected error, switching to next STT",
extra={"streamed": False},
)
raise
def _try_recovery(
self,
*,
stt: STT,
buffer: utils.AudioBuffer,
language: NotGivenOr[str],
conn_options: APIConnectOptions,
) -> None:
stt_status = self._status[self._stt_instances.index(stt)]
if (
stt_status.recovering_recognize_task is None
or stt_status.recovering_recognize_task.done()
):
async def _recover_stt_task(stt: STT) -> None:
try:
await self._try_recognize(
stt=stt,
buffer=buffer,
language=language,
conn_options=conn_options,
recovering=True,
)
stt_status.available = True
logger.info(f"{stt.label} recovered")
self.emit(
"stt_availability_changed",
AvailabilityChangedEvent(stt=stt, available=True),
)
except Exception as e:
logger.debug("%s recovery attempt failed: %s", stt.label, e)
return
stt_status.recovering_recognize_task = asyncio.create_task(_recover_stt_task(stt))
async def _recognize_impl(
self,
buffer: utils.AudioBuffer,
*,
language: NotGivenOr[str] = NOT_GIVEN,
conn_options: APIConnectOptions,
) -> SpeechEvent:
start_time = time.time()
all_failed = all(not stt_status.available for stt_status in self._status)
if all_failed:
logger.error("all STTs are unavailable, retrying..")
for i, stt in enumerate(self._stt_instances):
stt_status = self._status[i]
if stt_status.available or all_failed:
try:
return await self._try_recognize(
stt=stt,
buffer=buffer,
language=language,
conn_options=conn_options,
recovering=False,
)
except Exception: # exceptions already logged inside _try_recognize
if stt_status.available:
stt_status.available = False
self.emit(
"stt_availability_changed",
AvailabilityChangedEvent(stt=stt, available=False),
)
self._try_recovery(stt=stt, buffer=buffer, language=language, conn_options=conn_options)
raise APIConnectionError(
f"all STTs failed ({[stt.label for stt in self._stt_instances]}) after {time.time() - start_time} seconds" # noqa: E501
)
async def recognize(
self,
buffer: AudioBuffer,
*,
language: NotGivenOr[str] = NOT_GIVEN,
conn_options: APIConnectOptions = DEFAULT_FALLBACK_API_CONNECT_OPTIONS,
) -> SpeechEvent:
return await super().recognize(buffer, language=language, conn_options=conn_options)
def stream(
self,
*,
language: NotGivenOr[str] = NOT_GIVEN,
conn_options: APIConnectOptions = DEFAULT_FALLBACK_API_CONNECT_OPTIONS,
) -> RecognizeStream:
return FallbackRecognizeStream(stt=self, language=language, conn_options=conn_options)
async def aclose(self) -> None:
for stt_status in self._status:
if stt_status.recovering_recognize_task is not None:
await aio.cancel_and_wait(stt_status.recovering_recognize_task)
if stt_status.recovering_stream_task is not None:
await aio.cancel_and_wait(stt_status.recovering_stream_task)
for stt in self._stt_instances:
stt.off("metrics_collected", self._on_metrics_collected)
def _on_metrics_collected(self, *args: Any, **kwargs: Any) -> None:
self.emit("metrics_collected", *args, **kwargs)
class FallbackRecognizeStream(RecognizeStream):
def __init__(
self,
*,
stt: FallbackAdapter,
language: NotGivenOr[str] = NOT_GIVEN,
conn_options: APIConnectOptions,
):
super().__init__(stt=stt, conn_options=conn_options, sample_rate=NOT_GIVEN)
self._language = language
self._fallback_adapter = stt
self._recovering_streams: list[RecognizeStream] = []
async def _run(self) -> None:
start_time = time.time()
all_failed = all(not stt_status.available for stt_status in self._fallback_adapter._status)
if all_failed:
logger.error("all STTs are unavailable, retrying..")
main_stream: RecognizeStream | None = None
forward_input_task: asyncio.Task[None] | None = None
async def _forward_input_task() -> None:
async for data in self._input_ch:
for stream in list(self._recovering_streams):
try:
if isinstance(data, rtc.AudioFrame):
stream.push_frame(data)
elif isinstance(data, self._FlushSentinel):
stream.flush()
except Exception:
pass
if main_stream is not None:
try:
if isinstance(data, rtc.AudioFrame):
main_stream.push_frame(data)
elif isinstance(data, self._FlushSentinel):
main_stream.flush()
except Exception:
logger.exception(
"error happened in forwarding input", extra={"streamed": True}
)
if main_stream is not None:
with contextlib.suppress(RuntimeError):
main_stream.end_input()
for i, stt in enumerate(self._fallback_adapter._stt_instances):
stt_status = self._fallback_adapter._status[i]
if stt_status.available or all_failed:
try:
main_stream = stt.stream(
language=self._language,
conn_options=dataclasses.replace(
self._conn_options,
max_retry=self._fallback_adapter._max_retry_per_stt,
timeout=self._fallback_adapter._attempt_timeout,
retry_interval=self._fallback_adapter._retry_interval,
),
)
# update main_stream start time offset so transcript timestamps are properly adjusted
main_stream.start_time_offset = self.start_time_offset + (
time.time() - self._start_time
)
if forward_input_task is None or forward_input_task.done():
forward_input_task = asyncio.create_task(_forward_input_task())
try:
async with main_stream:
async for ev in main_stream:
self._event_ch.send_nowait(ev)
except asyncio.TimeoutError:
logger.warning(
f"{stt.label} timed out, switching to next STT",
extra={"streamed": True},
)
raise
except APIError as e:
logger.warning(
"%s failed, switching to next STT: %s",
stt.label,
e,
extra={"streamed": True},
)
raise
except Exception:
logger.exception(
f"{stt.label} unexpected error, switching to next STT",
extra={"streamed": True},
)
raise
return
except Exception:
if stt_status.available:
stt_status.available = False
self._stt.emit(
"stt_availability_changed",
AvailabilityChangedEvent(stt=stt, available=False),
)
self._try_recovery(stt)
if forward_input_task is not None:
await aio.cancel_and_wait(forward_input_task)
await asyncio.gather(*[stream.aclose() for stream in self._recovering_streams])
raise APIConnectionError(
f"all STTs failed ({[stt.label for stt in self._fallback_adapter._stt_instances]}) after {time.time() - start_time} seconds" # noqa: E501
)
def _try_recovery(self, stt: STT) -> None:
stt_status = self._fallback_adapter._status[
self._fallback_adapter._stt_instances.index(stt)
]
if stt_status.recovering_stream_task is None or stt_status.recovering_stream_task.done():
stream = stt.stream(
language=self._language,
conn_options=dataclasses.replace(
self._conn_options,
max_retry=0,
timeout=self._fallback_adapter._attempt_timeout,
),
)
self._recovering_streams.append(stream)
async def _recover_stt_task() -> None:
try:
nb_transcript = 0
async with stream:
async for ev in stream:
if ev.type == SpeechEventType.FINAL_TRANSCRIPT:
if not ev.alternatives or not ev.alternatives[0].text:
continue
nb_transcript += 1
break
if nb_transcript == 0:
return
stt_status.available = True
logger.info(f"stt.FallbackAdapter, {stt.label} recovered")
self._fallback_adapter.emit(
"stt_availability_changed",
AvailabilityChangedEvent(stt=stt, available=True),
)
except asyncio.TimeoutError:
logger.warning(
f"{stream._stt.label} recovery timed out",
extra={"streamed": True},
)
except APIError as e:
logger.warning(
"%s recovery failed: %s",
stream._stt.label,
e,
extra={"streamed": True},
)
except Exception:
logger.exception(
f"{stream._stt.label} recovery unexpected error",
extra={"streamed": True},
)
raise
stt_status.recovering_stream_task = task = asyncio.create_task(_recover_stt_task())
task.add_done_callback(lambda _: self._recovering_streams.remove(stream))
async def _metrics_monitor_task(self, event_aiter: AsyncIterable[SpeechEvent]) -> None:
async for _ in event_aiter:
pass