zsp / test /test_data.py
mgtotaro's picture
add fasta support; add unittest suite
92ea1b5
Raw
History Blame Contribute Delete
10.5 kB
import os
import sys
import tempfile
import unittest
from unittest.mock import MagicMock, patch
# Ensure project root is on path
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))
class FakeModel:
"Minimal model stub for testing without GPU"
def __init__(self):
self.alphabet = {c: i + 3 for i, c in enumerate("ACDEFGHIKLMNPQRSTVWY")}
self.alphabet.update({'<pad>': 1, '<eos>': 2, '<cls>': 0})
class TestParseSeq(unittest.TestCase):
"Sequence parsing (plain and FASTA)"
@patch('data.ModelFactory', return_value=FakeModel())
def test_plain_sequence(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.parse_seq("MVEQYLLEAI")
self.assertEqual(d.seq, "MVEQYLLEAI")
@patch('data.ModelFactory', return_value=FakeModel())
def test_fasta_single_line(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.parse_seq(">my protein\nMVEQYLLEAI")
self.assertEqual(d.seq, "MVEQYLLEAI")
self.assertFalse(d.seq.startswith('>'))
@patch('data.ModelFactory', return_value=FakeModel())
def test_fasta_multi_line(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.parse_seq(">seq1\nMV EQ YL LE AI")
self.assertEqual(d.seq, "MVEQYLLEAI")
@patch('data.ModelFactory', return_value=FakeModel())
def test_lowercase_converted_to_upper(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.parse_seq("mveqylleai")
self.assertEqual(d.seq, "MVEQYLLEAI")
@patch('data.ModelFactory', return_value=FakeModel())
def test_whitespace_stripped(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.parse_seq(" MV EQ YL \n LE AI ")
self.assertEqual(d.seq, "MVEQYLLEAI")
@patch('data.ModelFactory', return_value=FakeModel())
def test_invalid_characters_raises(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
with self.assertRaises(RuntimeError):
d.parse_seq("MVXYZ")
@patch('data.ModelFactory', return_value=FakeModel())
def test_empty_sequence_raises(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
with self.assertRaises(RuntimeError):
d.parse_seq("")
@patch('data.ModelFactory', return_value=FakeModel())
def test_whitespace_only_raises(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
with self.assertRaises(RuntimeError):
d.parse_seq(" \n ")
@patch('data.ModelFactory', return_value=FakeModel())
def test_fasta_header_only_raises(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
with self.assertRaises(RuntimeError):
d.parse_seq(">just a header")
class TestParseSub(unittest.TestCase):
"Substitution parsing and mode detection"
def setUp(self):
self.seq = "MVEQYLL"
@patch('data.ModelFactory', return_value=FakeModel())
def test_dms_mode(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = self.seq
d.parse_sub("2 5")
self.assertEqual(d.mode, 'SMS')
self.assertEqual(len(d.resi), 2)
self.assertIn(2, d.resi)
self.assertIn(5, d.resi)
# Each position has 19 alternatives (20 AA minus WT)
self.assertEqual(len(d.sub), 38)
@patch('data.ModelFactory', return_value=FakeModel())
def test_mut_explicit_mode(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = self.seq
d.parse_sub("V2A E3K")
self.assertEqual(d.mode, 'MUT')
self.assertEqual(len(d.sub), 2)
self.assertEqual(list(d.sub['0']), ['V2A', 'E3K'])
@patch('data.ModelFactory', return_value=FakeModel())
def test_mut_seq_vs_seq_mode(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = self.seq
d.parse_sub("MVEQYAL") # same length, differs at pos 6
self.assertEqual(d.mode, 'MUT')
self.assertTrue(any('L6A' in str(s) for s in d.sub['0']))
@patch('data.ModelFactory', return_value=FakeModel())
def test_tms_fallback_mode(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = self.seq
d.parse_sub("deep mutational scanning")
self.assertEqual(d.mode, 'DMS')
# All positions x all alternatives
self.assertEqual(len(d.resi), len(self.seq))
@patch('data.ModelFactory', return_value=FakeModel())
def test_dms_position_out_of_range(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = self.seq
with self.assertRaises(RuntimeError):
d.parse_sub("999")
@patch('data.ModelFactory', return_value=FakeModel())
def test_mut_wrong_wt_raises(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = self.seq
# V2A but position 2 is V — correct; try A2K where pos 2 is V not A
with self.assertRaises(RuntimeError):
d.parse_sub("A2K")
class TestSmsSort(unittest.TestCase):
"SMS output sorting and reshaping"
@patch('data.ModelFactory', return_value=FakeModel())
def test_sort_preserves_top_19_per_position(self, mock_factory):
from data import Data
import pandas as pd
d = object.__new__(Data)
d.model_name = 'test'
# _sort_sms requires exactly 19 rows per position; build minimal valid input
AA = "ACDEFGHIKLMNPQRSTVWY"
subs_pos2 = [f'V2{a}' for a in AA.replace('V', '')] # 19 alternatives (skip WT V)
subs_pos5 = [f'L5{a}' for a in AA.replace('L', '')] # 19 alternatives (skip WT L)
scores_p2 = list(range(19))
scores_p5 = list(range(18, -1, -1))
d.out = pd.DataFrame({
'0': subs_pos2 + subs_pos5,
'test': scores_p2[:19] + scores_p5
})
d.resi = [2, 5]
d._sort_sms()
self.assertEqual(d.out.shape[0], 19) # 19 rows after reshape
class TestResidueCycle(unittest.TestCase):
"resi_cycle helper"
@patch('data.ModelFactory', return_value=FakeModel())
def test_resi_cycles_correctly(self, _mock_factory):
from data import Data
import pandas as pd
d = object.__new__(Data)
d.model = FakeModel()
d.resi = [1, 3, 5]
d.out = pd.DataFrame({'x': range(9)})
cycle = d.resi_cycle()
self.assertEqual(cycle, [1, 3, 5, 1, 3, 5, 1, 3, 5])
@patch('data.ModelFactory', return_value=FakeModel())
def test_resi_truncates_to_length(self, _mock_factory):
from data import Data
import pandas as pd
d = object.__new__(Data)
d.model = FakeModel()
d.resi = [1, 2]
d.out = pd.DataFrame({'x': range(5)})
cycle = d.resi_cycle()
self.assertEqual(len(cycle), 5)
self.assertEqual(cycle[:4], [1, 2, 1, 2])
class TestStyleAndSave(unittest.TestCase):
"Table styling and CSV output"
@patch('data.ModelFactory', return_value=FakeModel())
def test_style_creates_styler_and_csv(self, _mock_factory):
from data import Data
import pandas as pd
with tempfile.TemporaryDirectory() as tmpdir:
d = object.__new__(Data)
d.model_name = 'test'
d.out = pd.DataFrame({'0': ['V2A'], 'test': [1.5]})
d.out_csv = os.path.join(tmpdir, 'out.csv')
d._style_and_save()
self.assertIsNotNone(d.out_table)
self.assertTrue(os.path.exists(d.out_csv))
class TestParseSubEdgeCases(unittest.TestCase):
"parse_sub edge cases"
@patch('data.ModelFactory', return_value=FakeModel())
def test_dms_duplicate_positions_allowed(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = "MVEQYLL"
d.parse_sub("3 3")
# Duplicate positions produce duplicate mutation sets
self.assertEqual(len(d.resi), 2)
self.assertEqual(len(d.sub), 38)
@patch('data.ModelFactory', return_value=FakeModel())
def test_mut_identical_sequences_no_muts(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = "MV"
d.parse_sub("MV") # identical — no differences
self.assertEqual(d.mode, 'MUT')
self.assertEqual(len(d.sub), 0)
@patch('data.ModelFactory', return_value=FakeModel())
def test_single_residue_sequence_tms(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.model = FakeModel()
d.seq = "A"
d.parse_sub("scan all")
self.assertEqual(d.mode, 'DMS')
self.assertEqual(len(d.resi), 1)
self.assertEqual(len(d.sub), 19)
class TestAppReturns(unittest.TestCase):
"app.py callback return values per mode"
@patch('data.ModelFactory', return_value=FakeModel())
def test_image_property_returns_path_for_tms(self, _mock_factory):
from data import Data
d = object.__new__(Data)
d.out_table = None
d.out_img_path = 'out.png'
self.assertIsInstance(d.image, str)
self.assertEqual(d.image, 'out.png')
@patch('data.ModelFactory', return_value=FakeModel())
def test_image_property_returns_styler_for_dms(self, _mock_factory):
from data import Data
import pandas as pd
d = object.__new__(Data)
d.model_name = 'test'
d.out_csv = '/tmp/x.csv'
d.out = pd.DataFrame({'0': ['V2A'], 'test': [1.5]})
d._style_and_save()
# After styling, image returns Styler, not path
self.assertNotIsInstance(d.image, str)
if __name__ == '__main__':
unittest.main()