test(auth): cover concurrent key cache misses

This commit is contained in:
Pin Zhu 2026-09-28 19:23:19 +08:00 • committed by GitHub
parent d5b2c586d2
commit 8ae022d0c0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

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