Files
AI-Profile-Router/dev/test_athena_operator.py
T

266 lines
13 KiB
Python

#!/usr/bin/env python3
"""Offline safety and workflow tests for the Athena Operator executor."""
from __future__ import annotations
import importlib.util
import json
import tempfile
import time
import unittest
from pathlib import Path
SOURCE = Path(__file__).parents[1] / "platform" / "operator" / "athena_operatord.py"
def load_module():
spec = importlib.util.spec_from_file_location("athena_operatord_tested", SOURCE)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
return module
class OperatorTests(unittest.TestCase):
def setUp(self):
self.module = load_module()
self.temporary = tempfile.TemporaryDirectory()
root = Path(self.temporary.name)
self.repo, self.stack = root / "repository", root / "stack"
self.models, self.state = root / "models", root / "state"
self.staging = root / "staging"
for path in (self.repo, self.stack, self.models, self.state, self.staging):
path.mkdir()
(self.repo / ".git").mkdir()
(self.repo / "docs").mkdir()
(self.stack / "docs").mkdir()
for base in (self.repo, self.stack):
(base / "docs" / "test.md").write_text("before\n", encoding="utf-8")
self.module.REPOSITORY = self.repo.resolve()
self.module.STACK = self.stack.resolve()
self.module.MODELS = self.models.resolve()
self.module.STATE = self.state.resolve()
self.module.STAGING_ROOTS = (self.staging.resolve(),)
self.module.audit = lambda *args, **kwargs: None
self.module.run = lambda argv, **kwargs: {"argv": argv, "exit_code": 0, "output": "ok"}
def tearDown(self):
self.temporary.cleanup()
def prepare_file(self):
before = self.module.sha((self.repo / "docs" / "test.md").read_bytes())
return self.module.prepare({
"operation": "file_update",
"payload": {"files": [{"path": "docs/test.md", "content": "after\n", "expected_sha256": before}]},
})
def test_preview_changes_nothing(self):
proposal = self.prepare_file()
self.assertEqual((self.repo / "docs" / "test.md").read_text(), "before\n")
self.assertIn("-before", proposal["preview"])
self.assertIn("+after", proposal["preview"])
def test_exact_later_confirmation_updates_repo_and_deploy_tree(self):
proposal = self.prepare_file()
result = self.module.execute({"ticket": proposal["ticket"], "confirmation": proposal["required_confirmation"]})
self.assertEqual((self.repo / "docs" / "test.md").read_text(), "after\n")
self.assertEqual((self.stack / "docs" / "test.md").read_text(), "after\n")
self.assertEqual(result["operation"], "file_update")
with self.assertRaises(ValueError):
self.module.execute({"ticket": proposal["ticket"], "confirmation": proposal["required_confirmation"]})
def test_direct_change_applies_authorized_work_without_ticket(self):
before = self.module.sha((self.repo / "docs" / "test.md").read_bytes())
result = self.module.change({
"operation": "patch_update",
"payload": {"files": [{
"path": "docs/test.md", "expected_sha256": before,
"patch": "@@ -1 +1 @@\n-before\n+after\n",
}]},
})
self.assertEqual(result["operation"], "patch_update")
self.assertEqual((self.repo / "docs" / "test.md").read_text(), "after\n")
self.assertEqual((self.stack / "docs" / "test.md").read_text(), "after\n")
def test_single_worktree_is_written_only_once(self):
self.module.STACK = self.repo
before = self.module.sha((self.repo / "docs" / "test.md").read_bytes())
result = self.module.change({
"operation": "patch_update",
"payload": {"files": [{
"path": "docs/test.md", "expected_sha256": before,
"patch": "@@ -1 +1 @@\n-before\n+after\n",
}]},
})
self.assertEqual(result["operation"], "patch_update")
self.assertEqual((self.repo / "docs" / "test.md").read_text(), "after\n")
def test_wrong_confirmation_and_expired_ticket_are_rejected(self):
proposal = self.prepare_file()
with self.assertRaises(PermissionError):
self.module.execute({"ticket": proposal["ticket"], "confirmation": "yes"})
path = self.state / "pending" / f"{proposal['ticket']}.json"
record = json.loads(path.read_text())
record["expires"] = int(time.time()) - 1
self.module.json_write(path, record)
with self.assertRaises(PermissionError):
self.module.execute({"ticket": proposal["ticket"], "confirmation": proposal["required_confirmation"]})
def test_source_drift_blocks_apply(self):
proposal = self.prepare_file()
(self.repo / "docs" / "test.md").write_text("drift\n")
with self.assertRaises(RuntimeError):
self.module.execute({"ticket": proposal["ticket"], "confirmation": proposal["required_confirmation"]})
def test_patch_update_is_compact_and_updates_both_trees(self):
before = self.module.sha((self.repo / "docs" / "test.md").read_bytes())
proposal = self.module.prepare({
"operation": "patch_update",
"payload": {"files": [{
"path": "docs/test.md", "expected_sha256": before,
"patch": "--- a/docs/test.md\n+++ b/docs/test.md\n@@ -1 +1 @@\n-before\n+after\n",
}]},
})
self.assertIn("+after", proposal["preview"])
result = self.module.execute({"ticket": proposal["ticket"], "confirmation": proposal["required_confirmation"]})
self.assertEqual(result["operation"], "patch_update")
self.assertEqual((self.repo / "docs" / "test.md").read_text(), "after\n")
self.assertEqual((self.stack / "docs" / "test.md").read_text(), "after\n")
def test_patch_update_rejects_wrong_context_and_drift(self):
before = self.module.sha((self.repo / "docs" / "test.md").read_bytes())
with self.assertRaises(RuntimeError):
self.module.normalise_operation("patch_update", {"files": [{
"path": "docs/test.md", "expected_sha256": before,
"patch": "@@ -1 +1 @@\n-wrong\n+after\n",
}]})
with self.assertRaises(RuntimeError):
self.module.normalise_operation("patch_update", {"files": [{
"path": "docs/test.md", "expected_sha256": "0" * 64,
"patch": "@@ -1 +1 @@\n-before\n+after\n",
}]})
def test_mcp_release_normalises_one_complete_workflow(self):
before = self.module.sha((self.repo / "docs" / "test.md").read_bytes())
payload, preview = self.module.normalise_operation("mcp_release", {
"files": [{"path": "docs/test.md", "expected_sha256": before, "patch": "@@ -1 +1 @@\n-before\n+after\n"}],
"services": ["mcp-example"], "message": "Add example MCP service",
"checks": ["operator-tests"], "create_recovery": False,
})
self.assertEqual(payload["services"], ["mcp-example"])
self.assertEqual(payload["paths"], ["docs/test.md"])
self.assertIn("ONE MCP RELEASE", preview)
def test_mcp_release_imports_reviewed_staging_file_by_sha(self):
source = self.staging / "server.py"
source.write_text("print('reviewed')\n", encoding="utf-8")
digest = self.module.sha(source.read_bytes())
payload, preview = self.module.normalise_operation("mcp_release", {
"imports": [{"source": str(source), "path": "server.py", "expected_source_sha256": digest}],
"services": ["mcp-example"], "message": "Import reviewed MCP source",
"checks": ["operator-tests"], "create_recovery": False, "hermes_sync": True,
})
self.assertEqual(payload["files"][0]["content"], "print('reviewed')\n")
self.assertTrue(payload["hermes_sync"])
self.assertIn(f"sha256={digest}", preview)
def test_staged_import_rejects_wrong_sha_and_unapproved_path(self):
source = self.staging / "server.py"
source.write_text("safe\n", encoding="utf-8")
with self.assertRaises(RuntimeError):
self.module.staged_file({"source": str(source), "path": "server.py", "expected_source_sha256": "0" * 64})
outside = Path(self.temporary.name) / "outside.py"
outside.write_text("unsafe\n", encoding="utf-8")
with self.assertRaises(PermissionError):
self.module.staged_file({"source": str(outside), "path": "server.py", "expected_source_sha256": self.module.sha(outside.read_bytes())})
def test_search_source_has_python_fallback_without_rg(self):
(self.repo / "docs" / "needle.md").write_text("alpha\nneedle here\nomega\n", encoding="utf-8")
original = self.module.shutil.which
self.module.shutil.which = lambda name: None
try:
result = self.module.search_source({"query": "needle"})
finally:
self.module.shutil.which = original
self.assertEqual(result["engine"], "python")
self.assertIn("docs/needle.md:2:needle here", result["matches"])
def test_structured_power_operations_do_not_exist(self):
self.assertNotIn("shell", self.module.ALLOWED_OPERATIONS)
for operation in ("shutdown", "reboot", "ssh", "network", "command"):
with self.assertRaises(ValueError):
self.module.normalise_operation(operation, {})
def test_general_terminal_runs_normal_work(self):
calls = []
self.module.run = lambda argv, **kwargs: calls.append((argv, kwargs)) or {
"argv": argv, "exit_code": 0, "output": "ok"
}
result = self.module.terminal({"command": "docker ps", "cwd": str(self.stack)})
self.assertEqual(result["exit_code"], 0)
self.assertEqual(calls[0][0], ["/bin/bash", "-lc", "docker ps"])
self.assertEqual(result["reachability_guard"], "active")
def test_general_terminal_blocks_reachability_changes(self):
blocked = (
"shutdown -h now", "systemctl restart ssh", "iptables -F",
"ip route del default", "umount /data", "nano /etc/ssh/sshd_config",
"docker restart mike-ai-wireguard-gateway",
)
for command in blocked:
with self.subTest(command=command), self.assertRaises(PermissionError):
self.module.terminal({"command": command, "cwd": str(self.stack)})
def test_wireguard_gateway_cannot_be_stopped(self):
with self.assertRaises(PermissionError):
self.module.normalise_operation("container_action", {"action": "stop", "containers": ["mike-ai-wireguard-gateway"]})
def test_only_huggingface_model_downloads_are_accepted(self):
with self.assertRaises(PermissionError):
self.module.normalise_operation("model_download", {"url": "https://evil.invalid/model.gguf", "destination": "model.gguf"})
payload, _ = self.module.normalise_operation("model_download", {"url": "https://huggingface.co/org/repo/resolve/main/model.gguf", "destination": "qwen/model.gguf"})
self.assertEqual(payload["destination"], "qwen/model.gguf")
def test_openwebui_sync_is_structured_and_requires_a_ticket(self):
payload, preview = self.module.normalise_operation("openwebui_sync", {})
self.assertEqual(payload, {})
self.assertIn("Synchronise", preview)
def test_git_publish_preview_is_bound_to_reviewed_worktree(self):
calls = []
def fake_run(argv, **kwargs):
calls.append(argv)
if argv[:3] == ["git", "status", "--short"]:
return {"argv": argv, "exit_code": 0, "output": " M compose.yaml"}
if argv[:3] == ["git", "diff", "--stat"]:
return {"argv": argv, "exit_code": 0, "output": " compose.yaml | 2 +-"}
return {"argv": argv, "exit_code": 0, "output": "ok"}
self.module.run = fake_run
payload, preview = self.module.normalise_operation("git_publish", {"message": "Update Athena platform", "paths": ["compose.yaml"]})
self.assertEqual(payload["reviewed_status"], "M compose.yaml")
self.assertEqual(payload["paths"], ["compose.yaml"])
self.assertIn("compose.yaml | 2 +-", preview)
def test_git_publish_rejects_empty_repository(self):
self.module.run = lambda argv, **kwargs: {"argv": argv, "exit_code": 0, "output": ""}
with self.assertRaises(RuntimeError):
self.module.normalise_operation("git_publish", {"message": "Update Athena platform", "paths": ["compose.yaml"]})
def test_git_publish_requires_explicit_safe_paths(self):
with self.assertRaises(ValueError):
self.module.normalise_operation("git_publish", {"message": "Update Athena platform"})
with self.assertRaises(ValueError):
self.module.normalise_operation("git_publish", {"message": "Update Athena platform", "paths": ["compose.yaml", "compose.yaml"]})
with self.assertRaises((ValueError, PermissionError)):
self.module.normalise_operation("git_publish", {"message": "Update Athena platform", "paths": [".git/config"]})
def test_path_traversal_and_protected_paths_are_rejected(self):
for path in ("../etc/passwd", ".git/config", "/etc/passwd", "secrets/key"):
with self.subTest(path=path), self.assertRaises((ValueError, PermissionError)):
self.module.safe_relative(path)
if __name__ == "__main__":
unittest.main(verbosity=2)