From 49e9b73fcb272d376d457c5837319fe2b69f87e1 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Mon, 14 Jul 2025 21:42:25 -0700 Subject: [PATCH] Claude 4 Bedrock /invoke route support + Bedrock application inference profile tool choice support (#12599) * docs(config_settings.md): document enable_json_schema_validation Closes https://github.com/BerriAI/litellm/issues/12518 * fix(utils.py): add claude-sonnet-4 on bedrock support Fixes https://github.com/BerriAI/litellm/issues/12366 * refactor(utils.py): move list to getter in function more maintainable * fix(utils.py): handle bedrock_converse in provider check Fixes https://github.com/BerriAI/litellm/issues/11751 --- docs/my-website/docs/proxy/config_settings.md | 1 + litellm/llms/bedrock/base_aws_llm.py | 14 +++-- .../bedrock/chat/converse_transformation.py | 45 ++++++++-------- .../anthropic_claude2_transformation.py | 8 +++ litellm/utils.py | 51 ++++++++++--------- tests/test_litellm/test_utils.py | 36 ++++++++++++- 6 files changed, 101 insertions(+), 54 deletions(-) diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index de507c87b28..e62dc29be94 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -141,6 +141,7 @@ general_settings: | key_generation_settings | object | Restricts who can generate keys. [Further docs](./virtual_keys.md#restricting-key-generation) | | disable_add_transform_inline_image_block | boolean | For Fireworks AI models - if true, turns off the auto-add of `#transform=inline` to the url of the image_url, if the model is not a vision model. | | disable_hf_tokenizer_download | boolean | If true, it defaults to using the openai tokenizer for all models (including huggingface models). | +| enable_json_schema_validation | boolean | If true, enables json schema validation for all requests. | ### general_settings - Reference diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 4c113097190..df6e0d19f7f 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -10,9 +10,9 @@ from typing import ( Literal, Optional, Tuple, + Union, cast, get_args, - Union, ) import httpx @@ -679,12 +679,14 @@ class BaseAWSLLM: aws_bearer_token: Optional[str] = api_key else: aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK") - + if aws_bearer_token: try: from botocore.awsrequest import AWSRequest except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + raise ImportError( + "Missing boto3 to call bedrock. Run 'pip install boto3'." + ) headers["Authorization"] = f"Bearer {aws_bearer_token}" request = AWSRequest( method="POST", url=endpoint_url, data=data, headers=headers @@ -694,7 +696,9 @@ class BaseAWSLLM: from botocore.auth import SigV4Auth from botocore.awsrequest import AWSRequest except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + raise ImportError( + "Missing boto3 to call bedrock. Run 'pip install boto3'." + ) sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name) request = AWSRequest( method="POST", url=endpoint_url, data=data, headers=headers @@ -730,7 +734,7 @@ class BaseAWSLLM: aws_bearer_token: Optional[str] = api_key else: aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK") - + # If aws bearer token is set, use it directly in the header if aws_bearer_token: headers = headers or {} diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index bd6a29172db..ec378ddbb85 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -146,14 +146,10 @@ class AmazonConverseConfig(BaseConfig): ): supported_params.append("tools") - if ( - litellm.utils.supports_tool_choice( - model=model, custom_llm_provider=self.custom_llm_provider - ) - or litellm.utils.supports_tool_choice( - model=base_model, - custom_llm_provider=self.custom_llm_provider - ) + if litellm.utils.supports_tool_choice( + model=model, custom_llm_provider=self.custom_llm_provider + ) or litellm.utils.supports_tool_choice( + model=base_model, custom_llm_provider=self.custom_llm_provider ): # only anthropic and mistral support tool choice config. otherwise (E.g. cohere) will fail the call - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html supported_params.append("tool_choice") @@ -168,8 +164,7 @@ class AmazonConverseConfig(BaseConfig): custom_llm_provider=self.custom_llm_provider, ) or supports_reasoning( - model=base_model, - custom_llm_provider=self.custom_llm_provider + model=base_model, custom_llm_provider=self.custom_llm_provider ) ): supported_params.append("thinking") @@ -324,12 +319,14 @@ class AmazonConverseConfig(BaseConfig): optional_params = self._add_tools_to_optional_params( optional_params=optional_params, tools=[_tool] ) + if ( litellm.utils.supports_tool_choice( model=model, custom_llm_provider=self.custom_llm_provider ) and not is_thinking_enabled ): + optional_params["tool_choice"] = ToolChoiceValuesBlock( tool=SpecificToolChoiceBlock( name=schema_name if schema_name != "" else "json_tool_call" @@ -812,9 +809,7 @@ class AmazonConverseConfig(BaseConfig): return message, returned_finish_reason - def _translate_message_content( - self, content_blocks: List[ContentBlock] - ) -> Tuple[ + def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[ str, List[ChatCompletionToolCallChunk], Optional[List[BedrockConverseReasoningContentBlock]], @@ -829,9 +824,9 @@ class AmazonConverseConfig(BaseConfig): """ content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[ - List[BedrockConverseReasoningContentBlock] - ] = None + reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( + None + ) for idx, content in enumerate(content_blocks): """ - Content is either a tool response or text @@ -952,9 +947,9 @@ class AmazonConverseConfig(BaseConfig): chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[ - List[BedrockConverseReasoningContentBlock] - ] = None + reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( + None + ) if message is not None: ( @@ -967,12 +962,12 @@ class AmazonConverseConfig(BaseConfig): chat_completion_message["provider_specific_fields"] = { "reasoningContentBlocks": reasoningContentBlocks, } - chat_completion_message[ - "reasoning_content" - ] = self._transform_reasoning_content(reasoningContentBlocks) - chat_completion_message[ - "thinking_blocks" - ] = self._transform_thinking_blocks(reasoningContentBlocks) + chat_completion_message["reasoning_content"] = ( + self._transform_reasoning_content(reasoningContentBlocks) + ) + chat_completion_message["thinking_blocks"] = ( + self._transform_thinking_blocks(reasoningContentBlocks) + ) chat_completion_message["content"] = content_str if json_mode is True and tools is not None and len(tools) == 1: # to support 'json_schema' logic on bedrock models diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py index d0d06ef2b2c..9cc6195cfbb 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude2_transformation.py @@ -59,6 +59,14 @@ class AmazonAnthropicConfig(AmazonInvokeConfig): and v is not None } + @staticmethod + def get_legacy_anthropic_model_names(): + return [ + "anthropic.claude-v2", + "anthropic.claude-instant-v1", + "anthropic.claude-v2:1", + ] + def get_supported_openai_params(self, model: str): return [ "max_tokens", diff --git a/litellm/utils.py b/litellm/utils.py index 47653e1ad0b..a92f75e079f 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -539,9 +539,9 @@ def function_setup( # noqa: PLR0915 function_id: Optional[str] = kwargs["id"] if "id" in kwargs else None ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Optional[ - List[Union[str, Callable, CustomLogger]] - ] = kwargs.pop("callbacks", None) + dynamic_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = ( + kwargs.pop("callbacks", None) + ) all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) if len(all_callbacks) > 0: @@ -1261,9 +1261,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs[ - "retry_policy" - ] = reset_retry_policy() # prevent infinite loops + kwargs["retry_policy"] = ( + reset_retry_policy() + ) # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -2996,10 +2996,10 @@ def pre_process_non_default_params( if "response_format" in non_default_params: if provider_config is not None: - non_default_params[ - "response_format" - ] = provider_config.get_json_schema_from_pydantic_object( - response_format=non_default_params["response_format"] + non_default_params["response_format"] = ( + provider_config.get_json_schema_from_pydantic_object( + response_format=non_default_params["response_format"] + ) ) else: non_default_params["response_format"] = type_to_response_format_param( @@ -3126,16 +3126,16 @@ def pre_process_optional_params( True # so that main.py adds the function call to the prompt ) if "tools" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("tools") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("tools") + ) non_default_params.pop( "tool_choice", None ) # causes ollama requests to hang elif "functions" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("functions") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("functions") + ) elif ( litellm.add_function_to_prompt ): # if user opts to add it to prompt instead @@ -4218,9 +4218,9 @@ def _count_characters(text: str) -> int: def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) -> str: - _choices: Union[ - List[Union[Choices, StreamingChoices]], List[StreamingChoices] - ] = response_obj.choices + _choices: Union[List[Union[Choices, StreamingChoices]], List[StreamingChoices]] = ( + response_obj.choices + ) response_str = "" for choice in _choices: @@ -4385,7 +4385,7 @@ def _strip_openai_finetune_model_name(model_name: str) -> str: def _strip_model_name(model: str, custom_llm_provider: Optional[str]) -> str: - if custom_llm_provider and custom_llm_provider == "bedrock": + if custom_llm_provider and custom_llm_provider in ["bedrock", "bedrock_converse"]: stripped_bedrock_model = _get_base_bedrock_model(model_name=model) return stripped_bedrock_model elif custom_llm_provider and ( @@ -6660,6 +6660,7 @@ class ProviderConfigManager: """ Returns the provider config for a given provider. """ + if ( provider == LlmProviders.OPENAI and litellm.openaiOSeriesConfig.is_model_o_series_model(model=model) @@ -6824,6 +6825,7 @@ class ProviderConfigManager: bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider( model=model ) + base_model = BedrockModelInfo.get_base_model(model) if bedrock_route == "converse" or bedrock_route == "converse_like": @@ -6837,10 +6839,13 @@ class ProviderConfigManager: elif bedrock_invoke_provider == "amazon": # amazon titan llms return litellm.AmazonTitanConfig() elif bedrock_invoke_provider == "anthropic": - if base_model.startswith("anthropic.claude-3"): - return litellm.AmazonAnthropicClaude3Config() - else: + if ( + base_model + in litellm.AmazonAnthropicConfig.get_legacy_anthropic_model_names() + ): return litellm.AmazonAnthropicConfig() + else: + return litellm.AmazonAnthropicClaude3Config() elif ( bedrock_invoke_provider == "meta" or bedrock_invoke_provider == "llama" ): # amazon / meta llms diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index feed6300b15..fa80f284937 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -11,7 +11,12 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm -from litellm.types.utils import Delta, LlmProviders, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + Delta, + LlmProviders, + ModelResponseStream, + StreamingChoices, +) from litellm.utils import ( ProviderConfigManager, TextCompletionStreamWrapper, @@ -2110,6 +2115,35 @@ def test_reasoning_content_preserved_in_text_completion_wrapper(): assert choice["reasoning_content"] == "Here's my chain of thought..." +def test_anthropic_claude_4_invoke_chat_provider_config(): + """Test that the Anthropic Claude 4 Invoke chat provider config is correct.""" + from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( + AmazonAnthropicClaude3Config, + ) + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_chat_config( + model="invoke/us.anthropic.claude-sonnet-4-20250514-v1:0", + provider=LlmProviders.BEDROCK, + ) + print(config) + assert isinstance(config, AmazonAnthropicClaude3Config) + + +def test_bedrock_application_inference_profile(): + model = "arn:aws:bedrock:us-east-2::inference-profile/us.anthropic.claude-3-5-haiku-20241022-v1:0" + from pydantic import BaseModel + + from litellm import completion + from litellm.utils import supports_tool_choice + + result = supports_tool_choice(model, custom_llm_provider="bedrock") + result_2 = supports_tool_choice(model, custom_llm_provider="bedrock_converse") + print(result) + assert result == result_2 + assert result is True + + if __name__ == "__main__": # Allow running this test file directly for debugging pytest.main([__file__, "-v"])