From 015bf5e0a6139629e11b8947518ae116cf50049d Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 6 Dec 2024 08:17:50 -0800 Subject: [PATCH] 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"}