From 85fc26db43f9d9ba79b1105e48178f5b4ddefa15 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 6 Dec 2024 08:03:26 -0800 Subject: [PATCH 1/3] lib: Support returning mixed list of commands and state updates --- libs/langgraph/langgraph/graph/state.py | 37 ++++++++++++++----------- 1 file changed, 21 insertions(+), 16 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 6e37c6150..0c4d0c15c 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -613,21 +613,23 @@ class CompiledStateGraph(CompiledGraph): ] def _get_root(input: Any) -> Optional[Sequence[tuple[str, Any]]]: - if ( - isinstance(input, (list, tuple)) - and input - and all(isinstance(i, Command) for i in input) - ): - updates: list[tuple[str, Any]] = [] - for i in input: - if i.graph == Command.PARENT: - continue - updates.extend(i._update_as_tuples()) - return updates - elif isinstance(input, Command): + if isinstance(input, Command): if input.graph == Command.PARENT: return () return input._update_as_tuples() + elif ( + isinstance(input, (list, tuple)) + and input + and any(isinstance(i, Command) for i in input) + ): + updates: list[tuple[str, Any]] = [] + for i in input: + if isinstance(i, Command): + if i.graph == Command.PARENT: + continue + updates.extend(i._update_as_tuples()) + else: + updates.append(("__root__", i)) elif input is not None: return [("__root__", input)] @@ -645,13 +647,16 @@ class CompiledStateGraph(CompiledGraph): elif ( isinstance(input, (list, tuple)) and input - and all(isinstance(i, Command) for i in input) + and any(isinstance(i, Command) for i in input) ): updates: list[tuple[str, Any]] = [] for i in input: - if i.graph == Command.PARENT: - continue - updates.extend(i._update_as_tuples()) + if isinstance(i, Command): + if i.graph == Command.PARENT: + continue + updates.extend(i._update_as_tuples()) + else: + updates.extend(_get_updates(i) or ()) return updates elif get_type_hints(type(input)): return [ From 015bf5e0a6139629e11b8947518ae116cf50049d Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 6 Dec 2024 08:17:50 -0800 Subject: [PATCH 2/3] Add tests, missing return stmt --- libs/langgraph/langgraph/graph/state.py | 1 + libs/langgraph/tests/test_pregel.py | 28 +++++++++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 0c4d0c15c..c416d5f6a 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -630,6 +630,7 @@ class CompiledStateGraph(CompiledGraph): updates.extend(i._update_as_tuples()) else: updates.append(("__root__", i)) + return updates elif input is not None: return [("__root__", input)] diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 2515270a1..a3f5a6d48 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -14807,3 +14807,31 @@ def test_interrupt_loop(request: pytest.FixtureRequest, checkpointer_name: str): assert [event for event in graph.stream(Command(resume="19"), thread1)] == [ {"node": {"age": 19}}, ] + + +def test_root_mixed_return() -> None: + def my_node(state: list[str]): + return [Command(update=["a"]), ["b"]] + + graph = StateGraph(Annotated[list[str], operator.add]) + + graph.add_node(my_node) + graph.add_edge(START, "my_node") + graph = graph.compile() + + assert graph.invoke([]) == ["a", "b"] + + +def test_dict_mixed_return() -> None: + class State(TypedDict): + foo: Annotated[str, operator.add] + + def my_node(state: State): + return [Command(update={"foo": "a"}), {"foo": "b"}] + + graph = StateGraph(State) + graph.add_node(my_node) + graph.add_edge(START, "my_node") + graph = graph.compile() + + assert graph.invoke({"foo": ""}) == {"foo": "ab"} From 4e0e9a4effe38e150089aa433c6c43c8bd0b2219 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 6 Dec 2024 08:20:41 -0800 Subject: [PATCH 3/3] Fix --- libs/langgraph/tests/test_pregel_async.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 514703781..b278d9f77 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -9670,14 +9670,14 @@ async def test_stream_subgraphs_during_execution(checkpointer_name: str) -> None ), (FloatBetween(0.2, 0.4), ((), {"outer_1": {"my_key": " and parallel"}})), ( - FloatBetween(0.5, 0.7), + FloatBetween(0.5, 0.8), ( (AnyStr("inner:"),), {"inner_2": {"my_key": " and there", "my_other_key": "got here"}}, ), ), - (FloatBetween(0.5, 0.7), ((), {"inner": {"my_key": "got here and there"}})), - (FloatBetween(0.5, 0.7), ((), {"outer_2": {"my_key": " and back again"}})), + (FloatBetween(0.5, 0.8), ((), {"inner": {"my_key": "got here and there"}})), + (FloatBetween(0.5, 0.8), ((), {"outer_2": {"my_key": " and back again"}})), ]