perf(auth): read user, team, membership, org, project and spend counters in one MGET, one query and one pipeline (#40834)

* perf(auth): prefetch user, team, membership, org and project in one MGET, one query and one pipeline

Auth read each object with its own Redis GET and, on a miss, its own DB
query, then the admission spend counters with one GET each. The prefetch
warms every entry the checks read with one MGET, one raw query for the
Redis misses and one pipeline write, and a per-request batch serves the
spend counter reads from one MGET. The per-object getters stay the
readers and the fallback, so enforcement does not depend on the prefetch

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(auth): keep prefetch and spend batch collections immutable

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* perf(auth): let the cold spend-counter reseed reuse the admission MGET instead of one GET per counter

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* perf(auth): prefetch referenced auth objects only after the key's model access check passes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(auth): give the prefetch-ordering test's patches their test-quality reasons

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(auth): move the real-Postgres prefetch join test to the proxy_behavior shard

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(auth): read NULL nested permission and budget lists as [] in the prefetch join

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(caching): assert async_set_cache_pipeline_with_ttls keeps per-entry TTLs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(auth): map the model table's aliases column to model_aliases in the prefetch join and read user memberships the way get_user_object does

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-12 08:52:23 -07:00 • committed by GitHub
parent 9071ca503e
commit f2e0a5db1e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1513 additions and 19 deletions

View file

@ -1138,6 +1138,46 @@ class RedisCache(BaseCache):
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
@_redis_circuit_breaker_guard
async def async_set_cache_pipeline_with_ttls(self, cache_list: Sequence[tuple[str, object, float | None]]) -> None:
"""One round trip for writes whose TTLs differ; a ``None`` TTL falls back to the default TTL."""
if len(cache_list) == 0:
return
commands: Final = tuple(
(self.check_and_fix_namespace(key=cache_key), json.dumps(cache_value), self.get_ttl(ttl=ttl))
for cache_key, cache_value, ttl in cache_list
)
start_time: Final = time.time()
try:
async with self.init_async_client().pipeline(transaction=False) as pipe:
for cache_key, json_cache_value, ttl in commands:
pipe.set(name=cache_key, value=json_cache_value, ex=None if ttl is None else timedelta(seconds=ttl))
await pipe.execute()
asyncio.create_task(
self.service_logger_obj.async_service_success_hook(
service=ServiceTypes.REDIS,
duration=time.time() - start_time,
call_type=f"async_set_cache_pipeline_with_ttls <- {_get_call_stack_info()}",
start_time=start_time,
end_time=time.time(),
)
)
except Exception as e:
asyncio.create_task(
self.service_logger_obj.async_service_failure_hook(
service=ServiceTypes.REDIS,
duration=time.time() - start_time,
error=e,
call_type=f"async_set_cache_pipeline_with_ttls <- {_get_call_stack_info()}",
start_time=start_time,
end_time=time.time(),
)
)
verbose_logger.error(
"LiteLLM Redis Caching: async_set_cache_pipeline_with_ttls() - Got exception from REDIS %s", str(e)
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
async def _set_cache_sadd_helper(
self,
redis_client: async_redis_client,

View file

@ -0,0 +1,307 @@
"""Warm the user, team, membership, org and project cache entries auth reads: one MGET, one DB query, one
pipeline write instead of one Redis GET (and one DB query when cold) per object. The per-object getters stay
the readers and the fallback, so enforcement never depends on this running."""
from __future__ import annotations
import time
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Literal, Protocol, TypeAlias
from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_cache import RedisCache
from litellm.constants import DEFAULT_IN_MEMORY_TTL
from litellm.models.organization import LiteLLM_OrganizationTable
from litellm.models.team import LiteLLM_TeamTableCachedObj
from litellm.models.team_membership import LiteLLM_TeamMembership
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy._types import LiteLLM_ProjectTableCachedObj, UserAPIKeyAuth
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
get_management_object_ttl,
team_membership_auth_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.utils import PrismaClient
_RowKind: TypeAlias = Literal["user_row", "team_row", "membership_row", "organization_row", "project_row"]
_TEAM_MEMBERSHIP_AUTH_TTL: Final = 5
_RowValues: Final = TypeAdapter(dict[str, object])
_NO_ROWS: Final[Mapping[str, object]] = MappingProxyType({})
_TEAM_BOUND_ROWS: Final = frozenset({"team_row", "membership_row"})
_REFRESH_STAMPED_ROWS: Final = frozenset({"team_row", "project_row"})
def _lists_as_json(alias: str, columns: Sequence[str]) -> str:
"""Prisma reads a NULL scalar list as ``[]``; ``to_jsonb`` reads it as ``null``, which the models reject."""
return ", ".join(f"'{column}', COALESCE(to_jsonb({alias}.{column}), '[]'::jsonb)" for column in columns)
_USER_LISTS: Final = _lists_as_json("u", ("teams", "models", "allowed_cache_controls", "policies"))
_TEAM_LISTS: Final = _lists_as_json(
"t",
(
"admins",
"members",
"models",
"team_member_permissions",
"access_group_ids",
"policies",
"default_team_member_models",
),
)
_ORG_LISTS: Final = _lists_as_json("o", ("models",))
_PROJECT_LISTS: Final = _lists_as_json("p", ("models",))
_PERMISSION_LISTS: Final = _lists_as_json(
"op",
(
"mcp_servers",
"mcp_access_groups",
"mcp_toolsets",
"blocked_tools",
"vector_stores",
"agents",
"agent_access_groups",
"models",
"search_tools",
"skills",
),
)
_BUDGET_LISTS: Final = _lists_as_json("b", ("allowed_models",))
def _budget_json(owner_alias: str) -> str:
return (
f"(SELECT to_jsonb(b) || jsonb_build_object({_BUDGET_LISTS}) "
f'FROM "LiteLLM_BudgetTable" b WHERE b.budget_id = {owner_alias}.budget_id)'
)
def _permission_json(owner_alias: str) -> str:
return (
f"(SELECT to_jsonb(op) || jsonb_build_object({_PERMISSION_LISTS}) "
f'FROM "LiteLLM_ObjectPermissionTable" op WHERE op.object_permission_id = {owner_alias}.object_permission_id)'
)
_SQL: Final = f"""
SELECT
(
SELECT to_jsonb(u) || jsonb_build_object(
{_USER_LISTS},
'organization_memberships',
COALESCE((
SELECT jsonb_agg(to_jsonb(om)) FROM "LiteLLM_OrganizationMembership" om WHERE om.user_id = u.user_id
), '[]'::jsonb)
)
FROM "LiteLLM_UserTable" u WHERE u.user_id = $1
) AS user_row,
(
SELECT to_jsonb(t) || jsonb_build_object(
{_TEAM_LISTS},
'litellm_model_table', (
SELECT (to_jsonb(m) - 'aliases') || jsonb_build_object('model_aliases', m.aliases)
FROM "LiteLLM_ModelTable" m WHERE m.id = t.model_id
),
'object_permission', {_permission_json("t")}
)
FROM "LiteLLM_TeamTable" t WHERE t.team_id = $2
) AS team_row,
(
SELECT to_jsonb(tm) || jsonb_build_object('litellm_budget_table', {_budget_json("tm")})
FROM "LiteLLM_TeamMembership" tm WHERE tm.user_id = $3 AND tm.team_id = $2
) AS membership_row,
(
SELECT to_jsonb(o) || jsonb_build_object(
{_ORG_LISTS},
'litellm_budget_table', {_budget_json("o")},
'object_permission', {_permission_json("o")}
)
FROM "LiteLLM_OrganizationTable" o WHERE o.organization_id = $4
) AS organization_row,
(
SELECT to_jsonb(p) || jsonb_build_object(
{_PROJECT_LISTS},
'litellm_budget_table', {_budget_json("p")},
'object_permission', {_permission_json("p")}
)
FROM "LiteLLM_ProjectTable" p WHERE p.project_id = $5
) AS project_row
"""
@dataclass(frozen=True, slots=True)
class AuthObjectRefs:
"""Ids of the objects a request's auth checks will read. ``None`` means not referenced."""
user_id: str | None = None
team_id: str | None = None
membership_user_id: str | None = None
organization_id: str | None = None
project_id: str | None = None
@classmethod
def from_token(cls, token: UserAPIKeyAuth) -> AuthObjectRefs:
has_membership: Final = token.team_id is not None and token.user_id is not None
return cls(
user_id=token.user_id,
team_id=token.team_id,
membership_user_id=token.user_id if has_membership else None,
organization_id=token.org_id,
project_id=token.project_id,
)
class _InMemoryCache(Protocol):
def get_cache(self, key: str) -> object: ...
def set_cache(self, key: str, value: object, *, ttl: float | None = ...) -> None: ...
@dataclass(frozen=True, slots=True)
class _CacheEntry:
cache_key: str
row: _RowKind
model_type: type[BaseModel]
ttl: float | None
def _iter_entries(refs: AuthObjectRefs, management_ttl: float) -> Iterator[_CacheEntry]:
if refs.user_id is not None:
yield _CacheEntry(refs.user_id, "user_row", LiteLLM_UserTable, management_ttl)
if refs.team_id is not None:
yield _CacheEntry(f"team_id:{refs.team_id}", "team_row", LiteLLM_TeamTableCachedObj, management_ttl)
if refs.team_id is not None and refs.membership_user_id is not None:
yield _CacheEntry(
team_membership_auth_cache_key(team_id=refs.team_id, user_id=refs.membership_user_id),
"membership_row",
LiteLLM_TeamMembership,
_TEAM_MEMBERSHIP_AUTH_TTL,
)
yield _CacheEntry(
team_membership_reservation_cache_key(user_id=refs.membership_user_id, team_id=refs.team_id),
"membership_row",
LiteLLM_TeamMembership,
None,
)
if refs.organization_id is not None:
yield _CacheEntry(
f"org_id:{refs.organization_id}", "organization_row", LiteLLM_OrganizationTable, DEFAULT_IN_MEMORY_TTL
)
yield _CacheEntry(
f"org_id:{refs.organization_id}:with_budget",
"organization_row",
LiteLLM_OrganizationTable,
DEFAULT_IN_MEMORY_TTL,
)
if refs.project_id is not None:
yield _CacheEntry(f"project_id:{refs.project_id}", "project_row", LiteLLM_ProjectTableCachedObj, management_ttl)
def _entries(refs: AuthObjectRefs, cache: UserApiKeyCache) -> tuple[_CacheEntry, ...]:
return tuple(_iter_entries(refs, get_management_object_ttl(cache)))
def _missing_in_memory(entries: Sequence[_CacheEntry], memory: _InMemoryCache) -> tuple[_CacheEntry, ...]:
return tuple(entry for entry in entries if memory.get_cache(key=entry.cache_key) is None)
def _set_in_memory(memory: _InMemoryCache, cache_key: str, value: object, ttl: float | None) -> None:
if ttl is None:
memory.set_cache(key=cache_key, value=value)
else:
memory.set_cache(key=cache_key, value=value, ttl=ttl)
async def _fill_from_redis(entries: Sequence[_CacheEntry], redis_cache: RedisCache, memory: _InMemoryCache) -> None:
if not entries:
return
found: Final = _RowValues.validate_python(
await redis_cache.async_batch_get_cache(key_list=sorted(entry.cache_key for entry in entries)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API
)
for entry, value in ((entry, found.get(entry.cache_key)) for entry in entries):
if value is not None:
_set_in_memory(memory, entry.cache_key, value, entry.ttl)
def _validate_row(
row_value: object, model_type: type[BaseModel], row: _RowKind, refreshed_at: float
) -> BaseModel | None:
if row_value is None:
return None
try:
columns: Final = _RowValues.validate_python(row_value)
if row in _REFRESH_STAMPED_ROWS:
stamped: Final = {**columns, "last_refreshed_at": refreshed_at} # mutable-ok: validators write into it
return model_type.model_validate(stamped)
return model_type.model_validate(columns)
except ValidationError as e:
verbose_proxy_logger.warning("auth prefetch: %s did not validate as %s: %s", row, model_type.__name__, e)
return None
async def _fetch_rows(
refs: AuthObjectRefs, kinds: frozenset[_RowKind], prisma_client: PrismaClient
) -> Mapping[str, object]:
row: Final[object] = await prisma_client.db.query_first( # pyright: ignore[reportAny] # prisma types query_first as Any
_SQL,
refs.user_id if "user_row" in kinds else None,
refs.team_id if kinds & _TEAM_BOUND_ROWS else None,
refs.membership_user_id if "membership_row" in kinds else None,
refs.organization_id if "organization_row" in kinds else None,
refs.project_id if "project_row" in kinds else None,
)
return _RowValues.validate_python(row) if row is not None else _NO_ROWS
async def _write_back(entries: Sequence[tuple[_CacheEntry, BaseModel]], cache: UserApiKeyCache) -> None:
payloads: Final = tuple(
(entry.cache_key, CacheCodec.serialize(value, model_type=entry.model_type), entry.ttl)
for entry, value in entries
)
memory: Final[_InMemoryCache] = cache.in_memory_cache
for cache_key, payload, ttl in payloads:
_set_in_memory(memory, cache_key, payload, cache.default_in_memory_ttl if ttl is None else ttl)
if cache.redis_cache is not None:
await cache.redis_cache.async_set_cache_pipeline_with_ttls(payloads)
async def _fill_from_db(
refs: AuthObjectRefs, entries: Sequence[_CacheEntry], cache: UserApiKeyCache, prisma_client: PrismaClient
) -> None:
if not entries:
return
model_for: Final[Mapping[_RowKind, type[BaseModel]]] = MappingProxyType(
{entry.row: entry.model_type for entry in entries}
)
rows: Final = await _fetch_rows(refs, frozenset(model_for), prisma_client)
refreshed_at: Final = time.time()
objects: Final[Mapping[_RowKind, BaseModel | None]] = MappingProxyType(
{row: _validate_row(rows.get(row), model_type, row, refreshed_at) for row, model_type in model_for.items()}
)
writes: Final = tuple((entry, value) for entry in entries if (value := objects[entry.row]) is not None)
if writes:
await _write_back(writes, cache)
async def prefetch_auth_objects(
refs: AuthObjectRefs,
user_api_key_cache: UserApiKeyCache,
prisma_client: PrismaClient | None,
) -> None:
"""Best effort: any failure leaves the per-object getters to fetch as before."""
try:
memory: Final[_InMemoryCache] = user_api_key_cache.in_memory_cache
missing: Final = _missing_in_memory(_entries(refs, user_api_key_cache), memory)
if user_api_key_cache.redis_cache is not None:
await _fill_from_redis(missing, user_api_key_cache.redis_cache, memory)
if prisma_client is None:
return
await _fill_from_db(refs, _missing_in_memory(missing, memory), user_api_key_cache, prisma_client)
except Exception as e: # noqa: BLE001 # warm-up only; the getters enforce and fail closed on their own
verbose_proxy_logger.warning("auth prefetch skipped, falling back to per-object lookups: %s", e)

View file

@ -24,6 +24,7 @@ from starlette.exceptions import WebSocketException
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm.caching.redis_cache import RedisCache
from litellm.constants import (
GLOBAL_PROXY_SPEND_CACHE_KEY,
INVALID_VIRTUAL_KEY_ERROR_MARKER,
@ -63,6 +64,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
from litellm.proxy.auth.auth_method import AuthMethod
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
from litellm.proxy.auth.auth_utils import (
abbreviate_api_key,
get_end_user_id_from_request_body,
@ -101,6 +103,11 @@ from litellm.proxy.common_utils.user_api_key_cache import (
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.spend_tracking.spend_counter_batch import (
bind_admission_counter_keys,
release_spend_counter_batch,
spend_counter_batch_scope,
)
from litellm.proxy.utils import (
PrismaClient,
ProxyLogging,
@ -1970,6 +1977,9 @@ async def _user_api_key_auth_builder(
llm_model_list=llm_model_list,
llm_router=llm_router,
)
await _prefetch_referenced_auth_objects(
valid_token, end_user_id=end_user_id, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client
)
# Check 2. If user_id for this token is in budget - done in common_checks()
if valid_token.user_id is not None:
@ -2672,21 +2682,25 @@ async def _run_centralized_common_checks(
user_api_key_dict=user_api_key_auth_obj,
)
_ = await common_checks(
request=request,
request_body=request_data,
team_object=team_object,
user_object=user_object,
end_user_object=end_user_object,
general_settings=general_settings,
global_proxy_spend=global_proxy_spend,
route=route,
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=user_api_key_auth_obj,
skip_budget_checks=skip_budget_checks,
project_object=project_object,
)
bind_admission_counter_keys(user_api_key_auth_obj, end_user_id=end_user_id)
try:
_ = await common_checks(
request=request,
request_body=request_data,
team_object=team_object,
user_object=user_object,
end_user_object=end_user_object,
general_settings=general_settings,
global_proxy_spend=global_proxy_spend,
route=route,
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=user_api_key_auth_obj,
skip_budget_checks=skip_budget_checks,
project_object=project_object,
)
finally:
release_spend_counter_batch()
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
@ -2864,6 +2878,28 @@ async def _authorize_authenticated_request(
return None
def _spend_counter_redis_cache() -> RedisCache | None:
from litellm.proxy.proxy_server import spend_counter_cache
return spend_counter_cache.redis_cache
async def _prefetch_referenced_auth_objects(
valid_token: UserAPIKeyAuth,
end_user_id: str | None,
user_api_key_cache: UserApiKeyCache,
prisma_client: PrismaClient | None,
) -> None:
"""Warm every object and spend counter the checks below will read, in one MGET each (one DB query when cold).
Runs after the key's model access check so a denied request costs no more than it did before."""
bind_admission_counter_keys(valid_token, end_user_id=end_user_id or None)
await prefetch_auth_objects(
refs=AuthObjectRefs.from_token(valid_token),
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
def _seed_request_destinations(user_api_key_dict: UserAPIKeyAuth, request: Request | None = None) -> None:
"""Anchor the OTLP destinations this key or team overrides its traces to.
@ -2928,7 +2964,7 @@ async def user_api_key_auth(
# Run the whole auth phase inside a live ``auth`` span so the DB lookups it
# triggers (key/user/team object reads) nest under it instead of flattening
# onto the server span. No-op when OTel V2 isn't active.
with phase_span(f"auth {route}"):
with phase_span(f"auth {route}"), spend_counter_batch_scope(_spend_counter_redis_cache()):
try:
user_api_key_auth_obj: Final = await _user_api_key_auth_builder(
request=request,

View file

@ -24,6 +24,7 @@ from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import Litellm_EntityType
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
from litellm.proxy.spend_tracking.spend_counter_batch import active_spend_counter_batch
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.table_repositories import (
BudgetWindowSpendRepository,
@ -194,6 +195,12 @@ class SpendCounterReseed:
return True
return False
@staticmethod
async def _read_active_batch(counter_key: str) -> tuple[float | None, bool] | None:
"""The request's admission MGET already answered for this counter; a Redis miss there is authoritative."""
batch: Final = active_spend_counter_batch()
return None if batch is None else await batch.read(counter_key)
@staticmethod
async def coalesced(
prisma_client: Optional["PrismaClient"],
@ -211,10 +218,13 @@ class SpendCounterReseed:
"""
lock: Final = await SpendCounterReseed._get_lock(counter_key)
async with lock:
batched: Final = await SpendCounterReseed._read_active_batch(counter_key)
if batched is not None and batched[0] is not None:
return batched[0]
# Re-check after acquiring the lock. Skip in-memory on a clean
# Redis miss - in-memory is per-pod-stale.
redis_clean_miss = False
if spend_counter_cache.redis_cache is not None:
redis_clean_miss = batched is not None
if spend_counter_cache.redis_cache is not None and not redis_clean_miss:
try:
val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)
if val is not None:

View file

@ -659,6 +659,7 @@ from litellm.proxy.route_priority import hot_routes_first
from litellm.proxy.search_endpoints.endpoints import router as search_router
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.spend_counter_batch import active_spend_counter_batch
from litellm.proxy.spend_tracking.spend_management_endpoints import (
router as spend_management_router,
)
@ -2724,6 +2725,12 @@ async def read_spend_counter_cache_value(counter_key: str) -> tuple[float | None
"""Return (value, authoritative) for the live counter, None when absent. A clean
Redis miss is final: the per-pod in-memory copy outlives the Redis TTL and only
holds this pod's writes, so it is consulted only when Redis is unreachable."""
batch: Final = active_spend_counter_batch()
if batch is not None:
batched: Final = await batch.read(counter_key)
if batched is not None:
return batched
if spend_counter_cache.redis_cache is not None:
try:
redis_val: Final = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)

View file

@ -0,0 +1,130 @@
"""One Redis MGET for every spend counter the admission checks read, instead of one GET per counter."""
import asyncio
from collections.abc import Iterator, Mapping
from contextvars import ContextVar, Token
from types import MappingProxyType, TracebackType
from typing import Final
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import UserAPIKeyAuth
_CounterValues: Final = TypeAdapter(dict[str, float | None])
_NO_VALUES: Final[Mapping[str, float | None]] = MappingProxyType({})
class SpendCounterBatch:
"""Bound counters are read with one MGET on first use; counters bound later join the next MGET.
``async_batch_get_cache`` maps a clean miss to ``None`` and drops keys only when Redis failed, so an absent
key means "read it yourself" and a present ``None`` is an authoritative miss."""
__slots__ = ("_fetched", "_keys", "_loaded", "_lock", "_open", "_redis_cache")
def __init__(self, redis_cache: RedisCache) -> None:
self._redis_cache: Final = redis_cache
self._lock: Final = asyncio.Lock()
self._open = True
self._keys: frozenset[str] = frozenset()
self._fetched: frozenset[str] = frozenset()
self._loaded: Mapping[str, float | None] = _NO_VALUES
@property
def counter_keys(self) -> frozenset[str]:
return self._keys
def bind(self, counter_keys: frozenset[str]) -> None:
if self._open:
self._keys = self._keys | counter_keys
def close(self) -> None:
"""Later reads go to Redis directly; call before any read-then-write on the counters."""
self._open = False
async def read(self, counter_key: str) -> tuple[float | None, bool] | None:
"""(value, authoritative) for a bound counter, None when the caller must read Redis itself."""
if not self._open or counter_key not in self._keys:
return None
loaded: Final = await self._load()
if counter_key not in loaded:
return None
return loaded[counter_key], True
async def _load(self) -> Mapping[str, float | None]:
async with self._lock:
pending: Final = self._keys - self._fetched
if pending:
self._fetched = self._fetched | pending
self._loaded = MappingProxyType({**self._loaded, **await self._fetch(pending)})
return self._loaded
async def _fetch(self, keys: frozenset[str]) -> Mapping[str, float | None]:
try:
return _CounterValues.validate_python(
await self._redis_cache.async_batch_get_cache(key_list=sorted(keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API
)
except Exception as e: # noqa: BLE001 # per-key reads take over and apply their own Redis fallback
verbose_proxy_logger.debug("spend counter batch read failed, falling back to per-key reads: %s", e)
return _NO_VALUES
_active_batch: Final[ContextVar[SpendCounterBatch | None]] = ContextVar("spend_counter_batch", default=None)
def active_spend_counter_batch() -> SpendCounterBatch | None:
return _active_batch.get()
class spend_counter_batch_scope:
"""Reads inside the scope share one MGET once ``bind_admission_counter_keys`` has run."""
__slots__ = ("_redis_cache", "_token")
def __init__(self, redis_cache: RedisCache | None) -> None:
self._redis_cache: Final = redis_cache
self._token: Token[SpendCounterBatch | None] | None = None
def __enter__(self) -> None:
if self._redis_cache is not None:
self._token = _active_batch.set(SpendCounterBatch(self._redis_cache))
def __exit__(
self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None
) -> None:
if self._token is not None:
_active_batch.reset(self._token)
def release_spend_counter_batch() -> None:
batch: Final = _active_batch.get()
if batch is not None:
batch.close()
def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> Iterator[str]:
if token.token is not None:
yield f"spend:key:{token.token}"
if token.team_id is not None:
yield f"spend:team:{token.team_id}"
if token.user_id is not None:
yield f"spend:team_member:{token.user_id}:{token.team_id}"
if token.user_id is not None:
yield f"spend:user:{token.user_id}"
if end_user_id is not None:
yield f"spend:end_user:{end_user_id}"
if token.org_id is not None:
yield f"spend:org:{token.org_id}"
def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]:
return frozenset(_iter_admission_counter_keys(token, end_user_id))
def bind_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> None:
"""Idempotent: call again after the token gains ids (end user, team org) so those counters join the MGET."""
batch: Final = _active_batch.get()
if batch is None:
return
batch.bind(admission_counter_keys(token, end_user_id))

View file

View file

@ -0,0 +1,20 @@
"""Session-scoped PrismaClient for auth behavior tests that run raw SQL against a real Postgres."""
import os
from unittest.mock import MagicMock
import pytest
import pytest_asyncio
from litellm.proxy.utils import PrismaClient
@pytest_asyncio.fixture(scope="session", loop_scope="session")
async def prisma():
database_url = os.environ.get("DATABASE_URL")
if not database_url:
pytest.skip("DATABASE_URL not set") # test-quality-ok: this suite exists to run SQL on a real Postgres
client = PrismaClient(database_url=database_url, proxy_logging_obj=MagicMock())
await client.connect()
yield client
await client.disconnect()

View file

@ -0,0 +1,174 @@
"""Runs the auth prefetch's raw SQL against a real Postgres: the join must bind the membership to the requested
team and hand the getters rows they validate. The per-regime round-trip counts are unit-tested with fakes in
tests/test_litellm/proxy/auth/test_auth_object_prefetch.py."""
import json
from unittest.mock import AsyncMock, MagicMock
from uuid import uuid4
import pytest
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.proxy.auth.auth_checks import (
get_org_object,
get_team_membership,
get_team_object,
get_user_object,
)
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
from litellm.proxy.auth.team_grants import team_model_aliases
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
pytestmark = pytest.mark.asyncio(loop_scope="session")
def _dead_db() -> MagicMock:
prisma = MagicMock(name="prisma_client")
prisma.db.query_first = AsyncMock(return_value=None)
return prisma
async def test_join_binds_the_membership_to_the_requested_team(prisma):
"""A user in two teams with different member budgets must get the requested team's row."""
run = uuid4().hex
user_id, team_a, team_b, org_id = (f"pf-user-{run}", f"pf-team-a-{run}", f"pf-team-b-{run}", f"pf-org-{run}")
try:
await prisma.db.litellm_budgettable.create(
data={"budget_id": f"a-{run}", "max_budget": 11.0, "created_by": "t", "updated_by": "t"}
)
await prisma.db.litellm_budgettable.create(
data={"budget_id": f"b-{run}", "max_budget": 22.0, "created_by": "t", "updated_by": "t"}
)
await prisma.db.litellm_organizationtable.create(
data={
"organization_id": org_id,
"organization_alias": "pf",
"created_by": "t",
"updated_by": "t",
"litellm_budget_table": {"connect": {"budget_id": f"b-{run}"}},
}
)
await prisma.db.litellm_usertable.create(data={"user_id": user_id, "max_budget": 33.0})
await prisma.db.litellm_teamtable.create(data={"team_id": team_a, "organization_id": org_id, "max_budget": 1.0})
await prisma.db.litellm_teamtable.create(data={"team_id": team_b, "max_budget": 2.0})
await prisma.db.litellm_teammembership.create(
data={"user_id": user_id, "team_id": team_a, "litellm_budget_table": {"connect": {"budget_id": f"a-{run}"}}}
)
await prisma.db.litellm_teammembership.create(
data={"user_id": user_id, "team_id": team_b, "litellm_budget_table": {"connect": {"budget_id": f"b-{run}"}}}
)
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(), redis_cache=None)
refs = AuthObjectRefs(user_id=user_id, team_id=team_a, membership_user_id=user_id, organization_id=org_id)
await prefetch_auth_objects(refs=refs, user_api_key_cache=cache, prisma_client=prisma)
dead_db = _dead_db()
membership = await get_team_membership(
user_id=user_id, team_id=team_a, prisma_client=dead_db, user_api_key_cache=cache
)
team = await get_team_object(team_id=team_a, prisma_client=dead_db, user_api_key_cache=cache)
user = await get_user_object(
user_id=user_id, prisma_client=dead_db, user_api_key_cache=cache, user_id_upsert=False
)
org = await get_org_object(org_id=org_id, prisma_client=dead_db, user_api_key_cache=cache)
assert dead_db.db.mock_calls == [], "getters must be served from the prefetched cache"
assert membership is not None and membership.litellm_budget_table is not None
assert (membership.team_id, membership.litellm_budget_table.max_budget) == (team_a, 11.0)
assert (team.team_id, team.max_budget, team.organization_id, team.models) == (team_a, 1.0, org_id, [])
assert user is not None and user.max_budget == 33.0
assert org is not None and (org.organization_id, org.models) == (org_id, [])
finally:
await prisma.db.litellm_teammembership.delete_many(where={"user_id": user_id})
await prisma.db.litellm_teamtable.delete_many(where={"team_id": {"in": [team_a, team_b]}})
await prisma.db.litellm_usertable.delete_many(where={"user_id": user_id})
await prisma.db.litellm_organizationtable.delete_many(where={"organization_id": org_id})
await prisma.db.litellm_budgettable.delete_many(where={"budget_id": {"in": [f"a-{run}", f"b-{run}"]}})
async def test_join_reads_team_model_aliases_from_the_mapped_column(prisma):
"""The model table stores aliases in a column named ``aliases``; the cached team must expose ``model_aliases``."""
run = uuid4().hex
team_id = f"pf-team-{run}"
aliases = {"gpt-4o": f"gpt-4o-{run}"}
model_table = await prisma.db.litellm_modeltable.create(
data={"model_aliases": json.dumps(aliases), "created_by": "t", "updated_by": "t"}
)
try:
await prisma.db.litellm_teamtable.create(data={"team_id": team_id, "model_id": model_table.id})
expected_team = await prisma.db.litellm_teamtable.find_unique(
where={"team_id": team_id}, include={"litellm_model_table": True}
)
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(), redis_cache=None)
refs = AuthObjectRefs(user_id=None, team_id=team_id, membership_user_id=None, organization_id=None)
await prefetch_auth_objects(refs=refs, user_api_key_cache=cache, prisma_client=prisma)
dead_db = _dead_db()
team = await get_team_object(team_id=team_id, prisma_client=dead_db, user_api_key_cache=cache)
assert dead_db.db.mock_calls == [], "getters must be served from the prefetched cache"
assert expected_team is not None and expected_team.litellm_model_table is not None
assert team.litellm_model_table is not None
assert team.litellm_model_table.model_aliases == expected_team.litellm_model_table.model_aliases == aliases
assert team_model_aliases(team) == aliases
finally:
await prisma.db.litellm_teamtable.delete_many(where={"team_id": team_id})
await prisma.db.litellm_modeltable.delete_many(where={"id": model_table.id})
async def test_join_reads_null_nested_lists_the_way_prisma_does(prisma):
"""Prisma reads a NULL scalar list as []; the nested permission and budget rows must match, not carry null."""
run = uuid4().hex
user_id, team_id, permission_id, budget_id = (f"pf-user-{run}", f"pf-team-{run}", f"pf-perm-{run}", f"pf-bud-{run}")
try:
await prisma.db.litellm_objectpermissiontable.create(data={"object_permission_id": permission_id})
await prisma.db.litellm_budgettable.create(data={"budget_id": budget_id, "created_by": "t", "updated_by": "t"})
await prisma.db.execute_raw(
'UPDATE "LiteLLM_ObjectPermissionTable" SET mcp_servers = NULL, models = NULL '
"WHERE object_permission_id = $1",
permission_id,
)
await prisma.db.execute_raw(
'UPDATE "LiteLLM_BudgetTable" SET allowed_models = NULL WHERE budget_id = $1', budget_id
)
await prisma.db.litellm_usertable.create(data={"user_id": user_id})
await prisma.db.litellm_teamtable.create(data={"team_id": team_id, "object_permission_id": permission_id})
await prisma.db.litellm_teammembership.create(
data={"user_id": user_id, "team_id": team_id, "litellm_budget_table": {"connect": {"budget_id": budget_id}}}
)
expected_team = await prisma.db.litellm_teamtable.find_unique(
where={"team_id": team_id}, include={"object_permission": True}
)
expected_membership = await prisma.db.litellm_teammembership.find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, include={"litellm_budget_table": True}
)
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(), redis_cache=None)
refs = AuthObjectRefs(user_id=user_id, team_id=team_id, membership_user_id=user_id, organization_id=None)
await prefetch_auth_objects(refs=refs, user_api_key_cache=cache, prisma_client=prisma)
dead_db = _dead_db()
team = await get_team_object(team_id=team_id, prisma_client=dead_db, user_api_key_cache=cache)
membership = await get_team_membership(
user_id=user_id, team_id=team_id, prisma_client=dead_db, user_api_key_cache=cache
)
assert dead_db.db.mock_calls == [], "getters must be served from the prefetched cache"
assert expected_team is not None and expected_team.object_permission is not None
assert team.object_permission is not None
assert team.object_permission.mcp_servers == expected_team.object_permission.mcp_servers == []
assert team.object_permission.models == expected_team.object_permission.models == []
assert expected_membership is not None and expected_membership.litellm_budget_table is not None
assert membership is not None and membership.litellm_budget_table is not None
assert (
membership.litellm_budget_table.allowed_models
== expected_membership.litellm_budget_table.allowed_models
== []
)
finally:
await prisma.db.litellm_teammembership.delete_many(where={"user_id": user_id})
await prisma.db.litellm_teamtable.delete_many(where={"team_id": team_id})
await prisma.db.litellm_usertable.delete_many(where={"user_id": user_id})
await prisma.db.litellm_objectpermissiontable.delete_many(where={"object_permission_id": permission_id})
await prisma.db.litellm_budgettable.delete_many(where={"budget_id": budget_id})

View file

@ -1,10 +1,12 @@
import asyncio
import time
from collections.abc import Iterator
from datetime import timedelta
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm._service_logger import ServiceLogging
from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError
@ -679,7 +681,6 @@ def test_sync_batch_get_cache_survives_a_service_callback_that_raises(
from concurrent.futures import ThreadPoolExecutor
import litellm
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
cache, service_logger = sync_batch_cache_with_service_logger
@ -1202,3 +1203,45 @@ async def test_a_probe_overtaken_by_a_later_outage_leaves_the_breaker_to_the_new
new_probe_release.set()
assert await new_probe == "new probe"
assert breaker._state == breaker.CLOSED
class _SetRecordingPipeline:
def __init__(self) -> None:
self.sets: list[tuple[str, str, timedelta | None]] = []
self.executes = 0
async def __aenter__(self) -> "_SetRecordingPipeline":
return self
async def __aexit__(self, *exc_info: object) -> None:
return None
def set(self, name: str, value: str, ex: timedelta | None) -> None:
self.sets.append((name, value, ex))
async def execute(self) -> list[bool]:
self.executes += 1
return [True] * len(self.sets)
@pytest.mark.asyncio
async def test_async_set_cache_pipeline_with_ttls_keeps_each_entry_ttl(monkeypatch, redis_no_ping):
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
monkeypatch.setattr(litellm, "default_redis_ttl", 300)
redis_cache = RedisCache(namespace="ns")
pipe = _SetRecordingPipeline()
client = MagicMock()
client.pipeline = MagicMock(return_value=pipe)
with patch.object(redis_cache, "init_async_client", return_value=client):
await redis_cache.async_set_cache_pipeline_with_ttls(
(("team_id:t1", {"team_id": "t1"}, 60), ("u1", {"user_id": "u1"}, 7), ("org_id:o1", {"a": 1}, None))
)
client.pipeline.assert_called_once_with(transaction=False)
assert pipe.executes == 1
assert pipe.sets == [
("ns:team_id:t1", '{"team_id": "t1"}', timedelta(seconds=60)),
("ns:u1", '{"user_id": "u1"}', timedelta(seconds=7)),
("ns:org_id:o1", '{"a": 1}', timedelta(seconds=300)),
]

View file

@ -0,0 +1,338 @@
"""Counts the Redis round trips and DB queries auth object reads cost per cache regime, and checks that the
per-object getters still enforce on their own when the prefetch cannot help."""
import json
from collections.abc import Sequence
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import (
LiteLLM_OrganizationTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTableCachedObj,
LiteLLM_UserTable,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import (
get_org_object,
get_team_membership,
get_team_object,
get_user_object,
)
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
USER_ID = "prefetch-user"
TEAM_ID = "prefetch-team"
ORG_ID = "prefetch-org"
USER_ROW = {
"user_id": USER_ID,
"max_budget": 50.0,
"spend": 1.0,
"models": ["gpt-5.4-mini"],
"organization_memberships": [],
}
TEAM_ROW = {
"team_id": TEAM_ID,
"organization_id": ORG_ID,
"max_budget": 500.0,
"spend": 2.0,
"models": [],
"blocked": False,
"members_with_roles": {},
}
MEMBERSHIP_ROW = {
"user_id": USER_ID,
"team_id": TEAM_ID,
"spend": 3.0,
"budget_id": "b1",
"litellm_budget_table": {"budget_id": "b1", "max_budget": 20.0},
}
ORG_ROW = {
"organization_id": ORG_ID,
"organization_alias": "org",
"budget_id": "b2",
"created_by": "admin",
"updated_by": "admin",
"models": [],
"spend": 4.0,
"litellm_budget_table": {"budget_id": "b2", "max_budget": 1000.0},
}
ALL_ROWS = {
"user_row": USER_ROW,
"team_row": TEAM_ROW,
"membership_row": MEMBERSHIP_ROW,
"organization_row": ORG_ROW,
"project_row": None,
}
class CountingRedis(RedisCache):
"""Redis fake that counts commands and round trips (an MGET or a pipeline is one round trip)."""
def __init__(self, store: dict[str, str] | None = None, fail: bool = False) -> None:
self.store: dict[str, str] = dict(store or {})
self.fail = fail
self.round_trips = 0
self.commands: list[str] = []
def _trip(self, *commands: str) -> None:
if self.fail:
raise ConnectionError("redis down")
self.round_trips += 1
self.commands.extend(commands)
async def async_get_cache(self, key: str, **kwargs: object) -> object:
self._trip(f"GET {key}")
raw = self.store.get(key)
return json.loads(raw) if raw is not None else None
async def async_batch_get_cache(self, key_list: Sequence[str], **kwargs: object) -> dict[str, object]:
self._trip(f"MGET {' '.join(key_list)}")
return {key: (json.loads(self.store[key]) if key in self.store else None) for key in key_list}
async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None:
self._trip(f"SET {key}")
self.store[key] = json.dumps(value)
async def async_set_cache_pipeline(self, cache_list: Sequence[tuple[str, object]], **kwargs: object) -> None:
self._trip(*(f"SET {key}" for key, _ in cache_list))
for key, value in cache_list:
self.store[key] = json.dumps(value)
async def async_set_cache_pipeline_with_ttls(self, cache_list: Sequence[tuple[str, object, float | None]]) -> None:
self._trip(*(f"SET {key} ttl={ttl}" for key, _, ttl in cache_list))
for key, value, _ in cache_list:
self.store[key] = json.dumps(value)
async def async_delete_cache(self, key: str) -> None:
self._trip(f"DEL {key}")
self.store.pop(key, None)
def _prisma(rows: dict[str, object] | None = ALL_ROWS) -> MagicMock:
prisma = MagicMock(name="prisma_client")
prisma.db.query_first = AsyncMock(return_value=rows)
return prisma
def _non_prefetch_db_calls(prisma: MagicMock) -> list[str]:
return [str(call) for call in prisma.db.mock_calls if not str(call).startswith("call.query_first(")]
def _cache(redis: RedisCache | None) -> UserApiKeyCache:
return UserApiKeyCache(in_memory_cache=InMemoryCache(), redis_cache=redis)
def _refs() -> AuthObjectRefs:
return AuthObjectRefs.from_token(UserAPIKeyAuth(token="t", user_id=USER_ID, team_id=TEAM_ID, org_id=ORG_ID))
async def _read_all_through_getters(
cache: UserApiKeyCache, prisma: MagicMock
) -> tuple[
LiteLLM_UserTable | None,
LiteLLM_TeamTableCachedObj,
LiteLLM_TeamMembership | None,
LiteLLM_OrganizationTable | None,
]:
return (
await get_user_object(user_id=USER_ID, prisma_client=prisma, user_api_key_cache=cache, user_id_upsert=False),
await get_team_object(team_id=TEAM_ID, prisma_client=prisma, user_api_key_cache=cache),
await get_team_membership(user_id=USER_ID, team_id=TEAM_ID, prisma_client=prisma, user_api_key_cache=cache),
await get_org_object(org_id=ORG_ID, prisma_client=prisma, user_api_key_cache=cache, include_budget_table=True),
)
def test_refs_from_token_only_names_membership_when_both_ids_present():
assert AuthObjectRefs.from_token(UserAPIKeyAuth(token="t", team_id=TEAM_ID)).membership_user_id is None
assert AuthObjectRefs.from_token(UserAPIKeyAuth(token="t", user_id=USER_ID)).membership_user_id is None
assert AuthObjectRefs.from_token(UserAPIKeyAuth(token="t", user_id=USER_ID, team_id=TEAM_ID)) == AuthObjectRefs(
user_id=USER_ID, team_id=TEAM_ID, membership_user_id=USER_ID
)
@pytest.mark.asyncio
async def test_cold_regime_is_one_mget_one_query_and_the_getters_never_touch_io_again():
redis = CountingRedis()
prisma = _prisma()
cache = _cache(redis)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
assert prisma.db.query_first.await_count == 1
assert prisma.db.query_first.await_args.args[1:] == (USER_ID, TEAM_ID, USER_ID, ORG_ID, None)
mgets = [c for c in redis.commands if c.startswith("MGET")]
assert len(mgets) == 1
assert set(mgets[0].split()[1:]) == {
USER_ID,
f"team_id:{TEAM_ID}",
f"{TEAM_ID}_{USER_ID}",
f"team_membership:{USER_ID}:{TEAM_ID}",
f"org_id:{ORG_ID}",
f"org_id:{ORG_ID}:with_budget",
}
sets = sorted(c for c in redis.commands if c.startswith("SET"))
assert sets == sorted(
[
f"SET {TEAM_ID}_{USER_ID} ttl=5",
f"SET org_id:{ORG_ID} ttl=5",
f"SET org_id:{ORG_ID}:with_budget ttl=5",
f"SET {USER_ID} ttl=60",
f"SET team_id:{TEAM_ID} ttl=60",
f"SET team_membership:{USER_ID}:{TEAM_ID} ttl=None",
]
)
assert redis.round_trips == 2, "one MGET, one pipeline"
before = (redis.round_trips, prisma.db.query_first.await_count)
user, team, membership, org = await _read_all_through_getters(cache, prisma)
assert (redis.round_trips, prisma.db.query_first.await_count) == before
assert _non_prefetch_db_calls(prisma) == []
assert isinstance(user, LiteLLM_UserTable) and user.max_budget == 50.0
assert isinstance(team, LiteLLM_TeamTableCachedObj) and team.organization_id == ORG_ID
assert team.last_refreshed_at is not None
assert isinstance(membership, LiteLLM_TeamMembership) and membership.litellm_budget_table is not None
assert membership.litellm_budget_table.max_budget == 20.0
assert isinstance(org, LiteLLM_OrganizationTable) and org.litellm_budget_table is not None
assert org.litellm_budget_table.max_budget == 1000.0
@pytest.mark.asyncio
async def test_redis_warm_regime_is_exactly_one_mget_and_zero_queries():
seeded = CountingRedis()
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=_cache(seeded), prisma_client=_prisma())
redis = CountingRedis(store=seeded.store)
prisma = _prisma()
cache = _cache(redis)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
assert redis.round_trips == 1
assert redis.commands[0].startswith("MGET")
assert prisma.db.query_first.await_count == 0
user, team, membership, org = await _read_all_through_getters(cache, prisma)
assert redis.round_trips == 1
assert prisma.db.mock_calls == []
assert (user.user_id, team.team_id, membership.team_id, org.organization_id) == (USER_ID, TEAM_ID, TEAM_ID, ORG_ID)
@pytest.mark.asyncio
async def test_hot_regime_costs_nothing():
redis = CountingRedis()
prisma = _prisma()
cache = _cache(redis)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
redis.round_trips, redis.commands = 0, []
prisma.db.reset_mock()
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
await _read_all_through_getters(cache, prisma)
assert redis.round_trips == 0
assert prisma.db.mock_calls == []
@pytest.mark.asyncio
async def test_partial_redis_hit_queries_only_the_missing_objects():
seeded = CountingRedis()
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=_cache(seeded), prisma_client=_prisma())
for key in (f"team_id:{TEAM_ID}", f"{TEAM_ID}_{USER_ID}", f"team_membership:{USER_ID}:{TEAM_ID}"):
del seeded.store[key]
prisma = _prisma()
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=_cache(seeded), prisma_client=prisma)
assert prisma.db.query_first.await_count == 1
assert prisma.db.query_first.await_args.args[1:] == (None, TEAM_ID, USER_ID, None, None)
del seeded.store[f"org_id:{ORG_ID}:with_budget"]
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=_cache(seeded), prisma_client=prisma)
assert prisma.db.query_first.await_args.args[1:] == (None, None, None, ORG_ID, None)
@pytest.mark.asyncio
async def test_deleted_team_cache_entry_is_refetched_and_the_update_is_visible():
redis = CountingRedis()
prisma = _prisma()
cache = _cache(redis)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
await cache.async_delete_cache(f"team_id:{TEAM_ID}")
prisma.db.query_first.return_value = {**ALL_ROWS, "team_row": {**TEAM_ROW, "blocked": True}}
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
team = await get_team_object(team_id=TEAM_ID, prisma_client=prisma, user_api_key_cache=cache)
assert team.blocked is True
assert prisma.db.query_first.await_count == 2
assert prisma.db.query_first.await_args.args[1:] == (None, TEAM_ID, None, None, None)
@pytest.mark.asyncio
async def test_row_missing_a_required_column_is_not_cached_so_the_getter_still_fails_closed():
prisma = _prisma({**ALL_ROWS, "team_row": {"max_budget": 1.0}})
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
cache = _cache(CountingRedis())
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
with pytest.raises(HTTPException) as exc:
await get_team_object(team_id=TEAM_ID, prisma_client=prisma, user_api_key_cache=cache)
assert exc.value.status_code == 404
assert prisma.db.litellm_teamtable.find_unique.await_count == 1
assert cache.in_memory_cache.get_cache(USER_ID) is not None
@pytest.mark.asyncio
async def test_absent_rows_are_not_cached_as_present():
redis = CountingRedis()
prisma = _prisma({key: None for key in ALL_ROWS})
cache = _cache(redis)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
assert [c for c in redis.commands if c.startswith("SET")] == []
assert cache.in_memory_cache.get_cache(f"team_id:{TEAM_ID}") is None
@pytest.mark.asyncio
async def test_redis_failure_is_swallowed_and_getters_fall_back_to_their_own_reads():
prisma = _prisma()
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
cache = _cache(CountingRedis(fail=True))
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
assert prisma.db.query_first.await_count == 0
assert cache.in_memory_cache.get_cache(USER_ID) is None
@pytest.mark.asyncio
async def test_no_prisma_still_uses_redis_but_never_queries():
seeded = CountingRedis()
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=_cache(seeded), prisma_client=_prisma())
redis = CountingRedis(store=seeded.store)
cache = _cache(redis)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=None)
assert redis.round_trips == 1
assert cache.in_memory_cache.get_cache(f"org_id:{ORG_ID}") is not None
@pytest.mark.asyncio
async def test_no_redis_goes_straight_to_one_query():
prisma = _prisma()
cache = _cache(None)
await prefetch_auth_objects(refs=_refs(), user_api_key_cache=cache, prisma_client=prisma)
assert prisma.db.query_first.await_count == 1
assert cache.in_memory_cache.get_cache(f"team_membership:{USER_ID}:{TEAM_ID}") is not None

View file

@ -1646,6 +1646,95 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
@pytest.mark.parametrize("model_allowed", [True, False])
async def test_auth_prefetches_referenced_objects_only_after_the_key_may_call_the_model(model_allowed):
"""A request denied by the key's model list must not pay for the team/user/org MGET or DB join."""
from fastapi import Request
from starlette.datastructures import URL
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
from litellm.proxy.proxy_server import hash_token
api_key = "sk-prefetch-order-test"
valid_token = UserAPIKeyAuth(api_key=api_key, token=hash_token(api_key), user_id="u1", team_id="t1")
mock_cache = AsyncMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.delete_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
_attrs_to_set = {
"prisma_client": MagicMock(),
"user_api_key_cache": mock_cache,
"proxy_logging_obj": mock_proxy_logging_obj,
"master_key": "sk-master-key",
"general_settings": {},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
denied = ProxyException(
message="Key not allowed to access model",
type=ProxyErrorTypes.key_model_access_denied,
param="model",
code=401,
)
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with (
patch( # test-quality-ok: the builder has no DI seam for the key lookup; stands in for the DB
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
return_value=valid_token,
),
patch( # test-quality-ok: the observable is whether the prefetch runs before or after this check
"litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access",
new_callable=AsyncMock,
side_effect=None if model_allowed else denied,
),
patch( # test-quality-ok: counting prefetch calls on a denied request IS the regression being pinned
"litellm.proxy.auth.user_api_key_auth.prefetch_auth_objects", new_callable=AsyncMock
) as mock_prefetch,
patch( # test-quality-ok: no DB in this test; the user lookup must not fail the allowed path
"litellm.proxy.auth.user_api_key_auth.get_user_object", new_callable=AsyncMock, return_value=None
),
):
call = _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={"model": "gpt-4o"},
)
if model_allowed:
assert isinstance(await call, UserAPIKeyAuth)
mock_prefetch.assert_awaited_once()
assert mock_prefetch.await_args.kwargs["refs"].team_id == "t1"
else:
with pytest.raises(ProxyException) as exc:
await call
assert exc.value.type == ProxyErrorTypes.key_model_access_denied
mock_prefetch.assert_not_awaited()
finally:
for attr, val in _original_values.items():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
async def test_return_user_api_key_auth_obj_user_spend_and_budget():
"""

View file

@ -0,0 +1,300 @@
"""Exact Redis round-trip counts for the spend counters admission reads within one auth scope."""
from __future__ import annotations
import asyncio
from collections.abc import Sequence
from unittest.mock import AsyncMock, MagicMock
import pytest
import litellm.proxy.proxy_server as ps
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
from litellm.proxy.spend_tracking.spend_counter_batch import (
SpendCounterBatch,
active_spend_counter_batch,
admission_counter_keys,
bind_admission_counter_keys,
release_spend_counter_batch,
spend_counter_batch_scope,
)
TOKEN = UserAPIKeyAuth(token="hashed", team_id="team", user_id="user", org_id="org")
TOKEN_KEYS = frozenset(
{
"spend:key:hashed",
"spend:team:team",
"spend:team_member:user:team",
"spend:user:user",
"spend:end_user:eu",
"spend:org:org",
}
)
class CountingRedis(RedisCache):
def __init__(self, store: dict[str, object] | None = None, fail: bool = False) -> None:
self.store: dict[str, object] = dict(store or {})
self.fail = fail
self.commands: list[str] = []
async def async_get_cache(self, key: str, **kwargs: object) -> object:
if self.fail:
raise ConnectionError("redis down")
self.commands.append(f"GET {key}")
return self.store.get(key)
async def async_batch_get_cache(self, key_list: Sequence[str], **kwargs: object) -> dict[str, object]:
if self.fail:
raise ConnectionError("redis down")
self.commands.append(f"MGET {' '.join(key_list)}")
return {key: self.store.get(key) for key in key_list}
def _spend_counter_cache(redis: RedisCache | None, in_memory: dict[str, float] | None = None) -> MagicMock:
cache = MagicMock()
cache.redis_cache = redis
cache.in_memory_cache.get_cache = MagicMock(side_effect=lambda key: (in_memory or {}).get(key))
return cache
def test_admission_counter_keys_cover_every_entity_the_checks_read():
assert admission_counter_keys(TOKEN, end_user_id="eu") == TOKEN_KEYS
assert admission_counter_keys(UserAPIKeyAuth(token="hashed"), end_user_id=None) == {"spend:key:hashed"}
assert "spend:team_member:user:team" not in admission_counter_keys(
UserAPIKeyAuth(token="hashed", user_id="user"), end_user_id=None
)
@pytest.mark.asyncio
async def test_bound_counters_share_one_mget_and_a_clean_miss_is_authoritative():
redis = CountingRedis({"spend:key:hashed": 1.5, "spend:team:team": 2.5})
batch = SpendCounterBatch(redis)
batch.bind(TOKEN_KEYS)
reads = await asyncio.gather(*(batch.read(key) for key in sorted(TOKEN_KEYS)))
assert len(redis.commands) == 1
assert set(redis.commands[0].split()[1:]) == TOKEN_KEYS
assert dict(zip(sorted(TOKEN_KEYS), reads)) == {
"spend:end_user:eu": (None, True),
"spend:key:hashed": (1.5, True),
"spend:org:org": (None, True),
"spend:team:team": (2.5, True),
"spend:team_member:user:team": (None, True),
"spend:user:user": (None, True),
}
@pytest.mark.asyncio
async def test_unbound_counter_and_closed_batch_leave_the_read_to_the_caller():
redis = CountingRedis({"spend:key:hashed": 1.0})
batch = SpendCounterBatch(redis)
batch.bind(frozenset({"spend:key:hashed"}))
assert await batch.read("spend:tag:prod") is None
assert redis.commands == []
assert await batch.read("spend:key:hashed") == (1.0, True)
batch.close()
batch.bind(frozenset({"spend:org:org"}))
assert await batch.read("spend:key:hashed") is None
assert await batch.read("spend:org:org") is None
assert len(redis.commands) == 1
@pytest.mark.asyncio
async def test_keys_bound_after_the_first_read_join_one_more_mget_for_only_the_new_keys():
redis = CountingRedis({"spend:key:hashed": 1.0, "spend:org:org": 9.0})
batch = SpendCounterBatch(redis)
batch.bind(frozenset({"spend:key:hashed"}))
assert await batch.read("spend:key:hashed") == (1.0, True)
batch.bind(frozenset({"spend:org:org", "spend:key:hashed"}))
assert await batch.read("spend:org:org") == (9.0, True)
assert await batch.read("spend:key:hashed") == (1.0, True)
assert redis.commands == ["MGET spend:key:hashed", "MGET spend:org:org"]
@pytest.mark.asyncio
async def test_failed_mget_hands_every_counter_back_to_the_caller():
batch = SpendCounterBatch(CountingRedis(fail=True))
batch.bind(TOKEN_KEYS)
assert await batch.read("spend:key:hashed") is None
@pytest.mark.asyncio
async def test_non_numeric_counter_payload_hands_the_batch_back_to_the_caller():
redis = CountingRedis()
redis.store["spend:key:hashed"] = "garbage"
batch = SpendCounterBatch(redis)
batch.bind(frozenset({"spend:key:hashed"}))
assert await batch.read("spend:key:hashed") is None
def test_scope_installs_a_batch_only_when_redis_exists_and_release_closes_without_clearing():
assert active_spend_counter_batch() is None
with spend_counter_batch_scope(None):
assert active_spend_counter_batch() is None
bind_admission_counter_keys(TOKEN, end_user_id=None)
with spend_counter_batch_scope(CountingRedis()):
batch = active_spend_counter_batch()
assert batch is not None
bind_admission_counter_keys(TOKEN, end_user_id="eu")
assert batch.counter_keys == TOKEN_KEYS
release_spend_counter_batch()
assert active_spend_counter_batch() is batch
batch.bind(frozenset({"spend:tag:x"}))
assert batch.counter_keys == TOKEN_KEYS
assert active_spend_counter_batch() is None
@pytest.mark.asyncio
async def test_get_current_spend_inside_the_scope_costs_one_mget_for_all_admission_counters(monkeypatch):
redis = CountingRedis({"spend:key:hashed": 3.0, "spend:team:team": 4.0, "spend:org:org": 5.0})
monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis))
monkeypatch.setattr(ps, "prisma_client", None)
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(TOKEN, end_user_id="eu")
key_spend = await ps.get_current_spend(counter_key="spend:key:hashed", fallback_spend=0.0)
team_spend = await ps.get_current_spend(counter_key="spend:team:team", fallback_spend=0.0)
org_spend = await ps.get_current_spend(counter_key="spend:org:org", fallback_spend=0.0)
user_spend = await ps.get_current_spend(counter_key="spend:user:user", fallback_spend=7.0)
assert (key_spend, team_spend, org_spend, user_spend) == (3.0, 4.0, 5.0, 7.0)
assert [c for c in redis.commands if c.startswith("GET ")] == [], "the cold reseed reuses the MGET miss"
assert [c for c in redis.commands if c.startswith("MGET ")] == [f"MGET {' '.join(sorted(TOKEN_KEYS))}"]
@pytest.mark.asyncio
async def test_get_current_spend_outside_the_scope_still_reads_redis_per_counter(monkeypatch):
redis = CountingRedis({"spend:key:hashed": 3.0})
monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis))
assert await ps.get_current_spend(counter_key="spend:key:hashed", fallback_spend=0.0) == 3.0
assert redis.commands == ["GET spend:key:hashed"]
@pytest.mark.asyncio
async def test_after_release_a_read_goes_to_redis_directly_so_read_then_write_sees_fresh_values(monkeypatch):
redis = CountingRedis({"spend:key:hashed": 3.0})
monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis))
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(TOKEN, end_user_id=None)
assert await ps.read_spend_counter_cache_value("spend:key:hashed") == (3.0, True)
redis.store["spend:key:hashed"] = 8.0
assert await ps.read_spend_counter_cache_value("spend:key:hashed") == (3.0, True)
release_spend_counter_batch()
assert await ps.read_spend_counter_cache_value("spend:key:hashed") == (8.0, True)
assert [c.split()[0] for c in redis.commands] == ["MGET", "GET"]
@pytest.mark.asyncio
async def test_batched_clean_miss_does_not_fall_back_to_the_per_pod_in_memory_copy(monkeypatch):
redis = CountingRedis()
monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis, in_memory={"spend:key:hashed": 99.0}))
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(TOKEN, end_user_id=None)
assert await ps.read_spend_counter_cache_value("spend:key:hashed") == (None, True)
@pytest.mark.asyncio
async def test_batched_redis_failure_falls_back_to_the_per_pod_in_memory_copy(monkeypatch):
redis = CountingRedis(fail=True)
monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis, in_memory={"spend:key:hashed": 99.0}))
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(TOKEN, end_user_id=None)
assert await ps.read_spend_counter_cache_value("spend:key:hashed") == (99.0, False)
@pytest.mark.asyncio
async def test_scope_is_per_task_so_concurrent_requests_do_not_share_a_batch():
redis = CountingRedis({"spend:key:a": 1.0, "spend:key:b": 2.0})
async def request(token: str) -> tuple[float | None, bool] | None:
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(UserAPIKeyAuth(token=token), end_user_id=None)
batch = active_spend_counter_batch()
assert batch is not None
await asyncio.sleep(0)
return await batch.read(f"spend:key:{token}")
assert await asyncio.gather(request("a"), request("b")) == [(1.0, True), (2.0, True)]
assert sorted(redis.commands) == ["MGET spend:key:a", "MGET spend:key:b"]
@pytest.mark.asyncio
async def test_batch_reads_never_touch_a_prisma_client_when_redis_answers(monkeypatch):
redis = CountingRedis({"spend:key:hashed": 3.0})
prisma = MagicMock()
prisma.db.litellm_verificationtoken.find_unique = AsyncMock()
monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis))
monkeypatch.setattr(ps, "prisma_client", prisma)
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(TOKEN, end_user_id=None)
assert await ps.get_current_spend(counter_key="spend:key:hashed", fallback_spend=0.0) == 3.0
assert prisma.db.mock_calls == []
def _reseed_prisma(spend: float) -> MagicMock:
prisma = MagicMock()
prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=MagicMock(spend=spend))
return prisma
@pytest.mark.asyncio
async def test_reseed_reuses_the_admission_mget_instead_of_its_own_get():
redis = CountingRedis({"spend:key:hashed": 7.5})
prisma = _reseed_prisma(spend=1.0)
cache = _spend_counter_cache(redis)
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(TOKEN, end_user_id=None)
value = await SpendCounterReseed.coalesced(prisma, cache, counter_key="spend:key:hashed")
assert value == 7.5
assert redis.commands == [
"MGET spend:key:hashed spend:org:org spend:team:team spend:team_member:user:team spend:user:user"
]
assert prisma.db.mock_calls == []
@pytest.mark.asyncio
async def test_reseed_treats_a_batched_clean_miss_as_authoritative_and_seeds_from_the_db():
redis = CountingRedis()
redis.async_set_cache = AsyncMock(return_value=True)
prisma = _reseed_prisma(spend=2.25)
cache = _spend_counter_cache(redis, in_memory={"spend:key:hashed": 99.0})
with spend_counter_batch_scope(redis):
bind_admission_counter_keys(TOKEN, end_user_id=None)
value = await SpendCounterReseed.coalesced(prisma, cache, counter_key="spend:key:hashed")
assert value == 2.25
assert [c for c in redis.commands if c.startswith("GET")] == []
redis.async_set_cache.assert_awaited_once_with(key="spend:key:hashed", value=2.25, nx=True)
@pytest.mark.asyncio
async def test_reseed_outside_the_scope_still_re_checks_redis_itself():
redis = CountingRedis({"spend:key:hashed": 4.0})
value = await SpendCounterReseed.coalesced(
_reseed_prisma(spend=1.0), _spend_counter_cache(redis), "spend:key:hashed"
)
assert value == 4.0
assert redis.commands == ["GET spend:key:hashed"]