"""Summarize the same pedigree COI by distinct pairing and sampled offspring.

Input: campbell-2016-pedigree-depth.csv, generated from the public Campbell et al.
pedigree by analyze-campbell-pedigree-depth.py. Uses the full-record COI column.
"""

from __future__ import annotations

import argparse
import csv
from decimal import Decimal
from pathlib import Path


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("pair_csv", type=Path, help="30-pair pedigree-depth CSV")
    parser.add_argument("--output", type=Path, required=True, help="Summary CSV path")
    args = parser.parse_args()

    with args.pair_csv.open(newline="") as file:
        pairs = list(csv.DictReader(file))
    if len(pairs) != 30:
        raise ValueError("expected 30 distinct recorded pairings")
    keys = [(row["sire_sample_id"], row["dam_sample_id"]) for row in pairs]
    if len(set(keys)) != len(keys):
        raise ValueError("duplicate sire–dam pairing")

    offspring = [int(row["recorded_offspring"]) for row in pairs]
    coi = [Decimal(row["coi_full_record"]) for row in pairs]
    if any(n <= 0 for n in offspring) or any(not 0 <= f <= 1 for f in coi):
        raise ValueError("invalid offspring count or COI")
    if sum(offspring) != 207:
        raise ValueError("expected 207 sampled offspring with two recorded parents")

    summaries = []
    for unit, weights in (("distinct_recorded_pairing", [1] * len(pairs)),
                          ("sampled_offspring_with_two_parents", offspring)):
        denominator = sum(weights)
        nonzero = sum(weight for weight, f in zip(weights, coi) if f > 0)
        weighted_coi = sum(Decimal(weight) * f for weight, f in zip(weights, coi))
        summaries.append({
            "analysis_unit": unit,
            "denominator": denominator,
            "nonzero_coi_count": nonzero,
            "nonzero_coi_share_percent": f"{Decimal(100) * nonzero / denominator:.6f}",
            "mean_expected_offspring_pedigree_coi_percent":
                f"{Decimal(100) * weighted_coi / denominator:.6f}",
        })

    args.output.parent.mkdir(parents=True, exist_ok=True)
    with args.output.open("w", newline="") as file:
        writer = csv.DictWriter(file, fieldnames=list(summaries[0]))
        writer.writeheader()
        writer.writerows(summaries)
    for row in summaries:
        print(row)


if __name__ == "__main__":
    main()
