diff --git a/litellm/llms/soniox/audio_transcription/handler.py b/litellm/llms/soniox/audio_transcription/handler.py index a335caa65c2..1725132829e 100644 --- a/litellm/llms/soniox/audio_transcription/handler.py +++ b/litellm/llms/soniox/audio_transcription/handler.py @@ -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, diff --git a/litellm/llms/soniox/audio_transcription/transformation.py b/litellm/llms/soniox/audio_transcription/transformation.py index 8507ae73305..9f224944073 100644 --- a/litellm/llms/soniox/audio_transcription/transformation.py +++ b/litellm/llms/soniox/audio_transcription/transformation.py @@ -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.""" diff --git a/litellm/llms/soniox/types.py b/litellm/llms/soniox/types.py new file mode 100644 index 00000000000..64d22f965ee --- /dev/null +++ b/litellm/llms/soniox/types.py @@ -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] diff --git a/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py b/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py index 45753d4ee7b..077b9142d6e 100644 --- a/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py +++ b/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py @@ -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. diff --git a/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py b/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py index eadf870cb61..99d7c3d4411 100644 --- a/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py +++ b/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py @@ -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"]'