From 1e3953d1e05907e26a7de8680c58bbd59323c20a Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 1 Nov 2024 13:23:10 -0700 Subject: [PATCH] Test order of update application after Send - updates from inside Send tasks are applied in the order the Sends were created, if when you fan out, and have each task write results to a list with reducer, the final list is in the order you used when triggering --- .../langgraph/checkpoint/postgres/base.py | 2 +- libs/langgraph/tests/test_pregel.py | 6 +-- libs/langgraph/tests/test_pregel_async.py | 49 ++++++++++++++----- 3 files changed, 40 insertions(+), 17 deletions(-) diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py index e4f930294..5f6a2ab1b 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py @@ -84,7 +84,7 @@ select and cw.checkpoint_id = checkpoints.checkpoint_id ) as pending_writes, ( - select array_agg(array[cw.type::bytea, cw.blob] order by cw.idx) + select array_agg(array[cw.type::bytea, cw.blob] order by cw.task_id, cw.idx) from checkpoint_writes cw where cw.thread_id = checkpoints.thread_id and cw.checkpoint_ns = checkpoints.checkpoint_ns diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 13bc25236..f8a5b49a8 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -8689,14 +8689,14 @@ def test_stream_subgraphs_during_execution( ), (FloatBetween(0.2, 0.3), ((), {"outer_1": {"my_key": " and parallel"}})), ( - FloatBetween(0.5, 0.6), + FloatBetween(0.5, 0.8), ( (AnyStr("inner:"),), {"inner_2": {"my_key": " and there", "my_other_key": "got here"}}, ), ), - (FloatBetween(0.5, 0.6), ((), {"inner": {"my_key": "got here and there"}})), - (FloatBetween(0.5, 0.6), ((), {"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"}})), ] diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index fe4bee325..bedbcef4d 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -1,5 +1,6 @@ import asyncio import operator +import random import re import sys import uuid @@ -1922,7 +1923,8 @@ async def test_cond_edge_after_send() -> None: assert await graph.ainvoke(["0"]) == ["0", "1", "2", "2", "3"] -async def test_max_concurrency() -> None: +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) +async def test_max_concurrency(checkpointer_name: str) -> None: class Node: def __init__(self, name: str): self.name = name @@ -1934,27 +1936,33 @@ async def test_max_concurrency() -> None: self.currently += 1 if self.currently > self.max_currently: self.max_currently = self.currently - await asyncio.sleep(0.1) + await asyncio.sleep(random.random() / 10) self.currently -= 1 - return [self.name] + return [state] + + def one(state): + return ["1"] + + def three(state): + return ["3"] async def send_to_many(state): - return [Send("2", state)] * 100 + return [Send("2", idx) for idx in range(100)] async def route_to_three(state) -> Literal["3"]: return "3" node2 = Node("2") builder = StateGraph(Annotated[list, operator.add]) - builder.add_node(Node("1")) + builder.add_node("1", one) builder.add_node(node2) - builder.add_node(Node("3")) + builder.add_node("3", three) builder.add_edge(START, "1") builder.add_conditional_edges("1", send_to_many) builder.add_conditional_edges("2", route_to_three) graph = builder.compile() - assert await graph.ainvoke(["0"]) == ["0", "1", *(["2"] * 100), "3"] + assert await graph.ainvoke(["0"]) == ["0", "1", *range(100), "3"] assert node2.max_currently == 100 assert node2.currently == 0 node2.max_currently = 0 @@ -1962,16 +1970,24 @@ async def test_max_concurrency() -> None: assert await graph.ainvoke(["0"], {"max_concurrency": 10}) == [ "0", "1", - *(["2"] * 100), + *range(100), "3", ] assert node2.max_currently == 10 assert node2.currently == 0 + async with awith_checkpointer(checkpointer_name) as checkpointer: + graph = builder.compile(checkpointer=checkpointer, interrupt_before=["2"]) + thread1 = {"max_concurrency": 10, "configurable": {"thread_id": "1"}} -async def test_max_concurrency_control() -> None: + assert await graph.ainvoke(["0"], thread1) == ["0", "1"] + assert await graph.ainvoke(None, thread1) == ["0", "1", *range(100), "3"] + + +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) +async def test_max_concurrency_control(checkpointer_name: str) -> None: async def node1(state) -> Control[Literal["2"]]: - return Control(update_state=["1"], send=[Send("2", state)] * 100) + return Control(update_state=["1"], send=[Send("2", idx) for idx in range(100)]) node2_currently = 0 node2_max_currently = 0 @@ -1984,7 +2000,7 @@ async def test_max_concurrency_control() -> None: await asyncio.sleep(0.1) node2_currently -= 1 - return Control(update_state=["2"], trigger="3") + return Control(update_state=[state], trigger="3") async def node3(state) -> Literal["3"]: return ["3"] @@ -2013,7 +2029,7 @@ graph TD; """ ) - assert await graph.ainvoke(["0"], debug=True) == ["0", "1", *(["2"] * 100), "3"] + assert await graph.ainvoke(["0"], debug=True) == ["0", "1", *range(100), "3"] assert node2_max_currently == 100 assert node2_currently == 0 node2_max_currently = 0 @@ -2021,12 +2037,19 @@ graph TD; assert await graph.ainvoke(["0"], {"max_concurrency": 10}) == [ "0", "1", - *(["2"] * 100), + *range(100), "3", ] assert node2_max_currently == 10 assert node2_currently == 0 + async with awith_checkpointer(checkpointer_name) as checkpointer: + graph = builder.compile(checkpointer=checkpointer, interrupt_before=["2"]) + thread1 = {"max_concurrency": 10, "configurable": {"thread_id": "1"}} + + assert await graph.ainvoke(["0"], thread1) == ["0", "1"] + assert await graph.ainvoke(None, thread1) == ["0", "1", *range(100), "3"] + @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) async def test_invoke_checkpoint_three(