diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 8ade4a011f7..2c05bfce601 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -677,25 +677,72 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): except Exception: return model.partition("/")[0] + @staticmethod + def _router_allows_bedrock(request_data: Mapping[str, object]) -> bool | None: + model: Final[object | None] = request_data.get("model") + if not isinstance(model, str): + return False + + try: + from litellm.proxy.proxy_server import llm_router + except ImportError: + return None + if llm_router is None: + return None + + metadata: Final = request_data.get("metadata") + litellm_metadata: Final = request_data.get("litellm_metadata") + team_id: Final[object | None] = ( + metadata.get("user_api_key_team_id") + if isinstance(metadata, Mapping) + else litellm_metadata.get("user_api_key_team_id") + if isinstance(litellm_metadata, Mapping) + else None + ) + try: + deployments: Final = llm_router.get_model_list( + model_name=model, + team_id=team_id if isinstance(team_id, str) else None, + ) or [] + except Exception: + return False + if not deployments: + return False + + providers: list[str] = [] + for deployment in deployments: + params: object = deployment.get("litellm_params") if isinstance(deployment, Mapping) else None + provider: object = params.get("custom_llm_provider") if isinstance(params, Mapping) else None + if not isinstance(provider, str) and isinstance(params, Mapping): + deployment_model: Final[object | None] = params.get("model") + provider = ( + BedrockGuardrail._resolve_model_provider(deployment_model) + if isinstance(deployment_model, str) + else None + ) + if not isinstance(provider, str): + return False + providers.append(provider) + return bool(providers) and all(provider in ("bedrock", "bedrock_converse") for provider in providers) + @staticmethod def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None: if not request_data: return None - top_level_provider: Final[object | None] = request_data.get("custom_llm_provider") - request_provider: Final[str | None] = ( - top_level_provider if isinstance(top_level_provider, str) else None - ) + api_key: Final[object | None] = request_data.get("api_key") + if not isinstance(api_key, str): + 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 + model: Final[object | None] = request_data.get("model") model_provider: Final[str | None] = ( BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None ) - custom_llm_provider: Final[str | None] = request_provider or model_provider - if custom_llm_provider not in ("bedrock", "bedrock_converse"): - return None - - api_key: Final[object | None] = request_data.get("api_key") - return api_key if isinstance(api_key, str) else None + return api_key if model_provider in ("bedrock", "bedrock_converse") else None def _load_credentials( self, 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 a1dbbaa5953..bf8b48625e8 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 @@ -1167,6 +1167,46 @@ async def test_make_apply_guardrail_request_scopes_api_key_to_bedrock_provider( ) +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.