mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 5d08d5e309 into 955b26ac08
This commit is contained in:
commit
34349a0aee
2 changed files with 93 additions and 6 deletions
|
|
@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal
|
|||
|
||||
import httpx
|
||||
from openai.types.responses import ResponseReasoningItem
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
|
|
@ -21,6 +22,9 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
_JSON_VALUE_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
|
||||
|
||||
class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
# Parameters not supported by Azure Responses API
|
||||
AZURE_UNSUPPORTED_PARAMS = ["context_management"]
|
||||
|
|
@ -96,18 +100,46 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
|
||||
# Then filter out status from message items
|
||||
if isinstance(validated_input, list):
|
||||
filtered_input: Final[list[Any]] = []
|
||||
filtered_input: Final[list[object]] = []
|
||||
for item in validated_input:
|
||||
if isinstance(item, dict) and item.get("type") == "message":
|
||||
# Filter out status field from message items
|
||||
filtered_item = {k: v for k, v in item.items() if k != "status"}
|
||||
filtered_input.append(filtered_item)
|
||||
filtered_item = (
|
||||
{k: v for k, v in item.items() if k != "status"}
|
||||
if isinstance(item, dict) and item.get("type") == "message"
|
||||
else item
|
||||
)
|
||||
if filtered_item.get("type") == "additional_tools":
|
||||
validated_additional_tools_item: JsonValue = _JSON_VALUE_ADAPTER.validate_python(filtered_item)
|
||||
filtered_input.append(self._normalize_additional_tools_item(validated_additional_tools_item))
|
||||
else:
|
||||
filtered_input.append(item)
|
||||
filtered_input.append(filtered_item)
|
||||
return cast(ResponseInputParam, filtered_input)
|
||||
|
||||
return validated_input
|
||||
|
||||
@staticmethod
|
||||
def _normalize_namespace_tool_description(tool: JsonValue) -> JsonValue:
|
||||
if not isinstance(tool, dict) or tool.get("type") != "namespace":
|
||||
return tool
|
||||
|
||||
description: Final = tool.get("description")
|
||||
if not isinstance(description, str) or description.strip():
|
||||
return tool
|
||||
|
||||
name: Final = tool.get("name")
|
||||
fallback_description: Final = name.strip() if isinstance(name, str) and name.strip() else "namespace"
|
||||
return {**tool, "description": fallback_description}
|
||||
|
||||
@classmethod
|
||||
def _normalize_additional_tools_item(cls, item: JsonValue) -> JsonValue:
|
||||
if not isinstance(item, dict) or item.get("type") != "additional_tools":
|
||||
return item
|
||||
|
||||
tools: Final = item.get("tools")
|
||||
if not isinstance(tools, list):
|
||||
return item
|
||||
|
||||
return {**item, "tools": [cls._normalize_namespace_tool_description(tool) for tool in tools]}
|
||||
|
||||
def transform_responses_api_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -529,6 +529,61 @@ class TestAzureResponsesAPIConfig:
|
|||
|
||||
assert "tools" not in response_api_params
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("description", "expected_description"),
|
||||
[
|
||||
("", "functions"),
|
||||
(" \t ", "functions"),
|
||||
("Existing description", "Existing description"),
|
||||
],
|
||||
)
|
||||
def test_azure_responses_api_normalizes_additional_tools_namespace_description(
|
||||
self, description, expected_description
|
||||
):
|
||||
namespace_tool = {
|
||||
"type": "namespace",
|
||||
"name": "functions",
|
||||
"description": description,
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "ping",
|
||||
"description": "Return pong.",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
"strict": False,
|
||||
}
|
||||
],
|
||||
}
|
||||
function_tool = {
|
||||
"type": "function",
|
||||
"name": "empty_description_is_allowed_here",
|
||||
"description": "",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
}
|
||||
input_items = [
|
||||
{
|
||||
"type": "additional_tools",
|
||||
"role": "developer",
|
||||
"tools": [namespace_tool, function_tool],
|
||||
}
|
||||
]
|
||||
original_input_items = deepcopy(input_items)
|
||||
|
||||
result = self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input=input_items,
|
||||
response_api_optional_request_params={},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["input"][0]["tools"][0] == {
|
||||
**namespace_tool,
|
||||
"description": expected_description,
|
||||
}
|
||||
assert result["input"][0]["tools"][1] == function_tool
|
||||
assert input_items == original_input_items
|
||||
|
||||
def test_azure_responses_api_context_management_unsupported(self):
|
||||
"""Test that context_management is not in Azure supported params.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue