feat(proxy): reject wildcard project models under enforce_project_model_quota

Project auth expands all-proxy-models, * patterns, and access-group names
to many concrete models, but the rate limiter looks quotas up by the exact
requested model name, so a quota keyed on one of those entries is never
applied. Fail loudly with a 400 instead of storing an unenforceable quota
This commit is contained in:
ryan-crabbe-berri 2026-08-27 20:49:22 -07:00
parent 2e06762fc4
commit bc127a0b82
2 changed files with 158 additions and 39 deletions

View file

@ -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"}))

View file

@ -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)