fix: identity module

This commit is contained in:
Yassin Kortam 2026-06-08 11:05:44 -07:00
parent f5b11b72a6
commit cfc57d709f
21 changed files with 1130 additions and 0 deletions

View file

@ -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",
]

151
litellm/identity/adapter.py Normal file
View file

@ -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})

View file

@ -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)

View file

@ -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.
"""

View file

@ -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))

View file

@ -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,
)

View file

@ -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
)

View file

@ -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

View file

@ -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,
)

View file

@ -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,
]

View file

@ -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,
)

View file

@ -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

View file

@ -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"
)

View file

@ -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"

View file

@ -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"

View file

@ -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

View file

@ -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"]}

View file

@ -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"

View file

@ -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 == []

View file

@ -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]

View file

@ -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"