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
This commit is contained in:
Nuno Campos
2024-11-01 13:23:10 -07:00
parent 18e71469e1
commit 1e3953d1e0
3 changed files with 40 additions and 17 deletions
@@ -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
+3 -3
View File
@@ -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"}})),
]
+36 -13
View File
@@ -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(