This commit is contained in:
ryan-crabbe-berri 2026-09-12 08:22:28 -04:00 committed by GitHub
commit 4b5da72aea
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 949 additions and 51 deletions

View file

@ -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())

View file

@ -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

View file

@ -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,
)

View file

@ -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)

View file

@ -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, ...).

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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