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>
This commit is contained in:
mrinal 2026-10-01 21:53:02 +00:00
parent d7e6154863
commit 5197236444

View file

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