From 91040b6b56c05f4cf29166fce38ada7f9510b09d Mon Sep 17 00:00:00 2001 From: Fang Gong Date: Wed, 20 Aug 2025 18:46:00 -0700 Subject: [PATCH] fix the api key support for bedrock guardrail in proxy --- .../guardrail_hooks/bedrock_guardrails.py | 58 ++++-- .../guardrails/test_guardrail_endpoints.py | 178 ++++++++++++++++++ 2 files changed, 217 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index d02671d69d8..6222233d502 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -15,7 +15,7 @@ sys.path.insert( import json import sys from typing import Any, AsyncGenerator, List, Literal, Optional, Tuple, Union - +from litellm.secret_managers.main import get_secret_str import httpx from fastapi import HTTPException @@ -249,31 +249,46 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): data: dict, optional_params: dict, aws_region_name: str, + api_key: Optional[str] = None, extra_headers: Optional[dict] = 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'.") - - sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name) - api_base = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com/guardrail/{self.guardrailIdentifier}/version/{self.guardrailVersion}/apply" - - encoded_data = json.dumps(data).encode("utf-8") headers = {"Content-Type": "application/json"} if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} + api_base = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com/guardrail/{self.guardrailIdentifier}/version/{self.guardrailVersion}/apply" + encoded_data = json.dumps(data).encode("utf-8") + + # first check api-key, if none, fall back to sigV4 + if api_key is not None: + aws_bearer_token: Optional[str] = api_key + else: + aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK") - request = AWSRequest( - method="POST", url=api_base, data=encoded_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 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=api_base, data=encoded_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=api_base, data=encoded_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 = request.prepare() return prepped_request @@ -298,15 +313,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_guardrail_response: BedrockGuardrailResponse = ( BedrockGuardrailResponse() ) + api_key: Optional[str] = None if request_data: bedrock_request_data.update( self.get_guardrail_dynamic_request_body_params(request_data=request_data) ) + if request_data.get("api_key") is not None: + api_key = request_data["api_key"] + prepared_request = self._prepare_request( credentials=credentials, data=bedrock_request_data, optional_params=self.optional_params, aws_region_name=aws_region_name, + api_key=api_key, ) verbose_proxy_logger.debug( "Bedrock AI request body: %s, url %s, headers: %s", diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 8567e84950f..da052e71b58 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -300,3 +300,181 @@ def test_optional_params_returned_when_properly_overridden(): print("FIELDS", fields) assert "optional_params" in fields + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_prepare_request_with_api_key(): + """Test _prepare_request method uses Bearer token when api_key is provided in data""" + from unittest.mock import Mock, patch + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + + # Setup guardrail hook + guardrail_hook = BedrockGuardrail( + guardrailIdentifier="test-guardrail-id", + guardrailVersion="1" + ) + mock_credentials = Mock() + test_data = { + "source": "INPUT", + "content": [{"text": {"text": "test content"}}] + } + + prepared_request = guardrail_hook._prepare_request( + credentials=mock_credentials, + data=test_data, + optional_params={}, + aws_region_name="us-east-1", + api_key="test-bearer-token-123" + ) + + # Verify Bearer token is used in Authorization header + assert "Authorization" in prepared_request.headers + assert prepared_request.headers["Authorization"] == "Bearer test-bearer-token-123" + + # Verify URL is correct + expected_url = "https://bedrock-runtime.us-east-1.amazonaws.com/guardrail/test-guardrail-id/version/1/apply" + assert prepared_request.url == expected_url + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_prepare_request_without_api_key(): + """Test _prepare_request method falls back to SigV4 when no api_key is provided""" + from unittest.mock import Mock, patch + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + + # Setup guardrail hook + guardrail_hook = BedrockGuardrail( + guardrailIdentifier="test-guardrail-id", + guardrailVersion="1" + ) + + # Mock credentials + mock_credentials = Mock() + + # Test data without api_key + test_data = { + "source": "INPUT", + "content": [{"text": {"text": "test content"}}] + } + + with patch("litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.get_secret_str") as mock_get_secret, \ + patch("botocore.auth.SigV4Auth") as mock_sigv4_auth, \ + patch("botocore.awsrequest.AWSRequest") as mock_aws_request: + + # Mock no AWS_BEARER_TOKEN_BEDROCK + mock_get_secret.return_value = None + + # Mock SigV4Auth + mock_sigv4_instance = Mock() + mock_sigv4_auth.return_value = mock_sigv4_instance + + # Mock AWSRequest + mock_request_instance = Mock() + mock_request_instance.prepare.return_value = Mock() + mock_aws_request.return_value = mock_request_instance + + # Call _prepare_request + prepared_request = guardrail_hook._prepare_request( + credentials=mock_credentials, + data=test_data, + optional_params={}, + aws_region_name="us-east-1" + ) + + # Verify SigV4 auth was used + mock_sigv4_auth.assert_called_once_with(mock_credentials, "bedrock", "us-east-1") + mock_sigv4_instance.add_auth.assert_called_once() + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_prepare_request_with_bearer_token_env(): + """Test _prepare_request method uses Bearer token from environment when available""" + from unittest.mock import Mock, patch + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + + # Setup guardrail hook + guardrail_hook = BedrockGuardrail( + guardrailIdentifier="test-guardrail-id", + guardrailVersion="1" + ) + + # Mock credentials + mock_credentials = Mock() + + # Test data without api_key + test_data = { + "source": "INPUT", + "content": [{"text": {"text": "test content"}}] + } + + with patch("litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.get_secret_str") as mock_get_secret, \ + patch("botocore.awsrequest.AWSRequest") as mock_aws_request: + + mock_get_secret.return_value = "env-bearer-token-456" + mock_request_instance = Mock() + mock_request_instance.prepare.return_value = Mock() + mock_aws_request.return_value = mock_request_instance + + prepared_request = guardrail_hook._prepare_request( + credentials=mock_credentials, + data=test_data, + optional_params={}, + aws_region_name="us-east-1" + ) + + # Verify Bearer token from environment is used + mock_aws_request.assert_called_once() + call_args = mock_aws_request.call_args + headers = call_args[1]["headers"] + assert headers["Authorization"] == "Bearer env-bearer-token-456" + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_make_api_request_passes_api_key(): + """Test make_bedrock_api_request method correctly passes api_key from request_data""" + from unittest.mock import Mock, patch, AsyncMock + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + + guardrail_hook = BedrockGuardrail( + guardrailIdentifier="test-guardrail-id", + guardrailVersion="1" + ) + + guardrail_hook.async_handler = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = {"action": "NONE", "outputs": []} + guardrail_hook.async_handler.post = AsyncMock(return_value=mock_response) + + test_request_data = { + "api_key": "test-api-key-789" + } + + with patch.object(guardrail_hook, "_load_credentials") as mock_load_creds, \ + patch.object(guardrail_hook, "convert_to_bedrock_format") as mock_convert, \ + patch.object(guardrail_hook, "get_guardrail_dynamic_request_body_params") as mock_get_params, \ + patch.object(guardrail_hook, "add_standard_logging_guardrail_information_to_request_data"), \ + patch("botocore.awsrequest.AWSRequest") as mock_aws_request: + + mock_load_creds.return_value = (Mock(), "us-east-1") + mock_convert.return_value = {"source": "INPUT", "content": []} + mock_get_params.return_value = {} + + mock_request_instance = Mock() + mock_request_instance.url = "test-url" + mock_request_instance.body = b"test-body" + mock_request_instance.headers = {"Content-Type": "application/json", "Authorization": "Bearer test-api-key-789"} + mock_request_instance.prepare.return_value = Mock() + mock_aws_request.return_value = mock_request_instance + + await guardrail_hook.make_bedrock_api_request( + source="INPUT", + messages=[{"role": "user", "content": "test"}], + request_data=test_request_data + ) + + # Verify _prepare_request was invoked and used the api_key + mock_aws_request.assert_called_once() + call_args = mock_aws_request.call_args + headers = call_args[1]["headers"] + assert headers["Authorization"] == "Bearer test-api-key-789" \ No newline at end of file