diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index d055789f3..e47f5f50b 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -695,19 +695,18 @@ class RemoteGraph(PregelProtocol): Returns: The output of the graph. """ - sync_client = self._validate_sync_client() - merged_config = merge_configs(self.config, config) - sanitized_config = self._sanitize_config(merged_config) - - return sync_client.runs.wait( - thread_id=sanitized_config["configurable"].get("thread_id"), - assistant_id=self.name, - input=input, - config=sanitized_config, + for chunk in self.stream( + input, + config=config, interrupt_before=interrupt_before, interrupt_after=interrupt_after, - if_not_exists="create", - ) + stream_mode="values", + ): + pass + try: + return chunk + except UnboundLocalError: + return None async def ainvoke( self, @@ -732,16 +731,15 @@ class RemoteGraph(PregelProtocol): Returns: The output of the graph. """ - client = self._validate_client() - merged_config = merge_configs(self.config, config) - sanitized_config = self._sanitize_config(merged_config) - - return await client.runs.wait( - thread_id=sanitized_config["configurable"].get("thread_id"), - assistant_id=self.name, - input=input, - config=sanitized_config, + async for chunk in self.astream( + input, + config=config, interrupt_before=interrupt_before, interrupt_after=interrupt_after, - if_not_exists="create", - ) + stream_mode="values", + ): + pass + try: + return chunk + except UnboundLocalError: + return None diff --git a/libs/langgraph/tests/test_remote_graph.py b/libs/langgraph/tests/test_remote_graph.py index 8ee0a3f9a..83e5f913b 100644 --- a/libs/langgraph/tests/test_remote_graph.py +++ b/libs/langgraph/tests/test_remote_graph.py @@ -650,9 +650,13 @@ async def test_astream(): def test_invoke(): # set up test mock_sync_client = MagicMock() - mock_sync_client.runs.wait.return_value = { - "values": {"messages": [{"type": "human", "content": "world"}]} - } + mock_sync_client.runs.stream.return_value = [ + StreamPart(event="values", data={"chunk": "data1"}), + StreamPart(event="values", data={"chunk": "data2"}), + StreamPart( + event="values", data={"messages": [{"type": "human", "content": "world"}]} + ), + ] # call method / assertions remote_pregel = RemoteGraph( @@ -665,16 +669,22 @@ def test_invoke(): {"input": {"messages": [{"type": "human", "content": "hello"}]}}, config ) - assert result == {"values": {"messages": [{"type": "human", "content": "world"}]}} + assert result == {"messages": [{"type": "human", "content": "world"}]} @pytest.mark.anyio async def test_ainvoke(): # set up test - mock_async_client = AsyncMock() - mock_async_client.runs.wait.return_value = { - "values": {"messages": [{"type": "human", "content": "world"}]} - } + mock_async_client = MagicMock() + async_iter = MagicMock() + async_iter.__aiter__.return_value = [ + StreamPart(event="values", data={"chunk": "data1"}), + StreamPart(event="values", data={"chunk": "data2"}), + StreamPart( + event="values", data={"messages": [{"type": "human", "content": "world"}]} + ), + ] + mock_async_client.runs.stream.return_value = async_iter # call method / assertions remote_pregel = RemoteGraph( @@ -687,7 +697,7 @@ async def test_ainvoke(): {"input": {"messages": [{"type": "human", "content": "hello"}]}}, config ) - assert result == {"values": {"messages": [{"type": "human", "content": "world"}]}} + assert result == {"messages": [{"type": "human", "content": "world"}]} @pytest.mark.skip("Unskip this test to manually test the LangGraph Cloud integration")