diff --git a/litellm/identity/__init__.py b/litellm/identity/__init__.py index c35355a1f45..59a02e1e15f 100644 --- a/litellm/identity/__init__.py +++ b/litellm/identity/__init__.py @@ -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", ] diff --git a/litellm/identity/cache.py b/litellm/identity/cache.py new file mode 100644 index 00000000000..8bfc770c3c2 --- /dev/null +++ b/litellm/identity/cache.py @@ -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) diff --git a/litellm/identity/invalidation.py b/litellm/identity/invalidation.py new file mode 100644 index 00000000000..e470cef054c --- /dev/null +++ b/litellm/identity/invalidation.py @@ -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) + ) diff --git a/litellm/identity/jwt.py b/litellm/identity/jwt.py new file mode 100644 index 00000000000..7089e5767fe --- /dev/null +++ b/litellm/identity/jwt.py @@ -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 diff --git a/litellm/identity/oauth2.py b/litellm/identity/oauth2.py new file mode 100644 index 00000000000..c6fd8b9aba3 --- /dev/null +++ b/litellm/identity/oauth2.py @@ -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), + ) diff --git a/litellm/identity/resolver.py b/litellm/identity/resolver.py index a2005e38627..30281447ee8 100644 --- a/litellm/identity/resolver.py +++ b/litellm/identity/resolver.py @@ -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, + ) diff --git a/litellm/identity/runtime.py b/litellm/identity/runtime.py new file mode 100644 index 00000000000..44d08a41b29 --- /dev/null +++ b/litellm/identity/runtime.py @@ -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 diff --git a/litellm/identity/store.py b/litellm/identity/store.py new file mode 100644 index 00000000000..b4652161002 --- /dev/null +++ b/litellm/identity/store.py @@ -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 diff --git a/litellm/integrations/otel/runtime.py b/litellm/integrations/otel/runtime.py index ac3b991c971..0455faf8ce6 100644 --- a/litellm/integrations/otel/runtime.py +++ b/litellm/integrations/otel/runtime.py @@ -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 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 45007861d55..9d10eb06719 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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( diff --git a/litellm/proxy/auth/oauth2_check.py b/litellm/proxy/auth/oauth2_check.py index 10b1759b77e..860f85e8959 100644 --- a/litellm/proxy/auth/oauth2_check.py +++ b/litellm/proxy/auth/oauth2_check.py @@ -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}") diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a970e0ddee8..3bb42a24cb9 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index b3a5c66e9e1..7ecc97210f2 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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 diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 99659121b27..49bca589264 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -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 diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 7a784ee4622..d1f1a297935 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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 diff --git a/tests/test_litellm/identity/test_cache.py b/tests/test_litellm/identity/test_cache.py new file mode 100644 index 00000000000..7ecd5bb509a --- /dev/null +++ b/tests/test_litellm/identity/test_cache.py @@ -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:") diff --git a/tests/test_litellm/identity/test_invalidation.py b/tests/test_litellm/identity/test_invalidation.py new file mode 100644 index 00000000000..e21e00cf34e --- /dev/null +++ b/tests/test_litellm/identity/test_invalidation.py @@ -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 diff --git a/tests/test_litellm/identity/test_jwt_builder.py b/tests/test_litellm/identity/test_jwt_builder.py new file mode 100644 index 00000000000..9da31b04f47 --- /dev/null +++ b/tests/test_litellm/identity/test_jwt_builder.py @@ -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" diff --git a/tests/test_litellm/identity/test_oauth2_builder.py b/tests/test_litellm/identity/test_oauth2_builder.py new file mode 100644 index 00000000000..84d5e24df42 --- /dev/null +++ b/tests/test_litellm/identity/test_oauth2_builder.py @@ -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" diff --git a/tests/test_litellm/identity/test_observability.py b/tests/test_litellm/identity/test_observability.py new file mode 100644 index 00000000000..c576061efeb --- /dev/null +++ b/tests/test_litellm/identity/test_observability.py @@ -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 diff --git a/tests/test_litellm/identity/test_resolver.py b/tests/test_litellm/identity/test_resolver.py index 7b333bab7d9..63b05ae1ba3 100644 --- a/tests/test_litellm/identity/test_resolver.py +++ b/tests/test_litellm/identity/test_resolver.py @@ -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" diff --git a/tests/test_litellm/identity/test_store.py b/tests/test_litellm/identity/test_store.py new file mode 100644 index 00000000000..6f9135b9ceb --- /dev/null +++ b/tests/test_litellm/identity/test_store.py @@ -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