diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 0c31ad00200..b0a075c89a0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -667,16 +667,13 @@ 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: + except Exception: # noqa: BLE001 # provider resolution has a safe prefix fallback return model.partition("/")[0] @staticmethod @@ -709,9 +706,29 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) or [] ) + model_group_aliases: Final = getattr(llm_router, "model_group_alias", None) + team_filtered_deployments: Final = ( + [ + deployment + for deployment in listed_deployments + if ( + ( + deployment.get("model_info") + if isinstance(deployment, Mapping) + else getattr(deployment, "model_info", None) + ) + or {} + ).get("team_id") + in (None, resolved_team_id) + ] + if resolved_team_id is not None + and isinstance(model_group_aliases, Mapping) + and model in model_group_aliases + else listed_deployments + ) model_id_deployment: Final = ( llm_router.get_deployment(model_id=model) - if not listed_deployments and llm_router.has_model_id(model) is True + if not team_filtered_deployments and llm_router.has_model_id(model) is True else None ) model_id_deployment_row: Final = ( @@ -720,7 +737,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): else model_id_deployment ) candidate_deployments: Final = ( - [model_id_deployment_row] if model_id_deployment_row is not None else listed_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) @@ -737,7 +756,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): deployments: Final = ( filtered_deployments if isinstance(filtered_deployments, list) else candidate_deployments ) - except Exception: + except Exception: # noqa: BLE001 # optional router state must not break guardrail auth return False active_deployments = [] for deployment in deployments: @@ -793,6 +812,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): 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 + model: Final[object | None] = request_data.get("model") model_provider: Final[str | None] = ( BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py index 7e670dfb366..4916d47d31b 100644 --- a/tests/guardrails_tests/test_bedrock_guardrails.py +++ b/tests/guardrails_tests/test_bedrock_guardrails.py @@ -10,84 +10,6 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.caching import DualCache from unittest.mock import MagicMock, AsyncMock, patch - -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_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 - - @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 bf8b48625e8..13bbb15ef30 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,126 @@ 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_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"} + router.get_model_list.return_value = [ + { + "litellm_params": {"custom_llm_provider": "openai"}, + "model_info": {"team_id": "other-team"}, + }, + { + "litellm_params": {"custom_llm_provider": "bedrock"}, + "model_info": {"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"""