"""Unit tests for agent_team.nodes.clarifier (design §3.3, §7.1 P1). Covers the clarifier's two responsibilities: the **98% confidence loop** driven by LangGraph ``interrupt()``/resume (the §7.1 P1 "riskiest mechanic"), and the **human gate** — only a run clearing the bar advances to planning, a run that hits the turn cap parks instead of spinning. The loop tests drive a real compiled ``StateGraph`` with an in-memory checkpointer so the suspend/resume + replay semantics are exercised exactly as they will be on the box (the box swaps in the SQLite checkpointer; the node itself is checkpointer-agnostic). Confidence and question generation are injected as deterministic callables, so the tests assert control flow without a live Claude call. """ from __future__ import annotations from collections.abc import Sequence import pytest from langgraph.checkpoint.memory import MemorySaver from langgraph.graph import END, START, StateGraph from langgraph.types import Command from agent_team.nodes.clarifier import ( DEFAULT_CONFIDENCE_THRESHOLD, DEFAULT_MAX_TURNS, ClarifierConfig, build_question_set, make_clarifier_node, ) from agent_team.task_model import ( Phase, PipelineState, TaskStatus, new_thread_id, ) from agent_team.transport import QuestionSet # --------------------------------------------------------------------------- # # Test helpers: deterministic injected assessor / generator. # --------------------------------------------------------------------------- # def _confidence_after_n_answers(target_turns: int) -> object: """Assessor that clears the 98% bar once ``target_turns`` answers exist. Confidence is 0.0 with no answers and steps to 1.0 once enough answers have been collected. Lets a test pin exactly how many interrupt turns the loop should take. """ def assess(qa_history: Sequence[object], state: PipelineState) -> float: return 1.0 if len(qa_history) >= target_turns else 0.0 return assess def _always_questions(*prompts: str): """Generator returning a fixed question-set every turn.""" def generate(qa_history: Sequence[object], state: PipelineState) -> list[str]: return list(prompts) or ["What is the goal?"] return generate def _build_graph(node): """Compile a one-node graph around ``node`` with an in-memory checkpointer.""" graph = StateGraph(PipelineState) graph.add_node("clarify", node) graph.add_edge(START, "clarify") graph.add_edge("clarify", END) return graph.compile(checkpointer=MemorySaver()) def _initial_state(thread_id: str, transport: str = "slack") -> PipelineState: return { "thread_id": thread_id, "status": TaskStatus.ACTIVE.value, "current_phase": Phase.CLARIFY.value, "qa_history": [], "transport": transport, } # --------------------------------------------------------------------------- # # build_question_set — the §3.3.1 interrupt payload. # --------------------------------------------------------------------------- # def test_build_question_set_returns_foundation_questionset() -> None: qs = build_question_set(thread_id="t1", turn=2, questions=["a", "b"]) assert isinstance(qs, QuestionSet) assert qs.thread_id == "t1" assert qs.turn == 2 assert qs.questions == ["a", "b"] assert qs.context == {} def test_build_question_set_mints_unique_question_ids() -> None: a = build_question_set(thread_id="t", turn=0, questions=["x"]) b = build_question_set(thread_id="t", turn=0, questions=["x"]) assert a.question_id != b.question_id assert len(a.question_id) == 32 # uuid4 hex def test_build_question_set_honours_explicit_question_id() -> None: qs = build_question_set( thread_id="t", turn=0, questions=["x"], question_id="fixed-id" ) assert qs.question_id == "fixed-id" def test_build_question_set_copies_inputs_defensively() -> None: questions = ["x"] context = {"repo": "r"} qs = build_question_set(thread_id="t", turn=0, questions=questions, context=context) questions.append("y") context["repo"] = "mutated" assert qs.questions == ["x"] assert qs.context == {"repo": "r"} # --------------------------------------------------------------------------- # # ClarifierConfig validation. # --------------------------------------------------------------------------- # def test_config_defaults_match_design_constants() -> None: cfg = ClarifierConfig() assert cfg.confidence_threshold == DEFAULT_CONFIDENCE_THRESHOLD == 0.98 assert cfg.max_turns == DEFAULT_MAX_TURNS @pytest.mark.parametrize("bad", [0.0, -0.1, 1.5]) def test_config_rejects_threshold_out_of_range(bad: float) -> None: with pytest.raises(ValueError, match="confidence_threshold"): ClarifierConfig(confidence_threshold=bad) def test_config_accepts_threshold_of_one() -> None: assert ClarifierConfig(confidence_threshold=1.0).confidence_threshold == 1.0 @pytest.mark.parametrize("bad", [0, -3]) def test_config_rejects_nonpositive_max_turns(bad: int) -> None: with pytest.raises(ValueError, match="max_turns"): ClarifierConfig(max_turns=bad) # --------------------------------------------------------------------------- # # Confidence loop: already-confident, no interrupt (human gate opens at once). # --------------------------------------------------------------------------- # def test_no_interrupt_when_already_confident() -> None: node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(0), # confident immediately generate_questions=_always_questions("q"), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} result = app.invoke(_initial_state("t1"), cfg) assert "__interrupt__" not in result assert result["current_phase"] == Phase.PLAN.value assert result["status"] == TaskStatus.ACTIVE.value assert result["qa_history"] == [] # --------------------------------------------------------------------------- # # Confidence loop: interrupts until the 98% bar is cleared, then human gate. # --------------------------------------------------------------------------- # def test_single_turn_loop_clears_bar_and_opens_gate() -> None: node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(1), generate_questions=_always_questions("What is the goal?"), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} first = app.invoke(_initial_state("t1"), cfg) assert "__interrupt__" in first # suspended on the question-set final = app.invoke(Command(resume="ship feature X"), cfg) assert "__interrupt__" not in final assert final["qa_history"] == ["ship feature X"] assert final["current_phase"] == Phase.PLAN.value assert final["status"] == TaskStatus.ACTIVE.value def test_multi_turn_loop_accumulates_answers_until_confident() -> None: node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(3), generate_questions=_always_questions("q?"), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} assert "__interrupt__" in app.invoke(_initial_state("t1"), cfg) assert "__interrupt__" in app.invoke(Command(resume="a1"), cfg) assert "__interrupt__" in app.invoke(Command(resume="a2"), cfg) final = app.invoke(Command(resume="a3"), cfg) assert "__interrupt__" not in final assert final["qa_history"] == ["a1", "a2", "a3"] assert final["current_phase"] == Phase.PLAN.value assert final["status"] == TaskStatus.ACTIVE.value # --------------------------------------------------------------------------- # # Interrupt payload shape (§3.3.1): {thread_id, question_id, turn, ...}. # --------------------------------------------------------------------------- # def test_interrupt_payload_carries_design_fields() -> None: node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(1), generate_questions=_always_questions("Q1", "Q2"), config=ClarifierConfig(transport="slack"), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "tid-123"}} app.invoke(_initial_state("tid-123", transport="ignored"), cfg) state = app.get_state(cfg) assert len(state.interrupts) == 1 payload = state.interrupts[0].value assert payload["thread_id"] == "tid-123" assert payload["turn"] == 0 assert payload["transport"] == "slack" # config wins over state assert payload["deadline"] is None # owned by the ledger/timer seam assert isinstance(payload["question_id"], str) and payload["question_id"] qs = payload["question_set"] assert isinstance(qs, QuestionSet) assert qs.questions == ["Q1", "Q2"] assert qs.question_id == payload["question_id"] assert qs.thread_id == "tid-123" def test_interrupt_payload_transport_falls_back_to_state() -> None: node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(1), generate_questions=_always_questions("q"), config=ClarifierConfig(), # empty transport ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} app.invoke(_initial_state("t1", transport="github"), cfg) payload = app.get_state(cfg).interrupts[0].value assert payload["transport"] == "github" def test_turn_index_increments_across_interrupts() -> None: node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(2), generate_questions=_always_questions("q"), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} app.invoke(_initial_state("t1"), cfg) assert app.get_state(cfg).interrupts[0].value["turn"] == 0 app.invoke(Command(resume="a1"), cfg) assert app.get_state(cfg).interrupts[0].value["turn"] == 1 # --------------------------------------------------------------------------- # # Turn cap (§7.1): park rather than spin; never advance to planning. # --------------------------------------------------------------------------- # def test_turn_cap_parks_instead_of_advancing() -> None: # Never confident; cap at 2 turns. node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(999), generate_questions=_always_questions("q"), config=ClarifierConfig(max_turns=2), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} assert "__interrupt__" in app.invoke(_initial_state("t1"), cfg) assert "__interrupt__" in app.invoke(Command(resume="a1"), cfg) final = app.invoke(Command(resume="a2"), cfg) # cap reached assert "__interrupt__" not in final assert final["current_phase"] == Phase.PARKED.value assert final["status"] == TaskStatus.PARKED.value assert final["qa_history"] == ["a1", "a2"] # Human gate stayed shut: never advanced to PLAN. assert final["current_phase"] != Phase.PLAN.value def test_clearing_bar_on_the_cap_turn_still_opens_gate() -> None: # Confident exactly on the last allowed turn — gate must open, not park. node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(2), generate_questions=_always_questions("q"), config=ClarifierConfig(max_turns=2), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} app.invoke(_initial_state("t1"), cfg) app.invoke(Command(resume="a1"), cfg) final = app.invoke(Command(resume="a2"), cfg) assert final["current_phase"] == Phase.PLAN.value assert final["status"] == TaskStatus.ACTIVE.value # --------------------------------------------------------------------------- # # Durable resume across a "restart" (§7.1 P1 exit (a)): a fresh app object # bound to the same checkpointer resumes mid-wait to the right thread. # --------------------------------------------------------------------------- # def test_resume_after_restart_uses_same_checkpoint() -> None: saver = MemorySaver() def compile_app(): node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(1), generate_questions=_always_questions("q"), ) graph = StateGraph(PipelineState) graph.add_node("clarify", node) graph.add_edge(START, "clarify") graph.add_edge("clarify", END) return graph.compile(checkpointer=saver) cfg = {"configurable": {"thread_id": "t1"}} app1 = compile_app() assert "__interrupt__" in app1.invoke(_initial_state("t1"), cfg) # Simulate a process restart: a brand-new app object, same checkpointer. app2 = compile_app() state = app2.get_state(cfg) assert state.next == ("clarify",) # still suspended at the node final = app2.invoke(Command(resume="late answer"), cfg) assert final["qa_history"] == ["late answer"] assert final["current_phase"] == Phase.PLAN.value # --------------------------------------------------------------------------- # # Parallel-task isolation (§7.1 P1 exit (d)): two threads resume independently. # --------------------------------------------------------------------------- # def test_two_threads_resume_independently() -> None: node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(1), generate_questions=_always_questions("q"), ) app = _build_graph(node) cfg_a = {"configurable": {"thread_id": "task-a"}} cfg_b = {"configurable": {"thread_id": "task-b"}} # Suspend both. app.invoke(_initial_state("task-a"), cfg_a) app.invoke(_initial_state("task-b"), cfg_b) # Resume out of order; each carries its own answer + thread_id. final_b = app.invoke(Command(resume="answer-b"), cfg_b) final_a = app.invoke(Command(resume="answer-a"), cfg_a) assert final_a["qa_history"] == ["answer-a"] assert final_b["qa_history"] == ["answer-b"] assert final_a["current_phase"] == Phase.PLAN.value assert final_b["current_phase"] == Phase.PLAN.value # --------------------------------------------------------------------------- # # Resumed task keeps prior Q&A history (a re-opened parked task adds context). # --------------------------------------------------------------------------- # def test_existing_qa_history_is_preserved_and_extended() -> None: node = make_clarifier_node( # Need 2 total answers; one already exists, so one more interrupt. assess_confidence=_confidence_after_n_answers(2), generate_questions=_always_questions("q"), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": "t1"}} seeded = _initial_state("t1") seeded["qa_history"] = ["prior answer"] assert "__interrupt__" in app.invoke(seeded, cfg) final = app.invoke(Command(resume="new answer"), cfg) assert final["qa_history"] == ["prior answer", "new answer"] assert final["current_phase"] == Phase.PLAN.value # --------------------------------------------------------------------------- # # The node is callable factory output and integrates with new_thread_id intake. # --------------------------------------------------------------------------- # def test_node_runs_against_freshly_minted_thread_id() -> None: tid = new_thread_id() node = make_clarifier_node( assess_confidence=_confidence_after_n_answers(0), generate_questions=_always_questions("q"), ) app = _build_graph(node) cfg = {"configurable": {"thread_id": tid}} result = app.invoke(_initial_state(tid), cfg) assert result["current_phase"] == Phase.PLAN.value