From a277b86fcb32b1d80fd5fb1f115a0d37854b61fc Mon Sep 17 00:00:00 2001 From: Andrew Nguonly Date: Fri, 11 Oct 2024 18:53:45 -0700 Subject: [PATCH] Fix unit test. --- libs/langgraph/langgraph/pregel/remote.py | 26 ++++++++++++----------- libs/langgraph/tests/test_remote_graph.py | 4 ++-- 2 files changed, 16 insertions(+), 14 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index 4b7e87421..b0df12bb7 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -400,13 +400,14 @@ class RemoteGraph(PregelProtocol, Runnable): interrupt_after=interrupt_after, # type: ignore stream_subgraphs=subgraphs, ): - if chunk.event == INTERRUPT: - raise GraphInterrupt() + if chunk.event == "updates": + if INTERRUPT in chunk.data: + raise GraphInterrupt() - # Don't emit 'updates' events if the original list of stream modes - # didn't include it. - if chunk.event == "updates" and not include_updates: - continue + # Don't emit 'updates' events if the original list of stream + # modes didn't include it. + if not include_updates: + continue yield chunk @@ -434,13 +435,14 @@ class RemoteGraph(PregelProtocol, Runnable): interrupt_after=interrupt_after, # type: ignore stream_subgraphs=subgraphs, ): - if chunk.event == INTERRUPT: - raise GraphInterrupt() + if chunk.event == "updates": + if INTERRUPT in chunk.data: + raise GraphInterrupt() - # Don't emit 'updates' events if the original list of stream modes - # didn't include it. - if chunk.event == "updates" and not include_updates: - continue + # Don't emit 'updates' events if the original list of stream + # modes didn't include it. + if not include_updates: + continue yield chunk diff --git a/libs/langgraph/tests/test_remote_graph.py b/libs/langgraph/tests/test_remote_graph.py index dd07d27bb..a6a208b27 100644 --- a/libs/langgraph/tests/test_remote_graph.py +++ b/libs/langgraph/tests/test_remote_graph.py @@ -479,7 +479,7 @@ def test_stream(): StreamPart(event="values", data={"chunk": "data2"}), StreamPart(event="values", data={"chunk": "data3"}), StreamPart(event="updates", data={"chunk": "data4"}), - StreamPart(event="__interrupt__", data={}), + StreamPart(event="updates", data={"__interrupt__": ()}), ] # call method / assertions @@ -527,7 +527,7 @@ async def test_astream(): StreamPart(event="values", data={"chunk": "data2"}), StreamPart(event="values", data={"chunk": "data3"}), StreamPart(event="updates", data={"chunk": "data4"}), - StreamPart(event="__interrupt__", data={}), + StreamPart(event="updates", data={"__interrupt__": ()}), ] mock_async_client.runs.stream.return_value = async_iter