Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/auth/auth_object_prefetch.py: 74%
128 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Warm the user, team, membership, org and project cache entries auth reads: one MGET, one DB query, one
2pipeline write instead of one Redis GET (and one DB query when cold) per object. The per-object getters stay
3the readers and the fallback, so enforcement never depends on this running."""
5from __future__ import annotations
7import time
8from collections.abc import Iterator, Mapping, Sequence
9from dataclasses import dataclass
10from types import MappingProxyType
11from typing import Final, Literal, Protocol, TypeAlias
13from pydantic import BaseModel, TypeAdapter, ValidationError
15from litellm._logging import verbose_proxy_logger
16from litellm.caching.redis_cache import RedisCache
17from litellm.constants import DEFAULT_IN_MEMORY_TTL
18from litellm.models.organization import LiteLLM_OrganizationTable
19from litellm.models.team import LiteLLM_TeamTableCachedObj
20from litellm.models.team_membership import LiteLLM_TeamMembership
21from litellm.models.user import LiteLLM_UserTable
22from litellm.proxy._types import LiteLLM_ProjectTableCachedObj, UserAPIKeyAuth
23from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
24from litellm.proxy.common_utils.user_api_key_cache import (
25 UserApiKeyCache,
26 get_management_object_ttl,
27 team_membership_auth_cache_key,
28 team_membership_reservation_cache_key,
29)
30from litellm.proxy.utils import PrismaClient
32_RowKind: TypeAlias = Literal["user_row", "team_row", "membership_row", "organization_row", "project_row"]
34_TEAM_MEMBERSHIP_AUTH_TTL: Final = 5
35_RowValues: Final = TypeAdapter(dict[str, object])
36_NO_ROWS: Final[Mapping[str, object]] = MappingProxyType({})
37_TEAM_BOUND_ROWS: Final = frozenset({"team_row", "membership_row"})
38_REFRESH_STAMPED_ROWS: Final = frozenset({"team_row", "project_row"})
41def _lists_as_json(alias: str, columns: Sequence[str]) -> str:
42 """Prisma reads a NULL scalar list as ``[]``; ``to_jsonb`` reads it as ``null``, which the models reject."""
43 return ", ".join(f"'{column}', COALESCE(to_jsonb({alias}.{column}), '[]'::jsonb)" for column in columns)
46_USER_LISTS: Final = _lists_as_json("u", ("teams", "models", "allowed_cache_controls", "policies"))
47_TEAM_LISTS: Final = _lists_as_json(
48 "t",
49 (
50 "admins",
51 "members",
52 "models",
53 "team_member_permissions",
54 "access_group_ids",
55 "policies",
56 "default_team_member_models",
57 ),
58)
59_ORG_LISTS: Final = _lists_as_json("o", ("models",))
60_PROJECT_LISTS: Final = _lists_as_json("p", ("models",))
61_PERMISSION_LISTS: Final = _lists_as_json(
62 "op",
63 (
64 "mcp_servers",
65 "mcp_access_groups",
66 "mcp_toolsets",
67 "blocked_tools",
68 "vector_stores",
69 "agents",
70 "agent_access_groups",
71 "models",
72 "search_tools",
73 "skills",
74 ),
75)
76_BUDGET_LISTS: Final = _lists_as_json("b", ("allowed_models",))
79def _budget_json(owner_alias: str) -> str:
80 return (
81 f"(SELECT to_jsonb(b) || jsonb_build_object({_BUDGET_LISTS}) "
82 f'FROM "LiteLLM_BudgetTable" b WHERE b.budget_id = {owner_alias}.budget_id)'
83 )
86def _permission_json(owner_alias: str) -> str:
87 return (
88 f"(SELECT to_jsonb(op) || jsonb_build_object({_PERMISSION_LISTS}) "
89 f'FROM "LiteLLM_ObjectPermissionTable" op WHERE op.object_permission_id = {owner_alias}.object_permission_id)'
90 )
93_SQL: Final = f"""
94SELECT
95 (
96 SELECT to_jsonb(u) || jsonb_build_object(
97 {_USER_LISTS},
98 'organization_memberships',
99 COALESCE((
100 SELECT jsonb_agg(to_jsonb(om)) FROM "LiteLLM_OrganizationMembership" om WHERE om.user_id = u.user_id
101 ), '[]'::jsonb)
102 )
103 FROM "LiteLLM_UserTable" u WHERE u.user_id = $1
104 ) AS user_row,
105 (
106 SELECT to_jsonb(t) || jsonb_build_object(
107 {_TEAM_LISTS},
108 'litellm_model_table', (
109 SELECT (to_jsonb(m) - 'aliases') || jsonb_build_object('model_aliases', m.aliases)
110 FROM "LiteLLM_ModelTable" m WHERE m.id = t.model_id
111 ),
112 'object_permission', {_permission_json("t")}
113 )
114 FROM "LiteLLM_TeamTable" t WHERE t.team_id = $2
115 ) AS team_row,
116 (
117 SELECT to_jsonb(tm) || jsonb_build_object('litellm_budget_table', {_budget_json("tm")})
118 FROM "LiteLLM_TeamMembership" tm WHERE tm.user_id = $3 AND tm.team_id = $2
119 ) AS membership_row,
120 (
121 SELECT to_jsonb(o) || jsonb_build_object(
122 {_ORG_LISTS},
123 'litellm_budget_table', {_budget_json("o")},
124 'object_permission', {_permission_json("o")}
125 )
126 FROM "LiteLLM_OrganizationTable" o WHERE o.organization_id = $4
127 ) AS organization_row,
128 (
129 SELECT to_jsonb(p) || jsonb_build_object(
130 {_PROJECT_LISTS},
131 'litellm_budget_table', {_budget_json("p")},
132 'object_permission', {_permission_json("p")}
133 )
134 FROM "LiteLLM_ProjectTable" p WHERE p.project_id = $5
135 ) AS project_row
136"""
139@dataclass(frozen=True, slots=True)
140class AuthObjectRefs:
141 """Ids of the objects a request's auth checks will read. ``None`` means not referenced."""
143 user_id: str | None = None
144 team_id: str | None = None
145 membership_user_id: str | None = None
146 organization_id: str | None = None
147 project_id: str | None = None
149 @classmethod
150 def from_token(cls, token: UserAPIKeyAuth) -> AuthObjectRefs:
151 has_membership: Final = token.team_id is not None and token.user_id is not None
152 return cls(
153 user_id=token.user_id,
154 team_id=token.team_id,
155 membership_user_id=token.user_id if has_membership else None,
156 organization_id=token.org_id,
157 project_id=token.project_id,
158 )
161class _InMemoryCache(Protocol):
162 def get_cache(self, key: str) -> object: ... 162 ↛ exitline 162 didn't return from function 'get_cache' because
163 def set_cache(self, key: str, value: object, *, ttl: float | None = ...) -> None: ... 163 ↛ exitline 163 didn't return from function 'set_cache' because
166@dataclass(frozen=True, slots=True)
167class _CacheEntry:
168 cache_key: str
169 row: _RowKind
170 model_type: type[BaseModel]
171 ttl: float | None
174def _iter_entries(refs: AuthObjectRefs, management_ttl: float) -> Iterator[_CacheEntry]:
175 if refs.user_id is not None:
176 yield _CacheEntry(refs.user_id, "user_row", LiteLLM_UserTable, management_ttl)
177 if refs.team_id is not None: 177 ↛ 178line 177 didn't jump to line 178 because the condition on line 177 was never true
178 yield _CacheEntry(f"team_id:{refs.team_id}", "team_row", LiteLLM_TeamTableCachedObj, management_ttl)
179 if refs.team_id is not None and refs.membership_user_id is not None: 179 ↛ 180line 179 didn't jump to line 180 because the condition on line 179 was never true
180 yield _CacheEntry(
181 team_membership_auth_cache_key(team_id=refs.team_id, user_id=refs.membership_user_id),
182 "membership_row",
183 LiteLLM_TeamMembership,
184 _TEAM_MEMBERSHIP_AUTH_TTL,
185 )
186 yield _CacheEntry(
187 team_membership_reservation_cache_key(user_id=refs.membership_user_id, team_id=refs.team_id),
188 "membership_row",
189 LiteLLM_TeamMembership,
190 None,
191 )
192 if refs.organization_id is not None: 192 ↛ 193line 192 didn't jump to line 193 because the condition on line 192 was never true
193 yield _CacheEntry(
194 f"org_id:{refs.organization_id}", "organization_row", LiteLLM_OrganizationTable, DEFAULT_IN_MEMORY_TTL
195 )
196 yield _CacheEntry(
197 f"org_id:{refs.organization_id}:with_budget",
198 "organization_row",
199 LiteLLM_OrganizationTable,
200 DEFAULT_IN_MEMORY_TTL,
201 )
202 if refs.project_id is not None: 202 ↛ 203line 202 didn't jump to line 203 because the condition on line 202 was never true
203 yield _CacheEntry(f"project_id:{refs.project_id}", "project_row", LiteLLM_ProjectTableCachedObj, management_ttl)
206def _entries(refs: AuthObjectRefs, cache: UserApiKeyCache) -> tuple[_CacheEntry, ...]:
207 return tuple(_iter_entries(refs, get_management_object_ttl(cache)))
210def _missing_in_memory(entries: Sequence[_CacheEntry], memory: _InMemoryCache) -> tuple[_CacheEntry, ...]:
211 return tuple(entry for entry in entries if memory.get_cache(key=entry.cache_key) is None)
214def _set_in_memory(memory: _InMemoryCache, cache_key: str, value: object, ttl: float | None) -> None:
215 if ttl is None: 215 ↛ 216line 215 didn't jump to line 216 because the condition on line 215 was never true
216 memory.set_cache(key=cache_key, value=value)
217 else:
218 memory.set_cache(key=cache_key, value=value, ttl=ttl)
221async def _fill_from_redis(entries: Sequence[_CacheEntry], redis_cache: RedisCache, memory: _InMemoryCache) -> None:
222 if not entries:
223 return
224 found: Final = _RowValues.validate_python(
225 await redis_cache.async_batch_get_cache(key_list=sorted(entry.cache_key for entry in entries)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API
226 )
227 for entry, value in ((entry, found.get(entry.cache_key)) for entry in entries):
228 if value is not None:
229 _set_in_memory(memory, entry.cache_key, value, entry.ttl)
232def _validate_row(
233 row_value: object, model_type: type[BaseModel], row: _RowKind, refreshed_at: float
234) -> BaseModel | None:
235 if row_value is None: 235 ↛ 236line 235 didn't jump to line 236 because the condition on line 235 was never true
236 return None
237 try:
238 columns: Final = _RowValues.validate_python(row_value)
239 if row in _REFRESH_STAMPED_ROWS: 239 ↛ 240line 239 didn't jump to line 240 because the condition on line 239 was never true
240 stamped: Final = {**columns, "last_refreshed_at": refreshed_at} # mutable-ok: validators write into it
241 return model_type.model_validate(stamped)
242 return model_type.model_validate(columns)
243 except ValidationError as e:
244 verbose_proxy_logger.warning("auth prefetch: %s did not validate as %s: %s", row, model_type.__name__, e)
245 return None
248async def _fetch_rows(
249 refs: AuthObjectRefs, kinds: frozenset[_RowKind], prisma_client: PrismaClient
250) -> Mapping[str, object]:
251 row: Final[object] = await prisma_client.db.query_first( # pyright: ignore[reportAny] # prisma types query_first as Any
252 _SQL,
253 refs.user_id if "user_row" in kinds else None,
254 refs.team_id if kinds & _TEAM_BOUND_ROWS else None,
255 refs.membership_user_id if "membership_row" in kinds else None,
256 refs.organization_id if "organization_row" in kinds else None,
257 refs.project_id if "project_row" in kinds else None,
258 )
259 return _RowValues.validate_python(row) if row is not None else _NO_ROWS
262async def _write_back(entries: Sequence[tuple[_CacheEntry, BaseModel]], cache: UserApiKeyCache) -> None:
263 payloads: Final = tuple(
264 (entry.cache_key, CacheCodec.serialize(value, model_type=entry.model_type), entry.ttl)
265 for entry, value in entries
266 )
267 memory: Final[_InMemoryCache] = cache.in_memory_cache
268 for cache_key, payload, ttl in payloads:
269 _set_in_memory(memory, cache_key, payload, cache.default_in_memory_ttl if ttl is None else ttl)
270 if cache.redis_cache is not None: 270 ↛ 271line 270 didn't jump to line 271 because the condition on line 270 was never true
271 await cache.redis_cache.async_set_cache_pipeline_with_ttls(payloads)
274async def _fill_from_db(
275 refs: AuthObjectRefs, entries: Sequence[_CacheEntry], cache: UserApiKeyCache, prisma_client: PrismaClient
276) -> None:
277 if not entries:
278 return
279 model_for: Final[Mapping[_RowKind, type[BaseModel]]] = MappingProxyType(
280 {entry.row: entry.model_type for entry in entries}
281 )
282 rows: Final = await _fetch_rows(refs, frozenset(model_for), prisma_client)
283 refreshed_at: Final = time.time()
284 objects: Final[Mapping[_RowKind, BaseModel | None]] = MappingProxyType(
285 {row: _validate_row(rows.get(row), model_type, row, refreshed_at) for row, model_type in model_for.items()}
286 )
287 writes: Final = tuple((entry, value) for entry in entries if (value := objects[entry.row]) is not None)
288 if writes: 288 ↛ exitline 288 didn't return from function '_fill_from_db' because the condition on line 288 was always true
289 await _write_back(writes, cache)
292async def prefetch_auth_objects(
293 refs: AuthObjectRefs,
294 user_api_key_cache: UserApiKeyCache,
295 prisma_client: PrismaClient | None,
296) -> None:
297 """Best effort: any failure leaves the per-object getters to fetch as before."""
298 try:
299 memory: Final[_InMemoryCache] = user_api_key_cache.in_memory_cache
300 missing: Final = _missing_in_memory(_entries(refs, user_api_key_cache), memory)
301 if user_api_key_cache.redis_cache is not None: 301 ↛ 302line 301 didn't jump to line 302 because the condition on line 301 was never true
302 await _fill_from_redis(missing, user_api_key_cache.redis_cache, memory)
303 if prisma_client is None: 303 ↛ 304line 303 didn't jump to line 304 because the condition on line 303 was never true
304 return
305 await _fill_from_db(refs, _missing_in_memory(missing, memory), user_api_key_cache, prisma_client)
306 except Exception as e: # noqa: BLE001 # warm-up only; the getters enforce and fail closed on their own
307 verbose_proxy_logger.warning("auth prefetch skipped, falling back to per-object lookups: %s", e)