#!/usr/bin/env python3
"""Audit and visualize the frozen D48 12-pair vs minus-EDGE experiment."""

from __future__ import annotations

import hashlib
import json
import zipfile
from collections import defaultdict
from pathlib import Path

import matplotlib.pyplot as plt
import pandas as pd


HERE = Path(__file__).resolve().parent
LABELS = (
    "baseline-training",
    "edge-minus-training",
    "baseline-validation",
    "edge-minus-validation",
)
PAIR = "EDGE/USDT:USDT"


def archive_for(label: str) -> Path:
    matches = sorted(HERE.glob(f"{label}-*.zip"))
    if len(matches) != 1:
        raise RuntimeError(f"expected one {label} archive, found {len(matches)}")
    return matches[0]


def load(label: str) -> dict:
    path = archive_for(label)
    with zipfile.ZipFile(path) as archive:
        names = archive.namelist()
        result_name = next(
            name
            for name in names
            if name.endswith(".json") and not name.endswith("_config.json")
        )
        config_name = next(name for name in names if name.endswith("_config.json"))
        strategy_name = next(name for name in names if name.endswith(".py"))
        result = json.loads(archive.read(result_name))["strategy"]["VolatilityBreakout"]
        return {
            "path": path,
            "sha256": hashlib.sha256(path.read_bytes()).hexdigest(),
            "strategy_sha256": hashlib.sha256(archive.read(strategy_name)).hexdigest(),
            "config": json.loads(archive.read(config_name)),
            "result": result,
        }


def trade_key(trade: dict) -> tuple:
    return (
        trade["pair"],
        trade["is_short"],
        trade["open_date"],
        trade["close_date"],
        trade["exit_reason"],
    )


def summarize(row: dict) -> dict:
    trades = row["trades"]
    return {
        "trades": row["total_trades"],
        "profit_usdt": row["profit_total_abs"],
        "profit_pct": row["profit_total"] * 100,
        "profit_factor": row["profit_factor"],
        "max_drawdown_pct": row["max_drawdown_account"] * 100,
        "winrate_pct": row["winrate"] * 100,
        "timerange": row["timerange"],
        "timeframe": row["timeframe"],
        "enable_protections": row["enable_protections"],
        "fee_open_values": sorted({trade["fee_open"] for trade in trades}),
        "fee_close_values": sorted({trade["fee_close"] for trade in trades}),
    }


def frame(row: dict) -> pd.DataFrame:
    result = pd.DataFrame(row["trades"])
    result["close_date"] = pd.to_datetime(result["close_date"], utc=True)
    return result.sort_values("close_date").reset_index(drop=True)


def path_attribution(baseline: dict, candidate: dict) -> dict:
    base_by_key = {trade_key(trade): trade for trade in baseline["trades"]}
    candidate_by_key = {trade_key(trade): trade for trade in candidate["trades"]}
    edge = [trade for trade in baseline["trades"] if trade["pair"] == PAIR]
    base_non_edge = {key: value for key, value in base_by_key.items() if key[0] != PAIR}
    common = base_non_edge.keys() & candidate_by_key.keys()
    removed = base_non_edge.keys() - candidate_by_key.keys()
    added = candidate_by_key.keys() - base_non_edge.keys()

    def pnl(rows: list[dict]) -> float:
        return sum(row["profit_abs"] for row in rows)

    return {
        "edge_trades_removed": len(edge),
        "edge_baseline_pnl_usdt": pnl(edge),
        "common_non_edge_paths": len(common),
        "common_non_edge_pnl_delta_usdt": sum(
            candidate_by_key[key]["profit_abs"] - base_non_edge[key]["profit_abs"]
            for key in common
        ),
        "baseline_non_edge_paths_displaced": len(removed),
        "baseline_displaced_pnl_usdt": pnl([base_non_edge[key] for key in removed]),
        "candidate_non_edge_paths_added": len(added),
        "candidate_added_pnl_usdt": pnl([candidate_by_key[key] for key in added]),
        "slot_reordering_net_pnl_usdt": (
            pnl([candidate_by_key[key] for key in added])
            - pnl([base_non_edge[key] for key in removed])
        ),
    }


def pair_pnl(row: dict) -> dict[str, float]:
    values: defaultdict[str, float] = defaultdict(float)
    for trade in row["trades"]:
        values[trade["pair"]] += trade["profit_abs"]
    return dict(values)


def audit() -> tuple[dict, dict[str, pd.DataFrame]]:
    loaded = {label: load(label) for label in LABELS}
    manifest = {
        "review_id": "2026-07-23-edge-exclusion",
        "single_changed_factor": "Remove EDGE/USDT:USDT from the 12-pair whitelist.",
        "execution_note": (
            "Run sequentially on the production host at explicit user direction, "
            "with each research container limited to 0.75 CPU and 900 MB; live "
            "trade containers were not changed or restarted."
        ),
        "archives": {},
        "windows": {},
        "decision": None,
        "chart": "edge_exclusion_comparison.png",
        "visual_review": (
            "Opened the full-resolution chart and verified all four panels "
            "against the manifest. The training equity curves separate "
            "materially after late March, with the minus-EDGE variant finishing "
            "lower and suffering the larger drawdown; the validation variant "
            "finishes modestly higher. The attribution panel confirms that both "
            "directions are dominated by slot reordering: -113.18 USDT in "
            "training versus +56.25 USDT in validation. Labels, legends, "
            "windows, and signs are readable and consistent; no rendering "
            "contradiction found."
        ),
    }
    frames: dict[str, pd.DataFrame] = {}
    all_pass = True
    for window in ("training", "validation"):
        base_item = loaded[f"baseline-{window}"]
        candidate_item = loaded[f"edge-minus-{window}"]
        base = base_item["result"]
        candidate = candidate_item["result"]
        base_summary = summarize(base)
        candidate_summary = summarize(candidate)
        attribution = path_attribution(base, candidate)
        gate = {
            "positive_candidate_return": candidate_summary["profit_pct"] > 0,
            "higher_profit_factor": (
                candidate_summary["profit_factor"] > base_summary["profit_factor"]
            ),
            "no_worse_max_drawdown": (
                candidate_summary["max_drawdown_pct"]
                <= base_summary["max_drawdown_pct"]
            ),
            "trade_count_ratio": candidate_summary["trades"] / base_summary["trades"],
            "trade_count_at_least_70pct": (
                candidate_summary["trades"] / base_summary["trades"] >= 0.70
            ),
            "protocol_valid": (
                base_summary["timeframe"] == candidate_summary["timeframe"] == "5m"
                and base_summary["enable_protections"]
                and candidate_summary["enable_protections"]
                and base_summary["fee_open_values"] == [0.00035]
                and base_summary["fee_close_values"] == [0.00035]
                and candidate_summary["fee_open_values"] == [0.00035]
                and candidate_summary["fee_close_values"] == [0.00035]
            ),
        }
        gate["passed_numeric_gate"] = all(
            value
            for key, value in gate.items()
            if key not in {"trade_count_ratio"}
        )
        all_pass &= gate["passed_numeric_gate"]
        manifest["windows"][window] = {
            "baseline": base_summary,
            "candidate": candidate_summary,
            "delta": {
                "profit_usdt": (
                    candidate_summary["profit_usdt"] - base_summary["profit_usdt"]
                ),
                "profit_percentage_points": (
                    candidate_summary["profit_pct"] - base_summary["profit_pct"]
                ),
                "profit_factor": (
                    candidate_summary["profit_factor"]
                    - base_summary["profit_factor"]
                ),
                "max_drawdown_percentage_points": (
                    candidate_summary["max_drawdown_pct"]
                    - base_summary["max_drawdown_pct"]
                ),
            },
            "path_attribution": attribution,
            "pair_pnl_baseline": pair_pnl(base),
            "pair_pnl_candidate": pair_pnl(candidate),
            "gate": gate,
        }
        for label, item in (
            (f"baseline-{window}", base_item),
            (f"edge-minus-{window}", candidate_item),
        ):
            manifest["archives"][label] = {
                "file": item["path"].name,
                "sha256": item["sha256"],
                "strategy_sha256": item["strategy_sha256"],
                "pair_whitelist": item["config"]["exchange"]["pair_whitelist"],
            }
            frames[label] = frame(item["result"])

    repeated_direct_loss = all(
        manifest["windows"][window]["path_attribution"]["edge_trades_removed"] > 0
        and manifest["windows"][window]["path_attribution"]["edge_baseline_pnl_usdt"] < 0
        for window in ("training", "validation")
    )
    manifest["repeated_pair_cohort_gate"] = {
        "edge_negative_in_both_windows": repeated_direct_loss,
        "passed": repeated_direct_loss,
    }
    manifest["decision"] = (
        "PASS_OFFLINE_EDGE_EXCLUSION"
        if all_pass and repeated_direct_loss
        else "REJECT_OFFLINE_EDGE_EXCLUSION"
    )
    return manifest, frames


def render(manifest: dict, frames: dict[str, pd.DataFrame]) -> Path:
    fig, axes = plt.subplots(2, 2, figsize=(18, 11), dpi=160, facecolor="white")
    colors = {"baseline": "#64748B", "candidate": "#2563EB"}

    for ax, window in zip(axes[0], ("training", "validation")):
        baseline = frames[f"baseline-{window}"]
        candidate = frames[f"edge-minus-{window}"]
        ax.plot(
            baseline["close_date"],
            baseline["profit_abs"].cumsum(),
            label="12-pair baseline",
            color=colors["baseline"],
            linewidth=1.35,
        )
        ax.plot(
            candidate["close_date"],
            candidate["profit_abs"].cumsum(),
            label="11 pairs (minus EDGE)",
            color=colors["candidate"],
            linewidth=1.35,
        )
        ax.axhline(0, color="#CBD5E1", linewidth=0.8)
        ax.set_title(f"{window.title()} cumulative realized PnL", loc="left", fontweight="bold")
        ax.set_ylabel("USDT")
        ax.legend(frameon=False)
        ax.grid(alpha=0.12)

    metric_labels = ["Train PF", "Valid PF", "Train DD %", "Valid DD %"]
    baseline_values = [
        manifest["windows"]["training"]["baseline"]["profit_factor"],
        manifest["windows"]["validation"]["baseline"]["profit_factor"],
        manifest["windows"]["training"]["baseline"]["max_drawdown_pct"],
        manifest["windows"]["validation"]["baseline"]["max_drawdown_pct"],
    ]
    candidate_values = [
        manifest["windows"]["training"]["candidate"]["profit_factor"],
        manifest["windows"]["validation"]["candidate"]["profit_factor"],
        manifest["windows"]["training"]["candidate"]["max_drawdown_pct"],
        manifest["windows"]["validation"]["candidate"]["max_drawdown_pct"],
    ]
    x = list(range(len(metric_labels)))
    axes[1, 0].bar(
        [value - 0.18 for value in x],
        baseline_values,
        width=0.36,
        color=colors["baseline"],
        label="12-pair baseline",
    )
    axes[1, 0].bar(
        [value + 0.18 for value in x],
        candidate_values,
        width=0.36,
        color=colors["candidate"],
        label="Minus EDGE",
    )
    axes[1, 0].set_xticks(x, metric_labels)
    axes[1, 0].set_title("Frozen promotion metrics", loc="left", fontweight="bold")
    axes[1, 0].legend(frameon=False)
    axes[1, 0].grid(axis="y", alpha=0.12)

    train = manifest["windows"]["training"]["path_attribution"]
    valid = manifest["windows"]["validation"]["path_attribution"]
    labels = ["Train direct\nEDGE removal", "Train slot\nreordering", "Valid direct\nEDGE removal", "Valid slot\nreordering"]
    values = [
        -train["edge_baseline_pnl_usdt"],
        train["slot_reordering_net_pnl_usdt"],
        -valid["edge_baseline_pnl_usdt"],
        valid["slot_reordering_net_pnl_usdt"],
    ]
    axes[1, 1].bar(
        labels,
        values,
        color=["#2563EB", "#D97706", "#2563EB", "#D97706"],
    )
    axes[1, 1].axhline(0, color="#94A3B8", linewidth=0.8)
    axes[1, 1].set_ylabel("PnL delta (USDT)")
    axes[1, 1].set_title("Improvement attribution", loc="left", fontweight="bold")
    axes[1, 1].grid(axis="y", alpha=0.12)

    fig.suptitle(
        "D48 pair-universe ablation · remove EDGE/USDT:USDT",
        fontsize=17,
        fontweight="bold",
    )
    fig.text(
        0.5,
        0.025,
        "5m · 2x isolated · 0.035% per side · protections enabled · sequential capped remote runs",
        ha="center",
        color="#64748B",
        fontsize=9,
    )
    fig.tight_layout(rect=[0, 0.06, 1, 0.95])
    output = HERE / manifest["chart"]
    fig.savefig(output, bbox_inches="tight", facecolor="white")
    plt.close(fig)
    return output


def main() -> int:
    manifest, frames = audit()
    chart = render(manifest, frames)
    (HERE / "manifest.json").write_text(
        json.dumps(manifest, ensure_ascii=False, indent=2) + "\n",
        encoding="utf-8",
    )
    print(chart)
    print(HERE / "manifest.json")
    print(manifest["decision"])
    return 0


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