Compare commits

..
Author SHA1 Message Date
Quanzheng Long 9800e80296 add_test 2026-04-28 11:19:23 -07:00
Sydney Runkle 0ae81f3cff format, lint, restructure 2026-04-24 07:46:57 -04:00
10 changed files with 269 additions and 606 deletions
+40 -74
View File
@@ -184,97 +184,63 @@ def add_messages(
```
"""
# 1. Coerce scalars to lists.
remove_all_idx = None
# coerce to list
if not isinstance(left, list):
left = [left] # type: ignore[assignment]
if not isinstance(right, list):
right = [right] # type: ignore[assignment]
left = cast(list, left)
# 2. Normalize `left`. After the first call, `left` is the previous return
# value of `add_messages` — a list of fully-resolved BaseMessages with
# IDs — and needs no work. Fresh user input (dicts, tuples, message
# chunks, BaseMessages without IDs) falls through to full conversion.
left_msgs: list[BaseMessage]
if (
left
and isinstance(left[0], BaseMessage)
and not isinstance(left[0], BaseMessageChunk)
and left[0].id is not None
):
left_msgs = left
else:
left_msgs = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(left)
]
for m in left_msgs:
if m.id is None:
m.id = str(uuid.uuid4())
# 3. Normalize `right` — always fresh input. Assign missing IDs and detect
# any RemoveMessage sentinels in a single pass.
right_msgs: list[BaseMessage] = [
# coerce to message
left = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(left)
]
right = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(right)
]
remove_all_idx: int | None = None
has_remove = False
for idx, m in enumerate(right_msgs):
# assign missing ids
for m in left:
if m.id is None:
m.id = str(uuid.uuid4())
if isinstance(m, RemoveMessage):
has_remove = True
if m.id == REMOVE_ALL_MESSAGES:
remove_all_idx = idx
for idx, m in enumerate(right):
if m.id is None:
m.id = str(uuid.uuid4())
if isinstance(m, RemoveMessage) and m.id == REMOVE_ALL_MESSAGES:
remove_all_idx = idx
# 4. REMOVE_ALL_MESSAGES: discard everything up to and including the sentinel.
if remove_all_idx is not None:
return right_msgs[remove_all_idx + 1 :]
return right[remove_all_idx + 1 :]
# 5. Decide fast vs. slow path. The fast path (pure append) is only valid
# when `right` has no removals, no intra-right duplicate IDs, and no IDs
# that overlap with `left` — any of those would force the indexed merge
# below to update or dedup.
pure_append = False
if not has_remove:
left_ids = {m.id for m in left_msgs}
right_ids = {m.id for m in right_msgs}
pure_append = len(right_ids) == len(right_msgs) and right_ids.isdisjoint(
left_ids
)
if pure_append:
merged = left_msgs + right_msgs
else:
# 6. Slow path: build id→index map over `left`, then replay `right`.
# In-place replacement for matching IDs, append for new IDs, and a
# deferred removal pass so RemoveMessages can target either side.
merged = left_msgs.copy()
merged_by_id = {m.id: i for i, m in enumerate(merged)}
ids_to_remove = set()
for m in right_msgs:
if (existing_idx := merged_by_id.get(m.id)) is not None:
if isinstance(m, RemoveMessage):
ids_to_remove.add(m.id)
else:
ids_to_remove.discard(m.id)
merged[existing_idx] = m
# merge
merged = left.copy()
merged_by_id = {m.id: i for i, m in enumerate(merged)}
ids_to_remove = set()
for m in right:
if (existing_idx := merged_by_id.get(m.id)) is not None:
if isinstance(m, RemoveMessage):
ids_to_remove.add(m.id)
else:
if isinstance(m, RemoveMessage):
raise ValueError(
f"Attempting to delete a message with an ID that doesn't exist ('{m.id}')"
)
merged_by_id[m.id] = len(merged)
merged.append(m)
merged = [m for m in merged if m.id not in ids_to_remove]
ids_to_remove.discard(m.id)
merged[existing_idx] = m
else:
if isinstance(m, RemoveMessage):
raise ValueError(
f"Attempting to delete a message with an ID that doesn't exist ('{m.id}')"
)
merged_by_id[m.id] = len(merged)
merged.append(m)
merged = [m for m in merged if m.id not in ids_to_remove]
# 7. Apply optional output format.
if format == "langchain-openai":
return _format_messages(merged)
if format:
merged = _format_messages(merged)
elif format:
msg = f"Unrecognized {format=}. Expected one of 'langchain-openai', None."
raise ValueError(msg)
else:
pass
return merged
+1 -1
View File
@@ -47,9 +47,9 @@ from langgraph._internal._fields import (
from langgraph._internal._pydantic import create_model
from langgraph._internal._runnable import coerce_to_runnable
from langgraph._internal._typing import EMPTY_SEQ, MISSING, DeprecatedKwargs
from langgraph.channels._delta import DeltaChannel
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import BinaryOperatorAggregate, _strip_extras
from langgraph.channels._delta import DeltaChannel
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue, LastValueAfterFinish
from langgraph.channels.named_barrier_value import (
@@ -8,8 +8,8 @@ from langgraph.checkpoint.base import DELTA_SENTINEL, BaseCheckpointSaver, Check
from langgraph.checkpoint.base.id import uuid6
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.channels._delta import DeltaChannel
from langgraph.channels.base import BaseChannel
from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec
LATEST_VERSION = 4
@@ -1,290 +0,0 @@
"""Benchmark: add_messages fast-path optimizations.
Both implementations are inlined so the benchmark is self-contained and
immune to import-cache or installed-vs-local confusion.
Run directly:
python tests/test_add_messages_benchmark.py
Or via pytest (correctness only, numbers printed to stdout):
pytest tests/test_add_messages_benchmark.py -s -v
"""
import statistics
import time
import tracemalloc
import uuid
from typing import cast
from langchain_core.messages import (
AIMessage,
BaseMessage,
BaseMessageChunk,
HumanMessage,
RemoveMessage,
convert_to_messages,
message_chunk_to_message,
)
from langgraph.graph.message import REMOVE_ALL_MESSAGES
# ── original implementation (pre-optimisation) ────────────────────────────────
def _add_messages_original(left, right):
remove_all_idx = None
if not isinstance(left, list):
left = [left]
if not isinstance(right, list):
right = [right]
left = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(left)
]
right = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(right)
]
for m in left:
if m.id is None:
m.id = str(uuid.uuid4())
for idx, m in enumerate(right):
if m.id is None:
m.id = str(uuid.uuid4())
if isinstance(m, RemoveMessage) and m.id == REMOVE_ALL_MESSAGES:
remove_all_idx = idx
if remove_all_idx is not None:
return right[remove_all_idx + 1 :]
merged = left.copy()
merged_by_id = {m.id: i for i, m in enumerate(merged)}
ids_to_remove = set()
for m in right:
if (existing_idx := merged_by_id.get(m.id)) is not None:
if isinstance(m, RemoveMessage):
ids_to_remove.add(m.id)
else:
ids_to_remove.discard(m.id)
merged[existing_idx] = m
else:
if isinstance(m, RemoveMessage):
raise ValueError(
f"Attempting to delete a message with an ID that doesn't exist ('{m.id}')"
)
merged_by_id[m.id] = len(merged)
merged.append(m)
return [m for m in merged if m.id not in ids_to_remove]
# ── optimised implementation ──────────────────────────────────────────────────
def _add_messages_optimized(left, right):
if not isinstance(left, list):
left = [left]
if not isinstance(right, list):
right = [right]
# Optimisation 1: skip conversion + ID assignment on left when it already
# contains fully-resolved BaseMessage objects (the common case after the
# first call, since add_messages always returns list[BaseMessage] with IDs).
if (
left
and isinstance(left[0], BaseMessage)
and not isinstance(left[0], BaseMessageChunk)
):
left = cast(list[BaseMessage], left)
else:
left = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(left)
]
for m in left:
if m.id is None:
m.id = str(uuid.uuid4())
# always normalise right — it's fresh external input
right = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(right)
]
remove_all_idx = None
has_remove = False
for idx, m in enumerate(right):
if m.id is None:
m.id = str(uuid.uuid4())
if isinstance(m, RemoveMessage):
has_remove = True
if m.id == REMOVE_ALL_MESSAGES:
remove_all_idx = idx
if remove_all_idx is not None:
return right[remove_all_idx + 1 :]
# Optimisation 2: pure-append fast path — no removals and no ID overlaps.
# Builds one set over left instead of copying left + building a full dict.
if not has_remove:
left_ids = {m.id for m in left}
if not any(m.id in left_ids for m in right):
return left + right
# slow path: updates or removals present — full indexed merge
merged = left.copy()
merged_by_id = {m.id: i for i, m in enumerate(merged)}
ids_to_remove = set()
for m in right:
if (existing_idx := merged_by_id.get(m.id)) is not None:
if isinstance(m, RemoveMessage):
ids_to_remove.add(m.id)
else:
ids_to_remove.discard(m.id)
merged[existing_idx] = m
else:
if isinstance(m, RemoveMessage):
raise ValueError(
f"Attempting to delete a message with an ID that doesn't exist ('{m.id}')"
)
merged_by_id[m.id] = len(merged)
merged.append(m)
return [m for m in merged if m.id not in ids_to_remove]
# ── helpers ───────────────────────────────────────────────────────────────────
def _make_messages(n: int) -> list[BaseMessage]:
return [
(HumanMessage if i % 2 == 0 else AIMessage)(
content=f"message {i}", id=str(uuid.uuid4())
)
for i in range(n)
]
def _bench_time(fn, left, right, *, iters: int = 2_000) -> float:
"""Return median latency in microseconds."""
for _ in range(100):
fn(list(left), list(right))
times = []
for _ in range(iters):
left_copy, right_copy = list(left), list(right)
t0 = time.perf_counter()
fn(left_copy, right_copy)
times.append(time.perf_counter() - t0)
return statistics.median(times) * 1e6
def _bench_memory(fn, left, right) -> int:
"""Return peak memory allocated during a single call (bytes)."""
# one warm-up so any lazy init is excluded
fn(list(left), list(right))
left_copy, right_copy = list(left), list(right)
tracemalloc.start()
tracemalloc.clear_traces()
fn(left_copy, right_copy)
_, peak = tracemalloc.get_traced_memory()
tracemalloc.stop()
return peak
# ── scenarios ─────────────────────────────────────────────────────────────────
SCENARIOS = [
("pure append 1 → 1 msg", 1, 1, "append"),
("pure append 10 → 1 msg", 10, 1, "append"),
("pure append 100 → 1 msg", 100, 1, "append"),
("pure append 1000 → 1 msg", 1000, 1, "append"),
("pure append 1000 → 5 msgs", 1000, 5, "append"),
("update existing 100 → 1 msg", 100, 1, "update"),
("remove message 100 → 1 msg", 100, 1, "remove"),
]
def _make_inputs(n_left, n_right, mode):
left = _make_messages(n_left)
right = _make_messages(n_right)
if mode == "update":
right[0] = AIMessage(content="updated", id=left[0].id)
elif mode == "remove":
right = [RemoveMessage(id=left[0].id)]
return left, right
# ── main output ───────────────────────────────────────────────────────────────
COL = 36
def run_benchmarks() -> None:
print()
print("=" * 88)
print("add_messages benchmark — time (µs, median of 2 000 iterations)")
print("=" * 88)
print(f"{'Scenario':<{COL}} {'Original':>10} {'Optimized':>11} {'Speedup':>8}")
print("-" * 88)
for label, n_left, n_right, mode in SCENARIOS:
left, right = _make_inputs(n_left, n_right, mode)
t_orig = _bench_time(_add_messages_original, left, right)
t_opt = _bench_time(_add_messages_optimized, left, right)
print(f"{label:<{COL}} {t_orig:>10.2f} {t_opt:>11.2f} {t_orig / t_opt:>7.2f}x")
print()
print("=" * 88)
print("add_messages benchmark — peak memory allocated per call (bytes)")
print("=" * 88)
print(f"{'Scenario':<{COL}} {'Original':>10} {'Optimized':>11} {'Reduction':>10}")
print("-" * 88)
for label, n_left, n_right, mode in SCENARIOS:
left, right = _make_inputs(n_left, n_right, mode)
m_orig = _bench_memory(_add_messages_original, left, right)
m_opt = _bench_memory(_add_messages_optimized, left, right)
reduction = (1 - m_opt / m_orig) * 100 if m_orig else 0.0
print(f"{label:<{COL}} {m_orig:>10,} {m_opt:>11,} {reduction:>9.1f}%")
print()
print("=" * 88)
print("Simulated long thread — 200 steps × 2 msgs appended per step")
print("=" * 88)
for name, fn in [
("original", _add_messages_original),
("optimized", _add_messages_optimized),
]:
state: list = []
t0 = time.perf_counter()
for step in range(200):
new_msgs = [
HumanMessage(content=f"step {step} human", id=str(uuid.uuid4())),
AIMessage(content=f"step {step} ai", id=str(uuid.uuid4())),
]
state = fn(state, new_msgs)
elapsed = (time.perf_counter() - t0) * 1_000
print(f" {name:<12} {elapsed:.2f} ms ({len(state)} messages)")
print()
# ── pytest entry-points ───────────────────────────────────────────────────────
def test_add_messages_correctness():
"""Optimised implementation must match original output for every scenario."""
for label, n_left, n_right, mode in SCENARIOS:
left, right = _make_inputs(n_left, n_right, mode)
expected = _add_messages_original(list(left), list(right))
actual = _add_messages_optimized(list(left), list(right))
assert len(actual) == len(expected), f"[{label}] length mismatch"
for a, b in zip(actual, expected):
assert type(a) is type(b), f"[{label}] type mismatch"
assert a.id == b.id, f"[{label}] id mismatch"
assert a.content == b.content, f"[{label}] content mismatch"
def test_add_messages_benchmark(capsys):
run_benchmarks()
out = capsys.readouterr().out
assert "Speedup" in out
assert "Optimized" in out
if __name__ == "__main__":
run_benchmarks()
+1 -1
View File
@@ -6,8 +6,8 @@ from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.base import DELTA_SENTINEL
from langgraph._internal._typing import MISSING
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels._delta import DeltaChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.last_value import LastValue
from langgraph.channels.topic import Topic
from langgraph.channels.untracked_value import UntrackedValue
@@ -49,8 +49,8 @@ import pytest
from langgraph.checkpoint.memory import InMemorySaver
from typing_extensions import TypedDict
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels._delta import DeltaChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.graph import END, START, StateGraph
pytestmark = pytest.mark.anyio
-118
View File
@@ -5,7 +5,6 @@ import langchain_core
import pytest
from langchain_core.messages import (
AIMessage,
AIMessageChunk,
AnyMessage,
HumanMessage,
RemoveMessage,
@@ -339,123 +338,6 @@ def test_remove_all_messages():
]
def test_fast_path_preserves_format_openai():
"""Pure-append fast path must still apply the `langchain-openai` formatter."""
left = [HumanMessage(content="prior", id="1")]
right = [
AIMessage(
content=[
{
"type": "tool_use",
"name": "foo",
"input": {"bar": "baz"},
"id": "t1",
}
],
id="2",
)
]
result = add_messages(left, right, format="langchain-openai")
assert isinstance(result[0], HumanMessage)
assert result[0].content == "prior"
assert isinstance(result[1], AIMessage)
# formatter collapses the tool_use content block into `tool_calls`
assert result[1].content == ""
assert len(result[1].tool_calls) == 1
assert result[1].tool_calls[0]["name"] == "foo"
assert result[1].tool_calls[0]["args"] == {"bar": "baz"}
assert result[1].tool_calls[0]["id"] == "t1"
def test_fast_path_rejects_invalid_format():
"""Pure-append fast path must validate the `format` arg like the slow path."""
left = [HumanMessage(content="prior", id="1")]
right = [AIMessage(content="new", id="2")]
with pytest.raises(ValueError, match="Unrecognized format="):
add_messages(left, right, format="bogus") # type: ignore[arg-type]
def test_left_starting_with_chunk_is_normalized():
"""Opt-1 guard: a `BaseMessageChunk` at left[0] must trigger full conversion."""
chunk = AIMessageChunk(content="chunk", id="c1")
result = add_messages([chunk], [HumanMessage(content="h", id="h1")])
assert len(result) == 2
# chunk must be converted to a non-chunk message
assert type(result[0]).__name__ == "AIMessage"
assert result[0].id == "c1"
assert result[1].id == "h1"
def test_left_as_dicts_is_normalized():
"""Opt-1 guard: dicts at left[0] must trigger full conversion."""
left = [{"role": "user", "content": "hi", "id": "d1"}]
right = [AIMessage(content="reply", id="a1")]
result = add_messages(left, right)
assert len(result) == 2
assert isinstance(result[0], HumanMessage)
assert result[0].id == "d1"
assert result[0].content == "hi"
def test_left_as_tuples_is_normalized():
"""Opt-1 guard: tuple-form messages must trigger full conversion."""
left = [("user", "hi")]
right = [AIMessage(content="reply", id="a1")]
result = add_messages(left, right)
assert len(result) == 2
assert isinstance(result[0], HumanMessage)
# id is auto-assigned
assert isinstance(result[0].id, str) and UUID(result[0].id, version=4)
def test_left_first_msg_missing_id_is_normalized():
"""Opt-1 guard: a BaseMessage without an id at left[0] falls to the else branch."""
left = [HumanMessage(content="hi")] # no id
right = [AIMessage(content="reply", id="a1")]
result = add_messages(left, right)
assert len(result) == 2
# left's id must have been auto-assigned
assert isinstance(result[0].id, str) and UUID(result[0].id, version=4)
def test_duplicate_ids_in_right_with_nonempty_left():
"""Opt-2 guard: intra-right duplicate ids must take slow path (dedup kept)."""
left = [HumanMessage(content="prior", id="1")]
right = [
AIMessage(content="first", id="2"),
AIMessage(content="second", id="2"),
]
result = add_messages(left, right)
assert len(result) == 2
assert result[0].id == "1"
assert result[1].id == "2"
assert result[1].content == "second"
def test_right_with_none_ids_pure_append():
"""Fast path still correct when right entries start with id=None (fresh uuids assigned)."""
left = [HumanMessage(content="prior", id="1")]
right = [AIMessage(content="a"), AIMessage(content="b")]
result = add_messages(left, right)
assert len(result) == 3
assert result[0].id == "1"
for m in result[1:]:
assert isinstance(m.id, str) and UUID(m.id, version=4)
# fresh uuids must be distinct
assert result[1].id != result[2].id
def test_fast_path_returns_fresh_list():
"""Fast path must return a new list object (not mutate or alias left)."""
left = [HumanMessage(content="prior", id="1")]
right = [AIMessage(content="new", id="2")]
result = add_messages(left, right)
assert result is not left
# left must be untouched
assert len(left) == 1
assert left[0].id == "1"
def test_push_messages_in_graph():
class MessagesState(TypedDict):
messages: Annotated[list[AnyMessage], add_messages]
+8 -28
View File
@@ -82,7 +82,6 @@ from langchain_core.tools.base import (
_is_injected_arg_type,
get_all_basemodel_annotations,
)
from langgraph._internal._constants import CONF, CONFIG_KEY_READ
from langgraph._internal._runnable import RunnableCallable
from langgraph.errors import GraphBubbleUp
from langgraph.graph.message import REMOVE_ALL_MESSAGES
@@ -801,7 +800,7 @@ class ToolNode(RunnableCallable):
# Construct ToolRuntime instances at the top level for each tool call
tool_runtimes = []
for call, cfg in zip(tool_calls, config_list, strict=False):
state = self._extract_state(input, cfg)
state = self._extract_state(input)
tool_runtime = ToolRuntime(
state=state,
tool_call_id=call["id"],
@@ -836,7 +835,7 @@ class ToolNode(RunnableCallable):
# Construct ToolRuntime instances at the top level for each tool call
tool_runtimes = []
for call, cfg in zip(tool_calls, config_list, strict=False):
state = self._extract_state(input, cfg)
state = self._extract_state(input)
tool_runtime = ToolRuntime(
state=state,
tool_call_id=call["id"],
@@ -1274,37 +1273,18 @@ class ToolNode(RunnableCallable):
return None
def _extract_state(
self,
input: list[AnyMessage] | dict[str, Any] | BaseModel,
config: RunnableConfig,
self, input: list[AnyMessage] | dict[str, Any] | BaseModel
) -> list[AnyMessage] | dict[str, Any] | BaseModel:
"""Extract state from input.
"""Extract state from input, handling ToolCallWithContext if present.
Three input shapes:
Args:
input: The input which may be raw state or ToolCallWithContext.
- `ToolCallWithContext` dict — legacy Send payload carrying an inlined
state snapshot; return `input["state"]`.
- list of `ToolCall` dicts — new Send payload with no inlined state;
hydrate state from channels via `CONFIG_KEY_READ`.
- regular graph state (dict/list/BaseModel) — return `input` as-is.
Returns:
The actual state to pass to wrap_tool_call wrappers.
"""
if isinstance(input, dict) and input.get("__type") == "tool_call_with_context":
return input["state"]
if (
isinstance(input, list)
and input
and isinstance(input[-1], dict)
and input[-1].get("type") == "tool_call"
):
read = config.get(CONF, {}).get(CONFIG_KEY_READ)
if read is None:
return {}
# Pregel installs CONFIG_KEY_READ as
# `functools.partial(local_read, scratchpad, channels, managed, task)`.
# Match the previous inlined-state contract by reading channels only;
# managed values have their own injection path (`ToolRuntime.context`).
channels = read.args[1]
return cast("dict[str, Any]", read(list(channels), False))
return input
def _inject_tool_args(
-92
View File
@@ -1320,98 +1320,6 @@ async def test_state_extraction_with_tool_call_with_context_async() -> None:
assert "tool_call" not in state_seen[0]
def _config_with_channel_read(
channel_values: dict[str, object],
store: BaseStore | None = None,
) -> RunnableConfig:
"""Build a config that mimics `CONFIG_KEY_READ` as Pregel installs it.
Pregel always installs a `functools.partial(local_read, scratchpad,
channels, managed, task)`, and `ToolNode` introspects that partial to
learn channel names. The stub matches the shape: partial whose second and
third positional args are `channels` and `managed` mappings.
"""
import functools
channels_stub = {k: None for k in channel_values}
managed_stub: dict[str, object] = {}
# Shape matches pregel's real partial:
# functools.partial(local_read, scratchpad, channels, managed, task)
def _read(scratchpad, channels, managed, task, select, fresh): # noqa: ARG001
if isinstance(select, str):
return channel_values[select]
return {k: channel_values[k] for k in select if k in channel_values}
read = functools.partial(_read, None, channels_stub, managed_stub, None)
cfg = _create_config_with_runtime(store)
cfg["configurable"]["__pregel_read"] = read
return cfg
def test_list_form_send_hydrates_state_from_channel_read() -> None:
"""Send('tools', [tool_call]) with no inlined state should hydrate
ToolRuntime.state from CONFIG_KEY_READ (full state read)."""
state_seen = []
def state_inspector_handler(
request: ToolCallRequest,
execute: Callable[[ToolCallRequest], ToolMessage | Command],
) -> ToolMessage | Command:
state_seen.append(request.state)
return execute(request)
channel_values = {
"messages": [AIMessage("from channels")],
"files": {"/a.md": "body"},
}
tool_node = ToolNode([add], wrap_tool_call=state_inspector_handler)
tool_call: ToolCall = {
"name": "add",
"args": {"a": 1, "b": 2},
"id": "call_1",
"type": "tool_call",
}
tool_node.invoke([tool_call], config=_config_with_channel_read(channel_values))
assert len(state_seen) == 1
got = state_seen[0]
assert got == channel_values
assert "messages" in got and "files" in got
async def test_list_form_send_hydrates_state_async() -> None:
state_seen = []
def state_inspector_handler(
request: ToolCallRequest,
execute: Callable[[ToolCallRequest], ToolMessage | Command],
) -> ToolMessage | Command:
state_seen.append(request.state)
return execute(request)
channel_values = {"messages": [AIMessage("from channels")], "files": {}}
tool_node = ToolNode([add], wrap_tool_call=state_inspector_handler)
tool_call: ToolCall = {
"name": "add",
"args": {"a": 1, "b": 2},
"id": "call_1",
"type": "tool_call",
}
await tool_node.ainvoke(
[tool_call], config=_config_with_channel_read(channel_values)
)
assert len(state_seen) == 1
assert state_seen[0] == channel_values
def test_tool_call_request_is_frozen() -> None:
"""Test that ToolCallRequest raises deprecation warnings on direct attribute reassignment."""
tool_call: ToolCall = {"name": "add", "args": {"a": 1, "b": 2}, "id": "call_1"}
+217
View File
@@ -0,0 +1,217 @@
"""Verify that SEND values land in checkpoint_blobs (TASKS channel),
not just in checkpoint_writes.
Strategy:
1. Build a tiny graph: node "1" emits a Send to node "2" via
conditional edges. interrupt_before=["2"] freezes the loop after
super-step 1 completes (cp_1 persisted) but before "2" runs.
2. Run graph.invoke against PostgresSaver.
3. Read checkpoint_blobs raw via psycopg.
4. Find the row with channel = '__pregel_tasks' and decode the blob
using the saver's serde. Assert it contains a non-empty list of
Send objects.
5. As a control, also dump checkpoint_writes to confirm the *same*
value is independently present in the writes table.
"""
from __future__ import annotations
import operator
from typing import Annotated
from typing_extensions import TypedDict
from uuid import uuid4
from psycopg import Connection
from psycopg.rows import dict_row
from langgraph.checkpoint.postgres import PostgresSaver
from langgraph.checkpoint.serde.types import TASKS
from langgraph.graph import START, StateGraph
from langgraph.types import Send
DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5441/"
class State(TypedDict):
history: Annotated[list[str], operator.add]
def make_graph(checkpointer: PostgresSaver, *, with_interrupt: bool = True):
def node_one(state: State) -> State:
return {"history": ["1"]}
def node_two(state: dict) -> State:
return {"history": [f"2:{state['payload']}"]}
def fanout(state: State):
return [Send("two", {"payload": x}) for x in ("a", "b", "c")]
builder = StateGraph(State)
builder.add_node("one", node_one)
builder.add_node("two", node_two)
builder.add_edge(START, "one")
builder.add_conditional_edges("one", fanout, ["two"])
return builder.compile(
checkpointer=checkpointer,
interrupt_before=["two"] if with_interrupt else [],
)
def main() -> None:
database = f"verify_{uuid4().hex[:12]}"
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"CREATE DATABASE {database}")
try:
uri = DEFAULT_POSTGRES_URI + database
with Connection.connect(
uri, autocommit=True, prepare_threshold=0, row_factory=dict_row
) as conn:
saver = PostgresSaver(conn)
saver.setup()
graph = make_graph(saver)
thread_id = "t1"
cfg = {"configurable": {"thread_id": thread_id}}
result = graph.invoke({"history": []}, cfg)
print("=== invoke result (interrupted before 'two') ===")
print(result)
print()
# Inspect raw checkpoint_blobs.
rows = conn.execute(
"""
SELECT thread_id, checkpoint_ns, channel, version, type,
octet_length(blob) AS nbytes, blob
FROM checkpoint_blobs
WHERE thread_id = %s
ORDER BY channel, version
""",
(thread_id,),
).fetchall()
print(f"=== checkpoint_blobs ({len(rows)} rows) ===")
for r in rows:
preview = (
"(empty)"
if r["blob"] is None
else f"{r['blob'][:60]!r}{'...' if len(r['blob']) > 60 else ''}"
)
print(
f" channel={r['channel']!r:30s} v={r['version']!s:18s} "
f"type={r['type']:8s} nbytes={r['nbytes']!s:>5} blob={preview}"
)
print()
tasks_rows = [r for r in rows if r["channel"] == TASKS]
assert tasks_rows, (
f"expected at least one row in checkpoint_blobs for {TASKS!r}, "
f"got channels={[r['channel'] for r in rows]}"
)
# Decode the latest TASKS blob using the saver's serde.
tasks_row = tasks_rows[-1]
decoded = saver.serde.loads_typed((tasks_row["type"], tasks_row["blob"]))
print(f"=== decoded {TASKS!r} blob ({tasks_row['type']}) ===")
print(f" type(decoded) = {type(decoded).__name__}")
print(f" len(decoded) = {len(decoded)}")
for i, v in enumerate(decoded):
print(f" [{i}] type={type(v).__name__} value={v!r}")
print()
assert isinstance(decoded, list), (
f"expected list of Sends, got {type(decoded).__name__}"
)
assert len(decoded) == 3, f"expected 3 Sends (a, b, c), got {len(decoded)}"
for v in decoded:
assert isinstance(v, Send), (
f"expected Send instance in blob, got {type(v).__name__}"
)
assert v.node == "two"
print("✅ ASSERTION PASSED: checkpoint_blobs contains the Sends.")
print()
# Control: confirm the same Sends are also in checkpoint_writes.
wrows = conn.execute(
"""
SELECT task_id, idx, channel, type, octet_length(blob) AS nbytes, blob
FROM checkpoint_writes
WHERE thread_id = %s AND channel = %s
ORDER BY task_id, idx
""",
(thread_id, TASKS),
).fetchall()
print(
f"=== checkpoint_writes (channel={TASKS!r}, {len(wrows)} rows) ==="
)
for r in wrows:
v = saver.serde.loads_typed((r["type"], r["blob"]))
print(
f" task={r['task_id']} idx={r['idx']} type={r['type']} "
f"nbytes={r['nbytes']} -> {v!r}"
)
assert len(wrows) == 3, (
f"expected 3 writes-table rows for {TASKS!r}, got {len(wrows)}"
)
print(
"✅ ASSERTION PASSED: checkpoint_writes also contains the Sends."
)
print()
print(
"Conclusion (cp_1): SEND is double-stored. blobs row carries "
"the authoritative Send list; writes rows carry the same "
"payload as a redundant per-task record."
)
# ----------------------------------------------------------------
# Now resume: let 'two' actually run, so the Sends get consumed.
# We expect the next checkpoint's TASKS blob to be cleared,
# proving the accumulate=False lifecycle.
# ----------------------------------------------------------------
graph2 = make_graph(saver, with_interrupt=False)
result2 = graph2.invoke(None, cfg)
print()
print("=== invoke result (resumed; 'two' fanout completed) ===")
print(result2)
print()
rows2 = conn.execute(
"""
SELECT channel, version, type, octet_length(blob) AS nbytes, blob
FROM checkpoint_blobs
WHERE thread_id = %s AND channel = %s
ORDER BY version
""",
(thread_id, TASKS),
).fetchall()
print(f"=== checkpoint_blobs rows for {TASKS!r} across all cps ===")
for r in rows2:
decoded_v = saver.serde.loads_typed((r["type"], r["blob"]))
print(
f" v={r['version']} type={r['type']} nbytes={r['nbytes']} "
f"-> {decoded_v!r}"
)
assert len(rows2) >= 2, (
"expected at least 2 versions of TASKS blob across the run"
)
latest = saver.serde.loads_typed(
(rows2[-1]["type"], rows2[-1]["blob"])
)
assert latest == [], (
f"expected latest TASKS blob to be cleared (==[]) "
f"after Sends were consumed, got {latest!r}"
)
print()
print(
"✅ ASSERTION PASSED: after Sends were consumed, the latest "
"TASKS blob is [] (Topic with accumulate=False clears itself "
"via update(EMPTY_SEQ) in apply_writes line 304-311)."
)
finally:
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"DROP DATABASE {database}")
if __name__ == "__main__":
main()