File size: 2,075 Bytes
b24bc66
 
 
 
 
 
 
 
 
380cc0c
b24bc66
380cc0c
 
 
b24bc66
380cc0c
 
 
 
b24bc66
 
 
 
 
 
 
 
 
 
 
 
 
 
380cc0c
b24bc66
 
 
 
 
 
 
 
380cc0c
 
 
 
 
 
 
 
b24bc66
 
 
 
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
import unittest


class ModelManagerTests(unittest.TestCase):
    def test_manager_caches_same_key_and_replaces_changed_key(self):
        from model_loader import ModelManager

        calls = []
        manager = ModelManager(
            loader=lambda model, path, config: calls.append((model, path, config)) or object()
        )
        first = manager.get("owner/model", "a", "pi05_ur_demo_no_state")
        self.assertIs(manager.get("owner/model", "a", "pi05_ur_demo_no_state"), first)
        second = manager.get("owner/model", "a", "pi05_ur_demo_state")
        self.assertIsNot(second, first)
        self.assertEqual(calls, [
            ("owner/model", "a", "pi05_ur_demo_no_state"),
            ("owner/model", "a", "pi05_ur_demo_state"),
        ])

    def test_failed_load_is_reported_and_can_retry(self):
        from model_loader import ModelManager, ModelUnavailableError

        attempts = 0

        def loader(*_):
            nonlocal attempts
            attempts += 1
            raise RuntimeError("bad checkpoint")

        manager = ModelManager(loader=loader)
        for _ in range(2):
            with self.assertRaisesRegex(ModelUnavailableError, "bad checkpoint"):
                manager.get("owner/model", "a", "pi05_ur_demo_no_state")
        self.assertEqual(attempts, 2)
        self.assertIn("bad checkpoint", manager.health_message)

    def test_invalid_identity_is_rejected_before_loader(self):
        from model_loader import ModelManager

        manager = ModelManager(loader=lambda *_: self.fail("loader was called"))
        with self.assertRaises(ValueError):
            manager.get("", "checkpoint", "pi05_ur_demo_no_state")

    def test_unsupported_policy_config_is_rejected_before_loader(self):
        from model_loader import ModelManager

        manager = ModelManager(loader=lambda *_: self.fail("loader was called"))
        with self.assertRaisesRegex(ValueError, "unsupported policy config"):
            manager.get("owner/model", "checkpoint", "pi05_libero")


if __name__ == "__main__":
    unittest.main()