mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(convert_dict_to_response.py): only convert if response is the response_format tool call passed in
Fixes https://github.com/BerriAI/litellm/issues/8241
This commit is contained in:
parent
fbe3c58372
commit
d9d8c06668
3 changed files with 82 additions and 6 deletions
|
|
@ -7,6 +7,7 @@ from typing import Dict, Iterable, List, Literal, Optional, Union
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionDeltaToolCall,
|
||||
ChatCompletionMessageToolCall,
|
||||
|
|
@ -313,6 +314,23 @@ class LiteLLMResponseObjectHandler:
|
|||
return transformed_logprobs
|
||||
|
||||
|
||||
def _should_convert_tool_call_to_json_mode(
|
||||
tool_calls: Optional[List[ChatCompletionMessageToolCall]] = None,
|
||||
convert_tool_call_to_json_mode: Optional[bool] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if tool calls should be converted to JSON mode
|
||||
"""
|
||||
if (
|
||||
convert_tool_call_to_json_mode
|
||||
and tool_calls is not None
|
||||
and len(tool_calls) == 1
|
||||
and tool_calls[0]["function"]["name"] == RESPONSE_FORMAT_TOOL_NAME
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def convert_to_model_response_object( # noqa: PLR0915
|
||||
response_object: Optional[dict] = None,
|
||||
model_response_object: Optional[
|
||||
|
|
@ -397,10 +415,9 @@ def convert_to_model_response_object( # noqa: PLR0915
|
|||
|
||||
message: Optional[Message] = None
|
||||
finish_reason: Optional[str] = None
|
||||
if (
|
||||
convert_tool_call_to_json_mode
|
||||
and tool_calls is not None
|
||||
and len(tool_calls) == 1
|
||||
if _should_convert_tool_call_to_json_mode(
|
||||
tool_calls=tool_calls,
|
||||
convert_tool_call_to_json_mode=convert_tool_call_to_json_mode,
|
||||
):
|
||||
# to support 'json_schema' logic on older models
|
||||
json_mode_content_str: Optional[str] = tool_calls[0][
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, List, Optional, Union
|
|||
from httpx._models import Headers, Response
|
||||
|
||||
import litellm
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
convert_to_azure_openai_messages,
|
||||
)
|
||||
|
|
@ -201,14 +202,14 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
_tool_choice = ChatCompletionToolChoiceObjectParam(
|
||||
type="function",
|
||||
function=ChatCompletionToolChoiceFunctionParam(
|
||||
name=schema_name
|
||||
name=RESPONSE_FORMAT_TOOL_NAME
|
||||
),
|
||||
)
|
||||
|
||||
_tool = ChatCompletionToolParam(
|
||||
type="function",
|
||||
function=ChatCompletionToolParamFunctionChunk(
|
||||
name=schema_name, parameters=json_schema
|
||||
name=RESPONSE_FORMAT_TOOL_NAME, parameters=json_schema
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -283,3 +283,61 @@ def test_azure_openai_gpt_4o_naming(monkeypatch):
|
|||
print(mock_post.call_args.kwargs)
|
||||
|
||||
assert "tool_calls" not in mock_post.call_args.kwargs
|
||||
|
||||
|
||||
def test_azure_gpt_4o_with_tool_call_and_response_format():
|
||||
from litellm import completion
|
||||
from typing import Optional
|
||||
from pydantic import BaseModel
|
||||
import litellm
|
||||
|
||||
class InvestigationOutput(BaseModel):
|
||||
alert_explanation: Optional[str] = None
|
||||
investigation: Optional[str] = None
|
||||
conclusions_and_possible_root_causes: Optional[str] = None
|
||||
next_steps: Optional[str] = None
|
||||
related_logs: Optional[str] = None
|
||||
app_or_infra: Optional[str] = None
|
||||
external_links: Optional[str] = None
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_time",
|
||||
"description": "Returns the current date and time",
|
||||
"strict": True,
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"timezone": {
|
||||
"type": "string",
|
||||
"description": "The timezone to get the current time for (e.g., 'UTC', 'America/New_York')",
|
||||
}
|
||||
},
|
||||
"required": ["timezone"],
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
response = litellm.completion(
|
||||
model="azure/gpt-4o",
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a tool-calling AI assist provided with common devops and IT tools that you can use to troubleshoot problems or answer questions.\nWhenever possible you MUST first use tools to investigate then answer the question.",
|
||||
},
|
||||
{"role": "user", "content": "What is the current date and time in NYC?"},
|
||||
],
|
||||
drop_params=True,
|
||||
temperature=0.00000001,
|
||||
tools=tools,
|
||||
tool_choice="auto",
|
||||
response_format=InvestigationOutput, # commenting this line will cause the output to be correct
|
||||
)
|
||||
|
||||
assert response.choices[0].finish_reason == "tool_calls"
|
||||
|
||||
print(response.to_json())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue