File size: 6,320 Bytes
932bc69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
"""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()