mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-11 10:45:18 +02:00
Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b4991f1ba3 | ||
|
|
6aa0afba68 | ||
|
|
12aeb0fddb | ||
|
|
d05236f805 | ||
|
|
9d92f33cca | ||
|
|
26356227c4 | ||
|
|
a5dbacae0d |
@@ -14,8 +14,8 @@ async def memory_checkpointer():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_memory_base():
|
||||
"""InMemorySaver passes all base capability tests."""
|
||||
async def test_validate_memory():
|
||||
"""InMemorySaver passes the tests of every capability it implements."""
|
||||
report = await validate(memory_checkpointer)
|
||||
report.print_report()
|
||||
assert report.passed_all_base(), f"Base tests failed: {report.to_dict()}"
|
||||
assert report.passed_all(), f"Capability tests failed: {report.to_dict()}"
|
||||
|
||||
@@ -469,7 +469,7 @@ class PostgresSaver(BasePostgresSaver):
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = self.get_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
|
||||
@@ -417,7 +417,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = await self.aget_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
|
||||
@@ -164,6 +164,40 @@ def test_sync_walk_reads_nothing_newer_than_the_target(
|
||||
assert read == _ids_from_target_down(configs)
|
||||
|
||||
|
||||
def _empty_checkpoint_id(config: dict) -> dict:
|
||||
return {"configurable": {**config["configurable"], "checkpoint_id": ""}}
|
||||
|
||||
|
||||
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
|
||||
await saver.setup()
|
||||
configs = await _abuild_chain(saver)
|
||||
latest = await saver.aget_delta_channel_history(
|
||||
config=configs[-1], channels=[CHANNEL]
|
||||
)
|
||||
|
||||
result = await saver.aget_delta_channel_history(
|
||||
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
|
||||
)
|
||||
|
||||
assert latest[CHANNEL]["writes"]
|
||||
assert result == latest
|
||||
|
||||
|
||||
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
|
||||
saver.setup()
|
||||
configs = _build_chain(saver)
|
||||
latest = saver.get_delta_channel_history(config=configs[-1], channels=[CHANNEL])
|
||||
|
||||
result = saver.get_delta_channel_history(
|
||||
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
|
||||
)
|
||||
|
||||
assert latest[CHANNEL]["writes"]
|
||||
assert result == latest
|
||||
|
||||
|
||||
async def test_root_target_has_no_history_and_still_terminates(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
|
||||
@@ -539,7 +539,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
thread_id = str(config["configurable"]["thread_id"])
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = self.get_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
|
||||
@@ -653,7 +653,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
thread_id = str(config["configurable"]["thread_id"])
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = await self.aget_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
|
||||
@@ -65,6 +65,39 @@ async def test_async_walk_reaches_parent_whatever_the_id_order(
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
|
||||
|
||||
EMPTY_CHECKPOINT_ID: dict[str, Any] = {
|
||||
"configurable": {**CONFIG["configurable"], "checkpoint_id": ""}
|
||||
}
|
||||
|
||||
|
||||
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
root = saver.put(CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {})
|
||||
saver.put_writes(root, [(CHANNEL, "write-root")], "task")
|
||||
saver.put(root, _checkpoint("z-newer", {}), {}, {})
|
||||
|
||||
got = saver.get_delta_channel_history(
|
||||
config=EMPTY_CHECKPOINT_ID, channels=[CHANNEL]
|
||||
)
|
||||
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
|
||||
|
||||
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
|
||||
root = await saver.aput(
|
||||
CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {}
|
||||
)
|
||||
await saver.aput_writes(root, [(CHANNEL, "write-root")], "task")
|
||||
await saver.aput(root, _checkpoint("z-newer", {}), {}, {})
|
||||
|
||||
got = await saver.aget_delta_channel_history(
|
||||
config=EMPTY_CHECKPOINT_ID, channels=[CHANNEL]
|
||||
)
|
||||
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
|
||||
|
||||
def test_walk_reaches_root_of_long_chain_with_descending_ids() -> None:
|
||||
steps = 40
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
|
||||
@@ -170,8 +170,8 @@ class InMemorySaver(
|
||||
return {}
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = config["configurable"].get("checkpoint_id", "")
|
||||
ns_storage = self.storage.get(thread_id, {}).get(checkpoint_ns, {})
|
||||
checkpoint_id = get_checkpoint_id(config) or max(ns_storage, default="")
|
||||
|
||||
chain: list[str] = []
|
||||
target_entry = ns_storage.get(checkpoint_id)
|
||||
|
||||
@@ -1,37 +0,0 @@
|
||||
"""Run delta-channel conformance capabilities against InMemorySaver."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
conformance = pytest.importorskip(
|
||||
"langgraph.checkpoint.conformance",
|
||||
reason="langgraph-checkpoint-conformance not installed",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delta_channel_conformance():
|
||||
# Imported inside the test: the module-level importorskip above is what
|
||||
# makes these safe, so they cannot move to the top of the file.
|
||||
from langgraph.checkpoint.conformance import validate # noqa: PLC0415
|
||||
from langgraph.checkpoint.conformance.initializer import ( # noqa: PLC0415
|
||||
checkpointer_test,
|
||||
)
|
||||
|
||||
from langgraph.checkpoint.memory import InMemorySaver # noqa: PLC0415
|
||||
|
||||
@checkpointer_test(name="InMemorySaver")
|
||||
async def mem_saver():
|
||||
yield InMemorySaver()
|
||||
|
||||
report = await validate(
|
||||
mem_saver,
|
||||
capabilities={
|
||||
"delta_channel_history",
|
||||
},
|
||||
)
|
||||
for cap, result in report.results.items():
|
||||
if result.passed is False:
|
||||
details = "\n".join(result.failures or [])
|
||||
pytest.fail(f"Capability {cap} failed:\n{details}")
|
||||
@@ -420,6 +420,19 @@ class TestInMemorySaverDeltaChannel:
|
||||
assert "seed" not in result
|
||||
assert result["writes"] == []
|
||||
|
||||
def test_get_channel_writes_without_checkpoint_id_reads_the_latest(self) -> None:
|
||||
saver = InMemorySaver()
|
||||
thread: RunnableConfig = {
|
||||
"configurable": {"thread_id": "t1", "checkpoint_ns": ""}
|
||||
}
|
||||
parent = saver.put(thread, empty_checkpoint(), {}, {})
|
||||
saver.put_writes(parent, [("messages", "hi")], "task1")
|
||||
saver.put(parent, empty_checkpoint(), {}, {})
|
||||
|
||||
result = saver.get_delta_channel_history(config=thread, channels=["messages"])
|
||||
|
||||
assert result == {"messages": {"writes": [("task1", "messages", "hi")]}}
|
||||
|
||||
|
||||
class TestBaseFallbackGetChannelWrites:
|
||||
"""Exercises the `BaseCheckpointSaver.get_delta_channel_history` default
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from typing import Any, Literal, cast
|
||||
@@ -32,6 +32,7 @@ from langgraph._internal._constants import (
|
||||
)
|
||||
from langgraph._internal._typing import MISSING
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.channels.binop import _get_overwrite
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec
|
||||
|
||||
@@ -152,6 +153,23 @@ def delta_channels_with_pending_writes(
|
||||
}
|
||||
|
||||
|
||||
def delta_channels_overwritten(
|
||||
specs: Mapping[str, Any], writes: Iterable[tuple[str, Any]]
|
||||
) -> set[str]:
|
||||
"""Return the names of the DeltaChannels that `writes` set with an `Overwrite`.
|
||||
|
||||
`update_state` saves a full snapshot of these channels in the checkpoint it
|
||||
creates, like the loop does when a node returns an `Overwrite`. Otherwise,
|
||||
reading the channel later starts from an older snapshot and replays the
|
||||
writes the `Overwrite` threw away.
|
||||
"""
|
||||
return {
|
||||
ch
|
||||
for ch, value in writes
|
||||
if isinstance(specs.get(ch), DeltaChannel) and _get_overwrite(value)[0]
|
||||
}
|
||||
|
||||
|
||||
def checkpoint_superseded(
|
||||
saver: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
|
||||
) -> bool:
|
||||
@@ -256,6 +274,7 @@ def create_checkpoint(
|
||||
get_next_version: GetNextVersion | None = None,
|
||||
channels_to_snapshot: set[str] | None = None,
|
||||
stored_versions: ChannelVersions | None = None,
|
||||
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
|
||||
) -> Checkpoint:
|
||||
"""Build a new Checkpoint from the previous one and live channel state.
|
||||
|
||||
@@ -314,7 +333,9 @@ def create_checkpoint(
|
||||
id=id or str(uuid6(clock_seq=step)),
|
||||
channel_values=values,
|
||||
channel_versions=channel_versions,
|
||||
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
|
||||
versions_seen=_mark_bumps_seen(
|
||||
checkpoint["versions_seen"], bumped, trigger_to_nodes or {}
|
||||
),
|
||||
updated_channels=None if updated_channels is None else sorted(updated_channels),
|
||||
)
|
||||
|
||||
@@ -322,19 +343,27 @@ def create_checkpoint(
|
||||
def _mark_bumps_seen(
|
||||
versions_seen: dict[str, ChannelVersions],
|
||||
bumped: Mapping[str, tuple[Any, Any]],
|
||||
trigger_to_nodes: Mapping[str, Sequence[str]],
|
||||
) -> dict[str, ChannelVersions]:
|
||||
"""Advance whoever had seen a bumped channel's old version to the new one.
|
||||
|
||||
A bump that only stores a snapshot is not a write. Left unseen, it would
|
||||
re-fire `interrupt_before` and rerun the channel's subscribers. For each
|
||||
entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the one the
|
||||
node really read, so `versions_seen_without_bumps` can put the read back.
|
||||
re-fire `interrupt_before` and rerun the channel's subscribers. A channel
|
||||
bumped from no version was never written, so it also goes to the
|
||||
subscribers that never ran: they have no entry, and would start on the bump.
|
||||
For each entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the
|
||||
one the node really read, so `versions_seen_without_bumps` can put the read
|
||||
back.
|
||||
"""
|
||||
if not bumped:
|
||||
return versions_seen
|
||||
out = dict(versions_seen)
|
||||
for k, (old, _) in bumped.items():
|
||||
if old is None:
|
||||
for node in trigger_to_nodes.get(k, ()):
|
||||
out.setdefault(node, {})
|
||||
marks = dict(versions_seen.get(SNAPSHOT_BUMPS, {}))
|
||||
for node, seen in versions_seen.items():
|
||||
for node, seen in out.items():
|
||||
if node == SNAPSHOT_BUMPS:
|
||||
continue
|
||||
for k, (old, new) in bumped.items():
|
||||
|
||||
@@ -19,6 +19,7 @@ from typing import (
|
||||
TypeVar,
|
||||
cast,
|
||||
)
|
||||
from uuid import UUID, uuid5
|
||||
|
||||
from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
@@ -160,12 +161,6 @@ def DuplexStream(*streams: StreamProtocol) -> StreamProtocol:
|
||||
return StreamProtocol(__call__, {mode for s in streams for mode in s.modes})
|
||||
|
||||
|
||||
def _cacheable(writes: WritesT) -> bool:
|
||||
# A cache hit skips the node, so a cached error would skip it without
|
||||
# raising: only a node that finished goes to the cache.
|
||||
return not any(c in (INTERRUPT, ERROR) for c, _ in writes)
|
||||
|
||||
|
||||
class PregelLoop:
|
||||
config: RunnableConfig
|
||||
store: BaseStore | None
|
||||
@@ -275,6 +270,9 @@ class PregelLoop:
|
||||
# `_put_exit_delta_writes` uses this to decide between anchoring on
|
||||
# the existing parent (True) or creating a lazy stub (False).
|
||||
_has_persisted_parent: bool = False
|
||||
# True iff `__enter__` loaded the thread's latest checkpoint, not one a
|
||||
# `checkpoint_id` addressed, so nothing has been built on it yet.
|
||||
_loaded_latest: bool = False
|
||||
|
||||
managed: ManagedValueMapping
|
||||
checkpoint: Checkpoint
|
||||
@@ -445,9 +443,7 @@ class PregelLoop:
|
||||
return None
|
||||
return self._graph_lifecycle_events.popleft()
|
||||
|
||||
def put_writes(
|
||||
self, task_id: str, writes: WritesT, *, cached: bool = False
|
||||
) -> None:
|
||||
def put_writes(self, task_id: str, writes: WritesT) -> None:
|
||||
"""Put writes for a task, to be read by the next tick."""
|
||||
if not writes:
|
||||
return
|
||||
@@ -540,7 +536,7 @@ class PregelLoop:
|
||||
self._error_handler_write_futs.append(fut)
|
||||
# output writes
|
||||
if hasattr(self, "tasks"):
|
||||
self.output_writes(task_id, writes, cached=cached)
|
||||
self.output_writes(task_id, writes)
|
||||
|
||||
def _put_pending_writes(self) -> None:
|
||||
if self.checkpointer_put_writes is None:
|
||||
@@ -1125,16 +1121,26 @@ class PregelLoop:
|
||||
self._exit_delta_writes.append(
|
||||
(self.step, NULL_TASK_ID, "", c, v)
|
||||
)
|
||||
# Persist delta-channel input writes so sub-freq inputs are
|
||||
# recoverable via ancestor walk (mirrors the Command input path).
|
||||
# A DeltaChannel reads its input from the writes stored on the
|
||||
# checkpoint this run starts from, under a task id of their own:
|
||||
# readers apply a checkpoint's NULL_TASK_ID writes as its own state.
|
||||
# A new thread has no checkpoint to store them on, and one a
|
||||
# `checkpoint_id` addressed may have children that would read them,
|
||||
# so then the input checkpoint snapshots the channel instead.
|
||||
if self.durability != "exit":
|
||||
delta_input = [
|
||||
(c, v)
|
||||
for c, v in input_writes
|
||||
if isinstance(self.specs.get(c), DeltaChannel)
|
||||
]
|
||||
if delta_input:
|
||||
self.put_writes(NULL_TASK_ID, delta_input)
|
||||
if delta_input and self._has_persisted_parent and self._loaded_latest:
|
||||
self.put_writes(
|
||||
str(uuid5(UUID(self.checkpoint["id"]), INPUT)), delta_input
|
||||
)
|
||||
else:
|
||||
self._delta_channels_forced_snapshot.update(
|
||||
c for c, _ in delta_input
|
||||
)
|
||||
# save input checkpoint
|
||||
self.updated_channels = updated_channels
|
||||
self._put_checkpoint({"source": "input"})
|
||||
@@ -1264,6 +1270,7 @@ class PregelLoop:
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=self.checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
for k in channels_to_snapshot:
|
||||
new_counters[k] = (0, 0)
|
||||
@@ -1692,7 +1699,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
) -> PregelExecutableTask | None:
|
||||
if pushed := super().accept_push(task, write_idx, call):
|
||||
for task in self.match_cached_writes():
|
||||
self.put_writes(task.id, task.writes, cached=True)
|
||||
self.output_writes(task.id, task.writes, cached=True)
|
||||
return pushed
|
||||
|
||||
def schedule_error_handler(
|
||||
@@ -1729,18 +1736,16 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
if self._reapplies_pending_writes:
|
||||
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
|
||||
for task in self.match_cached_writes():
|
||||
self.put_writes(task.id, task.writes, cached=True)
|
||||
self.output_writes(task.id, task.writes, cached=True)
|
||||
return handler_task
|
||||
|
||||
def put_writes(
|
||||
self, task_id: str, writes: WritesT, *, cached: bool = False
|
||||
) -> None:
|
||||
def put_writes(self, task_id: str, writes: WritesT) -> None:
|
||||
"""Put writes for a task, to be read by the next tick."""
|
||||
super().put_writes(task_id, writes, cached=cached)
|
||||
if cached or not writes or self.cache is None or not hasattr(self, "tasks"):
|
||||
super().put_writes(task_id, writes)
|
||||
if not writes or self.cache is None or not hasattr(self, "tasks"):
|
||||
return
|
||||
task = self.tasks.get(task_id)
|
||||
if task is None or task.cache_key is None or not _cacheable(writes):
|
||||
if task is None or task.cache_key is None:
|
||||
return
|
||||
self.submit(
|
||||
self.cache.set,
|
||||
@@ -1784,6 +1789,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
# Normal case: fetch the most recent checkpoint for this
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = self.checkpointer.get_tuple(self.checkpoint_config)
|
||||
self._loaded_latest = True
|
||||
|
||||
# Capture before the synthetic-empty fallback below overwrites `saved`.
|
||||
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
|
||||
@@ -1947,7 +1953,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
) -> PregelExecutableTask | None:
|
||||
if pushed := super().accept_push(task, write_idx, call):
|
||||
for task in await self.amatch_cached_writes():
|
||||
self.put_writes(task.id, task.writes, cached=True)
|
||||
self.output_writes(task.id, task.writes, cached=True)
|
||||
return pushed
|
||||
|
||||
async def aschedule_error_handler(
|
||||
@@ -1984,18 +1990,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
if self._reapplies_pending_writes:
|
||||
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
|
||||
for task in await self.amatch_cached_writes():
|
||||
self.put_writes(task.id, task.writes, cached=True)
|
||||
self.output_writes(task.id, task.writes, cached=True)
|
||||
return handler_task
|
||||
|
||||
def put_writes(
|
||||
self, task_id: str, writes: WritesT, *, cached: bool = False
|
||||
) -> None:
|
||||
def put_writes(self, task_id: str, writes: WritesT) -> None:
|
||||
"""Put writes for a task, to be read by the next tick."""
|
||||
super().put_writes(task_id, writes, cached=cached)
|
||||
if cached or not writes or self.cache is None or not hasattr(self, "tasks"):
|
||||
super().put_writes(task_id, writes)
|
||||
if not writes or self.cache is None or not hasattr(self, "tasks"):
|
||||
return
|
||||
task = self.tasks.get(task_id)
|
||||
if task is None or task.cache_key is None or not _cacheable(writes):
|
||||
if task is None or task.cache_key is None:
|
||||
return
|
||||
if writes[0][0] in (INTERRUPT, ERROR):
|
||||
# only cache successful tasks
|
||||
return
|
||||
self.submit(
|
||||
self.cache.aset,
|
||||
@@ -2039,6 +2046,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
# Normal case: fetch the most recent checkpoint for this
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
|
||||
self._loaded_latest = True
|
||||
|
||||
# Capture before the synthetic-empty fallback below overwrites `saved`.
|
||||
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
|
||||
|
||||
@@ -7,7 +7,7 @@ import sys
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from collections.abc import Callable, Sequence
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime, timedelta, timezone
|
||||
@@ -686,6 +686,8 @@ async def arun_with_retry(
|
||||
task: PregelExecutableTask,
|
||||
retry_policy: Sequence[RetryPolicy] | None,
|
||||
stream: bool = False,
|
||||
match_cached_writes: Callable[[], Awaitable[Sequence[PregelExecutableTask]]]
|
||||
| None = None,
|
||||
configurable: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""Run a task asynchronously with retries."""
|
||||
@@ -709,6 +711,11 @@ async def arun_with_retry(
|
||||
)
|
||||
},
|
||||
)
|
||||
if match_cached_writes is not None and task.cache_key is not None:
|
||||
for t in await match_cached_writes():
|
||||
if t is task:
|
||||
# if the task is already cached, return
|
||||
return
|
||||
while True:
|
||||
runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME)
|
||||
if isinstance(runtime, Runtime):
|
||||
|
||||
@@ -137,6 +137,7 @@ from langgraph.pregel._checkpoint import (
|
||||
copy_checkpoint,
|
||||
create_checkpoint,
|
||||
create_checkpoint_plan_for_update_state_api,
|
||||
delta_channels_overwritten,
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
get_updated_channels_from_tasks,
|
||||
@@ -152,6 +153,7 @@ from langgraph.pregel._loop import (
|
||||
from langgraph.pregel._messages import (
|
||||
StreamMessagesHandler,
|
||||
StreamMessagesHandlerV2,
|
||||
ensure_message_ids,
|
||||
)
|
||||
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
|
||||
from langgraph.pregel._retry import RetryPolicy
|
||||
@@ -1766,6 +1768,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
@@ -1788,6 +1791,22 @@ class Pregel(
|
||||
)
|
||||
|
||||
if input_writes := deque(map_input(self.input_channels, values)):
|
||||
_store_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
input_writes,
|
||||
None,
|
||||
)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
updated_channels = apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -1795,6 +1814,9 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, input_writes
|
||||
)
|
||||
|
||||
# apply input write to channels
|
||||
next_step = (
|
||||
@@ -1822,6 +1844,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
@@ -1833,13 +1856,6 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
|
||||
# store the writes
|
||||
checkpointer.put_writes(
|
||||
next_config,
|
||||
input_writes,
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
)
|
||||
|
||||
return patch_checkpoint_map(
|
||||
next_config, saved.metadata if saved else None
|
||||
)
|
||||
@@ -2033,30 +2049,22 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||
# The base's other children replay whatever is stored on it, so an
|
||||
# edit of an older checkpoint stores none of its writes there: the
|
||||
# checkpoint written here carries them, its delta channels
|
||||
# snapshotted. Later supersteps address the checkpoint just written.
|
||||
if (
|
||||
is_first
|
||||
and saved is not None
|
||||
and checkpoint_superseded(checkpointer, config, saved)
|
||||
):
|
||||
fork_pending.update(
|
||||
ch
|
||||
for ch in updated_channels
|
||||
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||
)
|
||||
elif saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config,
|
||||
channel_writes,
|
||||
task_id,
|
||||
**_task_path_kwarg(checkpointer.put_writes, task),
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, (w for t in run_tasks for w in t.writes)
|
||||
)
|
||||
_store_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(task_id, [w for w in task.writes if w[0] != PUSH], task)
|
||||
for task_id, task in zip(run_task_ids, run_tasks)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2086,6 +2094,7 @@ class Pregel(
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
@@ -2258,6 +2267,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
@@ -2280,6 +2290,22 @@ class Pregel(
|
||||
)
|
||||
|
||||
if input_writes := deque(map_input(self.input_channels, values)):
|
||||
await _astore_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
input_writes,
|
||||
None,
|
||||
)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
updated_channels = apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2287,6 +2313,9 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, input_writes
|
||||
)
|
||||
|
||||
# apply input write to channels
|
||||
next_step = (
|
||||
@@ -2314,6 +2343,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
@@ -2325,13 +2355,6 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
|
||||
# store the writes
|
||||
await checkpointer.aput_writes(
|
||||
next_config,
|
||||
input_writes,
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
)
|
||||
|
||||
return patch_checkpoint_map(
|
||||
next_config, saved.metadata if saved else None
|
||||
)
|
||||
@@ -2524,30 +2547,22 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||
# The base's other children replay whatever is stored on it, so an
|
||||
# edit of an older checkpoint stores none of its writes there: the
|
||||
# checkpoint written here carries them, its delta channels
|
||||
# snapshotted. Later supersteps address the checkpoint just written.
|
||||
if (
|
||||
is_first
|
||||
and saved is not None
|
||||
and await acheckpoint_superseded(checkpointer, config, saved)
|
||||
):
|
||||
fork_pending.update(
|
||||
ch
|
||||
for ch in updated_channels
|
||||
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||
)
|
||||
elif saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
await checkpointer.aput_writes(
|
||||
checkpoint_config,
|
||||
channel_writes,
|
||||
task_id,
|
||||
**_task_path_kwarg(checkpointer.aput_writes, task),
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, (w for t in run_tasks for w in t.writes)
|
||||
)
|
||||
await _astore_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(task_id, [w for w in task.writes if w[0] != PUSH], task)
|
||||
for task_id, task in zip(run_task_ids, run_tasks)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2577,6 +2592,7 @@ class Pregel(
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
@@ -3038,7 +3054,7 @@ class Pregel(
|
||||
# with channel updates applied only at the transition between steps.
|
||||
while loop.tick():
|
||||
for task in loop.match_cached_writes():
|
||||
loop.put_writes(task.id, task.writes, cached=True)
|
||||
loop.output_writes(task.id, task.writes, cached=True)
|
||||
for _ in runner.tick(
|
||||
[t for t in loop.tasks.values() if not t.writes],
|
||||
timeout=self.step_timeout,
|
||||
@@ -3509,7 +3525,7 @@ class Pregel(
|
||||
try:
|
||||
while loop.tick():
|
||||
for task in await loop.amatch_cached_writes():
|
||||
loop.put_writes(task.id, task.writes, cached=True)
|
||||
loop.output_writes(task.id, task.writes, cached=True)
|
||||
async for _ in runner.atick(
|
||||
[t for t in loop.tasks.values() if not t.writes],
|
||||
timeout=self.step_timeout,
|
||||
@@ -4255,6 +4271,88 @@ def _task_path_kwarg(put_writes: Callable[..., Any], task: PregelTaskWrites) ->
|
||||
return {"task_path": task_path_str(task.path)}
|
||||
|
||||
|
||||
_UpdateWrites = Sequence[tuple[str, Sequence[tuple[str, Any]], PregelTaskWrites | None]]
|
||||
|
||||
|
||||
def _delta_writes(
|
||||
channels: Mapping[str, BaseChannel | ManagedValueSpec], writes: _UpdateWrites
|
||||
) -> list[tuple[str, Any]]:
|
||||
"""Give the update's DeltaChannel writes message ids, as the loop's
|
||||
`put_writes` does, so every read of them returns the same ids."""
|
||||
delta = [
|
||||
(ch, value)
|
||||
for _, task_writes, _ in writes
|
||||
for ch, value in task_writes
|
||||
if isinstance(channels.get(ch), DeltaChannel)
|
||||
]
|
||||
for _, value in delta:
|
||||
ensure_message_ids(value)
|
||||
return delta
|
||||
|
||||
|
||||
def _store_or_fork_delta_writes(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
checkpoint_config: RunnableConfig,
|
||||
channels: Mapping[str, BaseChannel | ManagedValueSpec],
|
||||
writes: _UpdateWrites,
|
||||
fork_pending: set[str],
|
||||
*,
|
||||
is_first: bool,
|
||||
) -> None:
|
||||
"""Save an update's writes where its DeltaChannels will read them.
|
||||
|
||||
A DeltaChannel rebuilds its value from the writes saved on a checkpoint's
|
||||
ancestors, so the writes go on the checkpoint the update builds on. If the
|
||||
thread already moved past that checkpoint, its other children would read
|
||||
them too, so the new checkpoint snapshots those channels instead. `writes`
|
||||
holds `(task_id, writes, task)` per task; `task` is `None` for input.
|
||||
"""
|
||||
delta = _delta_writes(channels, writes)
|
||||
if saved is None:
|
||||
return
|
||||
if is_first and checkpoint_superseded(checkpointer, config, saved):
|
||||
fork_pending.update(ch for ch, _ in delta)
|
||||
return
|
||||
for task_id, task_writes, task in writes:
|
||||
if task_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config,
|
||||
task_writes,
|
||||
task_id,
|
||||
**(_task_path_kwarg(checkpointer.put_writes, task) if task else {}),
|
||||
)
|
||||
|
||||
|
||||
async def _astore_or_fork_delta_writes(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
checkpoint_config: RunnableConfig,
|
||||
channels: Mapping[str, BaseChannel | ManagedValueSpec],
|
||||
writes: _UpdateWrites,
|
||||
fork_pending: set[str],
|
||||
*,
|
||||
is_first: bool,
|
||||
) -> None:
|
||||
"""Async `_store_or_fork_delta_writes`."""
|
||||
delta = _delta_writes(channels, writes)
|
||||
if saved is None:
|
||||
return
|
||||
if is_first and await acheckpoint_superseded(checkpointer, config, saved):
|
||||
fork_pending.update(ch for ch, _ in delta)
|
||||
return
|
||||
for task_id, task_writes, task in writes:
|
||||
if task_writes:
|
||||
await checkpointer.aput_writes(
|
||||
checkpoint_config,
|
||||
task_writes,
|
||||
task_id,
|
||||
**(_task_path_kwarg(checkpointer.aput_writes, task) if task else {}),
|
||||
)
|
||||
|
||||
|
||||
def _update_task_id(checkpoint_id: str, i: int) -> str:
|
||||
"""Task id for the `i`th update of a superstep that has no task to reuse.
|
||||
|
||||
|
||||
@@ -1,122 +0,0 @@
|
||||
"""A node served from the node cache saves its writes like a node that ran,
|
||||
without writing them back to the cache, and a node that fails isn't cached."""
|
||||
|
||||
import operator
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langgraph.cache.memory import InMemoryCache
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.types import CachePolicy, Durability
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
INPUT = {"log": [], "plain": []}
|
||||
|
||||
|
||||
def _append(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
log: Annotated[list, DeltaChannel(_append)]
|
||||
plain: Annotated[list, operator.add]
|
||||
|
||||
|
||||
class _CountsSets(InMemoryCache):
|
||||
sets = 0
|
||||
|
||||
def set(self, keys: Any) -> None:
|
||||
self.sets += 1
|
||||
super().set(keys)
|
||||
|
||||
|
||||
def _a_then_cached_b_then_c(runs: list[str], cache: InMemoryCache) -> Any:
|
||||
def node(name: str) -> Any:
|
||||
def run(state: _State) -> dict:
|
||||
runs.append(name)
|
||||
return {"log": [name], "plain": [name]}
|
||||
|
||||
return run
|
||||
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("a", node("a"))
|
||||
builder.add_node("b", node("b"), cache_policy=CachePolicy())
|
||||
builder.add_node("c", node("c"))
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("b", "c")
|
||||
return builder.compile(checkpointer=InMemorySaver(), cache=cache)
|
||||
|
||||
|
||||
def test_a_cache_hit_saves_its_writes_without_caching_them_again(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
runs: list[str] = []
|
||||
cache = _CountsSets()
|
||||
graph = _a_then_cached_b_then_c(runs, cache)
|
||||
graph.invoke(INPUT, {"configurable": {"thread_id": "1"}}, durability=durability)
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
|
||||
graph.invoke(INPUT, config, durability=durability)
|
||||
|
||||
assert runs == ["a", "b", "c", "a", "c"]
|
||||
assert cache.sets == 1, "the cache hit was written back to the cache"
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
async def test_a_cache_hit_saves_its_writes_without_caching_them_again_async(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
runs: list[str] = []
|
||||
cache = _CountsSets()
|
||||
graph = _a_then_cached_b_then_c(runs, cache)
|
||||
await graph.ainvoke(
|
||||
INPUT, {"configurable": {"thread_id": "1"}}, durability=durability
|
||||
)
|
||||
config = {"configurable": {"thread_id": "2"}}
|
||||
|
||||
await graph.ainvoke(INPUT, config, durability=durability)
|
||||
|
||||
assert runs == ["a", "b", "c", "a", "c"]
|
||||
assert cache.sets == 1, "the cache hit was written back to the cache"
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
def _cached_node_that_fails(runs: list[str]) -> Any:
|
||||
def fail(state: _State) -> dict:
|
||||
runs.append("b")
|
||||
raise ValueError("b failed")
|
||||
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("b", fail, cache_policy=CachePolicy())
|
||||
builder.add_edge(START, "b")
|
||||
return builder.compile(checkpointer=InMemorySaver(), cache=InMemoryCache())
|
||||
|
||||
|
||||
def test_a_node_that_fails_is_not_cached() -> None:
|
||||
runs: list[str] = []
|
||||
graph = _cached_node_that_fails(runs)
|
||||
|
||||
for thread in ("1", "2"):
|
||||
with pytest.raises(ValueError, match="b failed"):
|
||||
graph.invoke(INPUT, {"configurable": {"thread_id": thread}})
|
||||
|
||||
assert runs == ["b", "b"]
|
||||
|
||||
|
||||
async def test_a_node_that_fails_is_not_cached_async() -> None:
|
||||
runs: list[str] = []
|
||||
graph = _cached_node_that_fails(runs)
|
||||
|
||||
for thread in ("1", "2"):
|
||||
with pytest.raises(ValueError, match="b failed"):
|
||||
await graph.ainvoke(INPUT, {"configurable": {"thread_id": thread}})
|
||||
|
||||
assert runs == ["b", "b"]
|
||||
@@ -0,0 +1,103 @@
|
||||
"""A run's input to a DeltaChannel input channel reads back on the checkpoints
|
||||
built from it, and on no others."""
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
from langgraph.types import Durability
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _sorted_extend(current: list, writes: list) -> list:
|
||||
return sorted([*current, *(item for write in writes for item in write)])
|
||||
|
||||
|
||||
def _delta_input_graph() -> Pregel:
|
||||
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={
|
||||
"log": DeltaChannel(_sorted_extend),
|
||||
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
|
||||
"go": LastValue(int),
|
||||
},
|
||||
input_channels=["log", "plain", "go"],
|
||||
output_channels=["log", "plain"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
def test_each_run_input_reads_back_on_its_own_checkpoints(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
graph.invoke({"log": [0], "plain": [0], "go": 1}, config, durability=durability)
|
||||
graph.invoke({"log": [5], "plain": [5], "go": 1}, config, durability=durability)
|
||||
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
async def test_each_run_input_reads_back_on_its_own_checkpoints_async(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
await graph.ainvoke(
|
||||
{"log": [0], "plain": [0], "go": 1}, config, durability=durability
|
||||
)
|
||||
await graph.ainvoke(
|
||||
{"log": [5], "plain": [5], "go": 1}, config, durability=durability
|
||||
)
|
||||
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
OTHER_BRANCH_INPUTS = pytest.mark.parametrize(
|
||||
"other_branch_input",
|
||||
[{"go": 1}, {"log": [5], "plain": [5], "go": 1}],
|
||||
ids=["other-branch-without-delta-input", "other-branch-with-delta-input"],
|
||||
)
|
||||
|
||||
|
||||
@OTHER_BRANCH_INPUTS
|
||||
def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch(
|
||||
durability: Durability, other_branch_input: dict
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"go": 1}, config, durability=durability)
|
||||
older = graph.get_state(config).config
|
||||
graph.invoke(other_branch_input, config, durability=durability)
|
||||
|
||||
graph.invoke({"log": [7], "plain": [7], "go": 1}, older, durability=durability)
|
||||
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
@OTHER_BRANCH_INPUTS
|
||||
async def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch_async(
|
||||
durability: Durability, other_branch_input: dict
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"go": 1}, config, durability=durability)
|
||||
older = (await graph.aget_state(config)).config
|
||||
await graph.ainvoke(other_branch_input, config, durability=durability)
|
||||
|
||||
await graph.ainvoke(
|
||||
{"log": [7], "plain": [7], "go": 1}, older, durability=durability
|
||||
)
|
||||
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
@@ -0,0 +1,78 @@
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
CONFIG = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
|
||||
def _extend(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
def _graph(reads: list) -> Pregel:
|
||||
writer = NodeBuilder().subscribe_only("a").do(lambda _: [1]).write_to("d")
|
||||
reader = NodeBuilder().subscribe_only("d").do(lambda d: reads.append(list(d)))
|
||||
return Pregel(
|
||||
nodes={"writer": writer, "reader": reader},
|
||||
channels={"a": LastValue(str), "d": DeltaChannel(_extend)},
|
||||
input_channels="a",
|
||||
output_channels=["d"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
def test_an_input_update_from_before_the_first_write_starts_only_the_writer() -> None:
|
||||
reads: list = []
|
||||
graph = _graph(reads)
|
||||
graph.invoke("go", CONFIG)
|
||||
first = next(
|
||||
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
|
||||
)
|
||||
|
||||
fork = graph.update_state(first, {"a": "go"}, as_node="__input__")
|
||||
|
||||
assert graph.get_state(fork).next == ("writer",)
|
||||
graph.invoke(None, fork)
|
||||
assert reads == [[1], [1]]
|
||||
|
||||
|
||||
def test_a_replay_from_before_the_first_write_forks_with_only_the_writer_next() -> None:
|
||||
reads: list = []
|
||||
graph = _graph(reads)
|
||||
graph.invoke("go", CONFIG)
|
||||
first = next(
|
||||
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
|
||||
)
|
||||
|
||||
graph.invoke(None, first, durability="sync")
|
||||
|
||||
fork = next(
|
||||
s for s in graph.get_state_history(CONFIG) if s.metadata["source"] == "fork"
|
||||
)
|
||||
assert fork.next == ("writer",)
|
||||
graph.invoke(None, fork.config)
|
||||
assert reads == [[1], [1], [1]]
|
||||
|
||||
|
||||
async def test_an_ainput_update_from_before_the_first_write_starts_only_the_writer() -> (
|
||||
None
|
||||
):
|
||||
reads: list = []
|
||||
graph = _graph(reads)
|
||||
await graph.ainvoke("go", CONFIG)
|
||||
first = [
|
||||
s.config
|
||||
async for s in graph.aget_state_history(CONFIG)
|
||||
if s.metadata["step"] == -1
|
||||
][0]
|
||||
|
||||
fork = await graph.aupdate_state(first, {"a": "go"}, as_node="__input__")
|
||||
|
||||
assert (await graph.aget_state(fork)).next == ("writer",)
|
||||
await graph.ainvoke(None, fork)
|
||||
assert reads == [[1], [1]]
|
||||
@@ -0,0 +1,110 @@
|
||||
"""An `Overwrite` through `update_state` snapshots its DeltaChannel on the
|
||||
checkpoint the update saves, as a node's `Overwrite` does on the loop's."""
|
||||
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
from langgraph.types import Overwrite
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
messages: Annotated[list, DeltaChannel(_messages_delta_reducer)]
|
||||
|
||||
|
||||
def _messages_graph(saver: InMemorySaver) -> Any:
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("model", lambda state: {})
|
||||
builder.add_edge(START, "model")
|
||||
return builder.compile(checkpointer=saver)
|
||||
|
||||
|
||||
def _extend(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
def _delta_input_graph(saver: InMemorySaver) -> Pregel:
|
||||
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log")
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={"log": DeltaChannel(_extend), "go": LastValue(int)},
|
||||
input_channels=["log", "go"],
|
||||
output_channels=["log"],
|
||||
checkpointer=saver,
|
||||
)
|
||||
|
||||
|
||||
def test_update_state_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _messages_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="a", id="1")]}, config)
|
||||
|
||||
graph.update_state(
|
||||
config,
|
||||
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
|
||||
as_node="model",
|
||||
)
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
|
||||
assert [m.content for m in graph.get_state(config).values["messages"]] == ["b"]
|
||||
|
||||
|
||||
async def test_aupdate_state_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _messages_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="a", id="1")]}, config)
|
||||
|
||||
await graph.aupdate_state(
|
||||
config,
|
||||
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
|
||||
as_node="model",
|
||||
)
|
||||
|
||||
head = await saver.aget_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
|
||||
values = (await graph.aget_state(config)).values
|
||||
assert [m.content for m in values["messages"]] == ["b"]
|
||||
|
||||
|
||||
def test_update_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _delta_input_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"log": [0], "go": 1}, config)
|
||||
|
||||
graph.update_state(config, {"log": Overwrite([1])}, as_node="__input__")
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
|
||||
assert graph.get_state(config).values["log"] == [1]
|
||||
|
||||
|
||||
async def test_aupdate_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _delta_input_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"log": [0], "go": 1}, config)
|
||||
|
||||
await graph.aupdate_state(config, {"log": Overwrite([1])}, as_node="__input__")
|
||||
|
||||
head = await saver.aget_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
|
||||
assert (await graph.aget_state(config)).values["log"] == [1]
|
||||
@@ -25,9 +25,12 @@ from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
from langgraph.types import StateSnapshot, StateUpdate
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
@@ -538,3 +541,189 @@ def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
|
||||
|
||||
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
|
||||
assert graph.get_state(config).next == ()
|
||||
|
||||
|
||||
def _sorted_extend(current: list, writes: list) -> list:
|
||||
return sorted([*current, *(item for write in writes for item in write)])
|
||||
|
||||
|
||||
def _delta_input_graph(snapshot_frequency: int = 1000) -> Any:
|
||||
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={
|
||||
"log": DeltaChannel(_sorted_extend, snapshot_frequency=snapshot_frequency),
|
||||
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
|
||||
"go": LastValue(int),
|
||||
},
|
||||
input_channels=["log", "plain", "go"],
|
||||
output_channels=["log", "plain"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
|
||||
def test_update_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
|
||||
snapshot_frequency: int,
|
||||
) -> None:
|
||||
graph = _delta_input_graph(snapshot_frequency)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"log": [0], "plain": [0], "go": 1}, config)
|
||||
|
||||
graph.update_state(config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__")
|
||||
after_update = graph.get_state(config).values
|
||||
graph.invoke(None, config)
|
||||
after_run = graph.get_state(config).values
|
||||
|
||||
assert after_update["log"] == after_update["plain"]
|
||||
assert after_run["log"] == after_run["plain"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
|
||||
async def test_aupdate_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
|
||||
snapshot_frequency: int,
|
||||
) -> None:
|
||||
graph = _delta_input_graph(snapshot_frequency)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"log": [0], "plain": [0], "go": 1}, config)
|
||||
|
||||
await graph.aupdate_state(
|
||||
config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
|
||||
)
|
||||
after_update = (await graph.aget_state(config)).values
|
||||
await graph.ainvoke(None, config)
|
||||
after_run = (await graph.aget_state(config)).values
|
||||
|
||||
assert after_update["log"] == after_update["plain"]
|
||||
assert after_run["log"] == after_run["plain"]
|
||||
|
||||
|
||||
def test_update_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"go": 1}, config)
|
||||
older = graph.get_state(config).config
|
||||
graph.invoke({"go": 1}, config)
|
||||
other_branch = graph.get_state(config)
|
||||
|
||||
edited = graph.update_state(
|
||||
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
|
||||
)
|
||||
|
||||
values = graph.get_state(edited).values
|
||||
assert values["log"] == values["plain"]
|
||||
assert graph.get_state(other_branch.config).values == other_branch.values
|
||||
|
||||
|
||||
async def test_aupdate_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> (
|
||||
None
|
||||
):
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"go": 1}, config)
|
||||
older = (await graph.aget_state(config)).config
|
||||
await graph.ainvoke({"go": 1}, config)
|
||||
other_branch = await graph.aget_state(config)
|
||||
|
||||
edited = await graph.aupdate_state(
|
||||
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
|
||||
)
|
||||
|
||||
values = (await graph.aget_state(edited)).values
|
||||
assert values["log"] == values["plain"]
|
||||
assert (await graph.aget_state(other_branch.config)).values == other_branch.values
|
||||
|
||||
|
||||
def _message_ids(graph: Any, config: dict) -> list[str | None]:
|
||||
return [m.id for m in graph.get_state(config).values["messages"]]
|
||||
|
||||
|
||||
async def _amessage_ids(graph: Any, config: dict) -> list[str | None]:
|
||||
return [m.id for m in (await graph.aget_state(config)).values["messages"]]
|
||||
|
||||
|
||||
def _messages_input_graph() -> Any:
|
||||
node = (
|
||||
NodeBuilder()
|
||||
.subscribe_only("go")
|
||||
.do(lambda _: [HumanMessage("n", id="n")])
|
||||
.write_to("messages")
|
||||
)
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={
|
||||
"messages": DeltaChannel(_messages_delta_reducer),
|
||||
"go": LastValue(int),
|
||||
},
|
||||
input_channels=["messages", "go"],
|
||||
output_channels=["messages"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
def test_update_state_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _build_graph(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
|
||||
|
||||
graph.update_state(config, {"messages": [HumanMessage("b")]})
|
||||
|
||||
first, second = _message_ids(graph, config), _message_ids(graph, config)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
|
||||
async def test_aupdate_state_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _build_graph(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage("a", id="a")]}, config)
|
||||
|
||||
await graph.aupdate_state(config, {"messages": [HumanMessage("b")]})
|
||||
|
||||
first, second = (
|
||||
await _amessage_ids(graph, config),
|
||||
await _amessage_ids(graph, config),
|
||||
)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
|
||||
def test_update_state_on_an_older_checkpoint_gives_a_message_an_id() -> None:
|
||||
graph = _build_graph(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
|
||||
older = graph.get_state(config).config
|
||||
graph.invoke({"messages": [HumanMessage("c", id="c")]}, config)
|
||||
|
||||
branch = graph.update_state(older, {"messages": [HumanMessage("b")]})
|
||||
|
||||
assert _message_ids(graph, branch)[-1] is not None
|
||||
|
||||
|
||||
def test_update_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _messages_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
|
||||
|
||||
graph.update_state(config, {"messages": [HumanMessage("b")]}, as_node="__input__")
|
||||
|
||||
first, second = _message_ids(graph, config), _message_ids(graph, config)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
|
||||
async def test_aupdate_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _messages_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
|
||||
|
||||
await graph.aupdate_state(
|
||||
config, {"messages": [HumanMessage("b")]}, as_node="__input__"
|
||||
)
|
||||
|
||||
first, second = (
|
||||
await _amessage_ids(graph, config),
|
||||
await _amessage_ids(graph, config),
|
||||
)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
@@ -5932,9 +5932,9 @@ def test_no_redundant_put_writes_for_cached_task(
|
||||
put_writes_task_ids: list[str] = []
|
||||
orig = PregelLoop.put_writes
|
||||
|
||||
def spy(self, task_id, writes, **kwargs):
|
||||
def spy(self, task_id, writes):
|
||||
put_writes_task_ids.append(task_id)
|
||||
return orig(self, task_id, writes, **kwargs)
|
||||
return orig(self, task_id, writes)
|
||||
|
||||
with patch.object(PregelLoop, "put_writes", spy):
|
||||
result = workflow.invoke(Command(resume="ans"), config=config)
|
||||
|
||||
@@ -8157,9 +8157,9 @@ async def test_no_redundant_put_writes_for_cached_task(
|
||||
put_writes_task_ids: list[str] = []
|
||||
orig = PregelLoop.put_writes
|
||||
|
||||
def spy(self, task_id, writes, **kwargs):
|
||||
def spy(self, task_id, writes):
|
||||
put_writes_task_ids.append(task_id)
|
||||
return orig(self, task_id, writes, **kwargs)
|
||||
return orig(self, task_id, writes)
|
||||
|
||||
with patch.object(PregelLoop, "put_writes", spy):
|
||||
result = await workflow.ainvoke(Command(resume="ans"), config=config)
|
||||
|
||||
Reference in New Issue
Block a user