From 5197236444007881ed7d9f4cabe5d423cad903f5 Mon Sep 17 00:00:00 2001 From: mrinal Date: Thu, 1 Oct 2026 21:53:02 +0000 Subject: [PATCH] test(proxy): cover cross-worker cache eviction for user tpm/rpm limit updates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_internal_user_endpoints.py | 122 ++++++++++++++++++ 1 file changed, 122 insertions(+) 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 d6be455a321..7e05706b740 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -2340,6 +2340,128 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker: broadcast.assert_awaited_once_with(cache_key=saved_user.user_id) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("field", "new_limit"), + [("tpm_limit", 100), ("rpm_limit", 1)], + ids=["tpm_limit", "rpm_limit"], +) +async def test_user_rate_limit_update_reaches_cached_user_on_every_worker( + mocker: MockerFixture, field: str, new_limit: int +) -> None: + from redis.asyncio import Redis + + from litellm.proxy._types import LiteLLM_UserTable + 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 + + published: Final[list[tuple[str, str]]] = [] # mutable-ok: captures messages from the async Redis publisher + + class _RecordingRedisClient(Redis): + def __init__(self) -> None: + pass + + async def publish(self, channel: str, message: str) -> int: + published.append((channel, message)) + return 1 + + class _FakeRedisCache: + def __init__(self) -> None: + self.namespace = None + + def init_pubsub_client(self) -> object: + return _RecordingRedisClient() + + saved_user: Final = LiteLLM_UserTable(user_id="user-spruce", tpm_limit=100000, rpm_limit=1000) + updated_user: Final = saved_user.model_copy(update={field: new_limit}) + old_limit: Final = 100000 if field == "tpm_limit" else 1000 + + 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.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 + "litellm.proxy.proxy_server.prisma_client", prisma_client + ) + + handling_worker_cache: Final = UserApiKeyCache() + other_worker_cache: Final = UserApiKeyCache() + await handling_worker_cache.async_set_cache( + key=saved_user.user_id, + value=saved_user, + model_type=LiteLLM_UserTable, + ) + await other_worker_cache.async_set_cache( + key=saved_user.user_id, + value=saved_user, + model_type=LiteLLM_UserTable, + ) + mocker.patch( # test-quality-ok: exercise a real isolated cache for the endpoint's worker + "litellm.proxy.proxy_server.user_api_key_cache", handling_worker_cache + ) + mocker.patch( # test-quality-ok: inject an in-memory pub/sub client without live Redis + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", + return_value=_FakeRedisCache(), + ) + + handling_user_before: Final = await get_user_object( + user_id=saved_user.user_id, + prisma_client=prisma_client, + user_api_key_cache=handling_worker_cache, + user_id_upsert=False, + ) + other_user_before: Final = await get_user_object( + user_id=saved_user.user_id, + prisma_client=prisma_client, + user_api_key_cache=other_worker_cache, + user_id_upsert=False, + ) + assert handling_user_before is not None + assert handling_user_before.model_dump()[field] == old_limit + 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), + ) + await asyncio.sleep(0) + await asyncio.sleep(0) + + remote_subscriber: Final = AuthCacheInvalidationSubscriber( + redis_cache=_FakeRedisCache(), + user_api_key_cache=other_worker_cache, + ) + for _, message in published: + remote_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API + {"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, + user_api_key_cache=handling_worker_cache, + user_id_upsert=False, + ) + other_user_after: Final = await get_user_object( + user_id=saved_user.user_id, + prisma_client=prisma_client, + user_api_key_cache=other_worker_cache, + user_id_upsert=False, + ) + assert handling_user_after is not None + assert handling_user_after.model_dump()[field] == new_limit + assert other_user_after is not None + assert other_user_after.model_dump()[field] == new_limit, ( + "another worker still enforces the old limit; the update was never broadcast" + ) + + def test_generate_request_base_validator(): """ Test that GenerateRequestBase validator converts empty string to None for max_budget