mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(responses/handler): ensure azure model name is stripped before sending to provider
Fixes model name error
This commit is contained in:
parent
dda4532f10
commit
5f8b1e9fd2
5 changed files with 63 additions and 17 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 ##############
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue