Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritz 0d91f9ed5b fix(langgraph): don't cache a node that failed in the sync loop
The async loop skipped caching an interrupted or failed task, but the
sync loop didn't, so a cached error made the next run with the same input
skip the node without raising. Both loops now share the check, which
looks at every write, since a cancelled task's error comes after any
writes it already made.

Also drops `arun_with_retry`'s `match_cached_writes`, unused since #4691
moved cache matching into the loop.
2026-10-10 07:47:26 -04:00
Elior Nataf Lackritz d8d02f052f fix(langgraph): save a cached node's writes like a node that ran
A cache hit's writes were applied but never saved, so a DeltaChannel,
rebuilt from saved writes, lost them for good under every durability.
The loop now saves a hit with put_writes like a node that ran, in all
six places it matches cached writes. A cached flag keeps it streamed as
cached and stops put_writes from writing it back to the cache.
2026-10-09 10:53:45 -04:00
6 changed files with 155 additions and 31 deletions
+26 -17
View File
@@ -160,6 +160,12 @@ def DuplexStream(*streams: StreamProtocol) -> StreamProtocol:
return StreamProtocol(__call__, {mode for s in streams for mode in s.modes})
def _cacheable(writes: WritesT) -> bool:
# A cache hit skips the node, so a cached error would skip it without
# raising: only a node that finished goes to the cache.
return not any(c in (INTERRUPT, ERROR) for c, _ in writes)
class PregelLoop:
config: RunnableConfig
store: BaseStore | None
@@ -439,7 +445,9 @@ class PregelLoop:
return None
return self._graph_lifecycle_events.popleft()
def put_writes(self, task_id: str, writes: WritesT) -> None:
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
"""Put writes for a task, to be read by the next tick."""
if not writes:
return
@@ -532,7 +540,7 @@ class PregelLoop:
self._error_handler_write_futs.append(fut)
# output writes
if hasattr(self, "tasks"):
self.output_writes(task_id, writes)
self.output_writes(task_id, writes, cached=cached)
def _put_pending_writes(self) -> None:
if self.checkpointer_put_writes is None:
@@ -1684,7 +1692,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
) -> PregelExecutableTask | None:
if pushed := super().accept_push(task, write_idx, call):
for task in self.match_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return pushed
def schedule_error_handler(
@@ -1721,16 +1729,18 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in self.match_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(self, task_id: str, writes: WritesT) -> None:
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
if not writes or self.cache is None or not hasattr(self, "tasks"):
super().put_writes(task_id, writes, cached=cached)
if cached or not writes or self.cache is None or not hasattr(self, "tasks"):
return
task = self.tasks.get(task_id)
if task is None or task.cache_key is None:
if task is None or task.cache_key is None or not _cacheable(writes):
return
self.submit(
self.cache.set,
@@ -1937,7 +1947,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
) -> PregelExecutableTask | None:
if pushed := super().accept_push(task, write_idx, call):
for task in await self.amatch_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return pushed
async def aschedule_error_handler(
@@ -1974,19 +1984,18 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in await self.amatch_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(self, task_id: str, writes: WritesT) -> None:
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
if not writes or self.cache is None or not hasattr(self, "tasks"):
super().put_writes(task_id, writes, cached=cached)
if cached or not writes or self.cache is None or not hasattr(self, "tasks"):
return
task = self.tasks.get(task_id)
if task is None or task.cache_key is None:
return
if writes[0][0] in (INTERRUPT, ERROR):
# only cache successful tasks
if task is None or task.cache_key is None or not _cacheable(writes):
return
self.submit(
self.cache.aset,
+1 -8
View File
@@ -7,7 +7,7 @@ import sys
import threading
import time
import weakref
from collections.abc import Awaitable, Callable, Sequence
from collections.abc import Callable, Sequence
from contextlib import suppress
from dataclasses import dataclass, replace
from datetime import datetime, timedelta, timezone
@@ -686,8 +686,6 @@ async def arun_with_retry(
task: PregelExecutableTask,
retry_policy: Sequence[RetryPolicy] | None,
stream: bool = False,
match_cached_writes: Callable[[], Awaitable[Sequence[PregelExecutableTask]]]
| None = None,
configurable: dict[str, Any] | None = None,
) -> None:
"""Run a task asynchronously with retries."""
@@ -711,11 +709,6 @@ async def arun_with_retry(
)
},
)
if match_cached_writes is not None and task.cache_key is not None:
for t in await match_cached_writes():
if t is task:
# if the task is already cached, return
return
while True:
runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME)
if isinstance(runtime, Runtime):
+2 -2
View File
@@ -3038,7 +3038,7 @@ class Pregel(
# with channel updates applied only at the transition between steps.
while loop.tick():
for task in loop.match_cached_writes():
loop.output_writes(task.id, task.writes, cached=True)
loop.put_writes(task.id, task.writes, cached=True)
for _ in runner.tick(
[t for t in loop.tasks.values() if not t.writes],
timeout=self.step_timeout,
@@ -3509,7 +3509,7 @@ class Pregel(
try:
while loop.tick():
for task in await loop.amatch_cached_writes():
loop.output_writes(task.id, task.writes, cached=True)
loop.put_writes(task.id, task.writes, cached=True)
async for _ in runner.atick(
[t for t in loop.tasks.values() if not t.writes],
timeout=self.step_timeout,
@@ -0,0 +1,122 @@
"""A node served from the node cache saves its writes like a node that ran,
without writing them back to the cache, and a node that fails isn't cached."""
import operator
from typing import Annotated, Any
import pytest
from langgraph.cache.memory import InMemoryCache
from langgraph.checkpoint.memory import InMemorySaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.types import CachePolicy, Durability
pytestmark = pytest.mark.anyio
INPUT = {"log": [], "plain": []}
def _append(current: list, writes: list) -> list:
return [*current, *(item for write in writes for item in write)]
class _State(TypedDict):
log: Annotated[list, DeltaChannel(_append)]
plain: Annotated[list, operator.add]
class _CountsSets(InMemoryCache):
sets = 0
def set(self, keys: Any) -> None:
self.sets += 1
super().set(keys)
def _a_then_cached_b_then_c(runs: list[str], cache: InMemoryCache) -> Any:
def node(name: str) -> Any:
def run(state: _State) -> dict:
runs.append(name)
return {"log": [name], "plain": [name]}
return run
builder = StateGraph(_State)
builder.add_node("a", node("a"))
builder.add_node("b", node("b"), cache_policy=CachePolicy())
builder.add_node("c", node("c"))
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("b", "c")
return builder.compile(checkpointer=InMemorySaver(), cache=cache)
def test_a_cache_hit_saves_its_writes_without_caching_them_again(
durability: Durability,
) -> None:
runs: list[str] = []
cache = _CountsSets()
graph = _a_then_cached_b_then_c(runs, cache)
graph.invoke(INPUT, {"configurable": {"thread_id": "1"}}, durability=durability)
config = {"configurable": {"thread_id": "2"}}
graph.invoke(INPUT, config, durability=durability)
assert runs == ["a", "b", "c", "a", "c"]
assert cache.sets == 1, "the cache hit was written back to the cache"
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
async def test_a_cache_hit_saves_its_writes_without_caching_them_again_async(
durability: Durability,
) -> None:
runs: list[str] = []
cache = _CountsSets()
graph = _a_then_cached_b_then_c(runs, cache)
await graph.ainvoke(
INPUT, {"configurable": {"thread_id": "1"}}, durability=durability
)
config = {"configurable": {"thread_id": "2"}}
await graph.ainvoke(INPUT, config, durability=durability)
assert runs == ["a", "b", "c", "a", "c"]
assert cache.sets == 1, "the cache hit was written back to the cache"
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
def _cached_node_that_fails(runs: list[str]) -> Any:
def fail(state: _State) -> dict:
runs.append("b")
raise ValueError("b failed")
builder = StateGraph(_State)
builder.add_node("b", fail, cache_policy=CachePolicy())
builder.add_edge(START, "b")
return builder.compile(checkpointer=InMemorySaver(), cache=InMemoryCache())
def test_a_node_that_fails_is_not_cached() -> None:
runs: list[str] = []
graph = _cached_node_that_fails(runs)
for thread in ("1", "2"):
with pytest.raises(ValueError, match="b failed"):
graph.invoke(INPUT, {"configurable": {"thread_id": thread}})
assert runs == ["b", "b"]
async def test_a_node_that_fails_is_not_cached_async() -> None:
runs: list[str] = []
graph = _cached_node_that_fails(runs)
for thread in ("1", "2"):
with pytest.raises(ValueError, match="b failed"):
await graph.ainvoke(INPUT, {"configurable": {"thread_id": thread}})
assert runs == ["b", "b"]
+2 -2
View File
@@ -5932,9 +5932,9 @@ def test_no_redundant_put_writes_for_cached_task(
put_writes_task_ids: list[str] = []
orig = PregelLoop.put_writes
def spy(self, task_id, writes):
def spy(self, task_id, writes, **kwargs):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes)
return orig(self, task_id, writes, **kwargs)
with patch.object(PregelLoop, "put_writes", spy):
result = workflow.invoke(Command(resume="ans"), config=config)
+2 -2
View File
@@ -8157,9 +8157,9 @@ async def test_no_redundant_put_writes_for_cached_task(
put_writes_task_ids: list[str] = []
orig = PregelLoop.put_writes
def spy(self, task_id, writes):
def spy(self, task_id, writes, **kwargs):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes)
return orig(self, task_id, writes, **kwargs)
with patch.object(PregelLoop, "put_writes", spy):
result = await workflow.ainvoke(Command(resume="ans"), config=config)