fix: move remaining stuff

This commit is contained in:
Yassin Kortam 2026-06-08 12:14:00 -07:00
parent cfc57d709f
commit 043d2d65b1
22 changed files with 1629 additions and 166 deletions

View file

@ -10,6 +10,10 @@ The public surface is small on purpose; downstream code should depend on
extractor internals.
"""
from litellm.identity.cache import IdentityCache
from litellm.identity.jwt import build_user_api_key_auth_from_jwt_result
from litellm.identity.oauth2 import build_user_api_key_auth_from_oauth2_response
from litellm.identity.runtime import get_identity_cache
from litellm.identity.context import (
AuditInfo,
ClientInfo,
@ -24,16 +28,25 @@ from litellm.identity.principal import (
SSOPrincipal,
ServiceAccountPrincipal,
)
from litellm.identity.resolver import resolve_identity, resolve_user_api_key_auth
from litellm.identity.store import load_identity
__all__ = [
"AnonymousPrincipal",
"ApiKeyPrincipal",
"AuditInfo",
"ClientInfo",
"IdentityCache",
"IdentityContext",
"JWTPrincipal",
"Principal",
"RequestIds",
"SSOPrincipal",
"ServiceAccountPrincipal",
"build_user_api_key_auth_from_jwt_result",
"build_user_api_key_auth_from_oauth2_response",
"get_identity_cache",
"load_identity",
"resolve_identity",
"resolve_user_api_key_auth",
]

152
litellm/identity/cache.py Normal file
View file

@ -0,0 +1,152 @@
"""Three-layer identity cache.
Layer 1 (per-process, ~5s TTL): bounded ``InMemoryCache`` inside the
``DualCache`` we wrap. Bounds revocation staleness without round-tripping
to Redis on every request.
Layer 2 (Redis, cross-replica): the ``redis_cache`` on the same
``DualCache``. Writes go to both layers; reads fall through.
Layer 3 (Prisma): not owned here. ``store.load_identity`` calls the DB
when both cache layers miss.
Cross-table fan-out is handled via *generation counters*: when a team or
user changes, the counter for that team/user is bumped. Cached
identities carry the team/user generation they were minted under, and a
read that finds a stale generation treats the entry as a miss. This is
cheaper than enumerating every key that references a given team.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from litellm.integrations.otel.model.spans import SpanRole
from litellm.integrations.otel.runtime import traced
if TYPE_CHECKING:
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import UserAPIKeyAuth
IDENTITY_KEY_PREFIX = "identity:v1"
IDENTITY_GENERATION_PREFIX = "identity:gen:v1"
DEFAULT_IDENTITY_TTL_SECONDS = 5
def identity_cache_key(token_hash: str) -> str:
return f"{IDENTITY_KEY_PREFIX}:{token_hash}"
def team_generation_key(team_id: str) -> str:
return f"{IDENTITY_GENERATION_PREFIX}:team:{team_id}"
def user_generation_key(user_id: str) -> str:
return f"{IDENTITY_GENERATION_PREFIX}:user:{user_id}"
def org_generation_key(org_id: str) -> str:
return f"{IDENTITY_GENERATION_PREFIX}:org:{org_id}"
def _generation_attr_key(scope: str) -> str:
return f"identity_cache_generation_{scope}"
def _attach_generations(
uak: "UserAPIKeyAuth", generations: dict
) -> None:
"""Stash the generation counters this entry was minted under.
Stored on the model's ``metadata`` so the value survives Pydantic
serialization round-trips through Redis (``CacheCodec``).
"""
if not isinstance(uak.metadata, dict):
uak.metadata = {}
uak.metadata["__identity_cache_generations__"] = generations
def _read_generations(uak: "UserAPIKeyAuth") -> dict:
metadata = uak.metadata if isinstance(uak.metadata, dict) else {}
raw = metadata.get("__identity_cache_generations__")
return raw if isinstance(raw, dict) else {}
class IdentityCache:
"""Memory + Redis facade keyed by the hashed virtual key.
Stores ``UserAPIKeyAuth`` because the legacy carrier is what the rest
of the proxy consumes today, and ``UserApiKeyCache`` already knows
how to round-trip it through ``CacheCodec``. Callers that want an
``IdentityContext`` view should call ``uak.to_identity_context()``
at the read site.
"""
def __init__(
self,
dual_cache: "DualCache",
ttl_seconds: int = DEFAULT_IDENTITY_TTL_SECONDS,
) -> None:
self._cache = dual_cache
self._ttl_seconds = ttl_seconds
@traced("identity.cache.get", role=SpanRole.DB_CALL,
attrs=lambda result: {
"identity.cache.layer": (
"miss" if result is None else "memory_or_redis"
),
})
async def get(self, token_hash: str) -> Optional["UserAPIKeyAuth"]:
from litellm.proxy._types import UserAPIKeyAuth
cache_key = identity_cache_key(token_hash)
cached = await self._cache.async_get_cache(
key=cache_key, model_type=UserAPIKeyAuth
)
if cached is None:
return None
if await self._is_stale(cached):
await self._cache.async_delete_cache(cache_key)
return None
return cached
@traced("identity.cache.set", role=SpanRole.DB_CALL)
async def set(
self, token_hash: str, uak: "UserAPIKeyAuth"
) -> None:
generations = await self._snapshot_generations_for(uak)
_attach_generations(uak, generations)
await self._cache.async_set_cache(
key=identity_cache_key(token_hash),
value=uak,
ttl=self._ttl_seconds,
)
async def delete(self, token_hash: str) -> None:
await self._cache.async_delete_cache(identity_cache_key(token_hash))
async def _is_stale(self, uak: "UserAPIKeyAuth") -> bool:
stored = _read_generations(uak)
if not stored:
return False
current = await self._snapshot_generations_for(uak)
return any(stored.get(k) != current.get(k) for k in stored)
async def _snapshot_generations_for(
self, uak: "UserAPIKeyAuth"
) -> dict:
scopes: list[tuple[str, str]] = []
if uak.team_id:
scopes.append(("team", team_generation_key(uak.team_id)))
if uak.user_id:
scopes.append(("user", user_generation_key(uak.user_id)))
if uak.org_id:
scopes.append(("org", org_generation_key(uak.org_id)))
return {
scope: await self._cache.async_get_cache(key=key) or 0
for scope, key in scopes
}
async def bump_generation(self, scope_key: str) -> None:
await self._cache.async_increment_cache(key=scope_key, value=1)

View file

@ -0,0 +1,70 @@
"""Identity-cache invalidation hooks.
Two flavors:
- Per-token: a key was rotated, blocked, or deleted. We know the exact
token hash, so we drop the entry from both memory and Redis.
- Per-scope (team / user / org): a row that fans out to many keys
changed. Rather than enumerating every key that references the team,
we bump a generation counter for that scope. Cached identities carry
the scope generations they were minted under; reads compare and treat
a mismatch as a miss.
The legacy ``_delete_cache_key_object`` and the per-table cache deletes
in ``auth_checks.py`` are left in place by Phase 2. These hooks run
side-by-side so we don't strand a partially-deployed fleet that's still
reading from the legacy cache keys.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
from litellm.identity.cache import (
IdentityCache,
org_generation_key,
team_generation_key,
user_generation_key,
)
if TYPE_CHECKING:
from litellm.caching.dual_cache import DualCache
def _identity_cache_for(dual_cache: "DualCache") -> IdentityCache:
return IdentityCache(dual_cache=dual_cache)
async def invalidate_identity_for_token(
*, token_hash: str, dual_cache: "DualCache"
) -> None:
"""Drop a single key's cached identity."""
await _identity_cache_for(dual_cache).delete(token_hash)
async def invalidate_identity_for_team(
*, team_id: str, dual_cache: "DualCache"
) -> None:
"""Mark every identity that references this team as stale."""
await _identity_cache_for(dual_cache).bump_generation(
team_generation_key(team_id)
)
async def invalidate_identity_for_user(
*, user_id: str, dual_cache: "DualCache"
) -> None:
"""Mark every identity that references this user as stale."""
await _identity_cache_for(dual_cache).bump_generation(
user_generation_key(user_id)
)
async def invalidate_identity_for_org(
*, org_id: str, dual_cache: "DualCache"
) -> None:
"""Mark every identity that references this organization as stale."""
await _identity_cache_for(dual_cache).bump_generation(
org_generation_key(org_id)
)

121
litellm/identity/jwt.py Normal file
View file

@ -0,0 +1,121 @@
"""JWT identity construction.
Owns the translation from a ``JWTAuthManager.auth_builder`` result (or
any equivalent JWT-validated payload) into the proxy's carrier model
``UserAPIKeyAuth``. The JWT branch of ``_user_api_key_auth_builder``
used to inline this construction; centralizing it here means:
- Every JWT-derived ``UserAPIKeyAuth`` carries the same team / user /
membership fields, so downstream auth checks see one shape.
- The mapping from ``jwt_claims`` to a ``JWTPrincipal`` happens at the
same boundary, so callers that want the typed principal can read it
off ``UserAPIKeyAuth.to_identity_context()``.
This module does NOT perform JWT validation or policy checks. Signature
verification, RBAC, scope, email-domain enforcement, and
``custom_validate`` are all done by ``JWTAuthManager.auth_builder``
upstream. Here we just build the carrier.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Optional
if TYPE_CHECKING:
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
def build_user_api_key_auth_from_jwt_result(
*,
result: dict,
parent_otel_span: Any = None,
is_proxy_admin: bool,
) -> "UserAPIKeyAuth":
"""Build a ``UserAPIKeyAuth`` carrier from a JWT auth-builder result.
``result`` is the dict returned by
``JWTAuthManager.auth_builder``; its shape (``team_id``,
``team_object``, ``user_id``, ``user_object``, ``end_user_id``,
``org_id``, ``team_membership``, ``jwt_claims``) is the public
contract between the validation layer and this construction layer.
The ``is_proxy_admin`` flag is the caller's responsibility to pass
because the auth_builder result reports it but the call site already
branched on it; threading it through avoids re-checking.
"""
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
team_id = result["team_id"]
team_object = result["team_object"]
user_id = result["user_id"]
user_object = result["user_object"]
end_user_id = result["end_user_id"]
org_id = result["org_id"]
team_membership = result.get("team_membership")
jwt_claims = result.get("jwt_claims")
team_alias = team_object.team_alias if team_object is not None else None
team_tpm_limit = team_object.tpm_limit if team_object is not None else None
team_rpm_limit = team_object.rpm_limit if team_object is not None else None
team_models = team_object.models if team_object is not None else []
team_metadata = team_object.metadata if team_object is not None else None
team_object_permission = (
team_object.object_permission if team_object is not None else None
)
if is_proxy_admin:
return UserAPIKeyAuth(
api_key=None,
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id=user_id,
team_id=team_id,
team_alias=team_alias,
team_tpm_limit=team_tpm_limit,
team_rpm_limit=team_rpm_limit,
team_models=team_models,
team_metadata=team_metadata,
org_id=org_id,
end_user_id=end_user_id,
parent_otel_span=parent_otel_span,
jwt_claims=jwt_claims,
)
user_role = (
LitellmUserRoles(user_object.user_role)
if user_object is not None and user_object.user_role is not None
else LitellmUserRoles.INTERNAL_USER
)
user_tpm_limit = user_object.tpm_limit if user_object is not None else None
user_rpm_limit = user_object.rpm_limit if user_object is not None else None
team_member_rpm_limit = (
team_membership.safe_get_team_member_rpm_limit()
if team_membership is not None
else None
)
team_member_tpm_limit = (
team_membership.safe_get_team_member_tpm_limit()
if team_membership is not None
else None
)
valid_token = UserAPIKeyAuth(
api_key=None,
team_id=team_id,
team_alias=team_alias,
team_tpm_limit=team_tpm_limit,
team_rpm_limit=team_rpm_limit,
team_models=team_models,
user_role=user_role,
user_id=user_id,
org_id=org_id,
parent_otel_span=parent_otel_span,
end_user_id=end_user_id,
user_tpm_limit=user_tpm_limit,
user_rpm_limit=user_rpm_limit,
team_member_rpm_limit=team_member_rpm_limit,
team_member_tpm_limit=team_member_tpm_limit,
team_metadata=team_metadata,
jwt_claims=jwt_claims,
)
valid_token.team_object_permission = team_object_permission
return valid_token

View file

@ -0,0 +1,49 @@
"""OAuth2 identity construction.
Owns the translation from an OAuth2 introspection / userinfo response
into the proxy's carrier model ``UserAPIKeyAuth``. The HTTP plumbing —
introspection endpoint detection, request signing, error handling —
stays in ``litellm.proxy.auth.oauth2_check.Oauth2Handler``; this module
just maps a validated response payload into the carrier so the JWT and
OAuth2 paths converge on the same construction surface.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Optional, cast
if TYPE_CHECKING:
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
def build_user_api_key_auth_from_oauth2_response(
*,
token: str,
response_data: dict,
user_id_field_name: str = "sub",
user_role_field_name: str = "role",
user_team_id_field_name: str = "team_id",
) -> "UserAPIKeyAuth":
"""Build a ``UserAPIKeyAuth`` carrier from an OAuth2 response.
``response_data`` is the parsed JSON body of an OAuth2 introspection
or userinfo response. The three ``*_field_name`` parameters let
deployments map their IdP's claim names onto our canonical fields;
defaults match the OAuth2 introspection spec (``sub``) plus
``role`` / ``team_id`` for the role and team claims.
Active-token validation, scope checks, and token-not-active rejection
happen upstream in ``Oauth2Handler``; this builder trusts its input.
"""
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
user_id: Optional[str] = response_data.get(user_id_field_name)
user_role: Optional[str] = response_data.get(user_role_field_name)
user_team_id: Optional[str] = response_data.get(user_team_id_field_name)
return UserAPIKeyAuth(
api_key=token,
team_id=user_team_id,
user_id=user_id,
user_role=cast("LitellmUserRoles", user_role),
)

View file

@ -1,11 +1,20 @@
"""Compose extractors into a single ``IdentityContext`` per request.
"""Compose extractors + DB load into a single ``IdentityContext`` per request.
This entrypoint is intentionally not yet wired into the proxy auth chain;
it exists so Phase 2 can switch over and so unit tests can exercise the
full path end-to-end.
Two call shapes:
- ``resolve_identity_for_principal`` — given pre-extracted credentials, decide
the principal kind and resolve it. Used by the proxy auth chain after it
already pulled the api-key out of the request.
- ``resolve_identity`` — request-scoped composition for new entrypoints. Not
yet wired into ``user_api_key_auth.py``; lives here so callers without a
hashed-token-in-hand (CLI, MCP, background jobs) can still build a
``IdentityContext``.
"""
from typing import Any, Dict, Optional
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Dict, Optional
from litellm.constants import (
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
@ -13,7 +22,10 @@ from litellm.constants import (
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
)
from litellm.identity.context import AuditInfo, ClientInfo, IdentityContext
from litellm.identity.extractors.api_key import extract_api_key_principal
from litellm.identity.extractors.api_key import (
extract_api_key_principal,
hash_principal_token,
)
from litellm.identity.extractors.client import extract_client_info
from litellm.identity.extractors.end_user import extract_end_user_id
from litellm.identity.extractors.header import extract_audit_changed_by
@ -23,6 +35,15 @@ from litellm.identity.principal import (
Principal,
ServiceAccountPrincipal,
)
from litellm.integrations.otel.model.spans import SpanRole
from litellm.integrations.otel.runtime import traced
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.identity.cache import IdentityCache
_SERVICE_ACCOUNT_API_KEYS = frozenset(
{
@ -33,7 +54,7 @@ _SERVICE_ACCOUNT_API_KEYS = frozenset(
)
def _resolve_principal(api_key: Optional[str]) -> Principal:
def _principal_from_raw_key(api_key: Optional[str]) -> Principal:
if api_key and api_key in _SERVICE_ACCOUNT_API_KEYS:
return ServiceAccountPrincipal(name=api_key)
@ -48,7 +69,15 @@ def _resolve_principal(api_key: Optional[str]) -> Principal:
return AnonymousPrincipal()
def resolve_identity(
@traced(
"identity.resolve",
role=SpanRole.SERVICE,
attrs=lambda result: {
"identity.principal.kind": result.principal.kind,
"identity.end_user.present": result.end_user_id is not None,
},
)
async def resolve_identity(
*,
api_key: Optional[str] = None,
request: Any = None,
@ -56,7 +85,13 @@ def resolve_identity(
headers: Optional[Dict[str, Any]] = None,
general_settings: Optional[Dict[str, Any]] = None,
) -> IdentityContext:
principal = _resolve_principal(api_key)
"""Build an ``IdentityContext`` purely from request-side signals.
Does not touch the database. The hydrated-row variant lives in
``store.load_identity`` and is composed by callers that have a
``PrismaClient`` + ``IdentityCache`` in hand.
"""
principal = _principal_from_raw_key(api_key)
end_user_id = extract_end_user_id(body=body, headers=headers)
audit = AuditInfo(changed_by=extract_audit_changed_by(headers))
client: ClientInfo
@ -73,3 +108,32 @@ def resolve_identity(
audit=audit,
client=client,
)
async def resolve_user_api_key_auth(
*,
api_key: str,
prisma_client: "PrismaClient",
identity_cache: "IdentityCache",
user_api_key_cache: "UserApiKeyCache",
parent_otel_span=None,
proxy_logging_obj: Optional["ProxyLogging"] = None,
) -> "UserAPIKeyAuth":
"""Cache-or-DB resolve a hashed virtual key into ``UserAPIKeyAuth``.
This is the surface ``_user_api_key_auth_builder`` calls in place of
the legacy ``get_key_object``. The hashed-token computation is
centralized via ``hash_principal_token`` so the auth chain doesn't
re-implement the JWT-vs-key hashing rules.
"""
from litellm.identity.store import load_identity
hashed_token = hash_principal_token(api_key)
return await load_identity(
hashed_token=hashed_token,
prisma_client=prisma_client,
cache=identity_cache,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)

View file

@ -0,0 +1,50 @@
"""Process-wide accessor for the shared ``IdentityCache``.
The proxy already owns one ``DualCache`` for caller identity at
``litellm.proxy.proxy_server.user_api_key_cache``. We layer an
``IdentityCache`` on top of it so the new identity load path shares a
single in-memory/Redis backend with the legacy caches. This avoids
double-caching on a single deploy and keeps invalidation surfaces
aligned.
Off-proxy callers (CLI, tests) can pass their own ``DualCache``.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from litellm.identity.cache import IdentityCache
if TYPE_CHECKING:
from litellm.caching.dual_cache import DualCache
_identity_cache: Optional[IdentityCache] = None
def get_identity_cache(
dual_cache: Optional["DualCache"] = None,
) -> IdentityCache:
"""Return the shared ``IdentityCache``, building it on first call.
When ``dual_cache`` is omitted, the proxy's module-level cache is
used. The first call wins; subsequent calls ignore the argument so
that every consumer in a process sees the same instance.
"""
global _identity_cache
if _identity_cache is not None:
return _identity_cache
if dual_cache is None:
from litellm.proxy.proxy_server import user_api_key_cache as _proxy_cache
dual_cache = _proxy_cache
_identity_cache = IdentityCache(dual_cache=dual_cache)
return _identity_cache
def reset_identity_cache_for_tests() -> None:
global _identity_cache
_identity_cache = None

232
litellm/identity/store.py Normal file
View file

@ -0,0 +1,232 @@
"""Cold-path identity loader.
One async function, one I/O round trip on cache miss. The actual SQL
JOIN lives in ``litellm.proxy.utils.PrismaClient.get_data`` (the
``combined_view`` query) and is reused as-is so we don't duplicate the
schema-coupled SQL. Object-permission rows that are referenced but not
joined are lazily filled exactly as the legacy ``get_key_object`` did.
The cached payload is ``UserAPIKeyAuth`` — the proxy's existing carrier.
Callers that want an ``IdentityContext`` view should call
``uak.to_identity_context()`` at the consumption site.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from fastapi import status
from litellm.integrations.otel.model.spans import SpanRole
from litellm.integrations.otel.runtime import traced
from litellm.proxy._types import ProxyErrorTypes, ProxyException
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.identity.cache import IdentityCache
@traced(
"identity.db.combined_view",
role=SpanRole.DB_CALL,
attrs=lambda result: {
"identity.load.outcome": "found" if result is not None else "missing",
"db.system.name": "postgresql",
},
)
async def _fetch_from_db(
*,
hashed_token: str,
prisma_client: "PrismaClient",
user_api_key_cache: "UserApiKeyCache",
parent_otel_span,
proxy_logging_obj: Optional["ProxyLogging"],
) -> Optional["UserAPIKeyAuth"]:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
_fetch_key_object_from_db_with_reconnect,
get_object_permission,
get_user_object,
)
row = await _fetch_key_object_from_db_with_reconnect(
hashed_token=hashed_token,
prisma_client=prisma_client,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if row is None:
return None
uak = UserAPIKeyAuth(**row.model_dump(exclude_none=True))
if uak.object_permission_id and not uak.object_permission:
try:
uak.object_permission = await get_object_permission(
object_permission_id=uak.object_permission_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,
)
except Exception:
pass
# Bundle the user row into the identity cache so the auth chain's
# follow-up `get_user_object` lookup is a free in-memory read on the
# next request (the user object survives the same TTL as the rest of
# the identity entry). The lookup itself is wrapped: failures fall
# back to the legacy per-row cache that `get_user_object` populates.
if uak.user_id:
try:
uak.user = await get_user_object(
user_id=uak.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception:
uak.user = None
return uak
def _principal_kind(uak: "UserAPIKeyAuth") -> str:
if uak.api_key and uak.api_key == uak.team_id:
return "service_account"
if uak.jwt_claims:
return "jwt"
if uak.token:
return "api_key"
return "anonymous"
@traced(
"identity.load",
role=SpanRole.SERVICE,
attrs=lambda result: {
"identity.principal.kind": _principal_kind(result),
"identity.principal.has_team": bool(result.team_id),
"identity.principal.has_org": bool(result.org_id),
"identity.principal.has_project": bool(result.project_id),
"identity.principal.has_agent": bool(result.agent_id),
},
)
async def load_identity(
*,
hashed_token: str,
prisma_client: "PrismaClient",
cache: "IdentityCache",
user_api_key_cache: "UserApiKeyCache",
parent_otel_span=None,
proxy_logging_obj: Optional["ProxyLogging"] = None,
) -> "UserAPIKeyAuth":
"""Return the hydrated ``UserAPIKeyAuth`` for a hashed token.
Cache-hit returns immediately. Cache-miss issues exactly one combined
view SQL query. A missing key raises ``ProxyException`` matching the
legacy ``get_key_object`` contract so callers don't need to special
case the cutover.
"""
cached = await cache.get(hashed_token)
if cached is not None:
_rehydrate_bundled_user(cached)
return _hydrated_copy(cached)
if prisma_client is None:
raise Exception(
"No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys"
)
uak = await _fetch_from_db(
hashed_token=hashed_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 uak is None:
raise ProxyException(
message=(
"Authentication Error, Invalid proxy server token passed. "
f"key={hashed_token}, not found in db. Create key via "
"`/key/generate` call."
),
type=ProxyErrorTypes.token_not_found_in_db,
param="key",
code=status.HTTP_401_UNAUTHORIZED,
)
await cache.set(hashed_token, uak)
await _populate_legacy_cache(
hashed_token=hashed_token,
uak=uak,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
return _hydrated_copy(uak)
async def _populate_legacy_cache(
*,
hashed_token: str,
uak: "UserAPIKeyAuth",
user_api_key_cache: "UserApiKeyCache",
proxy_logging_obj: Optional["ProxyLogging"],
) -> None:
"""Keep the legacy ``user_api_key_cache`` populated.
The pre-DB cache peek in ``_user_api_key_auth_builder`` (the
``check_cache_only=True`` admin fast-path) reads from this cache, so
populating it on every cold load preserves that fast path with no
code change at the call site.
"""
from litellm.proxy.auth.auth_checks import _cache_key_object
try:
await _cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=uak,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
except Exception:
pass
def _rehydrate_bundled_user(uak: "UserAPIKeyAuth") -> None:
"""Coerce the bundled user dict back into ``LiteLLM_UserTable``.
``UserAPIKeyAuth.user`` is typed ``Any`` so the codec round-trips it
as a plain dict. Consumers read scalar attributes (`tpm_limit`,
`metadata`, `user_role`, …) off the model, so we restore the typed
form before handing the cached entry back.
"""
from litellm.proxy._types import LiteLLM_UserTable
raw = getattr(uak, "user", None)
if isinstance(raw, dict):
try:
uak.user = LiteLLM_UserTable(**raw)
except Exception:
uak.user = None
def _hydrated_copy(uak: "UserAPIKeyAuth") -> "UserAPIKeyAuth":
"""Return a copy safe to mutate per-request.
The cached entry is shared across requests; consumers mutate fields
like ``parent_otel_span``, ``request_route``, and ``end_user_id`` on
the returned object, so we hand each caller its own model copy with
those request-scoped fields cleared.
"""
copy = uak.model_copy()
copy.parent_otel_span = None
copy.request_route = None
copy.budget_reservation = None
return copy

View file

@ -1,14 +1,18 @@
"""SDK-free entrypoints for proxy-core call sites (auth, …).
"""SDK-free entrypoints for proxy-core call sites (auth, identity, …).
Proxy code may run without the OpenTelemetry SDK installed, so it must not import
``litellm.integrations.otel.logger`` (which imports the SDK at module scope) at
module load. These wrappers import it lazily and no-op when the SDK is absent or
V2 is not the active logger — so a call site can wrap a request phase or seed
identity unconditionally.
V2 is not the active logger — so a call site can wrap a request phase, seed
identity, or decorate an async function with ``@traced`` unconditionally.
"""
import functools
import inspect
from contextlib import contextmanager
from typing import Any, Iterator
from typing import Any, Callable, Iterator, Mapping, Optional
from litellm.integrations.otel.model.spans import SpanRole
@contextmanager
@ -36,3 +40,94 @@ def seed_request_identity(user_api_key_dict: Any, model: Any = None) -> None:
except Exception:
return
_seed_request_identity(user_api_key_dict, model=model)
def _apply_span_attrs(span: Any, attrs: Optional[Mapping[str, Any]]) -> None:
if span is None or not attrs:
return
for key, value in attrs.items():
if value is None:
continue
try:
span.set_attribute(key, value)
except Exception:
pass
def _resolve_attrs(
builder: Optional[Callable[..., Mapping[str, Any]]],
*,
args: tuple,
kwargs: dict,
result: Any,
) -> Mapping[str, Any]:
if builder is None:
return {}
try:
sig = inspect.signature(builder)
accepts_result = "result" in sig.parameters or any(
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
)
if accepts_result:
payload = builder(*args, result=result, **kwargs)
else:
payload = builder(result)
return payload or {}
except Exception:
return {}
def traced(
span_name: str,
*,
role: SpanRole = SpanRole.SERVICE,
attrs: Optional[Callable[..., Mapping[str, Any]]] = None,
) -> Callable:
"""Wrap an async function so it runs inside an OTel v2 span.
``SpanRole.SERVICE`` opens a phase span under the active root proxy request
span; nested DB/service calls become its children. ``SpanRole.DB_CALL`` (and
any other non-SERVICE role) does not open a new span — it just attaches
attributes to the current span, so we don't double-count for code paths
that already produce a DB_CALL via service-logger plumbing.
``attrs`` is an optional callback that receives the function's return value
(and the wrapped function's args/kwargs via ``**kwargs`` if it accepts
them) and returns a mapping of span attributes to set on success. Any
exception in ``attrs`` is swallowed — instrumentation must never break the
request.
"""
def _decorator(func: Callable) -> Callable:
if not inspect.iscoroutinefunction(func):
raise TypeError(
f"@traced requires an async function; got {func!r}"
)
@functools.wraps(func)
async def _wrapper(*args, **kwargs):
if role is SpanRole.SERVICE:
with phase_span(span_name) as span:
result = await func(*args, **kwargs)
_apply_span_attrs(
span,
_resolve_attrs(attrs, args=args, kwargs=kwargs, result=result),
)
return result
result = await func(*args, **kwargs)
try:
from opentelemetry import trace as _otel_trace
current_span = _otel_trace.get_current_span()
except Exception:
current_span = None
_apply_span_attrs(
current_span,
_resolve_attrs(attrs, args=args, kwargs=kwargs, result=result),
)
return result
return _wrapper
return _decorator

View file

@ -1869,6 +1869,12 @@ async def _delete_cache_key_object(
key=key
)
from litellm.identity.invalidation import invalidate_identity_for_token
await invalidate_identity_for_token(
token_hash=hashed_token, dual_cache=user_api_key_cache
)
@log_db_metrics
async def _get_team_db_check(

View file

@ -1,15 +1,16 @@
import base64
import os
from typing import Dict, Optional, Tuple, cast
from typing import Dict, Optional, Tuple
import httpx
from litellm._logging import verbose_proxy_logger
from litellm.identity import build_user_api_key_auth_from_oauth2_response
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
class Oauth2Handler:
@ -86,31 +87,6 @@ class Oauth2Handler:
"""
return {"Authorization": f"Bearer {token}", "Content-Type": "application/json"}
@staticmethod
def _extract_user_info(
response_data: Dict,
user_id_field_name: str,
user_role_field_name: str,
user_team_id_field_name: str,
) -> Tuple[Optional[str], Optional[str], Optional[str]]:
"""
Extract user information from OAuth2 response.
Args:
response_data: The response data from OAuth2 endpoint
user_id_field_name: Field name for user ID
user_role_field_name: Field name for user role
user_team_id_field_name: Field name for team ID
Returns:
Tuple of (user_id, user_role, user_team_id)
"""
user_id = response_data.get(user_id_field_name)
user_team_id = response_data.get(user_team_id_field_name)
user_role = response_data.get(user_role_field_name)
return user_id, user_role, user_team_id
@staticmethod
async def check_oauth2_token(token: str) -> UserAPIKeyAuth:
"""
@ -202,20 +178,13 @@ class Oauth2Handler:
if is_introspection_endpoint and not data.get("active", True):
raise ValueError("Token is not active")
# Extract user information from response
user_id, user_role, user_team_id = Oauth2Handler._extract_user_info(
return build_user_api_key_auth_from_oauth2_response(
token=token,
response_data=data,
user_id_field_name=user_id_field_name,
user_role_field_name=user_role_field_name,
user_team_id_field_name=user_team_id_field_name,
)
return UserAPIKeyAuth(
api_key=token,
team_id=user_team_id,
user_id=user_id,
user_role=cast(LitellmUserRoles, user_role),
)
except httpx.HTTPStatusError as e:
# This will catch any 4xx or 5xx errors
raise ValueError(f"Oauth 2.0 Token validation failed: {e}")

View file

@ -22,6 +22,7 @@ import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.identity import build_user_api_key_auth_from_jwt_result
from litellm.integrations.otel.model.config import is_otel_v2_enabled
from litellm.integrations.otel.runtime import phase_span, seed_request_identity
from litellm.litellm_core_utils.dd_tracing import tracer
@ -1141,10 +1142,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
+ CommonProxyErrors.not_premium_user.value
)
return await Oauth2Handler.check_oauth2_token(token=api_key)
with phase_span("identity.oauth2"):
return await Oauth2Handler.check_oauth2_token(token=api_key)
if general_settings.get("enable_oauth2_proxy_auth", False) is True:
return await handle_oauth2_proxy_request(request=request)
with phase_span("identity.oauth2_proxy"):
return await handle_oauth2_proxy_request(request=request)
if general_settings.get("enable_jwt_auth", False) is True:
from litellm.proxy.proxy_server import premium_user
@ -1191,7 +1194,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# standard JWT auth_builder below
if do_standard_jwt_auth:
with tracer.trace("litellm.proxy.auth.jwt_auth_builder"):
with phase_span("identity.jwt"):
result = await JWTAuthManager.auth_builder(
request_data=request_data,
general_settings=general_settings,
@ -1210,14 +1213,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
is_proxy_admin = result["is_proxy_admin"]
team_id = result["team_id"]
team_object = result["team_object"]
user_id = result["user_id"]
user_object = result["user_object"]
end_user_id = result["end_user_id"]
org_id = result["org_id"]
team_membership: Optional[LiteLLM_TeamMembership] = result.get(
"team_membership", None
)
jwt_claims = result.get("jwt_claims", None)
if is_proxy_admin:
@ -1234,90 +1232,16 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
value=_JWT_PROXY_ADMIN_SENTINEL,
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
)
return UserAPIKeyAuth(
api_key=None,
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id=user_id,
team_id=team_id,
team_alias=(
team_object.team_alias
if team_object is not None
else None
),
team_tpm_limit=(
team_object.tpm_limit
if team_object is not None
else None
),
team_rpm_limit=(
team_object.rpm_limit
if team_object is not None
else None
),
team_models=(
team_object.models if team_object is not None else []
),
team_metadata=(
team_object.metadata
if team_object is not None
else None
),
org_id=org_id,
end_user_id=end_user_id,
return build_user_api_key_auth_from_jwt_result(
result=result,
parent_otel_span=parent_otel_span,
jwt_claims=jwt_claims,
is_proxy_admin=True,
)
valid_token = UserAPIKeyAuth(
api_key=None,
team_id=team_id,
team_alias=(
team_object.team_alias if team_object is not None else None
),
team_tpm_limit=(
team_object.tpm_limit if team_object is not None else None
),
team_rpm_limit=(
team_object.rpm_limit if team_object is not None else None
),
team_models=(
team_object.models if team_object is not None else []
),
user_role=(
LitellmUserRoles(user_object.user_role)
if user_object is not None
and user_object.user_role is not None
else LitellmUserRoles.INTERNAL_USER
),
user_id=user_id,
org_id=org_id,
valid_token = build_user_api_key_auth_from_jwt_result(
result=result,
parent_otel_span=parent_otel_span,
end_user_id=end_user_id,
user_tpm_limit=(
user_object.tpm_limit if user_object is not None else None
),
user_rpm_limit=(
user_object.rpm_limit if user_object is not None else None
),
team_member_rpm_limit=(
team_membership.safe_get_team_member_rpm_limit()
if team_membership is not None
else None
),
team_member_tpm_limit=(
team_membership.safe_get_team_member_tpm_limit()
if team_membership is not None
else None
),
team_metadata=(
team_object.metadata if team_object is not None else None
),
jwt_claims=jwt_claims,
)
valid_token.team_object_permission = (
team_object.object_permission
if team_object is not None
else None
is_proxy_admin=False,
)
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
@ -1679,9 +1603,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
try:
with tracer.trace("litellm.proxy.auth.get_key_object_from_db"):
valid_token = await get_key_object(
from litellm.identity import (
get_identity_cache,
load_identity,
)
valid_token = await load_identity(
hashed_token=api_key,
prisma_client=prisma_client,
cache=get_identity_cache(user_api_key_cache),
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
@ -1737,23 +1667,27 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Check 2. If user_id for this token is in budget - done in common_checks()
if valid_token.user_id is not None:
try:
with tracer.trace("litellm.proxy.auth.get_user_object"):
user_obj = await get_user_object(
user_id=valid_token.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
bundled_user = getattr(valid_token, "user", None)
if isinstance(bundled_user, LiteLLM_UserTable):
user_obj = bundled_user
else:
try:
with tracer.trace("litellm.proxy.auth.get_user_object"):
user_obj = await get_user_object(
user_id=valid_token.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e:
verbose_logger.debug(
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {}".format(
str(e)
)
)
except Exception as e:
verbose_logger.debug(
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {}".format(
str(e)
)
)
user_obj = None
user_obj = None
if (
user_obj is not None

View file

@ -1373,6 +1373,16 @@ async def _update_single_user_helper(
detail={"error": "Failed to update user"},
)
_strip_password_from_response(response)
updated_user_id = response.get("user_id") if isinstance(response, dict) else None
if updated_user_id:
from litellm.identity.invalidation import invalidate_identity_for_user
from litellm.proxy.proxy_server import user_api_key_cache
await invalidate_identity_for_user(
user_id=updated_user_id, dual_cache=user_api_key_cache
)
return response
@ -2325,6 +2335,14 @@ async def delete_user(
where={"user_id": {"in": data.user_ids}}
)
from litellm.identity.invalidation import invalidate_identity_for_user
from litellm.proxy.proxy_server import user_api_key_cache
for user_id in data.user_ids:
await invalidate_identity_for_user(
user_id=user_id, dual_cache=user_api_key_cache
)
return deleted_users

View file

@ -589,6 +589,13 @@ async def update_organization(
include={"members": True, "teams": True, "litellm_budget_table": True},
)
from litellm.identity.invalidation import invalidate_identity_for_org
from litellm.proxy.proxy_server import user_api_key_cache
await invalidate_identity_for_org(
org_id=data.organization_id, dual_cache=user_api_key_cache
)
return response
@ -676,6 +683,14 @@ async def delete_organization(
)
deleted_orgs.append(deleted_org)
from litellm.identity.invalidation import invalidate_identity_for_org
from litellm.proxy.proxy_server import user_api_key_cache
for organization_id in data.organization_ids:
await invalidate_identity_for_org(
org_id=organization_id, dual_cache=user_api_key_cache
)
return deleted_orgs

View file

@ -171,6 +171,12 @@ async def _refresh_cached_team(
proxy_logging_obj=proxy_logging_obj,
)
from litellm.identity.invalidation import invalidate_identity_for_team
await invalidate_identity_for_team(
team_id=team_row.team_id, dual_cache=user_api_key_cache
)
async def _verify_team_access(
team_obj: LiteLLM_TeamTable,
@ -3294,6 +3300,15 @@ async def delete_team(
deleted_teams = await prisma_client.delete_data(
team_id_list=data.team_ids, table_name="team"
)
from litellm.identity.invalidation import invalidate_identity_for_team
from litellm.proxy.proxy_server import user_api_key_cache
for team_id in data.team_ids:
await invalidate_identity_for_team(
team_id=team_id, dual_cache=user_api_key_cache
)
return deleted_teams

View file

@ -0,0 +1,77 @@
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import pytest
from litellm.caching.dual_cache import DualCache
from litellm.identity.cache import (
IdentityCache,
identity_cache_key,
team_generation_key,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
def _user_api_key_cache() -> UserApiKeyCache:
return UserApiKeyCache()
@pytest.mark.asyncio
async def test_set_then_get_returns_same_uak():
cache = IdentityCache(dual_cache=_user_api_key_cache())
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t1")
token_hash = uak.token
assert token_hash is not None
await cache.set(token_hash, uak)
got = await cache.get(token_hash)
assert got is not None
assert got.user_id == "u1"
assert got.team_id == "t1"
@pytest.mark.asyncio
async def test_miss_returns_none():
cache = IdentityCache(dual_cache=_user_api_key_cache())
assert await cache.get("missing-hash") is None
@pytest.mark.asyncio
async def test_delete_clears_entry():
dual = _user_api_key_cache()
cache = IdentityCache(dual_cache=dual)
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1")
await cache.set(uak.token, uak)
assert await cache.get(uak.token) is not None
await cache.delete(uak.token)
assert await cache.get(uak.token) is None
@pytest.mark.asyncio
async def test_generation_bump_invalidates_team_scoped_entry():
dual = _user_api_key_cache()
cache = IdentityCache(dual_cache=dual)
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t-rotate")
await cache.set(uak.token, uak)
assert await cache.get(uak.token) is not None
await cache.bump_generation(team_generation_key("t-rotate"))
assert await cache.get(uak.token) is None
@pytest.mark.asyncio
async def test_generation_bump_for_unrelated_team_keeps_entry():
dual = _user_api_key_cache()
cache = IdentityCache(dual_cache=dual)
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t-keep")
await cache.set(uak.token, uak)
await cache.bump_generation(team_generation_key("t-other"))
assert await cache.get(uak.token) is not None
def test_key_format_is_versioned():
assert identity_cache_key("abc").startswith("identity:v1:")

View file

@ -0,0 +1,73 @@
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import pytest
from litellm.identity.cache import IdentityCache, team_generation_key
from litellm.identity.invalidation import (
invalidate_identity_for_org,
invalidate_identity_for_team,
invalidate_identity_for_token,
invalidate_identity_for_user,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@pytest.mark.asyncio
async def test_token_invalidation_drops_entry():
backend = UserApiKeyCache()
cache = IdentityCache(dual_cache=backend)
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1")
await cache.set(uak.token, uak)
assert await cache.get(uak.token) is not None
await invalidate_identity_for_token(
token_hash=uak.token, dual_cache=backend
)
assert await cache.get(uak.token) is None
@pytest.mark.asyncio
async def test_team_invalidation_bumps_generation():
backend = UserApiKeyCache()
cache = IdentityCache(dual_cache=backend)
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", team_id="t-rotate")
await cache.set(uak.token, uak)
await invalidate_identity_for_team(
team_id="t-rotate", dual_cache=backend
)
assert await cache.get(uak.token) is None
@pytest.mark.asyncio
async def test_user_invalidation_is_scoped_to_user():
backend = UserApiKeyCache()
cache = IdentityCache(dual_cache=backend)
uak_a = UserAPIKeyAuth(api_key="sk-a", user_id="u-rotate", team_id="t1")
uak_b = UserAPIKeyAuth(api_key="sk-b", user_id="u-keep", team_id="t1")
await cache.set(uak_a.token, uak_a)
await cache.set(uak_b.token, uak_b)
await invalidate_identity_for_user(
user_id="u-rotate", dual_cache=backend
)
assert await cache.get(uak_a.token) is None
assert await cache.get(uak_b.token) is not None
@pytest.mark.asyncio
async def test_org_invalidation_drops_org_scoped_entry():
backend = UserApiKeyCache()
cache = IdentityCache(dual_cache=backend)
uak = UserAPIKeyAuth(api_key="sk-x", user_id="u1", org_id="org-rotate")
await cache.set(uak.token, uak)
await invalidate_identity_for_org(
org_id="org-rotate", dual_cache=backend
)
assert await cache.get(uak.token) is None

View file

@ -0,0 +1,144 @@
import os
import sys
from types import SimpleNamespace
sys.path.insert(0, os.path.abspath("../.."))
from litellm.identity import build_user_api_key_auth_from_jwt_result
from litellm.identity.principal import JWTPrincipal
from litellm.proxy._types import (
LiteLLM_TeamTableCachedObj,
LiteLLM_UserTable,
LitellmUserRoles,
UserAPIKeyAuth,
)
def _team(team_id="t-jwt"):
return LiteLLM_TeamTableCachedObj(
team_id=team_id,
team_alias=f"alias-{team_id}",
tpm_limit=100,
rpm_limit=10,
models=["gpt-4o"],
metadata={"env": "prod"},
)
def _user(user_id="u-jwt", role="internal_user"):
return LiteLLM_UserTable(
user_id=user_id,
user_email="jwt@litellm.io",
user_role=role,
tpm_limit=50,
rpm_limit=5,
)
def _membership(tpm=42, rpm=7):
return SimpleNamespace(
safe_get_team_member_tpm_limit=lambda: tpm,
safe_get_team_member_rpm_limit=lambda: rpm,
)
def _auth_builder_result(
*,
is_proxy_admin: bool = False,
team=None,
user=None,
team_membership=None,
jwt_claims=None,
):
return {
"is_proxy_admin": is_proxy_admin,
"team_id": team.team_id if team is not None else None,
"team_object": team,
"user_id": user.user_id if user is not None else None,
"user_object": user,
"end_user_id": "eu-jwt",
"org_id": "org-jwt",
"team_membership": team_membership,
"jwt_claims": jwt_claims
or {"sub": "u-jwt", "iss": "https://idp", "scope": "read write"},
}
def test_proxy_admin_path_marks_role_and_drops_membership_fields():
result = _auth_builder_result(is_proxy_admin=True, team=_team(), user=_user())
uak = build_user_api_key_auth_from_jwt_result(
result=result, parent_otel_span=None, is_proxy_admin=True
)
assert isinstance(uak, UserAPIKeyAuth)
assert uak.user_role == LitellmUserRoles.PROXY_ADMIN
assert uak.team_id == "t-jwt"
assert uak.team_alias == "alias-t-jwt"
assert uak.team_tpm_limit == 100
assert uak.team_metadata == {"env": "prod"}
assert uak.team_member_rpm_limit is None
assert uak.team_member_tpm_limit is None
assert uak.user_tpm_limit is None
assert uak.jwt_claims == result["jwt_claims"]
def test_regular_path_layers_team_user_membership_limits():
membership = _membership(tpm=42, rpm=7)
result = _auth_builder_result(
team=_team(), user=_user(role="internal_user"), team_membership=membership
)
uak = build_user_api_key_auth_from_jwt_result(
result=result, parent_otel_span=None, is_proxy_admin=False
)
assert uak.user_role == LitellmUserRoles.INTERNAL_USER
assert uak.user_tpm_limit == 50
assert uak.user_rpm_limit == 5
assert uak.team_member_tpm_limit == 42
assert uak.team_member_rpm_limit == 7
assert uak.org_id == "org-jwt"
assert uak.end_user_id == "eu-jwt"
def test_missing_team_object_defaults_team_fields_safely():
result = _auth_builder_result(team=None, user=_user())
uak = build_user_api_key_auth_from_jwt_result(
result=result, parent_otel_span=None, is_proxy_admin=False
)
assert uak.team_alias is None
assert uak.team_tpm_limit is None
assert uak.team_models == []
assert uak.team_metadata is None
def test_missing_user_object_defaults_user_role_to_internal():
result = _auth_builder_result(team=_team(), user=None)
uak = build_user_api_key_auth_from_jwt_result(
result=result, parent_otel_span=None, is_proxy_admin=False
)
assert uak.user_role == LitellmUserRoles.INTERNAL_USER
assert uak.user_tpm_limit is None
def test_adapter_roundtrip_produces_jwt_principal():
result = _auth_builder_result(team=_team(), user=_user())
uak = build_user_api_key_auth_from_jwt_result(
result=result, parent_otel_span=None, is_proxy_admin=False
)
ctx = uak.to_identity_context()
assert isinstance(ctx.principal, JWTPrincipal)
assert ctx.principal.sub == "u-jwt"
assert ctx.principal.iss == "https://idp"
assert ctx.principal.scopes == ["read", "write"]
assert ctx.principal.mapped_user_id == "u-jwt"
assert ctx.principal.mapped_team_id == "t-jwt"
assert ctx.principal.mapped_org_id == "org-jwt"
def test_team_object_permission_propagates_when_present():
team = _team()
team.object_permission = SimpleNamespace(object_permission_id="perm-1")
result = _auth_builder_result(team=team, user=_user())
uak = build_user_api_key_auth_from_jwt_result(
result=result, parent_otel_span=None, is_proxy_admin=False
)
assert uak.team_object_permission is not None
assert uak.team_object_permission.object_permission_id == "perm-1"

View file

@ -0,0 +1,55 @@
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
from litellm.identity import build_user_api_key_auth_from_oauth2_response
from litellm.proxy._types import UserAPIKeyAuth
def test_default_field_names_extract_from_introspection_response():
response = {"sub": "u-1", "role": "internal_user", "team_id": "t-1"}
uak = build_user_api_key_auth_from_oauth2_response(
token="opaque-token", response_data=response
)
assert isinstance(uak, UserAPIKeyAuth)
assert uak.user_id == "u-1"
assert uak.user_role == "internal_user"
assert uak.team_id == "t-1"
def test_custom_field_names_override_defaults():
response = {
"preferred_username": "alice",
"groups": "proxy_admin",
"tenant_id": "tenant-9",
}
uak = build_user_api_key_auth_from_oauth2_response(
token="t",
response_data=response,
user_id_field_name="preferred_username",
user_role_field_name="groups",
user_team_id_field_name="tenant_id",
)
assert uak.user_id == "alice"
assert uak.user_role == "proxy_admin"
assert uak.team_id == "tenant-9"
def test_missing_fields_default_to_none():
uak = build_user_api_key_auth_from_oauth2_response(
token="t", response_data={}
)
assert uak.user_id is None
assert uak.user_role is None
assert uak.team_id is None
def test_token_is_hashed_into_token_field():
"""The api_key is hashed by the UserAPIKeyAuth validator; the
OAuth2 builder must not bypass that path."""
uak = build_user_api_key_auth_from_oauth2_response(
token="sk-oauth2-test", response_data={"sub": "u"}
)
assert uak.token is not None
assert uak.token != "sk-oauth2-test"

View file

@ -0,0 +1,74 @@
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import pytest
from litellm.integrations.otel.model.spans import SpanRole
from litellm.integrations.otel.runtime import traced
@pytest.mark.asyncio
async def test_decorator_is_passthrough_when_v2_disabled():
calls = []
@traced("identity.unit-test", role=SpanRole.SERVICE)
async def work(x):
calls.append(x)
return x * 2
result = await work(5)
assert result == 10
assert calls == [5]
@pytest.mark.asyncio
async def test_attrs_callback_invoked_with_result():
captured = {}
@traced(
"identity.attrs-test",
role=SpanRole.SERVICE,
attrs=lambda result: captured.setdefault("result", result) or {},
)
async def work():
return {"principal_kind": "api_key"}
await work()
assert captured["result"] == {"principal_kind": "api_key"}
@pytest.mark.asyncio
async def test_attrs_exception_is_swallowed():
@traced(
"identity.attrs-error",
role=SpanRole.SERVICE,
attrs=lambda result: 1 / 0,
)
async def work():
return "ok"
assert await work() == "ok"
@pytest.mark.asyncio
async def test_db_call_role_does_not_open_a_new_span():
"""DB_CALL spans piggyback on the current span instead of starting one.
With V2 disabled there is no current span; the decorator should still
run the function and swallow any attribute-application errors.
"""
@traced("identity.db-test", role=SpanRole.DB_CALL)
async def work():
return "value"
assert await work() == "value"
def test_decorator_rejects_sync_functions():
with pytest.raises(TypeError):
@traced("identity.sync-not-allowed", role=SpanRole.SERVICE)
def sync_fn():
return None

View file

@ -6,6 +6,8 @@ from types import SimpleNamespace
sys.path.insert(0, os.path.abspath("../.."))
import pytest
from litellm.constants import LITTELM_CLI_SERVICE_ACCOUNT_NAME
from litellm.identity.principal import (
AnonymousPrincipal,
@ -32,30 +34,35 @@ def _jwt(claims):
return f"{b({'alg':'HS256','typ':'JWT'})}.{b(claims)}.sig"
def test_anonymous_when_no_credentials():
ctx = resolve_identity()
@pytest.mark.asyncio
async def test_anonymous_when_no_credentials():
ctx = await resolve_identity()
assert isinstance(ctx.principal, AnonymousPrincipal)
def test_api_key_principal_for_sk_key():
ctx = resolve_identity(api_key="sk-test")
@pytest.mark.asyncio
async def test_api_key_principal_for_sk_key():
ctx = await resolve_identity(api_key="sk-test")
assert isinstance(ctx.principal, ApiKeyPrincipal)
def test_jwt_principal_for_jwt_shaped_key():
ctx = resolve_identity(api_key=_jwt({"sub": "u1"}))
@pytest.mark.asyncio
async def test_jwt_principal_for_jwt_shaped_key():
ctx = await resolve_identity(api_key=_jwt({"sub": "u1"}))
assert isinstance(ctx.principal, JWTPrincipal)
assert ctx.principal.sub == "u1"
def test_service_account_for_known_sentinel():
ctx = resolve_identity(api_key=LITTELM_CLI_SERVICE_ACCOUNT_NAME)
@pytest.mark.asyncio
async def test_service_account_for_known_sentinel():
ctx = await resolve_identity(api_key=LITTELM_CLI_SERVICE_ACCOUNT_NAME)
assert isinstance(ctx.principal, ServiceAccountPrincipal)
assert ctx.principal.name == LITTELM_CLI_SERVICE_ACCOUNT_NAME
def test_end_user_and_audit_propagate():
ctx = resolve_identity(
@pytest.mark.asyncio
async def test_end_user_and_audit_propagate():
ctx = await resolve_identity(
body={"user": "eu-42"},
headers={"litellm-changed-by": "admin"},
)
@ -63,8 +70,9 @@ def test_end_user_and_audit_propagate():
assert ctx.audit.changed_by == "admin"
def test_client_info_from_request():
@pytest.mark.asyncio
async def test_client_info_from_request():
req = _fake_request({"user-agent": "curl/8"}, client_host="127.0.0.1")
ctx = resolve_identity(request=req)
ctx = await resolve_identity(request=req)
assert ctx.client.ip == "127.0.0.1"
assert ctx.client.user_agent == "curl/8"

View file

@ -0,0 +1,229 @@
import os
import sys
from unittest.mock import AsyncMock, patch
sys.path.insert(0, os.path.abspath("../.."))
import pytest
from fastapi import status
from litellm.identity.cache import IdentityCache
from litellm.identity.store import load_identity
from litellm.proxy._types import (
LiteLLM_VerificationTokenView,
ProxyErrorTypes,
ProxyException,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
def _stub_prisma_client():
"""Minimal prisma_client stub with the get_data hook the store uses."""
class _Stub:
def __init__(self):
self.get_data = AsyncMock()
self.calls = 0
return _Stub()
def _verification_token_view(**fields) -> LiteLLM_VerificationTokenView:
return LiteLLM_VerificationTokenView(
token=fields.pop("token", "hash-x"),
user_id=fields.pop("user_id", "u1"),
team_id=fields.pop("team_id", "t1"),
org_id=fields.pop("org_id", None),
**fields,
)
@pytest.mark.asyncio
async def test_cache_hit_skips_db():
cache_backend = UserApiKeyCache()
identity_cache = IdentityCache(dual_cache=cache_backend)
prisma = _stub_prisma_client()
from litellm.proxy._types import UserAPIKeyAuth
seed = UserAPIKeyAuth(
token="hash-cached", user_id="u-cache", team_id="t-cache"
)
await identity_cache.set("hash-cached", seed)
result = await load_identity(
hashed_token="hash-cached",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
assert prisma.get_data.await_count == 0
assert result.user_id == "u-cache"
assert result.team_id == "t-cache"
@pytest.mark.asyncio
async def test_cache_miss_hits_db_once_and_caches():
cache_backend = UserApiKeyCache()
identity_cache = IdentityCache(dual_cache=cache_backend)
prisma = _stub_prisma_client()
prisma.get_data.return_value = _verification_token_view(
token="hash-db", user_id="u-db", team_id="t-db"
)
with patch(
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
):
result = await load_identity(
hashed_token="hash-db",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
assert prisma.get_data.await_count == 1
assert result.user_id == "u-db"
cached = await identity_cache.get("hash-db")
assert cached is not None
assert cached.user_id == "u-db"
@pytest.mark.asyncio
async def test_warm_load_after_cold_does_not_hit_db():
cache_backend = UserApiKeyCache()
identity_cache = IdentityCache(dual_cache=cache_backend)
prisma = _stub_prisma_client()
prisma.get_data.return_value = _verification_token_view(
token="hash-warm", user_id="u-warm", team_id="t-warm"
)
with patch(
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
):
await load_identity(
hashed_token="hash-warm",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
await load_identity(
hashed_token="hash-warm",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
assert prisma.get_data.await_count == 1
@pytest.mark.asyncio
async def test_missing_token_raises_proxy_exception_with_401():
cache_backend = UserApiKeyCache()
identity_cache = IdentityCache(dual_cache=cache_backend)
prisma = _stub_prisma_client()
prisma.get_data.return_value = None
with pytest.raises(ProxyException) as excinfo:
await load_identity(
hashed_token="missing-hash",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
assert int(excinfo.value.code) == status.HTTP_401_UNAUTHORIZED
assert excinfo.value.type == ProxyErrorTypes.token_not_found_in_db
@pytest.mark.asyncio
async def test_bundled_user_survives_cache_roundtrip_as_typed_model():
from litellm.identity.cache import IdentityCache
from litellm.identity.store import _rehydrate_bundled_user
from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth
cache_backend = UserApiKeyCache()
identity_cache = IdentityCache(dual_cache=cache_backend)
uak = UserAPIKeyAuth(
api_key="sk-bundle",
user_id="u-bundle",
user=LiteLLM_UserTable(
user_id="u-bundle", user_email="x@y.com", tpm_limit=10
),
)
await identity_cache.set(uak.token, uak)
got = await identity_cache.get(uak.token)
_rehydrate_bundled_user(got)
assert isinstance(got.user, LiteLLM_UserTable)
assert got.user.user_email == "x@y.com"
assert got.user.tpm_limit == 10
@pytest.mark.asyncio
async def test_cold_path_bundles_user_into_cache():
"""``load_identity`` must populate ``user`` on first DB hit so the
next request reads the user object from the identity cache."""
cache_backend = UserApiKeyCache()
identity_cache = IdentityCache(dual_cache=cache_backend)
prisma = _stub_prisma_client()
prisma.get_data.return_value = _verification_token_view(
token="hash-user", user_id="u-bundle", team_id="t-bundle"
)
from litellm.proxy._types import LiteLLM_UserTable
async def _fake_get_user_object(*, user_id, **kwargs):
return LiteLLM_UserTable(
user_id=user_id, user_email="bundle@litellm.io", tpm_limit=42
)
with patch(
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
), patch(
"litellm.proxy.auth.auth_checks.get_user_object",
side_effect=_fake_get_user_object,
):
result = await load_identity(
hashed_token="hash-user",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
assert result.user is not None
assert result.user.user_email == "bundle@litellm.io"
cached = await identity_cache.get("hash-user")
assert cached is not None
assert cached.user is not None
@pytest.mark.asyncio
async def test_hydrated_copy_is_request_scoped():
cache_backend = UserApiKeyCache()
identity_cache = IdentityCache(dual_cache=cache_backend)
prisma = _stub_prisma_client()
prisma.get_data.return_value = _verification_token_view(
token="hash-copy", user_id="u-copy", team_id="t-copy"
)
with patch(
"litellm.identity.store._populate_legacy_cache", new=AsyncMock()
):
a = await load_identity(
hashed_token="hash-copy",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
a.parent_otel_span = "request-A-span"
a.request_route = "/v1/chat/completions"
b = await load_identity(
hashed_token="hash-copy",
prisma_client=prisma,
cache=identity_cache,
user_api_key_cache=cache_backend,
)
assert b.parent_otel_span is None
assert b.request_route is None