import hashlib import json from pathlib import Path import tempfile import unittest import numpy as np import pyarrow as pa import pyarrow.parquet as pq from catalog import Catalog, ROOT class SampleTests(unittest.TestCase): @classmethod def setUpClass(cls): cls.catalog = Catalog() def test_sample_integrity_and_budget(self): manifest = self.catalog.manifest self.assertEqual(hashlib.sha256(self.catalog.path.read_bytes()).hexdigest(), manifest['sha256']) self.assertEqual(len(self.catalog.records), manifest['rows']) self.assertLessEqual(manifest['rows'], manifest['max_rows']) self.assertLessEqual(manifest['bases'], manifest['max_bp']) for i, record in enumerate(self.catalog.records): row = self.catalog.segment_table(i).to_pylist()[0] length = record['segment_end_bp'] - record['segment_start_bp'] for key in ('pred_prob_positive_strand_cds', 'pred_prob_negative_strand_cds'): values = np.asarray(row[key], dtype=float) self.assertEqual(len(values), length) self.assertTrue(np.all(np.isfinite(values))) self.assertTrue(np.all((values >= 0) & (values <= 1))) def test_real_accession_lookup(self): first = self.catalog.records[0] self.assertIn(0, self.catalog.lookup(first['source_key'])) self.assertIn(0, self.catalog.lookup(' ' + first['record_name'].lower() + ' ')) self.assertIn(0, self.catalog.lookup(first['record_name'].rsplit('.', 1)[0])) expected = [i for i, r in enumerate(self.catalog.records) if r['assembly_accession'] == first['assembly_accession']] self.assertCountEqual(self.catalog.lookup(first['assembly_accession']), expected) self.assertEqual(self.catalog.lookup(first['record_name'].rsplit('.', 1)[0] + '.999999'), []) self.assertEqual(self.catalog.lookup(''), []) self.assertEqual(self.catalog.lookup('not-an-accession'), []) class CoordinateTests(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() directory = Path(self.temp.name) rows = [] for version in (1, 2): rows.append({'assembly_accession': f'GCA_000001.{version}', 'record_name': f'TEST000001.{version}', 'source_key': f'GCA_000001.{version}|TEST000001.{version}', 'segment_index': 1, 'segment_start_bp': 100, 'segment_end_bp': 105, 'pred_prob_positive_strand_cds': [0., 0.2, 0.4, 0.6, 0.8], 'pred_prob_negative_strand_cds': [1., 0.8, 0.6, 0.4, 0.2]}) pq.write_table(pa.Table.from_pylist(rows), directory / 'sample.parquet') (directory / 'manifest.json').write_text('{}') self.catalog = Catalog(directory) def tearDown(self): self.temp.cleanup() def test_version_resolution(self): self.assertEqual(self.catalog.lookup('GCA_000001'), [0, 1]) self.assertEqual(self.catalog.lookup('GCA_000001.2'), [1]) self.assertEqual(self.catalog.lookup('TEST000001'), [0, 1]) self.assertEqual(self.catalog.lookup('GCA_000001.3'), []) def test_exact_window_keeps_both_strands_per_base(self): frame, step = self.catalog.window(0, 103, 105) self.assertEqual(step, 1) self.assertEqual(frame['Position (bp)'].tolist(), [103, 104, 103, 104]) np.testing.assert_allclose(frame[frame.Strand == '+ strand']['P(CDS)'], [0.6, 0.8]) np.testing.assert_allclose(frame[frame.Strand == '− strand']['P(CDS)'], [0.4, 0.2]) def test_binned_window_keeps_the_distribution_not_the_mean(self): """A mean hides a part-coding bin; the histogram has to keep both bands.""" rows = 4 frame, step = self.catalog.window(0, 100, 105, max_points=2, hist_rows=rows) self.assertEqual(step, 3) self.assertNotIn('Strand', frame.columns) self.assertEqual(frame['Position (bp)'].tolist()[::rows], [100, 103]) bases = frame['Bases'].to_numpy().reshape(-1, rows) # Every base is counted exactly once, and only once. self.assertEqual(bases.sum(), 5) self.assertEqual(bases.sum(axis=1).tolist(), [3, 2]) # max(P_pos, P_neg) over 100..104 is [1.0, 0.8, 0.6, 0.6, 0.8]; with four # bands those land in the top band, top, third, third, top. self.assertEqual(bases.tolist(), [[0, 0, 1, 2], [0, 0, 1, 1]]) np.testing.assert_allclose(frame['Mean P'].to_numpy()[::rows], [0.8, 0.7], atol=1e-6) def test_binned_window_conserves_every_base(self): frame, step = self.catalog.window(0, max_points=2, hist_rows=8) width = self.catalog.records[0]['segment_end_bp'] - self.catalog.records[0]['segment_start_bp'] self.assertGreater(step, 1) self.assertEqual(frame['Bases'].sum(), width) self.assertTrue((frame['P(CDS)'] > 0).all() and (frame['P(CDS)'] < 1).all()) def test_invalid_ranges(self): for start, end in [(99, 103), (100, 106), (103, 103), (104, 102)]: with self.assertRaises(ValueError): self.catalog.window(0, start, end) def test_binary_labels_preserve_exact_transitions_and_coordinates(self): frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.75) self.assertEqual(step, 1) self.assertEqual(frame.Strand.unique().tolist(), ['CDS (either strand)']) self.assertEqual(frame['Position (bp)'].tolist(), [100, 102, 104, 105]) self.assertEqual(frame['Predicted CDS'].tolist(), [1, 0, 1, 1]) self.assertTrue(set(frame['Predicted CDS']).issubset({0, 1})) def test_binary_overview_thresholds_before_binning(self): frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.85, max_points=2, max_transitions=0) self.assertEqual(step, 3) # First bin's combined mean is 0.8, but one base exceeds 0.85. self.assertEqual(frame['Position (bp)'].tolist(), [100, 103, 105]) self.assertEqual(frame['Predicted CDS'].tolist(), [1, 0, 0]) def test_threshold_limits_and_equality(self): frame, _ = self.catalog.window(0, mode='Binary labels', threshold=0) self.assertEqual(set(frame['Predicted CDS']), {1}) frame, _ = self.catalog.window(0, mode='Binary labels', threshold=1) self.assertEqual(set(frame['Predicted CDS']), {0}) table = pa.Table.from_pylist([{ 'pred_prob_positive_strand_cds': [0.5, 0.1, 0.51, 0.49, 0.5], 'pred_prob_negative_strand_cds': [0.1, 0.5, 0.1, 0.1, 0.5]}]) frame, _ = self.catalog.window(0, mode='Binary labels', threshold=0.5, table=table) self.assertEqual(frame['Position (bp)'].tolist(), [100, 102, 103, 105]) self.assertEqual(frame['Predicted CDS'].tolist(), [0, 1, 0, 0]) for threshold in [-0.1, 1.1, float('nan')]: with self.assertRaises(ValueError): self.catalog.window(0, mode='Binary labels', threshold=threshold) if __name__ == '__main__': unittest.main()