"""Usage of the outside services (AV-052): every request that leaves this
server, counted in one place -- mountain_twin.offline.urlopen, the single
function every client opens its URLs through (enforced by
tests/test_service_budgets.py) -- per UTC hour, service, application
feature, outcome and whether a cache answered instead.

Only numbers: no route, place, coordinate or user is ever written here.

The feature comes from the calling context (``feature(...)``; the weather's
call_source maps onto it); the weight from ``request_weight(...)`` where a
service counts more than one per request (Open-Meteo: per location and
variable), else 1. A cache that answers instead of a request says so with
``record_cache_hit`` -- the share of cache hits per service.

Nothing is counted until a store is configured (the server does at start;
tests and scripts that want numbers configure their own).
"""

from __future__ import annotations

import contextvars
import sqlite3
import threading
import time
from contextlib import contextmanager
from datetime import datetime, timezone
from pathlib import Path

FEATURES = (
    "weather_open",
    "weather_place",
    "weather_compare",
    "weather_prefetch",
    "weather_refresh",
    "place_search",
    "place_osm_tags",
    "radar",
    "clouds",
    "terrain",
    "routing",
    "live_tracking",
    "strava",
    "other",
)
# The weather's call sources (mountain_twin.weather.forecast_cache) as features.
_FROM_CALL_SOURCE = {
    "OPEN": "weather_open",
    "PLACE": "weather_place",
    "COMPARE": "weather_compare",
    "PREFETCH": "weather_prefetch",
    "REFRESH": "weather_refresh",
}
# Features the user did not ask for at that moment (the background).
BACKGROUND_FEATURES = frozenset({"weather_prefetch", "weather_refresh"})
EXPENSIVE_FEATURES = frozenset({"weather_compare"})

_feature: contextvars.ContextVar[str | None] = contextvars.ContextVar("usage_feature", default=None)
_weight: contextvars.ContextVar[float | None] = contextvars.ContextVar("usage_weight", default=None)


@contextmanager
def feature(name: str):
    if name not in FEATURES:
        raise ValueError(f"unknown feature {name}")
    token = _feature.set(name)
    try:
        yield
    finally:
        _feature.reset(token)


# Where no feature is named, the service says what it was for.
_BY_SERVICE = {
    "photon": "place_search",
    "nominatim": "place_osm_tags",
    "rainviewer": "radar",
    "eumetsat": "clouds",
    "copernicus_dem": "terrain",
    "brouter_public": "routing",
    "own_server": "routing",
    "strava": "strava",
}


def current_feature(service: str | None = None) -> str:
    named = _feature.get()
    if named:
        return named
    if service in _BY_SERVICE:
        return _BY_SERVICE[service]
    from mountain_twin.weather.forecast_cache import current_call_source

    return _FROM_CALL_SOURCE.get(current_call_source(), "other")


@contextmanager
def request_weight(weight: float):
    token = _weight.set(weight)
    try:
        yield
    finally:
        _weight.reset(token)


def current_weight() -> float:
    weight = _weight.get()
    return 1.0 if weight is None else weight


_SCHEMA = """
CREATE TABLE IF NOT EXISTS usage_hourly (
    hour TEXT NOT NULL,
    instance TEXT NOT NULL,
    service TEXT NOT NULL,
    feature TEXT NOT NULL,
    outcome TEXT NOT NULL,
    cache_hit INTEGER NOT NULL,
    requests INTEGER NOT NULL,
    weight REAL NOT NULL,
    PRIMARY KEY (hour, instance, service, feature, outcome, cache_hit)
);
CREATE TABLE IF NOT EXISTS usage_errors (
    at REAL NOT NULL,
    instance TEXT NOT NULL,
    service TEXT NOT NULL,
    feature TEXT NOT NULL,
    outcome TEXT NOT NULL,
    detail TEXT
);
CREATE TABLE IF NOT EXISTS alarms_sent (
    day TEXT NOT NULL,
    instance TEXT NOT NULL,
    service TEXT NOT NULL,
    kind TEXT NOT NULL,
    at REAL NOT NULL,
    PRIMARY KEY (day, instance, service, kind)
);
"""


def _hour(epoch: float) -> str:
    return datetime.fromtimestamp(epoch, timezone.utc).strftime("%Y-%m-%dT%H:00Z")


def _day(epoch: float) -> str:
    return datetime.fromtimestamp(epoch, timezone.utc).strftime("%Y-%m-%d")


class UsageStore:
    """``clock`` returns epoch seconds (a simulation moves it)."""

    def __init__(self, path: Path, *, instance: str = "local", clock=time.time):
        self.path, self.instance, self.clock = Path(path), instance, clock
        self._ready = False
        self._lock = threading.Lock()
        self.listeners: list = []  # called with (service, feature, outcome) after each record

    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

    def record(
        self,
        service: str,
        *,
        feature: str | None = None,
        outcome: str = "OK",
        weight: float = 1.0,
        requests: int = 1,
        cache_hit: bool = False,
        detail: str | None = None,
    ) -> None:
        now, feature = self.clock(), feature or current_feature(service)
        with self._connect() as connection:
            connection.execute(
                "INSERT INTO usage_hourly VALUES (?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT"
                " (hour, instance, service, feature, outcome, cache_hit) DO UPDATE SET"
                " requests = requests + excluded.requests, weight = weight + excluded.weight",
                (
                    _hour(now),
                    self.instance,
                    service,
                    feature,
                    outcome,
                    int(cache_hit),
                    requests,
                    weight,
                ),
            )
            if outcome != "OK" and not outcome.startswith("REFUSED"):
                connection.execute(
                    "INSERT INTO usage_errors VALUES (?, ?, ?, ?, ?, ?)",
                    (now, self.instance, service, feature, outcome, (detail or "")[:300] or None),
                )
        for listener in list(self.listeners):
            listener(service, feature, outcome)

    # --- reading ----------------------------------------------------------------
    def today(self, service: str) -> dict[str, float]:
        """The UTC day's network use of a service: requests and weight
        (cache hits left out)."""
        return self._sum(service, _day(self.clock()) + "T", 24)

    def this_hour(self, service: str) -> dict[str, float]:
        return self._sum(service, _hour(self.clock()), 1)

    def _sum(self, service: str, prefix: str, _hours: int) -> dict[str, float]:
        with self._connect() as connection:
            row = connection.execute(
                "SELECT COALESCE(SUM(requests), 0) AS requests, COALESCE(SUM(weight), 0) AS weight"
                " FROM usage_hourly WHERE service = ? AND instance = ? AND cache_hit = 0"
                " AND outcome = 'OK' AND hour LIKE ?",
                (service, self.instance, prefix + "%"),
            ).fetchone()
        return {"requests": int(row["requests"]), "weight": float(row["weight"])}

    def hourly(self, days: int = 7) -> list[dict]:
        since = _hour(self.clock() - days * 86400)
        with self._connect() as connection:
            rows = connection.execute(
                "SELECT hour, service, feature, outcome, cache_hit, requests, weight FROM usage_hourly"
                " WHERE instance = ? AND hour >= ? ORDER BY hour",
                (self.instance, since),
            ).fetchall()
        return [dict(row) for row in rows]

    def recent_errors(self, limit: int = 20) -> list[dict]:
        with self._connect() as connection:
            rows = connection.execute(
                "SELECT at, service, feature, outcome, detail FROM usage_errors WHERE instance = ?"
                " ORDER BY at DESC LIMIT ?",
                (self.instance, limit),
            ).fetchall()
        return [dict(row) for row in rows]

    def errors_since(self, service: str, outcome: str, seconds: float) -> int:
        with self._connect() as connection:
            return connection.execute(
                "SELECT COUNT(*) FROM usage_errors WHERE instance = ? AND service = ? AND outcome = ?"
                " AND at >= ?",
                (self.instance, service, outcome, self.clock() - seconds),
            ).fetchone()[0]

    def mean_hourly_weight(self, service: str, days: int = 7) -> float:
        """The mean network weight per hour over the past ``days`` (hours
        with no use count as zero)."""
        since = _hour(self.clock() - days * 86400)
        with self._connect() as connection:
            total = connection.execute(
                "SELECT COALESCE(SUM(weight), 0) FROM usage_hourly WHERE instance = ? AND service = ?"
                " AND cache_hit = 0 AND hour >= ? AND hour < ?",
                (self.instance, service, since, _hour(self.clock())),
            ).fetchone()[0]
        return float(total) / (days * 24)

    def mark_alarm(self, service: str, kind: str) -> bool:
        """True the first time today (this kind of alarm for this service)."""
        with self._connect() as connection:
            cursor = connection.execute(
                "INSERT OR IGNORE INTO alarms_sent VALUES (?, ?, ?, ?, ?)",
                (_day(self.clock()), self.instance, service, kind, self.clock()),
            )
        return cursor.rowcount == 1


_store: UsageStore | None = None


def configure(store: UsageStore | None) -> UsageStore | None:
    """The process's store (None: nothing is counted)."""
    global _store
    previous, _store = _store, store
    return previous


def store() -> UsageStore | None:
    return _store


def record_cache_hit(service: str, *, weight: float = 1.0, requests: int = 1) -> None:
    """A cache answered what would otherwise have been a request."""
    if _store is not None:
        _store.record(service, cache_hit=True, weight=weight, requests=requests)
