mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-17 23:27:56 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a07be124e8 | ||
|
|
6d9f9f5326 |
@@ -32,10 +32,14 @@ from langgraph._internal._constants import (
|
|||||||
CONFIG_KEY_SCRATCHPAD,
|
CONFIG_KEY_SCRATCHPAD,
|
||||||
ERROR,
|
ERROR,
|
||||||
ERROR_SOURCE_NODE,
|
ERROR_SOURCE_NODE,
|
||||||
|
INPUT,
|
||||||
INTERRUPT,
|
INTERRUPT,
|
||||||
NO_WRITES,
|
NO_WRITES,
|
||||||
|
NULL_TASK_ID,
|
||||||
|
PUSH,
|
||||||
RESUME,
|
RESUME,
|
||||||
RETURN,
|
RETURN,
|
||||||
|
TASKS,
|
||||||
)
|
)
|
||||||
from langgraph._internal._future import chain_future, run_coroutine_threadsafe
|
from langgraph._internal._future import chain_future, run_coroutine_threadsafe
|
||||||
from langgraph._internal._scratchpad import PregelScratchpad
|
from langgraph._internal._scratchpad import PregelScratchpad
|
||||||
@@ -585,6 +589,12 @@ class PregelRunner:
|
|||||||
if isinstance(exception, GraphInterrupt):
|
if isinstance(exception, GraphInterrupt):
|
||||||
# save interrupt to checkpointer
|
# save interrupt to checkpointer
|
||||||
if exception.args[0]:
|
if exception.args[0]:
|
||||||
|
if pending_writes := _writes_to_persist_on_interrupt(task.writes):
|
||||||
|
# GraphInterrupt is a controlled suspension point, not a
|
||||||
|
# task failure. Persist writes emitted before the
|
||||||
|
# suspension without associating them with task
|
||||||
|
# completion, so the interrupted task remains resumable.
|
||||||
|
self.put_writes()(NULL_TASK_ID, pending_writes) # type: ignore[misc]
|
||||||
writes = [(INTERRUPT, exception.args[0])]
|
writes = [(INTERRUPT, exception.args[0])]
|
||||||
if resumes := [w for w in task.writes if w[0] == RESUME]:
|
if resumes := [w for w in task.writes if w[0] == RESUME]:
|
||||||
writes.extend(resumes)
|
writes.extend(resumes)
|
||||||
@@ -939,3 +949,24 @@ async def _acall_impl(
|
|||||||
destination.set_exception(RuntimeError("Task not scheduled"))
|
destination.set_exception(RuntimeError("Task not scheduled"))
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
destination.set_exception(exc)
|
destination.set_exception(exc)
|
||||||
|
|
||||||
|
|
||||||
|
def _writes_to_persist_on_interrupt(
|
||||||
|
writes: Iterable[tuple[str, Any]],
|
||||||
|
) -> list[tuple[str, Any]]:
|
||||||
|
return [
|
||||||
|
write
|
||||||
|
for write in writes
|
||||||
|
if write[0]
|
||||||
|
not in (
|
||||||
|
ERROR,
|
||||||
|
ERROR_SOURCE_NODE,
|
||||||
|
INPUT,
|
||||||
|
INTERRUPT,
|
||||||
|
NO_WRITES,
|
||||||
|
PUSH,
|
||||||
|
RESUME,
|
||||||
|
RETURN,
|
||||||
|
TASKS,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|||||||
@@ -40,15 +40,25 @@ from pytest_mock import MockerFixture
|
|||||||
from syrupy import SnapshotAssertion
|
from syrupy import SnapshotAssertion
|
||||||
from typing_extensions import NotRequired, TypedDict
|
from typing_extensions import NotRequired, TypedDict
|
||||||
|
|
||||||
from langgraph._internal._constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL
|
from langgraph._internal._constants import (
|
||||||
|
CONFIG_KEY_NODE_FINISHED,
|
||||||
|
CONFIG_KEY_SEND,
|
||||||
|
ERROR,
|
||||||
|
PULL,
|
||||||
|
)
|
||||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||||
from langgraph.channels.delta import DeltaChannel
|
from langgraph.channels.delta import DeltaChannel
|
||||||
from langgraph.channels.ephemeral_value import EphemeralValue
|
from langgraph.channels.ephemeral_value import EphemeralValue
|
||||||
from langgraph.channels.last_value import LastValue
|
from langgraph.channels.last_value import LastValue
|
||||||
from langgraph.channels.topic import Topic
|
from langgraph.channels.topic import Topic
|
||||||
from langgraph.channels.untracked_value import UntrackedValue
|
from langgraph.channels.untracked_value import UntrackedValue
|
||||||
from langgraph.config import get_stream_writer
|
from langgraph.config import get_config, get_stream_writer
|
||||||
from langgraph.errors import GraphRecursionError, InvalidUpdateError, ParentCommand
|
from langgraph.errors import (
|
||||||
|
GraphInterrupt,
|
||||||
|
GraphRecursionError,
|
||||||
|
InvalidUpdateError,
|
||||||
|
ParentCommand,
|
||||||
|
)
|
||||||
from langgraph.func import entrypoint, task
|
from langgraph.func import entrypoint, task
|
||||||
from langgraph.graph import END, START, StateGraph
|
from langgraph.graph import END, START, StateGraph
|
||||||
from langgraph.graph.message import MessagesState, _messages_delta_reducer, add_messages
|
from langgraph.graph.message import MessagesState, _messages_delta_reducer, add_messages
|
||||||
@@ -4820,6 +4830,111 @@ def test_parent_command(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_pending_send_write_persists_on_interrupt(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
class State(TypedDict):
|
||||||
|
quickjs_checkpoint: dict
|
||||||
|
result: str
|
||||||
|
|
||||||
|
def node(state: State) -> dict[str, Any]:
|
||||||
|
send = get_config()["configurable"][CONFIG_KEY_SEND]
|
||||||
|
send([("quickjs_checkpoint", {"snapshot": "abc"})])
|
||||||
|
value = interrupt("trip")
|
||||||
|
return {"result": value}
|
||||||
|
|
||||||
|
builder = StateGraph(State)
|
||||||
|
builder.add_node("node", node)
|
||||||
|
builder.add_edge(START, "node")
|
||||||
|
builder.add_edge("node", END)
|
||||||
|
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||||
|
config = {"configurable": {"thread_id": "1"}}
|
||||||
|
|
||||||
|
first = graph.invoke({"quickjs_checkpoint": {}, "result": ""}, config)
|
||||||
|
assert "__interrupt__" in first
|
||||||
|
|
||||||
|
snapshot = graph.get_state(config)
|
||||||
|
assert snapshot.values["quickjs_checkpoint"] == {"snapshot": "abc"}
|
||||||
|
assert snapshot.next == ("node",)
|
||||||
|
assert snapshot.tasks[0].interrupts
|
||||||
|
|
||||||
|
second = graph.invoke(Command(resume="ok"), config)
|
||||||
|
assert second["quickjs_checkpoint"] == {"snapshot": "abc"}
|
||||||
|
assert second["result"] == "ok"
|
||||||
|
assert graph.get_state(config).values["quickjs_checkpoint"] == {"snapshot": "abc"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_pending_send_write_persists_on_tool_subgraph_interrupt(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
from langchain_core.tools import tool
|
||||||
|
|
||||||
|
class SubgraphState(TypedDict):
|
||||||
|
value: str
|
||||||
|
|
||||||
|
def subgraph_node(state: SubgraphState) -> dict[str, Any]:
|
||||||
|
value = interrupt("subgraph trip")
|
||||||
|
return {"value": value}
|
||||||
|
|
||||||
|
subgraph_builder = StateGraph(SubgraphState)
|
||||||
|
subgraph_builder.add_node("subgraph_node", subgraph_node)
|
||||||
|
subgraph_builder.add_edge(START, "subgraph_node")
|
||||||
|
subgraph = subgraph_builder.compile(checkpointer=True)
|
||||||
|
|
||||||
|
class ParentState(TypedDict):
|
||||||
|
messages: Annotated[list[AnyMessage], add_messages]
|
||||||
|
quickjs_checkpoint: dict
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def task_tool() -> str:
|
||||||
|
"""Run the interrupting subgraph task."""
|
||||||
|
send = get_config()["configurable"][CONFIG_KEY_SEND]
|
||||||
|
try:
|
||||||
|
response = subgraph.invoke({"value": ""})
|
||||||
|
except GraphInterrupt:
|
||||||
|
send([("quickjs_checkpoint", {"snapshot": "tool"})])
|
||||||
|
raise
|
||||||
|
return response["value"]
|
||||||
|
|
||||||
|
def call_tool(state: ParentState) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"messages": AIMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[
|
||||||
|
{
|
||||||
|
"name": "task_tool",
|
||||||
|
"args": {},
|
||||||
|
"id": "tool_call_1",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
builder = StateGraph(ParentState)
|
||||||
|
builder.add_node("agent", call_tool)
|
||||||
|
builder.add_node("tools", ToolNode([task_tool]))
|
||||||
|
builder.add_edge(START, "agent")
|
||||||
|
builder.add_edge("agent", "tools")
|
||||||
|
builder.add_edge("tools", END)
|
||||||
|
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||||
|
config = {"configurable": {"thread_id": "1"}}
|
||||||
|
|
||||||
|
first = graph.invoke(
|
||||||
|
{"messages": [HumanMessage(content="start")], "quickjs_checkpoint": {}}, config
|
||||||
|
)
|
||||||
|
assert "__interrupt__" in first
|
||||||
|
|
||||||
|
snapshot = graph.get_state(config)
|
||||||
|
assert snapshot.values["quickjs_checkpoint"] == {"snapshot": "tool"}
|
||||||
|
assert snapshot.next == ("tools",)
|
||||||
|
assert snapshot.tasks[0].interrupts
|
||||||
|
|
||||||
|
second = graph.invoke(Command(resume="ok"), config)
|
||||||
|
assert second["quickjs_checkpoint"] == {"snapshot": "tool"}
|
||||||
|
assert second["messages"][-1].content == "ok"
|
||||||
|
assert graph.get_state(config).values["quickjs_checkpoint"] == {"snapshot": "tool"}
|
||||||
|
|
||||||
|
|
||||||
def test_interrupt_subgraph(sync_checkpointer: BaseCheckpointSaver):
|
def test_interrupt_subgraph(sync_checkpointer: BaseCheckpointSaver):
|
||||||
class State(TypedDict):
|
class State(TypedDict):
|
||||||
baz: str
|
baz: str
|
||||||
|
|||||||
Reference in New Issue
Block a user