diff --git a/litellm/llms/soniox/audio_transcription/handler.py b/litellm/llms/soniox/audio_transcription/handler.py index a335caa65c2..da0ac7d7a06 100644 --- a/litellm/llms/soniox/audio_transcription/handler.py +++ b/litellm/llms/soniox/audio_transcription/handler.py @@ -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 diff --git a/litellm/llms/soniox/audio_transcription/transformation.py b/litellm/llms/soniox/audio_transcription/transformation.py index 8507ae73305..f5a2f7d67cc 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,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.""" diff --git a/litellm/types/llms/soniox.py b/litellm/types/llms/soniox.py new file mode 100644 index 00000000000..64d22f965ee --- /dev/null +++ b/litellm/types/llms/soniox.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..64b6fa8a018 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,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"]'