From 5f8b1e9fd224b29e0a9473b3a1cf88be4b09f2ab Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 1 Jul 2025 17:01:27 -0700 Subject: [PATCH] fix(responses/handler): ensure azure model name is stripped before sending to provider Fixes model name error --- .../handler.py | 5 ++- .../transformation.py | 2 ++ .../llms/azure/responses/transformation.py | 34 +++++++++++++++---- litellm/llms/custom_httpx/llm_http_handler.py | 21 ++++++++---- tests/test_litellm/test_router.py | 18 +++++++++- 5 files changed, 63 insertions(+), 17 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index ea5e8b4c8dd..f2eeaf04554 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -75,9 +75,7 @@ class ResponsesToCompletionBridgeHandler: custom_llm_provider=custom_llm_provider, ) - def completion( - self, *args, **kwargs - ) -> Union[ + def completion(self, *args, **kwargs) -> Union[ Coroutine[Any, Any, Union["ModelResponse", "CustomStreamWrapper"]], "ModelResponse", "CustomStreamWrapper", @@ -106,6 +104,7 @@ class ResponsesToCompletionBridgeHandler: litellm_params=litellm_params, headers=headers, litellm_logging_obj=logging_obj, + client=kwargs.get("client"), ) result = responses( diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 9ba42ffff58..7d10f6b2166 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -121,6 +121,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): litellm_params: dict, headers: dict, litellm_logging_obj: "LiteLLMLoggingObj", + client: Optional[Any] = None, ) -> dict: from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams @@ -186,6 +187,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): "input": input_items, "litellm_logging_obj": litellm_logging_obj, **litellm_params, + "client": client, } verbose_logger.debug( diff --git a/litellm/llms/azure/responses/transformation.py b/litellm/llms/azure/responses/transformation.py index e6f48179e49..e3d37c8a15a 100644 --- a/litellm/llms/azure/responses/transformation.py +++ b/litellm/llms/azure/responses/transformation.py @@ -20,8 +20,33 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] ) -> dict: return BaseAzureLLM._base_validate_azure_environment( - headers=headers, - litellm_params=litellm_params + headers=headers, litellm_params=litellm_params + ) + + def get_stripped_model_name(self, model: str) -> str: + # if "responses/" is in the model name, remove it + if "responses/" in model: + model = model.replace("responses/", "") + if "o_series" in model: + model = model.replace("o_series/", "") + return model + + def transform_responses_api_request( + self, + model: str, + input: Union[str, ResponseInputParam], + response_api_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + """No transform applied since inputs are in OpenAI spec already""" + stripped_model_name = self.get_stripped_model_name(model) + return dict( + ResponsesAPIRequestParams( + model=stripped_model_name, + input=input, + **response_api_optional_request_params, + ) ) def get_complete_url( @@ -46,11 +71,8 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): "https://litellm8397336933.openai.azure.com/openai/responses?api-version=2024-05-01-preview" """ return BaseAzureLLM._get_base_azure_url( - api_base=api_base, - litellm_params=litellm_params, - route="/openai/responses" + api_base=api_base, litellm_params=litellm_params, route="/openai/responses" ) - ######################################################### ########## DELETE RESPONSE API TRANSFORMATION ############## diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 75013aea83c..e08b909b2a9 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1008,12 +1008,12 @@ class BaseLLMHTTPHandler: """ Shared logic for preparing audio transcription requests. Returns: (headers, complete_url, data, files) - """ + """ # Handle the response based on type from litellm.llms.base_llm.audio_transcription.transformation import ( AudioTranscriptionRequestData, ) - + headers = provider_config.validate_environment( api_key=api_key, headers=headers or {}, @@ -1038,11 +1038,13 @@ class BaseLLMHTTPHandler: optional_params=optional_params, litellm_params=litellm_params, ) - + # All providers now return AudioTranscriptionRequestData if not isinstance(transformed_result, AudioTranscriptionRequestData): - raise ValueError(f"Provider {provider_config.__class__.__name__} must return AudioTranscriptionRequestData") - + raise ValueError( + f"Provider {provider_config.__class__.__name__} must return AudioTranscriptionRequestData" + ) + data = transformed_result.data files = transformed_result.files @@ -1143,7 +1145,9 @@ class BaseLLMHTTPHandler: headers=headers, data=data, files=files, - json=data if files is None and isinstance(data, dict) else None, # Use json param only when no files and data is dict + json=( + data if files is None and isinstance(data, dict) else None + ), # Use json param only when no files and data is dict timeout=timeout, ) except Exception as e: @@ -1214,7 +1218,9 @@ class BaseLLMHTTPHandler: headers=headers, data=data, files=files, - json=data if files is None and isinstance(data, dict) else None, # Use json param only when no files and data is dict + json=( + data if files is None and isinstance(data, dict) else None + ), # Use json param only when no files and data is dict timeout=timeout, ) except Exception as e: @@ -1432,6 +1438,7 @@ class BaseLLMHTTPHandler: Handles responses API requests. When _is_async=True, returns a coroutine instead of making the call directly. """ + if _is_async: # Return the async coroutine if called with _is_async=True return self.async_response_api_handler( diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 2ea9162e024..6b078804a31 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -608,7 +608,7 @@ def test_router_responses_api_bridge(): """ Test that router.responses_api_bridge returns the correct response """ - import respx + from litellm.llms.custom_httpx.http_handler import HTTPHandler router = litellm.Router( model_list=[ @@ -636,3 +636,19 @@ def test_router_responses_api_bridge(): messages=[{"role": "user", "content": "Hello, world!"}], ) assert mock_responses.call_count == 1 + + ## CONFIRM MODEL NAME IS STRIPPED + client = HTTPHandler() + + with patch.object(client, "post", return_value=AsyncMock()) as mock_post: + result = router.completion( + model="[IP-approved] o3-pro", + messages=[{"role": "user", "content": "Hello, world!"}], + client=client, + ) + assert mock_post.call_count == 1 + assert ( + mock_post.call_args.kwargs["url"] + == "https://webhook.site/fba79dae-220a-4bb7-9a3a-8caa49604e55/openai/v1/responses?api-version=preview" + ) + assert mock_post.call_args.kwargs["json"]["model"] == "webinterface-o3-pro"