86 lines
2.6 KiB
Python
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()
|