This commit is contained in:
Dan Lemon 2026-10-01 23:20:32 +00:00 • committed by GitHub
commit 9129e68651
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 303 additions and 36 deletions

View file

@ -19,9 +19,11 @@ import asyncio
import math
import time
from collections.abc import Coroutine, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
import httpx
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm.litellm_core_utils.audio_utils.utils import (
@ -35,7 +37,11 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
)
from litellm.llms.soniox.audio_transcription.transformation import (
SONIOX_HANDLER_ONLY_PARAMS,
SonioxAudioTranscriptionConfig,
SonioxInvalidBoolParam,
decode_soniox_form_params,
raise_soniox_form_error,
)
from litellm.llms.soniox.common_utils import (
SONIOX_DEFAULT_CLEANUP,
@ -58,6 +64,31 @@ else:
LiteLLMLoggingObj = Any
_CLEANUP_TARGETS: Final = TypeAdapter(tuple[str, ...])
_NOT_FORWARDED_TO_SONIOX: Final[frozenset[str]] = frozenset(
(*SONIOX_HANDLER_ONLY_PARAMS, "audio_url", "file_id", "response_format", "language")
)
def _optional_str(value: object) -> str | None:
return None if value is None else str(value)
def _float_or_default(value: object, default: float) -> float:
try:
parsed: Final = float(str(value))
except ValueError:
return default
return parsed if math.isfinite(parsed) else default
def _int_or_default(value: object, default: int) -> int:
try:
return int(str(value))
except (ValueError, OverflowError):
return default
class _TranscriptionMeta(TypedDict, total=False):
"""Fields the handler reads from a Soniox transcription object."""
@ -182,7 +213,7 @@ class SonioxAudioTranscriptionHandler:
) -> tuple[
dict[str, str], # auth headers
str, # api_base (no trailing slash)
dict[str, object], # body for POST /v1/transcriptions (without file_id/audio_url)
Mapping[str, object], # body for POST /v1/transcriptions (without file_id/audio_url)
_HandlerOptions, # handler-only options (poll interval, cleanup, ...)
]:
# Validate env -> auth headers.
@ -198,52 +229,43 @@ class SonioxAudioTranscriptionHandler:
base_url: Final = get_soniox_api_base(api_base)
# Operate on a local copy so we don't mutate the caller's dict
# (the caller may reuse `optional_params` for retries or logging).
params: Final = dict(optional_params)
# Pull handler-only kwargs out of params so they aren't sent
# to Soniox.
poll_interval = float(params.pop("soniox_polling_interval", SONIOX_DEFAULT_POLL_INTERVAL))
try:
max_attempts = int(params.pop("soniox_max_polling_attempts", SONIOX_DEFAULT_MAX_POLL_ATTEMPTS))
except (ValueError, OverflowError):
max_attempts = SONIOX_DEFAULT_MAX_POLL_ATTEMPTS
cleanup_raw: Final = params.pop("soniox_cleanup", SONIOX_DEFAULT_CLEANUP)
if cleanup_raw is None:
cleanup: list[str] = []
elif isinstance(cleanup_raw, str):
cleanup = [cleanup_raw]
else:
cleanup = list(cleanup_raw)
filename_override: Final = params.pop("filename", None)
decoded: Final = decode_soniox_form_params(optional_params)
if isinstance(decoded, SonioxInvalidBoolParam):
raise_soniox_form_error(decoded)
# Server-side clamps. Caller-supplied poll settings (from request kwargs)
# are bounded so an authenticated caller cannot force a worker into a
# tight poll loop (zero interval) or pin it indefinitely (huge attempt
# count). Total polling time is bounded by
# SONIOX_MAX_POLL_ATTEMPTS * SONIOX_MAX_POLL_INTERVAL.
if not math.isfinite(poll_interval):
poll_interval = SONIOX_DEFAULT_POLL_INTERVAL
clamped_poll_interval: Final = max(SONIOX_MIN_POLL_INTERVAL, min(poll_interval, SONIOX_MAX_POLL_INTERVAL))
clamped_max_attempts: Final = max(1, min(max_attempts, SONIOX_MAX_POLL_ATTEMPTS))
poll_interval: Final = _float_or_default(decoded.get("soniox_polling_interval"), SONIOX_DEFAULT_POLL_INTERVAL)
max_attempts: Final = _int_or_default(
decoded.get("soniox_max_polling_attempts"), SONIOX_DEFAULT_MAX_POLL_ATTEMPTS
)
cleanup_raw: Final = decoded.get("soniox_cleanup", SONIOX_DEFAULT_CLEANUP)
cleanup: Final[tuple[str, ...]] = (
()
if cleanup_raw is None
else (cleanup_raw,)
if isinstance(cleanup_raw, str)
else _CLEANUP_TARGETS.validate_python(cleanup_raw)
)
# response_format is handled by LiteLLM post-processing, not Soniox.
handler_opts: Final[_HandlerOptions] = {
"poll_interval": clamped_poll_interval,
"max_attempts": clamped_max_attempts,
"poll_interval": max(SONIOX_MIN_POLL_INTERVAL, min(poll_interval, SONIOX_MAX_POLL_INTERVAL)),
"max_attempts": max(1, min(max_attempts, SONIOX_MAX_POLL_ATTEMPTS)),
"cleanup": cleanup,
"filename_override": filename_override,
"audio_url": params.pop("audio_url", None),
"file_id": params.pop("file_id", None),
"response_format": params.pop("response_format", None),
"filename_override": _optional_str(decoded.get("filename")),
"audio_url": _optional_str(decoded.get("audio_url")),
"file_id": _optional_str(decoded.get("file_id")),
"response_format": _optional_str(decoded.get("response_format")),
}
# Soniox does not accept `language` directly; map_openai_params should
# already have translated it, but drop any leftover to be safe.
params.pop("language", None)
return auth_headers, base_url, params, handler_opts
provider_params: Final = MappingProxyType(
{key: value for key, value in decoded.items() if key not in _NOT_FORWARDED_TO_SONIOX}
)
return auth_headers, base_url, provider_params, handler_opts
def _build_create_body(
self,

View file

@ -9,9 +9,13 @@ async API requires multiple HTTP calls and does not fit the single-request
contract of `base_llm_http_handler.audio_transcriptions`.
"""
from typing import Any, Final
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Any, Final, NoReturn
from httpx import Headers, Response
from pydantic import TypeAdapter, ValidationError
from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
@ -58,6 +62,56 @@ SONIOX_HANDLER_ONLY_PARAMS: Final[list[str]] = [
]
SONIOX_JSON_PARAMS: Final[frozenset[str]] = frozenset({"context", "translation", "language_hints"})
SONIOX_BOOL_PARAMS: Final[frozenset[str]] = frozenset(
{"enable_speaker_diarization", "enable_language_identification", "language_hints_strict"}
)
_JSON_CONTAINER: Final = TypeAdapter[dict[str, object] | list[object]](dict[str, object] | list[object])
@dataclass(frozen=True, slots=True)
class SonioxInvalidBoolParam:
key: str
value: str
def raise_soniox_form_error(error: SonioxInvalidBoolParam) -> NoReturn:
raise SonioxException(
message=f"`{error.key}` must be a boolean, got {error.value!r}", status_code=400, headers=None
)
def _decode_form_value(key: str, value: object) -> object | SonioxInvalidBoolParam:
if not isinstance(value, str):
return value
if key in SONIOX_BOOL_PARAMS:
lowered: Final = value.strip().lower()
if lowered in ("true", "1"):
return True
if lowered in ("false", "0"):
return False
return SonioxInvalidBoolParam(key=key, value=value)
if key in SONIOX_JSON_PARAMS and value.lstrip()[:1] in ("{", "["):
try:
return _JSON_CONTAINER.validate_json(value)
except ValidationError:
return value
if key == "language_hints":
return tuple(hint.strip() for hint in value.split(",") if hint.strip())
return value
def decode_soniox_form_params(
optional_params: Mapping[str, object],
) -> Mapping[str, object] | SonioxInvalidBoolParam:
"""Multipart form fields reach the proxy as strings; restore the JSON types Soniox expects."""
decoded: Final = {key: _decode_form_value(key, value) for key, value in optional_params.items()}
return next(
(value for value in decoded.values() if isinstance(value, SonioxInvalidBoolParam)),
MappingProxyType(decoded),
)
class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
"""Configuration for Soniox async speech-to-text transcription."""

View file

@ -0,0 +1,32 @@
from collections.abc import Sequence
from typing import Literal
from typing_extensions import ReadOnly, TypedDict
class SonioxContextGeneralEntry(TypedDict):
key: ReadOnly[str]
value: ReadOnly[str]
class SonioxTranslationTerm(TypedDict):
source: ReadOnly[str]
target: ReadOnly[str]
class SonioxContext(TypedDict, total=False):
"""Soniox `context` request field: https://soniox.com/docs/stt/concepts/context"""
general: ReadOnly[Sequence[SonioxContextGeneralEntry]]
text: ReadOnly[str]
terms: ReadOnly[Sequence[str]]
translation_terms: ReadOnly[Sequence[SonioxTranslationTerm]]
class SonioxTranslation(TypedDict, total=False):
"""Soniox `translation` request field: https://soniox.com/docs/translation/stt-translation"""
type: ReadOnly[Literal["one_way", "two_way"]]
target_language: ReadOnly[str]
language_a: ReadOnly[str]
language_b: ReadOnly[str]

View file

@ -339,6 +339,24 @@ class TestPollLimitsClamping:
a worker on tight poll loops.
"""
def test_should_raise_400_when_form_bool_param_is_invalid(self):
handler = SonioxAudioTranscriptionHandler()
with pytest.raises(SonioxException) as exc_info:
handler._prepare(
audio_file=None,
optional_params={
"enable_speaker_diarization": "yes",
"audio_url": "https://example.com/a.wav",
},
litellm_params={},
api_key="sk-test",
api_base=None,
provider_config=SonioxAudioTranscriptionConfig(),
headers={},
)
assert exc_info.value.status_code == 400
assert "enable_speaker_diarization" in str(exc_info.value)
def test_should_clamp_poll_interval_to_minimum(self):
from litellm.llms.soniox.common_utils import SONIOX_MIN_POLL_INTERVAL
@ -780,6 +798,51 @@ class TestPassthroughBodyBuilding:
post_call = next(c for c in client.calls if c["method"] == "POST")
assert "context" not in post_call["json"]
def test_should_send_form_string_params_to_soniox_as_json_types(self, monkeypatch):
"""Simulates the proxy path, where every multipart form field is a string."""
monkeypatch.setattr("time.sleep", lambda *_: None)
responses = {
"POST https://api.soniox.com/v1/transcriptions": [
_make_response({"id": "tx_1"}),
],
"GET https://api.soniox.com/v1/transcriptions/tx_1": [
_make_response({"status": "completed"}),
],
"GET https://api.soniox.com/v1/transcriptions/tx_1/transcript": [
_make_response({"text": "ok", "tokens": []}),
],
"DELETE https://api.soniox.com/v1/transcriptions/tx_1": [
_make_response({}),
],
}
client = _MockSyncClient(responses)
form_params = {
"audio_url": "https://example.com/a.wav",
"context": '{"terms": ["Celebrex"], "general": [{"key": "domain", "value": "Healthcare"}]}',
"translation": '{"type": "one_way", "target_language": "es"}',
"language_hints": "en,es",
"language_hints_strict": "true",
"enable_speaker_diarization": "true",
"enable_language_identification": "false",
"soniox_polling_interval": "0.5",
}
SonioxAudioTranscriptionHandler().audio_transcriptions(
audio_file=None,
optional_params=form_params,
litellm_params={},
atranscription=False,
**_common_call_kwargs(client),
)
body = next(c for c in client.calls if c["method"] == "POST")["json"]
assert body["context"] == {"terms": ["Celebrex"], "general": [{"key": "domain", "value": "Healthcare"}]}
assert body["translation"] == {"type": "one_way", "target_language": "es"}
assert list(body["language_hints"]) == ["en", "es"]
assert body["language_hints_strict"] is True
assert body["enable_speaker_diarization"] is True
assert body["enable_language_identification"] is False
assert "soniox_polling_interval" not in body
assert form_params["context"].startswith("{"), "caller's params must not be mutated"
class TestSecretRedaction:
"""Secret-bearing fields must be redacted before reaching logging callbacks.

View file

@ -9,8 +9,17 @@ import pytest
from litellm.llms.soniox.audio_transcription.transformation import (
SonioxAudioTranscriptionConfig,
SonioxInvalidBoolParam,
decode_soniox_form_params,
raise_soniox_form_error,
)
from litellm.llms.soniox.common_utils import SonioxException
from litellm.llms.soniox.types import (
SonioxContext,
SonioxContextGeneralEntry,
SonioxTranslation,
SonioxTranslationTerm,
)
from litellm.types.utils import TranscriptionResponse
@ -678,3 +687,90 @@ class TestGetErrorClass:
err = cfg.get_error_class(error_message="boom", status_code=500, headers={})
assert isinstance(err, SonioxException)
assert err.status_code == 500
class TestDecodeSonioxFormParams:
"""Proxy multipart form fields arrive as strings; Soniox needs real JSON types."""
def test_should_decode_json_string_context_with_all_four_sections(self):
context = {
"general": [{"key": "domain", "value": "Healthcare"}],
"text": "Follow-up visit for a patient on blood thinners.",
"terms": ["Celebrex", "Zyrtec"],
"translation_terms": [{"source": "Mr. Smith", "target": "Sr. Smith"}],
}
result = decode_soniox_form_params({"context": json.dumps(context)})
assert result["context"] == context
def test_should_keep_plain_text_context_as_string(self):
result = decode_soniox_form_params({"context": "medical conversation"})
assert result["context"] == "medical conversation"
def test_should_leave_sdk_typed_values_untouched(self):
context = SonioxContext(
general=[SonioxContextGeneralEntry(key="domain", value="Healthcare")],
text="Follow-up visit notes.",
terms=["Celebrex"],
translation_terms=[SonioxTranslationTerm(source="Mr. Smith", target="Sr. Smith")],
)
translation = SonioxTranslation(type="one_way", target_language="es")
result = decode_soniox_form_params(
{
"context": context,
"translation": translation,
"enable_speaker_diarization": True,
"language_hints": ["en", "es"],
}
)
assert result["context"] is context
assert result["translation"] is translation
assert result["enable_speaker_diarization"] is True
assert result["language_hints"] == ["en", "es"]
@pytest.mark.parametrize("raw", ["[inaudible] consultation", "{unbalanced", '{"terms": ["Celebrex"'])
def test_should_keep_bracket_prefixed_free_form_context_as_string(self, raw: str):
assert decode_soniox_form_params({"context": raw})["context"] == raw
def test_should_decode_json_string_translation(self):
result = decode_soniox_form_params({"translation": '{"type": "one_way", "target_language": "es"}'})
assert result["translation"] == {"type": "one_way", "target_language": "es"}
@pytest.mark.parametrize(
"raw, expected",
[("true", True), ("True", True), ("1", True), ("false", False), ("FALSE", False), ("0", False)],
)
def test_should_decode_boolean_strings(self, raw: str, expected: bool):
result = decode_soniox_form_params(
{
"enable_speaker_diarization": raw,
"enable_language_identification": raw,
"language_hints_strict": raw,
}
)
assert result == {
"enable_speaker_diarization": expected,
"enable_language_identification": expected,
"language_hints_strict": expected,
}
def test_should_return_error_value_on_non_boolean_string(self):
result = decode_soniox_form_params({"context": '{"general": []}', "enable_speaker_diarization": "yes"})
assert result == SonioxInvalidBoolParam(key="enable_speaker_diarization", value="yes")
def test_should_map_invalid_bool_error_to_400(self):
with pytest.raises(SonioxException) as exc_info:
raise_soniox_form_error(SonioxInvalidBoolParam(key="enable_speaker_diarization", value="yes"))
assert exc_info.value.status_code == 400
assert "enable_speaker_diarization" in str(exc_info.value)
assert "yes" in str(exc_info.value)
@pytest.mark.parametrize(
"raw, expected",
[('["en", "es"]', ["en", "es"]), ("en,es", ["en", "es"]), ("en, es ,", ["en", "es"]), ("en", ["en"])],
)
def test_should_decode_language_hints_strings(self, raw: str, expected: list):
assert list(decode_soniox_form_params({"language_hints": raw})["language_hints"]) == expected
def test_should_not_json_decode_unrelated_string_params(self):
result = decode_soniox_form_params({"client_reference_id": '["not", "json-decoded"]'})
assert result["client_reference_id"] == '["not", "json-decoded"]'