Files
call-center/services/streaming_asr_sidecar_service/app.py
T
2026-05-09 09:16:06 +05:00

603 lines
20 KiB
Python

from __future__ import annotations
import base64
import logging
import os
import threading
import time
from dataclasses import dataclass, field
from typing import Any, Protocol
from fastapi import Depends, FastAPI, HTTPException
from pydantic import BaseModel
from services.shared.core import Role, new_id
from services.shared.audioop_compat import audioop
from services.shared.models import HealthResponse
from services.shared.security import require_roles
LOGGER = logging.getLogger("uvicorn.error")
app = FastAPI(title="streaming-asr-sidecar-service", version="1.0.0")
def _int_env(name: str, default: int) -> int:
raw = os.getenv(name)
if raw is None:
return default
try:
return int(str(raw).strip())
except ValueError:
return default
def _float_env(name: str, default: float) -> float:
raw = os.getenv(name)
if raw is None:
return default
try:
return float(str(raw).strip())
except ValueError:
return default
def _stream_idle_ttl_seconds() -> float:
return max(_float_env("AI_VOICE_V2_STREAMING_ASR_IDLE_TTL_SECONDS", 45.0), 5.0)
def _partial_min_audio_ms() -> int:
return max(_int_env("AI_VOICE_V2_STREAMING_ASR_MIN_AUDIO_MS", 320), 120)
def _partial_recompute_interval_ms() -> int:
return max(_int_env("AI_VOICE_V2_STREAMING_ASR_PARTIAL_RECOMPUTE_INTERVAL_MS", 450), 80)
def _partial_stability_hold_ms() -> int:
return max(_int_env("AI_VOICE_V2_STREAMING_ASR_PARTIAL_STABILITY_HOLD_MS", 400), 120)
def _partial_max_audio_ms() -> int:
return max(_int_env("AI_VOICE_V2_STREAMING_ASR_PARTIAL_MAX_AUDIO_MS", 1800), 400)
def _model_name() -> str:
return str(os.getenv("AI_VOICE_V2_STREAMING_ASR_MODEL", "base") or "base").strip() or "base"
def _compute_type() -> str:
return str(os.getenv("AI_VOICE_V2_STREAMING_ASR_COMPUTE_TYPE", "int8") or "int8").strip() or "int8"
def _device() -> str:
return str(os.getenv("AI_VOICE_V2_STREAMING_ASR_DEVICE", "cpu") or "cpu").strip() or "cpu"
def _cache_dir() -> str:
return str(os.getenv("AI_VOICE_V2_STREAMING_ASR_CACHE_DIR", "/models/faster_whisper") or "/models/faster_whisper").strip() or "/models/faster_whisper"
def _supported_languages() -> set[str]:
raw = str(os.getenv("AI_VOICE_V2_STREAMING_ASR_SUPPORTED_LANGUAGES", "ru") or "ru").strip()
values = {item.strip().lower() for item in raw.split(",") if item.strip()}
return values or {"ru"}
def _target_sample_rate_hz() -> int:
return max(_int_env("AI_VOICE_V2_STREAMING_ASR_TARGET_SAMPLE_RATE_HZ", 16000), 8000)
def _normalize_language_hint(language_hint: str | None) -> str:
raw = str(language_hint or "ru").strip().lower() or "ru"
if raw in {"ru", "ru-ru"}:
return "ru"
if raw in {"kk", "kz", "kk-kk"}:
return "kk"
return raw
def _pcm_duration_ms(pcm_bytes: bytes | bytearray, sample_rate_hz: int) -> int:
if sample_rate_hz <= 0:
return 0
sample_count = len(pcm_bytes) // 2
return int((sample_count / float(sample_rate_hz)) * 1000.0)
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
@dataclass(slots=True)
class SidecarTranscript:
text: str
language: str | None = None
confidence: float | None = None
class TranscriptionEngine(Protocol):
def transcribe_pcm(
self,
pcm_bytes: bytes,
*,
sample_rate_hz: int,
language_hint: str | None = None,
) -> SidecarTranscript: ...
class FasterWhisperEngine:
def __init__(
self,
*,
model_name: str,
compute_type: str,
device: str,
cache_dir: str,
target_sample_rate_hz: int,
) -> None:
self._model_name = model_name
self._compute_type = compute_type
self._device = device
self._cache_dir = cache_dir
self._target_sample_rate_hz = target_sample_rate_hz
self._model = None
self._lock = threading.Lock()
def _load_model(self):
with self._lock:
if self._model is not None:
return self._model
try:
from faster_whisper import WhisperModel
except Exception as exc: # noqa: BLE001
raise RuntimeError(
"faster-whisper is not installed; install requirements before starting the streaming ASR sidecar"
) from exc
self._model = WhisperModel(
self._model_name,
device=self._device,
compute_type=self._compute_type,
download_root=self._cache_dir,
)
return self._model
def transcribe_pcm(
self,
pcm_bytes: bytes,
*,
sample_rate_hz: int,
language_hint: str | None = None,
) -> SidecarTranscript:
if not pcm_bytes:
return SidecarTranscript(text="", language=language_hint, confidence=None)
try:
import numpy as np
except Exception as exc: # noqa: BLE001
raise RuntimeError("numpy is required for streaming ASR sidecar") from exc
model = self._load_model()
pcm_model_rate = _resample_pcm16le(
pcm_bytes,
input_rate_hz=sample_rate_hz,
output_rate_hz=self._target_sample_rate_hz,
)
waveform = np.frombuffer(pcm_model_rate, dtype=np.int16).astype(np.float32) / 32768.0
segments, info = model.transcribe(
waveform,
language=_normalize_language_hint(language_hint),
beam_size=1,
best_of=1,
temperature=0.0,
vad_filter=False,
word_timestamps=False,
condition_on_previous_text=False,
without_timestamps=True,
)
parts: list[str] = []
for segment in segments:
text = str(getattr(segment, "text", "") or "").strip()
if text:
parts.append(text)
text = " ".join(parts).strip()
language = str(getattr(info, "language", "") or language_hint or "").strip() or language_hint
confidence_raw = getattr(info, "language_probability", None)
confidence = float(confidence_raw) if isinstance(confidence_raw, (int, float)) else None
if confidence is not None:
confidence = max(0.0, min(confidence, 1.0))
return SidecarTranscript(text=text, language=language, confidence=confidence)
def _default_engine_factory() -> TranscriptionEngine:
return FasterWhisperEngine(
model_name=_model_name(),
compute_type=_compute_type(),
device=_device(),
cache_dir=_cache_dir(),
target_sample_rate_hz=_target_sample_rate_hz(),
)
_ENGINE_FACTORY = _default_engine_factory
_ENGINE_INSTANCE: TranscriptionEngine | None = None
_ENGINE_LOCK = threading.Lock()
def _get_engine() -> TranscriptionEngine:
global _ENGINE_INSTANCE
with _ENGINE_LOCK:
if _ENGINE_INSTANCE is None:
_ENGINE_INSTANCE = _ENGINE_FACTORY()
return _ENGINE_INSTANCE
def _warmup_engine() -> None:
try:
_get_engine()
LOGGER.info("streaming_asr.engine_ready")
except Exception as exc: # noqa: BLE001
LOGGER.warning("streaming_asr.engine_warmup_failed error=%s", str(exc)[:500])
@app.on_event("startup")
def _startup_warmup_engine() -> None:
threading.Thread(target=_warmup_engine, daemon=True).start()
@dataclass(slots=True)
class StreamState:
stream_id: str
session_id: str
language_hint: str
sample_rate_hz: int
encoding: str
pcm_buffer: bytearray = field(default_factory=bytearray)
created_monotonic: float = field(default_factory=time.monotonic)
last_activity_monotonic: float = field(default_factory=time.monotonic)
buffer_version: int = 0
partial_text: str = ""
partial_language: str | None = None
partial_confidence: float | None = None
partial_buffer_version: int = -1
partial_updated_monotonic: float = 0.0
partial_first_seen_monotonic: float = 0.0
partial_repeat_count: int = 0
final_transcript: SidecarTranscript | None = None
def touch(self) -> None:
self.last_activity_monotonic = time.monotonic()
class StreamStore:
def __init__(self, *, idle_ttl_seconds: float) -> None:
self._idle_ttl_seconds = idle_ttl_seconds
self._streams: dict[str, StreamState] = {}
self._lock = threading.Lock()
def cleanup_expired(self) -> list[str]:
now = time.monotonic()
removed: list[str] = []
with self._lock:
for stream_id, stream in list(self._streams.items()):
if now - stream.last_activity_monotonic < self._idle_ttl_seconds:
continue
removed.append(stream_id)
self._streams.pop(stream_id, None)
return removed
def create(self, *, session_id: str, language_hint: str, sample_rate_hz: int, encoding: str) -> StreamState:
stream = StreamState(
stream_id=new_id("sasr"),
session_id=session_id,
language_hint=language_hint,
sample_rate_hz=sample_rate_hz,
encoding=encoding,
)
with self._lock:
self._streams[stream.stream_id] = stream
return stream
def get(self, stream_id: str) -> StreamState | None:
with self._lock:
stream = self._streams.get(stream_id)
if stream is not None:
stream.touch()
return stream
def append_pcm(self, stream_id: str, pcm_chunk: bytes) -> StreamState:
with self._lock:
stream = self._streams.get(stream_id)
if stream is None:
raise KeyError(stream_id)
stream.pcm_buffer.extend(pcm_chunk)
stream.buffer_version += 1
stream.touch()
return stream
def update_partial(
self,
stream_id: str,
*,
transcript: SidecarTranscript,
buffer_version: int,
) -> StreamState | None:
now = time.monotonic()
with self._lock:
stream = self._streams.get(stream_id)
if stream is None:
return None
previous_text = stream.partial_text
stream.touch()
if transcript.text:
if transcript.text == previous_text:
stream.partial_repeat_count += 1
else:
stream.partial_first_seen_monotonic = now
stream.partial_repeat_count = 1
stream.partial_text = transcript.text
stream.partial_language = transcript.language
stream.partial_confidence = transcript.confidence
stream.partial_updated_monotonic = now
stream.partial_buffer_version = buffer_version
return stream
def set_final(self, stream_id: str, transcript: SidecarTranscript) -> StreamState | None:
with self._lock:
stream = self._streams.get(stream_id)
if stream is None:
return None
stream.final_transcript = transcript
stream.touch()
return stream
def delete(self, stream_id: str) -> bool:
with self._lock:
return self._streams.pop(stream_id, None) is not None
_STREAMS = StreamStore(idle_ttl_seconds=_stream_idle_ttl_seconds())
class OpenStreamIn(BaseModel):
session_id: str
language_hint: str | None = None
sample_rate_hz: int = 8000
encoding: str = "pcm_s16le"
class OpenStreamOut(BaseModel):
ok: bool = True
stream_id: str
language: str
class PushChunkIn(BaseModel):
pcm_b64: str
sample_rate_hz: int = 8000
encoding: str = "pcm_s16le"
class PushChunkOut(BaseModel):
ok: bool = True
stream_id: str
received_bytes: int
duration_ms: int
class PartialOut(BaseModel):
text: str
language: str | None = None
confidence: float | None = None
is_final: bool = False
is_stable: bool = False
class FinalizeOut(BaseModel):
text: str
language: str | None = None
confidence: float | None = None
def _require_internal_actor(_: dict = Depends(require_roles(Role.ADMIN))) -> None:
return None
def _cleanup_expired_streams() -> None:
removed = _STREAMS.cleanup_expired()
if removed:
LOGGER.info("streaming_asr.cleanup_expired removed=%s", len(removed))
def _current_partial_is_stable(stream: StreamState) -> bool:
if not stream.partial_text:
return False
if stream.partial_repeat_count >= 2:
return True
return (time.monotonic() - stream.partial_first_seen_monotonic) >= (_partial_stability_hold_ms() / 1000.0)
def _copy_stream_snapshot(stream: StreamState) -> tuple[bytes, int, int, str]:
return (
bytes(stream.pcm_buffer),
stream.buffer_version,
stream.sample_rate_hz,
stream.language_hint,
)
def _trim_partial_snapshot(pcm_bytes: bytes, *, sample_rate_hz: int) -> bytes:
if not pcm_bytes or sample_rate_hz <= 0:
return pcm_bytes
max_bytes = int((_partial_max_audio_ms() / 1000.0) * sample_rate_hz * 2)
if max_bytes <= 0 or len(pcm_bytes) <= max_bytes:
return pcm_bytes
return pcm_bytes[-max_bytes:]
def _transcribe_stream_snapshot(pcm_bytes: bytes, *, sample_rate_hz: int, language_hint: str) -> SidecarTranscript:
return _get_engine().transcribe_pcm(
pcm_bytes,
sample_rate_hz=sample_rate_hz,
language_hint=language_hint,
)
def _current_partial_payload(stream: StreamState) -> PartialOut | None:
if not stream.partial_text:
return None
return PartialOut(
text=stream.partial_text,
language=stream.partial_language,
confidence=stream.partial_confidence,
is_final=False,
is_stable=_current_partial_is_stable(stream),
)
def _language_or_error(language_hint: str | None) -> str:
normalized = _normalize_language_hint(language_hint)
if normalized not in _supported_languages():
raise HTTPException(
status_code=400,
detail=f"Unsupported streaming ASR language: {normalized}",
)
return normalized
def _validate_audio_params(sample_rate_hz: int, encoding: str) -> None:
if int(sample_rate_hz) != 8000:
raise HTTPException(status_code=400, detail="Only 8000 Hz streaming ASR input is supported")
if str(encoding or "").strip().lower() != "pcm_s16le":
raise HTTPException(status_code=400, detail="Only pcm_s16le encoding is supported")
@app.get("/health", response_model=HealthResponse)
def health() -> HealthResponse:
_cleanup_expired_streams()
with _ENGINE_LOCK:
loaded = _ENGINE_INSTANCE is not None
suffix = "loaded" if loaded else "cold"
return HealthResponse(status="ok", service=f"streaming-asr-sidecar ({suffix})")
@app.post("/internal/asr/streams", response_model=OpenStreamOut)
def open_stream(payload: OpenStreamIn, _: None = Depends(_require_internal_actor)) -> OpenStreamOut:
_cleanup_expired_streams()
session_id = str(payload.session_id or "").strip()
if not session_id:
raise HTTPException(status_code=400, detail="session_id is required")
language = _language_or_error(payload.language_hint)
_validate_audio_params(payload.sample_rate_hz, payload.encoding)
stream = _STREAMS.create(
session_id=session_id,
language_hint=language,
sample_rate_hz=payload.sample_rate_hz,
encoding=payload.encoding,
)
return OpenStreamOut(stream_id=stream.stream_id, language=language)
@app.post("/internal/asr/streams/{stream_id}/chunks", response_model=PushChunkOut)
def push_chunk(stream_id: str, payload: PushChunkIn, _: None = Depends(_require_internal_actor)) -> PushChunkOut:
_cleanup_expired_streams()
_validate_audio_params(payload.sample_rate_hz, payload.encoding)
try:
pcm_bytes = base64.b64decode(str(payload.pcm_b64 or "").encode("ascii"), validate=True)
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=400, detail="pcm_b64 must be valid base64") from exc
try:
stream = _STREAMS.append_pcm(stream_id, pcm_bytes)
except KeyError as exc:
raise HTTPException(status_code=404, detail="ASR stream not found") from exc
return PushChunkOut(
stream_id=stream_id,
received_bytes=len(pcm_bytes),
duration_ms=_pcm_duration_ms(stream.pcm_buffer, stream.sample_rate_hz),
)
@app.get("/internal/asr/streams/{stream_id}/partial", response_model=PartialOut | dict[str, Any])
def poll_partial(stream_id: str, _: None = Depends(_require_internal_actor)) -> PartialOut | dict[str, Any]:
_cleanup_expired_streams()
stream = _STREAMS.get(stream_id)
if stream is None:
raise HTTPException(status_code=404, detail="ASR stream not found")
if stream.final_transcript is not None:
return PartialOut(
text=stream.final_transcript.text,
language=stream.final_transcript.language,
confidence=stream.final_transcript.confidence,
is_final=True,
is_stable=True,
)
if _pcm_duration_ms(stream.pcm_buffer, stream.sample_rate_hz) < _partial_min_audio_ms():
payload = _current_partial_payload(stream)
return payload.model_dump() if payload else {}
now = time.monotonic()
if (
stream.partial_text
and stream.partial_buffer_version == stream.buffer_version
and (now - stream.partial_updated_monotonic) < (_partial_recompute_interval_ms() / 1000.0)
):
payload = _current_partial_payload(stream)
return payload.model_dump() if payload else {}
pcm_bytes, buffer_version, sample_rate_hz, language_hint = _copy_stream_snapshot(stream)
transcript = _transcribe_stream_snapshot(
_trim_partial_snapshot(pcm_bytes, sample_rate_hz=sample_rate_hz),
sample_rate_hz=sample_rate_hz,
language_hint=language_hint,
)
stream = _STREAMS.update_partial(
stream_id,
transcript=transcript,
buffer_version=buffer_version,
)
if stream is None:
raise HTTPException(status_code=404, detail="ASR stream not found")
payload = _current_partial_payload(stream)
return payload.model_dump() if payload else {}
@app.post("/internal/asr/streams/{stream_id}/finalize", response_model=FinalizeOut)
def finalize_stream(stream_id: str, _: None = Depends(_require_internal_actor)) -> FinalizeOut:
_cleanup_expired_streams()
stream = _STREAMS.get(stream_id)
if stream is None:
raise HTTPException(status_code=404, detail="ASR stream not found")
if stream.final_transcript is not None:
return FinalizeOut(
text=stream.final_transcript.text,
language=stream.final_transcript.language,
confidence=stream.final_transcript.confidence,
)
pcm_bytes, _buffer_version, sample_rate_hz, language_hint = _copy_stream_snapshot(stream)
transcript = _transcribe_stream_snapshot(
pcm_bytes,
sample_rate_hz=sample_rate_hz,
language_hint=language_hint,
)
_STREAMS.set_final(stream_id, transcript)
return FinalizeOut(
text=transcript.text,
language=transcript.language,
confidence=transcript.confidence,
)
@app.delete("/internal/asr/streams/{stream_id}")
def close_stream(stream_id: str, _: None = Depends(_require_internal_actor)) -> dict[str, Any]:
_cleanup_expired_streams()
if not _STREAMS.delete(stream_id):
raise HTTPException(status_code=404, detail="ASR stream not found")
return {"ok": True, "stream_id": stream_id}