"""Test one-link pedigree omissions against fixed recorded dog pairings.

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

Only PLINK pedigree columns are used. This is a record-sensitivity experiment,
not a reconstruction of when records changed or an analysis of genotype calls.
"""

from __future__ import annotations

import argparse
import csv
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 pairing_coi(position, matrix, sire_id, dam_id):
    return 0.5 * matrix[position[sire_id]][position[dam_id]]


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("--links-output", type=Path, required=True)
    parser.add_argument("--pair-changes-output", type=Path, required=True)
    args = parser.parse_args()

    pedigree = read_pedigree(args.fam)
    pairs = sorted({(sire_id, dam_id) for sire_id, dam_id in pedigree.values()
                    if sire_id and dam_id})
    if len(pedigree) != 237 or len(pairs) != 30:
        raise ValueError("source cohort size differs from the published analysis")
    position, matrix = relationship_matrix(pedigree)
    full_coi = {pair: pairing_coi(position, matrix, *pair) for pair in pairs}
    for child_id, pair in pedigree.items():
        if all(pair):
            assert abs((matrix[position[child_id]][position[child_id]] - 1.0)
                       - full_coi[pair]) < 1e-12

    links = []
    pair_changes = []
    for child_id, (sire_id, dam_id) in sorted(pedigree.items()):
        for side, parent_id in enumerate((sire_id, dam_id)):
            if parent_id is None:
                continue
            role = ("sire", "dam")[side]
            omitted = pedigree.copy()
            omitted[child_id] = ((None, dam_id) if side == 0 else (sire_id, None))
            omitted_position, omitted_matrix = relationship_matrix(omitted)
            affected = []
            apparent_zeros = 0
            for pair_sire, pair_dam in pairs:
                pair = (pair_sire, pair_dam)
                hidden = pairing_coi(omitted_position, omitted_matrix, *pair)
                full = full_coi[pair]
                difference = full - hidden
                if difference < -1e-12:
                    raise AssertionError("omitting a parent link increased a COI")
                if difference <= 1e-12:
                    continue
                affected.append(difference)
                apparent_zero = hidden < 1e-12 and full > 1e-12
                apparent_zeros += apparent_zero
                pair_changes.append({
                    "child_id": child_id,
                    "omitted_parent_id": parent_id,
                    "parent_role": role,
                    "pair_sire_id": pair_sire,
                    "pair_dam_id": pair_dam,
                    "coi_with_link_percent": f"{full * 100:.8f}",
                    "coi_without_link_percent": f"{hidden * 100:.8f}",
                    "change_percentage_points": f"{difference * 100:.8f}",
                    "apparent_zero_without_link": apparent_zero,
                })
            links.append({
                "child_id": child_id,
                "omitted_parent_id": parent_id,
                "parent_role": role,
                "affected_recorded_pairs": len(affected),
                "apparent_zeros": apparent_zeros,
                "largest_change_percentage_points": f"{max(affected, default=0) * 100:.8f}",
            })

    if len(links) != 419:
        raise ValueError("parent-link count differs from the published analysis")
    write_csv(args.links_output, list(links[0]), links)
    write_csv(args.pair_changes_output, list(pair_changes[0]), pair_changes)
    print(f"Analyzed {len(links)} one-link omissions and {len(pairs)} fixed pairings; "
          f"{sum(int(row['affected_recorded_pairs']) > 0 for row in links)} links changed "
          f"at least one pairing, producing {len(pair_changes)} pair-level changes.")


if __name__ == "__main__":
    main()
