"""Tests for Lab 01 — The Agent Kernel.

    pytest test_lab.py -v
    LAB_MODULE=solution pytest test_lab.py -v    # the reference — must be green
"""

import importlib
import os

import pytest

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


def ticking_clock(step=0.1):
    """Deterministic clock: 0.0, 0.1, 0.2, ... — never the wall clock."""
    state = {"t": -step}

    def now():
        state["t"] += step
        return round(state["t"], 6)

    return now


def frozen_clock(value=0.0):
    return lambda: value


# ======================================================================================
# 1. Lifecycle
# ======================================================================================


def test_happy_path_transitions():
    s = lab.RunState.CREATED
    s = lab.transition(s, lab.Event.START)
    assert s is lab.RunState.PLANNING
    s = lab.transition(s, lab.Event.PROPOSE)
    assert s is lab.RunState.ACTING
    s = lab.transition(s, lab.Event.OBSERVE)
    assert s is lab.RunState.PLANNING
    s = lab.transition(s, lab.Event.FINISH)
    assert s is lab.RunState.COMPLETED


def test_terminal_states_accept_nothing():
    for terminal in (lab.RunState.COMPLETED, lab.RunState.FAILED, lab.RunState.CANCELLED):
        for event in lab.Event:
            with pytest.raises(lab.IllegalTransition):
                lab.transition(terminal, event)


def test_undeclared_transition_is_rejected():
    with pytest.raises(lab.IllegalTransition):
        lab.transition(lab.RunState.CREATED, lab.Event.OBSERVE)
    with pytest.raises(lab.IllegalTransition):
        lab.transition(lab.RunState.ACTING, lab.Event.FINISH)   # must observe first
    with pytest.raises(lab.IllegalTransition):
        lab.transition(lab.RunState.SUSPENDED, lab.Event.PROPOSE)


def test_every_non_terminal_state_can_reach_a_terminal_state():
    """An invariant worth asserting: no state is a trap."""
    reachable_terminals = {
        state: {TRANS for (s, e), TRANS in lab.TRANSITIONS.items() if s is state}
        for state in lab.RunState if state not in lab.TERMINAL_STATES
    }
    for state, targets in reachable_terminals.items():
        assert targets & lab.TERMINAL_STATES, f"{state} cannot terminate"


def test_hitl_round_trip():
    s = lab.transition(lab.RunState.PLANNING, lab.Event.NEED_INPUT)
    assert s is lab.RunState.WAITING_INPUT
    assert lab.transition(s, lab.Event.RESUME_INPUT) is lab.RunState.PLANNING


def test_suspend_resume_round_trip():
    s = lab.transition(lab.RunState.ACTING, lab.Event.SUSPEND)
    assert s is lab.RunState.SUSPENDED
    assert lab.transition(s, lab.Event.RESUME) is lab.RunState.PLANNING


# ======================================================================================
# 2. Memory
# ======================================================================================


def test_estimate_tokens_is_deterministic_and_monotone():
    assert lab.estimate_tokens("") == 0
    assert lab.estimate_tokens("abcd") == 1
    assert lab.estimate_tokens("abcde") == 2
    assert lab.estimate_tokens("x" * 400) == 100
    assert lab.estimate_tokens("hello") == lab.estimate_tokens("hello")


def make_step(i, obs="ok"):
    return lab.Step(index=i, thought=f"thinking about step {i}", tool="t",
                    arguments=(("k", "v"),), observation=obs)


def test_scratchpad_grows_and_renders_in_order():
    pad = lab.Scratchpad(max_tokens=10_000, summarize=lab.default_summarizer)
    for i in range(1, 4):
        pad.append(make_step(i))
    rendered = pad.render()
    assert rendered.index("[1]") < rendered.index("[2]") < rendered.index("[3]")
    assert pad.token_count() > 0


def test_scratchpad_does_not_compact_below_the_threshold():
    pad = lab.Scratchpad(max_tokens=10_000, summarize=lab.default_summarizer)
    pad.append(make_step(1))
    assert pad.compact_if_needed() is False
    assert pad.compactions == 0
    assert pad.summary == ""


def test_scratchpad_compacts_and_keeps_the_recent_steps():
    pad = lab.Scratchpad(max_tokens=40, summarize=lab.default_summarizer, keep_recent=2)
    for i in range(1, 7):
        pad.append(make_step(i, obs="a fairly long observation string " * 2))
    assert pad.compact_if_needed() is True
    assert pad.compactions == 1
    assert len(pad.steps) == 2
    assert [s.index for s in pad.steps] == [5, 6]
    assert pad.summary != ""
    assert "[summary of earlier steps]" in pad.render()


def test_scratchpad_will_not_compact_away_the_recent_window():
    """Even over budget, the last keep_recent steps survive — otherwise the model loses
    the context it needs to make the next decision."""
    pad = lab.Scratchpad(max_tokens=1, summarize=lab.default_summarizer, keep_recent=2)
    pad.append(make_step(1))
    pad.append(make_step(2))
    assert pad.compact_if_needed() is False
    assert len(pad.steps) == 2


def test_compaction_reduces_token_count():
    pad = lab.Scratchpad(max_tokens=40, summarize=lab.default_summarizer, keep_recent=1)
    for i in range(1, 9):
        pad.append(make_step(i, obs="observation text " * 6))
    before = pad.token_count()
    pad.compact_if_needed()
    assert pad.token_count() < before


def test_semantic_memory_is_scoped():
    mem = lab.SemanticMemory()
    mem.put(lab.Fact("a", "1", "tenant", "wholesale", ("x",)))
    mem.put(lab.Fact("b", "2", "tenant", "retail", ("x",)))
    seen = mem.search(scopes={"tenant": "wholesale"}, tags=["x"])
    assert [f.key for f in seen] == ["a"]


def test_semantic_memory_cannot_be_widened_by_asking_nicely():
    """A caller who holds no scope sees nothing, regardless of tags."""
    mem = lab.SemanticMemory()
    mem.put(lab.Fact("a", "1", "tenant", "wholesale", ("x",)))
    assert mem.search(scopes={}, tags=["x"]) == []
    assert mem.search(scopes={"tenant": "retail"}, tags=["x"]) == []


def test_semantic_memory_ranks_by_tag_overlap_then_key():
    mem = lab.SemanticMemory()
    mem.put(lab.Fact("zz", "v", "user", "u1", ("a", "b")))
    mem.put(lab.Fact("aa", "v", "user", "u1", ("a",)))
    mem.put(lab.Fact("bb", "v", "user", "u1", ("a",)))
    got = mem.search(scopes={"user": "u1"}, tags=["a", "b"])
    assert [f.key for f in got] == ["zz", "aa", "bb"]


def test_semantic_memory_rejects_unknown_scope():
    mem = lab.SemanticMemory()
    with pytest.raises(ValueError):
        mem.put(lab.Fact("a", "1", "galaxy", "x"))


def test_semantic_memory_get_is_exact():
    mem = lab.SemanticMemory()
    mem.put(lab.Fact("a", "1", "tenant", "wholesale"))
    assert mem.get("tenant", "wholesale", "a").value == "1"
    assert mem.get("tenant", "retail", "a") is None


def test_episodic_recall_prefers_overlap_then_recency():
    mem = lab.EpisodicMemory()
    mem.record(lab.Episode("e1", "s", "g", "completed", 3, ("pay",)))
    mem.record(lab.Episode("e2", "s", "g", "failed", 5, ("pay", "sanctions")))
    mem.record(lab.Episode("e3", "s", "g", "completed", 2, ("pay",)))
    got = mem.recall(["pay", "sanctions"], limit=3)
    assert [e.episode_id for e in got] == ["e2", "e3", "e1"]


def test_episodic_recall_ignores_unrelated_episodes():
    mem = lab.EpisodicMemory()
    mem.record(lab.Episode("e1", "s", "g", "completed", 1, ("kyc",)))
    assert mem.recall(["pay"]) == []


# ======================================================================================
# 3. Session store
# ======================================================================================


def snapshot(session_id="s", state=None):
    return lab.SessionSnapshot(
        session_id=session_id, tenant="wholesale", user_id="u", goal="g",
        state=state or lab.RunState.CREATED, version=0, steps=(), summary="",
        tokens_used=0, cost_micros=0,
    )


def test_create_assigns_version_one():
    store = lab.SessionStore()
    assert store.create(snapshot()).version == 1


def test_create_rejects_duplicates():
    store = lab.SessionStore()
    store.create(snapshot())
    with pytest.raises(KeyError):
        store.create(snapshot())


def test_load_missing_session_raises():
    with pytest.raises(KeyError):
        lab.SessionStore().load("nope")


def test_save_increments_version():
    store = lab.SessionStore()
    s = store.create(snapshot())
    s2 = store.save(s, expected_version=1)
    assert s2.version == 2
    assert store.load("s").version == 2


def test_stale_write_is_rejected():
    """Two workers advanced the same session; exactly one wins."""
    store = lab.SessionStore()
    s = store.create(snapshot())
    store.save(s, expected_version=1)
    with pytest.raises(lab.ConcurrentModification):
        store.save(s, expected_version=1)


def test_exists():
    store = lab.SessionStore()
    assert not store.exists("s")
    store.create(snapshot())
    assert store.exists("s")


# ======================================================================================
# 4. Session affinity
# ======================================================================================


def test_routing_is_deterministic():
    r1 = lab.AffinityRouter(["a", "b", "c"])
    r2 = lab.AffinityRouter(["c", "b", "a"])   # different insertion order
    for i in range(100):
        sid = f"s-{i}"
        assert r1.route(sid) == r2.route(sid)


def test_routing_is_stable_across_processes():
    """Uses a stable digest, not Python's per-process-salted hash(). If this fails,
    every pod restart reshuffles every session."""
    router = lab.AffinityRouter(["pod-a", "pod-b", "pod-c"])
    # These are fixed by the blake2b ring; they must not drift between runs.
    first = [router.route(f"s-{i}") for i in range(20)]
    second = [lab.AffinityRouter(["pod-a", "pod-b", "pod-c"]).route(f"s-{i}") for i in range(20)]
    assert first == second


def test_distribution_is_reasonably_balanced():
    router = lab.AffinityRouter(["a", "b", "c"], virtual_nodes=128)
    counts = router.distribution([f"s-{i}" for i in range(900)])
    assert sum(counts.values()) == 900
    for value in counts.values():
        assert 200 < value < 400        # within ±1/3 of the 300 ideal


def test_adding_a_replica_moves_about_one_over_n():
    sessions = [f"s-{i}" for i in range(600)]
    router = lab.AffinityRouter(["a", "b", "c"], virtual_nodes=128)
    before = {s: router.route(s) for s in sessions}
    router.add("d")
    after = {s: router.route(s) for s in sessions}
    moved = sum(1 for s in sessions if before[s] != after[s])
    # Consistent hashing moves ~1/4 of keys for 3 -> 4; modulo would move ~3/4.
    assert 0.10 < moved / len(sessions) < 0.45


def test_moved_sessions_only_go_to_the_new_replica():
    """The other invariant of consistent hashing: no churn between existing replicas."""
    sessions = [f"s-{i}" for i in range(400)]
    router = lab.AffinityRouter(["a", "b", "c"], virtual_nodes=128)
    before = {s: router.route(s) for s in sessions}
    router.add("d")
    for s in sessions:
        after = router.route(s)
        if after != before[s]:
            assert after == "d"


def test_draining_stops_new_routing_without_removing():
    router = lab.AffinityRouter(["a", "b", "c"])
    router.drain("a")
    routes = {router.route(f"s-{i}") for i in range(200)}
    assert "a" not in routes
    assert routes <= {"b", "c"}


def test_draining_every_replica_is_an_error():
    router = lab.AffinityRouter(["a"])
    router.drain("a")
    with pytest.raises(RuntimeError):
        router.route("s-1")


def test_remove_unknown_replica_raises():
    router = lab.AffinityRouter(["a"])
    with pytest.raises(KeyError):
        router.remove("zzz")
    with pytest.raises(KeyError):
        router.drain("zzz")


def test_empty_router_raises_on_route():
    router = lab.AffinityRouter([])
    with pytest.raises(RuntimeError):
        router.route("s")


def test_virtual_nodes_must_be_positive():
    with pytest.raises(ValueError):
        lab.AffinityRouter(["a"], virtual_nodes=0)


# ======================================================================================
# 5. Budgets and decisions
# ======================================================================================


@pytest.mark.parametrize("kwargs", [
    {"max_steps": 0}, {"max_tokens": 0}, {"max_cost_micros": 0}, {"deadline_seconds": 0},
])
def test_budgets_reject_non_positive(kwargs):
    with pytest.raises(ValueError):
        lab.Budgets(**kwargs)


def test_decision_validates_its_shape():
    with pytest.raises(ValueError):
        lab.Decision("teleport", "hmm")
    with pytest.raises(ValueError):
        lab.Decision("act", "no tool named")
    with pytest.raises(ValueError):
        lab.Decision("ask", "no question")
    lab.Decision("finish", "done", answer="42")     # answer is optional-but-sensible


# ======================================================================================
# 6. The kernel
# ======================================================================================


def echo_tool(args):
    return lab.ToolResult(True, f"echo:{args.get('x', '')}", tokens=10, cost_micros=100)


def failing_tool(args):
    return lab.ToolResult(False, "downstream 503", tokens=5, cost_micros=50)


def build_kernel(policy, *, budgets=None, tools=None, clock=None, pad_tokens=10_000):
    store = lab.SessionStore()
    kernel = lab.AgentKernel(
        store=store,
        tools=tools if tools is not None else {"echo": echo_tool, "boom": failing_tool},
        policy=policy,
        now=clock or ticking_clock(),
        budgets=budgets or lab.Budgets(max_steps=10, deadline_seconds=1e6),
        scratchpad_max_tokens=pad_tokens,
    )
    kernel.create_session(session_id="s", tenant="wholesale", user_id="u", goal="g")
    return kernel


def two_step_policy(rendered):
    if "echo:" not in rendered:
        return lab.Decision("act", "call echo", tool="echo",
                            arguments=(("x", "hi"),), tokens_in=10, tokens_out=5)
    return lab.Decision("finish", "done", answer="the answer", tokens_in=10, tokens_out=5)


def test_kernel_completes_a_simple_run():
    kernel = build_kernel(two_step_policy)
    result = kernel.run("s")
    assert result.succeeded
    assert result.state is lab.RunState.COMPLETED
    assert result.answer == "the answer"
    assert [s.index for s in result.steps] == [1, 2]
    assert result.steps[0].tool == "echo"


def test_kernel_accounts_tokens_and_cost():
    kernel = build_kernel(two_step_policy)
    result = kernel.run("s")
    # policy: 15 + 15; tool: 10 tokens, 100 micros
    assert result.tokens_used == 40
    assert result.cost_micros == 100


def test_unknown_tool_is_recoverable_not_fatal():
    calls = {"n": 0}

    def policy(rendered):
        calls["n"] += 1
        if calls["n"] == 1:
            return lab.Decision("act", "try", tool="nope", arguments=())
        return lab.Decision("finish", "recovered", answer="ok")

    kernel = build_kernel(policy)
    result = kernel.run("s")
    assert result.succeeded
    assert result.steps[0].error == "unknown tool: 'nope'"
    assert result.steps[0].observation is None


def test_tool_failure_is_recorded_as_an_error_not_an_exception():
    def policy(rendered):
        if "503" not in rendered:
            return lab.Decision("act", "call boom", tool="boom", arguments=())
        return lab.Decision("finish", "gave up", answer="degraded")

    kernel = build_kernel(policy)
    result = kernel.run("s")
    assert result.succeeded
    assert result.steps[0].error == "downstream 503"


def test_step_budget_stops_a_looping_agent():
    looper = lambda r: lab.Decision("act", "again", tool="echo", arguments=(("x", "1"),))
    kernel = build_kernel(looper, budgets=lab.Budgets(max_steps=3, deadline_seconds=1e6))
    result = kernel.run("s")
    assert result.state is lab.RunState.FAILED
    assert result.breach.kind == "steps"
    assert len(result.steps) == 3


def test_step_budget_boundary_allows_exactly_max_steps():
    calls = {"n": 0}

    def policy(rendered):
        calls["n"] += 1
        if calls["n"] < 3:
            return lab.Decision("act", "go", tool="echo", arguments=(("x", str(calls["n"])),))
        return lab.Decision("finish", "done", answer="ok")

    kernel = build_kernel(policy, budgets=lab.Budgets(max_steps=3, deadline_seconds=1e6))
    result = kernel.run("s")
    assert result.succeeded
    assert len(result.steps) == 3


def test_token_budget_breach():
    heavy = lambda r: lab.Decision("act", "big", tool="echo", arguments=(("x", "1"),),
                                   tokens_in=5000, tokens_out=5000)
    kernel = build_kernel(heavy, budgets=lab.Budgets(max_steps=50, max_tokens=12_000,
                                                     deadline_seconds=1e6))
    result = kernel.run("s")
    assert result.state is lab.RunState.FAILED
    assert result.breach.kind == "tokens"


def test_cost_budget_breach():
    def pricey(args):
        return lab.ToolResult(True, "spent", tokens=1, cost_micros=400_000)

    kernel = build_kernel(
        lambda r: lab.Decision("act", "spend", tool="pricey", arguments=()),
        budgets=lab.Budgets(max_steps=50, max_cost_micros=500_000, deadline_seconds=1e6),
        tools={"pricey": pricey},
    )
    result = kernel.run("s")
    assert result.state is lab.RunState.FAILED
    assert result.breach.kind == "cost"


def test_deadline_breach():
    kernel = build_kernel(
        lambda r: lab.Decision("act", "slow", tool="echo", arguments=(("x", "1"),)),
        budgets=lab.Budgets(max_steps=100, deadline_seconds=0.5),
        clock=ticking_clock(0.2),
    )
    result = kernel.run("s")
    assert result.state is lab.RunState.FAILED
    assert result.breach.kind == "deadline"


def test_budget_is_checked_before_the_expensive_call():
    """A kernel that checks budgets after calling the model has already paid for the
    step it is about to reject."""
    calls = {"n": 0}

    def counting_policy(rendered):
        calls["n"] += 1
        return lab.Decision("act", "again", tool="echo", arguments=(("x", "1"),))

    kernel = build_kernel(counting_policy, budgets=lab.Budgets(max_steps=2, deadline_seconds=1e6))
    kernel.run("s")
    assert calls["n"] == 2      # not 3


def test_human_in_the_loop_pauses_and_resumes():
    def policy(rendered):
        if "awaiting human input" not in rendered:
            return lab.Decision("ask", "need approval", question="approve?")
        return lab.Decision("finish", "approved", answer="released")

    kernel = build_kernel(policy)
    first = kernel.run("s")
    assert first.state is lab.RunState.WAITING_INPUT
    assert first.question == "approve?"
    assert kernel.store.load("s").pending_question == "approve?"

    second = kernel.run("s", resume_answer="approved by u-99")
    assert second.succeeded
    assert second.answer == "released"
    assert any(s.observation == "approved by u-99" for s in second.steps)


def test_resuming_without_an_answer_is_an_error():
    kernel = build_kernel(lambda r: lab.Decision("ask", "?", question="approve?"))
    kernel.run("s")
    with pytest.raises(ValueError):
        kernel.run("s")


def test_a_terminal_session_cannot_be_run_again():
    kernel = build_kernel(two_step_policy)
    kernel.run("s")
    with pytest.raises(lab.IllegalTransition):
        kernel.run("s")


def test_state_is_checkpointed_every_step():
    versions = []

    def policy(rendered):
        versions.append(kernel.store.load("s").version)
        if len(versions) < 4:
            return lab.Decision("act", "go", tool="echo", arguments=(("x", "1"),))
        return lab.Decision("finish", "done", answer="ok")

    kernel = build_kernel(policy)
    kernel.run("s")
    assert versions == sorted(versions)
    assert versions[-1] > versions[0]        # the store advanced during the run


def test_execution_chain_reads_from_the_store_not_the_scratchpad():
    """Compaction is lossy for the model and never for the audit record."""
    def policy(rendered):
        n = rendered.count("echo:")
        if n < 6:
            return lab.Decision("act", "go " * 10, tool="echo",
                                arguments=(("x", "y" * 40),))
        return lab.Decision("finish", "done", answer="ok")

    kernel = build_kernel(policy, pad_tokens=60)
    result = kernel.run("s")
    assert result.compactions > 0
    chain = kernel.execution_chain("s")
    assert len(chain) == len(result.steps)
    # The scratchpad dropped early steps, but the store kept every one it checkpointed.
    assert chain[0]["tenant"] == "wholesale"
    assert chain[0]["user_id"] == "u"


def test_execution_chain_carries_identity_and_outcome():
    kernel = build_kernel(two_step_policy)
    kernel.run("s")
    chain = kernel.execution_chain("s")
    assert chain[0]["tool"] == "echo"
    assert chain[0]["ok"] is True
    assert chain[0]["arguments"] == {"x": "hi"}
    assert chain[0]["tenant"] == "wholesale"


def test_terminal_run_records_an_episode():
    kernel = build_kernel(two_step_policy)
    assert len(kernel.episodic) == 0
    kernel.run("s")
    assert len(kernel.episodic) == 1
    assert kernel.episodic.recall(["echo"])[0].outcome == "completed"


def test_failed_run_records_the_breach_as_a_lesson():
    looper = lambda r: lab.Decision("act", "again", tool="echo", arguments=(("x", "1"),))
    kernel = build_kernel(looper, budgets=lab.Budgets(max_steps=2, deadline_seconds=1e6))
    kernel.run("s")
    episode = kernel.episodic.recall(["echo"])[0]
    assert episode.outcome == "failed"
    assert "max_steps" in episode.lesson


# ======================================================================================
# 7. Determinism
# ======================================================================================


def test_two_identical_kernels_produce_identical_runs():
    def build():
        return build_kernel(two_step_policy, clock=ticking_clock())

    a, b = build(), build()
    ra, rb = a.run("s"), b.run("s")
    assert ra.steps == rb.steps
    assert ra.tokens_used == rb.tokens_used
    assert ra.cost_micros == rb.cost_micros
    assert a.execution_chain("s") == b.execution_chain("s")


def test_scratchpad_render_is_stable():
    pad = lab.Scratchpad(max_tokens=10_000, summarize=lab.default_summarizer)
    for i in range(1, 5):
        pad.append(make_step(i))
    assert pad.render() == pad.render()
