From a079857948b8ab23d92890c4ac61d4862291b820 Mon Sep 17 00:00:00 2001 From: Yan Zhu Date: Thu, 13 Aug 2026 18:03:19 +0800 Subject: [PATCH] fix(hosted_vllm): honor custom_llm_provider and /v1 api_base on transcriptions atranscription/aspeech ignored custom_llm_provider, so unprefixed proxy models failed with LLM Provider NOT provided. Hosted vLLM URL joining also doubled /v1 when api_base already ended with it. Co-authored-by: Cursor --- .../transcriptions/transformation.py | 20 ++-- litellm/main.py | 14 ++- .../test_hosted_vllm_audio_transcription.py | 107 ++++++++++++++++++ 3 files changed, 131 insertions(+), 10 deletions(-) create mode 100644 tests/test_litellm/llms/hosted_vllm/transcriptions/test_hosted_vllm_audio_transcription.py diff --git a/litellm/llms/hosted_vllm/transcriptions/transformation.py b/litellm/llms/hosted_vllm/transcriptions/transformation.py index 9a8c67fadc0..edec37b8d86 100644 --- a/litellm/llms/hosted_vllm/transcriptions/transformation.py +++ b/litellm/llms/hosted_vllm/transcriptions/transformation.py @@ -39,13 +39,19 @@ class HostedVLLMAudioTranscriptionConfig(OpenAIWhisperAudioTranscriptionConfig): litellm_params: dict, stream: bool | None = None, ) -> str: - if api_base: - # Remove trailing slashes and ensure clean base URL - api_base = api_base.rstrip("/") - if not api_base.endswith("/v1/audio/transcriptions"): - api_base = f"{api_base}/v1/audio/transcriptions" - return api_base - raise ValueError("api_base must be provided for Hosted VLLM rerank") + if not api_base: + raise ValueError("api_base must be provided for Hosted VLLM transcriptions") + normalized_base: Final = api_base.rstrip("/") + if normalized_base.endswith("/v1/audio/transcriptions"): + return normalized_base + return super().get_complete_url( + api_base=normalized_base, + api_key=api_key, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + stream=stream, + ) def transform_audio_transcription_request( self, diff --git a/litellm/main.py b/litellm/main.py index 16eff5a0f3e..a0da99713d2 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -7541,7 +7541,7 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse: ### PASS ARGS TO Image Generation ### kwargs["atranscription"] = True file: Final = kwargs.get("file", None) - custom_llm_provider = None + custom_llm_provider = kwargs.get("custom_llm_provider", None) try: # Use a partial function to pass your keyword arguments func: Final = partial(transcription, *args, **kwargs) @@ -7550,7 +7550,11 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse: ctx: Final = contextvars.copy_context() func_with_context: Final = partial(ctx.run, func) - _, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None)) + _, custom_llm_provider, _, _ = get_llm_provider( + model=model, + custom_llm_provider=custom_llm_provider, + api_base=kwargs.get("api_base", None), + ) # Await normally init_response: Final = await loop.run_in_executor(None, func_with_context) @@ -7869,7 +7873,11 @@ async def aspeech(*args, **kwargs) -> HttpxBinaryResponseContent: ctx: Final = contextvars.copy_context() func_with_context: Final = partial(ctx.run, func) - _, custom_llm_provider, _, _ = get_llm_provider(model=model, api_base=kwargs.get("api_base", None)) + _, custom_llm_provider, _, _ = get_llm_provider( + model=model, + custom_llm_provider=custom_llm_provider, + api_base=kwargs.get("api_base", None), + ) # Await normally init_response: Final = await loop.run_in_executor(None, func_with_context) diff --git a/tests/test_litellm/llms/hosted_vllm/transcriptions/test_hosted_vllm_audio_transcription.py b/tests/test_litellm/llms/hosted_vllm/transcriptions/test_hosted_vllm_audio_transcription.py new file mode 100644 index 00000000000..5caea4e30a0 --- /dev/null +++ b/tests/test_litellm/llms/hosted_vllm/transcriptions/test_hosted_vllm_audio_transcription.py @@ -0,0 +1,107 @@ +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../../..")) + +import litellm +from litellm.llms.hosted_vllm.transcriptions.transformation import ( + HostedVLLMAudioTranscriptionConfig, +) +from litellm.types.utils import TranscriptionResponse + + +def _complete_url(api_base: str | None) -> str: + return HostedVLLMAudioTranscriptionConfig().get_complete_url( + api_base=api_base, + api_key=None, + model="whisper-1", + optional_params={}, + litellm_params={}, + ) + + +class TestHostedVLLMTranscriptionUrl: + @pytest.mark.parametrize( + "api_base,expected", + [ + ( + "http://vllm.example.com", + "http://vllm.example.com/v1/audio/transcriptions", + ), + ( + "http://vllm.example.com/", + "http://vllm.example.com/v1/audio/transcriptions", + ), + ( + "http://vllm.example.com/v1", + "http://vllm.example.com/v1/audio/transcriptions", + ), + ( + "http://vllm.example.com/v1/", + "http://vllm.example.com/v1/audio/transcriptions", + ), + ( + "http://proxy.example.com/qwen3-asr/v1", + "http://proxy.example.com/qwen3-asr/v1/audio/transcriptions", + ), + ( + "http://vllm.example.com/v1/audio/transcriptions", + "http://vllm.example.com/v1/audio/transcriptions", + ), + ], + ) + def test_get_complete_url(self, api_base: str, expected: str) -> None: + assert _complete_url(api_base) == expected + + def test_get_complete_url_requires_api_base(self) -> None: + with pytest.raises(ValueError, match="api_base must be provided"): + _complete_url(None) + + +class TestAtranscriptionCustomLlmProvider: + @pytest.mark.asyncio + async def test_unprefixed_model_uses_custom_llm_provider(self) -> None: + """ + Proxy/router deployments often register model=qwen3-asr-0.6b with + custom_llm_provider=hosted_vllm (no hosted_vllm/ prefix). Chat honors + that field; atranscription used to ignore it and raise + 'LLM Provider NOT provided'. + """ + with patch( + "litellm.main.transcription", + return_value=TranscriptionResponse(text="hello"), + ) as mock_transcription: + response = await litellm.atranscription( + model="qwen3-asr-0.6b", + file=b"fake-audio", + custom_llm_provider="hosted_vllm", + api_base="http://vllm.example.com/v1", + ) + + assert response.text == "hello" + mock_transcription.assert_called() + + @pytest.mark.asyncio + async def test_unprefixed_model_without_provider_still_fails(self) -> None: + with pytest.raises(Exception, match="LLM Provider NOT provided"): + await litellm.atranscription( + model="qwen3-asr-0.6b", + file=b"fake-audio", + api_base="http://vllm.example.com/v1", + ) + + @pytest.mark.asyncio + async def test_aspeech_unprefixed_model_uses_custom_llm_provider(self) -> None: + with patch("litellm.main.speech", return_value=MagicMock()) as mock_speech: + await litellm.aspeech( + model="tts-model", + input="hello", + voice="alloy", + custom_llm_provider="hosted_vllm", + api_base="http://vllm.example.com/v1", + ) + + mock_speech.assert_called()