"""Tests for circuit breakers in src.resilience.breakers (TEXTIQ-394 / GAP-031).

Covers SYS-NFR-20: 5 consecutive failures -> open 60s -> half-open trial.

These tests use isolated pybreaker.CircuitBreaker instances rather than the
production ``milvus_breaker``/``s3_breaker``/``llm_breaker`` singletons so
they can use a short ``reset_timeout`` and run in <1s. They share the
production ``_MetricsListener``, which proves that listener wiring updates
the Prometheus gauge correctly. State assertions read pybreaker's
``current_state`` directly to avoid races on time-based transitions.
"""

from __future__ import annotations

import time

import pybreaker
import pytest

from src.resilience.breakers import (
    _MetricsListener,
    async_breaker_call,
    breaker_call,
    circuit_breaker_state,
    llm_breaker,
    milvus_breaker,
    s3_breaker,
)


_LISTENER = _MetricsListener()


def _gauge_value(upstream: str) -> float | None:
    for fam in circuit_breaker_state.collect():
        for sample in fam.samples:
            if sample.labels.get("upstream") == upstream:
                return sample.value
    return None


@pytest.fixture
def breaker():
    """Fresh test breaker with short reset_timeout so tests don't sleep 60s."""
    name = "test_breaker"
    b = pybreaker.CircuitBreaker(
        fail_max=5,
        reset_timeout=0.05,
        listeners=[_LISTENER],
        name=name,
        throw_new_error_on_trip=False,
    )
    yield b
    # Reset gauge label so subsequent tests start clean
    circuit_breaker_state.labels(upstream=name).set(0)


def _boom():
    raise RuntimeError("upstream down")


def _ok():
    return "ok"


# ---------------------------------------------------------------------------
# Production singletons exist with the right config
# ---------------------------------------------------------------------------


class TestProductionBreakers:
    def test_three_breakers_configured_per_sys_nfr_20(self):
        for b in (milvus_breaker, s3_breaker, llm_breaker):
            assert b.fail_max == 5
            assert b.reset_timeout == 60

    def test_breakers_named_per_upstream(self):
        assert milvus_breaker.name == "milvus"
        assert s3_breaker.name == "s3"
        assert llm_breaker.name == "llm"

    def test_initial_gauge_seeded_to_closed(self):
        for upstream in ("milvus", "s3", "llm"):
            assert _gauge_value(upstream) is not None, f"{upstream} gauge not registered"


# ---------------------------------------------------------------------------
# State transitions
# ---------------------------------------------------------------------------


class TestTransitions:
    def test_closed_state_passes_through(self, breaker):
        assert breaker.current_state == "closed"
        assert breaker_call(breaker, _ok) == "ok"
        assert breaker.current_state == "closed"

    def test_opens_after_5_consecutive_failures(self, breaker):
        for _ in range(5):
            with pytest.raises(RuntimeError):
                breaker_call(breaker, _boom)
        assert breaker.current_state == "open"

    def test_does_not_open_at_4_failures(self, breaker):
        for _ in range(4):
            with pytest.raises(RuntimeError):
                breaker_call(breaker, _boom)
        assert breaker.current_state == "closed"

    def test_open_blocks_calls_with_CircuitBreakerError(self, breaker):
        for _ in range(5):
            with pytest.raises(RuntimeError):
                breaker_call(breaker, _boom)
        assert breaker.current_state == "open"
        # Subsequent call short-circuits — _ok is NOT invoked
        called = []

        def _flag():
            called.append(True)
            return "should_not_run"

        with pytest.raises(pybreaker.CircuitBreakerError):
            breaker_call(breaker, _flag)
        assert called == []

    def test_half_open_success_closes(self, breaker):
        for _ in range(5):
            with pytest.raises(RuntimeError):
                breaker_call(breaker, _boom)
        assert breaker.current_state == "open"
        time.sleep(breaker.reset_timeout + 0.02)
        # First call after reset_timeout is the half-open trial; success closes.
        assert breaker_call(breaker, _ok) == "ok"
        assert breaker.current_state == "closed"

    def test_half_open_failure_reopens(self, breaker):
        for _ in range(5):
            with pytest.raises(RuntimeError):
                breaker_call(breaker, _boom)
        assert breaker.current_state == "open"
        time.sleep(breaker.reset_timeout + 0.02)
        # Half-open trial fails -> immediate reopen.
        with pytest.raises(RuntimeError):
            breaker_call(breaker, _boom)
        assert breaker.current_state == "open"


# ---------------------------------------------------------------------------
# Gauge metric reflects state
# ---------------------------------------------------------------------------


class TestMetricGauge:
    def test_gauge_set_to_open_after_trip(self, breaker):
        assert _gauge_value(breaker.name) == 0
        for _ in range(5):
            with pytest.raises(RuntimeError):
                breaker_call(breaker, _boom)
        assert _gauge_value(breaker.name) == 2

    def test_gauge_returns_to_closed_after_recovery(self, breaker):
        for _ in range(5):
            with pytest.raises(RuntimeError):
                breaker_call(breaker, _boom)
        assert _gauge_value(breaker.name) == 2
        time.sleep(breaker.reset_timeout + 0.02)
        breaker_call(breaker, _ok)
        assert _gauge_value(breaker.name) == 0


# ---------------------------------------------------------------------------
# async_breaker_call
# ---------------------------------------------------------------------------


class TestAsyncBreakerCall:
    @pytest.mark.asyncio
    async def test_async_call_passes_through(self, breaker):
        async def _aok():
            return 42

        result = await async_breaker_call(breaker, _aok)
        assert result == 42
        assert breaker.current_state == "closed"

    @pytest.mark.asyncio
    async def test_async_call_propagates_failure(self, breaker):
        async def _aboom():
            raise ValueError("nope")

        with pytest.raises(ValueError):
            await async_breaker_call(breaker, _aboom)
