fix: address Bedrock guardrail review feedback

This commit is contained in:
aiedwardyi 2026-08-24 22:47:46 +09:00
parent 08cfef6791
commit 34e041e9f7
No known key found for this signature in database
3 changed files with 150 additions and 85 deletions

View file

@ -667,16 +667,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return merged_messages
# NOTE: Consider moving these helpers to CustomGuardrail when the filtering
# 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:
except Exception: # noqa: BLE001 # provider resolution has a safe prefix fallback
return model.partition("/")[0]
@staticmethod
@ -709,9 +706,29 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
or []
)
model_group_aliases: Final = getattr(llm_router, "model_group_alias", None)
team_filtered_deployments: Final = (
[
deployment
for deployment in listed_deployments
if (
(
deployment.get("model_info")
if isinstance(deployment, Mapping)
else getattr(deployment, "model_info", None)
)
or {}
).get("team_id")
in (None, resolved_team_id)
]
if resolved_team_id is not None
and isinstance(model_group_aliases, Mapping)
and model in model_group_aliases
else listed_deployments
)
model_id_deployment: Final = (
llm_router.get_deployment(model_id=model)
if not listed_deployments and llm_router.has_model_id(model) is True
if not team_filtered_deployments and llm_router.has_model_id(model) is True
else None
)
model_id_deployment_row: Final = (
@ -720,7 +737,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
else model_id_deployment
)
candidate_deployments: Final = (
[model_id_deployment_row] if model_id_deployment_row is not None else listed_deployments
[model_id_deployment_row]
if model_id_deployment_row is not None
else team_filtered_deployments
)
filter_deployments: Final = getattr(llm_router, "_filter_deployments_by_model_access_groups", None)
@ -737,7 +756,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
deployments: Final = (
filtered_deployments if isinstance(filtered_deployments, list) else candidate_deployments
)
except Exception:
except Exception: # noqa: BLE001 # optional router state must not break guardrail auth
return False
active_deployments = []
for deployment in deployments:
@ -793,6 +812,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if router_allows_bedrock is not None:
return api_key if router_allows_bedrock else None
explicit_provider: Final[object | None] = request_data.get("custom_llm_provider")
if isinstance(explicit_provider, str):
return api_key if explicit_provider in ("bedrock", "bedrock_converse") else None
model: Final[object | None] = request_data.get("model")
model_provider: Final[str | None] = (
BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None

View file

@ -10,84 +10,6 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.caching import DualCache
from unittest.mock import MagicMock, AsyncMock, patch
def test_bedrock_guardrail_uses_active_metadata_bucket_for_team_id():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
]
request_data = {
"model": "team-alias",
"metadata": {"user_api_key_team_id": "legacy-team"},
"litellm_metadata": {"user_api_key_team_id": "active-team"},
}
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="active-team")
def test_bedrock_guardrail_uses_proxy_team_when_alternate_metadata_is_empty():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
]
request_data = {
"model": "team-alias",
"metadata": {"user_api_key_team_id": "proxy-team"},
"litellm_metadata": {},
}
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="proxy-team")
def test_bedrock_guardrail_resolves_router_model_id():
router = MagicMock()
router.get_model_list.return_value = []
router.has_model_id.return_value = True
deployment = MagicMock()
deployment.model_dump.return_value = {
"litellm_params": {"custom_llm_provider": "bedrock"},
"model_info": {},
}
router.get_deployment.return_value = deployment
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock({"model": "deployment-id"}) is True
router.get_deployment.assert_called_once_with(model_id="deployment-id")
deployment.model_dump.assert_called_once_with(exclude_none=True)
def test_bedrock_guardrail_filters_access_group_deployments():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}},
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
]
router._filter_deployments_by_model_access_groups.return_value = [
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
]
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock({"model": "scoped-alias"}) is True
def test_bedrock_guardrail_ignores_blocked_deployments():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {"blocked": True}},
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
]
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock({"model": "alias"}) is True
@pytest.mark.asyncio
async def test_bedrock_guardrails_pii_masking():
# Create proper mock objects

View file

@ -30,6 +30,126 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
from litellm.types.utils import CallTypes, ModelResponse
def test_bedrock_guardrail_uses_active_metadata_bucket_for_team_id():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
]
request_data = {
"model": "team-alias",
"metadata": {"user_api_key_team_id": "legacy-team"},
"litellm_metadata": {"user_api_key_team_id": "active-team"},
}
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="active-team")
def test_bedrock_guardrail_uses_proxy_team_when_alternate_metadata_is_empty():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
]
request_data = {
"model": "team-alias",
"metadata": {"user_api_key_team_id": "proxy-team"},
"litellm_metadata": {},
}
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="proxy-team")
def test_bedrock_guardrail_resolves_router_model_id():
router = MagicMock()
router.get_model_list.return_value = []
router.has_model_id.return_value = True
deployment = MagicMock()
deployment.model_dump.return_value = {
"litellm_params": {"custom_llm_provider": "bedrock"},
"model_info": {},
}
router.get_deployment.return_value = deployment
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock({"model": "deployment-id"}) is True
router.get_deployment.assert_called_once_with(model_id="deployment-id")
deployment.model_dump.assert_called_once_with(exclude_none=True)
def test_bedrock_guardrail_filters_access_group_deployments():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}},
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
]
router._filter_deployments_by_model_access_groups.return_value = [
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
]
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock({"model": "scoped-alias"}) is True
def test_bedrock_guardrail_ignores_blocked_deployments():
router = MagicMock()
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {"blocked": True}},
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
]
with patch("litellm.proxy.proxy_server.llm_router", router):
assert BedrockGuardrail._router_allows_bedrock({"model": "alias"}) is True
def test_bedrock_guardrail_filters_alias_deployments_by_team():
router = MagicMock()
router.model_group_alias = {"team-alias": "shared-group"}
router.get_model_list.return_value = [
{
"litellm_params": {"custom_llm_provider": "openai"},
"model_info": {"team_id": "other-team"},
},
{
"litellm_params": {"custom_llm_provider": "bedrock"},
"model_info": {"team_id": "active-team"},
},
]
with patch("litellm.proxy.proxy_server.llm_router", router):
assert (
BedrockGuardrail._router_allows_bedrock(
{
"model": "team-alias",
"litellm_metadata": {"user_api_key_team_id": "active-team"},
}
)
is True
)
def test_bedrock_guardrail_honors_explicit_provider_without_router(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "llm_router", None)
assert (
BedrockGuardrail._get_bedrock_api_key(
{
"model": "bedrock/anthropic.claude-3-haiku",
"custom_llm_provider": "openai",
"api_key": "openai-key",
}
)
is None
)
@pytest.mark.asyncio
async def test__redact_pii_matches_function():
"""Test the _redact_pii_matches function directly"""