"""The shared forecast cache (AV-048, docs/reports/AV-048_raport.md section 4).

One SQLite file for every Journey: what Open-Meteo answered for a place is
kept per *grid cell* -- the route point's coordinates rounded to about the
model's own resolution, its elevation to ELEVATION_BUCKET_M -- with the model,
the variables, the timezone and the dates it covers. Two routes through the
same valley, or the same route opened again tomorrow, ask for the same cells:
no new request while the entry is fresh.

Freshness follows each model's update cycle (TTL_SECONDS_BY_MODEL). An entry
past its TTL is still *served* -- marked stale, with the time it was fetched,
for the page to say "prognoza z HH:MM" -- while a fresh copy is fetched in the
background (stale-while-revalidate, mountain_twin.weather.live), up to
MAX_STALE_SECONDS; older than that it is fetched before answering.

The same file counts the provider calls of each UTC day (the quota's own
weight: per location, per model, variables in tens, days in fortnights) for
the log and the prefetch's breaker, and remembers when a Journey's weather
was last opened (the prefetch's "recently opened" -- a Journey id and a time,
nothing else).

Only provider answers are stored, as they came (ADR-003): no value is
derived, interpolated or merged here.
"""

from __future__ import annotations

import contextvars
import json
import os
import sqlite3
import threading
import time
import zlib
from contextlib import contextmanager
from dataclasses import dataclass
from datetime import date, datetime, timezone
from pathlib import Path
from typing import Any

# Grid cells: coordinates rounded to about the model's own grid. Rounding
# finer than the grid changes nothing the model can tell apart; coarser
# would blur it. The comparison models' grids are ICON (2-13 km), ECMWF IFS
# (0.25 deg) and GFS (0.25 deg).
CELL_DEGREES_BY_MODEL: dict[str | None, float] = {
    None: 0.02,  # best_match -- by region, see best_match_zone()
    "icon_seamless": 0.02,
    "ecmwf_ifs025": 0.1,
    "gfs_seamless": 0.1,
}
DEFAULT_CELL_DEGREES = 0.02
# AV-051: best_match picks the finest model covering a place (Open-Meteo's
# documentation, checked 2026-10-05): ICON-D2 (0.02 deg, a run every 3 h)
# over central Europe, ICON-EU (0.0625 deg, 3 h) over the rest of Europe
# and North Africa, global models (ICON 0.125, ECMWF IFS 0.1, GFS 0.25 deg;
# a run every 6 h) elsewhere. The cell and the freshness follow that model
# -- before, every place got 0.01 deg cells and a 1 h freshness, so a
# Himalayan route asked for the same 10 km global cell dozens of times.
# (AROME over France and MET Nordic over Scandinavia are finer still and
# update more often; ICON-D2's cells and 1 h serve them as before.)
ICON_D2_DOMAIN = (43.18, 58.08, -3.94, 20.34)  # south, north, west, east
ICON_EU_DOMAIN = (29.5, 70.5, -23.5, 62.5)
BEST_MATCH_ZONES = {
    "ICON_D2": {"cell_degrees": 0.02, "ttl_seconds": 3 * 3600.0},
    "ICON_EU": {"cell_degrees": 0.06, "ttl_seconds": 3 * 3600.0},
    "GLOBAL": {"cell_degrees": 0.1, "ttl_seconds": 6 * 3600.0},
}


def best_match_zone(latitude: float | None, longitude: float | None) -> str:
    if latitude is None or longitude is None:
        return "ICON_D2"  # unknown place: the finest cells, the shortest freshness
    for name, (south, north, west, east) in (
        ("ICON_D2", ICON_D2_DOMAIN),
        ("ICON_EU", ICON_EU_DOMAIN),
    ):
        if south <= latitude <= north and west <= longitude <= east:
            return name
    return "GLOBAL"


# Open-Meteo corrects the temperature for the elevation asked for (a route
# point is not its grid cell's mean height). AV-051: 50 m buckets keep that
# within about 0.16 deg C at a standard lapse rate (below any model's error)
# while twice as many points of a steep route share a request.
ELEVATION_BUCKET_M = 50.0
# How long an answer counts as fresh, by model (Open-Meteo's model metadata,
# docs/design_reference/pogoda_spike_v0_1.md): best_match mixes models that
# publish every 1-3 h (MET Nordic hourly, AROME/ICON-D2 3-hourly) -- 1 h;
# ICON runs every 3 h (D2) to 6 h (global) -- 3 h; ECMWF IFS and GFS every
# 6 h -- 6 h. A combination of models is as fresh as its quickest member.
TTL_SECONDS_BY_MODEL: dict[str | None, float] = {
    None: 3600.0,
    "icon_seamless": 3 * 3600.0,
    "ecmwf_ifs025": 6 * 3600.0,
    "gfs_seamless": 6 * 3600.0,
}
DEFAULT_TTL_SECONDS = 3600.0
# Past its TTL an answer is still shown at once (stale) while a fresh one is
# fetched in the background -- for a day at most; older is fetched first.
MAX_STALE_SECONDS = 24 * 3600.0
# Open-Meteo's free tier (non-commercial use): 600 calls a minute, 5 000 an
# hour, 10 000 a day, counted per location and model, fractionally for more
# than 10 variables or more than 2 weeks.
PROVIDER_DAILY_CALL_LIMIT = 10_000.0
# AV-063: the background prefetch (mountain_twin/journey/weather_prefetch.py)
# leaves an answer younger than this alone, whatever its model's TTL: when a
# model's last run is not known, a 3-hour-old forecast is still the newest
# the prefetch can usefully have -- a restart of the server used to refetch
# every cell past 1 h (about 700 calls, 7 % of the day's limit).
PREFETCH_FRESH_SECONDS = 3 * 3600.0
_fresh_at_least: contextvars.ContextVar[float] = contextvars.ContextVar(
    "forecast_fresh_at_least", default=0.0
)


@contextmanager
def fresh_for_at_least(seconds: float):
    """Inside this block (this thread), answers younger than ``seconds``
    count as fresh -- the prefetch's rule; a page's request keeps the TTL."""
    token = _fresh_at_least.set(seconds)
    try:
        yield
    finally:
        _fresh_at_least.reset(token)


def effective_ttl(
    model: str | None, latitude: float | None = None, longitude: float | None = None
) -> float:
    return max(ttl_seconds(model, latitude, longitude), _fresh_at_least.get())


# AV-051 (the 4.10 lessons): every provider request is recorded with what
# asked for it -- a page opening a Journey (OPEN), a place's forecast
# (PLACE), the model comparison (COMPARE), the background prefetch
# (PREFETCH) or a stale answer refreshed in the background (REFRESH).
CALL_SOURCES = ("OPEN", "PLACE", "COMPARE", "PREFETCH", "REFRESH")
_call_source: contextvars.ContextVar[str] = contextvars.ContextVar(
    "provider_call_source", default="OPEN"
)


@contextmanager
def call_source(name: str):
    if name not in CALL_SOURCES:
        raise ValueError(f"unknown call source {name}")
    token = _call_source.set(name)
    try:
        yield
    finally:
        _call_source.reset(token)


def current_call_source() -> str:
    return _call_source.get()


# Which server wrote an event: the development server's port, a test's or a
# measurement's own name -- never mixed (ECHTRO_INSTANCE overrides).
INSTANCE_ENVIRONMENT = "ECHTRO_INSTANCE"


def cell_degrees(
    model: str | None, latitude: float | None = None, longitude: float | None = None
) -> float:
    """The grid-cell size for a model (best_match: its zone's model), or
    for a set of models ("a,b,c": the finest of them)."""
    if model and "," in model:
        return min(cell_degrees(part, latitude, longitude) for part in model.split(","))
    if model is None:
        return BEST_MATCH_ZONES[best_match_zone(latitude, longitude)]["cell_degrees"]
    return CELL_DEGREES_BY_MODEL.get(model, DEFAULT_CELL_DEGREES)


def ttl_seconds(
    model: str | None, latitude: float | None = None, longitude: float | None = None
) -> float:
    if model and "," in model:
        return min(ttl_seconds(part, latitude, longitude) for part in model.split(","))
    if model is None:
        zone = best_match_zone(latitude, longitude)
        # AV-052: every zone at its model's own run -- ICON-D2 publishes every
        # 3 h; the 1 h of AV-048 came from MET Nordic, which is Scandinavia's
        # (the ICON_EU zone here). The usage simulator showed the 1 h making
        # 75-82 % of all Open-Meteo use background refreshes of answers no
        # newer run had replaced.
        return BEST_MATCH_ZONES[zone]["ttl_seconds"]
    return TTL_SECONDS_BY_MODEL.get(model, DEFAULT_TTL_SECONDS)


def grid_cell(
    latitude: float, longitude: float, elevation_m: float | None, model: str | None
) -> tuple[float, float, float | None]:
    """The place actually asked for: the cell's centre and the elevation's
    bucket (None stays None -- the provider's own terrain height)."""
    step = cell_degrees(model, latitude, longitude)
    lat = round(round(latitude / step) * step, 4)
    lon = round(round(longitude / step) * step, 4)
    elevation = (
        None
        if elevation_m is None
        else round(elevation_m / ELEVATION_BUCKET_M) * ELEVATION_BUCKET_M
    )
    return lat, lon, elevation


def call_weight(locations: int, models: int, variables: int, days: int) -> float:
    """What a request weighs in the provider's quota."""
    return locations * models * max(1.0, variables / 10) * max(1.0, days / 14)


@dataclass(frozen=True)
class CachedAnswer:
    start_date: date
    end_date: date
    fetched_at: float
    raw: dict[str, Any]


class ForecastCellCache:
    """Thread-safe (one connection per call); the file is created on first
    use. ``clock`` returns epoch seconds."""

    def __init__(self, path: Path, *, clock=time.time, instance: str | None = None) -> None:
        self.path = Path(path)
        self.clock = clock
        self.instance = instance or os.environ.get(INSTANCE_ENVIRONMENT) or "local"
        self._lock = threading.Lock()
        self._ready = False

    def _connect(self) -> sqlite3.Connection:
        if not self._ready:
            with self._lock:
                if not self._ready:
                    self.path.parent.mkdir(parents=True, exist_ok=True)
                    with sqlite3.connect(self.path, timeout=10) as connection:
                        connection.execute("PRAGMA journal_mode=WAL")
                        connection.executescript(_SCHEMA)
                    self._ready = True
        connection = sqlite3.connect(self.path, timeout=10)
        connection.row_factory = sqlite3.Row
        return connection

    # --- answers ----------------------------------------------------------
    def get_many(self, keys: list[str]) -> dict[str, CachedAnswer]:
        if not keys:
            return {}
        found: dict[str, CachedAnswer] = {}
        with self._connect() as connection:
            for first in range(0, len(keys), 500):
                chunk = keys[first : first + 500]
                rows = connection.execute(
                    f"SELECT cell_key, start_date, end_date, fetched_at, raw FROM forecast_cells "
                    f"WHERE cell_key IN ({','.join('?' * len(chunk))})",
                    chunk,
                ).fetchall()
                for row in rows:
                    found[row["cell_key"]] = CachedAnswer(
                        date.fromisoformat(row["start_date"]),
                        date.fromisoformat(row["end_date"]),
                        row["fetched_at"],
                        json.loads(zlib.decompress(row["raw"])),
                    )
        return found

    def put_many(self, entries: list[tuple[str, date, date, float, dict[str, Any]]]) -> None:
        if not entries:
            return
        with self._connect() as connection:
            connection.executemany(
                "INSERT INTO forecast_cells (cell_key, start_date, end_date, fetched_at, raw) "
                "VALUES (?, ?, ?, ?, ?) ON CONFLICT(cell_key) DO UPDATE SET "
                "start_date=excluded.start_date, end_date=excluded.end_date, "
                "fetched_at=excluded.fetched_at, raw=excluded.raw",
                [
                    (
                        key,
                        start.isoformat(),
                        end.isoformat(),
                        fetched_at,
                        zlib.compress(json.dumps(raw, separators=(",", ":")).encode("utf-8"), 6),
                    )
                    for key, start, end, fetched_at, raw in entries
                ],
            )

    def prune(self, older_than_seconds: float = 2 * MAX_STALE_SECONDS) -> int:
        """Drop answers too old to be shown even as stale."""
        with self._connect() as connection:
            return connection.execute(
                "DELETE FROM forecast_cells WHERE fetched_at < ?",
                (self.clock() - older_than_seconds,),
            ).rowcount

    # --- the provider's quota -----------------------------------------------
    def record_calls(self, weight: float, requests: int = 1) -> dict[str, float]:
        day = _utc_day(self.clock())
        with self._connect() as connection:
            connection.execute(
                "INSERT INTO provider_calls (day, weight, requests) VALUES (?, ?, ?) "
                "ON CONFLICT(day) DO UPDATE SET weight = weight + excluded.weight, "
                "requests = requests + excluded.requests",
                (day, weight, requests),
            )
        return self.calls_today()

    def calls_today(self) -> dict[str, float]:
        day = _utc_day(self.clock())
        with self._connect() as connection:
            row = connection.execute(
                "SELECT weight, requests FROM provider_calls WHERE day = ?", (day,)
            ).fetchone()
        return {
            "day": day,
            "weight": float(row["weight"]) if row else 0.0,
            "requests": int(row["requests"]) if row else 0,
            "limit": PROVIDER_DAILY_CALL_LIMIT,
        }

    # --- every provider request (AV-051) ---------------------------------------
    def record_event(
        self,
        *,
        provider: str,
        outcome: str,
        weight: float,
        requests: int = 1,
        limit_period: str | None = None,
        message: str | None = None,
        source: str | None = None,
    ) -> None:
        """One provider request: OK, RATE_LIMITED (with the limit -- MINUTE,
        HOURLY, DAILY -- and the provider's own words) or ERROR, its weight
        (asked for, whether or not it was counted), the source that asked,
        this server instance."""
        with self._connect() as connection:
            connection.execute(
                "INSERT INTO provider_events (at, instance, provider, source, outcome,"
                " limit_period, message, weight, requests) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)",
                (
                    self.clock(),
                    self.instance,
                    provider,
                    source or current_call_source(),
                    outcome,
                    limit_period,
                    (message or "")[:500] or None,
                    weight,
                    requests,
                ),
            )

    def events_by_hour(self, hours: int = 24, instance: str | None = None) -> list[dict]:
        """Sums per UTC hour, provider, source and outcome over the last
        ``hours`` -- this instance's unless another is named."""
        with self._connect() as connection:
            rows = connection.execute(
                "SELECT strftime('%Y-%m-%dT%H:00Z', at, 'unixepoch') AS hour, provider, source,"
                " outcome, SUM(requests) AS requests, SUM(weight) AS weight"
                " FROM provider_events WHERE at >= ? AND instance = ?"
                " GROUP BY hour, provider, source, outcome ORDER BY hour",
                (self.clock() - hours * 3600, instance or self.instance),
            ).fetchall()
        return [dict(row) for row in rows]

    def last_rate_limit(self, provider: str = "open_meteo") -> dict | None:
        """The newest 429 of this instance: when, which limit, the
        provider's words."""
        with self._connect() as connection:
            row = connection.execute(
                "SELECT at, limit_period, message, source FROM provider_events"
                " WHERE provider = ? AND outcome = 'RATE_LIMITED' AND instance = ?"
                " ORDER BY at DESC LIMIT 1",
                (provider, self.instance),
            ).fetchone()
        return dict(row) if row else None

    # --- recently opened Journeys (the prefetch's list) -----------------------
    def note_opened(self, journey_id: str) -> None:
        with self._connect() as connection:
            connection.execute(
                "INSERT INTO opened_journeys (journey_id, opened_at) VALUES (?, ?) "
                "ON CONFLICT(journey_id) DO UPDATE SET opened_at = excluded.opened_at",
                (journey_id, self.clock()),
            )

    # --- the prefetch's last cycle (AV-063: a restart does not repeat it) ----
    def note_prefetch(self) -> None:
        with self._connect() as connection:
            connection.execute(
                "INSERT INTO prefetch_runs (name, finished_at) VALUES ('weather', ?) "
                "ON CONFLICT(name) DO UPDATE SET finished_at = excluded.finished_at",
                (self.clock(),),
            )

    def last_prefetch(self) -> float | None:
        with self._connect() as connection:
            row = connection.execute(
                "SELECT finished_at FROM prefetch_runs WHERE name = 'weather'"
            ).fetchone()
        return float(row["finished_at"]) if row else None

    def opened_since(self, seconds: float) -> list[str]:
        with self._connect() as connection:
            rows = connection.execute(
                "SELECT journey_id FROM opened_journeys WHERE opened_at >= ? "
                "ORDER BY opened_at DESC",
                (self.clock() - seconds,),
            ).fetchall()
        return [row["journey_id"] for row in rows]


def cell_key(
    *,
    endpoint: str,
    model: str | None,
    variables: tuple[str, ...],
    cell: tuple[float, float, float | None],
    timezone_name: str,
) -> str:
    return json.dumps(
        [endpoint, model, list(variables), list(cell), timezone_name], separators=(",", ":")
    )


def _utc_day(epoch_seconds: float) -> str:
    return datetime.fromtimestamp(epoch_seconds, timezone.utc).date().isoformat()


_SCHEMA = """
CREATE TABLE IF NOT EXISTS forecast_cells (
    cell_key TEXT PRIMARY KEY,
    start_date TEXT NOT NULL,
    end_date TEXT NOT NULL,
    fetched_at REAL NOT NULL,
    raw BLOB NOT NULL
);
CREATE TABLE IF NOT EXISTS provider_calls (
    day TEXT PRIMARY KEY,
    weight REAL NOT NULL,
    requests INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS opened_journeys (
    journey_id TEXT PRIMARY KEY,
    opened_at REAL NOT NULL
);
CREATE TABLE IF NOT EXISTS provider_events (
    id INTEGER PRIMARY KEY AUTOINCREMENT,
    at REAL NOT NULL,
    instance TEXT NOT NULL,
    provider TEXT NOT NULL,
    source TEXT NOT NULL,
    outcome TEXT NOT NULL,
    limit_period TEXT,
    message TEXT,
    weight REAL NOT NULL,
    requests INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS provider_events_at ON provider_events(at);
CREATE TABLE IF NOT EXISTS prefetch_runs (
    name TEXT PRIMARY KEY,
    finished_at REAL NOT NULL
);
"""
