"""Reproduce the pedigree-depth table in the AnimalTrace research note.

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

Only the six PLINK pedigree columns are read. No genotype calls are analyzed.
"""

from __future__ import annotations

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


def read_pedigree(path: Path) -> dict[str, tuple[str | None, str | None]]:
    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 animal_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 {animal_id}")
    return pedigree


def relationship_matrix(pedigree: dict[str, tuple[str | None, str | None]]):
    order = []
    state = {}

    def visit(animal_id):
        if state.get(animal_id) == 2:
            return
        if state.get(animal_id) == 1:
            raise ValueError(f"pedigree cycle at {animal_id}")
        state[animal_id] = 1
        for parent_id in pedigree[animal_id]:
            if parent_id is not None:
                visit(parent_id)
        state[animal_id] = 2
        order.append(animal_id)

    for animal_id in sorted(pedigree):
        visit(animal_id)
    position = {animal_id: i for i, animal_id in enumerate(order)}
    matrix = [[0.0] * len(order) for _ in order]
    for i, animal_id in enumerate(order):
        sire_id, dam_id = pedigree[animal_id]
        for j in range(i):
            matrix[i][j] = matrix[j][i] = 0.5 * (
                (matrix[position[sire_id]][j] if sire_id else 0.0)
                + (matrix[position[dam_id]][j] if dam_id else 0.0)
            )
        matrix[i][i] = 1.0 + (0.5 * matrix[position[sire_id]][position[dam_id]]
                               if sire_id and dam_id else 0.0)
    return position, matrix


def truncated_pedigree(pedigree, sire_id, dam_id, depth):
    # Depth 1 includes the two parents' own parents, depth 2 their grandparents, etc.
    seen_depth = {}
    pending = [(sire_id, 0), (dam_id, 0)]
    while pending:
        animal_id, generation = pending.pop()
        if generation >= seen_depth.get(animal_id, 10**9):
            continue
        seen_depth[animal_id] = generation
        if generation < depth:
            pending.extend((parent_id, generation + 1)
                           for parent_id in pedigree[animal_id] if parent_id)
    return {animal_id: pedigree[animal_id] if generation < depth else (None, None)
            for animal_id, generation in seen_depth.items()}


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="Derived CSV path")
    args = parser.parse_args()
    pedigree = read_pedigree(args.fam)
    position, matrix = relationship_matrix(pedigree)
    pairs = Counter((sire_id, dam_id) for sire_id, dam_id in pedigree.values()
                    if sire_id and dam_id)
    results = []
    for (sire_id, dam_id), offspring_count in sorted(pairs.items()):
        full = 0.5 * matrix[position[sire_id]][position[dam_id]]
        # Each recorded child's own F must equal the parental kinship calculation.
        children = [animal_id for animal_id, value in pedigree.items() if value == (sire_id, dam_id)]
        assert len(children) == offspring_count
        assert all(abs((matrix[position[child]][position[child]] - 1.0) - full) < 1e-12
                   for child in children)
        depths = []
        for depth in range(1, 6):
            limited = truncated_pedigree(pedigree, sire_id, dam_id, depth)
            limited_position, limited_matrix = relationship_matrix(limited)
            depths.append(0.5 * limited_matrix[limited_position[sire_id]][limited_position[dam_id]])
        assert all(left <= right + 1e-12 for left, right in zip(depths, depths[1:]))
        assert depths[-1] <= full + 1e-12
        results.append({
            "sire_sample_id": sire_id,
            "dam_sample_id": dam_id,
            "recorded_offspring": offspring_count,
            **{f"coi_depth_{i + 1}": f"{coi:.8f}" for i, coi in enumerate(depths)},
            "coi_full_record": f"{full:.8f}",
        })
    if len(pedigree) != 237 or len(pairs) != 30 or sum(pairs.values()) != 207:
        raise ValueError("source cohort size differs from the published analysis")
    args.output.parent.mkdir(parents=True, exist_ok=True)
    with args.output.open("w", newline="") as file:
        writer = csv.DictWriter(file, fieldnames=list(results[0]))
        writer.writeheader()
        writer.writerows(results)
    print(f"Wrote {len(results)} distinct recorded pairings to {args.output}")


if __name__ == "__main__":
    main()
