import io import json from pathlib import Path import sqlite3 import tempfile from types import SimpleNamespace import unittest from unittest.mock import Mock, patch import pyarrow as pa import pyarrow.parquet as pq from remote_catalog import RemoteCatalog, RemoteReadError from catalog import segment_metadata class CountingBuffer(io.BytesIO): def __init__(self, data): super().__init__(data) self.bytes_read = self.range_reads = 0 def read(self, n=-1): data = super().read(n) self.bytes_read += len(data) self.range_reads += 1 return data class RemoteCatalogTests(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() self.path = Path(self.temp.name) / 'catalog.sqlite' records = [] for i in range(3): version = 1 if i < 2 else 2 start = 100 if i == 1 else 0 records.append(dict(assembly_accession=f'GCA_123.{version}', record_name=f'ABC123.{version}', source_key=f'GCA_123.{version}|ABC123.{version}', organism_name='Test', division='fungi', segment_index=int(i == 1), segment_count=2 if i < 2 else 1, segment_start_bp=start, segment_end_bp=start + 5, pred_prob_positive_strand_cds=[0., .2, .4, .6, .8], pred_prob_negative_strand_cds=[1., .8, .6, .4, .2])) self.table = pa.Table.from_pylist(records) sink = pa.BufferOutputStream() pq.write_table(self.table, sink, row_group_size=1) self.parquet = sink.getvalue().to_pybytes() with sqlite3.connect(self.path) as conn: conn.executescript(''' CREATE TABLE segments(id INTEGER PRIMARY KEY, assembly_accession TEXT, record_name TEXT, segment_start_bp INTEGER, segment_end_bp INTEGER, object_path TEXT, object_hash TEXT, row_group INTEGER, row_in_group INTEGER, metadata_json TEXT); CREATE TABLE aliases(alias TEXT, segment_id INTEGER, PRIMARY KEY(alias,segment_id)); CREATE TABLE metadata(key TEXT PRIMARY KEY,value TEXT); ''') for i, record in enumerate(records): metadata = {k: v for k, v in record.items() if not k.startswith('pred_prob_')} conn.execute('INSERT INTO segments VALUES(?,?,?,?,?,?,?,?,?,?)', ( i, record['assembly_accession'], record['record_name'], record['segment_start_bp'], record['segment_end_bp'], 'annotations/test.parquet', 'original', i, 0, json.dumps(metadata))) aliases = {record['source_key'], 'GCA_123', 'ABC123', record['assembly_accession'], record['record_name']} conn.executemany('INSERT INTO aliases VALUES(?,?)', [(a, i) for a in aliases]) conn.execute('INSERT INTO metadata VALUES(?,?)', ('manifest', json.dumps({'rows': 3, 'bucket_id': 'test/bucket'}))) self.api, self.fs = Mock(), Mock() self.source = patch('remote_catalog.source_info', return_value=SimpleNamespace(xet_hash='original')) self.source_mock = self.source.start() self.reader = patch('remote_catalog.MeasuredFile', side_effect=lambda *args: CountingBuffer(self.parquet)) self.reader_mock = self.reader.start() self.catalog = RemoteCatalog(self.path, api=self.api, fs=self.fs) def tearDown(self): self.reader.stop() self.source.stop() self.temp.cleanup() def test_lookup_versions_segments_and_limit(self): self.assertEqual(self.catalog.find(' abc123.1 '), ([0, 1], 2)) self.assertEqual(self.catalog.find('GCA_123', limit=2), ([0, 1], 3)) self.assertEqual(self.catalog.lookup('ABC123.2'), [2]) self.assertEqual(self.catalog.lookup('ABC123.9'), []) self.assertEqual(self.catalog.lookup("' OR 1=1 --"), []) self.assertEqual(self.catalog.lookup('GCA_123.1|ABC123.1'), [0, 1]) def test_cached_read_and_absolute_window(self): table, cold = self.catalog.fetch(1) self.assertFalse(cold['cache_hit']) self.assertGreater(cold['bytes_read'], 0) cached, warm = self.catalog.fetch(1) self.assertTrue(warm['cache_hit']) self.assertEqual(warm['bytes_read'], 0) self.assertTrue(table.equals(cached)) self.assertEqual(self.reader_mock.call_count, 1) frame, step = self.catalog.window(1, table=table) self.assertEqual(frame['Position (bp)'].min(), 100) self.assertEqual(frame['Position (bp)'].max(), 104) self.assertEqual(self.reader_mock.call_count, 1) def test_cache_budget_and_eviction(self): self.catalog.cache_limit = pq.ParquetFile(io.BytesIO(self.parquet)).read_row_group(0).nbytes self.catalog.fetch(0) self.catalog.fetch(1) self.assertLessEqual(self.catalog.cache_bytes, self.catalog.cache_limit) self.assertEqual(len(self.catalog.cache), 1) _, stats = self.catalog.fetch(0) self.assertFalse(stats['cache_hit']) def test_stale_source_rejected_before_and_after_read(self): self.source_mock.return_value = SimpleNamespace(xet_hash='changed') with self.assertRaisesRegex(RemoteReadError, 'changed'): self.catalog.fetch(0) self.assertEqual(self.reader_mock.call_count, 0) self.source_mock.side_effect = [SimpleNamespace(xet_hash='original'), SimpleNamespace(xet_hash='changed')] with self.assertRaisesRegex(RemoteReadError, 'changed'): self.catalog.fetch(0) self.assertEqual(len(self.catalog.cache), 0) def test_oversized_group_and_invalid_id(self): self.catalog.max_group_bytes = 1 with self.assertRaisesRegex(RemoteReadError, 'read limit'): self.catalog.fetch(0) with self.assertRaises(RemoteReadError): self.catalog.fetch(-1) def test_auth_failure_is_not_a_lookup_miss(self): self.source_mock.side_effect = RuntimeError('sensitive upstream detail') with self.assertRaisesRegex(RemoteReadError, 'Could not retrieve') as error: self.catalog.fetch(0) self.assertNotIn('sensitive', str(error.exception)) self.assertEqual(self.catalog.lookup('ABC123.2'), [2]) def test_legacy_contig_coordinates(self): record = segment_metadata({'aligned_bp_length': 35000, 'record_name': 'OLD.1'}) self.assertEqual((record['segment_start_bp'], record['segment_end_bp']), (0, 35000)) self.assertEqual((record['segment_index'], record['segment_count']), (0, 1))