diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 18095aaafb4..f46cb66f7d1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -72,6 +72,7 @@ from litellm.proxy.auth.budget_throttle import ( ) from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation +from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, @@ -80,6 +81,7 @@ from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import ( END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL, MODEL_ACCESS_GROUP_REGISTRY_OVERFLOW_SENTINEL, + NO_TEAM_MEMBERSHIP_SENTINEL, TAG_REGISTRY_OVERFLOW_SENTINEL, UserApiKeyCache, end_user_cache_key, @@ -2163,10 +2165,10 @@ async def get_team_membership( _key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id) # check if in cache - cached_membership_obj: Final = await user_api_key_cache.async_get_cache( - key=_key, - model_type=LiteLLM_TeamMembership, - ) + cached: Final[object] = await user_api_key_cache.async_get_cache(key=_key) + if cached == NO_TEAM_MEMBERSHIP_SENTINEL: + return None + cached_membership_obj: Final = CacheCodec.deserialize(cached, model_type=LiteLLM_TeamMembership) if cached_membership_obj is not None: return cached_membership_obj @@ -2178,6 +2180,11 @@ async def get_team_membership( ) if response is None: + await user_api_key_cache.async_set_cache( + key=_key, + value=NO_TEAM_MEMBERSHIP_SENTINEL, + ttl=get_management_object_ttl(user_api_key_cache), + ) return None _response: Final = LiteLLM_TeamMembership.model_validate(response.dict()) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 69091ee8344..4304542fc83 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -52,6 +52,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import can_team_access_model +from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.team_grants import team_model_aliases from litellm.proxy.common_utils.user_api_key_cache import ( @@ -1656,9 +1657,7 @@ class JWTAuthManager: ``get_user_object`` resolved a legacy row with a different ``user_id``, use that row's id; otherwise keep the claim. GH #26789. """ - if user_object is not None and user_object.user_id: - return user_object.user_id - return user_id + return canonical_user_id(user_id=user_id, user_object=user_object) @staticmethod async def get_objects( @@ -1725,22 +1724,23 @@ class JWTAuthManager: code=403, ) - user_object: LiteLLM_UserTable | None = None - if user_id: - user_object = ( - await get_user_object( - user_id=user_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - user_id_upsert=jwt_handler.is_upsert_user_id(valid_user_email=valid_user_email), - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - user_email=user_email, - sso_user_id=user_id, - ) - if user_id - else None - ) + user_object, team_membership_object, effective_user_id = await GrantResolver( + prisma_client, + user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + load_user=get_user_object, + load_team=get_team_object, + load_membership=get_team_membership, + ).resolve_identity( + UserLookup( + user_id=user_id, + user_email=user_email, + sso_user_id=user_id, + upsert=jwt_handler.is_upsert_user_id(valid_user_email=valid_user_email), + ), + team_id=team_id, + ) end_user_object: LiteLLM_EndUserTable | None = None if end_user_id: @@ -1757,37 +1757,12 @@ class JWTAuthManager: else None ) - # Rebind to resolved DB user_id for team_membership + auth_builder (GH #26789). - effective_user_id: Final = JWTAuthManager._canonical_user_id_from_db(user_id=user_id, user_object=user_object) - if effective_user_id != user_id: - verbose_proxy_logger.debug( - "JWT Auth: rebinding user_id %r -> DB user_id %r (email/sso match)", - user_id, - effective_user_id, - ) - user_id = effective_user_id - - team_membership_object: LiteLLM_TeamMembership | None = None - if user_id and team_id: - team_membership_object = ( - await get_team_membership( - user_id=user_id, - team_id=team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - if user_id and team_id - else None - ) - return ( user_object, org_object, end_user_object, team_membership_object, - user_id, + effective_user_id, ) @staticmethod diff --git a/litellm/proxy/auth/resolvers/grants.py b/litellm/proxy/auth/resolvers/grants.py new file mode 100644 index 00000000000..eb39d2a6812 --- /dev/null +++ b/litellm/proxy/auth/resolvers/grants.py @@ -0,0 +1,267 @@ +"""Load a caller's user row, team row, and team membership from the database and validate them together. + +The virtual-key path reads these off the combined-view SQL join. Every other credential (an IdP JWT, a +``lite login`` session token) carries only identifiers, or a snapshot of grants taken when it was minted, so +it has to read the live rows on each request. Both of those paths resolve the same rows with the same +membership rule, and ``GrantResolver`` is the one place that rule lives. +""" + +from __future__ import annotations + +from collections.abc import Coroutine, Mapping +from dataclasses import dataclass +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, NoReturn, Protocol, TypeAlias + +from fastapi import HTTPException, status +from pydantic import BaseModel, ValidationError +from pydantic.main import IncEx +from typing_extensions import assert_never + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import ( + LiteLLM_TeamMembership, + LiteLLM_TeamTableCachedObj, + LiteLLM_UserTable, + ProxyErrorTypes, + ProxyException, +) +from litellm.proxy.auth.auth_checks import ( + TeamNotFoundError, + UserNotFoundError, + get_team_membership, + get_team_object, + get_user_object, +) + +if TYPE_CHECKING: + from litellm.proxy._types import Span + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import PrismaClient, ProxyLogging + + +class UserLoader(Protocol): + def __call__( + self, + *, + user_id: str | None, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + user_id_upsert: bool, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging | None, + sso_user_id: str | None, + user_email: str | None, + ) -> Coroutine[object, object, LiteLLM_UserTable | None]: ... + + +class TeamLoader(Protocol): + def __call__( + self, + *, + team_id: str, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging | None, + ) -> Coroutine[object, object, LiteLLM_TeamTableCachedObj]: ... + + +class MembershipLoader(Protocol): + def __call__( + self, + *, + user_id: str, + team_id: str, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging | None, + ) -> Coroutine[object, object, LiteLLM_TeamMembership | None]: ... + + +@dataclass(frozen=True, slots=True) +class UserLookup: + """The user a credential names, plus the hints ``get_user_object`` may fall back to when the id alone + matches no row.""" + + user_id: str | None + user_email: str | None = None + sso_user_id: str | None = None + upsert: bool = False + + +@dataclass(frozen=True, slots=True) +class ResolvedGrants: + """The live rows behind a credential. ``effective_user_id`` is the DB row's id when a fuzzy match found a + legacy row under a different id (GH #26789), otherwise the id the credential named.""" + + user_object: LiteLLM_UserTable | None + team_object: LiteLLM_TeamTableCachedObj | None + team_membership: LiteLLM_TeamMembership | None + effective_user_id: str | None + + +@dataclass(frozen=True, slots=True) +class UserGone: + user_id: str + + +@dataclass(frozen=True, slots=True) +class TeamGone: + team_id: str + + +@dataclass(frozen=True, slots=True) +class NotAMember: + user_id: str + team_id: str + + +@dataclass(frozen=True, slots=True) +class LookupDegraded: + """A row could not be read for a reason that says nothing about the caller: the database is down or a + loader failed. The caller decides whether a grant it already holds may stand in.""" + + error: Exception + + +GrantDenial: TypeAlias = UserGone | TeamGone | NotAMember +GrantOutcome: TypeAlias = ResolvedGrants | GrantDenial | LookupDegraded + + +_MODELS_COLUMN: Final[Mapping[str, IncEx | bool]] = MappingProxyType({"models": True}) + + +class _UserModelColumn(BaseModel): + """``LiteLLM_UserTable.models`` is a bare ``list``; re-read it with the shape a token's ``models`` takes.""" + + models: tuple[str, ...] = () + + +def user_models(user_object: LiteLLM_UserTable) -> tuple[str, ...]: + try: + return _UserModelColumn.model_validate(user_object.model_dump(include=_MODELS_COLUMN)).models + except ValidationError: + return () + + +def canonical_user_id(user_id: str | None, user_object: LiteLLM_UserTable | None) -> str | None: + if user_object is not None and user_object.user_id: + return user_object.user_id + return user_id + + +def raise_public(denial: GrantDenial) -> NoReturn: + match denial: + case UserGone(user_id=user_id): + raise ProxyException( + message=f"Authentication Error, user '{user_id}' no longer exists.", + type=ProxyErrorTypes.auth_error, + param="user_id", + code=status.HTTP_401_UNAUTHORIZED, + ) + case TeamGone(team_id=team_id): + raise TeamNotFoundError(team_id=team_id) + case NotAMember(team_id=team_id): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"Team '{team_id}' is not in your team memberships.", + ) + case _: + assert_never(denial) + + +class GrantResolver: + """Reads the user, membership, and team rows for a credential through injected loaders. + + The loaders default to the shared ``auth_checks`` readers. A caller passes its own module's names for them + so the reads stay interceptable where that module's callers already intercept them. ``resolve_identity`` + is the JWT half: user and membership only, since the JWT builder selects the team itself and lets loader + errors surface as they are. ``resolve`` also reads the team row and applies the membership rule, which is + what a credential carrying a grant snapshot needs to refresh it. + """ + + def __init__( + self, + prisma_client: PrismaClient | None, + cache: UserApiKeyCache, + *, + parent_otel_span: Span | None = None, + proxy_logging_obj: ProxyLogging | None = None, + load_user: UserLoader = get_user_object, + load_team: TeamLoader = get_team_object, + load_membership: MembershipLoader = get_team_membership, + ) -> None: + self._prisma = prisma_client + self._cache = cache + self._parent_otel_span = parent_otel_span + self._proxy_logging_obj = proxy_logging_obj + self._load_user = load_user + self._load_team = load_team + self._load_membership = load_membership + + async def resolve_identity( + self, lookup: UserLookup, team_id: str | None + ) -> tuple[LiteLLM_UserTable | None, LiteLLM_TeamMembership | None, str | None]: + user_object: Final = await self._user(lookup) if lookup.user_id else None + effective_user_id: Final = canonical_user_id(lookup.user_id, user_object) + if effective_user_id != lookup.user_id: + verbose_proxy_logger.debug( + "Auth: rebinding user_id %r -> DB user_id %r (email/sso match)", + lookup.user_id, + effective_user_id, + ) + membership: Final = ( + await self._membership(user_id=effective_user_id, team_id=team_id) + if effective_user_id and team_id + else None + ) + return user_object, membership, effective_user_id + + async def resolve(self, lookup: UserLookup, team_id: str | None) -> GrantOutcome: + try: + user_object, membership, effective_user_id = await self.resolve_identity(lookup, team_id) + except UserNotFoundError: + return UserGone(user_id=lookup.user_id or "") + except Exception as error: + return LookupDegraded(error=error) + if team_id is None: + return ResolvedGrants(user_object, None, membership, effective_user_id) + if user_object is not None and team_id not in user_object.teams: + return NotAMember(user_id=user_object.user_id, team_id=team_id) + try: + team_object: Final = await self._load_team( + team_id=team_id, + prisma_client=self._prisma, + user_api_key_cache=self._cache, + parent_otel_span=self._parent_otel_span, + proxy_logging_obj=self._proxy_logging_obj, + ) + except TeamNotFoundError: + return TeamGone(team_id=team_id) + except Exception as error: + return LookupDegraded(error=error) + return ResolvedGrants(user_object, team_object, membership, effective_user_id) + + async def _user(self, lookup: UserLookup) -> LiteLLM_UserTable | None: + return await self._load_user( + user_id=lookup.user_id, + prisma_client=self._prisma, + user_api_key_cache=self._cache, + user_id_upsert=lookup.upsert, + parent_otel_span=self._parent_otel_span, + proxy_logging_obj=self._proxy_logging_obj, + user_email=lookup.user_email, + sso_user_id=lookup.sso_user_id, + ) + + async def _membership(self, user_id: str, team_id: str) -> LiteLLM_TeamMembership | None: + return await self._load_membership( + user_id=user_id, + team_id=team_id, + prisma_client=self._prisma, + user_api_key_cache=self._cache, + parent_otel_span=self._parent_otel_span, + proxy_logging_obj=self._proxy_logging_obj, + ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 20ab9904f46..00d67a14565 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -55,6 +55,7 @@ from litellm.proxy.auth.auth_checks import ( get_jwt_key_mapping_object, get_object_permission, get_project_object, + get_team_membership, get_team_object, get_user_object, is_valid_fallback_model, @@ -80,6 +81,14 @@ from litellm.proxy.auth.network import TrustedProxyConfig, resolve_network_conte from litellm.proxy.auth.oauth2_check import Oauth2Handler from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request from litellm.proxy.auth.resolvers import CredentialRef, Principal +from litellm.proxy.auth.resolvers.grants import ( + GrantResolver, + LookupDegraded, + ResolvedGrants, + UserLookup, + raise_public, + user_models, +) from litellm.proxy.auth.resolvers.store import IdentityStore from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.team_grants import team_grants @@ -1222,6 +1231,52 @@ async def _record_unparsable_body_failure( verbose_proxy_logger.exception("Failed to log the request rejected for an unparsable body: %s", e) +async def _refresh_session_token_grants( + valid_token: UserAPIKeyAuth, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, +) -> UserAPIKeyAuth: + """Rebuild a ``lite login`` session token's grants from the live user and team rows. + + The blob only proves who logged in and which team they picked. Team models, aliases, the user's own model + list, and their role are re-read every request, so a `/team/update` or a demotion shows up without a + re-login, and a user removed from the team or deleted outright is refused. When a row cannot be read for + a reason unrelated to the caller, the minted grants stand in exactly as they did before this refresh. + """ + outcome: Final = await GrantResolver( + prisma_client, + user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + load_user=get_user_object, + load_team=get_team_object, + load_membership=get_team_membership, + ).resolve(UserLookup(user_id=valid_token.user_id), team_id=valid_token.team_id) + match outcome: + case ResolvedGrants( + user_object=LiteLLM_UserTable() as user_object, team_object=team_object, team_membership=team_membership + ): + return UserAPIKeyAuth.model_validate( + MappingProxyType( + { + **valid_token.model_dump(exclude_none=True), + **team_grants(team_object, team_membership, user_object.user_id), + "user_role": _get_user_role(user_object), + "models": () if team_object is not None else user_models(user_object), + } + ) + ) + case ResolvedGrants(): + return valid_token + case LookupDegraded(error=error): + verbose_proxy_logger.debug("Session token grants not refreshed, keeping minted grants: %s", error) + return valid_token + case _: + raise_public(outcome) + + async def _resolve_object_permission_for_unresolvable_team( object_permission_id: str | None, prisma_client: PrismaClient | None, @@ -1766,6 +1821,15 @@ async def _user_api_key_auth_builder( ): valid_token = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(api_key) + if valid_token is not None and valid_token.is_session_token and prisma_client is not None: + valid_token = await _refresh_session_token_grants( # rebind-ok: later checks read this name + valid_token=valid_token, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + if ( valid_token is not None and isinstance(valid_token, UserAPIKeyAuth) diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index cb72088ee4a..1c7a379897f 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -336,6 +336,14 @@ def team_membership_reservation_cache_key(user_id: str, team_id: str) -> str: return f"team_membership:{user_id}:{team_id}" +#: Cached under ``team_membership_reservation_cache_key`` when a member has no ``LiteLLM_TeamMembership`` +#: row, so a session-token member without a per-member budget costs no DB read per request. Lives beside +#: the key builder because it is part of the same cache protocol: every reader of the key must know that +#: a plain string here means "no row", distinct from a serialized membership. The two budget readers +#: already treat a non-model value as "no row", so they need no change to stay correct. +NO_TEAM_MEMBERSHIP_SENTINEL: Final = "__no_team_membership__" + + def get_management_object_ttl(cache: DualCache) -> float: """ In-memory TTL for management-object cache writes (keys, teams, users, budgets, ...). diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index e9e37540dd8..0b7f69bbb7f 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -14,7 +14,7 @@ import copy import json import math import traceback -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from collections.abc import Set as AbstractSet from datetime import datetime, timezone from types import MappingProxyType @@ -1903,6 +1903,39 @@ def validate_team_org_change( return True +def _member_user_ids(members_with_roles: Sequence[dict[str, object]]) -> tuple[str, ...]: + """Extract the string ``user_id`` of each team member, dropping rows without one. + + ``members_with_roles`` is a Prisma-deserialized JSON column, so its ``user_id`` is typed + ``object``; the ``isinstance`` narrows it to the ``str`` ``invalidate_team_member_spend_state`` needs. + """ + return tuple(user_id for member in members_with_roles if isinstance((user_id := member.get("user_id")), str)) + + +async def _evict_created_membership_caches( + user_ids: Iterable[str], + team_id: str, + user_api_key_cache: UserApiKeyCache, +) -> None: + """Evict the ``get_team_membership`` negative-cache sentinel for members whose row was just created. + + A session-token request caches ``NO_TEAM_MEMBERSHIP_SENTINEL`` for a member with no + ``LiteLLM_TeamMembership`` row. When a create path (``/team/member_add`` or the ``/team/update`` + budget backfill) later writes that row with a per-member budget, the stale sentinel keeps the + member's budget unenforced until the membership cache TTL expires, so it must be evicted here. + """ + await asyncio.gather( + *( + invalidate_team_member_spend_state( + user_id=user_id, + team_id=team_id, + user_api_key_cache=user_api_key_cache, + ) + for user_id in user_ids + ) + ) + + @router.post("/team/update", tags=["team management"], dependencies=[Depends(user_api_key_auth)]) @management_endpoint_wrapper async def update_team( @@ -2239,6 +2272,11 @@ async def update_team( team_member_budget_id=_backfill_budget_id, prisma_client=prisma_client, ) + await _evict_created_membership_caches( + user_ids=_member_user_ids(existing_team_row.members_with_roles), + team_id=data.team_id, + user_api_key_cache=user_api_key_cache, + ) elif _team_member_fields_in_request: updated_kv = await TeamMemberBudgetHandler.clear_team_member_budget_fields( team_table=existing_team_row, @@ -3191,6 +3229,12 @@ async def team_member_add( litellm_proxy_admin_name=litellm_proxy_admin_name, ) + await _evict_created_membership_caches( + user_ids=(tm.user_id for tm in updated_team_memberships), + team_id=data.team_id, + user_api_key_cache=user_api_key_cache, + ) + _emit_team_members_metric(complete_team_data) await _create_team_member_add_audit_logs( diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index dc614d18662..a027f4a9c15 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -6354,6 +6354,106 @@ async def test_get_team_membership_db_fetch_returns_validated_membership(): assert result.spend == 1.5 +@pytest.mark.asyncio +async def test_get_team_membership_negative_caches_a_missing_row(): + """ + Regression (LIT-7358): a member with no LiteLLM_TeamMembership row is the common lite-login case, + and the session-token refresh reads this loader on every request. Before the fix a missing row + returned None without caching, so every request re-queried the DB. The miss must be cached so the + second request serves from cache and never touches the DB. + """ + from litellm.proxy.auth.auth_checks import get_team_membership + from litellm.proxy.common_utils.user_api_key_cache import ( + NO_TEAM_MEMBERSHIP_SENTINEL, + team_membership_reservation_cache_key, + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) + + cache = UserApiKeyCache() + + first = await get_team_membership( + user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache + ) + second = await get_team_membership( + user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache + ) + + assert first is None + assert second is None + mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once() + cached = await cache.async_get_cache( + key=team_membership_reservation_cache_key(user_id="u-1", team_id="t-1") + ) + assert cached == NO_TEAM_MEMBERSHIP_SENTINEL + + +@pytest.mark.asyncio +async def test_get_team_membership_reads_sentinel_as_no_membership_not_a_model(): + """ + The negative-cache sentinel is a plain string sharing the key a serialized membership uses. + A pre-seeded sentinel must read back as None (no DB read), never be mistaken for a membership. + """ + from litellm.proxy.auth.auth_checks import get_team_membership + from litellm.proxy.common_utils.user_api_key_cache import ( + NO_TEAM_MEMBERSHIP_SENTINEL, + team_membership_reservation_cache_key, + ) + + cache = UserApiKeyCache() + await cache.async_set_cache( + key=team_membership_reservation_cache_key(user_id="u-1", team_id="t-1"), + value=NO_TEAM_MEMBERSHIP_SENTINEL, + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) + + result = await get_team_membership( + user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache + ) + + assert result is None + mock_prisma_client.db.litellm_teammembership.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sentinel(): + """ + A member who later gains a per-member budget writes a membership row and calls + invalidate_team_member_spend_state. That must drop a cached "no membership" sentinel so the next + request re-reads the DB and honors the new budget instead of serving the stale miss until TTL. + """ + from litellm.proxy.auth.auth_checks import get_team_membership, invalidate_team_member_spend_state + from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key + + cache = UserApiKeyCache() + membership_row = MagicMock() + membership_row.dict = lambda: {"user_id": "u-1", "team_id": "t-1", "spend": 0.0} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=[None, membership_row]) + + before = await get_team_membership( + user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache + ) + assert before is None + + await invalidate_team_member_spend_state(user_id="u-1", team_id="t-1", user_api_key_cache=cache) + assert ( + await cache.async_get_cache(key=team_membership_reservation_cache_key(user_id="u-1", team_id="t-1")) + is None + ) + + after = await get_team_membership( + user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache + ) + assert after is not None + assert after.user_id == "u-1" + assert mock_prisma_client.db.litellm_teammembership.find_unique.await_count == 2 + + @pytest.mark.asyncio async def test_get_access_object_db_fetch_returns_validated_access_group(): from litellm.proxy._types import LiteLLM_AccessGroupTable diff --git a/tests/test_litellm/proxy/auth/test_resolvers_grants.py b/tests/test_litellm/proxy/auth/test_resolvers_grants.py new file mode 100644 index 00000000000..3f6d943bf98 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_resolvers_grants.py @@ -0,0 +1,201 @@ +from fastapi import HTTPException +import pytest + +from litellm.proxy._types import ( + LiteLLM_TeamMembership, + LiteLLM_TeamTableCachedObj, + LiteLLM_UserTable, + ProxyException, +) +from litellm.proxy.auth.auth_checks import TeamNotFoundError, UserNotFoundError +from litellm.proxy.auth.resolvers.grants import ( + GrantResolver, + LookupDegraded, + NotAMember, + ResolvedGrants, + TeamGone, + UserGone, + UserLookup, + raise_public, + user_models, +) + +USER_ID = "user-1" +TEAM_ID = "team-1" + + +class _Loaders: + """Fake row readers standing in for the ``auth_checks`` loaders, recording every call they receive.""" + + def __init__(self, *, user=None, team=None, membership=None, user_error=None, team_error=None): + self._user = user + self._team = team + self._membership = membership + self._user_error = user_error + self._team_error = team_error + self.user_calls = [] + self.team_calls = [] + self.membership_calls = [] + + async def load_user(self, **kwargs): + self.user_calls.append(kwargs) + if self._user_error is not None: + raise self._user_error + return self._user + + async def load_team(self, **kwargs): + self.team_calls.append(kwargs) + if self._team_error is not None: + raise self._team_error + return self._team + + async def load_membership(self, **kwargs): + self.membership_calls.append(kwargs) + return self._membership + + def resolver(self) -> GrantResolver: + return GrantResolver( + object(), + object(), + load_user=self.load_user, + load_team=self.load_team, + load_membership=self.load_membership, + ) + + +def _user(teams=(TEAM_ID,), user_id=USER_ID) -> LiteLLM_UserTable: + return LiteLLM_UserTable(user_id=user_id, user_role="internal_user", teams=list(teams), models=["gpt-5.5"]) + + +def _team(models=("gpt-5.5",)) -> LiteLLM_TeamTableCachedObj: + return LiteLLM_TeamTableCachedObj(team_id=TEAM_ID, team_alias="alias", models=list(models)) + + +async def test_resolve_returns_live_rows_for_a_member(): + membership = LiteLLM_TeamMembership(user_id=USER_ID, team_id=TEAM_ID, spend=1.5) + loaders = _Loaders(user=_user(), team=_team(models=("new-a", "new-b")), membership=membership) + + outcome = await loaders.resolver().resolve(UserLookup(user_id=USER_ID), team_id=TEAM_ID) + + assert outcome == ResolvedGrants( + user_object=_user(), + team_object=_team(models=("new-a", "new-b")), + team_membership=membership, + effective_user_id=USER_ID, + ) + assert loaders.team_calls[0]["team_id"] == TEAM_ID + assert loaders.membership_calls[0]["user_id"] == USER_ID + assert loaders.membership_calls[0]["team_id"] == TEAM_ID + + +async def test_resolve_denies_a_user_removed_from_the_team_without_reading_the_team(): + loaders = _Loaders(user=_user(teams=("other-team",)), team=_team()) + + outcome = await loaders.resolver().resolve(UserLookup(user_id=USER_ID), team_id=TEAM_ID) + + assert outcome == NotAMember(user_id=USER_ID, team_id=TEAM_ID) + assert loaders.team_calls == [] + + +async def test_resolve_reports_a_deleted_user(): + loaders = _Loaders(user_error=UserNotFoundError(user_id=USER_ID), team=_team()) + + outcome = await loaders.resolver().resolve(UserLookup(user_id=USER_ID), team_id=TEAM_ID) + + assert outcome == UserGone(user_id=USER_ID) + assert loaders.team_calls == [] + + +async def test_resolve_reports_a_deleted_team(): + loaders = _Loaders(user=_user(), team_error=TeamNotFoundError(team_id=TEAM_ID)) + + outcome = await loaders.resolver().resolve(UserLookup(user_id=USER_ID), team_id=TEAM_ID) + + assert outcome == TeamGone(team_id=TEAM_ID) + + +@pytest.mark.parametrize( + "loaders", + [ + _Loaders(user_error=Exception("No db connected")), + _Loaders(user=_user(), team_error=HTTPException(status_code=500, detail="db timeout")), + ], + ids=["user-read-failed", "team-read-failed"], +) +async def test_resolve_marks_an_unreadable_row_as_degraded_not_denied(loaders): + outcome = await loaders.resolver().resolve(UserLookup(user_id=USER_ID), team_id=TEAM_ID) + + assert isinstance(outcome, LookupDegraded) + + +async def test_resolve_without_a_team_skips_team_and_membership_reads(): + loaders = _Loaders(user=_user(teams=())) + + outcome = await loaders.resolver().resolve(UserLookup(user_id=USER_ID), team_id=None) + + assert outcome == ResolvedGrants( + user_object=_user(teams=()), team_object=None, team_membership=None, effective_user_id=USER_ID + ) + assert loaders.team_calls == [] + assert loaders.membership_calls == [] + + +async def test_resolve_identity_reads_membership_under_the_matched_rows_id(): + legacy_uuid = "bb8ab11f-09aa-47ae-b063-6e80506ac3bc" + loaders = _Loaders(user=_user(user_id=legacy_uuid)) + + user_object, _membership, effective_user_id = await loaders.resolver().resolve_identity( + UserLookup(user_id="matt@example.com", user_email="matt@example.com", sso_user_id="matt@example.com"), + team_id=TEAM_ID, + ) + + assert user_object is not None and user_object.user_id == legacy_uuid + assert effective_user_id == legacy_uuid + assert loaders.membership_calls[0]["user_id"] == legacy_uuid + assert loaders.user_calls[0]["user_email"] == "matt@example.com" + + +async def test_resolve_identity_without_a_user_id_reads_nothing(): + loaders = _Loaders(user=_user()) + + outcome = await loaders.resolver().resolve_identity(UserLookup(user_id=None), team_id=TEAM_ID) + + assert outcome == (None, None, None) + assert loaders.user_calls == [] + assert loaders.membership_calls == [] + + +async def test_resolve_identity_lets_loader_errors_surface(): + loaders = _Loaders(user_error=UserNotFoundError(user_id=USER_ID)) + + with pytest.raises(UserNotFoundError): + await loaders.resolver().resolve_identity(UserLookup(user_id=USER_ID), team_id=None) + + +def test_raise_public_maps_a_deleted_user_to_401(): + with pytest.raises(ProxyException) as exc_info: + raise_public(UserGone(user_id=USER_ID)) + assert exc_info.value.code == "401" + assert USER_ID in exc_info.value.message + + +def test_raise_public_maps_a_removed_member_to_403(): + with pytest.raises(HTTPException) as exc_info: + raise_public(NotAMember(user_id=USER_ID, team_id=TEAM_ID)) + assert exc_info.value.status_code == 403 + assert TEAM_ID in str(exc_info.value.detail) + + +def test_raise_public_maps_a_deleted_team_to_404(): + with pytest.raises(TeamNotFoundError) as exc_info: + raise_public(TeamGone(team_id=TEAM_ID)) + assert exc_info.value.status_code == 404 + + +@pytest.mark.parametrize( + ("stored", "expected"), + [(["gpt-5.5", "claude-opus-5"], ("gpt-5.5", "claude-opus-5")), ([], ()), ([{"not": "a model"}], ())], + ids=["models", "empty", "unusable-column"], +) +def test_user_models_reads_the_column_as_a_tuple_of_names(stored, expected): + assert user_models(LiteLLM_UserTable(user_id=USER_ID, models=stored)) == expected 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..869e6d27fde 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 @@ -31,7 +31,7 @@ from litellm.proxy._types import ( JWTRoutingOverride, ) from litellm.proxy.auth.handle_jwt import JWTHandler -from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object +from litellm.proxy.auth.auth_checks import TeamNotFoundError, UserNotFoundError, get_key_object, _cache_key_object from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( _check_key_model_budget_with_fallback, @@ -6116,6 +6116,192 @@ async def test_non_admin_cli_session_token_reaches_production_auth_path(monkeypa assert result.is_session_token is True +SESSION_TEAM_ID = "team-abc" +SESSION_USER_ID = "member-1" + + +def _mint_session_token( + monkeypatch, + *, + role=LitellmUserRoles.INTERNAL_USER, + team_id=SESSION_TEAM_ID, + team_models=("stale-model",), + models=(), +): + """Mint a ``lite login`` token carrying the grants as they were at login time.""" + monkeypatch.delenv("EXPERIMENTAL_UI_LOGIN", raising=False) + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-cli-test") + from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken + + user_info = LiteLLM_UserTable( + user_id=SESSION_USER_ID, user_email="user@example.com", user_role=role.value, models=list(models) + ) + return ExperimentalUIJWTToken.get_cli_jwt_auth_token( + user_info, team_id=team_id, team_alias="stale-alias", team_models=list(team_models) + ) + + +def _session_user_row(*, teams=(SESSION_TEAM_ID,), role=LitellmUserRoles.INTERNAL_USER, models=()): + return LiteLLM_UserTable(user_id=SESSION_USER_ID, user_role=role.value, teams=list(teams), models=list(models)) + + +async def _authenticate_session_token_against_db( + cli_token, *, user_row=None, team_row=None, membership_row=None, user_error=None, team_error=None +): + """Drive the real builder for a session token with the DB row readers replaced by the given rows or + errors. Returns the ``_return_user_api_key_auth_obj`` mock so the caller can read the token it was + handed; a denial surfaces as the exception the builder raises.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import Request + from starlette.datastructures import URL + + attrs = _proxy_attrs_for_db_lookup() + attrs["prisma_client"].db.litellm_teammembership.find_first = AsyncMock(return_value=None) + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + assemble = AsyncMock(return_value=UserAPIKeyAuth(user_id=SESSION_USER_ID, is_session_token=True)) + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + with ( + patch( # test-quality-ok: the builder has no injection seam for its assembler yet + "litellm.proxy.auth.user_api_key_auth._return_user_api_key_auth_obj", assemble + ), + patch( # test-quality-ok: the builder reads its DB row loaders off module globals + "litellm.proxy.auth.user_api_key_auth.get_user_object", + AsyncMock(return_value=user_row, side_effect=user_error), + ), + patch( # test-quality-ok: the builder reads its DB row loaders off module globals + "litellm.proxy.auth.user_api_key_auth.get_team_object", + AsyncMock(return_value=team_row, side_effect=team_error), + ), + patch( # test-quality-ok: the builder reads its DB row loaders off module globals + "litellm.proxy.auth.user_api_key_auth.get_team_membership", + AsyncMock(return_value=membership_row), + ), + ): + await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {cli_token}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + return assemble + + +@pytest.mark.asyncio +async def test_session_token_reads_team_grants_from_the_live_team_row(monkeypatch): + """LIT-7358: a lite login token snapshots the team's models at login, so adding a model to the team did + nothing for that CLI until the user logged in again. The team row has to be re-read on every request.""" + from litellm.models.team import LiteLLM_ModelTable + from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj + + cli_token = _mint_session_token(monkeypatch, team_models=("stale-model",)) + live_team = LiteLLM_TeamTableCachedObj( + team_id=SESSION_TEAM_ID, + team_alias="renamed-team", + models=["gpt-5.5", "claude-opus-5"], + litellm_model_table=LiteLLM_ModelTable(model_aliases={"fast": "gpt-5.5"}, created_by="a", updated_by="a"), + ) + membership = LiteLLM_TeamMembership(user_id=SESSION_USER_ID, team_id=SESSION_TEAM_ID, spend=2.5) + + assemble = await _authenticate_session_token_against_db( + cli_token, user_row=_session_user_row(), team_row=live_team, membership_row=membership + ) + + token = assemble.call_args.kwargs["valid_token_dict"] + assert token["team_models"] == ["gpt-5.5", "claude-opus-5"] + assert token["team_alias"] == "renamed-team" + assert token["team_model_aliases"] == {"fast": "gpt-5.5"} + assert token["team_member_spend"] == 2.5 + assert token["is_session_token"] is True + + +@pytest.mark.asyncio +async def test_session_token_without_a_team_reads_models_from_the_live_user_row(monkeypatch): + cli_token = _mint_session_token(monkeypatch, team_id=None, team_models=(), models=("stale-model",)) + + assemble = await _authenticate_session_token_against_db( + cli_token, user_row=_session_user_row(teams=(), models=("gpt-5.5",)) + ) + + assert assemble.call_args.kwargs["valid_token_dict"]["models"] == ["gpt-5.5"] + + +@pytest.mark.asyncio +async def test_demoted_admin_session_token_loses_admin_on_the_next_request(monkeypatch): + """The role baked into the token used to send a former admin down the admin early return forever.""" + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + + cli_token = _mint_session_token(monkeypatch, role=LitellmUserRoles.PROXY_ADMIN) + + assemble = await _authenticate_session_token_against_db( + cli_token, + user_row=_session_user_row(role=LitellmUserRoles.INTERNAL_USER), + team_row=LiteLLM_TeamTableCachedObj(team_id=SESSION_TEAM_ID, models=["gpt-5.5"]), + ) + + assemble.assert_awaited_once() + assert assemble.call_args.kwargs["valid_token_dict"]["user_role"] == LitellmUserRoles.INTERNAL_USER + + +@pytest.mark.asyncio +async def test_session_token_is_refused_once_the_user_leaves_the_team(monkeypatch): + cli_token = _mint_session_token(monkeypatch) + + with pytest.raises(ProxyException) as exc_info: + await _authenticate_session_token_against_db(cli_token, user_row=_session_user_row(teams=("other-team",))) + + assert exc_info.value.code == str(status.HTTP_403_FORBIDDEN) + assert SESSION_TEAM_ID in exc_info.value.message + + +@pytest.mark.asyncio +async def test_session_token_is_refused_once_the_user_is_deleted(monkeypatch): + cli_token = _mint_session_token(monkeypatch) + + with pytest.raises(ProxyException) as exc_info: + await _authenticate_session_token_against_db(cli_token, user_error=UserNotFoundError(user_id=SESSION_USER_ID)) + + assert exc_info.value.code == str(status.HTTP_401_UNAUTHORIZED) + assert exc_info.value.type == ProxyErrorTypes.auth_error + + +@pytest.mark.asyncio +async def test_session_token_is_refused_once_the_team_is_deleted(monkeypatch): + cli_token = _mint_session_token(monkeypatch) + + with pytest.raises(ProxyException) as exc_info: + await _authenticate_session_token_against_db( + cli_token, user_row=_session_user_row(), team_error=TeamNotFoundError(team_id=SESSION_TEAM_ID) + ) + + assert exc_info.value.code == str(status.HTTP_404_NOT_FOUND) + + +@pytest.mark.asyncio +async def test_session_token_keeps_minted_grants_when_the_team_row_cannot_be_read(monkeypatch): + """A DB hiccup says nothing about the caller, so the grants minted at login stand for that request.""" + from fastapi import HTTPException + + cli_token = _mint_session_token(monkeypatch, team_models=("stale-model",)) + + assemble = await _authenticate_session_token_against_db( + cli_token, user_row=_session_user_row(), team_error=HTTPException(status_code=500, detail="db timeout") + ) + + token = assemble.call_args.kwargs["valid_token_dict"] + assert token["team_models"] == ["stale-model"] + assert token["team_alias"] == "stale-alias" + + @pytest.mark.asyncio async def test_cli_session_token_authenticates_when_jwt_auth_enabled_without_license(monkeypatch): """A lite login token is an encrypted (non-JWT) session blob. With diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 2fb496d6231..0e1831614ac 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -13941,6 +13941,52 @@ async def test_team_member_update_skips_invalidation_when_no_budget_fields_sent( assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:member-1:team-1") == 1.5 +@pytest.mark.asyncio +async def test_evict_created_membership_caches_drops_the_negative_sentinel(): + """ + Regression: a membership-create path (/team/member_add, the /team/update budget backfill) must + evict any cached "no membership" sentinel a prior session-token read left, so a per-member budget + attached at create time is enforced on the next request instead of after the membership cache TTL. + Uses a real cache so the assertion is that the sentinel is actually gone, not that a mock was called. + """ + from litellm.proxy.common_utils.user_api_key_cache import ( + NO_TEAM_MEMBERSHIP_SENTINEL, + UserApiKeyCache, + team_membership_reservation_cache_key, + ) + from litellm.proxy.management_endpoints.team_endpoints import _evict_created_membership_caches + + cache = UserApiKeyCache() + kept_key = team_membership_reservation_cache_key(user_id="carol", team_id="team-eviction") + evicted_key = team_membership_reservation_cache_key(user_id="bob", team_id="team-eviction") + await cache.async_set_cache(key=kept_key, value=NO_TEAM_MEMBERSHIP_SENTINEL) + await cache.async_set_cache(key=evicted_key, value=NO_TEAM_MEMBERSHIP_SENTINEL) + + await _evict_created_membership_caches(user_ids=("bob",), team_id="team-eviction", user_api_key_cache=cache) + + assert await cache.async_get_cache(key=evicted_key) is None + assert await cache.async_get_cache(key=kept_key) == NO_TEAM_MEMBERSHIP_SENTINEL + + +def test_member_user_ids_keeps_only_string_user_ids(): + """ + The /team/update backfill feeds Prisma-deserialized member dicts here; a row can be missing + user_id or carry a non-string value. Only real string ids may reach invalidate_team_member_spend_state, + so those get eviction and the malformed rows are dropped rather than crashing the update. + """ + from litellm.proxy.management_endpoints.team_endpoints import _member_user_ids + + members = [ + {"user_id": "alice", "role": "admin"}, + {"role": "user"}, + {"user_id": None, "role": "user"}, + {"user_id": 123, "role": "user"}, + {"user_id": "bob", "role": "user"}, + ] + + assert _member_user_ids(members) == ("alice", "bob") + + def _team_spend_by_user_team(team_id: str, team_alias: str, member: Member, permissions: list[str]) -> MagicMock: team = MagicMock(spec=LiteLLM_TeamTable) team.team_id = team_id