fix(responses/handler): ensure azure model name is stripped before sending to provider

Fixes model name error
This commit is contained in:
Krrish Dholakia 2025-07-01 17:01:27 -07:00
parent dda4532f10
commit 5f8b1e9fd2
5 changed files with 63 additions and 17 deletions

View file

@ -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(

View file

@ -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(

View file

@ -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 ##############

View file

@ -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(

View file

@ -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"