Files
call-center/services/ai_voice_runtime_service/providers/asr.py
T

260 lines
8.6 KiB
Python

from __future__ import annotations
import base64
import os
from dataclasses import dataclass
import httpx
from services.shared.security import issue_app_token
def _api_base() -> str:
return (os.getenv("AI_API_BASE", "https://api.openai.com/v1").strip() or "https://api.openai.com/v1").rstrip("/")
def _api_key() -> str:
return os.getenv("AI_API_KEY", "").strip()
def _timeout_seconds() -> float:
raw = os.getenv("AI_TIMEOUT_SECONDS", "20").strip()
try:
value = float(raw)
except ValueError:
value = 20.0
return max(value, 3.0)
def _openai_asr_model() -> str:
return os.getenv("AI_VOICE_ASR_MODEL", "gpt-4o-mini-transcribe").strip() or "gpt-4o-mini-transcribe"
def _streaming_asr_api_base() -> str:
return (
os.getenv("AI_VOICE_V2_STREAMING_ASR_BASE_URL", "http://127.0.0.1:8021").strip()
or "http://127.0.0.1:8021"
).rstrip("/")
def _streaming_asr_timeout_seconds() -> float:
raw = os.getenv("AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS", "4").strip()
try:
value = float(raw)
except ValueError:
value = 4.0
return max(value, 0.25)
def _service_headers() -> dict[str, str]:
token = issue_app_token(
subject="svc:ai-voice-runtime",
username="ai-voice-runtime",
role="admin",
auth_source="service",
provider="ai-voice-runtime",
ttl_seconds=300,
)
return {"Authorization": f"Bearer {token}"}
@dataclass(slots=True)
class ASRTranscription:
text: str
language: str | None = None
confidence: float | None = None
@dataclass(slots=True)
class StreamingASRPartial:
text: str
language: str | None = None
confidence: float | None = None
is_final: bool = False
is_stable: bool = False
class StreamingASRUnavailable(RuntimeError):
pass
class ASRProvider:
name = "stub"
def transcribe(self, audio_bytes: bytes, *, language_hint: str | None = None) -> ASRTranscription:
del audio_bytes
return ASRTranscription(text="", language=language_hint, confidence=None)
def transcribe_partial(self, audio_bytes: bytes, *, language_hint: str | None = None) -> ASRTranscription:
return self.transcribe(audio_bytes, language_hint=language_hint)
class StreamingASRProvider:
name = "streaming-stub"
supports_streaming = False
def open_stream(self, session_id: str, *, language_hint: str | None = None) -> str:
del session_id, language_hint
raise StreamingASRUnavailable("Streaming ASR backend is not configured")
def push_pcm(self, stream_id: str, pcm_8k_chunk: bytes) -> None:
del stream_id, pcm_8k_chunk
raise StreamingASRUnavailable("Streaming ASR backend is not configured")
def poll_partial(self, stream_id: str) -> StreamingASRPartial | None:
del stream_id
return None
def finalize(self, stream_id: str) -> ASRTranscription:
del stream_id
raise StreamingASRUnavailable("Streaming ASR backend is not configured")
def close_stream(self, stream_id: str) -> None:
del stream_id
class OpenAIASRProvider(ASRProvider):
name = "openai"
def __init__(self) -> None:
self._api_base = _api_base()
self._api_key = _api_key()
self._timeout_seconds = _timeout_seconds()
self._model = _openai_asr_model()
def transcribe(self, audio_bytes: bytes, *, language_hint: str | None = None) -> ASRTranscription:
if not audio_bytes:
return ASRTranscription(text="", language=language_hint, confidence=None)
if not self._api_key:
raise RuntimeError("AI_API_KEY is required for OpenAI ASR")
data: dict[str, str] = {"model": self._model}
language = str(language_hint or "").strip()
if language:
data["language"] = language
with httpx.Client(timeout=self._timeout_seconds) as client:
response = client.post(
f"{self._api_base}/audio/transcriptions",
headers={"Authorization": f"Bearer {self._api_key}"},
data=data,
files={"file": ("turn.wav", audio_bytes, "audio/wav")},
)
response.raise_for_status()
payload = response.json()
return ASRTranscription(
text=str(payload.get("text") or "").strip(),
language=str(payload.get("language") or language_hint or "").strip() or language_hint,
confidence=None,
)
def transcribe_partial(self, audio_bytes: bytes, *, language_hint: str | None = None) -> ASRTranscription:
return self.transcribe(audio_bytes, language_hint=language_hint)
class LocalSidecarStreamingASRProvider(StreamingASRProvider):
name = "local-sidecar"
supports_streaming = True
def __init__(
self,
*,
api_base: str | None = None,
timeout_seconds: float | None = None,
) -> None:
self._api_base = str(api_base or _streaming_asr_api_base()).strip().rstrip("/")
self._timeout_seconds = max(float(timeout_seconds or _streaming_asr_timeout_seconds()), 0.25)
def _request(
self,
method: str,
path: str,
*,
payload: dict[str, object] | None = None,
) -> dict[str, object]:
if not self._api_base:
raise StreamingASRUnavailable("Streaming ASR API base is not configured")
try:
with httpx.Client(timeout=self._timeout_seconds) as client:
response = client.request(
method,
f"{self._api_base}{path}",
json=payload,
headers=_service_headers(),
)
response.raise_for_status()
except httpx.HTTPError as exc:
raise StreamingASRUnavailable(str(exc)[:500] or "Streaming ASR sidecar is unavailable") from exc
body = response.json() if response.content else {}
return body if isinstance(body, dict) else {}
def open_stream(self, session_id: str, *, language_hint: str | None = None) -> str:
payload = {
"session_id": str(session_id or "").strip(),
"language_hint": str(language_hint or "").strip() or None,
"sample_rate_hz": 8000,
"encoding": "pcm_s16le",
}
body = self._request("POST", "/internal/asr/streams", payload=payload)
stream_id = str(body.get("stream_id") or "").strip()
if not stream_id:
raise StreamingASRUnavailable("Streaming ASR sidecar did not return stream_id")
return stream_id
def push_pcm(self, stream_id: str, pcm_8k_chunk: bytes) -> None:
if not pcm_8k_chunk:
return
self._request(
"POST",
f"/internal/asr/streams/{stream_id}/chunks",
payload={
"pcm_b64": base64.b64encode(pcm_8k_chunk).decode("ascii"),
"sample_rate_hz": 8000,
"encoding": "pcm_s16le",
},
)
def poll_partial(self, stream_id: str) -> StreamingASRPartial | None:
body = self._request("GET", f"/internal/asr/streams/{stream_id}/partial")
text = str(body.get("text") or "").strip()
if not text:
return None
language = str(body.get("language") or "").strip() or None
confidence_raw = body.get("confidence")
confidence = float(confidence_raw) if isinstance(confidence_raw, (int, float)) else None
return StreamingASRPartial(
text=text,
language=language,
confidence=confidence,
is_final=bool(body.get("is_final")),
is_stable=bool(body.get("is_stable")),
)
def finalize(self, stream_id: str) -> ASRTranscription:
body = self._request("POST", f"/internal/asr/streams/{stream_id}/finalize")
return ASRTranscription(
text=str(body.get("text") or "").strip(),
language=str(body.get("language") or "").strip() or None,
confidence=float(body["confidence"]) if isinstance(body.get("confidence"), (int, float)) else None,
)
def close_stream(self, stream_id: str) -> None:
try:
self._request("DELETE", f"/internal/asr/streams/{stream_id}")
except StreamingASRUnavailable:
return
def build_asr_provider(name: str) -> ASRProvider:
normalized = str(name or "stub").strip().lower()
if normalized == "openai":
return OpenAIASRProvider()
return ASRProvider()
def build_streaming_asr_provider(name: str) -> StreamingASRProvider:
normalized = str(name or "disabled").strip().lower()
if normalized in {"local_sidecar", "local-sidecar", "sidecar"}:
return LocalSidecarStreamingASRProvider()
return StreamingASRProvider()