| 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): |
| |
| 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 |
| 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 |
| 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() |
|
|