"""Reproduce ProteinIQ's GENCODE v45-v50 reference-chromosome comparison.

Python 3.11+; standard library only. Download the six GTF files listed in
sources.json, then run: python analyze.py --input-dir PATH --output-dir RESULTS
Run beside the supplied sources.json. Outputs are regenerated, not hand edited.
Only gene feature rows are counted. Stable IDs lose their numeric version but
retain any _PAR_Y suffix. Coding = protein_coding without readthrough_gene.
An absent identifier is NOT evidence of a new or lost biological gene. Span
overlaps are ambiguity flags, not inferred equivalences or confirmed mergers.
"""

import argparse
import csv
import gzip
import hashlib
import json
import platform
import re
from collections import Counter, defaultdict
from pathlib import Path


RELEASES = range(45, 51)
CHROMOSOMES = {f"chr{i}" for i in range(1, 23)} | {"chrX", "chrY", "chrM", "chrMT"}
DATES = {45: "2024-01", 46: "2024-05", 47: "2024-10", 48: "2025-05", 49: "2025-09", 50: "2026-06"}
ATTRIBUTES = re.compile(r'(\w+) "([^"]*)"')


def stable_id(identifier):
    return re.sub(r"\.\d+(?=_PAR_Y$|$)", "", identifier)


def read_genes(filename):
    genes = {}
    with gzip.open(filename, "rt", encoding="utf-8") as stream:
        for line in stream:
            if line.startswith("#"):
                continue
            fields = line.rstrip("\n").split("\t")
            if len(fields) != 9:
                raise ValueError("Malformed GTF row")
            if fields[2] != "gene":
                continue
            if fields[0] not in CHROMOSOMES:
                raise ValueError(f"Unexpected sequence in CHR file: {fields[0]}")
            pairs = ATTRIBUTES.findall(fields[8])
            attrs = dict(pairs)
            key = stable_id(attrs["gene_id"])
            if key in genes:
                raise ValueError(f"Duplicate stable gene identifier: {key}")
            genes[key] = {
                "id": attrs["gene_id"], "name": attrs["gene_name"],
                "biotype": attrs["gene_type"], "chr": fields[0],
                "start": int(fields[3]), "end": int(fields[4]), "strand": fields[6],
                "readthrough": ("tag", "readthrough_gene") in pairs,
            }
    if not genes:
        raise ValueError("No gene rows found")
    return genes


def coding(gene):
    return bool(gene and gene["biotype"] == "protein_coding" and not gene["readthrough"])


def state(gene):
    if gene is None:
        return "absent"
    if gene["biotype"] == "protein_coding" and gene["readthrough"]:
        return "protein_coding_readthrough"
    return gene["biotype"]


def compare(before, after, old_release, new_release):
    old_coding = {key for key, gene in before.items() if coding(gene)}
    new_coding = {key for key, gene in after.items() if coding(gene)}
    indexes = []
    for genes in (before, after):
        index = defaultdict(list)
        for key, gene in genes.items():
            index[(gene["chr"], gene["strand"])].append((key, gene))
        indexes.append(index)
    rows = []
    for key in sorted(old_coding ^ new_coding):
        old, new = before.get(key), after.get(key)
        entering = key in new_coding
        counterpart = old if entering else new
        if counterpart is None:
            category = "identifier_absent_in_other_release"
        elif counterpart["biotype"] == "protein_coding":
            category = "readthrough_status_change"
        else:
            category = "biotype_change"
        gene = new if entering else old
        # Review missing IDs against all same-strand gene spans, including introns.
        overlaps = []
        if counterpart is None:
            index = indexes[0 if entering else 1]
            overlaps = sorted(other for other, candidate in index[(gene["chr"], gene["strand"])]
                              if candidate["start"] <= gene["end"] and candidate["end"] >= gene["start"])
        row = {"from_release": old_release, "to_release": new_release,
               "stable_gene_id": key, "direction": "entering" if entering else "leaving",
               "category": category, "before_state": state(old), "after_state": state(new),
               "overlapping_ids_in_other_release": ";".join(overlaps)}
        for prefix, record in (("before", old), ("after", new)):
            for field in ("id", "name", "chr", "start", "end", "strand"):
                row[f"{prefix}_{field}"] = record[field] if record else ""
        rows.append(row)
    added, removed = len(new_coding - old_coding), len(old_coding - new_coding)
    retained = len(old_coding & new_coding)
    assert retained + added == len(new_coding)
    assert retained + removed == len(old_coding)
    assert added - removed == len(new_coding) - len(old_coding)
    summary = {"from_release": old_release, "to_release": new_release,
               "retained_coding_ids": retained, "entering_coding_ids": added,
               "leaving_coding_ids": removed, "net_change": added - removed}
    for direction in ("entering", "leaving"):
        counts = Counter(row["category"] for row in rows if row["direction"] == direction)
        for category in ("biotype_change", "readthrough_status_change", "identifier_absent_in_other_release"):
            summary[f"{direction}_{category}"] = counts[category]
    return summary, rows


def write_csv(filename, rows):
    if not rows:
        raise ValueError(f"No output rows for {filename}")
    with filename.open("w", newline="", encoding="utf-8") as stream:
        writer = csv.DictWriter(stream, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--input-dir", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    args = parser.parse_args()
    sources = json.loads(Path(__file__).with_name("sources.json").read_text(encoding="utf-8"))
    args.output_dir.mkdir(parents=True, exist_ok=True)
    releases, summaries = {}, []
    for source in sources:
        version = source["release"]
        filename = args.input_dir / f"gencode.v{version}.annotation.gtf.gz"
        with filename.open("rb") as stream:
            checksum = hashlib.file_digest(stream, "sha256").hexdigest()
        if checksum != source["sha256"]:
            raise ValueError(f"Source checksum mismatch: {filename}")
        genes = read_genes(filename)
        releases[version] = genes
        counts = Counter(gene["biotype"] for gene in genes.values())
        headline = sum(coding(gene) for gene in genes.values())
        readthrough = sum(gene["biotype"] == "protein_coding" and gene["readthrough"] for gene in genes.values())
        assert len(genes) == source["published_total_genes"]
        assert headline == source["published_coding_genes"]
        assert readthrough == source["published_readthrough_genes"]
        summaries.append({"release": version, "release_date": DATES[version], "assembly": "GRCh38.p14",
                          "all_gene_entries": len(genes), "headline_coding_genes": headline,
                          "protein_coding_biotype": counts["protein_coding"], "readthrough_genes": readthrough,
                          "lncRNA_biotype": counts["lncRNA"], "TEC_biotype": counts["TEC"]})
        print(f"v{version}: {headline} coding; {len(genes)} total; matches published statistics", flush=True)
    transitions, events = [], []
    for old, new in [(v, v + 1) for v in range(45, 50)] + [(45, 50)]:
        summary, rows = compare(releases[old], releases[new], old, new)
        transitions.append(summary)
        events.extend(rows)
    write_csv(args.output_dir / "release-counts.csv", summaries)
    write_csv(args.output_dir / "coding-transitions.csv", transitions)
    write_csv(args.output_dir / "changed-coding-identifiers.csv", events)
    provenance = {"python": platform.python_version(), "analysis_script_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
                  "sources": sources, "scope": "Gene rows on reference chromosomes; protein_coding excluding readthrough_gene",
                  "matching": "Remove numeric ID version only; preserve _PAR_Y; do not match by gene symbol",
                  "overlap": "Same chromosome and strand, inclusive gene-span overlap; ambiguity flag only"}
    (args.output_dir / "provenance.json").write_text(json.dumps(provenance, indent=2) + "\n", encoding="utf-8")
    print(json.dumps(transitions[-1], indent=2))


if __name__ == "__main__":
    main()
