voicerag / src /asr_api.py
menoone's picture
Add voice RAG over MSMARCO-XI, deployable without the GPU pod
11ecc5b
Raw
History Blame Contribute Delete
8.61 kB
#!/usr/bin/env python3
"""
Speech-to-text via Sarvam or ElevenLabs -- REQUIREMENT 1, which is not optional.
"Use either Sarvam or ElevenLabs for voice-to-text. Pick one."
An earlier draft of this system used Meta MMS because it covers all 14 languages
locally. That is non-compliant no matter how well it works, so MMS ASR is gone.
MMS **TTS** stays: the brief specifies the speech-to-TEXT provider and says
nothing about speech synthesis.
WHICH ONE TO PICK
-----------------
sarvam built for Indian languages; Indic accent and code-mixing handling
is its whole reason to exist. Verify its current language list
covers Assamese / Nepali / Sanskrit before committing -- if it does
not, those languages need ElevenLabs or a documented exception.
elevenlabs much wider nominal language coverage (Scribe), less Indic-specific.
Set VOICERAG_ASR=sarvam|elevenlabs and the matching key:
SARVAM_API_KEY=... or ELEVENLABS_API_KEY=...
VERIFY THE WIRE FORMAT BEFORE THE DEMO. The request shapes below follow each
provider's documented API, but both move faster than any hardcoded client. Run
`python src/asr_api.py --selftest` with a key set: it does one real round trip
and prints exactly what came back, so a field rename is caught now rather than
during the recording. Field names are overridable via the constructor for the
same reason.
NO SILENT FALLBACK. If the provider is misconfigured or down, this raises and
the caller surfaces the error. A local model quietly standing in for the
required provider would be a compliance failure that looks like a success.
"""
from __future__ import annotations
import json
import os
import time
from dataclasses import dataclass
SARVAM_URL = "https://api.sarvam.ai/speech-to-text"
ELEVEN_URL = "https://api.elevenlabs.io/v1/speech-to-text"
# Sarvam takes BCP-47-ish codes, not bare ISO-2.
SARVAM_LANG = {
"hi": "hi-IN", "bn": "bn-IN", "gu": "gu-IN", "kn": "kn-IN", "ml": "ml-IN",
"mr": "mr-IN", "or": "od-IN", "pa": "pa-IN", "ta": "ta-IN", "te": "te-IN",
"as": "as-IN", "ne": "ne-IN", "sa": "sa-IN", "ur": "ur-IN", "en": "en-IN",
}
class ASRError(RuntimeError):
"""Permanent: config, auth, or a response shape we do not understand.
`permanent` is read by the harness, which skips retries when it is set.
Retrying a missing API key three times with exponential backoff only makes
the same error arrive a second later, and it burns the latency budget doing
it. Transient failures (network, 5xx, timeout) raise plain RuntimeError and
DO retry."""
permanent = True
@dataclass
class ASRResult:
text: str
provider: str
latency_ms: float
raw: dict
class CloudASR:
"""One provider, chosen explicitly, with bounded retries."""
def __init__(self, provider: str | None = None, api_key: str | None = None,
model: str | None = None, timeout: float = 30.0,
retries: int = 2, text_field: str | None = None):
self.provider = (provider or os.environ.get("VOICERAG_ASR") or "sarvam").lower()
if self.provider not in ("sarvam", "elevenlabs"):
raise ASRError(f"provider must be sarvam or elevenlabs, got {self.provider!r}")
env_key = "SARVAM_API_KEY" if self.provider == "sarvam" else "ELEVENLABS_API_KEY"
self.api_key = api_key or os.environ.get(env_key, "")
self.model = model or ("saarika:v2" if self.provider == "sarvam" else "scribe_v1")
self.timeout, self.retries = timeout, retries
self.text_field = text_field or ("transcript" if self.provider == "sarvam" else "text")
self.ok = bool(self.api_key)
self.why = (f"{self.provider} ({self.model})" if self.ok
else f"{self.provider}: {env_key} is not set")
# -- wire format -----------------------------------------------------
def _request(self, wav: bytes, lang: str):
if self.provider == "sarvam":
fields = {"model": self.model,
"language_code": SARVAM_LANG.get(lang, "hi-IN")}
return SARVAM_URL, {"api-subscription-key": self.api_key}, fields, "file"
return ELEVEN_URL, {"xi-api-key": self.api_key}, {"model_id": self.model}, "file"
def transcribe(self, wav_bytes: bytes, lang: str) -> ASRResult:
if not self.ok:
raise ASRError(self.why)
url, headers, fields, file_field = self._request(wav_bytes, lang)
body, ctype = _multipart(fields, file_field, "audio.wav", wav_bytes, "audio/wav")
headers = {**headers, "Content-Type": ctype}
last = None
for attempt in range(self.retries + 1):
t0 = time.perf_counter()
try:
payload = _post(url, headers, body, self.timeout)
ms = (time.perf_counter() - t0) * 1000
text = payload.get(self.text_field)
if text is None:
# A rename is the most likely failure. Say which keys DID
# arrive instead of returning an empty transcript.
raise ASRError(
f"{self.provider} response has no {self.text_field!r}; "
f"keys present: {sorted(payload)}. Override with "
f"CloudASR(text_field=...) or re-check the provider docs.")
return ASRResult(str(text).strip(), self.provider, ms, payload)
except ASRError:
raise
except Exception as exc: # network / 5xx / timeout
last = exc
if attempt < self.retries:
time.sleep(0.4 * (2 ** attempt)) # 0.4s, 0.8s
raise ASRError(f"{self.provider} failed after {self.retries + 1} attempts: {last}")
# ------------------------------------------------------------------ http
def _multipart(fields: dict, file_field: str, filename: str,
content: bytes, content_type: str):
boundary = "----voicerag" + os.urandom(8).hex()
parts = []
for k, v in fields.items():
parts.append(f"--{boundary}\r\nContent-Disposition: form-data; name=\"{k}\"\r\n\r\n"
f"{v}\r\n".encode())
parts.append(
f"--{boundary}\r\nContent-Disposition: form-data; name=\"{file_field}\"; "
f"filename=\"{filename}\"\r\nContent-Type: {content_type}\r\n\r\n".encode())
parts.append(content)
parts.append(f"\r\n--{boundary}--\r\n".encode())
return b"".join(parts), f"multipart/form-data; boundary={boundary}"
def _post(url: str, headers: dict, body: bytes, timeout: float) -> dict:
import urllib.error
import urllib.request
req = urllib.request.Request(url, data=body, headers=headers, method="POST")
try:
with urllib.request.urlopen(req, timeout=timeout) as r:
return json.loads(r.read().decode("utf-8", "replace") or "{}")
except urllib.error.HTTPError as e:
detail = e.read().decode("utf-8", "replace")[:400]
raise RuntimeError(f"HTTP {e.code}: {detail}") from None
# ------------------------------------------------------------------ selftest
def main() -> int:
import argparse
ap = argparse.ArgumentParser()
ap.add_argument("--provider", default=None)
ap.add_argument("--lang", default="hi")
ap.add_argument("--wav", default=None, help="a real recording; omit for a tone")
ap.add_argument("--selftest", action="store_true")
a = ap.parse_args()
asr = CloudASR(a.provider)
print(f"provider : {asr.provider}\nmodel : {asr.model}\nkey set : {asr.ok}")
if not a.selftest:
return 0
if not asr.ok:
print(f"\n!! {asr.why}\n export the key and re-run.")
return 2
if a.wav:
wav = open(a.wav, "rb").read()
else:
import numpy as np
from src.voice import _to_wav
t = np.arange(16000) / 16000.0
wav = _to_wav((0.3 * np.sin(2 * np.pi * 220 * t)).astype("float32"), 16000)
print("\n(no --wav: sending a 1s tone. Expect an empty or nonsense transcript;\n"
" the point is to prove auth, the URL and the RESPONSE SHAPE.)")
try:
r = asr.transcribe(wav, a.lang)
print(f"\nOK {r.latency_ms:.0f} ms\ntranscript: {r.text!r}")
print(f"response keys: {sorted(r.raw)}")
except ASRError as exc:
print(f"\nFAILED: {exc}")
return 1
return 0
if __name__ == "__main__":
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
raise SystemExit(main())