"""Count direct offspring, unique descendants, and sire links in a public pedigree.

Input: new_pedigree_dogs.fam from dog_genotype_data.tar.gz at
https://github.com/cflerin/dog_recombination

Only the six PLINK pedigree columns are read. Counts describe sampled records,
not all dogs born in the colony or genetic contribution to descendants.
"""

from __future__ import annotations

import argparse
import csv
from collections import Counter
from functools import lru_cache
from pathlib import Path


def read_pedigree(path: Path):
    pedigree = {}
    sexes = {}
    for line_number, line in enumerate(path.read_text().splitlines(), start=1):
        fields = line.split()
        if len(fields) != 6:
            raise ValueError(f"{path}:{line_number}: expected six PLINK .fam fields")
        _, animal_id, sire_id, dam_id, sex, _ = fields
        if animal_id in pedigree:
            raise ValueError(f"duplicate sample ID: {animal_id}")
        if sex not in {"1", "2"}:
            raise ValueError(f"unknown sex code for {animal_id}: {sex}")
        pedigree[animal_id] = (None if sire_id == "0" else sire_id,
                               None if dam_id == "0" else dam_id)
        sexes[animal_id] = sex
    for child_id, (sire_id, dam_id) in pedigree.items():
        for parent_id, expected_sex in ((sire_id, "1"), (dam_id, "2")):
            if parent_id is not None and (parent_id not in pedigree or sexes[parent_id] != expected_sex):
                raise ValueError(f"missing or sex-inconsistent parent {parent_id} of {child_id}")
    return pedigree, sexes


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("fam", type=Path, help="Path to new_pedigree_dogs.fam")
    parser.add_argument("--output", type=Path, required=True, help="Per-animal CSV path")
    args = parser.parse_args()
    pedigree, sexes = read_pedigree(args.fam)

    sire_links = Counter(sire_id for sire_id, _ in pedigree.values() if sire_id)
    dam_links = Counter(dam_id for _, dam_id in pedigree.values() if dam_id)
    direct_offspring = sire_links + dam_links

    visiting = set()

    @lru_cache(None)
    def ancestors(animal_id):
        if animal_id in visiting:
            raise ValueError(f"pedigree cycle at {animal_id}")
        visiting.add(animal_id)
        found = set()
        for parent_id in pedigree[animal_id]:
            if parent_id:
                found.add(parent_id)
                found.update(ancestors(parent_id))
        visiting.remove(animal_id)
        return frozenset(found)

    unique_descendants = Counter()
    for animal_id in pedigree:
        for ancestor_id in ancestors(animal_id):
            unique_descendants[ancestor_id] += 1

    if (len(pedigree), sum(sire_links.values()), sum(dam_links.values()),
            len(sire_links), len(dam_links)) != (237, 212, 207, 18, 22):
        raise ValueError("source cohort differs from the published analysis")
    assert sum(direct_offspring.values()) == 419
    assert all(unique_descendants[animal_id] >= direct_offspring[animal_id]
               for animal_id in pedigree)

    rows = [{
        "sample_id": animal_id,
        "sex": "male" if sexes[animal_id] == "1" else "female",
        "direct_offspring": direct_offspring[animal_id],
        "unique_recorded_descendants": unique_descendants[animal_id],
        "recorded_sire_links": sire_links[animal_id],
        "recorded_dam_links": dam_links[animal_id],
    } for animal_id in sorted(pedigree)]
    args.output.parent.mkdir(parents=True, exist_ok=True)
    with args.output.open("w", newline="") as file:
        writer = csv.DictWriter(file, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)

    top_three_sires = sire_links.most_common(3)
    print(f"Wrote {len(rows)} sampled dogs to {args.output}")
    print(f"Top three sires: {top_three_sires}; "
          f"{sum(count for _, count in top_three_sires)} of {sum(sire_links.values())} sire links")
    for animal_id in ("PFZ13D08", "PFZ27D02", "PFZ26F05", "PFZ24C06"):
        print(f"{animal_id}: {direct_offspring[animal_id]} direct offspring, "
              f"{unique_descendants[animal_id]} unique recorded descendants")


if __name__ == "__main__":
    main()
