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:
Krrish Dholakia 2025-02-05 14:48:32 -08:00
parent fbe3c58372
commit d9d8c06668
3 changed files with 82 additions and 6 deletions

View file

@ -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][

View file

@ -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
),
)

View file

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