zsp / test /test_model.py
mgtotaro's picture
add fasta support; add unittest suite
92ea1b5
Raw
History Blame Contribute Delete
3.88 kB
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()