Files
realtime_voice_service/providers/stt_openai.py
T

189 lines
6.8 KiB
Python

from __future__ import annotations
import asyncio
import io
import logging
import os
import time
import wave
from typing import Any
from realtime_voice_service.providers.base import BaseSTT
LOGGER = logging.getLogger("uvicorn.error")
def _timeout_seconds() -> float:
raw = os.getenv("OPENAI_STT_TIMEOUT_SECONDS") or os.getenv("OPENAI_TIMEOUT_SECONDS")
if raw is None:
return 30.0
try:
return max(float(raw.strip()), 1.0)
except ValueError:
return 30.0
def _max_retries() -> int:
raw = os.getenv("OPENAI_STT_MAX_RETRIES")
if raw is None:
return 2
try:
return max(int(raw.strip()), 0)
except ValueError:
return 2
def _preview_text(text: str, *, limit: int = 120) -> str:
normalized = " ".join(str(text or "").split())
if len(normalized) <= limit:
return normalized
return f"{normalized[:limit]}..."
def _pcm16le_to_wav_bytes(
pcm_bytes: bytes,
*,
sample_rate_hz: int,
channels: int = 1,
sample_width_bytes: int = 2,
) -> bytes:
handle = io.BytesIO()
with wave.open(handle, "wb") as wav_file:
wav_file.setnchannels(channels)
wav_file.setsampwidth(sample_width_bytes)
wav_file.setframerate(sample_rate_hz)
wav_file.writeframes(pcm_bytes)
return handle.getvalue()
class OpenAISTT(BaseSTT):
name = "openai"
def __init__(
self,
*,
api_key: str | None = None,
model: str | None = None,
base_url: str | None = None,
input_sample_rate_hz: int = 8000,
prompt: str | None = None,
language: str | None = None,
timeout_seconds: float | None = None,
max_retries: int | None = None,
) -> None:
self._api_key = str(api_key if api_key is not None else os.getenv("OPENAI_API_KEY", "")).strip()
self._model = str(model or os.getenv("OPENAI_STT_MODEL", "whisper-1")).strip() or "whisper-1"
self._base_url = str(
base_url if base_url is not None else (os.getenv("OPENAI_STT_BASE_URL") or os.getenv("OPENAI_BASE_URL") or "")
).strip() or None
self._input_sample_rate_hz = max(int(input_sample_rate_hz), 1)
self._prompt = str(prompt if prompt is not None else os.getenv("STT_PROMPT", "")).strip()
self._language = str(language if language is not None else os.getenv("OPENAI_STT_LANGUAGE", "ru")).strip() or None
self._timeout_seconds = max(float(timeout_seconds if timeout_seconds is not None else _timeout_seconds()), 1.0)
self._max_retries = max(int(max_retries if max_retries is not None else _max_retries()), 0)
self._client: Any | None = None
self._openai_module: Any | None = None
LOGGER.info(
"OpenAI STT config: model=%s input_sample_rate=%s language=%s prompt_configured=%s base_url=%s timeout=%s",
self._model,
self._input_sample_rate_hz,
self._language,
bool(self._prompt),
self._base_url or "default",
self._timeout_seconds,
)
async def transcribe(self, audio_bytes: bytes) -> str:
if not audio_bytes:
return ""
if not self._api_key:
raise RuntimeError("OPENAI_API_KEY is required for OpenAI STT")
wav_bytes = _pcm16le_to_wav_bytes(audio_bytes, sample_rate_hz=self._input_sample_rate_hz)
audio_file = io.BytesIO(wav_bytes)
audio_file.name = "utterance.wav"
started_monotonic = time.perf_counter()
LOGGER.info(
"OpenAI STT request start: model=%s pcm_bytes=%s wav_bytes=%s sample_rate=%s language=%s prompt=%r",
self._model,
len(audio_bytes),
len(wav_bytes),
self._input_sample_rate_hz,
self._language,
_preview_text(self._prompt),
)
request: dict[str, object] = {
"file": audio_file,
"model": self._model,
"response_format": "json",
}
if self._prompt:
request["prompt"] = self._prompt
if self._language:
request["language"] = self._language
try:
response = await self._get_client().audio.transcriptions.create(**request)
except (TimeoutError, asyncio.TimeoutError) as exc:
raise RuntimeError("OpenAI STT request timed out") from exc
except Exception as exc: # noqa: BLE001
openai_module = self._openai_module
if openai_module is not None and isinstance(exc, getattr(openai_module, "APITimeoutError", ())):
raise RuntimeError("OpenAI STT request timed out") from exc
if openai_module is not None and isinstance(exc, getattr(openai_module, "RateLimitError", ())):
raise RuntimeError("OpenAI STT rate limit exceeded") from exc
if openai_module is not None and isinstance(exc, getattr(openai_module, "APIStatusError", ())):
status_code = getattr(exc, "status_code", "unknown")
raise RuntimeError(f"OpenAI STT returned HTTP {status_code}") from exc
if openai_module is not None and isinstance(exc, getattr(openai_module, "APIConnectionError", ())):
raise RuntimeError("OpenAI STT connection failed") from exc
raise RuntimeError("OpenAI STT request failed") from exc
transcript_value = getattr(response, "text", None)
if transcript_value is None and isinstance(response, dict):
transcript_value = response.get("text")
transcript = str(transcript_value or "").strip()
LOGGER.info(
"OpenAI STT transcript result: latency_ms=%s chars=%s transcript=%r",
int((time.perf_counter() - started_monotonic) * 1000.0),
len(transcript),
_preview_text(transcript),
)
return transcript
def _get_client(self):
if self._client is not None:
return self._client
try:
import httpx
import openai
from openai import AsyncOpenAI
except Exception as exc: # noqa: BLE001
raise RuntimeError("The `openai` package is required for OpenAI STT") from exc
timeout = httpx.Timeout(
self._timeout_seconds,
connect=min(self._timeout_seconds, 5.0),
write=min(self._timeout_seconds, 15.0),
read=self._timeout_seconds,
)
client_kwargs: dict[str, object] = {
"api_key": self._api_key,
"timeout": timeout,
"max_retries": self._max_retries,
}
if self._base_url:
client_kwargs["base_url"] = self._base_url
self._openai_module = openai
self._client = AsyncOpenAI(**client_kwargs)
return self._client
async def close(self) -> None:
client = self._client
self._client = None
if client is not None:
await client.close()