mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 22:15:11 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3129e521ea |
@@ -114,6 +114,26 @@ def create_metadata_for_update_state_api(
|
||||
return new_counters
|
||||
|
||||
|
||||
def advance_delta_counters(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
updated_channels: set[str],
|
||||
*,
|
||||
prev_metadata: Mapping[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
"""The `counters_since_delta_snapshot` entry for an update_state
|
||||
checkpoint saved one superstep after `prev_metadata`'s, for the paths
|
||||
that skip `create_checkpoint_plan_for_update_state_api`.
|
||||
|
||||
Without it, the next checkpoint restarts every delta channel's snapshot
|
||||
cadence from zero.
|
||||
"""
|
||||
counters = create_metadata_for_update_state_api(
|
||||
channels, updated_channels, prev_metadata=prev_metadata
|
||||
)
|
||||
non_zero = {k: v for k, v in counters.items() if v != (0, 0)}
|
||||
return {"counters_since_delta_snapshot": non_zero} if non_zero else {}
|
||||
|
||||
|
||||
def create_checkpoint_plan_for_update_state_api(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
updated_channels: set[str],
|
||||
|
||||
@@ -129,6 +129,7 @@ from langgraph.pregel._algo import (
|
||||
from langgraph.pregel._call import identifier
|
||||
from langgraph.pregel._checkpoint import (
|
||||
achannels_from_checkpoint,
|
||||
advance_delta_counters,
|
||||
channels_from_checkpoint,
|
||||
copy_checkpoint,
|
||||
create_checkpoint,
|
||||
@@ -1681,6 +1682,7 @@ class Pregel(
|
||||
"Cannot apply multiple updates when clearing state"
|
||||
)
|
||||
|
||||
updated_channels: set[str] = set()
|
||||
if saved is not None:
|
||||
# tasks for this checkpoint
|
||||
next_tasks = prepare_next_tasks(
|
||||
@@ -1703,7 +1705,7 @@ class Pregel(
|
||||
for w in saved.pending_writes or []
|
||||
if w[0] == NULL_TASK_ID
|
||||
]:
|
||||
apply_writes(
|
||||
updated_channels |= apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
[PregelTaskWrites((), INPUT, null_writes, [])],
|
||||
@@ -1718,7 +1720,7 @@ class Pregel(
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
# clear all current tasks
|
||||
apply_writes(
|
||||
updated_channels |= apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
next_tasks.values(),
|
||||
@@ -1733,6 +1735,13 @@ class Pregel(
|
||||
"source": "update",
|
||||
"step": step + 1,
|
||||
"parents": saved.metadata.get("parents", {}) if saved else {},
|
||||
**(
|
||||
advance_delta_counters(
|
||||
channels, updated_channels, prev_metadata=saved.metadata
|
||||
)
|
||||
if saved
|
||||
else {}
|
||||
),
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions,
|
||||
@@ -1751,7 +1760,7 @@ class Pregel(
|
||||
)
|
||||
|
||||
if input_writes := deque(map_input(self.input_channels, values)):
|
||||
apply_writes(
|
||||
updated_channels = apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
[PregelTaskWrites((), INPUT, input_writes, [])],
|
||||
@@ -1774,6 +1783,15 @@ class Pregel(
|
||||
"parents": saved.metadata.get("parents", {})
|
||||
if saved
|
||||
else {},
|
||||
**(
|
||||
advance_delta_counters(
|
||||
channels,
|
||||
updated_channels,
|
||||
prev_metadata=saved.metadata,
|
||||
)
|
||||
if saved
|
||||
else {}
|
||||
),
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions,
|
||||
@@ -1819,6 +1837,12 @@ class Pregel(
|
||||
"source": "fork",
|
||||
"step": step + 1,
|
||||
"parents": saved.metadata.get("parents", {}),
|
||||
# The copy has the same values and the same parent.
|
||||
**{
|
||||
k: v
|
||||
for k, v in saved.metadata.items()
|
||||
if k == "counters_since_delta_snapshot"
|
||||
},
|
||||
},
|
||||
{},
|
||||
)
|
||||
@@ -2145,6 +2169,7 @@ class Pregel(
|
||||
raise InvalidUpdateError(
|
||||
"Cannot apply multiple updates when clearing state"
|
||||
)
|
||||
updated_channels: set[str] = set()
|
||||
if saved is not None:
|
||||
# tasks for this checkpoint
|
||||
next_tasks = prepare_next_tasks(
|
||||
@@ -2167,7 +2192,7 @@ class Pregel(
|
||||
for w in saved.pending_writes or []
|
||||
if w[0] == NULL_TASK_ID
|
||||
]:
|
||||
apply_writes(
|
||||
updated_channels |= apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
[PregelTaskWrites((), INPUT, null_writes, [])],
|
||||
@@ -2182,7 +2207,7 @@ class Pregel(
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
# clear all current tasks
|
||||
apply_writes(
|
||||
updated_channels |= apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
next_tasks.values(),
|
||||
@@ -2197,6 +2222,13 @@ class Pregel(
|
||||
"source": "update",
|
||||
"step": step + 1,
|
||||
"parents": saved.metadata.get("parents", {}) if saved else {},
|
||||
**(
|
||||
advance_delta_counters(
|
||||
channels, updated_channels, prev_metadata=saved.metadata
|
||||
)
|
||||
if saved
|
||||
else {}
|
||||
),
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions, checkpoint["channel_versions"]
|
||||
@@ -2214,7 +2246,7 @@ class Pregel(
|
||||
)
|
||||
|
||||
if input_writes := deque(map_input(self.input_channels, values)):
|
||||
apply_writes(
|
||||
updated_channels = apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
[PregelTaskWrites((), INPUT, input_writes, [])],
|
||||
@@ -2237,6 +2269,15 @@ class Pregel(
|
||||
"parents": saved.metadata.get("parents", {})
|
||||
if saved
|
||||
else {},
|
||||
**(
|
||||
advance_delta_counters(
|
||||
channels,
|
||||
updated_channels,
|
||||
prev_metadata=saved.metadata,
|
||||
)
|
||||
if saved
|
||||
else {}
|
||||
),
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions,
|
||||
@@ -2282,6 +2323,12 @@ class Pregel(
|
||||
"source": "fork",
|
||||
"step": step + 1,
|
||||
"parents": saved.metadata.get("parents", {}),
|
||||
# The copy has the same values and the same parent.
|
||||
**{
|
||||
k: v
|
||||
for k, v in saved.metadata.items()
|
||||
if k == "counters_since_delta_snapshot"
|
||||
},
|
||||
},
|
||||
{},
|
||||
)
|
||||
|
||||
@@ -147,6 +147,60 @@ async def test_predicate_fires_on_supersteps_overflow() -> None:
|
||||
assert "x" not in result2
|
||||
|
||||
|
||||
def _delta_counters(saver: InMemorySaver, config: Any) -> dict[str, list[int]]:
|
||||
tup = saver.get_tuple(config)
|
||||
assert tup is not None
|
||||
counters = tup.metadata.get("counters_since_delta_snapshot") or {}
|
||||
return {ch: list(c) for ch, c in counters.items()}
|
||||
|
||||
|
||||
_UPDATE_PATHS_WITHOUT_THE_SNAPSHOT_PLAN = pytest.mark.parametrize(
|
||||
("values", "as_node", "supersteps"),
|
||||
[
|
||||
(None, END, 1),
|
||||
(None, "__copy__", 0),
|
||||
({"a": []}, "__input__", 1),
|
||||
],
|
||||
ids=["clear as END", "copy", "update as input"],
|
||||
)
|
||||
|
||||
|
||||
@_UPDATE_PATHS_WITHOUT_THE_SNAPSHOT_PLAN
|
||||
def test_update_state_path_keeps_delta_counters(
|
||||
values: Any, as_node: str, supersteps: int
|
||||
) -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _build_two_channel_graph(saver)
|
||||
config = {"configurable": {"thread_id": "counters"}}
|
||||
graph.invoke({"a": ["seed-a"], "b": ["seed-b"]}, config)
|
||||
before = _delta_counters(saver, config)
|
||||
assert set(before) == {"a", "b"}, f"both channels need live counters: {before}"
|
||||
|
||||
updated = graph.update_state(config, values, as_node=as_node)
|
||||
|
||||
assert _delta_counters(saver, updated) == {
|
||||
ch: [u, s + supersteps] for ch, (u, s) in before.items()
|
||||
}
|
||||
|
||||
|
||||
@_UPDATE_PATHS_WITHOUT_THE_SNAPSHOT_PLAN
|
||||
async def test_aupdate_state_path_keeps_delta_counters(
|
||||
values: Any, as_node: str, supersteps: int
|
||||
) -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _build_two_channel_graph(saver)
|
||||
config = {"configurable": {"thread_id": "counters"}}
|
||||
await graph.ainvoke({"a": ["seed-a"], "b": ["seed-b"]}, config)
|
||||
before = _delta_counters(saver, config)
|
||||
assert set(before) == {"a", "b"}, f"both channels need live counters: {before}"
|
||||
|
||||
updated = await graph.aupdate_state(config, values, as_node=as_node)
|
||||
|
||||
assert _delta_counters(saver, updated) == {
|
||||
ch: [u, s + supersteps] for ch, (u, s) in before.items()
|
||||
}
|
||||
|
||||
|
||||
async def test_counter_reset_after_supersteps_snapshot() -> None:
|
||||
"""After the supersteps bound triggers a snapshot, the counters for
|
||||
that channel reset. Verify by using a bound higher than one run's
|
||||
|
||||
Reference in New Issue
Block a user