#!/usr/bin/env python3
"""Causal backtest of station-observation knockout entries.

For each settled event, replay station METARs in observation-time order and
the bot's contemporaneous CLOB snapshots.  A bucket is eligible only after
the running daily extreme makes it mathematically impossible to win.  The
trade buys NO at the executable complement of the recorded YES bid, includes
the current Weather taker-fee curve, checks displayed depth, and permits one
entry per event.  No final temperature is used to create a signal; settlement
is consulted only to audit that every alleged knockout really lost.
"""

from __future__ import annotations

import argparse
import csv
import io
import math
import re
import sqlite3
import statistics
import time
import tomllib
import urllib.error
import urllib.parse
import urllib.request
from collections import defaultdict
from dataclasses import dataclass
from datetime import date, datetime, timedelta, timezone
from pathlib import Path
from random import Random
from zoneinfo import ZoneInfo

BUCKET_RE = re.compile(
    r"^\s*(-?\d+)(?:\s*-\s*(-?\d+))?\s*°\s*([CF])"
    r"(?:\s+or\s+(below|lower|above|higher))?\s*$"
)


def iem_target(code: str) -> tuple[str, str] | None:
    us = {
        "KATL": "GA_ASOS", "KAUS": "TX_ASOS", "KDAL": "TX_ASOS",
        "KHOU": "TX_ASOS", "KBKF": "CO_ASOS", "KLAX": "CA_ASOS",
        "KSFO": "CA_ASOS", "KLGA": "NY_ASOS", "KMIA": "FL_ASOS",
        "KORD": "IL_ASOS", "KSEA": "WA_ASOS",
    }
    if code in us:
        return us[code], code[1:]
    networks = {
        "CY": "CA_ON_ASOS", "ED": "DE__ASOS", "EF": "FI__ASOS",
        "EG": "GB__ASOS", "EH": "NL__ASOS", "EP": "PL__ASOS",
        "FA": "ZA__ASOS", "LE": "ES__ASOS", "LF": "FR__ASOS",
        "LI": "IT__ASOS", "LL": "IL__ASOS", "LT": "TR__ASOS",
        "MM": "MX__ASOS", "MP": "PA__ASOS", "NZ": "NZ__ASOS",
        "OE": "SA__ASOS", "OP": "PK__ASOS", "RC": "TW__ASOS",
        "RJ": "JP__ASOS", "RK": "KR__ASOS", "RP": "PH__ASOS",
        "SA": "AR__ASOS", "SB": "BR__ASOS", "UU": "RU__ASOS",
        "VH": "HK__ASOS", "VI": "IN__ASOS", "WM": "MY__ASOS",
        "WS": "SG__ASOS", "YM": "AU__ASOS", "YS": "AU__ASOS",
        "ZB": "CN__ASOS", "ZG": "CN__ASOS", "ZH": "CN__ASOS",
        "ZS": "CN__ASOS", "ZU": "CN__ASOS",
    }
    network = networks.get(code[:2])
    return (network, code) if network else None


def fetch_metar(code: str, start: date, end: date, cache: Path):
    target = iem_target(code)
    if not target:
        return []
    cache.mkdir(parents=True, exist_ok=True)
    path = cache / f"{code}-{start}-{end}.csv"
    if not path.exists():
        network, station = target
        params = {
            "network": network, "station": station, "data": "tmpf",
            "year1": start.year, "month1": start.month, "day1": start.day,
            "year2": end.year, "month2": end.month, "day2": end.day,
            "tz": "Etc/UTC", "format": "onlycomma", "latlon": "no",
            "elev": "no", "missing": "null", "trace": "null",
            "direct": "no", "report_type": [3, 4],
        }
        url = "https://mesonet.agron.iastate.edu/cgi-bin/request/asos.py?" + urllib.parse.urlencode(params, doseq=True)
        request = urllib.request.Request(url, headers={"User-Agent": "weatherbot-research/1.0"})
        for attempt in range(6):
            try:
                with urllib.request.urlopen(request, timeout=60) as response:
                    path.write_bytes(response.read())
                break
            except urllib.error.HTTPError as error:
                if error.code != 429 or attempt == 5:
                    raise
                delay = float(error.headers.get("Retry-After", 2 ** attempt))
                time.sleep(min(max(delay, 1.0), 30.0))
        # Be a polite archive client; all later runs are served from cache.
        time.sleep(1.0)
    out = []
    for row in csv.DictReader(io.StringIO(path.read_text(errors="replace"))):
        try:
            valid = datetime.strptime(row["valid"], "%Y-%m-%d %H:%M").replace(tzinfo=timezone.utc)
            tmpf = float(row["tmpf"])
        except (KeyError, TypeError, ValueError):
            continue
        out.append((valid, (tmpf - 32.0) * 5.0 / 9.0))
    return out


def market_int(c: float, unit: str) -> int:
    value = c if unit == "C" else c * 9.0 / 5.0 + 32.0
    # Rust f64::round: halves away from zero.
    return math.floor(value + 0.5) if value >= 0 else math.ceil(value - 0.5)


def knocked_out(label: str, kind: str, running_c: float) -> bool:
    match = BUCKET_RE.match(label)
    if not match:
        return False
    lo, hi, unit, tail = match.groups()
    lo, hi = int(lo), int(hi) if hi is not None else None
    value = market_int(running_c, unit)
    if kind == "highest":
        if tail in ("below", "lower"):
            return value > lo
        if hi is not None:
            return value > hi
        if tail in ("above", "higher"):
            return False
        return value > lo
    if tail in ("above", "higher"):
        return value < lo
    if hi is not None:
        return value < lo
    if tail in ("below", "lower"):
        return False
    return value < lo


def fee_per_share(price: float, rate: float = 0.05) -> float:
    return rate * price * (1.0 - price)


@dataclass
class Trade:
    event: str
    day: str
    ts: str
    station: str
    bucket: str
    price: float
    usd: float
    pnl: float
    audit_win: bool


def bootstrap(trades: list[Trade], rounds=5000):
    grouped = defaultdict(list)
    for trade in trades:
        grouped[trade.event].append(trade)
    events = list(grouped.values())
    if not events:
        return (0.0, 0.0)
    rng = Random(20260810)
    rois = []
    for _ in range(rounds):
        sample = [rng.choice(events) for _ in events]
        stake = sum(t.usd for event in sample for t in event)
        pnl = sum(t.pnl for event in sample for t in event)
        rois.append(100 * pnl / stake)
    rois.sort()
    return rois[int(rounds * 0.025)], rois[int(rounds * 0.975)]


def report(label: str, trades: list[Trade]):
    stake = sum(x.usd for x in trades)
    pnl = sum(x.pnl for x in trades)
    lo, hi = bootstrap(trades)
    print(f"{label:12s} n={len(trades):4d} stake=${stake:8.2f} pnl=${pnl:8.2f} roi={100*pnl/stake if stake else 0:6.2f}% event-bootstrap95=[{lo:.2f},{hi:.2f}] audit_wins={sum(x.audit_win for x in trades)}/{len(trades)}")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--db", required=True)
    ap.add_argument("--stations", default="stations.toml")
    ap.add_argument("--cache", default="/tmp/weatherbot-metar-cache")
    ap.add_argument("--publication-lag-min", type=int, default=5)
    ap.add_argument("--min-net-edge", type=float, default=0.02)
    ap.add_argument("--stake", type=float, default=5.0)
    ap.add_argument("--test-since", default="2026-08-01")
    args = ap.parse_args()

    stations = tomllib.loads(Path(args.stations).read_text())["stations"]
    conn = sqlite3.connect(f"file:{args.db}?mode=ro", uri=True)
    conn.row_factory = sqlite3.Row
    date_min, date_max = conn.execute("SELECT MIN(target_date),MAX(target_date) FROM settlements").fetchone()
    start, end = date.fromisoformat(date_min), date.fromisoformat(date_max) + timedelta(days=1)
    codes = [r[0] for r in conn.execute("SELECT DISTINCT station FROM forecast_batches") if r[0] in stations]
    observations = {code: fetch_metar(code, start, end, Path(args.cache)) for code in codes}

    rows = conn.execute("""SELECT s.ts,s.event_slug,s.kind,s.target_date,s.bucket,
                                  s.best_bid,s.best_bid_size,s.quote_source,
                                  fb.station,st.winning_bucket
                           FROM snapshots s
                           JOIN forecast_batches fb ON fb.id=s.forecast_batch_id
                           JOIN settlements st ON st.event_slug=s.event_slug
                           WHERE s.best_bid IS NOT NULL AND s.best_bid_size IS NOT NULL
                           ORDER BY s.ts,s.id""").fetchall()
    by_cycle = defaultdict(list)
    for row in rows:
        stamp = datetime.fromisoformat(row["ts"].replace("Z", "+00:00"))
        code = row["station"]
        try:
            local_day = stamp.astimezone(ZoneInfo(stations[code]["tz"])).date().isoformat()
        except (KeyError, ValueError):
            continue
        if local_day != row["target_date"] or row["quote_source"] != "clob":
            continue
        # All buckets in a scan share the minute; this avoids bucket iteration order.
        key = (row["event_slug"], stamp.replace(second=0, microsecond=0))
        by_cycle[key].append(row)

    entered = set()
    trades = []
    running = {}
    obs_pos = defaultdict(int)
    for (event, stamp), buckets in sorted(by_cycle.items(), key=lambda x: x[0][1]):
        if event in entered:
            continue
        code, kind, target_day = buckets[0]["station"], buckets[0]["kind"], buckets[0]["target_date"]
        cutoff = stamp - timedelta(minutes=args.publication_lag_min)
        key = (code, target_day, kind)
        obs = observations.get(code, [])
        pos = obs_pos[key]
        while pos < len(obs) and obs[pos][0] <= cutoff:
            valid, temp = obs[pos]
            local = valid.astimezone(ZoneInfo(stations[code]["tz"])).date().isoformat()
            if local == target_day:
                if key not in running:
                    running[key] = temp
                elif kind == "highest":
                    running[key] = max(running[key], temp)
                else:
                    running[key] = min(running[key], temp)
            pos += 1
        obs_pos[key] = pos
        if key not in running:
            continue
        candidates = []
        for row in buckets:
            if not knocked_out(row["bucket"], kind, running[key]):
                continue
            price = 1.0 - row["best_bid"]
            if not 0 < price < 1:
                continue
            shares = args.stake / price
            if row["best_bid_size"] + 1e-9 < shares:
                continue
            net_edge = 1.0 - price - fee_per_share(price)
            if net_edge + 1e-12 < args.min_net_edge:
                continue
            candidates.append((net_edge, row, price, shares))
        if not candidates:
            continue
        _, row, price, shares = max(candidates, key=lambda x: x[0])
        entry_fee = shares * fee_per_share(price)
        won = row["bucket"] != row["winning_bucket"]
        pnl = shares * ((1.0 if won else 0.0) - price) - entry_fee
        trades.append(Trade(event, target_day, stamp.isoformat(), code, row["bucket"], price, args.stake + entry_fee, pnl, won))
        entered.add(event)
    conn.close()

    split = args.test_since
    report("all", trades)
    report("train", [x for x in trades if x.day < split])
    report("test", [x for x in trades if x.day >= split])
    by_day = defaultdict(list)
    for trade in trades:
        by_day[trade.day].append(trade)
    for day in sorted(by_day):
        report(day, by_day[day])
    bad = [x for x in trades if not x.audit_win]
    if bad:
        print("AUDIT FAILURES")
        for trade in bad:
            print(trade)


if __name__ == "__main__":
    main()
