refactor(soniox): return invalid form bool as a value and raise once in the handler

This commit is contained in:
Dan Lemon 2026-09-15 19:29:42 +02:00
parent 79cda2c2d4
commit ef2a39b3f2
No known key found for this signature in database
GPG key ID: 63D20454AD4DAA0A
4 changed files with 55 additions and 7 deletions

View file

@ -39,7 +39,9 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.llms.soniox.audio_transcription.transformation import ( from litellm.llms.soniox.audio_transcription.transformation import (
SONIOX_HANDLER_ONLY_PARAMS, SONIOX_HANDLER_ONLY_PARAMS,
SonioxAudioTranscriptionConfig, SonioxAudioTranscriptionConfig,
SonioxInvalidBoolParam,
decode_soniox_form_params, decode_soniox_form_params,
raise_soniox_form_error,
) )
from litellm.llms.soniox.common_utils import ( from litellm.llms.soniox.common_utils import (
SONIOX_DEFAULT_CLEANUP, SONIOX_DEFAULT_CLEANUP,
@ -228,6 +230,8 @@ class SonioxAudioTranscriptionHandler:
base_url: Final = get_soniox_api_base(api_base) base_url: Final = get_soniox_api_base(api_base)
decoded: Final = decode_soniox_form_params(optional_params) 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) # Server-side clamps. Caller-supplied poll settings (from request kwargs)
# are bounded so an authenticated caller cannot force a worker into a # are bounded so an authenticated caller cannot force a worker into a

View file

@ -10,8 +10,9 @@ contract of `base_llm_http_handler.audio_transcriptions`.
""" """
from collections.abc import Mapping from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType from types import MappingProxyType
from typing import Any, Final from typing import Any, Final, NoReturn
from httpx import Headers, Response from httpx import Headers, Response
from pydantic import TypeAdapter, ValidationError from pydantic import TypeAdapter, ValidationError
@ -68,7 +69,19 @@ SONIOX_BOOL_PARAMS: Final[frozenset[str]] = frozenset(
_JSON_CONTAINER: Final = TypeAdapter[dict[str, object] | list[object]](dict[str, object] | list[object]) _JSON_CONTAINER: Final = TypeAdapter[dict[str, object] | list[object]](dict[str, object] | list[object])
def _decode_form_value(key: str, value: object) -> 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): if not isinstance(value, str):
return value return value
if key in SONIOX_BOOL_PARAMS: if key in SONIOX_BOOL_PARAMS:
@ -77,7 +90,7 @@ def _decode_form_value(key: str, value: object) -> object:
return True return True
if lowered in ("false", "0"): if lowered in ("false", "0"):
return False return False
raise SonioxException(message=f"`{key}` must be a boolean, got {value!r}", status_code=400, headers=None) return SonioxInvalidBoolParam(key=key, value=value)
if key in SONIOX_JSON_PARAMS and value.lstrip()[:1] in ("{", "["): if key in SONIOX_JSON_PARAMS and value.lstrip()[:1] in ("{", "["):
try: try:
return _JSON_CONTAINER.validate_json(value) return _JSON_CONTAINER.validate_json(value)
@ -88,9 +101,15 @@ def _decode_form_value(key: str, value: object) -> object:
return value return value
def decode_soniox_form_params(optional_params: Mapping[str, object]) -> Mapping[str, object]: 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.""" """Multipart form fields reach the proxy as strings; restore the JSON types Soniox expects."""
return MappingProxyType({key: _decode_form_value(key, value) for key, value in optional_params.items()}) 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): class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig):

View file

@ -339,6 +339,24 @@ class TestPollLimitsClamping:
a worker on tight poll loops. 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): def test_should_clamp_poll_interval_to_minimum(self):
from litellm.llms.soniox.common_utils import SONIOX_MIN_POLL_INTERVAL from litellm.llms.soniox.common_utils import SONIOX_MIN_POLL_INTERVAL

View file

@ -9,7 +9,9 @@ import pytest
from litellm.llms.soniox.audio_transcription.transformation import ( from litellm.llms.soniox.audio_transcription.transformation import (
SonioxAudioTranscriptionConfig, SonioxAudioTranscriptionConfig,
SonioxInvalidBoolParam,
decode_soniox_form_params, decode_soniox_form_params,
raise_soniox_form_error,
) )
from litellm.llms.soniox.common_utils import SonioxException from litellm.llms.soniox.common_utils import SonioxException
from litellm.llms.soniox.types import ( from litellm.llms.soniox.types import (
@ -751,11 +753,16 @@ class TestDecodeSonioxFormParams:
"language_hints_strict": expected, "language_hints_strict": expected,
} }
def test_should_raise_400_on_non_boolean_string(self): 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: with pytest.raises(SonioxException) as exc_info:
decode_soniox_form_params({"enable_speaker_diarization": "yes"}) raise_soniox_form_error(SonioxInvalidBoolParam(key="enable_speaker_diarization", value="yes"))
assert exc_info.value.status_code == 400 assert exc_info.value.status_code == 400
assert "enable_speaker_diarization" in str(exc_info.value) assert "enable_speaker_diarization" in str(exc_info.value)
assert "yes" in str(exc_info.value)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"raw, expected", "raw, expected",