feat(auth_v2): add standards-based auth and identity module

New additive litellm/auth_v2 package: a thin orchestration layer over PyJWT,
Authlib and scim2-models behind FastAPI's native Security() primitives that
normalizes every credential into one standards-shaped Principal carrying
org/team/user and network identity.

Authentication, identity resolution, authorization and enforcement are kept
as separate layers. Five authenticators cover the OpenAPI scheme types
(apiKey, http bearer-JWT/basic, oauth2 at+jwt + introspection, openIdConnect,
mutualTLS); a shared JwtVerifier enforces signature, issuer, audience and exp
on every JWT path via a cached PyJWKClient. RBAC is a flat Role enum plus
scope/role checks wired through SecurityScopes. Missing or invalid credentials
return 401 with an RFC 9110/6750 WWW-Authenticate challenge, scope failures
return 403 insufficient_scope. SCIM 2.0 Users/Groups/PATCH/discovery and an
Authlib OIDC login flow share one ProvisioningStore seam; the SAML SP is a
documented thin adapter pending pysaml2.

The module is unimported by the proxy app and depends on nothing in
litellm/proxy/auth.
This commit is contained in:
Yassin Kortam 2026-06-10 17:22:20 -07:00
parent 2bbf688613
commit a0a59a2197
12 changed files with 1220 additions and 0 deletions

View file

@ -0,0 +1,11 @@
from .config import AuthConfig
from .models import Principal
from .security import get_current_principal, install_auth, require_roles
__all__ = [
"Principal",
"AuthConfig",
"get_current_principal",
"require_roles",
"install_auth",
]

View file

@ -0,0 +1,347 @@
from __future__ import annotations
import base64
import binascii
from typing import Any, Dict, List, Optional, Protocol, runtime_checkable
import httpx
import jwt
from fastapi import Request
from jwt import PyJWKClient
from jwt import decode as jwt_decode
from . import errors
from .config import (
ApiKeySchemeConfig,
AuthConfig,
HttpBasicConfig,
MutualTlsConfig,
OAuth2IntrospectionConfig,
OidcProviderConfig,
)
from .models import (
AuthMethod,
ClientCertificate,
Credential,
CredentialRef,
SecuritySchemeType,
)
AT_JWT_TYPES = {"at+jwt", "application/at+jwt"}
@runtime_checkable
class Authenticator(Protocol):
scheme: SecuritySchemeType
async def authenticate(self, request: Request) -> Optional[Credential]: ...
def challenge(self) -> str: ...
def _extract_bearer(request: Request) -> Optional[str]:
header = request.headers.get("authorization")
if not header:
return None
scheme, _, value = header.partition(" ")
if scheme.lower() != "bearer" or not value:
return None
return value
def _looks_like_jwt(token: str) -> bool:
return token.count(".") == 2
def _normalize_audience(value: Any) -> List[str]:
if isinstance(value, str):
return [value]
if isinstance(value, list):
return [str(item) for item in value]
return []
def _split_scope(value: Any) -> List[str]:
return value.split() if isinstance(value, str) else []
def _credential_from_claims(
scheme: SecuritySchemeType,
method: AuthMethod,
token: str,
claims: Dict[str, Any],
) -> Credential:
header = jwt.get_unverified_header(token)
return Credential(
scheme=scheme,
method=method,
subject=str(claims.get("sub", "")),
issuer=claims.get("iss"),
audience=_normalize_audience(claims.get("aud")),
scopes=_split_scope(claims.get("scope")),
claims=claims,
credential_ref=CredentialRef(
key_id=header.get("kid"), token_id=claims.get("jti")
),
)
class JwtVerifier:
def __init__(
self,
provider: OidcProviderConfig,
jwks_client: Optional[PyJWKClient] = None,
) -> None:
self.provider = provider
if jwks_client is not None:
self._jwks_client = jwks_client
return
jwks_uri = (
str(provider.jwks_uri) if provider.jwks_uri else self._discover_jwks()
)
self._jwks_client = PyJWKClient(jwks_uri, cache_keys=True)
def _discover_jwks(self) -> str:
url = f"{self.provider.issuer.rstrip('/')}/.well-known/openid-configuration"
response = httpx.get(url, timeout=10.0)
response.raise_for_status()
jwks_uri = response.json().get("jwks_uri")
if not jwks_uri:
raise ValueError(f"discovery document missing jwks_uri: {url}")
return str(jwks_uri)
def verify(
self, token: str, *, require_at_jwt: Optional[bool] = None
) -> Dict[str, Any]:
enforce = (
self.provider.require_at_jwt if require_at_jwt is None else require_at_jwt
)
if enforce:
header = jwt.get_unverified_header(token)
if str(header.get("typ", "")).lower() not in AT_JWT_TYPES:
raise errors.invalid_token("token typ must be at+jwt")
try:
signing_key = self._jwks_client.get_signing_key_from_jwt(token)
return jwt_decode(
token,
signing_key.key,
algorithms=self.provider.algorithms,
audience=self.provider.audience,
issuer=self.provider.issuer,
options={"verify_exp": True, "require": ["exp", "iss", "aud"]},
)
except jwt.PyJWTError as exc:
raise errors.invalid_token(str(exc)) from exc
def _select_verifier(token: str, verifiers: List[JwtVerifier]) -> Optional[JwtVerifier]:
if not verifiers:
return None
try:
issuer = jwt.decode(token, options={"verify_signature": False}).get("iss")
except jwt.PyJWTError:
return None
for verifier in verifiers:
if verifier.provider.issuer == issuer:
return verifier
return None
class ApiKeyAuthenticator:
scheme = SecuritySchemeType.API_KEY
def __init__(self, config: ApiKeySchemeConfig) -> None:
self._header_name = config.header_name
async def authenticate(self, request: Request) -> Optional[Credential]:
raw = request.headers.get(self._header_name)
if not raw:
return None
return Credential(
scheme=self.scheme,
method=AuthMethod.API_KEY,
subject=raw,
credential_ref=CredentialRef(key_id=raw[:10]),
claims={"_raw_api_key": raw},
)
def challenge(self) -> str:
return ""
class HttpAuthenticator:
scheme = SecuritySchemeType.HTTP
def __init__(
self, basic: HttpBasicConfig, jwt_verifiers: List[JwtVerifier]
) -> None:
self._basic = basic
self._verifiers = jwt_verifiers
async def authenticate(self, request: Request) -> Optional[Credential]:
header = request.headers.get("authorization")
if not header:
return None
scheme, _, value = header.partition(" ")
scheme_lower = scheme.lower()
if scheme_lower == "bearer" and value:
return self._verify_bearer(value)
if scheme_lower == "basic" and self._basic.enabled and value:
return self._verify_basic(value)
return None
def _verify_bearer(self, token: str) -> Credential:
verifier = _select_verifier(token, self._verifiers)
if verifier is None:
raise errors.invalid_token("no issuer match")
claims = verifier.verify(token)
return _credential_from_claims(
self.scheme, AuthMethod.BEARER_JWT, token, claims
)
def _verify_basic(self, value: str) -> Credential:
try:
decoded = base64.b64decode(value).decode("utf-8")
except (binascii.Error, UnicodeDecodeError) as exc:
raise errors.unauthenticated(
errors.basic_challenge(self._basic.realm)
) from exc
username, _, password = decoded.partition(":")
if not username:
raise errors.unauthenticated(errors.basic_challenge(self._basic.realm))
return Credential(
scheme=self.scheme,
method=AuthMethod.HTTP_BASIC,
subject=username,
claims={"_basic_password": password},
)
def challenge(self) -> str:
bearer = errors.bearer_challenge()
if self._basic.enabled:
return f"{bearer}, {errors.basic_challenge(self._basic.realm)}"
return bearer
class OAuth2Authenticator:
scheme = SecuritySchemeType.OAUTH2
def __init__(
self,
jwt_verifiers: List[JwtVerifier],
introspection: Optional[OAuth2IntrospectionConfig],
) -> None:
self._verifiers = jwt_verifiers
self._introspection = introspection
async def authenticate(self, request: Request) -> Optional[Credential]:
token = _extract_bearer(request)
if token is None:
return None
if _looks_like_jwt(token):
return self._verify_at_jwt(token)
if self._introspection is not None:
return await self._introspect(token)
raise errors.invalid_token()
def _verify_at_jwt(self, token: str) -> Credential:
verifier = _select_verifier(token, self._verifiers)
if verifier is None:
raise errors.invalid_token("no issuer match")
claims = verifier.verify(token, require_at_jwt=True)
return _credential_from_claims(
self.scheme, AuthMethod.BEARER_JWT, token, claims
)
async def _introspect(self, token: str) -> Credential:
config = self._introspection
assert config is not None
async with httpx.AsyncClient(timeout=10.0) as client:
response = await client.post(
str(config.introspection_endpoint),
data={"token": token},
auth=(config.client_id, config.client_secret.get_secret_value()),
)
if response.status_code != 200:
raise errors.invalid_token("introspection failed")
body = response.json()
if not body.get("active"):
raise errors.invalid_token("token inactive")
return Credential(
scheme=self.scheme,
method=AuthMethod.OAUTH2_INTROSPECTION,
subject=str(body.get(config.subject_field, "")),
issuer=body.get("iss"),
audience=_normalize_audience(body.get("aud")),
scopes=_split_scope(body.get("scope")),
claims=body,
)
def challenge(self) -> str:
return errors.bearer_challenge()
class OidcAuthenticator:
scheme = SecuritySchemeType.OPENID_CONNECT
def __init__(self, jwt_verifiers: List[JwtVerifier]) -> None:
self._verifiers = jwt_verifiers
async def authenticate(self, request: Request) -> Optional[Credential]:
token = _extract_bearer(request)
if token is None:
return None
verifier = _select_verifier(token, self._verifiers)
if verifier is None:
raise errors.invalid_token("no issuer match")
claims = verifier.verify(token)
return _credential_from_claims(self.scheme, AuthMethod.OIDC, token, claims)
def challenge(self) -> str:
return errors.bearer_challenge()
class MutualTlsAuthenticator:
scheme = SecuritySchemeType.MUTUAL_TLS
def __init__(self, config: MutualTlsConfig) -> None:
self._config = config
async def authenticate(self, request: Request) -> Optional[Credential]:
cert = self._read_client_cert(request)
if cert is None:
return None
return Credential(
scheme=self.scheme,
method=AuthMethod.MUTUAL_TLS,
subject=cert.subject_dn,
client_certificate=cert,
)
def _read_client_cert(self, request: Request) -> Optional[ClientCertificate]:
if self._config.forwarded_subject_header:
dn = request.headers.get(self._config.forwarded_subject_header)
return ClientCertificate(subject_dn=dn) if dn else None
tls = request.scope.get("extensions", {}).get("tls", {})
dn = tls.get("client_cert_name")
return ClientCertificate(subject_dn=dn) if dn else None
def challenge(self) -> str:
return ""
def build_authenticators(config: AuthConfig) -> List[Authenticator]:
verifiers = [JwtVerifier(provider) for provider in config.oidc_providers]
by_scheme: Dict[SecuritySchemeType, Authenticator] = {}
if config.api_key is not None:
by_scheme[SecuritySchemeType.API_KEY] = ApiKeyAuthenticator(config.api_key)
by_scheme[SecuritySchemeType.HTTP] = HttpAuthenticator(config.http_basic, verifiers)
by_scheme[SecuritySchemeType.OPENID_CONNECT] = OidcAuthenticator(verifiers)
by_scheme[SecuritySchemeType.OAUTH2] = OAuth2Authenticator(
verifiers, config.oauth2_introspection
)
if config.mutual_tls.enabled:
by_scheme[SecuritySchemeType.MUTUAL_TLS] = MutualTlsAuthenticator(
config.mutual_tls
)
return [by_scheme[scheme] for scheme in config.scheme_order if scheme in by_scheme]

64
litellm/auth_v2/config.py Normal file
View file

@ -0,0 +1,64 @@
from __future__ import annotations
from typing import List, Optional
from pydantic import AnyHttpUrl, BaseModel, Field, SecretStr
from .models import SecuritySchemeType
class ApiKeySchemeConfig(BaseModel):
header_name: str = "x-litellm-api-key"
class HttpBasicConfig(BaseModel):
enabled: bool = False
realm: str = "litellm"
class OidcProviderConfig(BaseModel):
issuer: str
audience: List[str]
jwks_uri: Optional[AnyHttpUrl] = None
algorithms: List[str] = Field(default_factory=lambda: ["RS256"])
require_at_jwt: bool = False
client_id: Optional[str] = None
client_secret: Optional[SecretStr] = None
login_scopes: List[str] = Field(
default_factory=lambda: ["openid", "email", "profile"]
)
class OAuth2IntrospectionConfig(BaseModel):
introspection_endpoint: AnyHttpUrl
client_id: str
client_secret: SecretStr
subject_field: str = "sub"
class MutualTlsConfig(BaseModel):
enabled: bool = False
forwarded_subject_header: Optional[str] = None
class TrustedProxyConfig(BaseModel):
use_forwarded_for: bool = False
trusted_proxy_cidrs: List[str] = Field(default_factory=list)
class AuthConfig(BaseModel):
scheme_order: List[SecuritySchemeType] = Field(
default_factory=lambda: [
SecuritySchemeType.API_KEY,
SecuritySchemeType.HTTP,
SecuritySchemeType.OPENID_CONNECT,
SecuritySchemeType.OAUTH2,
SecuritySchemeType.MUTUAL_TLS,
]
)
api_key: Optional[ApiKeySchemeConfig] = Field(default_factory=ApiKeySchemeConfig)
http_basic: HttpBasicConfig = Field(default_factory=HttpBasicConfig)
oidc_providers: List[OidcProviderConfig] = Field(default_factory=list)
oauth2_introspection: Optional[OAuth2IntrospectionConfig] = None
mutual_tls: MutualTlsConfig = Field(default_factory=MutualTlsConfig)
network: TrustedProxyConfig = Field(default_factory=TrustedProxyConfig)

46
litellm/auth_v2/errors.py Normal file
View file

@ -0,0 +1,46 @@
from __future__ import annotations
from typing import Optional
from fastapi import HTTPException
class AuthError(HTTPException):
def __init__(
self, status_code: int, detail: str, challenge: Optional[str] = None
) -> None:
headers = {"WWW-Authenticate": challenge} if challenge else None
super().__init__(status_code=status_code, detail=detail, headers=headers)
def bearer_challenge(
error: Optional[str] = None, description: Optional[str] = None
) -> str:
parts = ['Bearer realm="litellm"']
if error:
parts.append(f'error="{error}"')
if description:
parts.append(f'error_description="{description}"')
return ", ".join(parts)
def basic_challenge(realm: str = "litellm") -> str:
return f'Basic realm="{realm}"'
def unauthenticated(challenge: str) -> AuthError:
return AuthError(401, "Not authenticated", challenge)
def invalid_token(description: Optional[str] = None) -> AuthError:
return AuthError(
401, "Invalid token", bearer_challenge("invalid_token", description)
)
def insufficient_scope() -> AuthError:
return AuthError(403, "Insufficient scope", bearer_challenge("insufficient_scope"))
def forbidden_role() -> AuthError:
return AuthError(403, "Insufficient role")

108
litellm/auth_v2/models.py Normal file
View file

@ -0,0 +1,108 @@
from __future__ import annotations
from enum import Enum
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, ConfigDict, Field
from .rbac import Role
class SecuritySchemeType(str, Enum):
API_KEY = "apiKey"
HTTP = "http"
OAUTH2 = "oauth2"
OPENID_CONNECT = "openIdConnect"
MUTUAL_TLS = "mutualTLS"
class AuthMethod(str, Enum):
API_KEY = "api_key"
HTTP_BASIC = "http_basic"
BEARER_JWT = "bearer_jwt"
OAUTH2_INTROSPECTION = "oauth2_introspection"
OIDC = "oidc"
MUTUAL_TLS = "mutual_tls"
class PrincipalType(str, Enum):
HUMAN = "human"
SERVICE_ACCOUNT = "service_account"
class TeamRole(str, Enum):
ADMIN = "admin"
MEMBER = "member"
class UserIdentity(BaseModel):
id: str
external_id: Optional[str] = None
user_name: Optional[str] = None
email: Optional[str] = None
display_name: Optional[str] = None
class OrganizationIdentity(BaseModel):
id: str
name: Optional[str] = None
class TeamIdentity(BaseModel):
id: str
name: Optional[str] = None
role: TeamRole = TeamRole.MEMBER
class CredentialRef(BaseModel):
key_id: Optional[str] = None
token_id: Optional[str] = None
class NetworkContext(BaseModel):
client_ip: Optional[str] = None
host: Optional[str] = None
via_trusted_proxy: bool = False
class ClientCertificate(BaseModel):
subject_dn: str
issuer_dn: Optional[str] = None
serial_number: Optional[str] = None
class Credential(BaseModel):
"""A verified credential, before identity resolution."""
model_config = ConfigDict(frozen=True)
scheme: SecuritySchemeType
method: AuthMethod
subject: str
issuer: Optional[str] = None
audience: List[str] = Field(default_factory=list)
scopes: List[str] = Field(default_factory=list)
claims: Dict[str, Any] = Field(default_factory=dict)
credential_ref: CredentialRef = Field(default_factory=CredentialRef)
client_certificate: Optional[ClientCertificate] = None
class Principal(BaseModel):
"""Normalized caller identity. Identity only, no policy/budget state."""
principal_type: PrincipalType
subject: str
issuer: Optional[str] = None
audience: List[str] = Field(default_factory=list)
user: Optional[UserIdentity] = None
organization: Optional[OrganizationIdentity] = None
teams: List[TeamIdentity] = Field(default_factory=list)
roles: List[Role] = Field(default_factory=list)
scopes: List[str] = Field(default_factory=list)
auth_method: AuthMethod
credential_ref: CredentialRef = Field(default_factory=CredentialRef)
network: NetworkContext = Field(default_factory=NetworkContext)
claims: Dict[str, Any] = Field(default_factory=dict)

View file

@ -0,0 +1,57 @@
from __future__ import annotations
import ipaddress
from typing import List, Optional, Tuple
from fastapi import Request
from .config import TrustedProxyConfig
from .models import NetworkContext
def _is_valid_ip(value: str) -> bool:
try:
ipaddress.ip_address(value)
return True
except ValueError:
return False
def _ip_in_cidrs(ip: Optional[str], cidrs: List[str]) -> bool:
if not ip or not _is_valid_ip(ip):
return False
address = ipaddress.ip_address(ip)
for cidr in cidrs:
try:
if address in ipaddress.ip_network(cidr, strict=False):
return True
except ValueError:
continue
return False
def resolve_client_ip(
request: Request, config: TrustedProxyConfig
) -> Tuple[Optional[str], bool]:
peer = request.client.host if request.client else None
if not config.use_forwarded_for or not _ip_in_cidrs(
peer, config.trusted_proxy_cidrs
):
return peer, False
forwarded = request.headers.get("x-forwarded-for", "")
hops = [h.strip() for h in forwarded.split(",") if h.strip()]
for hop in reversed(hops):
if not _ip_in_cidrs(hop, config.trusted_proxy_cidrs) and _is_valid_ip(hop):
return hop, True
return peer, True
def resolve_network_context(
request: Request, config: TrustedProxyConfig
) -> NetworkContext:
ip, via_proxy = resolve_client_ip(request, config)
return NetworkContext(
client_ip=ip,
host=request.headers.get("host"),
via_trusted_proxy=via_proxy,
)

65
litellm/auth_v2/oidc.py Normal file
View file

@ -0,0 +1,65 @@
from __future__ import annotations
import re
from typing import Any, Dict
from authlib.integrations.starlette_client import OAuth
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import JSONResponse
from scim2_models import User as ScimUser
from .config import AuthConfig, OidcProviderConfig
from .resolver import ProvisioningStore
def _provider_key(provider: OidcProviderConfig) -> str:
return re.sub(r"[^a-z0-9]+", "-", provider.issuer.lower()).strip("-")
def _user_from_userinfo(userinfo: Dict[str, Any]) -> ScimUser:
return ScimUser(
external_id=userinfo.get("sub"),
user_name=userinfo.get("preferred_username") or userinfo.get("email"),
display_name=userinfo.get("name"),
)
def build_oidc_router(config: AuthConfig) -> APIRouter:
oauth = OAuth()
for provider in config.oidc_providers:
oauth.register(
name=_provider_key(provider),
server_metadata_url=f"{provider.issuer.rstrip('/')}/.well-known/openid-configuration",
client_id=provider.client_id,
client_secret=(
provider.client_secret.get_secret_value()
if provider.client_secret
else None
),
client_kwargs={"scope": " ".join(provider.login_scopes)},
)
router = APIRouter(prefix="/auth/oidc", tags=["oidc"])
@router.get("/{provider}/login")
async def login(provider: str, request: Request) -> Any:
client = oauth.create_client(provider)
if client is None:
raise HTTPException(status_code=404, detail="unknown provider")
redirect_uri = request.url_for("oidc_callback", provider=provider)
return await client.authorize_redirect(request, str(redirect_uri))
@router.get("/{provider}/callback", name="oidc_callback")
async def callback(provider: str, request: Request) -> JSONResponse:
client = oauth.create_client(provider)
if client is None:
raise HTTPException(status_code=404, detail="unknown provider")
token = await client.authorize_access_token(request)
userinfo = token.get("userinfo")
if userinfo is None:
userinfo = await client.userinfo(token=token)
store: ProvisioningStore = request.app.state.auth_v2.resolver
stored = await store.upsert_user(_user_from_userinfo(dict(userinfo)))
return JSONResponse(content=stored.model_dump())
return router

28
litellm/auth_v2/rbac.py Normal file
View file

@ -0,0 +1,28 @@
from __future__ import annotations
from enum import Enum
from typing import TYPE_CHECKING, Tuple
from fastapi.security import SecurityScopes
if TYPE_CHECKING:
from .models import Principal
class Role(str, Enum):
PLATFORM_ADMIN = "platform_admin"
PLATFORM_VIEWER = "platform_viewer"
ORG_ADMIN = "org_admin"
ORG_VIEWER = "org_viewer"
TEAM_ADMIN = "team_admin"
TEAM_MEMBER = "team_member"
def has_required_scopes(
security_scopes: SecurityScopes, principal: "Principal"
) -> bool:
return set(security_scopes.scopes).issubset(set(principal.scopes))
def has_any_role(principal: "Principal", allowed: Tuple[Role, ...]) -> bool:
return any(role in allowed for role in principal.roles)

154
litellm/auth_v2/resolver.py Normal file
View file

@ -0,0 +1,154 @@
from __future__ import annotations
import hashlib
import uuid
from typing import Any, Dict, List, Optional, Protocol, runtime_checkable
from scim2_models import Group as ScimGroup
from scim2_models import User as ScimUser
from . import errors
from .models import (
AuthMethod,
Credential,
Principal,
PrincipalType,
TeamIdentity,
UserIdentity,
)
from .rbac import Role
@runtime_checkable
class IdentityResolver(Protocol):
async def resolve(self, credential: Credential) -> Principal: ...
@runtime_checkable
class ProvisioningStore(Protocol):
async def upsert_user(self, user: ScimUser) -> ScimUser: ...
async def get_user(self, resource_id: str) -> Optional[ScimUser]: ...
async def deactivate_user(self, resource_id: str) -> None: ...
async def list_users(self, filter_expr: Optional[str]) -> List[ScimUser]: ...
async def upsert_group(self, group: ScimGroup) -> ScimGroup: ...
async def get_group(self, resource_id: str) -> Optional[ScimGroup]: ...
async def delete_group(self, resource_id: str) -> None: ...
async def list_groups(self, filter_expr: Optional[str]) -> List[ScimGroup]: ...
def _hash_api_key(raw: str) -> str:
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
def _roles_from_claims(claims: Dict[str, Any]) -> List[Role]:
raw = claims.get("roles", [])
if not isinstance(raw, list):
return []
valid = {role.value for role in Role}
return [Role(value) for value in raw if value in valid]
def _teams_from_claims(claims: Dict[str, Any]) -> List[TeamIdentity]:
groups = claims.get("groups", [])
if not isinstance(groups, list):
return []
return [TeamIdentity(id=str(group), name=str(group)) for group in groups]
class InMemoryIdentityStore(IdentityResolver, ProvisioningStore):
def __init__(
self,
api_keys: Optional[Dict[str, Principal]] = None,
subjects: Optional[Dict[str, Principal]] = None,
users: Optional[Dict[str, ScimUser]] = None,
groups: Optional[Dict[str, ScimGroup]] = None,
) -> None:
self._api_keys = api_keys or {}
self._subjects = subjects or {}
self._users = users or {}
self._groups = groups or {}
async def resolve(self, credential: Credential) -> Principal:
if credential.method == AuthMethod.API_KEY:
return self._resolve_api_key(credential)
return self._resolve_subject(credential)
def _resolve_api_key(self, credential: Credential) -> Principal:
raw = credential.claims.get("_raw_api_key")
if not isinstance(raw, str):
raise errors.invalid_token()
principal = self._api_keys.get(_hash_api_key(raw))
if principal is None:
raise errors.invalid_token()
return principal
def _resolve_subject(self, credential: Credential) -> Principal:
stored = self._subjects.get(f"{credential.issuer}|{credential.subject}")
if stored is not None:
return stored
return self._principal_from_claims(credential)
def _principal_from_claims(self, credential: Credential) -> Principal:
claims = credential.claims
if credential.method == AuthMethod.MUTUAL_TLS:
return Principal(
principal_type=PrincipalType.SERVICE_ACCOUNT,
subject=credential.subject,
issuer=credential.issuer,
audience=list(credential.audience),
scopes=list(credential.scopes),
auth_method=credential.method,
credential_ref=credential.credential_ref,
claims=dict(claims),
)
return Principal(
principal_type=PrincipalType.HUMAN,
subject=credential.subject,
issuer=credential.issuer,
audience=list(credential.audience),
user=UserIdentity(
id=credential.subject,
external_id=credential.subject,
email=claims.get("email"),
user_name=claims.get("preferred_username"),
display_name=claims.get("name"),
),
teams=_teams_from_claims(claims),
roles=_roles_from_claims(claims),
scopes=list(credential.scopes),
auth_method=credential.method,
credential_ref=credential.credential_ref,
claims=dict(claims),
)
async def upsert_user(self, user: ScimUser) -> ScimUser:
if not user.id:
user.id = str(uuid.uuid4())
self._users[user.id] = user
return user
async def get_user(self, resource_id: str) -> Optional[ScimUser]:
return self._users.get(resource_id)
async def deactivate_user(self, resource_id: str) -> None:
user = self._users.get(resource_id)
if user is not None:
user.active = False
async def list_users(self, filter_expr: Optional[str]) -> List[ScimUser]:
return list(self._users.values())
async def upsert_group(self, group: ScimGroup) -> ScimGroup:
if not group.id:
group.id = str(uuid.uuid4())
self._groups[group.id] = group
return group
async def get_group(self, resource_id: str) -> Optional[ScimGroup]:
return self._groups.get(resource_id)
async def delete_group(self, resource_id: str) -> None:
self._groups.pop(resource_id, None)
async def list_groups(self, filter_expr: Optional[str]) -> List[ScimGroup]:
return list(self._groups.values())

35
litellm/auth_v2/saml.py Normal file
View file

@ -0,0 +1,35 @@
from __future__ import annotations
from typing import Optional
from fastapi import APIRouter, Request
from .models import Credential, SecuritySchemeType
class SamlAuthenticator:
"""Thin SAML SP seam. Full pysaml2 wiring (system libxmlsec1, pinned
xmlsec/lxml, multi-IdP metadata) is deferred; the ACS maps a SAML assertion's
NameID and attribute statements into the same scim2_models.User upsert as
OIDC and SCIM."""
scheme = SecuritySchemeType.HTTP
async def authenticate(self, request: Request) -> Optional[Credential]:
raise NotImplementedError("SAML SP deferred; see 03-design.md cut list")
def build_saml_router() -> APIRouter:
router = APIRouter(prefix="/auth/saml", tags=["saml"])
@router.get("/metadata")
async def metadata() -> None:
raise NotImplementedError("requires pysaml2 + system libxmlsec1")
@router.post("/acs")
async def assertion_consumer_service(request: Request) -> None:
raise NotImplementedError(
"parse assertion -> scim2_models.User -> store.upsert_user"
)
return router

216
litellm/auth_v2/scim.py Normal file
View file

@ -0,0 +1,216 @@
from __future__ import annotations
from typing import Any, Dict, Optional, Type, TypeVar
from fastapi import APIRouter, Request, Response, status
from fastapi.responses import JSONResponse
from pydantic import ValidationError
from scim2_models import (
Bulk,
ChangePassword,
Context,
Error,
Filter,
Group,
ListResponse,
Patch,
PatchOp,
Resource,
ResourceType,
ServiceProviderConfig,
Sort,
User,
)
from .resolver import ProvisioningStore
R = TypeVar("R", bound=Resource)
def _store(request: Request) -> ProvisioningStore:
return request.app.state.auth_v2.resolver
def _error(status_code: int, detail: str) -> JSONResponse:
return JSONResponse(
status_code=status_code,
content=Error(status=str(status_code), detail=detail).model_dump(),
)
async def _parse(request: Request, model: Type[R]) -> R:
body = await request.json()
return model.model_validate(body, scim_ctx=Context.RESOURCE_CREATION_REQUEST)
def _apply_patch(resource: R, patch: PatchOp) -> R:
data: Dict[str, Any] = resource.model_dump()
for op in patch.operations:
action = op.op.value if hasattr(op.op, "value") else str(op.op)
if action == "remove":
if op.path:
data.pop(op.path, None)
continue
if op.path is None and isinstance(op.value, dict):
data.update(op.value)
elif op.path is not None:
data[op.path] = op.value
return type(resource).model_validate(data)
def _dump(resource: Resource, ctx: Context) -> Dict[str, Any]:
return resource.model_dump(scim_ctx=ctx)
def build_scim_router() -> APIRouter:
router = APIRouter(prefix="/scim/v2", tags=["scim"])
@router.post("/Users", status_code=status.HTTP_201_CREATED)
async def create_user(request: Request) -> Response:
try:
user = await _parse(request, User)
except ValidationError as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
stored = await _store(request).upsert_user(user)
return JSONResponse(
status_code=status.HTTP_201_CREATED,
content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE),
)
@router.get("/Users/{resource_id}")
async def get_user(resource_id: str, request: Request) -> Response:
user = await _store(request).get_user(resource_id)
if user is None:
return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found")
return JSONResponse(content=_dump(user, Context.RESOURCE_QUERY_RESPONSE))
@router.patch("/Users/{resource_id}")
async def patch_user(resource_id: str, request: Request) -> Response:
store = _store(request)
user = await store.get_user(resource_id)
if user is None:
return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found")
try:
patch = PatchOp[User].model_validate(await request.json())
except ValidationError as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
updated = await store.upsert_user(_apply_patch(user, patch))
return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE))
@router.delete("/Users/{resource_id}", status_code=status.HTTP_204_NO_CONTENT)
async def deactivate_user(resource_id: str, request: Request) -> Response:
await _store(request).deactivate_user(resource_id)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@router.get("/Users")
async def list_users(request: Request, filter: Optional[str] = None) -> Response:
users = await _store(request).list_users(filter)
listing: ListResponse[User] = ListResponse[User](
total_results=len(users),
start_index=1,
items_per_page=len(users),
resources=users or None,
)
return JSONResponse(content=_dump(listing, Context.RESOURCE_QUERY_RESPONSE))
@router.post("/Groups", status_code=status.HTTP_201_CREATED)
async def create_group(request: Request) -> Response:
try:
group = await _parse(request, Group)
except ValidationError as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
stored = await _store(request).upsert_group(group)
return JSONResponse(
status_code=status.HTTP_201_CREATED,
content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE),
)
@router.get("/Groups/{resource_id}")
async def get_group(resource_id: str, request: Request) -> Response:
group = await _store(request).get_group(resource_id)
if group is None:
return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found")
return JSONResponse(content=_dump(group, Context.RESOURCE_QUERY_RESPONSE))
@router.patch("/Groups/{resource_id}")
async def patch_group(resource_id: str, request: Request) -> Response:
store = _store(request)
group = await store.get_group(resource_id)
if group is None:
return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found")
try:
patch = PatchOp[Group].model_validate(await request.json())
except ValidationError as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
updated = await store.upsert_group(_apply_patch(group, patch))
return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE))
@router.delete("/Groups/{resource_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_group(resource_id: str, request: Request) -> Response:
await _store(request).delete_group(resource_id)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@router.get("/Groups")
async def list_groups(request: Request, filter: Optional[str] = None) -> Response:
groups = await _store(request).list_groups(filter)
listing: ListResponse[Group] = ListResponse[Group](
total_results=len(groups),
start_index=1,
items_per_page=len(groups),
resources=groups or None,
)
return JSONResponse(content=_dump(listing, Context.RESOURCE_QUERY_RESPONSE))
@router.get("/ServiceProviderConfig")
async def service_provider_config() -> Response:
config = ServiceProviderConfig(
patch=Patch(supported=True),
bulk=Bulk(supported=False, max_operations=0, max_payload_size=0),
filter=Filter(supported=False, max_results=0),
change_password=ChangePassword(supported=False),
sort=Sort(supported=False),
etag=None,
authentication_schemes=[],
)
return JSONResponse(content=config.model_dump())
@router.get("/ResourceTypes")
async def resource_types() -> Response:
types = [
ResourceType(
id="User",
name="User",
endpoint="/Users",
schema="urn:ietf:params:scim:schemas:core:2.0:User",
),
ResourceType(
id="Group",
name="Group",
endpoint="/Groups",
schema="urn:ietf:params:scim:schemas:core:2.0:Group",
),
]
listing: ListResponse[ResourceType] = ListResponse[ResourceType](
total_results=len(types),
start_index=1,
items_per_page=len(types),
resources=types,
)
return JSONResponse(content=listing.model_dump())
@router.get("/Schemas")
async def schemas() -> Response:
return JSONResponse(
content={
"schemas": ["urn:ietf:params:scim:api:messages:2.0:ListResponse"],
"totalResults": 2,
"startIndex": 1,
"itemsPerPage": 2,
"Resources": [
User.to_schema().model_dump(),
Group.to_schema().model_dump(),
],
}
)
return router

View file

@ -0,0 +1,89 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Annotated, Callable, List
from fastapi import FastAPI, Request, Security
from fastapi.security import SecurityScopes
from . import errors
from .authenticators import Authenticator, build_authenticators
from .config import AuthConfig
from .models import Principal
from .network import resolve_network_context
from .rbac import Role, has_any_role, has_required_scopes
from .resolver import IdentityResolver
@dataclass
class AuthContext:
config: AuthConfig
authenticators: List[Authenticator]
resolver: IdentityResolver
def install_auth(
app: FastAPI,
config: AuthConfig,
resolver: IdentityResolver,
*,
mount_scim: bool = True,
mount_oidc: bool = True,
) -> AuthContext:
ctx = AuthContext(config, build_authenticators(config), resolver)
app.state.auth_v2 = ctx
if mount_scim:
from .scim import build_scim_router
app.include_router(build_scim_router())
if mount_oidc and config.oidc_providers:
from .oidc import build_oidc_router
app.include_router(build_oidc_router(config))
return ctx
def _ctx(request: Request) -> AuthContext:
return request.app.state.auth_v2
def _combined_challenge(authenticators: List[Authenticator]) -> str:
seen: List[str] = []
for authenticator in authenticators:
challenge = authenticator.challenge()
if challenge and challenge not in seen:
seen.append(challenge)
return ", ".join(seen)
async def get_current_principal(
security_scopes: SecurityScopes, request: Request
) -> Principal:
ctx = _ctx(request)
credential = None
for authenticator in ctx.authenticators:
credential = await authenticator.authenticate(request)
if credential is not None:
break
if credential is None:
raise errors.unauthenticated(_combined_challenge(ctx.authenticators))
resolved = await ctx.resolver.resolve(credential)
principal = resolved.model_copy(
update={"network": resolve_network_context(request, ctx.config.network)}
)
if not has_required_scopes(security_scopes, principal):
raise errors.insufficient_scope()
return principal
def require_roles(*allowed: Role) -> Callable[..., object]:
async def dependency(
principal: Annotated[Principal, Security(get_current_principal)],
) -> Principal:
if not has_any_role(principal, allowed):
raise errors.forbidden_role()
return principal
return dependency