mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix the api key support for bedrock guardrail in proxy
This commit is contained in:
parent
4df07a5060
commit
91040b6b56
2 changed files with 217 additions and 19 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue