#!/usr/bin/env python3
"""Hindsight opportunity cost of same_side_limit blocks.

Labels are hindsight_annotation only — not causal missed signals for promotion.

Candle sources (remote only, no public exchange pull in this script):
1. Freqtrade feather: user_data/data/binance/futures/*-5m-futures.feather
2. Running bot1 API pair_candles (fills the gap after feather end)
"""
from __future__ import annotations

import base64
import json
import re
import urllib.parse
import urllib.request
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from pathlib import Path

import numpy as np
import pandas as pd

LOG = Path("/data/freqtrade/user_data/logs/freqtrade.log")
DB = Path("/data/freqtrade/user_data/tradesv3-v5.sqlite")
CONFIG = Path("/data/freqtrade/user_data/config.json")
DATA_DIR = Path("/data/freqtrade/user_data/data/binance/futures")
OUT = Path("/data/freqtrade/user_data/backtest_results/same-side-opp-cost")
OUT.mkdir(parents=True, exist_ok=True)

LINE_RE = re.compile(
    r"^(?P<ts>\d{4}-\d{2}-\d{2} \d{2}:\d{2}:\d{2}),\d+ .* "
    r"entry_blocked reason=(?P<reason>\S+) pair=(?P<pair>\S+) "
    r"side=(?P<side>\S+) occupied=(?P<occupied>\S+)"
)

CLUSTER_GAP_S = 20 * 60
HORIZONS_H = (6, 24)
# Futures 2x: hard stoploss -5% equity ≈ 2.5% price; trail arm +5% profit ≈ 2.5% price.
HARD_STOP_PRICE = 0.025
TRAIL_ARM_PRICE = 0.025


def pair_to_feather(pair: str) -> Path:
    # BTC/USDT:USDT -> BTC_USDT_USDT-5m-futures.feather
    base_quote = pair.split(":")[0]  # BTC/USDT
    base, quote = base_quote.split("/")
    return DATA_DIR / f"{base}_{quote}_{quote}-5m-futures.feather"


def load_feather_klines(pair: str) -> list[dict]:
    path = pair_to_feather(pair)
    if not path.exists():
        print("missing feather", path)
        return []
    df = pd.read_feather(path)
    df["date"] = pd.to_datetime(df["date"], utc=True)
    out = []
    for r in df.itertuples(index=False):
        out.append(
            {
                "ts": r.date.to_pydatetime(),
                "open": float(r.open),
                "high": float(r.high),
                "low": float(r.low),
                "close": float(r.close),
            }
        )
    return out


def bot_auth_header() -> dict[str, str]:
    cfg = json.loads(CONFIG.read_text())
    api = cfg["api_server"]
    token = base64.b64encode(
        f"{api['username']}:{api['password']}".encode()
    ).decode()
    return {"Authorization": f"Basic {token}"}


def load_bot_klines(pair: str, limit: int = 1500) -> list[dict]:
    """Recent candles from live bot1 analyzed cache."""
    q = urllib.parse.urlencode(
        {"pair": pair, "timeframe": "5m", "limit": str(limit)}
    )
    url = f"http://127.0.0.1:8080/api/v1/pair_candles?{q}"
    req = urllib.request.Request(url, headers=bot_auth_header())
    try:
        data = json.loads(urllib.request.urlopen(req, timeout=30).read())
    except Exception as exc:  # noqa: BLE001
        print("bot candles fail", pair, exc)
        return []
    cols = data["columns"]
    idx = {c: i for i, c in enumerate(cols)}
    out = []
    for row in data["data"]:
        ds = row[idx["date"]]
        if isinstance(ds, (int, float)):
            ts = datetime.fromtimestamp(
                ds / 1000 if ds > 1e12 else ds, tz=timezone.utc
            )
        else:
            ts = datetime.fromisoformat(str(ds).replace("Z", "+00:00"))
            if ts.tzinfo is None:
                ts = ts.replace(tzinfo=timezone.utc)
        out.append(
            {
                "ts": ts,
                "open": float(row[idx["open"]]),
                "high": float(row[idx["high"]]),
                "low": float(row[idx["low"]]),
                "close": float(row[idx["close"]]),
            }
        )
    return out


def merge_klines(a: list[dict], b: list[dict]) -> list[dict]:
    by_ts: dict[datetime, dict] = {}
    for row in a + b:
        by_ts[row["ts"]] = row
    return [by_ts[k] for k in sorted(by_ts)]


def load_pair_klines(pair: str) -> list[dict]:
    feather = load_feather_klines(pair)
    bot = load_bot_klines(pair, limit=1500)
    merged = merge_klines(feather, bot)
    if merged:
        print(
            pair,
            "feather",
            len(feather),
            "bot",
            len(bot),
            "merged",
            len(merged),
            merged[0]["ts"],
            "->",
            merged[-1]["ts"],
        )
    else:
        print(pair, "NO DATA")
    return merged


def path_metrics(klines, t0: datetime, side: str, hours: int):
    t1 = t0 + timedelta(hours=hours)
    window = [k for k in klines if t0 <= k["ts"] < t1]
    if not window:
        after = [k for k in klines if k["ts"] >= t0]
        if not after:
            return None
        window = [k for k in after if k["ts"] < t1]
        if not window:
            return None
    entry = window[0]["open"]
    hi = max(k["high"] for k in window)
    lo = min(k["low"] for k in window)
    last = window[-1]["close"]
    if side == "short":
        mfe = (entry - lo) / entry
        mae = (hi - entry) / entry
        end_ret = (entry - last) / entry
    else:
        mfe = (hi - entry) / entry
        mae = (entry - lo) / entry
        end_ret = (last - entry) / entry
    return {
        "entry": entry,
        "mfe": mfe,
        "mae": mae,
        "end_ret": end_ret,
        "bars": len(window),
        "stop_hit": mae >= HARD_STOP_PRICE,
        "trail_armed": mfe >= TRAIL_ARM_PRICE,
        "hi": hi,
        "lo": lo,
    }


@dataclass
class Event:
    ts: datetime
    pair: str
    side: str
    occupied: str
    reason: str


def load_events() -> list[Event]:
    events: list[Event] = []
    for line in LOG.read_text(errors="replace").splitlines():
        m = LINE_RE.search(line)
        if not m or m.group("reason") != "same_side_limit":
            continue
        ts = datetime.strptime(m.group("ts"), "%Y-%m-%d %H:%M:%S").replace(
            tzinfo=timezone.utc
        )
        events.append(
            Event(
                ts=ts,
                pair=m.group("pair"),
                side=m.group("side"),
                occupied=m.group("occupied"),
                reason=m.group("reason"),
            )
        )
    return events


def cluster_events(events: list[Event]) -> list[dict]:
    by_key: dict[tuple[str, str], list[Event]] = defaultdict(list)
    for e in events:
        by_key[(e.pair, e.side)].append(e)
    clusters: list[list[Event]] = []
    for evs in by_key.values():
        evs.sort(key=lambda x: x.ts)
        cur = [evs[0]]
        for e in evs[1:]:
            if (e.ts - cur[-1].ts).total_seconds() <= CLUSTER_GAP_S:
                cur.append(e)
            else:
                clusters.append(cur)
                cur = [e]
        clusters.append(cur)
    opps = []
    for c in clusters:
        first = c[0]
        opps.append(
            {
                "ts": first.ts,
                "pair": first.pair,
                "side": first.side,
                "occupied": first.occupied,
                "n_logs": len(c),
                "span_min": (c[-1].ts - first.ts).total_seconds() / 60,
            }
        )
    opps.sort(key=lambda x: x["ts"])
    return opps


def load_trades() -> list[dict]:
    import sqlite3

    con = sqlite3.connect(f"file:{DB}?mode=ro", uri=True)
    con.row_factory = sqlite3.Row
    trades = [
        dict(r)
        for r in con.execute(
            "select id,pair,is_short,is_open,open_date,close_date,open_rate,"
            "close_rate,close_profit,close_profit_abs,exit_reason,stake_amount "
            "from trades order by id"
        )
    ]
    for t in trades:
        t["open_dt"] = datetime.fromisoformat(t["open_date"]).replace(tzinfo=timezone.utc)
        t["close_dt"] = (
            datetime.fromisoformat(t["close_date"]).replace(tzinfo=timezone.utc)
            if t["close_date"]
            else None
        )
        t["side"] = "short" if t["is_short"] else "long"
    return trades


def summarize(vals) -> dict:
    if not vals:
        return {}
    a = np.array(vals, dtype=float)
    return {
        "n": int(len(a)),
        "mean": float(a.mean()),
        "median": float(np.median(a)),
        "p25": float(np.percentile(a, 25)),
        "p75": float(np.percentile(a, 75)),
        "pos_rate": float((a > 0).mean()),
        "sum": float(a.sum()),
    }


def main() -> None:
    events = load_events()
    opps = cluster_events(events)
    trades = load_trades()
    print(f"raw_block_logs={len(events)} clusters/opps={len(opps)}")
    if not opps:
        return

    pairs_needed = sorted(
        {o["pair"] for o in opps} | {t["pair"] for t in trades}
    )
    print("pairs", pairs_needed)
    kcache: dict[str, list] = {}
    for pair in pairs_needed:
        kcache[pair] = load_pair_klines(pair)

    rows = []
    for o in opps:
        kl = kcache[o["pair"]]
        rec = {
            "ts": o["ts"].isoformat(),
            "pair": o["pair"],
            "side": o["side"],
            "occupied": o["occupied"],
            "n_logs": o["n_logs"],
            "span_min": o["span_min"],
            "label": "hindsight_annotation",
        }
        for h in HORIZONS_H:
            m = path_metrics(kl, o["ts"], o["side"], h)
            if m is None:
                rec[f"h{h}_mfe"] = None
                rec[f"h{h}_mae"] = None
                rec[f"h{h}_end"] = None
                rec[f"h{h}_stop_hit"] = None
                rec[f"h{h}_trail_armed"] = None
            else:
                rec[f"h{h}_mfe"] = m["mfe"]
                rec[f"h{h}_mae"] = m["mae"]
                rec[f"h{h}_end"] = m["end_ret"]
                rec[f"h{h}_stop_hit"] = m["stop_hit"]
                rec[f"h{h}_trail_armed"] = m["trail_armed"]
                rec[f"h{h}_entry"] = m["entry"]
        occ_ids = []
        for part in o["occupied"].split(","):
            if "#" in part:
                occ_ids.append(int(part.split("#")[1]))
        rec["occupied_ids"] = occ_ids
        rows.append(rec)

    (OUT / "blocked_opportunities.json").write_text(
        json.dumps(rows, indent=2, default=str)
    )

    def col(name: str):
        return [r[name] for r in rows if r.get(name) is not None]

    report: dict = {
        "window": {
            "first_block": opps[0]["ts"].isoformat(),
            "last_block": opps[-1]["ts"].isoformat(),
            "raw_logs": len(events),
            "unique_opportunities": len(opps),
            "note": (
                "Hindsight only. Anchor=first 5m open at/after block log. "
                "Not a causal miss signal for backtest promotion."
            ),
        },
        "by_pair": {},
        "horizons": {},
    }

    for h in HORIZONS_H:
        mfe = col(f"h{h}_mfe")
        mae = col(f"h{h}_mae")
        end = col(f"h{h}_end")
        stop = col(f"h{h}_stop_hit")
        trail = col(f"h{h}_trail_armed")
        sim = []
        for r in rows:
            if r.get(f"h{h}_mfe") is None:
                continue
            if r[f"h{h}_stop_hit"]:
                if r[f"h{h}_mfe"] >= HARD_STOP_PRICE:
                    sim.append(None)  # path-order ambiguous
                else:
                    sim.append(-HARD_STOP_PRICE)
            else:
                sim.append(r[f"h{h}_end"])
        sim_c = [x for x in sim if x is not None]
        valid = [r for r in rows if r.get(f"h{h}_mfe") is not None]
        report["horizons"][f"{h}h"] = {
            "mfe": summarize(mfe),
            "mae": summarize(mae),
            "end_ret": summarize(end),
            "stop_hit_rate": float(np.mean(stop)) if stop else None,
            "trail_arm_rate": float(np.mean(trail)) if trail else None,
            "sim_hold_or_stop": summarize(sim_c),
            "sim_ambiguous_n": sum(1 for x in sim if x is None),
            "attractive_mfe_gt_1pct_no_stop": float(
                np.mean(
                    [
                        (r[f"h{h}_mfe"] or 0) >= 0.01 and not r[f"h{h}_stop_hit"]
                        for r in valid
                    ]
                )
            )
            if valid
            else None,
            "attractive_mfe_gt_2_5pct": float(
                np.mean([(r[f"h{h}_mfe"] or 0) >= HARD_STOP_PRICE for r in valid])
            )
            if valid
            else None,
        }

    for p in sorted({r["pair"] for r in rows}):
        sub = [r for r in rows if r["pair"] == p]
        report["by_pair"][p] = {
            "n": len(sub),
            "h6_mfe_mean": float(
                np.mean([r["h6_mfe"] for r in sub if r.get("h6_mfe") is not None])
            )
            if any(r.get("h6_mfe") is not None for r in sub)
            else None,
            "h24_mfe_mean": float(
                np.mean([r["h24_mfe"] for r in sub if r.get("h24_mfe") is not None])
            )
            if any(r.get("h24_mfe") is not None for r in sub)
            else None,
            "h6_stop_rate": float(
                np.mean(
                    [r["h6_stop_hit"] for r in sub if r.get("h6_stop_hit") is not None]
                )
            )
            if any(r.get("h6_stop_hit") is not None for r in sub)
            else None,
            "h24_end_mean": float(
                np.mean([r["h24_end"] for r in sub if r.get("h24_end") is not None])
            )
            if any(r.get("h24_end") is not None for r in sub)
            else None,
        }

    taken_rows = []
    for t in trades:
        if t["side"] != "short":
            continue
        if t["pair"] not in kcache or not kcache[t["pair"]]:
            kcache[t["pair"]] = load_pair_klines(t["pair"])
        if not kcache[t["pair"]]:
            continue
        rec = {
            "id": t["id"],
            "pair": t["pair"],
            "open": t["open_dt"].isoformat(),
            "exit_reason": t["exit_reason"],
            "realized": t["close_profit"],
            "realized_abs": t["close_profit_abs"],
        }
        for h in HORIZONS_H:
            m = path_metrics(kcache[t["pair"]], t["open_dt"], "short", h)
            if m:
                rec[f"h{h}_mfe"] = m["mfe"]
                rec[f"h{h}_mae"] = m["mae"]
                rec[f"h{h}_end"] = m["end_ret"]
                rec[f"h{h}_stop_hit"] = m["stop_hit"]
        taken_rows.append(rec)

    report["taken_trades_path"] = {
        "n": len(taken_rows),
        "realized": summarize(
            [t["realized"] for t in taken_rows if t.get("realized") is not None]
        ),
        "realized_abs": summarize(
            [t["realized_abs"] for t in taken_rows if t.get("realized_abs") is not None]
        ),
        "h6_mfe": summarize(
            [t["h6_mfe"] for t in taken_rows if t.get("h6_mfe") is not None]
        ),
        "h24_mfe": summarize(
            [t["h24_mfe"] for t in taken_rows if t.get("h24_mfe") is not None]
        ),
        "h6_end": summarize(
            [t["h6_end"] for t in taken_rows if t.get("h6_end") is not None]
        ),
        "h24_end": summarize(
            [t["h24_end"] for t in taken_rows if t.get("h24_end") is not None]
        ),
    }

    occ_sets: dict[str, int] = defaultdict(int)
    for o in opps:
        occ_sets[o["occupied"]] += 1
    report["top_occupied_states"] = sorted(
        [{"occupied": k, "n": v} for k, v in occ_sets.items()],
        key=lambda x: -x["n"],
    )[:15]

    by_day: dict[str, int] = defaultdict(int)
    for o in opps:
        by_day[o["ts"].date().isoformat()] += 1
    report["opps_by_day"] = dict(sorted(by_day.items()))

    # Side-by-side comparison table for narrative
    b6 = report["horizons"]["6h"]
    b24 = report["horizons"]["24h"]
    report["comparison_snapshot"] = {
        "blocked_h6_mfe_mean": b6["mfe"].get("mean"),
        "blocked_h24_mfe_mean": b24["mfe"].get("mean"),
        "blocked_h6_end_mean": b6["end_ret"].get("mean"),
        "blocked_h24_end_mean": b24["end_ret"].get("mean"),
        "blocked_h6_stop_rate": b6["stop_hit_rate"],
        "blocked_h24_stop_rate": b24["stop_hit_rate"],
        "taken_realized_mean": report["taken_trades_path"]["realized"].get("mean"),
        "taken_h6_mfe_mean": report["taken_trades_path"]["h6_mfe"].get("mean"),
        "taken_h24_mfe_mean": report["taken_trades_path"]["h24_mfe"].get("mean"),
    }

    (OUT / "summary.json").write_text(json.dumps(report, indent=2))
    (OUT / "taken_paths.json").write_text(json.dumps(taken_rows, indent=2, default=str))
    print(json.dumps(report, indent=2))
    print("wrote", OUT)


if __name__ == "__main__":
    main()
