#!/usr/bin/env python3
"""Independent verifier for the public Kuhn Poker CFR experiment."""

from __future__ import annotations

import csv
import hashlib
import json
import math
from pathlib import Path


ROOT = Path(__file__).resolve().parent
RANKS = ("J", "Q", "K")
DEALS = tuple((left, right) for left in RANKS for right in RANKS if left != right)
CHECKPOINTS = (1, 2, 5, 10, 20, 50, 100, 200, 500, 1_000, 2_000, 5_000,
               10_000, 20_000, 50_000, 100_000)
MANIFEST_FILES = (
    "README.md",
    "experiment.json",
    "convergence.csv",
    "final-strategies.csv",
    "convergence.svg",
    "convergence-ja.svg",
    "game-tree.svg",
    "game-tree-ja.svg",
    "generate.mjs",
    "verify.py",
)


def key(player: int, card: str, history: str) -> str:
    return f"P{player}|{card}|{history or 'root'}"


def player_for(history: str) -> int:
    return 0 if history in ("", "cb") else 1


def utility_p0(cards: tuple[str, str], history: str) -> float | None:
    if history == "bf":
        return 1.0
    if history == "cbf":
        return -1.0
    if history not in ("cc", "bc", "cbc"):
        return None
    stake = 1.0 if history == "cc" else 2.0
    return stake if RANKS.index(cards[0]) > RANKS.index(cards[1]) else -stake


def player_keys(player: int) -> tuple[str, ...]:
    histories = ("", "cb") if player == 0 else ("c", "b")
    return tuple(key(player, card, history) for card in RANKS for history in histories)


ALL_KEYS = player_keys(0) + player_keys(1)


def regret_match(values: list[float]) -> list[float]:
    positive = [max(0.0, value) for value in values]
    total = sum(positive)
    return [value / total for value in positive] if total else [0.5, 0.5]


def policies(tables: dict[str, dict[str, list[float]]]) -> tuple[dict, dict]:
    current = {name: regret_match(row["regret"]) for name, row in tables.items()}
    average = {}
    for name, row in tables.items():
        total = sum(row["sum"])
        average[name] = [value / total for value in row["sum"]] if total else [0.5, 0.5]
    return current, average


def evaluate(policy: dict[str, list[float]], responder: int | None = None,
             pure: dict[str, int] | None = None) -> float:
    def walk(cards: tuple[str, str], history: str) -> float:
        terminal = utility_p0(cards, history)
        if terminal is not None:
            return terminal
        player = player_for(history)
        name = key(player, cards[player], history)
        actions = ("c", "b") if history in ("", "c") else ("f", "c")
        if player == responder:
            return walk(cards, history + actions[pure[name]])
        return sum(probability * walk(cards, history + actions[index])
                   for index, probability in enumerate(policy[name]))

    return sum(walk(cards, "") for cards in DEALS) / len(DEALS)


def best_response(policy: dict[str, list[float]], responder: int) -> float:
    names = player_keys(responder)
    best = -math.inf
    for mask in range(1 << len(names)):
        pure = {name: (mask >> index) & 1 for index, name in enumerate(names)}
        value_p0 = evaluate(policy, responder, pure)
        best = max(best, value_p0 if responder == 0 else -value_p0)
    return best


def measure(policy: dict[str, list[float]]) -> dict[str, float]:
    value = evaluate(policy)
    response0 = best_response(policy, 0)
    response1 = best_response(policy, 1)
    nash_conv = response0 + response1
    return {
        "on_policy_value_p0": value,
        "best_response_value_p0": response0,
        "best_response_value_p1": response1,
        "unilateral_improvement_p0": response0 - value,
        "unilateral_improvement_p1": response1 + value,
        "nash_conv": nash_conv,
        "exploitability": nash_conv / 2,
    }


def run_cfr() -> tuple[list[dict], dict, dict]:
    tables = {name: {"regret": [0.0, 0.0], "sum": [0.0, 0.0]} for name in ALL_KEYS}
    retained = []

    def traverse(cards, history, reach0, reach1, policy, deltas):
        terminal = utility_p0(cards, history)
        if terminal is not None:
            return terminal
        player = player_for(history)
        name = key(player, cards[player], history)
        strategy = policy[name]
        actions = ("c", "b") if history in ("", "c") else ("f", "c")
        utilities = [traverse(
            cards,
            history + action,
            reach0 * strategy[index] if player == 0 else reach0,
            reach1 * strategy[index] if player == 1 else reach1,
            policy,
            deltas,
        ) for index, action in enumerate(actions)]
        node = sum(strategy[index] * utilities[index] for index in range(2))
        chance = 1 / len(DEALS)
        for index in range(2):
            if player == 0:
                deltas[name][index] += chance * reach1 * (utilities[index] - node)
                tables[name]["sum"][index] += chance * reach0 * strategy[index]
            else:
                deltas[name][index] += chance * reach0 * (node - utilities[index])
                tables[name]["sum"][index] += chance * reach1 * strategy[index]
        return node

    for iteration in range(1, CHECKPOINTS[-1] + 1):
        current, _ = policies(tables)
        deltas = {name: [0.0, 0.0] for name in ALL_KEYS}
        for cards in DEALS:
            traverse(cards, "", 1.0, 1.0, current, deltas)
        for name in ALL_KEYS:
            for index in range(2):
                tables[name]["regret"][index] += deltas[name][index]
        if iteration in CHECKPOINTS:
            current, average = policies(tables)
            retained.append({
                "iteration": iteration,
                "average": measure(average),
                "current": measure(current),
            })
    current, average = policies(tables)
    return retained, current, average


def require(condition: bool, message: str) -> None:
    if not condition:
        raise AssertionError(message)


def assert_close(actual: float, expected: float, label: str, tolerance: float = 2e-12):
    require(
        math.isclose(actual, expected, rel_tol=0.0, abs_tol=tolerance),
        f"{label}: expected {expected}, got {actual}",
    )


def main() -> None:
    data = json.loads((ROOT / "experiment.json").read_text(encoding="utf-8"))
    expected_deals = tuple("".join(deal) for deal in DEALS)
    require(
        tuple(data["experiment"]["chance_deals"]) == expected_deals,
        "experiment chance deals do not match the six ordered Kuhn Poker deals",
    )
    require(
        data["experiment"]["best_response_method"].startswith("exhaustive enumeration"),
        "experiment does not declare exhaustive best-response enumeration",
    )
    require(
        len(player_keys(0)) == len(player_keys(1)) == 6,
        "each player must have exactly six information sets",
    )
    require(
        all(name.count("|") == 2 for name in ALL_KEYS),
        "information-set keys must have exactly three fields",
    )
    require(
        all(opponent not in name for name in ALL_KEYS
            for opponent in ("JK", "JQ", "QJ", "QK", "KJ", "KQ")),
        "information-set keys must not encode the opponent's hidden card",
    )

    retained, final_current, final_average = run_cfr()
    require(
        [row["iteration"] for row in retained] == list(CHECKPOINTS),
        "independent rerun retained an unexpected checkpoint sequence",
    )
    require(
        len(data["checkpoints"]) == len(retained),
        "published checkpoint count does not match the independent rerun",
    )
    for actual, expected in zip(retained, data["checkpoints"], strict=True):
        require(
            actual["iteration"] == expected["iteration"],
            "published checkpoint iteration does not match the independent rerun",
        )
        for profile in ("average", "current"):
            for metric, value in actual[profile].items():
                assert_close(value, expected[profile][metric],
                             f"{actual['iteration']}/{profile}/{metric}")

    rows = data["final_strategies"]
    require(len(rows) == 24, "published final strategies must contain 24 rows")
    for row in rows:
        policy = final_average if row["profile"] == "reach_weighted_average" else final_current
        assert_close(policy[row["information_set"]][0], row["probability_0"],
                     f"{row['profile']}/{row['information_set']}/p0")
        assert_close(policy[row["information_set"]][1], row["probability_1"],
                     f"{row['profile']}/{row['information_set']}/p1")

    assert_close(-1 / 18, data["experiment"]["game_value_p0"], "game value")
    final = retained[-1]
    require(
        abs(final["average"]["on_policy_value_p0"] + 1 / 18) < 2e-6,
        "average strategy did not approach the Kuhn Poker game value",
    )
    require(
        final["average"]["exploitability"] < 0.001,
        "average-strategy exploitability did not fall below 0.001",
    )
    require(
        final["current"]["exploitability"] > 0.2,
        "last-iterate exploitability no longer demonstrates oscillation",
    )

    with (ROOT / "convergence.csv").open(encoding="utf-8", newline="") as handle:
        trace_rows = list(csv.DictReader(handle))
    require(trace_rows, "convergence CSV must contain at least one data row")
    require(trace_rows[0]["iteration"] == "1", "convergence CSV must begin at iteration 1")
    require(
        trace_rows[-1]["iteration"] == "100000",
        "convergence CSV must end at iteration 100000",
    )
    require(len(trace_rows) == 1_006, "convergence CSV must contain 1,006 data rows")

    manifest = {}
    for line in (ROOT / "MANIFEST.sha256").read_text(encoding="utf-8").splitlines():
        digest, name = line.split("  ", 1)
        manifest[name] = digest
    require(
        tuple(manifest) == MANIFEST_FILES,
        "manifest entries do not match the fixed public artifact allow-list",
    )
    for name in MANIFEST_FILES:
        actual = hashlib.sha256((ROOT / name).read_bytes()).hexdigest()
        require(actual == manifest[name], f"SHA-256 mismatch: {name}")

    print("PASS: independent Kuhn Poker CFR, exact best responses, artifacts and manifest verified")


if __name__ == "__main__":
    main()
