From 80699f0e95e6f6764f05624ca858ca6ce37c0ff4 Mon Sep 17 00:00:00 2001 From: AIOSAI Date: Tue, 28 Apr 2026 18:07:43 -0700 Subject: [PATCH] feat(spawn): fix(spawn): atomic registry writes + file locking + path containment (issue #490) Co-Authored-By: @spawn --- .../spawn/apps/handlers/json/json_handler.py | 20 ++++- src/aipass/spawn/apps/handlers/registry.py | 88 +++++++++++++------ .../builder/.spawn/.template_registry.json | 2 +- src/aipass/spawn/tests/test_contracts.py | 2 +- src/aipass/spawn/tests/test_spawn.py | 67 ++++++++++++-- 5 files changed, 140 insertions(+), 39 deletions(-) diff --git a/src/aipass/spawn/apps/handlers/json/json_handler.py b/src/aipass/spawn/apps/handlers/json/json_handler.py index 42a541a8..41654f1a 100644 --- a/src/aipass/spawn/apps/handlers/json/json_handler.py +++ b/src/aipass/spawn/apps/handlers/json/json_handler.py @@ -14,6 +14,8 @@ for the three-JSON system (config, data, log). import inspect import json +import os +import tempfile from datetime import datetime from pathlib import Path from typing import Any, Dict, Optional @@ -42,10 +44,24 @@ def read_json(file_path: Path) -> Optional[dict]: def write_json(file_path: Path, data: Any, indent: int = 2) -> bool: - """Write data to a JSON file.""" + """Write data to a JSON file atomically (temp file + os.replace).""" try: file_path.parent.mkdir(parents=True, exist_ok=True) - file_path.write_text(json.dumps(data, indent=indent) + "\n", encoding="utf-8") + content = json.dumps(data, indent=indent) + "\n" + fd, tmp_path = tempfile.mkstemp(dir=file_path.parent, suffix=".tmp") + closed = False + try: + os.write(fd, content.encode("utf-8")) + os.fsync(fd) + os.close(fd) + closed = True + os.replace(tmp_path, file_path) + except BaseException: + if not closed: + os.close(fd) + if os.path.exists(tmp_path): + os.unlink(tmp_path) + raise return True except OSError as e: logger.error("Failed to write JSON to %s: %s", file_path, e) diff --git a/src/aipass/spawn/apps/handlers/registry.py b/src/aipass/spawn/apps/handlers/registry.py index 720d1d94..5e55edec 100644 --- a/src/aipass/spawn/apps/handlers/registry.py +++ b/src/aipass/spawn/apps/handlers/registry.py @@ -8,6 +8,7 @@ """*_REGISTRY.json discovery and CRUD operations.""" +import fcntl import os from datetime import datetime from pathlib import Path @@ -161,10 +162,26 @@ def get_next_citizen_number(registry_path): return len(_branches_as_list(branches)) + 1 +def _validate_path_containment(branch_path, registry_path): + """Reject paths that escape the project root (directory containing the registry).""" + registry_root = Path(registry_path).resolve().parent + bp = Path(branch_path) + resolved = (registry_root / bp).resolve() if not bp.is_absolute() else bp.resolve() + try: + resolved.relative_to(registry_root) + return True + except ValueError: + logger.warning("[registry] Path %s is outside project root %s", resolved, registry_root) + return False + + def add_to_registry(registry_path, branch_name, branch_path, profile, email, purpose=""): """ Add a new branch entry to the registry. + Uses file locking (fcntl.LOCK_EX) around the entire read-modify-write + cycle to prevent corruption from concurrent spawns. + Args: registry_path: Path to AIPASS_REGISTRY.json branch_name: Uppercase branch name (e.g. "MY_AGENT") @@ -176,41 +193,54 @@ def add_to_registry(registry_path, branch_name, branch_path, profile, email, pur Returns: True if added, False if already exists or error """ - registry = load_registry(registry_path) - branches = registry.get("branches", []) + registry_path = Path(registry_path) - # Check for duplicates — handle both dict and list formats - if isinstance(branches, dict): - if branch_name in branches: - return False - else: - for branch in branches: - if branch.get("name") == branch_name: + if not _validate_path_containment(branch_path, registry_path): + logger.error("[registry] Path containment violation: %s escapes project root", branch_path) + return False + + lock_path = registry_path.parent / f".{registry_path.stem}.lock" + lock_path.parent.mkdir(parents=True, exist_ok=True) + lock_fd = open(lock_path, "w", encoding="utf-8") # noqa: SIM115 + try: + fcntl.flock(lock_fd, fcntl.LOCK_EX) + + registry = load_registry(registry_path) + branches = registry.get("branches", []) + + if isinstance(branches, dict): + if branch_name in branches: return False + else: + for branch in branches: + if branch.get("name") == branch_name: + return False - today = datetime.now().strftime("%Y-%m-%d") - entry = { - "name": branch_name, - "path": str(branch_path), - "profile": profile, - "description": purpose or "New agent - purpose TBD", - "email": email, - "status": "active", - "created": today, - "last_active": today, - } + today = datetime.now().strftime("%Y-%m-%d") + entry = { + "name": branch_name, + "path": str(branch_path), + "profile": profile, + "description": purpose or "New agent - purpose TBD", + "email": email, + "status": "active", + "created": today, + "last_active": today, + } - # Add entry — handle both dict and list formats - if isinstance(branches, dict): - branches[branch_name] = entry - else: - branches.append(entry) - registry["branches"] = branches - registry["metadata"]["total_branches"] = len(_branches_as_list(branches)) + if isinstance(branches, dict): + branches[branch_name] = entry + else: + branches.append(entry) + registry["branches"] = branches + registry["metadata"]["total_branches"] = len(_branches_as_list(branches)) - json_handler.log_operation("registry_updated", data={"branch": branch_name}) + json_handler.log_operation("registry_updated", data={"branch": branch_name}) - return save_registry(registry_path, registry) + return save_registry(registry_path, registry) + finally: + fcntl.flock(lock_fd, fcntl.LOCK_UN) + lock_fd.close() def fix_passport_registry_id(branch_dir: Path, registry_path: Path) -> bool: diff --git a/src/aipass/spawn/templates/builder/.spawn/.template_registry.json b/src/aipass/spawn/templates/builder/.spawn/.template_registry.json index 82b32f2f..802ef0ed 100644 --- a/src/aipass/spawn/templates/builder/.spawn/.template_registry.json +++ b/src/aipass/spawn/templates/builder/.spawn/.template_registry.json @@ -1,7 +1,7 @@ { "metadata": { "version": "1.0.0", - "last_updated": "2026-04-26", + "last_updated": "2026-04-28", "description": "Template file tracking registry for ID-based updates" }, "files": { diff --git a/src/aipass/spawn/tests/test_contracts.py b/src/aipass/spawn/tests/test_contracts.py index 9cf54645..577f1606 100644 --- a/src/aipass/spawn/tests/test_contracts.py +++ b/src/aipass/spawn/tests/test_contracts.py @@ -45,7 +45,7 @@ class TestExceptionContracts: def test_invalid_write_caught(self, tmp_path): """write_json catches OSError and returns False, never raises.""" f = tmp_path / "test.json" - with patch.object(Path, "write_text", side_effect=OSError("disk full")): + with patch("os.write", side_effect=OSError("disk full")): result = write_json(f, {"data": True}) assert result is False diff --git a/src/aipass/spawn/tests/test_spawn.py b/src/aipass/spawn/tests/test_spawn.py index f483bc9b..77553e64 100644 --- a/src/aipass/spawn/tests/test_spawn.py +++ b/src/aipass/spawn/tests/test_spawn.py @@ -20,7 +20,13 @@ from aipass.spawn.apps.handlers.placeholders import ( replace_placeholders, validate_no_placeholders, ) -from aipass.spawn.apps.handlers.registry import load_registry, add_to_registry, save_registry, get_next_citizen_number +from aipass.spawn.apps.handlers.registry import ( + load_registry, + add_to_registry, + save_registry, + get_next_citizen_number, + _validate_path_containment, +) @pytest.fixture @@ -78,16 +84,20 @@ class TestRegistry: assert data["metadata"]["total_branches"] == 0 assert data["branches"] == [] - def test_add_and_load(self, tmp_registry): - result = add_to_registry(tmp_registry, "TEST", "/tmp/test", "Workshop", "@test", "A test") + def test_add_and_load(self, tmp_path, tmp_registry): + branch_path = tmp_path / "test" + branch_path.mkdir() + result = add_to_registry(tmp_registry, "TEST", str(branch_path), "Workshop", "@test", "A test") assert result is True data = load_registry(tmp_registry) assert len(data["branches"]) == 1 assert data["branches"][0]["name"] == "TEST" - def test_no_duplicates(self, tmp_registry): - add_to_registry(tmp_registry, "X", "/tmp/x", "W", "@x") - result = add_to_registry(tmp_registry, "X", "/tmp/x", "W", "@x") + def test_no_duplicates(self, tmp_path, tmp_registry): + branch_path = tmp_path / "x" + branch_path.mkdir() + add_to_registry(tmp_registry, "X", str(branch_path), "W", "@x") + result = add_to_registry(tmp_registry, "X", str(branch_path), "W", "@x") assert result is False @@ -220,3 +230,48 @@ class TestGetNextCitizenNumber: def test_missing_registry(self, tmp_path): reg_path = tmp_path / "NONEXISTENT_REGISTRY.json" assert get_next_citizen_number(reg_path) == 1 + + +class TestPathContainment: + """Tests for _validate_path_containment().""" + + def test_contained_path_accepted(self, tmp_path): + reg = tmp_path / "TEST_REGISTRY.json" + branch = tmp_path / "my_agent" + assert _validate_path_containment(str(branch), reg) is True + + def test_escaped_path_rejected(self, tmp_path): + reg = tmp_path / "TEST_REGISTRY.json" + assert _validate_path_containment("/tmp/evil", reg) is False + + def test_traversal_attack_rejected(self, tmp_path): + reg = tmp_path / "TEST_REGISTRY.json" + branch = str(tmp_path / ".." / ".." / "tmp" / "evil") + assert _validate_path_containment(branch, reg) is False + + +class TestAtomicWriteAndLocking: + """Tests for atomic writes and file locking in registry operations.""" + + def test_add_to_registry_creates_lock_file(self, tmp_path): + reg = tmp_path / "TEST_REGISTRY.json" + branch = tmp_path / "agent_a" + branch.mkdir() + add_to_registry(reg, "AGENT_A", str(branch), "W", "@a") + lock = tmp_path / ".TEST_REGISTRY.lock" + assert lock.exists() + + def test_atomic_write_not_corrupted(self, tmp_path): + reg = tmp_path / "TEST_REGISTRY.json" + branch = tmp_path / "agent_b" + branch.mkdir() + add_to_registry(reg, "AGENT_B", str(branch), "W", "@b") + data = json.loads(reg.read_text()) + assert data["branches"][0]["name"] == "AGENT_B" + + def test_path_containment_blocks_add(self, tmp_path): + reg = tmp_path / "TEST_REGISTRY.json" + result = add_to_registry(reg, "EVIL", "/tmp/evil", "W", "@evil") + assert result is False + data = json.loads(reg.read_text()) if reg.exists() else {"branches": []} + assert len(data.get("branches", [])) == 0