fix: resolve native Bedrock guardrail providers

This commit is contained in:
aiedwardyi 2026-08-24 14:01:02 +09:00
parent 99e4ba34c5
commit 31a529aee6
No known key found for this signature in database
4 changed files with 34 additions and 13 deletions

View file

@ -669,6 +669,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# logic becomes shared across providers.
#### CALL HOOKS - proxy only ####
@staticmethod
def _resolve_model_provider(model: str) -> str | None:
try:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
return custom_llm_provider
except Exception:
return model.partition("/")[0]
@staticmethod
def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None:
if not request_data:
@ -679,8 +687,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
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
model_provider: Final[str | None] = (
BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None
)
custom_llm_provider: Final[str | None] = request_provider or model_provider
if custom_llm_provider not in ("bedrock", "bedrock_converse"):
return None

View file

@ -1116,19 +1116,23 @@ async def test_make_apply_guardrail_request_skips_scan_without_credentials():
@pytest.mark.asyncio
@pytest.mark.parametrize(
("custom_llm_provider", "expected_auth_prefix"),
[("nvidia_nim", "AWS4-HMAC-SHA256"), ("bedrock", "Bearer bedrock-key")],
("model", "expected_auth_prefix"),
[
("nvidia_nim/test-model", "AWS4-HMAC-SHA256"),
("bedrock/test-model", "Bearer bedrock-key"),
("amazon.nova-lite-v1:0", "Bearer bedrock-key"),
],
)
async def test_make_apply_guardrail_request_scopes_api_key_to_bedrock_provider(
custom_llm_provider, expected_auth_prefix, monkeypatch
):
model: str, expected_auth_prefix: str, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
)
request_data = {
"model": f"{custom_llm_provider}/test-model",
"model": model,
"api_key": "bedrock-key",
}
mock_credentials = MagicMock()

View file

@ -397,16 +397,20 @@ 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")],
("model", "expected_auth_prefix"),
[
("nvidia_nim/test-model", "AWS4-HMAC-SHA256"),
("bedrock/test-model", "Bearer bedrock-key"),
("amazon.nova-lite-v1:0", "Bearer bedrock-key"),
],
)
async def test_request_scopes_api_key_to_bedrock_provider(
custom_llm_provider, expected_auth_prefix, monkeypatch
):
model: str, expected_auth_prefix: str, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS)
request_data = {
"model": f"{custom_llm_provider}/test-model",
"model": model,
"api_key": "bedrock-key",
}
mock_credentials = MagicMock()

View file

@ -889,7 +889,10 @@ async def test_bedrock_guardrail_make_api_request_passes_api_key():
mock_response.status_code = 200
mock_response.json.return_value = {"action": "NONE", "outputs": []}
test_request_data = {"api_key": "test-api-key-789"}
test_request_data = {
"model": "bedrock/test-model",
"api_key": "test-api-key-789",
}
with (
patch.object(