From 9fe37b1a3a8a86e0318e4807283493c6e28ba484 Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Tue, 25 Aug 2026 00:04:00 +0900 Subject: [PATCH] fix: align Bedrock router eligibility --- .../guardrail_hooks/bedrock_guardrails.py | 248 ++++++++++++++---- .../test_bedrock_guardrails.py | 137 ++++++++++ 2 files changed, 335 insertions(+), 50 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 5a5b0307fa9..5d045e4d64e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -53,6 +53,9 @@ from litellm.proxy.guardrails.anthropic_sse import ( is_raw_sse_stream, model_response_text, ) +from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs, is_valid_deployment_tag +from litellm.router_utils.common_utils import filter_team_based_models +from litellm.router_utils.cooldown_handlers import _get_cooldown_deployments from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import BedrockChecksConfigModel, GuardrailEventHooks from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage @@ -676,6 +679,82 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): except Exception: # noqa: BLE001 # provider resolution has a safe prefix fallback return model.partition("/")[0] + @staticmethod + def _filter_router_deployments_by_tags( + router: object, + deployments: list[object], + request_data: Mapping[str, object], + ) -> list[object]: + request_tags: Final = tuple(_get_tags_from_request_kwargs(request_data)) + deployment_tag_filtering: Final = any( + BedrockGuardrail._router_deployment_field(deployment, "enable_tag_filtering") is True + for deployment in deployments + ) + tag_filtering_enabled: Final = ( + request_data.get("enable_tag_filtering") is True + or getattr(router, "enable_tag_filtering", False) is True + or deployment_tag_filtering + ) + if not tag_filtering_enabled: + return deployments + + def _deployment_tags(deployment: object) -> tuple[str, ...]: + params: Final[object | None] = ( + deployment.get("litellm_params") + if isinstance(deployment, Mapping) + else getattr(deployment, "litellm_params", None) + ) + tags: Final[object | None] = ( + params.get("tags") if isinstance(params, Mapping) else getattr(params, "tags", None) + ) + return ( + tuple(tag for tag in tags if isinstance(tag, str)) + if isinstance(tags, Sequence) and not isinstance(tags, str) + else () + ) + + if not request_tags: + default_deployments: Final = [ + deployment for deployment in deployments if "default" in _deployment_tags(deployment) + ] + return default_deployments or deployments + + required_tags: Final = frozenset(tag[1:] for tag in request_tags if tag.startswith("&") and len(tag) > 1) + excluded_tags: Final = frozenset(tag[1:] for tag in request_tags if tag.startswith("!") and len(tag) > 1) + positive_tags: Final = tuple(tag for tag in request_tags if not tag.startswith(("&", "!"))) + allowed_deployments: Final = [ + deployment for deployment in deployments if not excluded_tags.intersection(_deployment_tags(deployment)) + ] + required_deployments: Final = [ + deployment for deployment in allowed_deployments if required_tags.issubset(_deployment_tags(deployment)) + ] + if not positive_tags: + return required_deployments + + match_any: Final[bool] = ( + getattr(router, "tag_filtering_match_any", True) + if isinstance(getattr(router, "tag_filtering_match_any", True), bool) + else True + ) + matched_deployments: Final = [ + deployment + for deployment in required_deployments + if is_valid_deployment_tag(_deployment_tags(deployment), positive_tags, match_any) + ] + if matched_deployments: + return matched_deployments + fallback_default_deployments: Final = [ + deployment for deployment in required_deployments if "default" in _deployment_tags(deployment) + ] + return fallback_default_deployments + + @staticmethod + def _router_deployment_field(deployment: object, field: str) -> object | None: + model_info: Final[object | None] = ( + deployment.get("model_info") if isinstance(deployment, Mapping) else getattr(deployment, "model_info", None) + ) + return model_info.get(field) if isinstance(model_info, Mapping) else getattr(model_info, field, None) + @staticmethod def _router_allows_bedrock(request_data: Mapping[str, object]) -> bool | None: model: Final[object | None] = request_data.get("model") @@ -699,76 +778,145 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) try: 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_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 team_filtered_deployments and llm_router.has_model_id(model) is True - else None + llm_router.get_deployment(model_id=model) if 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 team_filtered_deployments + raw_listed_deployments: Final = ( + [] + if model_id_deployment_row is not None + else ( + llm_router.get_model_list( + model_name=model, + team_id=resolved_team_id, + ) + or [] + ) + ) + model_group_aliases: Final = getattr(llm_router, "model_group_alias", None) + concrete_model_names: Final[object | None] = getattr(llm_router, "model_names", None) + is_concrete_model: Final = ( + isinstance(concrete_model_names, (list, tuple, set, frozenset)) and model in concrete_model_names + ) + is_model_alias: Final = isinstance(model_group_aliases, Mapping) and model in model_group_aliases + if model_id_deployment_row is None and not is_concrete_model and not is_model_alias: + pattern_router: Final[object | None] = getattr(llm_router, "pattern_router", None) + get_pattern_deployments: Final[object | None] = getattr( + pattern_router, "get_deployments_by_pattern", None + ) + global_pattern_deployments: Final = ( + get_pattern_deployments(model=model) if callable(get_pattern_deployments) else None + ) + team_pattern_router: Final[object | None] = ( + getattr(llm_router, "team_pattern_routers", {}).get(resolved_team_id) + if resolved_team_id is not None + and isinstance(getattr(llm_router, "team_pattern_routers", None), Mapping) + else None + ) + get_team_pattern_deployments: Final[object | None] = getattr( + team_pattern_router, "get_deployments_by_pattern", None + ) + team_pattern_deployments: Final = ( + get_team_pattern_deployments(model=model) if callable(get_team_pattern_deployments) else None + ) + selected_listed_deployments: Final = ( + global_pattern_deployments + if isinstance(global_pattern_deployments, list) and global_pattern_deployments + else ( + team_pattern_deployments + if isinstance(team_pattern_deployments, list) and team_pattern_deployments + else raw_listed_deployments + ) + ) + else: + selected_listed_deployments: Final = raw_listed_deployments + candidate_deployments: Final[list[object]] = ( + [model_id_deployment_row] + if model_id_deployment_row is not None + else [deployment for deployment in selected_listed_deployments if isinstance(deployment, Mapping)] + ) + router_matched: Final = bool(candidate_deployments) + team_filtered_result: Final = ( + candidate_deployments + if model_id_deployment_row is not None + else filter_team_based_models( + healthy_deployments=candidate_deployments, + request_kwargs=dict(request_data), + ) + ) + team_filtered_deployments: Final[list[object]] = ( + team_filtered_result if isinstance(team_filtered_result, list) else candidate_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, + healthy_deployments=team_filtered_deployments, request_kwargs=dict(request_data), request_team_id=resolved_team_id, ) - if callable(filter_deployments) and isinstance(candidate_deployments, list) + if callable(filter_deployments) + and isinstance(team_filtered_deployments, list) + and model_id_deployment_row is None else None ) - deployments: Final = ( - filtered_deployments if isinstance(filtered_deployments, list) else candidate_deployments + access_filtered_deployments: Final[list[object]] = ( + filtered_deployments if isinstance(filtered_deployments, list) else team_filtered_deployments + ) + health_filter: Final[object | None] = getattr( + llm_router, "_filter_health_check_unhealthy_deployments", None + ) + health_filtered_deployments: Final = ( + health_filter( + healthy_deployments=access_filtered_deployments, + parent_otel_span=None, + ) + if callable(health_filter) + else access_filtered_deployments + ) + healthy_deployments: Final[list[object]] = ( + health_filtered_deployments + if isinstance(health_filtered_deployments, list) + else access_filtered_deployments + ) + cooldown_cache: Final[object | None] = getattr(llm_router, "cooldown_cache", None) + cooldown_lookup: Final[object | None] = getattr(cooldown_cache, "get_active_cooldowns", None) + cooldown_deployments: Final = ( + _get_cooldown_deployments( + litellm_router_instance=llm_router, + parent_otel_span=None, + ) + if callable(cooldown_lookup) and callable(getattr(llm_router, "get_model_ids", None)) + else [] + ) + cooldown_ids: Final[frozenset[str]] = frozenset( + deployment_id for deployment_id in cooldown_deployments if isinstance(deployment_id, str) + ) + cooldown_filtered_deployments: Final = [ + deployment + for deployment in healthy_deployments + if BedrockGuardrail._router_deployment_field(deployment, "id") not in cooldown_ids + ] + unblocked_deployments: Final = [ + deployment + for deployment in cooldown_filtered_deployments + if BedrockGuardrail._router_deployment_field(deployment, "blocked") is not True + ] + deployments: Final = BedrockGuardrail._filter_router_deployments_by_tags( + router=llm_router, + deployments=unblocked_deployments, + request_data=request_data, ) except Exception: # noqa: BLE001 # optional router state must not break guardrail auth return False - active_deployments = [] - for deployment in deployments: - model_info = ( - deployment.get("model_info") - if isinstance(deployment, Mapping) - else getattr(deployment, "model_info", None) - ) - blocked = ( - model_info.get("blocked") if isinstance(model_info, Mapping) else getattr(model_info, "blocked", None) - ) - if blocked is not True: - active_deployments.append(deployment) - if not active_deployments: + if not deployments: + if router_matched: + return False router_settings = getattr(llm_router, "router_general_settings", None) if getattr(router_settings, "pass_through_all_models", False) is True: provider = request_data.get("custom_llm_provider") @@ -803,7 +951,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return False providers: list[str] = [] - for deployment in active_deployments: + for deployment in deployments: params: object = ( deployment.get("litellm_params") if isinstance(deployment, Mapping) 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 0340f83c9d3..ae11637a0a5 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 @@ -112,6 +112,143 @@ def test_bedrock_guardrail_accepts_bedrock_default_deployment(monkeypatch: pytes assert BedrockGuardrail._router_allows_bedrock({"model": "unlisted-model"}) is True +def test_bedrock_guardrail_rejects_blocked_model_with_pass_through(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {"blocked": True}} + ] + router.router_general_settings.pass_through_all_models = True + router.default_deployment = None + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._get_bedrock_api_key( + { + "model": "blocked-alias", + "custom_llm_provider": "bedrock", + "api_key": "bedrock-key", + } + ) + is None + ) + + +def test_bedrock_guardrail_rejects_access_filtered_model_with_pass_through( + monkeypatch: pytest.MonkeyPatch, +): + from litellm.proxy import proxy_server + + 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 = [] + router.router_general_settings.pass_through_all_models = True + router.default_deployment = None + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._get_bedrock_api_key( + { + "model": "scoped-alias", + "custom_llm_provider": "bedrock", + "api_key": "bedrock-key", + } + ) + is None + ) + + +def test_bedrock_guardrail_resolves_model_id_before_wildcards(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.has_model_id.return_value = True + deployment = MagicMock() + deployment.model_dump.return_value = { + "litellm_params": {"custom_llm_provider": "bedrock"}, + "model_info": {"id": "deployment-id"}, + } + router.get_deployment.return_value = deployment + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}} + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert BedrockGuardrail._router_allows_bedrock({"model": "deployment-id"}) is True + router.get_model_list.assert_not_called() + + +def test_bedrock_guardrail_ignores_cooling_non_bedrock_deployments(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.get_model_list.return_value = [ + { + "litellm_params": {"custom_llm_provider": "openai"}, + "model_info": {"id": "openai-deployment"}, + }, + { + "litellm_params": {"custom_llm_provider": "bedrock"}, + "model_info": {"id": "bedrock-deployment"}, + }, + ] + router.get_model_ids.return_value = ["openai-deployment", "bedrock-deployment"] + router.cooldown_cache.get_active_cooldowns.return_value = [("openai-deployment", 1.0)] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert BedrockGuardrail._router_allows_bedrock({"model": "shared-alias"}) is True + + +def test_bedrock_guardrail_matches_global_wildcard_precedence(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.model_names = [] + router.model_group_alias = {} + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}} + ] + router.pattern_router.get_deployments_by_pattern.return_value = [ + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}} + ] + team_pattern_router = MagicMock() + team_pattern_router.get_deployments_by_pattern.return_value = [ + {"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}} + ] + router.team_pattern_routers = {"team-id": team_pattern_router} + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "provider/model", "litellm_metadata": {"user_api_key_team_id": "team-id"}} + ) + is True + ) + + +def test_bedrock_guardrail_matches_request_tag_pool(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.enable_tag_filtering = True + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "openai", "tags": ["slow"]}, "model_info": {}}, + {"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["fast"]}, "model_info": {}}, + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "metadata": {"tags": ["fast"]}} + ) + is True + ) + + def test_bedrock_guardrail_explicit_non_bedrock_provider_wins_alias(monkeypatch: pytest.MonkeyPatch): from litellm.proxy import proxy_server