Download tests/test_remote_catalog.py from HuggingFaceBio/carbon-a-database-explorer: direct link, hf CLI and curl.
- Browser
- Download file 6.68 kB
-
https://huggingface.co/spaces/HuggingFaceBio/carbon-a-database-explorer/resolve/main/tests/test_remote_catalog.py
- Command line
-
hf download hf://spaces/HuggingFaceBio/carbon-a-database-explorer/tests/test_remote_catalog.py
-
curl -L -o test_remote_catalog.py https://huggingface.co/spaces/HuggingFaceBio/carbon-a-database-explorer/resolve/main/tests/test_remote_catalog.py
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)) | |