From 99e4ba34c578c1cc1c41b11145588c9c12a91bbe Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Mon, 24 Aug 2026 01:33:28 +0900 Subject: [PATCH] fix: scope Bedrock guardrail bearer keys --- .../guardrail_hooks/bedrock_guardrails.py | 24 +++++++-- .../test_bedrock_guardrails.py | 49 +++++++++++++++++++ .../test_bedrock_invoke_guardrail_checks.py | 42 ++++++++++++++++ 3 files changed, 111 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index c70a2ee8a74..06772f0ad07 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -669,6 +669,24 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # logic becomes shared across providers. #### CALL HOOKS - proxy only #### + @staticmethod + def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None: + if not request_data: + return None + + top_level_provider: Final[object | None] = request_data.get("custom_llm_provider") + request_provider: Final[str | None] = ( + top_level_provider if isinstance(top_level_provider, str) else None + ) + model: Final[object | None] = request_data.get("model") + model_provider: Final = model.partition("/")[0] if isinstance(model, str) else None + custom_llm_provider: Final = request_provider or model_provider + if custom_llm_provider not in ("bedrock", "bedrock_converse"): + return None + + api_key: Final[object | None] = request_data.get("api_key") + return api_key if isinstance(api_key, str) else None + def _load_credentials( self, ): @@ -841,7 +859,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_request_data: Final[dict] = dict( self.convert_to_bedrock_format(source=source, messages=messages, response=response) ) - api_key: str | None = None + api_key: Final = self._get_bedrock_api_key(request_data) if request_data: dynamic_request_body_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data) bedrock_request_data.update( @@ -851,8 +869,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if key not in _BEDROCK_DYNAMIC_BODY_DENYLIST } ) - if request_data.get("api_key") is not None: - api_key = request_data["api_key"] event_type: Final = ( logging_event_type @@ -1828,7 +1844,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): credentials, aws_region_name = self._load_credentials() body: Final[dict[str, Any]] = {"messages": checks_messages, "checks": self.checks} - api_key: Final[str | None] = request_data.get("api_key") if request_data else None + api_key: Final = self._get_bedrock_api_key(request_data) prepared_request: Final = self._prepare_request( credentials=credentials, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index dd339d4e51f..c111072a06a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -1114,6 +1114,55 @@ async def test_make_apply_guardrail_request_skips_scan_without_credentials(): mock_post.assert_not_called() +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("custom_llm_provider", "expected_auth_prefix"), + [("nvidia_nim", "AWS4-HMAC-SHA256"), ("bedrock", "Bearer bedrock-key")], +) +async def test_make_apply_guardrail_request_scopes_api_key_to_bedrock_provider( + custom_llm_provider, expected_auth_prefix, monkeypatch +): + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + guardrail = BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + ) + request_data = { + "model": f"{custom_llm_provider}/test-model", + "api_key": "bedrock-key", + } + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = {"action": "NONE", "assessments": []} + + with ( + patch.object( + guardrail, + "_load_credentials", + return_value=(mock_credentials, "us-east-1"), + ), + patch.object( + guardrail.async_handler, + "post", + new=AsyncMock(return_value=mock_response), + ) as mock_post, + ): + result = await guardrail.make_bedrock_api_request( + source="INPUT", + messages=[{"role": "user", "content": "hello"}], + request_data=request_data, + ) + + assert result["action"] == "NONE" + assert mock_post.await_args.kwargs["headers"]["Authorization"].startswith( + expected_auth_prefix + ) + + @pytest.mark.asyncio async def test_bedrock_apply_guardrail_response_uses_OUTPUT_source(): """input_type='response' must call Bedrock with source=OUTPUT and assistant content. diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py index d842a1ee5f9..634e2254187 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py @@ -395,6 +395,48 @@ async def test_request_uses_checks_path_and_body(): ] +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("custom_llm_provider", "expected_auth_prefix"), + [("nvidia_nim", "AWS4-HMAC-SHA256"), ("bedrock", "Bearer bedrock-key")], +) +async def test_request_scopes_api_key_to_bedrock_provider( + custom_llm_provider, expected_auth_prefix, monkeypatch +): + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS) + request_data = { + "model": f"{custom_llm_provider}/test-model", + "api_key": "bedrock-key", + } + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + mock_post = AsyncMock( + return_value=_mock_http_response(200, {"results": {}}) + ) + + with ( + patch.object( + g, + "_load_credentials", + return_value=(mock_credentials, "us-east-1"), + ), + patch.object(g.async_handler, "post", new=mock_post), + ): + result = await g.make_bedrock_api_request( + source="INPUT", + messages=[{"role": "user", "content": "hello"}], + request_data=request_data, + ) + + assert result == BedrockGuardrailResponse() + assert mock_post.await_args.kwargs["headers"]["Authorization"].startswith( + expected_auth_prefix + ) + + @pytest.mark.asyncio async def test_empty_messages_passes_without_api_call(): g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS)