From 6c94b8a627348e0b4fb170fee343bc037d37ee33 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Mon, 8 Jun 2026 14:41:58 -0700 Subject: [PATCH] perf(identity): batch generation-counter reads via async_batch_get_cache --- litellm/identity/cache.py | 35 ++++++++++--------- tests/test_litellm/identity/test_cache.py | 42 +++++++++++++++++++++++ 2 files changed, 60 insertions(+), 17 deletions(-) diff --git a/litellm/identity/cache.py b/litellm/identity/cache.py index 8bfc770c3c2..f9804d241df 100644 --- a/litellm/identity/cache.py +++ b/litellm/identity/cache.py @@ -54,9 +54,7 @@ def _generation_attr_key(scope: str) -> str: return f"identity_cache_generation_{scope}" -def _attach_generations( - uak: "UserAPIKeyAuth", generations: dict -) -> None: +def _attach_generations(uak: "UserAPIKeyAuth", generations: dict) -> None: """Stash the generation counters this entry was minted under. Stored on the model's ``metadata`` so the value survives Pydantic @@ -91,12 +89,13 @@ class IdentityCache: self._cache = dual_cache self._ttl_seconds = ttl_seconds - @traced("identity.cache.get", role=SpanRole.DB_CALL, - attrs=lambda result: { - "identity.cache.layer": ( - "miss" if result is None else "memory_or_redis" - ), - }) + @traced( + "identity.cache.get", + role=SpanRole.DB_CALL, + attrs=lambda result: { + "identity.cache.layer": ("miss" if result is None else "memory_or_redis"), + }, + ) async def get(self, token_hash: str) -> Optional["UserAPIKeyAuth"]: from litellm.proxy._types import UserAPIKeyAuth @@ -112,9 +111,7 @@ class IdentityCache: return cached @traced("identity.cache.set", role=SpanRole.DB_CALL) - async def set( - self, token_hash: str, uak: "UserAPIKeyAuth" - ) -> None: + async def set(self, token_hash: str, uak: "UserAPIKeyAuth") -> None: generations = await self._snapshot_generations_for(uak) _attach_generations(uak, generations) await self._cache.async_set_cache( @@ -133,9 +130,7 @@ class IdentityCache: current = await self._snapshot_generations_for(uak) return any(stored.get(k) != current.get(k) for k in stored) - async def _snapshot_generations_for( - self, uak: "UserAPIKeyAuth" - ) -> dict: + async def _snapshot_generations_for(self, uak: "UserAPIKeyAuth") -> dict: scopes: list[tuple[str, str]] = [] if uak.team_id: scopes.append(("team", team_generation_key(uak.team_id))) @@ -143,9 +138,15 @@ class IdentityCache: scopes.append(("user", user_generation_key(uak.user_id))) if uak.org_id: scopes.append(("org", org_generation_key(uak.org_id))) + if not scopes: + return {} + + values = await self._cache.async_batch_get_cache( + keys=[key for _, key in scopes] + ) return { - scope: await self._cache.async_get_cache(key=key) or 0 - for scope, key in scopes + scope: (values[index] or 0) if index < len(values) else 0 + for index, (scope, _) in enumerate(scopes) } async def bump_generation(self, scope_key: str) -> None: diff --git a/tests/test_litellm/identity/test_cache.py b/tests/test_litellm/identity/test_cache.py index 7ecd5bb509a..e7f03df9fce 100644 --- a/tests/test_litellm/identity/test_cache.py +++ b/tests/test_litellm/identity/test_cache.py @@ -10,6 +10,7 @@ from litellm.identity.cache import ( IdentityCache, identity_cache_key, team_generation_key, + user_generation_key, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -75,3 +76,44 @@ async def test_generation_bump_for_unrelated_team_keeps_entry(): def test_key_format_is_versioned(): assert identity_cache_key("abc").startswith("identity:v1:") + + +class _CountingCache: + def __init__(self): + self.batch_get_calls = 0 + self.single_get_calls = 0 + self.values: dict = {} + + async def async_batch_get_cache(self, keys, **kwargs): + self.batch_get_calls += 1 + return [self.values.get(key) for key in keys] + + async def async_get_cache(self, key, **kwargs): + self.single_get_calls += 1 + return self.values.get(key) + + +@pytest.mark.asyncio +async def test_snapshot_generations_uses_single_batch_read(): + fake = _CountingCache() + cache = IdentityCache(dual_cache=fake) # type: ignore[arg-type] + uak = UserAPIKeyAuth(token="hash-x", user_id="u1", team_id="t1", org_id="o1") + + snapshot = await cache._snapshot_generations_for(uak) + + assert fake.batch_get_calls == 1 + assert fake.single_get_calls == 0 + assert snapshot == {"team": 0, "user": 0, "org": 0} + + +@pytest.mark.asyncio +async def test_snapshot_generations_maps_returned_values_by_scope(): + fake = _CountingCache() + fake.values[team_generation_key("t1")] = 7 + fake.values[user_generation_key("u1")] = 3 + cache = IdentityCache(dual_cache=fake) # type: ignore[arg-type] + uak = UserAPIKeyAuth(token="hash-x", user_id="u1", team_id="t1") + + snapshot = await cache._snapshot_generations_for(uak) + + assert snapshot == {"team": 7, "user": 3}