"""Telemetry.

Most of this is format correctness — a series a scraper cannot parse is a series
nobody sees — plus the two decisions that stop a metrics system becoming a
problem of its own: bounded label values, and a scrape that cannot take the
endpoint down.
"""

from __future__ import annotations

import json
import urllib.request

import pytest

from app.runtime.metrics import RATIO_BUCKETS, Metrics
from app.runtime.server import HttpServer, Request, metrics_route
from app.runtime.workflow_stats import WorkflowStats


@pytest.fixture
def metrics() -> Metrics:
    return Metrics()


class TestCounters:
    def test_a_counter_accumulates(self, metrics):
        metrics.counter("things_total")
        metrics.counter("things_total")

        assert metrics.value("things_total") == 2

    def test_labels_separate_series(self, metrics):
        metrics.counter("calls_total", outcome="ok")
        metrics.counter("calls_total", outcome="failed")
        metrics.counter("calls_total", outcome="ok")

        assert metrics.value("calls_total", outcome="ok") == 2
        assert metrics.value("calls_total", outcome="failed") == 1

    def test_label_order_does_not_create_a_second_series(self, metrics):
        metrics.counter("calls_total", a="1", b="2")
        metrics.counter("calls_total", b="2", a="1")

        # Sorted keys, or the same measurement recorded two ways would appear as
        # two unrelated series that never add up.
        assert metrics.value("calls_total", a="1", b="2") == 2

    def test_an_unrecorded_series_reads_as_zero(self, metrics):
        assert metrics.value("never_touched") == 0.0


class TestGauges:
    def test_a_gauge_replaces_rather_than_adds(self, metrics):
        metrics.gauge("queue_depth", 5)
        metrics.gauge("queue_depth", 2)

        assert metrics.value("queue_depth") == 2


class TestHistograms:
    def test_observations_are_counted(self, metrics):
        for value in (0.1, 0.2, 0.3):
            metrics.observe("latency_seconds", value)

        assert metrics.value("latency_seconds") == 3

    def test_buckets_are_cumulative(self, metrics):
        for value in (0.01, 0.2, 5.0):
            metrics.observe("latency_seconds", value)

        rendered = metrics.render()

        # Cumulative: le="1" counts everything at or under a second, not just
        # what fell in that band.
        #
        # Whole numbers render without a decimal point — le="1", not le="1.0" —
        # matching Prometheus' own client. Either parses, but a bucket boundary
        # that rendered differently between scrapes would create two series, so
        # the formatting has to be deterministic.
        assert 'latency_seconds_bucket{le="0.05"} 1' in rendered
        assert 'latency_seconds_bucket{le="1"} 2' in rendered
        assert 'latency_seconds_bucket{le="+Inf"} 3' in rendered

    def test_sum_and_count_are_emitted(self, metrics):
        metrics.observe("latency_seconds", 1.5)
        metrics.observe("latency_seconds", 2.5)

        rendered = metrics.render()

        assert "latency_seconds_sum 4" in rendered
        assert "latency_seconds_count 2" in rendered

    def test_a_fractional_sum_keeps_its_decimals(self, metrics):
        metrics.observe("latency_seconds", 0.25)
        metrics.observe("latency_seconds", 0.5)

        # Only whole numbers lose the point; precision that matters is kept.
        assert "latency_seconds_sum 0.75" in metrics.render()

    def test_confidence_uses_a_ratio_scale(self, metrics):
        metrics.observe("confidence", 0.94, buckets=RATIO_BUCKETS)

        # Second-scale buckets would put every confidence value in the first
        # bucket and show nothing at all.
        assert 'confidence_bucket{le="0.95"} 1' in metrics.render()

    def test_timing_records_a_failed_block(self, metrics):
        with pytest.raises(RuntimeError), metrics.time("work_seconds"):
            raise RuntimeError("boom")

        # 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/stop records nothing for it.
        assert metrics.value("work_seconds") == 1


class TestFormat:
    def test_help_and_type_precede_a_series(self, metrics):
        metrics.counter("things_total", "How many things.")

        rendered = metrics.render()

        assert "# HELP things_total How many things." in rendered
        assert "# TYPE things_total counter" in rendered

    def test_label_values_are_escaped(self, metrics):
        # Label values come from things like error reasons, and a stray quote
        # produces output no scraper can read — which loses every series, not
        # just this one.
        metrics.counter("errors_total", reason='he said "no"')

        assert r'reason="he said \"no\""' in metrics.render()

    def test_a_newline_in_a_label_is_escaped(self, metrics):
        metrics.counter("errors_total", reason="line one\nline two")

        rendered = metrics.render()

        assert "\\n" in rendered
        # One line per series is the format's only structural rule.
        assert len([ln for ln in rendered.splitlines() if ln.startswith("errors_total")]) == 1

    def test_integers_render_without_a_decimal_point(self, metrics):
        metrics.counter("things_total")

        assert "things_total 1" in metrics.render()

    def test_output_ends_with_a_newline(self, metrics):
        metrics.counter("things_total")

        # A scraper reads line-oriented; a truncated final line is a lost series.
        assert metrics.render().endswith("\n")


class TestInstrumentation:
    """The measurements that were actually wired in."""

    def test_ocr_records_confidence_and_outcome(self, tmp_path):
        from app.ocr.engine import OcrResult, TextBlock
        from app.runtime.metrics import METRICS
        from app.tools.document_tools import OcrTool

        METRICS.reset()

        class Engine:
            name = "fake"

            def read(self, path):
                return OcrResult(blocks=[TextBlock(text="hello", confidence=0.9)], engine="fake")

        document = tmp_path / "doc.jpg"
        document.write_bytes(b"x")

        OcrTool(Engine()).run(path=str(document))

        assert METRICS.value("taxpilot_ocr_documents_total", outcome="read") == 1
        assert METRICS.value("taxpilot_ocr_confidence") == 1

    def test_an_unreadable_document_is_counted_separately(self, tmp_path):
        from app.ocr.engine import OcrResult
        from app.runtime.metrics import METRICS
        from app.tools.document_tools import OcrTool

        METRICS.reset()

        class Engine:
            name = "fake"

            def read(self, path):
                return OcrResult(blocks=[], engine="fake")

        document = tmp_path / "doc.jpg"
        document.write_bytes(b"x")

        OcrTool(Engine()).run(path=str(document))

        # The ratio of these two is the number worth watching.
        assert METRICS.value("taxpilot_ocr_documents_total", outcome="unreadable") == 1

    def test_a_failing_job_is_counted_as_failed(self):
        from app.runtime.daemon import Daemon
        from app.runtime.metrics import METRICS

        METRICS.reset()
        daemon = Daemon(max_wait_seconds=0.01)

        def explode():
            raise RuntimeError("boom")

        daemon.add("broken", explode, interval_seconds=0)
        daemon.run(max_iterations=2)

        assert METRICS.value("taxpilot_jobs_total", job="broken", outcome="failed") == 2
        assert METRICS.value("taxpilot_jobs_total", job="broken", outcome="ok") == 0

    def test_a_failed_send_is_counted(self):
        from app.runtime.metrics import METRICS
        from app.whatsapp.messages import OutboundMessage, SendResult
        from app.whatsapp.provider import MessageSender

        METRICS.reset()

        class Provider:
            name = "fake"

            def send(self, message):
                return SendResult(ok=False, error="rejected")

        sender = MessageSender(Provider())
        sender.window.record_inbound("92300")
        sender.send(OutboundMessage(recipient="92300", text="hi"))

        # A message that fails to send is invisible to whoever was waiting for it.
        assert METRICS.value("taxpilot_whatsapp_sent_total", outcome="failed") == 1


class TestWorkflowStats:
    def test_success_rate_is_computed_from_finished_runs(self):
        stats = WorkflowStats(7, started=10, completed=6, failed=2, awaiting=2,
                              median_seconds=30, retried=1)

        # Of what finished, not of what started: runs still awaiting a reviewer
        # have not succeeded or failed, and counting them as failures would
        # blame the AI for a queue nobody has worked through.
        assert stats.success_rate == 75.0

    def test_nothing_finished_reports_nothing_rather_than_zero(self):
        stats = WorkflowStats(7, started=3, completed=0, failed=0, awaiting=3,
                              median_seconds=None, retried=0)

        # "0% success" and "nothing has run" render identically on a dashboard
        # unless the difference survives to it.
        assert stats.success_rate is None

    def test_no_database_means_no_statistics(self):
        from app.runtime.workflow_stats import workflow_stats

        assert workflow_stats(None) is None


class TestMetricsEndpoint:
    @pytest.fixture
    def server(self):
        from app.runtime.metrics import METRICS

        METRICS.reset()
        METRICS.counter("taxpilot_test_total", "A test counter.")

        server = HttpServer("127.0.0.1", 0)
        metrics_route(server)
        server.start()

        yield server

        server.stop()

    def test_it_serves_the_prometheus_format(self, server):
        with urllib.request.urlopen(f"http://127.0.0.1:{server.port}/metrics", timeout=5) as r:
            body = r.read().decode()

            assert r.status == 200
            # text/plain, not JSON: a scraper reads the exposition format.
            assert r.headers["Content-Type"].startswith("text/plain")

        assert "taxpilot_test_total 1" in body

    def test_a_broken_database_does_not_take_the_endpoint_down(self):
        class Container:
            @property
            def connect(self):
                def bad():
                    raise RuntimeError("no database")

                return bad

        server = HttpServer("127.0.0.1", 0)
        metrics_route(server, Container())
        server.start()

        try:
            with urllib.request.urlopen(f"http://127.0.0.1:{server.port}/metrics", timeout=5) as r:
                # Losing one series beats losing the monitoring of everything
                # else — including the alert that would say the database is down.
                assert r.status == 200
        finally:
            server.stop()


class TestLabelCardinality:
    def test_cms_calls_are_labelled_by_endpoint_not_by_path(self):
        """One series per client id would be unbounded.

        The classic way to bring down a monitoring system is a label whose values
        come from user data. `clients/42` and `clients/77` are the same
        operation, so both are recorded as `clients`.
        """
        from app.api.client import CmsClient
        from app.config.settings import Settings
        from app.runtime.metrics import METRICS

        METRICS.reset()
        client = CmsClient(Settings("https://cms.invalid", "k", "s" * 64))

        for client_id in (42, 77, 99):
            try:
                client.get_client(client_id)
            except Exception:  # noqa: BLE001, S110 - unreachable host is the point
                pass

        assert METRICS.value("taxpilot_cms_requests_total",
                             endpoint="clients", outcome="unreachable") == 3
