mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-11 02:35:28 +02:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0d91f9ed5b | ||
|
|
d8d02f052f |
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user