- 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
784 lines
29 KiB
Python
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")
|