#!/usr/bin/env python3
"""Build the deterministic evidence summary and comparison chart."""

from __future__ import annotations

import json
import math
import zipfile
from datetime import datetime
from pathlib import Path

import matplotlib.dates as mdates
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd


HERE = Path(__file__).resolve().parent
TRAIN = HERE / "oss-round2-training-sanitized.zip"
VALIDATION = HERE / "oss-adx-risk-validation-sanitized.zip"
ROUND1 = HERE / "oss-round1-training-sanitized.zip"
ROBUSTNESS = HERE / "oss-robustness-summary.json"
SENSITIVITY_TRAIN = HERE / "oss-sensitivity-training-sanitized.zip"
SENSITIVITY_VALIDATION = HERE / "oss-sensitivity-validation-sanitized.zip"
BASE = "VolatilityBreakoutRiskCapMatchedBaseline"
CAND = "VolatilityBreakoutRiskCapAdxScaledRisk"
NOOP = "VolatilityBreakoutRiskCapPairDailyNoOp"


def load(path: Path) -> dict:
    with zipfile.ZipFile(path) as archive:
        member = next(
            name
            for name in archive.namelist()
            if name.endswith(".json") and not name.endswith(".meta.json")
        )
        return json.loads(archive.read(member))["strategy"]


def metrics(item: dict) -> dict:
    start = datetime.fromisoformat(item["backtest_start"])
    end = datetime.fromisoformat(item["backtest_end"])
    days = (end - start).total_seconds() / 86400
    return {
        "timerange": item["timerange"],
        "days": days,
        "trades": item["total_trades"],
        "return": item["profit_total"],
        "monthlyEquivalent": math.pow(1 + item["profit_total"], 30 / days) - 1,
        "profitFactor": item["profit_factor"],
        "maxDrawdown": item["max_drawdown_account"],
        "winRate": item["winrate"],
        "averageStake": item["avg_stake_amount"],
        "longs": item["trade_count_long"],
        "shorts": item["trade_count_short"],
        "timeframe": item["timeframe"],
        "protections": item["enable_protections"],
    }


def normalized_paths(item: dict) -> list[tuple]:
    keys = (
        "pair",
        "open_timestamp",
        "close_timestamp",
        "open_rate",
        "close_rate",
        "is_short",
        "exit_reason",
    )
    return [tuple(trade[key] for key in keys) for trade in item["trades"]]


def equity(item: dict) -> pd.Series:
    trades = pd.DataFrame(item["trades"])
    pnl = trades.groupby(pd.to_datetime(trades["close_date"], utc=True))["profit_abs"].sum()
    start = pd.Timestamp(item["backtest_start"], tz="UTC")
    increments = pd.concat([pd.Series({start: item["starting_balance"]}), pnl]).sort_index()
    return increments.cumsum()


train = load(TRAIN)
validation = load(VALIDATION)
round1 = load(ROUND1)
robustness = json.loads(ROBUSTNESS.read_text())
sensitivity_train = load(SENSITIVITY_TRAIN)
sensitivity_validation = load(SENSITIVITY_VALIDATION)

train_noop_equal = normalized_paths(train[BASE]) == normalized_paths(train[NOOP])
train_paths_equal = normalized_paths(train[BASE]) == normalized_paths(train[CAND])
validation_paths_equal = normalized_paths(validation[BASE]) == normalized_paths(validation[CAND])

stake_ratios = {}
for label, loaded in (("training", train), ("validation", validation)):
    base_stakes = np.array([trade["stake_amount"] for trade in loaded[BASE]["trades"]])
    cand_stakes = np.array([trade["stake_amount"] for trade in loaded[CAND]["trades"]])
    ratios = cand_stakes / base_stakes
    stake_ratios[label] = {
        "min": float(ratios.min()),
        "median": float(np.median(ratios)),
        "mean": float(ratios.mean()),
        "max": float(ratios.max()),
        "belowBaselineCount": int((ratios < 1).sum()),
        "aboveBaselineCount": int((ratios > 1).sum()),
    }

summary = {
    "protocol": {
        "universe": "active-11 (VANRY inactive and ignored by exchange metadata)",
        "startingBalance": 1000,
        "leverage": "2x isolated",
        "maxOpenTrades": 4,
        "feePerSide": 0.00035,
        "timeframe": "5m",
        "protections": True,
        "cache": "none",
    },
    "candidate": {
        "id": "ADX-SCALED-RISK-V1",
        "formula": "multiplier = clip(0.75 + (ADX - 25) / 40, 0.75, 1.25)",
        "changedFactor": "position risk only; entry, exit and portfolio gates unchanged",
    },
    "training": {"baseline": metrics(train[BASE]), "candidate": metrics(train[CAND])},
    "validation": {
        "baseline": metrics(validation[BASE]),
        "candidate": metrics(validation[CAND]),
    },
    "isolation": {
        "pairDailyNoOpExactPaths": train_noop_equal,
        "trainingCandidateExactPaths": train_paths_equal,
        "validationCandidateExactPaths": validation_paths_equal,
        "stakeRatios": stake_ratios,
    },
    "rejectedTrainingCandidates": {
        name: metrics(item)
        for name, item in {**round1, **train}.items()
        if name not in {BASE, CAND, NOOP}
    },
    "robustness": robustness,
    "parameterSensitivity": {
        "purpose": "neighborhood stability only; alternatives are not eligible to replace frozen V1",
        "training": {name: metrics(item) for name, item in sensitivity_train.items()},
        "validation": {name: metrics(item) for name, item in sensitivity_validation.items()},
    },
    "biasAudit": {
        "lookahead": "No bias; 20 signals, 0 biased entries, 0 biased exits",
        "recursive": "ATR/ATR baseline/ADX effectively 0%; EMA200 max 0.035% at 210 startup",
    },
}
(HERE / "deterministic-summary.json").write_text(
    json.dumps(summary, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
)

fig, axes = plt.subplots(2, 2, figsize=(15, 10), constrained_layout=True)
for axis, label, loaded in (
    (axes[0, 0], "Training", train),
    (axes[0, 1], "Held-out validation", validation),
):
    for name, color, display in (
        (BASE, "#2563EB", "Current D48 RiskCap"),
        (CAND, "#D97706", "ADX-scaled risk"),
    ):
        curve = equity(loaded[name])
        axis.step(curve.index, curve.values, where="post", color=color, linewidth=1.7, label=display)
    axis.set_title(label, loc="left", fontweight="bold")
    axis.set_ylabel("Realized wallet (USDT)")
    axis.grid(alpha=0.2)
    axis.legend(frameon=False)
    locator = mdates.AutoDateLocator(minticks=4, maxticks=7)
    axis.xaxis.set_major_locator(locator)
    axis.xaxis.set_major_formatter(mdates.ConciseDateFormatter(locator))

width = 0.36
axes[1, 0].axis("off")
table_rows = [
    ["Training return", f"{train[BASE]['profit_total']:.2%}", f"{train[CAND]['profit_total']:.2%}"],
    ["Validation return", f"{validation[BASE]['profit_total']:.2%}", f"{validation[CAND]['profit_total']:.2%}"],
    ["Training PF", f"{train[BASE]['profit_factor']:.3f}", f"{train[CAND]['profit_factor']:.3f}"],
    ["Validation PF", f"{validation[BASE]['profit_factor']:.3f}", f"{validation[CAND]['profit_factor']:.3f}"],
    ["Training max DD", f"{train[BASE]['max_drawdown_account']:.2%}", f"{train[CAND]['max_drawdown_account']:.2%}"],
    ["Validation max DD", f"{validation[BASE]['max_drawdown_account']:.2%}", f"{validation[CAND]['max_drawdown_account']:.2%}"],
    ["Training avg stake", f"{train[BASE]['avg_stake_amount']:.2f}", f"{train[CAND]['avg_stake_amount']:.2f}"],
    ["Validation avg stake", f"{validation[BASE]['avg_stake_amount']:.2f}", f"{validation[CAND]['avg_stake_amount']:.2f}"],
]
table = axes[1, 0].table(
    cellText=table_rows,
    colLabels=["Metric", "Current", "Candidate"],
    colColours=["#E2E8F0", "#DBEAFE", "#FFEDD5"],
    cellLoc="center",
    loc="center",
    colWidths=[0.46, 0.25, 0.25],
)
table.auto_set_font_size(False)
table.set_fontsize(10)
table.scale(1, 1.65)
axes[1, 0].set_title("Aligned performance and risk metrics", loc="left", fontweight="bold")

stress = [row for row in robustness if row["case"].startswith("stress/")]
stress_labels = [row["case"].split("/", 1)[1] for row in stress]
return_delta = [row["delta"]["return"] * 100 for row in stress]
dd_delta = [row["delta"]["maxDrawdown"] * 100 for row in stress]
sx = np.arange(len(stress_labels))
axes[1, 1].bar(sx - width / 2, return_delta, width, color="#D97706", label="Return delta (pp)")
axes[1, 1].bar(sx + width / 2, dd_delta, width, color="#0F766E", label="DD delta (pp)")
axes[1, 1].axhline(0, color="#64748B", linewidth=0.8)
axes[1, 1].set_xticks(sx, stress_labels, rotation=25, ha="right")
axes[1, 1].set_title("Eight stress windows · candidate minus current", loc="left", fontweight="bold")
axes[1, 1].grid(axis="y", alpha=0.2)
axes[1, 1].legend(frameon=False)

fig.suptitle("D48 RiskCap → ADX-scaled risk · deterministic offline validation", fontsize=16, fontweight="bold")
fig.savefig(HERE / "strategy-comparison.png", dpi=180, facecolor="white")
