#!/usr/bin/env python3
"""Replay entry-rule variants over the full candidate stream.

The counterfactuals that get run against `paper_orders` answer the wrong
question: which orders actually opened was decided by `max_open_exposure_usd`
freeing room, not by candidate quality, so filtering that table measures the
filter *and* the exposure cap's selection at the same time. This replays the
stream of signals the scanner produced (first signal per market/side, from
`snapshots`), so a filter is scored on everything it would have seen.

Held to settlement, taker entry at the recorded signal price, current Weather
fee curve on entry (settlement exit is free: price*(1-price) = 0 at 0 or 1).

Confidence intervals come from a bootstrap resampled over *events*, not
orders. Buckets inside one station/date are mutually exclusive and share one
forecast error, so per-order intervals would overstate precision by roughly
the average number of orders per event.
"""

from __future__ import annotations

import argparse
import dataclasses
import random
import sqlite3
from collections import defaultdict

FEE_RATE = 0.05
EDGE_THRESHOLD = 0.08


@dataclasses.dataclass(frozen=True)
class Candidate:
    event_slug: str
    target_date: str
    city: str
    kind: str
    bucket: str
    side: str
    price: float
    stake: float
    prob_raw: float
    prob_corr: float | None
    won: bool

    def net(self, stake: float) -> float:
        """Net USD on `stake` dollars: gross settlement P&L minus entry fee."""
        shares = stake / self.price
        fee = shares * FEE_RATE * self.price * (1.0 - self.price)
        gross = shares * (1.0 - self.price) if self.won else -stake
        return gross - fee


def load(db: str, since: str | None) -> list[Candidate]:
    conn = sqlite3.connect(f"file:{db}?mode=ro", uri=True)
    rows = conn.execute(
        """
        WITH first_signal AS (
          SELECT market_id, signal_side, MIN(id) id
          FROM snapshots
          WHERE signal_price IS NOT NULL AND signal_usd > 0
          GROUP BY market_id, signal_side
        )
        SELECT c.event_slug, c.target_date, c.city, c.kind, c.bucket, c.signal_side,
               c.signal_price, c.signal_usd, c.model_prob, c.model_prob_corr,
               s.winning_bucket
        FROM first_signal f
        JOIN snapshots c ON c.id=f.id
        JOIN settlements s ON s.event_slug = c.event_slug
        WHERE (? IS NULL OR c.target_date >= ?)
        """,
        (since, since),
    ).fetchall()
    out = []
    for slug, date, city, kind, bucket, side, price, usd, prob, prob_corr, winner in rows:
        if side != "NO":
            # YES has been off since 2026-07-14; no live YES stream to replay.
            continue
        out.append(
            Candidate(
                event_slug=slug,
                target_date=date,
                city=city,
                kind=kind,
                bucket=bucket,
                side=side,
                price=price,
                stake=usd,
                # snapshots store the YES-side probability of the bucket.
                prob_raw=1.0 - prob,
                prob_corr=None if prob_corr is None else 1.0 - prob_corr,
                won=bucket != winner,
            )
        )
    return out


@dataclasses.dataclass
class Rule:
    name: str
    corr_gate: bool = False
    min_price: float = 0.0
    max_price: float = 1.0
    stake_cap: float | None = None
    flat_stake: float | None = None
    kinds: tuple[str, ...] | None = None

    def admits(self, c: Candidate) -> bool:
        if self.kinds is not None and c.kind not in self.kinds:
            return False
        if not (self.min_price <= c.price <= self.max_price):
            return False
        if self.corr_gate:
            # No fresh bias row means no gate, matching the live fallback.
            prob = c.prob_corr if c.prob_corr is not None else c.prob_raw
            if prob - c.price < EDGE_THRESHOLD:
                return False
        return True

    def stake(self, c: Candidate) -> float:
        if self.flat_stake is not None:
            return self.flat_stake
        return c.stake if self.stake_cap is None else min(c.stake, self.stake_cap)


@dataclasses.dataclass
class Result:
    n: int
    staked: float
    net: float
    wins: int
    worst_day: float

    @property
    def roi(self) -> float:
        return 100.0 * self.net / self.staked if self.staked else 0.0

    @property
    def win_pct(self) -> float:
        return 100.0 * self.wins / self.n if self.n else 0.0


def evaluate(rule: Rule, cands: list[Candidate]) -> Result:
    n = wins = 0
    staked = net = 0.0
    by_day: dict[str, float] = defaultdict(float)
    for c in cands:
        if not rule.admits(c):
            continue
        stake = rule.stake(c)
        n += 1
        wins += c.won
        staked += stake
        pnl = c.net(stake)
        net += pnl
        by_day[c.target_date] += pnl
    return Result(n, staked, net, wins, min(by_day.values(), default=0.0))


def bootstrap_roi(rule: Rule, cands: list[Candidate], iters: int, seed: int) -> tuple[float, float]:
    """Percentile CI for ROI, resampling whole events with replacement."""
    events: dict[str, list[Candidate]] = defaultdict(list)
    for c in cands:
        events[c.event_slug].append(c)
    groups = list(events.values())
    rng = random.Random(seed)
    rois = []
    for _ in range(iters):
        sample = [c for _ in groups for c in rng.choice(groups)]
        r = evaluate(rule, sample)
        if r.staked:
            rois.append(r.roi)
    rois.sort()
    if not rois:
        return (0.0, 0.0)
    return (rois[int(0.025 * len(rois))], rois[int(0.975 * len(rois))])


RULES = [
    Rule("as-recorded (baseline)"),
    Rule("corr gate", corr_gate=True),
    Rule("corr gate + cap $20", corr_gate=True, stake_cap=20.0),
    Rule("corr gate + cap $10", corr_gate=True, stake_cap=10.0),
    Rule("corr gate + flat $8", corr_gate=True, flat_stake=8.0),
    Rule("corr gate + price>=0.50", corr_gate=True, min_price=0.50),
    Rule("corr gate + price>=0.50 + cap $20", corr_gate=True, min_price=0.50, stake_cap=20.0),
    Rule("cap $20 only", stake_cap=20.0),
    Rule("price>=0.50 only", min_price=0.50),
    Rule("current cap $20: highest only", corr_gate=True, stake_cap=20.0, kinds=("highest",)),
    Rule("current cap $20: lowest only", corr_gate=True, stake_cap=20.0, kinds=("lowest",)),
]


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--db", default="replay.db")
    ap.add_argument("--since", default=None, help="earliest target_date")
    ap.add_argument("--boot", type=int, default=2000)
    ap.add_argument("--seed", type=int, default=20260731)
    args = ap.parse_args()

    cands = load(args.db, args.since)
    events = len({c.event_slug for c in cands})
    print(f"candidates: {len(cands)} over {events} settled events")
    print(
        f"{'rule':38} {'n':>5} {'staked':>9} {'net':>9} {'ROI%':>7} "
        f"{'win%':>6} {'ROI 95% CI':>18} {'worst day':>10}"
    )
    for rule in RULES:
        r = evaluate(rule, cands)
        lo, hi = bootstrap_roi(rule, cands, args.boot, args.seed)
        print(
            f"{rule.name:38} {r.n:5d} {r.staked:9.0f} {r.net:9.1f} {r.roi:7.2f} "
            f"{r.win_pct:6.1f} {lo:8.2f} {hi:8.2f} {r.worst_day:10.1f}"
        )


if __name__ == "__main__":
    main()
