From 935e580b0d1b6245500fca10936901c96a3934cb Mon Sep 17 00:00:00 2001 From: WaelRabah11 Date: Wed, 29 Apr 2026 01:50:52 +0200 Subject: [PATCH] feat(audio-transcription): add streaming support for whisper + hosted_vllm Thread `stream` through `get_optional_params_transcription` and the OpenAI whisper / hosted_vllm transformations so audio transcription endpoints return SSE chunks (`text/event-stream`) instead of buffering into a single TranscriptionResponse. - Whisper / hosted_vllm transformations: filter optional_params to supported OpenAI params, stringify bools and join lists for httpx multipart compatibility, normalize audio via process_audio_file. - Skip the verbose_json override when stream=True (verbose_json is incompatible with streaming); preserve the override otherwise so cost calc still gets `duration`. - BaseLLMHTTPHandler: branch on data["stream"] to return a TranscriptionStreamingResponse iterator wired through transform_audio_transcription_streaming_chunk. - Proxy: coerce form-data "true"/"false" strings to bool before dispatch. - Tests: 17 unit tests covering transformation branches, param threading, default chunk pass-through, and sync/streaming HTTP handler paths. Signed-off-by: WaelRabah11 --- litellm/constants.py | 1 + .../audio_transcription/transformation.py | 20 +- litellm/llms/custom_httpx/llm_http_handler.py | 87 ++++++- .../transcriptions/transformation.py | 29 ++- .../transcriptions/gpt_transformation.py | 1 + .../transcriptions/whisper_transformation.py | 50 +++- litellm/main.py | 97 ++++++-- litellm/proxy/proxy_server.py | 15 ++ litellm/types/llms/openai.py | 1 + litellm/types/utils.py | 33 +++ litellm/utils.py | 3 + .../test_audio_transcription_streaming.py | 228 ++++++++++++++++++ 12 files changed, 522 insertions(+), 43 deletions(-) create mode 100644 tests/test_litellm/llms/openai/transcriptions/test_audio_transcription_streaming.py diff --git a/litellm/constants.py b/litellm/constants.py index 32e016cd7a0..393962719d1 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -668,6 +668,7 @@ OPENAI_TRANSCRIPTION_PARAMS = [ "language", "response_format", "timestamp_granularities", + "stream", ] OPENAI_EMBEDDING_PARAMS = ["dimensions", "encoding_format", "user"] diff --git a/litellm/llms/base_llm/audio_transcription/transformation.py b/litellm/llms/base_llm/audio_transcription/transformation.py index 3574996e48e..761cec06537 100644 --- a/litellm/llms/base_llm/audio_transcription/transformation.py +++ b/litellm/llms/base_llm/audio_transcription/transformation.py @@ -1,6 +1,14 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, List, Optional, Union +from typing import ( + TYPE_CHECKING, + Any, + AsyncIterator, + Iterator, + List, + Optional, + Union, +) import httpx @@ -80,6 +88,16 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC): "AudioTranscriptionConfig does not need a response transformation for audio transcription models" ) + def transform_audio_transcription_streaming_chunk( + self, + chunk: bytes, + ) -> bytes: + """ + Hook for providers to translate streaming chunks to OpenAI-compatible + SSE bytes. Default: pass-through (provider already speaks OpenAI SSE). + """ + return chunk + def transform_request( self, model: str, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index b9ada079f6b..2de0d89c93d 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -7,6 +7,7 @@ from typing import ( AsyncIterator, Coroutine, Dict, + Iterator, List, Literal, Optional, @@ -108,6 +109,7 @@ from litellm.types.utils import ( FileTypes, LiteLLMBatch, TranscriptionResponse, + TranscriptionStreamingResponse, ) from litellm.types.vector_store_files import ( VectorStoreFileContentResponse, @@ -1179,6 +1181,53 @@ class BaseLLMHTTPHandler: raw_response=response, ) + def _build_audio_transcription_streaming_response( + self, + provider_config: BaseAudioTranscriptionConfig, + response: httpx.Response, + is_async: bool, + ) -> TranscriptionStreamingResponse: + """ + Wrap an upstream streaming httpx response so callers (proxy / SDK) can + iterate raw SSE bytes. The provider config gets a hook + (`transform_audio_transcription_streaming_chunk`) to translate + non-OpenAI-compatible chunks; default is pass-through. + """ + if is_async: + + async def _aiter() -> AsyncIterator[bytes]: + try: + async for chunk in response.aiter_bytes(): + if not chunk: + continue + yield provider_config.transform_audio_transcription_streaming_chunk( + chunk + ) + finally: + await response.aclose() + + iterator: Any = _aiter() + else: + + def _iter() -> Iterator[bytes]: + try: + for chunk in response.iter_bytes(): + if not chunk: + continue + yield provider_config.transform_audio_transcription_streaming_chunk( + chunk + ) + finally: + response.close() + + iterator = _iter() + + return TranscriptionStreamingResponse( + iterator=iterator, + is_async=is_async, + response_headers=dict(response.headers), + ) + def audio_transcriptions( self, model: str, @@ -1197,7 +1246,13 @@ class BaseLLMHTTPHandler: headers: Optional[Dict[str, Any]] = None, provider_config: Optional[BaseAudioTranscriptionConfig] = None, shared_session: Optional["ClientSession"] = None, - ) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]: + ) -> Union[ + TranscriptionResponse, + TranscriptionStreamingResponse, + Coroutine[ + Any, Any, Union[TranscriptionResponse, TranscriptionStreamingResponse] + ], + ]: if provider_config is None: raise ValueError( f"No provider config found for model: {model} and provider: {custom_llm_provider}" @@ -1243,6 +1298,8 @@ class BaseLLMHTTPHandler: if client is None or not isinstance(client, HTTPHandler): client = _get_httpx_client() + is_streaming = bool(isinstance(data, dict) and data.get("stream")) + try: # Make the POST request - clean and simple, always use data and files response = client.post( @@ -1251,13 +1308,23 @@ class BaseLLMHTTPHandler: data=data, files=files, json=( - data if files is None and isinstance(data, dict) else None + data + if files is None and isinstance(data, dict) and not is_streaming + else None ), # Use json param only when no files and data is dict timeout=timeout, + stream=is_streaming, ) except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + if is_streaming: + return self._build_audio_transcription_streaming_response( + provider_config=provider_config, + response=response, + is_async=False, + ) + return self._transform_audio_transcription_response( provider_config=provider_config, model=model, @@ -1285,7 +1352,7 @@ class BaseLLMHTTPHandler: headers: Optional[Dict[str, Any]] = None, provider_config: Optional[BaseAudioTranscriptionConfig] = None, shared_session: Optional["ClientSession"] = None, - ) -> TranscriptionResponse: + ) -> Union[TranscriptionResponse, TranscriptionStreamingResponse]: if provider_config is None: raise ValueError( f"No provider config found for model: {model} and provider: {custom_llm_provider}" @@ -1318,6 +1385,8 @@ class BaseLLMHTTPHandler: else: async_httpx_client = client + is_streaming = bool(isinstance(data, dict) and data.get("stream")) + try: # Make the async POST request - clean and simple, always use data and files response = await async_httpx_client.post( @@ -1326,13 +1395,23 @@ class BaseLLMHTTPHandler: data=data, files=files, json=( - data if files is None and isinstance(data, dict) else None + data + if files is None and isinstance(data, dict) and not is_streaming + else None ), # Use json param only when no files and data is dict timeout=timeout, + stream=is_streaming, ) except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + if is_streaming: + return self._build_audio_transcription_streaming_response( + provider_config=provider_config, + response=response, + is_async=True, + ) + return self._transform_audio_transcription_response( provider_config=provider_config, model=model, diff --git a/litellm/llms/hosted_vllm/transcriptions/transformation.py b/litellm/llms/hosted_vllm/transcriptions/transformation.py index e726ee33abf..ea1525d07a7 100644 --- a/litellm/llms/hosted_vllm/transcriptions/transformation.py +++ b/litellm/llms/hosted_vllm/transcriptions/transformation.py @@ -6,6 +6,7 @@ from typing import Optional, Union import httpx +from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.audio_transcription.transformation import ( AudioTranscriptionRequestData, ) @@ -55,11 +56,29 @@ class HostedVLLMAudioTranscriptionConfig(OpenAIWhisperAudioTranscriptionConfig): litellm_params: dict, ) -> AudioTranscriptionRequestData: """ - Transform the audio transcription request + Transform the audio transcription request into multipart form-data. + + vLLM speaks the OpenAI transcription protocol but does not support + verbose_json output, so we skip the parent's verbose_json override and + pass through whatever response_format the caller specified. + Filtering / coercion of optional_params follows the same rules as the + parent (only supported params; bools → str; lists → comma-joined). """ + data: dict = {"model": model} + for key in self.get_supported_openai_params(model): + value = optional_params.get(key) + if value is None: + continue + if isinstance(value, bool): + data[key] = "true" if value else "false" + elif isinstance(value, (list, tuple)): + data[key] = ",".join(str(v) for v in value) + else: + data[key] = value - data = {"model": model, "file": audio_file, **optional_params} + processed = process_audio_file(audio_file) + files = { + "file": (processed.filename, processed.file_content, processed.content_type) + } - return AudioTranscriptionRequestData( - data=data, - ) + return AudioTranscriptionRequestData(data=data, files=files) diff --git a/litellm/llms/openai/transcriptions/gpt_transformation.py b/litellm/llms/openai/transcriptions/gpt_transformation.py index 34621c44e22..84eecf96259 100644 --- a/litellm/llms/openai/transcriptions/gpt_transformation.py +++ b/litellm/llms/openai/transcriptions/gpt_transformation.py @@ -22,6 +22,7 @@ class OpenAIGPTAudioTranscriptionConfig(OpenAIWhisperAudioTranscriptionConfig): "response_format", "temperature", "include", + "stream", ] def transform_audio_transcription_request( diff --git a/litellm/llms/openai/transcriptions/whisper_transformation.py b/litellm/llms/openai/transcriptions/whisper_transformation.py index fa507e1bc26..97e310fd99d 100644 --- a/litellm/llms/openai/transcriptions/whisper_transformation.py +++ b/litellm/llms/openai/transcriptions/whisper_transformation.py @@ -2,6 +2,7 @@ from typing import List, Optional, Union from httpx import Headers, Response +from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.audio_transcription.transformation import ( AudioTranscriptionRequestData, BaseAudioTranscriptionConfig, @@ -58,6 +59,7 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): "response_format", "temperature", "timestamp_granularities", + "stream", ] def map_openai_params( @@ -103,20 +105,46 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): litellm_params: dict, ) -> AudioTranscriptionRequestData: """ - Transform the audio transcription request + Transform the audio transcription request into multipart form-data. + + Form fields (model + supported OpenAI params) go in `data`; the audio + file is normalized into `files={"file": (filename, bytes, content_type)}` + so the shared HTTP handler can build a proper multipart upload. + + Only supported OpenAI params are forwarded — httpx multipart only + accepts str/bytes/int/float/None, so unrelated dicts (e.g. router + metadata) and bools (e.g. `stream=True`) would otherwise blow up + encoding. Bools get stringified to "true"/"false"; lists are + joined with commas (OpenAI accepts both forms). """ - data = {"model": model, "file": audio_file, **optional_params} + # 'verbose_json' provides 'duration' for cost calc but is incompatible + # with streaming. Skip the override when stream=True; cost calc falls + # back to file-derived duration in litellm.main.transcription(). + is_streaming = bool(optional_params.get("stream")) + effective_params = dict(optional_params) + if not is_streaming: + existing_format = effective_params.get("response_format") + if existing_format in (None, "text", "json"): + effective_params["response_format"] = "verbose_json" - if "response_format" not in data or ( - data["response_format"] == "text" or data["response_format"] == "json" - ): - data["response_format"] = ( - "verbose_json" # ensures 'duration' is received - used for cost calculation - ) + data: dict = {"model": model} + for key in self.get_supported_openai_params(model): + value = effective_params.get(key) + if value is None: + continue + if isinstance(value, bool): + data[key] = "true" if value else "false" + elif isinstance(value, (list, tuple)): + data[key] = ",".join(str(v) for v in value) + else: + data[key] = value - return AudioTranscriptionRequestData( - data=data, - ) + processed = process_audio_file(audio_file) + files = { + "file": (processed.filename, processed.file_content, processed.content_type) + } + + return AudioTranscriptionRequestData(data=data, files=files) def get_error_class( self, error_message: str, status_code: int, headers: Union[dict, Headers] diff --git a/litellm/main.py b/litellm/main.py index daa0fb063aa..a6bca7237e8 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -127,6 +127,7 @@ from litellm.utils import ( TextCompletionResponse, TextCompletionStreamWrapper, TranscriptionResponse, + TranscriptionStreamingResponse, Usage, _get_model_info_helper, add_provider_specific_params_to_optional_params, @@ -6352,7 +6353,9 @@ async def amoderation( @client -async def atranscription(*args, **kwargs) -> TranscriptionResponse: +async def atranscription( + *args, **kwargs +) -> Union[TranscriptionResponse, TranscriptionStreamingResponse]: """ Calls openai + azure whisper endpoints. @@ -6380,16 +6383,20 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse: init_response = await loop.run_in_executor(None, func_with_context) if isinstance(init_response, dict): response = TranscriptionResponse(**init_response) - elif isinstance(init_response, TranscriptionResponse): ## CACHING SCENARIO + elif isinstance( + init_response, (TranscriptionResponse, TranscriptionStreamingResponse) + ): ## CACHING / STREAMING SCENARIO response = init_response elif asyncio.iscoroutine(init_response): response = await init_response # type: ignore else: # Call the synchronous function using run_in_executor response = await loop.run_in_executor(None, func_with_context) - if not isinstance(response, TranscriptionResponse): + if not isinstance( + response, (TranscriptionResponse, TranscriptionStreamingResponse) + ): raise ValueError( - f"Invalid response from transcription provider, expected TranscriptionResponse, but got {type(response)}" + f"Invalid response from transcription provider, expected TranscriptionResponse or TranscriptionStreamingResponse, but got {type(response)}" ) # Store duration in _hidden_params for cost calculation without @@ -6433,6 +6440,7 @@ def transcription( ] = None, timestamp_granularities: Optional[List[Literal["word", "segment"]]] = None, temperature: Optional[int] = None, # openai defaults this to 0 + stream: Optional[bool] = None, ## LITELLM PARAMS ## user: Optional[str] = None, timeout=600, # default to 10 minutes @@ -6442,7 +6450,11 @@ def transcription( max_retries: Optional[int] = None, custom_llm_provider=None, **kwargs, -) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]: +) -> Union[ + TranscriptionResponse, + TranscriptionStreamingResponse, + Coroutine[Any, Any, Union[TranscriptionResponse, TranscriptionStreamingResponse]], +]: """ Calls openai + azure whisper endpoints. @@ -6493,6 +6505,7 @@ def transcription( response_format=response_format, timestamp_granularities=timestamp_granularities, temperature=temperature, + stream=stream, custom_llm_provider=custom_llm_provider, **non_default_params, ) @@ -6516,7 +6529,15 @@ def transcription( ) response: Optional[ - Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]] + Union[ + TranscriptionResponse, + TranscriptionStreamingResponse, + Coroutine[ + Any, + Any, + Union[TranscriptionResponse, TranscriptionStreamingResponse], + ], + ] ] = None provider_config = ProviderConfigManager.get_provider_audio_transcription_config( @@ -6579,22 +6600,54 @@ def transcription( # set API KEY api_key = api_key or litellm.api_key or litellm.openai_key or get_secret("OPENAI_API_KEY") # type: ignore - response = openai_audio_transcriptions.audio_transcriptions( - model=model, - audio_file=file, - optional_params=optional_params, - model_response=model_response, - atranscription=atranscription, - client=client, - timeout=timeout, - logging_obj=litellm_logging_obj, - max_retries=max_retries, - api_base=api_base, - api_key=api_key, - provider_config=provider_config, - litellm_params=litellm_params_dict, - shared_session=shared_session, - ) + + # Streaming transcription bypasses the openai SDK (which surfaces + # parsed delta events) and goes through the http_handler so callers + # get raw SSE bytes from the upstream provider. + if optional_params.get("stream") and provider_config is not None: + response = base_llm_http_handler.audio_transcriptions( + model=model, + audio_file=file, + optional_params=optional_params, + litellm_params=litellm_params_dict, + model_response=model_response, + atranscription=atranscription, + client=( + client + if client is not None + and ( + isinstance(client, HTTPHandler) + or isinstance(client, AsyncHTTPHandler) + ) + else None + ), + timeout=timeout, + max_retries=max_retries, + logging_obj=litellm_logging_obj, + api_base=api_base, + api_key=api_key, + custom_llm_provider=custom_llm_provider, + headers={}, + provider_config=provider_config, + shared_session=shared_session, + ) + else: + response = openai_audio_transcriptions.audio_transcriptions( + model=model, + audio_file=file, + optional_params=optional_params, + model_response=model_response, + atranscription=atranscription, + client=client, + timeout=timeout, + logging_obj=litellm_logging_obj, + max_retries=max_retries, + api_base=api_base, + api_key=api_key, + provider_config=provider_config, + litellm_params=litellm_params_dict, + shared_session=shared_session, + ) elif provider_config is not None: response = base_llm_http_handler.audio_transcriptions( model=model, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3982d096e69..953bab8079a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -103,6 +103,7 @@ from litellm.types.utils import ( ModelResponseStream, TextCompletionResponse, TokenCountResponse, + TranscriptionStreamingResponse, ) from litellm.utils import ( _invalidate_model_cost_lowercase_map, @@ -8055,6 +8056,11 @@ async def audio_transcriptions( form_data = await get_form_data(request) data = {key: value for key, value in form_data.items() if key != "file"} + # Form fields arrive as strings; coerce `stream` to bool so downstream + # provider transformation sees a real boolean. + if "stream" in data and isinstance(data["stream"], str): + data["stream"] = data["stream"].strip().lower() == "true" + # Include original request and headers in the data data = await add_litellm_data_to_request( data=data, @@ -8161,6 +8167,15 @@ async def audio_transcriptions( if callback_headers: fastapi_response.headers.update(callback_headers) + # Streaming transcription: wrap raw-byte iterator in StreamingResponse + # so FastAPI flushes SSE chunks instead of trying to JSON-serialize. + if isinstance(response, TranscriptionStreamingResponse): + return StreamingResponse( + content=response, + media_type="text/event-stream", + headers=dict(fastapi_response.headers), + ) + return response except Exception as e: await proxy_logging_obj.post_call_failure_hook( diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 2fd0c4ea970..07131a1f245 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1051,6 +1051,7 @@ OpenAIAudioTranscriptionOptionalParams = Literal[ "response_format", "timestamp_granularities", "include", + "stream", ] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a212d56c1ae..367a3d5a4d3 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2406,6 +2406,39 @@ class TranscriptionUsageTokensObject(BaseModel): input_token_details: TranscriptionUsageInputTokenDetailsObject +class TranscriptionStreamingResponse: + """ + Wraps a streaming audio transcription response. + + Holds an async iterator of raw SSE byte chunks from the upstream provider + so the proxy / SDK caller can stream them through to the client. + """ + + def __init__( + self, + iterator: Any, + is_async: bool = True, + hidden_params: Optional[dict] = None, + response_headers: Optional[dict] = None, + ): + self._iterator = iterator + self._is_async = is_async + self._hidden_params: dict = hidden_params or {} + self._response_headers: Optional[dict] = response_headers + + def __aiter__(self): + return self._iterator.__aiter__() + + async def __anext__(self): + return await self._iterator.__anext__() + + def __iter__(self): + return self._iterator.__iter__() + + def __next__(self): + return self._iterator.__next__() + + class TranscriptionResponse(OpenAIObject): text: Optional[str] = None usage: Optional[ diff --git a/litellm/utils.py b/litellm/utils.py index e1ad1db63ef..fdbb3d19c30 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -196,6 +196,7 @@ from litellm.types.utils import ( TextChoices, TextCompletionResponse, TranscriptionResponse, + TranscriptionStreamingResponse, Usage, all_litellm_params, ) @@ -2985,6 +2986,7 @@ def get_optional_params_transcription( response_format: Optional[str] = None, temperature: Optional[int] = None, timestamp_granularities: Optional[List[Literal["word", "segment"]]] = None, + stream: Optional[bool] = None, drop_params: Optional[bool] = None, **kwargs, ): @@ -3006,6 +3008,7 @@ def get_optional_params_transcription( "response_format": None, "temperature": None, # openai defaults this to 0 "timestamp_granularities": None, + "stream": None, } non_default_params = { diff --git a/tests/test_litellm/llms/openai/transcriptions/test_audio_transcription_streaming.py b/tests/test_litellm/llms/openai/transcriptions/test_audio_transcription_streaming.py new file mode 100644 index 00000000000..3d62d85176a --- /dev/null +++ b/tests/test_litellm/llms/openai/transcriptions/test_audio_transcription_streaming.py @@ -0,0 +1,228 @@ +""" +Tests for streaming audio transcription support. + +Covers: +- whisper transformation skips verbose_json override on stream=True +- get_supported_openai_params includes "stream" for whisper + gpt +- get_optional_params_transcription threads `stream` through +- BaseAudioTranscriptionConfig.transform_audio_transcription_streaming_chunk + is pass-through by default +- BaseLLMHTTPHandler.audio_transcriptions returns TranscriptionStreamingResponse + when data["stream"] is truthy +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from litellm.llms.base_llm.audio_transcription.transformation import ( + BaseAudioTranscriptionConfig, +) +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.llms.openai.transcriptions.gpt_transformation import ( + OpenAIGPTAudioTranscriptionConfig, +) +from litellm.llms.openai.transcriptions.whisper_transformation import ( + OpenAIWhisperAudioTranscriptionConfig, +) +from litellm.types.utils import ( + TranscriptionResponse, + TranscriptionStreamingResponse, +) +from litellm.utils import get_optional_params_transcription + + +class TestWhisperTransformation: + """The verbose_json override must NOT fire when stream=True.""" + + def test_streaming_skips_verbose_json_override(self): + config = OpenAIWhisperAudioTranscriptionConfig() + result = config.transform_audio_transcription_request( + model="whisper-1", + audio_file=("audio.wav", b"fake-bytes", "audio/wav"), + optional_params={"stream": True, "response_format": "json"}, + litellm_params={}, + ) + assert isinstance(result.data, dict) + # bools are stringified for httpx multipart compatibility + assert result.data["stream"] == "true" + assert result.data["response_format"] == "json" + + def test_non_streaming_applies_verbose_json_override(self): + """When stream is absent, override fires for json/text/missing format.""" + config = OpenAIWhisperAudioTranscriptionConfig() + result = config.transform_audio_transcription_request( + model="whisper-1", + audio_file=("audio.wav", b"fake-bytes", "audio/wav"), + optional_params={"response_format": "json"}, + litellm_params={}, + ) + assert isinstance(result.data, dict) + assert result.data["response_format"] == "verbose_json" + + def test_streaming_preserves_explicit_format(self): + """User-specified srt under stream=True must not be clobbered.""" + config = OpenAIWhisperAudioTranscriptionConfig() + result = config.transform_audio_transcription_request( + model="whisper-1", + audio_file=("a.wav", b"x", "audio/wav"), + optional_params={"stream": True, "response_format": "srt"}, + litellm_params={}, + ) + assert isinstance(result.data, dict) + assert result.data["response_format"] == "srt" + + def test_supported_params_includes_stream(self): + config = OpenAIWhisperAudioTranscriptionConfig() + assert "stream" in config.get_supported_openai_params(model="whisper-1") + + +class TestGPTTranscriptionTransformation: + def test_supported_params_includes_stream(self): + config = OpenAIGPTAudioTranscriptionConfig() + assert "stream" in config.get_supported_openai_params(model="gpt-4o-transcribe") + + +class TestGetOptionalParamsTranscription: + """The `stream` kwarg must flow through into the optional_params dict.""" + + def test_stream_threaded_through_for_openai(self): + params = get_optional_params_transcription( + model="whisper-1", + custom_llm_provider="openai", + stream=True, + ) + assert params.get("stream") is True + + def test_stream_default_none_not_emitted(self): + params = get_optional_params_transcription( + model="whisper-1", + custom_llm_provider="openai", + ) + # default None means non-default-params filter drops it + assert "stream" not in params or params.get("stream") is None + + +class TestBaseStreamingChunkHook: + """Default hook is pass-through; subclass override path works.""" + + def test_default_pass_through(self): + # BaseAudioTranscriptionConfig is abstract — use a concrete impl + config = OpenAIWhisperAudioTranscriptionConfig() + chunk = b'data: {"text":"hi"}\n\n' + assert config.transform_audio_transcription_streaming_chunk(chunk) == chunk + + def test_method_is_defined_on_base(self): + # Defensive: the hook must exist on the base class so providers + # that don't override still inherit pass-through. + assert hasattr( + BaseAudioTranscriptionConfig, + "transform_audio_transcription_streaming_chunk", + ) + + +class TestHTTPHandlerStreamingBranch: + """audio_transcriptions sync path returns TranscriptionStreamingResponse.""" + + def test_sync_streaming_returns_streaming_response(self): + handler = BaseLLMHTTPHandler() + provider_config = OpenAIWhisperAudioTranscriptionConfig() + + # Mock httpx.Response with iter_bytes + mock_response = MagicMock() + mock_response.headers = {"content-type": "text/event-stream"} + mock_response.iter_bytes = MagicMock( + return_value=iter([b'data: {"text":"hi"}\n\n', b"data: [DONE]\n\n"]) + ) + mock_response.close = MagicMock() + + mock_client = MagicMock() + mock_client.post = MagicMock(return_value=mock_response) + + with patch( + "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + return_value=mock_client, + ): + result = handler.audio_transcriptions( + model="whisper-1", + audio_file=("a.wav", b"x", "audio/wav"), + optional_params={"stream": True}, + litellm_params={}, + model_response=TranscriptionResponse(), + timeout=60.0, + max_retries=0, + logging_obj=MagicMock(), + api_key="sk-test", + api_base="https://api.openai.com", + custom_llm_provider="openai", + client=None, + atranscription=False, + provider_config=provider_config, + ) + + assert isinstance(result, TranscriptionStreamingResponse) + # client was invoked with stream=True + _, kwargs = mock_client.post.call_args + assert kwargs.get("stream") is True + + # Iterator yields the upstream chunks unchanged (default pass-through) + chunks = list(result) + assert chunks == [b'data: {"text":"hi"}\n\n', b"data: [DONE]\n\n"] + + def test_sync_non_streaming_returns_transcription_response(self): + handler = BaseLLMHTTPHandler() + provider_config = OpenAIWhisperAudioTranscriptionConfig() + + mock_response = MagicMock() + mock_response.headers = {"content-type": "application/json"} + mock_response.json = MagicMock(return_value={"text": "hello world"}) + + mock_client = MagicMock() + mock_client.post = MagicMock(return_value=mock_response) + + with patch( + "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + return_value=mock_client, + ): + result = handler.audio_transcriptions( + model="whisper-1", + audio_file=("a.wav", b"x", "audio/wav"), + optional_params={}, + litellm_params={}, + model_response=TranscriptionResponse(), + timeout=60.0, + max_retries=0, + logging_obj=MagicMock(), + api_key="sk-test", + api_base="https://api.openai.com", + custom_llm_provider="openai", + client=None, + atranscription=False, + provider_config=provider_config, + ) + + assert isinstance(result, TranscriptionResponse) + _, kwargs = mock_client.post.call_args + assert kwargs.get("stream") is False + + +class TestProxyStreamCoercion: + """Proxy must coerce string 'true'/'false' from form data to bool.""" + + @pytest.mark.parametrize( + "raw,expected", + [ + ("true", True), + ("True", True), + ("TRUE", True), + ("false", False), + ("False", False), + ("0", False), + ], + ) + def test_stream_coercion_logic(self, raw, expected): + # Mirror the inline coercion in proxy_server.audio_transcriptions + data = {"stream": raw} + if "stream" in data and isinstance(data["stream"], str): + data["stream"] = data["stream"].strip().lower() == "true" + assert data["stream"] is expected