diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index e4702364a3..4f3e51a2a3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2753,6 +2753,8 @@ class NewProjectRequest(LiteLLM_BudgetTable): budget_id: Optional[str] = None metadata: Optional[dict] = None tags: Optional[List[str]] = None + guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None models: List[str] = [] model_rpm_limit: Optional[dict] = None model_tpm_limit: Optional[dict] = None @@ -2785,6 +2787,8 @@ class UpdateProjectRequest(LiteLLM_BudgetTable): team_id: Optional[str] = None metadata: Optional[dict] = None tags: Optional[List[str]] = None + guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None models: Optional[List[str]] = None model_rpm_limit: Optional[dict] = None model_tpm_limit: Optional[dict] = None diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index ba9577f35d..fcb8b6db80 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1472,17 +1472,19 @@ def _add_guardrails_from_key_or_team_metadata( team_metadata: Optional[dict], data: dict, metadata_variable_name: str, + project_metadata: Optional[dict] = None, ) -> None: """ - Helper add guardrails from key or team metadata to request data + Helper add guardrails from key, team, or project metadata to request data - Key guardrails are set first, then team guardrails are appended (without duplicates). + Key guardrails are set first, then team and project guardrails are appended (without duplicates). Args: key_metadata: The key metadata dictionary to check for guardrails team_metadata: The team metadata dictionary to check for guardrails data: The request data to update metadata_variable_name: The name of the metadata field in data + project_metadata: The project metadata dictionary to check for guardrails """ from litellm.proxy.utils import _premium_user_check @@ -1508,6 +1510,15 @@ def _add_guardrails_from_key_or_team_metadata( _premium_user_check() combined_guardrails.update(team_metadata["guardrails"]) + # Add project-level guardrails (set automatically handles duplicates) + if project_metadata and "guardrails" in project_metadata: + if ( + isinstance(project_metadata["guardrails"], list) + and len(project_metadata["guardrails"]) > 0 + ): + _premium_user_check() + combined_guardrails.update(project_metadata["guardrails"]) + # Set combined guardrails in metadata as list if combined_guardrails: data[metadata_variable_name]["guardrails"] = list(combined_guardrails) @@ -1518,12 +1529,13 @@ def _add_guardrails_from_policies_in_metadata( team_metadata: Optional[dict], data: dict, metadata_variable_name: str, + project_metadata: Optional[dict] = None, ) -> None: """ - Helper to resolve guardrails from policies attached to key/team metadata. + Helper to resolve guardrails from policies attached to key/team/project metadata. This function: - 1. Gets policy names from key and team metadata + 1. Gets policy names from key, team, and project metadata 2. Resolves guardrails from those policies (including inheritance) 3. Adds resolved guardrails to request metadata @@ -1532,6 +1544,7 @@ def _add_guardrails_from_policies_in_metadata( team_metadata: The team metadata dictionary to check for policies data: The request data to update metadata_variable_name: The name of the metadata field in data + project_metadata: The project metadata dictionary to check for policies """ from litellm._logging import verbose_proxy_logger from litellm.proxy.policy_engine.policy_registry import get_policy_registry @@ -1560,6 +1573,15 @@ def _add_guardrails_from_policies_in_metadata( _premium_user_check() policy_names.update(team_metadata["policies"]) + # Add project-level policies + if project_metadata and "policies" in project_metadata: + if ( + isinstance(project_metadata["policies"], list) + and len(project_metadata["policies"]) > 0 + ): + _premium_user_check() + policy_names.update(project_metadata["policies"]) + if not policy_names: return @@ -1641,6 +1663,7 @@ async def move_guardrails_to_metadata( # Early-out: skip all guardrails processing when nothing is configured key_metadata = user_api_key_dict.metadata team_metadata = user_api_key_dict.team_metadata + project_metadata = user_api_key_dict.project_metadata or {} has_key_config = key_metadata and ( "guardrails" in key_metadata or "policies" in key_metadata @@ -1648,12 +1671,15 @@ async def move_guardrails_to_metadata( has_team_config = team_metadata and ( "guardrails" in team_metadata or "policies" in team_metadata ) + has_project_config = project_metadata and ( + "guardrails" in project_metadata or "policies" in project_metadata + ) has_request_config = ( "guardrails" in data or "guardrail_config" in data or "policies" in data ) # Only check policy engine if no local config (avoid import + registry lookup) - if not (has_key_config or has_team_config or has_request_config): + if not (has_key_config or has_team_config or has_project_config or has_request_config): from litellm.proxy.policy_engine.policy_registry import get_policy_registry if not get_policy_registry().is_initialized(): @@ -1661,20 +1687,22 @@ async def move_guardrails_to_metadata( data.pop("policies", None) return - # Check key-level guardrails + # Check key/team/project-level guardrails _add_guardrails_from_key_or_team_metadata( key_metadata=user_api_key_dict.metadata, team_metadata=user_api_key_dict.team_metadata, + project_metadata=project_metadata, data=data, metadata_variable_name=_metadata_variable_name, ) ######################################################################################### - # Add guardrails from policies attached to key/team metadata + # Add guardrails from policies attached to key/team/project metadata ######################################################################################### _add_guardrails_from_policies_in_metadata( key_metadata=user_api_key_dict.metadata, team_metadata=user_api_key_dict.team_metadata, + project_metadata=project_metadata, data=data, metadata_variable_name=_metadata_variable_name, ) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index bc13cea939..04af5cd008 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -1363,6 +1363,101 @@ async def test_request_guardrails_do_not_override_key_guardrails(): assert len(requested_guardrails) == 1 +@pytest.mark.asyncio +async def test_project_guardrails_merge_with_key_and_team(): + """ + Test that project guardrails are merged with key and team guardrails (union semantics). + All three levels should contribute to the final guardrails list without duplicates. + """ + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + metadata={"guardrails": ["key-guardrail-1"]}, + team_metadata={"guardrails": ["team-guardrail-1", "key-guardrail-1"]}, + project_metadata={"guardrails": ["project-guardrail-1", "team-guardrail-1"]}, + ) + + with patch("litellm.proxy.utils._premium_user_check"): + updated_data = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + metadata = updated_data.get("metadata", {}) + guardrails = metadata.get("guardrails", []) + + # All three sources contribute + assert "key-guardrail-1" in guardrails + assert "team-guardrail-1" in guardrails + assert "project-guardrail-1" in guardrails + # No duplicates + assert guardrails.count("key-guardrail-1") == 1 + assert guardrails.count("team-guardrail-1") == 1 + + +@pytest.mark.asyncio +async def test_project_guardrails_only(): + """ + Test that project guardrails work when key and team have no guardrails configured. + """ + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + metadata={}, + team_metadata={}, + project_metadata={"guardrails": ["project-guardrail-1", "project-guardrail-2"]}, + ) + + with patch("litellm.proxy.utils._premium_user_check"): + updated_data = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + metadata = updated_data.get("metadata", {}) + guardrails = metadata.get("guardrails", []) + + assert "project-guardrail-1" in guardrails + assert "project-guardrail-2" in guardrails + assert len(guardrails) == 2 + + def test_update_model_if_key_alias_exists(): """ Test that _update_model_if_key_alias_exists properly updates the model when a key alias exists.