#!/usr/bin/env python3
"""Backtest the F1 equal-model-family shadow probability from stored forecasts.

The script is deliberately read-only. Run it against a consistent SQLite copy,
not the live production database. Historical corrected F1 is intentionally not
estimated because bias calibrations are not versioned by snapshot timestamp.
"""

from __future__ import annotations

import argparse
import bisect
import json
import math
import re
import sqlite3
import statistics
import tomllib
from collections import defaultdict
from dataclasses import dataclass, field
from datetime import datetime, timedelta
from pathlib import Path
from zoneinfo import ZoneInfo


REQUIRED_MODELS = ("ecmwf_ifs025", "gfs_seamless", "icon_seamless")
ALL_MODELS = REQUIRED_MODELS + ("cma_grapes_global",)
BUCKET_RE = re.compile(
    r"^\s*(-?\d+)(?:\s*-\s*(-?\d+))?\s*°\s*([CF])"
    r"(?:\s+or\s+(below|lower|above|higher))?\s*$"
)
MAX_FORECAST_AGE = timedelta(hours=12)


def parse_time(value: str) -> datetime:
    return datetime.fromisoformat(value)


def rust_round(value: float) -> int:
    """Match f64::round(): halfway cases go away from zero."""
    return math.floor(value + 0.5) if value >= 0 else math.ceil(value - 0.5)


def parse_bucket(label: str):
    match = BUCKET_RE.match(label)
    if not match:
        return None
    lo = int(match.group(1))
    hi = int(match.group(2)) if match.group(2) is not None else None
    unit = match.group(3)
    direction = match.group(4)
    if hi is not None and direction is None:
        bound = ("range", lo, hi)
    elif direction in ("below", "lower"):
        bound = ("below", lo)
    elif direction in ("above", "higher"):
        bound = ("above", lo)
    elif hi is None and direction is None:
        bound = ("exact", lo)
    else:
        return None
    return bound, unit


def bucket_probability(members: list[float], parsed) -> float:
    bound, unit = parsed
    hits = 0
    for celsius in members:
        value = celsius if unit == "C" else celsius * 9.0 / 5.0 + 32.0
        market_int = rust_round(value)
        if bound[0] == "range":
            matched = bound[1] <= market_int <= bound[2]
        elif bound[0] == "below":
            matched = market_int <= bound[1]
        elif bound[0] == "above":
            matched = market_int >= bound[1]
        else:
            matched = market_int == bound[1]
        hits += matched
    return hits / len(members)


@dataclass
class Batch:
    available_at: datetime
    models: dict[str, list[float]]
    cache: dict[str, tuple[float, float, float]] = field(default_factory=dict)

    def probabilities(self, bucket: str) -> tuple[float, float, float] | None:
        if bucket in self.cache:
            return self.cache[bucket]
        parsed = parse_bucket(bucket)
        if parsed is None:
            return None
        family = statistics.fmean(
            bucket_probability(self.models[model], parsed) for model in REQUIRED_MODELS
        )
        required_members = [
            member for model in REQUIRED_MODELS for member in self.models[model]
        ]
        pooled_members = [
            member
            for model in ALL_MODELS
            for member in self.models.get(model, [])
        ]
        pooled = bucket_probability(pooled_members, parsed)
        pooled_required = bucket_probability(required_members, parsed)
        self.cache[bucket] = (family, pooled, pooled_required)
        return family, pooled, pooled_required


@dataclass
class ErrorAccumulator:
    rows: int = 0
    raw: float = 0.0
    family: float = 0.0
    market: float = 0.0
    market_rows: int = 0

    def add(self, raw: float, family: float, truth: float, market: float | None):
        self.rows += 1
        self.raw += (raw - truth) ** 2
        self.family += (family - truth) ** 2
        if market is not None:
            self.market += (market - truth) ** 2
            self.market_rows += 1

    def result(self) -> dict:
        raw = self.raw / self.rows
        family = self.family / self.rows
        return {
            "rows": self.rows,
            "raw_brier": raw,
            "family_brier": family,
            "family_vs_raw_pct": (family / raw - 1.0) * 100.0 if raw else None,
            "market_rows": self.market_rows,
            "market_brier": self.market / self.market_rows if self.market_rows else None,
        }


def load_stations(path: Path):
    with path.open("rb") as handle:
        data = tomllib.load(handle)["stations"]
    by_city = {row["city"]: (code, ZoneInfo(row["tz"])) for code, row in data.items()}
    return by_city


def load_forecasts(conn: sqlite3.Connection):
    grouped = defaultdict(list)
    columns = {row[1] for row in conn.execute("PRAGMA table_info(forecasts)")}
    batch_column = (
        "forecast_batch_id" if "forecast_batch_id" in columns else "NULL"
    )
    query = f"""
        SELECT ts, station, target_date, kind, model, members_json, {batch_column}
        FROM forecasts
        WHERE model IN (?, ?, ?, ?)
        ORDER BY station, target_date, kind, ts, id
    """
    by_id_records = defaultdict(list)
    for ts, station, target_date, kind, model, members_json, batch_id in conn.execute(
        query, ALL_MODELS
    ):
        record = (parse_time(ts), model, json.loads(members_json))
        grouped[(station, target_date, kind)].append(
            record
        )
        if batch_id is not None:
            by_id_records[batch_id].append(record)

    timelines = {}
    for key, rows in grouped.items():
        batches = []
        start = None
        records = []
        for at, model, members in rows:
            if start is None or at - start <= timedelta(seconds=5):
                start = at if start is None else start
                records.append((at, model, members))
                continue
            _append_complete_batch(batches, records)
            start = at
            records = [(at, model, members)]
        _append_complete_batch(batches, records)
        timelines[key] = ([batch.available_at for batch in batches], batches)
    by_id = {}
    for batch_id, records in by_id_records.items():
        models = {}
        for _, model, members in records:
            models.setdefault(model, members)
        if all(models.get(model) for model in REQUIRED_MODELS):
            by_id[batch_id] = Batch(max(row[0] for row in records), models)
    return timelines, by_id


def _append_complete_batch(batches: list[Batch], records):
    if not records:
        return
    models = {}
    for _, model, members in records:
        models.setdefault(model, members)
    if all(models.get(model) for model in REQUIRED_MODELS):
        batches.append(Batch(max(row[0] for row in records), models))


def find_batch(timelines, key, at: datetime) -> Batch | None:
    timeline = timelines.get(key)
    if timeline is None:
        return None
    times, batches = timeline
    index = bisect.bisect_right(times, at) - 1
    if index < 0:
        return None
    batch = batches[index]
    return batch if at - batch.available_at <= MAX_FORECAST_AGE else None


def market_midpoint(bid, ask):
    if bid is None or ask is None or not (0 <= bid <= ask <= 1):
        return None
    return (bid + ask) / 2.0


def mean_event_result(events: dict[str, ErrorAccumulator]) -> dict:
    per_event = [value.result() for value in events.values() if value.rows]
    raw = statistics.fmean(row["raw_brier"] for row in per_event)
    family = statistics.fmean(row["family_brier"] for row in per_event)
    deltas = [row["family_brier"] - row["raw_brier"] for row in per_event]
    delta_mean = statistics.fmean(deltas)
    delta_se = (
        statistics.stdev(deltas) / math.sqrt(len(deltas)) if len(deltas) > 1 else 0.0
    )
    return {
        "events": len(per_event),
        "raw_brier": raw,
        "family_brier": family,
        "family_vs_raw_pct": (family / raw - 1.0) * 100.0 if raw else None,
        "events_family_better": sum(delta < 0 for delta in deltas),
        "events_tied": sum(abs(delta) < 1e-15 for delta in deltas),
        "paired_delta_family_minus_raw": delta_mean,
        "paired_delta_95pct_ci": [
            delta_mean - 1.96 * delta_se,
            delta_mean + 1.96 * delta_se,
        ],
    }


def run_snapshot_backtest(conn, stations, timelines, batches_by_id):
    total = ErrorAccumulator()
    by_event = defaultdict(ErrorAccumulator)
    exact_by_event = defaultdict(ErrorAccumulator)
    last_rows = {}
    exact_last_rows = {}
    exact_by_city = defaultdict(ErrorAccumulator)
    skipped = defaultdict(int)
    reconstruction_abs = 0.0
    reconstruction_required_abs = 0.0
    reconstruction_rows = 0
    exact_reconstruction = ErrorAccumulator()
    min_ts = None
    max_ts = None
    snapshot_columns = {row[1] for row in conn.execute("PRAGMA table_info(snapshots)")}
    batch_column = (
        "s.forecast_batch_id"
        if "forecast_batch_id" in snapshot_columns
        else "NULL"
    )
    query = f"""
        SELECT s.ts, s.event_slug, s.city, s.kind, s.target_date, s.bucket,
               s.model_prob, s.reference_bid, s.reference_ask, t.winning_bucket,
               {batch_column}
        FROM snapshots s
        JOIN settlements t ON t.event_slug = s.event_slug
        ORDER BY s.id
    """
    for row in conn.execute(query):
        ts, event, city, kind, target_date, bucket, raw, bid, ask, winner, batch_id = row
        station_info = stations.get(city)
        if station_info is None:
            skipped["unknown_station"] += 1
            continue
        station, timezone = station_info
        at = parse_time(ts)
        if at.astimezone(timezone).date().isoformat() >= target_date:
            skipped["same_day_or_later"] += 1
            continue
        batch = batches_by_id.get(batch_id) if batch_id is not None else None
        if batch is None and batch_id is None:
            batch = find_batch(timelines, (station, target_date, kind), at)
        if batch is None:
            skipped[
                "missing_linked_forecast_batch"
                if batch_id is not None
                else "no_fresh_complete_forecast"
            ] += 1
            continue
        probabilities = batch.probabilities(bucket)
        if probabilities is None:
            skipped["unparsed_bucket"] += 1
            continue
        family, pooled, pooled_required = probabilities
        truth = float(bucket == winner)
        market = market_midpoint(bid, ask)
        total.add(raw, family, truth, market)
        by_event[event].add(raw, family, truth, market)
        last_rows[(event, bucket)] = (raw, family, truth, market)
        reconstruction_abs += abs(raw - pooled)
        reconstruction_required_abs += abs(raw - pooled_required)
        reconstruction_rows += 1
        if min(abs(raw - pooled), abs(raw - pooled_required)) < 1e-12:
            exact_reconstruction.add(raw, family, truth, market)
            exact_by_event[event].add(raw, family, truth, market)
            exact_by_city[city].add(raw, family, truth, market)
            exact_last_rows[(event, bucket)] = (raw, family, truth, market)
        min_ts = ts if min_ts is None else min(min_ts, ts)
        max_ts = ts if max_ts is None else max(max_ts, ts)

    last = ErrorAccumulator()
    for raw, family, truth, market in last_rows.values():
        last.add(raw, family, truth, market)
    exact_last = ErrorAccumulator()
    for raw, family, truth, market in exact_last_rows.values():
        exact_last.add(raw, family, truth, market)
    city_results = []
    for city, acc in exact_by_city.items():
        result = acc.result()
        result["city"] = city
        city_results.append(result)
    city_results.sort(key=lambda row: row["family_brier"] - row["raw_brier"])
    return {
        "time_range": [min_ts, max_ts],
        "row_weighted_10min": total.result(),
        "event_equal_weighted": mean_event_result(by_event),
        "last_eligible_per_event_bucket": last.result(),
        "stored_raw_reconstruction_mae": (
            reconstruction_abs / reconstruction_rows if reconstruction_rows else None
        ),
        "stored_raw_vs_required_only_mae": (
            reconstruction_required_abs / reconstruction_rows
            if reconstruction_rows
            else None
        ),
        "exact_reconstruction_subset": exact_reconstruction.result()
        if exact_reconstruction.rows
        else None,
        "exact_event_equal_weighted": mean_event_result(exact_by_event),
        "exact_last_eligible_per_event_bucket": exact_last.result(),
        "skipped": dict(skipped),
        "best_cities": city_results[:8],
        "worst_cities": city_results[-8:][::-1],
    }


def run_order_backtest(conn, stations, timelines, batches_by_id):
    acc = ErrorAccumulator()
    exact = ErrorAccumulator()
    exact_by_event = defaultdict(ErrorAccumulator)
    retained = ErrorAccumulator()
    retained_orders = 0
    retained_stake = 0.0
    retained_settlement_pnl = 0.0
    skipped = defaultdict(int)
    order_columns = {row[1] for row in conn.execute("PRAGMA table_info(paper_orders)")}
    batch_column = (
        "o.forecast_batch_id"
        if "forecast_batch_id" in order_columns
        else "NULL"
    )
    query = f"""
        SELECT o.event_slug, o.ts, o.city, o.kind, o.target_date, o.bucket, o.side, o.price,
               o.shares, o.usd, o.entry_fee, o.model_prob, s.winning_bucket,
               {batch_column}
        FROM paper_orders o
        JOIN settlements s ON s.event_slug = o.event_slug
        ORDER BY o.id
    """
    for row in conn.execute(query):
        event, ts, city, kind, target_date, bucket, side, price, shares, usd, fee, raw, winner, batch_id = row
        station_info = stations.get(city)
        if station_info is None:
            skipped["unknown_station"] += 1
            continue
        station, _ = station_info
        batch = batches_by_id.get(batch_id) if batch_id is not None else None
        if batch is None and batch_id is None:
            batch = find_batch(timelines, (station, target_date, kind), parse_time(ts))
        if batch is None:
            skipped[
                "missing_linked_forecast_batch"
                if batch_id is not None
                else "no_fresh_complete_forecast"
            ] += 1
            continue
        probabilities = batch.probabilities(bucket)
        if probabilities is None:
            skipped["unparsed_bucket"] += 1
            continue
        family_yes, pooled_yes, pooled_required_yes = probabilities
        family = family_yes if side == "YES" else 1.0 - family_yes
        pooled = pooled_yes if side == "YES" else 1.0 - pooled_yes
        pooled_required = (
            pooled_required_yes if side == "YES" else 1.0 - pooled_required_yes
        )
        truth_yes = bucket == winner
        truth = float(truth_yes if side == "YES" else not truth_yes)
        acc.add(raw, family, truth, None)
        if min(abs(raw - pooled), abs(raw - pooled_required)) >= 1e-12:
            skipped["raw_probability_not_reconstructed"] += 1
            continue
        exact.add(raw, family, truth, None)
        exact_by_event[event].add(raw, family, truth, None)
        family_edge = family - price
        if 0.08 <= family_edge < 0.30 and 0.10 <= price <= 0.98:
            retained.add(raw, family, truth, None)
            retained_orders += 1
            retained_stake += usd
            retained_settlement_pnl += shares * truth - usd - fee
    result = acc.result()
    result.update(
        {
            "exact_reconstruction_scores": exact.result() if exact.rows else None,
            "exact_event_equal_weighted": mean_event_result(exact_by_event)
            if exact_by_event
            else None,
            "retained_by_same_entry_rules": retained_orders,
            "retained_pct_of_exact": (
                retained_orders / exact.rows * 100.0 if exact.rows else None
            ),
            "retained_stake": retained_stake,
            "retained_hypothetical_settlement_pnl": retained_settlement_pnl,
            "retained_hypothetical_roi_pct": (
                retained_settlement_pnl / retained_stake * 100.0 if retained_stake else None
            ),
            "retained_scores": retained.result() if retained.rows else None,
            "skipped": dict(skipped),
        }
    )
    return result


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--db", required=True, type=Path)
    parser.add_argument("--stations", required=True, type=Path)
    args = parser.parse_args()

    stations = load_stations(args.stations)
    uri = f"file:{args.db}?mode=ro"
    conn = sqlite3.connect(uri, uri=True)
    conn.execute("PRAGMA query_only=ON")
    timelines, batches_by_id = load_forecasts(conn)
    output = {
        "method": {
            "required_models": list(REQUIRED_MODELS),
            "max_forecast_age_hours": MAX_FORECAST_AGE.total_seconds() / 3600,
            "next_day_only": True,
            "corrected_family_backtest": False,
            "reason": "historical bias-calibration versions were not persisted",
        },
        "snapshots": run_snapshot_backtest(conn, stations, timelines, batches_by_id),
        "hold_orders": run_order_backtest(
            conn, stations, timelines, batches_by_id
        ),
    }
    print(json.dumps(output, ensure_ascii=False, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
