From 32bf81c559e4bf916d5bd0321b9f896e8c554c5d Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 16 Jan 2025 14:10:16 -0800 Subject: [PATCH] Lint --- libs/langgraph/langgraph/func/__init__.py | 96 +++++++++++---------- libs/langgraph/langgraph/pregel/__init__.py | 10 --- 2 files changed, 51 insertions(+), 55 deletions(-) diff --git a/libs/langgraph/langgraph/func/__init__.py b/libs/langgraph/langgraph/func/__init__.py index 73d308365..0e7e53871 100644 --- a/libs/langgraph/langgraph/func/__init__.py +++ b/libs/langgraph/langgraph/func/__init__.py @@ -16,6 +16,7 @@ from typing import ( ) from langchain_core.runnables.base import Runnable +from langchain_core.runnables.config import RunnableConfig from langchain_core.runnables.graph import Graph, Node from typing_extensions import ParamSpec @@ -157,7 +158,7 @@ def entrypoint( else Any ) - return Pregel( + return EntrypointPregel( nodes={ func.__name__: PregelNode( bound=bound, @@ -178,14 +179,60 @@ def entrypoint( checkpointer=checkpointer, store=store, config_type=config_schema, - graph=_entrypoint_graph(bound), ) return _imp +class EntrypointPregel(Pregel): + def get_graph( + self, + config: Optional[RunnableConfig] = None, + *, + xray: int | bool = False, + ) -> Graph: + name, entrypoint = next(iter(self.nodes.items())) + graph = Graph() + node = Node(f"__{name}", name, entrypoint, None) + graph.nodes[node.id] = node + candidates: list[tuple[Node, Union[Callable, PregelProtocol]]] = [ + *_find_children(entrypoint, node) + ] + seen: set[Union[Callable, PregelProtocol]] = set() + for parent, child in candidates: + if child in seen: + continue + else: + seen.add(child) + if callable(child): + node = Node(f"__{child.__name__}", child.__name__, child, None) # type: ignore[arg-type] + graph.nodes[node.id] = node + graph.add_edge(parent, node, conditional=True) + graph.add_edge(node, parent) + candidates.extend(_find_children(child, node)) + elif isinstance(child, Runnable): + if xray > 0: + graph = child.get_graph(config, xray=xray - 1 if xray else 0) + graph.trim_first_node() + graph.trim_last_node() + s, e = graph.extend(graph, prefix=child.name or "") + if s is None: + raise ValueError( + f"Could not extend subgraph '{child.name}' due to missing entrypoint" + ) + else: + graph.add_edge(parent, s, conditional=True) + if e is not None: + graph.add_edge(e, parent) + else: + node = graph.add_node(child, child.name) + graph.add_edge(parent, node, conditional=True) + graph.add_edge(node, parent) + return graph + + def _find_children( - candidate: Runnable, parent: Node + candidate: Union[Callable, Runnable], parent: Node ) -> Iterator[tuple[Node, Union[Callable, PregelProtocol]]]: from langchain_core.runnables.utils import get_function_nonlocals @@ -196,10 +243,9 @@ def _find_children( RunnableSequence, ) - candidates: list[Runnable] = [candidate] + candidates: list[Union[Callable, Runnable]] = [candidate] for c in candidates: - print(c, type(c)) if callable(c) and getattr(c, "_is_pregel_task", False) is True: yield (parent, c) elif isinstance(c, PregelProtocol): @@ -219,43 +265,3 @@ def _find_children( nl.__self__ if hasattr(nl, "__self__") else nl for nl in get_function_nonlocals(c.afunc) ) - - -def _entrypoint_graph(entrypoint: Runnable, xray: int = 0) -> Graph: - graph = Graph() - node = Node(f"__{entrypoint.name}", entrypoint.name, entrypoint, None) - graph.nodes[node.id] = node - candidates: list[tuple[Node, Union[Callable, PregelProtocol]]] = [ - *_find_children(entrypoint, node) - ] - seen: set[Callable] = set() - for parent, child in candidates: - if child in seen: - continue - else: - seen.add(child) - if callable(child): - node = Node(f"__{child.__name__}", child.__name__, child, None) - graph.nodes[node.id] = node - graph.add_edge(parent, node, conditional=True) - graph.add_edge(node, parent) - candidates.extend(_find_children(child, node)) - elif isinstance(child, Runnable): - if xray > 0: - graph = child.get_graph(xray=xray - 1 if xray else 0) - graph.trim_first_node() - graph.trim_last_node() - s, e = graph.extend(graph, prefix=child.name) - if s is None: - raise ValueError( - f"Could not extend subgraph '{child.name}' due to missing entrypoint" - ) - else: - graph.add_edge(parent, s, conditional=True) - if e is not None: - graph.add_edge(e, parent) - else: - node = graph.add_node(child, child.name) - graph.add_edge(parent, node, conditional=True) - graph.add_edge(node, parent) - return graph diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index fe6e56f5d..877ddcf4e 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -238,8 +238,6 @@ class Pregel(PregelProtocol): config: Optional[RunnableConfig] = None - graph: Optional[Graph] = None - name: str = "LangGraph" def __init__( @@ -262,7 +260,6 @@ class Pregel(PregelProtocol): retry_policy: Optional[RetryPolicy] = None, config_type: Optional[Type[Any]] = None, config: Optional[RunnableConfig] = None, - graph: Optional[Graph] = None, name: str = "LangGraph", ) -> None: self.nodes = nodes @@ -281,7 +278,6 @@ class Pregel(PregelProtocol): self.retry_policy = retry_policy self.config_type = config_type self.config = config - self.graph = graph self.name = name if auto_validate: self.validate() @@ -289,17 +285,11 @@ class Pregel(PregelProtocol): def get_graph( self, config: RunnableConfig | None = None, *, xray: int | bool = False ) -> Graph: - if self.graph is not None: - # TODO xray - return self.graph raise NotImplementedError async def aget_graph( self, config: RunnableConfig | None = None, *, xray: int | bool = False ) -> Graph: - if self.graph is not None: - # TODO xray - return self.graph raise NotImplementedError def copy(self, update: dict[str, Any] | None = None) -> Self: