mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
fix(responses): keep a custom tool_choice type in the bridged echo and type the new tests
This commit is contained in:
parent
cce1d2087b
commit
f5a410585a
3 changed files with 32 additions and 17 deletions
|
|
@ -27,9 +27,10 @@ from openai.types.chat.chat_completion_named_tool_choice_param import (
|
|||
)
|
||||
from openai.types.responses import ResponseFunctionToolCall
|
||||
from openai.types.responses.response_create_params import ResponseInputParam
|
||||
from openai.types.responses.tool_choice_custom_param import ToolChoiceCustomParam
|
||||
from openai.types.responses.tool_choice_function_param import ToolChoiceFunctionParam
|
||||
from openai.types.responses.tool_param import FunctionToolParam
|
||||
from pydantic import TypeAdapter
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -272,13 +273,21 @@ class LiteLLMCompletionResponsesConfig:
|
|||
|
||||
@staticmethod
|
||||
def _transform_tool_choice_for_responses_api_response(tool_choice: object) -> ToolChoice:
|
||||
normalized: Final = LiteLLMCompletionResponsesConfig._transform_tool_choice(tool_choice)
|
||||
match normalized:
|
||||
case None:
|
||||
return "auto"
|
||||
case {"type": "function", "function": {"name": str(function_name)}}:
|
||||
if tool_choice is None:
|
||||
return "auto"
|
||||
try:
|
||||
return _RESPONSES_API_TOOL_CHOICE_ADAPTER.validate_python(tool_choice)
|
||||
except ValidationError:
|
||||
return LiteLLMCompletionResponsesConfig._chat_tool_choice_as_responses_api_tool_choice(tool_choice)
|
||||
|
||||
@staticmethod
|
||||
def _chat_tool_choice_as_responses_api_tool_choice(tool_choice: object) -> ToolChoice:
|
||||
match tool_choice, LiteLLMCompletionResponsesConfig._transform_tool_choice(tool_choice):
|
||||
case {"type": "custom"}, {"function": {"name": str(custom_name)}}:
|
||||
return ToolChoiceCustomParam(type="custom", name=custom_name)
|
||||
case _, {"type": "function", "function": {"name": str(function_name)}}:
|
||||
return ToolChoiceFunctionParam(type="function", name=function_name)
|
||||
case _:
|
||||
case _, normalized:
|
||||
return _RESPONSES_API_TOOL_CHOICE_ADAPTER.validate_python(normalized)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -1426,7 +1427,9 @@ class TestToolChoiceTransformation:
|
|||
[
|
||||
({"type": "function", "name": "run_command"}, {"type": "function", "name": "run_command"}),
|
||||
({"type": "function", "function": {"name": "run_command"}}, {"type": "function", "name": "run_command"}),
|
||||
({"type": "custom", "name": "ApplyPatch"}, {"type": "function", "name": "ApplyPatch"}),
|
||||
({"type": "custom", "name": "ApplyPatch"}, {"type": "custom", "name": "ApplyPatch"}),
|
||||
({"type": "custom", "custom": {"name": "ApplyPatch"}}, {"type": "custom", "name": "ApplyPatch"}),
|
||||
({"type": "function"}, "required"),
|
||||
({"type": "tool"}, "required"),
|
||||
({"type": "auto"}, "auto"),
|
||||
("required", "required"),
|
||||
|
|
@ -1434,14 +1437,16 @@ class TestToolChoiceTransformation:
|
|||
(None, "auto"),
|
||||
],
|
||||
)
|
||||
def test_transform_tool_choice_for_responses_api_response(self, request_tool_choice, expected):
|
||||
result = LiteLLMCompletionResponsesConfig._transform_tool_choice_for_responses_api_response(
|
||||
def test_transform_tool_choice_for_responses_api_response(
|
||||
self, request_tool_choice: object, expected: str | dict[str, str]
|
||||
) -> None:
|
||||
result: Final = LiteLLMCompletionResponsesConfig._transform_tool_choice_for_responses_api_response(
|
||||
request_tool_choice
|
||||
)
|
||||
assert result == expected
|
||||
|
||||
def test_non_streamed_response_echoes_named_tool_choice_in_responses_api_shape(self):
|
||||
chat_completion_response = ModelResponse(
|
||||
def test_non_streamed_response_echoes_named_tool_choice_in_responses_api_shape(self) -> None:
|
||||
chat_completion_response: Final = ModelResponse(
|
||||
id="chatcmpl-named-tool-choice",
|
||||
created=1748575031,
|
||||
model="claude-haiku-4-5",
|
||||
|
|
@ -1465,7 +1470,7 @@ class TestToolChoiceTransformation:
|
|||
],
|
||||
)
|
||||
|
||||
responses_api_response = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
|
||||
responses_api_response: Final = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
|
||||
request_input="Run the command pwd.",
|
||||
responses_api_request={"tool_choice": {"type": "function", "name": "run_command"}},
|
||||
chat_completion_response=chat_completion_response,
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ spend tracking stores, so a follow-up previous_response_id still finds the conve
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -657,8 +658,8 @@ def _tool_call_chunk(finish_reason: str | None = None) -> ModelResponseStream:
|
|||
)
|
||||
|
||||
|
||||
def test_streamed_named_tool_choice_is_echoed_in_responses_api_shape():
|
||||
iterator = LiteLLMCompletionStreamingIterator(
|
||||
def test_streamed_named_tool_choice_is_echoed_in_responses_api_shape() -> None:
|
||||
iterator: Final = LiteLLMCompletionStreamingIterator(
|
||||
model="claude-haiku-4-5",
|
||||
litellm_custom_stream_wrapper=_FakeStreamWrapper([_tool_call_chunk(finish_reason="tool_calls")]),
|
||||
request_input="Run the command pwd.",
|
||||
|
|
@ -670,9 +671,9 @@ def test_streamed_named_tool_choice_is_echoed_in_responses_api_shape():
|
|||
litellm_metadata={},
|
||||
)
|
||||
|
||||
events = list(iterator)
|
||||
events: Final = list(iterator)
|
||||
|
||||
response_events = [event for event in events if getattr(event, "type", None) in RESPONSE_ID_EVENT_TYPES]
|
||||
response_events: Final = [event for event in events if getattr(event, "type", None) in RESPONSE_ID_EVENT_TYPES]
|
||||
assert [event.type for event in response_events] == [
|
||||
"response.created",
|
||||
"response.in_progress",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue