From b0003bd03cc01a48e2d9d8e48db2d92953a04531 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 10 Jul 2025 15:12:17 -0700 Subject: [PATCH] Feat(bedrock): support api key authentication for AWS Bedrock API (#12426) (#12495) * add support of bearer token for bedrock integration * fix linting issue * fix type checking issue * reoder arguments to address type checking issue * switch to use get_secret_str to fetch env variable Co-authored-by: 0x-fang --- docs/my-website/docs/providers/bedrock.md | 21 ++ litellm/images/main.py | 2 + .../anthropic_messages/transformation.py | 1 + litellm/llms/base_llm/chat/transformation.py | 1 + litellm/llms/bedrock/base_aws_llm.py | 64 ++++-- litellm/llms/bedrock/chat/converse_handler.py | 10 +- .../chat/invoke_agent/transformation.py | 2 + .../base_invoke_transformation.py | 2 + litellm/llms/bedrock/embed/embedding.py | 82 +++---- litellm/llms/bedrock/image/image_handler.py | 36 ++- .../anthropic_claude3_transformation.py | 2 + litellm/llms/custom_httpx/llm_http_handler.py | 2 + litellm/llms/sagemaker/chat/transformation.py | 1 + litellm/main.py | 2 + .../bedrock/embed/test_bedrock_embedding.py | 153 +++++++++++++ .../image/test_bedrock_image_bearer_token.py | 130 +++++++++++ .../llms/bedrock/test_base_aws_llm.py | 214 +++++++++++++++++- 17 files changed, 634 insertions(+), 91 deletions(-) create mode 100644 tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py create mode 100644 tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 8217f429ff..21eb3ee686 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -25,11 +25,32 @@ For **Amazon Nova Models**: Bump to v1.53.5+ ::: +## Authentication + :::info LiteLLM uses boto3 to handle authentication. All these options are supported - https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html#credentials. ::: + +LiteLLM supports API key authentication in addition to traditional boto3 authentication methods. For additional API key details, refer to [docs](https://docs.aws.amazon.com/bedrock/latest/userguide/api-keys.html). + +Option 1: use the AWS_BEARER_TOKEN_BEDROCK environment variable + +```bash +export AWS_BEARER_TOKEN_BEDROCK="your-api-key" +``` + +Option 2: use the api_key parameter to pass in API key for completion, embedding, image_generation API calls. + +```python +response = completion( + model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{ "content": "Hello, how are you?","role": "user"}], + api_key="your-api-key" +) +``` + ## Usage diff --git a/litellm/images/main.py b/litellm/images/main.py index 8270879ba8..8da4ce3d59 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -302,6 +302,8 @@ def image_generation( # noqa: PLR0915 model_response=model_response, aimg_generation=aimg_generation, client=client, + api_base=api_base, + api_key=api_key ) elif custom_llm_provider == "vertex_ai": vertex_ai_project = ( diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index 63f4f23003..fdad1633e8 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -89,6 +89,7 @@ class BaseAnthropicMessagesConfig(ABC): optional_params: dict, request_data: dict, api_base: str, + api_key: Optional[str] = None, model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 0f19de6170..1867abde31 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -288,6 +288,7 @@ class BaseConfig(ABC): optional_params: dict, request_data: dict, api_base: str, + api_key: Optional[str] = None, model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index ce3f66339e..4c11309719 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -12,6 +12,7 @@ from typing import ( Tuple, cast, get_args, + Union, ) import httpx @@ -21,7 +22,7 @@ from litellm._logging import verbose_logger from litellm.caching.caching import DualCache from litellm.constants import BEDROCK_INVOKE_PROVIDERS_LITERAL, BEDROCK_MAX_POLICY_SIZE from litellm.litellm_core_utils.dd_tracing import tracer -from litellm.secret_managers.main import get_secret +from litellm.secret_managers.main import get_secret, get_secret_str if TYPE_CHECKING: from botocore.awsrequest import AWSPreparedRequest @@ -670,25 +671,39 @@ class BaseAWSLLM: aws_region_name: str, extra_headers: Optional[dict], endpoint_url: str, - data: str, + data: Union[str, bytes], headers: dict, + api_key: Optional[str] = None, ) -> AWSPreparedRequest: - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - 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 - ) - sigv4.add_auth(request) - if ( - extra_headers is not None and "Authorization" in extra_headers - ): # prevent sigv4 from overwriting the auth header - request.headers["Authorization"] = extra_headers["Authorization"] + if api_key is not None: + 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'.") + headers["Authorization"] = f"Bearer {aws_bearer_token}" + request = AWSRequest( + method="POST", url=endpoint_url, data=data, headers=headers + ) + else: + try: + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + except ImportError: + 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 + ) + sigv4.add_auth(request) + if ( + extra_headers is not None and "Authorization" in extra_headers + ): # prevent sigv4 from overwriting the auth header + request.headers["Authorization"] = extra_headers["Authorization"] prepped = request.prepare() return prepped @@ -703,6 +718,7 @@ class BaseAWSLLM: model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, + api_key: Optional[str] = None, ) -> Tuple[dict, Optional[bytes]]: """ Sign a request for Bedrock or Sagemaker @@ -710,7 +726,19 @@ class BaseAWSLLM: Returns: Tuple[dict, Optional[str]]: A tuple containing the headers and the json str body of the request """ + if api_key is not None: + 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 {} + headers["Content-Type"] = "application/json" + headers["Authorization"] = f"Bearer {aws_bearer_token}" + return headers, json.dumps(request_data).encode() + # If no bearer token is set, proceed with the existing SigV4 authentication try: from botocore.auth import SigV4Auth from botocore.awsrequest import AWSRequest diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 7f529c637a..900fad3d04 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -112,6 +112,7 @@ class BedrockConverseLLM(BaseAWSLLM): client: Optional[AsyncHTTPHandler] = None, fake_stream: bool = False, json_mode: Optional[bool] = False, + api_key: Optional[str] = None, ) -> CustomStreamWrapper: request_data = await litellm.AmazonConverseConfig()._async_transform_request( model=model, @@ -128,6 +129,7 @@ class BedrockConverseLLM(BaseAWSLLM): endpoint_url=api_base, data=data, headers=headers, + api_key=api_key ) ## LOGGING @@ -176,6 +178,7 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=None, headers: dict = {}, client: Optional[AsyncHTTPHandler] = None, + api_key: Optional[str] = None, ) -> Union[ModelResponse, CustomStreamWrapper]: request_data = await litellm.AmazonConverseConfig()._async_transform_request( model=model, @@ -184,7 +187,6 @@ class BedrockConverseLLM(BaseAWSLLM): litellm_params=litellm_params, ) data = json.dumps(request_data) - prepped = self.get_request_headers( credentials=credentials, aws_region_name=litellm_params.get("aws_region_name") or "us-west-2", @@ -192,6 +194,7 @@ class BedrockConverseLLM(BaseAWSLLM): endpoint_url=api_base, data=data, headers=headers, + api_key=api_key ) ## LOGGING @@ -261,6 +264,7 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=None, extra_headers: Optional[dict] = None, client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, + api_key: Optional[str] = None, ): ## SETUP ## stream = optional_params.pop("stream", None) @@ -353,6 +357,7 @@ class BedrockConverseLLM(BaseAWSLLM): json_mode=json_mode, fake_stream=fake_stream, credentials=credentials, + api_key=api_key ) # type: ignore ### ASYNC COMPLETION return self.async_completion( @@ -370,6 +375,7 @@ class BedrockConverseLLM(BaseAWSLLM): timeout=timeout, client=client, credentials=credentials, + api_key=api_key ) # type: ignore ## TRANSFORMATION ## @@ -381,7 +387,6 @@ class BedrockConverseLLM(BaseAWSLLM): litellm_params=litellm_params, ) data = json.dumps(_data) - prepped = self.get_request_headers( credentials=credentials, aws_region_name=aws_region_name, @@ -389,6 +394,7 @@ class BedrockConverseLLM(BaseAWSLLM): endpoint_url=proxy_endpoint_url, data=data, headers=headers, + api_key=api_key ) ## LOGGING diff --git a/litellm/llms/bedrock/chat/invoke_agent/transformation.py b/litellm/llms/bedrock/chat/invoke_agent/transformation.py index aa57bb7feb..e4ff6d398e 100644 --- a/litellm/llms/bedrock/chat/invoke_agent/transformation.py +++ b/litellm/llms/bedrock/chat/invoke_agent/transformation.py @@ -102,6 +102,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, + api_key: Optional[str] = None, model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, @@ -115,6 +116,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): model=model, stream=stream, fake_stream=fake_stream, + api_key=api_key, ) def _get_agent_id_and_alias_id(self, model: str) -> tuple[str, str]: diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index 4c977af2fd..16f146206b 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -118,6 +118,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, + api_key: Optional[str] = None, model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, @@ -128,6 +129,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): optional_params=optional_params, request_data=request_data, api_base=api_base, + api_key=api_key, model=model, stream=stream, fake_stream=fake_stream, diff --git a/litellm/llms/bedrock/embed/embedding.py b/litellm/llms/bedrock/embed/embedding.py index 9e4e4e22d0..91c71e86f1 100644 --- a/litellm/llms/bedrock/embed/embedding.py +++ b/litellm/llms/bedrock/embed/embedding.py @@ -156,28 +156,23 @@ class BedrockEmbedding(BaseAWSLLM): aws_region_name: str, model: str, logging_obj: Any, + api_key: Optional[str] = None, ): - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - responses: List[dict] = [] for data in batch_data: - sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name) headers = {"Content-Type": "application/json"} if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} - request = AWSRequest( - method="POST", url=endpoint_url, data=json.dumps(data), headers=headers - ) - sigv4.add_auth(request) - if ( - extra_headers is not None and "Authorization" in extra_headers - ): # prevent sigv4 from overwriting the auth header - request.headers["Authorization"] = extra_headers["Authorization"] - prepped = request.prepare() + + prepped = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=endpoint_url, + data=json.dumps(data), + headers=headers, + api_key=api_key + ) ## LOGGING logging_obj.pre_call( @@ -245,28 +240,23 @@ class BedrockEmbedding(BaseAWSLLM): aws_region_name: str, model: str, logging_obj: Any, + api_key: Optional[str] = None, ): - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - responses: List[dict] = [] for data in batch_data: - sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name) headers = {"Content-Type": "application/json"} if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} - request = AWSRequest( - method="POST", url=endpoint_url, data=json.dumps(data), headers=headers - ) - sigv4.add_auth(request) - if ( - extra_headers is not None and "Authorization" in extra_headers - ): # prevent sigv4 from overwriting the auth header - request.headers["Authorization"] = extra_headers["Authorization"] - prepped = request.prepare() + + prepped = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=endpoint_url, + data=json.dumps(data), + headers=headers, + api_key=api_key, + ) ## LOGGING logging_obj.pre_call( @@ -338,13 +328,8 @@ class BedrockEmbedding(BaseAWSLLM): extra_headers: Optional[dict], optional_params: dict, litellm_params: dict, + api_key: Optional[str] = None, ) -> EmbeddingResponse: - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - credentials, aws_region_name = self._load_credentials(optional_params) ### TRANSFORMATION ### @@ -428,6 +413,7 @@ class BedrockEmbedding(BaseAWSLLM): aws_region_name=aws_region_name, model=model, logging_obj=logging_obj, + api_key=api_key, ) return self._single_func_embeddings( client=( @@ -443,24 +429,24 @@ class BedrockEmbedding(BaseAWSLLM): aws_region_name=aws_region_name, model=model, logging_obj=logging_obj, + api_key=api_key, ) elif data is None: raise Exception("Unable to map Bedrock request to provider") - sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name) headers = {"Content-Type": "application/json"} if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} - - request = AWSRequest( - method="POST", url=endpoint_url, data=json.dumps(data), headers=headers + + prepped = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=endpoint_url, + data=json.dumps(data), + headers=headers, + api_key=api_key, ) - sigv4.add_auth(request) - if ( - extra_headers is not None and "Authorization" in extra_headers - ): # prevent sigv4 from overwriting the auth header - request.headers["Authorization"] = extra_headers["Authorization"] - prepped = request.prepare() ## ROUTING ## return cohere_embedding( diff --git a/litellm/llms/bedrock/image/image_handler.py b/litellm/llms/bedrock/image/image_handler.py index 27258aa20f..55d94675d1 100644 --- a/litellm/llms/bedrock/image/image_handler.py +++ b/litellm/llms/bedrock/image/image_handler.py @@ -54,6 +54,7 @@ class BedrockImageGeneration(BaseAWSLLM): api_base: Optional[str] = None, extra_headers: Optional[dict] = None, client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_key: Optional[str] = None, ): prepared_request = self._prepare_request( model=model, @@ -62,6 +63,7 @@ class BedrockImageGeneration(BaseAWSLLM): extra_headers=extra_headers, logging_obj=logging_obj, prompt=prompt, + api_key=api_key ) if aimg_generation is True: @@ -148,6 +150,7 @@ class BedrockImageGeneration(BaseAWSLLM): extra_headers: Optional[dict], logging_obj: LitellmLogging, prompt: str, + api_key: Optional[str], ) -> BedrockImagePreparedRequest: """ Prepare the request body, headers, and endpoint URL for the Bedrock Image Generation API @@ -167,11 +170,6 @@ class BedrockImageGeneration(BaseAWSLLM): prepped (httpx.Request): The prepared request object body (bytes): The request body """ - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") boto3_credentials_info = self._get_boto_credentials_from_optional_params( optional_params, model ) @@ -184,32 +182,26 @@ class BedrockImageGeneration(BaseAWSLLM): aws_region_name=boto3_credentials_info.aws_region_name, ) proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke" - sigv4 = SigV4Auth( - boto3_credentials_info.credentials, - "bedrock", - boto3_credentials_info.aws_region_name, - ) - data = self._get_request_body( model=model, prompt=prompt, optional_params=optional_params ) # Make POST Request body = json.dumps(data).encode("utf-8") - headers = {"Content-Type": "application/json"} if extra_headers is not None: - headers = {"Content-Type": "application/json", **extra_headers} - request = AWSRequest( - method="POST", url=proxy_endpoint_url, data=body, headers=headers - ) - sigv4.add_auth(request) - if ( - extra_headers is not None and "Authorization" in extra_headers - ): # prevent sigv4 from overwriting the auth header - request.headers["Authorization"] = extra_headers["Authorization"] - prepped = request.prepare() + headers = {"Content-Type": "application/json", **extra_headers} + prepped = self.get_request_headers( + credentials=boto3_credentials_info.credentials, + aws_region_name=boto3_credentials_info.aws_region_name, + extra_headers=extra_headers, + endpoint_url=proxy_endpoint_url, + data=body, + headers=headers, + api_key=api_key, + ) + ## LOGGING logging_obj.pre_call( input=prompt, diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index ba9b478b29..09c6673cc5 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -57,6 +57,7 @@ class AmazonAnthropicClaude3MessagesConfig( optional_params: dict, request_data: dict, api_base: str, + api_key: Optional[str] = None, model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, @@ -67,6 +68,7 @@ class AmazonAnthropicClaude3MessagesConfig( optional_params=optional_params, request_data=request_data, api_base=api_base, + api_key=api_key, model=model, stream=stream, fake_stream=fake_stream, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index e08b909b2a..1d78be0e97 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -339,6 +339,7 @@ class BaseLLMHTTPHandler: optional_params=optional_params, request_data=data, api_base=api_base, + api_key=api_key, stream=stream, fake_stream=fake_stream, model=model, @@ -1324,6 +1325,7 @@ class BaseLLMHTTPHandler: ), # dynamic aws_* params are passed under litellm_params request_data=request_body, api_base=request_url, + api_key=api_key, stream=stream, fake_stream=False, model=model, diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index 14dde144af..2b458fbc43 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -93,6 +93,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): optional_params: dict, request_data: dict, api_base: str, + api_key: Optional[str] = None, model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, diff --git a/litellm/main.py b/litellm/main.py index 98dc298e36..27374c9550 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2894,6 +2894,7 @@ def completion( # type: ignore # noqa: PLR0915 acompletion=acompletion, client=client, api_base=api_base, + api_key=api_key ) elif bedrock_route == "converse_like": model = model.replace("converse_like/", "") @@ -3918,6 +3919,7 @@ def embedding( # noqa: PLR0915 api_base=api_base, print_verbose=print_verbose, extra_headers=extra_headers, + api_key=api_key, ) elif custom_llm_provider == "triton": if api_base is None: diff --git a/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py b/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py new file mode 100644 index 0000000000..aec0b5fc6c --- /dev/null +++ b/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py @@ -0,0 +1,153 @@ +import json +import os +import sys +from unittest.mock import Mock, patch +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path +import litellm +from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler + +# Mock responses for different embedding models +titan_embedding_response = { + "embedding": [0.1, 0.2, 0.3], + "inputTextTokenCount": 10 +} + +cohere_embedding_response = { + "embeddings": [[0.1, 0.2, 0.3]], + "inputTextTokenCount": 10 +} + +# Test data +test_input = "Hello world from litellm" +test_image_base64 = "data:image/png,test_image_base64_data" + + +@pytest.mark.parametrize( + "model,input_type,embed_response", + [ + ("bedrock/amazon.titan-embed-text-v1", "text", titan_embedding_response), + ("bedrock/amazon.titan-embed-text-v2:0", "text", titan_embedding_response), + ("bedrock/amazon.titan-embed-image-v1", "image", titan_embedding_response), + ("bedrock/cohere.embed-english-v3", "text", cohere_embedding_response), + ("bedrock/cohere.embed-multilingual-v3", "text", cohere_embedding_response), + ], +) +def test_bedrock_embedding_with_api_key_bearer_token(model, input_type, embed_response): + """Test embedding functionality with bearer token authentication""" + litellm.set_verbose = True + client = HTTPHandler() + test_api_key = "test-bearer-token-12345" + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(embed_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + input_data = test_image_base64 if input_type == "image" else test_input + + response = litellm.embedding( + model=model, + input=input_data, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key=test_api_key + ) + + assert isinstance(response, litellm.EmbeddingResponse) + assert isinstance(response.data[0]['embedding'], list) + assert len(response.data[0]['embedding']) == 3 # Based on mock response + + headers = mock_post.call_args.kwargs.get("headers", {}) + assert "Authorization" in headers + assert headers["Authorization"] == f"Bearer {test_api_key}" + + +@pytest.mark.parametrize( + "model,input_type,embed_response", + [ + ("bedrock/amazon.titan-embed-text-v1", "text", titan_embedding_response), + ], +) +def test_bedrock_embedding_with_env_variable_bearer_token(model, input_type, embed_response): + """Test embedding functionality with bearer token from environment variable""" + litellm.set_verbose = True + client = HTTPHandler() + test_api_key = "env-bearer-token-12345" + + with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": test_api_key}), \ + patch.object(client, "post") as mock_post: + + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(embed_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + response = litellm.embedding( + model=model, + input=test_input, + client=client, + aws_region_name="us-west-2", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-west-2.amazonaws.com", + ) + + assert isinstance(response, litellm.EmbeddingResponse) + headers = mock_post.call_args.kwargs.get("headers", {}) + assert "Authorization" in headers + assert headers["Authorization"] == f"Bearer {test_api_key}" + + +@pytest.mark.asyncio +async def test_async_bedrock_embedding_with_bearer_token(): + """Test async embedding functionality with bearer token authentication""" + litellm.set_verbose = True + client = AsyncHTTPHandler() + test_api_key = "async-bearer-token-12345" + model = "bedrock/amazon.titan-embed-text-v1" + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(titan_embedding_response) + mock_response.json = Mock(return_value=titan_embedding_response) + mock_post.return_value = mock_response + + response = await litellm.aembedding( + model=model, + input=test_input, + client=client, + aws_region_name="us-west-2", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-west-2.amazonaws.com", + api_key=test_api_key + ) + + assert isinstance(response, litellm.EmbeddingResponse) + + headers = mock_post.call_args.kwargs.get("headers", {}) + assert "Authorization" in headers + assert headers["Authorization"] == f"Bearer {test_api_key}" + + +def test_bedrock_embedding_with_sigv4(): + """Test embedding falls back to SigV4 auth when no bearer token is provided""" + litellm.set_verbose = True + model = "bedrock/amazon.titan-embed-text-v1" + + with patch("litellm.llms.bedrock.embed.embedding.BedrockEmbedding.embeddings") as mock_bedrock_embed: + mock_embedding_response = litellm.EmbeddingResponse() + mock_embedding_response.data = [{"embedding": [0.1, 0.2, 0.3]}] + mock_bedrock_embed.return_value = mock_embedding_response + + response = litellm.embedding( + model=model, + input=test_input, + aws_region_name="us-west-2", + ) + + assert isinstance(response, litellm.EmbeddingResponse) + mock_bedrock_embed.assert_called_once() \ No newline at end of file diff --git a/tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py b/tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py new file mode 100644 index 0000000000..b348c1193c --- /dev/null +++ b/tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py @@ -0,0 +1,130 @@ +import json +import os +import sys +from unittest.mock import Mock, patch +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path + +import litellm +from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler + +# Mock response for Bedrock image generation +mock_image_response = { + "images": ["base64_encoded_image_data"], + "error": None +} + +class TestBedrockImageGeneration: + def test_image_generation_with_api_key_bearer_token(self): + """Test image generation with bearer token authentication""" + litellm.set_verbose = True + test_api_key = "test-bearer-token-12345" + model = "bedrock/stability.sd3-large-v1:0" + prompt = "A cute baby sea otter" + + with patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: + # Setup mock response + mock_image_response_obj = litellm.ImageResponse() + mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] + mock_bedrock_image_gen.return_value = mock_image_response_obj + + response = litellm.image_generation( + model=model, + prompt=prompt, + aws_region_name="us-west-2", + api_key=test_api_key + ) + + assert response is not None + assert len(response.data) > 0 + + mock_bedrock_image_gen.assert_called_once() + for call in mock_bedrock_image_gen.call_args_list: + if "headers" in call.kwargs: + headers = call.kwargs["headers"] + if "Authorization" in headers and headers["Authorization"] == f"Bearer {test_api_key}": + break + + def test_image_generation_with_env_variable_bearer_token(self, monkeypatch): + """Test image generation with bearer token from environment variable""" + litellm.set_verbose = True + test_api_key = "env-bearer-token-12345" + model = "bedrock/stability.sd3-large-v1:0" + prompt = "A cute baby sea otter" + + # Mock the environment variable + with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": test_api_key}), \ + patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: + + mock_image_response_obj = litellm.ImageResponse() + mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] + mock_bedrock_image_gen.return_value = mock_image_response_obj + + response = litellm.image_generation( + model=model, + prompt=prompt, + aws_region_name="us-west-2" + ) + + assert response is not None + assert len(response.data) > 0 + + mock_bedrock_image_gen.assert_called_once() + for call in mock_bedrock_image_gen.call_args_list: + if "headers" in call.kwargs: + headers = call.kwargs["headers"] + if "Authorization" in headers and headers["Authorization"] == f"Bearer {test_api_key}": + break + + @pytest.mark.asyncio + async def test_async_image_generation_with_bearer_token(self): + """Test async image generation with bearer token authentication""" + litellm.set_verbose = True + test_api_key = "async-bearer-token-12345" + model = "bedrock/stability.sd3-large-v1:0" + prompt = "A cute baby sea otter" + + with patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.async_image_generation") as mock_async_bedrock_image_gen: + mock_image_response_obj = litellm.ImageResponse() + mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] + mock_async_bedrock_image_gen.return_value = mock_image_response_obj + + # Call async image generation with api_key parameter + response = await litellm.aimage_generation( + model=model, + prompt=prompt, + aws_region_name="us-west-2", + api_key=test_api_key + ) + + assert response is not None + assert len(response.data) > 0 + + mock_async_bedrock_image_gen.assert_called_once() + for call in mock_async_bedrock_image_gen.call_args_list: + if "headers" in call.kwargs: + headers = call.kwargs["headers"] + if "Authorization" in headers and headers["Authorization"] == f"Bearer {test_api_key}": + break + + def test_image_generation_with_sigv4(self): + """Test image generation falls back to SigV4 auth when no bearer token is provided""" + litellm.set_verbose = True + model = "bedrock/stability.sd3-large-v1:0" + prompt = "A cute baby sea otter" + + with patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: + mock_image_response_obj = litellm.ImageResponse() + mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] + mock_bedrock_image_gen.return_value = mock_image_response_obj + + response = litellm.image_generation( + model=model, + prompt=prompt, + aws_region_name="us-west-2" + ) + + assert response is not None + assert len(response.data) > 0 + mock_bedrock_image_gen.assert_called_once() \ No newline at end of file diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py index 01282ad8c1..20abe8ae61 100644 --- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py +++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py @@ -15,7 +15,7 @@ from typing import Any, Dict from unittest.mock import MagicMock, patch from botocore.credentials import Credentials - +from botocore.awsrequest import AWSRequest, AWSPreparedRequest import litellm from litellm.llms.bedrock.base_aws_llm import ( AwsAuthError, @@ -176,3 +176,215 @@ def test_get_aws_region_name_boto3_fallback(): assert result == "ap-southeast-1" mock_boto3_session.assert_not_called() + +def test_sign_request_with_env_var_bearer_token(): + # Create instance of actual class + llm = BaseAWSLLM() + + # Test data + service_name = "bedrock" + headers = {"Custom-Header": "test"} + optional_params = {} + request_data = {"prompt": "test"} + api_base = "https://api.example.com" + + # Mock environment variable + with patch.dict(os.environ, {'AWS_BEARER_TOKEN_BEDROCK': 'test_token'}): + # Execute + result_headers, result_body = llm._sign_request( + service_name=service_name, + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base + ) + + # Assert + assert result_headers["Authorization"] == "Bearer test_token" + assert result_headers["Content-Type"] == "application/json" + assert result_headers["Custom-Header"] == "test" + assert result_body == json.dumps(request_data).encode() + + +def test_sign_request_with_sigv4(): + llm = BaseAWSLLM() + + # Mock AWS credentials and SigV4 auth + mock_credentials = Credentials("test_key", "test_secret", "test_token") + mock_sigv4 = MagicMock() + mock_request = MagicMock() + mock_request.headers = { + "Authorization": "AWS4-HMAC-SHA256 Credential=test", + "Content-Type": "application/json" + } + mock_request.body = b'{"prompt": "test"}' + + # Test data + service_name = "bedrock" + headers = {"Custom-Header": "test"} + optional_params = { + "aws_access_key_id": "test_key", + "aws_secret_access_key": "test_secret", + "aws_region_name": "us-west-2" + } + request_data = {"prompt": "test"} + api_base = "https://api.example.com" + + # Mock the necessary components + with patch('botocore.auth.SigV4Auth', return_value=mock_sigv4), \ + patch('botocore.awsrequest.AWSRequest', return_value=mock_request), \ + patch.object(llm, 'get_credentials', return_value=mock_credentials), \ + patch.object(llm, '_get_aws_region_name', return_value="us-west-2"): + + result_headers, result_body = llm._sign_request( + service_name=service_name, + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base + ) + + # Assert + assert "Authorization" in result_headers + assert result_headers["Authorization"] != "Bearer test_token" + assert result_headers["Content-Type"] == "application/json" + assert result_body == mock_request.body + + +def test_sign_request_with_api_key_bearer_token(): + """ + Test that _sign_request uses the api_key parameter as a bearer token when provided + """ + llm = BaseAWSLLM() + + # Test data + service_name = "bedrock" + headers = {"Custom-Header": "test"} + optional_params = {} + request_data = {"prompt": "test"} + api_base = "https://api.example.com" + api_key = "test_api_key" + + # Execute with api_key parameter + result_headers, result_body = llm._sign_request( + service_name=service_name, + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base, + api_key=api_key + ) + + # Assert + assert result_headers["Authorization"] == f"Bearer {api_key}" + assert result_headers["Content-Type"] == "application/json" + assert result_headers["Custom-Header"] == "test" + assert result_body == json.dumps(request_data).encode() + + +def test_get_request_headers_with_env_var_bearer_token(): + # Setup + llm = BaseAWSLLM() + credentials = Credentials("test_key", "test_secret", "test_token") + headers = {"Content-Type": "application/json"} + headers_dict = headers.copy() + + # Create mock request + mock_prepared_request = MagicMock(spec=AWSPreparedRequest) + mock_request = MagicMock(spec=AWSRequest) + mock_request.headers = headers_dict + mock_request.prepare.return_value = mock_prepared_request + + def mock_aws_request_init(method, url, data, headers): + mock_request.headers.update(headers) + return mock_request + + # Test with bearer token + with patch.dict(os.environ, {'AWS_BEARER_TOKEN_BEDROCK': 'test_token'}), \ + patch('botocore.awsrequest.AWSRequest', side_effect=mock_aws_request_init): + + result = llm.get_request_headers( + credentials=credentials, + aws_region_name="us-west-2", + extra_headers=None, + endpoint_url="https://api.example.com", + data='{"prompt": "test"}', + headers=headers_dict + ) + + # Assert + assert mock_request.headers["Authorization"] == "Bearer test_token" + assert result == mock_prepared_request + + +def test_get_request_headers_with_sigv4(): + # Setup + llm = BaseAWSLLM() + credentials = Credentials("test_key", "test_secret", "test_token") + headers = {"Content-Type": "application/json"} + + # Create mock request and SigV4 instance + mock_request = MagicMock(spec=AWSRequest) + mock_request.headers = headers.copy() + mock_request.prepare.return_value = MagicMock(spec=AWSPreparedRequest) + + mock_sigv4 = MagicMock() + + # Test without bearer token (should use SigV4) + with patch.dict(os.environ, {}, clear=True), \ + patch('botocore.auth.SigV4Auth', return_value=mock_sigv4) as mock_sigv4_class, \ + patch('botocore.awsrequest.AWSRequest', return_value=mock_request): + + result = llm.get_request_headers( + credentials=credentials, + aws_region_name="us-west-2", + extra_headers=None, + endpoint_url="https://api.example.com", + data='{"prompt": "test"}', + headers=headers + ) + + # Verify SigV4 authentication and result + mock_sigv4_class.assert_called_once_with(credentials, "bedrock", "us-west-2") + mock_sigv4.add_auth.assert_called_once_with(mock_request) + assert result == mock_request.prepare.return_value + + +def test_get_request_headers_with_api_key_bearer_token(): + """ + Test that get_request_headers uses the api_key parameter as a bearer token when provided + """ + # Setup + llm = BaseAWSLLM() + credentials = Credentials("test_key", "test_secret", "test_token") + headers = {"Content-Type": "application/json"} + headers_dict = headers.copy() + api_key = "test_api_key" + + # Create mock request + mock_prepared_request = MagicMock(spec=AWSPreparedRequest) + mock_request = MagicMock(spec=AWSRequest) + mock_request.headers = headers_dict + mock_request.prepare.return_value = mock_prepared_request + + def mock_aws_request_init(method, url, data, headers): + mock_request.headers.update(headers) + return mock_request + + # Test with api_key parameter + with patch.dict(os.environ, {}, clear=True), \ + patch('botocore.awsrequest.AWSRequest', side_effect=mock_aws_request_init): + + result = llm.get_request_headers( + credentials=credentials, + aws_region_name="us-west-2", + extra_headers=None, + endpoint_url="https://api.example.com", + data='{"prompt": "test"}', + headers=headers_dict, + api_key=api_key + ) + + # Assert + assert mock_request.headers["Authorization"] == f"Bearer {api_key}" + assert result == mock_prepared_request