From 3626478029cd61155560fdacb65d55f646179b5d Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 15 Jan 2025 13:31:49 -0800 Subject: [PATCH] Fix unexpected re-use of null resume value by subgraphs - Also stop exposing writes in config, in favor of scratchpad --- libs/langgraph/langgraph/constants.py | 4 +- libs/langgraph/langgraph/pregel/algo.py | 55 +++++---- libs/langgraph/langgraph/pregel/loop.py | 10 ++ libs/langgraph/langgraph/pregel/runner.py | 22 ++-- libs/langgraph/langgraph/types.py | 46 +++---- libs/langgraph/tests/test_pregel.py | 130 ++++++++++++++++++-- libs/langgraph/tests/test_pregel_async.py | 139 ++++++++++++++++++++-- 7 files changed, 321 insertions(+), 85 deletions(-) diff --git a/libs/langgraph/langgraph/constants.py b/libs/langgraph/langgraph/constants.py index cb6834f5f..85c4a890c 100644 --- a/libs/langgraph/langgraph/constants.py +++ b/libs/langgraph/langgraph/constants.py @@ -75,9 +75,7 @@ CONFIG_KEY_CHECKPOINT_ID = sys.intern("checkpoint_id") CONFIG_KEY_CHECKPOINT_NS = sys.intern("checkpoint_ns") # holds the current checkpoint_ns, "" for root graph CONFIG_KEY_NODE_FINISHED = sys.intern("__pregel_node_finished") -# holds the value that "answers" an interrupt() call -CONFIG_KEY_WRITES = sys.intern("__pregel_writes") -# read-only list of existing task writes +# holds a callback to be called when a node is finished CONFIG_KEY_SCRATCHPAD = sys.intern("__pregel_scratchpad") # holds a mutable dict for temporary storage scoped to the current task diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 12a53b8ad..c46138550 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -42,10 +42,10 @@ from langgraph.constants import ( CONFIG_KEY_SEND, CONFIG_KEY_STORE, CONFIG_KEY_TASK_ID, - CONFIG_KEY_WRITES, EMPTY_SEQ, ERROR, INTERRUPT, + MISSING, NO_WRITES, NS_END, NS_SEP, @@ -71,6 +71,7 @@ from langgraph.types import ( All, LoopProtocol, PregelExecutableTask, + PregelScratchpad, PregelTask, RetryPolicy, ) @@ -502,13 +503,10 @@ def prepare_single_task( }, CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, - CONFIG_KEY_WRITES: [ - w - for w in pending_writes - + configurable.get(CONFIG_KEY_WRITES, []) - if w[0] in (NULL_TASK_ID, task_id) - ], - CONFIG_KEY_SCRATCHPAD: {}, + CONFIG_KEY_SCRATCHPAD: _scratchpad( + pending_writes, + task_id, + ), }, ), triggers, @@ -614,13 +612,10 @@ def prepare_single_task( }, CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, - CONFIG_KEY_WRITES: [ - w - for w in pending_writes - + configurable.get(CONFIG_KEY_WRITES, []) - if w[0] in (NULL_TASK_ID, task_id) - ], - CONFIG_KEY_SCRATCHPAD: {}, + CONFIG_KEY_SCRATCHPAD: _scratchpad( + pending_writes, + task_id, + ), }, ), triggers, @@ -738,13 +733,10 @@ def prepare_single_task( }, CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, - CONFIG_KEY_WRITES: [ - w - for w in pending_writes - + configurable.get(CONFIG_KEY_WRITES, []) - if w[0] in (NULL_TASK_ID, task_id) - ], - CONFIG_KEY_SCRATCHPAD: {}, + CONFIG_KEY_SCRATCHPAD: _scratchpad( + pending_writes, + task_id, + ), }, ), triggers, @@ -758,6 +750,25 @@ def prepare_single_task( return PregelTask(task_id, name, task_path[:3]) +def _scratchpad( + pending_writes: Sequence[PendingWrite], + task_id: str, +) -> PregelScratchpad: + return PregelScratchpad( + # call + call_counter=0, + # interrupt + interrupt_counter=-1, + resume=next( + (w[2] for w in pending_writes if w[0] == task_id and w[1] == RESUME), [] + ), + null_resume=next( + (w[2] for w in pending_writes if w[0] == NULL_TASK_ID and w[1] == RESUME), + MISSING, + ), + ) + + def _proc_input( proc: PregelNode, managed: ManagedValueMapping, diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 5fd9b5d96..9a745a2e7 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -47,12 +47,14 @@ from langgraph.constants import ( CONFIG_KEY_DELEGATE, CONFIG_KEY_ENSURE_LATEST, CONFIG_KEY_RESUMING, + CONFIG_KEY_SCRATCHPAD, CONFIG_KEY_STREAM, CONFIG_KEY_TASK_ID, EMPTY_SEQ, ERROR, INPUT, INTERRUPT, + MISSING, NS_SEP, NULL_TASK_ID, PUSH, @@ -556,8 +558,16 @@ class PregelLoop(LoopProtocol): ) ) + # take resume value from parent + if scratchpad := configurable.get(CONFIG_KEY_SCRATCHPAD): + if scratchpad["null_resume"] is not MISSING: + self.put_writes(NULL_TASK_ID, [(RESUME, scratchpad["null_resume"])]) # map command to writes if isinstance(self.input, Command): + if self.input.resume is not None and not self.checkpointer: + raise RuntimeError( + "Cannot use Command(resume=...) without checkpointer" + ) writes: defaultdict[str, list[tuple[str, Any]]] = defaultdict(list) # group writes by task ID for tid, c, v in map_command(self.input, self.checkpoint_pending_writes): diff --git a/libs/langgraph/langgraph/pregel/runner.py b/libs/langgraph/langgraph/pregel/runner.py index 790d22a06..d354c7866 100644 --- a/libs/langgraph/langgraph/pregel/runner.py +++ b/libs/langgraph/langgraph/pregel/runner.py @@ -107,12 +107,11 @@ class PregelRunner: elif next_task.writes: # if it already ran, return the result fut = concurrent.futures.Future() - if ( - val := next( - (v for c, v in next_task.writes if c == RETURN), MISSING - ) - ) and val is not MISSING: - fut.set_result(val) + ret = next( + (v for c, v in next_task.writes if c == RETURN), MISSING + ) + if ret is not MISSING: + fut.set_result(ret) elif exc := next( (v for c, v in next_task.writes if c == ERROR), None ): @@ -295,12 +294,11 @@ class PregelRunner: elif next_task.writes: # if it already ran, return the result fut = asyncio.Future() - if ( - val := next( - (v for c, v in next_task.writes if c == RETURN), MISSING - ) - ) and val is not MISSING: - fut.set_result(val) + ret = next( + (v for c, v in next_task.writes if c == RETURN), MISSING + ) + if ret is not MISSING: + fut.set_result(ret) elif exc := next( (v for c, v in next_task.writes if c == ERROR), None ): diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 0b9fb9b1b..9a94ad85a 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -21,11 +21,7 @@ from typing import ( from langchain_core.runnables import Runnable, RunnableConfig from typing_extensions import Self, TypedDict -from langgraph.checkpoint.base import ( - BaseCheckpointSaver, - CheckpointMetadata, - PendingWrite, -) +from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata if TYPE_CHECKING: from langgraph.store.base import BaseStore @@ -341,13 +337,13 @@ class LoopProtocol: self.stop = stop -class PregelScratchpad(TypedDict, total=False): - # interrupt - interrupt_counter: int - used_null_resume: bool - resume: list[Any] +class PregelScratchpad(TypedDict): # call call_counter: int + # interrupt + interrupt_counter: int + resume: list[Any] + null_resume: Any def interrupt(value: Any) -> Any: @@ -449,10 +445,8 @@ def interrupt(value: Any) -> Any: CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_SCRATCHPAD, CONFIG_KEY_SEND, - CONFIG_KEY_TASK_ID, - CONFIG_KEY_WRITES, + MISSING, NS_SEP, - NULL_TASK_ID, RESUME, ) from langgraph.errors import GraphInterrupt @@ -461,29 +455,21 @@ def interrupt(value: Any) -> Any: conf = get_config()["configurable"] # track interrupt index scratchpad: PregelScratchpad = conf[CONFIG_KEY_SCRATCHPAD] - if "interrupt_counter" not in scratchpad: - scratchpad["interrupt_counter"] = 0 - else: - scratchpad["interrupt_counter"] += 1 + print("interrupt", scratchpad) + scratchpad["interrupt_counter"] += 1 idx = scratchpad["interrupt_counter"] # find previous resume values - task_id = conf[CONFIG_KEY_TASK_ID] - writes: list[PendingWrite] = conf[CONFIG_KEY_WRITES] - scratchpad.setdefault( - "resume", next((w[2] for w in writes if w[0] == task_id and w[1] == RESUME), []) - ) if scratchpad["resume"]: if idx < len(scratchpad["resume"]): return scratchpad["resume"][idx] # find current resume value - if not scratchpad.get("used_null_resume"): - scratchpad["used_null_resume"] = True - for tid, c, v in sorted(writes, key=lambda x: x[0], reverse=True): - if tid == NULL_TASK_ID and c == RESUME: - assert len(scratchpad["resume"]) == idx, (scratchpad["resume"], idx) - scratchpad["resume"].append(v) - conf[CONFIG_KEY_SEND]([(RESUME, scratchpad["resume"])]) - return v + if scratchpad["null_resume"] is not MISSING: + assert len(scratchpad["resume"]) == idx, (scratchpad["resume"], idx) + v = scratchpad["null_resume"] + scratchpad["null_resume"] = MISSING + scratchpad["resume"].append(v) + conf[CONFIG_KEY_SEND]([(RESUME, scratchpad["resume"])]) + return v # no resume value found raise GraphInterrupt( ( diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 12b7c6863..fe1a603da 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -5262,9 +5262,10 @@ def test_multiple_updates() -> None: ] -def test_falsy_return_from_task() -> None: +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) +def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_name: str): """Test with a falsy return from a task.""" - checkpointer = MemorySaver() + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") @task def falsy_task() -> bool: @@ -5276,17 +5277,18 @@ def test_falsy_return_from_task() -> None: falsy_task().result() interrupt("test") - configurable = {"configurable": {"thread_id": uuid.uuid4()}} + configurable = {"configurable": {"thread_id": str(uuid.uuid4())}} graph.invoke({"a": 5}, configurable) graph.invoke(Command(resume="123"), configurable) -def test_multiple_interrupts_imperative() -> None: +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) +def test_multiple_interrupts_imperative( + request: pytest.FixtureRequest, checkpointer_name: str +): """Test multiple interrupts with an imperative API.""" - from langgraph.checkpoint.memory import MemorySaver - from langgraph.func import entrypoint, task + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") - checkpointer = MemorySaver() counter = 0 @task @@ -5307,7 +5309,7 @@ def test_multiple_interrupts_imperative() -> None: return {"values": values} - configurable = {"configurable": {"thread_id": uuid.uuid4()}} + configurable = {"configurable": {"thread_id": str(uuid.uuid4())}} graph.invoke({}, configurable) graph.invoke(Command(resume="a"), configurable) graph.invoke(Command(resume="b"), configurable) @@ -5317,3 +5319,115 @@ def test_multiple_interrupts_imperative() -> None: "values": [2, "a", 4, "b", 6, "c"], } assert counter == 3 + + +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) +def test_double_interrupt_subgraph( + request: pytest.FixtureRequest, checkpointer_name: str +) -> None: + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") + + class AgentState(TypedDict): + input: str + + def node_1(state: AgentState): + result = interrupt("interrupt node 1") + return {"input": result} + + def node_2(state: AgentState): + result = interrupt("interrupt node 2") + return {"input": result} + + subgraph_builder = ( + StateGraph(AgentState) + .add_node("node_1", node_1) + .add_node("node_2", node_2) + .add_edge(START, "node_1") + .add_edge("node_1", "node_2") + .add_edge("node_2", END) + ) + + # invoke the sub graph + subgraph = subgraph_builder.compile(checkpointer=checkpointer) + thread = {"configurable": {"thread_id": str(uuid.uuid4())}} + assert [c for c in subgraph.stream({"input": "test"}, thread)] == [ + { + "__interrupt__": ( + Interrupt( + value="interrupt node 1", + resumable=True, + ns=[AnyStr("node_1:")], + when="during", + ), + ) + }, + ] + # resume from the first interrupt + assert [c for c in subgraph.stream(Command(resume="123"), thread)] == [ + { + "node_1": {"input": "123"}, + }, + { + "__interrupt__": ( + Interrupt( + value="interrupt node 2", + resumable=True, + ns=[AnyStr("node_2:")], + when="during", + ), + ) + }, + ] + # resume from the second interrupt + assert [c for c in subgraph.stream(Command(resume="123"), thread)] == [ + { + "node_2": {"input": "123"}, + }, + ] + + subgraph = subgraph_builder.compile() + + def invoke_sub_agent(state: AgentState): + return subgraph.invoke(state) + + parent_agent = ( + StateGraph(AgentState) + .add_node("invoke_sub_agent", invoke_sub_agent) + .add_edge(START, "invoke_sub_agent") + .add_edge("invoke_sub_agent", END) + .compile(checkpointer=checkpointer) + ) + + assert [c for c in parent_agent.stream({"input": "test"}, thread)] == [ + { + "__interrupt__": ( + Interrupt( + value="interrupt node 1", + resumable=True, + ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_1:")], + when="during", + ), + ) + }, + ] + + # resume from the first interrupt + assert [c for c in parent_agent.stream(Command(resume=True), thread)] == [ + { + "__interrupt__": ( + Interrupt( + value="interrupt node 2", + resumable=True, + ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_2:")], + when="during", + ), + ) + } + ] + + # resume from 2nd interrupt + assert [c for c in parent_agent.stream(Command(resume=True), thread)] == [ + { + "invoke_sub_agent": {"input": True}, + }, + ] diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index c05682f7e..2f7d7a1c4 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -6696,23 +6696,25 @@ async def test_multiple_updates() -> None: sys.version_info < (3, 11), reason="Python 3.11+ is required for async contextvars support", ) -async def test_falsy_return_from_task() -> None: +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) +async def test_falsy_return_from_task(checkpointer_name: str) -> None: """Test with a falsy return from a task.""" - checkpointer = MemorySaver() @task async def falsy_task() -> bool: return False - @entrypoint(checkpointer=checkpointer) - async def graph(state: dict) -> dict: - """React tool.""" - await falsy_task() - interrupt("test") + async with awith_checkpointer(checkpointer_name) as checkpointer: - configurable = {"configurable": {"thread_id": uuid.uuid4()}} - await graph.ainvoke({"a": 5}, configurable) - await graph.ainvoke(Command(resume="123"), configurable) + @entrypoint(checkpointer=checkpointer) + async def graph(state: dict) -> dict: + """React tool.""" + await falsy_task() + interrupt("test") + + configurable = {"configurable": {"thread_id": str(uuid.uuid4())}} + await graph.ainvoke({"a": 5}, configurable) + await graph.ainvoke(Command(resume="123"), configurable) @pytest.mark.skipif( @@ -6756,3 +6758,120 @@ async def test_multiple_interrupts_imperative(checkpointer_name: str) -> None: "values": [2, "a", 4, "b", 6, "c"], } assert counter == 3 + + +@pytest.mark.skipif( + sys.version_info < (3, 11), + reason="Python 3.11+ is required for async contextvars support", +) +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) +async def test_double_interrupt_subgraph(checkpointer_name: str) -> None: + class AgentState(TypedDict): + input: str + + def node_1(state: AgentState): + result = interrupt("interrupt node 1") + return {"input": result} + + def node_2(state: AgentState): + result = interrupt("interrupt node 2") + return {"input": result} + + subgraph_builder = ( + StateGraph(AgentState) + .add_node("node_1", node_1) + .add_node("node_2", node_2) + .add_edge(START, "node_1") + .add_edge("node_1", "node_2") + .add_edge("node_2", END) + ) + + async with awith_checkpointer(checkpointer_name) as checkpointer: + # invoke the sub graph + subgraph = subgraph_builder.compile(checkpointer=checkpointer) + thread = {"configurable": {"thread_id": str(uuid.uuid4())}} + assert [c async for c in subgraph.astream({"input": "test"}, thread)] == [ + { + "__interrupt__": ( + Interrupt( + value="interrupt node 1", + resumable=True, + ns=[AnyStr("node_1:")], + when="during", + ), + ) + }, + ] + # resume from the first interrupt + assert [c async for c in subgraph.astream(Command(resume="123"), thread)] == [ + { + "node_1": {"input": "123"}, + }, + { + "__interrupt__": ( + Interrupt( + value="interrupt node 2", + resumable=True, + ns=[AnyStr("node_2:")], + when="during", + ), + ) + }, + ] + # resume from the second interrupt + assert [c async for c in subgraph.astream(Command(resume="123"), thread)] == [ + { + "node_2": {"input": "123"}, + }, + ] + + subgraph = subgraph_builder.compile() + + def invoke_sub_agent(state: AgentState): + return subgraph.invoke(state) + + parent_agent = ( + StateGraph(AgentState) + .add_node("invoke_sub_agent", invoke_sub_agent) + .add_edge(START, "invoke_sub_agent") + .add_edge("invoke_sub_agent", END) + .compile(checkpointer=checkpointer) + ) + + assert [c async for c in parent_agent.astream({"input": "test"}, thread)] == [ + { + "__interrupt__": ( + Interrupt( + value="interrupt node 1", + resumable=True, + ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_1:")], + when="during", + ), + ) + }, + ] + + # resume from the first interrupt + assert [ + c async for c in parent_agent.astream(Command(resume=True), thread) + ] == [ + { + "__interrupt__": ( + Interrupt( + value="interrupt node 2", + resumable=True, + ns=[AnyStr("invoke_sub_agent:"), AnyStr("node_2:")], + when="during", + ), + ) + } + ] + + # resume from 2nd interrupt + assert [ + c async for c in parent_agent.astream(Command(resume=True), thread) + ] == [ + { + "invoke_sub_agent": {"input": True}, + }, + ]