From b71fd5092b647a2fd4eb61ef2e52073dd719b300 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Mon, 4 Nov 2024 16:03:34 -0800 Subject: [PATCH] lib: When getting latest state, alst make `next` reflect pending writes - ie. tasks already executed should not show up in `next` list --- libs/langgraph/langgraph/pregel/__init__.py | 18 ++++++++++++------ libs/langgraph/tests/test_pregel.py | 3 ++- libs/langgraph/tests/test_pregel_async.py | 8 +++----- 3 files changed, 17 insertions(+), 12 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 9aa3bc915..0755a4473 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -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"], diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index c9de54ec8..dd3a9f714 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -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 diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 1235d5cc0..87d2679e3 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -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