mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
perf(identity): batch generation-counter reads via async_batch_get_cache
This commit is contained in:
parent
3d2927c315
commit
6c94b8a627
2 changed files with 60 additions and 17 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue