From a0a59a2197dd7e2efd194eaa37141c75b5169e5a Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 10 Jun 2026 17:22:20 -0700 Subject: [PATCH] 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. --- litellm/auth_v2/__init__.py | 11 + litellm/auth_v2/authenticators.py | 347 ++++++++++++++++++++++++++++++ litellm/auth_v2/config.py | 64 ++++++ litellm/auth_v2/errors.py | 46 ++++ litellm/auth_v2/models.py | 108 ++++++++++ litellm/auth_v2/network.py | 57 +++++ litellm/auth_v2/oidc.py | 65 ++++++ litellm/auth_v2/rbac.py | 28 +++ litellm/auth_v2/resolver.py | 154 +++++++++++++ litellm/auth_v2/saml.py | 35 +++ litellm/auth_v2/scim.py | 216 +++++++++++++++++++ litellm/auth_v2/security.py | 89 ++++++++ 12 files changed, 1220 insertions(+) create mode 100644 litellm/auth_v2/__init__.py create mode 100644 litellm/auth_v2/authenticators.py create mode 100644 litellm/auth_v2/config.py create mode 100644 litellm/auth_v2/errors.py create mode 100644 litellm/auth_v2/models.py create mode 100644 litellm/auth_v2/network.py create mode 100644 litellm/auth_v2/oidc.py create mode 100644 litellm/auth_v2/rbac.py create mode 100644 litellm/auth_v2/resolver.py create mode 100644 litellm/auth_v2/saml.py create mode 100644 litellm/auth_v2/scim.py create mode 100644 litellm/auth_v2/security.py diff --git a/litellm/auth_v2/__init__.py b/litellm/auth_v2/__init__.py new file mode 100644 index 00000000000..7373bc860c4 --- /dev/null +++ b/litellm/auth_v2/__init__.py @@ -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", +] diff --git a/litellm/auth_v2/authenticators.py b/litellm/auth_v2/authenticators.py new file mode 100644 index 00000000000..54bd9c15a6f --- /dev/null +++ b/litellm/auth_v2/authenticators.py @@ -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] diff --git a/litellm/auth_v2/config.py b/litellm/auth_v2/config.py new file mode 100644 index 00000000000..c4a51afbabb --- /dev/null +++ b/litellm/auth_v2/config.py @@ -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) diff --git a/litellm/auth_v2/errors.py b/litellm/auth_v2/errors.py new file mode 100644 index 00000000000..8c0dbbdd1bf --- /dev/null +++ b/litellm/auth_v2/errors.py @@ -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") diff --git a/litellm/auth_v2/models.py b/litellm/auth_v2/models.py new file mode 100644 index 00000000000..d23306abf83 --- /dev/null +++ b/litellm/auth_v2/models.py @@ -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) diff --git a/litellm/auth_v2/network.py b/litellm/auth_v2/network.py new file mode 100644 index 00000000000..c8aee43fffb --- /dev/null +++ b/litellm/auth_v2/network.py @@ -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, + ) diff --git a/litellm/auth_v2/oidc.py b/litellm/auth_v2/oidc.py new file mode 100644 index 00000000000..e463981f6e6 --- /dev/null +++ b/litellm/auth_v2/oidc.py @@ -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 diff --git a/litellm/auth_v2/rbac.py b/litellm/auth_v2/rbac.py new file mode 100644 index 00000000000..ee094cac500 --- /dev/null +++ b/litellm/auth_v2/rbac.py @@ -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) diff --git a/litellm/auth_v2/resolver.py b/litellm/auth_v2/resolver.py new file mode 100644 index 00000000000..55a5883689f --- /dev/null +++ b/litellm/auth_v2/resolver.py @@ -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()) diff --git a/litellm/auth_v2/saml.py b/litellm/auth_v2/saml.py new file mode 100644 index 00000000000..706880b7d2b --- /dev/null +++ b/litellm/auth_v2/saml.py @@ -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 diff --git a/litellm/auth_v2/scim.py b/litellm/auth_v2/scim.py new file mode 100644 index 00000000000..d9a2739bbfa --- /dev/null +++ b/litellm/auth_v2/scim.py @@ -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 diff --git a/litellm/auth_v2/security.py b/litellm/auth_v2/security.py new file mode 100644 index 00000000000..0e016f3a5c2 --- /dev/null +++ b/litellm/auth_v2/security.py @@ -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