Add tests, missing return stmt

This commit is contained in:
Nuno Campos
2024-12-06 08:17:50 -08:00
parent 85fc26db43
commit 015bf5e0a6
2 changed files with 29 additions and 0 deletions
+1
View File
@@ -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)]
+28
View File
@@ -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"}