From a0ee31edf8858e84c7906a53aa064af2902acb92 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 8 May 2025 21:16:47 -0700 Subject: [PATCH] [Feat] Add support for using Bedrock Invoke models in /v1/messages format (#10681) * fix: add transform_anthropic_messages_request * fix: add get_requested_response_api_optional_param * fix: use base llm http handler for anthropic messages * fix: add anthropic transform response * fix: transform_anthropic_messages_response * fix: fixes for anthropic messages * fix: code qa fixes * fix: pass thinking to anthropic * fix: linting * fixes * feat: add folder for bedrock invoke messages * feat: init bedrock invoke messages for anthropic claude family * test: add bedrock invoke test for us anthropic * test: test_anthropic_messages_non_streaming_bedrock_invokec * feat: update anthropic messages transforms * feat: update anthropic messages transforms * Update litellm/utils.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * fix: test_anthropic_messages_non_streaming * fix: linting override * fix: linting error --------- Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- litellm/__init__.py | 3 + .../messages/transformation.py | 16 ++- .../anthropic_messages/transformation.py | 40 +++++- .../anthropic_claude3_transformation.py | 122 ++++++++++++++++++ litellm/llms/bedrock/messages/readme.md | 3 + litellm/llms/custom_httpx/llm_http_handler.py | 23 +++- litellm/utils.py | 4 + .../test_anthropic_messages_passthrough.py | 27 ++++ 8 files changed, 231 insertions(+), 7 deletions(-) create mode 100644 litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py create mode 100644 litellm/llms/bedrock/messages/readme.md diff --git a/litellm/__init__.py b/litellm/__init__.py index fddcf76f4f9..eb712ac870f 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -858,6 +858,9 @@ from .llms.meta_llama.chat.transformation import LlamaAPIConfig from .llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, ) +from .llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import ( + AmazonAnthropicClaude3MessagesConfig, +) from .llms.together_ai.chat import TogetherAIConfig from .llms.together_ai.completion.transformation import TogetherAITextCompletionConfig from .llms.cloudflare.chat.transformation import CloudflareChatConfig diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index 551aadd3cba..c153b4fbcf1 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -1,4 +1,4 @@ -from typing import Dict, List, Optional +from typing import Any, Dict, List, Optional import httpx @@ -36,7 +36,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): # "metadata", ] - def get_complete_url(self, api_base: Optional[str], model: str) -> str: + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: api_base = api_base or DEFAULT_ANTHROPIC_API_BASE if not api_base.endswith("/v1/messages"): api_base = f"{api_base}/v1/messages" @@ -46,7 +54,11 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): self, headers: dict, model: str, + messages: List[Any], + optional_params: dict, + litellm_params: dict, api_key: Optional[str] = None, + api_base: Optional[str] = None, ) -> dict: if "x-api-key" not in headers: headers["x-api-key"] = api_key diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index 2ca5bdc0725..29ac0cf2f28 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -22,12 +22,29 @@ class BaseAnthropicMessagesConfig(ABC): self, headers: dict, model: str, + messages: List[Any], + optional_params: dict, + litellm_params: dict, api_key: Optional[str] = None, + api_base: Optional[str] = None, ) -> dict: - pass + """ + OPTIONAL + + Validate the environment for the request + """ + return headers @abstractmethod - def get_complete_url(self, api_base: Optional[str], model: str) -> str: + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: """ OPTIONAL @@ -60,3 +77,22 @@ class BaseAnthropicMessagesConfig(ABC): logging_obj: LiteLLMLoggingObj, ) -> AnthropicMessagesResponse: pass + + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> dict: + """ + OPTIONAL + + Sign the request, providers like Bedrock need to sign the request before sending it to the API + + For all other providers, this is a no-op and we just return the headers + """ + return headers diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py new file mode 100644 index 00000000000..5eba423bc38 --- /dev/null +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -0,0 +1,122 @@ +from typing import TYPE_CHECKING, Any, Dict, List, Optional + +from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, +) +from litellm.llms.base_llm.anthropic_messages.transformation import ( + BaseAnthropicMessagesConfig, +) +from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( + AmazonInvokeConfig, +) +from litellm.types.router import GenericLiteLLMParams + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class AmazonAnthropicClaude3MessagesConfig( + AnthropicMessagesConfig, + AmazonInvokeConfig, +): + """ + Call Claude model family in the /v1/messages API spec + """ + + DEFAULT_BEDROCK_ANTHROPIC_API_VERSION = "bedrock-2023-05-31" + + def __init__(self, **kwargs): + BaseAnthropicMessagesConfig.__init__(self, **kwargs) + AmazonInvokeConfig.__init__(self, **kwargs) + + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> dict: + return AmazonInvokeConfig.sign_request( + self=self, + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base, + model=model, + stream=stream, + fake_stream=fake_stream, + ) + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[Any], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + return headers + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + return AmazonInvokeConfig.get_complete_url( + self=self, + api_base=api_base, + api_key=api_key, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + stream=stream, + ) + + def transform_anthropic_messages_request( + self, + model: str, + messages: List[Dict], + anthropic_messages_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + anthropic_messages_request = AnthropicMessagesConfig.transform_anthropic_messages_request( + self=self, + model=model, + messages=messages, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + litellm_params=litellm_params, + headers=headers, + ) + + ######################################################### + ############## BEDROCK Invoke SPECIFIC TRANSFORMATION ### + ######################################################### + + # 1. anthropic_version is required for all claude models + if "anthropic_version" not in anthropic_messages_request: + anthropic_messages_request[ + "anthropic_version" + ] = self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION + + # 2. `stream` is not allowed in request body for bedrock invoke + if "stream" in anthropic_messages_request: + anthropic_messages_request.pop("stream", None) + + # 3. `model` is not allowed in request body for bedrock invoke + if "model" in anthropic_messages_request: + anthropic_messages_request.pop("model", None) + return anthropic_messages_request diff --git a/litellm/llms/bedrock/messages/readme.md b/litellm/llms/bedrock/messages/readme.md new file mode 100644 index 00000000000..5d8d386accb --- /dev/null +++ b/litellm/llms/bedrock/messages/readme.md @@ -0,0 +1,3 @@ +# /v1/messages + +This folder contains transformation logic for calling bedrock models in the Anthropic /v1/messages API spec. \ No newline at end of file diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 70ec04a5e13..81739970bd1 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1059,7 +1059,11 @@ class BaseLLMHTTPHandler: headers = anthropic_messages_provider_config.validate_environment( headers=extra_headers or {}, model=model, + messages=messages, + optional_params=anthropic_messages_optional_request_params, + litellm_params=dict(litellm_params), api_key=api_key, + api_base=api_base, ) logging_obj.update_environment_variables( @@ -1081,14 +1085,27 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, headers=headers, ) - request_body["stream"] = stream - request_body["model"] = model logging_obj.stream = stream logging_obj.model_call_details.update(request_body) # Make the request request_url = anthropic_messages_provider_config.get_complete_url( - api_base=api_base, model=model + api_base=api_base, + api_key=api_key, + model=model, + optional_params=anthropic_messages_optional_request_params, + litellm_params=dict(litellm_params), + stream=stream, + ) + + headers = anthropic_messages_provider_config.sign_request( + headers=headers, + optional_params=anthropic_messages_optional_request_params, + request_data=request_body, + api_base=request_url, + stream=stream, + fake_stream=False, + model=model, ) logging_obj.pre_call( diff --git a/litellm/utils.py b/litellm/utils.py index 0bc579b1182..4c2c36ed8e1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6429,6 +6429,10 @@ class ProviderConfigManager: ) -> Optional[BaseAnthropicMessagesConfig]: if litellm.LlmProviders.ANTHROPIC == provider: return litellm.AnthropicMessagesConfig() + # The 'BEDROCK' provider corresponds to Amazon's implementation of Anthropic Claude v3. + # This mapping ensures that the correct configuration is returned for BEDROCK. + elif litellm.LlmProviders.BEDROCK == provider: + return litellm.AmazonAnthropicClaude3MessagesConfig() return None @staticmethod diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index ef288ba41ce..739c24fc424 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -100,6 +100,33 @@ async def test_anthropic_messages_non_streaming(): print(f"Non-streaming response: {json.dumps(response, indent=2)}") return response +@pytest.mark.asyncio +async def test_anthropic_messages_non_streaming_bedrock_invoke(): + """ + Test the anthropic_messages with non-streaming request + """ + litellm._turn_on_debug() + + # Set up test parameters + messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] + + # Call the handler + response = await litellm.anthropic.messages.acreate( + messages=messages, + model="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + max_tokens=100, + ) + + print("non-streaming bedrock invoke response: ", response) + + # Verify response + assert "id" in response + assert "content" in response + assert "model" in response + assert response["role"] == "assistant" + + print(f"Non-streaming response: {json.dumps(response, indent=2)}") + return response @pytest.mark.asyncio async def test_anthropic_messages_streaming():