fix(proxy): keep file usage counters in their own store and gate batch limits on user creation

This commit is contained in:
mateo-berri 2026-09-28 17:09:08 -07:00
parent 64c22e0cdd
commit 3b5f1bb122
7 changed files with 146 additions and 2 deletions

View file

@ -36,6 +36,7 @@ from litellm.proxy.auth.auth_checks import (
get_team_object,
get_user_object,
)
from litellm.proxy.auth.auth_utils import enforce_batch_limits_are_admin_only
from litellm.proxy.auth.password_policy import (
validate_password_not_breached,
validate_password_policy,
@ -581,6 +582,9 @@ async def new_user(
detail=f"Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). Attempted to create user with role: {data.user_role}. Your role: {user_api_key_dict.user_role}",
)
if data.auto_create_key and isinstance(user_api_key_dict, UserAPIKeyAuth):
enforce_batch_limits_are_admin_only(data, None, user_api_key_dict, "key")
_check_permissions_caller_permission(
data=data,
user_api_key_dict=user_api_key_dict,

View file

@ -29,6 +29,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state
from litellm.proxy.auth.auth_utils import enforce_batch_limits_are_admin_only
from litellm.proxy.auth.litellm_license import LicenseCheck
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
@ -215,6 +216,8 @@ def _row_error(item: BulkNewUserItem, user_api_key_dict: UserAPIKeyAuth) -> str
try:
validate_budget_duration(item.budget_duration)
_check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict)
if item.auto_create_key:
enforce_batch_limits_are_admin_only(item, None, user_api_key_dict, "key")
except Exception as exc: # noqa: BLE001 # any validation failure is reported on this row only
return _error_message(exc)
return None

View file

@ -682,7 +682,7 @@ async def create_file(
if batch_file_failure is not None:
raise_batch_file_validation_failure(batch_file_failure)
await enforce_batch_file_upload_limit(
proxy_logging_obj.internal_usage_cache, user_api_key_dict, general_settings
proxy_logging_obj.file_usage_cache, user_api_key_dict, general_settings
)
data = {"passthrough": True} if passthrough else {}
@ -988,7 +988,7 @@ async def get_file_content(
managed_files_obj=proxy_logging_obj.get_proxy_hook("managed_files"),
)
await enforce_file_download_limit(
proxy_logging_obj.internal_usage_cache, user_api_key_dict, general_settings, file_id
proxy_logging_obj.file_usage_cache, user_api_key_dict, general_settings, file_id
)
# Include original request and headers in the data

View file

@ -1206,6 +1206,7 @@ class ProxyLogging:
self.internal_usage_cache: InternalUsageCache = InternalUsageCache(
dual_cache=DualCache(default_in_memory_ttl=1) # ping redis cache every 1s
)
self.file_usage_cache: Final = InternalUsageCache(dual_cache=DualCache())
self.max_parallel_request_limiter = _PROXY_MaxParallelRequestsHandler(self.internal_usage_cache)
self.cache_control_check = _PROXY_CacheControlCheck()
self.alerting: list[str] | None = None
@ -1347,6 +1348,7 @@ class ProxyLogging:
if redis_cache is not None:
self.internal_usage_cache.dual_cache.redis_cache = redis_cache
self.file_usage_cache.dual_cache.redis_cache = redis_cache
self.db_spend_update_writer.redis_update_buffer.redis_cache = redis_cache
self.db_spend_update_writer.pod_lock_manager.redis_cache = redis_cache

View file

@ -1241,6 +1241,69 @@ async def test_new_user_admin_can_set_permissions(mocker):
assert result is not None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"limit_key",
["max_batch_file_records", "max_batch_file_uploads_per_day", "max_file_downloads_per_minute"],
)
async def test_new_user_only_proxy_admin_sets_batch_limits_on_the_created_key(mocker, limit_key):
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mock_prisma_client = mocker.MagicMock()
async def mock_count(*args, **kwargs):
return 5
mock_prisma_client.db.litellm_usertable.count = mock_count
async def mock_check(*_args, **_kwargs):
return None
mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email",
mock_check,
)
mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id",
mock_check,
)
mock_license_check = mocker.MagicMock()
mock_license_check.is_over_limit.return_value = False
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check)
created_with: list[dict[str, object]] = []
async def stub_helper(**kwargs):
created_with.append(kwargs)
return {"user_id": "alice", "key": "sk-alice", "expires": None}
mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints.generate_key_helper_fn",
stub_helper,
)
org_admin = UserAPIKeyAuth(user_id="org-admin", user_role=LitellmUserRoles.ORG_ADMIN)
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
def request(auto_create_key: bool) -> NewUserRequest:
return NewUserRequest(
user_email="alice@example.com",
user_role=LitellmUserRoles.INTERNAL_USER,
metadata={limit_key: 1000},
auto_create_key=auto_create_key,
)
with pytest.raises(ProxyException) as exc_info:
await new_user(data=request(auto_create_key=True), user_api_key_dict=org_admin)
assert str(exc_info.value.code) == "403"
assert f"Only proxy admins can set {limit_key} on a key" in str(exc_info.value.message)
assert created_with == []
await new_user(data=request(auto_create_key=False), user_api_key_dict=org_admin)
await new_user(data=request(auto_create_key=True), user_api_key_dict=admin)
assert [call["metadata"] for call in created_with] == [{limit_key: 1000}, {limit_key: 1000}]
@pytest.mark.asyncio
async def test_update_single_user_non_admin_permissions_rejected(mocker):
"""`_update_single_user_helper` rejects a non-admin when `permissions`

View file

@ -390,6 +390,37 @@ async def test_non_admin_cannot_create_admin_users_but_other_rows_proceed():
assert set(prisma.db.litellm_usertable.rows) == {"u2"}
@pytest.mark.asyncio
async def test_only_proxy_admin_sets_batch_limits_on_a_created_key():
limits = {"max_file_downloads_per_minute": 1000}
calls: list[dict[str, object]] = []
async def generate_key(**kwargs: object) -> dict[str, object]:
calls.append(kwargs)
return {"token": f"sk-{kwargs['user_id']}"}
rows = [
{"user_id": "u1", "auto_create_key": True, "metadata": limits},
{"user_id": "u2", "auto_create_key": False, "metadata": limits},
{"user_id": "u3", "auto_create_key": True},
]
prisma = _FakePrisma()
response = await _run(prisma, rows, caller=INTERNAL, generate_key=generate_key)
assert [r.success for r in response.data] == [False, True, True]
assert "Only proxy admins can set max_file_downloads_per_minute on a key" in (response.data[0].error or "")
assert set(prisma.db.litellm_usertable.rows) == {"u2", "u3"}
assert [call["user_id"] for call in calls] == ["u3"]
admin_prisma = _FakePrisma()
admin_response = await _run(admin_prisma, rows[:1], caller=ADMIN, generate_key=generate_key)
assert [r.key for r in admin_response.data] == ["sk-u1"]
assert calls[-1]["user_id"] == "u1"
assert calls[-1]["metadata"] == limits
@pytest.mark.asyncio
async def test_license_is_checked_once_against_the_whole_batch():
prisma = _FakePrisma()

View file

@ -4351,6 +4351,47 @@ def test_get_file_content_over_per_minute_download_limit_gets_429_per_file(monke
assert provider_calls == ["file-out", "file-out", "file-other"]
def test_download_counters_for_many_files_do_not_evict_rate_limit_state(monkeypatch, llm_router: Router):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
proxy_logging = setup_proxy_logging_object(monkeypatch, llm_router)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setitem(ps.general_settings, "max_file_downloads_per_minute", 1000)
_pin_file_usage_clock(monkeypatch)
async def _mock_afile_content(**kwargs):
async def _stream():
yield b"output"
return FileContentStreamingResult(stream_iterator=_stream(), headers={"content-length": "6"})
monkeypatch.setattr(litellm, "afile_content", _mock_afile_content)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
AsyncMock(return_value=(False, None, None, None)),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="hashed-caller-key", user_role=LitellmUserRoles.INTERNAL_USER, user_id="test-user"
)
rate_limit_cache: Final = proxy_logging.internal_usage_cache.dual_cache
rate_limit_window_key: Final = "{api_key:hashed-caller-key}:window"
rate_limit_cache.set_cache(key=rate_limit_window_key, value=12345, local_only=True, ttl=60)
distinct_files: Final = rate_limit_cache.in_memory_cache.max_size_in_memory + 1
try:
statuses = [
client.get(f"/v1/files/file-{index}/content", headers={"Authorization": "Bearer test-key"}).status_code
for index in range(distinct_files)
]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
assert set(statuses) == {200}
assert rate_limit_cache.get_cache(key=rate_limit_window_key, local_only=True) == 12345
def test_create_file_batch_wrong_extension_rejected_before_forwarding(monkeypatch, llm_router: Router):
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)