diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index dd76a27c80f..efcd278a5d2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -30,7 +30,10 @@ from litellm.caching import DualCache from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS from litellm.exceptions import ModifyResponseException from litellm.integrations.custom_guardrail import CustomGuardrail -from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, + redact_nested_match_and_regex_keys, +) from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import bedrock_guardrail_cost from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler from litellm.llms.base_llm.guardrail_translation.utils import ( @@ -51,6 +54,18 @@ from litellm.proxy.guardrails.anthropic_sse import ( is_raw_sse_stream, model_response_text, ) +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 from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import BedrockChecksConfigModel, GuardrailEventHooks from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage @@ -106,6 +121,19 @@ _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() + + +class _RouterCandidates(NamedTuple): + """What the router would consider for a request, before the eligibility filters.""" + + effective_model: str + common_result: tuple[object, object] | None + model_id_deployment_row: object | None + candidate_deployments: Sequence[object] + router_matched: bool + + _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 @@ -665,10 +693,803 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): return merged_messages - # NOTE: Consider moving these helpers to CustomGuardrail when the filtering - # logic becomes shared across providers. - #### CALL HOOKS - proxy only #### + @staticmethod + def _resolve_model_provider(model: str) -> str | None: + try: + _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model) + return custom_llm_provider + 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], + model: str | None = None, + ) -> list[object]: + 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 + ) + router_settings_override: Final[object] = request_data.get("router_settings_override") + override_tag_filtering: Final[object] = ( + router_settings_override.get("enable_tag_filtering") + if isinstance(router_settings_override, Mapping) + else None + ) + effective_tag_filtering: Final = ( + True + if override_tag_filtering is True + else chain_tag_filtering + if isinstance(chain_tag_filtering, bool) + else getattr(router, "enable_tag_filtering", False) + ) + tag_filtering_enabled: Final = effective_tag_filtering is True + 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 () + ) + + metadata_name: Final = get_metadata_variable_name_from_kwargs(request_data) + metadata: Final[object] = request_data.get(metadata_name) + if not isinstance(metadata, Mapping): + default_deployments: Final = [ + deployment for deployment in deployments if "default" in _deployment_tags(deployment) + ] + return default_deployments or deployments + + request_tags: Sequence[str] = tuple(_get_tags_from_request_kwargs(request_data)) + if isinstance(model, str): + 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, 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)) + ] + candidate_deployments: Final = [ + deployment for deployment in allowed_deployments if required_set.issubset(_deployment_tags(deployment)) + ] + + user_agent: Final[object] = metadata.get("user_agent") + header_strings: Final = [f"User-Agent: {user_agent}"] if isinstance(user_agent, str) and user_agent else [] + has_regex_deployments: Final = any( + isinstance(deployment, Mapping) and bool((deployment.get("litellm_params") or {}).get("tag_regex")) + for deployment in candidate_deployments + ) + 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 or _fail_open_deployments() + default_deployments: Final = [ + deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment) + ] + return default_deployments or candidate_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 candidate_deployments + if isinstance(deployment, Mapping) + and _match_deployment( + deployment=deployment, + request_tags=positive_tags, + header_strings=header_strings, + match_any=match_any, + ) + is not None + ] + if matched_deployments: + return matched_deployments + 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]: + router_kwargs: Final = dict(request_data) + router_kwargs.pop("enable_tag_filtering", None) + for metadata_name in ("metadata", "litellm_metadata"): + metadata: Final[object] = router_kwargs.get(metadata_name) + if isinstance(metadata, Mapping): + router_kwargs[metadata_name] = { + key: value for key, value in metadata.items() if key != "routing_decision" + } + router_settings_override: Final[object] = router_kwargs.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 + return router_kwargs + + @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_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_deployments_for_provider_check(deployments: Sequence[object]) -> list[object]: + if not deployments: + return [] + + def _weight(deployment: object, weight_by: str) -> object | None: + params: Final[object] = ( + deployment.get("litellm_params") + if isinstance(deployment, Mapping) + else getattr(deployment, "litellm_params", None) + ) + return params.get(weight_by) if isinstance(params, Mapping) else getattr(params, weight_by, None) + + for weight_by in ("weight", "rpm", "tpm"): + first_weight: Final[object | None] = _weight(deployments[0], weight_by) + if first_weight is None: + continue + try: + weights: Final[list[object]] = [ + 0 if (value := _weight(deployment, weight_by)) is None else value for deployment in deployments + ] + if sum(weights) > 0: + return [deployment for deployment, weight in zip(deployments, weights) if weight > 0] + except (TypeError, ValueError): + return list(deployments) + return list(deployments) + + @staticmethod + def _router_candidate_deployments( + llm_router: object, + request_data: Mapping[str, object], + router_kwargs: dict[str, object], # mutable-ok: the router pops its own routing keys off these kwargs + model: str, + resolved_team_id: str | None, + ) -> _RouterCandidates: + """Deployments the router would consider, before any of the eligibility filters.""" + effective_model: str = model + 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=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 ( + isinstance(raw_common_result, tuple) + and len(raw_common_result) == 2 + 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] + model_id_deployment_row: Final[object | None] = ( + raw_deployments if isinstance(raw_deployments, Mapping) else None + ) + candidate_deployments: Final[list[object]] = ( + [raw_deployments] + if isinstance(raw_deployments, Mapping) + else [deployment for deployment in raw_deployments if isinstance(deployment, Mapping)] + ) + router_matched: Final = bool(candidate_deployments) + else: + model_id_deployment: Final = ( + 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) + if model_id_deployment is not None and hasattr(model_id_deployment, "model_dump") + else model_id_deployment + ) + specific_deployment_rows: list[object] | None = None + deployment_names: Final[object] = getattr(llm_router, "deployment_names", None) + specific_lookup: Final[object] = getattr(llm_router, "_get_deployment_by_litellm_model", None) + 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 effective_model in concrete_model_names + ) + is_model_alias: Final = isinstance(model_group_aliases, Mapping) and effective_model in model_group_aliases + raw_listed_deployments = ( + [] + if model_id_deployment_row is not None + else ( + llm_router.get_model_list(model_name=effective_model, team_id=resolved_team_id) or [] + if is_concrete_model or is_model_alias + else [] + ) + ) + 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=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) + 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=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 + elif isinstance(team_pattern_deployments, list) and team_pattern_deployments: + raw_listed_deployments = team_pattern_deployments + else: + default_deployment = getattr(llm_router, "default_deployment", None) + if isinstance(default_deployment, Mapping): + raw_listed_deployments = [default_deployment] + elif ( + isinstance(deployment_names, Sequence) + and not isinstance(deployment_names, (str, bytes)) + and effective_model in deployment_names + and callable(specific_lookup) + ): + 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=effective_model, team_id=resolved_team_id) or [] + ) + candidate_deployments = ( + [model_id_deployment_row] + if model_id_deployment_row is not None + else [deployment for deployment in raw_listed_deployments if isinstance(deployment, Mapping)] + ) + router_matched = bool(candidate_deployments) + return _RouterCandidates( + effective_model, + common_result, + model_id_deployment_row, + candidate_deployments, + router_matched, + ) + + @staticmethod + def _filter_router_deployments( + llm_router: object, + request_data: Mapping[str, object], + router_kwargs: dict[str, object], # mutable-ok: the router pops its own routing keys off these kwargs + *, + effective_model: str, + resolved_team_id: str | None, + common_result: tuple[object, object] | None, + model_id_deployment_row: object | None, + candidate_deployments: Sequence[object], + cooldown_deployments: Sequence[str] | None | object, + apply_tag_filtering: bool, + ) -> Sequence[object]: + """The router's own eligibility chain, in the router's order. + + Order filtering runs before the weighted-failover exclusion, matching Router + so a guardrail verdict cannot disagree with the deployment actually picked. + """ + 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=router_kwargs, + ) + ) + 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=effective_model, + healthy_deployments=team_filtered_deployments, + request_kwargs=dict(request_data), + request_team_id=resolved_team_id, + ) + if callable(filter_deployments) + and isinstance(team_filtered_deployments, list) + and model_id_deployment_row is None + else None + ) + 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 (common_result is None or model_id_deployment_row is None) and 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 + ) + 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) + resolved_cooldown_deployments: Final = ( + _get_cooldown_deployments( + litellm_router_instance=llm_router, + parent_otel_span=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 (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 + else: + 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 = ( + 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( + healthy_deployments=deployments, + request_kwargs=router_kwargs, + ) + deployments = web_search_deployments if isinstance(web_search_deployments, list) else deployments + plugin_filter: Final[object] = getattr(llm_router, "_filter_by_routing_plugin_candidates", None) + if callable(plugin_filter): + plugin_deployments: Final = plugin_filter( + healthy_deployments=deployments, + request_kwargs=router_kwargs, + ) + 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), + ) + return deployments + + @staticmethod + def _router_verdict_without_deployments( + llm_router: object, + request_data: Mapping[str, object], + *, + effective_model: str, + resolved_team_id: str | None, + router_matched: bool, + apply_tag_filtering: bool, + ) -> bool | None: + """Verdict when the filters left nothing: pass-through, default deployment, or fallback.""" + if router_matched: + return False + router_settings: Final[object] = getattr(llm_router, "router_general_settings", None) + if getattr(router_settings, "pass_through_all_models", False) is True: + requested_provider: Final[object] = request_data.get("custom_llm_provider") + passthrough_provider: Final[object] = ( + requested_provider + if isinstance(requested_provider, str) + else BedrockGuardrail._resolve_model_provider(effective_model) + ) + return ( + passthrough_provider in ("bedrock", "bedrock_converse") + if isinstance(passthrough_provider, str) + else None + ) + + default_deployment: Final[object] = getattr(llm_router, "default_deployment", None) + if isinstance(default_deployment, Mapping): + default_params: Final[object] = default_deployment.get("litellm_params") + configured_provider: Final[object] = ( + default_params.get("custom_llm_provider") + if isinstance(default_params, Mapping) + else getattr(default_params, "custom_llm_provider", None) + ) + default_model: Final[object] = ( + default_params.get("model") + if isinstance(default_params, Mapping) + else getattr(default_params, "model", None) + ) + default_provider: Final[object] = ( + configured_provider + if isinstance(configured_provider, str) + else BedrockGuardrail._resolve_model_provider(default_model) + if isinstance(default_model, str) + else None + ) + return default_provider in ("bedrock", "bedrock_converse") if isinstance(default_provider, str) else None + + default_fallback_lookup: Final[object] = getattr(llm_router, "_get_first_default_fallback", None) + 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 != effective_model: + fallback_deployments: Final = ( + llm_router.get_model_list( + model_name=default_fallback_model, + team_id=resolved_team_id, + ) + or [] + ) + 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, + apply_tag_filtering=apply_tag_filtering, + ) + + return False + + @staticmethod + def _router_allows_bedrock( + 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): + return False + + selected_deployment: Final[object | None] = request_data.get("deployment") + if selected_deployment is not None and not isinstance(selected_deployment, Mapping): + 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: + return None + if llm_router is None: + return None + + team_id: Final[object | None] = next( + ( + metadata.get("user_api_key_team_id") + for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata")) + if isinstance(metadata, Mapping) and isinstance(metadata.get("user_api_key_team_id"), str) + ), + None, + ) + try: + resolved_team_id: Final = team_id if isinstance(team_id, str) else None + router_kwargs: Final = BedrockGuardrail._get_trusted_router_request_kwargs(request_data) + candidates: Final = BedrockGuardrail._router_candidate_deployments( + llm_router, request_data, router_kwargs, model, resolved_team_id + ) + deployments: Final = BedrockGuardrail._filter_router_deployments( + llm_router, + request_data, + router_kwargs, + effective_model=candidates.effective_model, + resolved_team_id=resolved_team_id, + common_result=candidates.common_result, + model_id_deployment_row=candidates.model_id_deployment_row, + candidate_deployments=candidates.candidate_deployments, + cooldown_deployments=cooldown_deployments, + apply_tag_filtering=apply_tag_filtering, + ) + except Exception: # noqa: BLE001 # optional router state must not break guardrail auth + return False + if not deployments: + return BedrockGuardrail._router_verdict_without_deployments( + llm_router, + request_data, + effective_model=candidates.effective_model, + resolved_team_id=resolved_team_id, + router_matched=candidates.router_matched, + apply_tag_filtering=apply_tag_filtering, + ) + + providers: list[str] = [] + for deployment in deployments: + 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) + + @staticmethod + async def _async_router_bedrock_verdict( + llm_router: object, + request_data: Mapping[str, object], + router_request_kwargs: dict[str, object], # mutable-ok: the pre-routing hook writes resolved params back here + routing_strategy: object | None, + ) -> bool | None: + """Whether the router's async healthy-deployment set is all-Bedrock. + + None means no verdict, so the caller falls through to the sync path. + """ + async_lookup: Final[object] = getattr(llm_router, "async_get_healthy_deployments", None) + if not callable(async_lookup): + return None + model: Final[object | None] = request_data.get("model") + if not isinstance(model, str): + return None + + effective_model = model + effective_messages: object | None = ( + router_request_kwargs.get("messages") if isinstance(router_request_kwargs.get("messages"), list) else None + ) + effective_input: object | None = ( + router_request_kwargs.get("input") if isinstance(router_request_kwargs.get("input"), (str, list)) else None + ) + try: + pre_routing_lookup: Final[object] = getattr(llm_router, "async_pre_routing_hook", None) + if callable(pre_routing_lookup): + pre_routing_result = pre_routing_lookup( + model=model, + request_kwargs=router_request_kwargs, + messages=effective_messages, + input=effective_input, + specific_deployment=request_data.get("specific_deployment") is True, + ) + if asyncio.iscoroutine(pre_routing_result): + pre_routing_result = await pre_routing_result + routed_model: Final[object] = getattr(pre_routing_result, "model", None) + if isinstance(routed_model, str): + effective_model = routed_model + routed_messages: Final[object] = getattr(pre_routing_result, "messages", None) + effective_messages = routed_messages if isinstance(routed_messages, list) else None + routed_params: Final[object] = getattr(pre_routing_result, "litellm_params", None) + if isinstance(routed_params, Mapping): + router_request_kwargs.update(routed_params) + healthy_deployments: Final = await async_lookup( + model=effective_model, + request_kwargs=router_request_kwargs, + messages=effective_messages, + input=effective_input, + specific_deployment=request_data.get("specific_deployment") is True, + ) + except Exception as exc: # noqa: BLE001 # fall back to the sync compatibility path + verbose_proxy_logger.debug("Bedrock guardrail: async router lookup failed, using the sync path: %s", exc) + return None + + deployments: list[object] = ( + [healthy_deployments] + if isinstance(healthy_deployments, Mapping) + else healthy_deployments + if isinstance(healthy_deployments, list) + else [] + ) + if routing_strategy == "simple-shuffle": + deployments = BedrockGuardrail._router_deployments_for_provider_check(deployments) + if not deployments: + return None + + providers: list[str] = [] + for deployment in deployments: + provider: Final[str | None] = BedrockGuardrail._router_deployment_provider(deployment) + if provider is None: + return False + providers.append(provider) + return 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: + router_request_kwargs: Final = BedrockGuardrail._get_trusted_router_request_kwargs(request_data) + routing_strategy: object | None = getattr(llm_router, "routing_strategy", None) + 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_verdict: Final = await BedrockGuardrail._async_router_bedrock_verdict( + llm_router, + request_data, + router_request_kwargs, + routing_strategy, + ) + if async_verdict is not None: + return api_key if async_verdict else None + + 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: + 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 + + router_allows_bedrock: Final = BedrockGuardrail._router_allows_bedrock(request_data) + 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 + def _load_credentials( self, ): @@ -843,7 +1664,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_request_data: Final[dict] = dict( self.convert_to_bedrock_format(source=source, messages=messages, response=response) ) - api_key: str | None = None + 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( @@ -853,8 +1674,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if key not in _BEDROCK_DYNAMIC_BODY_DENYLIST } ) - if request_data.get("api_key") is not None: - api_key = request_data["api_key"] event_type: Final = ( logging_event_type @@ -1830,7 +2649,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[str | None] = request_data.get("api_key") if request_data else None + api_key: Final = await self._async_get_bedrock_api_key(request_data) prepared_request: Final = self._prepare_request( credentials=credentials, diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py index 43d088268eb..4916d47d31b 100644 --- a/tests/guardrails_tests/test_bedrock_guardrails.py +++ b/tests/guardrails_tests/test_bedrock_guardrails.py @@ -10,7 +10,6 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.caching import DualCache from unittest.mock import MagicMock, AsyncMock, patch - @pytest.mark.asyncio async def test_bedrock_guardrails_pii_masking(): # Create proper mock objects 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 36b356e34d0..6965ea1fad1 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 @@ -30,6 +30,767 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( from litellm.types.utils import CallTypes, ModelResponse +def test_bedrock_guardrail_uses_active_metadata_bucket_for_team_id(): + router = MagicMock() + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}} + ] + request_data = { + "model": "team-alias", + "metadata": {"user_api_key_team_id": "legacy-team"}, + "litellm_metadata": {"user_api_key_team_id": "active-team"}, + } + + with patch("litellm.proxy.proxy_server.llm_router", router): + assert BedrockGuardrail._router_allows_bedrock(request_data) is True + + router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="active-team") + + +def test_bedrock_guardrail_uses_proxy_team_when_alternate_metadata_is_empty(): + router = MagicMock() + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}} + ] + request_data = { + "model": "team-alias", + "metadata": {"user_api_key_team_id": "proxy-team"}, + "litellm_metadata": {}, + } + + with patch("litellm.proxy.proxy_server.llm_router", router): + assert BedrockGuardrail._router_allows_bedrock(request_data) is True + + router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="proxy-team") + + +def test_bedrock_guardrail_resolves_router_model_id(): + router = MagicMock() + router.get_model_list.return_value = [] + router.has_model_id.return_value = True + deployment = MagicMock() + deployment.model_dump.return_value = { + "litellm_params": {"custom_llm_provider": "bedrock"}, + "model_info": {}, + } + router.get_deployment.return_value = deployment + + with patch("litellm.proxy.proxy_server.llm_router", router): + assert BedrockGuardrail._router_allows_bedrock({"model": "deployment-id"}) is True + + router.get_deployment.assert_called_once_with(model_id="deployment-id") + deployment.model_dump.assert_called_once_with(exclude_none=True) + + +def test_bedrock_guardrail_accepts_pass_through_bedrock_provider(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.get_model_list.return_value = [] + router.router_general_settings.pass_through_all_models = True + router.default_deployment = None + monkeypatch.setattr(proxy_server, "llm_router", router) + + request_data = { + "model": "amazon.nova-lite-v1:0", + "custom_llm_provider": "bedrock", + "api_key": "bedrock-key", + } + assert BedrockGuardrail._router_allows_bedrock(request_data) is True + assert BedrockGuardrail._get_bedrock_api_key(request_data) == "bedrock-key" + + +def test_bedrock_guardrail_accepts_bedrock_default_deployment(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.get_model_list.return_value = [] + router.router_general_settings.pass_through_all_models = False + router.default_deployment = {"litellm_params": {"custom_llm_provider": "bedrock"}} + monkeypatch.setattr(proxy_server, "llm_router", router) + + 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_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_honors_router_settings_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" + ) + + +def test_bedrock_guardrail_keeps_all_required_tag_matches(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": ["required", "default"]}, + "model_info": {}, + }, + { + "litellm_params": {"custom_llm_provider": "openai", "tags": ["required"]}, + "model_info": {}, + }, + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "metadata": {"tags": ["&required"]}} + ) + is False + ) + + +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 + + 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() + + +@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() + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_async_uses_pre_routed_model(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.routing_strategy = "usage-based-routing-v2" + router.async_pre_routing_hook = AsyncMock( + return_value=MagicMock(model="bedrock-model", messages=[{"role": "user", "content": "hi"}]) + ) + 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": "router-alias", + "api_key": "bedrock-key", + "messages": [{"role": "user", "content": "hi"}], + } + ) + == "bedrock-key" + ) + assert router.async_get_healthy_deployments.await_args.kwargs["model"] == "bedrock-model" + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_async_honors_nested_tag_override(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) + + assert ( + await BedrockGuardrail._async_get_bedrock_api_key( + { + "model": "router-alias", + "api_key": "bedrock-key", + "router_settings_override": {"enable_tag_filtering": True}, + } + ) + == "bedrock-key" + ) + assert router.async_get_healthy_deployments.await_args.kwargs["request_kwargs"]["enable_tag_filtering"] is True + + +@pytest.mark.asyncio +async def test_bedrock_guardrail_async_ignores_zero_weight_provider(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.routing_strategy = "simple-shuffle" + router.async_get_healthy_deployments = AsyncMock( + return_value=[ + {"litellm_params": {"custom_llm_provider": "openai", "weight": 0}, "model_info": {}}, + {"litellm_params": {"custom_llm_provider": "bedrock", "weight": 1}, "model_info": {}}, + ] + ) + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + await BedrockGuardrail._async_get_bedrock_api_key( + {"model": "router-alias", "api_key": "bedrock-key"} + ) + == "bedrock-key" + ) + + +def test_bedrock_guardrail_matches_regex_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"}, "model_info": {}}, + { + "litellm_params": { + "custom_llm_provider": "bedrock", + "tag_regex": [r"^User-Agent: claude-code/"], + }, + "model_info": {}, + }, + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "metadata": {"user_agent": "claude-code/1.0"}} + ) + is True + ) + + +def test_bedrock_guardrail_honors_false_chain_tag_filtering_override(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.enable_tag_filtering = True + deployments = [ + { + "litellm_params": {"custom_llm_provider": "openai", "tags": ["slow"]}, + "model_info": {"enable_tag_filtering": False}, + }, + {"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["fast"]}, "model_info": {}}, + ] + router.get_model_list.return_value = deployments + router._get_all_deployments.return_value = deployments + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "metadata": {"tags": ["fast"]}} + ) + is False + ) + + +def test_bedrock_guardrail_resolves_specific_deployment_name(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.deployment_names = ["bedrock-deployment"] + router._get_deployment_by_litellm_model.return_value = [ + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {"id": "bedrock-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": "bedrock-deployment"}) is True + router._get_deployment_by_litellm_model.assert_called_once_with(model="bedrock-deployment") + router.get_model_list.assert_not_called() + + +def test_bedrock_guardrail_ignores_user_agent_without_regex_route(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"}, "model_info": {}} + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "metadata": {"user_agent": "client/1.0"}} + ) is True + + +def test_bedrock_guardrail_applies_router_post_filters(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + # Equal order: the min-order filter runs before exclusion, so an openai row ordered + # ahead of bedrock would decide the verdict on its own and never exercise exclusion. + deployments = [ + { + "litellm_params": {"custom_llm_provider": "openai", "order": 1}, + "model_info": {"id": "openai"}, + }, + { + "litellm_params": {"custom_llm_provider": "bedrock", "order": 1}, + "model_info": {"id": "bedrock"}, + }, + ] + router._common_checks_available_deployment.return_value = ("shared-alias", deployments) + router._filter_health_check_unhealthy_deployments.return_value = deployments + router.routing_plugins = [] + router.default_deployment = None + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert BedrockGuardrail._router_allows_bedrock({"model": "shared-alias"}) is False + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "_excluded_deployment_ids": ["openai"]} + ) + is True + ) + assert ( + BedrockGuardrail._router_allows_bedrock({"model": "shared-alias", "_target_order": 2}) is False + ) + + +def test_bedrock_guardrail_applies_web_search_filter(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + deployments = [ + { + "litellm_params": {"custom_llm_provider": "openai"}, + "model_info": {"id": "openai", "supports_web_search": True}, + }, + { + "litellm_params": {"custom_llm_provider": "bedrock"}, + "model_info": {"id": "bedrock", "supports_web_search": False}, + }, + ] + router._common_checks_available_deployment.return_value = ("shared-alias", deployments) + router._filter_health_check_unhealthy_deployments.return_value = deployments + router.routing_plugins = [] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._router_allows_bedrock( + {"model": "shared-alias", "tools": [{"type": "web_search"}]} + ) + is False + ) + + +def test_bedrock_guardrail_follows_default_fallback_group(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router._get_first_default_fallback.return_value = "bedrock-fallback" + router.default_deployment = None + + def _get_model_list(model_name: str, team_id: str | None = None) -> list[dict[str, object]]: + if model_name == "bedrock-fallback": + return [{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}] + return [] + + router.get_model_list.side_effect = _get_model_list + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert BedrockGuardrail._router_allows_bedrock({"model": "unknown-model"}) is True + + +def test_bedrock_guardrail_explicit_non_bedrock_provider_wins_alias(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + router = MagicMock() + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}} + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._get_bedrock_api_key( + { + "model": "bedrock-alias", + "custom_llm_provider": "openai", + "api_key": "openai-key", + } + ) + is None + ) + + +def test_bedrock_guardrail_filters_access_group_deployments(): + 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 = [ + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}} + ] + + with patch("litellm.proxy.proxy_server.llm_router", router): + assert BedrockGuardrail._router_allows_bedrock({"model": "scoped-alias"}) is True + + +def test_bedrock_guardrail_ignores_blocked_deployments(): + router = MagicMock() + router.get_model_list.return_value = [ + {"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {"blocked": True}}, + {"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}, + ] + + with patch("litellm.proxy.proxy_server.llm_router", router): + assert BedrockGuardrail._router_allows_bedrock({"model": "alias"}) is True + + +def test_bedrock_guardrail_filters_alias_deployments_by_team(): + router = MagicMock() + router.model_group_alias = {"team-alias": "shared-group"} + # filter_team_based_models drops by model_info.id, so a row without one takes every + # other id-less row down with it. + router.get_model_list.return_value = [ + { + "litellm_params": {"custom_llm_provider": "openai"}, + "model_info": {"id": "openai", "team_id": "other-team"}, + }, + { + "litellm_params": {"custom_llm_provider": "bedrock"}, + "model_info": {"id": "bedrock", "team_id": "active-team"}, + }, + ] + + with patch("litellm.proxy.proxy_server.llm_router", router): + assert ( + BedrockGuardrail._router_allows_bedrock( + { + "model": "team-alias", + "litellm_metadata": {"user_api_key_team_id": "active-team"}, + } + ) + is True + ) + + +def test_bedrock_guardrail_honors_explicit_provider_without_router(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_router", None) + + assert ( + BedrockGuardrail._get_bedrock_api_key( + { + "model": "bedrock/anthropic.claude-3-haiku", + "custom_llm_provider": "openai", + "api_key": "openai-key", + } + ) + is None + ) + + @pytest.mark.asyncio async def test__redact_pii_matches_function(): """Test the _redact_pii_matches function directly""" @@ -1114,6 +1875,99 @@ async def test_make_apply_guardrail_request_skips_scan_without_credentials(): mock_post.assert_not_called() +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "expected_auth_prefix"), + [ + ("nvidia_nim/test-model", "AWS4-HMAC-SHA256"), + ("bedrock/test-model", "Bearer bedrock-key"), + ("amazon.nova-lite-v1:0", "Bearer bedrock-key"), + ], +) +async def test_make_apply_guardrail_request_scopes_api_key_to_bedrock_provider( + model: str, expected_auth_prefix: str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + guardrail = BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + ) + request_data = { + "model": model, + "api_key": "bedrock-key", + } + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = {"action": "NONE", "assessments": []} + + with ( + patch.object( + guardrail, + "_load_credentials", + return_value=(mock_credentials, "us-east-1"), + ), + patch.object( + guardrail.async_handler, + "post", + new=AsyncMock(return_value=mock_response), + ) as mock_post, + ): + result = await guardrail.make_bedrock_api_request( + source="INPUT", + messages=[{"role": "user", "content": "hello"}], + request_data=request_data, + ) + + assert result["action"] == "NONE" + assert mock_post.await_args.kwargs["headers"]["Authorization"].startswith( + expected_auth_prefix + ) + + +def test_bedrock_api_key_rejects_caller_provider_spoofing(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + router = MagicMock() + router.get_model_list.return_value = [ + {"litellm_params": {"model": "gpt-4o", "custom_llm_provider": "openai"}} + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._get_bedrock_api_key( + { + "model": "shared-alias", + "custom_llm_provider": "bedrock", + "api_key": "bedrock-key", + } + ) + is None + ) + + +def test_bedrock_api_key_accepts_alias_with_only_bedrock_deployments( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + + router = MagicMock() + router.get_model_list.return_value = [ + {"litellm_params": {"model": "amazon.nova-lite-v1:0", "custom_llm_provider": "bedrock"}} + ] + monkeypatch.setattr(proxy_server, "llm_router", router) + + assert ( + BedrockGuardrail._get_bedrock_api_key( + {"model": "bedrock-alias", "api_key": "bedrock-key"} + ) + == "bedrock-key" + ) + + @pytest.mark.asyncio async def test_bedrock_apply_guardrail_response_uses_OUTPUT_source(): """input_type='response' must call Bedrock with source=OUTPUT and assistant content. diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py index d842a1ee5f9..a50e4bcffa0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py @@ -395,6 +395,52 @@ async def test_request_uses_checks_path_and_body(): ] +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "expected_auth_prefix"), + [ + ("nvidia_nim/test-model", "AWS4-HMAC-SHA256"), + ("bedrock/test-model", "Bearer bedrock-key"), + ("amazon.nova-lite-v1:0", "Bearer bedrock-key"), + ], +) +async def test_request_scopes_api_key_to_bedrock_provider( + model: str, expected_auth_prefix: str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS) + request_data = { + "model": model, + "api_key": "bedrock-key", + } + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + mock_post = AsyncMock( + return_value=_mock_http_response(200, {"results": {}}) + ) + + with ( + patch.object( + g, + "_load_credentials", + return_value=(mock_credentials, "us-east-1"), + ), + patch.object(g.async_handler, "post", new=mock_post), + ): + result = await g.make_bedrock_api_request( + source="INPUT", + messages=[{"role": "user", "content": "hello"}], + request_data=request_data, + ) + + assert result == BedrockGuardrailResponse() + assert mock_post.await_args.kwargs["headers"]["Authorization"].startswith( + expected_auth_prefix + ) + + @pytest.mark.asyncio async def test_empty_messages_passes_without_api_call(): g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 45f5afef1bc..ab7a8f8acac 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -889,7 +889,10 @@ async def test_bedrock_guardrail_make_api_request_passes_api_key(): mock_response.status_code = 200 mock_response.json.return_value = {"action": "NONE", "outputs": []} - test_request_data = {"api_key": "test-api-key-789"} + test_request_data = { + "model": "bedrock/test-model", + "api_key": "test-api-key-789", + } with ( patch.object(