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