"""Compare direct-family links with recorded relationship in 207 sampled dogs.

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

Only pedigree columns are read. Pairs are comparisons, not proposed matings.
"""

from __future__ import annotations

import argparse
import csv
from itertools import combinations
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 write_csv(path: Path, fieldnames: list[str], rows: list[dict]):
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", newline="") as file:
        writer = csv.DictWriter(file, fieldnames=fieldnames)
        writer.writeheader()
        writer.writerows(rows)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("fam", type=Path, help="Path to new_pedigree_dogs.fam")
    parser.add_argument("--pairs-output", type=Path, required=True)
    parser.add_argument("--summary-output", type=Path, required=True)
    args = parser.parse_args()
    pedigree = read_pedigree(args.fam)
    cohort = sorted(animal_id for animal_id, parents in pedigree.items() if all(parents))
    if len(pedigree) != 237 or len(cohort) != 207:
        raise ValueError("source cohort differs from the published analysis")
    position, matrix = relationship_matrix(pedigree)

    rows = []
    counts = {"immediate_family_link": [0, 0], "no_immediate_family_link": [0, 0]}
    for animal_a, animal_b in combinations(cohort, 2):
        sire_a, dam_a = pedigree[animal_a]
        sire_b, dam_b = pedigree[animal_b]
        immediate = (animal_a in pedigree[animal_b] or animal_b in pedigree[animal_a]
                     or sire_a == sire_b or dam_a == dam_b)
        category = "immediate_family_link" if immediate else "no_immediate_family_link"
        relationship = matrix[position[animal_a]][position[animal_b]]
        nonzero = relationship > 1e-12
        counts[category][0] += 1
        counts[category][1] += nonzero
        rows.append({
            "sample_id_a": animal_a,
            "sample_id_b": animal_b,
            "family_link_category": category,
            "recorded_relationship_percent": f"{relationship * 100:.8f}",
            "nonzero_recorded_relationship": nonzero,
        })

    if (len(rows), sum(v[1] for v in counts.values()),
            counts["no_immediate_family_link"][1]) != (21321, 10386, 7620):
        raise ValueError("pairwise results differ from the published analysis")
    assert counts["immediate_family_link"][0] == counts["immediate_family_link"][1]
    summary = []
    for category, (total, nonzero) in counts.items():
        summary.append({
            "family_link_category": category,
            "pair_comparisons": total,
            "nonzero_recorded_relationship": nonzero,
            "zero_within_recorded_pedigree": total - nonzero,
        })
    summary.append({
        "family_link_category": "all_207_sample_comparisons",
        "pair_comparisons": len(rows),
        "nonzero_recorded_relationship": sum(v[1] for v in counts.values()),
        "zero_within_recorded_pedigree": len(rows) - sum(v[1] for v in counts.values()),
    })
    write_csv(args.pairs_output, list(rows[0]), rows)
    write_csv(args.summary_output, list(summary[0]), summary)
    for result in summary:
        print(result)


if __name__ == "__main__":
    main()
