diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 7ea27498c1b..5ec482c385a 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -12,7 +12,7 @@ Endpoints for /project operations import json from collections.abc import Sequence -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Final from fastapi import APIRouter, Depends, HTTPException, Request @@ -35,6 +35,8 @@ if TYPE_CHECKING: LiteLLM_VerificationTokenActions, ) + from litellm import Router + router = APIRouter() @@ -225,7 +227,40 @@ def _project_models_missing_positive_quota( return [model for model in (models or []) if not _is_positive(rpm.get(model)) or not _is_positive(tpm.get(model))] -def _raise_on_missing_project_model_quota(data: NewProjectRequest | UpdateProjectRequest) -> None: +def _router_access_group_names(llm_router: "Router | None") -> frozenset[str]: + return frozenset(llm_router.get_model_access_groups()) if llm_router is not None else frozenset() + + +def _project_models_expanding_at_request_time( + models: Sequence[str] | None, access_group_names: frozenset[str] +) -> tuple[str, ...]: + """Entries project auth expands to many concrete models (`all-proxy-models`, `*` patterns, + access groups). The rate limiter looks quotas up by the exact requested model name, so a + quota keyed on one of these entries is never applied.""" + return tuple( + model + for model in (models or ()) + if model == SpecialModelNames.all_proxy_models.value or "*" in model or model in access_group_names + ) + + +def _raise_on_project_models_expanding_at_request_time( + models: Sequence[str] | None, access_group_names: frozenset[str] +) -> None: + expanding: Final = _project_models_expanding_at_request_time(models, access_group_names) + if not expanding: + return + raise HTTPException( + status_code=400, + detail={ + "error": f"models {list(expanding)} expand to multiple models at request time, so a per-model rpm/tpm quota cannot be enforced for them while 'enforce_project_model_quota' is enabled. List concrete model names instead." + }, + ) + + +def _raise_on_missing_project_model_quota( + data: NewProjectRequest | UpdateProjectRequest, access_group_names: frozenset[str] = frozenset() +) -> None: """Require a positive `rpm`/`tpm` quota for every model on project CREATE. `model_rpm_limit`/`model_tpm_limit` are relocated into `metadata` by the request @@ -234,6 +269,7 @@ def _raise_on_missing_project_model_quota(data: NewProjectRequest | UpdateProjec Only invoked when `general_settings.enforce_project_model_quota` is enabled (default off), so it is opt-in and does not change behavior for existing users. """ + _raise_on_project_models_expanding_at_request_time(data.models, access_group_names) metadata = data.metadata or {} missing = _project_models_missing_positive_quota( data.models, metadata.get("model_rpm_limit"), metadata.get("model_tpm_limit") @@ -248,7 +284,9 @@ def _raise_on_missing_project_model_quota(data: NewProjectRequest | UpdateProjec ) -def _raise_on_missing_project_model_quota_on_update(data: UpdateProjectRequest, existing_project: object) -> None: +def _raise_on_missing_project_model_quota_on_update( + data: UpdateProjectRequest, existing_project: object, access_group_names: frozenset[str] = frozenset() +) -> None: """Require a positive `rpm`/`tpm` quota over the RESULTING state on project UPDATE. `/project/update` replaces `models` and `metadata` when they are provided, so the @@ -260,7 +298,10 @@ def _raise_on_missing_project_model_quota_on_update(data: UpdateProjectRequest, (default off), so it is opt-in and does not change behavior for existing users. """ resulting_models = data.models if data.models is not None else (getattr(existing_project, "models", None) or []) - resulting_metadata = data.metadata if data.metadata is not None else (getattr(existing_project, "metadata", None) or {}) + resulting_metadata = ( + data.metadata if data.metadata is not None else (getattr(existing_project, "metadata", None) or {}) + ) + _raise_on_project_models_expanding_at_request_time(resulting_models, access_group_names) missing = _project_models_missing_positive_quota( resulting_models, resulting_metadata.get("model_rpm_limit"), resulting_metadata.get("model_tpm_limit") ) @@ -423,6 +464,7 @@ async def new_project( from litellm.proxy.proxy_server import ( general_settings, litellm_proxy_admin_name, + llm_router, premium_user, prisma_client, ) @@ -471,7 +513,7 @@ async def new_project( # Opt-in (default off): require rpm/tpm for every model added to the project. if general_settings.get("enforce_project_model_quota", False): - _raise_on_missing_project_model_quota(data) + _raise_on_missing_project_model_quota(data, _router_access_group_names(llm_router)) # Check if user has permission to create projects for this team # only team admins can create projects for their team @@ -614,6 +656,7 @@ async def update_project( from litellm.proxy.proxy_server import ( general_settings, litellm_proxy_admin_name, + llm_router, premium_user, prisma_client, user_api_key_cache, @@ -719,7 +762,9 @@ async def update_project( # Opt-in (default off): require rpm/tpm for every model the update would leave on the project. if general_settings.get("enforce_project_model_quota", False): - _raise_on_missing_project_model_quota_on_update(data, existing_project) + _raise_on_missing_project_model_quota_on_update( + data, existing_project, _router_access_group_names(llm_router) + ) # Prepare update data update_data = _jsonified(prisma_client, data.model_dump(exclude_none=True, exclude={"project_id"})) diff --git a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index 362b7908052..d7dd2d9adc4 100644 --- a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -4,7 +4,7 @@ from litellm._uuid import uuid from unittest import mock from dotenv import load_dotenv -from fastapi import Request +from fastapi import HTTPException, Request load_dotenv() import time @@ -1080,8 +1080,7 @@ def test_enforce_project_model_quota_all_present_passes(): model_rpm_limit={"gpt-5.5": 100}, model_tpm_limit={"gpt-5.5": 1000}, ) - # Should not raise. - _raise_on_missing_project_model_quota(data) + assert _raise_on_missing_project_model_quota(data) is None def test_enforce_project_model_quota_no_models_passes(): @@ -1091,8 +1090,7 @@ def test_enforce_project_model_quota_no_models_passes(): ) data = NewProjectRequest(team_id="test-team") - # Should not raise. - _raise_on_missing_project_model_quota(data) + assert _raise_on_missing_project_model_quota(data) is None def test_enforce_project_model_quota_zero_rejected(): @@ -1158,8 +1156,7 @@ def test_update_quota_adds_model_with_quota_passes(): model_rpm_limit={"gpt-5.5": 100}, model_tpm_limit={"gpt-5.5": 1000}, ) - # Should not raise. - _raise_on_missing_project_model_quota_on_update(data, existing) + assert _raise_on_missing_project_model_quota_on_update(data, existing) is None def test_update_quota_partial_update_keeps_existing_valid_passes(): @@ -1176,8 +1173,7 @@ def test_update_quota_partial_update_keeps_existing_valid_passes(): metadata={"model_rpm_limit": {"gpt-5.5": 100}, "model_tpm_limit": {"gpt-5.5": 1000}}, ) data = UpdateProjectRequest(project_id="p", description="unrelated change") - # Should not raise (existing quota is valid, update doesn't touch it). - _raise_on_missing_project_model_quota_on_update(data, existing) + assert _raise_on_missing_project_model_quota_on_update(data, existing) is None def test_update_quota_existing_quotaless_model_rejected(): @@ -1195,34 +1191,34 @@ def test_update_quota_existing_quotaless_model_rejected(): _raise_on_missing_project_model_quota_on_update(data, existing) -@pytest.mark.asyncio -async def test_new_project_flag_on_missing_rpm_tpm_returns_400(): - """End-to-end: with the flag on, POST /project/new rejects a model added without rpm/tpm.""" - from unittest.mock import AsyncMock, MagicMock, patch - - from fastapi import Request - +def _enforced_new_project_mocks(monkeypatch, team_models: list[str], llm_router: mock.MagicMock | None) -> None: from litellm.proxy._types import LiteLLM_TeamTable from litellm_enterprise.proxy.management_endpoints import project_endpoints as pe - team = LiteLLM_TeamTable(team_id="test-team", models=["gpt-5.5"]) - data = NewProjectRequest(team_id="test-team", models=["gpt-5.5"]) # no rpm/tpm + team = LiteLLM_TeamTable(team_id="test-team", models=team_models) + monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock.MagicMock()) + monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", llm_router) + monkeypatch.setattr(litellm.proxy.proxy_server, "general_settings", {"enforce_project_model_quota": True}) + monkeypatch.setattr(pe, "_validate_team_exists", mock.AsyncMock(return_value=team)) + monkeypatch.setattr(pe, "_check_user_permission_for_project", mock.AsyncMock(return_value=True)) - with ( - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - patch("litellm.proxy.proxy_server.premium_user", True), - patch("litellm.proxy.proxy_server.general_settings", {"enforce_project_model_quota": True}), - patch.object(pe, "_validate_team_exists", AsyncMock(return_value=team)), - patch.object(pe, "_check_user_permission_for_project", AsyncMock(return_value=True)), - ): - with pytest.raises(Exception) as exc_info: - await pe.new_project( - data=data, - http_request=Request(scope={"type": "http"}), - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234" - ), - ) + +async def _run_new_project(data: NewProjectRequest) -> None: + await new_project( + data=data, + http_request=Request(scope={"type": "http"}), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234"), + ) + + +@pytest.mark.asyncio +async def test_new_project_flag_on_missing_rpm_tpm_returns_400(monkeypatch): + """End-to-end: with the flag on, POST /project/new rejects a model added without rpm/tpm.""" + _enforced_new_project_mocks(monkeypatch, team_models=["gpt-5.5"], llm_router=None) + + with pytest.raises(Exception) as exc_info: + await _run_new_project(NewProjectRequest(team_id="test-team", models=["gpt-5.5"])) # new_project re-wraps the HTTPException, so assert on the string form. assert "rpm/tpm quota" in str(exc_info.value) @@ -1293,3 +1289,81 @@ async def test_update_project_leaves_metadata_untouched_when_no_limit_is_sent(mo await _run_project_update(project_id, description="renamed only") assert "metadata" not in _written_project_data(mock_prisma) + + +@pytest.mark.parametrize("entry", ["all-proxy-models", "*", "azure/*"]) +def test_enforce_project_model_quota_rejects_entries_that_expand_at_request_time(entry): + """A quota keyed on a wildcard entry is never applied by the limiter, so it fails loudly.""" + from litellm_enterprise.proxy.management_endpoints.project_endpoints import ( + _raise_on_missing_project_model_quota, + ) + + data = NewProjectRequest( + team_id="test-team", + models=[entry], + model_rpm_limit={entry: 10}, + model_tpm_limit={entry: 1000}, + ) + with pytest.raises(HTTPException) as exc_info: + _raise_on_missing_project_model_quota(data) + assert exc_info.value.status_code == 400 + assert entry in str(exc_info.value.detail) + assert "expand to multiple models at request time" in str(exc_info.value.detail) + + +def test_enforce_project_model_quota_rejects_access_group_only_when_router_defines_it(): + """A plain model name passes; the same name is rejected once the router reports it as an access group.""" + from litellm_enterprise.proxy.management_endpoints.project_endpoints import ( + _raise_on_missing_project_model_quota, + ) + + data = NewProjectRequest( + team_id="test-team", + models=["prod-models"], + model_rpm_limit={"prod-models": 10}, + model_tpm_limit={"prod-models": 1000}, + ) + assert _raise_on_missing_project_model_quota(data, access_group_names=frozenset()) is None + with pytest.raises(HTTPException) as exc_info: + _raise_on_missing_project_model_quota(data, access_group_names=frozenset({"prod-models"})) + assert "prod-models" in str(exc_info.value.detail) + assert "expand to multiple models at request time" in str(exc_info.value.detail) + + +def test_update_quota_rejects_wildcard_left_on_project(): + """An update that leaves a wildcard entry on the project is rejected even when it carries a quota.""" + import types + + from litellm_enterprise.proxy.management_endpoints.project_endpoints import ( + _raise_on_missing_project_model_quota_on_update, + ) + + existing = types.SimpleNamespace( + models=["all-proxy-models"], + metadata={"model_rpm_limit": {"all-proxy-models": 10}, "model_tpm_limit": {"all-proxy-models": 1000}}, + ) + data = UpdateProjectRequest(project_id="p", description="unrelated change") + with pytest.raises(HTTPException) as exc_info: + _raise_on_missing_project_model_quota_on_update(data, existing) + assert "all-proxy-models" in str(exc_info.value.detail) + assert "expand to multiple models at request time" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_new_project_flag_on_access_group_model_returns_400(monkeypatch): + """End-to-end: the router's access groups reach the check, so an access-group entry is rejected.""" + llm_router = mock.MagicMock() + llm_router.get_model_access_groups.return_value = {"prod-models": ["gpt-5.5"]} + _enforced_new_project_mocks(monkeypatch, team_models=["prod-models"], llm_router=llm_router) + data = NewProjectRequest( + team_id="test-team", + models=["prod-models"], + model_rpm_limit={"prod-models": 10}, + model_tpm_limit={"prod-models": 1000}, + ) + + with pytest.raises(Exception) as exc_info: + await _run_new_project(data) + + assert "prod-models" in str(exc_info.value) + assert "expand to multiple models at request time" in str(exc_info.value)