mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
test(auth): cover concurrent key cache misses
This commit is contained in:
parent
d5b2c586d2
commit
8ae022d0c0
1 changed files with 91 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue