Merge pull request #128 from langchain-ai/nc/20feb/state-api

Add get_state() and update_state() methods to get and update checkpoint in between runs, Add debug arg to Graph/StateGraph
This commit is contained in:
Nuno Campos
2024-02-25 08:48:45 -08:00
committed by GitHub
14 changed files with 1545 additions and 103 deletions
+1 -1
View File
@@ -18,7 +18,7 @@ test:
poetry run pytest
test_watch:
poetry run ptw --snapshot-update --now . -- -vv -x --ff tests
poetry run ptw .
######################
# LINTING AND FORMATTING
+8 -8
View File
@@ -116,16 +116,16 @@ def create_checkpoint(
checkpoint: Checkpoint, channels: Mapping[str, BaseChannel]
) -> Checkpoint:
"""Create a checkpoint for the given channels."""
checkpoint = Checkpoint(
values: dict[str, Any] = {}
for k, v in channels.items():
try:
values[k] = v.checkpoint()
except EmptyChannelError:
pass
return Checkpoint(
v=1,
ts=datetime.now(timezone.utc).isoformat(),
channel_values=checkpoint["channel_values"],
channel_values=values,
channel_versions=checkpoint["channel_versions"],
versions_seen=checkpoint["versions_seen"],
)
for k, v in channels.items():
try:
checkpoint["channel_values"][k] = v.checkpoint()
except EmptyChannelError:
pass
return checkpoint
+11
View File
@@ -1,6 +1,7 @@
import asyncio
from abc import ABC, abstractmethod
from collections import defaultdict
from copy import deepcopy
from datetime import datetime, timezone
from typing import Any, Optional, TypedDict
@@ -33,6 +34,16 @@ def empty_checkpoint() -> Checkpoint:
)
def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
return Checkpoint(
v=checkpoint["v"],
ts=checkpoint["ts"],
channel_values=checkpoint["channel_values"].copy(),
channel_versions=checkpoint["channel_versions"].copy(),
versions_seen=deepcopy(checkpoint["versions_seen"]),
)
class CheckpointAt(StrEnum):
END_OF_STEP = "end_of_step"
END_OF_RUN = "end_of_run"
+1
View File
@@ -1,2 +1,3 @@
CONFIG_KEY_SEND = "__pregel_send"
CONFIG_KEY_READ = "__pregel_read"
INTERRUPT = "__interrupt__"
+12 -4
View File
@@ -12,6 +12,7 @@ from langchain_core.runnables.base import (
from langchain_core.runnables.config import RunnableConfig
from langchain_core.runnables.graph import Graph as RunnableGraph
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.checkpoint import BaseCheckpointSaver
from langgraph.pregel import Channel, Pregel
@@ -177,6 +178,7 @@ class Graph:
checkpointer: Optional[BaseCheckpointSaver] = None,
interrupt_before: Optional[Sequence[str]] = None,
interrupt_after: Optional[Sequence[str]] = None,
debug: bool = False,
) -> "CompiledGraph":
interrupt_before = interrupt_before or []
interrupt_after = interrupt_after or []
@@ -190,6 +192,11 @@ class Graph:
key: (Channel.subscribe_to(f"{key}:inbox") | node | Channel.write_to(key))
for key, node in self.nodes.items()
}
node_outboxes = {
# we clear outbox channels after each step
key: EphemeralValue(Any)
for key in self.nodes
}
for key in self.nodes:
outgoing = outgoing_edges[key]
@@ -216,14 +223,15 @@ class Graph:
return CompiledGraph(
graph=self,
nodes=nodes,
channels={**node_outboxes},
input=f"{self.entry_point}:inbox" if self.entry_point else START,
output=END,
hidden=[f"{node}:inbox" for node in self.nodes],
snapshot_channels=list(self.nodes),
checkpointer=checkpointer,
interrupt=(
[f"{node}:inbox" for node in interrupt_before]
+ [node for node in interrupt_after]
),
interrupt_before_nodes=[f"{node}:inbox" for node in interrupt_before],
interrupt_after_nodes=interrupt_after,
debug=debug,
)
+5 -4
View File
@@ -39,6 +39,7 @@ class StateGraph(Graph):
checkpointer: Optional[BaseCheckpointSaver] = None,
interrupt_before: Optional[Sequence[str]] = None,
interrupt_after: Optional[Sequence[str]] = None,
debug: bool = False,
) -> CompiledGraph:
interrupt_before = interrupt_before or []
interrupt_after = interrupt_after or []
@@ -146,11 +147,11 @@ class StateGraph(Graph):
input=f"{START}:inbox",
output=END,
hidden=[f"{node}:inbox" for node in self.nodes] + [START] + state_keys,
snapshot_channels=state_keys_read,
checkpointer=checkpointer,
interrupt=(
[f"{node}:inbox" for node in interrupt_before]
+ [node for node in interrupt_after]
),
interrupt_before_nodes=[f"{node}:inbox" for node in interrupt_before],
interrupt_after_nodes=interrupt_after,
debug=debug,
)
+299 -66
View File
@@ -11,6 +11,7 @@ from typing import (
Callable,
Iterator,
Mapping,
NamedTuple,
Optional,
Sequence,
Type,
@@ -41,6 +42,7 @@ from langchain_core.runnables.utils import (
)
from langchain_core.tracers.log_stream import LogStreamCallbackHandler
from langgraph.channels.any_value import AnyValue
from langgraph.channels.base import (
AsyncChannelsManager,
BaseChannel,
@@ -49,19 +51,21 @@ from langgraph.channels.base import (
InvalidUpdateError,
create_checkpoint,
)
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
Checkpoint,
CheckpointAt,
copy_checkpoint,
empty_checkpoint,
)
from langgraph.constants import CONFIG_KEY_READ, CONFIG_KEY_SEND
from langgraph.constants import CONFIG_KEY_READ, CONFIG_KEY_SEND, INTERRUPT
from langgraph.pregel.debug import print_checkpoint, print_step_start
from langgraph.pregel.io import map_input, map_output
from langgraph.pregel.log import logger
from langgraph.pregel.read import ChannelBatch, ChannelInvoke
from langgraph.pregel.reserved import ReservedChannels
from langgraph.pregel.reserved import AllReservedChannels, ReservedChannels
from langgraph.pregel.validate import validate_graph, validate_keys
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
@@ -158,6 +162,13 @@ class Channel:
)
class StateSnapshot(NamedTuple):
values: dict[str, Any] | Any
"""Current values of channels"""
next: tuple[str]
"""Nodes to execute in the next step, if any"""
class Pregel(
RunnableSerializable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]
):
@@ -165,12 +176,19 @@ class Pregel(
channels: Mapping[str, BaseChannel] = Field(default_factory=dict)
# TODO Rename to `output_channels`
output: Union[str, Sequence[str]] = "output"
# TODO Replace with `stream_channels`
hidden: Sequence[str] = Field(default_factory=list)
interrupt: Sequence[str] = Field(default_factory=list)
snapshot_channels: Union[str, Sequence[str]] = Field(default_factory=list)
interrupt_after_nodes: Sequence[str] = Field(default_factory=list)
interrupt_before_nodes: Sequence[str] = Field(default_factory=list)
# TODO Rename to `input_channels`
input: Union[str, Sequence[str]] = "input"
step_timeout: Optional[float] = None
@@ -192,8 +210,12 @@ class Pregel(
values["input"],
values["output"],
values["hidden"],
values["interrupt"],
values["interrupt_after_nodes"],
values["interrupt_before_nodes"],
)
if values["interrupt_after_nodes"] or values["interrupt_before_nodes"]:
if not values["checkpointer"]:
raise ValueError("Interrupts require a checkpointer")
return values
@property
@@ -247,31 +269,157 @@ class Pregel(
**{k: (self.channels[k].ValueType, None) for k in self.output},
)
@property
def snapshot_channels_list(self) -> Sequence[str]:
return (
[self.snapshot_channels]
if isinstance(self.snapshot_channels, str)
else self.snapshot_channels
or [k for k in self.channels if k not in AllReservedChannels]
)
def get_state(self, config: RunnableConfig) -> StateSnapshot:
if not self.checkpointer:
raise ValueError("No checkpointer set")
checkpoint = self.checkpointer.get(config)
checkpoint = checkpoint or empty_checkpoint()
with ChannelsManager(self.channels, checkpoint) as channels:
_, next_tasks = _prepare_next_tasks(
checkpoint, self.nodes, channels, update_seen=False
)
values = {
k: _read_channel(channels, k)
for k in channels
if k in self.snapshot_channels_list
}
return StateSnapshot(
values[self.snapshot_channels]
if isinstance(self.snapshot_channels, str)
else values,
tuple(name for _, _, name in next_tasks),
)
async def aget_state(self, config: RunnableConfig) -> StateSnapshot:
if not self.checkpointer:
raise ValueError("No checkpointer set")
checkpoint = await self.checkpointer.aget(config)
checkpoint = checkpoint or empty_checkpoint()
async with AsyncChannelsManager(self.channels, checkpoint) as channels:
_, next_tasks = _prepare_next_tasks(
checkpoint, self.nodes, channels, update_seen=False
)
values = {
k: _read_channel(channels, k)
for k in channels
if k in self.snapshot_channels_list
}
return StateSnapshot(
values[self.snapshot_channels]
if isinstance(self.snapshot_channels, str)
else values,
tuple(name for _, _, name in next_tasks),
)
def update_state(
self, config: RunnableConfig, values: dict[str, Any] | Any
) -> None:
if not self.checkpointer:
raise ValueError("No checkpointer set")
values = (
{self.snapshot_channels: values}
if isinstance(self.snapshot_channels, str)
else values
)
checkpoint = self.checkpointer.get(config)
checkpoint = copy_checkpoint(checkpoint) if checkpoint else empty_checkpoint()
with ChannelsManager(self.channels, checkpoint) as channels:
for k, v in values.items():
channels[k].update([v])
checkpoint["channel_versions"][k] += 1
for k in self.snapshot_channels or self.channels:
version = checkpoint["channel_versions"][k]
checkpoint["versions_seen"][INTERRUPT][k] = version
self.checkpointer.put(config, create_checkpoint(checkpoint, channels))
async def aupdate_state(
self, config: RunnableConfig, values: dict[str, Any] | Any
) -> None:
if not self.checkpointer:
raise ValueError("No checkpointer set")
values = (
{self.snapshot_channels: values}
if isinstance(self.snapshot_channels, str)
else values
)
checkpoint = await self.checkpointer.aget(config)
checkpoint = copy_checkpoint(checkpoint) if checkpoint else empty_checkpoint()
async with AsyncChannelsManager(self.channels, checkpoint) as channels:
for k, v in values.items():
channels[k].update([v])
checkpoint["channel_versions"][k] += 1
for k in self.snapshot_channels or self.channels:
version = checkpoint["channel_versions"][k]
checkpoint["versions_seen"][INTERRUPT][k] = version
await self.checkpointer.aput(
config, create_checkpoint(checkpoint, channels)
)
def _defaults(
self,
*,
input_keys: Optional[Union[str, Sequence[str]]] = None,
output_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt_before_nodes: Optional[Sequence[str]] = None,
interrupt_after_nodes: Optional[Sequence[str]] = None,
debug: Optional[bool] = None,
) -> tuple[
bool,
Union[str, Sequence[str]],
Union[str, Sequence[str]],
Optional[Sequence[str]],
Optional[Sequence[str]],
]:
debug = debug if debug is not None else self.debug
if output_keys is None:
output_keys = [chan for chan in self.channels if chan not in self.hidden]
else:
validate_keys(output_keys, self.channels)
if input_keys is None:
input_keys = self.input
else:
validate_keys(input_keys, self.channels)
interrupt_before_nodes = interrupt_before_nodes or self.interrupt_before_nodes
interrupt_after_nodes = interrupt_after_nodes or self.interrupt_after_nodes
return (
debug,
input_keys,
output_keys,
interrupt_before_nodes,
interrupt_after_nodes,
)
def _transform(
self,
input: Iterator[Union[dict[str, Any], Any]],
run_manager: CallbackManagerForChainRun,
config: RunnableConfig,
*,
input_keys: Optional[Union[str, Sequence[str]]] = None,
output_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt: Optional[Sequence[str]] = None,
**kwargs: Any,
) -> Iterator[Union[dict[str, Any], Any]]:
try:
if config["recursion_limit"] < 1:
raise ValueError("recursion_limit must be at least 1")
# assign defaults
if output_keys is None:
output_keys = [
chan for chan in self.channels if chan not in self.hidden
]
else:
validate_keys(output_keys, self.channels)
if input_keys is None:
input_keys = self.input
else:
validate_keys(input_keys, self.channels)
interrupt = interrupt or self.interrupt
(
debug,
input_keys,
output_keys,
interrupt_before_nodes,
interrupt_after_nodes,
) = self._defaults(**kwargs)
# copy nodes to ignore mutations during execution
processes = {**self.nodes}
# get checkpoint from saver, or create an empty one
@@ -286,7 +434,7 @@ class Pregel(
w for c in input for w in map_input(input_keys, c)
):
# discard any unfinished tasks from previous checkpoint
_prepare_next_tasks(checkpoint, processes, channels)
checkpoint, _ = _prepare_next_tasks(checkpoint, processes, channels)
# apply input writes
_apply_writes(
checkpoint,
@@ -304,7 +452,9 @@ class Pregel(
# channels are guaranteed to be immutable for the duration of the step,
# with channel updates applied only at the transition between steps
for step in range(config["recursion_limit"] + 1):
next_tasks = _prepare_next_tasks(checkpoint, processes, channels)
checkpoint, next_tasks = _prepare_next_tasks(
checkpoint, processes, channels
)
# if no more tasks, we're done
if not next_tasks:
@@ -316,7 +466,7 @@ class Pregel(
"by setting the `recursion_limit` config key."
)
if self.debug:
if debug:
print_step_start(step, next_tasks)
# collect all writes to channels, without applying them yet
@@ -354,15 +504,15 @@ class Pregel(
timeout=self.step_timeout,
)
# interrupt on failure or timeout
_interrupt_or_proceed(done, inflight, step)
# panic on failure or timeout
_panic_or_proceed(done, inflight, step)
# apply writes to channels
_apply_writes(
checkpoint, channels, pending_writes, config, step + 1
)
if self.debug:
if debug:
print_checkpoint(step, channels)
# yield current value and checkpoint view
@@ -373,22 +523,37 @@ class Pregel(
# if view was updated, apply writes to channels
_apply_writes_from_view(checkpoint, channels, step_output)
# with previous step's checkpoint
if do_interrupt_before := _should_interrupt(
checkpoint,
interrupt_before_nodes,
self.snapshot_channels_list,
pending_writes,
):
break
# save end of step checkpoint
if (
self.checkpointer is not None
and self.checkpointer.at == CheckpointAt.END_OF_STEP
if self.checkpointer is not None and (
self.checkpointer.at == CheckpointAt.END_OF_STEP
or interrupt_before_nodes
):
checkpoint = create_checkpoint(checkpoint, channels)
self.checkpointer.put(config, checkpoint)
# interrupt if any channel written to is in interrupt list
if any(chan for chan, _ in pending_writes if chan in interrupt):
# with this step's checkpoint,
if _should_interrupt(
checkpoint,
interrupt_after_nodes,
self.snapshot_channels_list,
pending_writes,
):
break
# save end of run checkpoint
if (
self.checkpointer is not None
and self.checkpointer.at == CheckpointAt.END_OF_RUN
and not do_interrupt_before
):
checkpoint = create_checkpoint(checkpoint, channels)
self.checkpointer.put(config, checkpoint)
@@ -405,10 +570,7 @@ class Pregel(
input: AsyncIterator[Union[dict[str, Any], Any]],
run_manager: AsyncCallbackManagerForChainRun,
config: RunnableConfig,
*,
input_keys: Optional[Union[str, Sequence[str]]] = None,
output_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt: Optional[Sequence[str]] = None,
**kwargs: Any,
) -> AsyncIterator[Union[dict[str, Any], Any]]:
try:
if config["recursion_limit"] < 1:
@@ -423,17 +585,13 @@ class Pregel(
None,
)
# assign defaults
if output_keys is None:
output_keys = [
chan for chan in self.channels if chan not in self.hidden
]
else:
validate_keys(output_keys, self.channels)
if input_keys is None:
input_keys = self.input
else:
validate_keys(input_keys, self.channels)
interrupt = interrupt or self.interrupt
(
debug,
input_keys,
output_keys,
interrupt_before_nodes,
interrupt_after_nodes,
) = self._defaults(**kwargs)
# copy nodes to ignore mutations during execution
processes = {**self.nodes}
# get checkpoint from saver, or create an empty one
@@ -448,7 +606,7 @@ class Pregel(
[w async for c in input for w in map_input(input_keys, c)]
):
# discard any unfinished tasks from previous checkpoint
_prepare_next_tasks(checkpoint, processes, channels)
checkpoint, _ = _prepare_next_tasks(checkpoint, processes, channels)
# apply input writes
_apply_writes(
checkpoint,
@@ -466,7 +624,9 @@ class Pregel(
# channels are guaranteed to be immutable for the duration of the step,
# channel updates being applied only at the transition between steps
for step in range(config["recursion_limit"] + 1):
next_tasks = _prepare_next_tasks(checkpoint, processes, channels)
checkpoint, next_tasks = _prepare_next_tasks(
checkpoint, processes, channels
)
# if no more tasks, we're done
if not next_tasks:
@@ -478,7 +638,7 @@ class Pregel(
"by setting the `recursion_limit` config key."
)
if self.debug:
if debug:
print_step_start(step, next_tasks)
# collect all writes to channels, without applying them yet
@@ -523,15 +683,15 @@ class Pregel(
timeout=self.step_timeout,
)
# interrupt on failure or timeout
_interrupt_or_proceed(done, inflight, step)
# panic on failure or timeout
_panic_or_proceed(done, inflight, step)
# apply writes to channels
_apply_writes(
checkpoint, channels, pending_writes, config, step + 1
)
if self.debug:
if debug:
print_checkpoint(step, channels)
# yield current value and checkpoint view
@@ -542,6 +702,15 @@ class Pregel(
# if view was updated, apply writes to channels
_apply_writes_from_view(checkpoint, channels, step_output)
# with previous step's checkpoint
if do_interrupt_before := _should_interrupt(
checkpoint,
interrupt_before_nodes,
self.snapshot_channels_list,
pending_writes,
):
break
# save end of step checkpoint
if (
self.checkpointer is not None
@@ -550,14 +719,20 @@ class Pregel(
checkpoint = create_checkpoint(checkpoint, channels)
await self.checkpointer.aput(config, checkpoint)
# interrupt if any channel written to is in interrupt list
if any(chan for chan, _ in pending_writes if chan in interrupt):
# with this step's checkpoint
if _should_interrupt(
checkpoint,
interrupt_after_nodes,
self.snapshot_channels_list,
pending_writes,
):
break
# save end of run checkpoint
if (
self.checkpointer is not None
and self.checkpointer.at == CheckpointAt.END_OF_RUN
and not do_interrupt_before
):
checkpoint = create_checkpoint(checkpoint, channels)
await self.checkpointer.aput(config, checkpoint)
@@ -576,6 +751,9 @@ class Pregel(
*,
output_keys: Optional[Union[str, Sequence[str]]] = None,
input_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt_before_nodes: Optional[Sequence[str]] = None,
interrupt_after_nodes: Optional[Sequence[str]] = None,
debug: Optional[bool] = None,
**kwargs: Any,
) -> Union[dict[str, Any], Any]:
latest: Union[dict[str, Any], Any] = None
@@ -584,6 +762,9 @@ class Pregel(
config,
output_keys=output_keys if output_keys is not None else self.output,
input_keys=input_keys,
interrupt_before_nodes=interrupt_before_nodes,
interrupt_after_nodes=interrupt_after_nodes,
debug=debug,
**kwargs,
):
latest = chunk
@@ -596,6 +777,9 @@ class Pregel(
*,
output_keys: Optional[Union[str, Sequence[str]]] = None,
input_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt_before_nodes: Optional[Sequence[str]] = None,
interrupt_after_nodes: Optional[Sequence[str]] = None,
debug: Optional[bool] = None,
**kwargs: Any,
) -> Iterator[Union[dict[str, Any], Any]]:
return self.transform(
@@ -603,6 +787,9 @@ class Pregel(
config,
output_keys=output_keys,
input_keys=input_keys,
interrupt_before_nodes=interrupt_before_nodes,
interrupt_after_nodes=interrupt_after_nodes,
debug=debug,
**kwargs,
)
@@ -613,6 +800,9 @@ class Pregel(
*,
output_keys: Optional[Union[str, Sequence[str]]] = None,
input_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt_before_nodes: Optional[Sequence[str]] = None,
interrupt_after_nodes: Optional[Sequence[str]] = None,
debug: Optional[bool] = None,
**kwargs: Any,
) -> Iterator[Union[dict[str, Any], Any]]:
for chunk in self._transform_stream_with_config(
@@ -621,6 +811,9 @@ class Pregel(
config,
output_keys=output_keys,
input_keys=input_keys,
interrupt_before_nodes=interrupt_before_nodes,
interrupt_after_nodes=interrupt_after_nodes,
debug=debug,
**kwargs,
):
yield chunk
@@ -632,6 +825,9 @@ class Pregel(
*,
output_keys: Optional[Union[str, Sequence[str]]] = None,
input_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt_before_nodes: Optional[Sequence[str]] = None,
interrupt_after_nodes: Optional[Sequence[str]] = None,
debug: Optional[bool] = None,
**kwargs: Any,
) -> Union[dict[str, Any], Any]:
latest: Union[dict[str, Any], Any] = None
@@ -640,6 +836,9 @@ class Pregel(
config,
output_keys=output_keys if output_keys is not None else self.output,
input_keys=input_keys,
interrupt_before_nodes=interrupt_before_nodes,
interrupt_after_nodes=interrupt_after_nodes,
debug=debug,
**kwargs,
):
latest = chunk
@@ -652,6 +851,9 @@ class Pregel(
*,
output_keys: Optional[Union[str, Sequence[str]]] = None,
input_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt_before_nodes: Optional[Sequence[str]] = None,
interrupt_after_nodes: Optional[Sequence[str]] = None,
debug: Optional[bool] = None,
**kwargs: Any,
) -> AsyncIterator[Union[dict[str, Any], Any]]:
async def input_stream() -> AsyncIterator[Union[dict[str, Any], Any]]:
@@ -662,6 +864,9 @@ class Pregel(
config,
output_keys=output_keys,
input_keys=input_keys,
interrupt_before_nodes=interrupt_before_nodes,
interrupt_after_nodes=interrupt_after_nodes,
debug=debug,
**kwargs,
):
yield chunk
@@ -673,6 +878,9 @@ class Pregel(
*,
output_keys: Optional[Union[str, Sequence[str]]] = None,
input_keys: Optional[Union[str, Sequence[str]]] = None,
interrupt_before_nodes: Optional[Sequence[str]] = None,
interrupt_after_nodes: Optional[Sequence[str]] = None,
debug: Optional[bool] = None,
**kwargs: Any,
) -> AsyncIterator[Union[dict[str, Any], Any]]:
async for chunk in self._atransform_stream_with_config(
@@ -681,12 +889,15 @@ class Pregel(
config,
output_keys=output_keys,
input_keys=input_keys,
interrupt_before_nodes=interrupt_before_nodes,
interrupt_after_nodes=interrupt_after_nodes,
debug=debug,
**kwargs,
):
yield chunk
def _interrupt_or_proceed(
def _panic_or_proceed(
done: Union[set[concurrent.futures.Future[Any]], set[asyncio.Task[Any]]],
inflight: Union[set[concurrent.futures.Future[Any]], set[asyncio.Task[Any]]],
step: int,
@@ -710,6 +921,24 @@ def _interrupt_or_proceed(
raise TimeoutError(f"Timed out at step {step}")
def _should_interrupt(
checkpoint: Checkpoint,
interrupt_nodes: Sequence[str],
snapshot_channels: Sequence[str],
pending_writes: Sequence[tuple[str, Any]],
) -> bool:
return (
# interrupt if any of snapshopt_channels has been updated since last interrupt
any(
checkpoint["channel_versions"][chan]
> checkpoint["versions_seen"][INTERRUPT][chan]
for chan in snapshot_channels
)
# and any channel written to is in interrupt_nodes list
and any(chan for chan, _ in pending_writes if chan in interrupt_nodes)
)
def _read_channel(
channels: Mapping[str, BaseChannel], chan: str, catch: bool = True
) -> Any:
@@ -732,7 +961,7 @@ def _apply_writes(
pending_writes_by_channel: dict[str, list[Any]] = defaultdict(list)
# Group writes by channel
for chan, val in pending_writes:
if chan in [c.value for c in ReservedChannels]:
if chan in AllReservedChannels:
raise ValueError(f"Can't write to reserved channel {chan}")
pending_writes_by_channel[chan].append(val)
@@ -764,11 +993,12 @@ def _apply_writes(
def _apply_writes_from_view(
checkpoint: Checkpoint, channels: Mapping[str, BaseChannel], values: dict[str, Any]
) -> None:
# Apply writes to channels
for chan, value in values.items():
if value == _read_channel(channels, chan):
continue
assert isinstance(channels[chan], LastValue), (
assert isinstance(channels[chan], (LastValue, EphemeralValue, AnyValue)), (
f"Can't modify channel {chan} of type "
f"{channels[chan].__class__.__name__}"
)
@@ -780,7 +1010,9 @@ def _prepare_next_tasks(
checkpoint: Checkpoint,
processes: Mapping[str, Union[ChannelInvoke, ChannelBatch]],
channels: Mapping[str, BaseChannel],
) -> list[tuple[Runnable, Any, str]]:
update_seen: bool = True,
) -> tuple[Checkpoint, list[tuple[Runnable, Any, str]]]:
checkpoint = copy_checkpoint(checkpoint) if update_seen else checkpoint
tasks: list[tuple[Runnable, Any, str]] = []
# Check if any processes should be run in next step
# If so, prepare the values to be passed to them
@@ -814,12 +1046,13 @@ def _prepare_next_tasks(
val = val[None]
# update seen versions
seen.update(
{
chan: checkpoint["channel_versions"][chan]
for chan in proc.triggers
}
)
if update_seen:
seen.update(
{
chan: checkpoint["channel_versions"][chan]
for chan in proc.triggers
}
)
# skip if condition is not met
if proc.when is None or proc.when(val):
@@ -836,9 +1069,9 @@ def _prepare_next_tasks(
val = [{proc.key: v} for v in val]
tasks.append((proc, val, name))
seen[proc.channel] = checkpoint["channel_versions"][proc.channel]
return tasks
if update_seen:
seen[proc.channel] = checkpoint["channel_versions"][proc.channel]
return checkpoint, tasks
async def _aconsume(iterator: AsyncIterator[Any]) -> None:
+3
View File
@@ -6,3 +6,6 @@ class ReservedChannels(StrEnum):
is_last_step = "is_last_step"
"""A channel that is True if the current step is the last step, False otherwise."""
AllReservedChannels = {channel.value for channel in ReservedChannels}
+8 -3
View File
@@ -2,6 +2,7 @@ from typing import Any, Mapping, Sequence, Union
from langgraph.channels.base import BaseChannel
from langgraph.channels.last_value import LastValue
from langgraph.constants import INTERRUPT
from langgraph.pregel.read import ChannelBatch, ChannelInvoke
from langgraph.pregel.reserved import ReservedChannels
@@ -12,10 +13,13 @@ def validate_graph(
input: Union[str, Sequence[str]],
output: Union[str, Sequence[str]],
hidden: Sequence[str],
interrupt: Sequence[str],
interrupt_after: Sequence[str],
interrupt_before: Sequence[str],
) -> None:
subscribed_channels = set[str]()
for node in nodes.values():
for name, node in nodes.items():
if name == INTERRUPT:
raise ValueError(f"Node name {INTERRUPT} is reserved")
if isinstance(node, ChannelInvoke):
subscribed_channels.update(node.channels.values())
elif isinstance(node, ChannelBatch):
@@ -56,7 +60,8 @@ def validate_graph(
channels[chan] = LastValue(Any) # type: ignore[arg-type]
validate_keys(hidden, channels)
validate_keys(interrupt, channels)
validate_keys(interrupt_after, channels)
validate_keys(interrupt_before, channels)
def validate_keys(
Generated
+5 -6
View File
@@ -1,4 +1,4 @@
# This file is automatically @generated by Poetry 1.7.1 and should not be changed by hand.
# This file is automatically @generated by Poetry 1.6.1 and should not be changed by hand.
[[package]]
name = "aiohttp"
@@ -2520,13 +2520,13 @@ dev = ["pre-commit", "pytest-asyncio", "tox"]
[[package]]
name = "pytest-watcher"
version = "0.3.5"
version = "0.4.1"
description = "Automatically rerun your tests on file modifications"
optional = false
python-versions = ">=3.7.0,<4.0.0"
files = [
{file = "pytest_watcher-0.3.5-py3-none-any.whl", hash = "sha256:af00ca52c7be22dc34c0fd3d7ffef99057207a73b05dc5161fe3b2fe91f58130"},
{file = "pytest_watcher-0.3.5.tar.gz", hash = "sha256:8896152460ba2b1a8200c12117c6611008ec96c8b2d811f0a05ab8a82b043ff8"},
{file = "pytest_watcher-0.4.1-py3-none-any.whl", hash = "sha256:29435669cb0124fb32d6de649fe9b1350f6dac94176313fff559ee4c2a66fd6e"},
{file = "pytest_watcher-0.4.1.tar.gz", hash = "sha256:5a793c4c883e3a55ab2abbfa3a8cd6fa6495b3767d5f6644052cc5f3236f511a"},
]
[package.dependencies]
@@ -2635,7 +2635,6 @@ files = [
{file = "PyYAML-6.0.1-cp311-cp311-win_amd64.whl", hash = "sha256:bf07ee2fef7014951eeb99f56f39c9bb4af143d8aa3c21b1677805985307da34"},
{file = "PyYAML-6.0.1-cp312-cp312-macosx_10_9_x86_64.whl", hash = "sha256:855fb52b0dc35af121542a76b9a84f8d1cd886ea97c84703eaa6d88e37a2ad28"},
{file = "PyYAML-6.0.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:40df9b996c2b73138957fe23a16a4f0ba614f4c0efce1e9406a184b6d07fa3a9"},
{file = "PyYAML-6.0.1-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:a08c6f0fe150303c1c6b71ebcd7213c2858041a7e01975da3a99aed1e7a378ef"},
{file = "PyYAML-6.0.1-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:6c22bec3fbe2524cde73d7ada88f6566758a8f7227bfbf93a408a9d86bcc12a0"},
{file = "PyYAML-6.0.1-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:8d4e9c88387b0f5c7d5f281e55304de64cf7f9c0021a3525bd3b1c542da3b0e4"},
{file = "PyYAML-6.0.1-cp312-cp312-win32.whl", hash = "sha256:d483d2cdf104e7c9fa60c544d92981f12ad66a457afae824d146093b8c294c54"},
@@ -3760,4 +3759,4 @@ testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "p
[metadata]
lock-version = "2.0"
python-versions = ">=3.9.0,<4.0"
content-hash = "0e7777d77d3b34acbfdead224a2b5c5e65ecbf890c57c29bf43f5ebff09d4c0d"
content-hash = "2d35e923bf3902e0e11a305f58d17b0efc3fbb444dff8d6cb92e070a993115c9"
+7 -1
View File
@@ -23,7 +23,7 @@ pytest-asyncio = "^0.20.3"
pytest-mock = "^3.10.0"
syrupy = "^4.0.2"
httpx = "^0.26.0"
pytest-watcher = "^0.3.4"
pytest-watcher = "^0.4.1"
langchain = "^0.1.0"
aiosqlite = "^0.19.0"
grandalf = "^0.8"
@@ -55,6 +55,12 @@ exclude = ["notebooks", "examples", "example_data"]
[tool.coverage.run]
omit = ["tests/*"]
[tool.pytest-watcher]
now = true
delay = 0.1
runner_args = ["-x", "--ff", "-vv", "--snapshot-update"]
patterns = ["*.py"]
[build-system]
requires = ["poetry-core>=1.0.0"]
build-backend = "poetry.core.masonry.api"
+19
View File
@@ -0,0 +1,19 @@
from langchain_core.pydantic_v1 import Field
from langgraph.checkpoint.base import Checkpoint, CheckpointAt, copy_checkpoint
from langgraph.checkpoint.memory import MemorySaver
class MemorySaverAssertImmutable(MemorySaver):
storage_for_copies: dict[str, Checkpoint] = Field(default_factory=dict)
at = CheckpointAt.END_OF_STEP
def put(self, config: dict, checkpoint: dict) -> None:
# assert checkpoint hasn't been modified since last written
thread_id = config["configurable"]["thread_id"]
if saved := super().get(config):
assert self.storage_for_copies[thread_id] == saved
self.storage_for_copies[thread_id] = copy_checkpoint(checkpoint)
# call super to write checkpoint
super().put(config, checkpoint)
+578 -5
View File
@@ -16,7 +16,6 @@ from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.context import Context
from langgraph.channels.last_value import LastValue
from langgraph.channels.topic import Topic
from langgraph.checkpoint.memory import MemorySaver
from langgraph.checkpoint.sqlite import SqliteSaver
from langgraph.graph import END, Graph
from langgraph.graph.message import MessageGraph
@@ -26,8 +25,9 @@ from langgraph.prebuilt.chat_agent_executor import (
create_tool_calling_executor,
)
from langgraph.prebuilt.tool_executor import ToolExecutor
from langgraph.pregel import Channel, GraphRecursionError, Pregel
from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot
from langgraph.pregel.reserved import ReservedChannels
from tests.memory_assert import MemorySaverAssertImmutable
def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
@@ -254,9 +254,11 @@ def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture) -> None:
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
memory = MemorySaver()
memory = MemorySaverAssertImmutable()
app = Pregel(
nodes={"one": one, "two": two}, checkpointer=memory, interrupt=["inbox"]
nodes={"one": one, "two": two},
checkpointer=memory,
interrupt_after_nodes=["inbox"],
)
# start execution, stop at inbox
@@ -282,6 +284,22 @@ def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture) -> None:
assert app.invoke(3, {"configurable": {"thread_id": 1}}) is None
assert app.invoke(None, {"configurable": {"thread_id": 1}}) == 5
# start execution again, stopping at inbox
assert app.invoke(20, {"configurable": {"thread_id": 2}}) is None
# inbox == 21
snapshot = app.get_state({"configurable": {"thread_id": 2}})
assert snapshot.values["inbox"] == 21
assert snapshot.next == ("two",)
# update the state, resume
app.update_state({"configurable": {"thread_id": 2}}, {"inbox": 25})
assert app.invoke(None, {"configurable": {"thread_id": 2}}) == 26
# no pending tasks
snapshot = app.get_state({"configurable": {"thread_id": 2}})
assert snapshot.next == ()
def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
add_one = mocker.Mock(side_effect=lambda x: x + 1)
@@ -431,7 +449,7 @@ def test_invoke_checkpoint(mocker: MockerFixture) -> None:
| raise_if_above_10
)
memory = MemorySaver()
memory = MemorySaverAssertImmutable()
app = Pregel(
nodes={"one": one},
@@ -761,6 +779,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
),
}
# deepcopy because the nodes mutate the data
assert [deepcopy(c) for c in app.stream({"input": "what is weather in sf"})] == [
{
"agent": {
@@ -882,6 +901,297 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
]
# test state get/update methods with interrupt_after
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_after=["agent"]
)
config = {"configurable": {"thread_id": "1"}}
assert [
c for c in app_w_interrupt.stream({"input": "what is weather in sf"}, config)
] == [
{
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
},
"tools": None,
},
next=("agent:edges",),
)
app_w_interrupt.update_state(
config,
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
},
)
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
"tools": None,
},
next=("agent:edges",),
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"tools": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
app_w_interrupt.update_state(
config,
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
# test state get/update methods with interrupt_before
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_before=["tools"]
)
config = {"configurable": {"thread_id": "2"}}
llm.i = 0 # reset the llm
assert [
c for c in app_w_interrupt.stream({"input": "what is weather in sf"}, config)
] == [
{
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
},
"tools": None,
},
next=("agent:edges",),
)
app_w_interrupt.update_state(
config,
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
},
)
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
"tools": None,
},
next=("agent:edges",),
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"tools": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
app_w_interrupt.update_state(
config,
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
def test_conditional_graph_state(snapshot: SnapshotAssertion) -> None:
from langchain.llms.fake import FakeStreamingListLLM
@@ -1073,6 +1383,233 @@ def test_conditional_graph_state(snapshot: SnapshotAssertion) -> None:
},
]
# test state get/update methods with interrupt_after
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_after=["agent"]
)
config = {"configurable": {"thread_id": "1"}}
assert [
c for c in app_w_interrupt.stream({"input": "what is weather in sf"}, config)
] == [
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
app_w_interrupt.update_state(
config,
{
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
},
)
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
app_w_interrupt.update_state(
config,
{
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
# test state get/update methods with interrupt_before
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(),
interrupt_before=["tools"],
debug=True,
)
config = {"configurable": {"thread_id": "2"}}
llm.i = 0 # reset the llm
assert [
c for c in app_w_interrupt.stream({"input": "what is weather in sf"}, config)
] == [
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
app_w_interrupt.update_state(
config,
{
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
},
)
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
app_w_interrupt.update_state(
config,
{
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
def test_conditional_entrypoint_graph(snapshot: SnapshotAssertion) -> None:
def left(data: str) -> str:
@@ -1749,6 +2286,42 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None:
},
]
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_after=["agent"]
)
config = {"configurable": {"thread_id": "1"}}
assert [
c
for c in app_w_interrupt.stream(
HumanMessage(content="what is weather in sf"), config
)
] == [
{
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
)
}
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values=[
HumanMessage(content="what is weather in sf"),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
),
],
next=("agent:edges",),
)
# TODO use update_state once we have message ids
def test_in_one_fan_out_out_one_graph_state() -> None:
def sorted_add(x: list[str], y: list[str]) -> list[str]:
+588 -5
View File
@@ -23,7 +23,6 @@ from langgraph.channels.context import Context
from langgraph.channels.last_value import LastValue
from langgraph.channels.topic import Topic
from langgraph.checkpoint.aiosqlite import AsyncSqliteSaver
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import END, Graph, StateGraph
from langgraph.graph.message import MessageGraph
from langgraph.prebuilt.chat_agent_executor import (
@@ -31,8 +30,9 @@ from langgraph.prebuilt.chat_agent_executor import (
create_tool_calling_executor,
)
from langgraph.prebuilt.tool_executor import ToolExecutor
from langgraph.pregel import Channel, GraphRecursionError, Pregel
from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot
from langgraph.pregel.reserved import ReservedChannels
from tests.memory_assert import MemorySaverAssertImmutable
async def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
@@ -261,9 +261,11 @@ async def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture) -> N
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
memory = MemorySaver()
memory = MemorySaverAssertImmutable()
app = Pregel(
nodes={"one": one, "two": two}, checkpointer=memory, interrupt=["inbox"]
nodes={"one": one, "two": two},
checkpointer=memory,
interrupt_after_nodes=["inbox"],
)
# start execution, stop at inbox
@@ -289,6 +291,22 @@ async def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture) -> N
assert await app.ainvoke(3, {"configurable": {"thread_id": 1}}) is None
assert await app.ainvoke(None, {"configurable": {"thread_id": 1}}) == 5
# start execution again, stopping at inbox
assert await app.ainvoke(20, {"configurable": {"thread_id": 2}}) is None
# inbox == 21
snapshot = await app.aget_state({"configurable": {"thread_id": 2}})
assert snapshot.values["inbox"] == 21
assert snapshot.next == ("two",)
# update the state, resume
await app.aupdate_state({"configurable": {"thread_id": 2}}, {"inbox": 25})
assert await app.ainvoke(None, {"configurable": {"thread_id": 2}}) == 26
# no pending tasks
snapshot = await app.aget_state({"configurable": {"thread_id": 2}})
assert snapshot.next == ()
async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
add_one = mocker.Mock(side_effect=lambda x: x + 1)
@@ -445,7 +463,7 @@ async def test_invoke_checkpoint(mocker: MockerFixture) -> None:
| raise_if_above_10
)
memory = MemorySaver()
memory = MemorySaverAssertImmutable()
app = Pregel(
nodes={"one": one},
@@ -798,6 +816,7 @@ async def test_conditional_graph() -> None:
),
}
# deepcopy because the nodes mutate the data
assert [
deepcopy(c) async for c in app.astream({"input": "what is weather in sf"})
] == [
@@ -927,6 +946,303 @@ async def test_conditional_graph() -> None:
# Check that agent (one of the nodes) has its output streamed to the logs
assert "/logs/agent/streamed_output/-" in patch_paths
# test state get/update methods with interrupt_after
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_after=["agent"]
)
config = {"configurable": {"thread_id": "1"}}
assert [
c
async for c in app_w_interrupt.astream(
{"input": "what is weather in sf"}, config
)
] == [
{
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
},
"tools": None,
},
next=("agent:edges",),
)
await app_w_interrupt.aupdate_state(
config,
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
},
)
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
"tools": None,
},
next=("agent:edges",),
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"tools": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
await app_w_interrupt.aupdate_state(
config,
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
# test state get/update methods with interrupt_before
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_before=["tools"]
)
config = {"configurable": {"thread_id": "2"}}
llm.i = 0
assert [
c
async for c in app_w_interrupt.astream(
{"input": "what is weather in sf"}, config
)
] == [
{
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
},
"tools": None,
},
next=("agent:edges",),
)
await app_w_interrupt.aupdate_state(
config,
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
},
)
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
"tools": None,
},
next=("agent:edges",),
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"tools": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
await app_w_interrupt.aupdate_state(
config,
{
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
async def test_conditional_graph_state() -> None:
from langchain.llms.fake import FakeStreamingListLLM
@@ -1113,6 +1429,237 @@ async def test_conditional_graph_state() -> None:
},
]
# test state get/update methods with interrupt_after
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_after=["agent"]
)
config = {"configurable": {"thread_id": "1"}}
assert [
c
async for c in app_w_interrupt.astream(
{"input": "what is weather in sf"}, config
)
] == [
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
await app_w_interrupt.aupdate_state(
config,
{
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
},
)
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
await app_w_interrupt.aupdate_state(
config,
{
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
# test state get/update methods with interrupt_before
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_before=["tools"]
)
config = {"configurable": {"thread_id": "2"}}
llm.i = 0 # reset the llm
assert [
c
async for c in app_w_interrupt.astream(
{"input": "what is weather in sf"}, config
)
] == [
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
}
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
await app_w_interrupt.aupdate_state(
config,
{
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
},
)
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values={
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"intermediate_steps": [],
},
next=("agent:edges",),
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
},
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
]
await app_w_interrupt.aupdate_state(
config,
{
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{
"__end__": {
"input": "what is weather in sf",
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
}
}
]
async def test_conditional_entrypoint_graph() -> None:
async def left(data: str) -> str:
@@ -1772,6 +2319,42 @@ async def test_message_graph() -> None:
},
]
app_w_interrupt = workflow.compile(
checkpointer=MemorySaverAssertImmutable(), interrupt_after=["agent"]
)
config = {"configurable": {"thread_id": "1"}}
assert [
c
async for c in app_w_interrupt.astream(
HumanMessage(content="what is weather in sf"), config
)
] == [
{
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
)
}
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
values=[
HumanMessage(content="what is weather in sf"),
AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
),
],
next=("agent:edges",),
)
# TODO use update_state once we have message ids
async def test_in_one_fan_out_out_one_graph_state() -> None:
def sorted_add(x: list[str], y: list[str]) -> list[str]: