| 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 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) |
| |
| 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") |
| 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') |
| |
| 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 |
| |
| 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' |
| |
| AA = "ACDEFGHIKLMNPQRSTVWY" |
| subs_pos2 = [f'V2{a}' for a in AA.replace('V', '')] |
| subs_pos5 = [f'L5{a}' for a in AA.replace('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) |
|
|
|
|
| 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") |
| |
| 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") |
| 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() |
| |
| self.assertNotIsInstance(d.image, str) |
|
|
|
|
| if __name__ == '__main__': |
| unittest.main() |
|
|