From 5375af782725b382935fd9f752e6787117d97119 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 15 Jan 2025 15:15:44 -0800 Subject: [PATCH 1/4] Add checkpointer=True mode for subgraphs that want to keep state betweenn turns --- libs/langgraph/langgraph/pregel/__init__.py | 9 ++++ libs/langgraph/langgraph/types.py | 8 +-- libs/langgraph/tests/test_pregel.py | 60 +++++++++++++++++++++ 3 files changed, 74 insertions(+), 3 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index cc8786968..7fd3ce646 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -1455,6 +1455,8 @@ class Pregel(PregelProtocol): checkpointer: Optional[BaseCheckpointSaver] = None elif CONFIG_KEY_CHECKPOINTER in config.get(CONF, {}): checkpointer = config[CONF][CONFIG_KEY_CHECKPOINTER] + elif self.checkpointer is True: + raise RuntimeError("checkpointer=True cannot be used for root graphs.") else: checkpointer = self.checkpointer if checkpointer and not config.get(CONF): @@ -1598,6 +1600,12 @@ class Pregel(PregelProtocol): interrupt_after=interrupt_after, debug=debug, ) + # set up subgraph checkpointing + if self.checkpointer is True: + ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS]) + config[CONF][CONFIG_KEY_CHECKPOINT_NS] = NS_SEP.join( + part.split(NS_END)[0] for part in ns.split(NS_SEP) + ) # set up messages stream mode if "messages" in stream_modes: run_manager.inheritable_handlers.append( @@ -1622,6 +1630,7 @@ class Pregel(PregelProtocol): interrupt_after=interrupt_after_, manager=run_manager, debug=debug, + check_subgraphs=self.checkpointer is not True, ) as loop: # create runner runner = PregelRunner( diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 076ee82f7..3b1bcd213 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -38,9 +38,11 @@ except ImportError: All = Literal["*"] """Special value to indicate that graph should interrupt on all nodes.""" -Checkpointer = Union[None, Literal[False], BaseCheckpointSaver] -"""Type of the checkpointer to use for a subgraph. False disables checkpointing, -even if the parent graph has a checkpointer. None inherits checkpointer.""" +Checkpointer = Union[None, bool, BaseCheckpointSaver] +"""Type of the checkpointer to use for a subgraph. +- True enables persistent checkpointing for this subgraph. +- False disables checkpointing, even if the parent graph has a checkpointer. +- None inherits checkpointer from the parent graph.""" StreamMode = Literal["values", "updates", "debug", "messages", "custom"] """How the stream method should emit outputs. diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index fe1a603da..67536deed 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -3189,6 +3189,66 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: ] +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) +def test_subgraph_checkpoint_true( + request: pytest.FixtureRequest, checkpointer_name: str +) -> None: + checkpointer = request.getfixturevalue("checkpointer_" + checkpointer_name) + + class InnerState(TypedDict): + my_key: Annotated[str, operator.add] + my_other_key: str + + def inner_1(state: InnerState): + return {"my_key": " got here", "my_other_key": state["my_key"]} + + def inner_2(state: InnerState): + return {"my_key": " and there"} + + inner = StateGraph(InnerState) + inner.add_node("inner_1", inner_1) + inner.add_node("inner_2", inner_2) + inner.add_edge("inner_1", "inner_2") + inner.set_entry_point("inner_1") + inner.set_finish_point("inner_2") + + class State(TypedDict): + my_key: str + + graph = StateGraph(State) + graph.add_node("inner", inner.compile(checkpointer=True)) + graph.add_edge(START, "inner") + graph.add_conditional_edges( + "inner", lambda s: "inner" if s["my_key"].count("there") < 2 else END + ) + app = graph.compile(checkpointer=checkpointer) + + config = {"configurable": {"thread_id": "2"}} + assert [c for c in app.stream({"my_key": ""}, config, subgraphs=True)] == [ + (("inner",), {"inner_1": {"my_key": " got here", "my_other_key": ""}}), + (("inner",), {"inner_2": {"my_key": " and there"}}), + ((), {"inner": {"my_key": " got here and there"}}), + ( + ("inner",), + { + "inner_1": { + "my_key": " got here", + "my_other_key": " got here and there got here and there", + } + }, + ), + (("inner",), {"inner_2": {"my_key": " and there"}}), + ( + (), + { + "inner": { + "my_key": " got here and there got here and there got here and there" + } + }, + ), + ] + + @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_stream_subgraphs_during_execution( request: pytest.FixtureRequest, checkpointer_name: str From 71fbd6a8b489301e0b17b0937177bd407535f9b8 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 15 Jan 2025 15:29:34 -0800 Subject: [PATCH 2/4] Lint --- libs/langgraph/langgraph/pregel/__init__.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 7fd3ce646..b249927b5 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -494,7 +494,9 @@ class Pregel(PregelProtocol): saved.metadata.get("step", -1) + 1, for_execution=True, store=self.store, - checkpointer=self.checkpointer or None, + checkpointer=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, manager=None, ) # get the subgraphs @@ -606,7 +608,9 @@ class Pregel(PregelProtocol): saved.metadata.get("step", -1) + 1, for_execution=True, store=self.store, - checkpointer=self.checkpointer or None, + checkpointer=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, manager=None, ) # get the subgraphs From be8b4a1d7f51ce26f7a2dd0c49161b428e22bd91 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 15 Jan 2025 15:30:29 -0800 Subject: [PATCH 3/4] Lint --- libs/langgraph/langgraph/pregel/__init__.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index b249927b5..d373e2440 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -930,7 +930,9 @@ class Pregel(PregelProtocol): saved.metadata.get("step", -1) + 1, for_execution=True, store=self.store, - checkpointer=self.checkpointer or None, + checkpointer=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, manager=None, ) # apply null writes @@ -1024,7 +1026,9 @@ class Pregel(PregelProtocol): saved.metadata.get("step", -1) + 1, for_execution=True, store=self.store, - checkpointer=self.checkpointer or None, + checkpointer=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, manager=None, ) # apply null writes @@ -1213,7 +1217,9 @@ class Pregel(PregelProtocol): saved.metadata.get("step", -1) + 1, for_execution=True, store=self.store, - checkpointer=self.checkpointer or None, + checkpointer=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, manager=None, ) # apply null writes @@ -1307,7 +1313,9 @@ class Pregel(PregelProtocol): saved.metadata.get("step", -1) + 1, for_execution=True, store=self.store, - checkpointer=self.checkpointer or None, + checkpointer=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, manager=None, ) # apply null writes From 38d9b39f6eaad310d456aaf00ba039605b39916e Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 15 Jan 2025 15:36:03 -0800 Subject: [PATCH 4/4] Add flag --- libs/langgraph/langgraph/pregel/__init__.py | 1 + 1 file changed, 1 insertion(+) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index d373e2440..102e68be8 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -1870,6 +1870,7 @@ class Pregel(PregelProtocol): interrupt_after=interrupt_after_, manager=run_manager, debug=debug, + check_subgraphs=self.checkpointer is not True, ) as loop: # create runner runner = PregelRunner(