From cfc57d709f4e6688380f1a3c37517cac4562d187 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Mon, 8 Jun 2026 11:05:44 -0700 Subject: [PATCH] fix: identity module --- litellm/identity/__init__.py | 39 +++++ litellm/identity/adapter.py | 151 ++++++++++++++++++ litellm/identity/context.py | 47 ++++++ litellm/identity/extractors/__init__.py | 7 + litellm/identity/extractors/api_key.py | 36 +++++ litellm/identity/extractors/client.py | 71 ++++++++ litellm/identity/extractors/end_user.py | 24 +++ litellm/identity/extractors/header.py | 22 +++ litellm/identity/extractors/jwt.py | 46 ++++++ litellm/identity/principal.py | 66 ++++++++ litellm/identity/resolver.py | 75 +++++++++ litellm/proxy/_types.py | 11 ++ .../identity/extractors/test_api_key.py | 52 ++++++ .../identity/extractors/test_client.py | 47 ++++++ .../identity/extractors/test_end_user.py | 45 ++++++ .../identity/extractors/test_header.py | 23 +++ .../identity/extractors/test_jwt.py | 65 ++++++++ tests/test_litellm/identity/test_adapter.py | 148 +++++++++++++++++ tests/test_litellm/identity/test_context.py | 46 ++++++ tests/test_litellm/identity/test_principal.py | 39 +++++ tests/test_litellm/identity/test_resolver.py | 70 ++++++++ 21 files changed, 1130 insertions(+) create mode 100644 litellm/identity/__init__.py create mode 100644 litellm/identity/adapter.py create mode 100644 litellm/identity/context.py create mode 100644 litellm/identity/extractors/__init__.py create mode 100644 litellm/identity/extractors/api_key.py create mode 100644 litellm/identity/extractors/client.py create mode 100644 litellm/identity/extractors/end_user.py create mode 100644 litellm/identity/extractors/header.py create mode 100644 litellm/identity/extractors/jwt.py create mode 100644 litellm/identity/principal.py create mode 100644 litellm/identity/resolver.py create mode 100644 tests/test_litellm/identity/extractors/test_api_key.py create mode 100644 tests/test_litellm/identity/extractors/test_client.py create mode 100644 tests/test_litellm/identity/extractors/test_end_user.py create mode 100644 tests/test_litellm/identity/extractors/test_header.py create mode 100644 tests/test_litellm/identity/extractors/test_jwt.py create mode 100644 tests/test_litellm/identity/test_adapter.py create mode 100644 tests/test_litellm/identity/test_context.py create mode 100644 tests/test_litellm/identity/test_principal.py create mode 100644 tests/test_litellm/identity/test_resolver.py diff --git a/litellm/identity/__init__.py b/litellm/identity/__init__.py new file mode 100644 index 00000000000..c35355a1f45 --- /dev/null +++ b/litellm/identity/__init__.py @@ -0,0 +1,39 @@ +"""Caller-identity module. + +Phase 1: domain types + extractors + bidirectional adapter to +``UserAPIKeyAuth``. The proxy still drives identity through +``litellm/proxy/auth/`` today; this module is the new home those flows +will migrate to. + +The public surface is small on purpose; downstream code should depend on +``IdentityContext`` and the ``Principal`` union, not on individual +extractor internals. +""" + +from litellm.identity.context import ( + AuditInfo, + ClientInfo, + IdentityContext, + RequestIds, +) +from litellm.identity.principal import ( + AnonymousPrincipal, + ApiKeyPrincipal, + JWTPrincipal, + Principal, + SSOPrincipal, + ServiceAccountPrincipal, +) + +__all__ = [ + "AnonymousPrincipal", + "ApiKeyPrincipal", + "AuditInfo", + "ClientInfo", + "IdentityContext", + "JWTPrincipal", + "Principal", + "RequestIds", + "SSOPrincipal", + "ServiceAccountPrincipal", +] diff --git a/litellm/identity/adapter.py b/litellm/identity/adapter.py new file mode 100644 index 00000000000..5185eea6973 --- /dev/null +++ b/litellm/identity/adapter.py @@ -0,0 +1,151 @@ +"""Bidirectional bridge between ``IdentityContext`` and ``UserAPIKeyAuth``. + +Phase 1 keeps the legacy Pydantic model as the universal carrier. These +two pure functions let new code work in terms of ``IdentityContext`` +without forcing call sites to migrate today. + +Invariants: +- ``identity_context_to_user_api_key_auth(uak.to_identity_context())`` + preserves every identity-relevant field on ``uak``. +- ``ApiKeyPrincipal.token_hash`` is treated as already-hashed; the + Pydantic ``check_api_key`` validator does not re-hash it. +""" + +from typing import TYPE_CHECKING + +from litellm.constants import ( + LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, + LITTELM_CLI_SERVICE_ACCOUNT_NAME, + LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, +) +from litellm.identity.context import AuditInfo, ClientInfo, IdentityContext, RequestIds +from litellm.identity.principal import ( + AnonymousPrincipal, + ApiKeyPrincipal, + JWTPrincipal, + Principal, + ServiceAccountPrincipal, +) + +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + +_SERVICE_ACCOUNT_NAMES = frozenset( + { + LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, + LITTELM_CLI_SERVICE_ACCOUNT_NAME, + LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, + } +) + + +def _principal_from_uak(uak: "UserAPIKeyAuth") -> Principal: + if uak.api_key in _SERVICE_ACCOUNT_NAMES or uak.key_alias in _SERVICE_ACCOUNT_NAMES: + return ServiceAccountPrincipal( + name=uak.api_key if uak.api_key in _SERVICE_ACCOUNT_NAMES else uak.key_alias # type: ignore[arg-type] + ) + + if uak.jwt_claims: + claims = uak.jwt_claims + scope_claim = claims.get("scope") or claims.get("scp") or "" + if isinstance(scope_claim, list): + scopes = [str(s) for s in scope_claim if s] + elif isinstance(scope_claim, str): + scopes = [s for s in scope_claim.split(" ") if s] + else: + scopes = [] + return JWTPrincipal( + sub=claims.get("sub"), + iss=claims.get("iss"), + aud=claims.get("aud"), + scopes=scopes, + claims=dict(claims), + mapped_user_id=uak.user_id, + mapped_team_id=uak.team_id, + mapped_org_id=uak.org_id, + ) + + if uak.token: + return ApiKeyPrincipal( + token_hash=uak.token, + key_alias=uak.key_alias, + user_id=uak.user_id, + team_id=uak.team_id, + org_id=uak.org_id, + project_id=uak.project_id, + agent_id=uak.agent_id, + ) + + return AnonymousPrincipal() + + +def user_api_key_auth_to_identity_context( + uak: "UserAPIKeyAuth", +) -> IdentityContext: + principal = _principal_from_uak(uak) + return IdentityContext( + principal=principal, + end_user_id=uak.end_user_id, + tags=[], + access_group_ids=list(uak.access_group_ids or []), + request=RequestIds(), + client=ClientInfo(), + audit=AuditInfo(), + ) + + +def identity_context_to_user_api_key_auth( + ctx: IdentityContext, +) -> "UserAPIKeyAuth": + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + principal = ctx.principal + kwargs: dict = { + "end_user_id": ctx.end_user_id, + "access_group_ids": list(ctx.access_group_ids) if ctx.access_group_ids else None, + } + + if isinstance(principal, ApiKeyPrincipal): + kwargs.update( + { + "token": principal.token_hash, + "key_alias": principal.key_alias, + "user_id": principal.user_id, + "team_id": principal.team_id, + "org_id": principal.org_id, + "project_id": principal.project_id, + "agent_id": principal.agent_id, + } + ) + elif isinstance(principal, JWTPrincipal): + kwargs.update( + { + "jwt_claims": dict(principal.claims), + "user_id": principal.mapped_user_id, + "team_id": principal.mapped_team_id, + "org_id": principal.mapped_org_id, + } + ) + elif isinstance(principal, ServiceAccountPrincipal): + if principal.name == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME: + kwargs.update( + { + "api_key": principal.name, + "team_id": "system", + "key_alias": principal.name, + "team_alias": "system", + "user_id": "system", + "user_role": LitellmUserRoles.PROXY_ADMIN, + } + ) + else: + kwargs.update( + { + "api_key": principal.name, + "team_id": principal.name, + "key_alias": principal.name, + "team_alias": principal.name, + } + ) + + return UserAPIKeyAuth(**{k: v for k, v in kwargs.items() if v is not None}) diff --git a/litellm/identity/context.py b/litellm/identity/context.py new file mode 100644 index 00000000000..386c3e051b8 --- /dev/null +++ b/litellm/identity/context.py @@ -0,0 +1,47 @@ +"""The per-request identity bundle. + +``IdentityContext`` is what downstream consumers (auth, spend, guardrails, +logging, audit) should read identity from once Phase 2 migration is done. +In Phase 1 it travels alongside the legacy ``UserAPIKeyAuth`` via the +adapter functions in ``litellm.identity.adapter``. + +The bundle is mutable on purpose: identity fields like ``end_user_id`` are +sometimes resolved or overridden after initial extraction, and the +existing ``UserAPIKeyAuth`` mutation patterns must keep working. +""" + +from dataclasses import dataclass, field +from typing import List, Optional + +from litellm.identity.principal import AnonymousPrincipal, Principal + + +@dataclass +class RequestIds: + request_id: Optional[str] = None + trace_id: Optional[str] = None + session_id: Optional[str] = None + mcp_session_id: Optional[str] = None + + +@dataclass +class ClientInfo: + ip: Optional[str] = None + user_agent: Optional[str] = None + forwarded_chain: List[str] = field(default_factory=list) + + +@dataclass +class AuditInfo: + changed_by: Optional[str] = None + + +@dataclass +class IdentityContext: + principal: Principal = field(default_factory=AnonymousPrincipal) + end_user_id: Optional[str] = None + tags: List[str] = field(default_factory=list) + access_group_ids: List[str] = field(default_factory=list) + request: RequestIds = field(default_factory=RequestIds) + client: ClientInfo = field(default_factory=ClientInfo) + audit: AuditInfo = field(default_factory=AuditInfo) diff --git a/litellm/identity/extractors/__init__.py b/litellm/identity/extractors/__init__.py new file mode 100644 index 00000000000..116b27a9190 --- /dev/null +++ b/litellm/identity/extractors/__init__.py @@ -0,0 +1,7 @@ +"""Identity extractors. + +Each extractor wraps an existing helper in ``litellm/proxy/auth/`` and +returns a piece of an ``IdentityContext``. Extractors must not introduce +new behavior. If you need to change *how* a field is resolved, change the +underlying helper and update the extractor's tests. +""" diff --git a/litellm/identity/extractors/api_key.py b/litellm/identity/extractors/api_key.py new file mode 100644 index 00000000000..e84e4e01e38 --- /dev/null +++ b/litellm/identity/extractors/api_key.py @@ -0,0 +1,36 @@ +"""API-key principal extraction. + +Wraps the existing ``get_api_key`` (header extraction) and +``UserAPIKeyAuth._safe_hash_litellm_api_key`` (hashing, including the +"hashed-jwt-..." prefix for JWT-shaped values). +""" + +from typing import Optional + +from litellm.identity.principal import ApiKeyPrincipal + + +def hash_principal_token(api_key: str) -> str: + """Hash an API key the same way the legacy auth path does. + + Centralized so future principal types (e.g. SSO-issued ephemeral + keys) can reuse the same hashing without re-importing the Pydantic + model. + """ + from litellm.proxy._types import UserAPIKeyAuth + + return UserAPIKeyAuth._safe_hash_litellm_api_key(api_key) + + +def extract_api_key_principal(api_key: Optional[str]) -> Optional[ApiKeyPrincipal]: + """Build an ``ApiKeyPrincipal`` from a raw API key string. + + Returns ``None`` when no key is supplied. Callers that need to pull + the key out of a FastAPI request should call ``get_api_key`` from + ``litellm.proxy.auth.user_api_key_auth`` first; this extractor stays + framework-agnostic on purpose so it can be reused by non-FastAPI + entrypoints (CLI, MCP, background jobs). + """ + if not api_key: + return None + return ApiKeyPrincipal(token_hash=hash_principal_token(api_key)) diff --git a/litellm/identity/extractors/client.py b/litellm/identity/extractors/client.py new file mode 100644 index 00000000000..53da72dc12a --- /dev/null +++ b/litellm/identity/extractors/client.py @@ -0,0 +1,71 @@ +"""Client/network identity extraction. + +Builds a ``ClientInfo`` from a FastAPI request. ``X-Forwarded-For`` is +only honored when the direct peer is in a configured trusted-proxy CIDR. +The trust logic is delegated to ``IPAddressUtils.is_request_from_trusted_proxy`` +so we stay in sync with the rest of the proxy. +""" + +from typing import Any, Dict, List, Optional + +from litellm.identity.context import ClientInfo + + +def _split_forwarded_chain(raw: Optional[str]) -> List[str]: + if not raw or not isinstance(raw, str): + return [] + return [hop.strip() for hop in raw.split(",") if hop.strip()] + + +def _direct_client_host(request: Any) -> Optional[str]: + client = getattr(request, "client", None) + host = getattr(client, "host", None) + if isinstance(host, str) and host: + return host + return None + + +def extract_client_info( + request: Any, + general_settings: Optional[Dict[str, Any]] = None, +) -> ClientInfo: + headers = getattr(request, "headers", {}) or {} + forwarded_chain: List[str] = [] + ip: Optional[str] = None + + # Headers in FastAPI are case-insensitive; normalize defensively for dicts. + def _get_header(name: str) -> Optional[str]: + try: + value = headers.get(name) + except AttributeError: + return None + if value is not None: + return value + try: + for key, val in headers.items(): + if isinstance(key, str) and key.lower() == name: + return val + except Exception: + return None + return None + + xff = _get_header("x-forwarded-for") + if xff: + forwarded_chain = _split_forwarded_chain(xff) + + from litellm.proxy.auth.ip_address_utils import IPAddressUtils + + if forwarded_chain and IPAddressUtils.is_request_from_trusted_proxy( + request=request, general_settings=general_settings + ): + ip = forwarded_chain[0] + else: + ip = _direct_client_host(request) + + user_agent = _get_header("user-agent") + + return ClientInfo( + ip=ip, + user_agent=user_agent if isinstance(user_agent, str) else None, + forwarded_chain=forwarded_chain, + ) diff --git a/litellm/identity/extractors/end_user.py b/litellm/identity/extractors/end_user.py new file mode 100644 index 00000000000..22acd98527c --- /dev/null +++ b/litellm/identity/extractors/end_user.py @@ -0,0 +1,24 @@ +"""End-user extraction. + +Thin wrapper over the existing six-check chain in +``litellm.proxy.auth.auth_utils.get_end_user_id_from_request_body``. +Validation against the DB stays in ``resolve_and_validate_end_user_id`` +and runs from the legacy auth path; this extractor returns the raw +identifier only. +""" + +from typing import Optional + + +def extract_end_user_id( + body: Optional[dict], + headers: Optional[dict] = None, +) -> Optional[str]: + if body is None: + body = {} + + from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body + + return get_end_user_id_from_request_body( + request_body=body, request_headers=headers + ) diff --git a/litellm/identity/extractors/header.py b/litellm/identity/extractors/header.py new file mode 100644 index 00000000000..f47cfd943c5 --- /dev/null +++ b/litellm/identity/extractors/header.py @@ -0,0 +1,22 @@ +"""Header-driven identity extractors. + +These pull non-credential identity fields out of request headers. They +do not perform authorization decisions; that stays in the auth chain. +""" + +from typing import Optional + + +AUDIT_CHANGED_BY_HEADER = "litellm-changed-by" + + +def extract_audit_changed_by(headers: Optional[dict]) -> Optional[str]: + """Read the ``litellm-changed-by`` header used for management-API audit.""" + if not headers: + return None + + for key, value in headers.items(): + if isinstance(key, str) and key.lower() == AUDIT_CHANGED_BY_HEADER: + if isinstance(value, str) and value: + return value + return None diff --git a/litellm/identity/extractors/jwt.py b/litellm/identity/extractors/jwt.py new file mode 100644 index 00000000000..31bfa5dec5d --- /dev/null +++ b/litellm/identity/extractors/jwt.py @@ -0,0 +1,46 @@ +"""JWT principal extraction. + +Uses ``JWTHandler.is_jwt`` for shape detection and +``JWTHandler.get_unverified_claims`` for claim peek. Signature +verification and DB-backed claim mapping live in ``JWTHandler.auth_builder``; +that path runs from the resolver when DB access is available. +""" + +from typing import Optional + +from litellm.identity.principal import JWTPrincipal + + +def extract_jwt_principal(token: Optional[str]) -> Optional[JWTPrincipal]: + """Decode JWT claims without verification and build a ``JWTPrincipal``. + + Returns ``None`` when the token is missing or not JWT-shaped. The + caller is responsible for invoking ``auth_builder`` (or another + verifier) before trusting the principal for authorization. + """ + if not token: + return None + + from litellm.proxy.auth.handle_jwt import JWTHandler + + if not JWTHandler.is_jwt(token=token): + return None + + claims = JWTHandler.get_unverified_claims(token=token) or {} + + aud = claims.get("aud") + scope_claim = claims.get("scope") or claims.get("scp") or "" + if isinstance(scope_claim, list): + scopes = [str(s) for s in scope_claim if s] + elif isinstance(scope_claim, str): + scopes = [s for s in scope_claim.split(" ") if s] + else: + scopes = [] + + return JWTPrincipal( + sub=claims.get("sub"), + iss=claims.get("iss"), + aud=aud, + scopes=scopes, + claims=claims, + ) diff --git a/litellm/identity/principal.py b/litellm/identity/principal.py new file mode 100644 index 00000000000..41c69c09877 --- /dev/null +++ b/litellm/identity/principal.py @@ -0,0 +1,66 @@ +"""Caller-identity primitives. + +A ``Principal`` answers "who is making this request" using only the fields +that uniquely identify the caller. Per-row enrichment (budgets, team rows, +object permissions) is intentionally not modeled here; that data continues +to ride on ``UserAPIKeyAuth`` in Phase 1. + +Each subtype is a frozen dataclass with a ``kind`` discriminator suitable +for ``match``-style dispatch. +""" + +from dataclasses import dataclass, field +from typing import Any, Dict, List, Literal, Optional, Union + + +@dataclass(frozen=True) +class ApiKeyPrincipal: + kind: Literal["api_key"] = field(default="api_key", init=False) + token_hash: str + key_alias: Optional[str] = None + user_id: Optional[str] = None + team_id: Optional[str] = None + org_id: Optional[str] = None + project_id: Optional[str] = None + agent_id: Optional[str] = None + + +@dataclass(frozen=True) +class JWTPrincipal: + kind: Literal["jwt"] = field(default="jwt", init=False) + sub: Optional[str] = None + iss: Optional[str] = None + aud: Optional[Union[str, List[str]]] = None + scopes: List[str] = field(default_factory=list) + claims: Dict[str, Any] = field(default_factory=dict) + mapped_user_id: Optional[str] = None + mapped_team_id: Optional[str] = None + mapped_org_id: Optional[str] = None + + +@dataclass(frozen=True) +class SSOPrincipal: + kind: Literal["sso"] = field(default="sso", init=False) + sso_user_id: str + email: Optional[str] = None + provider: Optional[str] = None + + +@dataclass(frozen=True) +class ServiceAccountPrincipal: + kind: Literal["service_account"] = field(default="service_account", init=False) + name: str + + +@dataclass(frozen=True) +class AnonymousPrincipal: + kind: Literal["anonymous"] = field(default="anonymous", init=False) + + +Principal = Union[ + ApiKeyPrincipal, + JWTPrincipal, + SSOPrincipal, + ServiceAccountPrincipal, + AnonymousPrincipal, +] diff --git a/litellm/identity/resolver.py b/litellm/identity/resolver.py new file mode 100644 index 00000000000..a2005e38627 --- /dev/null +++ b/litellm/identity/resolver.py @@ -0,0 +1,75 @@ +"""Compose extractors 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. +""" + +from typing import Any, Dict, Optional + +from litellm.constants import ( + LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, + LITTELM_CLI_SERVICE_ACCOUNT_NAME, + 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.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 +from litellm.identity.extractors.jwt import extract_jwt_principal +from litellm.identity.principal import ( + AnonymousPrincipal, + Principal, + ServiceAccountPrincipal, +) + +_SERVICE_ACCOUNT_API_KEYS = frozenset( + { + LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, + LITTELM_CLI_SERVICE_ACCOUNT_NAME, + LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, + } +) + + +def _resolve_principal(api_key: Optional[str]) -> Principal: + if api_key and api_key in _SERVICE_ACCOUNT_API_KEYS: + return ServiceAccountPrincipal(name=api_key) + + jwt_principal = extract_jwt_principal(api_key) + if jwt_principal is not None: + return jwt_principal + + api_key_principal = extract_api_key_principal(api_key) + if api_key_principal is not None: + return api_key_principal + + return AnonymousPrincipal() + + +def resolve_identity( + *, + api_key: Optional[str] = None, + request: Any = None, + body: Optional[Dict[str, Any]] = None, + headers: Optional[Dict[str, Any]] = None, + general_settings: Optional[Dict[str, Any]] = None, +) -> IdentityContext: + principal = _resolve_principal(api_key) + end_user_id = extract_end_user_id(body=body, headers=headers) + audit = AuditInfo(changed_by=extract_audit_changed_by(headers)) + client: ClientInfo + if request is not None: + client = extract_client_info( + request=request, general_settings=general_settings + ) + else: + client = ClientInfo() + + return IdentityContext( + principal=principal, + end_user_id=end_user_id, + audit=audit, + client=client, + ) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 57a7d860baa..a2ca9c12382 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2539,6 +2539,17 @@ class UserAPIKeyAuth( user_role=LitellmUserRoles.PROXY_ADMIN, ) + def to_identity_context(self): + from litellm.identity.adapter import user_api_key_auth_to_identity_context + + return user_api_key_auth_to_identity_context(self) + + @classmethod + def from_identity_context(cls, ctx) -> "UserAPIKeyAuth": + from litellm.identity.adapter import identity_context_to_user_api_key_auth + + return identity_context_to_user_api_key_auth(ctx) + def user_api_key_has_admin_view(user_api_key_dict: UserAPIKeyAuth) -> bool: """Return True if the caller's role grants unscoped read access to all diff --git a/tests/test_litellm/identity/extractors/test_api_key.py b/tests/test_litellm/identity/extractors/test_api_key.py new file mode 100644 index 00000000000..a7cbd30b47e --- /dev/null +++ b/tests/test_litellm/identity/extractors/test_api_key.py @@ -0,0 +1,52 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.identity.extractors.api_key import ( + extract_api_key_principal, + hash_principal_token, +) +from litellm.proxy._types import UserAPIKeyAuth + + +def test_returns_none_for_empty_input(): + assert extract_api_key_principal(None) is None + assert extract_api_key_principal("") is None + + +def test_sk_key_is_hashed_with_legacy_helper(): + raw = "sk-abc123" + principal = extract_api_key_principal(raw) + assert principal is not None + assert principal.token_hash == UserAPIKeyAuth._safe_hash_litellm_api_key(raw) + assert principal.token_hash != raw + + +def test_bearer_prefix_normalized_then_hashed(): + raw = "Bearer sk-abc123" + principal = extract_api_key_principal(raw) + assert principal is not None + assert principal.token_hash == UserAPIKeyAuth._safe_hash_litellm_api_key(raw) + sk_only_hash = UserAPIKeyAuth._safe_hash_litellm_api_key("sk-abc123") + assert principal.token_hash == sk_only_hash + + +def test_jwt_shaped_token_gets_hashed_jwt_prefix(): + fake_jwt = "aaaa.bbbb.cccc" + principal = extract_api_key_principal(fake_jwt) + assert principal is not None + assert principal.token_hash.startswith("hashed-jwt-") + + +def test_non_sk_non_jwt_returned_unhashed(): + raw = "custom-key-without-prefix" + principal = extract_api_key_principal(raw) + assert principal is not None + assert principal.token_hash == raw + + +def test_hash_helper_delegates_to_legacy_path(): + assert hash_principal_token("sk-x") == UserAPIKeyAuth._safe_hash_litellm_api_key( + "sk-x" + ) diff --git a/tests/test_litellm/identity/extractors/test_client.py b/tests/test_litellm/identity/extractors/test_client.py new file mode 100644 index 00000000000..e0d9c2f0cbe --- /dev/null +++ b/tests/test_litellm/identity/extractors/test_client.py @@ -0,0 +1,47 @@ +import os +import sys +from types import SimpleNamespace + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.identity.extractors.client import extract_client_info + + +def _fake_request(headers, client_host=None): + return SimpleNamespace( + headers=headers, + client=SimpleNamespace(host=client_host) if client_host else None, + ) + + +def test_falls_back_to_direct_peer_when_not_trusted(): + req = _fake_request({"x-forwarded-for": "9.9.9.9"}, client_host="10.0.0.1") + info = extract_client_info(req, general_settings={}) + assert info.ip == "10.0.0.1" + assert info.forwarded_chain == ["9.9.9.9"] + + +def test_uses_xff_first_hop_when_proxy_trusted(): + req = _fake_request( + {"x-forwarded-for": "1.2.3.4, 10.0.0.1"}, client_host="10.0.0.1" + ) + settings = { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["10.0.0.0/8"], + } + info = extract_client_info(req, general_settings=settings) + assert info.ip == "1.2.3.4" + assert info.forwarded_chain == ["1.2.3.4", "10.0.0.1"] + + +def test_no_xff_returns_direct_peer(): + req = _fake_request({}, client_host="127.0.0.1") + info = extract_client_info(req, general_settings={}) + assert info.ip == "127.0.0.1" + assert info.forwarded_chain == [] + + +def test_user_agent_passthrough(): + req = _fake_request({"user-agent": "curl/8"}, client_host="127.0.0.1") + info = extract_client_info(req, general_settings={}) + assert info.user_agent == "curl/8" diff --git a/tests/test_litellm/identity/extractors/test_end_user.py b/tests/test_litellm/identity/extractors/test_end_user.py new file mode 100644 index 00000000000..8eda186e41f --- /dev/null +++ b/tests/test_litellm/identity/extractors/test_end_user.py @@ -0,0 +1,45 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +import pytest + +from litellm.identity.extractors.end_user import extract_end_user_id + + +def test_user_body_field(): + assert extract_end_user_id({"user": "eu-1"}, {}) == "eu-1" + + +def test_litellm_metadata_user(): + assert ( + extract_end_user_id({"litellm_metadata": {"user": "eu-meta"}}, {}) == "eu-meta" + ) + + +def test_metadata_user_id(): + assert extract_end_user_id({"metadata": {"user_id": "eu-md"}}, {}) == "eu-md" + + +def test_safety_identifier_fallback(): + assert extract_end_user_id({"safety_identifier": "eu-safety"}, {}) == "eu-safety" + + +def test_returns_none_when_empty(): + assert extract_end_user_id(None, None) is None + assert extract_end_user_id({}, {}) is None + + +def test_user_field_wins_over_metadata(): + body = {"user": "eu-primary", "metadata": {"user_id": "eu-secondary"}} + assert extract_end_user_id(body, {}) == "eu-primary" + + +def test_anthropic_standard_customer_id_header(monkeypatch): + from litellm.constants import STANDARD_CUSTOMER_ID_HEADERS + + if not STANDARD_CUSTOMER_ID_HEADERS: + pytest.skip("no standard customer headers configured") + header_name = STANDARD_CUSTOMER_ID_HEADERS[0] + assert extract_end_user_id({}, {header_name: "eu-hdr"}) == "eu-hdr" diff --git a/tests/test_litellm/identity/extractors/test_header.py b/tests/test_litellm/identity/extractors/test_header.py new file mode 100644 index 00000000000..cbc9ff965fc --- /dev/null +++ b/tests/test_litellm/identity/extractors/test_header.py @@ -0,0 +1,23 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.identity.extractors.header import extract_audit_changed_by + + +def test_returns_value_when_present(): + assert extract_audit_changed_by({"litellm-changed-by": "alice"}) == "alice" + + +def test_case_insensitive(): + assert extract_audit_changed_by({"Litellm-Changed-By": "bob"}) == "bob" + + +def test_returns_none_when_missing(): + assert extract_audit_changed_by({}) is None + assert extract_audit_changed_by(None) is None + + +def test_empty_string_value_is_none(): + assert extract_audit_changed_by({"litellm-changed-by": ""}) is None diff --git a/tests/test_litellm/identity/extractors/test_jwt.py b/tests/test_litellm/identity/extractors/test_jwt.py new file mode 100644 index 00000000000..b3ab9aac6ba --- /dev/null +++ b/tests/test_litellm/identity/extractors/test_jwt.py @@ -0,0 +1,65 @@ +import base64 +import json +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.identity.extractors.jwt import extract_jwt_principal + + +def _b64url(payload: dict) -> str: + raw = json.dumps(payload).encode() + return base64.urlsafe_b64encode(raw).rstrip(b"=").decode() + + +def _build_unverified_jwt(claims: dict) -> str: + header = _b64url({"alg": "HS256", "typ": "JWT"}) + body = _b64url(claims) + return f"{header}.{body}.sig" + + +def test_returns_none_for_non_jwt(): + assert extract_jwt_principal("sk-abc") is None + assert extract_jwt_principal(None) is None + assert extract_jwt_principal("") is None + + +def test_extracts_sub_iss_aud(): + token = _build_unverified_jwt( + {"sub": "user-1", "iss": "https://idp.example", "aud": "litellm"} + ) + p = extract_jwt_principal(token) + assert p is not None + assert p.sub == "user-1" + assert p.iss == "https://idp.example" + assert p.aud == "litellm" + + +def test_string_scope_is_split(): + token = _build_unverified_jwt({"sub": "u", "scope": "read write admin"}) + p = extract_jwt_principal(token) + assert p is not None + assert p.scopes == ["read", "write", "admin"] + + +def test_list_scope_is_preserved(): + token = _build_unverified_jwt({"sub": "u", "scope": ["a", "b"]}) + p = extract_jwt_principal(token) + assert p is not None + assert p.scopes == ["a", "b"] + + +def test_scp_claim_supported(): + token = _build_unverified_jwt({"sub": "u", "scp": "read"}) + p = extract_jwt_principal(token) + assert p is not None + assert p.scopes == ["read"] + + +def test_raw_claims_preserved(): + claims = {"sub": "u", "custom": {"groups": ["g1"]}} + token = _build_unverified_jwt(claims) + p = extract_jwt_principal(token) + assert p is not None + assert p.claims["custom"] == {"groups": ["g1"]} diff --git a/tests/test_litellm/identity/test_adapter.py b/tests/test_litellm/identity/test_adapter.py new file mode 100644 index 00000000000..fd34e1b2af9 --- /dev/null +++ b/tests/test_litellm/identity/test_adapter.py @@ -0,0 +1,148 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.constants import ( + LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, + LITTELM_CLI_SERVICE_ACCOUNT_NAME, + LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, +) +from litellm.identity import ( + AnonymousPrincipal, + ApiKeyPrincipal, + IdentityContext, + JWTPrincipal, + ServiceAccountPrincipal, +) +from litellm.identity.adapter import ( + identity_context_to_user_api_key_auth, + user_api_key_auth_to_identity_context, +) +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + +def test_roundtrip_preserves_api_key_identity_fields(): + uak = UserAPIKeyAuth( + api_key="sk-abc", + user_id="u1", + team_id="t1", + org_id="o1", + project_id="p1", + agent_id="a1", + key_alias="alias-1", + end_user_id="eu-1", + access_group_ids=["g1", "g2"], + ) + ctx = user_api_key_auth_to_identity_context(uak) + assert isinstance(ctx.principal, ApiKeyPrincipal) + assert ctx.principal.user_id == "u1" + assert ctx.principal.team_id == "t1" + assert ctx.principal.org_id == "o1" + assert ctx.principal.project_id == "p1" + assert ctx.principal.agent_id == "a1" + assert ctx.principal.key_alias == "alias-1" + assert ctx.end_user_id == "eu-1" + assert ctx.access_group_ids == ["g1", "g2"] + + back = identity_context_to_user_api_key_auth(ctx) + assert back.user_id == "u1" + assert back.team_id == "t1" + assert back.org_id == "o1" + assert back.project_id == "p1" + assert back.agent_id == "a1" + assert back.key_alias == "alias-1" + assert back.end_user_id == "eu-1" + assert back.access_group_ids == ["g1", "g2"] + assert back.token == uak.token + + +def test_token_hash_not_double_hashed(): + ctx = IdentityContext(principal=ApiKeyPrincipal(token_hash="abc123")) + back = identity_context_to_user_api_key_auth(ctx) + assert back.token == "abc123" + + +def test_jwt_principal_roundtrip(): + uak = UserAPIKeyAuth( + api_key="aaaa.bbbb.cccc", + user_id="jwt-user", + team_id="jwt-team", + org_id="jwt-org", + jwt_claims={"sub": "jwt-user", "iss": "idp", "scope": "read write"}, + ) + ctx = user_api_key_auth_to_identity_context(uak) + assert isinstance(ctx.principal, JWTPrincipal) + assert ctx.principal.sub == "jwt-user" + assert ctx.principal.iss == "idp" + assert ctx.principal.scopes == ["read", "write"] + assert ctx.principal.mapped_user_id == "jwt-user" + assert ctx.principal.mapped_team_id == "jwt-team" + assert ctx.principal.mapped_org_id == "jwt-org" + + back = identity_context_to_user_api_key_auth(ctx) + assert back.user_id == "jwt-user" + assert back.team_id == "jwt-team" + assert back.org_id == "jwt-org" + assert back.jwt_claims is not None + assert back.jwt_claims.get("sub") == "jwt-user" + + +def test_service_account_jobs_principal_roundtrip(): + uak = UserAPIKeyAuth.get_litellm_internal_jobs_user_api_key_auth() + ctx = user_api_key_auth_to_identity_context(uak) + assert isinstance(ctx.principal, ServiceAccountPrincipal) + assert ctx.principal.name == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + + back = identity_context_to_user_api_key_auth(ctx) + assert back.user_id == "system" + assert back.team_id == "system" + assert back.user_role == LitellmUserRoles.PROXY_ADMIN + assert back.key_alias == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + + +def test_service_account_health_check_roundtrip(): + uak = UserAPIKeyAuth.get_litellm_internal_health_check_user_api_key_auth() + ctx = user_api_key_auth_to_identity_context(uak) + assert isinstance(ctx.principal, ServiceAccountPrincipal) + assert ctx.principal.name == LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME + + back = identity_context_to_user_api_key_auth(ctx) + assert back.team_id == LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME + assert back.team_alias == LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME + + +def test_service_account_cli_roundtrip(): + uak = UserAPIKeyAuth.get_litellm_cli_user_api_key_auth() + ctx = user_api_key_auth_to_identity_context(uak) + assert isinstance(ctx.principal, ServiceAccountPrincipal) + assert ctx.principal.name == LITTELM_CLI_SERVICE_ACCOUNT_NAME + + +def test_anonymous_principal_when_no_token(): + uak = UserAPIKeyAuth() + ctx = user_api_key_auth_to_identity_context(uak) + assert isinstance(ctx.principal, AnonymousPrincipal) + back = identity_context_to_user_api_key_auth(ctx) + assert back.token is None + assert back.user_id is None + + +def test_end_user_does_not_leak_into_principal(): + ctx = IdentityContext( + principal=ApiKeyPrincipal(token_hash="t", user_id="u"), + end_user_id="customer-99", + ) + back = identity_context_to_user_api_key_auth(ctx) + assert back.user_id == "u" + assert back.end_user_id == "customer-99" + assert ctx.principal.user_id == "u" + + +def test_uak_methods_delegate_to_adapter(): + uak = UserAPIKeyAuth(api_key="sk-test", user_id="u", team_id="t") + ctx = uak.to_identity_context() + assert isinstance(ctx, IdentityContext) + back = UserAPIKeyAuth.from_identity_context(ctx) + assert back.user_id == "u" + assert back.team_id == "t" diff --git a/tests/test_litellm/identity/test_context.py b/tests/test_litellm/identity/test_context.py new file mode 100644 index 00000000000..297876bcc97 --- /dev/null +++ b/tests/test_litellm/identity/test_context.py @@ -0,0 +1,46 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.identity.context import ( + AuditInfo, + ClientInfo, + IdentityContext, + RequestIds, +) +from litellm.identity.principal import AnonymousPrincipal, ApiKeyPrincipal + + +def test_default_principal_is_anonymous(): + ctx = IdentityContext() + assert isinstance(ctx.principal, AnonymousPrincipal) + assert ctx.end_user_id is None + assert ctx.tags == [] + assert ctx.access_group_ids == [] + assert ctx.request == RequestIds() + assert ctx.client == ClientInfo() + assert ctx.audit == AuditInfo() + + +def test_context_is_mutable(): + ctx = IdentityContext() + ctx.end_user_id = "eu-1" + ctx.tags.append("env:prod") + assert ctx.end_user_id == "eu-1" + assert "env:prod" in ctx.tags + + +def test_context_carries_principal(): + p = ApiKeyPrincipal(token_hash="abc", user_id="u1", team_id="t1") + ctx = IdentityContext(principal=p) + assert ctx.principal is p + + +def test_subobjects_are_independent_between_instances(): + a = IdentityContext() + b = IdentityContext() + a.tags.append("x") + assert b.tags == [] + a.client.forwarded_chain.append("1.2.3.4") + assert b.client.forwarded_chain == [] diff --git a/tests/test_litellm/identity/test_principal.py b/tests/test_litellm/identity/test_principal.py new file mode 100644 index 00000000000..d3ef95888ec --- /dev/null +++ b/tests/test_litellm/identity/test_principal.py @@ -0,0 +1,39 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../..")) + +import pytest + +from litellm.identity.principal import ( + AnonymousPrincipal, + ApiKeyPrincipal, + JWTPrincipal, + SSOPrincipal, + ServiceAccountPrincipal, +) + + +def test_principal_kind_discriminators_are_fixed(): + assert ApiKeyPrincipal(token_hash="x").kind == "api_key" + assert JWTPrincipal().kind == "jwt" + assert SSOPrincipal(sso_user_id="s").kind == "sso" + assert ServiceAccountPrincipal(name="n").kind == "service_account" + assert AnonymousPrincipal().kind == "anonymous" + + +def test_principals_are_frozen(): + p = ApiKeyPrincipal(token_hash="x", user_id="u1") + with pytest.raises(Exception): + p.user_id = "u2" # type: ignore[misc] + + +def test_principals_are_hashable(): + a = ApiKeyPrincipal(token_hash="x", user_id="u1") + b = ApiKeyPrincipal(token_hash="x", user_id="u1") + assert {a, b} == {a} + + +def test_kind_is_not_constructor_arg(): + with pytest.raises(TypeError): + ApiKeyPrincipal(kind="api_key", token_hash="x") # type: ignore[call-arg] diff --git a/tests/test_litellm/identity/test_resolver.py b/tests/test_litellm/identity/test_resolver.py new file mode 100644 index 00000000000..7b333bab7d9 --- /dev/null +++ b/tests/test_litellm/identity/test_resolver.py @@ -0,0 +1,70 @@ +import base64 +import json +import os +import sys +from types import SimpleNamespace + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.constants import LITTELM_CLI_SERVICE_ACCOUNT_NAME +from litellm.identity.principal import ( + AnonymousPrincipal, + ApiKeyPrincipal, + JWTPrincipal, + ServiceAccountPrincipal, +) +from litellm.identity.resolver import resolve_identity + + +def _fake_request(headers=None, client_host=None): + return SimpleNamespace( + headers=headers or {}, + client=SimpleNamespace(host=client_host) if client_host else None, + ) + + +def _jwt(claims): + def b(d): + return ( + base64.urlsafe_b64encode(json.dumps(d).encode()).rstrip(b"=").decode() + ) + + return f"{b({'alg':'HS256','typ':'JWT'})}.{b(claims)}.sig" + + +def test_anonymous_when_no_credentials(): + ctx = resolve_identity() + assert isinstance(ctx.principal, AnonymousPrincipal) + + +def test_api_key_principal_for_sk_key(): + ctx = 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"})) + 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) + assert isinstance(ctx.principal, ServiceAccountPrincipal) + assert ctx.principal.name == LITTELM_CLI_SERVICE_ACCOUNT_NAME + + +def test_end_user_and_audit_propagate(): + ctx = resolve_identity( + body={"user": "eu-42"}, + headers={"litellm-changed-by": "admin"}, + ) + assert ctx.end_user_id == "eu-42" + assert ctx.audit.changed_by == "admin" + + +def test_client_info_from_request(): + req = _fake_request({"user-agent": "curl/8"}, client_host="127.0.0.1") + ctx = resolve_identity(request=req) + assert ctx.client.ip == "127.0.0.1" + assert ctx.client.user_agent == "curl/8"