218 lines
6.6 KiB
Python
218 lines
6.6 KiB
Python
"""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",
|
|
]
|