Compare commits

..
Author SHA1 Message Date
Sydney RunkleandClaude Opus 4.7 18cbe46baf feat(prebuilt): hydrate ToolNode state from channels via CONFIG_KEY_READ
When ToolNode receives a bare `[tool_call]` list via the Send API (the
new dispatch shape that create_agent uses after langchain-ai/langchain
drops the ToolCallWithContext wrapper), hydrate ToolRuntime.state from
the current channel values instead of requiring the dispatcher to
inline the full agent state dict in the Send payload.

Implementation stays entirely in tool_node:

- Pregel installs CONFIG_KEY_READ as
  `functools.partial(local_read, scratchpad, channels, managed, task)`.
  Introspect the partial's positional args to learn channel + managed
  names, then read them all via `ChannelRead.do_read` with an explicit
  list. No changes to the pregel read machinery.
- Gracefully falls back to {} when invoked outside a Pregel context
  (e.g. direct ToolNode(...).invoke(...) from test harnesses).
- Legacy ToolCallWithContext path is preserved for external dispatchers.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-24 07:30:51 -04:00
Sydney RunkleandClaude Opus 4.7 e7af9869bb refactor(add_messages): clearer step-by-step structure
Restructure add_messages so fast vs. slow path, REMOVE_ALL handling,
and format application each live in a single numbered section with a
short lead-in comment. No behavior changes; the previous commits'
optimizations are preserved.

- fold the two path branches into one `if pure_append else slow_path`
  so format handling happens at a single exit instead of being
  duplicated between fast and slow paths
- drop the now-redundant `left_seq` local; index `left` directly after
  the coerce step (with a single `cast(list, left)` for the type
  checker)
- use `set.isdisjoint` for the overlap check

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-24 07:30:51 -04:00
Sydney RunkleandClaude Opus 4.7 e5bfae0f9d test(add_messages): cover fast-path guards and format handling
Nine new tests pin the behavioral boundaries introduced by the
optimization: chunk / dict / tuple / missing-id left inputs must fall
through to full conversion, duplicate right ids and None right ids
must still be handled correctly, format="langchain-openai" and invalid
format must round-trip through the fast path, and the fast path must
return a fresh list rather than aliasing left.

Also picks up a ruff-format reflow in test_time_travel.py.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-24 07:30:51 -04:00
Sydney RunkleandClaude Sonnet 4.6 3ef54c1ec8 fix(add_messages): guard fast path against None IDs and intra-right duplicates
Two bugs in the fast-path optimisation:

1. The left-side type guard checked isinstance(BaseMessage) but not
   id is not None. Messages without IDs (e.g. HumanMessage(content="hi"))
   would skip ID assignment and return None IDs.

2. The pure-append short-circuit only checked for overlaps between
   right and left, not duplicates within right itself. A right list
   containing two messages with the same ID would bypass the slow-path
   deduplication and return both.

Fixes:
- Add left_seq[0].id is not None to the type-guard condition.
- Replace the any() overlap check with a set-intersection check that
  also verifies len(right_id_set) == len(right_msgs) (no intra-right
  duplicates) before taking the fast return.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-24 07:30:51 -04:00
Sydney RunkleandClaude Sonnet 4.6 2a974d1b73 perf(add_messages): skip left-side conversion and fast-path pure appends
Two optimizations for the hot path in add_messages, which is called on
every write to a messages channel:

1. Skip conversion of left: when left is already list[BaseMessage] with
   IDs assigned (true for every call after the first), skip
   convert_to_messages + message_chunk_to_message + the ID-None loop.
   These are O(n) no-ops on already-resolved messages that allocate two
   intermediate lists.

2. Pure-append short-circuit: when right contains no RemoveMessage and
   no ID overlaps with left, return left + right directly. Replaces the
   O(n) copy + dict build + filter with a single set-membership check.

Benchmarks (median of 2000 iterations, pure-append scenario):
  10-msg thread:   2.9x faster
  100-msg thread:  6.6x faster
  1000-msg thread: 7.3x faster
  200-step simulation (2 msgs/step): 3.4x faster end-to-end

Also adds tests/test_add_messages_benchmark.py with correctness tests
for all scenarios (append, update, remove) and a runnable benchmark.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-24 07:30:51 -04:00
10 changed files with 606 additions and 269 deletions
+74 -40
View File
@@ -184,63 +184,97 @@ def add_messages(
```
"""
remove_all_idx = None
# coerce to list
# 1. Coerce scalars to lists.
if not isinstance(left, list):
left = [left] # type: ignore[assignment]
if not isinstance(right, list):
right = [right] # type: ignore[assignment]
# coerce to message
left = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(left)
]
right = [
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] = [
message_chunk_to_message(cast(BaseMessageChunk, m))
for m in convert_to_messages(right)
]
# assign missing ids
for m in left:
remove_all_idx: int | None = None
has_remove = False
for idx, m in enumerate(right_msgs):
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 isinstance(m, RemoveMessage):
has_remove = True
if 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[remove_all_idx + 1 :]
return right_msgs[remove_all_idx + 1 :]
# 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)
# 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
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)
merged = [m for m in merged if m.id not in ids_to_remove]
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":
merged = _format_messages(merged)
elif format:
return _format_messages(merged)
if 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._delta import DeltaChannel
from langgraph.channels.base import BaseChannel
from langgraph.channels._delta import DeltaChannel
from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec
LATEST_VERSION = 4
@@ -0,0 +1,290 @@
"""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._delta import DeltaChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels._delta import DeltaChannel
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._delta import DeltaChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels._delta import DeltaChannel
from langgraph.graph import END, START, StateGraph
pytestmark = pytest.mark.anyio
+118
View File
@@ -5,6 +5,7 @@ import langchain_core
import pytest
from langchain_core.messages import (
AIMessage,
AIMessageChunk,
AnyMessage,
HumanMessage,
RemoveMessage,
@@ -338,6 +339,123 @@ 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]
+28 -8
View File
@@ -82,6 +82,7 @@ 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
@@ -800,7 +801,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)
state = self._extract_state(input, cfg)
tool_runtime = ToolRuntime(
state=state,
tool_call_id=call["id"],
@@ -835,7 +836,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)
state = self._extract_state(input, cfg)
tool_runtime = ToolRuntime(
state=state,
tool_call_id=call["id"],
@@ -1273,18 +1274,37 @@ class ToolNode(RunnableCallable):
return None
def _extract_state(
self, input: list[AnyMessage] | dict[str, Any] | BaseModel
self,
input: list[AnyMessage] | dict[str, Any] | BaseModel,
config: RunnableConfig,
) -> list[AnyMessage] | dict[str, Any] | BaseModel:
"""Extract state from input, handling ToolCallWithContext if present.
"""Extract state from input.
Args:
input: The input which may be raw state or ToolCallWithContext.
Three input shapes:
Returns:
The actual state to pass to wrap_tool_call wrappers.
- `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.
"""
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,6 +1320,98 @@ 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
@@ -1,217 +0,0 @@
"""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()