feat: oidc renewal, groups, logout, scim

This commit is contained in:
Alex committed 2026-06-10 00:57:18 +01:00
1 parent 8bc777a428
commit 054f0f1b8b
29 files changed
+3392 -155

No files matched your search

+10
View File
@@ -51,3 +51,13 @@ MICROSOFT_AUTHORITY=https://{tenantId}.ciamlogin.com/{tenantId}
# OIDC_USER_ID_CLAIM=sub
# OIDC_REDIRECT_URI=<override callback URL when behind a reverse proxy>
# OIDC_SESSION_LIFETIME_SECONDS=28800
# OIDC_PROVIDER_NAME=<sign-in button label, e.g. Acme SSO; unset shows "SSO">
# OIDC_ALLOWED_GROUPS=<comma-separated IdP group allowlist; unset = any authenticated user>
# OIDC_GROUPS_CLAIM=groups
# Add offline_access to OIDC_SCOPES for silent session renewal on IdPs that
# require it for refresh tokens (Authentik does; Keycloak does not).
# SCIM 2.0 provisioning (IdP-driven user create/deactivate at /scim/v2;
# pair with OIDC_USER_ID_CLAIM=email so SCIM userName matches the OIDC user id)
# SCIM_ENABLED=false
# SCIM_TOKEN=<long random bearer token presented by the IdP's SCIM client>
@@ -0,0 +1,45 @@
"""0017 oidc scim — users.active flag + auth_events audit table.
``users.active`` backs SCIM deprovisioning: deactivated users are refused new
OIDC sessions and their live sessions are denylisted until they expire.
``auth_events`` is an append-only audit trail of login / logout / provisioning
events keyed by ``user_id``.
Revision ID: 0017_oidc_scim
Revises: 0016_conversation_visibility
"""
from typing import Sequence, Union
from alembic import op
revision: str = "0017_oidc_scim"
down_revision: Union[str, None] = "0016_conversation_visibility"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.execute("ALTER TABLE users ADD COLUMN active BOOLEAN NOT NULL DEFAULT TRUE;")
op.execute(
"""
CREATE TABLE auth_events (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
user_id TEXT NOT NULL,
event TEXT NOT NULL,
ip TEXT,
user_agent TEXT,
metadata JSONB NOT NULL DEFAULT '{}',
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
"""
)
op.execute(
"CREATE INDEX auth_events_user_idx ON auth_events (user_id, created_at DESC);"
)
def downgrade() -> None:
op.execute("DROP TABLE IF EXISTS auth_events;")
op.execute("ALTER TABLE users DROP COLUMN IF EXISTS active;")
+5
View File
@@ -24,6 +24,7 @@ from starlette.requests import Request
from starlette.responses import JSONResponse, StreamingResponse
from starlette.routing import Route
from application.api.oidc.denylist import is_denied as oidc_session_denied
from application.auth import handle_auth
from application.core.settings import settings
from application.events.keys import connection_counter_key
@@ -184,6 +185,10 @@ async def stream_message_events(request: Request) -> JSONResponse | StreamingRes
user_id = decoded.get("sub") if isinstance(decoded, dict) else None
if not user_id:
return _json("Authentication required", 401)
if settings.AUTH_TYPE == "oidc" and await anyio.to_thread.run_sync(
oidc_session_denied, decoded
):
return _json("Authentication error: session revoked", 401)
message_id = request.path_params["message_id"]
if not _MESSAGE_ID_RE.match(message_id):
+97
View File
@@ -0,0 +1,97 @@
"""Redis-backed session denylist for OIDC revocation.
Back-channel logout and SCIM deactivation drop identifiers here; the
request path refuses any session token whose identifiers match. Entries
live slightly longer than ``OIDC_SESSION_LIFETIME_SECONDS`` — every
session minted before the revocation expires before its denylist entry
does, so nothing needs to be stored durably.
Revocation is best-effort by design: if Redis is unreachable the check
fails open (sessions keep working) rather than taking the whole API down.
"""
from __future__ import annotations
import logging
from application.cache import get_redis_instance
from application.core.settings import settings
logger = logging.getLogger(__name__)
_USER_PREFIX = "oidc:deny:user:"
_SUB_PREFIX = "oidc:deny:sub:"
_SID_PREFIX = "oidc:deny:sid:"
def _ttl_seconds() -> int:
return settings.OIDC_SESSION_LIFETIME_SECONDS + 60
def _set(key: str) -> bool:
redis = get_redis_instance()
if redis is None:
logger.error("Redis unavailable — could not denylist %s", key)
return False
try:
redis.set(key, "1", ex=_ttl_seconds())
return True
except Exception:
logger.error("Failed to denylist %s", key, exc_info=True)
return False
def deny_user(user_id: str) -> bool:
"""Revoke every live session of the DocsGPT user ``user_id``."""
return _set(_USER_PREFIX + user_id)
def deny_idp_sub(sub: str) -> bool:
"""Revoke sessions by IdP ``sub`` (back-channel logout tokens carry this)."""
return _set(_SUB_PREFIX + sub)
def deny_sid(sid: str) -> bool:
"""Revoke sessions of one IdP session id (``sid``-only logout tokens)."""
return _set(_SID_PREFIX + sid)
def allow_user(user_id: str) -> None:
"""Clear a user-level denylist entry (SCIM reactivation)."""
_delete(_USER_PREFIX + user_id)
def allow_idp_sub(sub: str) -> None:
"""Clear an IdP-sub denylist entry (fresh login supersedes a back-channel logout)."""
_delete(_SUB_PREFIX + sub)
def _delete(key: str) -> None:
redis = get_redis_instance()
if redis is None:
return
try:
redis.delete(key)
except Exception:
logger.warning("Failed to clear denylist key %s", key, exc_info=True)
def is_denied(decoded_token: dict) -> bool:
"""True when any identifier in a decoded session token is denylisted."""
keys = []
if decoded_token.get("sub"):
keys.append(_USER_PREFIX + str(decoded_token["sub"]))
if decoded_token.get("oidc_sub"):
keys.append(_SUB_PREFIX + str(decoded_token["oidc_sub"]))
if decoded_token.get("oidc_sid"):
keys.append(_SID_PREFIX + str(decoded_token["oidc_sid"]))
if not keys:
return False
redis = get_redis_instance()
if redis is None:
return False
try:
return any(value is not None for value in redis.mget(keys))
except Exception:
logger.warning("Denylist check failed — allowing request", exc_info=True)
return False
+150 -49
View File
@@ -1,4 +1,4 @@
"""OIDC provider client: discovery, JWKS, code exchange, ID-token validation."""
"""OIDC provider client: discovery, JWKS, token grants, ID/logout-token validation."""
from __future__ import annotations
@@ -8,14 +8,17 @@ import time
import requests
from jose import jwt
from jose.exceptions import ExpiredSignatureError, JWTClaimsError
from application.core.settings import settings
logger = logging.getLogger(__name__)
DISCOVERY_TTL_SECONDS = 3600
FORCE_REFETCH_COOLDOWN_SECONDS = 10
LEEWAY_SECONDS = 60
HTTP_TIMEOUT_SECONDS = 10
BACKCHANNEL_LOGOUT_EVENT = "http://schemas.openid.net/event/backchannel-logout"
# Asymmetric algorithms only: a symmetric alg here would let an attacker
# forge ID tokens signed with the (public) JWKS material.
ALLOWED_ID_TOKEN_ALGS = [
@@ -25,7 +28,13 @@ ALLOWED_ID_TOKEN_ALGS = [
]
_lock = threading.Lock()
_cache: dict = {"discovery": None, "discovery_at": 0.0, "jwks": None, "jwks_at": 0.0}
_cache: dict = {
"discovery": None,
"discovery_at": 0.0,
"jwks": None,
"jwks_at": 0.0,
"jwks_force_at": 0.0,
}
class OIDCError(Exception):
@@ -35,7 +44,9 @@ class OIDCError(Exception):
def reset_cache() -> None:
"""Clear the cached discovery document and JWKS (used by tests)."""
with _lock:
_cache.update({"discovery": None, "discovery_at": 0.0, "jwks": None, "jwks_at": 0.0})
_cache.update(
{"discovery": None, "discovery_at": 0.0, "jwks": None, "jwks_at": 0.0, "jwks_force_at": 0.0}
)
def _fetch_json(url: str) -> dict:
@@ -64,12 +75,18 @@ def get_discovery() -> dict:
def get_jwks(force: bool = False) -> dict:
"""Return the IdP JWKS; ``force=True`` bypasses the cache (key rotation)."""
with _lock:
if (
not force
and _cache["jwks"] is not None
fresh = (
_cache["jwks"] is not None
and time.time() - _cache["jwks_at"] < DISCOVERY_TTL_SECONDS
):
)
if not force and fresh:
return _cache["jwks"]
if force and fresh:
# Rate-limit forced refetches: unauthenticated callers (back-channel
# logout) must not be able to hammer the IdP through us.
if time.time() - _cache["jwks_force_at"] < FORCE_REFETCH_COOLDOWN_SECONDS:
return _cache["jwks"]
_cache["jwks_force_at"] = time.time()
jwks = _fetch_json(get_discovery()["jwks_uri"])
with _lock:
_cache["jwks"] = jwks
@@ -84,14 +101,14 @@ def _find_key(kid: str | None) -> dict | None:
return next((key for key in keys if key.get("kid") == kid), None)
def validate_id_token(id_token: str, nonce: str) -> dict:
"""Verify the ID token's signature, iss, aud, exp, and nonce; return claims."""
def _resolve_signing_key(token: str) -> dict:
"""Return the JWKS key matching the token header, refetching once on unknown kid."""
try:
header = jwt.get_unverified_header(id_token)
header = jwt.get_unverified_header(token)
except Exception as exc:
raise OIDCError(f"Malformed id_token: {exc}") from exc
raise OIDCError(f"Malformed token: {exc}") from exc
if header.get("alg") not in ALLOWED_ID_TOKEN_ALGS:
raise OIDCError(f"Disallowed id_token alg: {header.get('alg')}")
raise OIDCError(f"Disallowed token alg: {header.get('alg')}")
key = _find_key(header.get("kid"))
if key is None:
@@ -99,53 +116,137 @@ def validate_id_token(id_token: str, nonce: str) -> dict:
key = _find_key(header.get("kid"))
if key is None:
raise OIDCError("No matching key in IdP JWKS")
return key
def _decode_verified(token: str, options: dict) -> dict:
"""Decode ``token`` against the JWKS, retrying once if the IdP re-keyed.
A signature failure can mean the IdP replaced its signing key while
reusing the same kid — the kid-miss refetch never triggers then, so
retry once against a freshly fetched JWKS (rate-limited in get_jwks).
"""
key = _resolve_signing_key(token)
decode_kwargs = {
"algorithms": ALLOWED_ID_TOKEN_ALGS,
"audience": settings.OIDC_CLIENT_ID,
# Compare against the discovery document's own issuer value —
# some IdPs (Authentik) use a trailing slash the operator may
# not have typed into OIDC_ISSUER.
"issuer": get_discovery()["issuer"],
"options": options,
}
try:
claims = jwt.decode(
id_token,
key,
algorithms=ALLOWED_ID_TOKEN_ALGS,
audience=settings.OIDC_CLIENT_ID,
# Compare against the discovery document's own issuer value —
# some IdPs (Authentik) use a trailing slash the operator may
# not have typed into OIDC_ISSUER.
issuer=get_discovery()["issuer"],
options={
"verify_at_hash": False,
"leeway": LEEWAY_SECONDS,
"require_iss": True,
"require_aud": True,
"require_exp": True,
"require_sub": True,
},
)
except Exception as exc:
raise OIDCError(f"id_token validation failed: {exc}") from exc
if claims.get("nonce") != nonce:
return jwt.decode(token, key, **decode_kwargs)
except (ExpiredSignatureError, JWTClaimsError) as exc:
raise OIDCError(f"token validation failed: {exc}") from exc
except Exception:
get_jwks(force=True)
key = _find_key(jwt.get_unverified_header(token).get("kid"))
if key is None:
raise OIDCError("No matching key in IdP JWKS")
try:
return jwt.decode(token, key, **decode_kwargs)
except Exception as exc:
raise OIDCError(f"token validation failed: {exc}") from exc
def validate_id_token(id_token: str, nonce: str | None = None) -> dict:
"""Verify the ID token's signature, iss, aud, exp, and (when given) nonce; return claims."""
claims = _decode_verified(
id_token,
options={
"verify_at_hash": False,
"leeway": LEEWAY_SECONDS,
"require_iss": True,
"require_aud": True,
"require_exp": True,
"require_sub": True,
},
)
# Refresh-issued id_tokens carry no nonce; callers pass None to skip the check.
if nonce is not None and claims.get("nonce") != nonce:
raise OIDCError("nonce mismatch")
return claims
def exchange_code(code: str, code_verifier: str, redirect_uri: str) -> dict:
"""Exchange the authorization code at the IdP token endpoint."""
data = {
"grant_type": "authorization_code",
"code": code,
"redirect_uri": redirect_uri,
"client_id": settings.OIDC_CLIENT_ID,
"code_verifier": code_verifier,
}
def validate_logout_token(logout_token: str) -> dict:
"""Verify a back-channel logout token per OIDC Back-Channel Logout 1.0; return claims."""
claims = _decode_verified(
logout_token,
options={
"verify_at_hash": False,
"leeway": LEEWAY_SECONDS,
"require_iss": True,
"require_aud": True,
"require_iat": True,
"require_exp": False,
},
)
events = claims.get("events")
if not isinstance(events, dict) or BACKCHANNEL_LOGOUT_EVENT not in events:
raise OIDCError("logout_token missing the backchannel-logout event")
if "nonce" in claims:
raise OIDCError("logout_token must not contain a nonce")
if not claims.get("sub") and not claims.get("sid"):
raise OIDCError("logout_token must contain sub or sid")
return claims
def _token_request(data: dict) -> dict:
"""POST to the token endpoint using the discovery-advertised client auth method."""
discovery = get_discovery()
data = {**data, "client_id": settings.OIDC_CLIENT_ID}
post_kwargs: dict = {"data": data, "timeout": HTTP_TIMEOUT_SECONDS}
if settings.OIDC_CLIENT_SECRET:
data["client_secret"] = settings.OIDC_CLIENT_SECRET
# Absent metadata means the RFC 8414 default, client_secret_basic.
methods = discovery.get("token_endpoint_auth_methods_supported") or ["client_secret_basic"]
if "client_secret_post" in methods:
data["client_secret"] = settings.OIDC_CLIENT_SECRET
else:
post_kwargs["auth"] = (settings.OIDC_CLIENT_ID, settings.OIDC_CLIENT_SECRET)
try:
response = requests.post(
get_discovery()["token_endpoint"], data=data, timeout=HTTP_TIMEOUT_SECONDS
)
response = requests.post(discovery["token_endpoint"], **post_kwargs)
except requests.RequestException as exc:
raise OIDCError(f"Token exchange request failed: {exc}") from exc
raise OIDCError(f"Token request failed: {exc}") from exc
if response.status_code != 200:
logger.error(
"OIDC token exchange failed (%s): %s", response.status_code, response.text[:500]
"OIDC token request failed (%s): %s", response.status_code, response.text[:500]
)
raise OIDCError(f"Token exchange returned {response.status_code}")
raise OIDCError(f"Token endpoint returned {response.status_code}")
return response.json()
def exchange_code(code: str, code_verifier: str, redirect_uri: str) -> dict:
"""Exchange the authorization code at the IdP token endpoint."""
return _token_request(
{
"grant_type": "authorization_code",
"code": code,
"redirect_uri": redirect_uri,
"code_verifier": code_verifier,
}
)
def refresh_grant(refresh_token: str) -> dict:
"""Redeem a refresh token at the IdP token endpoint."""
return _token_request({"grant_type": "refresh_token", "refresh_token": refresh_token})
def fetch_userinfo(access_token: str) -> dict:
"""Fetch claims from the IdP userinfo endpoint with a Bearer access token."""
endpoint = get_discovery().get("userinfo_endpoint")
if not endpoint:
raise OIDCError("No userinfo_endpoint in discovery document")
try:
response = requests.get(
endpoint,
headers={"Authorization": f"Bearer {access_token}"},
timeout=HTTP_TIMEOUT_SECONDS,
)
except requests.RequestException as exc:
raise OIDCError(f"userinfo request failed: {exc}") from exc
if response.status_code != 200:
raise OIDCError(f"userinfo endpoint returned {response.status_code}")
return response.json()
+336 -14
View File
@@ -1,4 +1,4 @@
"""Login, callback, and session-token endpoints for AUTH_TYPE=oidc.
"""Login, callback, session-token, logout, and refresh endpoints for AUTH_TYPE=oidc.
Flow: the backend redirects the browser to the IdP (Authorization Code +
PKCE), validates the ID token at the callback, mints a local HS256 session
@@ -15,19 +15,26 @@ import json
import logging
import secrets
import time
import uuid
from urllib.parse import quote, urlencode
from flask import Blueprint, jsonify, make_response, redirect, request
from flask import Blueprint, Response, jsonify, make_response, redirect, request
from jose import jwt
from application.api.oidc import provider
from application.api.oidc import denylist, provider
from application.auth import handle_auth
from application.cache import get_redis_instance
from application.core.settings import settings
from application.storage.db.repositories.auth_events import AuthEventsRepository
from application.storage.db.repositories.users import UsersRepository
from application.storage.db.session import db_readonly, db_session
logger = logging.getLogger(__name__)
STATE_TTL_SECONDS = 600
HANDOFF_TTL_SECONDS = 60
LOGOUT_JTI_TTL_SECONDS = 600
MAX_PICTURE_CLAIM_CHARS = 2048
def _state_key(state: str) -> str:
@@ -38,6 +45,14 @@ def _handoff_key(code: str) -> str:
return f"oidc:handoff:{code}"
def _refresh_key(jti: str) -> str:
return f"oidc:refresh:{jti}"
def _logout_jti_key(jti: str) -> str:
return f"oidc:bcl:jti:{jti}"
def _pkce_challenge(verifier: str) -> str:
digest = hashlib.sha256(verifier.encode("ascii")).digest()
return base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
@@ -56,6 +71,128 @@ def _frontend_redirect(fragment: str):
return redirect(f"{base}/#{fragment}", code=302)
def _no_store(payload, status: int = 200) -> Response:
"""Build a response marked non-cacheable (back-channel logout requirement)."""
response = make_response(payload, status)
response.headers["Cache-Control"] = "no-store"
return response
def _allowed_groups() -> list[str]:
"""Parse the comma-separated group allowlist; empty/unset means everyone."""
raw = settings.OIDC_ALLOWED_GROUPS or ""
return [group.strip() for group in raw.split(",") if group.strip()]
def _claim_groups(claims: dict) -> list[str]:
"""Read the groups claim as a list of strings; missing means no groups."""
value = claims.get(settings.OIDC_GROUPS_CLAIM)
if value is None:
return []
if isinstance(value, str):
return [value]
if isinstance(value, (list, tuple)):
return [str(member) for member in value]
return [str(value)]
def _effective_claims(tokens: dict, claims: dict) -> dict:
"""Merge userinfo into the id_token claims when required claims are missing."""
effective = dict(claims)
need_user_id = not effective.get(settings.OIDC_USER_ID_CLAIM)
need_groups = bool(_allowed_groups()) and settings.OIDC_GROUPS_CLAIM not in effective
if not (need_user_id or need_groups) or not tokens.get("access_token"):
return effective
try:
userinfo = provider.fetch_userinfo(tokens["access_token"])
except provider.OIDCError:
logger.warning("OIDC userinfo fetch failed; continuing with id_token claims", exc_info=True)
return effective
if userinfo.get("sub") != claims.get("sub"):
raise provider.OIDCError("userinfo sub does not match id_token sub")
for key, value in userinfo.items():
effective.setdefault(key, value)
return effective
def _mint_session_token(identity: dict) -> tuple[str, str]:
"""Mint the local HS256 session JWT for ``identity``; returns (token, jti)."""
now = int(time.time())
jti = str(uuid.uuid4())
payload = {
"sub": str(identity["sub"]),
"jti": jti,
"iat": now,
"exp": now + settings.OIDC_SESSION_LIFETIME_SECONDS,
}
if identity.get("oidc_sub"):
payload["oidc_sub"] = str(identity["oidc_sub"])
if identity.get("oidc_sid"):
payload["oidc_sid"] = str(identity["oidc_sid"])
for claim in ("email", "name"):
if identity.get(claim):
payload[claim] = identity[claim]
picture = identity.get("picture")
if picture and isinstance(picture, str) and len(picture) < MAX_PICTURE_CLAIM_CHARS:
payload["picture"] = picture
return jwt.encode(payload, settings.JWT_SECRET_KEY, algorithm="HS256"), jti
def _record_login_denied(user_id: str, metadata: dict) -> None:
"""Best-effort audit of a denied login; never raises."""
try:
with db_session() as conn:
AuthEventsRepository(conn).insert(
user_id,
"oidc_login_denied",
ip=request.remote_addr,
user_agent=request.headers.get("User-Agent"),
metadata=metadata,
)
except Exception:
logger.warning("Failed to record oidc_login_denied for %s", user_id, exc_info=True)
def _gate_and_audit_login(user_id: str, effective: dict, groups: list[str]) -> bool:
"""Reject disabled users, provision new ones, audit the login.
Returns False only when the user row was readable and marked inactive;
a DB outage logs an error and lets the login proceed.
"""
disabled = False
try:
with db_session() as conn:
users = UsersRepository(conn)
row = users.get(user_id)
if row is not None and row.get("active") is False:
disabled = True
AuthEventsRepository(conn).insert(
user_id,
"oidc_login_denied",
ip=request.remote_addr,
user_agent=request.headers.get("User-Agent"),
metadata={"reason": "account_disabled"},
)
else:
if row is None:
users.upsert(user_id)
AuthEventsRepository(conn).insert(
user_id,
"oidc_login",
ip=request.remote_addr,
user_agent=request.headers.get("User-Agent"),
metadata={"email": effective.get("email"), "groups": groups or None},
)
except Exception:
logger.error(
"OIDC provisioning/audit failed for %s%s",
user_id,
"" if disabled else "; continuing login",
exc_info=True,
)
return not disabled
def oidc_login():
"""Start the Authorization Code + PKCE flow with a 302 to the IdP."""
redis = get_redis_instance()
@@ -110,25 +247,50 @@ def oidc_callback():
try:
tokens = provider.exchange_code(code, stored["code_verifier"], _redirect_uri())
claims = provider.validate_id_token(tokens["id_token"], stored["nonce"])
effective = _effective_claims(tokens, claims)
except (provider.OIDCError, KeyError):
logger.error("OIDC callback failed", exc_info=True)
return _frontend_redirect("oidc_error=auth_failed")
user_id = claims.get(settings.OIDC_USER_ID_CLAIM)
user_id = effective.get(settings.OIDC_USER_ID_CLAIM)
if not user_id:
logger.error("OIDC id_token missing user id claim %r", settings.OIDC_USER_ID_CLAIM)
return _frontend_redirect("oidc_error=missing_claim")
user_id = str(user_id)
now = int(time.time())
payload = {
"sub": str(user_id),
"iat": now,
"exp": now + settings.OIDC_SESSION_LIFETIME_SECONDS,
}
for claim in ("email", "name"):
if claims.get(claim):
payload[claim] = claims[claim]
session_token = jwt.encode(payload, settings.JWT_SECRET_KEY, algorithm="HS256")
allowed = _allowed_groups()
groups = _claim_groups(effective)
if allowed and not set(groups) & set(allowed):
logger.info("OIDC login denied for %s: groups %s not in allowlist", user_id, groups)
_record_login_denied(user_id, {"reason": "not_authorized", "groups": groups})
return _frontend_redirect("oidc_error=not_authorized")
if not _gate_and_audit_login(user_id, effective, groups):
return _frontend_redirect("oidc_error=account_disabled")
# A fresh IdP-blessed authentication supersedes session-level revocations
# (back-channel logout denylists the IdP sub; without this, re-login would
# stay blocked until the denylist TTL ran out).
denylist.allow_user(str(user_id))
denylist.allow_idp_sub(str(claims["sub"]))
session_token, jti = _mint_session_token(
{
"sub": user_id,
"email": effective.get("email"),
"name": effective.get("name"),
"picture": effective.get("picture"),
"oidc_sub": claims["sub"],
"oidc_sid": claims.get("sid"),
}
)
refresh_token = tokens.get("refresh_token")
if refresh_token:
try:
redis.set(_refresh_key(jti), refresh_token, ex=settings.OIDC_SESSION_LIFETIME_SECONDS)
except Exception:
logger.warning("Failed to store OIDC refresh token", exc_info=True)
handoff = secrets.token_urlsafe(32)
redis.set(_handoff_key(handoff), session_token, ex=HANDOFF_TTL_SECONDS, nx=True)
@@ -151,6 +313,157 @@ def oidc_token():
return jsonify({"token": token})
def oidc_refresh():
"""Rotate the stored IdP refresh token and mint a fresh session JWT."""
decoded = handle_auth(request)
if (
not isinstance(decoded, dict)
or "error" in decoded
or not decoded.get("sub")
or not decoded.get("jti")
):
error = "invalid_token"
if isinstance(decoded, dict) and decoded.get("error") == "token_expired":
error = "token_expired"
return make_response(jsonify({"error": error}), 401)
if denylist.is_denied(decoded):
return make_response(jsonify({"error": "token_revoked"}), 401)
try:
with db_readonly() as conn:
row = UsersRepository(conn).get(str(decoded["sub"]))
except Exception:
logger.error("User lookup failed during OIDC refresh", exc_info=True)
row = None
if row is not None and row.get("active") is False:
return make_response(jsonify({"error": "account_disabled"}), 401)
redis = get_redis_instance()
if redis is None:
return make_response(jsonify({"error": "redis_unavailable"}), 503)
raw = redis.getdel(_refresh_key(str(decoded["jti"])))
if raw is None:
return make_response(jsonify({"error": "no_refresh_token"}), 404)
refresh_token = raw.decode("utf-8") if isinstance(raw, bytes) else str(raw)
try:
tokens = provider.refresh_grant(refresh_token)
except provider.OIDCError:
logger.warning("OIDC refresh grant failed", exc_info=True)
return make_response(jsonify({"error": "refresh_failed"}), 401)
identity = {
"sub": str(decoded["sub"]),
"email": decoded.get("email"),
"name": decoded.get("name"),
"picture": decoded.get("picture"),
"oidc_sub": decoded.get("oidc_sub"),
"oidc_sid": decoded.get("oidc_sid"),
}
id_token = tokens.get("id_token")
if id_token:
try:
claims = provider.validate_id_token(id_token, nonce=None)
effective = _effective_claims(tokens, claims)
except provider.OIDCError:
logger.warning("Refresh-issued id_token failed validation", exc_info=True)
return make_response(jsonify({"error": "refresh_failed"}), 401)
# Re-gate group membership on every renewal that carries fresh
# claims — otherwise removal from the allowlist would never bite
# while silent renewal keeps extending the session.
allowed = _allowed_groups()
groups = _claim_groups(effective)
if allowed and not (set(groups) & set(allowed)):
denied_user = str(effective.get(settings.OIDC_USER_ID_CLAIM) or decoded["sub"])
logger.info("OIDC refresh denied for %s: groups %s not allowed", denied_user, groups)
_record_login_denied(
denied_user,
{"reason": "not_authorized", "via": "refresh", "groups": groups},
)
return make_response(jsonify({"error": "not_authorized"}), 401)
user_id = effective.get(settings.OIDC_USER_ID_CLAIM)
if user_id:
identity["sub"] = str(user_id)
for claim in ("email", "name", "picture"):
if effective.get(claim):
identity[claim] = effective[claim]
identity["oidc_sub"] = effective.get("sub") or identity["oidc_sub"]
if effective.get("sid"):
identity["oidc_sid"] = effective["sid"]
new_token, new_jti = _mint_session_token(identity)
new_refresh = tokens.get("refresh_token") or refresh_token
try:
redis.set(_refresh_key(new_jti), new_refresh, ex=settings.OIDC_SESSION_LIFETIME_SECONDS)
except Exception:
logger.warning("Failed to store rotated OIDC refresh token", exc_info=True)
try:
with db_session() as conn:
AuthEventsRepository(conn).insert(
identity["sub"],
"oidc_refresh",
ip=request.remote_addr,
user_agent=request.headers.get("User-Agent"),
)
except Exception:
logger.warning("Failed to record oidc_refresh for %s", identity["sub"], exc_info=True)
return jsonify({"token": new_token})
def oidc_backchannel_logout():
"""Revoke sessions named by a signed IdP back-channel logout token."""
logout_token = request.form.get("logout_token")
if not logout_token:
body = request.get_json(silent=True)
if isinstance(body, dict):
logout_token = body.get("logout_token")
if not logout_token or not isinstance(logout_token, str):
return _no_store(jsonify({"error": "missing_logout_token"}), 400)
try:
claims = provider.validate_logout_token(logout_token)
except provider.OIDCError:
logger.warning("Rejected OIDC back-channel logout token", exc_info=True)
return _no_store(jsonify({"error": "invalid_logout_token"}), 400)
jti = claims.get("jti")
if jti:
redis = get_redis_instance()
if redis is not None:
try:
fresh = redis.set(_logout_jti_key(str(jti)), "1", ex=LOGOUT_JTI_TTL_SECONDS, nx=True)
except Exception:
logger.warning("Logout-token jti replay check failed; accepting token", exc_info=True)
fresh = True
if not fresh:
return _no_store(jsonify({"error": "invalid_logout_token"}), 400)
sub = claims.get("sub")
sid = claims.get("sid")
if sub:
denylist.deny_idp_sub(str(sub))
if sid:
denylist.deny_sid(str(sid))
user_id = str(sub) if sub else f"sid:{sid}"
try:
with db_session() as conn:
AuthEventsRepository(conn).insert(
user_id,
"backchannel_logout",
ip=request.remote_addr,
user_agent=request.headers.get("User-Agent"),
metadata={"sid": str(sid)} if sid else None,
)
except Exception:
logger.warning("Failed to record backchannel_logout for %s", user_id, exc_info=True)
return _no_store("", 200)
def oidc_logout():
"""Redirect to the IdP end-session endpoint, falling back to the frontend."""
frontend = _frontend_url()
@@ -178,6 +491,15 @@ def register(bp: Blueprint) -> None:
bp.add_url_rule(
"/api/auth/oidc/token", view_func=oidc_token, methods=["POST"], endpoint="token"
)
bp.add_url_rule(
"/api/auth/oidc/refresh", view_func=oidc_refresh, methods=["POST"], endpoint="refresh"
)
bp.add_url_rule(
"/api/auth/oidc/backchannel-logout",
view_func=oidc_backchannel_logout,
methods=["POST"],
endpoint="backchannel_logout",
)
bp.add_url_rule(
"/api/auth/oidc/logout", view_func=oidc_logout, methods=["GET"], endpoint="logout"
)
+11
View File
@@ -0,0 +1,11 @@
"""Flask blueprint for SCIM 2.0 user provisioning (/scim/v2)."""
from __future__ import annotations
from flask import Blueprint
from .routes import register as register_routes
scim_bp = Blueprint("scim", __name__)
register_routes(scim_bp)
+432
View File
@@ -0,0 +1,432 @@
"""SCIM 2.0 user-provisioning endpoints (RFC 7643/7644 subset for IdP clients).
IdPs (Okta, Authentik, Entra) push user lifecycle into DocsGPT through
``/scim/v2``: create users ahead of first login and deactivate them on
offboarding. Deactivation also revokes live sessions via the Redis
denylist; login refuses inactive users elsewhere. Only ``userName`` and
``active`` are honored — everything else IdPs send is ignored.
"""
from __future__ import annotations
import hmac
import json
import logging
import re
from typing import Any, Optional
from flask import Blueprint, Response, request
from sqlalchemy import Connection
from application.api.oidc.denylist import allow_user, deny_user
from application.core.settings import settings
from application.storage.db.repositories.auth_events import AuthEventsRepository
from application.storage.db.repositories.users import UsersRepository
from application.storage.db.session import db_readonly, db_session
logger = logging.getLogger(__name__)
_SCIM_MEDIA_TYPE = "application/scim+json"
_ERROR_URN = "urn:ietf:params:scim:api:messages:2.0:Error"
_LIST_RESPONSE_URN = "urn:ietf:params:scim:api:messages:2.0:ListResponse"
_USER_URN = "urn:ietf:params:scim:schemas:core:2.0:User"
_DEFAULT_COUNT = 100
_MAX_COUNT = 200
# The only filter IdPs need for provisioning: exact userName lookup.
_USERNAME_EQ_FILTER = re.compile(r'^\s*userName\s+eq\s+"([^"]*)"\s*$', re.IGNORECASE)
_SERVICE_PROVIDER_CONFIG = {
"schemas": ["urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig"],
"patch": {"supported": True},
"bulk": {"supported": False},
"filter": {"supported": True, "maxResults": _MAX_COUNT},
"changePassword": {"supported": False},
"sort": {"supported": False},
"etag": {"supported": False},
"authenticationSchemes": [
{
"type": "oauthbearertoken",
"name": "Bearer Token",
"description": "Authorization header carrying the configured SCIM bearer token",
}
],
}
_USER_RESOURCE_TYPE = {
"schemas": ["urn:ietf:params:scim:schemas:core:2.0:ResourceType"],
"id": "User",
"name": "User",
"endpoint": "/scim/v2/Users",
"schema": _USER_URN,
"meta": {"resourceType": "ResourceType", "location": "/scim/v2/ResourceTypes/User"},
}
_USER_SCHEMA = {
"id": _USER_URN,
"name": "User",
"description": "DocsGPT user account",
"attributes": [
{
"name": "userName",
"type": "string",
"multiValued": False,
"required": True,
"caseExact": True,
"mutability": "immutable",
"returned": "default",
"uniqueness": "server",
},
{
"name": "active",
"type": "boolean",
"multiValued": False,
"required": False,
"mutability": "readWrite",
"returned": "default",
},
],
"meta": {"resourceType": "Schema", "location": f"/scim/v2/Schemas/{_USER_URN}"},
}
# ----------------------------------------------------------------------
# Response helpers
# ----------------------------------------------------------------------
def _scim_response(payload: Optional[dict], status: int, headers: Optional[dict] = None) -> Response:
"""Build a response with the SCIM media type; ``None`` payload means empty body."""
body = "" if payload is None else json.dumps(payload)
response = Response(body, status=status, mimetype=_SCIM_MEDIA_TYPE)
for key, value in (headers or {}).items():
response.headers[key] = value
return response
def _scim_error(status: int, detail: str, scim_type: Optional[str] = None) -> Response:
"""Build an RFC 7644 error response."""
payload: dict = {"schemas": [_ERROR_URN], "status": str(status), "detail": detail}
if scim_type:
payload["scimType"] = scim_type
return _scim_response(payload, status)
def _static_list_response(resources: list) -> dict:
"""Wrap fixed resources in a SCIM ListResponse."""
return {
"schemas": [_LIST_RESPONSE_URN],
"totalResults": len(resources),
"startIndex": 1,
"itemsPerPage": len(resources),
"Resources": resources,
}
def _iso(value: Any) -> Optional[str]:
"""Return an ISO-8601 string (repository rows may carry str or datetime)."""
if value is None:
return None
return value if isinstance(value, str) else value.isoformat()
def _serialize_user(row: dict) -> dict:
"""Map a ``users`` row to a SCIM User resource."""
pk = str(row["id"])
user_name = row["user_id"]
resource = {
"schemas": [_USER_URN],
"id": pk,
"userName": user_name,
"active": bool(row["active"]),
}
if "@" in user_name:
resource["emails"] = [{"value": user_name, "primary": True}]
resource["meta"] = {
"resourceType": "User",
"created": _iso(row.get("created_at")),
"lastModified": _iso(row.get("updated_at")),
"location": f"/scim/v2/Users/{pk}",
}
return resource
# ----------------------------------------------------------------------
# Request parsing helpers
# ----------------------------------------------------------------------
def _coerce_active(value: Any) -> Optional[bool]:
"""Coerce a SCIM ``active`` value to bool; ``None`` when invalid (Okta sends strings)."""
if isinstance(value, bool):
return value
if isinstance(value, str) and value.lower() in ("true", "false"):
return value.lower() == "true"
return None
def _int_arg(name: str, default: int) -> int:
"""Read an integer query parameter, falling back to ``default``."""
raw = request.args.get(name)
if raw is None:
return default
try:
return int(raw)
except ValueError:
return default
def _parse_filter(raw: Optional[str]) -> tuple[Optional[str], Optional[Response]]:
"""Parse the ``filter`` query param; only ``userName eq "value"`` is supported."""
if raw is None or not raw.strip():
return None, None
match = _USERNAME_EQ_FILTER.match(raw)
if match is None:
return None, _scim_error(400, 'Only the filter userName eq "value" is supported', "invalidFilter")
return match.group(1), None
# ----------------------------------------------------------------------
# Side effects
# ----------------------------------------------------------------------
def _audit(conn: Connection, user_id: str, event: str) -> None:
"""Best-effort audit insert in a savepoint; failure never fails the request."""
try:
with conn.begin_nested():
AuthEventsRepository(conn).insert(user_id, event, metadata={"via": "scim"})
except Exception:
logger.error("SCIM audit insert failed for user %s event %s", user_id, event, exc_info=True)
def _apply_active(conn: Connection, row: dict, desired: bool) -> dict:
"""Apply an ``active`` transition; side effects run only when the value changes."""
if bool(row["active"]) == desired:
return row
updated = UsersRepository(conn).set_active(str(row["id"]), desired) or row
user_id = row["user_id"]
if desired:
allow_user(user_id)
_audit(conn, user_id, "scim_reactivated")
else:
deny_user(user_id)
_audit(conn, user_id, "scim_deactivated")
return updated
# ----------------------------------------------------------------------
# Bearer-token gate
# ----------------------------------------------------------------------
def _enforce_scim_auth() -> Optional[Response]:
"""Gate every SCIM request on SCIM_ENABLED and the shared bearer token."""
if not settings.SCIM_ENABLED:
return _scim_error(404, "SCIM provisioning is not enabled")
token = settings.SCIM_TOKEN
if not token:
logger.error("SCIM is enabled but SCIM_TOKEN is not configured — rejecting request")
return _scim_error(503, "SCIM is enabled but no SCIM_TOKEN is configured")
scheme, _, presented = request.headers.get("Authorization", "").partition(" ")
if scheme.lower() != "bearer" or not hmac.compare_digest(
presented.strip().encode("utf-8"), token.encode("utf-8")
):
return _scim_error(401, "Invalid or missing bearer token")
return None
# ----------------------------------------------------------------------
# Discovery endpoints
# ----------------------------------------------------------------------
def service_provider_config():
"""Static service-provider capabilities document."""
return _scim_response(_SERVICE_PROVIDER_CONFIG, 200)
def resource_types():
"""Advertise the User resource type."""
return _scim_response(_static_list_response([_USER_RESOURCE_TYPE]), 200)
def schemas():
"""Advertise the User schema."""
return _scim_response(_static_list_response([_USER_SCHEMA]), 200)
# ----------------------------------------------------------------------
# Users
# ----------------------------------------------------------------------
def list_users():
"""List users with optional exact userName filter and 1-based pagination."""
user_name, error = _parse_filter(request.args.get("filter"))
if error is not None:
return error
start_index = max(1, _int_arg("startIndex", 1))
count = min(max(0, _int_arg("count", _DEFAULT_COUNT)), _MAX_COUNT)
with db_readonly() as conn:
total, rows = UsersRepository(conn).list_paginated(user_name, start_index - 1, count)
return _scim_response(
{
"schemas": [_LIST_RESPONSE_URN],
"totalResults": total,
"startIndex": start_index,
"itemsPerPage": len(rows),
"Resources": [_serialize_user(row) for row in rows],
},
200,
)
def create_user():
"""Create a user from ``userName`` (+ optional ``active``); 409 on duplicates."""
body = request.get_json(force=True, silent=True)
if not isinstance(body, dict):
return _scim_error(400, "Request body must be a JSON object", "invalidValue")
user_name = body.get("userName")
if not isinstance(user_name, str) or not user_name.strip():
return _scim_error(400, "userName is required", "invalidValue")
active = _coerce_active(body.get("active", True))
if active is None:
return _scim_error(400, "active must be a boolean", "invalidValue")
with db_session() as conn:
row = UsersRepository(conn).create(user_name, active=active)
if row is None:
return _scim_error(409, f"User {user_name} already exists", "uniqueness")
_audit(conn, user_name, "scim_created")
resource = _serialize_user(row)
return _scim_response(resource, 201, headers={"Location": resource["meta"]["location"]})
def get_user(user_pk: str):
"""Fetch one user by primary key."""
with db_readonly() as conn:
row = UsersRepository(conn).get_by_pk(user_pk)
if not row:
return _scim_error(404, "User not found")
return _scim_response(_serialize_user(row), 200)
def replace_user(user_pk: str):
"""Full replace; only ``active`` is honored and ``userName`` is immutable."""
body = request.get_json(force=True, silent=True)
if not isinstance(body, dict):
return _scim_error(400, "Request body must be a JSON object", "invalidValue")
desired: Optional[bool] = None
if "active" in body:
desired = _coerce_active(body["active"])
if desired is None:
return _scim_error(400, "active must be a boolean", "invalidValue")
with db_session() as conn:
row = UsersRepository(conn).get_by_pk(user_pk)
if not row:
return _scim_error(404, "User not found")
if "userName" in body and body["userName"] != row["user_id"]:
return _scim_error(400, "userName is immutable", "mutability")
if desired is not None:
row = _apply_active(conn, row, desired)
return _scim_response(_serialize_user(row), 200)
def patch_user(user_pk: str):
"""Apply PatchOp replace operations targeting ``active``."""
body = request.get_json(force=True, silent=True)
if not isinstance(body, dict):
return _scim_error(400, "Request body must be a JSON object", "invalidValue")
operations = body.get("Operations")
if not isinstance(operations, list) or not operations:
return _scim_error(400, "PatchOp body with Operations is required", "invalidValue")
desired: Optional[bool] = None
for operation in operations:
if not isinstance(operation, dict) or str(operation.get("op", "")).strip().lower() != "replace":
return _scim_error(400, "Only the replace operation is supported", "invalidPath")
path = str(operation.get("path") or "").strip()
if not path:
value = operation.get("value")
if not isinstance(value, dict):
return _scim_error(400, "replace without path requires an object value", "invalidValue")
if "active" not in value:
continue # Only "active" is honored; other attributes are ignored.
candidate = value["active"]
elif path.lower() == "active":
candidate = operation.get("value")
else:
return _scim_error(400, f"Unsupported path: {path}", "invalidPath")
coerced = _coerce_active(candidate)
if coerced is None:
return _scim_error(400, "active must be a boolean", "invalidValue")
desired = coerced
with db_session() as conn:
row = UsersRepository(conn).get_by_pk(user_pk)
if not row:
return _scim_error(404, "User not found")
if desired is not None:
row = _apply_active(conn, row, desired)
return _scim_response(_serialize_user(row), 200)
def delete_user(user_pk: str):
"""Soft delete: deactivate the user and revoke live sessions."""
with db_session() as conn:
row = UsersRepository(conn).get_by_pk(user_pk)
if not row:
return _scim_error(404, "User not found")
_apply_active(conn, row, False)
return _scim_response(None, 204)
# ----------------------------------------------------------------------
# Groups (not supported — answered so IdP probes don't hard-fail)
# ----------------------------------------------------------------------
def list_groups():
"""Groups are not provisioned; always an empty ListResponse."""
return _scim_response(_static_list_response([]), 200)
def create_group():
"""Group creation is not supported."""
return _scim_error(501, "Group provisioning is not supported")
def group_detail(group_id: str):
"""Individual groups never exist; mutations are unsupported."""
if request.method == "GET":
return _scim_error(404, "Group not found")
return _scim_error(501, "Group provisioning is not supported")
def register(bp: Blueprint) -> None:
"""Attach the SCIM routes and bearer-token gate to ``bp``."""
bp.before_request(_enforce_scim_auth)
bp.add_url_rule(
"/scim/v2/ServiceProviderConfig", view_func=service_provider_config, methods=["GET"],
endpoint="service_provider_config",
)
bp.add_url_rule(
"/scim/v2/ResourceTypes", view_func=resource_types, methods=["GET"], endpoint="resource_types",
)
bp.add_url_rule(
"/scim/v2/Schemas", view_func=schemas, methods=["GET"], endpoint="schemas",
)
bp.add_url_rule(
"/scim/v2/Users", view_func=list_users, methods=["GET"], endpoint="list_users",
)
bp.add_url_rule(
"/scim/v2/Users", view_func=create_user, methods=["POST"], endpoint="create_user",
)
bp.add_url_rule(
"/scim/v2/Users/<user_pk>", view_func=get_user, methods=["GET"], endpoint="get_user",
)
bp.add_url_rule(
"/scim/v2/Users/<user_pk>", view_func=replace_user, methods=["PUT"], endpoint="replace_user",
)
bp.add_url_rule(
"/scim/v2/Users/<user_pk>", view_func=patch_user, methods=["PATCH"], endpoint="patch_user",
)
bp.add_url_rule(
"/scim/v2/Users/<user_pk>", view_func=delete_user, methods=["DELETE"], endpoint="delete_user",
)
bp.add_url_rule(
"/scim/v2/Groups", view_func=list_groups, methods=["GET"], endpoint="list_groups",
)
bp.add_url_rule(
"/scim/v2/Groups", view_func=create_group, methods=["POST"], endpoint="create_group",
)
bp.add_url_rule(
"/scim/v2/Groups/<group_id>", view_func=group_detail,
methods=["GET", "PUT", "PATCH", "DELETE"], endpoint="group_detail",
)
+20
View File
@@ -20,6 +20,8 @@ from application.api.devices import devices_bp # noqa: E402
from application.api.events.routes import events # noqa: E402
from application.api.internal.routes import internal # noqa: E402
from application.api.oidc import oidc_bp # noqa: E402
from application.api.oidc.denylist import is_denied as oidc_session_denied # noqa: E402
from application.api.scim import scim_bp # noqa: E402
from application.api.user.routes import user # noqa: E402
from application.api.connector.routes import connector # noqa: E402
from application.api.v1 import v1_bp # noqa: E402
@@ -63,6 +65,7 @@ app.register_blueprint(internal)
app.register_blueprint(connector)
app.register_blueprint(devices_bp)
app.register_blueprint(oidc_bp)
app.register_blueprint(scim_bp)
app.register_blueprint(v1_bp)
app.config.update(
UPLOAD_FOLDER="inputs",
@@ -123,6 +126,7 @@ def get_config():
response["oidc"] = {
"login_path": "/api/auth/oidc/login",
"logout_path": "/api/auth/oidc/logout",
"provider_name": settings.OIDC_PROVIDER_NAME,
}
return jsonify(response)
@@ -216,11 +220,27 @@ def authenticate_request():
if request.path.startswith("/api/auth/oidc/"):
request.decoded_token = None
return None
# SCIM provisioning authenticates with its own bearer token (SCIM_TOKEN),
# validated inside the blueprint.
if request.path.startswith("/scim/"):
request.decoded_token = None
return None
decoded_token = handle_auth(request)
if not decoded_token:
request.decoded_token = None
elif "error" in decoded_token:
return jsonify(decoded_token), 401
elif settings.AUTH_TYPE == "oidc" and oidc_session_denied(decoded_token):
# Back-channel logout / SCIM deactivation revoked this session.
return (
jsonify(
{
"message": "Authentication error: session revoked",
"error": "token_revoked",
}
),
401,
)
else:
request.decoded_token = decoded_token
+7
View File
@@ -28,6 +28,13 @@ class Settings(BaseSettings):
OIDC_FRONTEND_URL: Optional[str] = None # browser-facing app origin, e.g. http://localhost:5173
OIDC_REDIRECT_URI: Optional[str] = None # override; default <request host>/api/auth/oidc/callback
OIDC_SESSION_LIFETIME_SECONDS: int = 28800 # minted session JWT lifetime (8h)
OIDC_PROVIDER_NAME: Optional[str] = None # sign-in button label, e.g. "Acme SSO"
OIDC_ALLOWED_GROUPS: Optional[str] = None # comma-separated allowlist; unset = any authenticated user
OIDC_GROUPS_CLAIM: str = "groups" # ID-token/userinfo claim carrying group membership
# SCIM 2.0 provisioning (IdP-driven user create/deactivate at /scim/v2)
SCIM_ENABLED: bool = False
SCIM_TOKEN: Optional[str] = None # bearer token for IdP SCIM clients (required when enabled)
LLM_PROVIDER: str = "docsgpt"
LLM_NAME: Optional[str] = None # if LLM_PROVIDER is openai, LLM_NAME can be gpt-4 or gpt-3.5-turbo
+13
View File
@@ -49,10 +49,23 @@ users_table = Table(
server_default='{"pinned": [], "shared_with_me": []}',
),
Column("tool_preferences", JSONB, nullable=False, server_default="{}"),
Column("active", Boolean, nullable=False, server_default="true"),
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
)
auth_events_table = Table(
"auth_events",
metadata,
Column("id", UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid()),
Column("user_id", Text, nullable=False),
Column("event", Text, nullable=False),
Column("ip", Text),
Column("user_agent", Text),
Column("metadata", JSONB, nullable=False, server_default="{}"),
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
)
prompts_table = Table(
"prompts",
metadata,
@@ -0,0 +1,59 @@
"""Repository for the ``auth_events`` audit table."""
from __future__ import annotations
import json
from typing import Optional
from sqlalchemy import Connection, text
from application.storage.db.base_repository import row_to_dict
class AuthEventsRepository:
"""Append-only audit trail of login / logout / provisioning events."""
def __init__(self, conn: Connection) -> None:
self._conn = conn
def insert(
self,
user_id: str,
event: str,
ip: Optional[str] = None,
user_agent: Optional[str] = None,
metadata: Optional[dict] = None,
) -> dict:
"""Record one auth event and return the inserted row."""
result = self._conn.execute(
text(
"""
INSERT INTO auth_events (user_id, event, ip, user_agent, metadata)
VALUES (:user_id, :event, :ip, :user_agent, CAST(:metadata AS jsonb))
RETURNING *
"""
),
{
"user_id": user_id,
"event": event,
"ip": ip,
"user_agent": user_agent,
"metadata": json.dumps(metadata or {}),
},
)
return row_to_dict(result.fetchone())
def list_recent(self, user_id: str, limit: int = 50) -> list[dict]:
"""Return the newest events for ``user_id``, newest first."""
result = self._conn.execute(
text(
"""
SELECT * FROM auth_events
WHERE user_id = :user_id
ORDER BY created_at DESC
LIMIT :limit
"""
),
{"user_id": user_id, "limit": limit},
)
return [row_to_dict(row) for row in result.fetchall()]
@@ -24,6 +24,7 @@ rollback-per-test connection (tests).
from __future__ import annotations
from typing import Iterable, Optional
from uuid import UUID
from sqlalchemy import Connection, text
@@ -33,6 +34,14 @@ from application.storage.db.base_repository import row_to_dict
_DEFAULT_PREFERENCES = '{"pinned": [], "shared_with_me": []}'
def _canonical_uuid(value: str) -> Optional[str]:
"""Return the canonical UUID string for ``value``, or ``None`` when malformed."""
try:
return str(UUID(str(value)))
except (TypeError, ValueError):
return None
class UsersRepository:
"""Postgres-backed replacement for Mongo ``users_collection`` writes/reads."""
@@ -236,6 +245,81 @@ class UsersRepository:
{"user_id": user_id, "tool_name": tool_name},
)
# ------------------------------------------------------------------
# SCIM provisioning
# ------------------------------------------------------------------
def create(self, user_id: str, active: bool = True) -> Optional[dict]:
"""Insert a new user row; ``None`` means ``user_id`` already exists."""
result = self._conn.execute(
text(
"""
INSERT INTO users (user_id, agent_preferences, active)
VALUES (:user_id, CAST(:default_prefs AS jsonb), :active)
ON CONFLICT (user_id) DO NOTHING
RETURNING *
"""
),
{"user_id": user_id, "default_prefs": _DEFAULT_PREFERENCES, "active": active},
)
row = result.fetchone()
return row_to_dict(row) if row is not None else None
def get_by_pk(self, pk: str) -> Optional[dict]:
"""Return the user row by primary-key ``id``, or ``None`` (including malformed UUIDs)."""
canonical = _canonical_uuid(pk)
if canonical is None:
return None
result = self._conn.execute(
text("SELECT * FROM users WHERE id = CAST(:pk AS uuid)"),
{"pk": canonical},
)
row = result.fetchone()
return row_to_dict(row) if row is not None else None
def set_active(self, pk: str, active: bool) -> Optional[dict]:
"""Set ``active`` on the row with primary-key ``id`` and return the updated row."""
canonical = _canonical_uuid(pk)
if canonical is None:
return None
result = self._conn.execute(
text(
"""
UPDATE users
SET active = :active, updated_at = now()
WHERE id = CAST(:pk AS uuid)
RETURNING *
"""
),
{"pk": canonical, "active": active},
)
row = result.fetchone()
return row_to_dict(row) if row is not None else None
def list_paginated(
self, user_name: Optional[str], offset: int, limit: int
) -> tuple[int, list[dict]]:
"""Return ``(total, page)`` ordered by ``created_at, id``; optional exact ``user_id`` filter."""
where = ""
filter_params: dict = {}
if user_name is not None:
where = "WHERE user_id = :user_name"
filter_params = {"user_name": user_name}
total = self._conn.execute(
text(f"SELECT count(*) FROM users {where}"), filter_params
).scalar_one()
result = self._conn.execute(
text(
f"""
SELECT * FROM users
{where}
ORDER BY created_at, id
LIMIT :limit OFFSET :offset
"""
),
{**filter_params, "limit": limit, "offset": offset},
)
return int(total), [row_to_dict(row) for row in result.fetchall()]
# ------------------------------------------------------------------
# Private helpers
# ------------------------------------------------------------------
@@ -277,6 +277,7 @@ JWT_SECRET_KEY=your_secret_key_here
- The frontend redirects users to your identity provider to sign in (OAuth2 Authorization Code + PKCE).
- After a successful sign-in, DocsGPT issues its own session JWT; API requests carry it in the `Authorization` header like the other modes.
- Stable per-user identities come from the provider — see the full setup guide: [SSO with OIDC](/Deploying/OIDC-SSO).
- The same guide covers the optional access controls: group allowlists, silent session renewal, back-channel logout, SCIM provisioning, and login auditing.
#### Security Notes
+141 -11
View File
@@ -1,12 +1,14 @@
---
title: SSO with OIDC
description: Sign users into DocsGPT through any OpenID Connect identity provider — Authentik, Keycloak, Okta, and others.
description: Sign users into DocsGPT through any OpenID Connect identity provider (Authentik, Keycloak, Okta, ...) — with group allowlists, silent session renewal, back-channel logout, and SCIM provisioning.
---
# SSO with OIDC
Setting `AUTH_TYPE=oidc` makes DocsGPT delegate sign-in to an external OpenID Connect identity provider (IdP). Any spec-compliant IdP with a discovery document works; this guide uses [Authentik](https://goauthentik.io/) as the reference provider and includes a short note for Keycloak.
Beyond basic sign-in, this page covers the optional access controls: [group allowlists](#restricting-sign-in-by-group), [silent session renewal](#silent-session-renewal), [back-channel logout](#back-channel-logout), [SCIM user provisioning](#scim-user-provisioning), and [login auditing](#login-auditing).
## How the flow works
1. A user opens DocsGPT without a session. The frontend redirects the browser to `GET /api/auth/oidc/login` on the DocsGPT API.
@@ -17,7 +19,14 @@ Setting `AUTH_TYPE=oidc` makes DocsGPT delegate sign-in to an external OpenID Co
The user's identity (`sub` claim by default) becomes the DocsGPT `user_id`, so every user gets their own conversations, sources, agents, and settings.
> Redis must be reachable by the API (it stores the short-lived login state and handoff codes). Redis is already a required DocsGPT dependency, so no extra infrastructure is needed.
Sessions last `OIDC_SESSION_LIFETIME_SECONDS` (8 hours by default) and renew without interrupting the user — see [Silent session renewal](#silent-session-renewal).
> Redis must be reachable by the API — it stores the short-lived login state, handoff codes, server-side refresh tokens, and the session revocation denylist. Redis is already a required DocsGPT dependency, so no extra infrastructure is needed.
### IdP compatibility notes
- **Token-endpoint authentication** follows the IdP's discovery document (`token_endpoint_auth_methods_supported`): `client_secret_post` when the IdP advertises it, otherwise HTTP Basic (the RFC default). Okta's default web-app configuration works without extra toggles.
- **Userinfo fallback**: when the ID token lacks the user-id claim (`OIDC_USER_ID_CLAIM`) — or the groups claim while a group allowlist is configured — the backend fetches the IdP's userinfo endpoint and merges the missing claims. ID-token values win on conflict, and the userinfo `sub` must match the ID token's.
## Settings reference
@@ -28,12 +37,17 @@ The user's identity (`sub` claim by default) becomes the DocsGPT `user_id`, so e
| `OIDC_CLIENT_ID` | yes | — | Client ID registered at the IdP. |
| `OIDC_FRONTEND_URL` | yes | — | Browser-facing URL of the DocsGPT frontend (where users land after login/logout), e.g. `https://docsgpt.example.com`. |
| `OIDC_CLIENT_SECRET` | no | — | Set when the IdP client is *confidential*. PKCE is always used, so *public* clients work without a secret. |
| `OIDC_SCOPES` | no | `openid profile email` | Scopes requested at the IdP. |
| `OIDC_USER_ID_CLAIM` | no | `sub` | ID-token claim used as the DocsGPT user id. Set to `email` or `preferred_username` for human-readable ids. |
| `OIDC_SCOPES` | no | `openid profile email` | Scopes requested at the IdP. Add `offline_access` when your IdP requires it for refresh tokens (Authentik does). |
| `OIDC_USER_ID_CLAIM` | no | `sub` | ID-token claim used as the DocsGPT user id. Set to `email` or `preferred_username` for human-readable ids; use `email` when provisioning over [SCIM](#scim-user-provisioning). |
| `OIDC_REDIRECT_URI` | no | derived | Full callback URL registered at the IdP. Defaults to `<request host>/api/auth/oidc/callback`; set it explicitly when the API runs behind a reverse proxy. |
| `OIDC_SESSION_LIFETIME_SECONDS` | no | `28800` (8h) | Lifetime of the DocsGPT session JWT. After expiry the user is silently redirected through the IdP again. |
| `OIDC_SESSION_LIFETIME_SECONDS` | no | `28800` (8h) | Lifetime of the DocsGPT session JWT. Sessions renew before expiry — see [Silent session renewal](#silent-session-renewal). |
| `OIDC_PROVIDER_NAME` | no | — | Display name on the sign-in button: `Acme SSO` renders "Sign in with Acme SSO". Unset, the button shows a generic "SSO". |
| `OIDC_ALLOWED_GROUPS` | no | — | Comma-separated group allowlist. Unset, any authenticated IdP user may sign in — see [Restricting sign-in by group](#restricting-sign-in-by-group). |
| `OIDC_GROUPS_CLAIM` | no | `groups` | ID-token/userinfo claim carrying the user's group membership. |
| `JWT_SECRET_KEY` | recommended | auto-generated | Signs DocsGPT session tokens. Set it explicitly in production — required when running multiple API replicas. |
`SCIM_ENABLED` and `SCIM_TOKEN` are listed in the [SCIM section](#scim-user-provisioning).
## Setting up with Authentik
1. **Create a provider.** In the Authentik admin UI go to **Applications → Providers → Create** and pick **OAuth2/OpenID Provider**:
@@ -57,9 +71,11 @@ The user's identity (`sub` claim by default) becomes the DocsGPT `user_id`, so e
JWT_SECRET_KEY=<long random string>
```
> Planning to use [silent session renewal](#silent-session-renewal)? Authentik only issues refresh tokens when the `offline_access` scope is requested — set `OIDC_SCOPES=openid profile email offline_access`.
### Which claim becomes the user id?
Authentik's provider setting **Subject mode** controls what lands in the `sub` claim (the default is a hashed user ID — stable but opaque). If you'd rather key DocsGPT users on something readable, either change Subject mode (e.g. *based on username*) or leave Authentik alone and set `OIDC_USER_ID_CLAIM=email` in DocsGPT. Pick one strategy before going live: changing it later gives existing users fresh, empty accounts.
Authentik's provider setting **Subject mode** controls what lands in the `sub` claim (the default is a hashed user ID — stable but opaque). If you'd rather key DocsGPT users on something readable, either change Subject mode (e.g. *based on username*) or leave Authentik alone and set `OIDC_USER_ID_CLAIM=email` in DocsGPT. Pick one strategy before going live: changing it later gives existing users fresh, empty accounts. If you plan to provision users over [SCIM](#scim-user-provisioning), use `OIDC_USER_ID_CLAIM=email` — SCIM matches users by `userName`, which IdPs typically send as the email.
## Keycloak (and other IdPs)
@@ -72,12 +88,126 @@ OIDC_CLIENT_ID=<client id>
Create the client with *Standard flow* enabled and PKCE method `S256`; register the same `/api/auth/oidc/callback` redirect URI.
The feature sections below carry their own per-IdP notes — group claims, refresh tokens, back-channel logout, and SCIM each need one IdP-side setting.
## Restricting sign-in by group
By default any user who can authenticate at the IdP may use DocsGPT. To restrict access to specific IdP groups:
```env
OIDC_ALLOWED_GROUPS=docsgpt-users,platform-admins
# OIDC_GROUPS_CLAIM=groups # only if your IdP uses a different claim name
```
At login the backend reads the `OIDC_GROUPS_CLAIM` claim (default `groups`) from the ID token, falling back to the userinfo endpoint when the claim is absent. A user whose groups share no entry with the allowlist is rejected with a clean "not authorized" screen (`oidc_error=not_authorized`), and the denial lands in the [audit log](#login-auditing).
Group changes take effect at the next sign-in **or** the next [silent renewal](#silent-session-renewal): whenever the IdP returns a fresh ID token during renewal, the allowlist is re-checked — so removing a user from the allowed group cuts off their session at the next renewal instead of whenever they happen to sign in again.
Getting groups into the token:
- **Authentik** includes group names in the `groups` claim through its default `profile` scope — no extra configuration needed.
- **Keycloak** does not emit groups by default. On the client, open **Client scopes → the client's dedicated scope → Add mapper → By configuration → Group Membership**, set the claim name to `groups`, and turn **Full group path** off so the claim carries plain names (`devs`) rather than paths (`/devs`).
## Silent session renewal
The DocsGPT session JWT lives for `OIDC_SESSION_LIFETIME_SECONDS` (default 8 hours). Sessions renew without user-visible interruptions, in one of two ways:
- **With a refresh token.** When the IdP issues one, the backend stores it server-side (in Redis — never in the browser) and the frontend calls `POST /api/auth/oidc/refresh` about 15 minutes before the session expires. The backend redeems the refresh token at the IdP, re-validates the fresh ID token (including the [group allowlist](#restricting-sign-in-by-group)), mints a new session JWT, and rotates the stored refresh token. The user notices nothing.
- **Without a refresh token.** The frontend lets the session run to expiry and then redirects through the IdP again. While the IdP session is still alive, this round-trip is also silent; the user only sees a sign-in page once the IdP session is gone too.
Getting a refresh token:
- **Keycloak** issues refresh tokens for the authorization-code flow by default — nothing to change.
- **Authentik** only issues refresh tokens when the `offline_access` scope is requested:
```env
OIDC_SCOPES=openid profile email offline_access
```
Revoking the user's consent or sessions at the IdP makes the next renewal fail, and the user must sign in again. For revocation that doesn't wait for the next renewal, configure [back-channel logout](#back-channel-logout).
## Back-channel logout
DocsGPT implements [OIDC Back-Channel Logout 1.0](https://openid.net/specs/openid-connect-backchannel-1_0.html). The IdP POSTs a signed `logout_token` to:
```
POST https://<your-docsgpt-api>/api/auth/oidc/backchannel-logout
```
DocsGPT validates the token (signature via JWKS, issuer, audience, replay protection) and immediately revokes the user's live sessions through a Redis denylist — revoked requests get `401` with `error: token_revoked`. Signing the user out at the IdP, or an admin revoking their sessions there, takes effect on their next DocsGPT request instead of at session expiry.
The endpoint is called server-to-server, so it must be reachable from the IdP (it is not a browser redirect).
- **Keycloak**: open the client → **Settings** and set **Backchannel logout URL** to `https://<your-docsgpt-api>/api/auth/oidc/backchannel-logout`.
- **Authentik** (2025.8.0 and later; marked Preview): on the OAuth2/OpenID provider set **Logout Method** to *Back-channel* and **Logout URI** to the same URL — see the [Authentik logout docs](https://docs.goauthentik.io/add-secure-apps/providers/oauth2/frontchannel_and_backchannel_logout/). Authentik sends the logout token when a user logs out, an admin deletes their session, the account is deactivated, or the session is revoked. On older Authentik versions back-channel logout is unavailable — revocation latency then falls back to the session lifetime, or use [SCIM deactivation](#scim-user-provisioning), which also revokes sessions instantly.
## SCIM user provisioning
DocsGPT exposes a [SCIM 2.0](https://datatracker.ietf.org/doc/html/rfc7644) endpoint so your IdP can drive the user lifecycle: create accounts ahead of first login and — more importantly — deactivate them on offboarding. Deactivating a user revokes their live sessions immediately and blocks future sign-ins (they see an "account disabled" screen); reactivating restores access.
| Setting | Required | Default | Description |
| --- | --- | --- | --- |
| `SCIM_ENABLED` | yes | `false` | Set to `true` to serve the `/scim/v2` endpoints. |
| `SCIM_TOKEN` | yes | — | Bearer token the IdP's SCIM client must present. Use a long random string. |
The base URL is `https://<your-docsgpt-api>/scim/v2`; every request must carry `Authorization: Bearer <SCIM_TOKEN>`.
### Match the SCIM userName to the OIDC user id
SCIM identifies users by `userName`, which DocsGPT matches against its user id — the value of `OIDC_USER_ID_CLAIM`. With the default `sub` claim, the `userName` your IdP sends (typically the email) would never line up with the opaque `sub` of the same user signing in, and DocsGPT would treat them as two unrelated accounts. **When using SCIM, set `OIDC_USER_ID_CLAIM=email` and have the IdP send the email as the SCIM `userName`.**
### What the endpoint supports
| Operation | Support |
| --- | --- |
| `GET /scim/v2/ServiceProviderConfig`, `/ResourceTypes`, `/Schemas` | Discovery documents. |
| `GET /scim/v2/Users` | List, with the exact filter `userName eq "..."` and `startIndex`/`count` pagination (1-based, max 200 per page). |
| `POST /scim/v2/Users` | Create; returns `409` when the `userName` already exists. |
| `GET /scim/v2/Users/<id>` | Read. |
| `PUT` / `PATCH /scim/v2/Users/<id>` | Activate/deactivate via the `active` attribute (Okta's string `"true"`/`"false"` values are accepted). `userName` is immutable; other attributes are ignored. |
| `DELETE /scim/v2/Users/<id>` | Soft delete — deactivates the account instead of removing data. |
| `/scim/v2/Groups` | Group provisioning is **not** supported: listing returns an empty result so IdP probes don't fail, and mutations return `501`. Use the [group allowlist](#restricting-sign-in-by-group) for group-based access control instead. |
### IdP setup pointers
- **Okta**: add SCIM provisioning to the app integration with **SCIM connector base URL** = `https://<your-docsgpt-api>/scim/v2` and authentication mode **HTTP Header** carrying the bearer token. Enable creating and deactivating users; skip group push.
- **Authentik**: create a **SCIM provider** with the same base URL and the token, and attach it to the application as a backchannel provider. Sync users only — leave group mappings out, since DocsGPT answers group provisioning with `501`.
## Login auditing
Authentication activity is recorded in `auth_events`, an append-only Postgres table carrying the user id, event name, IP address, user agent, a JSONB `metadata` column, and a timestamp:
| Event | Recorded when |
| --- | --- |
| `oidc_login` | A user signs in successfully. |
| `oidc_login_denied` | A sign-in is rejected — `metadata.reason` is `not_authorized` (group allowlist) or `account_disabled`. |
| `oidc_refresh` | A session is silently renewed. |
| `backchannel_logout` | The IdP revokes sessions via back-channel logout. |
| `scim_created` / `scim_deactivated` / `scim_reactivated` | SCIM lifecycle changes. |
There is no UI for these events yet — query the table directly:
```sql
SELECT created_at, event, user_id, ip, metadata
FROM auth_events
ORDER BY created_at DESC
LIMIT 50;
```
## Troubleshooting
- **Redirected back with `oidc_error=auth_failed`** — check the API logs. The most common causes:
- *Issuer mismatch*: `OIDC_ISSUER` must be the URL the discovery document itself reports as `issuer` (for Authentik this includes the application slug and trailing slash).
- *Clock skew*: ID-token validation allows 60 seconds of skew; sync clocks if the API host drifts more than that.
- **`oidc_error=missing_claim`** — the ID token doesn't contain `OIDC_USER_ID_CLAIM`. Make sure the matching scope is requested (`OIDC_SCOPES`) and the IdP actually emits the claim, or switch the setting back to `sub`.
When sign-in fails, the browser lands back on the frontend with an `#oidc_error=<code>` fragment and the sign-in screen shows a matching message:
| Code | Cause |
| --- | --- |
| `invalid_state` | The login attempt expired (the state is held for 10 minutes) or was replayed. Retrying the sign-in usually fixes it. |
| `auth_failed` | Token exchange or ID-token validation failed — check the API logs. Most common: `OIDC_ISSUER` doesn't match the issuer the discovery document reports (for Authentik this includes the application slug and trailing slash), or clock skew beyond the allowed 60 seconds. |
| `missing_claim` | Neither the ID token nor userinfo contains `OIDC_USER_ID_CLAIM`. Make sure the matching scope is requested (`OIDC_SCOPES`) and the IdP actually emits the claim, or switch the setting back to `sub`. |
| `not_authorized` | The user's groups don't intersect `OIDC_ALLOWED_GROUPS` — see [Restricting sign-in by group](#restricting-sign-in-by-group). |
| `account_disabled` | The account was deactivated via [SCIM](#scim-user-provisioning) or by an operator. Reactivate it over SCIM to restore access. |
Other issues:
- **IdP shows a redirect URI error** — the callback URL registered at the IdP must match exactly. Behind a reverse proxy, set `OIDC_REDIRECT_URI` to the public callback URL instead of relying on the derived default.
- **Revoked users can still access DocsGPT** — DocsGPT sessions outlive IdP revocation for up to `OIDC_SESSION_LIFETIME_SECONDS`. Lower it if you need tighter revocation latency.
- **Revoked users can still access DocsGPT** — without back-channel logout, sessions outlive IdP-side revocation until the next renewal or expiry. Configure [back-channel logout](#back-channel-logout) for instant revocation, deactivate the user over [SCIM](#scim-user-provisioning), or lower `OIDC_SESSION_LIFETIME_SECONDS`.
- **SCIM requests fail** — `404`: `SCIM_ENABLED` is not `true`. `503`: SCIM is enabled but `SCIM_TOKEN` is unset. `401`: the presented bearer token doesn't match `SCIM_TOKEN`.
- **Login endpoints return 503** — Redis is unreachable or the IdP discovery document can't be fetched from the API host.
+16 -4
View File
@@ -25,15 +25,27 @@ import ToolApprovalToast from './notifications/ToolApprovalToast';
function AuthWrapper({ children }: { children: React.ReactNode }) {
const { t } = useTranslation();
const { isAuthLoading, oidcFailed, retryOidcLogin } = useTokenAuth();
const {
isAuthLoading,
oidcFailed,
oidcErrorCode,
oidcProviderName,
retryOidcLogin,
} = useTokenAuth();
useDataInitializer(isAuthLoading);
if (oidcFailed) {
const message =
oidcErrorCode === 'not_authorized'
? t('auth.notAuthorized')
: oidcErrorCode === 'account_disabled'
? t('auth.accountDisabled')
: t('auth.signInToContinue');
return (
<div className="flex h-screen flex-col items-center justify-center gap-6">
<img src={DocsGPT3} alt="DocsGPT" className="size-14" />
<p className="text-foreground text-sm dark:text-white">
{t('auth.signInToContinue')}
<p className="text-foreground max-w-md px-6 text-center text-sm dark:text-white">
{message}
</p>
<Button
type="button"
@@ -41,7 +53,7 @@ function AuthWrapper({ children }: { children: React.ReactNode }) {
className="rounded-3xl px-5"
data-testid="oidc-signin"
>
{t('auth.signInWithSSO')}
{t('auth.signInWith', { provider: oidcProviderName || 'SSO' })}
</Button>
</div>
);
+44 -18
View File
@@ -79,8 +79,15 @@ export default function Navigation({ navOpen, setNavOpen }: NavigationProps) {
const selectedAgent = useSelector(selectSelectedAgent);
const { isMobile, isTablet } = useMediaQuery();
const { showTokenModal, handleTokenSubmit, authType, logout, userEmail } =
useTokenAuth();
const {
showTokenModal,
handleTokenSubmit,
authType,
logout,
userEmail,
userName,
userPicture,
} = useTokenAuth();
const [isDeletingConversation, setIsDeletingConversation] = useState(false);
const [uploadModalState, setUploadModalState] =
@@ -650,27 +657,46 @@ export default function Navigation({ navOpen, setNavOpen }: NavigationProps) {
</p>
</NavLink>
{authType === 'oidc' && (
<button
onClick={logout}
data-testid="oidc-signout"
className="hover:bg-sidebar-accent mx-4 my-auto flex cursor-pointer items-center gap-2.5 rounded-3xl py-1.5 pl-3"
>
<LogOut
className="text-muted-foreground size-5 shrink-0"
strokeWidth={1.75}
aria-label="Sign out"
/>
<span className="flex min-w-0 flex-col items-start">
<p className="text-foreground text-sm dark:text-white">
{t('auth.signOut')}
</p>
<div className="mx-4 my-auto flex items-center gap-2.5 py-0.5 pr-1 pl-3">
{userPicture ? (
<Avatar
src={userPicture}
alt={userName || userEmail || 'User avatar'}
className="size-6"
imgClassName="size-6 rounded-full object-cover"
/>
) : (
<Avatar className="size-6">
<span className="bg-sidebar-accent text-foreground flex size-6 items-center justify-center rounded-full text-xs font-medium dark:text-white">
{(userName || userEmail || '?').charAt(0).toUpperCase()}
</span>
</Avatar>
)}
<span className="flex min-w-0 flex-1 flex-col items-start">
{userName && (
<p className="text-foreground max-w-[160px] truncate text-sm dark:text-white">
{userName}
</p>
)}
{userEmail && (
<p className="text-muted-foreground max-w-[170px] truncate text-xs">
<p className="text-muted-foreground max-w-[160px] truncate text-xs">
{userEmail}
</p>
)}
</span>
</button>
<Button
type="button"
variant="ghost"
size="icon"
onClick={logout}
data-testid="oidc-signout"
aria-label={t('auth.signOut')}
title={t('auth.signOut')}
className="text-muted-foreground hover:text-foreground shrink-0 rounded-full"
>
<LogOut className="size-5" strokeWidth={1.75} />
</Button>
</div>
)}
</div>
<div className="text-foreground flex flex-col justify-end dark:text-white">
+1
View File
@@ -4,6 +4,7 @@ const endpoints = {
NEW_TOKEN: '/api/generate_token',
OIDC_LOGIN: '/api/auth/oidc/login',
OIDC_TOKEN: '/api/auth/oidc/token',
OIDC_REFRESH: '/api/auth/oidc/refresh',
OIDC_LOGOUT: '/api/auth/oidc/logout',
MODELS: '/api/models',
DOCS: '/api/sources',
+2
View File
@@ -11,6 +11,8 @@ const userService = {
// to interfere with redeeming the one-time OIDC handoff code.
exchangeOidcCode: (code: string): Promise<any> =>
apiClient.post(endpoints.USER.OIDC_TOKEN, { code }, null),
refreshOidcSession: (token: string | null): Promise<any> =>
apiClient.post(endpoints.USER.OIDC_REFRESH, {}, token),
getDocs: (token: string | null): Promise<any> =>
apiClient.get(`${endpoints.USER.DOCS}`, token),
getDocsWithPagination: (query: string, token: string | null): Promise<any> =>
+141 -9
View File
@@ -5,17 +5,44 @@ import { baseURL } from '../api/client';
import endpoints from '../api/endpoints';
import userService from '../api/services/userService';
import { selectToken, setToken } from '../preferences/preferenceSlice';
import { decodeJwtPayload, isJwtExpired } from '../utils/jwtUtils';
import {
decodeJwtPayload,
getJwtRemainingMs,
isJwtExpired,
} from '../utils/jwtUtils';
const OIDC_ATTEMPT_KEY = 'oidc_login_attempted';
const OIDC_RETURN_TO_KEY = 'oidc_return_to';
// Renew the OIDC session when less than this much lifetime remains.
const OIDC_RENEWAL_THRESHOLD_MS = 15 * 60 * 1000;
// Delay before retrying a renewal that failed transiently (network/503).
const OIDC_RENEWAL_RETRY_MS = 60 * 1000;
// setTimeout treats delays above 2^31 - 1 ms as 0 — clamp to avoid firing
// immediately for far-future expiries.
const MAX_TIMER_DELAY_MS = 2 ** 31 - 1;
// Module-level so the two hook instances (AuthWrapper + Navigation) and
// StrictMode's double-invoked effects share one exchange/redirect — the
// handoff code is single-use server-side, so a second POST would fail.
let oidcExchangePromise: Promise<string | null> | null = null;
let oidcRedirectStarted = false;
// Renewal state is module-level for the same reason: at most one renewal
// timer and one in-flight refresh may exist app-wide, because the server
// rotates the refresh token on every renewal — concurrent renewals would
// invalidate each other.
let oidcRenewalTimer: ReturnType<typeof setTimeout> | null = null;
let oidcRenewalPromise: Promise<void> | null = null;
// Set when the server reports 404 no_refresh_token: the IdP issued no
// refresh token for this session, so silent renewal is impossible until a
// new token is stored. Expiry then falls back to the redirect-login path,
// which re-authenticates silently while the IdP session is still alive.
let oidcRenewalUnavailable = false;
const claimString = (value: unknown): string | undefined =>
typeof value === 'string' && value !== '' ? value : undefined;
function exchangeOidcCodeOnce(code: string): Promise<string | null> {
if (!oidcExchangePromise) {
oidcExchangePromise = (async () => {
@@ -49,6 +76,8 @@ export default function useAuth() {
const [showTokenModal, setShowTokenModal] = useState(false);
const [isAuthLoading, setIsAuthLoading] = useState(true);
const [oidcFailed, setOidcFailed] = useState(false);
const [oidcErrorCode, setOidcErrorCode] = useState<string | null>(null);
const [oidcProviderName, setOidcProviderName] = useState<string | null>(null);
const isGeneratingToken = useRef(false);
const generateNewToken = async () => {
@@ -83,7 +112,11 @@ export default function useAuth() {
const code = hash.slice('#oidc_code='.length);
const newToken = await exchangeOidcCodeOnce(code);
if (newToken) {
// A fresh session may carry a refresh token even if the previous
// one did not.
oidcRenewalUnavailable = false;
localStorage.setItem('authToken', newToken);
setOidcErrorCode(null);
sessionStorage.removeItem(OIDC_ATTEMPT_KEY);
const returnTo = sessionStorage.getItem(OIDC_RETURN_TO_KEY);
sessionStorage.removeItem(OIDC_RETURN_TO_KEY);
@@ -103,10 +136,9 @@ export default function useAuth() {
return;
}
if (hash.startsWith('#oidc_error=')) {
console.error(
'OIDC login failed:',
decodeURIComponent(hash.slice('#oidc_error='.length)),
);
const errorCode = decodeURIComponent(hash.slice('#oidc_error='.length));
console.error('OIDC login failed:', errorCode);
setOidcErrorCode(errorCode);
stripUrlFragment();
setOidcFailed(true);
setIsAuthLoading(false);
@@ -143,6 +175,7 @@ export default function useAuth() {
const config = await configRes.json();
resolvedAuthType = config.auth_type;
setAuthType(resolvedAuthType);
setOidcProviderName(config.oidc?.provider_name ?? null);
}
if (resolvedAuthType === 'oidc') {
@@ -163,6 +196,99 @@ export default function useAuth() {
initializeAuth();
}, [token, authType]);
useEffect(() => {
// Silent session renewal: refresh the OIDC session JWT shortly before
// it expires so users never hit a mid-session login redirect. Re-runs
// on every token change, which is what reschedules the timer after a
// successful renewal.
if (authType !== 'oidc') return;
const remaining = token ? getJwtRemainingMs(token) : null;
if (
!token ||
remaining === null ||
isJwtExpired(token) ||
oidcRenewalUnavailable
) {
// No renewable session (initializeAuth owns expired-token recovery)
// — drop any pending renewal/retry timer.
if (oidcRenewalTimer !== null) {
clearTimeout(oidcRenewalTimer);
oidcRenewalTimer = null;
}
return;
}
const scheduleRetry = () => {
if (oidcRenewalTimer !== null) clearTimeout(oidcRenewalTimer);
oidcRenewalTimer = setTimeout(renew, OIDC_RENEWAL_RETRY_MS);
};
const renew = () => {
if (oidcRenewalPromise) return; // single renewal in flight app-wide
oidcRenewalPromise = (async () => {
try {
// Another tab may have renewed first — its rotation invalidated
// our refresh token, so adopt the stored session instead of
// making a doomed network call.
const stored = localStorage.getItem('authToken');
if (stored !== token) {
dispatch(setToken(stored));
return;
}
const response = await userService.refreshOidcSession(token);
if (response.ok) {
const { token: newToken } = await response.json();
if (!newToken) {
scheduleRetry();
return;
}
oidcRenewalUnavailable = false;
localStorage.setItem('authToken', newToken);
// The token-keyed effect re-runs and schedules the next one.
dispatch(setToken(newToken));
return;
}
if (response.status === 404) {
// no_refresh_token: stop scheduling for this session.
oidcRenewalUnavailable = true;
return;
}
if (response.status === 401) {
// Session unusable (expired/revoked/disabled) — drop the
// token so initializeAuth walks through a fresh login.
localStorage.removeItem('authToken');
dispatch(setToken(null));
return;
}
// 503 or unexpected status: transient — retry once in a minute.
scheduleRetry();
} catch {
scheduleRetry(); // network error: transient
} finally {
oidcRenewalPromise = null;
}
})();
};
// One timer app-wide: replace whatever an earlier run (or the other
// hook instance) scheduled. A zero delay renews immediately.
if (oidcRenewalTimer !== null) clearTimeout(oidcRenewalTimer);
const delay = Math.min(
Math.max(remaining - OIDC_RENEWAL_THRESHOLD_MS, 0),
MAX_TIMER_DELAY_MS,
);
const timer = setTimeout(renew, delay);
oidcRenewalTimer = timer;
return () => {
// Clear only the timer this effect run set — another instance may
// have replaced it with its own since.
if (oidcRenewalTimer === timer) {
clearTimeout(timer);
oidcRenewalTimer = null;
}
};
}, [authType, token, dispatch]);
const handleTokenSubmit = (enteredToken: string) => {
localStorage.setItem('authToken', enteredToken);
dispatch(setToken(enteredToken));
@@ -171,6 +297,7 @@ export default function useAuth() {
const retryOidcLogin = () => {
sessionStorage.removeItem(OIDC_ATTEMPT_KEY);
setOidcErrorCode(null);
setOidcFailed(false);
setIsAuthLoading(true);
redirectToOidcLogin();
@@ -185,10 +312,11 @@ export default function useAuth() {
window.location.href = `${baseURL}${endpoints.USER.OIDC_LOGOUT}`;
};
const userEmail =
authType === 'oidc' && token
? (decodeJwtPayload(token)?.email as string | undefined)
: undefined;
const oidcClaims =
authType === 'oidc' && token ? decodeJwtPayload(token) : null;
const userEmail = claimString(oidcClaims?.email);
const userName = claimString(oidcClaims?.name);
const userPicture = claimString(oidcClaims?.picture);
return {
authType,
@@ -197,8 +325,12 @@ export default function useAuth() {
token,
handleTokenSubmit,
oidcFailed,
oidcErrorCode,
oidcProviderName,
retryOidcLogin,
logout,
userEmail,
userName,
userPicture,
};
}
+3
View File
@@ -33,6 +33,9 @@
"auth": {
"signInToContinue": "Sign in to continue to DocsGPT",
"signInWithSSO": "Sign in with SSO",
"signInWith": "Sign in with {{provider}}",
"notAuthorized": "Your account isn't authorized to access this workspace. Sign in with a different account or contact your administrator.",
"accountDisabled": "Your account has been disabled. Contact your administrator.",
"signOut": "Sign out"
},
"settings": {
+32 -1
View File
@@ -1,6 +1,6 @@
import { describe, expect, it } from 'vitest';
import { decodeJwtPayload, isJwtExpired } from './jwtUtils';
import { decodeJwtPayload, getJwtRemainingMs, isJwtExpired } from './jwtUtils';
const b64url = (value: object) =>
btoa(JSON.stringify(value))
@@ -50,3 +50,34 @@ describe('isJwtExpired', () => {
expect(isJwtExpired('garbage')).toBe(false);
});
});
describe('getJwtRemainingMs', () => {
it('returns the milliseconds left until expiry', () => {
const exp = Math.floor(Date.now() / 1000) + 3600;
const remaining = getJwtRemainingMs(makeToken({ sub: 'u', exp }));
expect(remaining).not.toBeNull();
// exp is truncated to whole seconds, so allow up to 1s of slack.
expect(remaining!).toBeGreaterThan(3599_000 - 1000);
expect(remaining!).toBeLessThanOrEqual(3600_000);
});
it('returns a negative value for an already-expired token', () => {
const exp = Math.floor(Date.now() / 1000) - 60;
const remaining = getJwtRemainingMs(makeToken({ sub: 'u', exp }));
expect(remaining).not.toBeNull();
expect(remaining!).toBeLessThan(0);
});
it('returns null when the exp claim is absent', () => {
expect(getJwtRemainingMs(makeToken({ sub: 'u' }))).toBeNull();
});
it('returns null when the exp claim is not a number', () => {
expect(getJwtRemainingMs(makeToken({ sub: 'u', exp: 'soon' }))).toBeNull();
});
it('returns null for undecodable tokens', () => {
expect(getJwtRemainingMs('garbage')).toBeNull();
expect(getJwtRemainingMs('')).toBeNull();
});
});
+6
View File
@@ -21,3 +21,9 @@ export function isJwtExpired(token: string, skewMs = 30000): boolean {
if (typeof exp !== 'number') return false;
return exp * 1000 <= Date.now() + skewMs;
}
export function getJwtRemainingMs(token: string): number | null {
const exp = decodeJwtPayload(token)?.exp;
if (typeof exp !== 'number') return null;
return exp * 1000 - Date.now();
}
+13 -1
View File
@@ -52,11 +52,23 @@ export ENCRYPTION_SECRET_KEY="e2e-fixed-encryption-key-never-use-in-prod"
# OIDC mode (AUTH_TYPE=oidc) — points at the mock IdP that oidc.spec.ts
# spawns on demand (scripts/e2e/mock_oidc_idp.py, port 7999). Discovery is
# lazy, so Flask boots fine before the IdP is up.
# lazy, so Flask boots fine before the IdP is up. Every OIDC_* var is pinned
# here because the app's load_dotenv() walks up and reads the repo .env —
# whatever a developer keeps there must not leak into the e2e stack.
if [[ "${AUTH_TYPE}" == "oidc" ]]; then
export OIDC_ISSUER="${OIDC_ISSUER:-http://127.0.0.1:7999}"
export OIDC_CLIENT_ID="${OIDC_CLIENT_ID:-docsgpt-e2e}"
export OIDC_FRONTEND_URL="${OIDC_FRONTEND_URL:-http://127.0.0.1:5179}"
export OIDC_CLIENT_SECRET="${OIDC_CLIENT_SECRET:-}"
export OIDC_SCOPES="${OIDC_SCOPES:-openid profile email}"
export OIDC_USER_ID_CLAIM="${OIDC_USER_ID_CLAIM:-sub}"
export OIDC_REDIRECT_URI="${OIDC_REDIRECT_URI:-}"
export OIDC_SESSION_LIFETIME_SECONDS="${OIDC_SESSION_LIFETIME_SECONDS:-28800}"
export OIDC_PROVIDER_NAME="${OIDC_PROVIDER_NAME:-}"
export OIDC_ALLOWED_GROUPS="${OIDC_ALLOWED_GROUPS:-}"
export OIDC_GROUPS_CLAIM="${OIDC_GROUPS_CLAIM:-groups}"
export SCIM_ENABLED="${SCIM_ENABLED:-false}"
export SCIM_TOKEN="${SCIM_TOKEN:-}"
fi
# -----------------------------------------------------------------------------
+136 -25
View File
@@ -4,17 +4,23 @@ Speaks the minimum OIDC surface the backend's ``AUTH_TYPE=oidc`` flow needs:
* ``GET /.well-known/openid-configuration`` (discovery)
* ``GET /authorize`` (auto-approves and redirects back with a code)
* ``POST /token`` (single-use code + PKCE S256 check, RS256 ``id_token``)
* ``POST /token`` (single-use code + PKCE S256 check, RS256 ``id_token``,
single-use rotating ``refresh_token``; also ``grant_type=refresh_token``)
* ``GET /jwks`` (public key for ID-token verification)
* ``GET /userinfo`` (Bearer access token from ``/token``)
* ``GET /end-session`` (honors ``post_logout_redirect_uri``)
* ``POST /trigger-backchannel-logout`` (test hook: signs and delivers a
back-channel logout token to a given URL)
* ``GET /healthz`` (liveness probe)
There is no login form: every ``/authorize`` request is approved as the user
configured via ``MOCK_OIDC_SUB`` / ``MOCK_OIDC_EMAIL`` (overridable per
request with ``?sub=``/``?email=`` for multi-user tests).
request with ``?sub=``/``?email=`` for multi-user tests). Group membership
comes from ``MOCK_OIDC_GROUPS`` (comma-separated).
Run standalone (does NOT import anything from ``application/``). Dependencies
(flask, python-jose, cryptography) are all in ``application/requirements.txt``.
(flask, python-jose, cryptography, requests) are all in
``application/requirements.txt``.
Usage::
@@ -34,6 +40,7 @@ import sys
import time
from urllib.parse import urlencode
import requests
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from flask import Flask, Response, jsonify, redirect, request
@@ -46,8 +53,19 @@ ISSUER = f"http://{HOST}:{PORT}"
DEFAULT_SUB = os.environ.get("MOCK_OIDC_SUB", "mock-oidc-user")
DEFAULT_EMAIL = os.environ.get("MOCK_OIDC_EMAIL", "mock-oidc-user@example.com")
DEFAULT_NAME = os.environ.get("MOCK_OIDC_NAME", "Mock OIDC User")
DEFAULT_GROUPS = [
group.strip()
for group in os.environ.get("MOCK_OIDC_GROUPS", "docsgpt-users").split(",")
if group.strip()
]
DEFAULT_CLIENT_ID = os.environ.get("MOCK_OIDC_CLIENT_ID", "docsgpt-e2e")
ID_TOKEN_TTL_SECONDS = 300
KID = "mock-oidc-key-1"
LOGOUT_TOKEN_TTL_SECONDS = 120
BACKCHANNEL_LOGOUT_EVENT = "http://schemas.openid.net/event/backchannel-logout"
# Random per process: each restart generates a fresh RSA key, and a fresh
# kid lets relying parties detect the change via their kid-miss refetch
# path instead of failing signature checks against a stale cached JWKS.
KID = f"mock-oidc-key-{secrets.token_hex(4)}"
app = Flask(__name__)
@@ -63,8 +81,13 @@ PUBLIC_JWK = {
"use": "sig",
}
# code -> {client_id, redirect_uri, code_challenge, nonce, sub, email, name}
# code -> {client_id, redirect_uri, code_challenge, nonce, sub, email, name, groups}
_codes: dict[str, dict] = {}
# access_token -> {client_id, sub, email, name, groups}
_access_tokens: dict[str, dict] = {}
# refresh_token -> {client_id, sub, email, name, groups} (single-use, rotated)
_refresh_tokens: dict[str, dict] = {}
_last_client_id: str | None = None
def _log(message: str) -> None:
@@ -77,6 +100,46 @@ def _pkce_challenge(verifier: str) -> str:
return base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
def _user_record(record: dict) -> dict:
"""Identity fields carried from /authorize through tokens and userinfo."""
return {
"client_id": record["client_id"],
"sub": record["sub"],
"email": record["email"],
"name": record["name"],
"groups": list(record.get("groups") or DEFAULT_GROUPS),
}
def _issue_tokens(record: dict, nonce: str | None) -> dict:
"""Mint an id_token (+ tracked access/refresh tokens) for ``record``."""
now = int(time.time())
claims = {
"iss": ISSUER,
"aud": record["client_id"],
"sub": record["sub"],
"email": record["email"],
"name": record["name"],
"groups": list(record.get("groups") or DEFAULT_GROUPS),
"iat": now,
"exp": now + ID_TOKEN_TTL_SECONDS,
}
if nonce:
claims["nonce"] = nonce
id_token = jose_jwt.encode(claims, PRIVATE_PEM, algorithm="RS256", headers={"kid": KID})
access_token = secrets.token_urlsafe(24)
refresh_token = secrets.token_urlsafe(24)
_access_tokens[access_token] = _user_record(record)
_refresh_tokens[refresh_token] = _user_record(record)
return {
"access_token": access_token,
"token_type": "Bearer",
"expires_in": ID_TOKEN_TTL_SECONDS,
"id_token": id_token,
"refresh_token": refresh_token,
}
@app.get("/.well-known/openid-configuration")
def discovery() -> Response:
return jsonify(
@@ -85,9 +148,12 @@ def discovery() -> Response:
"authorization_endpoint": f"{ISSUER}/authorize",
"token_endpoint": f"{ISSUER}/token",
"jwks_uri": f"{ISSUER}/jwks",
"userinfo_endpoint": f"{ISSUER}/userinfo",
"end_session_endpoint": f"{ISSUER}/end-session",
"backchannel_logout_supported": True,
"response_types_supported": ["code"],
"grant_types_supported": ["authorization_code"],
"grant_types_supported": ["authorization_code", "refresh_token"],
"token_endpoint_auth_methods_supported": ["client_secret_post", "client_secret_basic"],
"id_token_signing_alg_values_supported": ["RS256"],
"code_challenge_methods_supported": ["S256"],
"scopes_supported": ["openid", "profile", "email"],
@@ -124,6 +190,7 @@ def authorize() -> Response:
"sub": args.get("sub") or DEFAULT_SUB,
"email": args.get("email") or DEFAULT_EMAIL,
"name": DEFAULT_NAME,
"groups": list(DEFAULT_GROUPS),
}
_log(f"authorize: auto-approved sub={_codes[code]['sub']}")
separator = "&" if "?" in args["redirect_uri"] else "?"
@@ -133,8 +200,20 @@ def authorize() -> Response:
@app.post("/token")
def token() -> Response:
global _last_client_id
form = request.form
if form.get("grant_type") != "authorization_code":
grant_type = form.get("grant_type")
if grant_type == "refresh_token":
record = _refresh_tokens.pop(form.get("refresh_token", ""), None) # single-use
if record is None:
_log("token: unknown or reused refresh_token")
return jsonify({"error": "invalid_grant"}), 400
_last_client_id = record["client_id"]
_log(f"token: refreshed tokens for sub={record['sub']}")
return jsonify(_issue_tokens(record, nonce=None))
if grant_type != "authorization_code":
return jsonify({"error": "unsupported_grant_type"}), 400
record = _codes.pop(form.get("code", ""), None) # single-use
if record is None:
@@ -149,28 +228,60 @@ def token() -> Response:
_log("token: PKCE verification failed")
return jsonify({"error": "invalid_grant", "error_description": "PKCE failed"}), 400
_last_client_id = record["client_id"]
_log(f"token: issued id_token for sub={record['sub']}")
return jsonify(_issue_tokens(record, nonce=record["nonce"]))
@app.get("/userinfo")
def userinfo() -> Response:
header = request.headers.get("Authorization", "")
access_token = header[len("Bearer "):] if header.startswith("Bearer ") else ""
record = _access_tokens.get(access_token)
if record is None:
_log("userinfo: unknown access token")
return jsonify({"error": "invalid_token"}), 401
return jsonify(
{
"sub": record["sub"],
"email": record["email"],
"name": record["name"],
"groups": record["groups"],
}
)
@app.post("/trigger-backchannel-logout")
def trigger_backchannel_logout() -> Response:
"""Test hook: sign a back-channel logout token and POST it to ``url``."""
body = request.get_json(silent=True) or {}
url = body.get("url")
sub = body.get("sub")
sid = body.get("sid")
if not url or not (sub or sid):
return jsonify({"error": "url and sub (or sid) required"}), 400
now = int(time.time())
claims = {
"iss": ISSUER,
"aud": record["client_id"],
"sub": record["sub"],
"email": record["email"],
"name": record["name"],
"aud": body.get("client_id") or _last_client_id or DEFAULT_CLIENT_ID,
"iat": now,
"exp": now + ID_TOKEN_TTL_SECONDS,
"exp": now + LOGOUT_TOKEN_TTL_SECONDS,
"jti": secrets.token_urlsafe(16),
"events": {BACKCHANNEL_LOGOUT_EVENT: {}},
}
if record["nonce"]:
claims["nonce"] = record["nonce"]
id_token = jose_jwt.encode(claims, PRIVATE_PEM, algorithm="RS256", headers={"kid": KID})
_log(f"token: issued id_token for sub={record['sub']}")
return jsonify(
{
"access_token": secrets.token_urlsafe(24),
"token_type": "Bearer",
"expires_in": ID_TOKEN_TTL_SECONDS,
"id_token": id_token,
}
)
if sub:
claims["sub"] = sub
if sid:
claims["sid"] = sid
logout_token = jose_jwt.encode(claims, PRIVATE_PEM, algorithm="RS256", headers={"kid": KID})
try:
downstream = requests.post(url, data={"logout_token": logout_token}, timeout=10)
except requests.RequestException as exc:
_log(f"trigger-backchannel-logout: delivery to {url} failed: {exc}")
return jsonify({"error": "delivery_failed", "detail": str(exc)}), 502
_log(f"trigger-backchannel-logout: {url} responded {downstream.status_code}")
return jsonify({"status": downstream.status_code})
@app.get("/end-session")
@@ -188,7 +299,7 @@ def healthz() -> Response:
def main() -> None:
_log(f"listening on {ISSUER} (sub={DEFAULT_SUB})")
_log(f"listening on {ISSUER} (sub={DEFAULT_SUB}, groups={','.join(DEFAULT_GROUPS)})")
app.run(host=HOST, port=PORT, debug=False, use_reloader=False, threaded=True)
+184
View File
@@ -0,0 +1,184 @@
"""Integration tests for the SCIM 2.0 endpoints against a live Postgres.
These tests drive the real Flask app through its test client and verify
row state in the database configured by ``POSTGRES_URI``. Redis is not
required — the denylist functions are stubbed. They are skipped by the
default ``pytest`` run (``--ignore=tests/integration`` in ``pytest.ini``)
and marked ``@pytest.mark.integration``. Run them locally with::
.venv/bin/python -m pytest tests/integration/test_scim.py -q --no-cov \\
-p no:cacheprovider --override-ini "addopts="
"""
from __future__ import annotations
import uuid
from types import SimpleNamespace
from unittest.mock import patch
import pytest
from sqlalchemy import text
from application.core.settings import settings
from application.storage.db.repositories.auth_events import AuthEventsRepository
from application.storage.db.repositories.users import UsersRepository
from application.storage.db.session import db_readonly, db_session
SCIM_TOKEN = "scim-test-token"
AUTH = {"Authorization": f"Bearer {SCIM_TOKEN}"}
ERROR_URN = "urn:ietf:params:scim:api:messages:2.0:Error"
PATCH_OP_URN = "urn:ietf:params:scim:api:messages:2.0:PatchOp"
pytestmark = [
pytest.mark.integration,
pytest.mark.skipif(
not settings.POSTGRES_URI,
reason="POSTGRES_URI not set — skipping SCIM integration tests",
),
]
@pytest.fixture(scope="module")
def app():
"""Real Flask app; /scim/ paths bypass JWT auth so no handle_auth patching is needed."""
from application.app import app as flask_app
flask_app.config["TESTING"] = True
return flask_app
@pytest.fixture
def client(app):
return app.test_client()
@pytest.fixture
def scim_env(monkeypatch):
"""Enable SCIM on the settings singleton and stub the Redis denylist."""
monkeypatch.setattr(settings, "SCIM_ENABLED", True)
monkeypatch.setattr(settings, "SCIM_TOKEN", SCIM_TOKEN)
with patch("application.api.scim.routes.deny_user") as deny_user_mock, patch(
"application.api.scim.routes.allow_user"
) as allow_user_mock:
yield SimpleNamespace(deny_user=deny_user_mock, allow_user=allow_user_mock)
@pytest.fixture
def scim_user_name():
"""Unique per-run userName; deletes the user + audit rows afterwards."""
name = f"scim-it-{uuid.uuid4().hex[:10]}@example.com"
yield name
with db_session() as conn:
conn.execute(text("DELETE FROM auth_events WHERE user_id = :user_id"), {"user_id": name})
conn.execute(text("DELETE FROM users WHERE user_id = :user_id"), {"user_id": name})
def _fetch_user(user_name: str):
with db_readonly() as conn:
return UsersRepository(conn).get(user_name)
def _fetch_events(user_name: str):
with db_readonly() as conn:
return AuthEventsRepository(conn).list_recent(user_name)
@pytest.mark.integration
class TestScimLifecycle:
def test_full_user_lifecycle(self, client, scim_env, scim_user_name):
# Create
response = client.post("/scim/v2/Users", headers=AUTH, json={"userName": scim_user_name})
assert response.status_code == 201, response.get_data(as_text=True)
created = response.get_json()
pk = created["id"]
assert created["userName"] == scim_user_name
assert created["active"] is True
assert created["emails"] == [{"value": scim_user_name, "primary": True}]
assert response.headers["Location"].endswith(f"/scim/v2/Users/{pk}")
row = _fetch_user(scim_user_name)
assert row is not None
assert row["active"] is True
assert [event["event"] for event in _fetch_events(scim_user_name)] == ["scim_created"]
# GET by id
response = client.get(f"/scim/v2/Users/{pk}", headers=AUTH)
assert response.status_code == 200
assert response.get_json()["id"] == pk
# List with exact userName filter
response = client.get(
"/scim/v2/Users",
headers=AUTH,
query_string={"filter": f'userName eq "{scim_user_name}"'},
)
assert response.status_code == 200
body = response.get_json()
assert body["totalResults"] == 1
assert body["itemsPerPage"] == 1
assert body["Resources"][0]["id"] == pk
# Deactivate via PATCH (Okta no-path form with string value)
response = client.patch(
f"/scim/v2/Users/{pk}",
headers=AUTH,
json={
"schemas": [PATCH_OP_URN],
"Operations": [{"op": "replace", "value": {"active": "False"}}],
},
)
assert response.status_code == 200
assert response.get_json()["active"] is False
assert _fetch_user(scim_user_name)["active"] is False
deactivations = [e for e in _fetch_events(scim_user_name) if e["event"] == "scim_deactivated"]
assert len(deactivations) == 1
assert deactivations[0]["metadata"] == {"via": "scim"}
scim_env.deny_user.assert_called_once_with(scim_user_name)
# Reactivate via PUT (full replace; only active is honored)
response = client.put(
f"/scim/v2/Users/{pk}",
headers=AUTH,
json={"userName": scim_user_name, "active": True},
)
assert response.status_code == 200
assert response.get_json()["active"] is True
assert _fetch_user(scim_user_name)["active"] is True
scim_env.allow_user.assert_called_once_with(scim_user_name)
assert any(e["event"] == "scim_reactivated" for e in _fetch_events(scim_user_name))
# Soft delete deactivates again
response = client.delete(f"/scim/v2/Users/{pk}", headers=AUTH)
assert response.status_code == 204
assert response.data == b""
assert _fetch_user(scim_user_name)["active"] is False
assert scim_env.deny_user.call_count == 2
# Duplicate create conflicts (row still exists after soft delete)
response = client.post("/scim/v2/Users", headers=AUTH, json={"userName": scim_user_name})
assert response.status_code == 409
body = response.get_json()
assert body["schemas"] == [ERROR_URN]
assert body["scimType"] == "uniqueness"
@pytest.mark.integration
class TestScimAccess:
def test_wrong_bearer_token_rejected(self, client, scim_env):
response = client.get("/scim/v2/Users", headers={"Authorization": "Bearer wrong"})
assert response.status_code == 401
assert response.get_json()["schemas"] == [ERROR_URN]
def test_get_user_with_malformed_uuid_returns_404(self, client, scim_env):
response = client.get("/scim/v2/Users/not-a-uuid", headers=AUTH)
assert response.status_code == 404
assert response.get_json()["status"] == "404"
def test_service_provider_config_served(self, client, scim_env):
response = client.get("/scim/v2/ServiceProviderConfig", headers=AUTH)
assert response.status_code == 200
assert response.content_type.startswith("application/scim+json")
assert response.get_json()["patch"] == {"supported": True}
+6 -1
View File
@@ -43,7 +43,10 @@ class TestConfigRoute:
@pytest.mark.unit
def test_returns_auth_config(self, client):
response = client.get("/api/config")
# Pin AUTH_TYPE so the assertion doesn't depend on the dev .env.
with patch("application.app.settings") as mock_settings:
mock_settings.AUTH_TYPE = None
response = client.get("/api/config")
assert response.status_code == 200
data = json.loads(response.data)
assert "auth_type" in data
@@ -54,6 +57,7 @@ class TestConfigRoute:
def test_oidc_config_exposes_login_paths(self, client):
with patch("application.app.settings") as mock_settings:
mock_settings.AUTH_TYPE = "oidc"
mock_settings.OIDC_PROVIDER_NAME = "Test SSO"
response = client.get("/api/config")
assert response.status_code == 200
data = json.loads(response.data)
@@ -62,6 +66,7 @@ class TestConfigRoute:
assert data["oidc"] == {
"login_path": "/api/auth/oidc/login",
"logout_path": "/api/auth/oidc/logout",
"provider_name": "Test SSO",
}
+834 -22
View File
@@ -4,6 +4,8 @@ import base64
import hashlib
import json
import time
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import Mock, patch
from urllib.parse import parse_qs, urlparse
@@ -20,12 +22,14 @@ CLIENT_ID = "docsgpt-test"
FRONTEND_URL = "http://frontend.test"
JWT_SECRET = "test-oidc-secret"
KID = "test-key-1"
BCL_EVENT = "http://schemas.openid.net/event/backchannel-logout"
DISCOVERY = {
"issuer": ISSUER,
"authorization_endpoint": "https://idp.test/authorize",
"token_endpoint": "https://idp.test/token",
"jwks_uri": "https://idp.test/jwks",
"userinfo_endpoint": "https://idp.test/userinfo",
"end_session_endpoint": "https://idp.test/end-session",
}
@@ -67,6 +71,34 @@ def id_token_claims(**overrides):
return claims
def logout_token_claims(**overrides):
now = int(time.time())
claims = {
"iss": ISSUER,
"aud": CLIENT_ID,
"sub": "oidc-user-1",
"iat": now,
"jti": "bcl-jti-1",
"events": {BCL_EVENT: {}},
}
claims.update(overrides)
return claims
def make_session_token(**overrides):
"""Mint a session JWT the way the callback does, for refresh-route tests."""
now = int(time.time())
payload = {
"sub": "oidc-user-1",
"jti": "jti-1",
"iat": now,
"exp": now + 3600,
"oidc_sub": "oidc-user-1",
}
payload.update(overrides)
return jose_jwt.encode(payload, JWT_SECRET, algorithm="HS256")
class FakeRedis:
"""Minimal stand-in for the redis-py client used by the oidc routes."""
@@ -83,16 +115,20 @@ class FakeRedis:
return self.store.pop(key, None)
def make_fake_get(jwks_keys):
"""requests.get stub serving discovery + JWKS; ``jwks_keys`` is mutable."""
def make_fake_get(jwks_keys, discovery=None, userinfo=None, userinfo_status=200):
"""requests.get stub serving discovery + JWKS (+ optional userinfo); ``jwks_keys`` is mutable."""
document = dict(DISCOVERY if discovery is None else discovery)
def fake_get(url, timeout=None):
def fake_get(url, timeout=None, headers=None):
resp = Mock()
resp.status_code = 200
if "openid-configuration" in url:
resp.json.return_value = dict(DISCOVERY)
elif url == DISCOVERY["jwks_uri"]:
resp.json.return_value = dict(document)
elif url == document["jwks_uri"]:
resp.json.return_value = {"keys": list(jwks_keys)}
elif userinfo is not None and url == document.get("userinfo_endpoint"):
resp.status_code = userinfo_status
resp.json.return_value = dict(userinfo)
else:
resp.status_code = 404
return resp
@@ -111,6 +147,9 @@ def oidc_settings(monkeypatch):
monkeypatch.setattr(settings, "OIDC_FRONTEND_URL", FRONTEND_URL)
monkeypatch.setattr(settings, "OIDC_REDIRECT_URI", None)
monkeypatch.setattr(settings, "OIDC_SESSION_LIFETIME_SECONDS", 28800)
monkeypatch.setattr(settings, "OIDC_PROVIDER_NAME", None)
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", None)
monkeypatch.setattr(settings, "OIDC_GROUPS_CLAIM", "groups")
monkeypatch.setattr(settings, "JWT_SECRET_KEY", JWT_SECRET)
@@ -214,6 +253,29 @@ class TestValidateIdToken:
with pytest.raises(provider.OIDCError):
self._validate(forged)
def test_rekeyed_idp_with_reused_kid_recovers(self):
# The IdP replaced its signing key but kept the kid (mock IdP
# restarts, sloppy rotations): the cached key fails the signature,
# one forced JWKS refetch picks up the new key and validation
# succeeds.
from application.api.oidc import provider
new_pem = _generate_rsa_pem()
new_jwk = {
**jwk.construct(new_pem, algorithm="RS256").public_key().to_dict(),
"kid": KID,
"use": "sig",
}
token = sign_id_token(id_token_claims(), key=new_pem)
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
provider.get_jwks() # prime the cache with the OLD key
mock_requests.get.side_effect = make_fake_get([new_jwk])
claims = provider.validate_id_token(token, "test-nonce")
assert claims["sub"] == "oidc-user-1"
def test_unknown_kid_triggers_single_jwks_refetch(self):
from application.api.oidc import provider
@@ -274,12 +336,13 @@ class TestExchangeCode:
assert sent["client_id"] == CLIENT_ID
assert "client_secret" not in sent
def test_includes_client_secret_when_configured(self, monkeypatch):
def test_includes_client_secret_when_post_method_supported(self, monkeypatch):
from application.api.oidc import provider
monkeypatch.setattr(settings, "OIDC_CLIENT_SECRET", "s3cret")
discovery = {**DISCOVERY, "token_endpoint_auth_methods_supported": ["client_secret_post"]}
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK], discovery=discovery)
mock_requests.post.return_value = Mock(
status_code=200, json=Mock(return_value={"id_token": "x"})
)
@@ -297,6 +360,195 @@ class TestExchangeCode:
provider.exchange_code("auth-code", "verifier-123", "https://app.test/cb")
@pytest.mark.unit
class TestTokenEndpointAuthMethod:
def _exchange(self, discovery=None):
from application.api.oidc import provider
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK], discovery=discovery)
mock_requests.post.return_value = Mock(
status_code=200, json=Mock(return_value={"id_token": "x"})
)
provider.exchange_code("auth-code", "verifier-123", "https://app.test/cb")
return mock_requests.post.call_args
def test_basic_auth_when_discovery_omits_methods(self, monkeypatch):
monkeypatch.setattr(settings, "OIDC_CLIENT_SECRET", "s3cret")
call = self._exchange()
assert call.kwargs["auth"] == (CLIENT_ID, "s3cret")
assert "client_secret" not in call.kwargs["data"]
def test_basic_auth_when_only_client_secret_basic_supported(self, monkeypatch):
monkeypatch.setattr(settings, "OIDC_CLIENT_SECRET", "s3cret")
discovery = {**DISCOVERY, "token_endpoint_auth_methods_supported": ["client_secret_basic"]}
call = self._exchange(discovery)
assert call.kwargs["auth"] == (CLIENT_ID, "s3cret")
assert "client_secret" not in call.kwargs["data"]
def test_post_auth_when_client_secret_post_supported(self, monkeypatch):
monkeypatch.setattr(settings, "OIDC_CLIENT_SECRET", "s3cret")
discovery = {
**DISCOVERY,
"token_endpoint_auth_methods_supported": ["client_secret_post", "client_secret_basic"],
}
call = self._exchange(discovery)
assert call.kwargs["data"]["client_secret"] == "s3cret"
assert "auth" not in call.kwargs
def test_no_auth_kwarg_when_no_secret(self):
call = self._exchange()
assert "auth" not in call.kwargs
assert "client_secret" not in call.kwargs["data"]
@pytest.mark.unit
class TestFetchUserinfo:
def test_sends_bearer_token_and_returns_claims(self):
from application.api.oidc import provider
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get(
[PUBLIC_JWK], userinfo={"sub": "oidc-user-1", "groups": ["devs"]}
)
info = provider.fetch_userinfo("at-123")
assert info == {"sub": "oidc-user-1", "groups": ["devs"]}
userinfo_calls = [
call
for call in mock_requests.get.call_args_list
if call.args[0] == DISCOVERY["userinfo_endpoint"]
]
assert len(userinfo_calls) == 1
assert userinfo_calls[0].kwargs["headers"]["Authorization"] == "Bearer at-123"
def test_missing_endpoint_raises(self):
from application.api.oidc import provider
discovery = {k: v for k, v in DISCOVERY.items() if k != "userinfo_endpoint"}
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK], discovery=discovery)
with pytest.raises(provider.OIDCError):
provider.fetch_userinfo("at-123")
def test_non_200_raises(self):
from application.api.oidc import provider
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get(
[PUBLIC_JWK], userinfo={"sub": "x"}, userinfo_status=500
)
with pytest.raises(provider.OIDCError):
provider.fetch_userinfo("at-123")
@pytest.mark.unit
class TestRefreshGrant:
def test_posts_refresh_token_grant(self):
from application.api.oidc import provider
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
mock_requests.post.return_value = Mock(
status_code=200, json=Mock(return_value={"access_token": "at-2"})
)
tokens = provider.refresh_grant("rt-old")
assert tokens == {"access_token": "at-2"}
sent = mock_requests.post.call_args.kwargs["data"]
assert sent["grant_type"] == "refresh_token"
assert sent["refresh_token"] == "rt-old"
assert sent["client_id"] == CLIENT_ID
@pytest.mark.unit
class TestValidateIdTokenNonceOptional:
def test_nonce_none_skips_nonce_check(self):
from application.api.oidc import provider
claims = id_token_claims()
del claims["nonce"]
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
decoded = provider.validate_id_token(sign_id_token(claims), nonce=None)
assert decoded["sub"] == "oidc-user-1"
def test_nonce_none_accepts_token_that_still_has_nonce(self):
from application.api.oidc import provider
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
decoded = provider.validate_id_token(sign_id_token(id_token_claims()), nonce=None)
assert decoded["sub"] == "oidc-user-1"
@pytest.mark.unit
class TestValidateLogoutToken:
def _validate(self, token, jwks_keys=None):
from application.api.oidc import provider
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get(
jwks_keys if jwks_keys is not None else [PUBLIC_JWK]
)
return provider.validate_logout_token(token)
def test_valid_sub_token_returns_claims(self):
claims = self._validate(sign_id_token(logout_token_claims()))
assert claims["sub"] == "oidc-user-1"
def test_valid_sid_only_token(self):
claims = logout_token_claims(sid="sess-9")
del claims["sub"]
assert self._validate(sign_id_token(claims))["sid"] == "sess-9"
def test_missing_events_claim_rejected(self):
from application.api.oidc import provider
claims = logout_token_claims()
del claims["events"]
with pytest.raises(provider.OIDCError):
self._validate(sign_id_token(claims))
def test_wrong_event_uri_rejected(self):
from application.api.oidc import provider
claims = logout_token_claims(events={"http://other.event/uri": {}})
with pytest.raises(provider.OIDCError):
self._validate(sign_id_token(claims))
def test_nonce_prohibited(self):
from application.api.oidc import provider
with pytest.raises(provider.OIDCError):
self._validate(sign_id_token(logout_token_claims(nonce="n-1")))
def test_missing_sub_and_sid_rejected(self):
from application.api.oidc import provider
claims = logout_token_claims()
del claims["sub"]
with pytest.raises(provider.OIDCError):
self._validate(sign_id_token(claims))
def test_wrong_audience_rejected(self):
from application.api.oidc import provider
with pytest.raises(provider.OIDCError):
self._validate(sign_id_token(logout_token_claims(aud="someone-else")))
@pytest.fixture(scope="module")
def app():
with patch("application.app.handle_auth", return_value={"sub": "test_user"}):
@@ -318,17 +570,49 @@ def fake_redis():
yield redis
def _mint_id_token_response(stored_nonce):
return Mock(
status_code=200,
json=Mock(
return_value={
"access_token": "at",
"token_type": "Bearer",
"id_token": sign_id_token(id_token_claims(nonce=stored_nonce)),
}
),
@pytest.fixture
def db_mocks():
"""Patch the DB seams of the oidc routes; default: unknown user, active on upsert."""
users_repo = Mock()
users_repo.get.return_value = None
users_repo.upsert.side_effect = lambda user_id: {"user_id": user_id, "active": True}
events_repo = Mock()
@contextmanager
def fake_session():
yield Mock()
with patch("application.api.oidc.routes.db_session", fake_session), patch(
"application.api.oidc.routes.db_readonly", fake_session
), patch(
"application.api.oidc.routes.UsersRepository", return_value=users_repo
), patch(
"application.api.oidc.routes.AuthEventsRepository", return_value=events_repo
):
yield SimpleNamespace(users=users_repo, events=events_repo)
def _seed_state(fake_redis, state="state-1", nonce="nonce-1"):
fake_redis.set(
f"oidc:state:{state}",
json.dumps({"code_verifier": "verifier-1", "nonce": nonce}),
)
return state, nonce
def _signed_token_response(claims, **extra):
"""Token-endpoint response carrying an id_token signed over ``claims``."""
body = {
"access_token": "at",
"token_type": "Bearer",
"id_token": sign_id_token(claims),
}
body.update(extra)
return Mock(status_code=200, json=Mock(return_value=body))
def _mint_id_token_response(stored_nonce, **extra):
return _signed_token_response(id_token_claims(nonce=stored_nonce), **extra)
@pytest.mark.unit
@@ -376,12 +660,12 @@ class TestLoginRoute:
@pytest.mark.unit
class TestCallbackRoute:
@pytest.fixture(autouse=True)
def _db(self, db_mocks):
self.db = db_mocks
def _seed_state(self, fake_redis, state="state-1", nonce="nonce-1"):
fake_redis.set(
f"oidc:state:{state}",
json.dumps({"code_verifier": "verifier-1", "nonce": nonce}),
)
return state, nonce
return _seed_state(fake_redis, state=state, nonce=nonce)
def test_happy_path_mints_session_and_redirects_with_handoff(self, client, fake_redis):
state, nonce = self._seed_state(fake_redis)
@@ -516,3 +800,531 @@ class TestLogoutRoute:
assert response.status_code == 302
assert response.headers["Location"] == FRONTEND_URL
@pytest.mark.unit
class TestCallbackGroups:
@pytest.fixture(autouse=True)
def _db(self, db_mocks):
self.db = db_mocks
def _callback(self, client, fake_redis, claims, userinfo=None, userinfo_status=200):
state, nonce = _seed_state(fake_redis)
claims = {**claims, "nonce": nonce}
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get(
[PUBLIC_JWK], userinfo=userinfo, userinfo_status=userinfo_status
)
mock_requests.post.return_value = _signed_token_response(claims)
return client.get(f"/api/auth/oidc/callback?code=abc&state={state}")
def test_login_allowed_when_group_matches(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "admins, devs")
response = self._callback(client, fake_redis, id_token_claims(groups=["devs", "qa"]))
assert "#oidc_code=" in response.headers["Location"]
def test_login_denied_when_no_group_matches(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "admins, devs")
response = self._callback(client, fake_redis, id_token_claims(groups=["qa"]))
assert response.headers["Location"] == f"{FRONTEND_URL}/#oidc_error=not_authorized"
call = self.db.events.insert.call_args
assert call.args == ("oidc-user-1", "oidc_login_denied")
assert call.kwargs["metadata"] == {"reason": "not_authorized", "groups": ["qa"]}
def test_single_string_group_claim_coerced_to_list(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "devs")
response = self._callback(client, fake_redis, id_token_claims(groups="devs"))
assert "#oidc_code=" in response.headers["Location"]
def test_allow_everyone_when_allowlist_unset(self, client, fake_redis):
response = self._callback(client, fake_redis, id_token_claims())
assert "#oidc_code=" in response.headers["Location"]
def test_custom_groups_claim_name(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "admin")
monkeypatch.setattr(settings, "OIDC_GROUPS_CLAIM", "roles")
response = self._callback(client, fake_redis, id_token_claims(roles=["admin"]))
assert "#oidc_code=" in response.headers["Location"]
def test_missing_groups_claim_recovered_via_userinfo(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "devs")
response = self._callback(
client,
fake_redis,
id_token_claims(),
userinfo={"sub": "oidc-user-1", "groups": ["devs"]},
)
assert "#oidc_code=" in response.headers["Location"]
def test_userinfo_sub_mismatch_fails_auth(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "devs")
response = self._callback(
client,
fake_redis,
id_token_claims(),
userinfo={"sub": "intruder", "groups": ["devs"]},
)
assert response.headers["Location"] == f"{FRONTEND_URL}/#oidc_error=auth_failed"
def test_userinfo_failure_is_nonfatal_then_groups_denied(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "devs")
response = self._callback(
client,
fake_redis,
id_token_claims(),
userinfo={"sub": "oidc-user-1", "groups": ["devs"]},
userinfo_status=500,
)
assert response.headers["Location"] == f"{FRONTEND_URL}/#oidc_error=not_authorized"
def test_missing_user_id_claim_recovered_via_userinfo(self, client, fake_redis, monkeypatch):
monkeypatch.setattr(settings, "OIDC_USER_ID_CLAIM", "preferred_username")
response = self._callback(
client,
fake_redis,
id_token_claims(),
userinfo={"sub": "oidc-user-1", "preferred_username": "alice"},
)
handoff = response.headers["Location"].split("#oidc_code=", 1)[1]
session_token = fake_redis.store[f"oidc:handoff:{handoff}"].decode("utf-8")
decoded = jose_jwt.decode(session_token, JWT_SECRET, algorithms=["HS256"])
assert decoded["sub"] == "alice"
@pytest.mark.unit
class TestCallbackUserGate:
@pytest.fixture(autouse=True)
def _db(self, db_mocks):
self.db = db_mocks
def _callback(self, client, fake_redis):
state, nonce = _seed_state(fake_redis)
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
mock_requests.post.return_value = _mint_id_token_response(nonce)
return client.get(f"/api/auth/oidc/callback?code=abc&state={state}")
def test_disabled_user_rejected(self, client, fake_redis):
self.db.users.get.return_value = {"user_id": "oidc-user-1", "active": False}
response = self._callback(client, fake_redis)
assert response.headers["Location"] == f"{FRONTEND_URL}/#oidc_error=account_disabled"
assert not any(key.startswith("oidc:handoff:") for key in fake_redis.store)
call = self.db.events.insert.call_args
assert call.args == ("oidc-user-1", "oidc_login_denied")
assert call.kwargs["metadata"] == {"reason": "account_disabled"}
def test_new_user_provisioned_and_login_audited(self, client, fake_redis):
response = self._callback(client, fake_redis)
assert "#oidc_code=" in response.headers["Location"]
self.db.users.upsert.assert_called_once_with("oidc-user-1")
call = self.db.events.insert.call_args
assert call.args == ("oidc-user-1", "oidc_login")
assert call.kwargs["metadata"] == {"email": "user@example.com", "groups": None}
assert call.kwargs["user_agent"]
def test_existing_active_user_not_upserted(self, client, fake_redis):
self.db.users.get.return_value = {"user_id": "oidc-user-1", "active": True}
response = self._callback(client, fake_redis)
assert "#oidc_code=" in response.headers["Location"]
self.db.users.upsert.assert_not_called()
def test_db_outage_does_not_block_login(self, client, fake_redis):
with patch(
"application.api.oidc.routes.db_session",
side_effect=RuntimeError("db down"),
):
response = self._callback(client, fake_redis)
assert "#oidc_code=" in response.headers["Location"]
def test_audit_insert_failure_does_not_block_login(self, client, fake_redis):
self.db.events.insert.side_effect = RuntimeError("insert failed")
response = self._callback(client, fake_redis)
assert "#oidc_code=" in response.headers["Location"]
def test_successful_login_clears_session_revocations(self, client, fake_redis):
# A back-channel logout denylists the IdP sub; a fresh IdP-blessed
# login must lift session-level revocations or re-login would stay
# blocked until the denylist TTL expires.
with patch("application.api.oidc.routes.denylist") as deny:
response = self._callback(client, fake_redis)
assert "#oidc_code=" in response.headers["Location"]
deny.allow_user.assert_called_once_with("oidc-user-1")
deny.allow_idp_sub.assert_called_once_with("oidc-user-1")
def test_denied_login_does_not_clear_revocations(self, client, fake_redis):
self.db.users.get.return_value = {"user_id": "oidc-user-1", "active": False}
with patch("application.api.oidc.routes.denylist") as deny:
response = self._callback(client, fake_redis)
assert "account_disabled" in response.headers["Location"]
deny.allow_user.assert_not_called()
deny.allow_idp_sub.assert_not_called()
@pytest.mark.unit
class TestSessionTokenMint:
@pytest.fixture(autouse=True)
def _db(self, db_mocks):
self.db = db_mocks
def _login_decoded(self, client, fake_redis, claims=None, **token_extra):
state, nonce = _seed_state(fake_redis)
claims = {**(claims or id_token_claims()), "nonce": nonce}
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
mock_requests.post.return_value = _signed_token_response(claims, **token_extra)
response = client.get(f"/api/auth/oidc/callback?code=abc&state={state}")
handoff = response.headers["Location"].split("#oidc_code=", 1)[1]
session_token = fake_redis.store[f"oidc:handoff:{handoff}"].decode("utf-8")
return jose_jwt.decode(session_token, JWT_SECRET, algorithms=["HS256"])
def test_contains_jti_and_oidc_sub(self, client, fake_redis):
decoded = self._login_decoded(client, fake_redis)
assert len(decoded["jti"]) == 36
assert decoded["oidc_sub"] == "oidc-user-1"
assert "oidc_sid" not in decoded
def test_oidc_sid_included_when_id_token_has_sid(self, client, fake_redis):
decoded = self._login_decoded(client, fake_redis, claims=id_token_claims(sid="sess-42"))
assert decoded["oidc_sid"] == "sess-42"
def test_picture_claim_passthrough(self, client, fake_redis):
decoded = self._login_decoded(
client, fake_redis, claims=id_token_claims(picture="https://img.test/me.png")
)
assert decoded["picture"] == "https://img.test/me.png"
def test_overlong_picture_dropped(self, client, fake_redis):
decoded = self._login_decoded(
client, fake_redis, claims=id_token_claims(picture="https://img.test/" + "x" * 2048)
)
assert "picture" not in decoded
def test_refresh_token_stored_under_jti(self, client, fake_redis):
decoded = self._login_decoded(client, fake_redis, refresh_token="rt-1")
assert fake_redis.store[f"oidc:refresh:{decoded['jti']}"] == b"rt-1"
def test_no_refresh_key_when_idp_sends_none(self, client, fake_redis):
self._login_decoded(client, fake_redis)
assert not any(key.startswith("oidc:refresh:") for key in fake_redis.store)
@pytest.mark.unit
class TestBackchannelLogoutRoute:
URL = "/api/auth/oidc/backchannel-logout"
@pytest.fixture(autouse=True)
def _seams(self, db_mocks):
self.db = db_mocks
with patch("application.api.oidc.routes.denylist") as deny:
self.denylist = deny
yield
def _post(self, client, claims=None, **kwargs):
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
if claims is not None:
kwargs.setdefault("data", {"logout_token": sign_id_token(claims)})
return client.post(self.URL, **kwargs)
def test_sub_logout_denylists_sub(self, client, fake_redis):
response = self._post(client, logout_token_claims())
assert response.status_code == 200
assert response.headers["Cache-Control"] == "no-store"
self.denylist.deny_idp_sub.assert_called_once_with("oidc-user-1")
self.denylist.deny_sid.assert_not_called()
call = self.db.events.insert.call_args
assert call.args == ("oidc-user-1", "backchannel_logout")
assert call.kwargs["metadata"] is None
def test_sid_only_logout_denylists_sid(self, client, fake_redis):
claims = logout_token_claims(sid="sess-7", jti="bcl-jti-2")
del claims["sub"]
response = self._post(client, claims)
assert response.status_code == 200
self.denylist.deny_sid.assert_called_once_with("sess-7")
self.denylist.deny_idp_sub.assert_not_called()
call = self.db.events.insert.call_args
assert call.args == ("sid:sess-7", "backchannel_logout")
assert call.kwargs["metadata"] == {"sid": "sess-7"}
def test_sub_and_sid_denylists_both(self, client, fake_redis):
response = self._post(client, logout_token_claims(sid="sess-8", jti="bcl-jti-3"))
assert response.status_code == 200
self.denylist.deny_idp_sub.assert_called_once_with("oidc-user-1")
self.denylist.deny_sid.assert_called_once_with("sess-8")
def test_json_body_accepted(self, client, fake_redis):
token = sign_id_token(logout_token_claims(jti="bcl-jti-4"))
response = self._post(client, json={"logout_token": token})
assert response.status_code == 200
def test_missing_token_rejected(self, client, fake_redis):
response = client.post(self.URL)
assert response.status_code == 400
assert response.headers["Cache-Control"] == "no-store"
def test_missing_events_rejected(self, client, fake_redis):
claims = logout_token_claims()
del claims["events"]
response = self._post(client, claims)
assert response.status_code == 400
assert response.get_json() == {"error": "invalid_logout_token"}
assert response.headers["Cache-Control"] == "no-store"
self.denylist.deny_idp_sub.assert_not_called()
def test_nonce_present_rejected(self, client, fake_redis):
response = self._post(client, logout_token_claims(nonce="n-1"))
assert response.status_code == 400
self.denylist.deny_idp_sub.assert_not_called()
def test_bad_signature_rejected(self, client, fake_redis):
rogue_pem = _generate_rsa_pem()
token = sign_id_token(logout_token_claims(), key=rogue_pem)
response = self._post(client, data={"logout_token": token})
assert response.status_code == 400
self.denylist.deny_idp_sub.assert_not_called()
def test_jti_replay_rejected(self, client, fake_redis):
token = sign_id_token(logout_token_claims(jti="bcl-jti-replay"))
first = self._post(client, data={"logout_token": token})
replay = self._post(client, data={"logout_token": token})
assert first.status_code == 200
assert replay.status_code == 400
assert self.denylist.deny_idp_sub.call_count == 1
@pytest.mark.unit
class TestRefreshRoute:
URL = "/api/auth/oidc/refresh"
@pytest.fixture(autouse=True)
def _seams(self, db_mocks):
self.db = db_mocks
with patch("application.api.oidc.routes.denylist") as deny:
deny.is_denied.return_value = False
self.denylist = deny
yield
def _auth(self, token):
return {"Authorization": f"Bearer {token}"}
def _refresh(self, client, token, idp_response=None, idp_status=200):
with patch("application.api.oidc.provider.requests") as mock_requests:
mock_requests.get.side_effect = make_fake_get([PUBLIC_JWK])
mock_requests.post.return_value = Mock(
status_code=idp_status,
text="error",
json=Mock(return_value=idp_response or {}),
)
response = client.post(self.URL, headers=self._auth(token))
return response, mock_requests
def test_happy_path_rotates_refresh_token(self, client, fake_redis):
token = make_session_token()
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
id_claims = id_token_claims(sid="sess-1")
del id_claims["nonce"]
response, mock_requests = self._refresh(
client,
token,
idp_response={
"access_token": "at-2",
"refresh_token": "rt-new",
"id_token": sign_id_token(id_claims),
},
)
assert response.status_code == 200
decoded = jose_jwt.decode(
response.get_json()["token"], JWT_SECRET, algorithms=["HS256"]
)
assert decoded["sub"] == "oidc-user-1"
assert decoded["jti"] != "jti-1"
assert decoded["oidc_sub"] == "oidc-user-1"
assert decoded["oidc_sid"] == "sess-1"
assert "oidc:refresh:jti-1" not in fake_redis.store
assert fake_redis.store[f"oidc:refresh:{decoded['jti']}"] == b"rt-new"
sent = mock_requests.post.call_args.kwargs["data"]
assert sent["grant_type"] == "refresh_token"
assert sent["refresh_token"] == "rt-old"
call = self.db.events.insert.call_args
assert call.args == ("oidc-user-1", "oidc_refresh")
def test_identity_reused_when_no_id_token_returned(self, client, fake_redis):
token = make_session_token(
email="user@example.com", name="OIDC User", oidc_sid="sess-2"
)
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
response, _ = self._refresh(client, token, idp_response={"access_token": "at-2"})
assert response.status_code == 200
decoded = jose_jwt.decode(
response.get_json()["token"], JWT_SECRET, algorithms=["HS256"]
)
assert decoded["sub"] == "oidc-user-1"
assert decoded["email"] == "user@example.com"
assert decoded["oidc_sid"] == "sess-2"
# IdP kept the old refresh token, so it is re-stored under the new jti.
assert fake_redis.store[f"oidc:refresh:{decoded['jti']}"] == b"rt-old"
def test_refresh_denied_when_groups_no_longer_allowed(
self, client, fake_redis, monkeypatch
):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "admins")
token = make_session_token()
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
id_claims = id_token_claims(groups=["users"])
del id_claims["nonce"]
response, _ = self._refresh(
client,
token,
idp_response={
"access_token": "at-2",
"refresh_token": "rt-new",
"id_token": sign_id_token(id_claims),
},
)
assert response.status_code == 401
assert response.get_json() == {"error": "not_authorized"}
# Denial is audited and no renewed session/refresh token exists.
call = self.db.events.insert.call_args
assert call.args[1] == "oidc_login_denied"
assert call.kwargs["metadata"]["via"] == "refresh"
assert not any(k.startswith("oidc:refresh:") for k in fake_redis.store)
def test_refresh_allowed_when_group_still_matches(
self, client, fake_redis, monkeypatch
):
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "users, admins")
token = make_session_token()
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
id_claims = id_token_claims(groups=["users"])
del id_claims["nonce"]
response, _ = self._refresh(
client,
token,
idp_response={
"access_token": "at-2",
"refresh_token": "rt-new",
"id_token": sign_id_token(id_claims),
},
)
assert response.status_code == 200
def test_refresh_without_id_token_skips_group_check(
self, client, fake_redis, monkeypatch
):
# No fresh claims to evaluate — membership was checked at login and
# will be re-checked the next time the IdP returns an id_token.
monkeypatch.setattr(settings, "OIDC_ALLOWED_GROUPS", "admins")
token = make_session_token()
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
response, _ = self._refresh(client, token, idp_response={"access_token": "at-2"})
assert response.status_code == 200
def test_expired_session_rejected(self, client, fake_redis):
token = make_session_token(exp=int(time.time()) - 120)
response = client.post(self.URL, headers=self._auth(token))
assert response.status_code == 401
assert response.get_json() == {"error": "token_expired"}
def test_garbage_token_rejected(self, client, fake_redis):
response = client.post(self.URL, headers=self._auth("not-a-jwt"))
assert response.status_code == 401
assert response.get_json() == {"error": "invalid_token"}
def test_missing_authorization_rejected(self, client, fake_redis):
response = client.post(self.URL)
assert response.status_code == 401
assert response.get_json() == {"error": "invalid_token"}
def test_session_without_jti_rejected(self, client, fake_redis):
token = make_session_token(jti=None)
response = client.post(self.URL, headers=self._auth(token))
assert response.status_code == 401
assert response.get_json() == {"error": "invalid_token"}
def test_denylisted_session_rejected(self, client, fake_redis):
self.denylist.is_denied.return_value = True
token = make_session_token()
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
response = client.post(self.URL, headers=self._auth(token))
assert response.status_code == 401
assert response.get_json() == {"error": "token_revoked"}
assert "oidc:refresh:jti-1" in fake_redis.store
def test_disabled_user_rejected(self, client, fake_redis):
self.db.users.get.return_value = {"user_id": "oidc-user-1", "active": False}
token = make_session_token()
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
response = client.post(self.URL, headers=self._auth(token))
assert response.status_code == 401
assert response.get_json() == {"error": "account_disabled"}
def test_no_stored_refresh_token_404(self, client, fake_redis):
token = make_session_token()
response = client.post(self.URL, headers=self._auth(token))
assert response.status_code == 404
assert response.get_json() == {"error": "no_refresh_token"}
def test_idp_refresh_failure_401(self, client, fake_redis):
token = make_session_token()
fake_redis.store["oidc:refresh:jti-1"] = b"rt-old"
response, _ = self._refresh(client, token, idp_status=400)
assert response.status_code == 401
assert response.get_json() == {"error": "refresh_failed"}
def test_503_when_redis_unavailable(self, client):
token = make_session_token()
with patch("application.api.oidc.routes.get_redis_instance", return_value=None):
response = client.post(self.URL, headers=self._auth(token))
assert response.status_code == 503
assert response.get_json() == {"error": "redis_unavailable"}
+563
View File
@@ -0,0 +1,563 @@
"""Unit tests for the SCIM 2.0 provisioning endpoints (application/api/scim/)."""
import json
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from application.core.settings import settings
TOKEN = "test-scim-token"
AUTH = {"Authorization": f"Bearer {TOKEN}"}
USER_PK = "3f1f2bd6-8a87-4f7e-9c2b-2f76c1a4f0d3"
ERROR_URN = "urn:ietf:params:scim:api:messages:2.0:Error"
LIST_RESPONSE_URN = "urn:ietf:params:scim:api:messages:2.0:ListResponse"
PATCH_OP_URN = "urn:ietf:params:scim:api:messages:2.0:PatchOp"
USER_URN = "urn:ietf:params:scim:schemas:core:2.0:User"
def _user_row(pk=USER_PK, user_id="alice@example.com", active=True):
return {
"id": pk,
"user_id": user_id,
"active": active,
"agent_preferences": {"pinned": [], "shared_with_me": []},
"tool_preferences": {},
"created_at": "2026-06-09T10:00:00+00:00",
"updated_at": "2026-06-09T11:00:00+00:00",
}
@pytest.fixture(scope="module")
def app():
"""Import the Flask app with auth mocked to avoid JWT setup issues."""
with patch("application.app.handle_auth", return_value={"sub": "test_user"}):
from application.app import app as flask_app
flask_app.config["TESTING"] = True
yield flask_app
@pytest.fixture
def client(app):
return app.test_client()
@pytest.fixture
def scim_settings(monkeypatch):
"""Enable SCIM with a known bearer token on the real settings singleton."""
monkeypatch.setattr(settings, "SCIM_ENABLED", True)
monkeypatch.setattr(settings, "SCIM_TOKEN", TOKEN)
@pytest.fixture
def scim_mocks():
"""Patch DB plumbing, repositories, and denylist functions used by the routes."""
with patch("application.api.scim.routes.db_session") as db_session_mock, patch(
"application.api.scim.routes.db_readonly"
) as db_readonly_mock, patch(
"application.api.scim.routes.UsersRepository"
) as users_cls, patch(
"application.api.scim.routes.AuthEventsRepository"
) as audit_cls, patch(
"application.api.scim.routes.deny_user"
) as deny_user_mock, patch(
"application.api.scim.routes.allow_user"
) as allow_user_mock:
conn = MagicMock(name="conn")
db_session_mock.return_value.__enter__.return_value = conn
db_readonly_mock.return_value.__enter__.return_value = conn
yield SimpleNamespace(
users=users_cls.return_value,
audit=audit_cls.return_value,
deny_user=deny_user_mock,
allow_user=allow_user_mock,
conn=conn,
)
@pytest.mark.unit
class TestScimGate:
def test_disabled_returns_scim_404_everywhere(self, client, scim_mocks, monkeypatch):
monkeypatch.setattr(settings, "SCIM_ENABLED", False)
monkeypatch.setattr(settings, "SCIM_TOKEN", TOKEN)
for path in ("/scim/v2/Users", "/scim/v2/ServiceProviderConfig", "/scim/v2/Groups"):
response = client.get(path, headers=AUTH)
assert response.status_code == 404
body = response.get_json()
assert body["schemas"] == [ERROR_URN]
assert body["status"] == "404"
scim_mocks.users.list_paginated.assert_not_called()
def test_enabled_without_token_returns_503(self, client, scim_mocks, monkeypatch):
monkeypatch.setattr(settings, "SCIM_ENABLED", True)
monkeypatch.setattr(settings, "SCIM_TOKEN", None)
response = client.get("/scim/v2/Users", headers=AUTH)
assert response.status_code == 503
body = response.get_json()
assert body["schemas"] == [ERROR_URN]
assert body["status"] == "503"
def test_enabled_with_empty_token_returns_503(self, client, scim_mocks, monkeypatch):
monkeypatch.setattr(settings, "SCIM_ENABLED", True)
monkeypatch.setattr(settings, "SCIM_TOKEN", "")
response = client.get("/scim/v2/Users", headers=AUTH)
assert response.status_code == 503
def test_wrong_bearer_returns_401(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/Users", headers={"Authorization": "Bearer wrong-token"})
assert response.status_code == 401
assert response.get_json()["schemas"] == [ERROR_URN]
scim_mocks.users.list_paginated.assert_not_called()
def test_missing_authorization_returns_401(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/Users")
assert response.status_code == 401
assert response.get_json()["status"] == "401"
def test_non_bearer_scheme_returns_401(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/Users", headers={"Authorization": f"Basic {TOKEN}"})
assert response.status_code == 401
def test_errors_use_scim_content_type(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/Users", headers={"Authorization": "Bearer nope"})
assert response.content_type.startswith("application/scim+json")
@pytest.mark.unit
class TestDiscoveryEndpoints:
def test_service_provider_config_shape(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/ServiceProviderConfig", headers=AUTH)
assert response.status_code == 200
assert response.content_type.startswith("application/scim+json")
body = response.get_json()
assert body["schemas"] == ["urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig"]
assert body["patch"] == {"supported": True}
assert body["bulk"]["supported"] is False
assert body["filter"] == {"supported": True, "maxResults": 200}
assert body["changePassword"] == {"supported": False}
assert body["sort"] == {"supported": False}
assert body["etag"] == {"supported": False}
scheme = body["authenticationSchemes"][0]
assert scheme["type"] == "oauthbearertoken"
assert scheme["name"] == "Bearer Token"
def test_resource_types_advertise_user_only(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/ResourceTypes", headers=AUTH)
assert response.status_code == 200
body = response.get_json()
assert body["schemas"] == [LIST_RESPONSE_URN]
assert body["totalResults"] == 1
resource = body["Resources"][0]
assert resource["name"] == "User"
assert resource["endpoint"] == "/scim/v2/Users"
assert resource["schema"] == USER_URN
def test_schemas_advertise_user_only(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/Schemas", headers=AUTH)
assert response.status_code == 200
body = response.get_json()
assert body["schemas"] == [LIST_RESPONSE_URN]
assert body["totalResults"] == 1
assert body["Resources"][0]["id"] == USER_URN
@pytest.mark.unit
class TestListUsers:
def test_filter_username_eq_passes_value_to_repo(self, client, scim_settings, scim_mocks):
scim_mocks.users.list_paginated.return_value = (1, [_user_row()])
response = client.get(
"/scim/v2/Users",
headers=AUTH,
query_string={"filter": 'userName eq "alice smith@example.com"'},
)
assert response.status_code == 200
scim_mocks.users.list_paginated.assert_called_once_with("alice smith@example.com", 0, 100)
def test_filter_keywords_are_case_insensitive(self, client, scim_settings, scim_mocks):
scim_mocks.users.list_paginated.return_value = (0, [])
response = client.get(
"/scim/v2/Users", headers=AUTH, query_string={"filter": 'UserName EQ "bob@example.com"'}
)
assert response.status_code == 200
scim_mocks.users.list_paginated.assert_called_once_with("bob@example.com", 0, 100)
@pytest.mark.parametrize(
"bad_filter",
[
'userName co "alice"',
'displayName eq "alice"',
'userName eq "a" and active eq true',
"userName eq alice",
],
)
def test_unsupported_filter_rejected(self, client, scim_settings, scim_mocks, bad_filter):
response = client.get("/scim/v2/Users", headers=AUTH, query_string={"filter": bad_filter})
assert response.status_code == 400
body = response.get_json()
assert body["schemas"] == [ERROR_URN]
assert body["scimType"] == "invalidFilter"
scim_mocks.users.list_paginated.assert_not_called()
@pytest.mark.parametrize(
("query", "expected_offset", "expected_limit"),
[
({}, 0, 100),
({"startIndex": "3", "count": "2"}, 2, 2),
({"startIndex": "0", "count": "999"}, 0, 200),
({"startIndex": "-4", "count": "-5"}, 0, 0),
],
)
def test_pagination_offset_limit_math(
self, client, scim_settings, scim_mocks, query, expected_offset, expected_limit
):
scim_mocks.users.list_paginated.return_value = (0, [])
response = client.get("/scim/v2/Users", headers=AUTH, query_string=query)
assert response.status_code == 200
scim_mocks.users.list_paginated.assert_called_once_with(None, expected_offset, expected_limit)
def test_list_response_shape(self, client, scim_settings, scim_mocks):
scim_mocks.users.list_paginated.return_value = (42, [_user_row()])
response = client.get(
"/scim/v2/Users", headers=AUTH, query_string={"startIndex": "5", "count": "1"}
)
assert response.status_code == 200
assert response.content_type.startswith("application/scim+json")
body = response.get_json()
assert body["schemas"] == [LIST_RESPONSE_URN]
assert body["totalResults"] == 42
assert body["startIndex"] == 5
assert body["itemsPerPage"] == 1
resource = body["Resources"][0]
assert resource["schemas"] == [USER_URN]
assert resource["id"] == USER_PK
assert resource["userName"] == "alice@example.com"
assert resource["active"] is True
assert resource["emails"] == [{"value": "alice@example.com", "primary": True}]
assert resource["meta"]["resourceType"] == "User"
assert resource["meta"]["location"] == f"/scim/v2/Users/{USER_PK}"
@pytest.mark.unit
class TestCreateUser:
def test_missing_username_returns_400(self, client, scim_settings, scim_mocks):
response = client.post("/scim/v2/Users", headers=AUTH, json={"active": True})
assert response.status_code == 400
assert response.get_json()["scimType"] == "invalidValue"
scim_mocks.users.create.assert_not_called()
def test_blank_username_returns_400(self, client, scim_settings, scim_mocks):
response = client.post("/scim/v2/Users", headers=AUTH, json={"userName": " "})
assert response.status_code == 400
assert response.get_json()["scimType"] == "invalidValue"
def test_conflict_returns_409(self, client, scim_settings, scim_mocks):
scim_mocks.users.create.return_value = None
response = client.post("/scim/v2/Users", headers=AUTH, json={"userName": "alice@example.com"})
assert response.status_code == 409
body = response.get_json()
assert body["schemas"] == [ERROR_URN]
assert body["scimType"] == "uniqueness"
scim_mocks.audit.insert.assert_not_called()
def test_create_happy_path(self, client, scim_settings, scim_mocks):
scim_mocks.users.create.return_value = _user_row()
response = client.post("/scim/v2/Users", headers=AUTH, json={"userName": "alice@example.com"})
assert response.status_code == 201
assert response.content_type.startswith("application/scim+json")
assert response.headers["Location"].endswith(f"/scim/v2/Users/{USER_PK}")
scim_mocks.users.create.assert_called_once_with("alice@example.com", active=True)
scim_mocks.audit.insert.assert_called_once_with(
"alice@example.com", "scim_created", metadata={"via": "scim"}
)
body = response.get_json()
assert body["id"] == USER_PK
assert body["userName"] == "alice@example.com"
assert body["active"] is True
assert body["emails"] == [{"value": "alice@example.com", "primary": True}]
assert body["meta"]["resourceType"] == "User"
def test_create_honors_active_false(self, client, scim_settings, scim_mocks):
scim_mocks.users.create.return_value = _user_row(active=False)
response = client.post(
"/scim/v2/Users", headers=AUTH, json={"userName": "alice@example.com", "active": False}
)
assert response.status_code == 201
scim_mocks.users.create.assert_called_once_with("alice@example.com", active=False)
assert response.get_json()["active"] is False
def test_create_accepts_scim_content_type(self, client, scim_settings, scim_mocks):
scim_mocks.users.create.return_value = _user_row()
response = client.post(
"/scim/v2/Users",
headers=AUTH,
data=json.dumps({"userName": "alice@example.com"}),
content_type="application/scim+json",
)
assert response.status_code == 201
def test_create_without_email_username_has_no_emails(self, client, scim_settings, scim_mocks):
scim_mocks.users.create.return_value = _user_row(user_id="ldap-user-1")
response = client.post("/scim/v2/Users", headers=AUTH, json={"userName": "ldap-user-1"})
assert response.status_code == 201
assert "emails" not in response.get_json()
def test_audit_failure_does_not_fail_request(self, client, scim_settings, scim_mocks):
scim_mocks.users.create.return_value = _user_row()
scim_mocks.audit.insert.side_effect = RuntimeError("audit down")
response = client.post("/scim/v2/Users", headers=AUTH, json={"userName": "alice@example.com"})
assert response.status_code == 201
@pytest.mark.unit
class TestGetUser:
def test_get_user_found(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row()
response = client.get(f"/scim/v2/Users/{USER_PK}", headers=AUTH)
assert response.status_code == 200
scim_mocks.users.get_by_pk.assert_called_once_with(USER_PK)
body = response.get_json()
assert body["id"] == USER_PK
assert body["userName"] == "alice@example.com"
assert body["meta"]["location"] == f"/scim/v2/Users/{USER_PK}"
def test_get_user_missing_returns_404(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = None
response = client.get(f"/scim/v2/Users/{USER_PK}", headers=AUTH)
assert response.status_code == 404
assert response.get_json()["schemas"] == [ERROR_URN]
def test_get_user_malformed_uuid_returns_404(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = None
response = client.get("/scim/v2/Users/not-a-uuid", headers=AUTH)
assert response.status_code == 404
assert response.get_json()["status"] == "404"
@pytest.mark.unit
class TestReplaceUser:
def test_username_change_rejected_as_mutability(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row()
response = client.put(
f"/scim/v2/Users/{USER_PK}",
headers=AUTH,
json={"userName": "bob@example.com", "active": True},
)
assert response.status_code == 400
assert response.get_json()["scimType"] == "mutability"
scim_mocks.users.set_active.assert_not_called()
def test_put_active_true_triggers_allow(self, client, scim_settings, scim_mocks):
row = _user_row(active=False)
scim_mocks.users.get_by_pk.return_value = row
scim_mocks.users.set_active.return_value = {**row, "active": True}
response = client.put(
f"/scim/v2/Users/{USER_PK}",
headers=AUTH,
json={"userName": "alice@example.com", "active": True},
)
assert response.status_code == 200
assert response.get_json()["active"] is True
scim_mocks.users.set_active.assert_called_once_with(USER_PK, True)
scim_mocks.allow_user.assert_called_once_with("alice@example.com")
scim_mocks.deny_user.assert_not_called()
scim_mocks.audit.insert.assert_called_once_with(
"alice@example.com", "scim_reactivated", metadata={"via": "scim"}
)
def test_put_active_false_triggers_deny(self, client, scim_settings, scim_mocks):
row = _user_row(active=True)
scim_mocks.users.get_by_pk.return_value = row
scim_mocks.users.set_active.return_value = {**row, "active": False}
response = client.put(
f"/scim/v2/Users/{USER_PK}",
headers=AUTH,
json={"userName": "alice@example.com", "active": False},
)
assert response.status_code == 200
assert response.get_json()["active"] is False
scim_mocks.users.set_active.assert_called_once_with(USER_PK, False)
scim_mocks.deny_user.assert_called_once_with("alice@example.com")
scim_mocks.audit.insert.assert_called_once_with(
"alice@example.com", "scim_deactivated", metadata={"via": "scim"}
)
def test_put_unchanged_active_has_no_side_effects(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row(active=True)
response = client.put(
f"/scim/v2/Users/{USER_PK}",
headers=AUTH,
json={"userName": "alice@example.com", "active": True},
)
assert response.status_code == 200
scim_mocks.users.set_active.assert_not_called()
scim_mocks.deny_user.assert_not_called()
scim_mocks.allow_user.assert_not_called()
scim_mocks.audit.insert.assert_not_called()
def test_put_missing_user_returns_404(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = None
response = client.put(
f"/scim/v2/Users/{USER_PK}", headers=AUTH, json={"userName": "x", "active": True}
)
assert response.status_code == 404
@pytest.mark.unit
class TestPatchUser:
def _patch(self, client, operations):
return client.patch(
f"/scim/v2/Users/{USER_PK}",
headers=AUTH,
json={"schemas": [PATCH_OP_URN], "Operations": operations},
)
def _expect_deactivated(self, scim_mocks, response):
assert response.status_code == 200
assert response.get_json()["active"] is False
scim_mocks.users.set_active.assert_called_once_with(USER_PK, False)
scim_mocks.deny_user.assert_called_once_with("alice@example.com")
scim_mocks.audit.insert.assert_called_once_with(
"alice@example.com", "scim_deactivated", metadata={"via": "scim"}
)
def test_replace_with_path_deactivates(self, client, scim_settings, scim_mocks):
row = _user_row(active=True)
scim_mocks.users.get_by_pk.return_value = row
scim_mocks.users.set_active.return_value = {**row, "active": False}
response = self._patch(client, [{"op": "replace", "path": "active", "value": False}])
self._expect_deactivated(scim_mocks, response)
def test_replace_without_path_object_value_deactivates(self, client, scim_settings, scim_mocks):
row = _user_row(active=True)
scim_mocks.users.get_by_pk.return_value = row
scim_mocks.users.set_active.return_value = {**row, "active": False}
response = self._patch(client, [{"op": "replace", "value": {"active": False}}])
self._expect_deactivated(scim_mocks, response)
def test_replace_with_string_false_deactivates(self, client, scim_settings, scim_mocks):
row = _user_row(active=True)
scim_mocks.users.get_by_pk.return_value = row
scim_mocks.users.set_active.return_value = {**row, "active": False}
response = self._patch(client, [{"op": "Replace", "path": "active", "value": "False"}])
self._expect_deactivated(scim_mocks, response)
def test_replace_with_string_true_reactivates(self, client, scim_settings, scim_mocks):
row = _user_row(active=False)
scim_mocks.users.get_by_pk.return_value = row
scim_mocks.users.set_active.return_value = {**row, "active": True}
response = self._patch(client, [{"op": "replace", "path": "active", "value": "true"}])
assert response.status_code == 200
assert response.get_json()["active"] is True
scim_mocks.users.set_active.assert_called_once_with(USER_PK, True)
scim_mocks.allow_user.assert_called_once_with("alice@example.com")
scim_mocks.audit.insert.assert_called_once_with(
"alice@example.com", "scim_reactivated", metadata={"via": "scim"}
)
def test_bogus_path_rejected(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row()
response = self._patch(client, [{"op": "replace", "path": "displayName", "value": "X"}])
assert response.status_code == 400
assert response.get_json()["scimType"] == "invalidPath"
scim_mocks.users.set_active.assert_not_called()
def test_bogus_op_rejected(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row()
response = self._patch(client, [{"op": "add", "path": "active", "value": True}])
assert response.status_code == 400
assert response.get_json()["scimType"] == "invalidPath"
def test_invalid_active_value_rejected(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row()
response = self._patch(client, [{"op": "replace", "path": "active", "value": "maybe"}])
assert response.status_code == 400
assert response.get_json()["scimType"] == "invalidValue"
def test_unchanged_active_is_noop(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row(active=True)
response = self._patch(client, [{"op": "replace", "path": "active", "value": True}])
assert response.status_code == 200
scim_mocks.users.set_active.assert_not_called()
scim_mocks.deny_user.assert_not_called()
scim_mocks.audit.insert.assert_not_called()
def test_no_path_value_without_active_is_noop(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row(active=True)
response = self._patch(client, [{"op": "replace", "value": {"displayName": "X"}}])
assert response.status_code == 200
scim_mocks.users.set_active.assert_not_called()
def test_missing_operations_rejected(self, client, scim_settings, scim_mocks):
response = client.patch(
f"/scim/v2/Users/{USER_PK}", headers=AUTH, json={"schemas": [PATCH_OP_URN]}
)
assert response.status_code == 400
def test_patch_missing_user_returns_404(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = None
response = self._patch(client, [{"op": "replace", "path": "active", "value": False}])
assert response.status_code == 404
@pytest.mark.unit
class TestDeleteUser:
def test_delete_deactivates_and_returns_204(self, client, scim_settings, scim_mocks):
row = _user_row(active=True)
scim_mocks.users.get_by_pk.return_value = row
scim_mocks.users.set_active.return_value = {**row, "active": False}
response = client.delete(f"/scim/v2/Users/{USER_PK}", headers=AUTH)
assert response.status_code == 204
assert response.data == b""
scim_mocks.users.set_active.assert_called_once_with(USER_PK, False)
scim_mocks.deny_user.assert_called_once_with("alice@example.com")
scim_mocks.audit.insert.assert_called_once_with(
"alice@example.com", "scim_deactivated", metadata={"via": "scim"}
)
def test_delete_already_inactive_has_no_side_effects(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = _user_row(active=False)
response = client.delete(f"/scim/v2/Users/{USER_PK}", headers=AUTH)
assert response.status_code == 204
scim_mocks.users.set_active.assert_not_called()
scim_mocks.deny_user.assert_not_called()
def test_delete_missing_user_returns_404(self, client, scim_settings, scim_mocks):
scim_mocks.users.get_by_pk.return_value = None
response = client.delete(f"/scim/v2/Users/{USER_PK}", headers=AUTH)
assert response.status_code == 404
@pytest.mark.unit
class TestGroups:
def test_get_groups_returns_empty_list(self, client, scim_settings, scim_mocks):
response = client.get("/scim/v2/Groups", headers=AUTH)
assert response.status_code == 200
body = response.get_json()
assert body["schemas"] == [LIST_RESPONSE_URN]
assert body["totalResults"] == 0
assert body["Resources"] == []
def test_post_groups_not_implemented(self, client, scim_settings, scim_mocks):
response = client.post("/scim/v2/Groups", headers=AUTH, json={"displayName": "Team"})
assert response.status_code == 501
body = response.get_json()
assert body["schemas"] == [ERROR_URN]
assert body["detail"] == "Group provisioning is not supported"
@pytest.mark.parametrize("method", ["put", "patch", "delete"])
def test_group_mutations_not_implemented(self, client, scim_settings, scim_mocks, method):
response = getattr(client, method)("/scim/v2/Groups/some-group", headers=AUTH)
assert response.status_code == 501
assert response.get_json()["detail"] == "Group provisioning is not supported"