"""Targeted tests for interrupted-graph-state recovery. These run against a real compiled LangGraph graph with a checkpointer, so they actually verify the two claims the recovery rests on: 1. After a mid-run crash, ``aupdate_state(config, None, as_node=END)`` clears the stuck ``next`` tuple while preserving channel values. 2. A legitimate human-in-the-loop ``interrupt()`` (also a non-empty ``next``) is left intact, so a pending question is never silently discarded. """ from typing import Annotated, TypedDict from langchain_core.messages import AIMessage, HumanMessage, ToolMessage from langgraph.checkpoint.memory import InMemorySaver from langgraph.graph import END, START, StateGraph from langgraph.graph.message import add_messages from langgraph.types import interrupt from EvoScientist.stream.events import _clear_interrupted_graph_state class _S(TypedDict): x: int class _MessageState(TypedDict): messages: Annotated[list, add_messages] def _crashing_app(): # Node 'b' crashes once, then succeeds — so a post-recovery run can complete # and prove the graph is genuinely unstuck (not replaying the dead step). crashed = {"v": False} def a(state): return {"x": state["x"] + 1} def b(state): if not crashed["v"]: crashed["v"] = True raise RuntimeError("boom") return {"x": state["x"] + 100} g = StateGraph(_S) g.add_node("a", a) g.add_node("b", b) g.add_edge(START, "a") g.add_edge("a", "b") g.add_edge("b", END) return g.compile(checkpointer=InMemorySaver()) def _interrupting_app(): def ask(state): interrupt({"question": "continue?"}) return {"x": state["x"] + 1} g = StateGraph(_S) g.add_node("ask", ask) g.add_edge(START, "ask") g.add_edge("ask", END) return g.compile(checkpointer=InMemorySaver()) def _invalid_tool_call_app(): def write_invalid_call(state): return { "messages": [ AIMessage( content="", invalid_tool_calls=[ { "type": "invalid_tool_call", "id": None, "name": "execute", "args": '{"command":', "error": "bad json", } ], ) ] } def crash(state): raise RuntimeError("provider stream failed") g = StateGraph(_MessageState) g.add_node("write_invalid_call", write_invalid_call) g.add_node("crash", crash) g.add_edge(START, "write_invalid_call") g.add_edge("write_invalid_call", "crash") g.add_edge("crash", END) return g.compile(checkpointer=InMemorySaver()) def _repetitive_tool_call_app(): messages = [HumanMessage(content="inspect")] for call_id in ("call-1", "call-2"): messages.extend( [ AIMessage( content="", tool_calls=[ { "id": call_id, "name": "execute", "args": {"command": "pwd"}, } ], ), ToolMessage( content="/workspace", tool_call_id=call_id, name="execute", ), ] ) def write_repetitive_history(state): return {"messages": messages} def crash(state): raise RuntimeError("provider rejected repetitive tool history") g = StateGraph(_MessageState) g.add_node("write_repetitive_history", write_repetitive_history) g.add_node("crash", crash) g.add_edge(START, "write_repetitive_history") g.add_edge("write_repetitive_history", "crash") g.add_edge("crash", END) return g.compile(checkpointer=InMemorySaver()) async def test_recovery_clears_stuck_state_after_crash(): app = _crashing_app() cfg = {"configurable": {"thread_id": "t1"}} try: app.invoke({"x": 0}, cfg) except Exception: pass # LangGraph re-raises the node error (wrapped); we only care about state # The crash left the graph frozen at node 'b'. assert app.get_state(cfg).next == ("b",) await _clear_interrupted_graph_state(app, cfg) snap = app.get_state(cfg) assert snap.next == () # stuck state actually cleared assert snap.values == {"x": 1} # channel values (history) preserved # And the graph is genuinely unstuck: a fresh run completes (a: +1, b: +100) # instead of replaying the dead node. assert app.invoke({"x": 41}, cfg)["x"] == 142 async def test_recovery_preserves_pending_hitl_interrupt(): app = _interrupting_app() cfg = {"configurable": {"thread_id": "t1"}} app.invoke({"x": 0}, cfg) # parks at interrupt() before = app.get_state(cfg) assert before.next == ("ask",) assert before.interrupts await _clear_interrupted_graph_state(app, cfg) after = app.get_state(cfg) assert after.next == ("ask",) # interrupt left intact, still resumable assert after.interrupts async def test_recovery_removes_invalid_tool_call_from_checkpoint(): app = _invalid_tool_call_app() cfg = {"configurable": {"thread_id": "tool-history"}} try: await app.ainvoke( {"messages": [HumanMessage(content="run the command")]}, cfg, ) except RuntimeError: pass before = await app.aget_state(cfg) assert before.next == ("crash",) assert any( isinstance(message, AIMessage) and message.invalid_tool_calls for message in before.values["messages"] ) await _clear_interrupted_graph_state(app, cfg) after = await app.aget_state(cfg) assert after.next == () assert [message.type for message in after.values["messages"]] == ["human"] async def test_recovery_preserves_complete_repetitive_tool_rounds_in_checkpoint(): app = _repetitive_tool_call_app() cfg = {"configurable": {"thread_id": "tool-loop-history"}} try: await app.ainvoke({"messages": []}, cfg) except RuntimeError: pass before = await app.aget_state(cfg) assert before.next == ("crash",) assert len(before.values["messages"]) == 5 await _clear_interrupted_graph_state(app, cfg) after = await app.aget_state(cfg) messages = after.values["messages"] assert after.next == () assert [message.type for message in messages] == ["human", "ai", "tool", "ai", "tool"] assert [messages[1].tool_calls[0]["id"], messages[3].tool_calls[0]["id"]] == [ "call-1", "call-2", ]