mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-08 02:37:52 +02:00
Port to new architecture
This commit is contained in:
@@ -246,7 +246,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
||||
# find the latest checkpoint for the thread_id
|
||||
if config["configurable"].get("thread_ts"):
|
||||
await cur.execute(
|
||||
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
|
||||
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts <= ? ORDER BY thread_ts DESC LIMIT 1",
|
||||
(
|
||||
str(config["configurable"]["thread_id"]),
|
||||
str(config["configurable"]["thread_ts"]),
|
||||
|
||||
@@ -245,7 +245,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
# find the latest checkpoint for the thread_id
|
||||
if config["configurable"].get("thread_ts"):
|
||||
cur.execute(
|
||||
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
|
||||
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts <= ? ORDER BY thread_ts DESC LIMIT 1",
|
||||
(
|
||||
str(config["configurable"]["thread_id"]),
|
||||
str(config["configurable"]["thread_ts"]),
|
||||
|
||||
@@ -3,9 +3,19 @@ from typing import Any
|
||||
INPUT = "__input__"
|
||||
CONFIG_KEY_SEND = "__pregel_send"
|
||||
CONFIG_KEY_READ = "__pregel_read"
|
||||
CONFIG_KEY_CHECKPOINTER = "__pregel_checkpointer"
|
||||
CONFIG_KEY_RESUMING = "__pregel_resuming"
|
||||
INTERRUPT = "__interrupt__"
|
||||
TASKS = "__pregel_tasks"
|
||||
RESERVED = {INTERRUPT, TASKS, CONFIG_KEY_SEND, CONFIG_KEY_READ, INPUT}
|
||||
RESERVED = {
|
||||
INTERRUPT,
|
||||
TASKS,
|
||||
CONFIG_KEY_SEND,
|
||||
CONFIG_KEY_READ,
|
||||
CONFIG_KEY_CHECKPOINTER,
|
||||
CONFIG_KEY_RESUMING,
|
||||
INPUT,
|
||||
}
|
||||
TAG_HIDDEN = "langsmith:hidden"
|
||||
|
||||
START = "__start__"
|
||||
|
||||
@@ -28,3 +28,15 @@ class InvalidUpdateError(Exception):
|
||||
"""Raised when attempting to update a channel with an invalid sequence of updates."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class GraphInterrupt(Exception):
|
||||
"""Raised when a subgraph is interrupted."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class EmptyInputError(Exception):
|
||||
"""Raised when graph receives an empty input."""
|
||||
|
||||
pass
|
||||
|
||||
@@ -63,7 +63,9 @@ from langgraph.checkpoint.base import (
|
||||
empty_checkpoint,
|
||||
)
|
||||
from langgraph.constants import (
|
||||
CONFIG_KEY_CHECKPOINTER,
|
||||
CONFIG_KEY_READ,
|
||||
CONFIG_KEY_RESUMING,
|
||||
CONFIG_KEY_SEND,
|
||||
INTERRUPT,
|
||||
)
|
||||
@@ -281,7 +283,13 @@ class Pregel(
|
||||
)
|
||||
)
|
||||
# these are provided by the Pregel class
|
||||
if spec.id not in [CONFIG_KEY_READ, CONFIG_KEY_SEND]
|
||||
if spec.id
|
||||
not in [
|
||||
CONFIG_KEY_READ,
|
||||
CONFIG_KEY_SEND,
|
||||
CONFIG_KEY_CHECKPOINTER,
|
||||
CONFIG_KEY_RESUMING,
|
||||
]
|
||||
]
|
||||
|
||||
@property
|
||||
@@ -699,6 +707,7 @@ class Pregel(
|
||||
Union[str, Sequence[str]],
|
||||
Optional[Sequence[str]],
|
||||
Optional[Sequence[str]],
|
||||
Optional[BaseCheckpointSaver],
|
||||
]:
|
||||
debug = debug if debug is not None else self.debug
|
||||
if output_keys is None:
|
||||
@@ -710,15 +719,24 @@ class Pregel(
|
||||
stream_mode = stream_mode if stream_mode is not None else self.stream_mode
|
||||
if not isinstance(stream_mode, list):
|
||||
stream_mode = [stream_mode]
|
||||
if config is not None and config.get("configurable", {}).get(CONFIG_KEY_READ):
|
||||
if config and config.get("configurable", {}).get(CONFIG_KEY_READ) is not None:
|
||||
# if being called as a node in another graph, always use values mode
|
||||
stream_mode = ["values"]
|
||||
if config is not None and config.get("configurable", {}).get(
|
||||
CONFIG_KEY_CHECKPOINTER
|
||||
):
|
||||
checkpointer: Optional[BaseCheckpointSaver] = config["configurable"][
|
||||
CONFIG_KEY_CHECKPOINTER
|
||||
]
|
||||
else:
|
||||
checkpointer = self.checkpointer
|
||||
return (
|
||||
debug,
|
||||
stream_mode,
|
||||
output_keys,
|
||||
interrupt_before,
|
||||
interrupt_after,
|
||||
checkpointer,
|
||||
)
|
||||
|
||||
def stream(
|
||||
@@ -820,6 +838,7 @@ class Pregel(
|
||||
output_keys,
|
||||
interrupt_before,
|
||||
interrupt_after,
|
||||
checkpointer,
|
||||
) = self._defaults(
|
||||
config,
|
||||
stream_mode=stream_mode,
|
||||
@@ -830,7 +849,7 @@ class Pregel(
|
||||
)
|
||||
|
||||
with SyncPregelLoop(
|
||||
input, config=config, checkpointer=self.checkpointer, graph=self
|
||||
input, config=config, checkpointer=checkpointer, graph=self
|
||||
) as loop:
|
||||
# Similarly to Bulk Synchronous Parallel / Pregel model
|
||||
# computation proceeds in steps, while there are channel updates
|
||||
@@ -1063,6 +1082,7 @@ class Pregel(
|
||||
output_keys,
|
||||
interrupt_before,
|
||||
interrupt_after,
|
||||
checkpointer,
|
||||
) = self._defaults(
|
||||
config,
|
||||
stream_mode=stream_mode,
|
||||
@@ -1072,7 +1092,7 @@ class Pregel(
|
||||
debug=debug,
|
||||
)
|
||||
async with AsyncPregelLoop(
|
||||
input, config=config, checkpointer=self.checkpointer, graph=self
|
||||
input, config=config, checkpointer=checkpointer, graph=self
|
||||
) as loop:
|
||||
aioloop = asyncio.get_event_loop()
|
||||
# Similarly to Bulk Synchronous Parallel / Pregel model
|
||||
@@ -1201,6 +1221,7 @@ class Pregel(
|
||||
read_channels(loop.channels, output_keys)
|
||||
)
|
||||
except BaseException as e:
|
||||
# TODO use on_chain_end if exc is GraphInterrupt
|
||||
await asyncio.shield(run_manager.on_chain_error(e))
|
||||
raise
|
||||
|
||||
|
||||
@@ -26,9 +26,11 @@ from langchain_core.runnables.config import (
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.channels.context import Context
|
||||
from langgraph.channels.manager import ChannelsManager, create_checkpoint
|
||||
from langgraph.checkpoint.base import Checkpoint, copy_checkpoint
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver, Checkpoint, copy_checkpoint
|
||||
from langgraph.constants import (
|
||||
CONFIG_KEY_CHECKPOINTER,
|
||||
CONFIG_KEY_READ,
|
||||
CONFIG_KEY_RESUMING,
|
||||
CONFIG_KEY_SEND,
|
||||
INTERRUPT,
|
||||
RESERVED,
|
||||
@@ -213,6 +215,8 @@ def prepare_next_tasks(
|
||||
config: RunnableConfig,
|
||||
step: int,
|
||||
for_execution: Literal[False],
|
||||
is_resuming: bool = False,
|
||||
checkpointer: Literal[None] = None,
|
||||
manager: Literal[None] = None,
|
||||
) -> list[PregelTaskDescription]:
|
||||
...
|
||||
@@ -227,6 +231,8 @@ def prepare_next_tasks(
|
||||
config: RunnableConfig,
|
||||
step: int,
|
||||
for_execution: Literal[True],
|
||||
is_resuming: bool,
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
manager: Union[None, ParentRunManager, AsyncParentRunManager],
|
||||
) -> list[PregelExecutableTask]:
|
||||
...
|
||||
@@ -241,6 +247,8 @@ def prepare_next_tasks(
|
||||
step: int,
|
||||
*,
|
||||
for_execution: bool,
|
||||
is_resuming: bool = False,
|
||||
checkpointer: Optional[BaseCheckpointSaver] = None,
|
||||
manager: Union[None, ParentRunManager, AsyncParentRunManager] = None,
|
||||
) -> Union[list[PregelTaskDescription], list[PregelExecutableTask]]:
|
||||
tasks: Union[list[PregelTaskDescription], list[PregelExecutableTask]] = []
|
||||
@@ -291,6 +299,8 @@ def prepare_next_tasks(
|
||||
PregelTaskWrites(packet.node, writes, triggers),
|
||||
config,
|
||||
),
|
||||
# in Send we can't checkpoint nested graphs
|
||||
# as they could be running in parallel
|
||||
},
|
||||
),
|
||||
triggers,
|
||||
@@ -332,6 +342,12 @@ def prepare_next_tasks(
|
||||
"langgraph_task_idx": len(tasks),
|
||||
}
|
||||
task_id = str(uuid5(UUID(checkpoint["id"]), json.dumps(metadata)))
|
||||
if parent_thread_id := config.get("configurable", {}).get(
|
||||
"thread_id"
|
||||
):
|
||||
thread_id: Optional[str] = f"{parent_thread_id}-{name}"
|
||||
else:
|
||||
thread_id = None
|
||||
writes = deque()
|
||||
tasks.append(
|
||||
PregelExecutableTask(
|
||||
@@ -363,6 +379,9 @@ def prepare_next_tasks(
|
||||
PregelTaskWrites(name, writes, triggers),
|
||||
config,
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINTER: checkpointer,
|
||||
CONFIG_KEY_RESUMING: is_resuming,
|
||||
"thread_id": thread_id,
|
||||
},
|
||||
),
|
||||
triggers,
|
||||
|
||||
@@ -6,7 +6,6 @@ from contextvars import copy_context
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
AsyncContextManager,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Iterator,
|
||||
Optional,
|
||||
@@ -18,6 +17,8 @@ from langchain_core.runnables import RunnableConfig
|
||||
from langchain_core.runnables.config import get_executor_for_config
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
from langgraph.errors import GraphInterrupt
|
||||
|
||||
P = ParamSpec("P")
|
||||
T = TypeVar("T")
|
||||
|
||||
@@ -42,6 +43,10 @@ def BackgroundExecutor(config: RunnableConfig) -> Iterator[Submit]:
|
||||
def done(task: concurrent.futures.Future) -> None:
|
||||
try:
|
||||
task.result()
|
||||
except GraphInterrupt:
|
||||
# This exception is an interruption signal, not an error
|
||||
# so we don't want to re-raise it on exit
|
||||
tasks.pop(task)
|
||||
except BaseException:
|
||||
pass
|
||||
else:
|
||||
@@ -79,7 +84,7 @@ class AsyncBackgroundExecutor(AsyncContextManager):
|
||||
|
||||
def submit(
|
||||
self,
|
||||
fn: Callable[P, Awaitable[T]],
|
||||
fn: Callable[P, T],
|
||||
*args: P.args,
|
||||
__name__: Optional[str] = None,
|
||||
__cancel_on_exit__: bool = False,
|
||||
@@ -97,12 +102,16 @@ class AsyncBackgroundExecutor(AsyncContextManager):
|
||||
def done(self, task: asyncio.Task) -> None:
|
||||
try:
|
||||
task.result()
|
||||
except GraphInterrupt:
|
||||
# This exception is an interruption signal, not an error
|
||||
# so we don't want to re-raise it on exit
|
||||
self.tasks.pop(task)
|
||||
except BaseException:
|
||||
pass
|
||||
else:
|
||||
self.tasks.pop(task)
|
||||
|
||||
async def __aenter__(self) -> "submit":
|
||||
async def __aenter__(self) -> Submit:
|
||||
return self.submit
|
||||
|
||||
async def exit(self) -> None:
|
||||
|
||||
@@ -38,7 +38,8 @@ from langgraph.checkpoint.base import (
|
||||
copy_checkpoint,
|
||||
empty_checkpoint,
|
||||
)
|
||||
from langgraph.constants import INPUT, INTERRUPT
|
||||
from langgraph.constants import CONFIG_KEY_READ, CONFIG_KEY_RESUMING, INPUT, INTERRUPT
|
||||
from langgraph.errors import EmptyInputError, GraphInterrupt
|
||||
from langgraph.managed.base import (
|
||||
AsyncManagedValuesManager,
|
||||
ManagedValueMapping,
|
||||
@@ -66,6 +67,7 @@ if TYPE_CHECKING:
|
||||
|
||||
V = TypeVar("V")
|
||||
INPUT_DONE = object()
|
||||
INPUT_RESUMING = object()
|
||||
|
||||
|
||||
class PregelLoop:
|
||||
@@ -95,6 +97,7 @@ class PregelLoop:
|
||||
]
|
||||
tasks: Sequence[PregelExecutableTask]
|
||||
stream: deque[Tuple[str, Any]]
|
||||
is_nested: bool
|
||||
|
||||
# public
|
||||
|
||||
@@ -133,7 +136,7 @@ class PregelLoop:
|
||||
if self.status != "pending":
|
||||
raise RuntimeError("Cannot tick when status is no longer 'pending'")
|
||||
|
||||
if self.input is not INPUT_DONE:
|
||||
if self.input not in (INPUT_DONE, INPUT_RESUMING):
|
||||
self._first()
|
||||
elif all(task.writes for task in self.tasks):
|
||||
writes = [w for t in self.tasks for w in t.writes]
|
||||
@@ -165,7 +168,10 @@ class PregelLoop:
|
||||
# after execution, check if we should interrupt
|
||||
if should_interrupt(self.checkpoint, interrupt_after, self.tasks):
|
||||
self.status = "interrupt_after"
|
||||
return False
|
||||
if self.is_nested:
|
||||
raise GraphInterrupt(self)
|
||||
else:
|
||||
return False
|
||||
else:
|
||||
return False
|
||||
|
||||
@@ -184,6 +190,8 @@ class PregelLoop:
|
||||
self.step,
|
||||
for_execution=True,
|
||||
manager=manager,
|
||||
checkpointer=self.checkpointer,
|
||||
is_resuming=self.input is INPUT_RESUMING,
|
||||
)
|
||||
|
||||
# if no more tasks, we're done
|
||||
@@ -204,7 +212,10 @@ class PregelLoop:
|
||||
# before execution, check if we should interrupt
|
||||
if should_interrupt(self.checkpoint, interrupt_before, self.tasks):
|
||||
self.status = "interrupt_before"
|
||||
return False
|
||||
if self.is_nested:
|
||||
raise GraphInterrupt(self)
|
||||
else:
|
||||
return False
|
||||
|
||||
# produce debug output
|
||||
self.stream.extend(("debug", v) for v in map_debug_tasks(self.step, self.tasks))
|
||||
@@ -214,8 +225,23 @@ class PregelLoop:
|
||||
# private
|
||||
|
||||
def _first(self) -> None:
|
||||
# resuming from previous checkpoint requires
|
||||
# - finding a previous checkpoint
|
||||
# - receiving None input (outer graph) or RESUMING flag (subgraph)
|
||||
is_resuming = bool(self.checkpoint["channel_versions"]) and bool(
|
||||
self.config.get("configurable", {}).get(CONFIG_KEY_RESUMING)
|
||||
or self.input is None
|
||||
)
|
||||
|
||||
# proceed past previous checkpoint
|
||||
if is_resuming:
|
||||
self.checkpoint["versions_seen"].setdefault(INTERRUPT, {})
|
||||
for k in self.channels:
|
||||
if k in self.checkpoint["channel_versions"]:
|
||||
version = self.checkpoint["channel_versions"][k]
|
||||
self.checkpoint["versions_seen"][INTERRUPT][k] = version
|
||||
# map inputs to channel updates
|
||||
if input_writes := deque(map_input(self.graph.input_channels, self.input)):
|
||||
elif input_writes := deque(map_input(self.graph.input_channels, self.input)):
|
||||
# discard any unfinished tasks from previous checkpoint
|
||||
discard_tasks = prepare_next_tasks(
|
||||
self.checkpoint,
|
||||
@@ -236,14 +262,9 @@ class PregelLoop:
|
||||
# save input checkpoint
|
||||
self._put_checkpoint({"source": "input", "writes": self.input})
|
||||
else:
|
||||
# no input is taken as signal to proceed past previous interrupt
|
||||
self.checkpoint["versions_seen"].setdefault(INTERRUPT, {})
|
||||
for k in self.channels:
|
||||
if k in self.checkpoint["channel_versions"]:
|
||||
version = self.checkpoint["channel_versions"][k]
|
||||
self.checkpoint["versions_seen"][INTERRUPT][k] = version
|
||||
raise EmptyInputError(f"Received no input for {self.graph.input_channels}")
|
||||
# done with input
|
||||
self.input = INPUT_DONE
|
||||
self.input = INPUT_RESUMING if is_resuming else INPUT_DONE
|
||||
|
||||
def _put_checkpoint(
|
||||
self,
|
||||
@@ -313,6 +334,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
# context manager
|
||||
|
||||
def __enter__(self) -> Self:
|
||||
self.is_nested = CONFIG_KEY_READ in self.config.get("configurable", {})
|
||||
saved = (
|
||||
self.checkpointer.get_tuple(self.config) if self.checkpointer else None
|
||||
) or CheckpointTuple(self.config, empty_checkpoint(), {"step": -2}, None, [])
|
||||
@@ -348,6 +370,20 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
exc_value: Optional[BaseException],
|
||||
traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
# handle interrupt
|
||||
if exc_type is GraphInterrupt:
|
||||
if exc_value.args[0] is self:
|
||||
# interrupt raised by this loop
|
||||
exc_value.args = (object(),)
|
||||
else:
|
||||
# interrupt raised by a nested loop, save interrupt checkpoint
|
||||
self._put_checkpoint({"source": "interrupt"})
|
||||
if not self.is_nested:
|
||||
# in outer graph, catch interrupt
|
||||
del self.graph
|
||||
return True and self.stack.__exit__(None, None, None)
|
||||
|
||||
# unwind stack
|
||||
del self.graph
|
||||
return self.stack.__exit__(exc_type, exc_value, traceback)
|
||||
|
||||
@@ -380,6 +416,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
# context manager
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
self.is_nested = CONFIG_KEY_READ in self.config.get("configurable", {})
|
||||
saved = (
|
||||
await self.checkpointer.aget_tuple(self.config)
|
||||
if self.checkpointer
|
||||
@@ -417,6 +454,22 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
exc_value: Optional[BaseException],
|
||||
traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
# handle interrupt
|
||||
if exc_type is GraphInterrupt:
|
||||
if exc_value.args[0] is self:
|
||||
# interrupt raised by this loop
|
||||
exc_value.args = (object(),)
|
||||
else:
|
||||
# interrupt raised by a nested loop, save interrupt checkpoint
|
||||
self._put_checkpoint({"source": "interrupt"})
|
||||
if not self.is_nested:
|
||||
# in outer graph, catch interrupt
|
||||
del self.graph
|
||||
return True and await asyncio.shield(
|
||||
self.stack.__aexit__(None, None, None)
|
||||
)
|
||||
|
||||
# unwind stack
|
||||
del self.graph
|
||||
return await asyncio.shield(
|
||||
self.stack.__aexit__(exc_type, exc_value, traceback)
|
||||
|
||||
@@ -7524,6 +7524,381 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None:
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
SqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
def test_nested_graph_interrupts(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: str
|
||||
my_other_key: str
|
||||
|
||||
def inner_1(state: InnerState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " here",
|
||||
"my_other_key": state["my_key"],
|
||||
}
|
||||
|
||||
def inner_2(state: InnerState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " and there",
|
||||
"my_other_key": state["my_key"],
|
||||
}
|
||||
|
||||
inner = StateGraph(InnerState)
|
||||
inner.add_node("inner_1", inner_1)
|
||||
inner.add_node("inner_2", inner_2)
|
||||
inner.add_edge("inner_1", "inner_2")
|
||||
inner.set_entry_point("inner_1")
|
||||
inner.set_finish_point("inner_2")
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: str
|
||||
|
||||
def outer_1(state: State):
|
||||
return {"my_key": "hi " + state["my_key"]}
|
||||
|
||||
def outer_2(state: State):
|
||||
return {"my_key": state["my_key"] + " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
graph.add_node("outer_1", outer_1)
|
||||
graph.add_node("inner", inner.compile(interrupt_before=["inner_2"]))
|
||||
graph.add_node("outer_2", outer_2)
|
||||
graph.set_entry_point("outer_1")
|
||||
graph.add_edge("outer_1", "inner")
|
||||
graph.add_edge("inner", "outer_2")
|
||||
graph.set_finish_point("outer_2")
|
||||
|
||||
app = graph.compile(checkpointer=checkpointer)
|
||||
|
||||
# test invoke w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
assert app.invoke({"my_key": "my value"}, config, debug=True) == {
|
||||
"my_key": "hi my value",
|
||||
}
|
||||
|
||||
assert app.invoke(None, config, debug=True) == {
|
||||
"my_key": "hi my value here and there and back again",
|
||||
}
|
||||
|
||||
# test stream updates w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
assert [*app.stream({"my_key": "my value"}, config)] == [
|
||||
{"outer_1": {"my_key": "hi my value"}},
|
||||
]
|
||||
assert [*app.stream(None, config)] == [
|
||||
{"inner": {"my_key": "hi my value here and there"}},
|
||||
{"outer_2": {"my_key": "hi my value here and there and back again"}},
|
||||
]
|
||||
|
||||
# test stream values w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "3"}}
|
||||
assert [*app.stream({"my_key": "my value"}, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# test interrupts BEFORE the node w/ interrupts
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_before=["inner"])
|
||||
config = {"configurable": {"thread_id": "4"}}
|
||||
assert [*app.stream({"my_key": "my value"}, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
# while we're waiting for the node w/ interrupt inside to finish
|
||||
assert [*app.stream(None, config, stream_mode="values")] == []
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# test interrupts AFTER the node w/ interrupts
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_after=["inner"])
|
||||
config = {"configurable": {"thread_id": "5"}}
|
||||
assert [*app.stream({"my_key": "my value"}, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
]
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
SqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
def test_nested_graph_interrupts_parallel(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: Annotated[str, operator.add]
|
||||
my_other_key: str
|
||||
|
||||
def inner_1(state: InnerState):
|
||||
return {"my_key": "got here", "my_other_key": state["my_key"]}
|
||||
|
||||
def inner_2(state: InnerState):
|
||||
return {
|
||||
"my_key": " and there",
|
||||
"my_other_key": state["my_key"],
|
||||
}
|
||||
|
||||
inner = StateGraph(InnerState)
|
||||
inner.add_node("inner_1", inner_1)
|
||||
inner.add_node("inner_2", inner_2)
|
||||
inner.add_edge("inner_1", "inner_2")
|
||||
inner.set_entry_point("inner_1")
|
||||
inner.set_finish_point("inner_2")
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: Annotated[str, operator.add]
|
||||
|
||||
def outer_1(state: State):
|
||||
return {"my_key": " and parallel"}
|
||||
|
||||
def outer_2(state: State):
|
||||
return {"my_key": " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
graph.add_node("inner", inner.compile(interrupt_before=["inner_2"]))
|
||||
graph.add_node("outer_1", outer_1)
|
||||
graph.add_node("outer_2", outer_2)
|
||||
|
||||
graph.add_edge(START, "inner")
|
||||
graph.add_edge(START, "outer_1")
|
||||
graph.add_edge(["inner", "outer_1"], "outer_2")
|
||||
graph.set_finish_point("outer_2")
|
||||
|
||||
app = graph.compile(checkpointer=checkpointer)
|
||||
|
||||
# test invoke w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
assert app.invoke({"my_key": ""}, config, debug=True) == {
|
||||
"my_key": "",
|
||||
}
|
||||
|
||||
assert app.invoke(None, config, debug=True) == {
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
}
|
||||
|
||||
# test stream updates w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
assert [*app.stream({"my_key": ""}, config)] == [
|
||||
# we got to parallel node first
|
||||
{"outer_1": {"my_key": " and parallel"}},
|
||||
]
|
||||
assert [*app.stream(None, config)] == [
|
||||
{"inner": {"my_key": "got here and there"}},
|
||||
{"outer_2": {"my_key": " and back again"}},
|
||||
]
|
||||
|
||||
# test stream values w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "3"}}
|
||||
assert [*app.stream({"my_key": ""}, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "",
|
||||
},
|
||||
]
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "got here and there and parallel",
|
||||
},
|
||||
{
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# test interrupts BEFORE the parallel node
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_before=["outer_1"])
|
||||
config = {"configurable": {"thread_id": "4"}}
|
||||
assert [*app.stream({"my_key": ""}, config, stream_mode="values")] == [
|
||||
{"my_key": ""}
|
||||
]
|
||||
# while we're waiting for the node w/ interrupt inside to finish
|
||||
assert [*app.stream(None, config, stream_mode="values")] == []
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "got here and there and parallel",
|
||||
},
|
||||
{
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# test interrupts AFTER the parallel node
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_after=["outer_1"])
|
||||
config = {"configurable": {"thread_id": "5"}}
|
||||
assert [*app.stream({"my_key": ""}, config, stream_mode="values")] == [
|
||||
{"my_key": ""}
|
||||
]
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{"my_key": "got here and there and parallel"},
|
||||
]
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
SqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
def test_doubly_nested_graph_interrupts(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: str
|
||||
|
||||
class ChildState(TypedDict):
|
||||
my_key: str
|
||||
|
||||
class GrandChildState(TypedDict):
|
||||
my_key: str
|
||||
|
||||
def grandchild_1(state: ChildState):
|
||||
return {"my_key": state["my_key"] + " here"}
|
||||
|
||||
def grandchild_2(state: ChildState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " and there",
|
||||
}
|
||||
|
||||
grandchild = StateGraph(GrandChildState)
|
||||
grandchild.add_node("grandchild_1", grandchild_1)
|
||||
grandchild.add_node("grandchild_2", grandchild_2)
|
||||
grandchild.add_edge("grandchild_1", "grandchild_2")
|
||||
grandchild.set_entry_point("grandchild_1")
|
||||
grandchild.set_finish_point("grandchild_2")
|
||||
|
||||
child = StateGraph(ChildState)
|
||||
child.add_node("child_1", grandchild.compile(interrupt_before=["grandchild_2"]))
|
||||
child.set_entry_point("child_1")
|
||||
child.set_finish_point("child_1")
|
||||
|
||||
def parent_1(state: State):
|
||||
return {"my_key": "hi " + state["my_key"]}
|
||||
|
||||
def parent_2(state: State):
|
||||
return {"my_key": state["my_key"] + " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
graph.add_node("parent_1", parent_1)
|
||||
graph.add_node("child", child.compile())
|
||||
graph.add_node("parent_2", parent_2)
|
||||
graph.set_entry_point("parent_1")
|
||||
graph.add_edge("parent_1", "child")
|
||||
graph.add_edge("child", "parent_2")
|
||||
graph.set_finish_point("parent_2")
|
||||
|
||||
app = graph.compile(checkpointer=checkpointer)
|
||||
|
||||
# test invoke w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
assert app.invoke({"my_key": "my value"}, config, debug=True) == {
|
||||
"my_key": "hi my value",
|
||||
}
|
||||
|
||||
assert app.invoke(None, config, debug=True) == {
|
||||
"my_key": "hi my value here and there and back again",
|
||||
}
|
||||
|
||||
# test stream updates w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
assert [*app.stream({"my_key": "my value"}, config)] == [
|
||||
{"parent_1": {"my_key": "hi my value"}},
|
||||
]
|
||||
assert [*app.stream(None, config)] == [
|
||||
{"child": {"my_key": "hi my value here and there"}},
|
||||
{"parent_2": {"my_key": "hi my value here and there and back again"}},
|
||||
]
|
||||
|
||||
# test stream values w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "3"}}
|
||||
assert [*app.stream({"my_key": "my value"}, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
assert [*app.stream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_repeat_condition(snapshot: SnapshotAssertion) -> None:
|
||||
class AgentState(TypedDict):
|
||||
hello: str
|
||||
|
||||
@@ -6016,6 +6016,407 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None:
|
||||
assert times_called == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
async def test_nested_graph_interrupts(checkpointer: BaseCheckpointSaver) -> None:
|
||||
try:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: str
|
||||
my_other_key: str
|
||||
|
||||
async def inner_1(state: InnerState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " here",
|
||||
"my_other_key": state["my_key"],
|
||||
}
|
||||
|
||||
async def inner_2(state: InnerState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " and there",
|
||||
"my_other_key": state["my_key"],
|
||||
}
|
||||
|
||||
inner = StateGraph(InnerState)
|
||||
inner.add_node("inner_1", inner_1)
|
||||
inner.add_node("inner_2", inner_2)
|
||||
inner.add_edge("inner_1", "inner_2")
|
||||
inner.set_entry_point("inner_1")
|
||||
inner.set_finish_point("inner_2")
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: str
|
||||
|
||||
async def outer_1(state: State):
|
||||
return {"my_key": "hi " + state["my_key"]}
|
||||
|
||||
async def outer_2(state: State):
|
||||
return {"my_key": state["my_key"] + " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
graph.add_node("outer_1", outer_1)
|
||||
graph.add_node("inner", inner.compile(interrupt_before=["inner_2"]))
|
||||
graph.add_node("outer_2", outer_2)
|
||||
graph.set_entry_point("outer_1")
|
||||
graph.add_edge("outer_1", "inner")
|
||||
graph.add_edge("inner", "outer_2")
|
||||
graph.set_finish_point("outer_2")
|
||||
|
||||
app = graph.compile(checkpointer=checkpointer)
|
||||
|
||||
# test invoke w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
assert await app.ainvoke({"my_key": "my value"}, config, debug=True) == {
|
||||
"my_key": "hi my value",
|
||||
}
|
||||
|
||||
assert await app.ainvoke(None, config, debug=True) == {
|
||||
"my_key": "hi my value here and there and back again",
|
||||
}
|
||||
|
||||
# test stream updates w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
assert [c async for c in app.astream({"my_key": "my value"}, config)] == [
|
||||
{"outer_1": {"my_key": "hi my value"}},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config)] == [
|
||||
{"inner": {"my_key": "hi my value here and there"}},
|
||||
{"outer_2": {"my_key": "hi my value here and there and back again"}},
|
||||
]
|
||||
|
||||
# test stream values w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "3"}}
|
||||
assert [
|
||||
c
|
||||
async for c in app.astream(
|
||||
{"my_key": "my value"}, config, stream_mode="values"
|
||||
)
|
||||
] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# test interrupts BEFORE the node w/ interrupts
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_before=["inner"])
|
||||
config = {"configurable": {"thread_id": "4"}}
|
||||
assert [
|
||||
c
|
||||
async for c in app.astream(
|
||||
{"my_key": "my value"}, config, stream_mode="values"
|
||||
)
|
||||
] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
# while we're waiting for the node w/ interrupt inside to finish
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == []
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# test interrupts AFTER the node w/ interrupts
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_after=["inner"])
|
||||
config = {"configurable": {"thread_id": "5"}}
|
||||
assert [
|
||||
c
|
||||
async for c in app.astream(
|
||||
{"my_key": "my value"}, config, stream_mode="values"
|
||||
)
|
||||
] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
async def test_nested_graph_interrupts_parallel(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
try:
|
||||
|
||||
class InnerState(TypedDict):
|
||||
my_key: Annotated[str, operator.add]
|
||||
my_other_key: str
|
||||
|
||||
async def inner_1(state: InnerState):
|
||||
return {"my_key": "got here", "my_other_key": state["my_key"]}
|
||||
|
||||
async def inner_2(state: InnerState):
|
||||
return {
|
||||
"my_key": " and there",
|
||||
"my_other_key": state["my_key"],
|
||||
}
|
||||
|
||||
inner = StateGraph(InnerState)
|
||||
inner.add_node("inner_1", inner_1)
|
||||
inner.add_node("inner_2", inner_2)
|
||||
inner.add_edge("inner_1", "inner_2")
|
||||
inner.set_entry_point("inner_1")
|
||||
inner.set_finish_point("inner_2")
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: Annotated[str, operator.add]
|
||||
|
||||
async def outer_1(state: State):
|
||||
return {"my_key": " and parallel"}
|
||||
|
||||
async def outer_2(state: State):
|
||||
return {"my_key": " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
graph.add_node("inner", inner.compile(interrupt_before=["inner_2"]))
|
||||
graph.add_node("outer_1", outer_1)
|
||||
graph.add_node("outer_2", outer_2)
|
||||
|
||||
graph.add_edge(START, "inner")
|
||||
graph.add_edge(START, "outer_1")
|
||||
graph.add_edge(["inner", "outer_1"], "outer_2")
|
||||
graph.set_finish_point("outer_2")
|
||||
|
||||
app = graph.compile(checkpointer=checkpointer)
|
||||
|
||||
# test invoke w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
assert await app.ainvoke({"my_key": ""}, config, debug=True) == {
|
||||
"my_key": "",
|
||||
}
|
||||
|
||||
assert await app.ainvoke(None, config, debug=True) == {
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
}
|
||||
|
||||
# test stream updates w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
assert [c async for c in app.astream({"my_key": ""}, config)] == [
|
||||
# we got to parallel node first
|
||||
{"outer_1": {"my_key": " and parallel"}},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config)] == [
|
||||
{"inner": {"my_key": "got here and there"}},
|
||||
{"outer_2": {"my_key": " and back again"}},
|
||||
]
|
||||
|
||||
# test stream values w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "3"}}
|
||||
assert [
|
||||
c async for c in app.astream({"my_key": ""}, config, stream_mode="values")
|
||||
] == [
|
||||
{
|
||||
"my_key": "",
|
||||
},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "got here and there and parallel",
|
||||
},
|
||||
{
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# # test interrupts BEFORE the parallel node
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_before=["outer_1"])
|
||||
config = {"configurable": {"thread_id": "4"}}
|
||||
assert [
|
||||
c async for c in app.astream({"my_key": ""}, config, stream_mode="values")
|
||||
] == [{"my_key": ""}]
|
||||
# while we're waiting for the node w/ interrupt inside to finish
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == []
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "got here and there and parallel",
|
||||
},
|
||||
{
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
|
||||
# test interrupts AFTER the parallel node
|
||||
app = graph.compile(checkpointer=checkpointer, interrupt_after=["outer_1"])
|
||||
config = {"configurable": {"thread_id": "5"}}
|
||||
assert [
|
||||
c async for c in app.astream({"my_key": ""}, config, stream_mode="values")
|
||||
] == [{"my_key": ""}]
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{"my_key": "got here and there and parallel"},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "got here and there and parallel and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
async def test_doubly_nested_graph_interrupts(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
try:
|
||||
|
||||
class State(TypedDict):
|
||||
my_key: str
|
||||
|
||||
class ChildState(TypedDict):
|
||||
my_key: str
|
||||
|
||||
class GrandChildState(TypedDict):
|
||||
my_key: str
|
||||
|
||||
async def grandchild_1(state: ChildState):
|
||||
return {"my_key": state["my_key"] + " here"}
|
||||
|
||||
async def grandchild_2(state: ChildState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " and there",
|
||||
}
|
||||
|
||||
grandchild = StateGraph(GrandChildState)
|
||||
grandchild.add_node("grandchild_1", grandchild_1)
|
||||
grandchild.add_node("grandchild_2", grandchild_2)
|
||||
grandchild.add_edge("grandchild_1", "grandchild_2")
|
||||
grandchild.set_entry_point("grandchild_1")
|
||||
grandchild.set_finish_point("grandchild_2")
|
||||
|
||||
child = StateGraph(ChildState)
|
||||
child.add_node("child_1", grandchild.compile(interrupt_before=["grandchild_2"]))
|
||||
child.set_entry_point("child_1")
|
||||
child.set_finish_point("child_1")
|
||||
|
||||
async def parent_1(state: State):
|
||||
return {"my_key": "hi " + state["my_key"]}
|
||||
|
||||
async def parent_2(state: State):
|
||||
return {"my_key": state["my_key"] + " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
graph.add_node("parent_1", parent_1)
|
||||
graph.add_node("child", child.compile())
|
||||
graph.add_node("parent_2", parent_2)
|
||||
graph.set_entry_point("parent_1")
|
||||
graph.add_edge("parent_1", "child")
|
||||
graph.add_edge("child", "parent_2")
|
||||
graph.set_finish_point("parent_2")
|
||||
|
||||
app = graph.compile(checkpointer=checkpointer)
|
||||
|
||||
# test invoke w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
assert await app.ainvoke({"my_key": "my value"}, config, debug=True) == {
|
||||
"my_key": "hi my value",
|
||||
}
|
||||
|
||||
assert await app.ainvoke(None, config, debug=True) == {
|
||||
"my_key": "hi my value here and there and back again",
|
||||
}
|
||||
|
||||
# test stream updates w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
assert [c async for c in app.astream({"my_key": "my value"}, config)] == [
|
||||
{"parent_1": {"my_key": "hi my value"}},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config)] == [
|
||||
{"child": {"my_key": "hi my value here and there"}},
|
||||
{"parent_2": {"my_key": "hi my value here and there and back again"}},
|
||||
]
|
||||
|
||||
# test stream values w/ nested interrupt
|
||||
config = {"configurable": {"thread_id": "3"}}
|
||||
assert [
|
||||
c
|
||||
async for c in app.astream(
|
||||
{"my_key": "my value"}, config, stream_mode="values"
|
||||
)
|
||||
] == [
|
||||
{
|
||||
"my_key": "my value",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value",
|
||||
},
|
||||
]
|
||||
assert [c async for c in app.astream(None, config, stream_mode="values")] == [
|
||||
{
|
||||
"my_key": "hi my value here and there",
|
||||
},
|
||||
{
|
||||
"my_key": "hi my value here and there and back again",
|
||||
},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
async def test_checkpoint_metadata() -> None:
|
||||
"""This test verifies that a run's configurable fields are merged with the
|
||||
previous checkpoint config for each step in the run.
|
||||
|
||||
Reference in New Issue
Block a user