"""Tests for the SRE console.

Run against your own work:      pytest
Run against the solution:       LAB_MODULE=solution pytest
"""

from __future__ import annotations

import importlib
import os
from typing import List, Sequence

import pytest

lab = importlib.import_module(os.environ.get("LAB_MODULE", "lab"))

Outcome = lab.Outcome
RequestEvent = lab.RequestEvent
ValidityPredicate = lab.ValidityPredicate
availability_sli = lab.availability_sli
latency_sli = lab.latency_sli
percentile = lab.percentile
Slo = lab.Slo
ErrorBudget = lab.ErrorBudget
allocate_budget = lab.allocate_budget
burn_rate_threshold = lab.burn_rate_threshold
observed_burn_rate = lab.observed_burn_rate
BurnRateRule = lab.BurnRateRule
STANDARD_LADDER = lab.STANDARD_LADDER
BurnRateAlerting = lab.BurnRateAlerting
SpanKind = lab.SpanKind
Span = lab.Span
Tracer = lab.Tracer
SpanTree = lab.SpanTree
LabelSpec = lab.LabelSpec
MetricSpec = lab.MetricSpec
FORBIDDEN_LABELS = lab.FORBIDDEN_LABELS
CardinalityBudget = lab.CardinalityBudget
DegradationStep = lab.DegradationStep
STANDARD_LADDER_STEPS = lab.STANDARD_LADDER_STEPS
DegradationLadder = lab.DegradationLadder
cost_report = lab.cost_report
CostCircuitBreaker = lab.CostCircuitBreaker
forecast_capacity = lab.forecast_capacity
Baseline = lab.Baseline
classify_regression = lab.classify_regression


# ======================================================================================
# helpers
# ======================================================================================


def clock(start: int = 0, step: int = 1):
    state = {"t": start - step}

    def now() -> int:
        state["t"] += step
        return state["t"]

    return now


def ev(tick: int = 0, outcome: Outcome = Outcome.SUCCESS, latency: int = 100,
       **kw) -> RequestEvent:
    base = dict(tick=tick, tenant="w", deployment="d", outcome=outcome,
                latency_ms=latency)
    base.update(kw)
    return RequestEvent(**base)


def run(successes: int, failures: int = 0, *, start: int = 0,
        outcome: Outcome = Outcome.PLATFORM_ERROR) -> List[RequestEvent]:
    events = [ev(start + i) for i in range(successes)]
    events += [ev(start + successes + i, outcome) for i in range(failures)]
    return events


# ======================================================================================
# 1. the validity predicate
# ======================================================================================


def test_the_default_predicate_excludes_client_errors():
    assert not ValidityPredicate().is_valid(ev(outcome=Outcome.CLIENT_ERROR))


def test_the_default_predicate_excludes_user_aborts():
    assert not ValidityPredicate().is_valid(ev(outcome=Outcome.USER_ABORTED))


def test_the_default_predicate_excludes_synthetic_probes():
    assert not ValidityPredicate().is_valid(ev(synthetic=True))


def test_a_safety_block_is_valid_even_though_it_is_not_good():
    assert ValidityPredicate().is_valid(ev(outcome=Outcome.SAFETY_BLOCKED))


def test_a_safety_block_is_not_counted_as_good():
    sli = availability_sli([ev(outcome=Outcome.SAFETY_BLOCKED)])
    assert sli.good == 0 and sli.valid == 1


def test_an_upstream_error_counts_against_us():
    sli = availability_sli([ev(outcome=Outcome.UPSTREAM_ERROR)])
    assert sli.valid == 1 and sli.good == 0


def test_exclusions_can_be_switched_off():
    predicate = ValidityPredicate(False, False, False)
    assert predicate.is_valid(ev(outcome=Outcome.CLIENT_ERROR))


def test_the_predicate_can_scope_to_tenants():
    predicate = ValidityPredicate(tenants=("retail",))
    assert not predicate.is_valid(ev())
    assert predicate.is_valid(ev(tenant="retail"))


def test_the_predicate_can_scope_to_deployments():
    predicate = ValidityPredicate(deployments=("other",))
    assert not predicate.is_valid(ev())


def test_the_predicate_describes_itself():
    assert "4xx" in ValidityPredicate().describe()


# ======================================================================================
# 2. SLIs
# ======================================================================================


def test_a_perfect_window_is_one():
    assert availability_sli(run(10)).ratio == 1.0


def test_the_ratio_is_good_over_valid():
    sli = availability_sli(run(9, 1))
    assert sli.ratio == pytest.approx(0.9) and sli.bad == 1


def test_an_empty_window_is_one_not_zero():
    assert availability_sli([]).ratio == 1.0


def test_a_window_of_only_excluded_events_is_one():
    events = [ev(outcome=Outcome.CLIENT_ERROR) for _ in range(5)]
    assert availability_sli(events).ratio == 1.0


def test_latency_sli_counts_requests_under_the_threshold():
    events = [ev(latency=100), ev(latency=3_000)]
    assert latency_sli(events, threshold_ms=1_000).ratio == pytest.approx(0.5)


def test_the_latency_threshold_is_inclusive():
    assert latency_sli([ev(latency=1_000)], threshold_ms=1_000).ratio == 1.0


def test_a_failed_request_is_not_good_however_fast_it_was():
    events = [ev(outcome=Outcome.PLATFORM_ERROR, latency=1)]
    assert latency_sli(events, threshold_ms=1_000).ratio == 0.0


def test_percentiles_use_nearest_rank():
    assert percentile([1, 2, 3, 4, 5, 6, 7, 8, 9, 10], 0.5) == 5.0
    assert percentile([1, 2, 3, 4, 5, 6, 7, 8, 9, 10], 0.9) == 9.0


def test_an_empty_percentile_is_zero():
    assert percentile([], 0.99) == 0.0


def test_the_sli_reports_the_predicate_it_used():
    assert availability_sli(run(1)).predicate


# ======================================================================================
# 3. error budgets
# ======================================================================================


SLO = Slo("availability", 0.99, window_ticks=1_000)


def test_a_perfect_window_leaves_the_whole_budget():
    assert ErrorBudget(SLO).state(run(1_000)).remaining_fraction == pytest.approx(1.0)


def test_spending_half_the_budget_leaves_half():
    state = ErrorBudget(SLO).state(run(995, 5))
    assert state.remaining_fraction == pytest.approx(0.5)


def test_a_budget_consumed_exactly_is_exhausted():
    state = ErrorBudget(SLO).state(run(990, 10))
    assert state.exhausted and state.remaining_fraction == pytest.approx(0.0)


def test_the_remaining_fraction_never_goes_negative():
    state = ErrorBudget(SLO).state(run(900, 100))
    assert state.remaining_fraction == 0.0


def test_overspend_is_reported_separately():
    state = ErrorBudget(SLO).state(run(900, 100))
    assert state.overspent_by == 90


def test_a_budget_within_target_reports_no_overspend():
    assert ErrorBudget(SLO).state(run(995, 5)).overspent_by == 0


def test_an_empty_window_is_not_exhausted():
    assert not ErrorBudget(SLO).state([]).exhausted


def test_a_latency_slo_budgets_on_latency():
    slo = Slo("latency", 0.99, window_ticks=1_000, threshold_ms=1_000)
    events = [ev(i, latency=100) for i in range(990)] + \
             [ev(990 + i, latency=5_000) for i in range(10)]
    assert ErrorBudget(slo).state(events).exhausted


def test_the_budget_state_reports_the_achieved_ratio():
    assert ErrorBudget(SLO).state(run(99, 1)).achieved == pytest.approx(0.99)


def test_allocation_sums_to_the_total():
    allocation = allocate_budget(0.01, {"a": 1, "b": 1, "c": 2})
    assert sum(allocation.values()) == pytest.approx(0.01)


def test_allocation_is_proportional_to_the_weights():
    allocation = allocate_budget(0.01, {"a": 1, "b": 3})
    assert allocation["b"] == pytest.approx(allocation["a"] * 3)


def test_allocation_rejects_zero_weights():
    with pytest.raises(ValueError):
        allocate_budget(0.01, {"a": 0})


# ======================================================================================
# 4. burn rate
# ======================================================================================


def test_the_famous_threshold_is_derived_not_magic():
    assert burn_rate_threshold(0.02, 1.0) == pytest.approx(14.4)


def test_the_six_hour_threshold_is_six():
    assert burn_rate_threshold(0.05, 6.0) == pytest.approx(6.0)


def test_the_one_day_threshold_is_three():
    assert burn_rate_threshold(0.10, 24.0) == pytest.approx(3.0)


def test_a_longer_window_for_the_same_budget_gives_a_lower_threshold():
    assert burn_rate_threshold(0.02, 24.0) < burn_rate_threshold(0.02, 1.0)


def test_a_zero_window_is_an_error():
    with pytest.raises(ValueError):
        burn_rate_threshold(0.02, 0)


def test_a_perfect_window_burns_nothing():
    assert observed_burn_rate(availability_sli(run(100)), SLO) == 0.0


def test_burning_exactly_at_target_is_one():
    assert observed_burn_rate(availability_sli(run(99, 1)), SLO) == pytest.approx(1.0)


def test_ten_times_the_error_rate_is_ten_times_the_burn():
    assert observed_burn_rate(availability_sli(run(90, 10)), SLO) == pytest.approx(10.0)


def test_an_empty_window_burns_nothing():
    assert observed_burn_rate(availability_sli([]), SLO) == 0.0


# ======================================================================================
# 5. multi-window alerting
# ======================================================================================


def alerting(target: float = 0.99, **kw) -> BurnRateAlerting:
    return BurnRateAlerting(Slo("a", target, window_ticks=43_200), **kw)


def sustained_failure(ticks: int = 70) -> List[RequestEvent]:
    return [ev(i, Outcome.PLATFORM_ERROR) for i in range(ticks)]


def test_a_healthy_window_fires_nothing():
    assert alerting().firing(run(500), now=500) == []


def test_a_sustained_outage_fires_the_fast_rule():
    firing = {d.rule for d in alerting().firing(sustained_failure(), now=70)}
    assert "fast-burn" in firing


def test_a_recovered_short_window_stops_the_alert():
    events = sustained_failure(30) + [ev(30 + i) for i in range(200)]
    firing = {d.rule for d in alerting().firing(events, now=230)}
    assert "fast-burn" not in firing


def test_the_long_window_alone_does_not_fire():
    events = sustained_failure(30) + [ev(30 + i) for i in range(200)]
    decisions = {d.rule: d for d in alerting().evaluate(events, now=230)}
    assert not decisions["medium-burn"].firing
    assert "recovered" in decisions["medium-burn"].reason


def test_a_low_traffic_window_does_not_fire_on_one_failure():
    events = [ev(0), ev(1, Outcome.PLATFORM_ERROR)]
    assert alerting().firing(events, now=1) == []


def test_the_low_traffic_reason_names_the_minimum():
    decisions = alerting().evaluate([ev(0), ev(1, Outcome.PLATFORM_ERROR)], now=1)
    assert "minimum" in decisions[0].reason


def test_enough_traffic_lets_the_alert_fire():
    events = [ev(i, Outcome.PLATFORM_ERROR) for i in range(20)]
    assert alerting().firing(events, now=20)


def test_every_rule_is_evaluated_even_when_not_firing():
    assert len(alerting().evaluate(run(500), now=500)) == len(STANDARD_LADDER)


def test_a_non_firing_decision_still_explains_itself():
    for decision in alerting().evaluate(run(500), now=500):
        assert decision.reason


def test_the_ladder_thresholds_decrease_with_window_length():
    thresholds = [r.threshold for r in STANDARD_LADDER]
    assert thresholds == sorted(thresholds, reverse=True)


def test_the_fast_rule_pages_and_the_slow_rule_tickets():
    by_name = {r.name: r for r in STANDARD_LADDER}
    assert by_name["fast-burn"].severity == "page"
    assert by_name["slow-burn"].severity == "ticket"


def test_a_latency_slo_alerts_on_latency():
    slo = Slo("latency", 0.99, window_ticks=43_200, threshold_ms=1_000)
    slow = [ev(i, latency=9_000) for i in range(70)]
    assert BurnRateAlerting(slo).firing(slow, now=70)


def test_alerting_is_deterministic():
    events = sustained_failure()
    a = alerting().evaluate(events, now=70)
    b = alerting().evaluate(events, now=70)
    assert [d.reason for d in a] == [d.reason for d in b]


# ======================================================================================
# 6. the span tree
# ======================================================================================


def span(sid: str, parent, start: int, end: int, kind=SpanKind.TOOL,
         **attrs) -> Span:
    return Span(sid, parent, "t", sid, kind, start, end, attrs)


def simple_tree() -> SpanTree:
    return SpanTree([
        span("root", None, 0, 100, SpanKind.AGENT),
        span("a", "root", 10, 40),
        span("b", "root", 50, 90),
    ])


def test_the_root_is_found():
    assert simple_tree().roots == ["root"]


def test_a_span_with_a_missing_parent_is_treated_as_a_root():
    tree = SpanTree([span("orphan", "gone", 0, 10)])
    assert tree.roots == ["orphan"]


def test_walking_yields_every_span():
    assert len(list(simple_tree().walk())) == 3


def test_walking_reports_depth():
    depths = {s.span_id: d for d, s in simple_tree().walk()}
    assert depths == {"root": 0, "a": 1, "b": 1}


def test_children_are_walked_in_start_order():
    order = [s.span_id for _, s in simple_tree().walk()]
    assert order == ["root", "a", "b"]


def test_a_leaf_span_is_all_self_time():
    assert simple_tree().self_time("a") == 30


def test_self_time_subtracts_children():
    assert simple_tree().self_time("root") == 100 - 30 - 40


def test_concurrent_children_are_not_double_counted():
    tree = SpanTree([
        span("root", None, 0, 100),
        span("a", "root", 10, 60),
        span("b", "root", 10, 60),          # fully concurrent with a
    ])
    assert tree.self_time("root") == 50     # not 100 - 50 - 50 == 0


def test_partially_overlapping_children_count_their_union():
    tree = SpanTree([
        span("root", None, 0, 100),
        span("a", "root", 10, 50),
        span("b", "root", 40, 70),          # overlaps a by 10
    ])
    assert tree.self_time("root") == 100 - 60


def test_self_time_is_never_negative():
    tree = SpanTree([
        span("root", None, 0, 10),
        span("a", "root", 0, 100),          # a child longer than its parent
    ])
    assert tree.self_time("root") >= 0


def test_time_by_kind_sums_self_time():
    tree = SpanTree([
        span("root", None, 0, 100, SpanKind.AGENT),
        span("m", "root", 10, 60, SpanKind.MODEL),
        span("t", "root", 60, 90, SpanKind.TOOL),
    ])
    assert tree.time_by_kind() == {"agent": 20, "model": 50, "tool": 30}


def test_the_total_self_time_equals_the_root_duration():
    tree = simple_tree()
    assert sum(tree.time_by_kind().values()) == 100


def test_cost_is_summed_from_gen_ai_attributes():
    tree = SpanTree([
        span("root", None, 0, 10, SpanKind.AGENT),
        span("m1", "root", 0, 5, SpanKind.MODEL, **{"gen_ai.usage.cost_micros": 300}),
        span("m2", "root", 5, 9, SpanKind.MODEL, **{"gen_ai.usage.cost_micros": 700}),
    ])
    assert tree.total_cost_micros() == 1_000


def test_the_critical_path_is_the_longest_chain():
    tree = SpanTree([
        span("root", None, 0, 100),
        span("short", "root", 0, 10),
        span("long", "root", 10, 95),
    ])
    assert tree.critical_path() == ["root", "long"]


def test_errored_spans_are_listed():
    tree = SpanTree([
        span("root", None, 0, 10),
        Span("bad", "root", "t", "bad", SpanKind.TOOL, 0, 5, {}, status="error"),
    ])
    assert [s.span_id for s in tree.errors()] == ["bad"]


def test_span_ids_from_the_tracer_are_derived():
    a = Tracer(now=lambda: 10)
    b = Tracer(now=lambda: 10)
    for tracer in (a, b):
        tracer.record(trace_id="t", parent_id=None, name="x", kind=SpanKind.AGENT,
                      start_tick=0)
    assert a.spans[0].span_id == b.spans[0].span_id


def test_the_tracer_filters_by_trace():
    tracer = Tracer(now=lambda: 10)
    tracer.record(trace_id="a", parent_id=None, name="x", kind=SpanKind.AGENT,
                  start_tick=0)
    tracer.record(trace_id="b", parent_id=None, name="y", kind=SpanKind.AGENT,
                  start_tick=0)
    assert len(tracer.trace("a")) == 1


# ======================================================================================
# 7. cardinality
# ======================================================================================


def metric(name: str, **labels) -> MetricSpec:
    return MetricSpec(name, tuple(LabelSpec(k, v) for k, v in labels.items()))


def test_series_is_the_product_of_cardinalities():
    assert metric("m", tenant=10, outcome=5, deployment=4).series_count == 200


def test_a_metric_with_no_labels_is_one_series():
    assert MetricSpec("m", ()).series_count == 1


def test_a_small_metric_is_within_budget():
    assert CardinalityBudget().check(metric("m", tenant=10, outcome=5)).ok


def test_a_large_metric_is_rejected():
    verdict = CardinalityBudget(max_series_per_metric=100).check(
        metric("m", tenant=50, outcome=10))
    assert not verdict.ok and "exceeds" in verdict.reason


def test_the_rejection_names_the_worst_label():
    verdict = CardinalityBudget(max_series_per_metric=100).check(
        metric("m", tenant=5, model=200))
    assert verdict.worst_label == "model"


def test_a_forbidden_label_is_rejected_regardless_of_size():
    verdict = CardinalityBudget().check(metric("m", user_id=2))
    assert not verdict.ok and "unbounded" in verdict.reason


def test_every_forbidden_label_is_rejected():
    for name in FORBIDDEN_LABELS:
        assert not CardinalityBudget().check(metric("m", **{name: 2})).ok


def test_trace_id_belongs_on_a_trace_not_a_metric():
    verdict = CardinalityBudget().check(metric("m", trace_id=5))
    assert "trace" in verdict.reason


def test_registering_a_valid_metric_succeeds():
    budget = CardinalityBudget()
    budget.register(metric("m", tenant=10))
    assert budget.total_series() == 10


def test_registering_an_invalid_metric_raises():
    with pytest.raises(ValueError):
        CardinalityBudget(max_series_per_metric=5).register(metric("m", tenant=10))


def test_a_rejected_metric_is_not_registered():
    budget = CardinalityBudget(max_series_per_metric=5)
    with pytest.raises(ValueError):
        budget.register(metric("m", tenant=10))
    assert budget.total_series() == 0


def test_the_total_budget_is_enforced_across_metrics():
    budget = CardinalityBudget(max_series_per_metric=1_000, max_total_series=150)
    budget.register(metric("a", tenant=100))
    verdict = budget.check(metric("b", tenant=100))
    assert not verdict.ok and "total" in verdict.reason


# ======================================================================================
# 8. the degradation ladder
# ======================================================================================


def ladder(**kw) -> DegradationLadder:
    kw.setdefault("now", clock())
    return DegradationLadder(**kw)


def test_a_healthy_burn_rate_degrades_nothing():
    state = ladder().evaluate(0.5)
    assert state.level == 0 and not state.degraded


def test_a_rising_burn_rate_engages_the_first_step():
    state = ladder().evaluate(3.0)
    assert state.level == 1 and state.active == ("disable-rerank",)


def test_a_severe_burn_rate_jumps_levels():
    state = ladder().evaluate(20.0)
    assert state.level == 3


def test_descent_does_not_step_one_rung_at_a_time():
    l = ladder()
    l.evaluate(1.0)
    assert l.evaluate(60.0).level >= 5


def test_recovery_holds_before_ascending():
    l = ladder(recovery_hold_ticks=5)
    l.evaluate(20.0)
    state = l.evaluate(0.5)
    assert state.level == 3 and "holding" in state.reason


def test_recovery_ascends_one_rung_after_the_hold():
    now = clock()
    l = DegradationLadder(recovery_hold_ticks=2, now=now)
    l.evaluate(20.0)
    for _ in range(5):
        now()
    assert l.evaluate(0.5).level == 2


def test_recovery_takes_several_evaluations():
    now = clock()
    l = DegradationLadder(recovery_hold_ticks=0, now=now)
    l.evaluate(20.0)
    levels = [l.evaluate(0.5).level for _ in range(4)]
    assert levels == [2, 1, 0, 0]


def test_the_active_steps_are_a_prefix_of_the_ladder():
    l = ladder()
    state = l.evaluate(20.0)
    assert list(state.active) == [s.name for s in STANDARD_LADDER_STEPS[:state.level]]


def test_user_visible_degradation_is_reported():
    l = ladder()
    l.evaluate(3.0)
    assert not l.user_visible_degradation
    l.evaluate(20.0)
    assert l.user_visible_degradation


def test_the_first_step_is_invisible_to_users():
    assert not STANDARD_LADDER_STEPS[0].user_visible


def test_the_ladder_sheds_in_the_declared_order():
    l = ladder()
    seen = []
    for burn in (3.0, 7.0, 20.0, 40.0, 70.0, 200.0):
        seen.append(l.evaluate(burn).level)
    assert seen == sorted(seen)


def test_transitions_are_recorded():
    l = ladder()
    l.evaluate(20.0)
    assert l.history and l.history[-1][1] == 3


def test_a_ladder_without_enough_thresholds_is_rejected():
    with pytest.raises(ValueError):
        DegradationLadder(STANDARD_LADDER_STEPS, thresholds=(2.0,))


def test_every_step_states_what_it_saves():
    for step in STANDARD_LADDER_STEPS:
        assert step.saves


# ======================================================================================
# 9. cost
# ======================================================================================


def costed(n_success: int, n_fail: int, *, cost: int = 1_000,
           tenant: str = "w") -> List[RequestEvent]:
    events = [ev(i, tenant=tenant, cost_micros=cost) for i in range(n_success)]
    events += [ev(n_success + i, Outcome.PLATFORM_ERROR, tenant=tenant,
                  cost_micros=cost) for i in range(n_fail)]
    return events


def test_cost_per_action_divides_by_successful_actions():
    report = cost_report(costed(10, 0), "w")
    assert report.cost_per_action_micros == 1_000


def test_failures_raise_the_cost_per_action():
    report = cost_report(costed(10, 10), "w")
    assert report.cost_per_action_micros == 2_000


def test_wasted_spend_is_the_cost_of_failures():
    assert cost_report(costed(10, 5), "w").wasted_micros == 5_000


def test_the_waste_fraction_is_reported():
    assert cost_report(costed(10, 10), "w").waste_fraction == pytest.approx(0.5)


def test_a_tenant_with_no_successes_reports_zero_per_action():
    assert cost_report(costed(0, 5), "w").cost_per_action_micros == 0


def test_cost_is_scoped_to_the_tenant():
    events = costed(10, 0, tenant="a") + costed(5, 0, tenant="b")
    assert cost_report(events, "a").successful_actions == 10


def test_a_breaker_allows_within_budget():
    breaker = CostCircuitBreaker(budgets_micros={"w": 1_000})
    breaker.record("w", 500)
    assert breaker.allow("w")


def test_a_breaker_trips_at_the_budget():
    breaker = CostCircuitBreaker(budgets_micros={"w": 1_000})
    breaker.record("w", 1_000)
    assert not breaker.allow("w")


def test_the_breaker_warns_before_it_trips():
    breaker = CostCircuitBreaker(budgets_micros={"w": 1_000}, warn_fraction=0.8)
    breaker.record("w", 850)
    verdict = breaker.check("w")
    assert not verdict.tripped and "warn" in verdict.reason


def test_one_tenant_tripping_does_not_affect_another():
    breaker = CostCircuitBreaker(budgets_micros={"a": 100, "b": 100})
    breaker.record("a", 200)
    assert not breaker.allow("a") and breaker.allow("b")


def test_a_tenant_with_no_budget_is_denied():
    assert not CostCircuitBreaker(budgets_micros={}).allow("ghost")


def test_spend_accumulates():
    breaker = CostCircuitBreaker(budgets_micros={"w": 1_000})
    for _ in range(3):
        breaker.record("w", 100)
    assert breaker.spent("w") == 300


def test_resetting_clears_the_trip():
    breaker = CostCircuitBreaker(budgets_micros={"w": 100})
    breaker.record("w", 200)
    breaker.reset("w")
    assert breaker.allow("w") and breaker.spent("w") == 0


# ======================================================================================
# 10. capacity forecasting
# ======================================================================================


def test_steady_growth_projects_a_limit():
    forecast = forecast_capacity([10, 20, 30, 40], limit=100, lead_time_periods=1)
    assert forecast.periods_to_limit == pytest.approx(6.0)


def test_flat_usage_projects_no_limit():
    assert forecast_capacity([50] * 5, limit=100,
                             lead_time_periods=1).periods_to_limit is None


def test_falling_usage_does_not_alert():
    assert not forecast_capacity([90, 80, 70, 60], limit=100,
                                 lead_time_periods=1).alert


def test_a_long_lead_time_alerts_earlier():
    samples = [30, 34, 38, 42, 46, 50]
    short = forecast_capacity(samples, limit=100, lead_time_periods=2)
    long = forecast_capacity(samples, limit=100, lead_time_periods=12)
    assert not short.alert and long.alert


def test_the_alert_names_the_horizon():
    forecast = forecast_capacity([30, 40, 50, 60], limit=100, lead_time_periods=12)
    assert "horizon" in forecast.reason


def test_the_safety_factor_widens_the_horizon():
    samples = [30, 34, 38, 42, 46, 50]
    tight = forecast_capacity(samples, limit=100, lead_time_periods=4,
                              safety_factor=1.0)
    loose = forecast_capacity(samples, limit=100, lead_time_periods=4,
                              safety_factor=4.0)          # horizon 16 > 12.5 periods
    assert not tight.alert and loose.alert


def test_headroom_is_reported():
    forecast = forecast_capacity([10, 20, 30, 40], limit=100, lead_time_periods=1)
    assert forecast.headroom_fraction == pytest.approx(0.6)


def test_already_at_the_limit_alerts():
    assert forecast_capacity([100, 100, 100], limit=100, lead_time_periods=1).alert


def test_too_few_samples_does_not_project():
    forecast = forecast_capacity([50], limit=100, lead_time_periods=1)
    assert forecast.periods_to_limit is None and not forecast.alert


def test_a_non_positive_limit_is_an_error():
    with pytest.raises(ValueError):
        forecast_capacity([1, 2], limit=0, lead_time_periods=1)


def test_the_growth_rate_is_the_slope():
    forecast = forecast_capacity([10, 20, 30, 40], limit=100, lead_time_periods=1)
    assert forecast.growth_per_period == pytest.approx(10.0)


# ======================================================================================
# 11. incident classification
# ======================================================================================


BASE = Baseline("model-1", "prompt-1", "corpus-1", "policy-1", 0.94, recorded_tick=0)


def confirmed(hypotheses) -> List[str]:
    return [h.cause for h in hypotheses if h.confidence == "confirmed"]


def test_no_regression_short_circuits():
    result = classify_regression(BASE, lab.replace(BASE, eval_score=0.93))
    assert [h.cause for h in result] == ["no-regression"]


def test_a_prompt_change_is_confirmed():
    current = lab.replace(BASE, prompt_version="prompt-2", eval_score=0.80)
    assert confirmed(classify_regression(BASE, current)) == ["prompt-changed"]


def test_a_model_change_is_confirmed():
    current = lab.replace(BASE, model_version="model-2", eval_score=0.80)
    assert confirmed(classify_regression(BASE, current)) == ["model-changed"]


def test_a_corpus_change_is_confirmed():
    current = lab.replace(BASE, corpus_version="corpus-2", eval_score=0.80)
    assert confirmed(classify_regression(BASE, current)) == ["corpus-changed"]


def test_unchanged_versions_are_explicitly_excluded():
    current = lab.replace(BASE, prompt_version="prompt-2", eval_score=0.80)
    excluded = {h.cause for h in classify_regression(BASE, current)
                if h.confidence == "excluded"}
    assert "model-changed" in excluded and "corpus-changed" in excluded


def test_a_regression_with_nothing_changed_is_a_provider_or_drift_hypothesis():
    current = lab.replace(BASE, eval_score=0.80)
    causes = {h.cause for h in classify_regression(BASE, current)}
    assert "provider-silent-change-or-input-drift" in causes


def test_two_changes_are_both_confirmed():
    current = lab.replace(BASE, prompt_version="p2", corpus_version="c2",
                          eval_score=0.80)
    assert set(confirmed(classify_regression(BASE, current))) == \
        {"prompt-changed", "corpus-changed"}


def test_the_drop_threshold_is_configurable():
    current = lab.replace(BASE, prompt_version="p2", eval_score=0.90)
    assert classify_regression(BASE, current, eval_drop_threshold=0.01)[0].cause != \
        "no-regression"


def test_confirmed_hypotheses_come_first():
    current = lab.replace(BASE, model_version="model-2", eval_score=0.80)
    result = classify_regression(BASE, current)
    assert result[0].confidence == "confirmed"


def test_every_hypothesis_carries_evidence():
    current = lab.replace(BASE, model_version="model-2", eval_score=0.80)
    for hypothesis in classify_regression(BASE, current):
        assert hypothesis.evidence


def test_classification_is_deterministic():
    current = lab.replace(BASE, model_version="model-2", eval_score=0.80)
    assert classify_regression(BASE, current) == classify_regression(BASE, current)
