"""Reproduce ProteinIQ's two-entry coordinate-coverage audit (Biopython 1.85).

Usage: python analyze.py --output-dir /path/to/results
Reads the adjacent gzip snapshots; never fetches changing live records.
"""
import argparse
import csv
import gzip
import hashlib
import io
import json
from pathlib import Path

import Bio
from Bio.PDB.MMCIF2Dict import MMCIF2Dict


def rows(data, category, fields):
    columns = [data.get(f"{category}.{field}", []) for field in fields]
    if len({len(column) for column in columns}) > 1:
        raise ValueError(f"Mismatched columns in {category}")
    return [dict(zip(fields, values)) for values in zip(*columns)]


def analyze(raw):
    data = MMCIF2Dict(io.StringIO(raw.decode("utf-8")))
    entities = rows(data, "_entity", ["id", "pdbx_description"])
    matches = [r for r in entities if "semaglutide" in r["pdbx_description"].lower()]
    if len(matches) != 1:
        raise ValueError("Expected exactly one semaglutide peptide entity")
    entity = matches[0]
    chains = [r["id"] for r in rows(data, "_struct_asym", ["id", "entity_id"])
              if r["entity_id"] == entity["id"]]
    if len(chains) != 1:
        raise ValueError("This audit requires a single copy of the peptide")
    chain = chains[0]
    sequence = {int(r["num"]): r["mon_id"] for r in rows(
        data, "_entity_poly_seq", ["entity_id", "num", "mon_id"])
        if r["entity_id"] == entity["id"]}
    atoms = [r for r in rows(data, "_atom_site", [
        "label_asym_id", "auth_asym_id", "label_seq_id", "auth_seq_id",
        "label_comp_id", "label_atom_id", "type_symbol", "occupancy",
        "pdbx_PDB_model_num"]) if r["pdbx_PDB_model_num"] == "1"
        and float(r["occupancy"]) > 0 and r["type_symbol"] not in {"H", "D"}]
    peptide = [r for r in atoms if r["label_asym_id"] == chain]
    if not peptide:
        raise ValueError("Peptide has no positive-occupancy heavy atoms")
    author_chain = peptide[0]["auth_asym_id"]
    modeled = {int(r["label_seq_id"]) for r in peptide}
    if not modeled <= sequence.keys():
        raise ValueError("Coordinate residues are outside the deposited sequence")
    numbering = {int(r["seq_id"]): r["pdb_seq_num"] for r in rows(
        data, "_pdbx_poly_seq_scheme", ["asym_id", "seq_id", "pdb_seq_num"])
        if r["asym_id"] == chain}
    missing = sorted(sequence.keys() - modeled)
    missing_author = [numbering[i] for i in missing]
    annotated_missing = {r["auth_seq_id"] for r in rows(
        data, "_pdbx_unobs_or_zero_occ_residues",
        ["auth_asym_id", "auth_seq_id", "PDB_model_num"])
        if r["auth_asym_id"] == author_chain and r["PDB_model_num"] == "1"}
    if set(missing_author) != annotated_missing:
        raise ValueError("Coordinate coverage disagrees with missing-residue annotations")
    components = {r["id"]: {"name": r["name"], "formula": r["formula"]}
                  for r in rows(data, "_chem_comp", ["id", "name", "formula"])}
    fields = ["conn_type_id"] + [f"ptnr{n}_{f}" for n in (1, 2) for f in
              ("label_asym_id", "label_comp_id", "label_seq_id", "label_atom_id")]
    linked = []
    for r in rows(data, "_struct_conn", fields):
        if r["conn_type_id"] != "covale":
            continue
        for n, other in [(1, 2), (2, 1)]:
            if r[f"ptnr{n}_label_asym_id"] != chain or r[f"ptnr{other}_label_asym_id"] == chain:
                continue
            other_chain = r[f"ptnr{other}_label_asym_id"]
            comp = r[f"ptnr{other}_label_comp_id"]
            linked_atoms = {a["label_atom_id"] for a in atoms
                            if a["label_asym_id"] == other_chain and a["label_comp_id"] == comp}
            linked.append({"component_id": comp, "label_chain": other_chain,
                           "peptide_label_position": int(r[f"ptnr{n}_label_seq_id"]),
                           "peptide_atom": r[f"ptnr{n}_label_atom_id"],
                           "component_atom": r[f"ptnr{other}_label_atom_id"],
                           "modeled_heavy_atoms": len(linked_atoms), **components[comp]})
    coverage = [{"label_position": i, "author_position": numbering[i],
                 "component_id": sequence[i], "has_coordinates": i in modeled}
                for i in sorted(sequence)]
    return {
        "pdb_id": data["_entry.id"][0], "title": data["_struct.title"][0],
        "sha256_uncompressed": hashlib.sha256(raw).hexdigest(),
        "latest_revision_date": max(data["_pdbx_audit_revision_history.revision_date"]),
        "method": data["_exptl.method"][0], "entity_id": entity["id"],
        "label_chain": chain, "author_chain": author_chain,
        "declared_residues": len(sequence), "modeled_residues": len(modeled),
        "missing_author_positions": missing_author,
        "aib_declared": "AIB" in sequence.values(),
        "aib_has_coordinates": any(a["label_comp_id"] == "AIB" for a in peptide),
        "external_covalent_components": linked, "residues": coverage,
    }


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output-dir", type=Path, required=True)
    args = parser.parse_args()
    root = Path(__file__).resolve().parent
    sources = json.loads((root / "sources.json").read_text(encoding="utf-8"))
    results = []
    for source in sources["records"]:
        raw = gzip.decompress((root / source["snapshot"]).read_bytes())
        if hashlib.sha256(raw).hexdigest() != source["sha256_uncompressed"]:
            raise ValueError(f"Snapshot checksum mismatch: {source['pdb_id']}")
        results.append(analyze(raw))
    args.output_dir.mkdir(parents=True, exist_ok=True)
    output = {"retrieved_on": sources["retrieved_on"], "biopython_version": Bio.__version__,
              "model": 1, "rule": "At least one positive-occupancy non-hydrogen atom per peptide residue",
              "records": results}
    (args.output_dir / "results.json").write_text(json.dumps(output, indent=2) + "\n", encoding="utf-8")
    with (args.output_dir / "coverage.csv").open("w", newline="", encoding="utf-8") as handle:
        writer = csv.writer(handle)
        writer.writerow(["pdb_id", "label_chain", "author_chain", "label_position", "author_position",
                         "component_id", "has_coordinates"])
        for result in results:
            for residue in result["residues"]:
                writer.writerow([result["pdb_id"], result["label_chain"], result["author_chain"],
                                 residue["label_position"], residue["author_position"],
                                 residue["component_id"], int(residue["has_coordinates"])])
    for result in results:
        print(f"{result['pdb_id']}: {result['modeled_residues']}/{result['declared_residues']} residues; "
              f"missing native positions {', '.join(result['missing_author_positions'])}")


if __name__ == "__main__":
    main()
