From 689010e9a92b29dc7e7d2b5bfd79d4d891a2a3ca Mon Sep 17 00:00:00 2001 From: Ahmed Allam Date: Wed, 26 Aug 2026 22:29:28 +0000 Subject: [PATCH] Mirror the run's threat models into its state dir so resume keeps them --- strix/core/runner.py | 2 + strix/tools/threat_model/tools.py | 86 +++++++++++++++++++++++++++---- tests/test_threat_model_tool.py | 29 +++++++++-- 3 files changed, 103 insertions(+), 14 deletions(-) diff --git a/strix/core/runner.py b/strix/core/runner.py index ee996cd3..28acebd6 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -233,11 +233,13 @@ async def run_strix_scan( from strix.tools.coverage.tools import hydrate_coverage_from_disk from strix.tools.notes.tools import hydrate_notes_from_disk + from strix.tools.threat_model.tools import hydrate_threat_models_from_disk from strix.tools.todo.tools import hydrate_todos_from_disk hydrate_todos_from_disk(state_dir) hydrate_notes_from_disk(state_dir) hydrate_coverage_from_disk(state_dir) + hydrate_threat_models_from_disk(state_dir) root_id: str | None = None if is_resume: diff --git a/strix/tools/threat_model/tools.py b/strix/tools/threat_model/tools.py index a3435481..1015f457 100644 --- a/strix/tools/threat_model/tools.py +++ b/strix/tools/threat_model/tools.py @@ -1,16 +1,17 @@ -"""Run-scoped threat models — held in memory for the duration of one scan. +"""Run-scoped threat models — mirrored to ``{state_dir}/threat_models.json``. A threat model is the scan's shared answer to who the attacker is, where the trust boundaries sit, and what counts as critical for the target. One agent derives it and every other agent on the same run reads it back instead of re-deriving trust boundaries from scratch. -It never outlives the run. Nothing is written to disk and nothing is shared -between scans: a new scan against the same host or checkout starts with no -model and derives its own. Agents do spell one target several ways within a -run — the URL they were handed, the page they happen to be testing, a checkout -path — so a model is keyed by a normalized target identity to keep them -converging on one document instead of each starting a fresh one. +It does not outlive the scan. The mirror lives in the run's own state directory +and exists only so a resumed scan keeps the baseline its earlier agents agreed +on; a new scan against the same host or checkout starts with no model and +derives its own. Agents do spell one target several ways within a run — the URL +they were handed, the page they happen to be testing, a checkout path — so a +model is keyed by a normalized target identity to keep them converging on one +document instead of each starting a fresh one. """ from __future__ import annotations @@ -20,6 +21,7 @@ import json import logging import re import subprocess +import tempfile import threading from datetime import UTC, datetime from pathlib import Path @@ -43,9 +45,10 @@ _DEFAULT_PORTS = {"http": "80", "https": "443"} _store_lock = threading.RLock() -# The whole store: target identity -> model. Process-local and never persisted, -# so it holds exactly the models this run derived and dies with it. +# The whole store: target identity -> model. It holds exactly the models this +# scan derived, and is mirrored to the run's state directory for resume. _MODELS: dict[str, dict[str, Any]] = {} +_store_path: Path | None = None _REQUIRED_SECTIONS = ( "overview", @@ -202,6 +205,67 @@ def _resolve_target( return (_snap_to_scan_target(raw, known) if known else raw), None +def hydrate_threat_models_from_disk(state_dir: Path) -> None: + """Point the store at this run's mirror and load whatever it already holds. + + A resumed scan is the same scan, so its agents have to keep the baseline + the earlier ones agreed on. The mirror lives under the run directory, so a + different scan never reads it. + """ + global _store_path # noqa: PLW0603 + _store_path = state_dir / "threat_models.json" + with _store_lock: + _MODELS.clear() + if not _store_path.is_file(): + return + try: + data = json.loads(_store_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + logger.exception( + "threat_models.json at %s is unreadable; starting with no models", + _store_path, + ) + return + if not isinstance(data, dict): + return + _MODELS.update( + { + identity: model + for identity, model in data.items() + if isinstance(identity, str) and isinstance(model, dict) + } + ) + logger.info("threat models hydrated from %s (%d)", _store_path, len(_MODELS)) + + +def _persist_locked() -> None: + """Mirror the store to disk. Callers must already hold ``_store_lock``. + + Serializing and renaming in one critical section keeps a writer holding an + older serialization from winning the rename and dropping a concurrent + agent's model or amendment. + """ + path = _store_path + if path is None: + return + try: + payload = json.dumps(_MODELS, ensure_ascii=False, default=str) + path.parent.mkdir(parents=True, exist_ok=True) + with tempfile.NamedTemporaryFile( + mode="w", + encoding="utf-8", + dir=str(path.parent), + prefix=f".{path.name}.", + suffix=".tmp", + delete=False, + ) as tmp: + tmp.write(payload) + tmp_path = Path(tmp.name) + tmp_path.replace(path) + except OSError: + logger.exception("threat model mirror to %s failed", path) + + def _missing_sections(content: str) -> list[str]: lowered = content.lower() return [section for section in _REQUIRED_SECTIONS if section not in lowered] @@ -221,7 +285,7 @@ def _not_found(identity: str) -> dict[str, Any]: "target": identity, "message": ( "No threat model for this target on this scan. Nothing carries over " - "from earlier runs, so derive one — from the code if you have it, from " + "from other scans, so derive one — from the code if you have it, from " "recon output if you do not — and share it with save_threat_model, so " "every agent on this scan works from one view of the trust boundaries " "instead of each inventing their own." @@ -306,6 +370,7 @@ def _save_impl( "written_by": agent_name, "content": body, } + _persist_locked() message = ( "Threat model shared with this scan. Subagents should call get_threat_model " @@ -346,6 +411,7 @@ def _append_amendment( if len(json.dumps(sized, ensure_ascii=False).encode("utf-8")) > _MAX_MODEL_BYTES: return None, "Threat model with this amendment exceeds 512KB; tighten it." model["amendments"] = candidate + _persist_locked() return candidate, None diff --git a/tests/test_threat_model_tool.py b/tests/test_threat_model_tool.py index d12d650f..5b378b07 100644 --- a/tests/test_threat_model_tool.py +++ b/tests/test_threat_model_tool.py @@ -15,6 +15,7 @@ from strix.tools.threat_model.tools import ( _save_impl, amend_threat_model, get_threat_model, + hydrate_threat_models_from_disk, save_threat_model, ) @@ -63,8 +64,9 @@ def _make_repo(tmp_path: Path, name: str = "repo") -> Path: @pytest.fixture(autouse=True) def _empty_store() -> None: - """Each test is its own run, so it starts with an empty store.""" + """Each test is its own run, so it starts with an empty, unmirrored store.""" threat_model_tools._MODELS.clear() + threat_model_tools._store_path = None def test_missing_model_reports_not_found(tmp_path: Path) -> None: @@ -87,8 +89,10 @@ def test_saved_model_round_trips(tmp_path: Path) -> None: assert "multi-tenant billing API" in result["content"] -def test_nothing_is_written_to_disk(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: - """The model must not outlive the run, so no file may be left behind.""" +def test_nothing_is_written_outside_the_run( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """The model must not outlive the scan, so nothing may land in the home dir.""" home = tmp_path / "home" home.mkdir() monkeypatch.setenv("HOME", str(home)) @@ -103,13 +107,30 @@ def test_nothing_is_written_to_disk(tmp_path: Path, monkeypatch: pytest.MonkeyPa def test_a_new_run_starts_without_the_model(tmp_path: Path) -> None: """A later scan of the same target inherits nothing from this one.""" repo = _make_repo(tmp_path) + hydrate_threat_models_from_disk(tmp_path / "first-run") _save_impl(str(repo), _MODEL, "root") - threat_model_tools._MODELS.clear() # what a fresh process starts from + hydrate_threat_models_from_disk(tmp_path / "second-run") # a different scan assert _get_impl(str(repo))["found"] is False +def test_resuming_the_same_run_keeps_the_model(tmp_path: Path) -> None: + """A resumed scan is the same scan, so its agents keep the shared baseline.""" + state_dir = tmp_path / "state" + repo = _make_repo(tmp_path) + hydrate_threat_models_from_disk(state_dir) + _save_impl(str(repo), _MODEL, "root") + _amend_impl(str(repo), _ADDENDUM, "agent-a") + + threat_model_tools._MODELS.clear() # what the resuming process starts from + hydrate_threat_models_from_disk(state_dir) + + result = _get_impl(str(repo)) + assert result["found"] is True + assert [a["content"] for a in result["amendments"]] == [_ADDENDUM] + + def test_model_survives_a_new_revision_within_the_run(tmp_path: Path) -> None: """The model is not pinned to a revision; a commit mid-run does not drop it.""" repo = _make_repo(tmp_path)