mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 22:15:11 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
22083959f4 |
@@ -108,6 +108,7 @@ from langgraph.callbacks import (
|
||||
get_sync_graph_callback_manager_for_config,
|
||||
)
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.topic import Topic
|
||||
from langgraph.config import get_config
|
||||
from langgraph.constants import END
|
||||
@@ -1649,6 +1650,7 @@ class Pregel(
|
||||
# get last checkpoint
|
||||
config = ensure_config(self.config, input_config)
|
||||
saved = checkpointer.get_tuple(config)
|
||||
first_superstep = fork_pending is None
|
||||
if fork_pending is None:
|
||||
fork_pending = delta_channels_with_pending_writes(
|
||||
self.channels, saved.pending_writes if saved else None
|
||||
@@ -2024,13 +2026,21 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||
if saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
)
|
||||
edited_delta_channels = {
|
||||
ch
|
||||
for ch in updated_channels
|
||||
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||
}
|
||||
# The base's other children replay whatever is stored on it, so an
|
||||
# edit of an older checkpoint snapshots its delta channels here
|
||||
# instead. Later supersteps address the checkpoint just written.
|
||||
if (
|
||||
first_superstep
|
||||
and saved is not None
|
||||
and edited_delta_channels
|
||||
and _is_older_checkpoint(checkpointer, config, saved)
|
||||
):
|
||||
fork_pending.update(edited_delta_channels)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2058,7 +2068,17 @@ class Pregel(
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
)
|
||||
sealed = fork_pending.intersection(checkpoint["channel_values"])
|
||||
fork_pending.difference_update(checkpoint["channel_values"])
|
||||
if saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
channel_writes = [
|
||||
w for w in task.writes if w[0] != PUSH and w[0] not in sealed
|
||||
]
|
||||
if channel_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
)
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
checkpoint,
|
||||
@@ -2141,6 +2161,7 @@ class Pregel(
|
||||
# get last checkpoint
|
||||
config = ensure_config(self.config, input_config)
|
||||
saved = await checkpointer.aget_tuple(config)
|
||||
first_superstep = fork_pending is None
|
||||
if fork_pending is None:
|
||||
fork_pending = delta_channels_with_pending_writes(
|
||||
self.channels, saved.pending_writes if saved else None
|
||||
@@ -2511,13 +2532,21 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||
if saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
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
|
||||
)
|
||||
edited_delta_channels = {
|
||||
ch
|
||||
for ch in updated_channels
|
||||
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||
}
|
||||
# The base's other children replay whatever is stored on it, so an
|
||||
# edit of an older checkpoint snapshots its delta channels here
|
||||
# instead. Later supersteps address the checkpoint just written.
|
||||
if (
|
||||
first_superstep
|
||||
and saved is not None
|
||||
and edited_delta_channels
|
||||
and await _ais_older_checkpoint(checkpointer, config, saved)
|
||||
):
|
||||
fork_pending.update(edited_delta_channels)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2545,7 +2574,17 @@ class Pregel(
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
)
|
||||
sealed = fork_pending.intersection(checkpoint["channel_values"])
|
||||
fork_pending.difference_update(checkpoint["channel_values"])
|
||||
if saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
channel_writes = [
|
||||
w for w in task.writes if w[0] != PUSH and w[0] not in sealed
|
||||
]
|
||||
if channel_writes:
|
||||
await checkpointer.aput_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
)
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
checkpoint,
|
||||
@@ -4235,6 +4274,30 @@ def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str
|
||||
return dict(trigger_to_nodes)
|
||||
|
||||
|
||||
def _is_older_checkpoint(
|
||||
checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
|
||||
) -> bool:
|
||||
"""Whether `config` addressed a checkpoint the thread has moved past."""
|
||||
if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
|
||||
return False
|
||||
latest = checkpointer.get_tuple(
|
||||
patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None})
|
||||
)
|
||||
return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"]
|
||||
|
||||
|
||||
async def _ais_older_checkpoint(
|
||||
checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
|
||||
) -> bool:
|
||||
"""Whether `config` addressed a checkpoint the thread has moved past."""
|
||||
if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
|
||||
return False
|
||||
latest = await checkpointer.aget_tuple(
|
||||
patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None})
|
||||
)
|
||||
return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"]
|
||||
|
||||
|
||||
def _output(
|
||||
stream_mode: StreamMode | Sequence[StreamMode],
|
||||
print_mode: StreamMode | Sequence[StreamMode],
|
||||
|
||||
@@ -228,6 +228,118 @@ async def test_afork_by_update_state(
|
||||
assert state.values["log"] == [*base.values["log"], "patched"]
|
||||
|
||||
|
||||
def _assert_branch_unchanged(state: StateSnapshot, expected: list, edit: str) -> None:
|
||||
assert state.values["log"] == state.values["plain"] == expected, (
|
||||
f"{edit!r} was written by an update_state on this branch's base, "
|
||||
f"but this branch now reads {state.values['log']}"
|
||||
)
|
||||
|
||||
|
||||
# The old checkpoint is either a finished turn, which saved no writes, or one
|
||||
# whose next node already ran there, so the edit reuses that task's id.
|
||||
@pytest.mark.parametrize("next_node_ran", [False, True])
|
||||
def test_update_state_on_an_old_checkpoint_leaves_its_other_branch_alone(
|
||||
sync_checkpointer: BaseCheckpointSaver, next_node_ran: bool
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
graph.invoke(_both("in-1"), config)
|
||||
_build(sync_checkpointer, "second").invoke(_both("in-2"), config)
|
||||
branch = graph.get_state(config)
|
||||
base = next(
|
||||
snapshot
|
||||
for snapshot in graph.get_state_history(config)
|
||||
if "in-2" not in snapshot.values["log"]
|
||||
and snapshot.next == (("n",) if next_node_ran else ())
|
||||
)
|
||||
|
||||
edited = graph.update_state(_at(config, base), _both("edit"), as_node="n")
|
||||
|
||||
_assert_branch_unchanged(
|
||||
graph.get_state(branch.config), branch.values["log"], "edit"
|
||||
)
|
||||
assert graph.get_state(edited).values["log"] == [*base.values["log"], "edit"]
|
||||
|
||||
_build(sync_checkpointer, "third").invoke(_both("in-3"), branch.config)
|
||||
_assert_branch_unchanged(
|
||||
graph.get_state(config),
|
||||
[*branch.values["log"], "in-3", "third-out"],
|
||||
"edit",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("next_node_ran", [False, True])
|
||||
async def test_aupdate_state_on_an_old_checkpoint_leaves_its_other_branch_alone(
|
||||
async_checkpointer: BaseCheckpointSaver, next_node_ran: bool
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(async_checkpointer, "first")
|
||||
await graph.ainvoke(_both("in-1"), config)
|
||||
await _build(async_checkpointer, "second").ainvoke(_both("in-2"), config)
|
||||
branch = await graph.aget_state(config)
|
||||
base = await anext(
|
||||
snapshot
|
||||
async for snapshot in graph.aget_state_history(config)
|
||||
if "in-2" not in snapshot.values["log"]
|
||||
and snapshot.next == (("n",) if next_node_ran else ())
|
||||
)
|
||||
|
||||
edited = await graph.aupdate_state(_at(config, base), _both("edit"), as_node="n")
|
||||
|
||||
_assert_branch_unchanged(
|
||||
await graph.aget_state(branch.config), branch.values["log"], "edit"
|
||||
)
|
||||
assert (await graph.aget_state(edited)).values["log"] == [
|
||||
*base.values["log"],
|
||||
"edit",
|
||||
]
|
||||
|
||||
await _build(async_checkpointer, "third").ainvoke(_both("in-3"), branch.config)
|
||||
_assert_branch_unchanged(
|
||||
await graph.aget_state(config),
|
||||
[*branch.values["log"], "in-3", "third-out"],
|
||||
"edit",
|
||||
)
|
||||
|
||||
|
||||
def test_bulk_update_on_an_old_checkpoint_leaves_its_other_branch_alone(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
graph.invoke(_both("in-1"), config)
|
||||
base = graph.get_state(config)
|
||||
_build(sync_checkpointer, "second").invoke(_both("in-2"), config)
|
||||
branch = graph.get_state(config)
|
||||
|
||||
edited = graph.bulk_update_state(
|
||||
_at(config, base),
|
||||
[[StateUpdate(_both("s1"), "n")], [StateUpdate(_both("s2"), "n")]],
|
||||
)
|
||||
|
||||
_assert_branch_unchanged(graph.get_state(branch.config), branch.values["log"], "s1")
|
||||
assert graph.get_state(edited).values["log"] == [*base.values["log"], "s1", "s2"]
|
||||
|
||||
|
||||
def test_update_state_with_the_head_checkpoint_id_stores_no_snapshot(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
graph.invoke(_both("in-1"), config)
|
||||
for i in range(3):
|
||||
graph.update_state(graph.get_state(config).config, _both(f"u{i}"))
|
||||
|
||||
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||
assert graph.get_state(config).values["log"] == [
|
||||
"in-1",
|
||||
"first-out",
|
||||
"u0",
|
||||
"u1",
|
||||
"u2",
|
||||
]
|
||||
|
||||
|
||||
def test_unaddressed_run_keeps_snapshot_cadence(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
|
||||
Reference in New Issue
Block a user