From 8baee043d5e1316c3414e4ed00ddd5ea3a518b47 Mon Sep 17 00:00:00 2001 From: mrinal Date: Thu, 1 Oct 2026 22:06:32 +0000 Subject: [PATCH] fix(proxy): broadcast user cache eviction when tpm_limit or rpm_limit changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../internal_user_endpoints.py | 6 ++-- .../test_internal_user_endpoints.py | 28 +++++++++++++------ 2 files changed, 23 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index f6fe58e25b9..c17e7acd794 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -114,7 +114,7 @@ if TYPE_CHECKING: router: Final = APIRouter() _USER_MODEL_BUDGET_ADAPTER: Final = TypeAdapter(dict[str, float | BudgetConfig]) _USER_BUDGET_CACHE_INVALIDATION_BATCH_SIZE: Final = 50 -_USER_BUDGET_CACHE_FIELDS: Final = frozenset({"max_budget", "model_max_budget"}) +_USER_LIMIT_CACHE_FIELDS: Final = frozenset({"max_budget", "model_max_budget", "tpm_limit", "rpm_limit"}) def _user_table( @@ -1619,7 +1619,7 @@ async def _update_single_user_helper( await _invalidate_user_spend_counter_if_changed(non_default_values) - if not _USER_BUDGET_CACHE_FIELDS.isdisjoint(non_default_values) or "metadata" in data_json: + if not _USER_LIMIT_CACHE_FIELDS.isdisjoint(non_default_values) or "metadata" in data_json: await evict_and_broadcast( cache_keys=(non_default_values["user_id"],), user_api_key_cache=user_api_key_cache, @@ -1977,7 +1977,7 @@ async def bulk_user_update( ), ) - if not _USER_BUDGET_CACHE_FIELDS.isdisjoint(non_default_values): + if not _USER_LIMIT_CACHE_FIELDS.isdisjoint(non_default_values): for start in range(0, len(all_users_in_db), _USER_BUDGET_CACHE_INVALIDATION_BATCH_SIZE): await asyncio.gather( *( diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 7e05706b740..cee1136c380 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -2346,8 +2346,9 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker: [("tpm_limit", 100), ("rpm_limit", 1)], ids=["tpm_limit", "rpm_limit"], ) +@pytest.mark.parametrize("all_users", [False, True], ids=["single-user", "bulk-all-users"]) async def test_user_rate_limit_update_reaches_cached_user_on_every_worker( - mocker: MockerFixture, field: str, new_limit: int + mocker: MockerFixture, field: str, new_limit: int, all_users: bool ) -> None: from redis.asyncio import Redis @@ -2355,7 +2356,8 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker( from litellm.proxy.auth.auth_checks import get_user_object from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache - from litellm.proxy.management_endpoints.internal_user_endpoints import user_update + from litellm.proxy.management_endpoints.internal_user_endpoints import bulk_user_update, user_update + from litellm.types.proxy.management_endpoints.internal_user_endpoints import BulkUpdateUserRequest published: Final[list[tuple[str, str]]] = [] # mutable-ok: captures messages from the async Redis publisher @@ -2381,6 +2383,8 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker( prisma_client: Final = mocker.MagicMock() prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user) prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=updated_user) + prisma_client.db.litellm_usertable.find_many = mocker.AsyncMock(return_value=[saved_user]) + prisma_client.db.litellm_usertable.update_many = mocker.AsyncMock(return_value=1) prisma_client.get_data = mocker.AsyncMock(return_value=saved_user) prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": updated_user}) mocker.patch( # test-quality-ok: substitute the database dependency @@ -2424,10 +2428,20 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker( assert other_user_before is not None assert other_user_before.model_dump()[field] == old_limit - await user_update( - data=UpdateUserRequest(user_id=saved_user.user_id, **{field: new_limit}), - user_api_key_dict=UserAPIKeyAuth(user_id="admin-spruce", user_role=LitellmUserRoles.PROXY_ADMIN), - ) + admin: Final = UserAPIKeyAuth(user_id="admin-spruce", user_role=LitellmUserRoles.PROXY_ADMIN) + if all_users: + await bulk_user_update( + data=BulkUpdateUserRequest(all_users=True, user_updates={field: new_limit}), + user_api_key_dict=admin, + litellm_changed_by=None, + ) + prisma_client.db.litellm_usertable.update_many.assert_awaited_once_with(where={}, data={field: new_limit}) + else: + await user_update( + data=UpdateUserRequest(user_id=saved_user.user_id, **{field: new_limit}), + user_api_key_dict=admin, + ) + assert prisma_client.update_data.call_args.kwargs["data"][field] == new_limit await asyncio.sleep(0) await asyncio.sleep(0) @@ -2440,8 +2454,6 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker( {"type": "message", "data": message} ) - assert prisma_client.update_data.call_args.kwargs["data"][field] == new_limit - handling_user_after: Final = await get_user_object( user_id=saved_user.user_id, prisma_client=prisma_client,