Implement get_graph for imperative api

This commit is contained in:
Nuno Campos
2025-01-16 13:44:18 -08:00
parent 47122ce88b
commit 86b6afc982
4 changed files with 192 additions and 8 deletions
+85 -3
View File
@@ -4,6 +4,7 @@ import concurrent.futures
import functools
import inspect
import types
from collections.abc import Iterator
from typing import (
Any,
Awaitable,
@@ -14,6 +15,8 @@ from typing import (
overload,
)
from langchain_core.runnables.base import Runnable
from langchain_core.runnables.graph import Graph, Node
from typing_extensions import ParamSpec
from langgraph.channels.ephemeral_value import EphemeralValue
@@ -22,6 +25,7 @@ from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.constants import CONF, END, START, TAG_HIDDEN
from langgraph.pregel import Pregel
from langgraph.pregel.call import get_runnable_for_func
from langgraph.pregel.protocol import PregelProtocol
from langgraph.pregel.read import PregelNode
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
from langgraph.store.base import BaseStore
@@ -96,9 +100,9 @@ def task(
def _tick(__allargs__: tuple) -> T:
return func(*__allargs__[0], **__allargs__[1])
return functools.update_wrapper(
functools.partial(call, _tick, retry=retry), func
)
wrapper = functools.partial(call, _tick, retry=retry)
object.__setattr__(wrapper, "_is_pregel_task", True)
return functools.update_wrapper(wrapper, func)
if __func_or_none__ is not None:
return decorator(__func_or_none__)
@@ -174,6 +178,84 @@ def entrypoint(
checkpointer=checkpointer,
store=store,
config_type=config_schema,
graph=_entrypoint_graph(bound),
)
return _imp
def _find_children(
candidate: Runnable, parent: Node
) -> Iterator[tuple[Node, Union[Callable, PregelProtocol]]]:
from langchain_core.runnables.utils import get_function_nonlocals
from langgraph.utils.runnable import (
RunnableCallable,
RunnableLambda,
RunnableSeq,
RunnableSequence,
)
candidates: list[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):
yield (parent, c)
elif isinstance(c, RunnableSequence) or isinstance(c, RunnableSeq):
candidates.extend(c.steps)
elif isinstance(c, RunnableLambda):
candidates.extend(c.deps)
elif isinstance(c, RunnableCallable):
if c.func is not None:
candidates.extend(
nl.__self__ if hasattr(nl, "__self__") else nl
for nl in get_function_nonlocals(c.func)
)
elif c.afunc is not None:
candidates.extend(
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,6 +238,8 @@ class Pregel(PregelProtocol):
config: Optional[RunnableConfig] = None
graph: Optional[Graph] = None
name: str = "LangGraph"
def __init__(
@@ -260,6 +262,7 @@ 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
@@ -278,6 +281,7 @@ 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()
@@ -285,11 +289,17 @@ 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:
File diff suppressed because one or more lines are too long
+21 -5
View File
@@ -1523,7 +1523,7 @@ def test_imp_task(request: pytest.FixtureRequest, checkpointer_name: str) -> Non
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_imp_stream_order(
request: pytest.FixtureRequest, checkpointer_name: str
request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion
) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
@@ -1546,6 +1546,9 @@ def test_imp_stream_order(
fut_baz = baz(fut_bar.result())
return fut_baz.result()
if checkpointer_name == "memory":
assert graph.get_graph().draw_mermaid() == snapshot
thread1 = {"configurable": {"thread_id": "1"}}
assert [c for c in graph.stream({"a": "0"}, thread1)] == [
{
@@ -4951,7 +4954,7 @@ def test_interrupt_loop(request: pytest.FixtureRequest, checkpointer_name: str):
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_interrupt_functional(
request: pytest.FixtureRequest, checkpointer_name: str
request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion
) -> None:
checkpointer: BaseCheckpointSaver = request.getfixturevalue(
f"checkpointer_{checkpointer_name}"
@@ -4973,6 +4976,8 @@ def test_interrupt_functional(
fut_bar = bar(bar_input)
return fut_bar.result()
assert graph.get_graph().draw_mermaid() == snapshot
config = {"configurable": {"thread_id": "1"}}
# First run, interrupted at bar
graph.invoke({"a": ""}, config)
@@ -4983,7 +4988,7 @@ def test_interrupt_functional(
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_interrupt_task_functional(
request: pytest.FixtureRequest, checkpointer_name: str
request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion
) -> None:
checkpointer: BaseCheckpointSaver = request.getfixturevalue(
f"checkpointer_{checkpointer_name}"
@@ -5004,6 +5009,9 @@ def test_interrupt_task_functional(
fut_bar = bar(fut_foo.result())
return fut_bar.result()
if checkpointer_name == "memory":
assert graph.get_graph().draw_mermaid() == snapshot
config = {"configurable": {"thread_id": "1"}}
# First run, interrupted at bar
graph.invoke({"a": ""}, config)
@@ -5432,7 +5440,9 @@ def test_multiple_updates() -> None:
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_name: str):
def test_falsy_return_from_task(
request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion
):
"""Test with a falsy return from a task."""
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
@@ -5446,6 +5456,9 @@ def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_nam
falsy_task().result()
interrupt("test")
if checkpointer_name == "memory":
assert graph.get_graph().draw_mermaid() == snapshot
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
graph.invoke({"a": 5}, configurable)
graph.invoke(Command(resume="123"), configurable)
@@ -5453,7 +5466,7 @@ def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_nam
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_multiple_interrupts_imperative(
request: pytest.FixtureRequest, checkpointer_name: str
request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion
):
"""Test multiple interrupts with an imperative API."""
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
@@ -5478,6 +5491,9 @@ def test_multiple_interrupts_imperative(
return {"values": values}
if checkpointer_name == "memory":
assert graph.get_graph().draw_mermaid() == snapshot
configurable = {"configurable": {"thread_id": str(uuid.uuid4())}}
graph.invoke({}, configurable)
graph.invoke(Command(resume="a"), configurable)