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:
Nuno Campos
2024-11-04 16:16:53 -08:00
committed by GitHub
3 changed files with 17 additions and 12 deletions
+12 -6
View File
@@ -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"],
+2 -1
View File
@@ -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
+3 -5
View File
@@ -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