Merge pull request #492 from AIOSAI/work/spawn
feat(spawn): fix(spawn): atomic registry writes + file locking + path containment (issue #490)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user