#!/usr/bin/env python3
"""
Formal event-study + backtest for short-horizon Polymarket lag hypothesis.

Hypothesis template:
  - Cross-venue spot dislocation filter (Binance vs Coinbase)
  - Order-book imbalance filter
  - Delayed entry after signal "open"
  - Exit on fixed holding horizon

This script is intentionally data-driven: you provide a timestamped CSV with the
required columns and it outputs:
  - trade-level CSV
  - summary JSON
"""

from __future__ import annotations

import argparse
import json
import math
from dataclasses import asdict, dataclass
from pathlib import Path

import numpy as np
import pandas as pd

REQUIRED_COLUMNS = {
    "ts",
    "binance_px",
    "coinbase_px",
    "chainlink_px",
    "market_mid",
    "orderbook_imbalance",
}


@dataclass
class Trade:
    signal_open_ts: str
    entry_ts: str
    exit_ts: str
    side: str
    entry_price: float
    exit_price: float
    pnl_gross: float
    pnl_net: float
    hold_seconds: float
    spot_delta_usd: float
    chainlink_gap_usd: float
    orderbook_imbalance: float


def _load_ticks(path: Path) -> pd.DataFrame:
    if not path.is_file():
        raise FileNotFoundError(f"input CSV not found: {path}")
    df = pd.read_csv(path)
    missing = REQUIRED_COLUMNS - set(df.columns)
    if missing:
        raise ValueError(f"missing required columns: {sorted(missing)}")
    df = df.copy()
    df["ts"] = pd.to_datetime(df["ts"], utc=True, errors="coerce")
    df = df.dropna(subset=["ts"]).sort_values("ts").reset_index(drop=True)
    for c in ["binance_px", "coinbase_px", "chainlink_px", "market_mid", "orderbook_imbalance"]:
        df[c] = pd.to_numeric(df[c], errors="coerce")
    df = df.dropna(subset=["binance_px", "coinbase_px", "chainlink_px", "market_mid", "orderbook_imbalance"]).reset_index(
        drop=True
    )
    return df


def _first_row_at_or_after(df: pd.DataFrame, ts: pd.Timestamp) -> int | None:
    idx = df["ts"].searchsorted(ts, side="left")
    if idx >= len(df):
        return None
    return int(idx)


def _first_row_at_or_before(df: pd.DataFrame, ts: pd.Timestamp) -> int | None:
    idx = df["ts"].searchsorted(ts, side="right") - 1
    if idx < 0:
        return None
    return int(idx)


def run_backtest(
    df: pd.DataFrame,
    *,
    min_spot_delta_usd: float,
    min_imbalance_abs: float,
    entry_delay_min_sec: float,
    entry_delay_max_sec: float,
    hold_sec: float,
    signal_cooldown_sec: float,
    fee_bps_per_side: float,
) -> list[Trade]:
    out: list[Trade] = []

    # Core features
    df = df.copy()
    df["spot_ref"] = 0.5 * (df["binance_px"] + df["coinbase_px"])
    df["spot_delta_usd"] = df["binance_px"] - df["coinbase_px"]
    df["chainlink_gap_usd"] = df["spot_ref"] - df["chainlink_px"]

    # Event condition: dislocation + imbalance. Direction comes from oracle gap sign.
    cond = (df["spot_delta_usd"].abs() >= float(min_spot_delta_usd)) & (
        df["orderbook_imbalance"].abs() >= float(min_imbalance_abs)
    )
    event_idx = np.flatnonzero(cond.values)
    if len(event_idx) == 0:
        return out

    last_signal_open: pd.Timestamp | None = None
    cooldown = pd.Timedelta(seconds=float(signal_cooldown_sec))

    for i in event_idx:
        signal_ts = pd.Timestamp(df.at[i, "ts"])
        if last_signal_open is not None and signal_ts < last_signal_open + cooldown:
            continue

        gap = float(df.at[i, "chainlink_gap_usd"])
        imb = float(df.at[i, "orderbook_imbalance"])
        if not math.isfinite(gap) or gap == 0:
            continue

        # Direction: if reference > chainlink, assume upward catch-up pressure.
        side = "long" if gap > 0 else "short"
        # Optional consistency gate: imbalance sign should agree with side.
        if side == "long" and imb <= 0:
            continue
        if side == "short" and imb >= 0:
            continue

        t_entry_min = signal_ts + pd.Timedelta(seconds=float(entry_delay_min_sec))
        t_entry_max = signal_ts + pd.Timedelta(seconds=float(entry_delay_max_sec))
        j = _first_row_at_or_after(df, t_entry_min)
        if j is None:
            continue
        if pd.Timestamp(df.at[j, "ts"]) > t_entry_max:
            continue

        entry_ts = pd.Timestamp(df.at[j, "ts"])
        entry_px = float(df.at[j, "market_mid"])
        if not math.isfinite(entry_px):
            continue

        t_exit = entry_ts + pd.Timedelta(seconds=float(hold_sec))
        k = _first_row_at_or_after(df, t_exit)
        if k is None:
            continue
        exit_ts = pd.Timestamp(df.at[k, "ts"])
        exit_px = float(df.at[k, "market_mid"])
        if not math.isfinite(exit_px):
            continue

        gross = (exit_px - entry_px) if side == "long" else (entry_px - exit_px)
        fee = (float(fee_bps_per_side) / 10_000.0) * (abs(entry_px) + abs(exit_px))
        net = gross - fee

        out.append(
            Trade(
                signal_open_ts=str(signal_ts),
                entry_ts=str(entry_ts),
                exit_ts=str(exit_ts),
                side=side,
                entry_price=entry_px,
                exit_price=exit_px,
                pnl_gross=float(gross),
                pnl_net=float(net),
                hold_seconds=float((exit_ts - entry_ts).total_seconds()),
                spot_delta_usd=float(df.at[i, "spot_delta_usd"]),
                chainlink_gap_usd=gap,
                orderbook_imbalance=imb,
            )
        )
        last_signal_open = signal_ts
    return out


def _summary(trades: list[Trade]) -> dict:
    if not trades:
        return {
            "n_trades": 0,
            "win_rate": 0.0,
            "total_pnl_net": 0.0,
            "avg_pnl_net": 0.0,
            "median_pnl_net": 0.0,
            "sharpe_per_trade": 0.0,
            "max_drawdown_net": 0.0,
        }
    pnls = np.array([t.pnl_net for t in trades], dtype=float)
    wins = float(np.mean(pnls > 0))
    cum = np.cumsum(pnls)
    peak = np.maximum.accumulate(cum)
    dd = cum - peak
    sd = float(np.std(pnls))
    sharpe_trade = float(np.mean(pnls) / sd) if sd > 0 else 0.0
    return {
        "n_trades": int(len(trades)),
        "win_rate": round(wins, 4),
        "total_pnl_net": round(float(np.sum(pnls)), 8),
        "avg_pnl_net": round(float(np.mean(pnls)), 8),
        "median_pnl_net": round(float(np.median(pnls)), 8),
        "sharpe_per_trade": round(sharpe_trade, 6),
        "max_drawdown_net": round(float(np.min(dd)), 8),
    }


def main() -> None:
    ap = argparse.ArgumentParser(description="Backtest Polymarket/Chainlink lag hypothesis on tick data.")
    ap.add_argument(
        "--input-csv",
        type=Path,
        default=Path("RenTech/data/logs/polymarket_chainlink_ticks.csv"),
        help="Tick CSV with required columns. See research markdown for schema.",
    )
    ap.add_argument("--min-spot-delta-usd", type=float, default=50.0)
    ap.add_argument("--min-imbalance-abs", type=float, default=1.8)
    ap.add_argument("--entry-delay-min-sec", type=float, default=60.0)
    ap.add_argument("--entry-delay-max-sec", type=float, default=180.0)
    ap.add_argument("--hold-sec", type=float, default=14.0)
    ap.add_argument("--signal-cooldown-sec", type=float, default=300.0)
    ap.add_argument("--fee-bps-per-side", type=float, default=5.0)
    ap.add_argument(
        "--out-trades",
        type=Path,
        default=Path("RenTech/data/logs/polymarket_chainlink_backtest_trades.csv"),
    )
    ap.add_argument(
        "--out-summary",
        type=Path,
        default=Path("RenTech/data/logs/polymarket_chainlink_backtest_summary.json"),
    )
    args = ap.parse_args()

    df = _load_ticks(args.input_csv.expanduser())
    trades = run_backtest(
        df,
        min_spot_delta_usd=float(args.min_spot_delta_usd),
        min_imbalance_abs=float(args.min_imbalance_abs),
        entry_delay_min_sec=float(args.entry_delay_min_sec),
        entry_delay_max_sec=float(args.entry_delay_max_sec),
        hold_sec=float(args.hold_sec),
        signal_cooldown_sec=float(args.signal_cooldown_sec),
        fee_bps_per_side=float(args.fee_bps_per_side),
    )
    summary = _summary(trades)

    args.out_trades.parent.mkdir(parents=True, exist_ok=True)
    pd.DataFrame([asdict(t) for t in trades]).to_csv(args.out_trades, index=False)

    args.out_summary.parent.mkdir(parents=True, exist_ok=True)
    payload = {
        "input_csv": str(args.input_csv.expanduser()),
        "params": {
            "min_spot_delta_usd": float(args.min_spot_delta_usd),
            "min_imbalance_abs": float(args.min_imbalance_abs),
            "entry_delay_min_sec": float(args.entry_delay_min_sec),
            "entry_delay_max_sec": float(args.entry_delay_max_sec),
            "hold_sec": float(args.hold_sec),
            "signal_cooldown_sec": float(args.signal_cooldown_sec),
            "fee_bps_per_side": float(args.fee_bps_per_side),
        },
        "summary": summary,
    }
    args.out_summary.write_text(json.dumps(payload, indent=2), encoding="utf-8")

    print(json.dumps(payload, indent=2))
    print(f"Wrote trades -> {args.out_trades}")
    print(f"Wrote summary -> {args.out_summary}")


if __name__ == "__main__":
    main()
