mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-06 00:15:09 +02:00
Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
03ef1a26e0 | ||
|
|
dc0d992b90 | ||
|
|
ed168deb97 | ||
|
|
fb6e5c2bce | ||
|
|
6dade64aa7 | ||
|
|
398d6cc59d | ||
|
|
d736564eb1 | ||
|
|
69f2d3a430 |
@@ -63,6 +63,7 @@ The suite tests **base** capabilities (required) and **extended** capabilities (
|
||||
| `delete_for_runs` | no | `adelete_for_runs` |
|
||||
| `copy_thread` | no | `acopy_thread` |
|
||||
| `prune` | no | `aprune` |
|
||||
| `delta_channel_history` | no | `aget_delta_channel_history` |
|
||||
|
||||
Extended capabilities are detected by checking whether the method is overridden from `BaseCheckpointSaver`. If not overridden, those tests are skipped.
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ class Capability(str, Enum):
|
||||
DELETE_FOR_RUNS = "delete_for_runs"
|
||||
COPY_THREAD = "copy_thread"
|
||||
PRUNE = "prune"
|
||||
DELTA_CHANNEL_HISTORY = "delta_channel_history"
|
||||
|
||||
|
||||
# Capabilities that every checkpointer must support.
|
||||
@@ -42,6 +43,7 @@ EXTENDED_CAPABILITIES = frozenset(
|
||||
Capability.DELETE_FOR_RUNS,
|
||||
Capability.COPY_THREAD,
|
||||
Capability.PRUNE,
|
||||
Capability.DELTA_CHANNEL_HISTORY,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -57,6 +59,7 @@ _CAPABILITY_METHOD_MAP: dict[Capability, str] = {
|
||||
Capability.DELETE_FOR_RUNS: "adelete_for_runs",
|
||||
Capability.COPY_THREAD: "acopy_thread",
|
||||
Capability.PRUNE: "aprune",
|
||||
Capability.DELTA_CHANNEL_HISTORY: "aget_delta_channel_history",
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -9,6 +9,9 @@ from langgraph.checkpoint.conformance.spec.test_delete_for_runs import (
|
||||
from langgraph.checkpoint.conformance.spec.test_delete_thread import (
|
||||
run_delete_thread_tests,
|
||||
)
|
||||
from langgraph.checkpoint.conformance.spec.test_delta_channel_history import (
|
||||
run_delta_channel_history_tests,
|
||||
)
|
||||
from langgraph.checkpoint.conformance.spec.test_get_tuple import run_get_tuple_tests
|
||||
from langgraph.checkpoint.conformance.spec.test_list import run_list_tests
|
||||
from langgraph.checkpoint.conformance.spec.test_prune import run_prune_tests
|
||||
@@ -24,4 +27,5 @@ __all__ = [
|
||||
"run_delete_for_runs_tests",
|
||||
"run_copy_thread_tests",
|
||||
"run_prune_tests",
|
||||
"run_delta_channel_history_tests",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
"""Shared fixtures for delta-channel conformance tests.
|
||||
|
||||
Builds a parent chain with `_DeltaSnapshot` blobs at known positions via
|
||||
direct `aput` / `aput_writes` calls. No langgraph or Pregel dependency.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver, Checkpoint
|
||||
from langgraph.checkpoint.base.id import uuid6
|
||||
|
||||
from langgraph.checkpoint.conformance.test_utils import generate_metadata
|
||||
|
||||
|
||||
async def build_delta_chain(
|
||||
saver: BaseCheckpointSaver,
|
||||
*,
|
||||
thread_id: str | None = None,
|
||||
checkpoint_ns: str = "",
|
||||
channel: str = "messages",
|
||||
snapshots_at_steps: Sequence[int] = (0,),
|
||||
total_steps: int = 6,
|
||||
write_value_fn: Any | None = None,
|
||||
) -> list[RunnableConfig]:
|
||||
"""Build a parent chain with `_DeltaSnapshot` at known positions.
|
||||
|
||||
Args:
|
||||
saver: Checkpointer instance.
|
||||
thread_id: Defaults to a random UUID.
|
||||
checkpoint_ns: Namespace (default root).
|
||||
channel: Channel name used for snapshots and writes.
|
||||
snapshots_at_steps: Steps at which a `_DeltaSnapshot` blob is stored
|
||||
in `channel_values[channel]`. Step 0 is the oldest checkpoint.
|
||||
total_steps: Number of checkpoints in the chain.
|
||||
write_value_fn: Callable(step) -> write value. Defaults to step index.
|
||||
|
||||
Returns:
|
||||
List of stored configs (oldest first), one per step.
|
||||
"""
|
||||
if write_value_fn is None:
|
||||
|
||||
def write_value_fn(step: int) -> Any:
|
||||
return step
|
||||
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
thread_id = thread_id or str(uuid4())
|
||||
snapshot_set = set(snapshots_at_steps)
|
||||
stored: list[RunnableConfig] = []
|
||||
parent_cfg: RunnableConfig | None = None
|
||||
|
||||
for step in range(total_steps):
|
||||
config: RunnableConfig = {
|
||||
"configurable": {
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
}
|
||||
}
|
||||
if parent_cfg:
|
||||
config["configurable"]["checkpoint_id"] = parent_cfg["configurable"][
|
||||
"checkpoint_id"
|
||||
]
|
||||
|
||||
channel_values: dict[str, Any] = {}
|
||||
channel_versions: dict[str, int] = {}
|
||||
if step in snapshot_set:
|
||||
channel_values[channel] = _DeltaSnapshot(
|
||||
write_value_fn(step),
|
||||
)
|
||||
channel_versions[channel] = step + 1
|
||||
|
||||
cp = Checkpoint(
|
||||
v=1,
|
||||
id=str(uuid6(clock_seq=-1)),
|
||||
ts="",
|
||||
channel_values=channel_values,
|
||||
channel_versions=channel_versions,
|
||||
versions_seen={},
|
||||
updated_channels=None,
|
||||
)
|
||||
new_versions = dict(channel_versions)
|
||||
parent_cfg = await saver.aput(
|
||||
config, cp, generate_metadata(step=step), new_versions
|
||||
)
|
||||
stored.append(parent_cfg)
|
||||
|
||||
# Write a pending write for non-snapshot steps so the walk has
|
||||
# something to collect.
|
||||
if step not in snapshot_set:
|
||||
await saver.aput_writes(
|
||||
parent_cfg, [(channel, write_value_fn(step))], str(uuid4())
|
||||
)
|
||||
|
||||
return stored
|
||||
+247
@@ -0,0 +1,247 @@
|
||||
"""DELTA_CHANNEL_HISTORY capability tests — aget_delta_channel_history contract."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from uuid import uuid4
|
||||
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
|
||||
from langgraph.checkpoint.conformance.spec._delta_fixtures import build_delta_chain
|
||||
|
||||
|
||||
async def test_history_returns_writes_oldest_first(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Writes are returned oldest-to-newest."""
|
||||
tid = str(uuid4())
|
||||
# 5 steps: snapshot at 0, writes at 1,2,3,4.
|
||||
# Head is step 4. Walk starts at step 3 (parent of head).
|
||||
# Collects writes from steps 1,2,3 (between snapshot at 0 and head's parent).
|
||||
configs = await build_delta_chain(
|
||||
saver, thread_id=tid, channel="ch", snapshots_at_steps=[0], total_steps=5
|
||||
)
|
||||
head = configs[-1]
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
||||
writes = result["ch"]["writes"]
|
||||
values = [w[2] for w in writes]
|
||||
assert values == [1, 2, 3], f"Expected [1,2,3], got {values}"
|
||||
|
||||
|
||||
async def test_history_seed_is_nearest_snapshot(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Seed is the value from the nearest ancestor with channel_values populated."""
|
||||
tid = str(uuid4())
|
||||
# 6 steps: snapshots at 0 and 3, writes at 1,2,4,5.
|
||||
# Head is step 5. Walk from step 4 backward stops at step 3 (snapshot).
|
||||
# Collects writes from step 4 only (between step 3 and head's parent step 4).
|
||||
configs = await build_delta_chain(
|
||||
saver,
|
||||
thread_id=tid,
|
||||
channel="ch",
|
||||
snapshots_at_steps=[0, 3],
|
||||
total_steps=6,
|
||||
)
|
||||
head = configs[-1]
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
||||
assert "seed" in result["ch"], "Expected seed from snapshot at step 3"
|
||||
seed = result["ch"]["seed"]
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
actual_value = seed.value if isinstance(seed, _DeltaSnapshot) else seed
|
||||
assert actual_value == 3, f"Expected seed value 3 (step 3), got {actual_value}"
|
||||
writes = result["ch"]["writes"]
|
||||
values = [w[2] for w in writes]
|
||||
assert values == [4], f"Expected [4], got {values}"
|
||||
|
||||
|
||||
async def test_history_excludes_target_pending_writes(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Target's own pending_writes are NOT included in the history."""
|
||||
tid = str(uuid4())
|
||||
configs = await build_delta_chain(
|
||||
saver, thread_id=tid, channel="ch", snapshots_at_steps=[0], total_steps=3
|
||||
)
|
||||
head = configs[-1]
|
||||
# Add writes directly to the head checkpoint
|
||||
await saver.aput_writes(head, [("ch", "extra")], str(uuid4()))
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
||||
writes = result["ch"]["writes"]
|
||||
values = [w[2] for w in writes]
|
||||
assert "extra" not in values, f"Target's writes should be excluded, got {values}"
|
||||
|
||||
|
||||
async def test_history_multi_channel(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Multiple channels have independent walk termination."""
|
||||
tid = str(uuid4())
|
||||
configs: list = []
|
||||
parent_cfg = None
|
||||
from langgraph.checkpoint.base import Checkpoint
|
||||
from langgraph.checkpoint.base.id import uuid6
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
from langgraph.checkpoint.conformance.test_utils import generate_metadata
|
||||
|
||||
for step in range(5):
|
||||
config = {"configurable": {"thread_id": tid, "checkpoint_ns": ""}}
|
||||
if parent_cfg:
|
||||
config["configurable"]["checkpoint_id"] = parent_cfg["configurable"][
|
||||
"checkpoint_id"
|
||||
]
|
||||
cv: dict = {}
|
||||
cvs: dict = {}
|
||||
if step == 1:
|
||||
cv["a"] = _DeltaSnapshot("snap_a")
|
||||
cvs["a"] = step + 1
|
||||
if step == 3:
|
||||
cv["b"] = _DeltaSnapshot("snap_b")
|
||||
cvs["b"] = step + 1
|
||||
cp = Checkpoint(
|
||||
v=1,
|
||||
id=str(uuid6(clock_seq=-1)),
|
||||
ts="",
|
||||
channel_values=cv,
|
||||
channel_versions=cvs,
|
||||
versions_seen={},
|
||||
updated_channels=None,
|
||||
)
|
||||
parent_cfg = await saver.aput(config, cp, generate_metadata(step=step), cvs)
|
||||
configs.append(parent_cfg)
|
||||
await saver.aput_writes(parent_cfg, [("a", step), ("b", step)], str(uuid4()))
|
||||
|
||||
head = configs[-1]
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["a", "b"])
|
||||
a_writes = [w[2] for w in result["a"]["writes"]]
|
||||
b_writes = [w[2] for w in result["b"]["writes"]]
|
||||
assert a_writes == [1, 2, 3], f"Expected a writes [1,2,3], got {a_writes}"
|
||||
assert b_writes == [3], f"Expected b writes [3], got {b_writes}"
|
||||
|
||||
|
||||
async def test_history_empty_channels_returns_empty(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Empty channels list returns empty mapping."""
|
||||
tid = str(uuid4())
|
||||
configs = await build_delta_chain(
|
||||
saver, thread_id=tid, channel="ch", snapshots_at_steps=[0], total_steps=3
|
||||
)
|
||||
result = await saver.aget_delta_channel_history(config=configs[-1], channels=[])
|
||||
assert result == {}
|
||||
|
||||
|
||||
async def test_history_walk_to_root_no_seed(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Walk reaches root without finding seed — no 'seed' key in result."""
|
||||
tid = str(uuid4())
|
||||
configs = await build_delta_chain(
|
||||
saver,
|
||||
thread_id=tid,
|
||||
channel="ch",
|
||||
snapshots_at_steps=[],
|
||||
total_steps=4,
|
||||
)
|
||||
head = configs[-1]
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
||||
assert "seed" not in result["ch"], f"Expected no seed, got {result['ch']}"
|
||||
|
||||
|
||||
async def test_history_migration_plain_value_as_seed(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Pre-delta plain value in channel_values acts as seed (migration case).
|
||||
|
||||
When a thread was originally using a regular channel (BinaryOperatorAggregate)
|
||||
and later switches to DeltaChannel, the old checkpoint has a plain value in
|
||||
channel_values[ch] (not a _DeltaSnapshot). The walk should treat it as the
|
||||
seed and terminate there.
|
||||
"""
|
||||
from langgraph.checkpoint.base import Checkpoint
|
||||
from langgraph.checkpoint.base.id import uuid6
|
||||
|
||||
from langgraph.checkpoint.conformance.test_utils import generate_metadata
|
||||
|
||||
tid = str(uuid4())
|
||||
configs: list = []
|
||||
parent_cfg = None
|
||||
|
||||
for step in range(4):
|
||||
config = {"configurable": {"thread_id": tid, "checkpoint_ns": ""}}
|
||||
if parent_cfg:
|
||||
config["configurable"]["checkpoint_id"] = parent_cfg["configurable"][
|
||||
"checkpoint_id"
|
||||
]
|
||||
cv: dict = {}
|
||||
cvs: dict = {}
|
||||
# Step 1: plain value (migration case — old checkpoint before delta)
|
||||
if step == 1:
|
||||
cv["ch"] = [10, 20, 30]
|
||||
cvs["ch"] = step + 1
|
||||
cp = Checkpoint(
|
||||
v=1,
|
||||
id=str(uuid6(clock_seq=-1)),
|
||||
ts="",
|
||||
channel_values=cv,
|
||||
channel_versions=cvs,
|
||||
versions_seen={},
|
||||
updated_channels=None,
|
||||
)
|
||||
parent_cfg = await saver.aput(config, cp, generate_metadata(step=step), cvs)
|
||||
configs.append(parent_cfg)
|
||||
if step != 1:
|
||||
await saver.aput_writes(parent_cfg, [("ch", step)], str(uuid4()))
|
||||
|
||||
head = configs[-1]
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
||||
# Seed should be the plain value from step 1
|
||||
assert "seed" in result["ch"], "Expected seed from migration plain value at step 1"
|
||||
seed = result["ch"]["seed"]
|
||||
assert seed == [10, 20, 30], f"Expected plain value [10,20,30], got {seed}"
|
||||
# Writes should be from step 2 only (between seed at step 1 and head's parent step 2)
|
||||
writes = result["ch"]["writes"]
|
||||
values = [w[2] for w in writes]
|
||||
assert values == [2], f"Expected [2], got {values}"
|
||||
|
||||
|
||||
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
||||
test_history_returns_writes_oldest_first,
|
||||
test_history_seed_is_nearest_snapshot,
|
||||
test_history_excludes_target_pending_writes,
|
||||
test_history_multi_channel,
|
||||
test_history_empty_channels_returns_empty,
|
||||
test_history_walk_to_root_no_seed,
|
||||
test_history_migration_plain_value_as_seed,
|
||||
]
|
||||
|
||||
|
||||
async def run_delta_channel_history_tests(
|
||||
saver: BaseCheckpointSaver,
|
||||
on_test_result: Callable[[str, str, bool, str | None], None] | None = None,
|
||||
) -> tuple[int, int, list[str]]:
|
||||
"""Run all delta_channel_history tests. Returns (passed, failed, failure_names)."""
|
||||
passed = 0
|
||||
failed = 0
|
||||
failures: list[str] = []
|
||||
for test_fn in ALL_DELTA_CHANNEL_HISTORY_TESTS:
|
||||
try:
|
||||
await test_fn(saver)
|
||||
passed += 1
|
||||
if on_test_result:
|
||||
on_test_result("delta_channel_history", test_fn.__name__, True, None)
|
||||
except Exception:
|
||||
failed += 1
|
||||
msg = f"{test_fn.__name__}: {traceback.format_exc()}"
|
||||
failures.append(msg)
|
||||
if on_test_result:
|
||||
on_test_result(
|
||||
"delta_channel_history",
|
||||
test_fn.__name__,
|
||||
False,
|
||||
traceback.format_exc(),
|
||||
)
|
||||
return passed, failed, failures
|
||||
@@ -19,6 +19,9 @@ from langgraph.checkpoint.conformance.spec.test_delete_for_runs import (
|
||||
from langgraph.checkpoint.conformance.spec.test_delete_thread import (
|
||||
run_delete_thread_tests,
|
||||
)
|
||||
from langgraph.checkpoint.conformance.spec.test_delta_channel_history import (
|
||||
run_delta_channel_history_tests,
|
||||
)
|
||||
from langgraph.checkpoint.conformance.spec.test_get_tuple import run_get_tuple_tests
|
||||
from langgraph.checkpoint.conformance.spec.test_list import run_list_tests
|
||||
from langgraph.checkpoint.conformance.spec.test_prune import run_prune_tests
|
||||
@@ -35,6 +38,7 @@ _RUNNERS = {
|
||||
Capability.DELETE_FOR_RUNS: run_delete_for_runs_tests,
|
||||
Capability.COPY_THREAD: run_copy_thread_tests,
|
||||
Capability.PRUNE: run_prune_tests,
|
||||
Capability.DELTA_CHANNEL_HISTORY: run_delta_channel_history_tests,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -43,7 +43,11 @@ asyncio_mode = "auto"
|
||||
# The extended methods (acopy_thread, adelete_for_runs, aprune) are checked
|
||||
# at runtime via capability detection and may not exist on the installed
|
||||
# base class. Dict literal inference is also overly strict for RunnableConfig.
|
||||
# Delta-channel tests import from `langgraph` (not a declared dep of this
|
||||
# package — at test time it is installed alongside); private `_DeltaSnapshot`
|
||||
# imports are intentional (beta surface).
|
||||
unresolved-attribute = "ignore"
|
||||
unresolved-import = "ignore"
|
||||
invalid-argument-type = "ignore"
|
||||
invalid-return-type = "ignore"
|
||||
|
||||
@@ -58,6 +62,9 @@ lint.select = [
|
||||
lint.ignore = ["E501", "B008"]
|
||||
target-version = "py310"
|
||||
|
||||
[tool.uv.sources]
|
||||
langgraph-checkpoint = {path = "../checkpoint", editable = true}
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "testpypi"
|
||||
url = "https://test.pypi.org/simple/"
|
||||
|
||||
Generated
+605
-504
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,35 @@
|
||||
"""Run delta-channel conformance capabilities against AsyncSqliteSaver."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip(
|
||||
"langgraph.checkpoint.conformance",
|
||||
reason="langgraph-checkpoint-conformance not installed",
|
||||
)
|
||||
pytest.importorskip("aiosqlite", reason="aiosqlite not installed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delta_channel_conformance():
|
||||
from langgraph.checkpoint.conformance import validate
|
||||
from langgraph.checkpoint.conformance.initializer import checkpointer_test
|
||||
|
||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
||||
|
||||
@checkpointer_test(name="AsyncSqliteSaver")
|
||||
async def sqlite_saver():
|
||||
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
|
||||
yield saver
|
||||
|
||||
report = await validate(
|
||||
sqlite_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}")
|
||||
@@ -327,6 +327,14 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
|
||||
Args:
|
||||
run_ids: The run IDs whose checkpoints should be deleted.
|
||||
|
||||
!!! warning "DeltaChannel"
|
||||
|
||||
Deleting a run that produced ancestor `checkpoint_writes` — or
|
||||
the only `_DeltaSnapshot` blob — for a still-live thread will
|
||||
break reconstruction of any `DeltaChannel` whose history
|
||||
depended on those rows. See the `DeltaChannel` note on `prune`
|
||||
for safe-recovery strategies.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -340,6 +348,17 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
Args:
|
||||
source_thread_id: The thread ID to copy from.
|
||||
target_thread_id: The thread ID to copy to.
|
||||
|
||||
!!! warning "DeltaChannel"
|
||||
|
||||
Implementations must copy the **complete** parent chain (all
|
||||
ancestor checkpoints and their `checkpoint_writes`) — copying
|
||||
only the head checkpoint will leave the target thread with
|
||||
`DeltaChannel` state that cannot be reconstructed (no path back
|
||||
to a `_DeltaSnapshot` ancestor). Equivalently, the copy must
|
||||
include enough ancestors that every `DeltaChannel`-backed key
|
||||
has either a `_DeltaSnapshot` in `channel_values` somewhere in
|
||||
the chain, or a complete write history back to the chain root.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -355,6 +374,34 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
thread_ids: The thread IDs to prune.
|
||||
strategy: The pruning strategy. `"keep_latest"` retains only the most
|
||||
recent checkpoint per namespace. `"delete"` removes all checkpoints.
|
||||
|
||||
!!! warning "DeltaChannel"
|
||||
|
||||
Custom implementations must be `DeltaChannel`-aware. `DeltaChannel`
|
||||
stores only a sentinel in `channel_values` for non-snapshot steps;
|
||||
reconstruction walks the parent chain via
|
||||
`get_delta_channel_history`, accumulating rows from
|
||||
`checkpoint_writes` until it reaches an ancestor whose
|
||||
`channel_values` contains a `_DeltaSnapshot` blob (written every
|
||||
`snapshot_frequency` updates).
|
||||
|
||||
A naive `"keep_latest"` that drops intermediate checkpoints and
|
||||
their writes can sever that chain: the surviving "latest"
|
||||
checkpoint is rarely a snapshot point itself, so its delta
|
||||
channels would silently reconstruct as empty (no error raised —
|
||||
`get_delta_channel_history` simply returns no `seed`). Safe
|
||||
options when the graph uses `DeltaChannel`:
|
||||
|
||||
* Walk back from each kept checkpoint and preserve every
|
||||
ancestor (plus its `checkpoint_writes`) up to the nearest one
|
||||
whose `channel_values` already contains a `_DeltaSnapshot` for
|
||||
every `DeltaChannel`-backed key.
|
||||
* Force a fresh snapshot on the kept checkpoint before deleting
|
||||
ancestors — rewrite `channel_values[k] = _DeltaSnapshot(value)`
|
||||
for each delta channel `k` (resolving `value` via the existing
|
||||
ancestor walk first), then prune.
|
||||
* Skip pruning threads whose graph uses `DeltaChannel` until one
|
||||
of the above is implemented.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -471,6 +518,13 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
|
||||
Args:
|
||||
run_ids: The run IDs whose checkpoints should be deleted.
|
||||
|
||||
!!! warning "DeltaChannel"
|
||||
|
||||
See `delete_for_runs` — deleting rows a still-live thread's
|
||||
`DeltaChannel` reconstruction depends on (writes between the
|
||||
head and its nearest `_DeltaSnapshot` ancestor) will silently
|
||||
corrupt that channel's state.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -484,6 +538,13 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
Args:
|
||||
source_thread_id: The thread ID to copy from.
|
||||
target_thread_id: The thread ID to copy to.
|
||||
|
||||
!!! warning "DeltaChannel"
|
||||
|
||||
See `copy_thread` — the copy must carry the complete parent
|
||||
chain (or at least back to a `_DeltaSnapshot` ancestor for every
|
||||
`DeltaChannel`) so the target thread can reconstruct delta
|
||||
state.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -499,6 +560,13 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
thread_ids: The thread IDs to prune.
|
||||
strategy: The pruning strategy. `"keep_latest"` retains only the most
|
||||
recent checkpoint per namespace. `"delete"` removes all checkpoints.
|
||||
|
||||
!!! warning "DeltaChannel"
|
||||
|
||||
See `prune` for the full `DeltaChannel` caveat. In short:
|
||||
`"keep_latest"` must not drop ancestor checkpoints / writes that
|
||||
sit between the kept checkpoint and the nearest `_DeltaSnapshot`
|
||||
ancestor, or delta channels will silently reconstruct as empty.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -44,7 +44,7 @@ if TYPE_CHECKING:
|
||||
AllowedMsgpackModules,
|
||||
)
|
||||
|
||||
LC_REVIVER = Reviver()
|
||||
LC_REVIVER = Reviver(allowed_objects="core")
|
||||
EMPTY_BYTES = b""
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
"""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():
|
||||
from langgraph.checkpoint.conformance import validate
|
||||
from langgraph.checkpoint.conformance.initializer import checkpointer_test
|
||||
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
@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}")
|
||||
@@ -34,28 +34,23 @@ def empty_checkpoint() -> Checkpoint:
|
||||
)
|
||||
|
||||
|
||||
def _should_snapshot_delta(
|
||||
name: str,
|
||||
ch: DeltaChannel,
|
||||
updates_since_snapshot: Mapping[str, int],
|
||||
*,
|
||||
force: bool,
|
||||
) -> bool:
|
||||
"""Decide whether `ch` should write a `_DeltaSnapshot` this step.
|
||||
def delta_channels_to_snapshot(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
counts: Mapping[str, int],
|
||||
) -> set[str]:
|
||||
"""Return the set of DeltaChannel names that should snapshot now.
|
||||
|
||||
Triggers:
|
||||
* `force` — always snapshot (used by `durability="exit"`).
|
||||
* Update-count: this channel has accumulated at least
|
||||
`snapshot_frequency` updates since its last snapshot. The count
|
||||
is supplied by the caller via `updates_since_snapshot[name]` and
|
||||
is reset to `0` whenever a snapshot fires.
|
||||
|
||||
Version-format-independent: works for `int`, `float`, and `str`
|
||||
versioning schemes alike.
|
||||
A channel snapshots when its accumulated update count (since the last
|
||||
snapshot) reaches or exceeds `snapshot_frequency`. This is a pure
|
||||
predicate — no mutation.
|
||||
"""
|
||||
if force:
|
||||
return True
|
||||
return updates_since_snapshot.get(name, 0) >= ch.snapshot_frequency
|
||||
return {
|
||||
name
|
||||
for name, ch in channels.items()
|
||||
if isinstance(ch, DeltaChannel)
|
||||
and ch.is_available()
|
||||
and counts.get(name, 0) >= ch.snapshot_frequency
|
||||
}
|
||||
|
||||
|
||||
def create_checkpoint(
|
||||
@@ -66,34 +61,19 @@ def create_checkpoint(
|
||||
id: str | None = None,
|
||||
updated_channels: set[str] | None = None,
|
||||
get_next_version: GetNextVersion | None = None,
|
||||
force_delta_snapshot: bool = False,
|
||||
updates_since_snapshot: Mapping[str, int] | None = None,
|
||||
new_updates_since_snapshot: dict[str, int] | None = None,
|
||||
channels_to_snapshot: set[str] | None = None,
|
||||
) -> Checkpoint:
|
||||
"""Create a checkpoint for the given channels.
|
||||
"""Build a new Checkpoint from the previous one and live channel state.
|
||||
|
||||
For each `DeltaChannel`, a `_DeltaSnapshot(value)` blob is written into
|
||||
`channel_values[k]` when this channel has accumulated at least
|
||||
`snapshot_frequency` updates since its last snapshot (counter supplied
|
||||
via `updates_since_snapshot`). Otherwise the channel is omitted from
|
||||
`channel_values`; its `channel_versions` entry still bumps so that the
|
||||
saver tracks the channel and the ancestor walk can replay writes.
|
||||
|
||||
Snapshots are eager: even if the channel had no write this step, a
|
||||
version bump is forced (via `get_next_version`) so `put()` includes
|
||||
the channel in `new_versions` and stores the blob.
|
||||
|
||||
`force_delta_snapshot` ignores the cadence and always snapshots —
|
||||
used by `durability="exit"` where intermediate writes are not stored
|
||||
as ancestor `checkpoint_writes`.
|
||||
|
||||
If `new_updates_since_snapshot` is provided, the function resets the
|
||||
counter to `0` for any channel that snapshotted this step. Counters
|
||||
for channels that did not snapshot are left untouched (the caller is
|
||||
responsible for incrementing them based on `updated_channels`).
|
||||
For each name in `channels_to_snapshot`, a `_DeltaSnapshot(value)` blob
|
||||
is written into `channel_values[k]`. Other delta channels are omitted
|
||||
from `channel_values` — the ancestor walk reconstructs their state
|
||||
from `checkpoint_writes`. Callers compute the set via
|
||||
`delta_channels_to_snapshot(channels, counts)`; defaults to empty
|
||||
(no snapshots) when not provided.
|
||||
"""
|
||||
ts = datetime.now(timezone.utc).isoformat()
|
||||
counts = updates_since_snapshot or {}
|
||||
channels_to_snapshot = channels_to_snapshot or set()
|
||||
if channels is None:
|
||||
values = checkpoint["channel_values"]
|
||||
channel_versions = checkpoint["channel_versions"]
|
||||
@@ -104,25 +84,23 @@ def create_checkpoint(
|
||||
if k not in channel_versions:
|
||||
continue
|
||||
ch = channels[k]
|
||||
if (
|
||||
isinstance(ch, DeltaChannel)
|
||||
and ch.is_available()
|
||||
and _should_snapshot_delta(
|
||||
k,
|
||||
ch,
|
||||
counts,
|
||||
force=force_delta_snapshot,
|
||||
)
|
||||
):
|
||||
# Eager snapshot: bump version if not already written this step
|
||||
# so put() includes this channel in new_versions and stores blob.
|
||||
if k in channels_to_snapshot:
|
||||
# In exit mode, the snapshot decision is deferred to exit
|
||||
# time (intermediate steps have do_checkpoint=False). The
|
||||
# channel's count may have reached snapshot_frequency over
|
||||
# several supersteps, but the LAST superstep may not have
|
||||
# written to this channel. In that case apply_writes()
|
||||
# (in _algo.py) didn't bump this channel's version, so
|
||||
# saver.put() wouldn't include it in new_versions and
|
||||
# the snapshot blob would be silently dropped. The manual
|
||||
# bump below closes the gap. In sync/async durability this
|
||||
# branch is effectively dead code (the step that pushes
|
||||
# the count to freq always writes the channel).
|
||||
if get_next_version is not None and (
|
||||
updated_channels is None or k not in updated_channels
|
||||
):
|
||||
channel_versions[k] = get_next_version(channel_versions[k], None)
|
||||
values[k] = _DeltaSnapshot(ch.get())
|
||||
if new_updates_since_snapshot is not None:
|
||||
new_updates_since_snapshot[k] = 0
|
||||
else:
|
||||
v = ch.checkpoint()
|
||||
if v is not MISSING:
|
||||
|
||||
@@ -100,6 +100,7 @@ from langgraph.pregel._checkpoint import (
|
||||
channels_from_checkpoint,
|
||||
copy_checkpoint,
|
||||
create_checkpoint,
|
||||
delta_channels_to_snapshot,
|
||||
empty_checkpoint,
|
||||
)
|
||||
from langgraph.pregel._executor import (
|
||||
@@ -194,8 +195,40 @@ class PregelLoop:
|
||||
_migrate_checkpoint: Callable[[Checkpoint], None] | None
|
||||
submit: Submit
|
||||
channels: Mapping[str, BaseChannel]
|
||||
# Only set on AsyncPregelLoop; sync loops keep this as None.
|
||||
# Futures from `checkpointer.put_writes` calls that produced delta-channel
|
||||
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
|
||||
# local `futs` then reset to `[]` and wait/gather) before putting the
|
||||
# next checkpoint, so a checkpoint never becomes durable before the
|
||||
# writes that produced it. Initialised to `[]` in both sync and async
|
||||
# `__enter__`; stays `None` only when no checkpointer.
|
||||
_delta_write_futs: list[Any] | None = None
|
||||
|
||||
# Exit-mode accumulator: every delta-channel write produced during this
|
||||
# run (input writes from `_first` + per-superstep writes captured in
|
||||
# `after_tick`). At exit, `_put_exit_delta_writes` filters out channels
|
||||
# that will snapshot, then persists the rest under an anchor parent.
|
||||
# `None` when not in exit mode (so the capture sites are no-ops).
|
||||
# Each tuple is `(step, task_id, channel, value)` — `step` drives the
|
||||
# synthetic step-prefixed task_id used to preserve chronological order
|
||||
# under the saver's `ORDER BY task_id, idx` sorting.
|
||||
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
|
||||
|
||||
# The checkpoint_config that points at the parent loaded at `__enter__`
|
||||
# (or the synthetic-empty checkpoint, on first run). We capture it
|
||||
# eagerly because every `_put_checkpoint` advances `self.checkpoint_config`
|
||||
# to the newly-saved checkpoint's id — by exit time the original parent
|
||||
# config would otherwise be lost. `_put_exit_delta_writes` uses this:
|
||||
# on resumed runs as the anchor for exit delta writes; on first runs
|
||||
# to derive the lazy stub's config (its `checkpoint_id` is the
|
||||
# synthetic-empty id we want the stub persisted under).
|
||||
_initial_checkpoint_config: RunnableConfig
|
||||
|
||||
# True iff the saver actually returned a tuple at `__enter__`. False
|
||||
# on the first-ever run for a thread (no parent persisted yet).
|
||||
# `_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
|
||||
|
||||
managed: ManagedValueMapping
|
||||
checkpoint: Checkpoint
|
||||
checkpoint_id_saved: str
|
||||
@@ -637,6 +670,11 @@ class PregelLoop:
|
||||
self._emit(
|
||||
"values", map_output_values, self.output_keys, writes, self.channels
|
||||
)
|
||||
# capture delta-channel writes for exit-mode accumulator before clearing
|
||||
if self._exit_delta_writes is not None:
|
||||
for tid, ch, v in self.checkpoint_pending_writes:
|
||||
if isinstance(self.specs.get(ch), DeltaChannel):
|
||||
self._exit_delta_writes.append((self.step, tid, ch, v))
|
||||
# clear pending writes
|
||||
self.checkpoint_pending_writes.clear()
|
||||
# only replay (re-execute) done tasks on the first tick
|
||||
@@ -854,6 +892,27 @@ class PregelLoop:
|
||||
self.checkpointer_get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# Input writes go through `apply_writes` directly (above) — they
|
||||
# never enter `checkpoint_pending_writes`, so the after_tick
|
||||
# capture site does not see them. In exit mode, capture them
|
||||
# here so `_exit_delta_writes` includes the input's delta writes
|
||||
# alongside per-superstep writes; otherwise the input would be
|
||||
# lost on read (it's not in final_checkpoint.channel_values for
|
||||
# sub-freq channels, and walks ignore target.pending_writes).
|
||||
if self._exit_delta_writes is not None:
|
||||
for c, v in input_writes:
|
||||
if isinstance(self.specs.get(c), DeltaChannel):
|
||||
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).
|
||||
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)
|
||||
# save input checkpoint
|
||||
self.updated_channels = updated_channels
|
||||
self._put_checkpoint({"source": "input"})
|
||||
@@ -905,36 +964,60 @@ class PregelLoop:
|
||||
return updated_channels
|
||||
|
||||
def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:
|
||||
# assign step and parents
|
||||
# `is` (object identity) — not `==`. Three of four call sites pass a
|
||||
# fresh dict ({"source":"input"|"loop"|"fork"}); only
|
||||
# `_suppress_interrupt`(will rename to _on_loop_exit soon)
|
||||
# at exit reuses the existing `self.checkpoint_metadata` instance. So
|
||||
# `metadata is self.checkpoint_metadata` is True only on the exit call,
|
||||
# which is what we use to gate exit-only behaviour (skip count-bump,
|
||||
# don't replace metadata). Could be replaced by an explicit
|
||||
# `exiting: bool = False` parameter; left as-is to match the existing
|
||||
# idiom in this file.
|
||||
# TODO: replace with an explicit `exiting: bool = False` parameter.
|
||||
exiting = metadata is self.checkpoint_metadata
|
||||
if exiting and self.checkpoint["id"] == self.checkpoint_id_saved:
|
||||
# checkpoint already saved
|
||||
return
|
||||
# Carry per-delta-channel update bookkeeping forward across
|
||||
# supersteps. Capture from the OLD metadata before potentially
|
||||
# replacing it with a fresh dict that wouldn't contain it. Then
|
||||
# increment for any delta channel updated this step (so the count
|
||||
# reflects "supersteps that wrote to this channel since last
|
||||
# snapshot"). create_checkpoint will reset entries to 0 for any
|
||||
# channel that fires a snapshot this step.
|
||||
prev_counts = dict(
|
||||
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
|
||||
)
|
||||
new_counts = dict(prev_counts)
|
||||
if self.updated_channels:
|
||||
for ch_name in self.updated_channels:
|
||||
ch_obj = self.channels.get(ch_name)
|
||||
if isinstance(ch_obj, DeltaChannel):
|
||||
new_counts[ch_name] = new_counts.get(ch_name, 0) + 1
|
||||
# Per-delta-channel update bookkeeping.
|
||||
#
|
||||
# `_put_checkpoint` is called once per superstep with a fresh
|
||||
# metadata dict (source="input"|"loop"|"fork") — those are the
|
||||
# intermediate calls that bump the count by +1 for each delta
|
||||
# channel touched that step. In exit mode,
|
||||
# `_suppress_interrupt`(will rename to _on_loop_exit soon)
|
||||
# additionally calls `_put_checkpoint(self.checkpoint_metadata)` AT
|
||||
# EXIT to commit the final checkpoint — this runs *after* the last
|
||||
# intermediate call already counted the last superstep. So the
|
||||
# exit call must NOT bump again or it would double-count the last
|
||||
# superstep. (Sync/async durability does not call `_put_checkpoint`
|
||||
# at exit, so the issue only surfaces in exit mode. force_delta_snapshot
|
||||
# used to mask this latent bug by resetting every count to 0.)
|
||||
if not exiting:
|
||||
prev_counts = dict(
|
||||
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
|
||||
)
|
||||
new_counts = dict(prev_counts)
|
||||
if self.updated_channels:
|
||||
for ch_name in self.updated_channels:
|
||||
if isinstance(self.channels.get(ch_name), DeltaChannel):
|
||||
new_counts[ch_name] = new_counts.get(ch_name, 0) + 1
|
||||
metadata["step"] = self.step
|
||||
metadata["parents"] = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})
|
||||
self.checkpoint_metadata = metadata
|
||||
else:
|
||||
new_counts = dict(
|
||||
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
|
||||
)
|
||||
# do checkpoint?
|
||||
do_checkpoint = self._checkpointer_put_after_previous is not None and (
|
||||
exiting or self.durability != "exit"
|
||||
)
|
||||
# create new checkpoint
|
||||
channels_to_snapshot = (
|
||||
delta_channels_to_snapshot(self.channels, new_counts)
|
||||
if do_checkpoint
|
||||
else set()
|
||||
)
|
||||
self.checkpoint = create_checkpoint(
|
||||
self.checkpoint,
|
||||
self.channels if do_checkpoint else None,
|
||||
@@ -944,10 +1027,10 @@ class PregelLoop:
|
||||
get_next_version=self.checkpointer_get_next_version
|
||||
if do_checkpoint
|
||||
else None,
|
||||
force_delta_snapshot=exiting and self.durability == "exit",
|
||||
updates_since_snapshot=new_counts,
|
||||
new_updates_since_snapshot=new_counts,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
)
|
||||
for k in channels_to_snapshot:
|
||||
new_counts[k] = 0
|
||||
if new_counts:
|
||||
self.checkpoint_metadata["delta_updates_since_snapshot"] = new_counts
|
||||
elif "delta_updates_since_snapshot" in self.checkpoint_metadata:
|
||||
@@ -1010,6 +1093,97 @@ class PregelLoop:
|
||||
# increment step
|
||||
self.step += 1
|
||||
|
||||
def _put_exit_delta_writes(self) -> None:
|
||||
"""Stage stub + accumulated delta writes so final_checkpoint's put
|
||||
waits on them (visibility invariant: both must be durable before
|
||||
final_checkpoint becomes visible to readers).
|
||||
|
||||
Stub is created lazily — only when no persisted parent exists AND at
|
||||
least one delta channel has writes that won't be snapshotted.
|
||||
"""
|
||||
if (
|
||||
not self._exit_delta_writes
|
||||
or self.checkpointer is None
|
||||
or self._checkpointer_put_after_previous is None
|
||||
or self.checkpointer_put_writes is None
|
||||
):
|
||||
return
|
||||
|
||||
counts = self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
|
||||
channels_to_snapshot = delta_channels_to_snapshot(self.channels, counts)
|
||||
|
||||
pending = [
|
||||
(step, tid, ch, v)
|
||||
for (step, tid, ch, v) in self._exit_delta_writes
|
||||
if ch not in channels_to_snapshot
|
||||
]
|
||||
if not pending:
|
||||
return
|
||||
|
||||
if self._has_persisted_parent:
|
||||
# _initial_checkpoint_config's checkpoint_id is the saved parent's
|
||||
# id (saver returned a real tuple at __enter__).
|
||||
anchor_config = self._initial_checkpoint_config
|
||||
else:
|
||||
stub_cp = empty_checkpoint()
|
||||
stub_cp["id"] = self.checkpoint_id_saved
|
||||
stub_cp["ts"] = datetime.now(timezone.utc).isoformat()
|
||||
# Stub has no parent (checkpoint_id=None in config).
|
||||
stub_put_config = patch_configurable(
|
||||
self._initial_checkpoint_config,
|
||||
{CONFIG_KEY_CHECKPOINT_ID: None},
|
||||
)
|
||||
# Anchor config for put_writes: checkpoint_id = stub's id.
|
||||
anchor_config = patch_configurable(
|
||||
self._initial_checkpoint_config,
|
||||
{CONFIG_KEY_CHECKPOINT_ID: stub_cp["id"]},
|
||||
)
|
||||
self._put_checkpoint_fut = self.submit(
|
||||
self._checkpointer_put_after_previous,
|
||||
getattr(self, "_put_checkpoint_fut", None),
|
||||
stub_put_config,
|
||||
stub_cp,
|
||||
{"step": -2},
|
||||
{},
|
||||
)
|
||||
# Set checkpoint_config so final_checkpoint's _put_checkpoint
|
||||
# sees the stub as its parent.
|
||||
self.checkpoint_config = anchor_config
|
||||
|
||||
# Step-prefixed synthetic task_id preserves chronological superstep
|
||||
# order under the saver's ORDER BY task_id, idx sorting.
|
||||
grouped: dict[tuple[int, str], list[tuple[str, Any]]] = {}
|
||||
for step, tid, ch, v in pending:
|
||||
grouped.setdefault((step, tid), []).append((ch, v))
|
||||
anchor_write_config = patch_configurable(
|
||||
anchor_config,
|
||||
{
|
||||
CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINT_NS, ""
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID],
|
||||
},
|
||||
)
|
||||
for (step, tid), entries in grouped.items():
|
||||
synth_tid = f"{step:08d}-{tid}"
|
||||
if self.checkpointer_put_writes_accepts_task_path:
|
||||
fut = self.submit(
|
||||
self.checkpointer_put_writes,
|
||||
anchor_write_config,
|
||||
entries,
|
||||
synth_tid,
|
||||
"",
|
||||
)
|
||||
else:
|
||||
fut = self.submit(
|
||||
self.checkpointer_put_writes,
|
||||
anchor_write_config,
|
||||
entries,
|
||||
synth_tid,
|
||||
)
|
||||
if self._delta_write_futs is not None:
|
||||
self._delta_write_futs.append(fut)
|
||||
|
||||
def _suppress_interrupt(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
@@ -1025,6 +1199,7 @@ class PregelLoop:
|
||||
# or a nested graph with checkpointer=True
|
||||
or all(NS_END not in part for part in self.checkpoint_ns)
|
||||
):
|
||||
self._put_exit_delta_writes()
|
||||
self._put_checkpoint(self.checkpoint_metadata)
|
||||
self._put_pending_writes()
|
||||
# suppress interrupt
|
||||
@@ -1230,6 +1405,9 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
if self._delta_write_futs:
|
||||
futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
concurrent.futures.wait(futs)
|
||||
try:
|
||||
if prev is not None:
|
||||
prev.result()
|
||||
@@ -1347,6 +1525,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = self.checkpointer.get_tuple(self.checkpoint_config)
|
||||
|
||||
# Capture before the synthetic-empty fallback below overwrites `saved`.
|
||||
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
|
||||
# to lazy-create a stub instead of anchoring delta writes on a parent.
|
||||
self._has_persisted_parent = saved is not None
|
||||
if saved is None:
|
||||
saved = CheckpointTuple(
|
||||
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
|
||||
@@ -1362,6 +1544,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
**saved.config.get(CONF, {}),
|
||||
},
|
||||
}
|
||||
self._initial_checkpoint_config = self.checkpoint_config
|
||||
self.prev_checkpoint_config = saved.parent_config
|
||||
self.checkpoint_id_saved = saved.checkpoint["id"]
|
||||
self.checkpoint = saved.checkpoint
|
||||
@@ -1371,6 +1554,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
if saved.pending_writes is not None
|
||||
else []
|
||||
)
|
||||
self._delta_write_futs = []
|
||||
self._exit_delta_writes = (
|
||||
[] if self.durability == "exit" and self.checkpointer is not None else None
|
||||
)
|
||||
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
|
||||
self.channels, self.managed = channels_from_checkpoint(
|
||||
self.specs,
|
||||
@@ -1596,6 +1783,10 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
|
||||
|
||||
# Capture before the synthetic-empty fallback below overwrites `saved`.
|
||||
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
|
||||
# to lazy-create a stub instead of anchoring delta writes on a parent.
|
||||
self._has_persisted_parent = saved is not None
|
||||
if saved is None:
|
||||
saved = CheckpointTuple(
|
||||
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
|
||||
@@ -1611,6 +1802,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
**saved.config.get(CONF, {}),
|
||||
},
|
||||
}
|
||||
self._initial_checkpoint_config = self.checkpoint_config
|
||||
self.prev_checkpoint_config = saved.parent_config
|
||||
self.checkpoint_id_saved = saved.checkpoint["id"]
|
||||
self.checkpoint = saved.checkpoint
|
||||
@@ -1621,6 +1813,9 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
else []
|
||||
)
|
||||
self._delta_write_futs = []
|
||||
self._exit_delta_writes = (
|
||||
[] if self.durability == "exit" and self.checkpointer is not None else None
|
||||
)
|
||||
self.submit = await self.stack.enter_async_context(
|
||||
AsyncBackgroundExecutor(self.config)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,365 @@
|
||||
"""Tests for exit-mode delta channel persistence redesign.
|
||||
|
||||
Validates that `durability="exit"` correctly persists delta-channel writes
|
||||
using count-based snapshot decisions (rather than force-snapshotting every
|
||||
channel), lazy stub creation when no parent exists, and proper read-path
|
||||
reconstruction via ancestor walks.
|
||||
"""
|
||||
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, 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.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _build_graph(
|
||||
checkpointer: InMemorySaver,
|
||||
*,
|
||||
freq: int = 1000,
|
||||
) -> Any:
|
||||
channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=freq)
|
||||
# Functional TypedDict form: class form can't reference `channel` (a
|
||||
# local variable) inside Annotated due to forward-ref evaluation rules.
|
||||
State = TypedDict("State", {"messages": Annotated[list, channel]}) # type: ignore[call-overload] # noqa: UP013
|
||||
|
||||
def respond(state: dict) -> dict:
|
||||
i = len(state["messages"])
|
||||
return {"messages": [AIMessage(content=f"reply-{i}", id=f"ai{i}")]}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("respond", respond)
|
||||
builder.add_edge(START, "respond")
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 8a. Write-path / structural tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_exit_first_run_no_delta_writes() -> None:
|
||||
"""Graph with delta channel invoked with input that doesn't touch it.
|
||||
Only one checkpoint row, no stub."""
|
||||
State = TypedDict( # noqa: UP013
|
||||
"State",
|
||||
{
|
||||
"messages": Annotated[list, DeltaChannel(_messages_delta_reducer)],
|
||||
"value": str,
|
||||
},
|
||||
) # type: ignore[call-overload]
|
||||
|
||||
def noop(state: dict) -> dict:
|
||||
return {"value": "done"}
|
||||
|
||||
saver = InMemorySaver()
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("noop", noop)
|
||||
builder.add_edge(START, "noop")
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
config = {"configurable": {"thread_id": "no-delta-writes"}}
|
||||
|
||||
graph.invoke({"value": "start"}, config, durability="exit")
|
||||
|
||||
checkpoints = list(saver.list(config))
|
||||
assert len(checkpoints) == 1
|
||||
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
|
||||
assert len(stubs) == 0
|
||||
|
||||
|
||||
async def test_exit_first_run_all_snapshot() -> None:
|
||||
"""snapshot_frequency=1 forces every channel to snapshot.
|
||||
No stub needed; final_checkpoint has _DeltaSnapshot."""
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver, freq=1)
|
||||
config = {"configurable": {"thread_id": "all-snapshot"}}
|
||||
|
||||
result = graph.invoke(
|
||||
{"messages": [HumanMessage(content="hi", id="h1")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
assert len(result["messages"]) == 2
|
||||
|
||||
checkpoints = list(saver.list(config))
|
||||
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
|
||||
assert len(stubs) == 0
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert [m.content for m in state.values["messages"]] == ["hi", "reply-1"]
|
||||
|
||||
|
||||
async def test_exit_first_run_sub_freq_with_writes() -> None:
|
||||
"""First run with default snapshot_frequency (1000), writes below threshold.
|
||||
A stub is created; writes are anchored under it; get_state reconstructs."""
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": "sub-freq-first"}}
|
||||
|
||||
result = graph.invoke(
|
||||
{"messages": [HumanMessage(content="hello", id="h1")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
assert [m.content for m in result["messages"]] == ["hello", "reply-1"]
|
||||
|
||||
checkpoints = list(saver.list(config))
|
||||
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
|
||||
assert len(stubs) == 1, f"Expected 1 stub, got {len(stubs)}"
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert "messages" not in head.checkpoint["channel_values"]
|
||||
assert "messages" in head.checkpoint["channel_versions"]
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert [m.content for m in state.values["messages"]] == ["hello", "reply-1"]
|
||||
|
||||
|
||||
async def test_exit_resumed_run_sub_freq() -> None:
|
||||
"""Two consecutive exit runs. Second run anchors on the first's
|
||||
final_checkpoint (no new stub). Ordering preserved."""
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": "resumed-sub-freq"}}
|
||||
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content="msg1", id="h1")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content="msg2", id="h2")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
|
||||
checkpoints = list(saver.list(config))
|
||||
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
|
||||
assert len(stubs) == 1
|
||||
|
||||
state = graph.get_state(config)
|
||||
contents = [m.content for m in state.values["messages"]]
|
||||
assert len(contents) == 4
|
||||
assert contents[0] == "msg1"
|
||||
assert contents[2] == "msg2"
|
||||
assert contents[0:4:2] == ["msg1", "msg2"]
|
||||
|
||||
|
||||
async def test_exit_count_parity_sync_vs_exit() -> None:
|
||||
"""Sync and exit durability produce the same delta_updates_since_snapshot
|
||||
after an equivalent run."""
|
||||
for durability in ("sync", "exit"):
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": f"parity-{durability}"}}
|
||||
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content="hi", id="h1")]},
|
||||
config,
|
||||
durability=durability,
|
||||
)
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
counts = head.metadata.get("delta_updates_since_snapshot", {})
|
||||
assert counts.get("messages") == 2, (
|
||||
f"durability={durability}: expected count=2, got {counts}"
|
||||
)
|
||||
|
||||
|
||||
async def test_exit_snapshot_fires_at_frequency() -> None:
|
||||
"""With snapshot_frequency=3, after 3 exit runs (each incrementing count
|
||||
by 2: input + superstep), the 2nd run hits count=4>=3, triggering snapshot.
|
||||
After that run, count resets to 0 and channel_values has _DeltaSnapshot."""
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver, freq=3)
|
||||
config = {"configurable": {"thread_id": "snapshot-at-freq"}}
|
||||
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content="m1", id="h1")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
count1 = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
|
||||
assert count1 == 2
|
||||
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content="m2", id="h2")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
count2 = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
|
||||
assert count2 == 0, f"Expected reset to 0 after snapshot, got {count2}"
|
||||
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
|
||||
|
||||
|
||||
async def test_exit_mixed_snapshot_and_non_snapshot() -> None:
|
||||
"""One delta channel at freq=1 (always snapshot) and one at freq=1000
|
||||
(never snapshot within this test). Verify correct behavior for both."""
|
||||
|
||||
fast_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1)
|
||||
slow_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1000)
|
||||
State = TypedDict( # noqa: UP013
|
||||
"State",
|
||||
{"fast": Annotated[list, fast_ch], "slow": Annotated[list, slow_ch]},
|
||||
) # type: ignore[call-overload]
|
||||
|
||||
def respond(state: dict) -> dict:
|
||||
return {
|
||||
"fast": [AIMessage(content="fast-reply", id="f1")],
|
||||
"slow": [AIMessage(content="slow-reply", id="s1")],
|
||||
}
|
||||
|
||||
saver = InMemorySaver()
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("respond", respond)
|
||||
builder.add_edge(START, "respond")
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
config = {"configurable": {"thread_id": "mixed-freq"}}
|
||||
|
||||
graph.invoke(
|
||||
{
|
||||
"fast": [HumanMessage(content="fast-in", id="fi")],
|
||||
"slow": [HumanMessage(content="slow-in", id="si")],
|
||||
},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("fast"), _DeltaSnapshot)
|
||||
assert "slow" not in head.checkpoint["channel_values"]
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert [m.content for m in state.values["fast"]] == ["fast-in", "fast-reply"]
|
||||
assert [m.content for m in state.values["slow"]] == ["slow-in", "slow-reply"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 8b. Read-path tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_exit_multi_run_replay_chain() -> None:
|
||||
"""K=4 consecutive exit runs, each adding a message. After each run,
|
||||
get_state returns all messages in chronological order."""
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": "replay-chain"}}
|
||||
|
||||
for i in range(4):
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content=f"user-{i}", id=f"h{i}")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
contents = [m.content for m in state.values["messages"]]
|
||||
user_msgs = [c for c in contents if c.startswith("user-")]
|
||||
assert user_msgs == [f"user-{j}" for j in range(i + 1)], (
|
||||
f"After run {i}: user messages out of order: {user_msgs}"
|
||||
)
|
||||
assert len(contents) == (i + 1) * 2
|
||||
|
||||
|
||||
async def test_exit_metadata_round_trip() -> None:
|
||||
"""K=5 consecutive exit runs with snapshot_frequency=5. Verify metadata
|
||||
delta_updates_since_snapshot increments correctly across runs."""
|
||||
freq = 5
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver, freq=freq)
|
||||
config = {"configurable": {"thread_id": "metadata-rt"}}
|
||||
|
||||
for i in range(1, 6):
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content=f"m{i}", id=f"h{i}")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
count = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
|
||||
cumulative = i * 2
|
||||
if cumulative >= freq:
|
||||
assert count == 0 or count == cumulative % freq or count < freq, (
|
||||
f"After run {i}: count={count} should have reset or be partial"
|
||||
)
|
||||
else:
|
||||
assert count == cumulative, (
|
||||
f"After run {i}: expected {cumulative}, got {count}"
|
||||
)
|
||||
|
||||
|
||||
async def test_exit_mixed_durability_round_trip() -> None:
|
||||
"""Alternate sync and exit durability; verify counts stay monotonic
|
||||
and state accumulates correctly."""
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": "mixed-durability"}}
|
||||
|
||||
for i, dur in enumerate(["sync", "exit", "sync", "exit"]):
|
||||
graph.invoke(
|
||||
{"messages": [HumanMessage(content=f"msg-{i}", id=f"h{i}")]},
|
||||
config,
|
||||
durability=dur,
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
contents = [m.content for m in state.values["messages"]]
|
||||
user_msgs = [c for c in contents if c.startswith("msg-")]
|
||||
assert user_msgs == [f"msg-{j}" for j in range(i + 1)], (
|
||||
f"After run {i} (durability={dur}): {user_msgs}"
|
||||
)
|
||||
assert len(contents) == (i + 1) * 2
|
||||
|
||||
|
||||
async def test_exit_snapshot_then_tail_deltas() -> None:
|
||||
"""Run 1 forces snapshot (freq=1). Run 2 at freq=1000 adds more writes
|
||||
that don't snapshot. Reading after run 2 must combine the snapshot seed
|
||||
with the tail deltas."""
|
||||
saver = InMemorySaver()
|
||||
|
||||
graph1 = _build_graph(saver, freq=1)
|
||||
config = {"configurable": {"thread_id": "snapshot-then-tail"}}
|
||||
graph1.invoke(
|
||||
{"messages": [HumanMessage(content="seed-msg", id="h1")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
|
||||
|
||||
graph2 = _build_graph(saver, freq=1000)
|
||||
graph2.invoke(
|
||||
{"messages": [HumanMessage(content="tail-msg", id="h2")]},
|
||||
config,
|
||||
durability="exit",
|
||||
)
|
||||
|
||||
state = graph2.get_state(config)
|
||||
contents = [m.content for m in state.values["messages"]]
|
||||
assert "seed-msg" in contents
|
||||
assert "tail-msg" in contents
|
||||
assert contents.index("seed-msg") < contents.index("tail-msg")
|
||||
@@ -1674,15 +1674,28 @@ async def test_arun_with_retry_timeout_observer_tracks_attempts():
|
||||
async def test_arun_with_retry_timeout_observer_emits_progress_on_heartbeat():
|
||||
events: list = []
|
||||
|
||||
# `_TimedAttemptScope.__init__` sets `_last_progress` to `time.monotonic()`,
|
||||
# but the watchdog itself doesn't start running until after `wrap_config`
|
||||
# and task scheduling — under CI load that gap can be large enough to eat
|
||||
# the entire idle window before the task body's first await even runs. We
|
||||
# defend against that by:
|
||||
# 1. Using a generous idle_timeout so scheduling slack stays well within it.
|
||||
# 2. Calling `runtime.heartbeat()` BEFORE the first sleep, which resets
|
||||
# `_last_progress` to "now" the moment the task body actually starts.
|
||||
idle_timeout_s = 1.0
|
||||
|
||||
class HeartbeatProc:
|
||||
async def ainvoke(self, input, config):
|
||||
runtime = config[CONF][CONFIG_KEY_RUNTIME]
|
||||
runtime.heartbeat() # reset the idle clock at task-body entry
|
||||
for _ in range(8):
|
||||
await asyncio.sleep(0.05)
|
||||
runtime.heartbeat()
|
||||
return "ok"
|
||||
|
||||
task = _make_task(HeartbeatProc(), timeout=_idle_timeout(0.2), name="heartbeat")
|
||||
task = _make_task(
|
||||
HeartbeatProc(), timeout=_idle_timeout(idle_timeout_s), name="heartbeat"
|
||||
)
|
||||
task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append
|
||||
assert await arun_with_retry(task, retry_policy=None) == "ok"
|
||||
|
||||
@@ -1691,13 +1704,13 @@ async def test_arun_with_retry_timeout_observer_emits_progress_on_heartbeat():
|
||||
assert by_event[-1] == "finish"
|
||||
progress = [ev for ev in events if ev.event == "progress"]
|
||||
assert progress, "expected at least one progress event from heartbeat"
|
||||
# Rate limit is `idle_timeout / 4` = 0.05s; with 8 heartbeats spaced ~0.05s
|
||||
# we should see at most ~one progress event per heartbeat (well below 8).
|
||||
# Rate limit is `idle_timeout / 4` = 0.25s; with the task running for
|
||||
# ~400ms we expect 1–2 progress events (well below the 9 heartbeats).
|
||||
assert len(progress) <= len(by_event)
|
||||
for ev in progress:
|
||||
assert ev.context.task_name == "heartbeat"
|
||||
assert ev.context.attempt == 1
|
||||
assert ev.context.idle_timeout_secs == 0.2
|
||||
assert ev.context.idle_timeout_secs == idle_timeout_s
|
||||
assert isinstance(ev.progress_at, datetime)
|
||||
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import (
|
||||
Any,
|
||||
Literal,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
get_type_hints,
|
||||
)
|
||||
@@ -125,6 +126,9 @@ Prompt = (
|
||||
| Runnable[StateSchema, LanguageModelInput]
|
||||
)
|
||||
|
||||
# A single hook or a list of hooks to be composed in order.
|
||||
HookLike = Union[RunnableLike, Sequence[RunnableLike]]
|
||||
|
||||
|
||||
def _get_state_value(state: StateSchema, key: str, default: Any = None) -> Any:
|
||||
return (
|
||||
@@ -134,6 +138,131 @@ def _get_state_value(state: StateSchema, key: str, default: Any = None) -> Any:
|
||||
)
|
||||
|
||||
|
||||
def _set_state_value(state: StateSchema, key: str, value: Any) -> None:
|
||||
"""Set a value in the state, supporting both dict and Pydantic model states."""
|
||||
if isinstance(state, dict):
|
||||
state[key] = value
|
||||
else:
|
||||
setattr(state, key, value)
|
||||
|
||||
|
||||
def _merge_state_update(state: StateSchema, update: dict) -> StateSchema:
|
||||
"""Return a shallow copy of *state* with *update* applied.
|
||||
|
||||
This is used when chaining multiple hooks: each hook receives the state
|
||||
as it would look after all previous hooks have run, so that hooks later
|
||||
in the chain can observe updates made by earlier ones.
|
||||
|
||||
Note: only simple key-level merging is performed here (no reducer logic).
|
||||
The full reducer logic is applied by the graph engine when the final
|
||||
combined update dict is written back to the state.
|
||||
"""
|
||||
if isinstance(state, dict):
|
||||
return {**state, **update} # type: ignore[return-value]
|
||||
else:
|
||||
# Pydantic / dataclass – make a shallow copy and patch fields
|
||||
try:
|
||||
merged = state.model_copy() # pydantic v2
|
||||
except AttributeError:
|
||||
merged = state.copy() # pydantic v1 / dataclass fallback
|
||||
for k, v in update.items():
|
||||
setattr(merged, k, v)
|
||||
return merged # type: ignore[return-value]
|
||||
|
||||
|
||||
def _coerce_to_runnable(hook: RunnableLike) -> RunnableCallable:
|
||||
"""Wrap a plain callable into a RunnableCallable if necessary."""
|
||||
if isinstance(hook, RunnableCallable):
|
||||
return hook
|
||||
if isinstance(hook, Runnable):
|
||||
# Already a Runnable – wrap so we get a uniform interface
|
||||
sync_fn = hook.invoke
|
||||
async_fn = hook.ainvoke
|
||||
return RunnableCallable(sync_fn, async_fn)
|
||||
if inspect.iscoroutinefunction(hook):
|
||||
return RunnableCallable(None, hook)
|
||||
if callable(hook):
|
||||
return RunnableCallable(hook)
|
||||
raise TypeError(f"Expected a callable or Runnable, got {type(hook)!r}")
|
||||
|
||||
|
||||
def _chain_hooks(hooks: Sequence[RunnableLike]) -> RunnableCallable:
|
||||
"""Compose multiple hook callables into a single hook.
|
||||
|
||||
Each hook is called in order. After each hook the returned update dict is
|
||||
merged into a running copy of the graph state so that subsequent hooks can
|
||||
observe the changes made by earlier ones. The accumulated update dict
|
||||
(union of all individual update dicts, with later hooks winning on key
|
||||
conflicts) is returned as the final state update.
|
||||
|
||||
Args:
|
||||
hooks: A sequence of :data:`RunnableLike` objects. Each must accept
|
||||
the graph state as its first positional argument and return a
|
||||
``dict`` of state updates.
|
||||
|
||||
Returns:
|
||||
A :class:`~langgraph._internal._runnable.RunnableCallable` that behaves
|
||||
like a single hook but applies all of *hooks* in sequence.
|
||||
"""
|
||||
if not hooks:
|
||||
raise ValueError("_chain_hooks requires at least one hook")
|
||||
if len(hooks) == 1:
|
||||
return _coerce_to_runnable(hooks[0])
|
||||
|
||||
runnables = [_coerce_to_runnable(h) for h in hooks]
|
||||
|
||||
def _sync_chained(state: Any, **kwargs: Any) -> dict:
|
||||
accumulated: dict = {}
|
||||
current_state = state
|
||||
for runnable in runnables:
|
||||
# Pass extra kwargs (e.g. config, store) through if the hook
|
||||
# accepts them; RunnableCallable handles introspection.
|
||||
update = runnable.invoke(current_state, **kwargs)
|
||||
if update:
|
||||
accumulated.update(update)
|
||||
current_state = _merge_state_update(current_state, update)
|
||||
return accumulated
|
||||
|
||||
async def _async_chained(state: Any, **kwargs: Any) -> dict:
|
||||
accumulated: dict = {}
|
||||
current_state = state
|
||||
for runnable in runnables:
|
||||
update = await runnable.ainvoke(current_state, **kwargs)
|
||||
if update:
|
||||
accumulated.update(update)
|
||||
current_state = _merge_state_update(current_state, update)
|
||||
return accumulated
|
||||
|
||||
return RunnableCallable(_sync_chained, _async_chained, name="chained_hooks")
|
||||
|
||||
|
||||
def _resolve_hook(hook: HookLike | None) -> RunnableLike | None:
|
||||
"""Normalise *hook* to a single ``RunnableLike`` (or ``None``).
|
||||
|
||||
* If *hook* is ``None`` → return ``None``.
|
||||
* If *hook* is already a ``RunnableLike`` → return it unchanged.
|
||||
* If *hook* is a non-empty :class:`~collections.abc.Sequence` of
|
||||
``RunnableLike`` → chain them with :func:`_chain_hooks`.
|
||||
"""
|
||||
if hook is None:
|
||||
return None
|
||||
# A Sequence[RunnableLike] but NOT a single Runnable/callable
|
||||
if (
|
||||
isinstance(hook, Sequence)
|
||||
and not isinstance(hook, str)
|
||||
and not isinstance(hook, Runnable)
|
||||
and not callable(hook)
|
||||
):
|
||||
hooks_list: list[RunnableLike] = list(hook)
|
||||
if not hooks_list:
|
||||
return None
|
||||
if len(hooks_list) == 1:
|
||||
return hooks_list[0]
|
||||
return _chain_hooks(hooks_list)
|
||||
# Single hook – return as-is
|
||||
return hook # type: ignore[return-value]
|
||||
|
||||
|
||||
def _get_prompt_runnable(prompt: Prompt | None) -> Runnable:
|
||||
prompt_runnable: Runnable
|
||||
if prompt is None:
|
||||
@@ -293,8 +422,8 @@ def create_react_agent(
|
||||
response_format: StructuredResponseSchema
|
||||
| tuple[str, StructuredResponseSchema]
|
||||
| None = None,
|
||||
pre_model_hook: RunnableLike | None = None,
|
||||
post_model_hook: RunnableLike | None = None,
|
||||
pre_model_hook: HookLike | None = None,
|
||||
post_model_hook: HookLike | None = None,
|
||||
state_schema: StateSchemaType | None = None,
|
||||
context_schema: type[Any] | None = None,
|
||||
checkpointer: Checkpointer | None = None,
|
||||
@@ -393,10 +522,21 @@ def create_react_agent(
|
||||
The graph will make a separate call to the LLM to generate the structured response after the agent loop is finished.
|
||||
This is not the only strategy to get structured responses, see more options in [this guide](https://langchain-ai.github.io/langgraph/how-tos/react-agent-structured-output/).
|
||||
|
||||
pre_model_hook: An optional node to add before the `agent` node (i.e., the node that calls the LLM).
|
||||
Useful for managing long message histories (e.g., message trimming, summarization, etc.).
|
||||
Pre-model hook must be a callable or a runnable that takes in current graph state and returns a state update in the form of
|
||||
```python
|
||||
pre_model_hook: An optional node (or list of nodes) to add before the
|
||||
``agent`` node (i.e., the node that calls the LLM). Useful for
|
||||
managing long message histories (e.g., message trimming,
|
||||
summarization, etc.) or for composing multiple pre-processing
|
||||
steps.
|
||||
|
||||
A single hook **or a list of hooks** may be provided. When a list
|
||||
is given the hooks are executed in order: each hook receives the
|
||||
graph state as updated by all preceding hooks, and the union of
|
||||
all their return dicts is applied to the graph state before the
|
||||
agent node runs.
|
||||
|
||||
Each hook must be a callable or a runnable that takes the current
|
||||
graph state and returns a state update::
|
||||
|
||||
# At least one of `messages` or `llm_input_messages` MUST be provided
|
||||
{
|
||||
# If provided, will UPDATE the `messages` in the state
|
||||
@@ -407,27 +547,63 @@ def create_react_agent(
|
||||
# Any other state keys that need to be propagated
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
!!! Important
|
||||
At least one of `messages` or `llm_input_messages` MUST be provided and will be used as an input to the `agent` node.
|
||||
The rest of the keys will be added to the graph state.
|
||||
At least one of `messages` or `llm_input_messages` MUST be
|
||||
provided (by at least one hook in the chain) and will be used
|
||||
as an input to the ``agent`` node. The rest of the keys will
|
||||
be added to the graph state.
|
||||
|
||||
!!! Warning
|
||||
If you are returning `messages` in the pre-model hook, you should OVERWRITE the `messages` key by doing the following:
|
||||
If you are returning `messages` in the pre-model hook, you
|
||||
should OVERWRITE the `messages` key::
|
||||
|
||||
{
|
||||
"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *new_messages]
|
||||
...
|
||||
}
|
||||
|
||||
!!! Example "Composing multiple pre-model hooks"
|
||||
|
||||
```python
|
||||
{
|
||||
"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *new_messages]
|
||||
...
|
||||
}
|
||||
from langchain_core.messages import RemoveMessage
|
||||
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||
|
||||
def trim_messages(state):
|
||||
# Keep only the last 10 messages
|
||||
return {
|
||||
"messages": [
|
||||
RemoveMessage(id=REMOVE_ALL_MESSAGES),
|
||||
*state["messages"][-10:],
|
||||
]
|
||||
}
|
||||
|
||||
def inject_system_prompt(state):
|
||||
return {
|
||||
"llm_input_messages": [
|
||||
SystemMessage("You are a helpful assistant."),
|
||||
*state["messages"],
|
||||
]
|
||||
}
|
||||
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
tools,
|
||||
pre_model_hook=[trim_messages, inject_system_prompt],
|
||||
)
|
||||
```
|
||||
post_model_hook: An optional node to add after the `agent` node (i.e., the node that calls the LLM).
|
||||
Useful for implementing human-in-the-loop, guardrails, validation, or other post-processing.
|
||||
Post-model hook must be a callable or a runnable that takes in current graph state and returns a state update.
|
||||
|
||||
post_model_hook: An optional node (or list of nodes) to add after the
|
||||
``agent`` node (i.e., the node that calls the LLM). Useful for
|
||||
implementing human-in-the-loop, guardrails, validation, or other
|
||||
post-processing steps.
|
||||
|
||||
Accepts the same single-hook-or-list-of-hooks form as
|
||||
``pre_model_hook``.
|
||||
|
||||
!!! Note
|
||||
Only available with `version="v2"`.
|
||||
Only available with ``version="v2"``.
|
||||
|
||||
state_schema: An optional state schema that defines graph state.
|
||||
Must have `messages` and `remaining_steps` keys.
|
||||
Defaults to `AgentState` that defines those two keys.
|
||||
@@ -551,6 +727,10 @@ def create_react_agent(
|
||||
else AgentState
|
||||
)
|
||||
|
||||
# Normalise hook arguments: a list of hooks is composed into a single hook.
|
||||
resolved_pre_model_hook: RunnableLike | None = _resolve_hook(pre_model_hook)
|
||||
resolved_post_model_hook: RunnableLike | None = _resolve_hook(post_model_hook)
|
||||
|
||||
llm_builtin_tools: list[dict] = []
|
||||
if isinstance(tools, ToolNode):
|
||||
tool_classes = list(tools.tools_by_name.values())
|
||||
@@ -634,7 +814,7 @@ def create_react_agent(
|
||||
return False
|
||||
|
||||
def _get_model_input_state(state: StateSchema) -> StateSchema:
|
||||
if pre_model_hook is not None:
|
||||
if resolved_pre_model_hook is not None:
|
||||
messages = (
|
||||
_get_state_value(state, "llm_input_messages")
|
||||
) or _get_state_value(state, "messages")
|
||||
@@ -721,7 +901,7 @@ def create_react_agent(
|
||||
return {"messages": [response]}
|
||||
|
||||
input_schema: StateSchemaType
|
||||
if pre_model_hook is not None:
|
||||
if resolved_pre_model_hook is not None:
|
||||
# Dynamically create a schema that inherits from state_schema and adds 'llm_input_messages'
|
||||
if isinstance(state_schema, type) and issubclass(state_schema, BaseModel):
|
||||
# For Pydantic schemas
|
||||
@@ -792,8 +972,8 @@ def create_react_agent(
|
||||
RunnableCallable(call_model, acall_model),
|
||||
input_schema=input_schema,
|
||||
)
|
||||
if pre_model_hook is not None:
|
||||
workflow.add_node("pre_model_hook", pre_model_hook) # type: ignore[arg-type]
|
||||
if resolved_pre_model_hook is not None:
|
||||
workflow.add_node("pre_model_hook", resolved_pre_model_hook) # type: ignore[arg-type]
|
||||
workflow.add_edge("pre_model_hook", "agent")
|
||||
entrypoint = "pre_model_hook"
|
||||
else:
|
||||
@@ -801,8 +981,8 @@ def create_react_agent(
|
||||
|
||||
workflow.set_entry_point(entrypoint)
|
||||
|
||||
if post_model_hook is not None:
|
||||
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
|
||||
if resolved_post_model_hook is not None:
|
||||
workflow.add_node("post_model_hook", resolved_post_model_hook) # type: ignore[arg-type]
|
||||
workflow.add_edge("agent", "post_model_hook")
|
||||
|
||||
if response_format is not None:
|
||||
@@ -813,7 +993,7 @@ def create_react_agent(
|
||||
agenerate_structured_response,
|
||||
),
|
||||
)
|
||||
if post_model_hook is not None:
|
||||
if resolved_post_model_hook is not None:
|
||||
workflow.add_edge("post_model_hook", "generate_structured_response")
|
||||
else:
|
||||
workflow.add_edge("agent", "generate_structured_response")
|
||||
@@ -833,7 +1013,7 @@ def create_react_agent(
|
||||
last_message = messages[-1]
|
||||
# If there is no function call, then we finish
|
||||
if not isinstance(last_message, AIMessage) or not last_message.tool_calls:
|
||||
if post_model_hook is not None:
|
||||
if resolved_post_model_hook is not None:
|
||||
return "post_model_hook"
|
||||
elif response_format is not None:
|
||||
return "generate_structured_response"
|
||||
@@ -844,7 +1024,7 @@ def create_react_agent(
|
||||
if version == "v1":
|
||||
return "tools"
|
||||
elif version == "v2":
|
||||
if post_model_hook is not None:
|
||||
if resolved_post_model_hook is not None:
|
||||
return "post_model_hook"
|
||||
return [
|
||||
Send(
|
||||
@@ -873,8 +1053,8 @@ def create_react_agent(
|
||||
|
||||
# Optionally add a pre-model hook node that will be called
|
||||
# every time before the "agent" (LLM-calling node)
|
||||
if pre_model_hook is not None:
|
||||
workflow.add_node("pre_model_hook", pre_model_hook) # type: ignore[arg-type]
|
||||
if resolved_pre_model_hook is not None:
|
||||
workflow.add_node("pre_model_hook", resolved_pre_model_hook) # type: ignore[arg-type]
|
||||
workflow.add_edge("pre_model_hook", "agent")
|
||||
entrypoint = "pre_model_hook"
|
||||
else:
|
||||
@@ -888,8 +1068,8 @@ def create_react_agent(
|
||||
post_model_hook_paths = [entrypoint, "tools"]
|
||||
|
||||
# Add a post model hook node if post_model_hook is provided
|
||||
if post_model_hook is not None:
|
||||
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
|
||||
if resolved_post_model_hook is not None:
|
||||
workflow.add_node("post_model_hook", resolved_post_model_hook) # type: ignore[arg-type]
|
||||
agent_paths.append("post_model_hook")
|
||||
workflow.add_edge("agent", "post_model_hook")
|
||||
else:
|
||||
@@ -904,17 +1084,17 @@ def create_react_agent(
|
||||
agenerate_structured_response,
|
||||
),
|
||||
)
|
||||
if post_model_hook is not None:
|
||||
if resolved_post_model_hook is not None:
|
||||
post_model_hook_paths.append("generate_structured_response")
|
||||
else:
|
||||
agent_paths.append("generate_structured_response")
|
||||
else:
|
||||
if post_model_hook is not None:
|
||||
if resolved_post_model_hook is not None:
|
||||
post_model_hook_paths.append(END)
|
||||
else:
|
||||
agent_paths.append(END)
|
||||
|
||||
if post_model_hook is not None:
|
||||
if resolved_post_model_hook is not None:
|
||||
|
||||
def post_model_hook_router(state: StateSchema) -> str | list[Send]:
|
||||
"""Route to the next node after post_model_hook.
|
||||
@@ -1012,4 +1192,5 @@ __all__ = [
|
||||
"AgentStatePydantic",
|
||||
"AgentStateWithStructuredResponse",
|
||||
"AgentStateWithStructuredResponsePydantic",
|
||||
"HookLike",
|
||||
]
|
||||
|
||||
@@ -18,6 +18,7 @@ from langgraph_sdk.schema import (
|
||||
CronSortBy,
|
||||
Durability,
|
||||
Input,
|
||||
Json,
|
||||
OnCompletionBehavior,
|
||||
QueryParamTypes,
|
||||
Run,
|
||||
@@ -413,6 +414,7 @@ class CronClient:
|
||||
assistant_id: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
enabled: bool | None = None,
|
||||
metadata: Json = None,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
sort_by: CronSortBy | None = None,
|
||||
@@ -427,6 +429,8 @@ class CronClient:
|
||||
assistant_id: The assistant ID or graph name to search for.
|
||||
thread_id: the thread ID to search for.
|
||||
enabled: The enabled status to search for.
|
||||
metadata: Metadata to filter by. Exact match filter for each KV pair.
|
||||
!!! version-added "Added in Agent Server version 0.9.0"
|
||||
limit: The maximum number of results to return.
|
||||
offset: The number of results to skip.
|
||||
headers: Optional custom headers to include with the request.
|
||||
@@ -481,6 +485,8 @@ class CronClient:
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
}
|
||||
if metadata:
|
||||
payload["metadata"] = metadata
|
||||
if sort_by:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
@@ -497,6 +503,7 @@ class CronClient:
|
||||
*,
|
||||
assistant_id: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
metadata: Json = None,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
) -> int:
|
||||
@@ -505,6 +512,8 @@ class CronClient:
|
||||
Args:
|
||||
assistant_id: Assistant ID to filter by.
|
||||
thread_id: Thread ID to filter by.
|
||||
metadata: Metadata to filter by. Exact match filter for each KV pair.
|
||||
!!! version-added "Added in Agent Server version 0.9.0"
|
||||
headers: Optional custom headers to include with the request.
|
||||
params: Optional query parameters to include with the request.
|
||||
|
||||
@@ -516,6 +525,8 @@ class CronClient:
|
||||
payload["assistant_id"] = assistant_id
|
||||
if thread_id:
|
||||
payload["thread_id"] = thread_id
|
||||
if metadata:
|
||||
payload["metadata"] = metadata
|
||||
return await self.http.post(
|
||||
"/runs/crons/count", json=payload, headers=headers, params=params
|
||||
)
|
||||
|
||||
@@ -18,6 +18,7 @@ from langgraph_sdk.schema import (
|
||||
CronSortBy,
|
||||
Durability,
|
||||
Input,
|
||||
Json,
|
||||
OnCompletionBehavior,
|
||||
QueryParamTypes,
|
||||
Run,
|
||||
@@ -402,6 +403,7 @@ class SyncCronClient:
|
||||
assistant_id: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
enabled: bool | None = None,
|
||||
metadata: Json = None,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
sort_by: CronSortBy | None = None,
|
||||
@@ -416,6 +418,8 @@ class SyncCronClient:
|
||||
assistant_id: The assistant ID or graph name to search for.
|
||||
thread_id: the thread ID to search for.
|
||||
enabled: Whether the cron job is enabled.
|
||||
metadata: Metadata to filter by. Exact match filter for each KV pair.
|
||||
!!! version-added "Added in Agent Server version 0.9.0"
|
||||
limit: The maximum number of results to return.
|
||||
offset: The number of results to skip.
|
||||
headers: Optional custom headers to include with the request.
|
||||
@@ -468,6 +472,8 @@ class SyncCronClient:
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
}
|
||||
if metadata:
|
||||
payload["metadata"] = metadata
|
||||
if sort_by:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
@@ -484,6 +490,7 @@ class SyncCronClient:
|
||||
*,
|
||||
assistant_id: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
metadata: Json = None,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
) -> int:
|
||||
@@ -492,6 +499,8 @@ class SyncCronClient:
|
||||
Args:
|
||||
assistant_id: Assistant ID to filter by.
|
||||
thread_id: Thread ID to filter by.
|
||||
metadata: Metadata to filter by. Exact match filter for each KV pair.
|
||||
!!! version-added "Added in Agent Server version 0.9.0"
|
||||
headers: Optional custom headers to include with the request.
|
||||
params: Optional query parameters to include with the request.
|
||||
|
||||
@@ -503,6 +512,8 @@ class SyncCronClient:
|
||||
payload["assistant_id"] = assistant_id
|
||||
if thread_id:
|
||||
payload["thread_id"] = thread_id
|
||||
if metadata:
|
||||
payload["metadata"] = metadata
|
||||
return self.http.post(
|
||||
"/runs/crons/count", json=payload, headers=headers, params=params
|
||||
)
|
||||
|
||||
@@ -485,3 +485,165 @@ def test_sync_update_with_enabled_parameter(enabled_value):
|
||||
)
|
||||
|
||||
assert result == cron
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_search_with_metadata():
|
||||
"""Test that CronClient.search forwards metadata in the request body."""
|
||||
cron = _cron_response()
|
||||
|
||||
async def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.method == "POST"
|
||||
assert request.url.path == "/runs/crons/search"
|
||||
|
||||
body = json.loads(request.content)
|
||||
assert body["metadata"] == {"owner": "alice"}
|
||||
assert body["limit"] == 10
|
||||
assert body["offset"] == 0
|
||||
|
||||
return httpx.Response(200, json=[cron])
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(
|
||||
transport=transport, base_url="https://example.com"
|
||||
) as client:
|
||||
http_client = HttpClient(client)
|
||||
cron_client = CronClient(http_client)
|
||||
result = await cron_client.search(metadata={"owner": "alice"})
|
||||
|
||||
assert result == [cron]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_search_omits_empty_metadata():
|
||||
"""Test that CronClient.search does not send metadata when not provided."""
|
||||
cron = _cron_response()
|
||||
|
||||
async def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = json.loads(request.content)
|
||||
assert "metadata" not in body
|
||||
return httpx.Response(200, json=[cron])
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(
|
||||
transport=transport, base_url="https://example.com"
|
||||
) as client:
|
||||
http_client = HttpClient(client)
|
||||
cron_client = CronClient(http_client)
|
||||
await cron_client.search()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_count_with_metadata():
|
||||
"""Test that CronClient.count forwards metadata in the request body."""
|
||||
|
||||
async def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.method == "POST"
|
||||
assert request.url.path == "/runs/crons/count"
|
||||
|
||||
body = json.loads(request.content)
|
||||
assert body["metadata"] == {"team": "infra"}
|
||||
|
||||
return httpx.Response(200, json=2)
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(
|
||||
transport=transport, base_url="https://example.com"
|
||||
) as client:
|
||||
http_client = HttpClient(client)
|
||||
cron_client = CronClient(http_client)
|
||||
result = await cron_client.count(metadata={"team": "infra"})
|
||||
|
||||
assert result == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_count_omits_empty_metadata():
|
||||
"""Test that CronClient.count does not send metadata when not provided."""
|
||||
|
||||
async def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = json.loads(request.content)
|
||||
assert "metadata" not in body
|
||||
return httpx.Response(200, json=0)
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(
|
||||
transport=transport, base_url="https://example.com"
|
||||
) as client:
|
||||
http_client = HttpClient(client)
|
||||
cron_client = CronClient(http_client)
|
||||
await cron_client.count()
|
||||
|
||||
|
||||
def test_sync_search_with_metadata():
|
||||
"""Test that SyncCronClient.search forwards metadata in the request body."""
|
||||
cron = _cron_response()
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.method == "POST"
|
||||
assert request.url.path == "/runs/crons/search"
|
||||
|
||||
body = json.loads(request.content)
|
||||
assert body["metadata"] == {"owner": "alice"}
|
||||
|
||||
return httpx.Response(200, json=[cron])
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
||||
http_client = SyncHttpClient(client)
|
||||
cron_client = SyncCronClient(http_client)
|
||||
result = cron_client.search(metadata={"owner": "alice"})
|
||||
|
||||
assert result == [cron]
|
||||
|
||||
|
||||
def test_sync_search_omits_empty_metadata():
|
||||
"""Test that SyncCronClient.search does not send metadata when not provided."""
|
||||
cron = _cron_response()
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = json.loads(request.content)
|
||||
assert "metadata" not in body
|
||||
return httpx.Response(200, json=[cron])
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
||||
http_client = SyncHttpClient(client)
|
||||
cron_client = SyncCronClient(http_client)
|
||||
cron_client.search()
|
||||
|
||||
|
||||
def test_sync_count_with_metadata():
|
||||
"""Test that SyncCronClient.count forwards metadata in the request body."""
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
assert request.method == "POST"
|
||||
assert request.url.path == "/runs/crons/count"
|
||||
|
||||
body = json.loads(request.content)
|
||||
assert body["metadata"] == {"team": "infra"}
|
||||
|
||||
return httpx.Response(200, json=2)
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
||||
http_client = SyncHttpClient(client)
|
||||
cron_client = SyncCronClient(http_client)
|
||||
result = cron_client.count(metadata={"team": "infra"})
|
||||
|
||||
assert result == 2
|
||||
|
||||
|
||||
def test_sync_count_omits_empty_metadata():
|
||||
"""Test that SyncCronClient.count does not send metadata when not provided."""
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = json.loads(request.content)
|
||||
assert "metadata" not in body
|
||||
return httpx.Response(200, json=0)
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
||||
http_client = SyncHttpClient(client)
|
||||
cron_client = SyncCronClient(http_client)
|
||||
cron_client.count()
|
||||
|
||||
Reference in New Issue
Block a user