File size: 3,879 Bytes
92ea1b5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import json
import os
import sys
import tempfile
import unittest
from unittest.mock import MagicMock, patch

sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))


class TestModelFactoryRegistry(unittest.TestCase):
    "ModelFactory register() and models()"

    def setUp(self):
        # Isolate _models for each test
        from model import ModelFactory
        self._orig = dict(ModelFactory._models)

    def tearDown(self):
        from model import ModelFactory
        ModelFactory._models.clear()
        ModelFactory._models.update(self._orig)

    @patch('model.HfApi')
    def test_register_adds_new_model(self, _mock_api):
        from model import ModelFactory
        class StubModel:
            pass
        ModelFactory.register("stub/model", StubModel)
        self.assertIn("stub/model", ModelFactory.models())
        self.assertEqual(ModelFactory._models["stub/model"], StubModel)

    @patch('model.HfApi')
    def test_models_returns_list_of_keys(self, _mock_api):
        from model import ModelFactory
        keys = ModelFactory.models()
        self.assertIsInstance(keys, list)
        self.assertTrue(len(keys) > 0)

    @patch('model.HfApi')
    def test_register_overwrites_existing(self, _mock_api):
        from model import ModelFactory
        existing_key = ModelFactory.models()[0]
        class Replacement:
            pass
        ModelFactory.register(existing_key, Replacement)
        self.assertIs(ModelFactory._models[existing_key], Replacement)


class TestFetchCache(unittest.TestCase):
    "_fetch_model_ids cache read/write/TTL"

    @patch('model.time')
    @patch('builtins.open', new_callable=unittest.mock.mock_open, read_data=json.dumps({"ts": 100, "ids": ["a/b"]}))
    def test_cache_hit_within_ttl(self, mock_file, mock_time):
        mock_time.return_value = 200  # 100s ago, within 3600 TTL
        with patch('os.path.exists', return_value=True):
            from model import _fetch_model_ids
            result = _fetch_model_ids()
            self.assertEqual(result, ["a/b"])

    @patch('model.time')
    @patch('builtins.open', new_callable=unittest.mock.mock_open, read_data=json.dumps({"ts": 100, "ids": ["a/b"]}))
    def test_cache_miss_expired(self, mock_file, mock_time):
        mock_time.return_value = 5000  # well past 3600 TTL
        fallback_ids = ["fallback/model"]
        mock_api = MagicMock()
        mock_api.list_models.side_effect = [
            [MagicMock(id="fallback/model")],
            []
        ]
        with patch('os.path.exists', return_value=True), \
             patch('model.HfApi', return_value=mock_api), \
             tempfile.TemporaryDirectory() as tmpdir:
            import model
            orig_cache = model._MODEL_CACHE
            model._MODEL_CACHE = os.path.join(tmpdir, "cache.json")
            try:
                result = model._fetch_model_ids()
                self.assertIn("fallback/model", result)
            finally:
                model._MODEL_CACHE = orig_cache


class TestScoringIndexExtraction(unittest.TestCase):
    "Verify digit extraction from mutation strings (used in vectorized scoring)"

    def _extract_idx(self, m):
        """Replicate the logic from ESMModel.run_model"""
        return int(''.join(c for c in m if c.isdigit())) - 1

    def test_simple_two_digit(self):
        self.assertEqual(self._extract_idx("V2A"), 1)

    def test_three_digit_position(self):
        self.assertEqual(self._extract_idx("R218K"), 217)

    def test_single_digit_position(self):
        self.assertEqual(self._extract_idx("M1D"), 0)

    def test_various_positions(self):
        cases = [("A10F", 9), ("L100W", 99), ("G500S", 499)]
        for mut, expected in cases:
            with self.subTest(mut=mut):
                self.assertEqual(self._extract_idx(mut), expected)


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