#!/usr/bin/env python3
"""Audit the frozen RiskCap backtest archives and render a promotion chart."""

from __future__ import annotations

import hashlib
import json
import zipfile
from pathlib import Path

import matplotlib.pyplot as plt
import pandas as pd


HERE = Path(__file__).resolve().parent
ARCHIVES = {
    "training": HERE / "train-12pairs.zip",
    "validation": HERE / "validation-12pairs.zip",
}
BASELINE = "VolatilityBreakout"
CANDIDATE = "VolatilityBreakoutRiskCap"


def load_archive(path: Path) -> dict:
    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_names = sorted(name for name in names if name.endswith(".py"))
        return {
            "sha256": hashlib.sha256(path.read_bytes()).hexdigest(),
            "result": json.loads(archive.read(result_name)),
            "config": json.loads(archive.read(config_name)),
            "strategy_sha256": {
                name: hashlib.sha256(archive.read(name)).hexdigest()
                for name in strategy_names
            },
        }


def strategy_summary(row: dict) -> dict:
    trades = row["trades"]
    return {
        "trades": row["total_trades"],
        "profit_pct": row["profit_total"] * 100,
        "profit_factor": row["profit_factor"],
        "max_drawdown_pct": row["max_drawdown_account"] * 100,
        "winrate": row["winrate"],
        "avg_stake": row["avg_stake_amount"],
        "worst_trade_usdt": min(trade["profit_abs"] for trade in trades),
        "best_trade_usdt": max(trade["profit_abs"] for trade in trades),
        "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 trade_frame(row: dict) -> pd.DataFrame:
    frame = pd.DataFrame(row["trades"])
    frame["close_date"] = pd.to_datetime(frame["close_date"], utc=True)
    return frame.sort_values("close_date").reset_index(drop=True)


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


def audit() -> tuple[dict, dict[str, tuple[pd.DataFrame, pd.DataFrame]]]:
    loaded = {label: load_archive(path) for label, path in ARCHIVES.items()}
    manifest = {
        "review_id": "2026-07-23-riskcap-promotion",
        "hypothesis": {
            "baseline": BASELINE,
            "candidate": CANDIDATE,
            "single_changed_factor": (
                "Position sizing denominator uses max(leveraged 3xATR risk, "
                "absolute hard-stop risk)."
            ),
            "pass_gate": (
                "Same trades; PF decline no worse than 0.02; max drawdown and "
                "worst single-trade loss both improve in training and validation."
            ),
        },
        "archives": {},
        "windows": {},
        "decision": None,
        "visual_review": None,
    }
    frames = {}
    passed = True
    for label, item in loaded.items():
        result = item["result"]["strategy"]
        baseline = result[BASELINE]
        candidate = result[CANDIDATE]
        base_frame = trade_frame(baseline)
        candidate_frame = trade_frame(candidate)
        base_keys = [trade_key(trade) for trade in baseline["trades"]]
        candidate_keys = [trade_key(trade) for trade in candidate["trades"]]
        exact_trade_paths = base_keys == candidate_keys
        base_summary = strategy_summary(baseline)
        candidate_summary = strategy_summary(candidate)
        gate = {
            "exact_trade_paths": exact_trade_paths,
            "pf_delta": (
                candidate_summary["profit_factor"] - base_summary["profit_factor"]
            ),
            "drawdown_delta_percentage_points": (
                candidate_summary["max_drawdown_pct"]
                - base_summary["max_drawdown_pct"]
            ),
            "worst_trade_delta_usdt": (
                candidate_summary["worst_trade_usdt"]
                - base_summary["worst_trade_usdt"]
            ),
        }
        gate["passed"] = (
            exact_trade_paths
            and gate["pf_delta"] >= -0.02
            and gate["drawdown_delta_percentage_points"] < 0
            and gate["worst_trade_delta_usdt"] > 0
            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]
            and base_summary["enable_protections"]
            and candidate_summary["enable_protections"]
            and base_summary["timeframe"] == candidate_summary["timeframe"] == "5m"
        )
        passed &= gate["passed"]
        manifest["archives"][label] = {
            "file": ARCHIVES[label].name,
            "sha256": item["sha256"],
            "strategy_sha256": item["strategy_sha256"],
        }
        manifest["windows"][label] = {
            "baseline": base_summary,
            "candidate": candidate_summary,
            "gate": gate,
        }
        frames[label] = (base_frame, candidate_frame)
    manifest["decision"] = (
        "PASS_OFFLINE_RISK_GATE" if passed else "REJECT_OFFLINE_RISK_GATE"
    )
    return manifest, frames


def render(manifest: dict, frames: dict[str, tuple[pd.DataFrame, pd.DataFrame]]) -> Path:
    fig, axes = plt.subplots(2, 2, figsize=(18, 11), dpi=160, facecolor="white")
    colors = {BASELINE: "#64748B", CANDIDATE: "#2563EB"}
    for ax, label in zip(axes[0], ("training", "validation")):
        baseline, candidate = frames[label]
        ax.plot(
            baseline["close_date"], baseline["profit_abs"].cumsum(),
            label="D48 baseline", color=colors[BASELINE], linewidth=1.4,
        )
        ax.plot(
            candidate["close_date"], candidate["profit_abs"].cumsum(),
            label="RiskCap", color=colors[CANDIDATE], linewidth=1.4,
        )
        ax.axhline(0, color="#CBD5E1", linewidth=.8)
        ax.set_title(f"{label.title()} cumulative realized PnL", loc="left", fontweight="bold")
        ax.set_ylabel("USDT")
        ax.legend(frameon=False)
        ax.grid(alpha=.12)

    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 = range(len(labels))
    axes[1, 0].bar(
        [value - .18 for value in x], baseline_values, width=.36,
        color=colors[BASELINE], label="D48 baseline",
    )
    axes[1, 0].bar(
        [value + .18 for value in x], candidate_values, width=.36,
        color=colors[CANDIDATE], label="RiskCap",
    )
    axes[1, 0].set_xticks(list(x), labels)
    axes[1, 0].set_title("Promotion-gate metrics", loc="left", fontweight="bold")
    axes[1, 0].legend(frameon=False)
    axes[1, 0].grid(axis="y", alpha=.12)

    base = pd.concat([frames["training"][0], frames["validation"][0]], ignore_index=True)
    candidate = pd.concat([frames["training"][1], frames["validation"][1]], ignore_index=True)
    axes[1, 1].scatter(
        base["stake_amount"], candidate["stake_amount"],
        s=14, alpha=.45, color=colors[CANDIDATE],
    )
    low = min(base["stake_amount"].min(), candidate["stake_amount"].min())
    high = max(base["stake_amount"].max(), candidate["stake_amount"].max())
    axes[1, 1].plot([low, high], [low, high], "--", color="#94A3B8", linewidth=1)
    axes[1, 1].set_xlabel("D48 stake (USDT)")
    axes[1, 1].set_ylabel("RiskCap stake (USDT)")
    axes[1, 1].set_title(
        "Same trade paths; hard-stop-bound stakes shrink",
        loc="left", fontweight="bold",
    )
    axes[1, 1].grid(alpha=.12)

    fig.suptitle(
        "D48 RiskCapFix · isolated position-sizing comparison",
        fontsize=17, fontweight="bold",
    )
    fig.text(
        .5, .025,
        "12 pairs · 5m · 2x isolated · 0.035% per side · protections enabled · exact trade-path match",
        ha="center", color="#64748B", fontsize=9,
    )
    fig.tight_layout(rect=[0, .06, 1, .95])
    output = HERE / "riskcap_promotion_comparison.png"
    fig.savefig(output, bbox_inches="tight", facecolor="white")
    plt.close(fig)
    return output


def main() -> int:
    manifest, frames = audit()
    chart = render(manifest, frames)
    manifest["chart"] = chart.name
    (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 manifest["decision"] == "PASS_OFFLINE_RISK_GATE" else 1


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