From e673f4b48ede77bfb9b04fc2a1d17609cc91e41a Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Wed, 26 Aug 2026 11:43:33 +0900 Subject: [PATCH] fix: harden Bedrock pre-call routing --- .../guardrail_hooks/bedrock_guardrails.py | 99 +++++++++++----- .../test_bedrock_guardrails.py | 109 ++++++++++++++++++ 2 files changed, 179 insertions(+), 29 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 2e0a2eb68ec..82706c04c1d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -704,8 +704,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if isinstance(chain_tag_filtering, bool) else getattr(router, "enable_tag_filtering", False) ) + router_settings_override: Final[object] = request_data.get("router_settings_override") + trusted_request_tag_filtering: Final[bool] = ( + isinstance(router_settings_override, Mapping) + and router_settings_override.get("enable_tag_filtering") is True + ) tag_filtering_enabled: Final = ( - request_data.get("enable_tag_filtering") is True or effective_tag_filtering is True + trusted_request_tag_filtering or effective_tag_filtering is True ) if not tag_filtering_enabled: return deployments @@ -796,6 +801,21 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return candidate_deployments or allowed_deployments return [] + @staticmethod + def _get_trusted_router_request_kwargs(request_data: Mapping[str, object]) -> dict[str, object]: + router_kwargs: Final = dict(request_data) + router_kwargs.pop("enable_tag_filtering", None) + router_settings_override: Final[object] = request_data.get("router_settings_override") + if isinstance(router_settings_override, Mapping) and router_settings_override.get("enable_tag_filtering") is True: + router_kwargs["enable_tag_filtering"] = True + for metadata_name in ("metadata", "litellm_metadata"): + metadata: Final[object] = router_kwargs.get(metadata_name) + if isinstance(metadata, Mapping) and "routing_decision" in metadata: + router_kwargs[metadata_name] = { + key: value for key, value in metadata.items() if key != "routing_decision" + } + return router_kwargs + @staticmethod def _router_deployment_field(deployment: object, field: str) -> object | None: model_info: Final[object | None] = ( @@ -808,6 +828,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): request_data: Mapping[str, object], *, cooldown_deployments: Sequence[str] | None | object = _ROUTER_COOLDOWNS_UNSET, + apply_tag_filtering: bool = True, ) -> bool | None: model: Final[object | None] = request_data.get("model") if not isinstance(model, str): @@ -830,15 +851,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) try: resolved_team_id: Final = team_id if isinstance(team_id, str) else None - router_kwargs: Final = dict(request_data) + router_kwargs: Final = BedrockGuardrail._get_trusted_router_request_kwargs(request_data) effective_model: str = model - for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata")): - if not isinstance(metadata, Mapping): - continue - routing_decision: Final[object] = metadata.get("routing_decision") - if isinstance(routing_decision, Mapping) and isinstance(routing_decision.get("routed_model"), str): - effective_model = routing_decision["routed_model"] - break common_result: tuple[object, object] | None = None common_lookup: Final[object] = getattr(llm_router, "_common_checks_available_deployment", None) if callable(common_lookup): @@ -1042,11 +1056,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): for deployment in cooldown_filtered_deployments if BedrockGuardrail._router_deployment_field(deployment, "blocked") is not True ] - deployments = BedrockGuardrail._filter_router_deployments_by_tags( - router=llm_router, - deployments=unblocked_deployments, - request_data=request_data, - model=effective_model, + deployments = ( + BedrockGuardrail._filter_router_deployments_by_tags( + router=llm_router, + deployments=unblocked_deployments, + request_data=request_data, + model=effective_model, + ) + if apply_tag_filtering + else unblocked_deployments ) if common_result is not None and model_id_deployment_row is None: web_search_deployments: Final = filter_web_search_deployments( @@ -1126,7 +1144,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if isinstance(fallback_deployments, list) and fallback_deployments: fallback_request_data: Final = dict(request_data) fallback_request_data["model"] = default_fallback_model - return BedrockGuardrail._router_allows_bedrock(fallback_request_data) + return BedrockGuardrail._router_allows_bedrock( + fallback_request_data, + apply_tag_filtering=apply_tag_filtering, + ) return False @@ -1175,29 +1196,49 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): llm_router = None if llm_router is not None: + router_request_kwargs: Final = BedrockGuardrail._get_trusted_router_request_kwargs(request_data) + routing_strategy: object | None = getattr(llm_router, "routing_strategy", None) + routing_context: Final[object] = getattr(llm_router, "_get_routing_context", None) + if callable(routing_context): + try: + context_result: Final = routing_context( + model=request_data.get("model"), request_kwargs=router_request_kwargs + ) + if isinstance(context_result, tuple) and context_result: + routing_strategy = context_result[0] + except Exception: # noqa: BLE001 # fall back to the router default + pass + if hasattr(routing_strategy, "value"): + routing_strategy = routing_strategy.value + if isinstance(routing_strategy, str) and routing_strategy not in { + "usage-based-routing-v2", + "simple-shuffle", + "cost-based-routing", + "latency-based-routing", + "least-busy", + }: + router_allows_bedrock: Final = BedrockGuardrail._router_allows_bedrock( + request_data, + cooldown_deployments=[], + apply_tag_filtering=False, + ) + if router_allows_bedrock is not None: + return api_key if router_allows_bedrock else None + async_lookup: Final[object] = getattr(llm_router, "async_get_healthy_deployments", None) if callable(async_lookup): model: Final[object | None] = request_data.get("model") if isinstance(model, str): effective_model = model - for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata")): - if not isinstance(metadata, Mapping): - continue - routing_decision: Final[object] = metadata.get("routing_decision") - if isinstance(routing_decision, Mapping) and isinstance( - routing_decision.get("routed_model"), str - ): - effective_model = routing_decision["routed_model"] - break try: healthy_deployments: Final = await async_lookup( model=effective_model, - request_kwargs=dict(request_data), - messages=request_data.get("messages") - if isinstance(request_data.get("messages"), list) + request_kwargs=router_request_kwargs, + messages=router_request_kwargs.get("messages") + if isinstance(router_request_kwargs.get("messages"), list) else None, - input=request_data.get("input") - if isinstance(request_data.get("input"), (str, list)) + input=router_request_kwargs.get("input") + if isinstance(router_request_kwargs.get("input"), (str, list)) else None, specific_deployment=request_data.get("specific_deployment") is True, ) 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 5249fc29a6e..e14d1b99cff 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 @@ -249,6 +249,115 @@ def test_bedrock_guardrail_matches_request_tag_pool(monkeypatch: pytest.MonkeyPa ) +def test_bedrock_guardrail_ignores_client_routing_decision(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": {}}, + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._get_bedrock_api_key( + { + "model": "shared-alias", + "api_key": "bedrock-key", + "metadata": {"routing_decision": {"routed_model": "bedrock-alias"}}, + } + ) + is None + ) + + +def test_bedrock_guardrail_ignores_client_tag_filtering_override(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.enable_tag_filtering = False + 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._get_bedrock_api_key( + { + "model": "shared-alias", + "api_key": "bedrock-key", + "enable_tag_filtering": True, + "metadata": {"tags": ["fast"]}, + } + ) + is None + ) + + +def test_bedrock_guardrail_uses_trusted_tag_filtering_override(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.enable_tag_filtering = False + 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._get_bedrock_api_key( + { + "model": "shared-alias", + "api_key": "bedrock-key", + "router_settings_override": {"enable_tag_filtering": True}, + "metadata": {"tags": ["fast"]}, + } + ) + == "bedrock-key" + ) + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_async_sync_strategy_uses_unfiltered_pool(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.routing_strategy = "usage-based-routing" + router.model_names = ["shared-alias"] + router.model_group_alias = {} + router.has_model_id.return_value = False + 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": {}}, + ] + router._filter_health_check_unhealthy_deployments.side_effect = lambda healthy_deployments, **_: healthy_deployments + router._filter_deployments_by_model_access_groups.side_effect = ( + lambda **kwargs: kwargs["healthy_deployments"] + ) + router.get_model_ids.return_value = [] + router.cooldown_cache.get_active_cooldowns.return_value = [] + router.pattern_router = None + router.default_deployment = None + router.router_general_settings.pass_through_all_models = False + router.async_get_healthy_deployments = AsyncMock( + return_value=[{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}] + ) + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + await BedrockGuardrail._async_get_bedrock_api_key( + { + "model": "shared-alias", + "api_key": "bedrock-key", + "metadata": {"tags": ["fast"]}, + } + ) + is None + ) + router.async_get_healthy_deployments.assert_not_awaited() + + def test_bedrock_guardrail_matches_regex_tag_pool(monkeypatch: pytest.MonkeyPatch): from litellm.proxy import proxy_server