mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-07 00:45:07 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
56fd3458cc | ||
|
|
a7a465877d | ||
|
|
b1e167401e | ||
|
|
43a227b584 | ||
|
|
a8e732c879 |
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
@@ -49,6 +50,15 @@ def empty_checkpoint() -> Checkpoint:
|
||||
)
|
||||
|
||||
|
||||
def put_writes_accepts_task_path(put_writes: Callable[..., Any]) -> bool:
|
||||
"""Whether a saver's `put_writes` or `aput_writes` takes `task_path`.
|
||||
|
||||
Savers written before the parameter existed don't, so it is passed only
|
||||
when this is true.
|
||||
"""
|
||||
return signature(put_writes).parameters.get("task_path") is not None
|
||||
|
||||
|
||||
def exit_delta_task_id(step: int, task_id: str) -> str:
|
||||
"""Synthetic task id for exit-mode DeltaChannel writes.
|
||||
|
||||
|
||||
@@ -12,7 +12,6 @@ from contextlib import (
|
||||
ExitStack,
|
||||
)
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
Any,
|
||||
@@ -107,6 +106,7 @@ from langgraph.pregel._checkpoint import (
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
exit_delta_task_id,
|
||||
put_writes_accepts_task_path,
|
||||
)
|
||||
from langgraph.pregel._executor import (
|
||||
AsyncBackgroundExecutor,
|
||||
@@ -1584,8 +1584,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
self.checkpointer_get_next_version = checkpointer.get_next_version
|
||||
self.checkpointer_put_writes = checkpointer.put_writes
|
||||
self.checkpointer_put_writes_accepts_task_path = (
|
||||
signature(checkpointer.put_writes).parameters.get("task_path")
|
||||
is not None
|
||||
put_writes_accepts_task_path(checkpointer.put_writes)
|
||||
)
|
||||
else:
|
||||
self.checkpointer_get_next_version = increment
|
||||
@@ -1840,8 +1839,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
self.checkpointer_get_next_version = checkpointer.get_next_version
|
||||
self.checkpointer_put_writes = checkpointer.aput_writes
|
||||
self.checkpointer_put_writes_accepts_task_path = (
|
||||
signature(checkpointer.aput_writes).parameters.get("task_path")
|
||||
is not None
|
||||
put_writes_accepts_task_path(checkpointer.aput_writes)
|
||||
)
|
||||
else:
|
||||
self.checkpointer_get_next_version = increment
|
||||
|
||||
@@ -126,6 +126,7 @@ from langgraph.pregel._algo import (
|
||||
apply_writes,
|
||||
local_read,
|
||||
prepare_next_tasks,
|
||||
task_path_str,
|
||||
)
|
||||
from langgraph.pregel._call import identifier
|
||||
from langgraph.pregel._checkpoint import (
|
||||
@@ -139,6 +140,7 @@ from langgraph.pregel._checkpoint import (
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
get_updated_channels_from_tasks,
|
||||
put_writes_accepts_task_path,
|
||||
versions_seen_without_bumps,
|
||||
)
|
||||
from langgraph.pregel._draw import draw_graph
|
||||
@@ -1983,13 +1985,13 @@ class Pregel(
|
||||
run_tasks: list[PregelTaskWrites] = []
|
||||
run_task_ids: list[str] = []
|
||||
|
||||
for as_node, values, provided_task_id in valid_updates:
|
||||
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
|
||||
# create task to run all writers of the chosen node
|
||||
writers = self.nodes[as_node].flat_writers
|
||||
if not writers:
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
|
||||
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
|
||||
# get the task ids that were prepared for this node
|
||||
# if a task id was provided in the StateUpdate, we use it
|
||||
# otherwise, we use the next available task id
|
||||
@@ -1997,7 +1999,7 @@ class Pregel(
|
||||
task_id = provided_task_id or (
|
||||
prepared_task_ids.popleft()
|
||||
if prepared_task_ids
|
||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
||||
else _update_task_id(checkpoint["id"], i)
|
||||
)
|
||||
run_tasks.append(task)
|
||||
run_task_ids.append(task_id)
|
||||
@@ -2050,7 +2052,10 @@ class Pregel(
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
checkpoint_config,
|
||||
channel_writes,
|
||||
task_id,
|
||||
**_task_path_kwarg(checkpointer.put_writes, task),
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
@@ -2471,13 +2476,13 @@ class Pregel(
|
||||
run_tasks: list[PregelTaskWrites] = []
|
||||
run_task_ids: list[str] = []
|
||||
|
||||
for as_node, values, provided_task_id in valid_updates:
|
||||
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
|
||||
# create task to run all writers of the chosen node
|
||||
writers = self.nodes[as_node].flat_writers
|
||||
if not writers:
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
|
||||
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
|
||||
# get the task ids that were prepared for this node
|
||||
# if a task id was provided in the StateUpdate, we use it
|
||||
# otherwise, we use the next available task id
|
||||
@@ -2485,7 +2490,7 @@ class Pregel(
|
||||
task_id = provided_task_id or (
|
||||
prepared_task_ids.popleft()
|
||||
if prepared_task_ids
|
||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
||||
else _update_task_id(checkpoint["id"], i)
|
||||
)
|
||||
run_tasks.append(task)
|
||||
run_task_ids.append(task_id)
|
||||
@@ -2538,7 +2543,10 @@ class Pregel(
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
await checkpointer.aput_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
checkpoint_config,
|
||||
channel_writes,
|
||||
task_id,
|
||||
**_task_path_kwarg(checkpointer.aput_writes, task),
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
@@ -4237,6 +4245,27 @@ class Pregel(
|
||||
await self.cache.aclear(namespaces)
|
||||
|
||||
|
||||
def _task_path_kwarg(put_writes: Callable[..., Any], task: PregelTaskWrites) -> dict:
|
||||
"""Pass the task's path to savers whose `put_writes` takes one.
|
||||
|
||||
Savers that replay a checkpoint's writes in task path order then give back
|
||||
updates applied together in the order they were given.
|
||||
"""
|
||||
if not put_writes_accepts_task_path(put_writes):
|
||||
return {}
|
||||
return {"task_path": task_path_str(task.path)}
|
||||
|
||||
|
||||
def _update_task_id(checkpoint_id: str, i: int) -> str:
|
||||
"""Task id for the `i`th update of a superstep that has no task to reuse.
|
||||
|
||||
Savers keep one write per `(task_id, idx)`, so updates sharing an id lose
|
||||
all but the first one's writes, which a `DeltaChannel` replays from. The
|
||||
first update keeps the id a lone update has always had.
|
||||
"""
|
||||
return str(uuid5(UUID(checkpoint_id), INTERRUPT if i == 0 else f"{INTERRUPT}:{i}"))
|
||||
|
||||
|
||||
def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str]]:
|
||||
"""Index from a trigger to nodes that depend on it."""
|
||||
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
|
||||
|
||||
@@ -20,6 +20,7 @@ from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import TypedDict
|
||||
@@ -27,16 +28,17 @@ 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
|
||||
from langgraph.types import StateUpdate
|
||||
from langgraph.types import StateSnapshot, StateUpdate
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _build_graph(
|
||||
checkpointer: InMemorySaver,
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
*,
|
||||
two_nodes: bool = False,
|
||||
snapshot_frequency: int = 1000,
|
||||
interrupt_before: list[str] | None = None,
|
||||
) -> Any:
|
||||
"""Compile a minimal DeltaChannel-backed `messages` graph.
|
||||
|
||||
@@ -63,7 +65,7 @@ def _build_graph(
|
||||
builder.set_finish_point("assistant")
|
||||
else:
|
||||
builder.set_finish_point("model")
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
return builder.compile(checkpointer=checkpointer, interrupt_before=interrupt_before)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -304,15 +306,14 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
||||
that each call `put_writes`. Guards the regression where moving
|
||||
`put_writes` outside the per-task loop would persist only the last
|
||||
task's writes.
|
||||
|
||||
Explicit `task_id`s are required to disambiguate writes belonging to
|
||||
different `StateUpdate`s targeting the same node — otherwise both share
|
||||
the deterministic interrupt-derived id and collide in the saver.
|
||||
"""
|
||||
|
||||
saver = InMemorySaver()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": "bulk-multi-task"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
base = saver.get_tuple(config)
|
||||
assert base is not None
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
@@ -332,13 +333,155 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
||||
],
|
||||
)
|
||||
|
||||
stored = saver.get_tuple(base.config)
|
||||
assert stored is not None
|
||||
assert {task_id for task_id, _, _ in stored.pending_writes or []} == {
|
||||
"task-1",
|
||||
"task-2",
|
||||
}, "explicit task ids must key the stored writes"
|
||||
state = graph.get_state(config)
|
||||
contents = [m.content for m in state.values["messages"]]
|
||||
ids = [m.id for m in state.values["messages"]]
|
||||
assert sorted(contents) == ["first", "second"], (
|
||||
assert sorted(contents) == ["first", "hi", "second"], (
|
||||
f"both updates' writes must persist; got {contents}"
|
||||
)
|
||||
assert sorted(ids) == ["m1", "m2"]
|
||||
assert sorted(ids) == ["hi", "m1", "m2"]
|
||||
|
||||
|
||||
def _update(content: str, as_node: str) -> StateUpdate:
|
||||
return StateUpdate(
|
||||
values={"messages": [HumanMessage(content=content, id=content)]},
|
||||
as_node=as_node,
|
||||
)
|
||||
|
||||
|
||||
def _contents(state: StateSnapshot) -> list[str]:
|
||||
return [m.content for m in state.values["messages"]]
|
||||
|
||||
|
||||
def test_bulk_update_state_keeps_every_update_without_task_ids(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_graph(sync_checkpointer, two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
[
|
||||
[
|
||||
_update("first", "model"),
|
||||
_update("second", "model"),
|
||||
_update("third", "assistant"),
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
contents = _contents(graph.get_state(config))
|
||||
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||
f"every update's writes must persist; got {contents}"
|
||||
)
|
||||
|
||||
|
||||
async def test_abulk_update_state_keeps_every_update_without_task_ids(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_graph(async_checkpointer, two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
await graph.abulk_update_state(
|
||||
config,
|
||||
[
|
||||
[
|
||||
_update("first", "model"),
|
||||
_update("second", "model"),
|
||||
_update("third", "assistant"),
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
contents = _contents(await graph.aget_state(config))
|
||||
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||
f"every update's writes must persist; got {contents}"
|
||||
)
|
||||
|
||||
|
||||
def test_bulk_update_state_keeps_every_update_next_to_a_pending_task(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_graph(
|
||||
sync_checkpointer, two_nodes=True, interrupt_before=["assistant"]
|
||||
)
|
||||
config = {"configurable": {"thread_id": "bulk-pending-task"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
assert graph.get_state(config).next == ("assistant",)
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
[
|
||||
[
|
||||
_update("first", "assistant"),
|
||||
_update("second", "model"),
|
||||
_update("third", "model"),
|
||||
]
|
||||
],
|
||||
)
|
||||
|
||||
contents = _contents(graph.get_state(config))
|
||||
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||
f"every update's writes must persist; got {contents}"
|
||||
)
|
||||
|
||||
|
||||
class _TaskPathOrderSaver(InMemorySaver):
|
||||
"""Replays each checkpoint's writes by `(task_path, task_id, idx)`."""
|
||||
|
||||
def get_tuple(self, config: Any) -> Any:
|
||||
tup = super().get_tuple(config)
|
||||
if tup is None or not tup.pending_writes:
|
||||
return tup
|
||||
conf = tup.config["configurable"]
|
||||
stored = self.writes[
|
||||
(conf["thread_id"], conf["checkpoint_ns"], conf["checkpoint_id"])
|
||||
]
|
||||
rows = sorted(
|
||||
zip(stored.items(), tup.pending_writes),
|
||||
key=lambda row: (row[0][1][3], *row[0][0]),
|
||||
)
|
||||
return tup._replace(pending_writes=[write for _, write in rows])
|
||||
|
||||
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
|
||||
aget_delta_channel_history = BaseCheckpointSaver.aget_delta_channel_history
|
||||
|
||||
|
||||
GIVEN = ["u1", "u2", "u3", "u4", "u5", "u6"]
|
||||
|
||||
|
||||
def _updates_in_given_order() -> list[list[StateUpdate]]:
|
||||
return [
|
||||
[_update(c, "assistant" if i % 2 else "model") for i, c in enumerate(GIVEN)]
|
||||
]
|
||||
|
||||
|
||||
def test_bulk_update_state_replays_updates_in_the_order_given() -> None:
|
||||
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-order"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
graph.bulk_update_state(config, _updates_in_given_order())
|
||||
|
||||
assert _contents(graph.get_state(config)) == ["hi", *GIVEN]
|
||||
|
||||
|
||||
async def test_abulk_update_state_replays_updates_in_the_order_given() -> None:
|
||||
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-order"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
await graph.abulk_update_state(config, _updates_in_given_order())
|
||||
|
||||
assert _contents(await graph.aget_state(config)) == ["hi", *GIVEN]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user