From dca200d6c4689f6fb8eb8b4b9a04e110dbb6a9a3 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 23 Oct 2024 13:23:42 -0700 Subject: [PATCH] Finish --- libs/langgraph/langgraph/pregel/remote.py | 97 ++++++++++------------- 1 file changed, 44 insertions(+), 53 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index 509f048d3..59150b7b1 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -19,7 +19,6 @@ from langchain_core.runnables.graph import ( from langchain_core.runnables.graph import ( Node as DrawableNode, ) -from langchain_core.runnables.schema import StandardStreamEvent, StreamEvent from langgraph_sdk.client import ( LangGraphClient, SyncLangGraphClient, @@ -39,6 +38,12 @@ from langgraph.types import Interrupt from langgraph.utils.config import merge_configs +class RemoteException(Exception): + """Exception raised when an error occurs in the remote graph.""" + + pass + + class RemoteGraph(PregelProtocol, Runnable): def __init__( self, @@ -326,7 +331,7 @@ class RemoteGraph(PregelProtocol, Runnable): response: dict = self.sync_client.threads.update_state( # type: ignore thread_id=merged_config["configurable"]["thread_id"], - values=values, # type: ignore + values=values, as_node=as_node, checkpoint=self._get_checkpoint(merged_config), ) @@ -342,7 +347,7 @@ class RemoteGraph(PregelProtocol, Runnable): response: dict = await self.client.threads.update_state( # type: ignore thread_id=merged_config["configurable"]["thread_id"], - values=values, # type: ignore + values=values, as_node=as_node, checkpoint=self._get_checkpoint(merged_config), ) @@ -392,7 +397,6 @@ class RemoteGraph(PregelProtocol, Runnable): merged_config = merge_configs(self.config, config) sanitized_config = self._sanitize_config(merged_config) stream_modes, req_updates, req_single = self._get_stream_modes(stream_mode) - # TODO if req_subgraphs transform chunk to match Pregel for chunk in self.sync_client.runs.stream( thread_id=sanitized_config["configurable"]["thread_id"], @@ -400,16 +404,28 @@ class RemoteGraph(PregelProtocol, Runnable): input=input, config=sanitized_config, stream_mode=stream_modes, - interrupt_before=interrupt_before, # type: ignore - interrupt_after=interrupt_after, # type: ignore + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, stream_subgraphs=subgraphs, ): - if chunk.event == "updates": + if chunk.event.startswith("updates"): if isinstance(chunk.data, dict) and INTERRUPT in chunk.data: raise GraphInterrupt(chunk.data[INTERRUPT]) if not req_updates: continue - if req_single: + elif chunk.event.startswith("error"): + raise RemoteException(chunk.data) + if subgraphs: + if "|" in chunk.event: + mode, ns_ = chunk.event.split("|", 1) + ns = tuple(ns_.split("|")) + else: + mode, ns = chunk.event, () + if req_single: + yield ns, chunk.data + else: + yield ns, mode, chunk.data + elif req_single: yield chunk.data else: yield chunk @@ -434,57 +450,32 @@ class RemoteGraph(PregelProtocol, Runnable): input=input, config=sanitized_config, stream_mode=stream_modes, - interrupt_before=interrupt_before, # type: ignore - interrupt_after=interrupt_after, # type: ignore + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, stream_subgraphs=subgraphs, ): - if chunk.event == "updates": + if chunk.event.startswith("updates"): if isinstance(chunk.data, dict) and INTERRUPT in chunk.data: raise GraphInterrupt(chunk.data[INTERRUPT]) if not req_updates: continue - if req_single: + elif chunk.event.startswith("error"): + raise RemoteException(chunk.data) + if subgraphs: + if "|" in chunk.event: + mode, ns_ = chunk.event.split("|", 1) + ns = tuple(ns_.split("|")) + else: + mode, ns = chunk.event, () + if req_single: + yield ns, chunk.data + else: + yield ns, mode, chunk.data + elif req_single: yield chunk.data else: yield chunk - async def astream_events( - self, - input: Any, - config: Optional[RunnableConfig] = None, - **kwargs: Any, - ) -> AsyncIterator[StreamEvent]: - merged_config = merge_configs(self.config, config) - sanitized_config = self._sanitize_config(merged_config) - - # manually add 'events' to stream modes list - stream_mode: Union[StreamMode, list[StreamMode]] = kwargs.get("stream_mode", []) - updated_stream_modes, include_updates, _ = self._get_stream_modes(stream_mode) - if "events" not in updated_stream_modes: - updated_stream_modes.append("events") - # TODO bundle main stream events back into StreamEvent - - async for chunk in self.client.runs.stream( - thread_id=sanitized_config["configurable"]["thread_id"], - assistant_id=self.graph_id, - input=input, - config=sanitized_config, - stream_mode=updated_stream_modes, - interrupt_before=kwargs.get("interrupt_before"), - interrupt_after=kwargs.get("interrupt_after"), - stream_subgraphs=kwargs.get("subgraphs", False), - ): - if chunk.event == "updates": - if INTERRUPT in chunk.data: - raise GraphInterrupt() - if not include_updates: - continue - - yield StandardStreamEvent( - event=chunk.event, - data=chunk.data, - ) - def invoke( self, input: Union[dict[str, Any], Any], @@ -501,8 +492,8 @@ class RemoteGraph(PregelProtocol, Runnable): assistant_id=self.graph_id, input=input, config=sanitized_config, - interrupt_before=interrupt_before, # type: ignore - interrupt_after=interrupt_after, # type: ignore + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, ) async def ainvoke( @@ -521,6 +512,6 @@ class RemoteGraph(PregelProtocol, Runnable): assistant_id=self.graph_id, input=input, config=sanitized_config, - interrupt_before=interrupt_before, # type: ignore - interrupt_after=interrupt_after, # type: ignore + interrupt_before=interrupt_before, + interrupt_after=interrupt_after, )