mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(anthropic/): fix handling response_format for anthropic messages with anthropic api
This commit is contained in:
parent
f2904cbb4e
commit
feefd18614
4 changed files with 243 additions and 22 deletions
|
|
@ -63,3 +63,4 @@ LITELLM_CHAT_PROVIDERS = [
|
|||
"lm_studio",
|
||||
"galadriel",
|
||||
]
|
||||
RESPONSE_FORMAT_TOOL_NAME = "json_tool_call" # default tool name used when converting response format to tool call
|
||||
|
|
|
|||
|
|
@ -19,9 +19,10 @@ import httpx
|
|||
import requests
|
||||
|
||||
import litellm
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.litellm_core_utils.core_helpers import map_finish_reason
|
||||
from litellm.llms.base_llm.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import anthropic_messages_pt
|
||||
from litellm.llms.base_llm.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.types.llms.anthropic import (
|
||||
AllAnthropicToolsValues,
|
||||
AnthropicComputerTool,
|
||||
|
|
@ -298,6 +299,18 @@ class AnthropicConfig(BaseConfig):
|
|||
new_stop = new_v
|
||||
return new_stop
|
||||
|
||||
def _add_tools_to_optional_params(
|
||||
self, optional_params: dict, tools: List[AllAnthropicToolsValues]
|
||||
) -> dict:
|
||||
if "tools" not in optional_params:
|
||||
optional_params["tools"] = tools
|
||||
else:
|
||||
optional_params["tools"] = [
|
||||
*optional_params["tools"],
|
||||
*tools,
|
||||
]
|
||||
return optional_params
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
|
|
@ -311,7 +324,11 @@ class AnthropicConfig(BaseConfig):
|
|||
if param == "max_completion_tokens":
|
||||
optional_params["max_tokens"] = value
|
||||
if param == "tools":
|
||||
optional_params["tools"] = self._map_tools(value)
|
||||
# check if optional params already has tools
|
||||
tool_value = self._map_tools(value)
|
||||
optional_params = self._add_tools_to_optional_params(
|
||||
optional_params=optional_params, tools=tool_value
|
||||
)
|
||||
if param == "tool_choice" or param == "parallel_tool_calls":
|
||||
_tool_choice: Optional[AnthropicMessagesToolChoice] = (
|
||||
self._map_tool_choice(
|
||||
|
|
@ -333,6 +350,7 @@ class AnthropicConfig(BaseConfig):
|
|||
if param == "top_p":
|
||||
optional_params["top_p"] = value
|
||||
if param == "response_format" and isinstance(value, dict):
|
||||
|
||||
json_schema: Optional[dict] = None
|
||||
if "response_schema" in value:
|
||||
json_schema = value["response_schema"]
|
||||
|
|
@ -344,11 +362,14 @@ class AnthropicConfig(BaseConfig):
|
|||
- You should set tool_choice (see Forcing tool use) to instruct the model to explicitly use that tool
|
||||
- Remember that the model will pass the input to the tool, so the name of the tool and description should be from the model’s perspective.
|
||||
"""
|
||||
_tool_choice = {"name": "json_tool_call", "type": "tool"}
|
||||
|
||||
_tool_choice = {"name": RESPONSE_FORMAT_TOOL_NAME, "type": "tool"}
|
||||
_tool = self._create_json_tool_call_for_response_format(
|
||||
json_schema=json_schema,
|
||||
)
|
||||
optional_params["tools"] = [_tool]
|
||||
optional_params = self._add_tools_to_optional_params(
|
||||
optional_params=optional_params, tools=[_tool]
|
||||
)
|
||||
optional_params["tool_choice"] = _tool_choice
|
||||
optional_params["json_mode"] = True
|
||||
if param == "user":
|
||||
|
|
@ -381,7 +402,9 @@ class AnthropicConfig(BaseConfig):
|
|||
else:
|
||||
_input_schema["properties"] = {"values": json_schema}
|
||||
|
||||
_tool = AnthropicMessagesTool(name="json_tool_call", input_schema=_input_schema)
|
||||
_tool = AnthropicMessagesTool(
|
||||
name=RESPONSE_FORMAT_TOOL_NAME, input_schema=_input_schema
|
||||
)
|
||||
return _tool
|
||||
|
||||
def is_cache_control_set(self, messages: List[AllMessageValues]) -> bool:
|
||||
|
|
@ -537,10 +560,6 @@ class AnthropicConfig(BaseConfig):
|
|||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
optional_params[k] = v
|
||||
|
||||
## Handle Tool Calling
|
||||
if "tools" in optional_params:
|
||||
_is_function_call = True
|
||||
|
||||
## Handle user_id in metadata
|
||||
_litellm_metadata = litellm_params.get("metadata", None)
|
||||
if (
|
||||
|
|
@ -558,6 +577,26 @@ class AnthropicConfig(BaseConfig):
|
|||
|
||||
return data
|
||||
|
||||
def _transform_response_for_json_mode(
|
||||
self,
|
||||
json_mode: Optional[bool],
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Optional[LitellmMessage]:
|
||||
_message: Optional[LitellmMessage] = None
|
||||
if json_mode is True and len(tool_calls) == 1:
|
||||
# check if tool name is the default tool name
|
||||
json_mode_content_str: Optional[str] = None
|
||||
if (
|
||||
"name" in tool_calls[0]["function"]
|
||||
and tool_calls[0]["function"]["name"] == RESPONSE_FORMAT_TOOL_NAME
|
||||
):
|
||||
json_mode_content_str = tool_calls[0]["function"].get("arguments")
|
||||
if json_mode_content_str is not None:
|
||||
_message = AnthropicConfig._convert_tool_response_to_message(
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
return _message
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -629,19 +668,14 @@ class AnthropicConfig(BaseConfig):
|
|||
)
|
||||
|
||||
## HANDLE JSON MODE - anthropic returns single function call
|
||||
if json_mode is True and len(tool_calls) == 1:
|
||||
json_mode_content_str: Optional[str] = tool_calls[0]["function"].get(
|
||||
"arguments"
|
||||
)
|
||||
if json_mode_content_str is not None:
|
||||
_converted_message = (
|
||||
AnthropicConfig._convert_tool_response_to_message(
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
)
|
||||
if _converted_message is not None:
|
||||
completion_response["stop_reason"] = "stop"
|
||||
_message = _converted_message
|
||||
json_mode_message = self._transform_response_for_json_mode(
|
||||
json_mode=json_mode,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
if json_mode_message is not None:
|
||||
completion_response["stop_reason"] = "stop"
|
||||
_message = json_mode_message
|
||||
|
||||
model_response.choices[0].message = _message # type: ignore
|
||||
model_response._hidden_params["original_response"] = completion_response[
|
||||
"content"
|
||||
|
|
|
|||
|
|
@ -263,6 +263,67 @@ class BaseLLMChatTest(ABC):
|
|||
assert content is not None
|
||||
assert len(content) > 0
|
||||
|
||||
def test_tool_call_and_json_response_format(self):
|
||||
"""
|
||||
Test that the tool call and JSON response format is supported by the LLM API
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
from pydantic import BaseModel
|
||||
from litellm.utils import supports_response_schema
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
class RFormat(BaseModel):
|
||||
question: str
|
||||
answer: str
|
||||
|
||||
base_completion_call_args = self.get_base_completion_call_args()
|
||||
if not supports_response_schema(base_completion_call_args["model"], None):
|
||||
pytest.skip("Model does not support response schema")
|
||||
|
||||
try:
|
||||
res = litellm.completion(
|
||||
**base_completion_call_args,
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "response user question with JSON object",
|
||||
},
|
||||
{"role": "user", "content": "Hey! What's the weather in NewYork?"},
|
||||
],
|
||||
tool_choice="required",
|
||||
response_format=RFormat,
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_weather",
|
||||
"description": "Get the current weather in a given location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "The city and state, e.g. San Francisco, CA",
|
||||
},
|
||||
"unit": {
|
||||
"type": "string",
|
||||
"enum": ["celsius", "fahrenheit"],
|
||||
},
|
||||
},
|
||||
"required": ["location"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
assert res is not None
|
||||
|
||||
assert res.choices[0].message.tool_calls is not None
|
||||
except litellm.InternalServerError:
|
||||
pytest.skip("Model is overloaded")
|
||||
|
||||
@pytest.fixture
|
||||
def tool_call_no_arguments(self):
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -829,3 +829,128 @@ def test_anthropic_tool_with_image():
|
|||
)
|
||||
|
||||
assert b64_data in json.dumps(result)
|
||||
|
||||
|
||||
def test_anthropic_map_openai_params_tools_and_json_schema():
|
||||
import json
|
||||
|
||||
args = {
|
||||
"non_default_params": {
|
||||
"response_format": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"schema": {
|
||||
"properties": {
|
||||
"question": {"title": "Question", "type": "string"},
|
||||
"answer": {"title": "Answer", "type": "string"},
|
||||
},
|
||||
"required": ["question", "answer"],
|
||||
"title": "RFormat",
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
},
|
||||
"name": "RFormat",
|
||||
"strict": True,
|
||||
},
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_weather",
|
||||
"description": "Get the current weather in a given location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "The city and state, e.g. San Francisco, CA",
|
||||
},
|
||||
"unit": {
|
||||
"type": "string",
|
||||
"enum": ["celsius", "fahrenheit"],
|
||||
},
|
||||
},
|
||||
"required": ["location"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
"tool_choice": "required",
|
||||
}
|
||||
}
|
||||
|
||||
mapped_params = litellm.AnthropicConfig().map_openai_params(
|
||||
non_default_params=args["non_default_params"],
|
||||
optional_params={},
|
||||
model="claude-3-5-sonnet-20240620",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "Question" in json.dumps(mapped_params)
|
||||
|
||||
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"json_mode, tool_calls, expect_null_response",
|
||||
[
|
||||
(
|
||||
True,
|
||||
[
|
||||
{
|
||||
"id": "toolu_013JszbnYBVygTxh6EGHEHia",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_weather",
|
||||
"arguments": '{"location": "New York, NY"}',
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
True,
|
||||
),
|
||||
(
|
||||
True,
|
||||
[
|
||||
{
|
||||
"id": "toolu_013JszbnYBVygTxh6EGHEHia",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": RESPONSE_FORMAT_TOOL_NAME,
|
||||
"arguments": '{"location": "New York, NY"}',
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
False,
|
||||
),
|
||||
(
|
||||
False,
|
||||
[
|
||||
{
|
||||
"id": "toolu_013JszbnYBVygTxh6EGHEHia",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": RESPONSE_FORMAT_TOOL_NAME,
|
||||
"arguments": '{"location": "New York, NY"}',
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
True,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_anthropic_json_mode_and_tool_call_response(
|
||||
json_mode, tool_calls, expect_null_response
|
||||
):
|
||||
result = litellm.AnthropicConfig()._transform_response_for_json_mode(
|
||||
json_mode=json_mode,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
|
||||
assert (
|
||||
result is None if expect_null_response else result is not None
|
||||
), f"Expected result to be {None if expect_null_response else 'not None'}, but got {result}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue