diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index f58d0a4d9..07d44cb66 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -654,9 +654,10 @@ class RemoteGraph(PregelProtocol): # raise interrupt or errors if chunk.event.startswith("updates"): if isinstance(chunk.data, dict) and INTERRUPT in chunk.data: - raise GraphInterrupt( - [Interrupt(**i) for i in chunk.data[INTERRUPT]] - ) + if caller_ns: + raise GraphInterrupt( + [Interrupt(**i) for i in chunk.data[INTERRUPT]] + ) elif chunk.event.startswith("error"): raise RemoteException(chunk.data) # filter for what was actually requested @@ -748,9 +749,10 @@ class RemoteGraph(PregelProtocol): # raise interrupt or errors if chunk.event.startswith("updates"): if isinstance(chunk.data, dict) and INTERRUPT in chunk.data: - raise GraphInterrupt( - [Interrupt(**i) for i in chunk.data[INTERRUPT]] - ) + if caller_ns: + raise GraphInterrupt( + [Interrupt(**i) for i in chunk.data[INTERRUPT]] + ) elif chunk.event.startswith("error"): raise RemoteException(chunk.data) # filter for what was actually requested diff --git a/libs/langgraph/tests/test_remote_graph.py b/libs/langgraph/tests/test_remote_graph.py index 4aea1fc32..878f5e6f1 100644 --- a/libs/langgraph/tests/test_remote_graph.py +++ b/libs/langgraph/tests/test_remote_graph.py @@ -437,15 +437,17 @@ def test_stream(): sync_client=mock_sync_client, ) - # stream modes doesn't include 'updates' - stream_parts = [] + # test raising graph interrupt if invoked as a subgraph with pytest.raises(GraphInterrupt) as exc: for stream_part in remote_pregel.stream( {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, + # pretend we invoked this as a subgraph + config={ + "configurable": {"thread_id": "thread_1", "checkpoint_ns": "some_ns"} + }, stream_mode="values", ): - stream_parts.append(stream_part) + pass assert exc.value.args[0] == [ Interrupt( @@ -456,6 +458,15 @@ def test_stream(): ) ] + # stream modes doesn't include 'updates' + stream_parts = [] + for stream_part in remote_pregel.stream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode="values", + ): + stream_parts.append(stream_part) + assert stream_parts == [ {"chunk": "data1"}, {"chunk": "data2"}, @@ -470,62 +481,62 @@ def test_stream(): # default stream_mode is updates stream_parts = [] - with pytest.raises(GraphInterrupt): - for stream_part in remote_pregel.stream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - ): - stream_parts.append(stream_part) + for stream_part in remote_pregel.stream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + ): + stream_parts.append(stream_part) assert stream_parts == [ {"chunk": "data3"}, {"chunk": "data4"}, + {"__interrupt__": ()}, ] # list stream_mode includes mode names stream_parts = [] - with pytest.raises(GraphInterrupt): - for stream_part in remote_pregel.stream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - stream_mode=["updates"], - ): - stream_parts.append(stream_part) + for stream_part in remote_pregel.stream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode=["updates"], + ): + stream_parts.append(stream_part) assert stream_parts == [ ("updates", {"chunk": "data3"}), ("updates", {"chunk": "data4"}), + ("updates", {"__interrupt__": ()}), ] # subgraphs + list modes stream_parts = [] - with pytest.raises(GraphInterrupt): - for stream_part in remote_pregel.stream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - stream_mode=["updates"], - subgraphs=True, - ): - stream_parts.append(stream_part) + for stream_part in remote_pregel.stream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode=["updates"], + subgraphs=True, + ): + stream_parts.append(stream_part) assert stream_parts == [ ((), "updates", {"chunk": "data3"}), ((), "updates", {"chunk": "data4"}), + ((), "updates", {"__interrupt__": ()}), ] # subgraphs + single mode stream_parts = [] - with pytest.raises(GraphInterrupt): - for stream_part in remote_pregel.stream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - subgraphs=True, - ): - stream_parts.append(stream_part) + for stream_part in remote_pregel.stream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + subgraphs=True, + ): + stream_parts.append(stream_part) assert stream_parts == [ ((), {"chunk": "data3"}), ((), {"chunk": "data4"}), + ((), {"__interrupt__": ()}), ] @@ -561,15 +572,17 @@ async def test_astream(): client=mock_async_client, ) - # stream modes doesn't include 'updates' - stream_parts = [] + # test raising graph interrupt if invoked as a subgraph with pytest.raises(GraphInterrupt) as exc: async for stream_part in remote_pregel.astream( {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, + # pretend we invoked this as a subgraph + config={ + "configurable": {"thread_id": "thread_1", "checkpoint_ns": "some_ns"} + }, stream_mode="values", ): - stream_parts.append(stream_part) + pass assert exc.value.args[0] == [ Interrupt( @@ -580,6 +593,15 @@ async def test_astream(): ) ] + # stream modes doesn't include 'updates' + stream_parts = [] + async for stream_part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode="values", + ): + stream_parts.append(stream_part) + assert stream_parts == [ {"chunk": "data1"}, {"chunk": "data2"}, @@ -596,62 +618,62 @@ async def test_astream(): # default stream_mode is updates stream_parts = [] - with pytest.raises(GraphInterrupt): - async for stream_part in remote_pregel.astream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - ): - stream_parts.append(stream_part) + async for stream_part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + ): + stream_parts.append(stream_part) assert stream_parts == [ {"chunk": "data3"}, {"chunk": "data4"}, + {"__interrupt__": ()}, ] # list stream_mode includes mode names stream_parts = [] - with pytest.raises(GraphInterrupt): - async for stream_part in remote_pregel.astream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - stream_mode=["updates"], - ): - stream_parts.append(stream_part) + async for stream_part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode=["updates"], + ): + stream_parts.append(stream_part) assert stream_parts == [ ("updates", {"chunk": "data3"}), ("updates", {"chunk": "data4"}), + ("updates", {"__interrupt__": ()}), ] # subgraphs + list modes stream_parts = [] - with pytest.raises(GraphInterrupt): - async for stream_part in remote_pregel.astream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - stream_mode=["updates"], - subgraphs=True, - ): - stream_parts.append(stream_part) + async for stream_part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode=["updates"], + subgraphs=True, + ): + stream_parts.append(stream_part) assert stream_parts == [ ((), "updates", {"chunk": "data3"}), ((), "updates", {"chunk": "data4"}), + ((), "updates", {"__interrupt__": ()}), ] # subgraphs + single mode stream_parts = [] - with pytest.raises(GraphInterrupt): - async for stream_part in remote_pregel.astream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - subgraphs=True, - ): - stream_parts.append(stream_part) + async for stream_part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + subgraphs=True, + ): + stream_parts.append(stream_part) assert stream_parts == [ ((), {"chunk": "data3"}), ((), {"chunk": "data4"}), + ((), {"__interrupt__": ()}), ] async_iter = MagicMock() @@ -664,33 +686,33 @@ async def test_astream(): # subgraphs + list modes stream_parts = [] - with pytest.raises(GraphInterrupt): - async for stream_part in remote_pregel.astream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - stream_mode=["updates"], - subgraphs=True, - ): - stream_parts.append(stream_part) + async for stream_part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + stream_mode=["updates"], + subgraphs=True, + ): + stream_parts.append(stream_part) assert stream_parts == [ (("my", "subgraph"), "updates", {"chunk": "data3"}), (("hello", "subgraph"), "updates", {"chunk": "data4"}), + (("bye", "subgraph"), "updates", {"__interrupt__": ()}), ] # subgraphs + single mode stream_parts = [] - with pytest.raises(GraphInterrupt): - async for stream_part in remote_pregel.astream( - {"input": "data"}, - config={"configurable": {"thread_id": "thread_1"}}, - subgraphs=True, - ): - stream_parts.append(stream_part) + async for stream_part in remote_pregel.astream( + {"input": "data"}, + config={"configurable": {"thread_id": "thread_1"}}, + subgraphs=True, + ): + stream_parts.append(stream_part) assert stream_parts == [ (("my", "subgraph"), {"chunk": "data3"}), (("hello", "subgraph"), {"chunk": "data4"}), + (("bye", "subgraph"), {"__interrupt__": ()}), ]