#!/usr/bin/env python3
"""Interleave the Java addendum without recalculating baseline rows."""

import json
import runpy
from pathlib import Path


ROOT = Path(__file__).resolve().parent
BASELINE = ROOT / "baseline-combined.json"
FORWARD = ROOT / "java-hyperfine-forward.json"
REVERSE = ROOT / "java-hyperfine-reverse.json"
OUTPUT = ROOT / "combined.json"

COMMANDS = {
    "JDK SAX JIT",
    "JDK SAX GraalVM Native",
    "JDK StAX JIT",
    "JDK StAX GraalVM Native",
    "Aalto JIT",
    "Aalto GraalVM Native",
    "Woodstox JIT",
    "Woodstox GraalVM Native",
}
EXCLUDED = {"Aalto JIT", "Aalto GraalVM Native"}
EXCLUSION_REASON = (
    "accepts overlong UTF-8 and out-of-range Unicode code point U+110000"
)

summarize = runpy.run_path(
    str(ROOT / "combine.py")
)["summarize"]


def load_addendum(path):
    results = json.loads(path.read_text(encoding="utf-8"))["results"]
    if len(results) != 8 or {row["command"] for row in results} != COMMANDS:
        raise ValueError(f"unexpected Java command set in {path}")
    if any(len(row["times"]) != 1000 for row in results):
        raise ValueError(f"expected 1000 times per command in {path}")
    if any(len(row.get("exit_codes", [])) != 1000 for row in results):
        raise ValueError(f"expected 1000 exit codes per command in {path}")
    if any(any(code != 0 for code in row["exit_codes"]) for row in results):
        raise ValueError(f"non-zero benchmark exit code in {path}")
    return results


def assert_old_rows_unchanged(old, merged):
    current = {row["command"]: row for row in merged["ranked"]}
    for old_row in old["ranked"]:
        new_row = current[old_row["command"]]
        assert {k: v for k, v in old_row.items() if k != "rank"} == {
            k: v for k, v in new_row.items() if k != "rank"
        }
    assert merged["excluded"][: len(old["excluded"])] == old["excluded"]
    assert merged["control"] == old["control"]


def main():
    old = json.loads(BASELINE.read_text(encoding="utf-8"))
    forward = load_addendum(FORWARD)
    reverse = load_addendum(REVERSE)
    if [row["command"] for row in reverse] != [
        row["command"] for row in reversed(forward)
    ]:
        raise ValueError("reverse order is not the exact inverse of forward order")

    reverse_by_command = {row["command"]: row for row in reverse}
    java_rows = [
        summarize(
            row["command"],
            row["times"] + reverse_by_command[row["command"]]["times"],
        )
        for row in forward
    ]
    ranked = [dict(row) for row in old["ranked"]]
    ranked.extend(row for row in java_rows if row["command"] not in EXCLUDED)
    ranked.sort(key=lambda row: row["mean_seconds"])
    for rank, row in enumerate(ranked, 1):
        row["rank"] = rank

    merged = {
        "methodology": old["methodology"]
        | {
            "java_addendum": (
                "8 Java parser-only rows measured later as a separate sequential "
                "series with the same input, CPU 0, shell=none, 10 warmups and "
                "1000 runs per forward/reverse order; 26 baseline rows were not rerun"
            )
        },
        "ranked": ranked,
        "excluded": old["excluded"]
        + [
            row | {"reason": EXCLUSION_REASON}
            for row in java_rows
            if row["command"] in EXCLUDED
        ],
        "control": old["control"],
    }

    assert len(merged["ranked"]) == 29
    assert len(merged["excluded"]) == 4
    assert [row["rank"] for row in ranked] == list(range(1, 30))
    assert [row["mean_seconds"] for row in ranked] == sorted(
        row["mean_seconds"] for row in ranked
    )
    assert_old_rows_unchanged(old, merged)

    OUTPUT.write_text(
        json.dumps(merged, ensure_ascii=False, indent=2) + "\n",
        encoding="utf-8",
        newline="\n",
    )
    print(f"wrote {OUTPUT.name}: 29 ranked, 4 excluded, 1 control")


if __name__ == "__main__":
    main()
