"""Tests for the action gateway.

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 Any, Dict, List, Mapping, Optional

import pytest

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

SideEffect = lab.SideEffect
EFFECT_POLICY = lab.EFFECT_POLICY
ToolContract = lab.ToolContract
Principal = lab.Principal
ActionRequest = lab.ActionRequest
ActionGateway = lab.ActionGateway
IdempotencyStore = lab.IdempotencyStore
IdemState = lab.IdemState
IdempotencyConflict = lab.IdempotencyConflict
InFlight = lab.InFlight
CircuitBreaker = lab.CircuitBreaker
BreakerState = lab.BreakerState
AuditLog = lab.AuditLog
Saga = lab.Saga
SagaStep = lab.SagaStep
DownstreamError = lab.DownstreamError
validate_schema = lab.validate_schema
redact = lab.redact
request_hash = lab.request_hash
check_dual_control = lab.check_dual_control


# ======================================================================================
# 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 frozen(value: int = 0):
    return lambda: value


PRINCIPAL = Principal("agent-1", "user-1", "wholesale", ("orchestrator",))

SIMPLE_SCHEMA = {
    "type": "object",
    "required": ["id"],
    "properties": {"id": {"type": "string"}, "amount": {"type": "integer"}},
}


class Recorder:
    """A scriptable downstream. Deterministic, and it counts."""

    def __init__(self, *, failures: int = 0, retryable: bool = True,
                 payload: Any = None) -> None:
        self.calls: List[tuple] = []
        self.failures = failures
        self.retryable = retryable
        self.payload = payload if payload is not None else {"ok": True}

    def __call__(self, tool_id: str, args: Mapping[str, Any]) -> Any:
        self.calls.append((tool_id, dict(args)))
        if self.failures > 0:
            self.failures -= 1
            raise DownstreamError("boom", retryable=self.retryable)
        return self.payload

    @property
    def count(self) -> int:
        return len(self.calls)


def gateway(*, downstream=None, effect=SideEffect.WRITE_IDEMPOTENT,
            invariants=(), now=None, fallback=None, breakers=None):
    down = downstream or Recorder()
    gw = ActionGateway(now=now or clock(), downstream=down, fallback=fallback,
                       breakers=breakers)
    gw.register(ToolContract("t.do", effect, SIMPLE_SCHEMA, invariants=invariants))
    return gw, down


def req(**kw) -> ActionRequest:
    base = dict(tool_id="t.do", arguments={"id": "x", "amount": 5},
                principal=PRINCIPAL, idempotency_key="k1",
                policy_version="v1", model_version="m1", trace_id="tr1",
                value_micros=0)
    base.update(kw)
    return ActionRequest(**base)


# ======================================================================================
# 1. schema validation
# ======================================================================================


def test_a_valid_payload_produces_no_errors():
    assert validate_schema(SIMPLE_SCHEMA, {"id": "x"}) == []


def test_a_missing_required_field_is_reported():
    assert validate_schema(SIMPLE_SCHEMA, {}) == ["$.id: required"]


def test_every_error_is_reported_not_just_the_first():
    schema = {"type": "object", "required": ["a", "b", "c"]}
    assert len(validate_schema(schema, {})) == 3


def test_errors_are_sorted_so_the_message_is_stable():
    schema = {"type": "object", "required": ["z", "a", "m"]}
    errors = validate_schema(schema, {})
    assert errors == sorted(errors)


def test_a_wrong_type_is_reported():
    assert validate_schema({"type": "integer"}, "5") == [
        "$: expected integer, got str"]


def test_a_boolean_is_not_an_integer():
    errors = validate_schema({"type": "integer"}, True)
    assert errors and "boolean" in errors[0]


def test_a_boolean_is_not_a_number_either():
    assert validate_schema({"type": "number"}, False)


def test_a_pattern_is_enforced():
    schema = {"type": "string", "pattern": r"PMT-\d+"}
    assert validate_schema(schema, "PMT-1") == []
    assert validate_schema(schema, "1") != []


def test_a_pattern_must_match_the_whole_string():
    schema = {"type": "string", "pattern": r"PMT-\d+"}
    assert validate_schema(schema, "xxPMT-1yy") != []


def test_an_enum_is_enforced():
    schema = {"type": "string", "enum": ["AED", "USD"]}
    assert validate_schema(schema, "AED") == []
    assert validate_schema(schema, "GBP") != []


def test_a_minimum_is_enforced():
    assert validate_schema({"type": "integer", "minimum": 1}, 0) != []
    assert validate_schema({"type": "integer", "minimum": 1}, 1) == []


def test_additional_properties_can_be_forbidden():
    schema = {"type": "object", "properties": {"a": {"type": "string"}},
              "additionalProperties": False}
    assert validate_schema(schema, {"a": "x", "b": "y"}) == ["$.b: not permitted"]


def test_nested_errors_carry_a_path():
    schema = {"type": "object", "properties": {
        "inner": {"type": "object", "required": ["deep"]}}}
    assert validate_schema(schema, {"inner": {}}) == ["$.inner.deep: required"]


def test_array_items_are_validated_with_an_index():
    schema = {"type": "array", "items": {"type": "integer"}}
    assert validate_schema(schema, [1, "x"]) == ["$[1]: expected integer, got str"]


# ======================================================================================
# 2. contracts and invariants
# ======================================================================================


def test_a_contract_runs_the_schema_first():
    contract = ToolContract("t", SideEffect.READ, SIMPLE_SCHEMA)
    assert contract.check(req(arguments={})) == ["$.id: required"]


def test_invariants_do_not_run_on_a_structurally_invalid_payload():
    def explodes(request):
        raise AssertionError("invariants must not see a bad payload")

    contract = ToolContract("t", SideEffect.READ, SIMPLE_SCHEMA,
                            invariants=(explodes,))
    assert contract.check(req(arguments={})) == ["$.id: required"]


def test_an_invariant_can_reject_a_structurally_valid_payload():
    contract = ToolContract("t", SideEffect.READ, SIMPLE_SCHEMA,
                            invariants=(lambda r: "nope",))
    assert contract.check(req()) == ["nope"]


def test_an_invariant_returning_none_passes():
    contract = ToolContract("t", SideEffect.READ, SIMPLE_SCHEMA,
                            invariants=(lambda r: None,))
    assert contract.check(req()) == []


def test_every_failing_invariant_is_reported():
    contract = ToolContract("t", SideEffect.READ, SIMPLE_SCHEMA,
                            invariants=(lambda r: "a", lambda r: "b"))
    assert contract.check(req()) == ["a", "b"]


def test_an_invariant_can_read_the_principal_not_just_the_arguments():
    contract = ToolContract(
        "t", SideEffect.READ, SIMPLE_SCHEMA,
        invariants=(lambda r: None if r.principal.user_id else "needs a user",))
    assert contract.check(req()) == []
    assert contract.check(req(principal=Principal("a", None, "t"))) == [
        "needs a user"]


# ======================================================================================
# 3. the effect policy table
# ======================================================================================


def test_every_side_effect_class_has_a_policy():
    assert set(EFFECT_POLICY) == set(SideEffect)


def test_reads_may_be_retried():
    assert EFFECT_POLICY[SideEffect.READ].max_attempts > 1


def test_irreversible_actions_are_never_retried_automatically():
    assert EFFECT_POLICY[SideEffect.IRREVERSIBLE].max_attempts == 1


def test_non_idempotent_writes_are_never_retried_automatically():
    assert EFFECT_POLICY[SideEffect.WRITE_NON_IDEMPOTENT].max_attempts == 1


def test_reads_do_not_need_an_idempotency_key():
    assert not EFFECT_POLICY[SideEffect.READ].requires_idempotency_key


def test_every_write_class_needs_an_idempotency_key():
    for effect in (SideEffect.WRITE_IDEMPOTENT, SideEffect.WRITE_NON_IDEMPOTENT,
                   SideEffect.IRREVERSIBLE):
        assert EFFECT_POLICY[effect].requires_idempotency_key


def test_only_irreversible_actions_require_dual_control():
    requiring = {e for e, p in EFFECT_POLICY.items()
                 if p.requires_dual_control_above is not None}
    assert requiring == {SideEffect.IRREVERSIBLE}


def test_irreversible_actions_are_not_compensable():
    assert not EFFECT_POLICY[SideEffect.IRREVERSIBLE].compensable


def test_a_write_without_a_key_is_refused():
    gw, down = gateway(effect=SideEffect.WRITE_IDEMPOTENT)
    result = gw.execute(req(idempotency_key=None))
    assert not result.ok and "idempotency key" in result.error
    assert down.count == 0


def test_a_read_without_a_key_is_fine():
    gw, down = gateway(effect=SideEffect.READ)
    assert gw.execute(req(idempotency_key=None)).ok


def test_a_read_is_retried_up_to_the_policy_limit():
    down = Recorder(failures=2)
    gw, _ = gateway(downstream=down, effect=SideEffect.READ)
    result = gw.execute(req(idempotency_key=None))
    assert result.ok and result.attempts == 3 and down.count == 3


def test_an_irreversible_action_is_attempted_once():
    down = Recorder(failures=1)
    gw, _ = gateway(downstream=down, effect=SideEffect.IRREVERSIBLE)
    result = gw.execute(req())
    assert not result.ok and down.count == 1


def test_a_non_retryable_error_stops_the_retry_loop():
    down = Recorder(failures=5, retryable=False)
    gw, _ = gateway(downstream=down, effect=SideEffect.READ)
    gw.execute(req(idempotency_key=None))
    assert down.count == 1


# ======================================================================================
# 4. idempotency
# ======================================================================================


def store(**kw) -> IdempotencyStore:
    kw.setdefault("now", clock())
    return IdempotencyStore(**kw)


def test_a_fresh_key_returns_none_and_reserves():
    s = store()
    assert s.begin("k", "h") is None
    assert s.get("k").state is IdemState.IN_FLIGHT


def test_a_second_begin_while_in_flight_raises():
    s = store()
    s.begin("k", "h")
    with pytest.raises(InFlight):
        s.begin("k", "h")


def test_a_replay_after_completion_returns_the_stored_response():
    s = store()
    s.begin("k", "h")
    s.complete("k", {"ref": "R1"})
    record = s.begin("k", "h")
    assert record is not None and record.response == {"ref": "R1"}


def test_a_different_hash_on_the_same_key_conflicts():
    s = store()
    s.begin("k", "h1")
    s.complete("k", {"ref": "R1"})
    with pytest.raises(IdempotencyConflict):
        s.begin("k", "h2")


def test_a_conflict_beats_the_in_flight_check():
    s = store()
    s.begin("k", "h1")
    with pytest.raises(IdempotencyConflict):
        s.begin("k", "h2")


def test_failing_releases_the_key_for_a_retry():
    s = store()
    s.begin("k", "h")
    s.fail("k")
    assert s.begin("k", "h") is None


def test_an_expired_record_is_treated_as_fresh():
    now = clock(start=0)
    s = IdempotencyStore(now=now, ttl_ticks=5)
    s.begin("k", "h")
    s.complete("k", "old")
    for _ in range(10):
        now()
    assert s.begin("k", "h") is None


def test_the_request_hash_ignores_the_trace_id():
    assert request_hash(req(trace_id="a")) == request_hash(req(trace_id="b"))


def test_the_request_hash_ignores_the_model_version():
    assert request_hash(req(model_version="m1")) == request_hash(req(model_version="m2"))


def test_the_request_hash_covers_the_arguments():
    assert request_hash(req(arguments={"id": "a"})) != request_hash(
        req(arguments={"id": "b"}))


def test_the_request_hash_covers_the_principal():
    other = Principal("agent-2", "user-1", "wholesale")
    assert request_hash(req()) != request_hash(req(principal=other))


def test_the_request_hash_is_insensitive_to_key_order():
    a = req(arguments={"id": "x", "amount": 5})
    b = req(arguments={"amount": 5, "id": "x"})
    assert request_hash(a) == request_hash(b)


def test_a_replayed_request_executes_exactly_once():
    gw, down = gateway()
    first = gw.execute(req())
    second = gw.execute(req())
    assert down.count == 1
    assert second.replayed and second.payload == first.payload


def test_a_conflicting_request_never_executes():
    gw, down = gateway()
    gw.execute(req())
    result = gw.execute(req(arguments={"id": "different"}))
    assert not result.ok and down.count == 1


def test_a_conflict_is_reported_as_such():
    gw, _ = gateway()
    gw.execute(req())
    result = gw.execute(req(arguments={"id": "different"}))
    assert "different request" in result.error


def test_a_failed_call_releases_the_key():
    down = Recorder(failures=1, retryable=False)
    gw, _ = gateway(downstream=down, effect=SideEffect.WRITE_IDEMPOTENT)
    assert not gw.execute(req()).ok
    down.failures = 0
    assert gw.execute(req()).ok


# ======================================================================================
# 5. dual control
# ======================================================================================


def dual(**kw):
    threshold = kw.pop("threshold", 100)
    return check_dual_control(req(**kw), threshold=threshold)


def test_below_the_threshold_no_approvers_are_needed():
    assert dual(value_micros=99, approvals=()) is None


def test_the_threshold_is_inclusive():
    assert dual(value_micros=100, approvals=()) is not None


def test_two_distinct_approvers_pass():
    assert dual(value_micros=100, approvals=("a", "b")) is None


def test_one_approver_twice_is_one_approver():
    assert dual(value_micros=100, approvals=("a", "a")) is not None


def test_the_agent_may_not_approve_its_own_action():
    assert dual(value_micros=100, approvals=("a", "agent-1")) is not None


def test_the_requesting_user_may_not_approve_their_own_action():
    assert dual(value_micros=100, approvals=("a", "user-1")) is not None


def test_an_agent_in_the_delegation_chain_may_not_approve():
    assert dual(value_micros=100, approvals=("a", "orchestrator")) is not None


def test_no_threshold_means_no_dual_control_ever():
    assert check_dual_control(req(value_micros=10 ** 12, approvals=()),
                              threshold=None) is None


def test_the_required_count_is_configurable():
    assert check_dual_control(req(value_micros=100, approvals=("a",)),
                              threshold=100, required=1) is None


def test_the_gateway_enforces_dual_control_before_executing():
    gw, down = gateway(effect=SideEffect.IRREVERSIBLE)
    result = gw.execute(req(value_micros=10 ** 12, approvals=()))
    assert not result.ok and down.count == 0


def test_a_rejected_approval_does_not_burn_the_idempotency_key():
    gw, down = gateway(effect=SideEffect.IRREVERSIBLE)
    gw.execute(req(value_micros=10 ** 12, approvals=()))
    result = gw.execute(req(value_micros=10 ** 12, approvals=("a", "b")))
    assert result.ok and down.count == 1


# ======================================================================================
# 6. the circuit breaker
# ======================================================================================


def breaker(**kw) -> CircuitBreaker:
    kw.setdefault("now", clock())
    kw.setdefault("failure_threshold", 0.5)
    kw.setdefault("minimum_throughput", 4)
    kw.setdefault("window_ticks", 1000)
    kw.setdefault("open_ticks", 10)
    kw.setdefault("half_open_successes", 2)
    return CircuitBreaker(**kw)


def test_a_new_breaker_is_closed():
    assert breaker().state() is BreakerState.CLOSED


def test_a_closed_breaker_allows():
    assert breaker().allow()


def test_failures_below_the_minimum_throughput_do_not_open_it():
    b = breaker(minimum_throughput=5)
    for _ in range(4):
        b.record(False)
    assert b.state() is BreakerState.CLOSED


def test_it_opens_exactly_at_the_threshold():
    b = breaker(minimum_throughput=4, failure_threshold=0.5)
    b.record(False)
    b.record(False)
    b.record(True)
    assert b.state() is BreakerState.CLOSED
    b.record(False)                       # 3/4 = 0.75 >= 0.5, and 4 >= 4
    assert b.state() is BreakerState.OPEN


def test_a_rate_below_the_threshold_does_not_open_it():
    b = breaker(minimum_throughput=4, failure_threshold=0.75)
    b.record(False)
    for _ in range(3):
        b.record(True)
    assert b.state() is BreakerState.CLOSED


def test_an_open_breaker_refuses():
    b = breaker(minimum_throughput=1)
    b.record(False)
    assert not b.allow()


def test_it_half_opens_after_the_open_interval():
    now = clock(start=0)
    b = breaker(now=now, minimum_throughput=1, open_ticks=5)
    b.record(False)
    assert b.state() is BreakerState.OPEN
    for _ in range(6):
        now()
    assert b.state() is BreakerState.HALF_OPEN


def test_it_does_not_half_open_early():
    now = clock(start=0)
    b = breaker(now=now, minimum_throughput=1, open_ticks=20)
    b.record(False)
    for _ in range(3):
        now()
    assert b.state() is BreakerState.OPEN


def test_a_half_open_breaker_admits_a_probe():
    now = clock(start=0)
    b = breaker(now=now, minimum_throughput=1, open_ticks=2)
    b.record(False)
    for _ in range(5):
        now()
    assert b.allow()


def test_one_failed_probe_reopens_it():
    now = clock(start=0)
    b = breaker(now=now, minimum_throughput=1, open_ticks=2)
    b.record(False)
    for _ in range(5):
        now()
    b.state()
    b.record(False)
    assert b.state() is BreakerState.OPEN


def test_it_closes_only_after_enough_successful_probes():
    now = clock(start=0)
    b = breaker(now=now, minimum_throughput=1, open_ticks=2, half_open_successes=3)
    b.record(False)
    for _ in range(5):
        now()
    b.state()
    b.record(True)
    b.record(True)
    assert b.state() is BreakerState.HALF_OPEN
    b.record(True)
    assert b.state() is BreakerState.CLOSED


def test_closing_clears_the_window_so_old_failures_do_not_reopen_it():
    now = clock(start=0)
    b = breaker(now=now, minimum_throughput=2, open_ticks=2, half_open_successes=1)
    b.record(False)
    b.record(False)
    for _ in range(5):
        now()
    b.state()
    b.record(True)
    assert b.state() is BreakerState.CLOSED
    b.record(False)
    assert b.state() is BreakerState.CLOSED


def test_events_outside_the_window_are_forgotten():
    now = clock(start=0)
    b = breaker(now=now, minimum_throughput=2, window_ticks=5, failure_threshold=0.5)
    b.record(False)
    b.record(False)
    for _ in range(10):
        now()
    b.record(True)
    b.record(True)
    assert b.state() is BreakerState.CLOSED


def test_transitions_are_recorded_for_the_incident_timeline():
    b = breaker(minimum_throughput=1)
    b.record(False)
    assert b.transitions and b.transitions[-1][1] == "open"


def test_the_gateway_refuses_when_the_breaker_is_open():
    b = breaker(minimum_throughput=1)
    down = Recorder()
    gw, _ = gateway(downstream=down, effect=SideEffect.READ,
                    breakers={"t.do": b})
    b.record(False)
    result = gw.execute(req(idempotency_key=None))
    assert not result.ok and "circuit" in result.error and down.count == 0


def test_an_open_breaker_does_not_consume_an_idempotency_key():
    b = breaker(minimum_throughput=1)
    down = Recorder()
    gw, _ = gateway(downstream=down, breakers={"t.do": b})
    b.record(False)
    gw.execute(req())
    assert gw.idempotency.get("k1") is None


def test_a_fallback_turns_an_open_breaker_into_a_degraded_answer():
    b = breaker(minimum_throughput=1)
    gw, down = gateway(effect=SideEffect.READ, breakers={"t.do": b},
                       fallback=lambda r: {"stale": True})
    b.record(False)
    result = gw.execute(req(idempotency_key=None))
    assert result.ok and result.payload == {"stale": True} and down.count == 0


# ======================================================================================
# 7. redaction
# ======================================================================================


def test_a_secret_key_is_removed_entirely():
    assert redact({"api_key": "sk-live-1"})["api_key"] == "[REDACTED]"


def test_secret_detection_is_case_insensitive():
    assert redact({"API_KEY": "x"})["API_KEY"] == "[REDACTED]"


def test_every_secretish_name_is_caught():
    for name in ("password", "secret", "token", "credential", "pin", "otp"):
        assert redact({name: "v"})[name] == "[REDACTED]"


def test_an_account_number_keeps_only_its_last_four_digits():
    assert redact({"acct": "AE070331234567890123456"})["acct"] == "****3456"


def test_redaction_recurses_into_nested_objects():
    out = redact({"a": {"b": {"password": "x"}}})
    assert out["a"]["b"]["password"] == "[REDACTED]"


def test_redaction_recurses_into_lists():
    out = redact({"xs": [{"secret": "a"}, {"secret": "b"}]})
    assert [x["secret"] for x in out["xs"]] == ["[REDACTED]", "[REDACTED]"]


def test_short_numbers_are_left_alone():
    assert redact({"count": "42"})["count"] == "42"


def test_non_sensitive_values_pass_through():
    assert redact({"status": "APPLIED"})["status"] == "APPLIED"


def test_the_gateway_redacts_before_writing_the_audit_record():
    gw, _ = gateway()
    gw.execute(req(arguments={"id": "AE070331234567890123456", "amount": 1}))
    assert gw.audit.records[-1].arguments["id"] == "****3456"


# ======================================================================================
# 8. the audit log
# ======================================================================================


def audit_fields(**kw):
    base = dict(trace_id="tr", tool_id="t", side_effect="read", actor_chain="u -> a",
                tenant="wholesale", arguments={}, outcome="success", error=None,
                policy_version="v1", model_version="m1", idempotency_key=None,
                approvals=(), value_micros=0)
    base.update(kw)
    return base


def test_the_first_record_chains_from_genesis():
    log = AuditLog(now=clock())
    record = log.append(**audit_fields())
    assert record.prev_hash == lab.GENESIS


def test_each_record_chains_to_the_previous():
    log = AuditLog(now=clock())
    a = log.append(**audit_fields())
    b = log.append(**audit_fields())
    assert b.prev_hash == a.this_hash


def test_sequence_numbers_start_at_one_and_increment():
    log = AuditLog(now=clock())
    for _ in range(3):
        log.append(**audit_fields())
    assert [r.seq for r in log.records] == [1, 2, 3]


def test_an_untouched_chain_verifies():
    log = AuditLog(now=clock())
    for i in range(5):
        log.append(**audit_fields(outcome=f"o{i}"))
    assert log.verify() == (True, None)


def test_an_empty_chain_verifies():
    assert AuditLog(now=clock()).verify() == (True, None)


def test_editing_a_records_content_breaks_verification():
    log = AuditLog(now=clock())
    for _ in range(3):
        log.append(**audit_fields())
    log.records[1] = lab.replace(log.records[1], outcome="tampered")
    ok, problem = log.verify()
    assert not ok and "record 2" in problem


def test_editing_a_records_arguments_breaks_verification():
    log = AuditLog(now=clock())
    log.append(**audit_fields(arguments={"amount": 100}))
    log.records[0] = lab.replace(log.records[0], arguments={"amount": 1})
    assert not log.verify()[0]


def test_rewriting_a_records_hash_breaks_the_next_link():
    log = AuditLog(now=clock())
    for _ in range(3):
        log.append(**audit_fields())
    edited = lab.replace(log.records[1], outcome="tampered")
    log.records[1] = lab.replace(edited, this_hash=edited.digest())
    ok, problem = log.verify()
    assert not ok and "record 3" in problem


def test_deleting_a_record_breaks_verification():
    log = AuditLog(now=clock())
    for _ in range(3):
        log.append(**audit_fields())
    del log.records[1]
    assert not log.verify()[0]


def test_the_head_hash_changes_with_every_append():
    log = AuditLog(now=clock())
    heads = set()
    for i in range(3):
        log.append(**audit_fields(outcome=f"o{i}"))
        heads.add(log.head())
    assert len(heads) == 3


def test_records_can_be_selected_by_trace():
    log = AuditLog(now=clock())
    log.append(**audit_fields(trace_id="a"))
    log.append(**audit_fields(trace_id="b"))
    log.append(**audit_fields(trace_id="a"))
    assert len(log.for_trace("a")) == 2


def test_a_successful_action_is_audited():
    gw, _ = gateway()
    gw.execute(req())
    assert gw.audit.records[-1].outcome == "success"


def test_a_refusal_is_audited_just_as_carefully():
    gw, _ = gateway()
    gw.execute(req(arguments={}))
    record = gw.audit.records[-1]
    assert record.outcome == "contract_violation" and record.policy_version == "v1"


def test_a_replay_is_audited_separately_from_the_original():
    gw, _ = gateway()
    gw.execute(req())
    gw.execute(req())
    assert [r.outcome for r in gw.audit.records] == ["success", "replayed"]


def test_the_audit_record_names_the_full_actor_chain():
    gw, _ = gateway()
    gw.execute(req())
    chain = gw.audit.records[-1].actor_chain
    assert "user-1" in chain and "orchestrator" in chain and "agent-1" in chain


def test_the_audit_record_carries_the_policy_and_model_versions():
    gw, _ = gateway()
    gw.execute(req())
    record = gw.audit.records[-1]
    assert record.policy_version == "v1" and record.model_version == "m1"


def test_an_unknown_tool_is_audited_and_never_executed():
    gw, down = gateway()
    result = gw.execute(req(tool_id="t.nope"))
    assert not result.ok and down.count == 0
    assert gw.audit.records[-1].outcome == "unknown_tool"


def test_the_gateways_chain_survives_a_mixed_workload():
    gw, _ = gateway()
    gw.execute(req())
    gw.execute(req())
    gw.execute(req(arguments={}))
    gw.execute(req(tool_id="t.nope"))
    assert gw.audit.verify() == (True, None)


# ======================================================================================
# 9. sagas
# ======================================================================================


def tracking_step(name: str, ledger: List[str], *, fail: bool = False,
                  comp_fails: bool = False, compensable: bool = True,
                  retries: int = 0, retryable: bool = True) -> SagaStep:
    def forward(ctx):
        if fail:
            raise DownstreamError(f"{name} failed", retryable=retryable)
        ledger.append(f"do:{name}")
        return f"{name}-ok"

    def compensate(ctx):
        if comp_fails:
            raise DownstreamError("compensation failed")
        ledger.append(f"undo:{name}")

    return SagaStep(name, forward, compensate if compensable else None,
                    retries=retries)


def test_a_saga_with_no_failures_completes_every_step():
    ledger: List[str] = []
    saga = Saga("s", [tracking_step("a", ledger), tracking_step("b", ledger)])
    outcome = saga.run()
    assert outcome.ok and outcome.completed == ("a", "b")
    assert ledger == ["do:a", "do:b"]


def test_a_successful_saga_compensates_nothing():
    ledger: List[str] = []
    saga = Saga("s", [tracking_step("a", ledger)])
    assert saga.run().compensated == ()


def test_a_failure_compensates_the_completed_steps():
    ledger: List[str] = []
    saga = Saga("s", [tracking_step("a", ledger), tracking_step("b", ledger),
                      tracking_step("c", ledger, fail=True)])
    outcome = saga.run()
    assert not outcome.ok and outcome.failed_step == "c"
    assert outcome.compensated == ("b", "a")


def test_compensations_run_in_reverse_order():
    ledger: List[str] = []
    Saga("s", [tracking_step("a", ledger), tracking_step("b", ledger),
               tracking_step("c", ledger, fail=True)]).run()
    assert ledger == ["do:a", "do:b", "undo:b", "undo:a"]


def test_the_failing_step_is_not_compensated():
    ledger: List[str] = []
    Saga("s", [tracking_step("a", ledger),
               tracking_step("b", ledger, fail=True)]).run()
    assert "undo:b" not in ledger


def test_a_failure_at_the_first_step_compensates_nothing():
    ledger: List[str] = []
    outcome = Saga("s", [tracking_step("a", ledger, fail=True),
                         tracking_step("b", ledger)]).run()
    assert outcome.compensated == () and ledger == []


def test_a_step_can_read_earlier_results_from_the_context():
    seen = {}

    def second(ctx):
        seen.update(ctx)
        return "ok"

    Saga("s", [SagaStep("a", lambda ctx: "a-result"),
               SagaStep("b", second)]).run()
    assert seen["a"] == "a-result"


def test_the_initial_context_is_available_to_the_first_step():
    seen = {}
    Saga("s", [SagaStep("a", lambda ctx: seen.update(ctx))]).run({"input": 42})
    assert seen["input"] == 42


def test_a_retryable_step_is_retried():
    attempts = {"n": 0}

    def flaky(ctx):
        attempts["n"] += 1
        if attempts["n"] < 3:
            raise DownstreamError("flaky")
        return "ok"

    outcome = Saga("s", [SagaStep("a", flaky, retries=3)]).run()
    assert outcome.ok and attempts["n"] == 3


def test_retries_are_bounded():
    attempts = {"n": 0}

    def always(ctx):
        attempts["n"] += 1
        raise DownstreamError("always")

    Saga("s", [SagaStep("a", always, retries=2)]).run()
    assert attempts["n"] == 3


def test_a_non_retryable_failure_is_not_retried():
    attempts = {"n": 0}

    def hard(ctx):
        attempts["n"] += 1
        raise DownstreamError("hard", retryable=False)

    Saga("s", [SagaStep("a", hard, retries=5)]).run()
    assert attempts["n"] == 1


def test_a_failing_compensation_is_reported_as_an_orphan():
    ledger: List[str] = []
    outcome = Saga("s", [tracking_step("a", ledger, comp_fails=True),
                         tracking_step("b", ledger, fail=True)]).run()
    assert outcome.orphaned == ("a",) and outcome.compensated == ()


def test_a_failing_compensation_does_not_stop_the_others():
    ledger: List[str] = []
    outcome = Saga("s", [tracking_step("a", ledger),
                         tracking_step("b", ledger, comp_fails=True),
                         tracking_step("c", ledger, fail=True)]).run()
    assert outcome.compensated == ("a",) and outcome.orphaned == ("b",)


def test_a_step_with_no_compensation_is_an_orphan_when_the_saga_unwinds():
    ledger: List[str] = []
    outcome = Saga("s", [tracking_step("a", ledger, compensable=False),
                         tracking_step("b", ledger, fail=True)]).run()
    assert outcome.orphaned == ("a",)


def test_the_outcome_names_the_error():
    ledger: List[str] = []
    outcome = Saga("s", [tracking_step("a", ledger, fail=True)]).run()
    assert "a failed" in outcome.error


def test_running_the_same_saga_twice_is_deterministic():
    a, b = [], []
    steps_a = [tracking_step("x", a), tracking_step("y", a, fail=True)]
    steps_b = [tracking_step("x", b), tracking_step("y", b, fail=True)]
    assert Saga("s", steps_a).run() == Saga("s", steps_b).run()
    assert a == b
