From 48865daf0220e7c4ca1023696a2d6ee1724d5a4b Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Mon, 6 May 2024 11:38:28 -0700 Subject: [PATCH] Fix checkpoint lineage for updates/resumes --- langgraph/checkpoint/base.py | 3 +- langgraph/pregel/__init__.py | 20 ++++----- tests/test_pregel.py | 54 +++++++++++++++++++++++- tests/test_pregel_async.py | 79 +++++++++++++++++++++++++++++++++++- 4 files changed, 141 insertions(+), 15 deletions(-) diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index 089d7cc19..5d197e19e 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -158,7 +158,7 @@ class BaseCheckpointSaver(ABC): async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: raise NotImplementedError - async def alist( + def alist( self, config: RunnableConfig, *, @@ -166,6 +166,7 @@ class BaseCheckpointSaver(ABC): limit: Optional[int] = None, ) -> AsyncIterator[CheckpointTuple]: raise NotImplementedError + yield async def aput( self, diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index 2d50e4496..75321bf46 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -333,7 +333,6 @@ class Pregel( saved = self.checkpointer.get_tuple(config) checkpoint = saved.checkpoint if saved else empty_checkpoint() - config = saved.config if saved else config with ChannelsManager(self.channels, checkpoint) as channels: _, next_tasks = _prepare_next_tasks( checkpoint, self.nodes, channels, for_execution=False @@ -341,8 +340,9 @@ class Pregel( return StateSnapshot( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), - config, + saved.config if saved else config, saved.metadata if saved else None, + saved.parent_config if saved else None, ) async def aget_state(self, config: RunnableConfig) -> StateSnapshot: @@ -352,7 +352,6 @@ class Pregel( saved = await self.checkpointer.aget_tuple(config) checkpoint = saved.checkpoint if saved else empty_checkpoint() - config = saved.config if saved else config async with AsyncChannelsManager(self.channels, checkpoint) as channels: _, next_tasks = _prepare_next_tasks( checkpoint, self.nodes, channels, for_execution=False @@ -360,8 +359,9 @@ class Pregel( return StateSnapshot( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), - config, + saved.config if saved else config, saved.metadata if saved else None, + saved.parent_config if saved else None, ) def get_state_history( @@ -650,13 +650,9 @@ class Pregel( # copy nodes to ignore mutations during execution processes = {**self.nodes} # get checkpoint from saver, or create an empty one - checkpoint_config = config - saved = ( - self.checkpointer.get_tuple(checkpoint_config) - if self.checkpointer - else None - ) + saved = self.checkpointer.get_tuple(config) if self.checkpointer else None checkpoint = saved.checkpoint if saved else empty_checkpoint() + checkpoint_config = saved.config if saved else config start = saved.metadata.get("step", -2) + 1 if saved else -1 # create channels from checkpoint with ChannelsManager( @@ -918,13 +914,13 @@ class Pregel( # copy nodes to ignore mutations during execution processes = {**self.nodes} # get checkpoint from saver, or create an empty one - checkpoint_config = config saved = ( - await self.checkpointer.aget_tuple(checkpoint_config) + await self.checkpointer.aget_tuple(config) if self.checkpointer else None ) checkpoint = saved.checkpoint if saved else empty_checkpoint() + checkpoint_config = saved.config if saved else config start = saved.metadata.get("step", -2) + 1 if saved else -1 # create channels from checkpoint async with AsyncChannelsManager(self.channels, checkpoint) as channels: diff --git a/tests/test_pregel.py b/tests/test_pregel.py index bf295801f..a4e84956d 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -717,7 +717,7 @@ def test_invoke_checkpoint_sqlite(mocker: MockerFixture) -> None: # total is now 2+5=7, so output would be 7+4=11, but raises ValueError with pytest.raises(ValueError): app.invoke(4, thread_1) - # checkpoint is not updated + # checkpoint is updated with new input state = app.get_state(thread_1) assert state is not None assert state.values.get("total") == 7 @@ -782,6 +782,11 @@ def test_invoke_checkpoint_sqlite(mocker: MockerFixture) -> None: thread_1_next_config["configurable"]["thread_ts"] > thread_1_history[0].config["configurable"]["thread_ts"] ) + # update makes new checkpoint child of the previous one + assert ( + app.get_state(thread_1_next_config).parent_config + == thread_1_history[1].config + ) # 1 more checkpoint in history assert len(list(app.get_state_history(thread_1))) == 8 assert Counter( @@ -3511,6 +3516,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, metadata={"source": "loop", "step": 0}, + parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3522,6 +3528,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=tool_two.checkpointer.get_tuple(thread1).config, metadata={"source": "loop", "step": 1}, + parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3535,6 +3542,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, metadata={"source": "loop", "step": 0}, + parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3546,6 +3554,42 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=tool_two.checkpointer.get_tuple(thread2).config, metadata={"source": "loop", "step": 1}, + parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config, + ) + + thread3 = {"configurable": {"thread_id": "3"}} + # stop when about to enter node + assert tool_two.invoke({"my_key": "value", "market": "US"}, thread3) == { + "my_key": "value", + "market": "US", + } + assert tool_two.get_state(thread3) == StateSnapshot( + values={"my_key": "value", "market": "US"}, + next=("tool_two_fast",), + config=tool_two.checkpointer.get_tuple(thread3).config, + metadata={"source": "loop", "step": 0}, + parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config, + ) + # update state + tool_two.update_state(thread3, {"my_key": "key"}) # appends to my_key + assert tool_two.get_state(thread3) == StateSnapshot( + values={"my_key": "valuekey", "market": "US"}, + next=("tool_two_fast",), + config=tool_two.checkpointer.get_tuple(thread3).config, + metadata={"source": "update", "step": 1}, + parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config, + ) + # resume, for same result as above + assert tool_two.invoke(None, thread3, debug=1) == { + "my_key": "valuekey fast", + "market": "US", + } + assert tool_two.get_state(thread3) == StateSnapshot( + values={"my_key": "valuekey fast", "market": "US"}, + next=(), + config=tool_two.checkpointer.get_tuple(thread3).config, + metadata={"source": "loop", "step": 2}, + parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config, ) @@ -3735,6 +3779,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, metadata={"source": "loop", "step": 1}, + parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3746,6 +3791,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=tool_two.checkpointer.get_tuple(thread1).config, metadata={"source": "loop", "step": 3}, + parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3759,6 +3805,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, metadata={"source": "loop", "step": 1}, + parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3770,6 +3817,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=tool_two.checkpointer.get_tuple(thread2).config, metadata={"source": "loop", "step": 3}, + parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config, ) with SqliteSaver.from_conn_string(":memory:") as saver: @@ -3792,6 +3840,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, metadata={"source": "loop", "step": 1}, + parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3803,6 +3852,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=tool_two.checkpointer.get_tuple(thread1).config, metadata={"source": "loop", "step": 3}, + parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3816,6 +3866,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, metadata={"source": "loop", "step": 1}, + parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3827,6 +3878,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=tool_two.checkpointer.get_tuple(thread2).config, metadata={"source": "loop", "step": 3}, + parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config, ) diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 4a01b1631..d0cf8a161 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -2926,7 +2926,7 @@ async def test_in_one_fan_out_out_one_graph_state() -> None: ] -async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: +async def test_start_branch_then() -> None: class State(TypedDict): my_key: Annotated[str, operator.add] market: str @@ -2968,6 +2968,9 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, metadata={"source": "loop", "step": 0}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread1, limit=2) + ][-1].config, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -2979,6 +2982,9 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, metadata={"source": "loop", "step": 1}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread1, limit=2) + ][-1].config, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -2992,6 +2998,9 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, metadata={"source": "loop", "step": 0}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread2, limit=2) + ][-1].config, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3003,6 +3012,50 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, metadata={"source": "loop", "step": 1}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread2, limit=2) + ][-1].config, + ) + + thread3 = {"configurable": {"thread_id": "3"}} + # stop when about to enter node + assert await tool_two.ainvoke({"my_key": "value", "market": "US"}, thread3) == { + "my_key": "value", + "market": "US", + } + assert await tool_two.aget_state(thread3) == StateSnapshot( + values={"my_key": "value", "market": "US"}, + next=("tool_two_fast",), + config=(await tool_two.checkpointer.aget_tuple(thread3)).config, + metadata={"source": "loop", "step": 0}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread3, limit=2) + ][-1].config, + ) + # update state + await tool_two.aupdate_state(thread3, {"my_key": "key"}) # appends to my_key + assert await tool_two.aget_state(thread3) == StateSnapshot( + values={"my_key": "valuekey", "market": "US"}, + next=("tool_two_fast",), + config=(await tool_two.checkpointer.aget_tuple(thread3)).config, + metadata={"source": "update", "step": 1}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread3, limit=2) + ][-1].config, + ) + # resume, for same result as above + assert await tool_two.ainvoke(None, thread3, debug=1) == { + "my_key": "valuekey fast", + "market": "US", + } + assert await tool_two.aget_state(thread3) == StateSnapshot( + values={"my_key": "valuekey fast", "market": "US"}, + next=(), + config=(await tool_two.checkpointer.aget_tuple(thread3)).config, + metadata={"source": "loop", "step": 2}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread3, limit=2) + ][-1].config, ) @@ -3177,6 +3230,9 @@ async def test_branch_then() -> None: next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, metadata={"source": "loop", "step": 1}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread1, limit=2) + ][-1].config, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -3188,6 +3244,9 @@ async def test_branch_then() -> None: next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, metadata={"source": "loop", "step": 3}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread1, limit=2) + ][-1].config, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3201,6 +3260,9 @@ async def test_branch_then() -> None: next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, metadata={"source": "loop", "step": 1}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread2, limit=2) + ][-1].config, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3212,6 +3274,9 @@ async def test_branch_then() -> None: next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, metadata={"source": "loop", "step": 3}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread2, limit=2) + ][-1].config, ) async with AsyncSqliteSaver.from_conn_string(":memory:") as saver: @@ -3234,6 +3299,9 @@ async def test_branch_then() -> None: next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, metadata={"source": "loop", "step": 1}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread1, limit=2) + ][-1].config, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -3245,6 +3313,9 @@ async def test_branch_then() -> None: next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, metadata={"source": "loop", "step": 3}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread1, limit=2) + ][-1].config, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3258,6 +3329,9 @@ async def test_branch_then() -> None: next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, metadata={"source": "loop", "step": 1}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread2, limit=2) + ][-1].config, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3269,6 +3343,9 @@ async def test_branch_then() -> None: next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, metadata={"source": "loop", "step": 3}, + parent_config=[ + c async for c in tool_two.checkpointer.alist(thread2, limit=2) + ][-1].config, )