perf(identity): batch generation-counter reads via async_batch_get_cache

This commit is contained in:
Yassin Kortam 2026-06-08 14:41:58 -07:00
parent 3d2927c315
commit 6c94b8a627
2 changed files with 60 additions and 17 deletions

View file

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

View file

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