mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
2bbf688613
commit
a0a59a2197
12 changed files with 1220 additions and 0 deletions
11
litellm/auth_v2/__init__.py
Normal file
11
litellm/auth_v2/__init__.py
Normal 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",
|
||||
]
|
||||
347
litellm/auth_v2/authenticators.py
Normal file
347
litellm/auth_v2/authenticators.py
Normal 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
64
litellm/auth_v2/config.py
Normal 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
46
litellm/auth_v2/errors.py
Normal 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
108
litellm/auth_v2/models.py
Normal 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)
|
||||
57
litellm/auth_v2/network.py
Normal file
57
litellm/auth_v2/network.py
Normal 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
65
litellm/auth_v2/oidc.py
Normal 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
28
litellm/auth_v2/rbac.py
Normal 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
154
litellm/auth_v2/resolver.py
Normal 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
35
litellm/auth_v2/saml.py
Normal 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
216
litellm/auth_v2/scim.py
Normal 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
|
||||
89
litellm/auth_v2/security.py
Normal file
89
litellm/auth_v2/security.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue