From 3b5f1bb122f8d7f38d3f095a89b02e48fc8d9958 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 17:09:08 -0700 Subject: [PATCH] fix(proxy): keep file usage counters in their own store and gate batch limits on user creation --- .../internal_user_endpoints.py | 4 ++ .../management_helpers/bulk_user_creation.py | 3 + .../openai_files_endpoints/files_endpoints.py | 4 +- litellm/proxy/utils.py | 2 + .../test_internal_user_endpoints.py | 63 +++++++++++++++++++ .../test_bulk_user_creation.py | 31 +++++++++ .../test_files_endpoint.py | 41 ++++++++++++ 7 files changed, 146 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 59d8dd821d8..6bd92cd78df 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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, diff --git a/litellm/proxy/management_helpers/bulk_user_creation.py b/litellm/proxy/management_helpers/bulk_user_creation.py index ec8fd312766..ef7ba7eb772 100644 --- a/litellm/proxy/management_helpers/bulk_user_creation.py +++ b/litellm/proxy/management_helpers/bulk_user_creation.py @@ -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 diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 83b3433130e..5e52fe1641b 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1fca50e24c9..203f9098e06 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index c663e63414c..75a98ff00a1 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -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` diff --git a/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py b/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py index b5349fc2387..cf1ac61c93a 100644 --- a/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py +++ b/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py @@ -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() diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 5b7b5514a0c..01eed0228a5 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -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)