mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
d7e6154863
commit
5197236444
1 changed files with 122 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue