diff --git a/litellm/constants.py b/litellm/constants.py index 8f969410252..66500432fa2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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 diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 860ba5eae85..30f87d54561 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -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" diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index 19b382347db..28b8e7c7a05 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -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 { diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index 967b2d27229..87227033270 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -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}"