From 8ae022d0c08a4133c07504d8d38ba424dc683708 Mon Sep 17 00:00:00 2001 From: Pin Zhu <113764596+47Elysia@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:23:19 +0800 Subject: [PATCH] test(auth): cover concurrent key cache misses --- .../proxy/auth/test_auth_checks.py | 91 +++++++++++++++++++ 1 file changed, 91 insertions(+) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index f014e9c26d1..220478423ae 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -40,6 +40,7 @@ from litellm.types.agents import AgentCaller from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, _cache_management_object, + _delete_cache_key_object, _can_object_call_model, _can_object_call_vector_stores, _check_agent_access_group_model_access, @@ -611,6 +612,96 @@ async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_e assert mock_prisma_client.get_data.await_count == 1 +@pytest.mark.asyncio +async def test_get_key_object_coalesces_parallel_cache_misses(): + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get_data(*args: object, **kwargs: object) -> UserAPIKeyAuth: + started.set() + await release.wait() + return UserAPIKeyAuth(token="hashed-token") + + get_data_mock: Final = AsyncMock(side_effect=slow_get_data) + prisma: Final = MagicMock() + prisma.get_data = get_data_mock + cache: Final = UserApiKeyCache() + + first: Final = asyncio.create_task( + get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + ) + await started.wait() + second: Final = asyncio.create_task( + get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + ) + await asyncio.sleep(0) + release.set() + + results: Final = await asyncio.gather(first, second) + + assert [result.token for result in results] == ["hashed-token", "hashed-token"] + assert get_data_mock.await_count == 1 + + +@pytest.mark.asyncio +async def test_delete_key_cache_waits_for_inflight_load_before_eviction(): + started: Final = asyncio.Event() + release: Final = asyncio.Event() + rows: Final = iter(("old-model", "new-model")) + + async def get_data(*args: object, **kwargs: object) -> UserAPIKeyAuth: + model: Final = next(rows) + if model == "old-model": + started.set() + await release.wait() + return UserAPIKeyAuth(token="hashed-token", models=[model]) + + get_data_mock: Final = AsyncMock(side_effect=get_data) + prisma: Final = MagicMock() + prisma.get_data = get_data_mock + cache: Final = UserApiKeyCache() + stale_load: Final = asyncio.create_task( + get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + ) + await started.wait() + + invalidation: Final = asyncio.create_task( + _delete_cache_key_object( + hashed_token="hashed-token", + user_api_key_cache=cache, + proxy_logging_obj=None, + ) + ) + await asyncio.sleep(0) + assert not invalidation.done() + + release.set() + await invalidation + stale: Final = await stale_load + assert stale.model_dump(mode="json")["models"] == ["old-model"] + assert await cache.async_get_cache(key="hashed-token", model_type=UserAPIKeyAuth) is None + + fresh: Final = await get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + assert fresh.model_dump(mode="json")["models"] == ["new-model"] + assert get_data_mock.await_count == 2 + + class _InFlightCountingPrisma: def __init__(self) -> None: self.in_flight = 0