This commit is contained in:
Nuno Campos
2025-01-16 14:10:16 -08:00
parent b98a7b09a3
commit 32bf81c559
2 changed files with 51 additions and 55 deletions
+51 -45
View File
@@ -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
@@ -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: