mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-26 17:42:24 +02:00
Merge pull request #2331 from langchain-ai/nc/4nov/get-state-latest-next
lib: When getting latest state, alst make `next` reflect pending writes
This commit is contained in:
@@ -65,9 +65,11 @@ from langgraph.constants import (
|
||||
CONFIG_KEY_STREAM,
|
||||
CONFIG_KEY_STREAM_WRITER,
|
||||
CONFIG_KEY_TASK_ID,
|
||||
ERROR,
|
||||
INTERRUPT,
|
||||
NS_END,
|
||||
NS_SEP,
|
||||
SCHEDULED,
|
||||
)
|
||||
from langgraph.errors import (
|
||||
ErrorCode,
|
||||
@@ -510,14 +512,16 @@ class Pregel(PregelProtocol):
|
||||
)
|
||||
# apply pending writes
|
||||
if apply_pending_writes and saved.pending_writes:
|
||||
for tid, *t in saved.pending_writes:
|
||||
next_tasks[tid].writes.append(t) # type: ignore[arg-type]
|
||||
for tid, k, v in saved.pending_writes:
|
||||
if k in (ERROR, INTERRUPT, SCHEDULED):
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
if tasks := [t for t in next_tasks.values() if t.writes]:
|
||||
apply_writes(saved.checkpoint, channels, tasks, None)
|
||||
# assemble the state snapshot
|
||||
return StateSnapshot(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(t.name for t in next_tasks.values()),
|
||||
tuple(t.name for t in next_tasks.values() if not t.writes),
|
||||
patch_checkpoint_map(saved.config, saved.metadata),
|
||||
saved.metadata,
|
||||
saved.checkpoint["ts"],
|
||||
@@ -608,14 +612,16 @@ class Pregel(PregelProtocol):
|
||||
)
|
||||
# apply pending writes
|
||||
if apply_pending_writes and saved.pending_writes:
|
||||
for tid, *t in saved.pending_writes:
|
||||
next_tasks[tid].writes.append(t) # type: ignore[arg-type]
|
||||
for tid, k, v in saved.pending_writes:
|
||||
if k in (ERROR, INTERRUPT, SCHEDULED):
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
if tasks := [t for t in next_tasks.values() if t.writes]:
|
||||
apply_writes(saved.checkpoint, channels, tasks, None)
|
||||
# assemble the state snapshot
|
||||
return StateSnapshot(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(t.name for t in next_tasks.values()),
|
||||
tuple(t.name for t in next_tasks.values() if not t.writes),
|
||||
patch_checkpoint_map(saved.config, saved.metadata),
|
||||
saved.metadata,
|
||||
saved.checkpoint["ts"],
|
||||
|
||||
@@ -1517,7 +1517,7 @@ def test_pending_writes_resume(
|
||||
state = graph.get_state(thread1)
|
||||
assert state is not None
|
||||
assert state.values == {"value": 3}
|
||||
assert state.next == ("one", "two")
|
||||
assert state.next == ("two",)
|
||||
assert state.tasks == (
|
||||
PregelTask(AnyStr(), "one", (PULL, "one"), result={"value": 2}),
|
||||
PregelTask(AnyStr(), "two", (PULL, "two"), 'ConnectionError("I\'m not good")'),
|
||||
@@ -1533,6 +1533,7 @@ def test_pending_writes_resume(
|
||||
state = graph.get_state(state.config)
|
||||
assert state is not None
|
||||
assert state.values == {"value": 1}
|
||||
assert state.next == ("one", "two")
|
||||
# should contain pending write of "one"
|
||||
checkpoint = checkpointer.get_tuple(thread1)
|
||||
assert checkpoint is not None
|
||||
|
||||
@@ -487,10 +487,7 @@ async def test_cancel_graph_astream(checkpointer_name: str) -> None:
|
||||
state = await graph.aget_state(thread1)
|
||||
assert state is not None
|
||||
assert state.values == {"value": 3} # 1 + 2
|
||||
assert state.next == (
|
||||
"aparallelwhile",
|
||||
"alittlewhile",
|
||||
)
|
||||
assert state.next == ("aparallelwhile",)
|
||||
assert state.metadata == {
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
@@ -1727,7 +1724,7 @@ async def test_pending_writes_resume(
|
||||
state = await graph.aget_state(thread1)
|
||||
assert state is not None
|
||||
assert state.values == {"value": 3}
|
||||
assert state.next == ("one", "two")
|
||||
assert state.next == ("two",)
|
||||
assert state.tasks == (
|
||||
PregelTask(AnyStr(), "one", (PULL, "one"), result={"value": 2}),
|
||||
PregelTask(
|
||||
@@ -1748,6 +1745,7 @@ async def test_pending_writes_resume(
|
||||
state = await graph.aget_state(state.config)
|
||||
assert state is not None
|
||||
assert state.values == {"value": 1}
|
||||
assert state.next == ("one", "two")
|
||||
# should contain pending write of "one"
|
||||
checkpoint = await checkpointer.aget_tuple(thread1)
|
||||
assert checkpoint is not None
|
||||
|
||||
Reference in New Issue
Block a user