fix(auth_v2): authenticate OIDC login sessions and adapt SCIM to AuthSecurity

The OIDC callback now mints a server-side session and sets the shared session
cookie (httponly, secure, samesite=lax) before redirecting, so login yields an
authenticated session that the shared SessionAuthenticator resolves; previously
it returned the user record as JSON and left the caller unauthenticated. The
login flow owns its CSRF protection without Starlette SessionMiddleware: it
stores state, nonce and the PKCE (S256) verifier in the short-lived
oauth_txn_store keyed by a temporary cookie, then on callback consumes the
transaction one-time, checks the returned state, exchanges the code with the
verifier and validates the nonce against the id_token. A replayed or expired
state finds no transaction and returns 400.

SCIM moves onto the AuthSecurity DI surface: build_scim_router(auth) reads
auth.resolver and guards writes with Security(auth.principal, scopes=
["scim:write"]) instead of app.state and get_current_principal.
This commit is contained in:
Yassin Kortam 2026-06-10 19:15:43 -07:00
parent 4ec48302a5
commit 5f02c88369
3 changed files with 220 additions and 143 deletions

View file

@ -10,9 +10,11 @@ from .config import (
TrustedProxyConfig,
)
from .models import Principal
from .oidc import build_oidc_router
from .rbac import Role
from .resolver import IdentityResolver, InMemoryIdentityStore, ProvisioningStore
from .saml import build_saml_router
from .scim import build_scim_router
from .security import AuthSecurity
__all__ = [
@ -32,4 +34,6 @@ __all__ = [
"SessionConfig",
"SAMLConfig",
"build_saml_router",
"build_scim_router",
"build_oidc_router",
]

View file

@ -1,18 +1,24 @@
from __future__ import annotations
import re
from typing import Any, Dict
from typing import TYPE_CHECKING, Any, Dict, cast
from authlib.integrations.starlette_client import OAuth
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import JSONResponse
from fastapi.responses import RedirectResponse
from scim2_models import User as ScimUser
from .config import AuthConfig, OidcProviderConfig
from .config import OIDCProviderConfig
from .resolver import ProvisioningStore
from .session import safe_relay_state
if TYPE_CHECKING:
from .security import AuthSecurity
_CLAIM_KEYS = ("email", "preferred_username", "name", "groups", "roles")
def _provider_key(provider: OidcProviderConfig) -> str:
def _provider_key(provider: OIDCProviderConfig) -> str:
return re.sub(r"[^a-z0-9]+", "-", provider.issuer.lower()).strip("-")
@ -24,9 +30,15 @@ def _user_from_userinfo(userinfo: Dict[str, Any]) -> ScimUser:
)
def build_oidc_router(config: AuthConfig) -> APIRouter:
def _mapped_claims(userinfo: Dict[str, Any]) -> Dict[str, Any]:
return {key: userinfo[key] for key in _CLAIM_KEYS if userinfo.get(key) is not None}
def build_oidc_router(auth: AuthSecurity) -> APIRouter:
session = auth.config.session
issuers = {_provider_key(p): p.issuer for p in auth.config.oidc_providers}
oauth = OAuth()
for provider in config.oidc_providers:
for provider in auth.config.oidc_providers:
oauth.register(
name=_provider_key(provider),
server_metadata_url=f"{provider.issuer.rstrip('/')}/.well-known/openid-configuration",
@ -36,30 +48,98 @@ def build_oidc_router(config: AuthConfig) -> APIRouter:
if provider.client_secret
else None
),
client_kwargs={"scope": " ".join(provider.login_scopes)},
client_kwargs={
"scope": " ".join(provider.login_scopes),
"code_challenge_method": "S256",
},
)
router = APIRouter(prefix="/auth/oidc", tags=["oidc"])
@router.get("/{provider}/login")
async def login(provider: str, request: Request) -> Any:
async def login(provider: str, request: Request) -> RedirectResponse:
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))
redirect_uri = str(request.url_for("oidc_callback", provider=provider))
relay = safe_relay_state(
request.query_params.get("next"), session.default_redirect_path
)
authorization = await client.create_authorization_url(redirect_uri)
txn_id = auth.oauth_txn_store.create_session(
{
"provider": provider,
"state": authorization["state"],
"nonce": authorization.get("nonce"),
"code_verifier": authorization.get("code_verifier"),
"redirect_uri": redirect_uri,
"relay": relay,
}
)
response = RedirectResponse(authorization["url"], status_code=303)
response.set_cookie(
session.login_cookie,
txn_id,
httponly=True,
samesite="lax",
secure=session.secure,
max_age=session.login_state_ttl,
)
return response
@router.get("/{provider}/callback", name="oidc_callback")
async def callback(provider: str, request: Request) -> JSONResponse:
async def callback(provider: str, request: Request) -> RedirectResponse:
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:
txn_id = request.cookies.get(session.login_cookie)
txn = auth.oauth_txn_store.pop(txn_id) if txn_id else None
if txn is None or txn.get("provider") != provider:
raise HTTPException(
status_code=400, detail="invalid or expired login state"
)
returned_state = request.query_params.get("state")
if not returned_state or returned_state != txn["state"]:
raise HTTPException(status_code=400, detail="state mismatch")
error = request.query_params.get("error")
if error:
raise HTTPException(status_code=400, detail=error)
code = request.query_params.get("code")
if not code:
raise HTTPException(status_code=400, detail="missing authorization code")
token = await client.fetch_access_token(
redirect_uri=txn["redirect_uri"],
code=code,
code_verifier=txn.get("code_verifier"),
state=txn["state"],
)
if token.get("id_token"):
userinfo = await client.parse_id_token(token, nonce=txn.get("nonce"))
else:
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())
info = dict(userinfo)
store = cast(ProvisioningStore, auth.resolver)
await store.upsert_user(_user_from_userinfo(info))
session_id = auth.session_store.create_session(
{
"method": "oidc",
"subject": info.get("sub"),
"issuer": info.get("iss") or issuers.get(provider),
"claims": _mapped_claims(info),
}
)
target = safe_relay_state(txn.get("relay"), session.default_redirect_path)
response = RedirectResponse(target, status_code=303)
response.set_cookie(
session.cookie,
session_id,
httponly=True,
samesite="lax",
secure=session.secure,
)
return response
return router

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from typing import Any, Dict, Optional, Type, TypeVar
from typing import TYPE_CHECKING, Any, Dict, Optional, Type, TypeVar, cast
from fastapi import APIRouter, Query, Request, Response, Security, status
from fastapi.responses import JSONResponse
@ -24,15 +24,13 @@ from scim2_models import (
)
from .resolver import ProvisioningStore
from .security import get_current_principal
if TYPE_CHECKING:
from .security import AuthSecurity
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,
@ -89,6 +87,119 @@ def _dump(resource: Resource, ctx: Context) -> Dict[str, Any]:
return resource.model_dump(scim_ctx=ctx)
def _build_protected_router(auth: AuthSecurity) -> APIRouter:
store = cast(ProvisioningStore, auth.resolver)
protected = APIRouter(
dependencies=[Security(auth.principal, scopes=["scim:write"])],
)
@protected.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.upsert_user(user)
return JSONResponse(
status_code=status.HTTP_201_CREATED,
content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE),
)
@protected.get("/Users/{resource_id}")
async def get_user(resource_id: str) -> Response:
user = await store.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))
@protected.patch("/Users/{resource_id}")
async def patch_user(resource_id: str, request: Request) -> Response:
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())
patched = _apply_patch(user, patch)
except (ValidationError, ValueError) as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
updated = await store.upsert_user(patched)
return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE))
@protected.delete("/Users/{resource_id}", status_code=status.HTTP_204_NO_CONTENT)
async def deactivate_user(resource_id: str) -> Response:
if await store.get_user(resource_id) is None:
return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found")
await store.deactivate_user(resource_id)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@protected.get("/Users")
async def list_users(
filter_expr: Optional[str] = Query(default=None, alias="filter"),
) -> Response:
users = await store.list_users(filter_expr)
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))
@protected.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.upsert_group(group)
return JSONResponse(
status_code=status.HTTP_201_CREATED,
content=_dump(stored, Context.RESOURCE_CREATION_RESPONSE),
)
@protected.get("/Groups/{resource_id}")
async def get_group(resource_id: str) -> Response:
group = await store.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))
@protected.patch("/Groups/{resource_id}")
async def patch_group(resource_id: str, request: Request) -> Response:
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())
patched = _apply_patch(group, patch)
except (ValidationError, ValueError) as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
updated = await store.upsert_group(patched)
return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE))
@protected.delete("/Groups/{resource_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_group(resource_id: str) -> Response:
if await store.get_group(resource_id) is None:
return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found")
await store.delete_group(resource_id)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@protected.get("/Groups")
async def list_groups(
filter_expr: Optional[str] = Query(default=None, alias="filter"),
) -> Response:
groups = await store.list_groups(filter_expr)
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))
return protected
def _build_discovery_router() -> APIRouter:
router = APIRouter()
@ -143,126 +254,8 @@ def _build_discovery_router() -> APIRouter:
return router
def _build_protected_router() -> APIRouter:
protected = APIRouter(
dependencies=[Security(get_current_principal, scopes=["scim:write"])],
)
@protected.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),
)
@protected.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))
@protected.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())
patched = _apply_patch(user, patch)
except (ValidationError, ValueError) as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
updated = await store.upsert_user(patched)
return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE))
@protected.delete("/Users/{resource_id}", status_code=status.HTTP_204_NO_CONTENT)
async def deactivate_user(resource_id: str, request: Request) -> Response:
store = _store(request)
if await store.get_user(resource_id) is None:
return _error(status.HTTP_404_NOT_FOUND, f"User {resource_id} not found")
await store.deactivate_user(resource_id)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@protected.get("/Users")
async def list_users(
request: Request,
filter_expr: Optional[str] = Query(default=None, alias="filter"),
) -> Response:
users = await _store(request).list_users(filter_expr)
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))
@protected.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),
)
@protected.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))
@protected.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())
patched = _apply_patch(group, patch)
except (ValidationError, ValueError) as exc:
return _error(status.HTTP_400_BAD_REQUEST, str(exc))
updated = await store.upsert_group(patched)
return JSONResponse(content=_dump(updated, Context.RESOURCE_PATCH_RESPONSE))
@protected.delete("/Groups/{resource_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_group(resource_id: str, request: Request) -> Response:
store = _store(request)
if await store.get_group(resource_id) is None:
return _error(status.HTTP_404_NOT_FOUND, f"Group {resource_id} not found")
await store.delete_group(resource_id)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@protected.get("/Groups")
async def list_groups(
request: Request,
filter_expr: Optional[str] = Query(default=None, alias="filter"),
) -> Response:
groups = await _store(request).list_groups(filter_expr)
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))
return protected
def build_scim_router() -> APIRouter:
def build_scim_router(auth: AuthSecurity) -> APIRouter:
router = APIRouter(prefix="/scim/v2", tags=["scim"])
router.include_router(_build_protected_router())
router.include_router(_build_protected_router(auth))
router.include_router(_build_discovery_router())
return router