Files
realtime_voice_service/providers/stt.py
T
2026-05-02 00:48:40 +05:00

1298 lines
50 KiB
Python

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)