From 8eea7ac401e013a10234c6389659cc214432fd7e Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Dec 2024 08:15:01 -0800 Subject: [PATCH 1/2] lib: Call ensure_config in state crud methods - this ensures that config from context vars is merged in --- libs/langgraph/langgraph/pregel/__init__.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 62506142d..2cc462b0f 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -673,6 +673,7 @@ class Pregel(PregelProtocol): self, config: RunnableConfig, *, subgraphs: bool = False ) -> StateSnapshot: """Get the current state of the graph.""" + config = ensure_config(config) checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) @@ -710,6 +711,7 @@ class Pregel(PregelProtocol): self, config: RunnableConfig, *, subgraphs: bool = False ) -> StateSnapshot: """Get the current state of the graph.""" + config = ensure_config(config) checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) @@ -751,6 +753,7 @@ class Pregel(PregelProtocol): before: Optional[RunnableConfig] = None, limit: Optional[int] = None, ) -> Iterator[StateSnapshot]: + config = ensure_config(config) """Get the history of the state of the graph.""" checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer @@ -800,6 +803,7 @@ class Pregel(PregelProtocol): before: Optional[RunnableConfig] = None, limit: Optional[int] = None, ) -> AsyncIterator[StateSnapshot]: + config = ensure_config(config) """Get the history of the state of the graph.""" checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer @@ -855,6 +859,7 @@ class Pregel(PregelProtocol): node `as_node`. If `as_node` is not provided, it will be set to the last node that updated the state, if not ambiguous. """ + config = ensure_config(config) checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) @@ -1130,6 +1135,7 @@ class Pregel(PregelProtocol): values: dict[str, Any] | Any, as_node: Optional[str] = None, ) -> RunnableConfig: + config = ensure_config(config) checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) From e5b00cdd1ef3f80a0aaf2ac9a2edeadcdab7763b Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Dec 2024 08:30:10 -0800 Subject: [PATCH 2/2] Fix --- libs/langgraph/langgraph/pregel/__init__.py | 16 ++++++---------- 1 file changed, 6 insertions(+), 10 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 2cc462b0f..e714afe21 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -673,8 +673,7 @@ class Pregel(PregelProtocol): self, config: RunnableConfig, *, subgraphs: bool = False ) -> StateSnapshot: """Get the current state of the graph.""" - config = ensure_config(config) - checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( + checkpointer: Optional[BaseCheckpointSaver] = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) if not checkpointer: @@ -711,8 +710,7 @@ class Pregel(PregelProtocol): self, config: RunnableConfig, *, subgraphs: bool = False ) -> StateSnapshot: """Get the current state of the graph.""" - config = ensure_config(config) - checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( + checkpointer: Optional[BaseCheckpointSaver] = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) if not checkpointer: @@ -755,7 +753,7 @@ class Pregel(PregelProtocol): ) -> Iterator[StateSnapshot]: config = ensure_config(config) """Get the history of the state of the graph.""" - checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( + checkpointer: Optional[BaseCheckpointSaver] = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) if not checkpointer: @@ -805,7 +803,7 @@ class Pregel(PregelProtocol): ) -> AsyncIterator[StateSnapshot]: config = ensure_config(config) """Get the history of the state of the graph.""" - checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( + checkpointer: Optional[BaseCheckpointSaver] = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) if not checkpointer: @@ -859,8 +857,7 @@ class Pregel(PregelProtocol): node `as_node`. If `as_node` is not provided, it will be set to the last node that updated the state, if not ambiguous. """ - config = ensure_config(config) - checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( + checkpointer: Optional[BaseCheckpointSaver] = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) if not checkpointer: @@ -1135,8 +1132,7 @@ class Pregel(PregelProtocol): values: dict[str, Any] | Any, as_node: Optional[str] = None, ) -> RunnableConfig: - config = ensure_config(config) - checkpointer: Optional[BaseCheckpointSaver] = config[CONF].get( + checkpointer: Optional[BaseCheckpointSaver] = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) if not checkpointer: