From 8205c85383feb122b86ca745885237e60e0b7b58 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 24 Sep 2024 09:38:24 -0700 Subject: [PATCH] Add namesapce filter to get_subgraphs - If passed it returns a single subgraph, mathching that namespace, if found --- libs/langgraph/langgraph/pregel/__init__.py | 138 +++++++++++--------- 1 file changed, 75 insertions(+), 63 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index b642f8007..cd3c8150d 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -382,13 +382,23 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): k for k in self.channels if isinstance(self.channels[k], BaseChannel) ] - def get_subgraphs(self, recurse: bool = False) -> Iterator[tuple[str, Pregel]]: + def get_subgraphs( + self, *, namespace: Optional[str] = None, recurse: bool = False + ) -> Iterator[tuple[str, Pregel]]: for name, node in self.nodes.items(): + # filter by prefix + if namespace is not None: + if not namespace.startswith(name): + continue # find the subgraph, if any graph: Optional[Pregel] = None candidates = [node.bound] for candidate in candidates: - if isinstance(candidate, Pregel): + if ( + isinstance(candidate, Pregel) + # subgraphs that disabled checkpointing are not considered + and candidate.checkpointer is not False + ): graph = candidate break elif isinstance(candidate, RunnableSequence): @@ -408,17 +418,25 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): ) # if found, yield recursively if graph: - yield name, graph + if name == namespace: + yield name, graph + return # we found it, stop searching + if namespace is None: + yield name, graph if recurse: + if namespace is not None: + namespace = namespace[len(name) + 1 :] yield from ( (f"{name}{NS_SEP}{n}", s) - for n, s in graph.get_subgraphs(recurse=recurse) + for n, s in graph.get_subgraphs( + namespace=namespace, recurse=recurse + ) ) async def aget_subgraphs( - self, recurse: bool = False + self, *, namespace: Optional[str] = None, recurse: bool = False ) -> AsyncIterator[tuple[str, Pregel]]: - for name, node in self.get_subgraphs(recurse=recurse): + for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse): yield name, node def _prepare_state_snapshot( @@ -588,14 +606,13 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) ) # find the subgraph with the matching name - for name, pregel in self.get_subgraphs(recurse=True): - if name == recast_checkpoint_ns: - return pregel.get_state( - patch_configurable( - config, {CONFIG_KEY_CHECKPOINTER: checkpointer} - ), - subgraphs=subgraphs, - ) + for _, pregel in self.get_subgraphs( + namespace=recast_checkpoint_ns, recurse=True + ): + return pregel.get_state( + patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), + subgraphs=subgraphs, + ) else: raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") @@ -623,14 +640,13 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) ) # find the subgraph with the matching name - async for name, pregel in self.aget_subgraphs(recurse=True): - if name == recast_checkpoint_ns: - return await pregel.aget_state( - patch_configurable( - config, {CONFIG_KEY_CHECKPOINTER: checkpointer} - ), - subgraphs=subgraphs, - ) + async for _, pregel in self.aget_subgraphs( + namespace=recast_checkpoint_ns, recurse=True + ): + return await pregel.aget_state( + patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), + subgraphs=subgraphs, + ) else: raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") @@ -663,17 +679,16 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) ) # find the subgraph with the matching name - for name, pregel in self.get_subgraphs(recurse=True): - if name == recast_checkpoint_ns: - yield from pregel.get_state_history( - patch_configurable( - config, {CONFIG_KEY_CHECKPOINTER: checkpointer} - ), - filter=filter, - before=before, - limit=limit, - ) - return + for _, pregel in self.get_subgraphs( + namespace=recast_checkpoint_ns, recurse=True + ): + yield from pregel.get_state_history( + patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), + filter=filter, + before=before, + limit=limit, + ) + return else: raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") @@ -713,18 +728,17 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) ) # find the subgraph with the matching name - async for name, pregel in self.aget_subgraphs(recurse=True): - if name == recast_checkpoint_ns: - async for state in pregel.aget_state_history( - patch_configurable( - config, {CONFIG_KEY_CHECKPOINTER: checkpointer} - ), - filter=filter, - before=before, - limit=limit, - ): - yield state - return + async for _, pregel in self.aget_subgraphs( + namespace=recast_checkpoint_ns, recurse=True + ): + async for state in pregel.aget_state_history( + patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), + filter=filter, + before=before, + limit=limit, + ): + yield state + return else: raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") @@ -769,15 +783,14 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) ) # find the subgraph with the matching name - for name, pregel in self.get_subgraphs(recurse=True): - if name == recast_checkpoint_ns: - return pregel.update_state( - patch_configurable( - config, {CONFIG_KEY_CHECKPOINTER: checkpointer} - ), - values, - as_node, - ) + for _, pregel in self.get_subgraphs( + namespace=recast_checkpoint_ns, recurse=True + ): + return pregel.update_state( + patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), + values, + as_node, + ) else: raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") @@ -917,15 +930,14 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) ) # find the subgraph with the matching name - async for name, pregel in self.aget_subgraphs(recurse=True): - if name == recast_checkpoint_ns: - return await pregel.aupdate_state( - patch_configurable( - config, {CONFIG_KEY_CHECKPOINTER: checkpointer} - ), - values, - as_node, - ) + async for _, pregel in self.aget_subgraphs( + namespace=recast_checkpoint_ns, recurse=True + ): + return await pregel.aupdate_state( + patch_configurable(config, {CONFIG_KEY_CHECKPOINTER: checkpointer}), + values, + as_node, + ) else: raise ValueError(f"Subgraph {recast_checkpoint_ns} not found")