refactor(auth_v2): replace install_auth with AuthSecurity DI

Per user direction, drop install_auth/AuthContext/app.state entirely; the
enforcement layer is now an AuthSecurity object whose bound methods are the
FastAPI Security() dependencies. The app constructs AuthSecurity(config, resolver)
once and passes auth.principal / auth.require_roles / auth.require_permission to
Security(); routers take the instance explicitly via build_*_router(auth) and read
auth.resolver/auth.config rather than request.app.state.

Browser sessions are unified behind a shared SessionStore + SessionAuthenticator
(session.py): one cookie, keyed on identity["method"], so SAML and the upcoming
OIDC login flow share one store. AuthSecurity owns the post-login session_store
and a short-TTL oauth_txn_store for OIDC state/nonce/PKCE; SessionConfig moves the
cookie/TTL/redirect settings off SAMLConfig. SAML keeps its protocol-specific
outstanding/replay state local.

Fold in the standing judge findings: collapse the triplicated http/oauth2/oidc
bearer paths into one _authenticate_bearer_jwt helper, inject the introspection
async-client via a factory instead of importing litellm inline, drop the dead
scheme attribute from the authenticators, and rename to PEP 8 acronym casing
(JWTVerifier, OIDCAuthenticator, OIDCProviderConfig, SAMLConfig, RBACEngine,
APIKeyAuthenticator, MutualTLSAuthenticator). __all__ now exports AuthSecurity,
Role and the resolver protocols.

scim.py and oidc.py move to build_*_router(auth) separately.
This commit is contained in:
Yassin Kortam 2026-06-10 19:08:13 -07:00
parent c512a49fa5
commit 8302f55995
7 changed files with 337 additions and 316 deletions

View file

@ -1,17 +1,35 @@
from .config import AuthConfig
from .models import Principal
from .security import (
get_current_principal,
install_auth,
require_permission,
require_roles,
from .config import (
ApiKeySchemeConfig,
AuthConfig,
HttpBasicConfig,
MutualTLSConfig,
OAuth2IntrospectionConfig,
OIDCProviderConfig,
SAMLConfig,
SessionConfig,
TrustedProxyConfig,
)
from .models import Principal
from .rbac import Role
from .resolver import IdentityResolver, InMemoryIdentityStore, ProvisioningStore
from .saml import build_saml_router
from .security import AuthSecurity
__all__ = [
"Principal",
"AuthSecurity",
"AuthConfig",
"get_current_principal",
"require_roles",
"require_permission",
"install_auth",
"Principal",
"Role",
"IdentityResolver",
"ProvisioningStore",
"InMemoryIdentityStore",
"ApiKeySchemeConfig",
"HttpBasicConfig",
"OIDCProviderConfig",
"OAuth2IntrospectionConfig",
"MutualTLSConfig",
"TrustedProxyConfig",
"SessionConfig",
"SAMLConfig",
"build_saml_router",
]

View file

@ -6,9 +6,8 @@ import functools
import hashlib
import hmac
import secrets
from typing import Any, Dict, List, Optional, Protocol, runtime_checkable
from typing import Any, Callable, Dict, List, Optional, Protocol, runtime_checkable
import httpx
import jwt
from fastapi import Request
from jwt import PyJWKClient
@ -20,9 +19,9 @@ from .config import (
ApiKeySchemeConfig,
AuthConfig,
HttpBasicConfig,
MutualTlsConfig,
MutualTLSConfig,
OAuth2IntrospectionConfig,
OidcProviderConfig,
OIDCProviderConfig,
TrustedProxyConfig,
)
from .models import (
@ -39,13 +38,40 @@ 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: ...
@runtime_checkable
class BasicAuthVerifier(Protocol):
def verify(self, username: str, password: str) -> bool: ...
def hash_basic_password(password: str, salt: Optional[str] = None) -> str:
salt = salt or secrets.token_hex(16)
digest = hashlib.sha256(bytes.fromhex(salt) + password.encode()).hexdigest()
return f"{salt}${digest}"
class InMemoryBasicAuthStore:
def __init__(self, credentials: Dict[str, str]) -> None:
self._credentials = credentials
def verify(self, username: str, password: str) -> bool:
stored = self._credentials.get(username)
if stored is None:
return False
salt, _, expected = stored.partition("$")
try:
candidate = hashlib.sha256(
bytes.fromhex(salt) + password.encode()
).hexdigest()
except ValueError:
return False
return hmac.compare_digest(candidate, expected)
def _extract_bearer(request: Request) -> Optional[str]:
header = request.headers.get("authorization")
if not header:
@ -93,39 +119,10 @@ def _credential_from_claims(
)
@runtime_checkable
class BasicAuthVerifier(Protocol):
def verify(self, username: str, password: str) -> bool: ...
def hash_basic_password(password: str, salt: Optional[str] = None) -> str:
salt = salt or secrets.token_hex(16)
digest = hashlib.sha256(bytes.fromhex(salt) + password.encode()).hexdigest()
return f"{salt}${digest}"
class InMemoryBasicAuthStore:
def __init__(self, credentials: Dict[str, str]) -> None:
self._credentials = credentials
def verify(self, username: str, password: str) -> bool:
stored = self._credentials.get(username)
if stored is None:
return False
salt, _, expected = stored.partition("$")
try:
candidate = hashlib.sha256(
bytes.fromhex(salt) + password.encode()
).hexdigest()
except ValueError:
return False
return hmac.compare_digest(candidate, expected)
class JwtVerifier:
class JWTVerifier:
def __init__(
self,
provider: OidcProviderConfig,
provider: OIDCProviderConfig,
jwks_client: Optional[PyJWKClient] = None,
) -> None:
self.provider = provider
@ -138,6 +135,8 @@ class JwtVerifier:
self._jwks_client = PyJWKClient(jwks_uri, cache_keys=True)
def _discover_jwks(self) -> str:
import httpx
url = f"{self.provider.issuer.rstrip('/')}/.well-known/openid-configuration"
response = httpx.get(url, timeout=10.0)
response.raise_for_status()
@ -171,14 +170,14 @@ class JwtVerifier:
async def _verify_jwt_off_loop(
verifier: JwtVerifier, token: str, *, require_at_jwt: Optional[bool] = None
verifier: JWTVerifier, token: str, *, require_at_jwt: Optional[bool] = None
) -> Dict[str, Any]:
return await run_in_threadpool(
functools.partial(verifier.verify, token, require_at_jwt=require_at_jwt)
)
def _select_verifier(token: str, verifiers: List[JwtVerifier]) -> Optional[JwtVerifier]:
def _select_verifier(token: str, verifiers: List[JWTVerifier]) -> Optional[JWTVerifier]:
if not verifiers:
return None
try:
@ -191,9 +190,22 @@ def _select_verifier(token: str, verifiers: List[JwtVerifier]) -> Optional[JwtVe
return None
class ApiKeyAuthenticator:
scheme = SecuritySchemeType.API_KEY
async def _authenticate_bearer_jwt(
token: str,
verifiers: List[JWTVerifier],
scheme: SecuritySchemeType,
method: AuthMethod,
*,
require_at_jwt: bool = False,
) -> Credential:
verifier = _select_verifier(token, verifiers)
if verifier is None:
raise errors.invalid_token("no issuer match")
claims = await _verify_jwt_off_loop(verifier, token, require_at_jwt=require_at_jwt)
return _credential_from_claims(scheme, method, token, claims)
class APIKeyAuthenticator:
def __init__(self, config: ApiKeySchemeConfig) -> None:
self._header_name = config.header_name
@ -202,7 +214,7 @@ class ApiKeyAuthenticator:
if not raw:
return None
return Credential(
scheme=self.scheme,
scheme=SecuritySchemeType.API_KEY,
method=AuthMethod.API_KEY,
subject=raw,
credential_ref=CredentialRef(key_id=raw[:10]),
@ -214,12 +226,10 @@ class ApiKeyAuthenticator:
class HttpAuthenticator:
scheme = SecuritySchemeType.HTTP
def __init__(
self,
basic: HttpBasicConfig,
jwt_verifiers: List[JwtVerifier],
jwt_verifiers: List[JWTVerifier],
basic_verifier: Optional[BasicAuthVerifier] = None,
) -> None:
self._basic = basic
@ -233,20 +243,13 @@ class HttpAuthenticator:
scheme, _, value = header.partition(" ")
scheme_lower = scheme.lower()
if scheme_lower == "bearer" and value:
return await self._verify_bearer(value)
return await _authenticate_bearer_jwt(
value, self._verifiers, SecuritySchemeType.HTTP, AuthMethod.BEARER_JWT
)
if scheme_lower == "basic" and self._basic.enabled and value:
return self._verify_basic(value)
return None
async 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 = await _verify_jwt_off_loop(verifier, token)
return _credential_from_claims(
self.scheme, AuthMethod.BEARER_JWT, token, claims
)
def _verify_basic(self, value: str) -> Credential:
challenge = errors.basic_challenge(self._basic.realm)
try:
@ -262,7 +265,7 @@ class HttpAuthenticator:
):
raise errors.unauthenticated(challenge)
return Credential(
scheme=self.scheme,
scheme=SecuritySchemeType.HTTP,
method=AuthMethod.HTTP_BASIC,
subject=username,
)
@ -274,46 +277,47 @@ class HttpAuthenticator:
return bearer
class OAuth2Authenticator:
scheme = SecuritySchemeType.OAUTH2
def _default_introspection_client() -> Any:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
return get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
class OAuth2Authenticator:
def __init__(
self,
jwt_verifiers: List[JwtVerifier],
jwt_verifiers: List[JWTVerifier],
introspection: Optional[OAuth2IntrospectionConfig],
client_factory: Optional[Callable[[], Any]] = None,
) -> None:
self._verifiers = jwt_verifiers
self._introspection = introspection
self._client_factory = client_factory or _default_introspection_client
async def authenticate(self, request: Request) -> Optional[Credential]:
token = _extract_bearer(request)
if token is None:
return None
if _looks_like_jwt(token):
return await self._verify_at_jwt(token)
return await _authenticate_bearer_jwt(
token,
self._verifiers,
SecuritySchemeType.OAUTH2,
AuthMethod.BEARER_JWT,
require_at_jwt=True,
)
if self._introspection is not None:
return await self._introspect(token)
raise errors.invalid_token()
async 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 = await _verify_jwt_off_loop(verifier, 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
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
basic = base64.b64encode(
f"{config.client_id}:{config.client_secret.get_secret_value()}".encode()
).decode()
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
client = self._client_factory()
response = await client.post(
str(config.introspection_endpoint),
data={"token": token},
@ -329,7 +333,7 @@ class OAuth2Authenticator:
if config.audience and not set(token_audience) & set(config.audience):
raise errors.invalid_token("audience mismatch")
return Credential(
scheme=self.scheme,
scheme=SecuritySchemeType.OAUTH2,
method=AuthMethod.OAUTH2_INTROSPECTION,
subject=str(body.get(config.subject_field, "")),
issuer=body.get("iss"),
@ -342,30 +346,24 @@ class OAuth2Authenticator:
return errors.bearer_challenge()
class OidcAuthenticator:
scheme = SecuritySchemeType.OPENID_CONNECT
def __init__(self, jwt_verifiers: List[JwtVerifier]) -> None:
class OIDCAuthenticator:
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 = await _verify_jwt_off_loop(verifier, token)
return _credential_from_claims(self.scheme, AuthMethod.OIDC, token, claims)
return await _authenticate_bearer_jwt(
token, self._verifiers, SecuritySchemeType.OPENID_CONNECT, AuthMethod.OIDC
)
def challenge(self) -> str:
return errors.bearer_challenge()
class MutualTlsAuthenticator:
scheme = SecuritySchemeType.MUTUAL_TLS
def __init__(self, config: MutualTlsConfig, network: TrustedProxyConfig) -> None:
class MutualTLSAuthenticator:
def __init__(self, config: MutualTLSConfig, network: TrustedProxyConfig) -> None:
self._config = config
self._network = network
@ -374,7 +372,7 @@ class MutualTlsAuthenticator:
if cert is None:
return None
return Credential(
scheme=self.scheme,
scheme=SecuritySchemeType.MUTUAL_TLS,
method=AuthMethod.MUTUAL_TLS,
subject=cert.subject_dn,
client_certificate=cert,
@ -398,19 +396,19 @@ class MutualTlsAuthenticator:
def build_authenticators(
config: AuthConfig, *, basic_verifier: Optional[BasicAuthVerifier] = None
) -> List[Authenticator]:
verifiers = [JwtVerifier(provider) for provider in config.oidc_providers]
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.API_KEY] = APIKeyAuthenticator(config.api_key)
by_scheme[SecuritySchemeType.HTTP] = HttpAuthenticator(
config.http_basic, verifiers, basic_verifier
)
by_scheme[SecuritySchemeType.OPENID_CONNECT] = OidcAuthenticator(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(
by_scheme[SecuritySchemeType.MUTUAL_TLS] = MutualTLSAuthenticator(
config.mutual_tls, config.network
)
return [by_scheme[scheme] for scheme in config.scheme_order if scheme in by_scheme]

View file

@ -29,7 +29,7 @@ class HttpBasicConfig(BaseModel):
realm: str = "litellm"
class OidcProviderConfig(BaseModel):
class OIDCProviderConfig(BaseModel):
issuer: str
audience: List[str]
jwks_uri: Optional[AnyHttpUrl] = None
@ -50,7 +50,7 @@ class OAuth2IntrospectionConfig(BaseModel):
audience: List[str] = Field(default_factory=list)
class MutualTlsConfig(BaseModel):
class MutualTLSConfig(BaseModel):
enabled: bool = False
forwarded_subject_header: Optional[str] = None
@ -60,7 +60,17 @@ class TrustedProxyConfig(BaseModel):
trusted_proxy_cidrs: List[str] = Field(default_factory=list)
class SamlConfig(BaseModel):
class SessionConfig(BaseModel):
cookie: str = "litellm_session"
secure: bool = True
ttl_seconds: int = 3600
max_size: int = 10000
default_redirect_path: str = "/"
login_cookie: str = "litellm_oidc_txn"
login_state_ttl: int = 300
class SAMLConfig(BaseModel):
enabled: bool = False
entity_id: str
acs_url: str
@ -68,18 +78,13 @@ class SamlConfig(BaseModel):
sp_key_file: Optional[str] = None
sp_cert_file: Optional[str] = None
allow_unsolicited: bool = False
session_cookie: str = "saml_session"
cookie_secure: bool = True
session_ttl_seconds: int = 3600
session_max_size: int = 10000
default_redirect_path: str = "/"
xmlsec_binary: Optional[str] = None
attribute_map: Dict[str, str] = Field(
default_factory=lambda: dict(DEFAULT_SAML_ATTRIBUTE_MAP)
)
@model_validator(mode="after")
def _require_idp_metadata(self) -> "SamlConfig":
def _require_idp_metadata(self) -> "SAMLConfig":
if self.enabled and not self.idp_metadata.strip():
raise ValueError(
"SAML enabled but idp_metadata is empty (inline XML, local path, or URL)"
@ -88,10 +93,6 @@ class SamlConfig(BaseModel):
class AuthConfig(BaseModel):
# First-match-wins precedence. HTTP precedes OPENID_CONNECT, so a bearer JWT
# is claimed by HttpAuthenticator (auth_method=bearer_jwt) and OidcAuthenticator
# never runs; both share the same JwtVerifiers and verify identically, so this
# only changes the auth_method label. Reorder if openIdConnect labeling matters.
scheme_order: List[SecuritySchemeType] = Field(
default_factory=lambda: [
SecuritySchemeType.API_KEY,
@ -103,9 +104,10 @@ class AuthConfig(BaseModel):
)
api_key: Optional[ApiKeySchemeConfig] = Field(default_factory=ApiKeySchemeConfig)
http_basic: HttpBasicConfig = Field(default_factory=HttpBasicConfig)
oidc_providers: List[OidcProviderConfig] = Field(default_factory=list)
oidc_providers: List[OIDCProviderConfig] = Field(default_factory=list)
oauth2_introspection: Optional[OAuth2IntrospectionConfig] = None
mutual_tls: MutualTlsConfig = Field(default_factory=MutualTlsConfig)
mutual_tls: MutualTLSConfig = Field(default_factory=MutualTLSConfig)
network: TrustedProxyConfig = Field(default_factory=TrustedProxyConfig)
saml: Optional[SamlConfig] = None
session: SessionConfig = Field(default_factory=SessionConfig)
saml: Optional[SAMLConfig] = None
casbin_policy_path: Optional[str] = None

View file

@ -57,7 +57,7 @@ _DEFAULT_POLICY: List[Tuple[str, str, str]] = [
]
class RbacEngine:
class RBACEngine:
def __init__(self, policy_path: Optional[str] = None) -> None:
model = casbin.Model()
model.load_model_from_text(_MODEL_TEXT)
@ -75,7 +75,7 @@ class RbacEngine:
self._enforcer.enforce(role.value, obj, act) for role in principal.roles
)
def has_role(self, principal: "Principal", allowed: Tuple[Role, ...]) -> bool:
def has_any_role(self, principal: "Principal", allowed: Tuple[Role, ...]) -> bool:
allowed_values = {role.value for role in allowed}
for role in principal.roles:
if role.value in allowed_values:

View file

@ -1,8 +1,7 @@
from __future__ import annotations
import secrets
import time
from typing import Any, Dict, List, Optional, Tuple
from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import RedirectResponse, Response
@ -13,9 +12,12 @@ from saml2.metadata import entity_descriptor
from scim2_models import Email, Name
from scim2_models import User as ScimUser
from .config import SamlConfig
from .models import AuthMethod, Credential, CredentialRef, SecuritySchemeType
from .config import SAMLConfig
from .resolver import ProvisioningStore
from .session import safe_relay_state
if TYPE_CHECKING:
from .security import AuthSecurity
_SINGLE_VALUE_TARGETS = {
"email",
@ -85,18 +87,6 @@ def _claims_from_mapped(mapped: Dict[str, Any]) -> Dict[str, Any]:
return claims
def _safe_relay_state(target: Optional[str], default: str) -> str:
if (
target
and target.startswith("/")
and not target.startswith("//")
and "://" not in target
and "\\" not in target
):
return target
return default
def _metadata_source(idp_metadata: str) -> Dict[str, Any]:
stripped = idp_metadata.strip()
if stripped.startswith("<"):
@ -106,7 +96,7 @@ def _metadata_source(idp_metadata: str) -> Dict[str, Any]:
return {"local": [idp_metadata]}
def _sp_config_dict(config: SamlConfig) -> Dict[str, Any]:
def _sp_config_dict(config: SAMLConfig) -> Dict[str, Any]:
cfg: Dict[str, Any] = {
"entityid": config.entity_id,
"service": {
@ -132,21 +122,19 @@ def _sp_config_dict(config: SamlConfig) -> Dict[str, Any]:
return cfg
def build_sp_client(config: SamlConfig) -> Saml2Client:
def build_sp_client(config: SAMLConfig) -> Saml2Client:
conf = SPConfig()
conf.load(_sp_config_dict(config))
return Saml2Client(config=conf)
class SamlSessionStore:
def __init__(self, ttl_seconds: int = 3600, max_size: int = 10000) -> None:
self._sessions: Dict[str, Tuple[float, Dict[str, Any]]] = {}
self._seen_assertions: Dict[str, float] = {}
class SAMLProtocolStore:
def __init__(self, replay_ttl_seconds: int) -> None:
self.outstanding: Dict[str, str] = {}
self._ttl = ttl_seconds
self._max_size = max_size
self._seen_assertions: Dict[str, float] = {}
self._replay_ttl = replay_ttl_seconds
def remember_request(self, request_id: str, relay_state: str = "/") -> None:
def remember_request(self, request_id: str, relay_state: str) -> None:
self.outstanding[request_id] = relay_state
def consume_assertion(self, assertion_id: str) -> bool:
@ -156,65 +144,16 @@ class SamlSessionStore:
}
if assertion_id in self._seen_assertions:
return False
self._seen_assertions[assertion_id] = now + self._ttl
self._seen_assertions[assertion_id] = now + self._replay_ttl
return True
def create_session(self, identity: Dict[str, Any]) -> str:
now = time.time()
self._evict(now)
session_id = secrets.token_urlsafe(32)
self._sessions[session_id] = (now + self._ttl, identity)
return session_id
def get(self, session_id: str) -> Optional[Dict[str, Any]]:
entry = self._sessions.get(session_id)
if entry is None:
return None
expires_at, identity = entry
if expires_at < time.time():
self._sessions.pop(session_id, None)
return None
return identity
def _evict(self, now: float) -> None:
for key in [k for k, (exp, _) in self._sessions.items() if exp < now]:
self._sessions.pop(key, None)
overflow = len(self._sessions) - self._max_size + 1
if overflow > 0:
oldest = sorted(self._sessions, key=lambda k: self._sessions[k][0])
for key in oldest[:overflow]:
self._sessions.pop(key, None)
class SamlAuthenticator:
scheme = SecuritySchemeType.HTTP
def __init__(self, config: SamlConfig, session_store: SamlSessionStore) -> None:
self._config = config
self._store = session_store
async def authenticate(self, request: Request) -> Optional[Credential]:
session_id = request.cookies.get(self._config.session_cookie)
if not session_id:
return None
identity = self._store.get(session_id)
if identity is None:
return None
return Credential(
scheme=self.scheme,
method=AuthMethod.SAML,
subject=identity["name_id"],
issuer=identity.get("issuer"),
claims=identity.get("claims", {}),
credential_ref=CredentialRef(token_id=session_id),
)
def challenge(self) -> str:
return ""
def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> APIRouter:
def build_saml_router(auth: "AuthSecurity") -> APIRouter:
config = auth.config.saml
assert config is not None
session = auth.config.session
client = build_sp_client(config)
protocol = SAMLProtocolStore(session.ttl_seconds)
router = APIRouter(prefix="/auth/saml", tags=["saml"])
@router.get("/metadata")
@ -226,11 +165,11 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP
@router.get("/login")
async def login(request: Request) -> RedirectResponse:
relay_state = _safe_relay_state(
request.query_params.get("next"), config.default_redirect_path
relay_state = safe_relay_state(
request.query_params.get("next"), session.default_redirect_path
)
request_id, info = client.prepare_for_authenticate(relay_state=relay_state)
session_store.remember_request(request_id, relay_state)
protocol.remember_request(request_id, relay_state)
location = dict(info["headers"]).get("Location")
if not location:
raise HTTPException(status_code=500, detail="no SAML redirect produced")
@ -246,7 +185,7 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP
authn_response = client.parse_authn_request_response(
saml_response,
BINDING_HTTP_POST,
outstanding=session_store.outstanding or None,
outstanding=protocol.outstanding or None,
)
except Exception as exc:
raise HTTPException(
@ -257,14 +196,12 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP
in_response_to = getattr(authn_response, "in_response_to", None)
bound_relay = (
session_store.outstanding.pop(in_response_to, None)
if in_response_to
else None
protocol.outstanding.pop(in_response_to, None) if in_response_to else None
)
assertion = getattr(authn_response, "assertion", None)
assertion_id = getattr(assertion, "id", None)
if assertion_id and not session_store.consume_assertion(assertion_id):
if assertion_id and not protocol.consume_assertion(assertion_id):
raise HTTPException(status_code=401, detail="SAML assertion replay")
name_id = authn_response.get_subject().text
@ -272,24 +209,25 @@ def build_saml_router(config: SamlConfig, session_store: SamlSessionStore) -> AP
mapped = _map_attributes(ava, config.attribute_map)
user = _user_from_mapped(name_id, mapped)
store: ProvisioningStore = request.app.state.auth_v2.resolver
store = cast(ProvisioningStore, auth.resolver)
await store.upsert_user(user)
session_id = session_store.create_session(
session_id = auth.session_store.create_session(
{
"name_id": name_id,
"method": "saml",
"subject": name_id,
"issuer": authn_response.issuer(),
"claims": _claims_from_mapped(mapped),
}
)
target = _safe_relay_state(bound_relay, config.default_redirect_path)
target = safe_relay_state(bound_relay, session.default_redirect_path)
response = RedirectResponse(target, status_code=303)
response.set_cookie(
config.session_cookie,
session.cookie,
session_id,
httponly=True,
samesite="lax",
secure=config.cookie_secure,
secure=session.secure,
)
return response

View file

@ -1,78 +1,20 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Annotated, Callable, List, Optional
from fastapi import FastAPI, Request, Security
from fastapi import Request, Security
from fastapi.security import SecurityScopes
from . import errors
from .authenticators import Authenticator, BasicAuthVerifier, build_authenticators
from .authenticators import (
Authenticator,
BasicAuthVerifier,
build_authenticators,
)
from .config import AuthConfig
from .models import Principal
from .network import resolve_network_context
from .rbac import RbacEngine, Role, has_required_scopes
from .rbac import RBACEngine, Role, has_required_scopes
from .resolver import IdentityResolver
@dataclass
class AuthContext:
config: AuthConfig
authenticators: List[Authenticator]
resolver: IdentityResolver
rbac: RbacEngine = field(default_factory=RbacEngine)
def install_auth(
app: FastAPI,
config: AuthConfig,
resolver: IdentityResolver,
*,
rbac: Optional[RbacEngine] = None,
basic_verifier: Optional[BasicAuthVerifier] = None,
mount_scim: bool = True,
mount_oidc: bool = True,
mount_saml: bool = True,
) -> AuthContext:
"""Wire the authenticators, resolver and optional routers onto the app.
Deployment requirement for trusted-proxy IP resolution: uvicorn's
``--proxy-headers`` (enabled by default) overwrites ``request.client`` from
``X-Forwarded-For`` before this module's ``trusted_proxy_cidrs`` check runs,
which silently bypasses it. Run uvicorn with ``--no-proxy-headers`` and let
this module resolve the client IP, or leave ``trusted_proxy_cidrs`` empty and
rely on uvicorn's own ``--forwarded-allow-ips``. Do not enable both.
"""
engine = rbac if rbac is not None else RbacEngine(config.casbin_policy_path)
ctx = AuthContext(
config,
build_authenticators(config, basic_verifier=basic_verifier),
resolver,
engine,
)
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))
if mount_saml and config.saml is not None and config.saml.enabled:
from .saml import SamlAuthenticator, SamlSessionStore, build_saml_router
session_store = SamlSessionStore(
ttl_seconds=config.saml.session_ttl_seconds,
max_size=config.saml.session_max_size,
)
ctx.authenticators.append(SamlAuthenticator(config.saml, session_store))
app.include_router(build_saml_router(config.saml, session_store))
return ctx
def _ctx(request: Request) -> AuthContext:
return request.app.state.auth_v2
from .session import SessionAuthenticator, SessionStore
def _combined_challenge(authenticators: List[Authenticator]) -> str:
@ -84,47 +26,82 @@ def _combined_challenge(authenticators: List[Authenticator]) -> str:
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))
class AuthSecurity:
"""Enforcement layer consumed purely through FastAPI ``Security()``.
resolved = await ctx.resolver.resolve(credential)
principal = resolved.model_copy(
update={"network": resolve_network_context(request, ctx.config.network)}
)
Construct once at the composition root and pass the bound methods
(``principal``, ``require_roles``, ``require_permission``) to ``Security()``;
routers receive the instance explicitly via ``build_*_router(auth)``. There is
no app mutation and no ``app.state``.
if not has_required_scopes(security_scopes, principal):
raise errors.insufficient_scope()
return principal
Deployment note for trusted-proxy IP resolution: uvicorn's ``--proxy-headers``
(on by default) overwrites ``request.client`` from ``X-Forwarded-For`` before
this module's ``trusted_proxy_cidrs`` check runs, silently bypassing it. Run
uvicorn with ``--no-proxy-headers`` and let this module resolve the client IP,
or leave ``trusted_proxy_cidrs`` empty and rely on uvicorn's
``--forwarded-allow-ips``. Do not enable both.
"""
def __init__(
self,
config: AuthConfig,
resolver: IdentityResolver,
rbac: Optional[RBACEngine] = None,
authenticators: Optional[List[Authenticator]] = None,
basic_verifier: Optional[BasicAuthVerifier] = None,
) -> None:
self.config = config
self.resolver = resolver
self.rbac = rbac or RBACEngine(config.casbin_policy_path)
self.session_store = SessionStore(
config.session.ttl_seconds, config.session.max_size
)
self.oauth_txn_store = SessionStore(
config.session.login_state_ttl, config.session.max_size
)
chain = (
list(authenticators)
if authenticators is not None
else build_authenticators(config, basic_verifier=basic_verifier)
)
chain.append(SessionAuthenticator(config.session.cookie, self.session_store))
self.authenticators = chain
def require_roles(*allowed: Role) -> Callable[..., object]:
async def dependency(
request: Request,
principal: Annotated[Principal, Security(get_current_principal)],
async def principal(
self, security_scopes: SecurityScopes, request: Request
) -> Principal:
if not _ctx(request).rbac.has_role(principal, allowed):
raise errors.forbidden_role()
credential = None
for authenticator in self.authenticators:
credential = await authenticator.authenticate(request)
if credential is not None:
break
if credential is None:
raise errors.unauthenticated(_combined_challenge(self.authenticators))
resolved = await self.resolver.resolve(credential)
principal = resolved.model_copy(
update={"network": resolve_network_context(request, self.config.network)}
)
if not has_required_scopes(security_scopes, principal):
raise errors.insufficient_scope()
return principal
return dependency
def require_roles(self, *allowed: Role) -> Callable[..., object]:
async def dependency(
principal: Annotated[Principal, Security(self.principal)],
) -> Principal:
if not self.rbac.has_any_role(principal, allowed):
raise errors.forbidden_role()
return principal
return dependency
def require_permission(obj: str, act: str) -> Callable[..., object]:
async def dependency(
request: Request,
principal: Annotated[Principal, Security(get_current_principal)],
) -> Principal:
if not _ctx(request).rbac.enforce(principal, obj, act):
raise errors.forbidden_permission()
return principal
def require_permission(self, obj: str, act: str) -> Callable[..., object]:
async def dependency(
principal: Annotated[Principal, Security(self.principal)],
) -> Principal:
if not self.rbac.enforce(principal, obj, act):
raise errors.forbidden_permission()
return principal
return dependency
return dependency

View file

@ -0,0 +1,88 @@
from __future__ import annotations
import secrets
import time
from typing import Any, Dict, Optional, Tuple
from fastapi import Request
from .models import AuthMethod, Credential, CredentialRef, SecuritySchemeType
def safe_relay_state(target: Optional[str], default: str) -> str:
if (
target
and target.startswith("/")
and not target.startswith("//")
and "://" not in target
and "\\" not in target
):
return target
return default
class SessionStore:
def __init__(self, ttl_seconds: int = 3600, max_size: int = 10000) -> None:
self._sessions: Dict[str, Tuple[float, Dict[str, Any]]] = {}
self._ttl = ttl_seconds
self._max_size = max_size
def create_session(self, identity: Dict[str, Any]) -> str:
now = time.time()
self._evict(now)
session_id = secrets.token_urlsafe(32)
self._sessions[session_id] = (now + self._ttl, identity)
return session_id
def get(self, session_id: str) -> Optional[Dict[str, Any]]:
entry = self._sessions.get(session_id)
if entry is None:
return None
expires_at, identity = entry
if expires_at < time.time():
self._sessions.pop(session_id, None)
return None
return identity
def pop(self, session_id: str) -> Optional[Dict[str, Any]]:
entry = self._sessions.pop(session_id, None)
if entry is None:
return None
expires_at, identity = entry
if expires_at < time.time():
return None
return identity
def _evict(self, now: float) -> None:
for key in [k for k, (exp, _) in self._sessions.items() if exp < now]:
self._sessions.pop(key, None)
overflow = len(self._sessions) - self._max_size + 1
if overflow > 0:
oldest = sorted(self._sessions, key=lambda k: self._sessions[k][0])
for key in oldest[:overflow]:
self._sessions.pop(key, None)
class SessionAuthenticator:
def __init__(self, cookie_name: str, store: SessionStore) -> None:
self._cookie_name = cookie_name
self._store = store
async def authenticate(self, request: Request) -> Optional[Credential]:
session_id = request.cookies.get(self._cookie_name)
if not session_id:
return None
identity = self._store.get(session_id)
if identity is None:
return None
return Credential(
scheme=SecuritySchemeType.API_KEY,
method=AuthMethod(identity["method"]),
subject=identity["subject"],
issuer=identity.get("issuer"),
claims=identity.get("claims", {}),
credential_ref=CredentialRef(token_id=session_id),
)
def challenge(self) -> str:
return ""