"""CPU-only tests: renewal uses salloc; cancellation/failure never renews.""" import shlex import subprocess import tempfile import unittest from pathlib import Path from unittest.mock import MagicMock, patch from ngram_interactive_supervisor import Supervisor class SupervisorTests(unittest.TestCase): def controller(self, root): item = Supervisor.__new__(Supervisor) item.root = root item.source = root / "source" item.job = "1" item.ssh = ["ssh"] item.login = "login" item.status = MagicMock() item.child = MagicMock() item.child.poll.return_value = 1 item.child.returncode = 1 item.maintenance = None item.start_segment = MagicMock(return_value=root / "segments/1") return item def test_timeout_renews(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) controller.job_state = MagicMock(return_value="TIMEOUT") controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal")) with self.assertRaisesRegex(RuntimeError, "reached renewal"): controller.run() controller.allocate.assert_called_once() def test_cancellation_does_not_renew(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) controller.job_state = MagicMock(return_value="CANCELLED") controller.allocate = MagicMock() with self.assertRaisesRegex(RuntimeError, "not renewing"): controller.run() controller.allocate.assert_not_called() def test_long_completing_then_timeout_renews(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) controller.job_state = MagicMock( side_effect=["COMPLETING"] * 20 + ["UNKNOWN", "TIMEOUT"] ) controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal")) with patch("ngram_interactive_supervisor.time.sleep") as sleep: with self.assertRaisesRegex(RuntimeError, "reached renewal"): controller.run() self.assertEqual(sleep.call_count, 21) controller.allocate.assert_called_once() def test_accounting_timeout_overrides_completing(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) controller.remote = MagicMock(side_effect=["COMPLETING", "TIMEOUT|"]) self.assertEqual(controller.job_state(), "TIMEOUT") def test_expired_job_missing_from_squeue_uses_accounting(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) controller.remote = MagicMock(side_effect=[ subprocess.CalledProcessError(1, "squeue"), "TIMEOUT|" ]) self.assertEqual(controller.job_state(), "TIMEOUT") def test_transient_query_error_retries(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) controller.job_state = MagicMock(side_effect=[ subprocess.TimeoutExpired("ssh", 30), "TIMEOUT" ]) with patch("ngram_interactive_supervisor.time.sleep"): self.assertEqual(controller.wait_for_allocation_exit(), "TIMEOUT") def test_dead_trainer_in_running_allocation_does_not_renew(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) controller.job_state = MagicMock(return_value="RUNNING") controller.allocate = MagicMock() with patch("ngram_interactive_supervisor.time.monotonic", side_effect=[0, 61]): self.assertEqual(controller.wait_for_allocation_exit(), "RUNNING") controller.allocate.assert_not_called() def test_restart_after_timeout_does_not_relaunch_old_segment(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) (controller.root / "segments" / controller.job).mkdir(parents=True) controller.job_state = MagicMock(return_value="TIMEOUT") controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal")) with self.assertRaisesRegex(RuntimeError, "reached renewal"): controller.run() controller.start_segment.assert_not_called() controller.allocate.assert_called_once() def test_restart_after_cancellation_does_not_renew(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) (controller.root / "segments" / controller.job).mkdir(parents=True) controller.job_state = MagicMock(return_value="CANCELLED") controller.allocate = MagicMock() with self.assertRaisesRegex(RuntimeError, "not renewing"): controller.run() controller.start_segment.assert_not_called() controller.allocate.assert_not_called() def test_allocation_command_and_fairshare(self): with tempfile.TemporaryDirectory() as directory: controller = self.controller(Path(directory)) (controller.root / "allocations").mkdir() controller.remote = MagicMock(return_value="low|user|||||0.1|\nbest|user|||||0.9|") def launch(command, **kwargs): allocation = shlex.split(command[-1]) self.assertEqual(allocation[0], "salloc") for flag in ("--account=best", "--nodes=1", "--gpus-per-node=4", "--time=04:00:00", "--qos=interactive"): self.assertIn(flag, allocation) (Path(allocation[-1]) / "job_id").write_text("90001\n") return MagicMock() with patch("ngram_interactive_supervisor.subprocess.Popen", side_effect=launch): controller.allocate() self.assertEqual(controller.job, "90001") if __name__ == "__main__": unittest.main()