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

86 lines
2.6 KiB
Python

from __future__ import annotations
import os
from dataclasses import dataclass
import httpx
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"
@dataclass(slots=True)
class ASRTranscription:
text: str
language: str | None = None
confidence: float | None = None
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)
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 build_asr_provider(name: str) -> ASRProvider:
normalized = str(name or "stub").strip().lower()
if normalized == "openai":
return OpenAIASRProvider()
return ASRProvider()