#!/usr/bin/env python3
"""從公開逐次判定表複算網站上的三種同向率。"""

import csv
import math
from collections import Counter, defaultdict
from pathlib import Path


SOURCE = Path(__file__).with_name("rollcalls.csv")
ALIGNMENTS = ("dpp", "kmt", "neither")


def main():
    with SOURCE.open(encoding="utf-8-sig", newline="") as handle:
        rows = list(csv.DictReader(handle))

    event_sizes = Counter(
        (row["term"], row["target_party"], row["event_id"]) for row in rows
    )
    totals = defaultdict(lambda: defaultdict(float))

    for row in rows:
        event_key = (row["term"], row["target_party"], row["event_id"])
        event_size = event_sizes[event_key]
        weights = {
            "unweighted": 1.0,
            "sqrt": 1.0 / math.sqrt(event_size),
            "event": 1.0 / event_size,
        }
        for method, weight in weights.items():
            totals[(row["term"], row["target_party"], method)][row["alignment"]] += weight

    print("term,target_party,method,total,dpp_rate,kmt_rate,neither_rate")
    for key in sorted(totals, key=lambda value: (int(value[0]), value[1], value[2])):
        counts = totals[key]
        total = sum(counts.values())
        rates = [counts[name] / total for name in ALIGNMENTS]
        print(",".join([*key, f"{total:.6f}", *(f"{rate:.6f}" for rate in rates)]))


if __name__ == "__main__":
    main()
