#!/usr/bin/env python3
"""Compare WeatherNext 2 ensemble mean with the stored fixed-lead baseline.

The Previous Runs API exposes the WeatherNext ensemble *mean*, not historical
members. This tool therefore evaluates station-level daily-extreme MAE only;
it deliberately does not claim a probability or trading-P&L backtest.
"""

from __future__ import annotations

import argparse
import json
import sqlite3
import statistics
import tomllib
import urllib.parse
import urllib.request
from collections import defaultdict, deque
from pathlib import Path

API = "https://previous-runs-api.open-meteo.com/v1/forecast"
MODEL = "google_weathernext2_ensemble_mean"
LEADS = (1, 2, 3)


def fetch_chunk(stations: list[tuple[str, dict]], past_days: int) -> list[dict]:
    variables = ",".join(f"temperature_2m_previous_day{lead}" for lead in LEADS)
    query = urllib.parse.urlencode(
        {
            "latitude": ",".join(str(value["lat"]) for _, value in stations),
            "longitude": ",".join(str(value["lon"]) for _, value in stations),
            "hourly": variables,
            "models": MODEL,
            "past_days": past_days,
            "forecast_days": 1,
            "timezone": "auto",
        }
    )
    request = urllib.request.Request(f"{API}?{query}", headers={"User-Agent": "weatherbot-research/0.1"})
    with urllib.request.urlopen(request, timeout=120) as response:
        payload = json.load(response)
    if isinstance(payload, dict):
        payload = [payload]
    if len(payload) != len(stations):
        raise RuntimeError(f"expected {len(stations)} locations, got {len(payload)}")
    return payload


def load_or_fetch(stations_path: str, past_days: int, cache_path: str | None) -> dict:
    if cache_path and Path(cache_path).exists():
        return json.loads(Path(cache_path).read_text())
    station_map = tomllib.loads(Path(stations_path).read_text())["stations"]
    items = list(station_map.items())
    output: dict[str, dict] = {}
    for start in range(0, len(items), 10):
        chunk = items[start : start + 10]
        for (code, _), payload in zip(chunk, fetch_chunk(chunk, past_days), strict=True):
            output[code] = payload
    if cache_path:
        Path(cache_path).write_text(json.dumps(output, separators=(",", ":")))
    return output


def daily_extremes(payloads: dict) -> dict[tuple[str, str, int, str], float]:
    result = {}
    for station, payload in payloads.items():
        hourly = payload["hourly"]
        times = hourly["time"]
        for lead in LEADS:
            values = hourly[f"temperature_2m_previous_day{lead}"]
            by_date: dict[str, list[float]] = defaultdict(list)
            for timestamp, value in zip(times, values, strict=True):
                if value is not None:
                    by_date[timestamp[:10]].append(float(value))
            for date, day_values in by_date.items():
                result[(station, "highest", lead, date)] = max(day_values)
                result[(station, "lowest", lead, date)] = min(day_values)
    return result


def add_walk_forward(rows: list[dict], field: str) -> None:
    history: dict[tuple[str, str, int], deque[float]] = defaultdict(lambda: deque(maxlen=30))
    for row in sorted(rows, key=lambda item: (item["target_date"], item["station"], item["kind"], item["lead"])):
        key = (row["station"], row["kind"], row["lead"])
        errors = history[key]
        row[f"{field}_corr"] = row[field] - statistics.fmean(errors) if len(errors) >= 15 else None
        errors.append(row[field] - row["actual"])


def metrics(rows: list[dict]) -> tuple:
    def mae(field: str) -> float | None:
        values = [abs(row[field] - row["actual"]) for row in rows if row.get(field) is not None]
        return statistics.fmean(values) if values else None

    def fmt(value: float | None) -> str:
        return "-" if value is None else f"{value:.3f}"

    return (
        len(rows),
        fmt(mae("baseline")),
        fmt(mae("weathernext")),
        fmt(mae("blend")),
        fmt(mae("baseline_corr")),
        fmt(mae("weathernext_corr")),
    )


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("db")
    parser.add_argument("--stations", default="stations.toml")
    parser.add_argument("--past-days", type=int, default=92)
    parser.add_argument("--cache")
    args = parser.parse_args()

    wn = daily_extremes(load_or_fetch(args.stations, args.past_days, args.cache))
    conn = sqlite3.connect(f"file:{args.db}?mode=ro", uri=True)
    rows = []
    for station, kind, lead, date, baseline, actual in conn.execute(
        """SELECT station,kind,lead_days,target_date,forecast,actual
           FROM calibration_observations ORDER BY target_date,station,kind,lead_days"""
    ):
        value = wn.get((station, kind, lead, date))
        if value is None:
            continue
        rows.append(
            {
                "station": station,
                "kind": kind,
                "lead": lead,
                "target_date": date,
                "baseline": baseline,
                "weathernext": value,
                "blend": (baseline + value) / 2,
                "actual": actual,
            }
        )
    add_walk_forward(rows, "baseline")
    add_walk_forward(rows, "weathernext")

    print("scope                 n  base_raw  wn_raw  blend  base_wf  wn_wf")
    scopes = [("all", rows)]
    scopes += [(f"{kind} D+{lead}", [r for r in rows if r["kind"] == kind and r["lead"] == lead])
               for kind in ("highest", "lowest") for lead in LEADS]
    scopes += [("holdout >=07-29", [r for r in rows if r["target_date"] >= "2026-07-29"])]
    for label, sample in scopes:
        n, *values = metrics(sample)
        print(f"{label:<20} {n:4d} " + "  ".join(f"{value:>7}" for value in values))


if __name__ == "__main__":
    main()
