This commit is contained in:
Dan Lemon 2026-09-12 08:25:10 -04:00 • committed by GitHub
commit 37b35991ba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 254 additions and 35 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,9 @@ 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,
decode_soniox_form_params,
)
from litellm.llms.soniox.common_utils import (
SONIOX_DEFAULT_CLEANUP,
@ -58,6 +62,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 +211,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 +227,41 @@ 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)
# 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,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,38 @@ 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:
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]:
"""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,8 +9,15 @@ 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.llms.soniox.types import (
SonioxContext,
SonioxContextGeneralEntry,
SonioxTranslation,
SonioxTranslationTerm,
)
from litellm.types.utils import TranscriptionResponse
@ -678,3 +685,85 @@ 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_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"]'