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>
This commit is contained in:
mrinal 2026-10-01 22:06:32 +00:00
parent 5197236444
commit 8baee043d5
2 changed files with 23 additions and 11 deletions

View file

@ -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(
*(

View file

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