diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 2c36995c4f8..400d30bc0a2 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -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, diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py new file mode 100644 index 00000000000..52e26e885c9 --- /dev/null +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -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) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 20ab9904f46..2481a7436a7 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 0131a67db8b..b3ddf9bb8cd 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -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: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ba1e2632489..3d751b2ad25 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py new file mode 100644 index 00000000000..52e6e98b500 --- /dev/null +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -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)) diff --git a/tests/proxy_behavior/auth/__init__.py b/tests/proxy_behavior/auth/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/proxy_behavior/auth/conftest.py b/tests/proxy_behavior/auth/conftest.py new file mode 100644 index 00000000000..21982fa25cd --- /dev/null +++ b/tests/proxy_behavior/auth/conftest.py @@ -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() diff --git a/tests/proxy_behavior/auth/test_auth_object_prefetch.py b/tests/proxy_behavior/auth/test_auth_object_prefetch.py new file mode 100644 index 00000000000..e2d947f4284 --- /dev/null +++ b/tests/proxy_behavior/auth/test_auth_object_prefetch.py @@ -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}) diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index bcae33b976e..3c93687d467 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -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)), + ] diff --git a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py new file mode 100644 index 00000000000..0fd0dda3017 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 6cce6d0316b..a257288ebe0 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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(): """ diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py new file mode 100644 index 00000000000..288c3d8c7b4 --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py @@ -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"]