#!/usr/bin/env python3
"""Tag exported backtest trades with causal short-horizon RV at entry.

The audit reads an existing Freqtrade backtest zip and the matching local 5m
futures candles.  It does not rerun the strategy or mutate its results.
"""

from __future__ import annotations

import argparse
import json
import math
import zipfile
from pathlib import Path

import numpy as np
import pandas as pd


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("backtest_zip", type=Path)
    parser.add_argument("--strategy", default="VolatilityBreakoutRiskCap")
    parser.add_argument(
        "--data-dir",
        type=Path,
        default=Path("user_data/data/binance/futures"),
    )
    parser.add_argument("--rolling-days", type=int, default=20)
    return parser.parse_args()


def load_backtest(path: Path, strategy: str) -> tuple[dict, pd.DataFrame]:
    with zipfile.ZipFile(path) as archive:
        result_name = next(
            name
            for name in archive.namelist()
            if name.endswith(".json") and not name.endswith("_config.json")
        )
        result = json.loads(archive.read(result_name))
    summary = result["strategy"][strategy]
    trades = pd.DataFrame(summary["trades"])
    trades["open_date"] = pd.to_datetime(trades["open_date"], utc=True).astype(
        "datetime64[ns, UTC]"
    )
    return summary, trades


def candle_path(data_dir: Path, pair: str) -> Path:
    stem = pair.replace("/", "_").replace(":", "_")
    return data_dir / f"{stem}-5m-futures.feather"


def add_rv_columns(candles: pd.DataFrame, rolling_days: int) -> pd.DataFrame:
    candles = candles.sort_values("date").copy()
    candles["date"] = pd.to_datetime(candles["date"], utc=True).astype(
        "datetime64[ns, UTC]"
    )

    # Exact implementation used by btc-yield-enhancer, before its 0.5%-5% clip.
    open_close_return = (candles["close"] - candles["open"]) / candles["open"]
    candles["rv_grid_raw"] = (
        open_close_return.pow(2).rolling(12, min_periods=12).mean().pow(0.5)
        * math.sqrt(24)
    )
    candles["rv_grid_clipped"] = candles["rv_grid_raw"].clip(0.005, 0.05)

    # Conventional completed-close realized volatility.  Annualization is not
    # needed because the experiment uses a rolling within-pair percentile.
    close_return = np.log(candles["close"] / candles["close"].shift(1))
    candles["rv_fast"] = close_return.pow(2).rolling(12, min_periods=12).sum().pow(0.5)

    window = rolling_days * 24 * 12
    minimum = max(12 * 24, window // 4)
    candles["rv_percentile"] = candles["rv_fast"].rolling(
        window, min_periods=minimum
    ).rank(pct=True)

    previous_close = candles["close"].shift(1)
    true_range = pd.concat(
        [
            candles["high"] - candles["low"],
            (candles["high"] - previous_close).abs(),
            (candles["low"] - previous_close).abs(),
        ],
        axis=1,
    ).max(axis=1)
    candles["atr_pct"] = true_range.rolling(14, min_periods=14).mean() / candles["close"]
    return candles[
        [
            "date",
            "rv_grid_raw",
            "rv_grid_clipped",
            "rv_fast",
            "rv_percentile",
            "atr_pct",
        ]
    ]


def tag_trades(trades: pd.DataFrame, data_dir: Path, rolling_days: int) -> pd.DataFrame:
    tagged: list[pd.DataFrame] = []
    for pair, pair_trades in trades.groupby("pair", sort=True):
        path = candle_path(data_dir, pair)
        candles = add_rv_columns(pd.read_feather(path), rolling_days)

        # Freqtrade fills at the next candle open.  Use only the last fully
        # completed candle before open_date to keep the feature causal.
        eligible = pair_trades.sort_values("open_date").copy()
        eligible["signal_close"] = eligible["open_date"] - pd.Timedelta(minutes=5)
        merged = pd.merge_asof(
            eligible,
            candles,
            left_on="signal_close",
            right_on="date",
            direction="backward",
            tolerance=pd.Timedelta(minutes=5),
        )
        tagged.append(merged)
    result = pd.concat(tagged, ignore_index=True)

    long_mfe = result["max_rate"] / result["open_rate"] - 1
    long_mae = result["min_rate"] / result["open_rate"] - 1
    short_mfe = 1 - result["min_rate"] / result["open_rate"]
    short_mae = 1 - result["max_rate"] / result["open_rate"]
    result["mfe"] = np.where(result["is_short"], short_mfe, long_mfe)
    result["mae"] = np.where(result["is_short"], short_mae, long_mae)
    return result


def profit_factor(frame: pd.DataFrame) -> float:
    gross_profit = frame.loc[frame["profit_abs"] > 0, "profit_abs"].sum()
    gross_loss = -frame.loc[frame["profit_abs"] < 0, "profit_abs"].sum()
    return float("inf") if gross_loss == 0 else gross_profit / gross_loss


def metrics(frame: pd.DataFrame) -> dict:
    return {
        "trades": int(len(frame)),
        "profit_abs": round(float(frame["profit_abs"].sum()), 4),
        "profit_factor": round(profit_factor(frame), 3),
        "win_rate_pct": round(float((frame["profit_abs"] > 0).mean() * 100), 2),
        "mean_mfe_pct": round(float(frame["mfe"].mean() * 100), 3),
        "mean_mae_pct": round(float(frame["mae"].mean() * 100), 3),
        "stoploss_pct": round(
            float(frame["exit_reason"].str.contains("stop", case=False).mean() * 100), 2
        ),
        "long": int((~frame["is_short"]).sum()),
        "short": int(frame["is_short"].sum()),
    }


def summarize(summary: dict, tagged: pd.DataFrame) -> dict:
    usable = tagged.dropna(subset=["rv_percentile"]).copy()
    usable["rv_bucket"] = pd.cut(
        usable["rv_percentile"],
        bins=[0, 0.2, 0.4, 0.6, 0.8, 1.0],
        labels=["0-20", "20-40", "40-60", "60-80", "80-100"],
        include_lowest=True,
    )
    buckets = {
        str(bucket): metrics(frame)
        for bucket, frame in usable.groupby("rv_bucket", observed=True)
    }
    filters = {
        "all_usable": metrics(usable),
        "exclude_bottom_20_hindsight": metrics(usable[usable["rv_percentile"] > 0.2]),
        "exclude_top_10_hindsight": metrics(usable[usable["rv_percentile"] <= 0.9]),
        "middle_20_to_90_hindsight": metrics(
            usable[
                (usable["rv_percentile"] > 0.2)
                & (usable["rv_percentile"] <= 0.9)
            ]
        ),
    }
    return {
        "backtest": {
            "strategy": summary["strategy_name"],
            "start": summary["backtest_start"],
            "end": summary["backtest_end"],
            "reported_trades": summary["total_trades"],
            "reported_profit_abs": summary["profit_total_abs"],
            "reported_profit_factor": summary["profit_factor"],
        },
        "coverage": {
            "tagged": int(len(tagged)),
            "usable_percentile": int(len(usable)),
            "grid_rv_floor_pct": round(
                float((tagged["rv_grid_raw"] <= 0.005).mean() * 100), 2
            ),
            "grid_rv_ceiling_pct": round(
                float((tagged["rv_grid_raw"] >= 0.05).mean() * 100), 2
            ),
            "rv_atr_correlation": round(
                float(tagged[["rv_fast", "atr_pct"]].corr().iloc[0, 1]), 3
            ),
        },
        "rv_percentile_buckets": buckets,
        "counterfactual_trade_subset_only": filters,
        "warning": (
            "Filtered subsets are diagnostic only. Removing entries changes portfolio "
            "slot occupancy, so promotion decisions require a full strategy backtest."
        ),
    }


def main() -> None:
    args = parse_args()
    summary, trades = load_backtest(args.backtest_zip, args.strategy)
    tagged = tag_trades(trades, args.data_dir, args.rolling_days)
    print(json.dumps(summarize(summary, tagged), indent=2, ensure_ascii=False))


if __name__ == "__main__":
    main()
