| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505 |
- """Configuration store: load, validate, atomically write, and hot-reload config.toml.
- Design notes (see AGENTS.md):
- * **tomlkit, not tomllib/tomli-w.** The admin panel writes this file back, and the file
- is heavily commented. tomlkit round-trips comments/ordering; a plain dict dump would
- destroy them on the first save.
- * **Immutable snapshots.** Config is never mutated in place. `apply()` builds a new
- frozen `Snapshot` and rebinds one attribute — atomic under the GIL, so readers need
- no lock. Views take one snapshot per request; the worker takes one per job. A job
- therefore runs under a single consistent config and an admin edit affects the *next*
- job, which removes every mid-job race without locking.
- * **Write-only secrets.** `redacted()` never emits a secret value; `apply()` treats a
- blank/absent secret as "leave unchanged".
- * **PPSQ_CONFIG** selects an alternate file. Tests rely on this to never touch the real
- config.toml.
- """
- from __future__ import annotations
- import copy
- import logging
- import os
- import shutil
- import tempfile
- import threading
- from dataclasses import dataclass
- from pathlib import Path
- from types import MappingProxyType
- from typing import Any, Mapping
- from urllib.parse import urlparse
- import tomlkit
- from pps_client import PPSClient
- log = logging.getLogger("pps.config")
- # Dotted keys the admin panel may write. An explicit ALLOWLIST — never a denylist.
- ADMIN_EDITABLE: frozenset[str] = frozenset(
- {
- "pps.base_url", "pps.username", "pps.password", "pps.verify_tls",
- "pps.timeout", "pps.client_cert", "pps.client_key",
- "quarantine.default_folder", "quarantine.folders", "quarantine.deleted_folder",
- "quarantine.default_limit", "quarantine.list_query",
- "quarantine.default_days_back", "quarantine.chunk_size",
- "quarantine.default_sort_field", "quarantine.default_sort_dir",
- "quarantine.report_release.steps", "quarantine.report_release.move_target",
- "quarantine.report_release.step_delay_seconds",
- "auth.denied_message", "auth.users",
- "app.log_level",
- }
- )
- # Written to disk but only picked up on restart; reported back so the UI can say so.
- RESTART_KEYS: frozenset[str] = frozenset(
- {
- "app.secret_key", "app.listen", "app.port", "app.db_path",
- "app.worker_log", "app.app_log", "app.cookie_secure",
- "auth.mode",
- }
- )
- # Never leave the process. `redacted()` emits a `<name>_set` boolean instead.
- SECRET_KEYS: frozenset[str] = frozenset(
- {"pps.password", "okta.client_secret", "auth.static.password", "app.secret_key"}
- )
- # Connection identity — a change here means the PPSClient must be rebuilt.
- _PPS_FINGERPRINT = (
- "base_url", "username", "password", "verify_tls", "timeout", "client_cert", "client_key",
- )
- _SORT_FIELDS = frozenset({"subject", "date", "from", "rcpt"})
- _SORT_DIRS = frozenset({"asc", "desc"})
- _ROLES = frozenset({"admin", "user"})
- class ConfigError(Exception):
- """Validation failure. Carries per-field messages for a 400 response."""
- def __init__(self, errors: list[dict[str, str]]):
- self.errors = errors
- super().__init__("; ".join(f"{e['field']}: {e['message']}" for e in errors))
- @dataclass(frozen=True)
- class Snapshot:
- version: int
- pps: Mapping[str, Any]
- quarantine: Mapping[str, Any]
- app: Mapping[str, Any]
- auth: Mapping[str, Any]
- okta: Mapping[str, Any]
- path: Path
- @dataclass(frozen=True)
- class ApplyResult:
- version: int
- applied: list[str]
- restart_required: list[str]
- def _freeze(d: Any) -> Any:
- """Deep-copy plain data and wrap mappings read-only, so a snapshot can't be mutated."""
- if isinstance(d, Mapping):
- return MappingProxyType({k: _freeze(v) for k, v in d.items()})
- if isinstance(d, (list, tuple)):
- return tuple(_freeze(v) for v in d)
- return d
- def _plain(doc: Any) -> Any:
- """tomlkit containers -> plain python (dict/list/scalars)."""
- if isinstance(doc, Mapping):
- return {k: _plain(v) for k, v in doc.items()}
- if isinstance(doc, (list, tuple)):
- return [_plain(v) for v in doc]
- return doc
- def _dig(d: Mapping, dotted: str, default=None):
- cur: Any = d
- for part in dotted.split("."):
- if not isinstance(cur, Mapping) or part not in cur:
- return default
- cur = cur[part]
- return cur
- def _flatten(d: Mapping, prefix: str = "") -> dict[str, Any]:
- """Flatten nested dicts to dotted keys. Lists are leaves (e.g. quarantine.folders)."""
- out: dict[str, Any] = {}
- for k, v in d.items():
- key = f"{prefix}{k}"
- if isinstance(v, Mapping):
- out.update(_flatten(v, f"{key}."))
- else:
- out[key] = v
- return out
- def resolve_config_path(explicit: str | Path | None = None) -> Path:
- """Explicit arg -> $PPSQ_CONFIG -> ./config.toml (next to this module)."""
- if explicit:
- return Path(explicit)
- env = os.environ.get("PPSQ_CONFIG")
- if env:
- return Path(env)
- return Path(__file__).with_name("config.toml")
- class ConfigStore:
- def __init__(self, path: str | Path | None = None):
- # Guard: a test that forgets PPSQ_CONFIG must never touch the real config.toml.
- if os.environ.get("PYTEST_CURRENT_TEST") and not (path or os.environ.get("PPSQ_CONFIG")):
- raise RuntimeError(
- "Refusing to open the default config.toml under pytest. "
- "Set PPSQ_CONFIG to a temp copy (see tests/conftest.py)."
- )
- self.path = resolve_config_path(path)
- if not self.path.exists():
- raise SystemExit(
- f"Missing {self.path}. Copy config.example.toml to config.toml and edit it."
- )
- self._lock = threading.RLock()
- self._version = 0
- doc = self._read_doc()
- data = _migrate(_plain(doc))
- _validate(data)
- self._snapshot = self._build_snapshot(data)
- self._pps_fp: tuple | None = None
- self._pps: PPSClient | None = None
- self._rebuild_pps(data)
- # ------------------------------------------------------------------ reading
- def snapshot(self) -> Snapshot:
- """Current config. Lock-free: attribute reads are atomic under the GIL."""
- return self._snapshot
- def pps(self) -> PPSClient:
- return self._pps # type: ignore[return-value]
- def redacted(self) -> dict:
- """Admin-panel payload. Secrets are replaced by a `<name>_set` boolean."""
- snap = self._snapshot
- data = {
- "pps": dict(_plain(snap.pps)),
- "quarantine": dict(_plain(snap.quarantine)),
- "app": dict(_plain(snap.app)),
- "auth": dict(_plain(snap.auth)),
- }
- for dotted in SECRET_KEYS:
- section, _, leaf = dotted.rpartition(".")
- parent = _dig(data, section) if section else data
- if isinstance(parent, dict) and leaf in parent:
- parent[f"{leaf}_set"] = bool(parent.pop(leaf))
- # okta is server-only: expose nothing but whether it is configured.
- data["okta"] = {"configured": bool(_dig(snap.okta, "client_id"))}
- return data
- # ------------------------------------------------------------------ writing
- def apply(self, patch: dict, *, actor: str = "system") -> ApplyResult:
- """Validate + atomically write a nested patch, then swap in a new snapshot."""
- with self._lock:
- flat = _flatten(patch)
- rejected = [k for k in flat if k not in ADMIN_EDITABLE]
- if rejected:
- raise ConfigError(
- [{"field": k, "message": "unknown or non-editable key"} for k in rejected]
- )
- doc = self._read_doc()
- current = _migrate(_plain(doc))
- merged = copy.deepcopy(current)
- applied: list[str] = []
- for dotted, value in flat.items():
- # Blank secret = leave unchanged (write-only fields).
- if dotted in SECRET_KEYS and (value is None or value == ""):
- continue
- if _dig(merged, dotted) == value:
- continue
- _set_dotted(merged, dotted, value)
- applied.append(dotted)
- if not applied:
- return ApplyResult(self._version, [], [])
- _validate(merged)
- for dotted in applied:
- _set_dotted_doc(doc, dotted, _dig(merged, dotted))
- self._write_atomic(doc)
- self._rebuild_pps(merged)
- self._snapshot = self._build_snapshot(merged)
- log.info("config updated by %s: %s", actor, ", ".join(sorted(applied)))
- return ApplyResult(
- version=self._snapshot.version,
- applied=sorted(applied),
- restart_required=sorted(k for k in applied if k in RESTART_KEYS),
- )
- def reload_from_disk(self) -> Snapshot:
- with self._lock:
- data = _migrate(_plain(self._read_doc()))
- _validate(data)
- self._rebuild_pps(data)
- self._snapshot = self._build_snapshot(data)
- log.info("config reloaded from %s (version %d)", self.path, self._snapshot.version)
- return self._snapshot
- # ------------------------------------------------------------------ internals
- def _read_doc(self):
- with self.path.open("r", encoding="utf-8") as fh:
- return tomlkit.parse(fh.read())
- def _build_snapshot(self, data: dict) -> Snapshot:
- self._version += 1
- return Snapshot(
- version=self._version,
- pps=_freeze(data.get("pps", {})),
- quarantine=_freeze(data.get("quarantine", {})),
- app=_freeze(data.get("app", {})),
- auth=_freeze(data.get("auth", {})),
- okta=_freeze(data.get("okta", {})),
- path=self.path,
- )
- def _rebuild_pps(self, data: dict) -> None:
- """Rebuild the PPS client only when connection identity changed."""
- p = data.get("pps", {})
- fp = tuple(p.get(k) for k in _PPS_FINGERPRINT)
- if fp == self._pps_fp and self._pps is not None:
- return
- # Rebind the client before the snapshot; in-flight requests keep their own ref.
- self._pps = PPSClient(
- base_url=p["base_url"],
- username=p["username"],
- password=p["password"],
- verify_tls=p.get("verify_tls", False),
- timeout=int(p.get("timeout", 120)),
- client_cert=p.get("client_cert") or None,
- client_key=p.get("client_key") or None,
- )
- self._pps_fp = fp
- def _write_atomic(self, doc) -> None:
- """Temp file in the same dir -> fsync -> backup -> atomic rename -> fsync dir."""
- parent = self.path.parent
- fd, tmp = tempfile.mkstemp(dir=parent, prefix=".config.", suffix=".toml.tmp")
- try:
- os.fchmod(fd, 0o600)
- with os.fdopen(fd, "w", encoding="utf-8") as fh:
- fh.write(tomlkit.dumps(doc))
- fh.flush()
- os.fsync(fh.fileno())
- if self.path.exists():
- shutil.copy2(self.path, self.path.with_name(self.path.name + ".bak"))
- os.replace(tmp, self.path) # same filesystem by construction
- tmp = None
- dfd = os.open(parent, os.O_DIRECTORY)
- try:
- os.fsync(dfd)
- finally:
- os.close(dfd)
- finally:
- if tmp:
- Path(tmp).unlink(missing_ok=True)
- def _set_dotted(d: dict, dotted: str, value) -> None:
- parts = dotted.split(".")
- cur = d
- for part in parts[:-1]:
- cur = cur.setdefault(part, {})
- cur[parts[-1]] = value
- def _set_dotted_doc(doc, dotted: str, value) -> None:
- """Set a key in a tomlkit document, creating intermediate tables as needed."""
- parts = dotted.split(".")
- cur = doc
- for part in parts[:-1]:
- if part not in cur:
- cur[part] = tomlkit.table()
- cur = cur[part]
- cur[parts[-1]] = value
- def _migrate(data: dict) -> dict:
- """In-memory back-compat. Never rewrites the user's file on load.
- `quarantine.report_release_folder` (PoC) -> `[quarantine.report_release]` with
- steps=["release","move"], which reproduces the old hardcoded behaviour exactly.
- """
- q = data.setdefault("quarantine", {})
- if "report_release" not in q:
- legacy = q.get("report_release_folder")
- q["report_release"] = {
- "steps": ["release", "move"] if legacy else ["release"],
- "move_target": legacy or "",
- }
- rr = q["report_release"]
- rr.setdefault("steps", ["release", "move"])
- rr.setdefault("move_target", q.get("report_release_folder", "") or "")
- rr.setdefault("step_delay_seconds", 60)
- q.setdefault("default_sort_field", "subject")
- q.setdefault("default_sort_dir", "asc")
- a = data.setdefault("auth", {})
- # Fail closed: a config with no explicit mode defaults to oidc (which then requires
- # [okta] and fails loudly if missing), never to shared-password static admin.
- a.setdefault("mode", "oidc")
- a.setdefault("users", [])
- a.setdefault(
- "denied_message",
- "Your account is not authorised to use the PPS Quarantine Manager.",
- )
- return data
- def _validate(data: dict) -> None:
- """Validate the merged config. Raises ConfigError with per-field messages."""
- import pipeline # local import: pipeline imports nothing from us
- errors: list[dict[str, str]] = []
- def bad(field: str, msg: str) -> None:
- errors.append({"field": field, "message": msg})
- p = data.get("pps", {})
- url = str(p.get("base_url", ""))
- parsed = urlparse(url)
- if parsed.scheme not in ("http", "https") or not parsed.netloc:
- bad("pps.base_url", "must be an http(s) URL with a host")
- elif parsed.scheme == "http" and not os.environ.get("PPSQ_ALLOW_INSECURE"):
- bad("pps.base_url", "must use https (set PPSQ_ALLOW_INSECURE=1 to override)")
- try:
- t = int(p.get("timeout", 120))
- if not 5 <= t <= 600:
- bad("pps.timeout", "must be between 5 and 600 seconds")
- except (TypeError, ValueError):
- bad("pps.timeout", "must be an integer")
- for key in ("client_cert", "client_key"):
- val = p.get(key)
- if val and not Path(val).exists():
- bad(f"pps.{key}", f"file not found: {val}")
- q = data.get("quarantine", {})
- folders = q.get("folders", [])
- if not isinstance(folders, (list, tuple)) or not folders:
- bad("quarantine.folders", "at least one folder is required")
- folders = []
- else:
- seen = set()
- for f in folders:
- if not isinstance(f, str) or not f.strip():
- bad("quarantine.folders", "folder names must be non-empty strings")
- elif "," in f:
- # localguids are joined with "," in the POST payload (pps_client.act).
- bad("quarantine.folders", f"folder name may not contain a comma: {f!r}")
- elif f != f.strip():
- bad("quarantine.folders", f"folder name has leading/trailing space: {f!r}")
- elif len(f) > 128:
- bad("quarantine.folders", f"folder name too long: {f[:20]!r}…")
- elif f in seen:
- bad("quarantine.folders", f"duplicate folder: {f!r}")
- seen.add(f)
- for field in ("default_folder", "deleted_folder"):
- val = q.get(field)
- if folders and val and val not in folders:
- bad(f"quarantine.{field}", f"{val!r} is not in the folder list")
- try:
- lim = int(q.get("default_limit", 200))
- if not 1 <= lim <= 1000:
- bad("quarantine.default_limit", "must be between 1 and 1000")
- except (TypeError, ValueError):
- bad("quarantine.default_limit", "must be an integer")
- try:
- cs = int(q.get("chunk_size", 25))
- if not 1 <= cs <= 500:
- bad("quarantine.chunk_size", "must be between 1 and 500")
- except (TypeError, ValueError):
- bad("quarantine.chunk_size", "must be an integer")
- try:
- db = int(q.get("default_days_back", 7))
- if db < 1:
- bad("quarantine.default_days_back", "must be at least 1")
- except (TypeError, ValueError):
- bad("quarantine.default_days_back", "must be an integer")
- if not str(q.get("list_query", "")).strip():
- bad("quarantine.list_query", "required (the PPS API rejects folder-only searches)")
- if q.get("default_sort_field") not in _SORT_FIELDS:
- bad("quarantine.default_sort_field", f"must be one of {sorted(_SORT_FIELDS)}")
- if q.get("default_sort_dir") not in _SORT_DIRS:
- bad("quarantine.default_sort_dir", f"must be one of {sorted(_SORT_DIRS)}")
- rr = q.get("report_release", {})
- steps = rr.get("steps", [])
- if not isinstance(steps, (list, tuple)) or not steps:
- bad("quarantine.report_release.steps", "select at least one action")
- else:
- unknown = [s for s in steps if s not in pipeline.ALLOWED_STEPS]
- if unknown:
- bad("quarantine.report_release.steps", f"unknown step(s): {unknown}")
- elif len(set(steps)) != len(steps):
- bad("quarantine.report_release.steps", "duplicate steps")
- elif not _is_subsequence(steps, pipeline.ALLOWED_STEPS):
- bad(
- "quarantine.report_release.steps",
- f"must follow the fixed order {list(pipeline.ALLOWED_STEPS)}",
- )
- if "move" in steps:
- target = rr.get("move_target")
- if not target:
- bad("quarantine.report_release.move_target", "required when 'move' is selected")
- elif folders and target not in folders:
- bad(
- "quarantine.report_release.move_target",
- f"{target!r} is not in the folder list",
- )
- if "delete" in steps and not q.get("deleted_folder"):
- bad("quarantine.deleted_folder", "required when 'delete' is selected")
- try:
- d = int(rr.get("step_delay_seconds", 60))
- if not 0 <= d <= 3600:
- bad("quarantine.report_release.step_delay_seconds", "must be between 0 and 3600")
- except (TypeError, ValueError):
- bad("quarantine.report_release.step_delay_seconds", "must be an integer")
- a = data.get("auth", {})
- if a.get("mode") not in ("oidc", "static"):
- bad("auth.mode", "must be 'oidc' or 'static'")
- users = a.get("users", [])
- if not isinstance(users, (list, tuple)):
- bad("auth.users", "must be a list")
- else:
- seen_emails = set()
- for i, u in enumerate(users):
- if not isinstance(u, Mapping):
- bad(f"auth.users[{i}]", "must be a table with email and role")
- continue
- email = str(u.get("email", "")).strip()
- if "@" not in email or " " in email or len(email) < 3:
- bad(f"auth.users[{i}].email", f"invalid email: {email!r}")
- elif email.casefold() in seen_emails:
- bad(f"auth.users[{i}].email", f"duplicate user: {email!r}")
- seen_emails.add(email.casefold())
- if u.get("role") not in _ROLES:
- bad(f"auth.users[{i}].role", f"must be one of {sorted(_ROLES)}")
- if errors:
- raise ConfigError(errors)
- def _is_subsequence(seq, universe) -> bool:
- it = iter(universe)
- return all(any(x == u for u in it) for x in seq)
|