mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
Merge 0f10c06241 into 9071ca503e
This commit is contained in:
commit
4b5da72aea
10 changed files with 949 additions and 51 deletions
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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, ...).
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue