mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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
This commit is contained in:
parent
7c392475e6
commit
49e9b73fcb
6 changed files with 101 additions and 54 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:<AWS-ACCOUNT-ID>: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"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue