carbon-a-database-explorer / tests /test_remote_catalog.py
cgeorgiaw's picture
cgeorgiaw HF Staff
Add SQLite accession index and bounded on-demand bucket retrieval
f0190da verified
Raw History Blame Contribute Delete
6.68 kB
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))