#!/usr/bin/env python3
"""Replay fixed take-profit exits from weatherbot's 10-minute snapshots.

This is deliberately a conservative taker replay: the exit uses the first
eligible CLOB bid, charges the current Weather fee curve on entry and exit,
and can require enough top-of-book depth to liquidate the whole position.
It never uses a later peak price to choose the fill.
"""

from __future__ import annotations

import argparse
import dataclasses
import sqlite3
from collections import defaultdict
from datetime import datetime, timedelta


FEE_RATE = 0.05
THRESHOLDS = (0.05, 0.10, 0.15, 0.20, 0.30, 0.50)


@dataclasses.dataclass(frozen=True)
class Order:
    id: int
    ts: str
    event_slug: str
    target_date: str
    bucket: str
    side: str
    price: float
    shares: float
    usd: float
    edge: float
    pnl: float


@dataclasses.dataclass
class Result:
    hold: float
    pnl: dict[float, float]
    triggered: set[float]


def fee(shares: float, price: float) -> float:
    return shares * FEE_RATE * price * (1.0 - price)


def position(order: Order, fixed_five: bool) -> tuple[float, float, float, float]:
    """Return shares, market stake, entry fee and total capital at risk."""
    if not fixed_five:
        entry_fee = fee(order.shares, order.price)
        return order.shares, order.usd, entry_fee, order.usd + entry_fee

    fee_per_share = fee(1.0, order.price)
    shares = 5.0 / (order.price + fee_per_share)
    stake = shares * order.price
    entry_fee = shares * fee_per_share
    return shares, stake, entry_fee, stake + entry_fee


def parse_ts(value: str) -> datetime:
    return datetime.fromisoformat(value.replace("Z", "+00:00"))


def meta(conn: sqlite3.Connection, key: str) -> str | None:
    row = conn.execute("SELECT value FROM bot_meta WHERE key=?", (key,)).fetchone()
    return row[0] if row else None


def load_orders(conn: sqlite3.Connection, sample: str) -> tuple[list[Order], bool]:
    rows = [
        Order(*row)
        for row in conn.execute(
            """
            SELECT id,ts,event_slug,target_date,bucket,side,price,shares,usd,edge,pnl
            FROM paper_orders
            WHERE status='settled' AND pnl IS NOT NULL
            ORDER BY ts,id
            """
        )
    ]
    fixed_five = False
    if sample == "yes":
        rows = [o for o in rows if o.side == "YES"]
    elif sample == "no":
        rows = [o for o in rows if o.side == "NO"]
    elif sample in (
        "shadow-a",
        "shadow-a-forward",
        "shadow-b-proxy",
        "shadow-b-v2-proxy",
    ):
        rows = [o for o in rows if o.side == "NO" and 0.08 <= o.edge < 0.20]
        if sample == "shadow-a-forward":
            boundary = meta(conn, "shadow_no_08_20_start")
            rows = [o for o in rows if boundary is not None and o.ts >= boundary]
        elif sample in ("shadow-b-proxy", "shadow-b-v2-proxy"):
            if sample == "shadow-b-v2-proxy":
                boundary = meta(conn, "allocation_v2_start")
                rows = [o for o in rows if boundary is not None and o.ts >= boundary]
            rows = [o for o in rows if o.price >= 0.70]
            by_event: dict[str, list[Order]] = defaultdict(list)
            for order in rows:
                by_event[order.event_slug].append(order)
            selected = []
            for event_orders in by_event.values():
                first = min(parse_ts(o.ts) for o in event_orders)
                same_cycle = [
                    o
                    for o in event_orders
                    if parse_ts(o.ts) <= first + timedelta(minutes=2)
                ]
                selected.append(max(same_cycle, key=lambda o: (o.edge, -o.id)))
            rows = sorted(selected, key=lambda o: (o.ts, o.id))
            fixed_five = True
    elif sample == "v2":
        boundary = meta(conn, "allocation_v2_start")
        rows = [o for o in rows if boundary is not None and o.ts >= boundary]
    return rows, fixed_five


def replay(
    conn: sqlite3.Connection,
    orders: list[Order],
    fixed_five: bool,
    require_depth: bool,
) -> dict[int, Result]:
    conn.execute("DROP TABLE IF EXISTS temp.replay_orders")
    conn.execute(
        "CREATE TEMP TABLE replay_orders(id INTEGER PRIMARY KEY,event_slug,bucket,entry_ts)"
    )
    conn.executemany(
        "INSERT INTO replay_orders VALUES(?,?,?,?)",
        [(o.id, o.event_slug, o.bucket, o.ts) for o in orders],
    )
    conn.execute(
        "CREATE INDEX temp.replay_order_key ON replay_orders(event_slug,bucket,entry_ts)"
    )

    results: dict[int, Result] = {}
    lookup = {o.id: o for o in orders}
    for order in orders:
        shares, stake, entry_fee, _capital = position(order, fixed_five)
        won = order.pnl > 0.0
        hold = shares * (1.0 if won else 0.0) - stake - entry_fee
        results[order.id] = Result(
            hold=hold,
            pnl={threshold: hold for threshold in THRESHOLDS},
            triggered=set(),
        )

    sql = """
        SELECT r.id,s.best_bid,s.best_ask,s.best_bid_size,s.best_ask_size
        FROM snapshots s
        JOIN replay_orders r
          ON r.event_slug=s.event_slug AND r.bucket=s.bucket AND s.ts>=r.entry_ts
        WHERE s.quote_source='clob'
        ORDER BY r.id,s.ts,s.id
    """
    for order_id, best_bid, best_ask, bid_size, ask_size in conn.execute(sql):
        result = results[order_id]
        if len(result.triggered) == len(THRESHOLDS):
            continue
        order = lookup[order_id]
        shares, stake, entry_fee, capital = position(order, fixed_five)
        if order.side == "YES":
            price, depth = best_bid, bid_size
        else:
            price = None if best_ask is None else 1.0 - best_ask
            depth = ask_size
        if price is None or not 0.0 < price < 1.0:
            continue
        if require_depth and (depth is None or depth + 1e-9 < shares):
            continue
        liquidation = (
            shares * price
            - stake
            - entry_fee
            - fee(shares, price)
        )
        for threshold in THRESHOLDS:
            if threshold not in result.triggered and liquidation >= threshold * capital:
                result.pnl[threshold] = liquidation
                result.triggered.add(threshold)
    return results


def summarize(
    orders: list[Order], results: dict[int, Result], threshold: float, dates: set[str]
) -> tuple[int, float, float, int]:
    chosen = [o for o in orders if o.target_date in dates]
    hold = sum(results[o.id].hold for o in chosen)
    policy = sum(results[o.id].pnl[threshold] for o in chosen)
    triggers = sum(threshold in results[o.id].triggered for o in chosen)
    return len(chosen), hold, policy, triggers


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("db")
    parser.add_argument(
        "--sample",
        choices=(
            "all",
            "yes",
            "no",
            "shadow-a",
            "shadow-a-forward",
            "shadow-b-proxy",
            "shadow-b-v2-proxy",
            "v2",
        ),
        default="shadow-b-proxy",
    )
    parser.add_argument(
        "--ignore-depth",
        action="store_true",
        help="optimistic top-of-book replay; default requires full top-level depth",
    )
    args = parser.parse_args()

    conn = sqlite3.connect(f"file:{args.db}?mode=ro", uri=True)
    orders, fixed_five = load_orders(conn, args.sample)
    if not orders:
        raise SystemExit("no settled orders in selected sample")
    results = replay(conn, orders, fixed_five, not args.ignore_depth)
    dates = sorted({o.target_date for o in orders})
    split = max(1, int(len(dates) * 0.6))
    train_dates, test_dates = set(dates[:split]), set(dates[split:])
    all_dates = set(dates)

    print(
        f"sample={args.sample} orders={len(orders)} dates={dates[0]}..{dates[-1]} "
        f"depth={'ignored' if args.ignore_depth else 'full-top-level'} "
        f"sizing={'$5/event' if fixed_five else 'original'} fee_rate={FEE_RATE:.2%}"
    )
    print("threshold  n_all triggers   hold_net  exit_net    delta")
    for threshold in THRESHOLDS:
        n, hold, policy, triggers = summarize(orders, results, threshold, all_dates)
        print(
            f"{threshold:>8.0%} {n:6d} {triggers:8d} "
            f"{hold:10.2f} {policy:9.2f} {policy-hold:+8.2f}"
        )

    train_scores = {
        threshold: summarize(orders, results, threshold, train_dates)[2]
        for threshold in THRESHOLDS
    }
    chosen = max(THRESHOLDS, key=lambda threshold: train_scores[threshold])
    n_train, hold_train, pnl_train, triggers_train = summarize(
        orders, results, chosen, train_dates
    )
    n_test, hold_test, pnl_test, triggers_test = summarize(
        orders, results, chosen, test_dates
    )
    print(
        f"chronological split: train={dates[0]}..{dates[split-1]} "
        f"test={dates[split] if split < len(dates) else 'none'}..{dates[-1]}"
    )
    print(
        f"train-chosen={chosen:.0%}: train n={n_train} triggers={triggers_train} "
        f"hold={hold_train:+.2f} exit={pnl_train:+.2f} delta={pnl_train-hold_train:+.2f}"
    )
    print(
        f"held-out:          test  n={n_test} triggers={triggers_test} "
        f"hold={hold_test:+.2f} exit={pnl_test:+.2f} delta={pnl_test-hold_test:+.2f}"
    )


if __name__ == "__main__":
    main()
