Files
flowdeck/app/services/sso_provisioning.py
T
bruno ffa1fa89ab
FlowDeck CI / lint (push) Successful in 1m51s
FlowDeck CI / docker (push) Canceled after 0s
FlowDeck CI / test (push) Canceled after 10m12s
fix: A25 + A21 (partiel) — plus d'exception muque, transaction protégée (v7.3.8)
- A25 — 84 `except Exception: pass/…` → `logger.exception("<fonction>")`
  (19 fichiers : api_v2 30, dashboard 10, board 7, sites 5, workspace 5,
  api_v2_helpers 5, …) ; `logger` ajouté là où il manquait (api_v2_helpers,
  sites + `import logging`)
- A25 critique — les `try` autour de `materialize_properties` supprimés dans
  `create_collection_v2` ET `apply_db_template_v2` : un échec interrompt la
  transaction au lieu de commiter une collection sans schéma
- test `test_collection_rollback_when_materialize_fails` (Bearer v2, monkeypatch
  qui lève, assertions : RuntimeError + 0 ligne commitée)
- A21 partiel — `PRAGMA busy_timeout=5000` dans `get_conn()` (point d'entrée
  unique) ; commentaire `ponytail:` : le wrapper async + les 510 call sites
  restent à migrer module par module
- suite **1028/1028** · `ruff check app tests` OK
2026-10-01 08:16:42 -04:00

784 lines
29 KiB
Python

"""v6.7.0 — SSO provisioning: config store, auto-provisioning, group mapping.
Single source of truth for the SSO configuration (``sso_config`` table, with
an ``SSO_*`` environment fallback for bootstrap installs) and for what
happens when an IdP says "this is [email protected]":
1. resolve the local account (by email → merge, else by login),
2. create it when ``auto_provision`` is on, else reject with an audit row,
3. sync attributes + map SSO groups to workspace roles,
4. hand back the user dict so the caller can mint a session.
Secrets at rest: ``client_secret`` and the generated SP private key are
encrypted with a Fernet key derived from ``app_secret_key``.
"""
from __future__ import annotations
import hashlib
import json
import logging
import secrets
import time
import urllib.parse
from app.config import settings
logger = logging.getLogger(__name__)
VALID_PROVIDER_TYPES = ("saml", "oidc")
class SSOConfigError(Exception):
"""Invalid SSO configuration payload (message shown to the admin)."""
class SSOProvisioningError(Exception):
"""A login was rejected (no local account, missing attributes…)."""
# ── Secrets at rest ────────────────────────────────────────────────────────
def _fernet():
from cryptography.fernet import Fernet
key = hashlib.sha256((settings.app_secret_key or "flowdeck").encode()).digest()
import base64
return Fernet(base64.urlsafe_b64encode(key))
def encrypt_secret(value: str) -> str:
if not value:
return ""
return _fernet().encrypt(value.encode()).decode()
def decrypt_secret(value: str) -> str:
if not value:
return ""
try:
return _fernet().decrypt(value.encode()).decode()
except Exception:
return "" # key rotated / not ours — treat as unset
# ── Config store ───────────────────────────────────────────────────────────
def _default_mapping(provider_type: str) -> dict:
if provider_type == "oidc":
from app.auth.providers.oidc_provider import DEFAULT_OIDC_MAPPING
return dict(DEFAULT_OIDC_MAPPING)
from app.auth.providers.saml_provider import DEFAULT_SAML_MAPPING
return dict(DEFAULT_SAML_MAPPING)
def _env_config() -> dict | None:
"""Bootstrap config from ``SSO_*`` env vars (design doc §3.3).
Only used when the table holds no active row — the Settings UI always
wins once an admin saved a configuration.
"""
provider_type = (settings.sso_provider or "").strip().lower()
if provider_type not in VALID_PROVIDER_TYPES:
return None
cfg = {
"id": 0,
"provider_type": provider_type,
"name": settings.sso_name or "Company SSO",
"entity_id": settings.sso_entity_id,
"sso_url": settings.sso_sso_url,
"slo_url": settings.sso_slo_url,
"x509_certificate": settings.sso_x509_certificate,
"issuer_url": settings.sso_issuer_url,
"client_id": settings.sso_client_id,
"client_secret": settings.sso_client_secret,
"scope": settings.sso_scope,
"attribute_mapping": settings.sso_attribute_mapping,
"groups_mapping": settings.sso_groups_mapping,
"auto_provision": int(settings.sso_auto_provision),
"sso_only": int(settings.sso_only),
"sign_requests": int(settings.sso_sign_requests),
"default_workspace_id": settings.sso_default_workspace_id,
"sp_private_key": "",
"sp_certificate": "",
"workspace_id": None,
"active": 1,
"_source": "env",
}
if provider_type == "saml" and (not cfg["entity_id"] or not cfg["sso_url"]):
return None
if provider_type == "oidc" and (not cfg["issuer_url"] or not cfg["client_id"]):
return None
return cfg
def get_sso_config(require_active: bool = True) -> dict | None:
"""Active SSO config as a dict (DB row, else env fallback)."""
row = _raw_row(require_active=require_active)
if row:
cfg = dict(row)
cfg["_source"] = "db"
return cfg
if not require_active:
return _env_config()
return _env_config()
def _raw_row(require_active: bool = True) -> dict | None:
"""Raw ``sso_config`` row (``client_secret`` still encrypted, ``_source`` unset)."""
from app.db import get_conn
try:
with get_conn() as conn:
where = "WHERE active=1" if require_active else ""
row = conn.execute(
f"SELECT * FROM sso_config {where} ORDER BY id LIMIT 1"
).fetchone()
except Exception: # table missing (very old install) → env only
return None
return dict(row) if row else None
def client_secret_value(cfg: dict) -> str:
"""Plaintext OIDC client secret (decrypted for DB rows, raw for env)."""
raw = (cfg or {}).get("client_secret") or ""
if not raw:
return ""
if (cfg or {}).get("_source") == "env":
return raw
return decrypt_secret(raw)
def _json_field(value, fallback):
if isinstance(value, (dict, list)):
return value
try:
parsed = json.loads(value or "")
return parsed if isinstance(parsed, type(fallback)) else fallback
except Exception:
return fallback
def _strict_json(value, expected, field: str):
"""Parse a payload field and reject wrong shapes (before normalize,
which would otherwise silently coerce ``"[]"`` → ``{}``)."""
if value is None or value == "":
return expected()
if isinstance(value, (dict, list)):
parsed = value
else:
try:
parsed = json.loads(value)
except Exception as exc:
raise SSOConfigError(f"{field} must be valid JSON") from exc
if not isinstance(parsed, expected):
kind = "object" if expected is dict else "array"
raise SSOConfigError(f"{field} must be a JSON {kind}")
return parsed
def normalize_config(cfg: dict) -> dict:
"""Parse JSON columns + fill defaults (single place for every consumer)."""
out = dict(cfg)
out["attribute_mapping"] = _json_field(out.get("attribute_mapping"), {})
out["groups_mapping"] = _json_field(out.get("groups_mapping"), [])
if not out["attribute_mapping"]:
out["attribute_mapping"] = _default_mapping(out.get("provider_type", "saml"))
for key in ("entity_id", "sso_url", "slo_url", "x509_certificate", "issuer_url",
"client_id", "client_secret", "scope", "name"):
out[key] = (out.get(key) or "").strip()
for key in ("auto_provision", "sso_only", "sign_requests", "active"):
out[key] = int(out.get(key) or 0)
return out
def validate_config_payload(payload: dict) -> dict:
"""Validate + sanitize an admin payload. Raises ``SSOConfigError``."""
provider_type = str(payload.get("provider_type") or "").strip().lower()
if provider_type not in VALID_PROVIDER_TYPES:
raise SSOConfigError(f"provider_type must be one of {', '.join(VALID_PROVIDER_TYPES)}")
payload = dict(payload)
payload["attribute_mapping"] = _strict_json(
payload.get("attribute_mapping"), dict, "attribute_mapping"
)
payload["groups_mapping"] = _strict_json(
payload.get("groups_mapping"), list, "groups_mapping"
)
cfg = normalize_config({**payload, "provider_type": provider_type})
if provider_type == "saml":
for field in ("entity_id", "sso_url"):
if not cfg[field]:
raise SSOConfigError(f"SAML requires '{field}'")
for field, url in (("sso_url", cfg["sso_url"]), ("slo_url", cfg["slo_url"])):
if url and not url.startswith(("http://", "https://")):
raise SSOConfigError(f"'{field}' must be an http(s) URL")
cert = cfg["x509_certificate"].strip()
if cert and "BEGIN CERTIFICATE" not in cert:
raise SSOConfigError("x509_certificate must be a PEM certificate")
if not cert:
raise SSOConfigError("SAML requires the IdP signing certificate (x509_certificate)")
cfg["x509_certificate"] = cert
else:
if not cfg["issuer_url"] or not cfg["client_id"]:
raise SSOConfigError("OIDC requires 'issuer_url' and 'client_id'")
if not cfg["issuer_url"].startswith(("http://", "https://")):
raise SSOConfigError("'issuer_url' must be an http(s) URL")
if not isinstance(cfg["attribute_mapping"], dict):
raise SSOConfigError("attribute_mapping must be a JSON object")
if not isinstance(cfg["groups_mapping"], list):
raise SSOConfigError("groups_mapping must be a JSON array")
for entry in cfg["groups_mapping"]:
if not isinstance(entry, dict) or "sso_group" not in entry:
raise SSOConfigError("groups_mapping entries need at least an 'sso_group' key")
ws = cfg.get("default_workspace_id")
cfg["default_workspace_id"] = int(ws) if ws not in (None, "", 0) else None
return cfg
def save_sso_config(payload: dict, created_by: int | None = None) -> dict:
"""Create or replace the single SSO configuration (idempotent)."""
from app.db import get_conn
cfg = validate_config_payload(payload)
columns = {
"provider_type": cfg["provider_type"],
"name": cfg.get("name") or "Company SSO",
"entity_id": cfg["entity_id"],
"sso_url": cfg["sso_url"],
"slo_url": cfg["slo_url"],
"x509_certificate": cfg["x509_certificate"],
"issuer_url": cfg["issuer_url"],
"client_id": cfg["client_id"],
"scope": cfg.get("scope") or "openid profile email",
"attribute_mapping": json.dumps(cfg["attribute_mapping"]),
"groups_mapping": json.dumps(cfg["groups_mapping"]),
"auto_provision": cfg["auto_provision"],
"sso_only": cfg["sso_only"],
"sign_requests": cfg["sign_requests"],
"default_workspace_id": cfg["default_workspace_id"],
"active": 1,
"updated_at": str(int(time.time())),
}
# Secret handling: a blank incoming secret keeps the stored one (the raw
# row still holds the Fernet blob — never re-encrypt a decrypted value).
existing = _raw_row() or {}
if "client_secret" in cfg:
incoming = str(cfg.get("client_secret") or "").strip()
if incoming:
columns["client_secret"] = encrypt_secret(incoming)
else:
columns["client_secret"] = existing.get("client_secret") or ""
# SP keypair: keep an existing one, generate one for SAML if missing.
sp_key = existing.get("sp_private_key") or ""
sp_cert = existing.get("sp_certificate") or ""
if cfg["provider_type"] == "saml" and not (sp_key and sp_cert):
sp_key, sp_cert = generate_sp_keypair()
columns["sp_private_key"] = sp_key
columns["sp_certificate"] = sp_cert
if created_by:
columns["created_by"] = created_by
with get_conn() as conn:
row = conn.execute("SELECT id FROM sso_config ORDER BY id LIMIT 1").fetchone()
if row:
sets = ", ".join(f"{k}=?" for k in columns)
conn.execute(f"UPDATE sso_config SET {sets} WHERE id=?", (*columns.values(), row["id"]))
cfg_id = row["id"]
else:
keys = ", ".join(columns)
placeholders = ", ".join("?" for _ in columns)
cur = conn.execute(
f"INSERT INTO sso_config ({keys}) VALUES ({placeholders})", tuple(columns.values())
)
cfg_id = cur.lastrowid
conn.commit()
saved = get_sso_config(require_active=False)
saved["id"] = cfg_id
return saved
def delete_sso_config() -> bool:
"""Disable SSO entirely (local logins keep working)."""
from app.db import get_conn
with get_conn() as conn:
cur = conn.execute("UPDATE sso_config SET active=0, updated_at=?", (str(int(time.time())),))
conn.commit()
return cur.rowcount > 0
def public_config_view(cfg: dict | None) -> dict:
"""Config for the admin UI — secrets never leave the server."""
if not cfg:
return {"configured": False}
cfg = normalize_config(cfg)
return {
"configured": True,
"id": cfg.get("id"),
"source": cfg.get("_source", "db"),
"provider_type": cfg["provider_type"],
"name": cfg.get("name") or "Company SSO",
"entity_id": cfg["entity_id"],
"sso_url": cfg["sso_url"],
"slo_url": cfg["slo_url"],
"x509_certificate": cfg["x509_certificate"],
"issuer_url": cfg["issuer_url"],
"client_id": cfg["client_id"],
"client_secret_set": bool(client_secret_value(cfg)),
"scope": cfg.get("scope") or "openid profile email",
"attribute_mapping": cfg["attribute_mapping"],
"groups_mapping": cfg["groups_mapping"],
"auto_provision": bool(cfg["auto_provision"]),
"sso_only": bool(cfg["sso_only"]),
"sign_requests": bool(cfg["sign_requests"]),
"default_workspace_id": cfg.get("default_workspace_id"),
"sp_certificate": cfg.get("sp_certificate") or "",
"active": bool(cfg.get("active", 1)),
"provisioned_users": provisioned_count(),
}
def is_sso_only(cfg: dict | None = None) -> bool:
"""True when local login must be refused (design §7.1 / §4.2)."""
cfg = cfg if cfg is not None else get_sso_config()
return bool(cfg and normalize_config(cfg).get("sso_only"))
def provisioned_count() -> int:
from app.db import get_conn
try:
with get_conn() as conn:
row = conn.execute(
"SELECT COUNT(*) AS n FROM users WHERE auth_method IN ('saml','oidc')"
).fetchone()
return int(row["n"] if row is not None else 0)
except Exception:
return 0
def generate_sp_keypair() -> tuple[str, str]:
"""RSA-2048 key + self-signed certificate for the SP (metadata + signing)."""
import datetime
from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.x509.oid import NameOID
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
name = x509.Name([
x509.NameAttribute(NameOID.COMMON_NAME, f"flowdeck-sp-{secrets.token_hex(4)}"),
])
now = datetime.datetime.now(datetime.UTC)
cert = (
x509.CertificateBuilder()
.subject_name(name)
.issuer_name(name)
.public_key(key.public_key())
.serial_number(x509.random_serial_number())
.not_valid_before(now - datetime.timedelta(days=1))
.not_valid_after(now + datetime.timedelta(days=3650))
.add_extension(x509.BasicConstraints(ca=False, path_length=None), critical=True)
.sign(key, hashes.SHA256())
)
priv = key.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
).decode()
public = cert.public_bytes(serialization.Encoding.PEM).decode()
return priv, public
def ensure_sp_keypair(cfg: dict) -> dict:
"""Guarantee the SAML config carries an SP keypair (generates + persists)."""
if cfg.get("provider_type") != "saml":
return cfg
if cfg.get("sp_private_key") and cfg.get("sp_certificate"):
return cfg
from app.db import get_conn
priv, cert = generate_sp_keypair()
try:
with get_conn() as conn:
conn.execute(
"UPDATE sso_config SET sp_private_key=?, sp_certificate=? WHERE id=?",
(priv, cert, cfg.get("id")),
)
conn.commit()
except Exception as err: # env-sourced config has no row to update
logger.debug("SP keypair not persisted: %s", err)
cfg = dict(cfg)
cfg["sp_private_key"], cfg["sp_certificate"] = priv, cert
return cfg
cfg = dict(cfg)
cfg["sp_private_key"], cfg["sp_certificate"] = priv, cert
return cfg
# ── Audit ──────────────────────────────────────────────────────────────────
def log_sso_login(
*,
user_id: int | None,
provider_type: str,
provider_name: str,
identifier: str,
request,
success: bool,
error: str = "",
) -> None:
"""Write one ``sso_login_history`` row (failures included — design §5.2)."""
ip = request.client.host if request is not None and getattr(request, "client", None) else ""
ua = (request.headers.get("user-agent", "") if request is not None else "")[:500]
try:
from app.db import get_conn
with get_conn() as conn:
conn.execute(
"""INSERT INTO sso_login_history
(user_id, provider_type, provider_name, sso_identifier,
ip_address, user_agent, success, error_message)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
(user_id, provider_type, provider_name, (identifier or "")[:320], ip, ua,
1 if success else 0, (error or "")[:500]),
)
conn.commit()
except Exception as err: # audit must never break the login path
logger.warning("sso_login_history write failed: %s", err)
# ── Group mapping ──────────────────────────────────────────────────────────
def sync_sso_groups(user_id: int, sso_groups: list[str], cfg: dict) -> list[int]:
"""Apply ``groups_mapping`` → ``workspace_members.role``. Returns touched ws ids."""
from app.db import get_conn
cfg = normalize_config(cfg)
mappings = cfg.get("groups_mapping") or []
wanted = {g.strip().lower() for g in sso_groups if g and str(g).strip()}
touched: list[int] = []
with get_conn() as conn:
for entry in mappings:
group_name = str(entry.get("sso_group") or "").strip().lower()
if not group_name or group_name not in wanted:
continue
ws_id = entry.get("workspace_id") or cfg.get("default_workspace_id")
if not ws_id:
continue
role = str(entry.get("workspace_role") or "editor").strip() or "editor"
if role not in ("owner", "admin", "editor", "viewer"):
role = "editor"
conn.execute(
"""INSERT INTO workspace_members (workspace_id, user_id, role)
VALUES (?, ?, ?)
ON CONFLICT(workspace_id, user_id) DO UPDATE SET role=excluded.role""",
(int(ws_id), user_id, role),
)
touched.append(int(ws_id))
# Default workspace: every SSO user lands there as a plain member.
default_ws = cfg.get("default_workspace_id")
if default_ws:
conn.execute(
"""INSERT OR IGNORE INTO workspace_members (workspace_id, user_id, role)
VALUES (?, ?, 'editor')""",
(int(default_ws), user_id),
)
if int(default_ws) not in touched:
touched.append(int(default_ws))
conn.commit()
return touched
def force_sync_all_groups() -> dict:
"""Re-apply the group mapping for every SSO user (``POST /api/v2/sso/sync``)."""
from app.db import get_conn
cfg = get_sso_config()
if not cfg:
raise SSOProvisioningError("No SSO configuration")
cfg = normalize_config(cfg)
mapping = cfg.get("attribute_mapping") or {}
groups_source = mapping.get("groups", "groups")
updated = 0
with get_conn() as conn:
rows = conn.execute(
"SELECT id, auth_method FROM users WHERE auth_method IN ('saml','oidc')"
).fetchall()
for row in rows:
groups = _stored_groups(row["id"], groups_source, cfg)
if sync_sso_groups(row["id"], groups, cfg):
updated += 1
return {"users": len(rows), "updated": updated}
def _stored_groups(user_id: int, source: str, cfg: dict) -> list[str]:
"""Groups seen at the last login of that user (stored in attribute sync)."""
try:
from app.db import get_conn
with get_conn() as conn:
row = conn.execute(
"SELECT sso_identifier FROM sso_login_history "
"WHERE user_id=? AND success=1 ORDER BY id DESC LIMIT 1",
(user_id,),
).fetchone()
if not row or not row["sso_identifier"]:
return []
raw = row["sso_identifier"]
if "|" in raw:
ident, _, groups_json = raw.partition("|")
groups = json.loads(groups_json or "[]")
return [str(g) for g in groups] if isinstance(groups, list) else []
return []
except Exception:
return []
# ── Auto-provisioning ──────────────────────────────────────────────────────
def _unique_login(conn, base: str) -> str:
candidate = base
n = 1
while conn.execute("SELECT 1 FROM users WHERE login=?", (candidate,)).fetchone():
n += 1
candidate = f"{base}_{n}"
return candidate
def identity_from_saml(identity, cfg: dict) -> dict:
"""Apply the SAML attribute mapping to a validated assertion."""
cfg = normalize_config(cfg)
mapping = cfg.get("attribute_mapping") or {}
out = {
"login": identity.resolve(mapping.get("login", "nameid")),
"email": identity.resolve(mapping.get("email", "nameid")),
"full_name": identity.resolve(mapping.get("full_name", "displayName")),
"avatar_url": identity.resolve(mapping.get("avatar_url", "avatar")),
"name_id": identity.name_id,
}
groups = identity.resolve(mapping.get("groups", "groups"))
if groups:
# Multi-valued SAML attribute: take every value of the resolved source.
source = mapping.get("groups", "groups")
values = identity.attributes.get(source) or identity.friendly_attributes.get(source) or [groups]
out["groups"] = [str(v).strip() for v in values if v and str(v).strip()]
else:
out["groups"] = []
out["email"] = (out["email"] or "").strip().lower()
if "@" not in out["email"]:
# NameID may be a persistent opaque id — fall back to login when it is
# an email, otherwise leave empty (login will carry the identity).
out["email"] = out["email"] if "@" in (out["login"] or "") else ""
if not out["full_name"]:
out["full_name"] = out["email"] or out["login"]
out["login"] = out["login"] or out["email"] or f"sso_{identity.name_id[:32]}"
return out
def handle_sso_login(identity: dict, *, provider_type: str, cfg: dict, request) -> dict:
"""Resolve/create the local user for an SSO identity. Returns the user dict.
Raises ``SSOProvisioningError`` when the login must be refused (the
caller writes the audit row).
"""
from app.db import get_conn
cfg = normalize_config(cfg)
email = (identity.get("email") or "").strip().lower()
login_hint = (identity.get("login") or "").strip()
if not email and not login_hint:
raise SSOProvisioningError(
"SSO assertion carries no usable email/login — check the attribute mapping"
)
with get_conn() as conn:
user = None
if email:
user = conn.execute(
"SELECT * FROM users WHERE lower(email)=? AND email!='' ORDER BY id LIMIT 1",
(email,),
).fetchone()
if not user and login_hint:
user = conn.execute("SELECT * FROM users WHERE login=?", (login_hint,)).fetchone()
if user:
# §7.1 — email match → merge: the existing account is reused and
# tagged with the SSO method (no duplicate account).
updates, params = [], []
if identity.get("full_name"):
updates.append("full_name=?")
params.append(identity["full_name"])
if email:
updates.append("email=?")
params.append(email)
if identity.get("avatar_url"):
updates.append("avatar_url=?")
params.append(identity["avatar_url"])
updates.append("auth_method=?")
params.append(provider_type)
updates.append("last_login=?")
params.append(str(time.time()))
params.append(user["id"])
conn.execute(f"UPDATE users SET {', '.join(updates)} WHERE id=?", params)
conn.commit()
row = conn.execute("SELECT * FROM users WHERE id=?", (user["id"],)).fetchone()
else:
if not cfg.get("auto_provision"):
raise SSOProvisioningError(
"No local account for this SSO identity and auto-provisioning is disabled"
)
base_login = login_hint or email
login = _unique_login(conn, base_login)
conn.execute(
"""INSERT INTO users
(login, full_name, email, avatar_url, auth_method, is_admin, last_login)
VALUES (?, ?, ?, ?, ?, 0, ?)""",
(
login,
identity.get("full_name") or email or login,
email,
identity.get("avatar_url") or "",
provider_type,
str(time.time()),
),
)
conn.commit()
row = conn.execute("SELECT * FROM users WHERE login=?", (login,)).fetchone()
if not row:
raise SSOProvisioningError("Could not create or load the SSO user")
user_dict = dict(row)
sync_sso_groups(user_dict["id"], identity.get("groups") or [], cfg)
return user_dict
def sso_identifier_field(identity: dict) -> str:
"""Audit identifier: ``nameid|["groups",...]`` (groups kept for re-sync)."""
ident = identity.get("name_id") or identity.get("email") or identity.get("login") or ""
groups = identity.get("groups") or []
if groups:
return f"{ident}|{json.dumps(groups)}"
return ident
def safe_next_path(candidate: str | None) -> str:
"""Sanitize the post-login redirect target (open-redirect guard)."""
if not candidate:
return "/workspaces"
candidate = str(candidate)
if not candidate.startswith("/") or candidate.startswith("//"):
return "/workspaces"
parsed = urllib.parse.urlsplit(candidate)
if parsed.scheme or parsed.netloc:
return "/workspaces"
return candidate
# ── Anti-replay request store ──────────────────────────────────────────────
REQUEST_TTL_SECONDS = 600 # AuthnRequest / OIDC state lifetime
def create_request(kind: str, *, request_id: str, relay_state: str = "",
code_verifier: str = "", next_path: str = "/workspaces") -> None:
"""Store a single-use SSO request (AuthnRequest id / OIDC state)."""
from app.db import get_conn
purge_stale_requests()
with get_conn() as conn:
conn.execute(
"""INSERT OR REPLACE INTO sso_requests
(id, kind, relay_state, code_verifier, next_path, used, created_at)
VALUES (?, ?, ?, ?, ?, 0, CURRENT_TIMESTAMP)""",
(request_id, kind, relay_state, code_verifier, safe_next_path(next_path)),
)
conn.commit()
def consume_request(kind: str, request_id: str, relay_state: str = "") -> dict | None:
"""Atomically consume a request. Returns the row, or None (replay/unknown)."""
if not request_id:
return None
from app.db import get_conn
with get_conn() as conn:
conn.execute(
"DELETE FROM sso_requests WHERE created_at < datetime('now', ?)",
(f"-{REQUEST_TTL_SECONDS} seconds",),
)
row = conn.execute(
"SELECT * FROM sso_requests WHERE id=? AND kind=? AND used=0",
(request_id, kind),
).fetchone()
if not row:
return None
if relay_state and row["relay_state"] and not secrets.compare_digest(
row["relay_state"], relay_state
):
return None
cur = conn.execute(
"UPDATE sso_requests SET used=1 WHERE id=? AND used=0", (request_id,)
)
conn.commit()
if cur.rowcount != 1:
return None
return dict(row)
def peek_request(kind: str, request_id: str) -> dict | None:
"""Read a request without consuming it (CSRF check before heavy validation)."""
if not request_id:
return None
from app.db import get_conn
with get_conn() as conn:
row = conn.execute(
"SELECT * FROM sso_requests WHERE id=? AND kind=? AND used=0",
(request_id, kind),
).fetchone()
return dict(row) if row else None
def was_consumed(kind: str, request_id: str) -> bool:
"""True when this single-use request id was already spent (replay)."""
if not request_id:
return False
from app.db import get_conn
with get_conn() as conn:
row = conn.execute(
"SELECT 1 FROM sso_requests WHERE id=? AND kind=? AND used=1",
(request_id, kind),
).fetchone()
return row is not None
def purge_stale_requests() -> None:
try:
from app.db import get_conn
with get_conn() as conn:
conn.execute(
"DELETE FROM sso_requests WHERE created_at < datetime('now', ?)",
(f"-{REQUEST_TTL_SECONDS} seconds",),
)
conn.commit()
except Exception:
logger.exception("purge_stale_requests")