From 31a7bcf75048096787c2c9273ce8c9cba1f4500a Mon Sep 17 00:00:00 2001 From: Vadym Barda Date: Thu, 20 Feb 2025 13:15:27 -0500 Subject: [PATCH] langgraph: handle non-overlapping subgraph updates in Command.PARENT (#3521) --- libs/langgraph/langgraph/graph/state.py | 8 +++- libs/langgraph/langgraph/pregel/loop.py | 14 ++++-- libs/langgraph/tests/test_pregel.py | 61 +++++++++++++++++++++++++ 3 files changed, 77 insertions(+), 6 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 872c64d95..a48356df5 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -675,7 +675,9 @@ class CompiledStateGraph(CompiledGraph): elif isinstance(input, Command): if input.graph == Command.PARENT: return None - return input._update_as_tuples() + return [ + (k, v) for k, v in input._update_as_tuples() if k in output_keys + ] elif ( isinstance(input, (list, tuple)) and input @@ -686,7 +688,9 @@ class CompiledStateGraph(CompiledGraph): if isinstance(i, Command): if i.graph == Command.PARENT: continue - updates.extend(i._update_as_tuples()) + updates.extend( + (k, v) for k, v in i._update_as_tuples() if k in output_keys + ) else: updates.extend(_get_updates(i) or ()) return updates diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 1c9ea1c8f..0f70c1499 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -888,7 +888,11 @@ class SyncPregelLoop(PregelLoop, ContextManager): ) def _update_mv(self, key: str, values: Sequence[Any]) -> None: - return self.submit(cast(WritableManagedValue, self.managed[key]).update, values) + managed_value = self.managed.get(key) + if managed_value is None: + return + + return self.submit(cast(WritableManagedValue, managed_value).update, values) # context manager @@ -1023,9 +1027,11 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): ) def _update_mv(self, key: str, values: Sequence[Any]) -> None: - return self.submit( - cast(WritableManagedValue, self.managed[key]).aupdate, values - ) + managed_value = self.managed.get(key) + if managed_value is None: + return + + return self.submit(cast(WritableManagedValue, managed_value).aupdate, values) # context manager diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 992b4664f..bb5c7d972 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -6260,6 +6260,67 @@ def test_merging_updates_command_parent(): ] +def test_merging_non_overlapping_updates_command_parent(): + # simple reducer + def append_unique(left, right): + combined = list(left) + for item in right: + if item in combined: + continue + else: + combined.append(item) + return combined + + class State(TypedDict): + foo: Annotated[list, append_unique] + + # Define subgraph + def subgraph_node_1(state: State): + return Command( + goto="subgraph_node_2", + update={ + "foo": ["bar"], + "bar": ["subgraph_node_1"], + }, + ) + + def subgraph_node_2(state: State): + return Command( + goto="node_3", + update={"bar": ["subgraph_node_2"]}, + graph=Command.PARENT, + ) + + subgraph_builder = StateGraph(State) + subgraph_builder.add_node(subgraph_node_1) + subgraph_builder.add_node(subgraph_node_2) + subgraph_builder.add_edge(START, "subgraph_node_1") + + # Define main graph + def node_1(state: State): + return Command( + goto="node_2", + update={"foo": ["foo"]}, + ) + + def node_3(state: State, store): + return Command( + update={"foo": ["baz"]}, + ) + + main_builder = StateGraph(State) + main_builder.add_node("node_1", node_1) + main_builder.add_node("node_2", subgraph_builder.compile()) + main_builder.add_node("node_3", node_3) + main_builder.add_edge(START, "node_1") + main_builder.add_edge("node_2", "node_3") + main_graph = main_builder.compile() + + assert main_graph.invoke({"foo": []}) == { + "foo": ["foo", "bar", "baz"], + } + + def test_entrypoint_output_schema_with_return_and_save() -> None: """Test output schema inference with entrypoint.final."""