badrex-endpoint / _audio_io.py
filipok's picture
redeploy: sync endpoint-badrex handler
afe40c9 verified
Raw
History Blame Contribute Delete
4.15 kB
"""Shared audio-input bounds + decode for the ASR endpoint handlers.
Byte-identical across the endpoint*/ dirs that decode audio to a numpy array
(MMS / badrex / seamless / whisper). Each endpoint dir is deployed standalone
(pushed to its own HF model repo), so it cannot import from src/ at runtime —
hence one copy per dir. Keep the copies in sync; a CI hash-diff guards it
(see docs/endpoint-adversarial-review-2026-06-20.md, Phase 2 smoke test).
All bounds are enforced BEFORE b64decode / torchaudio.load / model inference, so
an oversized or malformed payload is rejected up front rather than exhausting
memory/CPU/GPU. Bounds are env-overridable so a bigger box can raise them without
a code change.
"""
from __future__ import annotations
import base64
import io
import os
import numpy as np
SAMPLE_RATE = 16000
# Reject before decode. Encoded cap mirrors the HF Endpoints ~100 MB request cap;
# decoded-sample cap bounds a 90-minute 16 kHz mono broadcast.
MAX_ENCODED_BYTES = int(os.environ.get("ASR_MAX_ENCODED_BYTES", str(100 * 1024 * 1024)))
MAX_DECODED_SAMPLES = int(os.environ.get("ASR_MAX_DECODED_SAMPLES", str(SAMPLE_RATE * 60 * 90)))
MAX_LIST_LEN = int(os.environ.get("ASR_MAX_LIST_LEN", str(MAX_DECODED_SAMPLES * 2)))
class AudioRequestError(ValueError):
"""A request that fails a pre-decode bound. Handlers map this to a
{"error": ..., "status": "bad_request"} response so the runner can tell a
client error apart from a model failure."""
def decode_to_raw(audio_input):
"""Return raw bytes for the str/bytes path, or None for the list path
(the caller routes None to array_from_list). Enforces the encoded-size +
base64-validity bounds on the way."""
if isinstance(audio_input, str):
# base64 expands ~4/3; reject obviously-oversized strings before decode.
if len(audio_input) > MAX_ENCODED_BYTES * 4 // 3 + 4:
raise AudioRequestError("encoded payload exceeds size limit")
try:
raw = base64.b64decode(audio_input, validate=True) # fail closed on non-base64
except Exception as exc:
raise AudioRequestError(f"invalid base64: {exc}") from exc
if len(raw) > MAX_ENCODED_BYTES:
raise AudioRequestError("decoded payload exceeds size limit")
return raw
if isinstance(audio_input, list):
return None
raw = bytes(audio_input)
if len(raw) > MAX_ENCODED_BYTES:
raise AudioRequestError("payload exceeds size limit")
return raw
def array_from_list(audio_input):
"""float32 mono array from a JSON list, with shape + finiteness bounds."""
if len(audio_input) > MAX_LIST_LEN:
raise AudioRequestError("list input exceeds sample limit")
arr = np.asarray(audio_input, dtype=np.float32)
if not np.isfinite(arr).all():
raise AudioRequestError("list input contains NaN/Inf")
if arr.ndim > 1:
arr = arr.mean(axis=0) # (channels, samples) -> mono
return arr
def check_samples(arr):
"""Post-decode duration bound (covers compressed payloads that expand past
the encoded cap once decoded)."""
if arr.shape[0] > MAX_DECODED_SAMPLES:
raise AudioRequestError("decoded audio exceeds duration limit")
return arr
def load_and_validate_audio(audio_input):
"""Decode any accepted input (base64 str / raw bytes / float list) to a
bounded float32 mono 16 kHz numpy array, raising AudioRequestError for any
input that fails a bound. torchaudio decodes m4a/webm/opus/mp4 via ffmpeg, so
no client-side conversion is needed."""
raw = decode_to_raw(audio_input)
if raw is None: # list path
return check_samples(array_from_list(audio_input))
import torchaudio
waveform, sr = torchaudio.load(io.BytesIO(raw)) # (channels, samples)
if sr != SAMPLE_RATE:
waveform = torchaudio.transforms.Resample(orig_freq=sr, new_freq=SAMPLE_RATE)(waveform)
arr = waveform.numpy()
if arr.ndim > 1:
arr = arr.mean(axis=0)
return check_samples(arr.astype(np.float32))