config_store.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505
  1. """Configuration store: load, validate, atomically write, and hot-reload config.toml.
  2. Design notes (see AGENTS.md):
  3. * **tomlkit, not tomllib/tomli-w.** The admin panel writes this file back, and the file
  4. is heavily commented. tomlkit round-trips comments/ordering; a plain dict dump would
  5. destroy them on the first save.
  6. * **Immutable snapshots.** Config is never mutated in place. `apply()` builds a new
  7. frozen `Snapshot` and rebinds one attribute — atomic under the GIL, so readers need
  8. no lock. Views take one snapshot per request; the worker takes one per job. A job
  9. therefore runs under a single consistent config and an admin edit affects the *next*
  10. job, which removes every mid-job race without locking.
  11. * **Write-only secrets.** `redacted()` never emits a secret value; `apply()` treats a
  12. blank/absent secret as "leave unchanged".
  13. * **PPSQ_CONFIG** selects an alternate file. Tests rely on this to never touch the real
  14. config.toml.
  15. """
  16. from __future__ import annotations
  17. import copy
  18. import logging
  19. import os
  20. import shutil
  21. import tempfile
  22. import threading
  23. from dataclasses import dataclass
  24. from pathlib import Path
  25. from types import MappingProxyType
  26. from typing import Any, Mapping
  27. from urllib.parse import urlparse
  28. import tomlkit
  29. from pps_client import PPSClient
  30. log = logging.getLogger("pps.config")
  31. # Dotted keys the admin panel may write. An explicit ALLOWLIST — never a denylist.
  32. ADMIN_EDITABLE: frozenset[str] = frozenset(
  33. {
  34. "pps.base_url", "pps.username", "pps.password", "pps.verify_tls",
  35. "pps.timeout", "pps.client_cert", "pps.client_key",
  36. "quarantine.default_folder", "quarantine.folders", "quarantine.deleted_folder",
  37. "quarantine.default_limit", "quarantine.list_query",
  38. "quarantine.default_days_back", "quarantine.chunk_size",
  39. "quarantine.default_sort_field", "quarantine.default_sort_dir",
  40. "quarantine.report_release.steps", "quarantine.report_release.move_target",
  41. "quarantine.report_release.step_delay_seconds",
  42. "auth.denied_message", "auth.users",
  43. "app.log_level",
  44. }
  45. )
  46. # Written to disk but only picked up on restart; reported back so the UI can say so.
  47. RESTART_KEYS: frozenset[str] = frozenset(
  48. {
  49. "app.secret_key", "app.listen", "app.port", "app.db_path",
  50. "app.worker_log", "app.app_log", "app.cookie_secure",
  51. "auth.mode",
  52. }
  53. )
  54. # Never leave the process. `redacted()` emits a `<name>_set` boolean instead.
  55. SECRET_KEYS: frozenset[str] = frozenset(
  56. {"pps.password", "okta.client_secret", "auth.static.password", "app.secret_key"}
  57. )
  58. # Connection identity — a change here means the PPSClient must be rebuilt.
  59. _PPS_FINGERPRINT = (
  60. "base_url", "username", "password", "verify_tls", "timeout", "client_cert", "client_key",
  61. )
  62. _SORT_FIELDS = frozenset({"subject", "date", "from", "rcpt"})
  63. _SORT_DIRS = frozenset({"asc", "desc"})
  64. _ROLES = frozenset({"admin", "user"})
  65. class ConfigError(Exception):
  66. """Validation failure. Carries per-field messages for a 400 response."""
  67. def __init__(self, errors: list[dict[str, str]]):
  68. self.errors = errors
  69. super().__init__("; ".join(f"{e['field']}: {e['message']}" for e in errors))
  70. @dataclass(frozen=True)
  71. class Snapshot:
  72. version: int
  73. pps: Mapping[str, Any]
  74. quarantine: Mapping[str, Any]
  75. app: Mapping[str, Any]
  76. auth: Mapping[str, Any]
  77. okta: Mapping[str, Any]
  78. path: Path
  79. @dataclass(frozen=True)
  80. class ApplyResult:
  81. version: int
  82. applied: list[str]
  83. restart_required: list[str]
  84. def _freeze(d: Any) -> Any:
  85. """Deep-copy plain data and wrap mappings read-only, so a snapshot can't be mutated."""
  86. if isinstance(d, Mapping):
  87. return MappingProxyType({k: _freeze(v) for k, v in d.items()})
  88. if isinstance(d, (list, tuple)):
  89. return tuple(_freeze(v) for v in d)
  90. return d
  91. def _plain(doc: Any) -> Any:
  92. """tomlkit containers -> plain python (dict/list/scalars)."""
  93. if isinstance(doc, Mapping):
  94. return {k: _plain(v) for k, v in doc.items()}
  95. if isinstance(doc, (list, tuple)):
  96. return [_plain(v) for v in doc]
  97. return doc
  98. def _dig(d: Mapping, dotted: str, default=None):
  99. cur: Any = d
  100. for part in dotted.split("."):
  101. if not isinstance(cur, Mapping) or part not in cur:
  102. return default
  103. cur = cur[part]
  104. return cur
  105. def _flatten(d: Mapping, prefix: str = "") -> dict[str, Any]:
  106. """Flatten nested dicts to dotted keys. Lists are leaves (e.g. quarantine.folders)."""
  107. out: dict[str, Any] = {}
  108. for k, v in d.items():
  109. key = f"{prefix}{k}"
  110. if isinstance(v, Mapping):
  111. out.update(_flatten(v, f"{key}."))
  112. else:
  113. out[key] = v
  114. return out
  115. def resolve_config_path(explicit: str | Path | None = None) -> Path:
  116. """Explicit arg -> $PPSQ_CONFIG -> ./config.toml (next to this module)."""
  117. if explicit:
  118. return Path(explicit)
  119. env = os.environ.get("PPSQ_CONFIG")
  120. if env:
  121. return Path(env)
  122. return Path(__file__).with_name("config.toml")
  123. class ConfigStore:
  124. def __init__(self, path: str | Path | None = None):
  125. # Guard: a test that forgets PPSQ_CONFIG must never touch the real config.toml.
  126. if os.environ.get("PYTEST_CURRENT_TEST") and not (path or os.environ.get("PPSQ_CONFIG")):
  127. raise RuntimeError(
  128. "Refusing to open the default config.toml under pytest. "
  129. "Set PPSQ_CONFIG to a temp copy (see tests/conftest.py)."
  130. )
  131. self.path = resolve_config_path(path)
  132. if not self.path.exists():
  133. raise SystemExit(
  134. f"Missing {self.path}. Copy config.example.toml to config.toml and edit it."
  135. )
  136. self._lock = threading.RLock()
  137. self._version = 0
  138. doc = self._read_doc()
  139. data = _migrate(_plain(doc))
  140. _validate(data)
  141. self._snapshot = self._build_snapshot(data)
  142. self._pps_fp: tuple | None = None
  143. self._pps: PPSClient | None = None
  144. self._rebuild_pps(data)
  145. # ------------------------------------------------------------------ reading
  146. def snapshot(self) -> Snapshot:
  147. """Current config. Lock-free: attribute reads are atomic under the GIL."""
  148. return self._snapshot
  149. def pps(self) -> PPSClient:
  150. return self._pps # type: ignore[return-value]
  151. def redacted(self) -> dict:
  152. """Admin-panel payload. Secrets are replaced by a `<name>_set` boolean."""
  153. snap = self._snapshot
  154. data = {
  155. "pps": dict(_plain(snap.pps)),
  156. "quarantine": dict(_plain(snap.quarantine)),
  157. "app": dict(_plain(snap.app)),
  158. "auth": dict(_plain(snap.auth)),
  159. }
  160. for dotted in SECRET_KEYS:
  161. section, _, leaf = dotted.rpartition(".")
  162. parent = _dig(data, section) if section else data
  163. if isinstance(parent, dict) and leaf in parent:
  164. parent[f"{leaf}_set"] = bool(parent.pop(leaf))
  165. # okta is server-only: expose nothing but whether it is configured.
  166. data["okta"] = {"configured": bool(_dig(snap.okta, "client_id"))}
  167. return data
  168. # ------------------------------------------------------------------ writing
  169. def apply(self, patch: dict, *, actor: str = "system") -> ApplyResult:
  170. """Validate + atomically write a nested patch, then swap in a new snapshot."""
  171. with self._lock:
  172. flat = _flatten(patch)
  173. rejected = [k for k in flat if k not in ADMIN_EDITABLE]
  174. if rejected:
  175. raise ConfigError(
  176. [{"field": k, "message": "unknown or non-editable key"} for k in rejected]
  177. )
  178. doc = self._read_doc()
  179. current = _migrate(_plain(doc))
  180. merged = copy.deepcopy(current)
  181. applied: list[str] = []
  182. for dotted, value in flat.items():
  183. # Blank secret = leave unchanged (write-only fields).
  184. if dotted in SECRET_KEYS and (value is None or value == ""):
  185. continue
  186. if _dig(merged, dotted) == value:
  187. continue
  188. _set_dotted(merged, dotted, value)
  189. applied.append(dotted)
  190. if not applied:
  191. return ApplyResult(self._version, [], [])
  192. _validate(merged)
  193. for dotted in applied:
  194. _set_dotted_doc(doc, dotted, _dig(merged, dotted))
  195. self._write_atomic(doc)
  196. self._rebuild_pps(merged)
  197. self._snapshot = self._build_snapshot(merged)
  198. log.info("config updated by %s: %s", actor, ", ".join(sorted(applied)))
  199. return ApplyResult(
  200. version=self._snapshot.version,
  201. applied=sorted(applied),
  202. restart_required=sorted(k for k in applied if k in RESTART_KEYS),
  203. )
  204. def reload_from_disk(self) -> Snapshot:
  205. with self._lock:
  206. data = _migrate(_plain(self._read_doc()))
  207. _validate(data)
  208. self._rebuild_pps(data)
  209. self._snapshot = self._build_snapshot(data)
  210. log.info("config reloaded from %s (version %d)", self.path, self._snapshot.version)
  211. return self._snapshot
  212. # ------------------------------------------------------------------ internals
  213. def _read_doc(self):
  214. with self.path.open("r", encoding="utf-8") as fh:
  215. return tomlkit.parse(fh.read())
  216. def _build_snapshot(self, data: dict) -> Snapshot:
  217. self._version += 1
  218. return Snapshot(
  219. version=self._version,
  220. pps=_freeze(data.get("pps", {})),
  221. quarantine=_freeze(data.get("quarantine", {})),
  222. app=_freeze(data.get("app", {})),
  223. auth=_freeze(data.get("auth", {})),
  224. okta=_freeze(data.get("okta", {})),
  225. path=self.path,
  226. )
  227. def _rebuild_pps(self, data: dict) -> None:
  228. """Rebuild the PPS client only when connection identity changed."""
  229. p = data.get("pps", {})
  230. fp = tuple(p.get(k) for k in _PPS_FINGERPRINT)
  231. if fp == self._pps_fp and self._pps is not None:
  232. return
  233. # Rebind the client before the snapshot; in-flight requests keep their own ref.
  234. self._pps = PPSClient(
  235. base_url=p["base_url"],
  236. username=p["username"],
  237. password=p["password"],
  238. verify_tls=p.get("verify_tls", False),
  239. timeout=int(p.get("timeout", 120)),
  240. client_cert=p.get("client_cert") or None,
  241. client_key=p.get("client_key") or None,
  242. )
  243. self._pps_fp = fp
  244. def _write_atomic(self, doc) -> None:
  245. """Temp file in the same dir -> fsync -> backup -> atomic rename -> fsync dir."""
  246. parent = self.path.parent
  247. fd, tmp = tempfile.mkstemp(dir=parent, prefix=".config.", suffix=".toml.tmp")
  248. try:
  249. os.fchmod(fd, 0o600)
  250. with os.fdopen(fd, "w", encoding="utf-8") as fh:
  251. fh.write(tomlkit.dumps(doc))
  252. fh.flush()
  253. os.fsync(fh.fileno())
  254. if self.path.exists():
  255. shutil.copy2(self.path, self.path.with_name(self.path.name + ".bak"))
  256. os.replace(tmp, self.path) # same filesystem by construction
  257. tmp = None
  258. dfd = os.open(parent, os.O_DIRECTORY)
  259. try:
  260. os.fsync(dfd)
  261. finally:
  262. os.close(dfd)
  263. finally:
  264. if tmp:
  265. Path(tmp).unlink(missing_ok=True)
  266. def _set_dotted(d: dict, dotted: str, value) -> None:
  267. parts = dotted.split(".")
  268. cur = d
  269. for part in parts[:-1]:
  270. cur = cur.setdefault(part, {})
  271. cur[parts[-1]] = value
  272. def _set_dotted_doc(doc, dotted: str, value) -> None:
  273. """Set a key in a tomlkit document, creating intermediate tables as needed."""
  274. parts = dotted.split(".")
  275. cur = doc
  276. for part in parts[:-1]:
  277. if part not in cur:
  278. cur[part] = tomlkit.table()
  279. cur = cur[part]
  280. cur[parts[-1]] = value
  281. def _migrate(data: dict) -> dict:
  282. """In-memory back-compat. Never rewrites the user's file on load.
  283. `quarantine.report_release_folder` (PoC) -> `[quarantine.report_release]` with
  284. steps=["release","move"], which reproduces the old hardcoded behaviour exactly.
  285. """
  286. q = data.setdefault("quarantine", {})
  287. if "report_release" not in q:
  288. legacy = q.get("report_release_folder")
  289. q["report_release"] = {
  290. "steps": ["release", "move"] if legacy else ["release"],
  291. "move_target": legacy or "",
  292. }
  293. rr = q["report_release"]
  294. rr.setdefault("steps", ["release", "move"])
  295. rr.setdefault("move_target", q.get("report_release_folder", "") or "")
  296. rr.setdefault("step_delay_seconds", 60)
  297. q.setdefault("default_sort_field", "subject")
  298. q.setdefault("default_sort_dir", "asc")
  299. a = data.setdefault("auth", {})
  300. # Fail closed: a config with no explicit mode defaults to oidc (which then requires
  301. # [okta] and fails loudly if missing), never to shared-password static admin.
  302. a.setdefault("mode", "oidc")
  303. a.setdefault("users", [])
  304. a.setdefault(
  305. "denied_message",
  306. "Your account is not authorised to use the PPS Quarantine Manager.",
  307. )
  308. return data
  309. def _validate(data: dict) -> None:
  310. """Validate the merged config. Raises ConfigError with per-field messages."""
  311. import pipeline # local import: pipeline imports nothing from us
  312. errors: list[dict[str, str]] = []
  313. def bad(field: str, msg: str) -> None:
  314. errors.append({"field": field, "message": msg})
  315. p = data.get("pps", {})
  316. url = str(p.get("base_url", ""))
  317. parsed = urlparse(url)
  318. if parsed.scheme not in ("http", "https") or not parsed.netloc:
  319. bad("pps.base_url", "must be an http(s) URL with a host")
  320. elif parsed.scheme == "http" and not os.environ.get("PPSQ_ALLOW_INSECURE"):
  321. bad("pps.base_url", "must use https (set PPSQ_ALLOW_INSECURE=1 to override)")
  322. try:
  323. t = int(p.get("timeout", 120))
  324. if not 5 <= t <= 600:
  325. bad("pps.timeout", "must be between 5 and 600 seconds")
  326. except (TypeError, ValueError):
  327. bad("pps.timeout", "must be an integer")
  328. for key in ("client_cert", "client_key"):
  329. val = p.get(key)
  330. if val and not Path(val).exists():
  331. bad(f"pps.{key}", f"file not found: {val}")
  332. q = data.get("quarantine", {})
  333. folders = q.get("folders", [])
  334. if not isinstance(folders, (list, tuple)) or not folders:
  335. bad("quarantine.folders", "at least one folder is required")
  336. folders = []
  337. else:
  338. seen = set()
  339. for f in folders:
  340. if not isinstance(f, str) or not f.strip():
  341. bad("quarantine.folders", "folder names must be non-empty strings")
  342. elif "," in f:
  343. # localguids are joined with "," in the POST payload (pps_client.act).
  344. bad("quarantine.folders", f"folder name may not contain a comma: {f!r}")
  345. elif f != f.strip():
  346. bad("quarantine.folders", f"folder name has leading/trailing space: {f!r}")
  347. elif len(f) > 128:
  348. bad("quarantine.folders", f"folder name too long: {f[:20]!r}…")
  349. elif f in seen:
  350. bad("quarantine.folders", f"duplicate folder: {f!r}")
  351. seen.add(f)
  352. for field in ("default_folder", "deleted_folder"):
  353. val = q.get(field)
  354. if folders and val and val not in folders:
  355. bad(f"quarantine.{field}", f"{val!r} is not in the folder list")
  356. try:
  357. lim = int(q.get("default_limit", 200))
  358. if not 1 <= lim <= 1000:
  359. bad("quarantine.default_limit", "must be between 1 and 1000")
  360. except (TypeError, ValueError):
  361. bad("quarantine.default_limit", "must be an integer")
  362. try:
  363. cs = int(q.get("chunk_size", 25))
  364. if not 1 <= cs <= 500:
  365. bad("quarantine.chunk_size", "must be between 1 and 500")
  366. except (TypeError, ValueError):
  367. bad("quarantine.chunk_size", "must be an integer")
  368. try:
  369. db = int(q.get("default_days_back", 7))
  370. if db < 1:
  371. bad("quarantine.default_days_back", "must be at least 1")
  372. except (TypeError, ValueError):
  373. bad("quarantine.default_days_back", "must be an integer")
  374. if not str(q.get("list_query", "")).strip():
  375. bad("quarantine.list_query", "required (the PPS API rejects folder-only searches)")
  376. if q.get("default_sort_field") not in _SORT_FIELDS:
  377. bad("quarantine.default_sort_field", f"must be one of {sorted(_SORT_FIELDS)}")
  378. if q.get("default_sort_dir") not in _SORT_DIRS:
  379. bad("quarantine.default_sort_dir", f"must be one of {sorted(_SORT_DIRS)}")
  380. rr = q.get("report_release", {})
  381. steps = rr.get("steps", [])
  382. if not isinstance(steps, (list, tuple)) or not steps:
  383. bad("quarantine.report_release.steps", "select at least one action")
  384. else:
  385. unknown = [s for s in steps if s not in pipeline.ALLOWED_STEPS]
  386. if unknown:
  387. bad("quarantine.report_release.steps", f"unknown step(s): {unknown}")
  388. elif len(set(steps)) != len(steps):
  389. bad("quarantine.report_release.steps", "duplicate steps")
  390. elif not _is_subsequence(steps, pipeline.ALLOWED_STEPS):
  391. bad(
  392. "quarantine.report_release.steps",
  393. f"must follow the fixed order {list(pipeline.ALLOWED_STEPS)}",
  394. )
  395. if "move" in steps:
  396. target = rr.get("move_target")
  397. if not target:
  398. bad("quarantine.report_release.move_target", "required when 'move' is selected")
  399. elif folders and target not in folders:
  400. bad(
  401. "quarantine.report_release.move_target",
  402. f"{target!r} is not in the folder list",
  403. )
  404. if "delete" in steps and not q.get("deleted_folder"):
  405. bad("quarantine.deleted_folder", "required when 'delete' is selected")
  406. try:
  407. d = int(rr.get("step_delay_seconds", 60))
  408. if not 0 <= d <= 3600:
  409. bad("quarantine.report_release.step_delay_seconds", "must be between 0 and 3600")
  410. except (TypeError, ValueError):
  411. bad("quarantine.report_release.step_delay_seconds", "must be an integer")
  412. a = data.get("auth", {})
  413. if a.get("mode") not in ("oidc", "static"):
  414. bad("auth.mode", "must be 'oidc' or 'static'")
  415. users = a.get("users", [])
  416. if not isinstance(users, (list, tuple)):
  417. bad("auth.users", "must be a list")
  418. else:
  419. seen_emails = set()
  420. for i, u in enumerate(users):
  421. if not isinstance(u, Mapping):
  422. bad(f"auth.users[{i}]", "must be a table with email and role")
  423. continue
  424. email = str(u.get("email", "")).strip()
  425. if "@" not in email or " " in email or len(email) < 3:
  426. bad(f"auth.users[{i}].email", f"invalid email: {email!r}")
  427. elif email.casefold() in seen_emails:
  428. bad(f"auth.users[{i}].email", f"duplicate user: {email!r}")
  429. seen_emails.add(email.casefold())
  430. if u.get("role") not in _ROLES:
  431. bad(f"auth.users[{i}].role", f"must be one of {sorted(_ROLES)}")
  432. if errors:
  433. raise ConfigError(errors)
  434. def _is_subsequence(seq, universe) -> bool:
  435. it = iter(universe)
  436. return all(any(x == u for u in it) for x in seq)