"""SQLite metadata lookup with measured, bounded reads from the annotation bucket.""" from collections import OrderedDict from contextlib import closing from functools import lru_cache import json from pathlib import Path import sqlite3 import threading import time import zlib from huggingface_hub import HfApi, HfFileSystem, get_token from huggingface_hub.hf_file_system import HfFileSystemFile import pyarrow.parquet as pq from catalog import Catalog, ROOT, PROBS, normalize, unversioned, segment_metadata BUCKET = "HuggingFaceBio/genbank-annotations" def source_info(api, bucket, path): return next((e for e in api.list_bucket_tree(bucket, path, recursive=False) if e.path == path and e.type == "file"), None) class MeasuredFile(HfFileSystemFile): """Count returned range bytes, excluding HTTP headers and retries.""" def __init__(self, fs, path): self.bytes_read = 0 self.range_reads = 0 super().__init__(fs, path, mode="rb", block_size=1024 * 1024, cache_type="none") def _fetch_range(self, start, end): data = super()._fetch_range(start, end) self.bytes_read += len(data) self.range_reads += 1 return data class RemoteReadError(ValueError): pass class Records: def __init__(self, catalog): self.catalog = catalog def __len__(self): return self.catalog.manifest["rows"] def __getitem__(self, index): return self.catalog.record(int(index)) class RemoteCatalog(Catalog): def __init__(self, path=None, cache_bytes=512 * 1024**2, max_group_bytes=512 * 1024**2, api=None, fs=None): if path is None: from catalog_snapshot import catalog_path path = catalog_path() self.path = Path(path) self.cache_limit = cache_bytes self.max_group_bytes = max_group_bytes self.api = api or HfApi(token=get_token()) self.fs = fs or HfFileSystem(token=get_token()) self.cache = OrderedDict() self.cache_bytes = 0 self.lock = threading.Lock() with closing(self.connect()) as conn: self.manifest = json.loads(conn.execute("SELECT value FROM metadata WHERE key='manifest'").fetchone()[0]) self.records = Records(self) def connect(self): conn = sqlite3.connect(self.path.resolve().as_uri() + "?mode=ro&immutable=1", uri=True) conn.row_factory = sqlite3.Row return conn @lru_cache(maxsize=512) def entry(self, index): with closing(self.connect()) as conn: row = conn.execute("SELECT * FROM segments WHERE id=?", (int(index),)).fetchone() if row is None: raise RemoteReadError("Choose a segment from the current index.") return dict(row) def record(self, index): entry = self.entry(index) if self.manifest.get("schema_version", 1) >= 3: result = json.loads(entry["context_json"]) for key in ("record_name", "aligned_bp_length", "segment_start_bp", "segment_end_bp", "segment_index", "segment_count"): result[key] = entry[key] result["segment_bp_length"] = entry["segment_end_bp"] - entry["segment_start_bp"] result["source_key"] = entry["source_key"] or result["assembly_accession"] + "|" + result["record_name"] return result value = entry["metadata_json"] return json.loads(zlib.decompress(value) if isinstance(value, bytes) else value) def browse_ids(self, limit=200): with closing(self.connect()) as conn: return [r[0] for r in conn.execute("SELECT id FROM segments ORDER BY id LIMIT ?", (limit,))] def find(self, accession, limit=200): key = normalize(accession) if not key: return [], 0 if self.manifest.get("schema_version", 1) >= 3: return self.find_compact(key, limit) # Aliases include exact versions and versionless IDs; never strip a query's version. with closing(self.connect()) as conn: count = conn.execute("SELECT count(*) FROM aliases WHERE alias=?", (key,)).fetchone()[0] ids = [r[0] for r in conn.execute( "SELECT s.id FROM aliases a JOIN segments s ON s.id=a.segment_id " "WHERE a.alias=? ORDER BY s.assembly_accession,s.record_name,s.segment_start_bp,s.id LIMIT ?", (key, limit))] return ids, count def find_compact(self, key, limit): # Exact and versionless IDs use separate B-tree indexes. Assembly and # contig lookups never require a full scan of the segment table. with closing(self.connect()) as conn: if "|" in key: assembly, record = key.split("|", 1) query = ("SELECT s.id FROM segment_data s JOIN contexts c ON c.id=s.context_id " "WHERE c.assembly_accession=? AND s.record_name=? COLLATE NOCASE " "AND s.source_key IS NULL UNION SELECT id FROM segment_data " "WHERE source_key=? COLLATE NOCASE") params = (assembly, record, key) elif key.startswith(("GCA_", "GCF_")): column = "assembly_accession" if unversioned(key) != key else "assembly_base" query = (f"SELECT s.id FROM contexts c JOIN segment_data s ON s.context_id=c.id WHERE c.{column}=?") params = (key,) else: column = "record_name" if unversioned(key) != key else "record_base" query = f"SELECT id FROM segment_data WHERE {column}=? COLLATE NOCASE" # record_base is normalized; its index uses binary collation. if column == "record_base": query = "SELECT id FROM segment_data WHERE record_base=?" params = (key,) if "|" not in key: query += " UNION SELECT id FROM segment_data WHERE source_key=? COLLATE NOCASE" params += (key,) count = conn.execute(f"SELECT count(*) FROM ({query})", params).fetchone()[0] ordered = (f"SELECT s.id FROM ({query}) hits JOIN segment_data s ON s.id=hits.id " "JOIN contexts c ON c.id=s.context_id " "ORDER BY c.assembly_accession,s.record_name,s.segment_start_bp,s.id LIMIT ?") ids = [r[0] for r in conn.execute(ordered, params + (limit,))] return ids, count def lookup(self, accession): return self.find(accession)[0] def check_source(self, entry): source = source_info(self.api, self.manifest["bucket_id"], entry["object_path"]) if source is None or source.xet_hash != entry["object_hash"]: raise RemoteReadError("The bucket object has changed since indexing. Rebuild the index before retrieving this segment.") def fetch(self, index): began = time.perf_counter() entry = self.entry(int(index)) key = (entry["object_path"], entry["object_hash"], entry["row_group"]) stats = {"cache_hit": False, "bytes_read": 0, "range_reads": 0} # Serialize cache fills to bound memory and avoid duplicate remote downloads. with self.lock: if key in self.cache: table = self.cache.pop(key) self.cache[key] = table stats["cache_hit"] = True else: try: self.check_source(entry) path = f"buckets/{self.manifest['bucket_id']}/{entry['object_path']}" self.fs.invalidate_cache(path) with MeasuredFile(self.fs, path) as remote: parquet = pq.ParquetFile(remote) group = parquet.metadata.row_group(entry["row_group"]) if group.total_byte_size > self.max_group_bytes: raise RemoteReadError("This row group exceeds the 512 MiB read limit. It needs a smaller storage chunk.") table = parquet.read_row_group(entry["row_group"], use_threads=False) stats.update(bytes_read=remote.bytes_read, range_reads=remote.range_reads) self.check_source(entry) except RemoteReadError: raise except Exception as exc: raise RemoteReadError("Could not retrieve annotations from the bucket. Check bucket access and try again; this is not an accession-not-found result.") from exc if table.nbytes <= self.cache_limit: while self.cache and self.cache_bytes + table.nbytes > self.cache_limit: _, old = self.cache.popitem(last=False) self.cache_bytes -= old.nbytes self.cache[key] = table self.cache_bytes += table.nbytes result = table.slice(entry["row_in_group"], 1) if result.num_rows != 1: raise RemoteReadError("The retrieved segment does not match the index. Rebuild the index.") actual = segment_metadata(result.select([c for c in result.column_names if c not in PROBS and c != "sequence"]).to_pylist()[0]) if any(actual[k] != entry[k] for k in ("record_name", "assembly_accession", "segment_start_bp", "segment_end_bp")): raise RemoteReadError("The retrieved segment does not match the index. Rebuild the index.") stats["seconds"] = time.perf_counter() - began return result, stats def segment_table(self, index): return self.fetch(index)[0]