Add bounded workspace execution snapshots

This commit is contained in:
Hamza Ayed
2026-10-03 12:58:11 +03:00
parent 82bb5189a1
commit 3fa456698e
3 changed files with 301 additions and 0 deletions
@@ -0,0 +1,115 @@
"""Security and bounds tests for preparing an isolated command snapshot."""
from __future__ import annotations
import os
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
from app import execution_snapshot
class ExecutionSnapshotTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory(
dir=Path(__file__).resolve().parents[1]
)
self.root = Path(self.temporary.name).resolve()
(self.root / "src").mkdir()
(self.root / "src" / "main.py").write_bytes(b"print('safe')\n")
def tearDown(self) -> None:
self.temporary.cleanup()
def stage(self, paths: list[str]) -> execution_snapshot.StagedWorkspace:
with patch.dict(
os.environ,
{"SOVEREIGNAI_ALLOWED_WORKSPACES": str(self.root)},
clear=False,
), patch.object(tempfile, "tempdir", str(self.root.parent)):
return execution_snapshot.stage_selected_files(self.root, paths)
def test_stages_only_explicit_files_and_reports_hash_and_size(self) -> None:
(self.root / "src" / "notes.md").write_bytes(b"read-only input\n")
staged = self.stage(["src/main.py", "src/notes.md"])
try:
self.assertEqual(
(staged.root / "src" / "main.py").read_bytes(), b"print('safe')\n"
)
self.assertFalse((staged.root / "src" / "other.py").exists())
self.assertEqual([item["path"] for item in staged.files], [
"src/main.py", "src/notes.md"
])
self.assertEqual(staged.total_bytes, len(b"print('safe')\nread-only input\n"))
self.assertTrue(all(len(str(item["sha256"])) == 64 for item in staged.files))
staged_path = staged.root
finally:
staged.close()
self.assertFalse(staged_path.exists())
def test_rejects_paths_outside_hidden_ignored_secret_and_unsupported_files(self) -> None:
(self.root / ".env").write_text("TOKEN=x", encoding="utf-8")
(self.root / "src" / "api_token.py").write_text("TOKEN=x", encoding="utf-8")
(self.root / "src" / "secrets.json").write_text("{}", encoding="utf-8")
(self.root / "src" / "image.png").write_bytes(b"image")
(self.root / ".git").mkdir()
(self.root / ".git" / "config").write_text("secret", encoding="utf-8")
for path in (
"../outside.py", ".env", ".git/config", "src/api_token.py",
"src/secrets.json", "src/image.png", "src/CON.py", "src/trailing.py."
):
with self.subTest(path=path), self.assertRaises(ValueError):
self.stage([path])
def test_rejects_duplicate_empty_and_unallowlisted_workspace(self) -> None:
for paths in ([], ["src/main.py", "src/main.py"]):
with self.subTest(paths=paths), self.assertRaises(ValueError):
self.stage(paths)
outside = self.root.parent / (self.root.name + "-outside")
outside.mkdir()
try:
with self.assertRaises(ValueError):
with patch.dict(
os.environ,
{"SOVEREIGNAI_ALLOWED_WORKSPACES": str(self.root)},
clear=False,
):
execution_snapshot.stage_selected_files(outside, ["file.py"])
finally:
outside.rmdir()
def test_enforces_per_file_and_total_size_limits(self) -> None:
large = self.root / "src" / "large.txt"
large.write_bytes(b"x" * (execution_snapshot.MAX_SNAPSHOT_FILE_BYTES + 1))
with self.assertRaisesRegex(ValueError, "512 كيلوبايت"):
self.stage(["src/large.txt"])
(self.root / "src" / "a.txt").write_bytes(b"a" * 6)
(self.root / "src" / "b.txt").write_bytes(b"b" * 6)
with patch.object(execution_snapshot, "MAX_SNAPSHOT_TOTAL_BYTES", 10):
with self.assertRaisesRegex(ValueError, "10 ميغابايت"):
self.stage(["src/a.txt", "src/b.txt"])
def test_rejects_symlinked_input_when_supported(self) -> None:
outside = self.root.parent / (self.root.name + "-linked.py")
outside.write_text("outside = True\n", encoding="utf-8")
link = self.root / "src" / "linked.py"
try:
try:
link.symlink_to(outside)
except OSError as exc:
self.skipTest(f"symlinks unavailable in this Windows context: {exc}")
with self.assertRaisesRegex(ValueError, "الروابط الرمزية"):
self.stage(["src/linked.py"])
finally:
if link.exists() or link.is_symlink():
link.unlink()
outside.unlink(missing_ok=True)
if __name__ == "__main__":
unittest.main()