fix: scope Bedrock guardrail bearer keys

This commit is contained in:
aiedwardyi 2026-08-24 01:33:28 +09:00
parent f005afa146
commit 99e4ba34c5
No known key found for this signature in database
3 changed files with 111 additions and 4 deletions

View file

@ -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,

View file

@ -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.

View file

@ -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)