mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 04:55:09 +02:00
Implement get_graph for imperative api
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user