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