"""Operational telemetry, in the format every monitoring tool already reads.

Prometheus text format, emitted from the standard library. No client library:
the format is a documented handful of lines, and a dependency to produce it
would cost more than it saves in a codebase whose small surface has been worth
keeping.

## What belongs here, and what does not

This answers **"how is this process doing"** — latency, failures, retries,
durations. Operations.

It deliberately does *not* answer "how is the AI doing at its job". Approval
time, correction rate and rejection reasons live in the CMS, computed from
proposals, because that is the business question and the CMS is where the people
asking it already are. Duplicating them here would give two systems the same
number and two chances to disagree about it.

## Counters reset on restart

Which is correct for this format: a scraper samples every few seconds and
computes rates itself, and `rate()` handles a counter resetting. Durable history
comes from the workflow tables instead — see `workflow_stats`.
"""

from __future__ import annotations

import threading
import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass, field

#: Seconds. Chosen for what this actually does: an API call in tens of
#: milliseconds, OCR in seconds, a whole workflow in tens of seconds.
DEFAULT_BUCKETS = (0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0, 60.0)

#: 0–1, for confidence.
RATIO_BUCKETS = (0.1, 0.3, 0.5, 0.6, 0.7, 0.8, 0.85, 0.9, 0.95, 1.0)

Labels = tuple[tuple[str, str], ...]


def _labels(pairs: dict[str, str] | None) -> Labels:
    """Sorted, so the same labels in a different order are the same series."""
    return tuple(sorted((pairs or {}).items()))


@dataclass
class _Series:
    help_text: str
    kind: str
    values: dict[Labels, float] = field(default_factory=lambda: defaultdict(float))

    # Histograms only.
    buckets: tuple[float, ...] = ()
    counts: dict[Labels, list[int]] = field(default_factory=dict)
    sums: dict[Labels, float] = field(default_factory=lambda: defaultdict(float))
    totals: dict[Labels, int] = field(default_factory=lambda: defaultdict(int))


class Metrics:
    """A registry of series.

    Locked because the HTTP server, the daemon loop and the webhook thread all
    record into it. Without that, a scrape landing mid-update reads a torn value
    — which shows up as an impossible rate and is nearly impossible to
    reproduce.
    """

    def __init__(self) -> None:
        self._series: dict[str, _Series] = {}
        self._lock = threading.Lock()

    # ── Recording ─────────────────────────────────────────────────────────

    def counter(self, name: str, help_text: str = "", **labels: str) -> None:
        """Something happened once more."""
        self.add(name, 1.0, help_text, **labels)

    def add(self, name: str, amount: float, help_text: str = "", **labels: str) -> None:
        with self._lock:
            series = self._series.setdefault(name, _Series(help_text, "counter"))
            series.values[_labels(labels)] += amount

    def gauge(self, name: str, value: float, help_text: str = "", **labels: str) -> None:
        """Something is currently this much."""
        with self._lock:
            series = self._series.setdefault(name, _Series(help_text, "gauge"))
            series.values[_labels(labels)] = value

    def observe(
        self,
        name: str,
        value: float,
        help_text: str = "",
        buckets: tuple[float, ...] = DEFAULT_BUCKETS,
        **labels: str,
    ) -> None:
        """Record a measurement into a histogram."""
        key = _labels(labels)

        with self._lock:
            series = self._series.setdefault(name, _Series(help_text, "histogram", buckets=buckets))
            counts = series.counts.setdefault(key, [0] * len(series.buckets))

            for index, edge in enumerate(series.buckets):
                if value <= edge:
                    counts[index] += 1

            series.sums[key] += value
            series.totals[key] += 1

    @contextmanager
    def time(self, name: str, help_text: str = "", **labels: str):
        """Measure how long a block took, whether or not it succeeded.

        The failure path is the one worth timing: a call that hangs for thirty
        seconds and then errors is the interesting event, and a naive
        `start … call … stop` records nothing at all for it.
        """
        started = time.perf_counter()

        try:
            yield
        finally:
            self.observe(name, time.perf_counter() - started, help_text, **labels)

    # ── Reading ───────────────────────────────────────────────────────────

    def render(self) -> str:
        """The Prometheus text exposition format."""
        lines: list[str] = []

        with self._lock:
            for name in sorted(self._series):
                series = self._series[name]

                if series.help_text:
                    lines.append(f"# HELP {name} {series.help_text}")

                lines.append(f"# TYPE {name} {series.kind}")

                if series.kind == "histogram":
                    lines.extend(self._render_histogram(name, series))
                else:
                    for key, value in sorted(series.values.items()):
                        lines.append(f"{name}{_format_labels(key)} {_number(value)}")

        return "\n".join(lines) + "\n"

    @staticmethod
    def _render_histogram(name: str, series: _Series) -> list[str]:
        lines: list[str] = []

        for key in sorted(series.counts):
            counts = series.counts[key]

            for edge, count in zip(series.buckets, counts, strict=True):
                lines.append(f"{name}_bucket{_format_labels(key, le=_number(edge))} {count}")

            # +Inf is required by the format, and equals the total: everything
            # observed is at most infinity.
            lines.append(f"{name}_bucket{_format_labels(key, le='+Inf')} {series.totals[key]}")
            lines.append(f"{name}_sum{_format_labels(key)} {_number(series.sums[key])}")
            lines.append(f"{name}_count{_format_labels(key)} {series.totals[key]}")

        return lines

    def value(self, name: str, **labels: str) -> float:
        """One series' current value. For tests and for health reporting."""
        with self._lock:
            series = self._series.get(name)

            if series is None:
                return 0.0

            key = _labels(labels)

            if series.kind == "histogram":
                return float(series.totals.get(key, 0))

            return series.values.get(key, 0.0)

    def reset(self) -> None:
        with self._lock:
            self._series.clear()


def _format_labels(key: Labels, **extra: str) -> str:
    pairs = list(key) + sorted(extra.items())

    if not pairs:
        return ""

    inner = ",".join(f'{name}="{_escape(value)}"' for name, value in pairs)

    return "{" + inner + "}"


def _escape(value: str) -> str:
    """A stray quote or newline in a label would produce output no scraper can
    read, and label values come from things like error reasons."""
    return value.replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n")


def _number(value: float) -> str:
    return str(int(value)) if value == int(value) else repr(value)


#: The default registry.
#:
#: Global state, which the rest of this codebase avoids — accepted here because
#: telemetry is genuinely cross-cutting, and threading a registry through every
#: constructor would put an observability concern into the signature of code
#: that has nothing to do with it. It is injectable where it matters, and tests
#: reset it.
METRICS = Metrics()
