"""Lightweight PDB/mmCIF readers for Evo-IF inference. Only the information consumed by ``PDBDataset.load_chains`` and ``PDBDataset.load_assembly`` is retained: polymer type, atom coordinates and occupancies, and biological-assembly transforms. """ from __future__ import annotations import gzip import itertools import re from collections import OrderedDict, namedtuple from contextlib import contextmanager from pathlib import Path from typing import Iterator, Mapping, Sequence, TextIO import numpy as np from Bio.PDB import PDBParser as BioPDBParser from Bio.PDB.MMCIF2Dict import MMCIF2Dict Atom = namedtuple("Atom", ["name", "xyz", "occ", "bfac"]) Chain = namedtuple("Chain", ["id", "type", "sequence", "atoms"]) PROTEIN_RESTYPES = { "ALA", "ARG", "ASN", "ASP", "CYS", "GLN", "GLU", "GLY", "HIS", "ILE", "LEU", "LYS", "MET", "PHE", "PRO", "SER", "THR", "TRP", "TYR", "VAL", "UNK", } DNA_RESTYPES = {"DA", "DC", "DG", "DT", "DX"} RNA_RESTYPES = {"A", "C", "G", "U", "RX"} POLYMER_TYPES = { "polypeptide(L)", "polydeoxyribonucleotide", "polyribonucleotide", "polydeoxyribonucleotide/polyribonucleotide hybrid", } @contextmanager def _open_text(path: str | Path) -> Iterator[TextIO]: value = str(Path(path).expanduser()) if value.lower().endswith(".gz"): with gzip.open(value, "rt") as handle: yield handle else: with open(value, "rt") as handle: yield handle def _as_list(value) -> list[str]: if value is None: return [] if isinstance(value, list): return value if isinstance(value, tuple): return list(value) return [value] def _column( data: Mapping[str, object], name: str, *, length: int | None = None, default: str | None = None, ) -> list[str]: values = _as_list(data.get(name)) if values: return values if length is not None and default is not None: return [default] * length raise ValueError(f"mmCIF is missing required column {name}") def _safe_float(value: str, default: float = 0.0) -> float: if value in {"", ".", "?"}: return default return float(value) def _parse_operation_group(expression: str) -> list[str]: operations: list[str] = [] for item in expression.strip("() ").split(","): item = item.strip() if not item: continue match = re.fullmatch(r"(-?\d+)-(-?\d+)", item) if match is None: operations.append(item) continue start, stop = (int(value) for value in match.groups()) step = 1 if stop >= start else -1 operations.extend(str(value) for value in range(start, stop + step, step)) return operations def _operation_groups(expression: str) -> list[list[str]]: parenthesized = re.findall(r"\(([^()]*)\)", expression) raw_groups = parenthesized or [expression] groups = [_parse_operation_group(group) for group in raw_groups] if not groups or any(not group for group in groups): raise ValueError(f"Invalid mmCIF operation expression: {expression!r}") return groups def _read_operation_matrices(data: Mapping[str, object]) -> dict[str, np.ndarray]: operation_ids = _as_list(data.get("_pdbx_struct_oper_list.id")) if not operation_ids: return {} columns: dict[str, list[str]] = {} for row in range(3): columns[f"v{row}"] = _column( data, f"_pdbx_struct_oper_list.vector[{row + 1}]", ) for col in range(3): columns[f"m{row}{col}"] = _column( data, f"_pdbx_struct_oper_list.matrix[{row + 1}][{col + 1}]", ) matrices: dict[str, np.ndarray] = {} for index, operation_id in enumerate(operation_ids): matrix = np.eye(4, dtype=np.float64) for row in range(3): matrix[row, 3] = float(columns[f"v{row}"][index]) for col in range(3): matrix[row, col] = float(columns[f"m{row}{col}"][index]) matrices[operation_id] = matrix return matrices def _compose_operations( expression: str, matrices: Mapping[str, np.ndarray], ) -> list[np.ndarray]: groups = _operation_groups(expression) transforms: list[np.ndarray] = [] for operation_tuple in itertools.product(*groups): transform = np.eye(4, dtype=np.float64) for operation_id in operation_tuple: if operation_id not in matrices: raise ValueError( f"mmCIF assembly references missing operation {operation_id!r}" ) transform = transform @ matrices[operation_id] transforms.append(transform) return transforms def _identity_assembly(chains: Mapping[str, Chain]) -> dict[str, list[tuple[str, np.ndarray]]]: return {"1": [(chain_id, np.eye(4, dtype=np.float64)) for chain_id in chains]} def _read_mmcif_assemblies( data: Mapping[str, object], chains: Mapping[str, Chain], ) -> dict[str, list[tuple[str, np.ndarray]]]: assembly_ids = _as_list(data.get("_pdbx_struct_assembly_gen.assembly_id")) expressions = _as_list(data.get("_pdbx_struct_assembly_gen.oper_expression")) chain_lists = _as_list(data.get("_pdbx_struct_assembly_gen.asym_id_list")) if not assembly_ids or not expressions or not chain_lists: return _identity_assembly(chains) if not (len(assembly_ids) == len(expressions) == len(chain_lists)): raise ValueError("Inconsistent mmCIF biological-assembly columns") matrices = _read_operation_matrices(data) assemblies: dict[str, list[tuple[str, np.ndarray]]] = {} for assembly_id, expression, chain_list in zip( assembly_ids, expressions, chain_lists, ): transforms = _compose_operations(expression, matrices) selected_chains = [ chain_id.strip() for chain_id in chain_list.split(",") if chain_id.strip() in chains ] entries = assemblies.setdefault(assembly_id, []) entries.extend( (chain_id, transform) for chain_id in selected_chains for transform in transforms ) nonempty = {key: value for key, value in assemblies.items() if value} if not nonempty: raise ValueError( "mmCIF declares biological assemblies, but none reference a " "supported parsed polymer chain" ) return nonempty def _mmcif_metadata(data: Mapping[str, object], model_id: str) -> dict[str, object]: def first(name: str, default=None): values = _as_list(data.get(name)) return values[0] if values else default resolution = first("_refine.ls_d_res_high") if resolution in {None, ".", "?"}: resolution = first("_em_3d_reconstruction.resolution") try: resolution = float(resolution) if resolution not in {None, ".", "?"} else None except (TypeError, ValueError): resolution = None return { "method": first("_exptl.method"), "date": first("_pdbx_database_status.recvd_initial_deposition_date"), "resolution": resolution, "model_id": model_id, "assembly_source": "mmcif", } class CIFParser: """Read polymer coordinates and declared biological assemblies from mmCIF.""" def __init__( self, skip_res: Sequence[str] | None = None, randomize_nmr_model: int = 0, ) -> None: self.skip_res = set(skip_res or ()) self.randomize_nmr_model = int(randomize_nmr_model) def parse(self, filename: str | Path): with _open_text(filename) as handle: data = MMCIF2Dict(handle) entity_ids = _column(data, "_entity_poly.entity_id") entity_types = _column(data, "_entity_poly.type") entity_type = { entity_id: polymer_type for entity_id, polymer_type in zip(entity_ids, entity_types) if polymer_type in POLYMER_TYPES } chain_ids = _column(data, "_atom_site.label_asym_id") row_count = len(chain_ids) columns = { "chain": chain_ids, "entity": _column(data, "_atom_site.label_entity_id"), "residue_id": _column(data, "_atom_site.label_seq_id"), "residue_name": _column(data, "_atom_site.label_comp_id"), "atom_name": _column(data, "_atom_site.label_atom_id"), "x": _column(data, "_atom_site.Cartn_x"), "y": _column(data, "_atom_site.Cartn_y"), "z": _column(data, "_atom_site.Cartn_z"), "occupancy": _column( data, "_atom_site.occupancy", length=row_count, default="1.0", ), "bfactor": _column( data, "_atom_site.B_iso_or_equiv", length=row_count, default="0.0", ), "model": _column( data, "_atom_site.pdbx_PDB_model_num", length=row_count, default="1", ), } if any(len(values) != row_count for values in columns.values()): raise ValueError("Inconsistent mmCIF atom_site column lengths") model_ids = list(dict.fromkeys(columns["model"])) selected_model = model_ids[0] if self.randomize_nmr_model and len(model_ids) > 1: selected_model = str(np.random.choice(model_ids)) chain_data: OrderedDict[str, dict[str, object]] = OrderedDict() for row in zip(*(columns[name] for name in columns)): ( chain_id, entity_id, residue_id, residue_name, atom_name, x, y, z, occupancy, bfactor, model_id, ) = row polymer_type = entity_type.get(entity_id) if polymer_type is None or model_id != selected_model: continue if residue_id in {"", ".", "?"} or residue_name in self.skip_res: continue item = chain_data.setdefault( chain_id, {"type": polymer_type, "atoms": OrderedDict()}, ) atom_key = (chain_id, residue_id, residue_name, atom_name) atom = Atom( name=atom_key, xyz=[float(x), float(y), float(z)], occ=_safe_float(occupancy, 1.0), bfac=_safe_float(bfactor, 0.0), ) atoms = item["atoms"] previous = atoms.get(atom_key) if previous is None or atom.occ > previous.occ: atoms[atom_key] = atom chains = { chain_id: Chain( id=chain_id, type=item["type"], sequence=None, atoms=item["atoms"], ) for chain_id, item in chain_data.items() if item["atoms"] } if not chains: raise ValueError(f"No supported polymer chains found in {filename}") assemblies = _read_mmcif_assemblies(data, chains) metadata = _mmcif_metadata(data, selected_model) return chains, assemblies, [], metadata def _classify_pdb_chain(residue_names: set[str]) -> str | None: has_protein = bool(residue_names & PROTEIN_RESTYPES) has_dna = bool(residue_names & DNA_RESTYPES) has_rna = bool(residue_names & RNA_RESTYPES) if has_protein and not has_dna and not has_rna: return "polypeptide(L)" if has_dna and not has_protein and not has_rna: return "polydeoxyribonucleotide" if has_rna and not has_protein and not has_dna: return "polyribonucleotide" if has_dna and has_rna and not has_protein: return "polydeoxyribonucleotide/polyribonucleotide hybrid" if has_protein or has_dna or has_rna: raise ValueError("A PDB chain mixes protein and nucleic-acid residues") return None class PDBParser: """Read the coordinates present in a PDB file as one identity assembly.""" def __init__(self) -> None: self._parser = BioPDBParser(QUIET=True) def parse(self, filename: str | Path): with _open_text(filename) as handle: structure = self._parser.get_structure(Path(filename).stem, handle) try: model = next(structure.get_models()) except StopIteration as exc: raise ValueError(f"No coordinate model found in {filename}") from exc chains: dict[str, Chain] = {} for bio_chain in model: residues = list(bio_chain.get_residues()) residue_names = {residue.get_resname().strip() for residue in residues} polymer_type = _classify_pdb_chain(residue_names) if polymer_type is None: continue allowed = PROTEIN_RESTYPES | DNA_RESTYPES | RNA_RESTYPES atoms: OrderedDict[tuple[str, str, str, str], Atom] = OrderedDict() for residue in residues: residue_name = residue.get_resname().strip() if residue_name not in allowed: continue residue_id = str(int(residue.id[1])) for bio_atom in residue.get_atoms(): atom_name = bio_atom.get_name().strip() atom_key = (bio_chain.id, residue_id, residue_name, atom_name) occupancy = bio_atom.get_occupancy() bfactor = bio_atom.get_bfactor() atom = Atom( name=atom_key, xyz=bio_atom.get_coord().astype(float).tolist(), occ=float(occupancy) if occupancy is not None else 0.0, bfac=float(bfactor) if bfactor is not None else 0.0, ) previous = atoms.get(atom_key) if previous is None or atom.occ > previous.occ: atoms[atom_key] = atom if atoms: chains[bio_chain.id] = Chain( id=bio_chain.id, type=polymer_type, sequence=None, atoms=atoms, ) if not chains: raise ValueError(f"No supported polymer chains found in {filename}") metadata = { "model_id": str(model.id), "assembly_source": "coordinates_as_provided", } return chains, _identity_assembly(chains), [], metadata def parse_structure( filename: str | Path, *, skip_res: Sequence[str] | None = None, randomize_nmr_model: int = 0, ): """Dispatch to the lightweight PDB or mmCIF reader by filename suffix.""" value = str(filename).lower() if value.endswith((".pdb", ".pdb.gz")): return PDBParser().parse(filename) if value.endswith((".cif", ".cif.gz", ".mmcif", ".mmcif.gz")): return CIFParser( skip_res=skip_res, randomize_nmr_model=randomize_nmr_model, ).parse(filename) raise ValueError(f"Unsupported structure format: {filename}") __all__ = ["Atom", "Chain", "CIFParser", "PDBParser", "parse_structure"]