#!/usr/bin/env python3 """ Build the website sidecars from MuSProt.db. Outputs (into --out-dir): MuSProt-index.db compact, indexed lookup tables the website queries directly sequence one row per sequence_id (catalog: counts, representative, function) member one row per chain observation (light columns only), incl. its biological binding partners (JSON list of partner keys) partner_name partner key -> type, description, UniProt (from the partner table) EM entries carry an em_label (cryo-EM, negative-stain EM, MicroED, ...) from the em_specimen table; sequence.methods and the overview count that label instead of the bare PDB method, so negative-stain maps are not reported as cryo-EM. state one row per (sequence_id, state_id) state_pair mean similarity / fidelity between two states of a sequence function_fts FTS5 over UniProt / CATH / ECOD / Pfam / top ranked functions per sequence ecod_h, ecod_f ECOD homology / family names (from ecod.latest.domains.txt) ecod_group outer-graph groups: sequences sharing the same set of ECOD H-groups overview.json dataset-level counts and histograms for the home page ECOD names (homology / family) come from ECOD's own domain list via --ecod. The main DB is only read. Everything is derived in a few SQL scans, so the edge table (100M+ rows) is aggregated inside SQLite rather than in Python. Usage: python backend/scripts/build_index.py /path/MuSProt.db --out-dir /tmp/musprot-bucket \ --ecod /path/ecod.latest.domains.txt """ from __future__ import annotations import argparse import json import math import sqlite3 import time from collections import Counter, defaultdict from pathlib import Path FUNCTION_TEXT_TOP_N = 3 FUNCTION_TEXT_MAX_CHARS = 600 MAX_PAIR_STATES = 40 # state_pair rows are kept only among a sequence's largest states INNER_MAX_NODES = 60 # precomputed inner graph: observations per sequence … INNER_MAX_STATES = 7 # … drawn from its largest states (must match app/protein/catalog.py) FIDELITY_CODE = {"identical": 0, "low": 1, "medium": 2, "high": 3} def log(msg: str) -> None: print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True) def to_float(v): try: f = float(v) return None if math.isnan(f) else f except (TypeError, ValueError): return None def to_int(v): f = to_float(v) return None if f is None else int(f) def parse_list(raw): raw = (raw or "").strip() if not raw: return [] try: val = json.loads(raw) return val if isinstance(val, list) else [] except ValueError: import ast try: val = ast.literal_eval(raw) return val if isinstance(val, list) else [] except (ValueError, SyntaxError): return [] def hist(values, edges): """Counts of values in [edges[i], edges[i+1]); last bin is closed.""" counts = [0] * (len(edges) - 1) for v in values: if v is None: continue for i in range(len(edges) - 1): if v < edges[i + 1] or i == len(edges) - 2: if v >= edges[i]: counts[i] += 1 break return counts # ─────────────────────────────────────────────────────────────── domain inputs def split_ids(raw) -> list[str]: return [x for x in (raw or "").split(";") if x] def ecod_hset(fids: str) -> tuple: """ECOD F-ids ('292.2.1.1;11.1.1.3') -> sorted homology-level set ('11.1', '292.2').""" return tuple(sorted({".".join(f.split(".")[:2]) for f in split_ids(fids)})) def load_ecod_names(path: Path | None) -> tuple[dict, dict]: """f_id -> (t_name, f_name) and h_id -> (architecture, x_name, h_name).""" fam, hom = {}, {} if path is None: return fam, hom header = None with open(path, encoding="utf-8") as fh: for line in fh: if line.startswith("#"): cols = line.lstrip("#").rstrip("\n").split("\t") if "f_id" in cols: header = cols continue parts = line.rstrip("\n").split("\t") if header is None: if "f_id" in parts: header = parts continue r = dict(zip(header, parts)) f_id = r.get("f_id", "") if not f_id or f_id in fam: continue clean = lambda v: (v or "").strip('"') if (v or "").strip('"') not in ("NO_X_NAME", "NO_H_NAME", "NO_T_NAME", "F_UNCLASSIFIED") else "" fam[f_id] = (clean(r.get("t_name")), clean(r.get("f_name"))) h_id = ".".join(f_id.split(".")[:2]) if h_id not in hom: hom[h_id] = (clean(r.get("architecture_name")), clean(r.get("x_name")), clean(r.get("h_name"))) return fam, hom # ─────────────────────────────────────────────────────────────── binding partners # partner table (pipeline step 12): one row per (target chain, partner chain). # A contact counts as a biological partner when the target side has at least # PARTNER_MIN_RES interface residues and the contact exists in biological # assembly 1 (entries without assembly records: not only via a symmetry mate). # Must match app/protein/records.py. PARTNER_MIN_RES = 5 NUCLEIC_LABEL = {"dna": "DNA", "rna": "RNA", "hybrid": "DNA/RNA hybrid"} def partner_is_biological(n_res_target, via_symmetry, in_assembly1) -> bool: if (to_int(n_res_target) or 0) < PARTNER_MIN_RES: return False if in_assembly1 in ("0", "1"): return in_assembly1 == "1" return via_symmetry == "0" def _sequence_like(desc: str) -> bool: """Nucleic-acid descriptions that just spell the sequence, e.g. DNA (5'-D(*CP*GP*...)-3').""" d = desc.strip().upper() return "*" in d or d.startswith(("5'", "DNA (", "RNA (", "DNA(", "RNA(")) or d in ("DNA", "RNA") def partner_key(ptype: str, desc: str, uniprot: str) -> str: """UniProt accession(s) when known; otherwise ':', or just the polymer type for nucleic acids whose description only spells the sequence.""" if uniprot: return uniprot if not desc or (ptype in NUCLEIC_LABEL and _sequence_like(desc)): return ptype or "other" d = " ".join(desc.lower().split()).replace("ribosomal rna", "rrna") return f"{ptype}:{d}" def load_partners(src: sqlite3.Connection): """(pdb_lower, chain) -> [(key, n_res_target)] biological partners, largest interface first; 'self' marks another copy of the same entity. Also returns key -> Counter of (type, description, uniprot) for naming.""" has = src.execute("SELECT 1 FROM sqlite_master WHERE type='table' AND name='partner'").fetchone() if not has: log(" no partner table in the DB; skipping binding partners") return {}, {} per_chain: dict = defaultdict(dict) names: dict = defaultdict(Counter) n = 0 for pdb, chain, ptype, desc, unp, nres, same, sym, asm in src.execute( "SELECT pdb_id, auth_asym_id, partner_type, partner_description, partner_uniprot, " "n_res_target, same_entity, via_symmetry, in_assembly1 FROM partner" ): n += 1 if not partner_is_biological(nres, sym, asm): continue if same == "1": key = "self" else: key = partner_key(ptype or "", desc or "", unp or "") names[key][(ptype or "", desc or "", unp or "")] += 1 d = per_chain[(pdb.lower(), chain)] d[key] = max(d.get(key, 0), to_int(nres) or 0) log(f" partner rows: {n:,}; chains with a biological partner: {len(per_chain):,}") out = {k: sorted(v.items(), key=lambda kv: -kv[1]) for k, v in per_chain.items()} return out, names def write_partner_names(dst: sqlite3.Connection, names: dict, used: Counter) -> None: dst.executescript( """ DROP TABLE IF EXISTS partner_name; CREATE TABLE partner_name ( key TEXT PRIMARY KEY, type TEXT, description TEXT, uniprot TEXT, n_chains INTEGER ); """ ) rows = [] for key, n_chains in used.items(): if key == "self": continue (ptype, desc, unp), _ = names[key].most_common(1)[0] if key in NUCLEIC_LABEL: desc = NUCLEIC_LABEL[key] rows.append((key, ptype, desc, unp or None, n_chains)) dst.executemany("INSERT INTO partner_name VALUES (?,?,?,?,?)", rows) # ─────────────────────────────────────────────────────────────── EM specimen def load_em_labels(src: sqlite3.Connection) -> dict: """pdb_lower -> em_label: from the em_specimen table (step 13 patch) or, in a DB rebuilt from the merged step 1, from node.em_label; {} when neither exists.""" has = src.execute("SELECT 1 FROM sqlite_master WHERE type='table' AND name='em_specimen'").fetchone() if has: return {p.lower(): lab for p, lab in src.execute("SELECT pdb_id, em_label FROM em_specimen") if lab} if "em_label" in {r[1] for r in src.execute("PRAGMA table_info(node)")}: return {p.lower(): lab for p, lab in src.execute( "SELECT DISTINCT pdb_id, em_label FROM node WHERE em_label <> ''") if lab} log(" no EM specimen labels in the DB; EM entries keep the bare PDB method") return {} # ─────────────────────────────────────────────────────────────── node pass NODE_SQL = """ SELECT sequence_id, uniprot_id, pdb_id, auth_asym_id, sequence_length, n_resolved_aa, resolved_coverage, binders, binding_status, experimental_method, resolution, pH, temp_K, chain_composition, non_protein_polymer_binding, initial_release_date, cath_superfamily, state_id, ranked_functions, cath_id, pfam_id, ecod_id, ecod_fid FROM node """ def build_members(src: sqlite3.Connection, dst: sqlite3.Connection, limit: int | None): dst.executescript( """ DROP TABLE IF EXISTS member; CREATE TABLE member ( sequence_id TEXT NOT NULL, state_id INTEGER, pdb_id TEXT NOT NULL, auth_asym_id TEXT NOT NULL, uniprot_id TEXT, experimental_method TEXT, resolution REAL, binding_status TEXT, chain_composition TEXT, binders TEXT, release_date TEXT, resolved_coverage REAL, n_resolved_aa INTEGER, pH REAL, temp_K REAL, cath_id TEXT, pfam_id TEXT, ecod_id TEXT, ecod_fid TEXT, partners TEXT, em_label TEXT ); """ ) partners, partner_names = load_partners(src) em_labels = load_em_labels(src) seqs: dict[str, dict] = {} functions: dict[str, list[str]] = {} ov = { "method": Counter(), "binding": Counter(), "composition": Counter(), "npp": Counter(), "year": Counter(), "cath": Counter(), "resolution": [], "length": {}, "uniprot": set(), "entries": set(), "partner_used": Counter(), "partner_type": Counter(), "with_partner": 0, "with_hetero_partner": 0, } sql = NODE_SQL + (f" LIMIT {int(limit)}" if limit else "") batch = [] n = 0 for row in src.execute(sql): (seq_id, uniprot, pdb, chain, seq_len, n_res, res_cov, binders, binding, method, resolution, ph, temp, comp, npp, date, cath_sf, state, funcs, cath_id, pfam, ecod_ids, ecod_fid) = row state_i = to_int(state) res_f = to_float(resolution) em_label = em_labels.get(pdb.lower()) method_label = em_label or method # what the site shows and counts # another entity with the same UniProt accession is still a homo contact pkeys = [] for key, _ in partners.get((pdb.lower(), chain), ()): key = "self" if uniprot and key == uniprot else key if key not in pkeys: pkeys.append(key) if pkeys: ov["with_partner"] += 1 hetero = [k for k in pkeys if k != "self"] ov["with_hetero_partner"] += bool(hetero) ov["partner_used"].update(pkeys) ov["partner_type"].update({partner_names[k].most_common(1)[0][0][0] for k in hetero}) batch.append(( seq_id, state_i, pdb, chain, uniprot or None, method or None, res_f, binding or None, comp or None, binders or None, date or None, to_float(res_cov), to_int(n_res), to_float(ph), to_float(temp), cath_id or None, pfam or None, ecod_ids or None, ecod_fid or None, json.dumps(pkeys, separators=(",", ":")) if pkeys else None, em_label, )) s = seqs.get(seq_id) if s is None: s = seqs[seq_id] = { "uniprot": Counter(), "length": to_int(seq_len), "n_obs": 0, "states": Counter(), "n_holo": 0, "entries": set(), "methods": Counter(), "cath": Counter(), "dates": [], "rep": None, "ecod": Counter(), "pfam": Counter(), } s["n_obs"] += 1 if uniprot: s["uniprot"][uniprot] += 1 s["states"][state_i] += 1 s["n_holo"] += binding == "holo" s["entries"].add(pdb) if method_label: s["methods"][method_label] += 1 if cath_sf: s["cath"][cath_sf] += 1 hs = ecod_hset(ecod_fid) if hs: s["ecod"][hs] += 1 pf = tuple(sorted(set(split_ids(pfam)))) if pf: s["pfam"][pf] += 1 if date: s["dates"].append(date) # representative: best (lowest) resolution among well-resolved chains key = ((to_float(res_cov) or 0) < 0.9, res_f if res_f is not None else 99.0) if s["rep"] is None or key < s["rep"][0]: s["rep"] = (key, pdb, chain) if seq_id not in functions and funcs: fl = [str(f).strip() for f in parse_list(funcs) if len(str(f).strip()) > 10] if fl: functions[seq_id] = fl[:FUNCTION_TEXT_TOP_N] ov["method"][method_label or "Unknown"] += 1 ov["binding"][binding or "Unknown"] += 1 ov["composition"][comp or "Unknown"] += 1 ov["npp"][npp or "None"] += 1 if date: ov["year"][date[:4]] += 1 if res_f is not None: ov["resolution"].append(res_f) if uniprot: ov["uniprot"].add(uniprot) ov["entries"].add(pdb.lower()) ov["length"][seq_id] = to_int(seq_len) n += 1 if len(batch) >= 50000: dst.executemany(f"INSERT INTO member VALUES ({','.join('?' * 21)})", batch) batch.clear() log(f" node rows: {n:,}") if batch: dst.executemany(f"INSERT INTO member VALUES ({','.join('?' * 21)})", batch) log(f" node rows total: {n:,}; sequences: {len(seqs):,}") write_partner_names(dst, partner_names, ov["partner_used"]) top = [] for key, c in ov["partner_used"].most_common(41): if key == "self": continue (ptype, desc, unp), _ = partner_names[key].most_common(1)[0] top.append({"key": key, "type": ptype, "label": NUCLEIC_LABEL.get(key, desc), "uniprot": unp or None, "value": c}) ov["partners"] = { "observations_with_partner": ov["with_partner"], "observations_with_hetero_partner": ov["with_hetero_partner"], "observations_homo_only": ov["with_partner"] - ov["with_hetero_partner"], "distinct_partners": len([k for k in ov["partner_used"] if k != "self"]), "by_type": [{"label": k, "value": v} for k, v in ov["partner_type"].most_common()], "top": top[:40], } dst.executescript( """ CREATE INDEX idx_member_seq ON member(sequence_id, state_id); CREATE INDEX idx_member_pdb ON member(pdb_id COLLATE NOCASE, auth_asym_id); CREATE INDEX idx_member_uniprot ON member(uniprot_id COLLATE NOCASE); """ ) return seqs, functions, ov, n def consensus(counter: Counter) -> tuple: """Most common non-empty assignment set across a sequence's chains (ties: larger set).""" if not counter: return () return max(counter.items(), key=lambda kv: (kv[1], len(kv[0]), kv[0]))[0] def write_sequences(dst: sqlite3.Connection, seqs: dict, functions: dict, ecod_names: tuple): dst.executescript( """ DROP TABLE IF EXISTS sequence; CREATE TABLE sequence ( sequence_id TEXT PRIMARY KEY, uniprot_id TEXT, length INTEGER, n_obs INTEGER, n_states INTEGER, n_entries INTEGER, n_holo INTEGER, n_apo INTEGER, n_cross_pairs INTEGER, largest_state INTEGER, rep_pdb TEXT, rep_chain TEXT, cath_superfamily TEXT, methods TEXT, first_release TEXT, last_release TEXT, top_function TEXT, ecod_hset TEXT, pfam_ids TEXT ); DROP TABLE IF EXISTS state; CREATE TABLE state ( sequence_id TEXT NOT NULL, state_id INTEGER NOT NULL, n_members INTEGER, PRIMARY KEY (sequence_id, state_id) ); """ ) rows, state_rows = [], [] for seq_id, s in seqs.items(): sizes = list(s["states"].values()) n_obs = s["n_obs"] cross = n_obs * n_obs - sum(x * x for x in sizes) # directed cross-state pairs dates = sorted(s["dates"]) rows.append(( seq_id, s["uniprot"].most_common(1)[0][0] if s["uniprot"] else None, s["length"], n_obs, len(sizes), len(s["entries"]), s["n_holo"], n_obs - s["n_holo"], cross, max(sizes), s["rep"][1], s["rep"][2], s["cath"].most_common(1)[0][0] if s["cath"] else None, ";".join(m for m, _ in s["methods"].most_common()), dates[0] if dates else None, dates[-1] if dates else None, (functions.get(seq_id) or [None])[0], ";".join(consensus(s["ecod"])) or None, ";".join(consensus(s["pfam"])) or None, )) for st, cnt in s["states"].items(): state_rows.append((seq_id, st, cnt)) dst.executemany(f"INSERT INTO sequence VALUES ({','.join('?' * 19)})", rows) dst.executemany("INSERT INTO state VALUES (?,?,?)", state_rows) dst.executescript( """ CREATE INDEX idx_seq_states ON sequence(n_states DESC, n_obs DESC); CREATE INDEX idx_seq_obs ON sequence(n_obs DESC); CREATE INDEX idx_seq_uniprot ON sequence(uniprot_id COLLATE NOCASE); CREATE INDEX idx_seq_cath ON sequence(cath_superfamily); CREATE INDEX idx_seq_ecod ON sequence(ecod_hset, n_obs DESC); """ ) fam, hom = ecod_names dst.executescript( """ DROP TABLE IF EXISTS ecod_h; CREATE TABLE ecod_h (h_id TEXT PRIMARY KEY, architecture TEXT, x_name TEXT, h_name TEXT); DROP TABLE IF EXISTS ecod_f; CREATE TABLE ecod_f (f_id TEXT PRIMARY KEY, t_name TEXT, f_name TEXT); DROP TABLE IF EXISTS ecod_group; CREATE TABLE ecod_group ( hset TEXT PRIMARY KEY, label TEXT, n_sequences INTEGER, n_obs INTEGER, n_multistate INTEGER, n_states INTEGER ); """ ) dst.executemany("INSERT INTO ecod_h VALUES (?,?,?,?)", [(k, *v) for k, v in hom.items()]) dst.executemany("INSERT INTO ecod_f VALUES (?,?,?)", [(k, *v) for k, v in fam.items()]) groups: dict[tuple, list] = {} for s in seqs.values(): hs = consensus(s["ecod"]) if hs: g = groups.setdefault(hs, [0, 0, 0, 0]) g[0] += 1 g[1] += s["n_obs"] g[2] += len(s["states"]) > 1 g[3] += len(s["states"]) dst.executemany( "INSERT INTO ecod_group VALUES (?,?,?,?,?,?)", [ (";".join(hs), " + ".join(hom.get(h, ("", "", ""))[2] or h for h in hs), *v) for hs, v in groups.items() ], ) dst.execute("CREATE INDEX idx_group_size ON ecod_group(n_sequences DESC)") dst.executescript( """ DROP TABLE IF EXISTS function_fts; CREATE VIRTUAL TABLE function_fts USING fts5( sequence_id UNINDEXED, uniprot_id, cath_superfamily, domains, text, tokenize = 'porter unicode61' ); """ ) fts_rows = [] for seq_id, s in seqs.items(): text = " ".join(functions.get(seq_id, []))[:FUNCTION_TEXT_MAX_CHARS * FUNCTION_TEXT_TOP_N] uni = " ".join(s["uniprot"].keys()) cath = " ".join(s["cath"].keys()).replace(";", " ") hs = consensus(s["ecod"]) dom = " ".join( [*(h.replace(".", "_") for h in hs), *(hom.get(h, ("", "", ""))[2] for h in hs), *consensus(s["pfam"])] ) fts_rows.append((seq_id, uni, cath, dom, text)) dst.executemany("INSERT INTO function_fts VALUES (?,?,?,?,?)", fts_rows) # ─────────────────────────────────────────────────────────────── edge passes def build_state_pairs(src: sqlite3.Connection, dst: sqlite3.Connection, seqs: dict, limit: int | None) -> dict: """Store state-pair similarities (top states only); return fidelity counts over all pairs.""" keep = { seq_id: {st for st, _ in s["states"].most_common(MAX_PAIR_STATES)} for seq_id, s in seqs.items() if len(s["states"]) > MAX_PAIR_STATES } fid_counts: dict = defaultdict(lambda: {"state_pairs": 0, "observation_pairs": 0}) dst.executescript( """ DROP TABLE IF EXISTS state_pair; CREATE TABLE state_pair ( sequence_id TEXT NOT NULL, state_a INTEGER NOT NULL, state_b INTEGER NOT NULL, similarity REAL, fidelity TEXT, n_pairs INTEGER, PRIMARY KEY (sequence_id, state_a, state_b) ); """ ) src_tbl = f"(SELECT * FROM edge LIMIT {int(limit)})" if limit else "edge" # directed edges -> keep one direction (A <= B) and count undirected pairs sql = f""" SELECT sequence_id, CAST(state_id_A AS INTEGER) a, CAST(state_id_B AS INTEGER) b, MIN(CAST(state_similarity AS REAL)), MIN(state_fidelity), COUNT(*) FROM {src_tbl} WHERE CAST(state_id_A AS INTEGER) <= CAST(state_id_B AS INTEGER) GROUP BY sequence_id, a, b """ cur = src.execute(sql) n = 0 while True: rows = cur.fetchmany(100000) if not rows: break out = [] for s, a, b, sim, fid, c in rows: # within-state rows appear in both directions with A==B -> halve the count c = c // 2 if a == b else c if a != b: fc = fid_counts[fid or "NA"] fc["state_pairs"] += 1 fc["observation_pairs"] += c allowed = keep.get(s) if allowed is None or (a in allowed and b in allowed): out.append((s, a, b, sim, fid, c)) dst.executemany("INSERT INTO state_pair VALUES (?,?,?,?,?,?)", out) n += len(rows) log(f" state pairs: {n:,}") return dict(fid_counts) def edge_histograms(src: sqlite3.Connection, limit: int | None) -> dict: src_tbl = f"(SELECT * FROM edge LIMIT {int(limit)})" if limit else "edge" sql = f""" SELECT state_id_A = state_id_B AS same_state, COALESCE(NULLIF(pair_fidelity, ''), 'NA') AS fid, -- pairs without shared residues store '' similarity and 0.0 RMSD/coverage: keep them -- out of every histogram (they are counted under fidelity 'NA') CASE WHEN pair_similarity = '' THEN NULL ELSE MIN(CAST(CAST(pair_similarity AS REAL) * 50 AS INTEGER), 49) END AS sim_bin, CASE WHEN pair_similarity = '' THEN NULL ELSE MIN(CAST(CAST(RMSD AS REAL) * 2 AS INTEGER), 40) END AS rmsd_bin, CASE WHEN pair_similarity = '' THEN NULL ELSE MIN(CAST(CAST(coverage AS REAL) * 20 AS INTEGER), 19) END AS cov_bin, CASE WHEN tm_aln = '' THEN NULL ELSE MIN(CAST(CAST(tm_aln AS REAL) * 50 AS INTEGER), 49) END AS tm_bin, COUNT(*) FROM {src_tbl} GROUP BY 1, 2, 3, 4, 5, 6 """ out = { "total": 0, "cross_state": 0, "pair_fidelity": Counter(), "pair_fidelity_cross": Counter(), "similarity": [0] * 50, "similarity_cross": [0] * 50, "tm_aln": [0] * 50, "rmsd": [0] * 41, "coverage": [0] * 20, } for same, fid, sim_b, rmsd_b, cov_b, tm_b, c in src.execute(sql): out["total"] += c out["pair_fidelity"][fid] += c if not same: out["cross_state"] += c out["pair_fidelity_cross"][fid] += c if sim_b is not None and sim_b >= 0: out["similarity"][sim_b] += c if not same: out["similarity_cross"][sim_b] += c if tm_b is not None and tm_b >= 0: out["tm_aln"][tm_b] += c if rmsd_b is not None and rmsd_b >= 0: out["rmsd"][rmsd_b] += c if cov_b is not None and cov_b >= 0: out["coverage"][cov_b] += c return out def order_stats(src: sqlite3.Connection, limit: int | None) -> dict: """Counts of the observed order/disorder labels over all directed transitions.""" src_tbl = f"(SELECT * FROM edge LIMIT {int(limit)})" if limit else "edge" rows = src.execute(f""" SELECT order_invalid_reason, order_evidence, ordering_with_ligand_change, state_id_A = state_id_B, COUNT(*) FROM {src_tbl} GROUP BY 1, 2, 3, 4 """).fetchall() reason, evidence, evidence_cross, ligand = Counter(), Counter(), Counter(), Counter() for r, ev, lig, same, c in rows: reason[r or "NA"] += c if r == "valid": evidence[ev or "NA"] += c if not same: evidence_cross[ev or "NA"] += c if ev in ("low", "medium", "high"): ligand[(ev, lig == "True")] += c order = ["high", "medium", "low", "not_applicable"] return { "validity": [{"label": k, "value": v} for k, v in reason.most_common()], "evidence": [{"label": k, "value": evidence.get(k, 0), "cross_state": evidence_cross.get(k, 0)} for k in order], "with_ligand_change": [ {"label": k, "with": ligand.get((k, True), 0), "without": ligand.get((k, False), 0)} for k in order[:3] ], } def select_inner_nodes(members: list, states: list, max_nodes: int, max_states: int) -> list: """Round-robin over the largest states. Mirrors catalog.inner_graph_nodes exactly.""" order = [st for st, _ in states[:max_states]] keep = set(order) pool = [m for m in members if m[0] in keep] if len(pool) <= max_nodes: return pool by_state: dict = {} for m in pool: by_state.setdefault(m[0], []).append(m) chosen, depth = [], 0 while len(chosen) < max_nodes: added = False for st in order: bucket = by_state.get(st, []) if depth < len(bucket) and len(chosen) < max_nodes: chosen.append(bucket[depth]) added = True if not added: break depth += 1 return chosen def build_inner_graphs(src: sqlite3.Connection, dst: sqlite3.Connection) -> None: """Precompute each sequence's inner-graph sample and its pairwise similarities. Serving these live needs one edge-index lookup per node, which is slow when the DB sits on a network mount; stored here it is a single row read. """ import zlib states: dict = defaultdict(list) for seq, st, n in dst.execute( "SELECT sequence_id, state_id, n_members FROM state ORDER BY sequence_id, n_members DESC, state_id" ): states[seq].append((st, n)) members: dict = defaultdict(list) for seq, st, pdb, chain in dst.execute( "SELECT sequence_id, state_id, pdb_id, auth_asym_id FROM member " "ORDER BY sequence_id, state_id, resolution IS NULL, resolution" ): members[seq].append((st, pdb.lower(), chain)) src.execute("CREATE TEMP TABLE sel (pdb TEXT, chain TEXT, seq TEXT, idx INTEGER, PRIMARY KEY (pdb, chain))") nodes_by_seq = {} rows = [] for seq, mem in members.items(): chosen = select_inner_nodes(mem, states[seq], INNER_MAX_NODES, INNER_MAX_STATES) nodes_by_seq[seq] = [[pdb, chain] for _, pdb, chain in chosen] rows.extend((pdb, chain, seq, i) for i, (_, pdb, chain) in enumerate(chosen)) src.executemany("INSERT OR IGNORE INTO sel VALUES (?,?,?,?)", rows) log(f" {len(rows):,} sampled observations over {len(nodes_by_seq):,} sequences") edges: dict = defaultdict(list) seen: set = set() n = 0 cur = src.execute(""" SELECT a.seq, a.idx, b.idx, e.pair_similarity, e.pair_fidelity FROM temp.sel a JOIN edge e ON e.pdb_id_A = a.pdb COLLATE NOCASE AND e.auth_asym_id_A = a.chain JOIN temp.sel b ON b.pdb = lower(e.pdb_id_B) AND b.chain = e.auth_asym_id_B AND b.seq = a.seq WHERE a.idx < b.idx """) while True: batch = cur.fetchmany(200000) if not batch: break for seq, i, j, sim, fid in batch: if (seq, i, j) in seen: # the edge table holds some exact duplicate rows continue seen.add((seq, i, j)) simf = to_float(sim) edges[seq].append([i, j, -1 if simf is None else round(simf * 1000), FIDELITY_CODE.get(fid, 4)]) n += len(batch) log(f" inner edges: {n:,}") dst.executescript( """ DROP TABLE IF EXISTS inner_graph; CREATE TABLE inner_graph (sequence_id TEXT PRIMARY KEY, max_nodes INTEGER, max_states INTEGER, nodes TEXT, edges BLOB); """ ) dst.executemany( "INSERT INTO inner_graph VALUES (?,?,?,?,?)", ( (seq, INNER_MAX_NODES, INNER_MAX_STATES, json.dumps(nodes), zlib.compress(json.dumps(edges.get(seq, []), separators=(",", ":")).encode(), 6)) for seq, nodes in nodes_by_seq.items() ), ) # ─────────────────────────────────────────────────────────────── overview def outer_graph_stats(dst: sqlite3.Connection, n_sequences: int) -> dict: rows = dst.execute( "SELECT hset, label, n_sequences, n_obs, n_multistate FROM ecod_group ORDER BY n_sequences DESC" ).fetchall() covered = sum(r[2] for r in rows) size_bins = Counter() for r in rows: n = r[2] size_bins["1" if n == 1 else "2–5" if n <= 5 else "6–20" if n <= 20 else "21–100" if n <= 100 else ">100"] += 1 return { "level": "ECOD homology (H)", "groups": len(rows), "covered_sequences": covered, "coverage": covered / n_sequences if n_sequences else 0, "edges": sum(r[2] * (r[2] - 1) // 2 for r in rows), "group_sizes": [{"label": k, "value": size_bins[k]} for k in ["1", "2–5", "6–20", "21–100", ">100"]], "largest": [ {"hset": r[0], "label": r[1], "n_sequences": r[2], "n_obs": r[3], "n_multistate": r[4]} for r in rows[:12] ], } def build_overview(seqs, ov, n_nodes, edges, state_fid) -> dict: states_per_seq = Counter(len(s["states"]) for s in seqs.values()) obs_per_seq = [s["n_obs"] for s in seqs.values()] n_states_total = sum(len(s["states"]) for s in seqs.values()) multi_state = sum(1 for s in seqs.values() if len(s["states"]) > 1) lengths = [v for v in ov["length"].values() if v] def capped(counter, cap): out = Counter() for k, v in counter.items(): out[k if k < cap else cap] += v return [{"label": (f"{k}+" if k == cap else str(k)), "value": out[k]} for k in sorted(out)] obs_edges = [2, 3, 4, 5, 6, 11, 21, 51, 101, 1_000_000] obs_labels = ["2", "3", "4", "5", "6–10", "11–20", "21–50", "51–100", ">100"] len_edges = [0, 100, 200, 300, 400, 500, 750, 1000, 1_000_000] len_labels = ["<100", "100–199", "200–299", "300–399", "400–499", "500–749", "750–999", "≥1000"] res_edges = [0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 5.0, 1000] res_labels = ["<1.5", "1.5–2.0", "2.0–2.5", "2.5–3.0", "3.0–3.5", "3.5–4.0", "4.0–5.0", "≥5.0"] fid_order = ["identical", "low", "medium", "high", "no_shared_residues", "NA"] return { "generated": time.strftime("%Y-%m-%d"), "counts": { "observations": n_nodes, "sequences": len(seqs), "multi_state_sequences": multi_state, "state_clusters": n_states_total, "pdb_entries": len(ov["entries"]), "uniprot_accessions": len(ov["uniprot"]), "transitions": edges["total"], "cross_state_transitions": edges["cross_state"], }, "states_per_sequence": capped(states_per_seq, 10), "observations_per_sequence": [ {"label": l, "value": v} for l, v in zip(obs_labels, hist(obs_per_seq, obs_edges)) ], "sequence_length": [ {"label": l, "value": v} for l, v in zip(len_labels, hist(lengths, len_edges)) ], "resolution": [ {"label": l, "value": v} for l, v in zip(res_labels, hist(ov["resolution"], res_edges)) ], "experimental_method": [{"label": k, "value": v} for k, v in ov["method"].most_common()], "binding_status": [{"label": k, "value": v} for k, v in ov["binding"].most_common()], "chain_composition": [{"label": k, "value": v} for k, v in ov["composition"].most_common()], "nucleic_acid_binding": [{"label": k, "value": v} for k, v in ov["npp"].most_common()], "release_year": [{"label": k, "value": ov["year"][k]} for k in sorted(ov["year"])], "partners": ov.get("partners"), "pair_fidelity": [ {"label": k, "value": edges["pair_fidelity"].get(k, 0), "cross_state": edges["pair_fidelity_cross"].get(k, 0)} for k in fid_order if edges["pair_fidelity"].get(k) ], "state_fidelity": [ {"label": k, **state_fid[k]} for k in fid_order if k in state_fid ], "pair_similarity": { "bin_width": 0.02, "all": edges["similarity"], "cross_state": edges["similarity_cross"], }, "tm_aln": {"bin_width": 0.02, "all": edges["tm_aln"]}, "rmsd": {"bin_width": 0.5, "all": edges["rmsd"], "last_bin_open": True}, "coverage": {"bin_width": 0.05, "all": edges["coverage"]}, } def main(): ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("db", type=Path) ap.add_argument("--out-dir", type=Path, required=True) ap.add_argument("--limit", type=int, default=None, help="debug: only read the first N rows") ap.add_argument("--ecod", type=Path, default=None, help="ecod.latest.domains.txt (for ECOD names)") ap.add_argument("--inner-only", action="store_true", help="only (re)build the inner_graph table in an existing MuSProt-index.db") args = ap.parse_args() args.out_dir.mkdir(parents=True, exist_ok=True) index_path = args.out_dir / "MuSProt-index.db" tmp_path = index_path.with_suffix(".db.tmp") tmp_path.unlink(missing_ok=True) src = sqlite3.connect(f"file:{args.db}?mode=ro", uri=True) src.execute("PRAGMA temp_store = FILE") src.execute("PRAGMA cache_size = -2000000") if args.inner_only: dst = sqlite3.connect(index_path) log("inner graphs") build_inner_graphs(src, dst) dst.commit() dst.execute("VACUUM") dst.close() log(f"done → {index_path} ({index_path.stat().st_size / 1e6:.1f} MB)") return dst = sqlite3.connect(tmp_path) dst.execute("PRAGMA journal_mode = OFF") dst.execute("PRAGMA synchronous = OFF") ecod_names = load_ecod_names(args.ecod) log(f"ECOD names: {len(ecod_names[0]):,} families, {len(ecod_names[1]):,} homology groups") log("pass 1/4: node table") seqs, functions, ov, n_nodes = build_members(src, dst, args.limit) write_sequences(dst, seqs, functions, ecod_names) dst.commit() log("pass 2/4: state pairs (edge GROUP BY)") state_fid = build_state_pairs(src, dst, seqs, args.limit) dst.commit() log("pass 3/4: edge histograms") edges = edge_histograms(src, args.limit) log("pass 4/4: order / disorder labels") order = order_stats(src, args.limit) log("inner graphs") build_inner_graphs(src, dst) overview = build_overview(seqs, ov, n_nodes, edges, state_fid) overview["outer_graph"] = outer_graph_stats(dst, len(seqs)) overview["order"] = order dst.execute("CREATE TABLE meta (key TEXT PRIMARY KEY, value TEXT)") dst.execute("INSERT INTO meta VALUES ('overview', ?)", (json.dumps(overview),)) dst.commit() dst.execute("VACUUM") dst.close() tmp_path.replace(index_path) (args.out_dir / "overview.json").write_text(json.dumps(overview, indent=1)) log(f"done → {index_path} ({index_path.stat().st_size / 1e6:.1f} MB)") log(json.dumps(overview["counts"], indent=1)) if __name__ == "__main__": main()