SuperAI_Forecast / scripts /test_scratch_cleanup.py
Thang6822
feat: cancel in-flight forecast tasks immediately when switching symbol/interval and fix DXY sources config
4d10530
Raw
History Blame Contribute Delete
4.94 kB
from __future__ import annotations
import importlib.util
import io
import sys
import tempfile
import unittest
from contextlib import redirect_stdout
from pathlib import Path
from types import ModuleType
from unittest.mock import patch
PROJECT_ROOT = Path(__file__).resolve().parent.parent
MODULE_PATH = PROJECT_ROOT / "scripts" / "scratch_cleanup.py"
def _load_module() -> ModuleType:
if not MODULE_PATH.exists():
raise FileNotFoundError(f"Missing cleanup module: {MODULE_PATH}")
spec = importlib.util.spec_from_file_location("scratch_cleanup", MODULE_PATH)
if spec is None or spec.loader is None:
raise AssertionError(f"Unable to load module spec: {MODULE_PATH}")
module = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
class ScratchCleanupTests(unittest.TestCase):
def test_validate_scratch_root_requires_directory_named_scratch(self) -> None:
module = _load_module()
with tempfile.TemporaryDirectory() as temp_dir:
wrong_root = Path(temp_dir) / "not-scratch"
wrong_root.mkdir()
with self.assertRaises(ValueError):
module._validate_scratch_root(wrong_root)
def test_dry_run_collects_targets_without_mutating_filesystem(self) -> None:
module = _load_module()
with tempfile.TemporaryDirectory() as temp_dir:
scratch_root = Path(temp_dir) / "scratch"
nested_dir = scratch_root / "logs"
nested_dir.mkdir(parents=True)
artifact_file = scratch_root / "artifact.png"
artifact_file.write_text("png", encoding="utf-8")
nested_file = nested_dir / "server.log"
nested_file.write_text("log", encoding="utf-8")
result = module.cleanup_scratch(scratch_root, dry_run=True)
self.assertTrue(artifact_file.exists())
self.assertTrue(nested_file.exists())
self.assertEqual(len(result.targets), 2)
self.assertEqual(result.total_files, 2)
self.assertEqual(result.total_directories, 1)
self.assertEqual(result.error_paths, ())
def test_gitkeep_is_not_selected_as_cleanup_target(self) -> None:
module = _load_module()
with tempfile.TemporaryDirectory() as temp_dir:
scratch_root = Path(temp_dir) / "scratch"
scratch_root.mkdir()
keep_file = scratch_root / ".gitkeep"
keep_file.write_text("", encoding="utf-8")
extra_file = scratch_root / "note.txt"
extra_file.write_text("remove me", encoding="utf-8")
result = module.cleanup_scratch(scratch_root, dry_run=True)
target_names = [target.path.name for target in result.targets]
self.assertEqual(target_names, ["note.txt"])
self.assertTrue(keep_file.exists())
def test_apply_recycles_only_top_level_entries(self) -> None:
module = _load_module()
with tempfile.TemporaryDirectory() as temp_dir:
scratch_root = Path(temp_dir) / "scratch"
nested_dir = scratch_root / "nested"
nested_dir.mkdir(parents=True)
top_file = scratch_root / "preview.mp4"
top_file.write_text("video", encoding="utf-8")
nested_file = nested_dir / "inside.txt"
nested_file.write_text("inside", encoding="utf-8")
recycled_paths: list[Path] = []
def _fake_recycle(path: Path) -> None:
recycled_paths.append(path)
with patch.object(module, "_move_to_recycle_bin", side_effect=_fake_recycle):
result = module.cleanup_scratch(scratch_root, dry_run=False)
self.assertEqual(sorted(path.name for path in recycled_paths), ["nested", "preview.mp4"])
self.assertEqual(len(result.targets), 2)
def test_apply_aggregates_recycle_errors_and_returns_non_zero_exit(self) -> None:
module = _load_module()
with tempfile.TemporaryDirectory() as temp_dir:
scratch_root = Path(temp_dir) / "scratch"
scratch_root.mkdir()
(scratch_root / "a.log").write_text("a", encoding="utf-8")
(scratch_root / "b.log").write_text("b", encoding="utf-8")
def _fake_recycle(path: Path) -> None:
if path.name == "a.log":
raise OSError("locked")
stdout_buffer = io.StringIO()
with patch.object(module, "_move_to_recycle_bin", side_effect=_fake_recycle):
with redirect_stdout(stdout_buffer):
exit_code = module.main(["--root", str(scratch_root), "--apply"])
self.assertEqual(exit_code, 1)
output = stdout_buffer.getvalue()
self.assertIn("a.log", output)
self.assertIn("Errors", output)
if __name__ == "__main__":
unittest.main()