mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-03 11:11:58 +00:00
feat: oidc renewal, groups, logout, scim
This commit is contained in:
1 parent
8bc777a428
commit
054f0f1b8b
29 files changed
+3392
-155
No files matched your search
@@ -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;")
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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"
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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,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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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">
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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> =>
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
}
|
||||
@@ -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": {
|
||||
|
||||
@@ -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();
|
||||
});
|
||||
});
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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}
|
||||
@@ -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
@@ -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"}
|
||||
@@ -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"
|
||||
Reference in new issue
Block a user