diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index b0a075c89a0..5a5b0307fa9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -737,9 +737,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): else model_id_deployment ) candidate_deployments: Final = ( - [model_id_deployment_row] - if model_id_deployment_row is not None - else team_filtered_deployments + [model_id_deployment_row] if model_id_deployment_row is not None else team_filtered_deployments ) filter_deployments: Final = getattr(llm_router, "_filter_deployments_by_model_access_groups", None) @@ -771,6 +769,37 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): if blocked is not True: active_deployments.append(deployment) if not active_deployments: + router_settings = getattr(llm_router, "router_general_settings", None) + 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 + return provider in ("bedrock", "bedrock_converse") if isinstance(provider, str) else None + + default_deployment = getattr(llm_router, "default_deployment", None) + if default_deployment is not None: + default_params: object = ( + default_deployment.get("litellm_params") + if isinstance(default_deployment, Mapping) + else getattr(default_deployment, "litellm_params", None) + ) + provider: object = ( + default_params.get("custom_llm_provider") + if isinstance(default_params, Mapping) + else getattr(default_params, "custom_llm_provider", None) + ) + if not isinstance(provider, str): + default_model: object = ( + default_params.get("model") + if isinstance(default_params, Mapping) + else getattr(default_params, "model", None) + ) + provider = ( + BedrockGuardrail._resolve_model_provider(default_model) + if isinstance(default_model, str) + else None + ) + return provider in ("bedrock", "bedrock_converse") if isinstance(provider, str) else None return False providers: list[str] = [] @@ -808,11 +837,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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 - explicit_provider: Final[object | None] = request_data.get("custom_llm_provider") if isinstance(explicit_provider, str): return api_key if explicit_provider in ("bedrock", "bedrock_converse") else None 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 13bbb15ef30..0340f83c9d3 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 @@ -82,6 +82,57 @@ def test_bedrock_guardrail_resolves_router_model_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_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 = [