mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
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:
parent
8a4fae0e17
commit
ed90ff4a39
5 changed files with 739 additions and 46 deletions
|
|
@ -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
|
||||
|
|
|
|||
267
litellm/proxy/auth/resolvers/grants.py
Normal file
267
litellm/proxy/auth/resolvers/grants.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
201
tests/test_litellm/proxy/auth/test_resolvers_grants.py
Normal file
201
tests/test_litellm/proxy/auth/test_resolvers_grants.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue