diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 2fed75f2dbe..325431fd9f7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -62,7 +62,7 @@ from litellm.router_strategy.tag_based_routing import ( _split_tags, _strip_routing_prefix, ) -from litellm.router_utils.common_utils import filter_team_based_models +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 @@ -695,16 +695,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) -> list[object]: model: Final[object] = request_data.get("model") chain_tag_filtering: Final[object] = ( - _chain_tag_filtering_override(router, model, deployments) - if isinstance(model, str) - else None + _chain_tag_filtering_override(router, model, deployments) if isinstance(model, str) else None ) effective_tag_filtering: Final = ( chain_tag_filtering if isinstance(chain_tag_filtering, bool) else getattr(router, "enable_tag_filtering", False) ) - tag_filtering_enabled: Final = request_data.get("enable_tag_filtering") is True or effective_tag_filtering is True + tag_filtering_enabled: Final = ( + request_data.get("enable_tag_filtering") is True or effective_tag_filtering is True + ) if not tag_filtering_enabled: return deployments @@ -749,7 +749,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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 [] - if not positive_tags and not header_strings: + 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 + ) + if not has_positive_filter: default_deployments: Final = [ deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment) ] @@ -809,94 +816,112 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) try: resolved_team_id: Final = team_id if isinstance(team_id, str) else None - model_id_deployment: Final = ( - llm_router.get_deployment(model_id=model) if llm_router.has_model_id(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) - if ( - model_id_deployment_row is None - and isinstance(deployment_names, Sequence) - and not isinstance(deployment_names, (str, bytes)) - and model in deployment_names - and callable(specific_lookup) - ): - specific_result: Final = specific_lookup(model=model) - if isinstance(specific_result, list): - specific_deployment_rows = specific_result - raw_listed_deployments: Final = ( - [] - if model_id_deployment_row is not None - else specific_deployment_rows - if specific_deployment_rows is not None - else ( - llm_router.get_model_list( - model_name=model, - team_id=resolved_team_id, - ) - or [] + router_kwargs: Final = dict(request_data) + 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) + 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 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 ) - ) - 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 - ) - is_model_alias: Final = isinstance(model_group_aliases, Mapping) and model in model_group_aliases - if ( - model_id_deployment_row is None - and specific_deployment_rows 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=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=model) if callable(get_team_pattern_deployments) else None - ) - selected_listed_deployments: Final = ( - global_pattern_deployments - if isinstance(global_pattern_deployments, list) and global_pattern_deployments - else ( - team_pattern_deployments - if isinstance(team_pattern_deployments, list) and team_pattern_deployments - else raw_listed_deployments - ) + 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: - selected_listed_deployments: Final = raw_listed_deployments - candidate_deployments: Final[list[object]] = ( - [model_id_deployment_row] - if model_id_deployment_row is not None - else [deployment for deployment in selected_listed_deployments if isinstance(deployment, Mapping)] - ) - router_matched: Final = bool(candidate_deployments) + model_id_deployment: Final = ( + llm_router.get_deployment(model_id=model) if llm_router.has_model_id(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 model in concrete_model_names + ) + 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 [] + 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=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=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 model in deployment_names + and callable(specific_lookup) + ): + specific_result: Final = specific_lookup(model=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 [] + ) + 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) 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=dict(request_data), + request_kwargs=router_kwargs, ) ) team_filtered_deployments: Final[list[object]] = ( @@ -927,7 +952,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): healthy_deployments=access_filtered_deployments, parent_otel_span=None, ) - if callable(health_filter) + 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]] = ( @@ -942,27 +967,54 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): litellm_router_instance=llm_router, parent_otel_span=None, ) - if callable(cooldown_lookup) and callable(getattr(llm_router, "get_model_ids", None)) + if (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_ids: Final[frozenset[str]] = frozenset( deployment_id for deployment_id in cooldown_deployments if isinstance(deployment_id, str) ) - 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: Final = BedrockGuardrail._filter_router_deployments_by_tags( - router=llm_router, - deployments=unblocked_deployments, - request_data=request_data, - ) + 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, + ) + 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), + ) except Exception: # noqa: BLE001 # optional router state must not break guardrail auth return False if not deployments: @@ -975,25 +1027,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): provider = BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None return provider in ("bedrock", "bedrock_converse") if isinstance(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 != 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) - default_deployment = getattr(llm_router, "default_deployment", None) - if default_deployment is not None: + if isinstance(default_deployment, Mapping): default_params: object = ( default_deployment.get("litellm_params") if isinstance(default_deployment, Mapping) @@ -1016,6 +1051,24 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): else None ) return provider in ("bedrock", "bedrock_converse") if isinstance(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 != 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) + return False providers: list[str] = [] 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 e09a8a94f80..5249fc29a6e 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 @@ -316,6 +316,77 @@ def test_bedrock_guardrail_resolves_specific_deployment_name(monkeypatch: pytest 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() + deployments = [ + { + "litellm_params": {"custom_llm_provider": "openai", "order": 1}, + "model_info": {"id": "openai"}, + }, + { + "litellm_params": {"custom_llm_provider": "bedrock", "order": 2}, + "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 + ) + + +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