#!/usr/bin/env python3
"""Recompute the public cross-wave summary without private answer archives."""

import argparse
import csv
import io
import sys
from decimal import Decimal, ROUND_HALF_UP
from pathlib import Path


FIELDS = [
    "collection_date",
    "wave_id",
    "engine",
    "retained_answers",
    "answers_with_visible_links",
    "answer_url_occurrences",
    "distinct_normalized_urls",
    "distinct_source_hosts",
    "complete_repeat_pairs",
    "pairs_with_identical_url_sets",
    "pairs_with_no_shared_urls",
    "mean_jaccard",
    "mean_jaccard_percent",
    "scope",
]

WAVES = [
    {
        "collection_date": "2026-09-09",
        "wave_id": "consumer-citation-pilot-2026-09-09",
        "directory": "consumer-citation-pilot-2026-09-09",
    },
    {
        "collection_date": "2026-09-14",
        "wave_id": "consumer-repeatability-2026-09-14",
        "directory": "consumer-citation-wave-2026-09-14",
    },
]


def read_rows(path: Path) -> list[dict[str, str]]:
    with path.open(newline="", encoding="utf-8") as handle:
        return list(csv.DictReader(handle))


def summarize(resource_root: Path) -> list[dict[str, object]]:
    rows: list[dict[str, object]] = []
    for wave in WAVES:
        directory = resource_root / "resources" / wave["directory"]
        runs = read_rows(directory / "runs.csv")
        links = read_rows(directory / "visible-links.csv")
        overlaps = read_rows(directory / "repeat-overlap.csv")

        for source_name, source_rows in (
            ("runs.csv", runs),
            ("visible-links.csv", links),
            ("repeat-overlap.csv", overlaps),
        ):
            source_wave_ids = {
                row["wave_id"] for row in source_rows if row.get("wave_id")
            }
            if source_wave_ids and source_wave_ids != {wave["wave_id"]}:
                raise ValueError(
                    f"{wave['directory']}/{source_name}: expected canonical wave_id "
                    f"{wave['wave_id']}; found {sorted(source_wave_ids)}"
                )

        for engine in ("chatgpt", "claude"):
            engine_runs = [
                row
                for row in runs
                if row["engine"] == engine
                and row.get("status", "complete") == "complete"
            ]
            engine_links = [row for row in links if row["engine"] == engine]
            engine_pairs = [
                row
                for row in overlaps
                if row["engine"] == engine
                and row.get("overlap_status", "available") == "available"
                and int(row["union_urls"]) > 0
            ]
            jaccards = [
                Decimal(row["shared_urls"]) / Decimal(row["union_urls"])
                for row in engine_pairs
            ]

            if len(engine_runs) != 24 or len(engine_pairs) != 12:
                raise ValueError(
                    f"{wave['wave_id']} {engine}: expected 24 retained runs and "
                    f"12 complete pairs; found {len(engine_runs)} and {len(engine_pairs)}"
                )

            mean = (sum(jaccards) / Decimal(len(jaccards))).quantize(
                Decimal("0.000001"), rounding=ROUND_HALF_UP
            )
            rows.append(
                {
                    "collection_date": wave["collection_date"],
                    "wave_id": wave["wave_id"],
                    "engine": engine,
                    "retained_answers": len(engine_runs),
                    "answers_with_visible_links": sum(
                        int(row["visible_unique_urls"]) > 0 for row in engine_runs
                    ),
                    "answer_url_occurrences": len(engine_links),
                    "distinct_normalized_urls": len(
                        {row["url"] for row in engine_links}
                    ),
                    "distinct_source_hosts": len(
                        {row["domain"].removeprefix("www.") for row in engine_links}
                    ),
                    "complete_repeat_pairs": len(engine_pairs),
                    "pairs_with_identical_url_sets": sum(
                        value == Decimal("1") for value in jaccards
                    ),
                    "pairs_with_no_shared_urls": sum(
                        int(row["shared_urls"]) == 0 for row in engine_pairs
                    ),
                    "mean_jaccard": f"{mean:.6f}",
                    "mean_jaccard_percent": f"{mean * 100:.1f}",
                    "scope": "Separate dated consumer-interface convenience sample",
                }
            )
    return rows


def render_csv(rows: list[dict[str, object]]) -> str:
    output = io.StringIO(newline="")
    writer = csv.DictWriter(output, fieldnames=FIELDS, lineterminator="\n")
    writer.writeheader()
    writer.writerows(rows)
    return output.getvalue()


def main() -> int:
    default_root = Path(__file__).resolve().parents[2]
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--root", type=Path, default=default_root)
    parser.add_argument(
        "--check",
        action="store_true",
        help="compare the recomputation with the checked-in aggregate",
    )
    args = parser.parse_args()

    rendered = render_csv(summarize(args.root.resolve()))
    if not args.check:
        sys.stdout.write(rendered)
        return 0

    expected_path = (
        args.root.resolve()
        / "resources"
        / "consumer-citation-repeatability-study"
        / "study-summary.csv"
    )
    expected = expected_path.read_text(encoding="utf-8")
    if rendered != expected:
        print(f"recomputed aggregate differs from {expected_path}", file=sys.stderr)
        return 1
    print("verified 4 separate wave/interface rows from 96 public run records")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
