603 lines
20 KiB
Python
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}
|