fix(auth): refresh lite login session token grants from the live user and team rows

A lite login token carried a snapshot of the team's models, aliases and the
user's role taken at login, so team or role changes never reached that CLI
until the user logged in again.

Pull the user, membership and team row loading that the JWT path did inline in
JWTAuthManager.get_objects into a GrantResolver under auth/resolvers, and have
the session token branch of the auth builder resolve the same rows on every
request. A user removed from the team now gets 403, a deleted user 401, and a
demoted admin no longer takes the admin early return.
This commit is contained in:
ryan-crabbe-berri 2026-09-10 18:33:49 -07:00
parent 8a4fae0e17
commit ed90ff4a39
5 changed files with 739 additions and 46 deletions

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

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