fix: scope Bedrock guardrail routing

This commit is contained in:
aiedwardyi 2026-08-24 18:17:57 +09:00
parent 81c21fcb5b
commit 08cfef6791
No known key found for this signature in database
2 changed files with 107 additions and 16 deletions

View file

@ -31,7 +31,6 @@ from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS
from litellm.exceptions import ModifyResponseException
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
redact_nested_match_and_regex_keys,
)
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import bedrock_guardrail_cost
@ -693,16 +692,51 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if llm_router is None:
return None
metadata_key: Final = get_metadata_variable_name_from_kwargs(dict(request_data))
metadata: Final = request_data.get(metadata_key)
team_id: Final[object | None] = (
metadata.get("user_api_key_team_id") if isinstance(metadata, Mapping) else None
team_id: Final[object | None] = next(
(
metadata.get("user_api_key_team_id")
for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata"))
if isinstance(metadata, Mapping) and isinstance(metadata.get("user_api_key_team_id"), str)
),
None,
)
try:
deployments: Final = llm_router.get_model_list(
model_name=model,
team_id=team_id if isinstance(team_id, str) else None,
) or []
resolved_team_id: Final = team_id if isinstance(team_id, str) else None
listed_deployments: Final = (
llm_router.get_model_list(
model_name=model,
team_id=resolved_team_id,
)
or []
)
model_id_deployment: Final = (
llm_router.get_deployment(model_id=model)
if not listed_deployments and llm_router.has_model_id(model) is True
else None
)
model_id_deployment_row: Final = (
model_id_deployment.model_dump(exclude_none=True)
if model_id_deployment is not None and hasattr(model_id_deployment, "model_dump")
else model_id_deployment
)
candidate_deployments: Final = (
[model_id_deployment_row] if model_id_deployment_row is not None else listed_deployments
)
filter_deployments: Final = getattr(llm_router, "_filter_deployments_by_model_access_groups", None)
filtered_deployments: Final = (
filter_deployments(
model=model,
healthy_deployments=candidate_deployments,
request_kwargs=dict(request_data),
request_team_id=resolved_team_id,
)
if callable(filter_deployments) and isinstance(candidate_deployments, list)
else None
)
deployments: Final = (
filtered_deployments if isinstance(filtered_deployments, list) else candidate_deployments
)
except Exception:
return False
active_deployments = []
@ -713,9 +747,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
else getattr(deployment, "model_info", None)
)
blocked = (
model_info.get("blocked")
if isinstance(model_info, Mapping)
else getattr(model_info, "blocked", None)
model_info.get("blocked") if isinstance(model_info, Mapping) else getattr(model_info, "blocked", None)
)
if blocked is not True:
active_deployments.append(deployment)
@ -724,10 +756,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
providers: list[str] = []
for deployment in active_deployments:
params: object = deployment.get("litellm_params") if isinstance(deployment, Mapping) else None
provider: object = params.get("custom_llm_provider") if isinstance(params, Mapping) else None
if not isinstance(provider, str) and isinstance(params, Mapping):
deployment_model: Final[object | None] = params.get("model")
params: object = (
deployment.get("litellm_params")
if isinstance(deployment, Mapping)
else getattr(deployment, "litellm_params", None)
)
provider: object = (
params.get("custom_llm_provider")
if isinstance(params, Mapping)
else getattr(params, "custom_llm_provider", None)
)
if not isinstance(provider, str):
deployment_model: Final[object | None] = (
params.get("model") if isinstance(params, Mapping) else getattr(params, "model", None)
)
provider = (
BedrockGuardrail._resolve_model_provider(deployment_model)
if isinstance(deployment_model, str)

View file

@ -28,6 +28,55 @@ def test_bedrock_guardrail_uses_active_metadata_bucket_for_team_id():
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 = [