diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 325431fd9f7..2e0a2eb68ec 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -119,6 +119,7 @@ _BEDROCK_INVOKE_GUARDRAIL_CHECKS_PATH: Final = "/guardrail-checks/invoke" # more text blocks is split across multiple messages so ALL content is scanned -- # never truncated (truncation would let a user hide content past the limit). _BEDROCK_CHECKS_MAX_CONTENT_BLOCKS: Final = 10 +_ROUTER_COOLDOWNS_UNSET: Final = object() _BEDROCK_CHECKS_KNOWN_KEYS: Final = frozenset({"contentFilter", "promptAttack", "sensitiveInformation"}) # Keys in a sensitiveInformation result that pinpoint the PII location. They are # stripped before the response is handed to standard logging / telemetry so the @@ -692,8 +693,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): router: object, deployments: list[object], request_data: Mapping[str, object], + model: str | None = None, ) -> list[object]: - model: Final[object] = request_data.get("model") + model: Final[object] = model or request_data.get("model") chain_tag_filtering: Final[object] = ( _chain_tag_filtering_override(router, model, deployments) if isinstance(model, str) else None ) @@ -784,7 +786,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): fallback_default_deployments: Final = [ deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment) ] - return fallback_default_deployments + 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 + for deployment in allowed_deployments + ): + return candidate_deployments or allowed_deployments + return [] @staticmethod def _router_deployment_field(deployment: object, field: str) -> object | None: @@ -794,7 +804,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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: + def _router_allows_bedrock( + request_data: Mapping[str, object], + *, + cooldown_deployments: Sequence[str] | None | object = _ROUTER_COOLDOWNS_UNSET, + ) -> bool | None: model: Final[object | None] = request_data.get("model") if not isinstance(model, str): return False @@ -817,11 +831,19 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): try: resolved_team_id: Final = team_id if isinstance(team_id, str) else None router_kwargs: Final = dict(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): try: - raw_common_result: Final = common_lookup(model=model, request_kwargs=router_kwargs) + raw_common_result: Final = common_lookup(model=effective_model, request_kwargs=router_kwargs) except Exception: # noqa: BLE001 # fall back for lightweight router test doubles raw_common_result = None if ( @@ -830,6 +852,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): and isinstance(raw_common_result[1], (Mapping, list)) ): common_result = raw_common_result + if isinstance(raw_common_result[0], str): + effective_model = raw_common_result[0] if common_result is not None: raw_deployments: Final[object] = common_result[1] @@ -844,7 +868,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): router_matched: Final = bool(candidate_deployments) else: model_id_deployment: Final = ( - llm_router.get_deployment(model_id=model) if llm_router.has_model_id(model) is True else None + llm_router.get_deployment(model_id=effective_model) + if llm_router.has_model_id(effective_model) is True + else None ) model_id_deployment_row: Final = ( model_id_deployment.model_dump(exclude_none=True) @@ -857,14 +883,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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 + isinstance(concrete_model_names, (list, tuple, set, frozenset)) + and effective_model in concrete_model_names + ) + is_model_alias: Final = ( + isinstance(model_group_aliases, Mapping) and effective_model in model_group_aliases ) - is_model_alias: Final = isinstance(model_group_aliases, Mapping) and model in model_group_aliases raw_listed_deployments = ( [] if model_id_deployment_row is not None else ( - llm_router.get_model_list(model_name=model, team_id=resolved_team_id) or [] + llm_router.get_model_list(model_name=effective_model, team_id=resolved_team_id) or [] if is_concrete_model or is_model_alias else [] ) @@ -875,7 +904,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): pattern_router, "get_deployments_by_pattern", None ) global_pattern_deployments: Final = ( - get_pattern_deployments(model=model) if callable(get_pattern_deployments) else None + get_pattern_deployments(model=effective_model) if callable(get_pattern_deployments) else None ) team_pattern_router: Final[object | None] = ( getattr(llm_router, "team_pattern_routers", {}).get(resolved_team_id) @@ -887,7 +916,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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 + get_team_pattern_deployments(model=effective_model) + if callable(get_team_pattern_deployments) + else None ) if isinstance(global_pattern_deployments, list) and global_pattern_deployments: raw_listed_deployments = global_pattern_deployments @@ -900,15 +931,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): elif ( isinstance(deployment_names, Sequence) and not isinstance(deployment_names, (str, bytes)) - and model in deployment_names + and effective_model in deployment_names and callable(specific_lookup) ): - specific_result: Final = specific_lookup(model=model) + specific_result: Final = specific_lookup(model=effective_model) specific_deployment_rows = specific_result if isinstance(specific_result, list) else [] raw_listed_deployments = specific_deployment_rows else: raw_listed_deployments = ( - llm_router.get_model_list(model_name=model, team_id=resolved_team_id) or [] + llm_router.get_model_list(model_name=effective_model, team_id=resolved_team_id) or [] ) candidate_deployments = ( [model_id_deployment_row] @@ -931,7 +962,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): filter_deployments: Final = getattr(llm_router, "_filter_deployments_by_model_access_groups", None) filtered_deployments: Final = ( filter_deployments( - model=model, + model=effective_model, healthy_deployments=team_filtered_deployments, request_kwargs=dict(request_data), request_team_id=resolved_team_id, @@ -960,20 +991,43 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if isinstance(health_filtered_deployments, list) else access_filtered_deployments ) + pre_call_filter: Final[object] = getattr(llm_router, "_pre_call_checks", None) + request_messages: Final[object] = request_data.get("messages") + request_input: Final[object] = request_data.get("input") + if ( + model_id_deployment_row is None + and getattr(llm_router, "enable_pre_call_checks", False) is True + and (isinstance(request_messages, list) or isinstance(request_input, (str, list))) + and callable(pre_call_filter) + ): + pre_call_deployments: Final = pre_call_filter( + model=effective_model, + healthy_deployments=healthy_deployments, + messages=request_messages if isinstance(request_messages, list) else None, + input=request_input if isinstance(request_input, (str, list)) else None, + request_kwargs=router_kwargs, + ) + if isinstance(pre_call_deployments, list): + healthy_deployments = pre_call_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 = ( + resolved_cooldown_deployments: Final = ( _get_cooldown_deployments( litellm_router_instance=llm_router, parent_otel_span=None, ) - if (common_result is None or model_id_deployment_row is None) + if cooldown_deployments is _ROUTER_COOLDOWNS_UNSET + and (common_result is None or model_id_deployment_row is None) and callable(cooldown_lookup) and callable(getattr(llm_router, "get_model_ids", None)) + else cooldown_deployments + if cooldown_deployments is not _ROUTER_COOLDOWNS_UNSET else [] ) cooldown_ids: Final[frozenset[str]] = frozenset( - deployment_id for deployment_id in cooldown_deployments if isinstance(deployment_id, str) + deployment_id + for deployment_id in (resolved_cooldown_deployments or []) + if isinstance(deployment_id, str) ) if common_result is not None and model_id_deployment_row is not None: deployments = candidate_deployments @@ -992,6 +1046,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): router=llm_router, deployments=unblocked_deployments, request_data=request_data, + model=effective_model, ) if common_result is not None and model_id_deployment_row is None: web_search_deployments: Final = filter_web_search_deployments( @@ -1007,14 +1062,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) if isinstance(plugin_deployments, list): deployments = plugin_deployments - 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), ) + deployments = litellm.utils._get_order_filtered_deployments( + deployments, + target_order=router_kwargs.pop("_target_order", None), + ) except Exception: # noqa: BLE001 # optional router state must not break guardrail auth return False if not deployments: @@ -1024,7 +1079,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if getattr(router_settings, "pass_through_all_models", False) is True: provider = request_data.get("custom_llm_provider") if not isinstance(provider, str): - provider = BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None + provider = ( + BedrockGuardrail._resolve_model_provider(effective_model) + if isinstance(effective_model, str) + else None + ) return provider in ("bedrock", "bedrock_converse") if isinstance(provider, str) else None default_deployment = getattr(llm_router, "default_deployment", None) @@ -1056,7 +1115,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): default_fallback_model: Final[object] = ( default_fallback_lookup() if callable(default_fallback_lookup) else None ) - if isinstance(default_fallback_model, str) and default_fallback_model != model: + if isinstance(default_fallback_model, str) and default_fallback_model != effective_model: fallback_deployments: Final = ( llm_router.get_model_list( model_name=default_fallback_model, @@ -1097,6 +1156,107 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): providers.append(provider) return bool(providers) and all(provider in ("bedrock", "bedrock_converse") for provider in providers) + @staticmethod + async def _async_get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None: + if not request_data: + return None + + api_key: Final[object | None] = request_data.get("api_key") + if not isinstance(api_key, str): + return None + + explicit_provider: Final[object | None] = request_data.get("custom_llm_provider") + if isinstance(explicit_provider, str) and explicit_provider not in ("bedrock", "bedrock_converse"): + return None + + try: + from litellm.proxy.proxy_server import llm_router + except ImportError: + llm_router = None + + if llm_router is not 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) + else None, + input=request_data.get("input") + if isinstance(request_data.get("input"), (str, list)) + else None, + specific_deployment=request_data.get("specific_deployment") is True, + ) + deployments: Final[list[object]] = ( + [healthy_deployments] + if isinstance(healthy_deployments, Mapping) + else healthy_deployments + if isinstance(healthy_deployments, list) + else [] + ) + 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): + return None + providers.append(provider) + return api_key if all( + provider in ("bedrock", "bedrock_converse") for provider in providers + ) else None + except Exception: # noqa: BLE001 # fall back to the sync compatibility path + pass + + 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 + + if isinstance(explicit_provider, str): + return api_key if explicit_provider in ("bedrock", "bedrock_converse") else None + + model: Final[object | None] = request_data.get("model") + model_provider: Final[str | None] = ( + BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None + ) + return api_key if model_provider in ("bedrock", "bedrock_converse") else None + @staticmethod def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None: if not request_data: @@ -1295,7 +1455,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_request_data: Final[dict] = dict( self.convert_to_bedrock_format(source=source, messages=messages, response=response) ) - api_key: Final = self._get_bedrock_api_key(request_data) + api_key: Final = await self._async_get_bedrock_api_key(request_data) if request_data: dynamic_request_body_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data) bedrock_request_data.update( @@ -2280,7 +2440,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): credentials, aws_region_name = self._load_credentials() body: Final[dict[str, Any]] = {"messages": checks_messages, "checks": self.checks} - api_key: Final = self._get_bedrock_api_key(request_data) + api_key: Final = await self._async_get_bedrock_api_key(request_data) prepared_request: Final = self._prepare_request( credentials=credentials,