diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index ced66528c97..ba0ec63ff23 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1046,10 +1046,10 @@ def convert_to_gemini_tool_call_invoke( if tool_calls is not None: for tool in tool_calls: if "function" in tool: - gemini_function_call: Optional[ - VertexFunctionCall - ] = _gemini_tool_call_invoke_helper( - function_call_params=tool["function"] + gemini_function_call: Optional[VertexFunctionCall] = ( + _gemini_tool_call_invoke_helper( + function_call_params=tool["function"] + ) ) if gemini_function_call is not None: _parts_list.append( @@ -1144,7 +1144,7 @@ def convert_to_gemini_tool_call_result( def convert_to_anthropic_tool_result( - message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage] + message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], ) -> AnthropicMessagesToolResultParam: """ OpenAI message with a tool result looks like: @@ -1465,9 +1465,9 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_element[ - "cache_control" - ] = _content_element["cache_control"] + _anthropic_content_element["cache_control"] = ( + _content_element["cache_control"] + ) user_content.append(_anthropic_content_element) elif m.get("type", "") == "text": m = cast(ChatCompletionTextObject, m) @@ -1518,9 +1518,9 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_text_element[ - "cache_control" - ] = _content_element["cache_control"] + _anthropic_content_text_element["cache_control"] = ( + _content_element["cache_control"] + ) user_content.append(_anthropic_content_text_element) @@ -2260,6 +2260,7 @@ from litellm.types.llms.bedrock import ToolBlock as BedrockToolBlock from litellm.types.llms.bedrock import ( ToolInputSchemaBlock as BedrockToolInputSchemaBlock, ) +from litellm.types.llms.bedrock import ToolJsonSchemaBlock as BedrockToolJsonSchemaBlock from litellm.types.llms.bedrock import ToolResultBlock as BedrockToolResultBlock from litellm.types.llms.bedrock import ( ToolResultContentBlock as BedrockToolResultContentBlock, @@ -2515,7 +2516,7 @@ def _convert_to_bedrock_tool_call_invoke( def _convert_to_bedrock_tool_call_result( - message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage] + message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], ) -> BedrockContentBlock: """ OpenAI message with a tool result looks like: @@ -2688,7 +2689,7 @@ def get_user_message_block_or_continue_message( def return_assistant_continue_message( assistant_continue_message: Optional[ Union[str, ChatCompletionAssistantMessage] - ] = None + ] = None, ) -> ChatCompletionAssistantMessage: if assistant_continue_message and isinstance(assistant_continue_message, str): return ChatCompletionAssistantMessage( @@ -3522,7 +3523,13 @@ def _bedrock_tools_pt(tools: List) -> List[BedrockToolBlock]: for _, value in defs_copy.items(): unpack_defs(value, defs_copy) unpack_defs(parameters, defs_copy) - tool_input_schema = BedrockToolInputSchemaBlock(json=parameters) + tool_input_schema = BedrockToolInputSchemaBlock( + json=BedrockToolJsonSchemaBlock( + type=parameters.get("type", ""), + properties=parameters.get("properties", {}), + required=parameters.get("required", []), + ) + ) tool_spec = BedrockToolSpecBlock( inputSchema=tool_input_schema, name=name, description=description ) diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 6e02ff2ab72..834be071da1 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -125,8 +125,14 @@ class ConverseResponseBlock(TypedDict): usage: ConverseTokenUsageBlock +class ToolJsonSchemaBlock(TypedDict, total=False): + type: Literal["object"] + properties: dict + required: List[str] + + class ToolInputSchemaBlock(TypedDict): - json: Optional[dict] + json: Optional[ToolJsonSchemaBlock] class ToolSpecBlock(TypedDict, total=False): diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index 6dc6fea96d8..99da2c3b5d0 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -1058,6 +1058,7 @@ class BaseLLMChatTest(ABC): def test_function_calling_with_tool_response(self): from litellm.utils import supports_function_calling from litellm import completion + litellm._turn_on_debug() os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") @@ -1077,6 +1078,8 @@ class BaseLLMChatTest(ABC): "name": "get_weather", "description": "Get the weather in a city", "parameters": { + "$id": "https://some/internal/name", + "$schema": "https://json-schema.org/draft-07/schema", "type": "object", "properties": { "city": { diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 2c43a9237c8..6a1335d0500 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -1264,6 +1264,50 @@ def test_bedrock_tools_pt_invalid_names(): assert result[1]["toolSpec"]["name"] == "another_invalid_name" +def test_bedrock_tools_transformation_valid_params(): + from litellm.types.llms.bedrock import ToolJsonSchemaBlock + tools = [ + { + "type": "function", + "function": { + "name": "123-invalid@name", + "description": "Invalid name test", + "parameters": { + "$id": "https://some/internal/name", + "type": "object", + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "test": {"type": "string"}, + }, + "required": ["test"], + }, + }, + } + ] + + result = _bedrock_tools_pt(tools) + + print("bedrock tools after prompt formatting=", result) + # Ensure the keys for properties in the response is a subset of keys in ToolJsonSchemaBlock + toolJsonSchema = result[0]["toolSpec"]["inputSchema"]["json"] + assert toolJsonSchema is not None + print("transformed toolJsonSchema keys=", toolJsonSchema.keys()) + print("allowed ToolJsonSchemaBlock keys=", ToolJsonSchemaBlock.__annotations__.keys()) + assert set(toolJsonSchema.keys()).issubset(set(ToolJsonSchemaBlock.__annotations__.keys())) + + + assert isinstance(result, list) + assert len(result) == 1 + assert "toolSpec" in result[0] + assert result[0]["toolSpec"]["name"] == "a123_invalid_name" + assert result[0]["toolSpec"]["description"] == "Invalid name test" + assert "inputSchema" in result[0]["toolSpec"] + assert "json" in result[0]["toolSpec"]["inputSchema"] + assert result[0]["toolSpec"]["inputSchema"]["json"]["properties"]["test"]["type"] == "string" + assert "test" in result[0]["toolSpec"]["inputSchema"]["json"]["required"] + + + def test_not_found_error(): with pytest.raises(litellm.NotFoundError): completion( @@ -2226,6 +2270,33 @@ class TestBedrockConverseChatNormal(BaseLLMChatTest): """ pass +class TestBedrockConverseNovaTestSuite(BaseLLMChatTest): + def get_base_completion_call_args(self) -> dict: + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm.add_known_models() + return { + "model": "bedrock/us.amazon.nova-lite-v1:0", + "aws_region_name": "us-east-1", + } + + def test_tool_call_no_arguments(self, tool_call_no_arguments): + """Test that tool calls with no arguments is translated correctly. Relevant issue: https://github.com/BerriAI/litellm/issues/6833""" + pass + + def test_multilingual_requests(self): + """ + Bedrock API raises a 400 BadRequest error when the request contains invalid utf-8 sequences. + + Todo: if litellm.modify_params is True ensure it's a valid utf-8 sequence + """ + pass + + def test_prompt_caching(self): + """ + TODO: Ensure this test passes our base llm test suite + """ + class TestBedrockRerank(BaseLLMRerankTest): def get_custom_llm_provider(self) -> litellm.LlmProviders: diff --git a/tests/llm_translation/test_bedrock_invoke_tests.py b/tests/llm_translation/test_bedrock_invoke_tests.py index 93d7c73174e..2795d894f2b 100644 --- a/tests/llm_translation/test_bedrock_invoke_tests.py +++ b/tests/llm_translation/test_bedrock_invoke_tests.py @@ -32,7 +32,7 @@ class TestBedrockInvokeNovaJson(BaseLLMChatTest): def test_tool_call_no_arguments(self, tool_call_no_arguments): """Test that tool calls with no arguments is translated correctly. Relevant issue: https://github.com/BerriAI/litellm/issues/6833""" pass - + @pytest.fixture(autouse=True) def skip_non_json_tests(self, request): if not "json" in request.function.__name__.lower():