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