From 844373e9b3539a932af13be0a0b921c7a4e8c556 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Fri, 19 Dec 2025 05:58:51 +0900 Subject: [PATCH] chore: remotegraph - use wait enpoint for invoke/ainvoke --- libs/langgraph/langgraph/pregel/remote.py | 102 +++++++++++++--------- libs/langgraph/tests/test_remote_graph.py | 59 ++++++++----- 2 files changed, 97 insertions(+), 64 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index 2535d966a..fb9865cf4 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -631,6 +631,26 @@ class RemoteGraph(PregelProtocol): ) return self._get_config(response["checkpoint"]) + def _prepare_run_input( + self, + input: dict[str, Any] | Any, + config: RunnableConfig | None, + ) -> tuple[RunnableConfig, dict[str, Any] | Any, CommandSDK | None, str | None]: + """Prepare input for run calls. + + Returns: + Tuple of (sanitized_config, input, command, thread_id) + """ + merged_config = merge_configs(self.config, config) + sanitized_config = self._sanitize_config(merged_config) + if isinstance(input, Command): + command: CommandSDK | None = cast(CommandSDK, asdict(input)) + input = None + else: + command = None + thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None) + return sanitized_config, input, command, thread_id + def _get_stream_modes( self, stream_mode: StreamMode | list[StreamMode] | None, @@ -715,17 +735,12 @@ class RemoteGraph(PregelProtocol): 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) + sanitized_config, input, command, thread_id = self._prepare_run_input( + input, config + ) stream_modes, requested, req_single, stream = self._get_stream_modes( stream_mode, config ) - if isinstance(input, Command): - command: CommandSDK | None = cast(CommandSDK, asdict(input)) - input = None - else: - command = None - thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None) for chunk in sync_client.runs.stream( thread_id=thread_id, @@ -825,17 +840,12 @@ class RemoteGraph(PregelProtocol): The output of the graph. """ client = self._validate_client() - merged_config = merge_configs(self.config, config) - sanitized_config = self._sanitize_config(merged_config) + sanitized_config, input, command, thread_id = self._prepare_run_input( + input, config + ) stream_modes, requested, req_single, stream = self._get_stream_modes( stream_mode, config ) - if isinstance(input, Command): - command: CommandSDK | None = cast(CommandSDK, asdict(input)) - input = None - else: - command = None - thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None) async for chunk in client.runs.stream( thread_id=thread_id, @@ -937,27 +947,31 @@ class RemoteGraph(PregelProtocol): interrupt_before: Interrupt the graph before these nodes. interrupt_after: Interrupt the graph after these nodes. headers: Additional headers to pass to the request. - **kwargs: Additional params to pass to RemoteGraph.stream. + **kwargs: Additional params to pass to client.runs.wait. Returns: The output of the graph. """ - for chunk in self.stream( - input, - config=config, + sync_client = self._validate_sync_client() + sanitized_config, input, command, thread_id = self._prepare_run_input( + input, config + ) + + return sync_client.runs.wait( # type: ignore + thread_id=thread_id, + assistant_id=self.assistant_id, + input=input, + command=command, + config=sanitized_config, interrupt_before=interrupt_before, interrupt_after=interrupt_after, - headers=headers, - stream_mode="values", + if_not_exists="create", + headers=( + _merge_tracing_headers(headers) if self.distributed_tracing else headers + ), params=params, **kwargs, - ): - pass - try: - return chunk - except UnboundLocalError: - logger.warning("No events received from remote graph") - return None + ) async def ainvoke( self, @@ -978,27 +992,31 @@ class RemoteGraph(PregelProtocol): interrupt_before: Interrupt the graph before these nodes. interrupt_after: Interrupt the graph after these nodes. headers: Additional headers to pass to the request. - **kwargs: Additional params to pass to RemoteGraph.astream. + **kwargs: Additional params to pass to client.runs.wait. Returns: The output of the graph. """ - async for chunk in self.astream( - input, - config=config, + client = self._validate_client() + sanitized_config, input, command, thread_id = self._prepare_run_input( + input, config + ) + + return await client.runs.wait( + thread_id=thread_id, + assistant_id=self.assistant_id, + input=input, + command=command, + config=sanitized_config, interrupt_before=interrupt_before, interrupt_after=interrupt_after, - headers=headers, - stream_mode="values", + if_not_exists="create", + headers=( + _merge_tracing_headers(headers) if self.distributed_tracing else headers + ), params=params, **kwargs, - ): - pass - try: - return chunk - except UnboundLocalError: - logger.warning("No events received from remote graph") - return None + ) def _merge_tracing_headers(headers: dict[str, str] | None) -> dict[str, str] | None: diff --git a/libs/langgraph/tests/test_remote_graph.py b/libs/langgraph/tests/test_remote_graph.py index e0eacf006..8da0e9856 100644 --- a/libs/langgraph/tests/test_remote_graph.py +++ b/libs/langgraph/tests/test_remote_graph.py @@ -818,13 +818,9 @@ async def test_astream(): def test_invoke(): # set up test mock_sync_client = MagicMock() - 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"}]} - ), - ] + mock_sync_client.runs.wait.return_value = { + "messages": [{"type": "human", "content": "world"}] + } # call method / assertions remote_pregel = RemoteGraph( @@ -838,13 +834,19 @@ def test_invoke(): ) assert result == {"messages": [{"type": "human", "content": "world"}]} + # verify runs.wait was called with expected args + assert mock_sync_client.runs.wait.called + _, kwargs = mock_sync_client.runs.wait.call_args + assert kwargs.get("thread_id") == "thread_1" + assert kwargs.get("assistant_id") == "test_graph_id" + assert kwargs.get("if_not_exists") == "create" def test_invoke_sanitizes_thread_id(): # Ensure that invoking with thread_id passes thread_id as a top-level arg # and removes it from the config body. mock_sync_client = MagicMock() - mock_sync_client.runs.stream.return_value = [] + mock_sync_client.runs.wait.return_value = {} remote_pregel = RemoteGraph("test_graph_id", sync_client=mock_sync_client) config = {"configurable": {"thread_id": "thread_1"}} @@ -852,8 +854,8 @@ def test_invoke_sanitizes_thread_id(): {"input": {"messages": [{"type": "human", "content": "hello"}]}}, config ) - assert mock_sync_client.runs.stream.called - _, kwargs = mock_sync_client.runs.stream.call_args + assert mock_sync_client.runs.wait.called + _, kwargs = mock_sync_client.runs.wait.call_args assert kwargs.get("thread_id") == "thread_1" passed_config = kwargs.get("config") or {} assert "configurable" in passed_config @@ -883,16 +885,10 @@ def test_stream_sanitizes_thread_id(): @pytest.mark.anyio async def test_ainvoke(): # set up test - 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 + mock_async_client = AsyncMock() + mock_async_client.runs.wait.return_value = { + "messages": [{"type": "human", "content": "world"}] + } # call method / assertions remote_pregel = RemoteGraph( @@ -906,6 +902,12 @@ async def test_ainvoke(): ) assert result == {"messages": [{"type": "human", "content": "world"}]} + # verify runs.wait was called with expected args + assert mock_async_client.runs.wait.called + _, kwargs = mock_async_client.runs.wait.call_args + assert kwargs.get("thread_id") == "thread_1" + assert kwargs.get("assistant_id") == "test_graph_id" + assert kwargs.get("if_not_exists") == "create" @pytest.mark.skip( @@ -1241,12 +1243,18 @@ async def test_include_headers( async_iter.__aiter__.return_value = return_value astream_mock = mock_async_client.runs.stream astream_mock.return_value = async_iter + # Mock for ainvoke which uses runs.wait + await_mock = AsyncMock(return_value={"chunk": "data1"}) + mock_async_client.runs.wait = await_mock mock_sync_client = MagicMock() sync_iter = MagicMock() sync_iter.__iter__.return_value = return_value stream_mock = mock_sync_client.runs.stream stream_mock.return_value = async_iter + # Mock for invoke which uses runs.wait + wait_mock = MagicMock(return_value={"chunk": "data1"}) + mock_sync_client.runs.wait = wait_mock remote_pregel = RemoteGraph( "test_graph_id", @@ -1279,8 +1287,12 @@ async def test_include_headers( expected["langsmith-trace"] = AnyStr() expected["baggage"] = AnyStr("langsmith-metadata=") - assert astream_mock.call_args.kwargs["headers"] == expected + if stream: + assert astream_mock.call_args.kwargs["headers"] == expected + else: + assert await_mock.call_args.kwargs["headers"] == expected stream_mock.assert_not_called() + wait_mock.assert_not_called() with ls.tracing_context(enabled=True, client=MagicMock()): with ls.trace("foo"): @@ -1298,4 +1310,7 @@ async def test_include_headers( config, headers=headers, ) - assert stream_mock.call_args.kwargs["headers"] == expected + if stream: + assert stream_mock.call_args.kwargs["headers"] == expected + else: + assert wait_mock.call_args.kwargs["headers"] == expected