from __future__ import annotations import asyncio import audioop import base64 import contextlib import io import json import logging import os import time import wave from collections.abc import Sequence from urllib.parse import urlencode from realtime_voice_service.providers.base import BaseSTT from realtime_voice_service.providers.base import BaseSTTStream from realtime_voice_service.providers.base import PartialTranscriptCallback from realtime_voice_service.providers.stt_openai import OpenAISTT LOGGER = logging.getLogger("uvicorn.error") _STREAM_COMMIT = object() _STREAM_CANCEL = object() _REALTIME_MAX_KEYTERMS = 50 _REALTIME_MAX_KEYTERM_CHARS = 20 _BATCH_MAX_KEYTERMS = 1000 _BATCH_MAX_KEYTERM_CHARS = 50 def _normalize_keyterms( keyterms: Sequence[str] | None, *, max_count: int, max_chars: int, ) -> list[str]: if not keyterms: return [] seen: set[str] = set() normalized: list[str] = [] for raw in keyterms: term = " ".join(str(raw or "").split()) if not term or len(term) > max_chars: continue if term in seen: continue seen.add(term) normalized.append(term) if len(normalized) >= max_count: break return normalized def _normalize_elevenlabs_language_code(language_code: str | None) -> str | None: normalized = str(language_code or "").strip() if not normalized: return None lowered = normalized.lower().replace("_", "-") mapping = { "kz": "kaz", "kk": "kaz", "kaz": "kaz", "kk-kz": "kaz", "kz-kz": "kaz", "ru": "rus", "rus": "rus", "ru-ru": "rus", } return mapping.get(lowered, normalized) def _normalize_yandex_language_code(language_code: str | None) -> str | None: normalized = str(language_code or "").strip() if not normalized: return None lowered = normalized.lower().replace("_", "-") mapping = { "ru": "ru-RU", "rus": "ru-RU", "ru-ru": "ru-RU", "kk": "kk-KZ", "kaz": "kk-KZ", "kz": "kk-KZ", "kk-kz": "kk-KZ", "kz-kz": "kk-KZ", } return mapping.get(lowered, normalized) def _api_base() -> str: return (os.getenv("ELEVENLABS_API_BASE", "https://api.elevenlabs.io").strip() or "https://api.elevenlabs.io").rstrip("/") def _stt_provider_name() -> str: return ( os.getenv("STT_PROVIDER", "").strip() or os.getenv("REALTIME_VOICE_STT_PROVIDER", "").strip() or os.getenv("AI_VOICE_ASR_PROVIDER", "").strip() or "elevenlabs" ).lower() def _stt_fallback_provider_name() -> str: return os.getenv("REALTIME_VOICE_STT_FALLBACK_PROVIDER", "").strip().lower() def _yandex_api_base() -> str: return ( os.getenv("YANDEX_STT_API_BASE", "").strip() or os.getenv("AI_VOICE_ASR_YANDEX_API_BASE", "").strip() or "https://stt.api.cloud.yandex.net" ).rstrip("/") def _yandex_api_key() -> str: return ( os.getenv("YANDEX_STT_API_KEY", "").strip() or os.getenv("AI_VOICE_ASR_YANDEX_API_KEY", "").strip() or os.getenv("AI_VOICE_TTS_YANDEX_API_KEY", "").strip() ) def _yandex_iam_token() -> str: return ( os.getenv("YANDEX_STT_IAM_TOKEN", "").strip() or os.getenv("AI_VOICE_ASR_YANDEX_IAM_TOKEN", "").strip() or os.getenv("AI_VOICE_TTS_YANDEX_IAM_TOKEN", "").strip() ) def _yandex_folder_id() -> str: return ( os.getenv("YANDEX_STT_FOLDER_ID", "").strip() or os.getenv("AI_VOICE_ASR_YANDEX_FOLDER_ID", "").strip() or os.getenv("AI_VOICE_TTS_YANDEX_FOLDER_ID", "").strip() ) def _yandex_language() -> str: raw = ( os.getenv("YANDEX_STT_LANGUAGE", "").strip() or os.getenv("AI_VOICE_ASR_YANDEX_LANGUAGE", "").strip() or "ru-RU" ) return _normalize_yandex_language_code(raw) or "ru-RU" def _yandex_topic() -> str: return os.getenv("YANDEX_STT_TOPIC", "").strip() or os.getenv("AI_VOICE_ASR_YANDEX_TOPIC", "general").strip() or "general" def _ws_base(http_base: str) -> str: normalized = http_base.rstrip("/") if normalized.startswith("https://"): return f"wss://{normalized[len('https://') :]}" if normalized.startswith("http://"): return f"ws://{normalized[len('http://') :]}" return normalized def _timeout_seconds() -> float: raw = os.getenv("ELEVENLABS_TIMEOUT_SECONDS") if raw is None: return 30.0 try: return max(float(raw.strip()), 1.0) except ValueError: return 30.0 def _bool_env(name: str, default: bool) -> bool: raw = os.getenv(name) if raw is None: return default return str(raw).strip().lower() in {"1", "true", "yes", "on"} def _preview_text(text: str, *, limit: int = 120) -> str: normalized = " ".join(str(text or "").split()) if len(normalized) <= limit: return normalized return f"{normalized[:limit]}..." def _pcm16le_to_wav_bytes( pcm_bytes: bytes, *, sample_rate_hz: int, channels: int = 1, sample_width_bytes: int = 2, ) -> bytes: handle = io.BytesIO() with wave.open(handle, "wb") as wav_file: wav_file.setnchannels(channels) wav_file.setsampwidth(sample_width_bytes) wav_file.setframerate(sample_rate_hz) wav_file.writeframes(pcm_bytes) return handle.getvalue() def _resample_pcm16le( pcm_bytes: bytes, *, input_rate_hz: int, output_rate_hz: int, ) -> bytes: if not pcm_bytes or input_rate_hz == output_rate_hz: return pcm_bytes converted, _ = audioop.ratecv( pcm_bytes, 2, 1, input_rate_hz, output_rate_hz, None, ) return converted def _pcm_duration_ms(pcm_bytes: bytes, *, sample_rate_hz: int) -> int: if not pcm_bytes or sample_rate_hz <= 0: return 0 sample_count = len(pcm_bytes) // 2 return max(int((sample_count / float(sample_rate_hz)) * 1000.0), 0) def _pad_pcm16le_to_duration( pcm_bytes: bytes, *, sample_rate_hz: int, min_duration_ms: int, ) -> bytes: if not pcm_bytes: return pcm_bytes current_duration_ms = _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz) if current_duration_ms >= min_duration_ms: return pcm_bytes target_samples = max(int(sample_rate_hz * (min_duration_ms / 1000.0)), 1) target_bytes = target_samples * 2 if len(pcm_bytes) >= target_bytes: return pcm_bytes return pcm_bytes + (b"\x00" * (target_bytes - len(pcm_bytes))) class ElevenLabsRealtimeSTTStream(BaseSTTStream): def __init__( self, *, api_key: str, websocket_url: str, sample_rate_hz: int, timeout_seconds: float, partial_callback: PartialTranscriptCallback | None = None, ) -> None: self._api_key = api_key self._websocket_url = websocket_url self._sample_rate_hz = sample_rate_hz self._timeout_seconds = timeout_seconds self._partial_callback = partial_callback self._websocket = None self._sender_task: asyncio.Task[None] | None = None self._receiver_task: asyncio.Task[None] | None = None self._outgoing_queue: asyncio.Queue[bytes | object] = asyncio.Queue() self._final_transcript_future: asyncio.Future[str] = asyncio.get_running_loop().create_future() self._close_lock = asyncio.Lock() self._closed = False self._commit_requested = False self._audio_sent = False self._latest_partial = "" self._latest_committed = "" self._connect_started_monotonic: float | None = None self._sent_chunk_count = 0 self._sent_byte_count = 0 self._partial_count = 0 self._committed_count = 0 async def connect(self) -> None: try: from websockets.legacy.client import connect as websocket_connect except Exception as exc: # noqa: BLE001 raise RuntimeError("The `websockets` package is required for ElevenLabs realtime STT") from exc self._connect_started_monotonic = time.perf_counter() LOGGER.info( "ElevenLabs realtime STT connecting: sample_rate=%s timeout=%s", self._sample_rate_hz, self._timeout_seconds, ) self._websocket = await websocket_connect( self._websocket_url, extra_headers=[("xi-api-key", self._api_key)], open_timeout=min(self._timeout_seconds, 10.0), close_timeout=1.0, ping_interval=20.0, ping_timeout=20.0, max_size=None, ) try: await self._await_session_started() except Exception: await self._close_websocket() raise LOGGER.info( "ElevenLabs realtime STT connected: sample_rate=%s connect_ms=%s", self._sample_rate_hz, int((time.perf_counter() - self._connect_started_monotonic) * 1000.0), ) self._sender_task = asyncio.create_task(self._sender_loop(), name="elevenlabs-stt-sender") self._receiver_task = asyncio.create_task(self._receiver_loop(), name="elevenlabs-stt-receiver") async def push_audio(self, audio_chunk: bytes) -> None: if self._closed or not audio_chunk: return await self._outgoing_queue.put(audio_chunk) async def finish(self) -> str: if self._closed: return self._latest_committed or self._latest_partial self._commit_requested = True finished_started_monotonic = time.perf_counter() LOGGER.info( "ElevenLabs realtime STT finish requested: sent_chunks=%s sent_bytes=%s latest_partial=%r", self._sent_chunk_count, self._sent_byte_count, _preview_text(self._latest_partial), ) await self._outgoing_queue.put(_STREAM_COMMIT) try: transcript = await asyncio.wait_for( asyncio.shield(self._final_transcript_future), timeout=self._timeout_seconds, ) LOGGER.info( "ElevenLabs realtime STT final transcript: latency_ms=%s chars=%s transcript=%r", int((time.perf_counter() - finished_started_monotonic) * 1000.0), len(transcript), _preview_text(transcript), ) return transcript except asyncio.TimeoutError as exc: transcript = self._latest_committed or self._latest_partial if transcript: LOGGER.warning("ElevenLabs realtime STT finish timed out; returning best-effort transcript") return transcript raise RuntimeError("ElevenLabs realtime STT finish timed out") from exc finally: await self._shutdown() async def cancel(self) -> None: if self._closed: return LOGGER.info( "ElevenLabs realtime STT cancel requested: sent_chunks=%s sent_bytes=%s partials=%s commits=%s", self._sent_chunk_count, self._sent_byte_count, self._partial_count, self._committed_count, ) await self._outgoing_queue.put(_STREAM_CANCEL) await self._shutdown() async def _sender_loop(self) -> None: pending_chunk: bytes | None = None try: while True: item = await self._outgoing_queue.get() if item is _STREAM_CANCEL: return if item is _STREAM_COMMIT: if pending_chunk is not None: await self._send_audio_chunk(pending_chunk, commit=True) pending_chunk = None elif not self._final_transcript_future.done(): self._final_transcript_future.set_result(self._latest_committed or self._latest_partial) return audio_chunk = bytes(item) if not audio_chunk: continue self._audio_sent = True if pending_chunk is not None: await self._send_audio_chunk(pending_chunk, commit=False) pending_chunk = audio_chunk except asyncio.CancelledError: raise except Exception as exc: # noqa: BLE001 if not self._final_transcript_future.done(): self._final_transcript_future.set_exception(exc) raise async def _receiver_loop(self) -> None: try: while True: payload = await self._receive_payload() message_type = str(payload.get("message_type") or "").strip().lower() if message_type == "partial_transcript": self._latest_partial = str(payload.get("text") or "").strip() self._partial_count += 1 LOGGER.info( "ElevenLabs realtime STT partial: index=%s chars=%s text=%r", self._partial_count, len(self._latest_partial), _preview_text(self._latest_partial), ) await self._emit_partial(self._latest_partial) continue if message_type in {"committed_transcript", "committed_transcript_with_timestamps"}: self._latest_committed = str(payload.get("text") or "").strip() self._committed_count += 1 LOGGER.info( "ElevenLabs realtime STT committed: index=%s chars=%s text=%r", self._committed_count, len(self._latest_committed), _preview_text(self._latest_committed), ) if self._commit_requested and not self._final_transcript_future.done(): self._final_transcript_future.set_result(self._latest_committed) continue if message_type == "session_started": continue if self._is_error_message(message_type): raise RuntimeError(self._format_realtime_error(payload)) except asyncio.CancelledError: raise except Exception as exc: # noqa: BLE001 if not self._final_transcript_future.done(): transcript = self._latest_committed or self._latest_partial if transcript and self._commit_requested: self._final_transcript_future.set_result(transcript) else: self._final_transcript_future.set_exception(exc) raise async def _send_audio_chunk(self, audio_chunk: bytes, *, commit: bool) -> None: websocket = self._require_websocket() payload = { "message_type": "input_audio_chunk", "audio_base_64": base64.b64encode(audio_chunk).decode("ascii"), "sample_rate": self._sample_rate_hz, "commit": commit, } await websocket.send(json.dumps(payload)) self._sent_chunk_count += 1 self._sent_byte_count += len(audio_chunk) if commit or self._sent_chunk_count == 1 or self._sent_chunk_count % 10 == 0: LOGGER.info( "ElevenLabs realtime STT audio sent: chunks=%s bytes=%s last_chunk_bytes=%s commit=%s", self._sent_chunk_count, self._sent_byte_count, len(audio_chunk), commit, ) async def _await_session_started(self) -> None: while True: payload = await self._receive_payload() message_type = str(payload.get("message_type") or "").strip().lower() if message_type == "session_started": return if self._is_error_message(message_type): raise RuntimeError(self._format_realtime_error(payload)) async def _receive_payload(self) -> dict[str, object]: websocket = self._require_websocket() raw_message = await asyncio.wait_for(websocket.recv(), timeout=self._timeout_seconds) if isinstance(raw_message, bytes): raw_message = raw_message.decode("utf-8") try: payload = json.loads(raw_message) except json.JSONDecodeError as exc: raise RuntimeError("ElevenLabs realtime STT returned invalid JSON") from exc if not isinstance(payload, dict): raise RuntimeError("ElevenLabs realtime STT returned malformed payload") return payload async def _emit_partial(self, text: str) -> None: if not text or self._partial_callback is None: return maybe_awaitable = self._partial_callback(text) if maybe_awaitable is not None: await maybe_awaitable async def _shutdown(self) -> None: async with self._close_lock: if self._closed: return self._closed = True tasks = [self._sender_task, self._receiver_task] for task in tasks: if task is not None and not task.done(): task.cancel() for task in tasks: if task is not None: try: await task except asyncio.CancelledError: pass except Exception: LOGGER.debug("ignored realtime STT background task failure during shutdown", exc_info=True) await self._close_websocket() async def _close_websocket(self) -> None: websocket = self._websocket self._websocket = None if websocket is None: return with contextlib.suppress(Exception): await websocket.close() def _require_websocket(self): if self._websocket is None: raise RuntimeError("ElevenLabs realtime STT websocket is not connected") return self._websocket @staticmethod def _is_error_message(message_type: str) -> bool: return message_type in {"error", "auth_error", "input_error"} or message_type.endswith("_error") @staticmethod def _format_realtime_error(payload: dict[str, object]) -> str: message_type = str(payload.get("message_type") or "error") detail = str(payload.get("message") or payload.get("detail") or payload.get("error") or "").strip() if detail: return f"ElevenLabs realtime STT {message_type}: {detail}" return f"ElevenLabs realtime STT {message_type}" class ElevenLabsSTT(BaseSTT): def __init__( self, *, api_key: str | None = None, api_base: str | None = None, model_id: str | None = None, realtime_model_id: str | None = None, input_sample_rate_hz: int = 16000, target_sample_rate_hz: int = 16000, timeout_seconds: float | None = None, language_code: str | None = None, use_realtime: bool | None = None, allow_batch_fallback: bool | None = None, realtime_chunk_duration_ms: int = 120, ) -> None: self._api_key = str(api_key if api_key is not None else os.getenv("ELEVENLABS_API_KEY", "")).strip() self._api_base = str(api_base or _api_base()).strip().rstrip("/") self._ws_base = _ws_base(self._api_base) self._model_id = str(model_id or os.getenv("ELEVENLABS_STT_MODEL_ID", "scribe_v2")).strip() or "scribe_v2" self._realtime_model_id = ( str(realtime_model_id or os.getenv("ELEVENLABS_STT_REALTIME_MODEL_ID", "scribe_v2_realtime")).strip() or "scribe_v2_realtime" ) self._input_sample_rate_hz = max(int(input_sample_rate_hz), 1) self._target_sample_rate_hz = max(int(target_sample_rate_hz), 1) self._timeout_seconds = max(float(timeout_seconds if timeout_seconds is not None else _timeout_seconds()), 1.0) self._language_code = str(language_code or os.getenv("ELEVENLABS_STT_LANGUAGE_CODE", "ru")).strip() or "ru" self._use_realtime = ( bool(use_realtime) if use_realtime is not None else _bool_env("ELEVENLABS_STT_USE_REALTIME", True) ) self._allow_batch_fallback = ( bool(allow_batch_fallback) if allow_batch_fallback is not None else _bool_env("ELEVENLABS_STT_ALLOW_BATCH_FALLBACK", True) ) self._realtime_chunk_duration_ms = max(int(realtime_chunk_duration_ms), 40) self._batch_min_audio_ms = max(int(os.getenv("ELEVENLABS_STT_BATCH_MIN_AUDIO_MS", "800")), 200) self._client_session = None self._client_session_lock = asyncio.Lock() LOGGER.info( "ElevenLabs STT config: batch_model=%s realtime_model=%s input_sample_rate=%s " "target_sample_rate=%s use_realtime=%s allow_batch_fallback=%s language=%s " "realtime_chunk_ms=%s batch_min_audio_ms=%s", self._model_id, self._realtime_model_id, self._input_sample_rate_hz, self._target_sample_rate_hz, self._use_realtime, self._allow_batch_fallback, self._language_code, self._realtime_chunk_duration_ms, self._batch_min_audio_ms, ) async def transcribe( self, audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, language_code: str | None = None, force_batch: bool = False, ) -> str: if not audio_bytes: return "" if not self._api_key: raise RuntimeError("ELEVENLABS_API_KEY is required for ElevenLabs STT") effective_language_code = _normalize_elevenlabs_language_code(language_code) or self._language_code pcm_bytes, sample_rate_hz = self._extract_pcm(audio_bytes) LOGGER.info( "ElevenLabs STT transcribe start: input_bytes=%s extracted_pcm_bytes=%s sample_rate=%s duration_ms=%s " "use_realtime=%s force_batch=%s language=%s keyterms=%s", len(audio_bytes), len(pcm_bytes), sample_rate_hz, _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz), self._use_realtime, force_batch, effective_language_code, len(keyterms) if keyterms else 0, ) if sample_rate_hz != self._target_sample_rate_hz: before_rate_hz = sample_rate_hz before_bytes = len(pcm_bytes) pcm_bytes = _resample_pcm16le( pcm_bytes, input_rate_hz=sample_rate_hz, output_rate_hz=self._target_sample_rate_hz, ) sample_rate_hz = self._target_sample_rate_hz LOGGER.info( "ElevenLabs STT resampled input: from_rate=%s to_rate=%s before_bytes=%s after_bytes=%s", before_rate_hz, sample_rate_hz, before_bytes, len(pcm_bytes), ) if self._use_realtime and not force_batch: try: transcript = await self._transcribe_realtime( pcm_bytes=pcm_bytes, sample_rate_hz=sample_rate_hz, keyterms=keyterms, language_code=effective_language_code, ) LOGGER.info( "ElevenLabs STT realtime transcript result: chars=%s transcript=%r", len(transcript), _preview_text(transcript), ) return transcript except Exception: if not self._allow_batch_fallback: raise LOGGER.exception("ElevenLabs STT realtime failed; falling back to batch STT") transcript = await self._transcribe_batch( pcm_bytes=pcm_bytes, sample_rate_hz=sample_rate_hz, keyterms=keyterms, language_code=effective_language_code, ) LOGGER.info( "ElevenLabs STT batch transcript result: chars=%s transcript=%r", len(transcript), _preview_text(transcript), ) return transcript async def start_stream( self, *, partial_callback: PartialTranscriptCallback | None = None, keyterms: Sequence[str] | None = None, language_code: str | None = None, ) -> BaseSTTStream | None: if not self._use_realtime: return None if not self._api_key: raise RuntimeError("ELEVENLABS_API_KEY is required for ElevenLabs STT") effective_language_code = _normalize_elevenlabs_language_code(language_code) or self._language_code websocket_url = self._build_realtime_websocket_url( sample_rate_hz=self._target_sample_rate_hz, keyterms=keyterms, language_code=effective_language_code, ) normalized_keyterm_count = len( _normalize_keyterms( keyterms, max_count=_REALTIME_MAX_KEYTERMS, max_chars=_REALTIME_MAX_KEYTERM_CHARS, ) ) LOGGER.info( "ElevenLabs STT live stream starting: realtime_model=%s sample_rate=%s language=%s keyterms=%s", self._realtime_model_id, self._target_sample_rate_hz, effective_language_code, normalized_keyterm_count, ) stream = ElevenLabsRealtimeSTTStream( api_key=self._api_key, websocket_url=websocket_url, sample_rate_hz=self._target_sample_rate_hz, timeout_seconds=self._timeout_seconds, partial_callback=partial_callback, ) await stream.connect() LOGGER.info("ElevenLabs STT live stream started: sample_rate=%s", self._target_sample_rate_hz) return stream async def close(self) -> None: session = self._client_session self._client_session = None if session is not None and not session.closed: await session.close() async def _transcribe_realtime( self, *, pcm_bytes: bytes, sample_rate_hz: int, keyterms: Sequence[str] | None = None, language_code: str | None = None, ) -> str: try: from websockets.exceptions import WebSocketException from websockets.legacy.client import connect as websocket_connect except Exception as exc: # noqa: BLE001 raise RuntimeError("The `websockets` package is required for ElevenLabs realtime STT") from exc websocket_url = self._build_realtime_websocket_url( sample_rate_hz=sample_rate_hz, keyterms=keyterms, language_code=language_code, ) chunk_bytes = max(int(sample_rate_hz * self._realtime_chunk_duration_ms / 1000.0) * 2, 320) started_monotonic = time.perf_counter() sent_chunks = 0 sent_bytes = 0 LOGGER.info( "ElevenLabs STT realtime batch-style stream start: model=%s sample_rate=%s pcm_bytes=%s " "duration_ms=%s chunk_bytes=%s", self._realtime_model_id, sample_rate_hz, len(pcm_bytes), _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz), chunk_bytes, ) try: async with websocket_connect( websocket_url, extra_headers=[("xi-api-key", self._api_key)], open_timeout=min(self._timeout_seconds, 10.0), close_timeout=1.0, ping_interval=20.0, ping_timeout=20.0, max_size=None, ) as websocket: await self._await_realtime_message( websocket, allowed_message_types={"session_started"}, ) for offset in range(0, len(pcm_bytes), chunk_bytes): audio_chunk = pcm_bytes[offset : offset + chunk_bytes] if not audio_chunk: continue payload = { "message_type": "input_audio_chunk", "audio_base_64": base64.b64encode(audio_chunk).decode("ascii"), "sample_rate": sample_rate_hz, "commit": offset + chunk_bytes >= len(pcm_bytes), } await websocket.send(json.dumps(payload)) sent_chunks += 1 sent_bytes += len(audio_chunk) if sent_chunks == 1 or payload["commit"] or sent_chunks % 10 == 0: LOGGER.info( "ElevenLabs STT realtime batch-style audio sent: chunks=%s bytes=%s commit=%s", sent_chunks, sent_bytes, payload["commit"], ) transcript = "" partial_transcript = "" while True: message = await self._receive_realtime_payload(websocket) message_type = str(message.get("message_type") or "").strip().lower() if message_type in {"committed_transcript", "committed_transcript_with_timestamps"}: transcript = str(message.get("text") or "").strip() if transcript: LOGGER.info( "ElevenLabs STT realtime batch-style committed: latency_ms=%s chars=%s text=%r", int((time.perf_counter() - started_monotonic) * 1000.0), len(transcript), _preview_text(transcript), ) return transcript continue if message_type == "partial_transcript": partial_transcript = str(message.get("text") or "").strip() LOGGER.info( "ElevenLabs STT realtime batch-style partial: chars=%s text=%r", len(partial_transcript), _preview_text(partial_transcript), ) continue if message_type == "session_started": continue if self._is_error_message(message_type): raise RuntimeError(self._format_realtime_error(message)) if not message_type: continue if partial_transcript: LOGGER.info( "ElevenLabs STT realtime batch-style returning partial: latency_ms=%s chars=%s", int((time.perf_counter() - started_monotonic) * 1000.0), len(partial_transcript), ) return partial_transcript except asyncio.TimeoutError as exc: raise RuntimeError("ElevenLabs realtime STT request timed out") from exc except WebSocketException as exc: raise RuntimeError("ElevenLabs realtime STT stream failed") from exc except OSError as exc: raise RuntimeError("ElevenLabs realtime STT connection failed") from exc return "" async def _transcribe_batch( self, *, pcm_bytes: bytes, sample_rate_hz: int, keyterms: Sequence[str] | None = None, language_code: str | None = None, ) -> str: # TODO: Architectural Bottleneck: Рассмотреть замену STT на Deepgram WebSocket API для достижения true-streaming latency. try: import aiohttp except Exception as exc: # noqa: BLE001 raise RuntimeError("The `aiohttp` package is required for ElevenLabs STT") from exc pcm_bytes = _pad_pcm16le_to_duration( pcm_bytes, sample_rate_hz=sample_rate_hz, min_duration_ms=self._batch_min_audio_ms, ) wav_bytes = _pcm16le_to_wav_bytes(pcm_bytes, sample_rate_hz=sample_rate_hz) normalized_keyterms = _normalize_keyterms( keyterms, max_count=_BATCH_MAX_KEYTERMS, max_chars=_BATCH_MAX_KEYTERM_CHARS, ) started_monotonic = time.perf_counter() LOGGER.info( "ElevenLabs batch STT request start: model=%s sample_rate=%s pcm_bytes=%s wav_bytes=%s " "duration_ms=%s language=%s keyterms=%s", self._model_id, sample_rate_hz, len(pcm_bytes), len(wav_bytes), _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz), language_code or self._language_code, len(normalized_keyterms), ) form = aiohttp.FormData() form.add_field("model_id", self._model_id) form.add_field("timestamps_granularity", "none") form.add_field("diarize", "false") effective_language_code = language_code or self._language_code if effective_language_code: form.add_field("language_code", effective_language_code) if normalized_keyterms: form.add_field("keyterms", json.dumps(normalized_keyterms, ensure_ascii=False)) form.add_field( "file", wav_bytes, filename="turn.wav", content_type="audio/wav", ) session = await self._get_client_session() try: async with session.post( f"{self._api_base}/v1/speech-to-text", headers={"xi-api-key": self._api_key}, data=form, ) as response: payload_text = await response.text() LOGGER.info( "ElevenLabs batch STT response: status=%s latency_ms=%s response_bytes=%s", response.status, int((time.perf_counter() - started_monotonic) * 1000.0), len(payload_text), ) if response.status >= 400: if response.status == 400 and "audio_too_short" in payload_text: LOGGER.warning("ElevenLabs batch STT reported audio_too_short; ignoring utterance") return "" raise RuntimeError( f"ElevenLabs STT returned HTTP {response.status}: {payload_text[:300]}" ) except (TimeoutError, asyncio.TimeoutError) as exc: raise RuntimeError("ElevenLabs STT request timed out") from exc except aiohttp.ClientError as exc: raise RuntimeError("ElevenLabs STT request failed") from exc try: payload = json.loads(payload_text) except json.JSONDecodeError as exc: raise RuntimeError("ElevenLabs STT returned invalid JSON") from exc return str(payload.get("text") or payload.get("transcript") or "").strip() async def _get_client_session(self): try: import aiohttp except Exception as exc: # noqa: BLE001 raise RuntimeError("The `aiohttp` package is required for ElevenLabs STT") from exc if self._client_session is not None and not self._client_session.closed: return self._client_session async with self._client_session_lock: if self._client_session is not None and not self._client_session.closed: return self._client_session timeout = aiohttp.ClientTimeout(total=self._timeout_seconds) connector = aiohttp.TCPConnector(limit=32, ttl_dns_cache=300) self._client_session = aiohttp.ClientSession(timeout=timeout, connector=connector) return self._client_session async def _await_realtime_message( self, websocket, *, allowed_message_types: set[str], ) -> dict[str, object]: while True: message = await self._receive_realtime_payload(websocket) message_type = str(message.get("message_type") or "").strip().lower() if message_type in allowed_message_types: return message if self._is_error_message(message_type): raise RuntimeError(self._format_realtime_error(message)) async def _receive_realtime_payload(self, websocket) -> dict[str, object]: raw_message = await asyncio.wait_for(websocket.recv(), timeout=self._timeout_seconds) if isinstance(raw_message, bytes): raw_message = raw_message.decode("utf-8") try: payload = json.loads(raw_message) except json.JSONDecodeError as exc: raise RuntimeError("ElevenLabs realtime STT returned invalid JSON") from exc if not isinstance(payload, dict): raise RuntimeError("ElevenLabs realtime STT returned malformed payload") return payload @staticmethod def _is_error_message(message_type: str) -> bool: return message_type in {"error", "auth_error", "input_error"} or message_type.endswith("_error") @staticmethod def _format_realtime_error(payload: dict[str, object]) -> str: message_type = str(payload.get("message_type") or "error") detail = str(payload.get("message") or payload.get("detail") or payload.get("error") or "").strip() if detail: return f"ElevenLabs realtime STT {message_type}: {detail}" return f"ElevenLabs realtime STT {message_type}" def _build_realtime_websocket_url( self, *, sample_rate_hz: int, keyterms: Sequence[str] | None = None, language_code: str | None = None, ) -> str: query: list[tuple[str, str]] = [ ("model_id", self._realtime_model_id), ("audio_format", f"pcm_{sample_rate_hz}"), ("commit_strategy", "manual"), ("include_timestamps", "false"), ] effective_language_code = language_code or self._language_code if effective_language_code: query.append(("language_code", effective_language_code)) for term in _normalize_keyterms( keyterms, max_count=_REALTIME_MAX_KEYTERMS, max_chars=_REALTIME_MAX_KEYTERM_CHARS, ): query.append(("keyterms", term)) return f"{self._ws_base}/v1/speech-to-text/realtime?{urlencode(query)}" def _extract_pcm(self, audio_bytes: bytes) -> tuple[bytes, int]: try: with wave.open(io.BytesIO(audio_bytes), "rb") as wav_file: pcm_bytes = wav_file.readframes(wav_file.getnframes()) sample_width = wav_file.getsampwidth() channels = wav_file.getnchannels() sample_rate_hz = int(wav_file.getframerate() or self._input_sample_rate_hz) if sample_width != 2: return audio_bytes, self._input_sample_rate_hz if channels == 2: pcm_bytes = audioop.tomono(pcm_bytes, sample_width, 0.5, 0.5) return pcm_bytes, sample_rate_hz except (wave.Error, EOFError): return audio_bytes, self._input_sample_rate_hz class YandexSpeechKitSTT(BaseSTT): name = "yandex" def __init__( self, *, api_key: str | None = None, iam_token: str | None = None, folder_id: str | None = None, api_base: str | None = None, language: str | None = None, topic: str | None = None, input_sample_rate_hz: int = 16000, target_sample_rate_hz: int = 16000, timeout_seconds: float | None = None, ) -> None: self._api_key = str(api_key if api_key is not None else _yandex_api_key()).strip() self._iam_token = str(iam_token if iam_token is not None else _yandex_iam_token()).strip() self._folder_id = str(folder_id if folder_id is not None else _yandex_folder_id()).strip() self._api_base = str(api_base or _yandex_api_base()).strip().rstrip("/") self._language = str(language or _yandex_language()).strip() or "ru-RU" self._topic = str(topic or _yandex_topic()).strip() or "general" self._input_sample_rate_hz = max(int(input_sample_rate_hz), 1) self._target_sample_rate_hz = max(int(target_sample_rate_hz), 1) self._timeout_seconds = max(float(timeout_seconds if timeout_seconds is not None else _timeout_seconds()), 1.0) self._client_session = None self._client_session_lock = asyncio.Lock() LOGGER.info( "Yandex SpeechKit STT config: api_base=%s input_sample_rate=%s target_sample_rate=%s " "language=%s topic=%s folder_configured=%s auth=%s", self._api_base, self._input_sample_rate_hz, self._target_sample_rate_hz, self._language, self._topic, bool(self._folder_id), "iam" if self._iam_token else ("api-key" if self._api_key else "missing"), ) async def transcribe( self, audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, language_code: str | None = None, force_batch: bool = False, ) -> str: del keyterms, force_batch if not audio_bytes: return "" if not self._api_key and not self._iam_token: raise RuntimeError("YANDEX_STT_API_KEY or YANDEX_STT_IAM_TOKEN is required for Yandex STT") effective_language = _normalize_yandex_language_code(language_code) or self._language pcm_bytes, sample_rate_hz = self._extract_pcm(audio_bytes) LOGGER.info( "Yandex SpeechKit STT transcribe start: input_bytes=%s extracted_pcm_bytes=%s sample_rate=%s " "duration_ms=%s language=%s", len(audio_bytes), len(pcm_bytes), sample_rate_hz, _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz), effective_language, ) if sample_rate_hz != self._target_sample_rate_hz: before_rate_hz = sample_rate_hz before_bytes = len(pcm_bytes) pcm_bytes = _resample_pcm16le( pcm_bytes, input_rate_hz=sample_rate_hz, output_rate_hz=self._target_sample_rate_hz, ) sample_rate_hz = self._target_sample_rate_hz LOGGER.info( "Yandex SpeechKit STT resampled input: from_rate=%s to_rate=%s before_bytes=%s after_bytes=%s", before_rate_hz, sample_rate_hz, before_bytes, len(pcm_bytes), ) session = await self._get_client_session() params = { "lang": effective_language, "topic": self._topic, "format": "lpcm", "sampleRateHertz": str(sample_rate_hz), } if self._folder_id: params["folderId"] = self._folder_id headers = { "Content-Type": f"audio/x-pcm;bit=16;rate={sample_rate_hz}", "Authorization": f"Bearer {self._iam_token}" if self._iam_token else f"Api-Key {self._api_key}", } started_monotonic = time.perf_counter() try: async with session.post( f"{self._api_base}/speech/v1/stt:recognize", params=params, headers=headers, data=pcm_bytes, ) as response: payload_text = await response.text() LOGGER.info( "Yandex SpeechKit STT response: status=%s latency_ms=%s response_bytes=%s", response.status, int((time.perf_counter() - started_monotonic) * 1000.0), len(payload_text), ) if response.status >= 400: raise RuntimeError( f"Yandex SpeechKit STT returned HTTP {response.status}: {payload_text[:300]}" ) except (TimeoutError, asyncio.TimeoutError) as exc: raise RuntimeError("Yandex SpeechKit STT request timed out") from exc except Exception as exc: try: import aiohttp except Exception: # noqa: BLE001 aiohttp = None if aiohttp is not None and isinstance(exc, aiohttp.ClientError): raise RuntimeError("Yandex SpeechKit STT request failed") from exc raise try: payload = json.loads(payload_text) except json.JSONDecodeError as exc: raise RuntimeError("Yandex SpeechKit STT returned invalid JSON") from exc transcript = str(payload.get("result") or payload.get("text") or "").strip() LOGGER.info( "Yandex SpeechKit STT transcript result: chars=%s transcript=%r", len(transcript), _preview_text(transcript), ) return transcript async def close(self) -> None: session = self._client_session self._client_session = None if session is not None and not session.closed: await session.close() async def _get_client_session(self): try: import aiohttp except Exception as exc: # noqa: BLE001 raise RuntimeError("The `aiohttp` package is required for Yandex STT") from exc if self._client_session is not None and not self._client_session.closed: return self._client_session async with self._client_session_lock: if self._client_session is not None and not self._client_session.closed: return self._client_session timeout = aiohttp.ClientTimeout(total=self._timeout_seconds) connector = aiohttp.TCPConnector(limit=16, ttl_dns_cache=300) self._client_session = aiohttp.ClientSession(timeout=timeout, connector=connector) return self._client_session def _extract_pcm(self, audio_bytes: bytes) -> tuple[bytes, int]: try: with wave.open(io.BytesIO(audio_bytes), "rb") as wav_file: pcm_bytes = wav_file.readframes(wav_file.getnframes()) sample_width = wav_file.getsampwidth() channels = wav_file.getnchannels() sample_rate_hz = int(wav_file.getframerate() or self._input_sample_rate_hz) if sample_width != 2: return audio_bytes, self._input_sample_rate_hz if channels == 2: pcm_bytes = audioop.tomono(pcm_bytes, sample_width, 0.5, 0.5) return pcm_bytes, sample_rate_hz except (wave.Error, EOFError): return audio_bytes, self._input_sample_rate_hz class FallbackSTT(BaseSTT): def __init__(self, *, primary: BaseSTT, fallback: BaseSTT) -> None: self._primary = primary self._fallback = fallback LOGGER.warning( "STT fallback configured: primary=%s fallback=%s", type(primary).__name__, type(fallback).__name__, ) async def transcribe( self, audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, language_code: str | None = None, force_batch: bool = False, ) -> str: try: return await self._primary.transcribe( audio_bytes, keyterms=keyterms, language_code=language_code, force_batch=force_batch, ) except Exception: LOGGER.exception( "Primary STT provider failed; falling back: primary=%s fallback=%s", type(self._primary).__name__, type(self._fallback).__name__, ) return await self._fallback.transcribe( audio_bytes, keyterms=keyterms, language_code=language_code, force_batch=force_batch, ) async def start_stream( self, *, partial_callback: PartialTranscriptCallback | None = None, keyterms: Sequence[str] | None = None, language_code: str | None = None, ) -> BaseSTTStream | None: try: stream = await self._primary.start_stream( partial_callback=partial_callback, keyterms=keyterms, language_code=language_code, ) if stream is not None: return stream except Exception: LOGGER.exception( "Primary STT start_stream failed; falling back: primary=%s fallback=%s", type(self._primary).__name__, type(self._fallback).__name__, ) return await self._fallback.start_stream( partial_callback=partial_callback, keyterms=keyterms, language_code=language_code, ) async def close(self) -> None: for provider in (self._primary, self._fallback): close = getattr(provider, "close", None) if close is None: continue result = close() if asyncio.iscoroutine(result): await result def _build_single_stt_provider( provider: str, *, input_sample_rate_hz: int, target_sample_rate_hz: int, ) -> BaseSTT: normalized = provider.lower() if normalized in {"yandex", "yandex_speechkit", "speechkit"}: return YandexSpeechKitSTT( input_sample_rate_hz=input_sample_rate_hz, target_sample_rate_hz=target_sample_rate_hz, ) if normalized in {"elevenlabs", "eleven_labs", "scribe"}: return ElevenLabsSTT( input_sample_rate_hz=input_sample_rate_hz, target_sample_rate_hz=target_sample_rate_hz, ) if normalized in {"openai", "whisper"}: return OpenAISTT( input_sample_rate_hz=input_sample_rate_hz, ) raise RuntimeError(f"Unsupported STT provider: {provider}") def build_stt_provider( *, input_sample_rate_hz: int, target_sample_rate_hz: int, ) -> BaseSTT: provider = _stt_provider_name() primary = _build_single_stt_provider( provider, input_sample_rate_hz=input_sample_rate_hz, target_sample_rate_hz=target_sample_rate_hz, ) fallback_provider = _stt_fallback_provider_name() if not fallback_provider or fallback_provider == provider: return primary fallback = _build_single_stt_provider( fallback_provider, input_sample_rate_hz=input_sample_rate_hz, target_sample_rate_hz=target_sample_rate_hz, ) return FallbackSTT(primary=primary, fallback=fallback)