mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
9071ca503e
commit
f2e0a5db1e
13 changed files with 1513 additions and 19 deletions
|
|
@ -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,
|
||||
|
|
|
|||
307
litellm/proxy/auth/auth_object_prefetch.py
Normal file
307
litellm/proxy/auth/auth_object_prefetch.py
Normal 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)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
130
litellm/proxy/spend_tracking/spend_counter_batch.py
Normal file
130
litellm/proxy/spend_tracking/spend_counter_batch.py
Normal 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))
|
||||
0
tests/proxy_behavior/auth/__init__.py
Normal file
0
tests/proxy_behavior/auth/__init__.py
Normal file
20
tests/proxy_behavior/auth/conftest.py
Normal file
20
tests/proxy_behavior/auth/conftest.py
Normal 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()
|
||||
174
tests/proxy_behavior/auth/test_auth_object_prefetch.py
Normal file
174
tests/proxy_behavior/auth/test_auth_object_prefetch.py
Normal 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})
|
||||
|
|
@ -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)),
|
||||
]
|
||||
|
|
|
|||
338
tests/test_litellm/proxy/auth/test_auth_object_prefetch.py
Normal file
338
tests/test_litellm/proxy/auth/test_auth_object_prefetch.py
Normal 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
|
||||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
Loading…
Add table
Reference in a new issue