From 053a26a57ae3eee1b984a98179162f3b19da7243 Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Wed, 26 Aug 2026 16:18:59 +0900 Subject: [PATCH] fix: align Bedrock tag fallback --- .../guardrail_hooks/bedrock_guardrails.py | 57 ++++++++++---- .../test_bedrock_guardrails.py | 77 +++++++++++++++++++ 2 files changed, 120 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 5f5d2d231b1..4af516abae2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -57,10 +57,12 @@ from litellm.proxy.guardrails.anthropic_sse import ( from litellm.router_strategy.tag_based_routing import ( _chain_tag_filtering_override, _get_tags_from_request_kwargs, + _inherited_constraint_sets, _match_deployment, _request_tags_after_router_consumption, _split_tags, _strip_routing_prefix, + _unknown_required_tag_hides_an_answer, ) from litellm.router_utils.common_utils import filter_team_based_models, filter_web_search_deployments from litellm.router_utils.cooldown_handlers import _get_cooldown_deployments @@ -744,10 +746,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_tags = _request_tags_after_router_consumption(metadata, model) or () routing_prefix: Final[object] = getattr(router, "tag_routing_prefix", "") resolved_prefix: Final[str] = routing_prefix if isinstance(routing_prefix, str) else "" - rewritten_tags, _ = _strip_routing_prefix(request_tags, resolved_prefix) + rewritten_tags, routing_confirmed = _strip_routing_prefix(request_tags, resolved_prefix) required_tags, positive_tags, excluded_tags = _split_tags(rewritten_tags) required_set: Final = frozenset(required_tags) excluded_set: Final = frozenset(excluded_tags) + inherited_required_set, inherited_excluded_set = _inherited_constraint_sets( + metadata.get("inherited_tags"), resolved_prefix + ) allowed_deployments: Final = [ deployment for deployment in deployments if not excluded_set.intersection(_deployment_tags(deployment)) ] @@ -764,9 +769,41 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): has_positive_filter: Final = bool(positive_tags) or ( bool(header_strings) and has_regex_deployments and not required_set ) + + def _fail_open_deployments() -> list[object]: + if _unknown_required_tag_hides_an_answer( + deployments, + excluded_set, + required_set, + routing_confirmed, + ): + return [] + if not any( + isinstance(deployment, Mapping) + and (deployment.get("model_info") or {}).get("allow_fail_open") is True + for deployment in deployments + ): + return [] + trusted_excluded: Final = ( + frozenset() if inherited_excluded_set is None else inherited_excluded_set & excluded_set + ) + trusted_required: Final = ( + frozenset() if inherited_required_set is None else inherited_required_set & required_set + ) + trusted_deployments: Final = [ + deployment + for deployment in deployments + if not trusted_excluded.intersection(_deployment_tags(deployment)) + and trusted_required.issubset(_deployment_tags(deployment)) + ] + default_deployments: Final = [ + deployment for deployment in trusted_deployments if "default" in _deployment_tags(deployment) + ] + return default_deployments or trusted_deployments + if not has_positive_filter: if required_set or excluded_set: - return candidate_deployments + return candidate_deployments or _fail_open_deployments() default_deployments: Final = [ deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment) ] @@ -791,12 +828,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ] if matched_deployments: return matched_deployments - if any( - isinstance(deployment, Mapping) and (deployment.get("model_info") or {}).get("allow_fail_open") is True - for deployment in allowed_deployments - ): - return candidate_deployments or allowed_deployments - return [] + default_deployments: Final = [ + deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment) + ] + return default_deployments or _fail_open_deployments() @staticmethod def _get_trusted_router_request_kwargs(request_data: Mapping[str, object]) -> dict[str, object]: @@ -1250,12 +1285,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if provider is None: return None providers.append(provider) - router_allows_bedrock: Final = BedrockGuardrail._router_allows_bedrock( - request_data, - cooldown_deployments=[], - ) - if router_allows_bedrock is not None: - return api_key if router_allows_bedrock else None return ( api_key if all(provider in ("bedrock", "bedrock_converse") for provider in providers) 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 d70f7bbd94b..080c32ed317 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 @@ -349,6 +349,61 @@ def test_bedrock_guardrail_keeps_all_required_tag_matches(monkeypatch: pytest.Mo ) +def test_bedrock_guardrail_mirrors_router_fail_open_default(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": "bedrock", "tags": ["caller-only"]}, + "model_info": {"allow_fail_open": True}, + }, + { + "litellm_params": {"custom_llm_provider": "openai", "tags": ["default"]}, + "model_info": {}, + }, + ] + router._get_all_deployments.return_value = router.get_model_list.return_value + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + { + "model": "shared-alias", + "metadata": {"tags": ["&caller-only", "unmatched"], "inherited_tags": []}, + } + ) + is False + ) + + +def test_bedrock_guardrail_preserves_default_for_unknown_tag(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": ["other"]}, + "model_info": {}, + }, + { + "litellm_params": {"custom_llm_provider": "bedrock", "tags": ["default"]}, + "model_info": {}, + }, + ] + router._get_all_deployments.return_value = router.get_model_list.return_value + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "metadata": {"tags": ["unknown"]}} + ) + is True + ) + + @pytest.mark.asyncio async def test_bedrock_guardrail_async_sync_strategy_uses_unfiltered_pool(monkeypatch: pytest.MonkeyPatch): from litellm.proxy import proxy_server @@ -389,6 +444,28 @@ async def test_bedrock_guardrail_async_sync_strategy_uses_unfiltered_pool(monkey router.async_get_healthy_deployments.assert_not_awaited() +@pytest.mark.asyncio +async def test_bedrock_guardrail_async_uses_callback_filtered_pool(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.routing_strategy = "usage-based-routing-v2" + router.async_get_healthy_deployments = AsyncMock( + return_value=[{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}] + ) + monkeypatch.setattr(proxy_server, "llm_router", router) + + with patch.object(BedrockGuardrail, "_router_allows_bedrock", return_value=False) as router_allows: + assert ( + await BedrockGuardrail._async_get_bedrock_api_key( + {"model": "shared-alias", "api_key": "bedrock-key"} + ) + == "bedrock-key" + ) + + router_allows.assert_not_called() + + def test_bedrock_guardrail_matches_regex_tag_pool(monkeypatch: pytest.MonkeyPatch): from litellm.proxy import proxy_server