fix(anthropic/): fix handling response_format for anthropic messages with anthropic api

This commit is contained in:
Krrish Dholakia 2024-12-12 16:54:48 -08:00
parent f2904cbb4e
commit feefd18614
4 changed files with 243 additions and 22 deletions

View file

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

View file

@ -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"

View file

@ -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 {

View file

@ -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}"