From 4dc27b98f140bbe6ba8eca1d30c3b596cae195b5 Mon Sep 17 00:00:00 2001 From: Vadym Barda Date: Thu, 8 Aug 2024 11:55:55 -0400 Subject: [PATCH] langgraph, checkpoint-postgres: propagate new versions in update_state (#1270) * langgraph, checkpoint-postgres: propagate new versions in update_state --- .../langgraph/checkpoint/postgres/__init__.py | 1 - .../langgraph/checkpoint/postgres/aio.py | 1 - .../langgraph/checkpoint/postgres/base.py | 4 ---- libs/langgraph/langgraph/pregel/__init__.py | 17 +++++++++++++++-- libs/langgraph/langgraph/pregel/loop.py | 16 ++++------------ libs/langgraph/langgraph/pregel/utils.py | 19 +++++++++++++++++++ 6 files changed, 38 insertions(+), 20 deletions(-) create mode 100644 libs/langgraph/langgraph/pregel/utils.py diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index 3c12c4d24..ee7abe120 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -194,7 +194,6 @@ class PostgresSaver(BasePostgresSaver): thread_id, checkpoint_ns, copy.pop("channel_values"), - copy["channel_versions"], new_versions, ), ) diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index dd150535c..3622150e3 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -201,7 +201,6 @@ class AsyncPostgresSaver(BasePostgresSaver): thread_id, checkpoint_ns, copy.pop("channel_values"), - copy["channel_versions"], new_versions, ), ) diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py index 97b13ed30..e9ca12305 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py @@ -150,14 +150,10 @@ class BasePostgresSaver(BaseCheckpointSaver): checkpoint_ns: str, values: dict[str, Any], versions: dict[str, str], - new_versions: Optional[dict[str, str]], ) -> list[tuple[str, str, str, str, str, bytes]]: if not versions: return [] - if new_versions: - versions = new_versions - return [ ( thread_id, diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index b4c64af1f..8f9583964 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -100,6 +100,7 @@ from langgraph.pregel.types import ( StateSnapshot, StreamMode, ) +from langgraph.pregel.utils import get_new_channel_versions from langgraph.pregel.validate import validate_graph, validate_keys from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry @@ -517,6 +518,9 @@ class Pregel( # get last checkpoint saved = self.checkpointer.get_tuple(config) checkpoint = copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint() + checkpoint_previous_versions = ( + saved.checkpoint["channel_versions"] if saved else {} + ) step = saved.metadata.get("step", -1) if saved else -1 # merge configurable fields with previous checkpoint config checkpoint_config = { @@ -606,6 +610,9 @@ class Pregel( checkpoint, channels, [task], self.checkpointer.get_next_version ) + new_versions = get_new_channel_versions( + checkpoint_previous_versions, checkpoint["channel_versions"] + ) return self.checkpointer.put( checkpoint_config, create_checkpoint(checkpoint, channels, step + 1), @@ -614,7 +621,7 @@ class Pregel( "step": step + 1, "writes": {as_node: values}, }, - {}, + new_versions, ) async def aupdate_state( @@ -629,6 +636,9 @@ class Pregel( # get last checkpoint saved = await self.checkpointer.aget_tuple(config) checkpoint = copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint() + checkpoint_previous_versions = ( + saved.checkpoint["channel_versions"] if saved else {} + ) step = saved.metadata.get("step", -1) if saved else -1 # merge configurable fields with previous checkpoint config checkpoint_config = { @@ -716,6 +726,9 @@ class Pregel( checkpoint, channels, [task], self.checkpointer.get_next_version ) + new_versions = get_new_channel_versions( + checkpoint_previous_versions, checkpoint["channel_versions"] + ) return await self.checkpointer.aput( checkpoint_config, create_checkpoint(checkpoint, channels, step + 1), @@ -724,7 +737,7 @@ class Pregel( "step": step + 1, "writes": {as_node: values}, }, - {}, + new_versions, ) def _defaults( diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 4dccb6429..0508cd942 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -61,6 +61,7 @@ from langgraph.pregel.executor import ( ) from langgraph.pregel.io import map_input, map_output_updates, map_output_values, single from langgraph.pregel.types import PregelExecutableTask +from langgraph.pregel.utils import get_new_channel_versions if TYPE_CHECKING: from langgraph.pregel import Pregel @@ -336,18 +337,9 @@ class PregelLoop: } channel_versions = self.checkpoint["channel_versions"].copy() - if self.checkpoint_previous_versions: - new_versions = { - k: v - for k, v in channel_versions.items() - if k not in self.checkpoint_previous_versions - or ( - k in self.checkpoint_previous_versions - and v > self.checkpoint_previous_versions[k] - ) - } - else: - new_versions = channel_versions + new_versions = get_new_channel_versions( + self.checkpoint_previous_versions, channel_versions + ) self.checkpoint_previous_versions = channel_versions diff --git a/libs/langgraph/langgraph/pregel/utils.py b/libs/langgraph/langgraph/pregel/utils.py new file mode 100644 index 000000000..d3d0d989f --- /dev/null +++ b/libs/langgraph/langgraph/pregel/utils.py @@ -0,0 +1,19 @@ +from langgraph.checkpoint.base import ChannelVersions + + +def get_new_channel_versions( + previous_versions: ChannelVersions, current_versions: ChannelVersions +) -> ChannelVersions: + """Get new channel versions.""" + if previous_versions: + version_type = type(next(iter(current_versions.values()), None)) + null_version = version_type() + new_versions = { + k: v + for k, v in current_versions.items() + if v > previous_versions.get(k, null_version) + } + else: + new_versions = current_versions + + return new_versions