mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 05:55:14 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a8e732c879 |
@@ -1950,7 +1950,7 @@ class Pregel(
|
|||||||
run_tasks: list[PregelTaskWrites] = []
|
run_tasks: list[PregelTaskWrites] = []
|
||||||
run_task_ids: list[str] = []
|
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
|
# create task to run all writers of the chosen node
|
||||||
writers = self.nodes[as_node].flat_writers
|
writers = self.nodes[as_node].flat_writers
|
||||||
if not writers:
|
if not writers:
|
||||||
@@ -1964,7 +1964,7 @@ class Pregel(
|
|||||||
task_id = provided_task_id or (
|
task_id = provided_task_id or (
|
||||||
prepared_task_ids.popleft()
|
prepared_task_ids.popleft()
|
||||||
if prepared_task_ids
|
if prepared_task_ids
|
||||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
else _update_task_id(checkpoint["id"], i)
|
||||||
)
|
)
|
||||||
run_tasks.append(task)
|
run_tasks.append(task)
|
||||||
run_task_ids.append(task_id)
|
run_task_ids.append(task_id)
|
||||||
@@ -2410,7 +2410,7 @@ class Pregel(
|
|||||||
run_tasks: list[PregelTaskWrites] = []
|
run_tasks: list[PregelTaskWrites] = []
|
||||||
run_task_ids: list[str] = []
|
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
|
# create task to run all writers of the chosen node
|
||||||
writers = self.nodes[as_node].flat_writers
|
writers = self.nodes[as_node].flat_writers
|
||||||
if not writers:
|
if not writers:
|
||||||
@@ -2424,7 +2424,7 @@ class Pregel(
|
|||||||
task_id = provided_task_id or (
|
task_id = provided_task_id or (
|
||||||
prepared_task_ids.popleft()
|
prepared_task_ids.popleft()
|
||||||
if prepared_task_ids
|
if prepared_task_ids
|
||||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
else _update_task_id(checkpoint["id"], i)
|
||||||
)
|
)
|
||||||
run_tasks.append(task)
|
run_tasks.append(task)
|
||||||
run_task_ids.append(task_id)
|
run_task_ids.append(task_id)
|
||||||
@@ -4172,6 +4172,16 @@ class Pregel(
|
|||||||
await self.cache.aclear(namespaces)
|
await self.cache.aclear(namespaces)
|
||||||
|
|
||||||
|
|
||||||
|
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]]:
|
def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str]]:
|
||||||
"""Index from a trigger to nodes that depend on it."""
|
"""Index from a trigger to nodes that depend on it."""
|
||||||
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
|
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from typing import Annotated, Any
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from langchain_core.messages import HumanMessage
|
from langchain_core.messages import HumanMessage
|
||||||
|
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||||
from langgraph.checkpoint.memory import InMemorySaver
|
from langgraph.checkpoint.memory import InMemorySaver
|
||||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||||
from typing_extensions import TypedDict
|
from typing_extensions import TypedDict
|
||||||
@@ -27,16 +28,17 @@ from typing_extensions import TypedDict
|
|||||||
from langgraph.channels.delta import DeltaChannel
|
from langgraph.channels.delta import DeltaChannel
|
||||||
from langgraph.graph import START, StateGraph
|
from langgraph.graph import START, StateGraph
|
||||||
from langgraph.graph.message import _messages_delta_reducer
|
from langgraph.graph.message import _messages_delta_reducer
|
||||||
from langgraph.types import StateUpdate
|
from langgraph.types import StateSnapshot, StateUpdate
|
||||||
|
|
||||||
pytestmark = pytest.mark.anyio
|
pytestmark = pytest.mark.anyio
|
||||||
|
|
||||||
|
|
||||||
def _build_graph(
|
def _build_graph(
|
||||||
checkpointer: InMemorySaver,
|
checkpointer: BaseCheckpointSaver,
|
||||||
*,
|
*,
|
||||||
two_nodes: bool = False,
|
two_nodes: bool = False,
|
||||||
snapshot_frequency: int = 1000,
|
snapshot_frequency: int = 1000,
|
||||||
|
interrupt_before: list[str] | None = None,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""Compile a minimal DeltaChannel-backed `messages` graph.
|
"""Compile a minimal DeltaChannel-backed `messages` graph.
|
||||||
|
|
||||||
@@ -63,7 +65,7 @@ def _build_graph(
|
|||||||
builder.set_finish_point("assistant")
|
builder.set_finish_point("assistant")
|
||||||
else:
|
else:
|
||||||
builder.set_finish_point("model")
|
builder.set_finish_point("model")
|
||||||
return builder.compile(checkpointer=checkpointer)
|
return builder.compile(checkpointer=checkpointer, interrupt_before=interrupt_before)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -273,10 +275,6 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
|||||||
that each call `put_writes`. Guards the regression where moving
|
that each call `put_writes`. Guards the regression where moving
|
||||||
`put_writes` outside the per-task loop would persist only the last
|
`put_writes` outside the per-task loop would persist only the last
|
||||||
task's writes.
|
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()
|
saver = InMemorySaver()
|
||||||
@@ -310,6 +308,92 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
|||||||
assert sorted(ids) == ["m1", "m2"]
|
assert sorted(ids) == ["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}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Public-API observation of fresh-thread checkpoint shape
|
# Public-API observation of fresh-thread checkpoint shape
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user