Fix issue when cond edge visited after multiple executions of Send

- cond edge will run for each execution of Send, so target channels need to support multiple publishes
This commit is contained in:
Nuno Campos
2024-07-29 12:38:48 -07:00
parent ea071935fe
commit 466cb8acb5
3 changed files with 17 additions and 17 deletions
+3 -3
View File
@@ -610,7 +610,7 @@ class CompiledStateGraph(CompiledGraph):
if branch.then and branch.then != END:
writes.append(
ChannelWriteEntry(
f"branch:{start}:{name}:then",
f"branch:{start}:{name}::then",
WaitForNames(
{p.node if isinstance(p, Send) else p for p in filtered}
),
@@ -630,12 +630,12 @@ class CompiledStateGraph(CompiledGraph):
for end in ends:
if end != END:
channel_name = f"branch:{start}:{name}:{end}"
self.channels[channel_name] = EphemeralValue(Any)
self.channels[channel_name] = EphemeralValue(Any, guard=False)
self.nodes[end].triggers.append(channel_name)
# attach then subscriber
if branch.then and branch.then != END:
channel_name = f"branch:{start}:{name}:then"
channel_name = f"branch:{start}:{name}::then"
self.channels[channel_name] = DynamicBarrierValue(str)
self.nodes[branch.then].triggers.append(channel_name)
for end in ends:
+7 -7
View File
@@ -1134,15 +1134,15 @@ def test_cond_edge_after_send() -> None:
setattr(self, "__name__", name)
def __call__(self, state):
return state + [self.name]
return [self.name]
def send_for_fun(state):
return [Send("2", state)]
return [Send("2", state), Send("2", state)]
def route_to_three(state) -> Literal["3"]:
return "3"
builder = StateGraph(list)
builder = StateGraph(Annotated[list, operator.add])
builder.add_node(Node("1"))
builder.add_node(Node("2"))
builder.add_node(Node("3"))
@@ -1150,7 +1150,7 @@ def test_cond_edge_after_send() -> None:
builder.add_conditional_edges("1", send_for_fun)
builder.add_conditional_edges("2", route_to_three)
graph = builder.compile()
assert graph.invoke(["0"]) == ["0", "1", "2", "3"]
assert graph.invoke(["0"]) == ["0", "1", "2", "2", "3"]
async def test_checkpointer_null_pending_writes() -> None:
@@ -6536,10 +6536,10 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
"timestamp": AnyStr(),
"step": 3,
"payload": {
"id": "ceada3c5-5f25-59e4-9ea5-544599ce1d2f",
"id": "9b590c54-15ef-54b1-83a7-140d27b0bc52",
"name": "finish",
"input": {"my_key": "value prepared slow", "market": "DE"},
"triggers": ["branch:prepare:condition:then"],
"triggers": ["branch:prepare:condition::then"],
},
},
{
@@ -6547,7 +6547,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
"timestamp": AnyStr(),
"step": 3,
"payload": {
"id": "ceada3c5-5f25-59e4-9ea5-544599ce1d2f",
"id": "9b590c54-15ef-54b1-83a7-140d27b0bc52",
"name": "finish",
"result": [("my_key", " finished")],
},
+7 -7
View File
@@ -1257,15 +1257,15 @@ async def test_cond_edge_after_send() -> None:
setattr(self, "__name__", name)
async def __call__(self, state):
return state + [self.name]
return [self.name]
async def send_for_fun(state):
return [Send("2", state)]
return [Send("2", state), Send("2", state)]
async def route_to_three(state) -> Literal["3"]:
return "3"
builder = StateGraph(list)
builder = StateGraph(Annotated[list, operator.add])
builder.add_node(Node("1"))
builder.add_node(Node("2"))
builder.add_node(Node("3"))
@@ -1274,7 +1274,7 @@ async def test_cond_edge_after_send() -> None:
builder.add_conditional_edges("2", route_to_three)
graph = builder.compile()
assert await graph.ainvoke(["0"]) == ["0", "1", "2", "3"]
assert await graph.ainvoke(["0"]) == ["0", "1", "2", "2", "3"]
async def test_invoke_checkpoint_aiosqlite(mocker: MockerFixture) -> None:
@@ -5119,10 +5119,10 @@ async def test_branch_then() -> None:
"timestamp": AnyStr(),
"step": 3,
"payload": {
"id": "ceada3c5-5f25-59e4-9ea5-544599ce1d2f",
"id": "9b590c54-15ef-54b1-83a7-140d27b0bc52",
"name": "finish",
"input": {"my_key": "value prepared slow", "market": "DE"},
"triggers": ["branch:prepare:condition:then"],
"triggers": ["branch:prepare:condition::then"],
},
},
{
@@ -5130,7 +5130,7 @@ async def test_branch_then() -> None:
"timestamp": AnyStr(),
"step": 3,
"payload": {
"id": "ceada3c5-5f25-59e4-9ea5-544599ce1d2f",
"id": "9b590c54-15ef-54b1-83a7-140d27b0bc52",
"name": "finish",
"result": [("my_key", " finished")],
},