feat(soniox): decode proxy form fields for context, translation and flags

Multipart form fields reach the proxy as strings, so Soniox received
context and translation as JSON strings and booleans as "true".
Decode them at the handler entry point and add typed shapes for the
context and translation request fields.
This commit is contained in:
Dan Lemon 2026-09-10 21:19:04 +02:00
parent 6c69dd0f72
commit fe6d6d9b7e
No known key found for this signature in database
GPG key ID: 63D20454AD4DAA0A
5 changed files with 211 additions and 14 deletions

View file

@ -22,6 +22,7 @@ from collections.abc import Coroutine, Mapping, Sequence
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 (
@ -36,6 +37,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.llms.soniox.audio_transcription.transformation import (
SonioxAudioTranscriptionConfig,
decode_soniox_form_params,
)
from litellm.llms.soniox.common_utils import (
SONIOX_DEFAULT_CLEANUP,
@ -58,6 +60,13 @@ else:
LiteLLMLoggingObj = Any
_CLEANUP_TARGETS: Final = TypeAdapter(tuple[str, ...])
def _optional_str(value: object) -> str | None:
return None if value is None else str(value)
class _TranscriptionMeta(TypedDict, total=False):
"""Fields the handler reads from a Soniox transcription object."""
@ -198,25 +207,26 @@ 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
# Decoded copy so the caller's dict is never mutated
# (the caller may reuse `optional_params` for retries or logging).
params: Final = dict(optional_params)
params: Final = dict(decode_soniox_form_params(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))
poll_interval = float(str(params.pop("soniox_polling_interval", SONIOX_DEFAULT_POLL_INTERVAL)))
try:
max_attempts = int(params.pop("soniox_max_polling_attempts", SONIOX_DEFAULT_MAX_POLL_ATTEMPTS))
max_attempts = int(str(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)
cleanup: Final[tuple[str, ...]] = (
()
if cleanup_raw is None
else (cleanup_raw,)
if isinstance(cleanup_raw, str)
else _CLEANUP_TARGETS.validate_python(cleanup_raw)
)
filename_override: Final = _optional_str(params.pop("filename", None))
# Server-side clamps. Caller-supplied poll settings (from request kwargs)
# are bounded so an authenticated caller cannot force a worker into a
@ -234,9 +244,9 @@ class SonioxAudioTranscriptionHandler:
"max_attempts": clamped_max_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),
"audio_url": _optional_str(params.pop("audio_url", None)),
"file_id": _optional_str(params.pop("file_id", None)),
"response_format": _optional_str(params.pop("response_format", None)),
}
# Soniox does not accept `language` directly; map_openai_params should

View file

@ -9,9 +9,12 @@ async API requires multiple HTTP calls and does not fit the single-request
contract of `base_llm_http_handler.audio_transcriptions`.
"""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final
from httpx import Headers, Response
from pydantic import TypeAdapter, ValidationError
from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
@ -58,6 +61,40 @@ 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])
def _decode_form_value(key: str, value: object) -> object:
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
raise SonioxException(message=f"`{key}` must be a boolean, got {value!r}", status_code=400, headers=None)
if key in SONIOX_JSON_PARAMS and value.lstrip()[:1] in ("{", "["):
try:
return _JSON_CONTAINER.validate_json(value)
except ValidationError as exc:
raise SonioxException(
message=f"`{key}` is not valid JSON: {exc.errors()[0]['msg']}", status_code=400, headers=None
)
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]:
"""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()})
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

@ -780,6 +780,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,6 +9,7 @@ import pytest
from litellm.llms.soniox.audio_transcription.transformation import (
SonioxAudioTranscriptionConfig,
decode_soniox_form_params,
)
from litellm.llms.soniox.common_utils import SonioxException
from litellm.types.utils import TranscriptionResponse
@ -678,3 +679,75 @@ 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_already_typed_values_untouched(self):
context = {"terms": ["Celebrex"]}
result = decode_soniox_form_params(
{"context": context, "enable_speaker_diarization": True, "language_hints": ["en", "es"]}
)
assert result["context"] is context
assert result["enable_speaker_diarization"] is True
assert result["language_hints"] == ["en", "es"]
def test_should_raise_400_on_malformed_json_context(self):
with pytest.raises(SonioxException) as exc_info:
decode_soniox_form_params({"context": '{"terms": ["Celebrex"'})
assert exc_info.value.status_code == 400
assert "context" in str(exc_info.value)
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_raise_400_on_non_boolean_string(self):
with pytest.raises(SonioxException) as exc_info:
decode_soniox_form_params({"enable_speaker_diarization": "yes"})
assert exc_info.value.status_code == 400
assert "enable_speaker_diarization" 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"]'