#!/usr/bin/env python3
from __future__ import annotations

import gzip
import json
import math
import time
import urllib.parse
import urllib.request
from collections import defaultdict
from datetime import datetime, timedelta, timezone
from pathlib import Path

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

ROOT = Path(__file__).resolve().parent
EVIDENCE = ROOT / "evidence.json"
CACHE = ROOT / "binance-usdtm-5m.json.gz"
METRICS = ROOT / "path-metrics.json"
CHART = ROOT / "layered-cases.png"


def dt(value: str) -> datetime:
    parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
    if parsed.tzinfo is None:
        parsed = parsed.replace(tzinfo=timezone.utc)
    return parsed.astimezone(timezone.utc)


def symbol(pair: str) -> str:
    return pair.split(":")[0].replace("/", "")


def fetch_pair(pair: str, start: datetime, end: datetime) -> list[dict]:
    rows = []
    cursor = int(start.timestamp() * 1000)
    end_ms = int(end.timestamp() * 1000)
    while cursor < end_ms:
        query = urllib.parse.urlencode({
            "symbol": symbol(pair), "interval": "5m", "startTime": cursor,
            "endTime": end_ms, "limit": 1500,
        })
        url = "https://fapi.binance.com/fapi/v1/klines?" + query
        with urllib.request.urlopen(url, timeout=30) as response:
            batch = json.loads(response.read())
        if not batch:
            break
        for k in batch:
            rows.append({
                "time": datetime.fromtimestamp(k[0] / 1000, tz=timezone.utc).isoformat(),
                "open": float(k[1]), "high": float(k[2]), "low": float(k[3]),
                "close": float(k[4]), "volume": float(k[5]),
            })
        next_cursor = int(batch[-1][0]) + 300_000
        if next_cursor <= cursor:
            break
        cursor = next_cursor
        time.sleep(0.08)
    return rows


def load_or_fetch(evidence: dict) -> dict[str, list[dict]]:
    if CACHE.exists():
        with gzip.open(CACHE, "rt") as handle:
            return json.load(handle)
    pairs = sorted({t["pair"] for t in evidence["trades"]}
                   | {x["pair"] for x in evidence["cancelled_entries"]}
                   | {x["pair"] for x in evidence["block_clusters"]})
    start = dt(evidence["window"]["start"]) - timedelta(hours=12)
    end = dt(evidence["window"]["end"]) + timedelta(hours=24)
    data = {pair: fetch_pair(pair, start, end) for pair in pairs}
    with gzip.open(CACHE, "wt") as handle:
        json.dump(data, handle, separators=(",", ":"))
    return data


def frame(rows: list[dict]) -> pd.DataFrame:
    out = pd.DataFrame(rows)
    out["time"] = pd.to_datetime(out["time"], utc=True)
    return out.set_index("time").sort_index()


def path_metric(df: pd.DataFrame, when: datetime, side: str, hours: int) -> dict | None:
    start = pd.Timestamp(when).floor("5min")
    sub = df[(df.index >= start) & (df.index < start + pd.Timedelta(hours=hours))]
    if sub.empty:
        return None
    entry = float(sub.iloc[0].open)
    high, low, last = float(sub.high.max()), float(sub.low.min()), float(sub.iloc[-1].close)
    if side == "short":
        mfe, mae, end_ret = (entry - low) / entry, (high - entry) / entry, (entry - last) / entry
    else:
        mfe, mae, end_ret = (high - entry) / entry, (entry - low) / entry, (last - entry) / entry
    return {"entry": entry, "mfe": mfe, "mae": mae, "end_ret": end_ret, "bars": len(sub),
            "stop_hit": mae >= 0.025, "trail_armed": mfe >= 0.025}


def summarize(rows: list[dict], prefix: str) -> dict:
    valid = [r for r in rows if r]
    out = {"n": len(valid)}
    for key in ("mfe", "mae", "end_ret"):
        vals = np.array([r[key] for r in valid], dtype=float)
        out[key] = ({"mean": float(vals.mean()), "median": float(np.median(vals)),
                     "p25": float(np.percentile(vals, 25)), "p75": float(np.percentile(vals, 75))}
                    if len(vals) else {})
    out["stop_hit_rate"] = float(np.mean([r["stop_hit"] for r in valid])) if valid else None
    out["trail_arm_rate"] = float(np.mean([r["trail_armed"] for r in valid])) if valid else None
    out["label"] = prefix
    return out


def dedupe_24h(events: list[dict]) -> list[dict]:
    kept, last = [], {}
    for event in sorted(events, key=lambda x: x["time"]):
        key = (event["pair"], event["side"])
        when = dt(event["time"])
        if key not in last or when - last[key] >= timedelta(hours=24):
            kept.append(event)
            last[key] = when
    return kept


def candle_plot(ax, df: pd.DataFrame, title: str):
    if df.empty:
        ax.text(0.5, 0.5, "No candle data", ha="center", va="center")
        return
    xs = mdates.date2num(df.index.to_pydatetime())
    width = (5 / 1440) * 0.72
    for x, row in zip(xs, df.itertuples()):
        up = row.close >= row.open
        color = "#15803D" if up else "#B91C1C"
        ax.vlines(x, row.low, row.high, color=color, linewidth=0.55, alpha=0.8)
        bottom, height = min(row.open, row.close), abs(row.close - row.open)
        if height == 0:
            ax.hlines(row.close, x - width / 2, x + width / 2, color=color, linewidth=0.8)
        else:
            ax.add_patch(Rectangle((x - width / 2, bottom), width, height,
                                   facecolor=color, edgecolor=color, linewidth=0.3, alpha=0.72))
    ax.set_title(title, loc="left", fontsize=10, fontweight="bold")
    ax.grid(color="#E2E8F0", linewidth=0.5, alpha=0.8)
    ax.xaxis.set_major_formatter(mdates.DateFormatter("%m-%d\n%H:%M", tz=timezone.utc))
    ax.tick_params(labelsize=7)


def resample_15m(df: pd.DataFrame) -> pd.DataFrame:
    return df.resample("15min", label="left", closed="left").agg(
        {"open": "first", "high": "max", "low": "min", "close": "last", "volume": "sum"}
    ).dropna()


def select_trade(evidence: dict, trade_id: int) -> dict:
    return next(t for t in evidence["trades"] if t["id"] == trade_id)


def overlay_trade(ax, trade: dict, include_exit=True):
    opened = dt(trade["open_date"])
    ax.scatter(opened, trade["open_rate"], marker="v" if trade["is_short"] else "^",
               s=58, color="#2563EB", edgecolor="white", linewidth=0.6, zorder=5,
               label="strategy signal / entry")
    ax.scatter(opened, trade["open_rate"], marker="o", s=20, color="#111827", zorder=6, label="fill")
    ax.axhline(trade["initial_stop_loss"], color="#DC2626", linewidth=1.0, label="initial stop")
    if include_exit and trade.get("close_date"):
        ax.scatter(dt(trade["close_date"]), trade["close_rate"], marker="D", s=44,
                   color="#0F766E", edgecolor="white", linewidth=0.5, zorder=6, label="exit")


def render_chart(evidence: dict, frames: dict[str, pd.DataFrame]):
    cases = [
        ("BTC open short #50", select_trade(evidence, 50)),
        ("BTC hard-stop loss #42", select_trade(evidence, 42)),
        ("XRP trailing-stop winner #44", select_trade(evidence, 44)),
    ]
    fig, axes = plt.subplots(4, 2, figsize=(16, 17), constrained_layout=True)
    for row, (label, trade) in enumerate(cases):
        opened = dt(trade["open_date"]); closed = dt(trade["close_date"]) if trade.get("close_date") else dt(evidence["window"]["end"])
        full_start, full_end = opened - timedelta(hours=4), closed + timedelta(hours=4)
        entry_end = min(full_end, opened + timedelta(hours=16))
        df = frames[trade["pair"]]
        left = df[(df.index >= pd.Timestamp(full_start)) & (df.index <= pd.Timestamp(entry_end))]
        right = resample_15m(df[(df.index >= pd.Timestamp(full_start)) & (df.index <= pd.Timestamp(full_end))])
        candle_plot(axes[row, 0], left, label + " · 5m entry context")
        candle_plot(axes[row, 1], right, label + " · 15m full path")
        overlay_trade(axes[row, 0], trade, include_exit=closed <= entry_end)
        overlay_trade(axes[row, 1], trade, include_exit=True)
        if row == 0:
            axes[row, 0].legend(fontsize=7, loc="best")

    miss = evidence["cancelled_entries"][-1]
    pair = miss["pair"]; submitted = dt(miss["order_time"]); cancelled_at = dt(miss["time"])
    df = frames[pair]
    left = df[(df.index >= pd.Timestamp(submitted - timedelta(hours=4))) & (df.index <= pd.Timestamp(submitted + timedelta(hours=12)))]
    right = resample_15m(df[(df.index >= pd.Timestamp(submitted - timedelta(hours=4))) & (df.index <= pd.Timestamp(submitted + timedelta(hours=24)))])
    candle_plot(axes[3, 0], left, "BNB unfilled entry · 5m")
    candle_plot(axes[3, 1], right, "BNB unfilled entry · 15m / 24h hindsight path")
    for ax in axes[3]:
        ax.scatter(submitted, miss["intended_rate"], marker="v", s=58, color="#2563EB", zorder=5, label="strategy signal")
        ax.hlines(miss["intended_rate"], submitted, cancelled_at,
                  color="#D97706", linestyle="--", linewidth=1.6, label="passive limit pending")
        ax.scatter(cancelled_at, miss["intended_rate"], marker="x", s=55, color="#D97706", zorder=6, label="zero-fill cancel")
    axes[3, 0].legend(fontsize=7, loc="best")
    fig.suptitle("Freqtrade two-week order-chain sample (Binance USDT-M path; production ledger events)",
                 fontsize=14, fontweight="bold")
    fig.savefig(CHART, dpi=170, facecolor="white")
    plt.close(fig)


def main():
    evidence = json.loads(EVIDENCE.read_text())
    raw = load_or_fetch(evidence)
    frames = {pair: frame(rows) for pair, rows in raw.items()}

    blocked = evidence["block_clusters"]
    blocked_dedup = dedupe_24h(blocked)
    filled_events = [{"time": t["open_date"], "pair": t["pair"],
                      "side": "short" if t["is_short"] else "long"}
                     for t in evidence["trades"] if dt(t["open_date"]) >= dt(evidence["window"]["start"])]
    cancelled = evidence["cancelled_entries"]
    metrics = {
        "provenance": "Binance USDT-M public 5m klines; price-path/hindsight labels only, not authentic Freqtrade indicators",
        "coverage": {p: {"bars": len(f), "start": f.index.min().isoformat() if len(f) else None,
                         "end": f.index.max().isoformat() if len(f) else None} for p, f in frames.items()},
        "blocked": {}, "blocked_dedup_24h": {}, "filled": {}, "cancelled": {},
        "shadow": {}, "post_exit": {},
    }
    for hours in (6, 24):
        metrics["blocked"][f"{hours}h"] = summarize(
            [path_metric(frames[e["pair"]], dt(e["time"]), e["side"], hours) for e in blocked], "hindsight_annotation")
        metrics["blocked_dedup_24h"][f"{hours}h"] = summarize(
            [path_metric(frames[e["pair"]], dt(e["time"]), e["side"], hours) for e in blocked_dedup], "hindsight_annotation")
        metrics["filled"][f"{hours}h"] = summarize(
            [path_metric(frames[e["pair"]], dt(e["time"]), e["side"], hours) for e in filled_events], "outcome_path")
        metrics["cancelled"][f"{hours}h"] = summarize(
            [path_metric(frames[e["pair"]], dt(e.get("order_time", e["time"])), e["side"], hours) for e in cancelled], "hindsight_annotation")

    shadow_events = [e for e in blocked if e.get("shadow")]
    for eligible in (True, False):
        group = [e for e in shadow_events if bool(e["shadow"].get("eligible")) is eligible]
        metrics["shadow"]["eligible" if eligible else "ineligible"] = {
            "events": len(group),
            "24h": summarize([path_metric(frames[e["pair"]], dt(e["time"]), e["side"], 24) for e in group], "hindsight_annotation"),
        }

    for hours in (6, 24):
        rows = []
        for trade in evidence["trades"]:
            if not trade.get("close_date"):
                continue
            side = "short" if trade["is_short"] else "long"
            rows.append(path_metric(frames[trade["pair"]], dt(trade["close_date"]), side, hours))
        metrics["post_exit"][f"{hours}h"] = summarize(rows, "hindsight_annotation")

    METRICS.write_text(json.dumps(metrics, ensure_ascii=False, indent=2))
    render_chart(evidence, frames)
    print(json.dumps(metrics, ensure_ascii=False, indent=2))
    print(CHART)


if __name__ == "__main__":
    main()
