From 34ae7695901ed2162a908c639e6ebde7f4bf0d3b Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Wed, 26 Aug 2026 12:32:11 +0900 Subject: [PATCH] fix: align Bedrock router checks --- .../guardrail_hooks/bedrock_guardrails.py | 109 ++++++++---------- 1 file changed, 47 insertions(+), 62 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index e329875fb47..56cfb87dc6f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -781,11 +781,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ] if matched_deployments: return matched_deployments - fallback_default_deployments: Final = [ - deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment) - ] - if fallback_default_deployments: - return fallback_default_deployments if any( isinstance(deployment, Mapping) and (deployment.get("model_info") or {}).get("allow_fail_open") is True @@ -813,6 +808,29 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) return model_info.get(field) if isinstance(model_info, Mapping) else getattr(model_info, field, None) + @staticmethod + def _router_deployment_provider(deployment: object) -> str | None: + params: Final[object] = ( + deployment.get("litellm_params") + if isinstance(deployment, Mapping) + else getattr(deployment, "litellm_params", None) + ) + provider: Final[object] = ( + params.get("custom_llm_provider") + if isinstance(params, Mapping) + else getattr(params, "custom_llm_provider", None) + ) + if isinstance(provider, str): + return provider + deployment_model: Final[object] = ( + params.get("model") if isinstance(params, Mapping) else getattr(params, "model", None) + ) + return ( + BedrockGuardrail._resolve_model_provider(deployment_model) + if isinstance(deployment_model, str) + else None + ) + @staticmethod def _router_allows_bedrock( request_data: Mapping[str, object], @@ -824,6 +842,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if not isinstance(model, str): return False + selected_deployment: Final[object | None] = request_data.get("deployment") + if selected_deployment is not None: + selected_provider: Final[str | None] = BedrockGuardrail._router_deployment_provider(selected_deployment) + return selected_provider in ("bedrock", "bedrock_converse") if selected_provider is not None else False + try: from litellm.proxy.proxy_server import llm_router except ImportError: @@ -847,7 +870,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): common_lookup: Final[object] = getattr(llm_router, "_common_checks_available_deployment", None) if callable(common_lookup): try: - raw_common_result: Final = common_lookup(model=effective_model, request_kwargs=router_kwargs) + raw_common_result: Final = common_lookup( + model=effective_model, + request_kwargs=router_kwargs, + specific_deployment=request_data.get("specific_deployment") is True, + ) except Exception: # noqa: BLE001 # fall back for lightweight router test doubles raw_common_result = None if ( @@ -1070,14 +1097,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) if isinstance(plugin_deployments, list): deployments = plugin_deployments - deployments = litellm.utils._get_excluded_filtered_deployments( - deployments, - excluded_deployment_ids=router_kwargs.pop("_excluded_deployment_ids", None), - ) deployments = litellm.utils._get_order_filtered_deployments( deployments, target_order=router_kwargs.pop("_target_order", None), ) + deployments = litellm.utils._get_excluded_filtered_deployments( + deployments, + excluded_deployment_ids=router_kwargs.pop("_excluded_deployment_ids", None), + ) except Exception: # noqa: BLE001 # optional router state must not break guardrail auth return False if not deployments: @@ -1143,26 +1170,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): providers: list[str] = [] for deployment in deployments: - 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) - else None - ) - if not isinstance(provider, str): + provider: Final[str | None] = BedrockGuardrail._router_deployment_provider(deployment) + if provider is None: return False providers.append(provider) return bool(providers) and all(provider in ("bedrock", "bedrock_converse") for provider in providers) @@ -1188,16 +1197,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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 { @@ -1242,30 +1241,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if deployments: providers: list[str] = [] for deployment in deployments: - 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: object = ( - params.get("model") - if isinstance(params, Mapping) - else getattr(params, "model", None) - ) - provider = ( - BedrockGuardrail._resolve_model_provider(deployment_model) - if isinstance(deployment_model, str) - else None - ) - if not isinstance(provider, str): + provider: Final[str | None] = BedrockGuardrail._router_deployment_provider(deployment) + 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 ) else None