diff --git a/langgraph/checkpoint/aiosqlite.py b/langgraph/checkpoint/aiosqlite.py index 314448b01..c9a322561 100644 --- a/langgraph/checkpoint/aiosqlite.py +++ b/langgraph/checkpoint/aiosqlite.py @@ -1,7 +1,7 @@ import asyncio from contextlib import AbstractAsyncContextManager from types import TracebackType -from typing import Any, AsyncIterator, Optional +from typing import AsyncIterator, Optional import aiosqlite from langchain_core.runnables import RunnableConfig @@ -10,6 +10,7 @@ from typing_extensions import Self from langgraph.checkpoint.base import ( BaseCheckpointSaver, Checkpoint, + CheckpointMetadata, CheckpointTuple, SerializerProtocol, ) @@ -164,7 +165,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): return CheckpointTuple( config, self.serde.loads(value[0]), - self.serde.loads(value[2]) if value[2] is not None else None, + self.serde.loads(value[2]) if value[2] is not None else {}, { "configurable": { "thread_id": config["configurable"]["thread_id"], @@ -188,7 +189,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): } }, self.serde.loads(value[3]), - self.serde.loads(value[4]) if value[4] is not None else None, + self.serde.loads(value[4]) if value[4] is not None else {}, { "configurable": { "thread_id": value[0], @@ -242,7 +243,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): yield CheckpointTuple( {"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}}, self.serde.loads(value), - self.serde.loads(metadata) if metadata is not None else None, + self.serde.loads(metadata) if metadata is not None else {}, {"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}} if parent_ts else None, @@ -252,7 +253,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): self, config: RunnableConfig, checkpoint: Checkpoint, - metadata: Optional[dict[str, Any]] = None, + metadata: CheckpointMetadata, ) -> RunnableConfig: """Save a checkpoint to the database asynchronously. @@ -274,7 +275,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): checkpoint["ts"], config["configurable"].get("thread_ts"), self.serde.dumps(checkpoint), - self.serde.dumps(metadata) if metadata is not None else None, + self.serde.dumps(metadata), ), ): await self.conn.commit() diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index 70734f1e6..089d7cc19 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -5,6 +5,7 @@ from typing import ( Any, AsyncIterator, Iterator, + Literal, NamedTuple, Optional, TypedDict, @@ -16,6 +17,22 @@ from langgraph.serde.base import SerializerProtocol from langgraph.serde.jsonplus import JsonPlusSerializer +# Marked as total=False to allow for future expansion. +class CheckpointMetadata(TypedDict, total=False): + source: Literal["input", "loop", "update"] + """The source of the checkpoint. + - "input": The checkpoint was created from an input to invoke/stream/batch. + - "loop": The checkpoint was created from inside the pregel loop. + - "update": The checkpoint was created from a manual state update. + """ + step: int + """The step number of the checkpoint. + -1 for the first "input" checkpoint. + 0 for the first "loop" checkpoint. + ... for the nth checkpoint afterwards. + """ + + class Checkpoint(TypedDict): """State snapshot at a given point in time.""" @@ -73,7 +90,7 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint: class CheckpointTuple(NamedTuple): config: RunnableConfig checkpoint: Checkpoint - metadata: Optional[dict[str, Any]] + metadata: CheckpointMetadata parent_config: Optional[RunnableConfig] = None @@ -130,7 +147,7 @@ class BaseCheckpointSaver(ABC): self, config: RunnableConfig, checkpoint: Checkpoint, - metadata: dict[str, Any], + metadata: CheckpointMetadata, ) -> RunnableConfig: raise NotImplementedError @@ -154,6 +171,6 @@ class BaseCheckpointSaver(ABC): self, config: RunnableConfig, checkpoint: Checkpoint, - metadata: dict[str, Any], + metadata: CheckpointMetadata, ) -> RunnableConfig: raise NotImplementedError diff --git a/langgraph/checkpoint/memory.py b/langgraph/checkpoint/memory.py index 81eba2438..fb49bab58 100644 --- a/langgraph/checkpoint/memory.py +++ b/langgraph/checkpoint/memory.py @@ -1,12 +1,13 @@ import asyncio from collections import defaultdict -from typing import Any, AsyncIterator, Iterator, Optional +from typing import AsyncIterator, Iterator, Optional from langchain_core.runnables import RunnableConfig from langgraph.checkpoint.base import ( BaseCheckpointSaver, Checkpoint, + CheckpointMetadata, CheckpointTuple, SerializerProtocol, ) @@ -118,7 +119,7 @@ class MemorySaver(BaseCheckpointSaver): self, config: RunnableConfig, checkpoint: Checkpoint, - metadata: dict[str, Any] = None, + metadata: CheckpointMetadata, ) -> RunnableConfig: """Save a checkpoint to the in-memory storage. @@ -136,7 +137,7 @@ class MemorySaver(BaseCheckpointSaver): { checkpoint["ts"]: ( self.serde.dumps(checkpoint), - self.serde.dumps(metadata or {}), + self.serde.dumps(metadata), ) } ) @@ -184,8 +185,11 @@ class MemorySaver(BaseCheckpointSaver): return async def aput( - self, config: RunnableConfig, checkpoint: Checkpoint + self, + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: CheckpointMetadata, ) -> RunnableConfig: return await asyncio.get_running_loop().run_in_executor( - None, self.put, config, checkpoint + None, self.put, config, checkpoint, metadata ) diff --git a/langgraph/checkpoint/sqlite.py b/langgraph/checkpoint/sqlite.py index 3c2b63ca3..54feff576 100644 --- a/langgraph/checkpoint/sqlite.py +++ b/langgraph/checkpoint/sqlite.py @@ -10,6 +10,7 @@ from typing_extensions import Self from langgraph.checkpoint.base import ( BaseCheckpointSaver, Checkpoint, + CheckpointMetadata, CheckpointTuple, SerializerProtocol, ) @@ -220,7 +221,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): return CheckpointTuple( config, self.serde.loads(value[0]), - self.serde.loads(value[2]) if value[2] is not None else None, + self.serde.loads(value[2]) if value[2] is not None else {}, { "configurable": { "thread_id": config["configurable"]["thread_id"], @@ -244,7 +245,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): } }, self.serde.loads(value[3]), - self.serde.loads(value[4]) if value[4] is not None else None, + self.serde.loads(value[4]) if value[4] is not None else {}, { "configurable": { "thread_id": value[0], @@ -308,7 +309,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): yield CheckpointTuple( {"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}}, self.serde.loads(value), - self.serde.loads(metadata) if metadata is not None else None, + self.serde.loads(metadata) if metadata is not None else {}, { "configurable": { "thread_id": thread_id, @@ -323,7 +324,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): self, config: RunnableConfig, checkpoint: Checkpoint, - metadata: dict[str, Any] = None, + metadata: CheckpointMetadata, ) -> RunnableConfig: """Save a checkpoint to the database. @@ -355,7 +356,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): checkpoint["ts"], config["configurable"].get("thread_ts"), self.serde.dumps(checkpoint), - self.serde.dumps(metadata) if metadata is not None else None, + self.serde.dumps(metadata), ), ) return { diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index ffa14304c..de56f8925 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -342,7 +342,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, - saved.metadata or {}, + saved.metadata, ) async def aget_state(self, config: RunnableConfig) -> StateSnapshot: @@ -361,7 +361,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, - saved.metadata or {}, + saved.metadata, ) def get_state_history( @@ -386,7 +386,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, - metadata or {}, + metadata, parent_config, ) @@ -415,7 +415,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, - metadata or {}, + metadata, parent_config, ) @@ -433,8 +433,8 @@ class Pregel( raise ValueError("No checkpointer set") # get last checkpoint - checkpoint = self.checkpointer.get(config) - checkpoint = copy_checkpoint(checkpoint) if checkpoint else empty_checkpoint() + saved = self.checkpointer.get_tuple(config) + checkpoint = copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint() # find last node that updated the state, if not provided if as_node is None: last_seen_by_node = sorted( @@ -482,7 +482,14 @@ class Pregel( # apply to checkpoint and save _apply_writes(checkpoint, channels, task.writes) return self.checkpointer.put( - config, create_checkpoint(checkpoint, channels) + config, + create_checkpoint(checkpoint, channels), + { + "source": "update", + "step": saved.metadata.get("step", 0) + 1 + if saved.metadata + else None, + }, ) async def aupdate_state( @@ -495,8 +502,8 @@ class Pregel( raise ValueError("No checkpointer set") # get last checkpoint - checkpoint = await self.checkpointer.aget(config) - checkpoint = copy_checkpoint(checkpoint) if checkpoint else empty_checkpoint() + saved = await self.checkpointer.aget_tuple(config) + checkpoint = copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint() # find last node that updated the state, if not provided if as_node is None: last_seen_by_node = sorted( @@ -544,7 +551,12 @@ class Pregel( # apply to checkpoint and save _apply_writes(checkpoint, channels, task.writes) return await self.checkpointer.aput( - config, create_checkpoint(checkpoint, channels) + config, + create_checkpoint(checkpoint, channels), + { + "source": "update", + "step": saved.metadata.get("step", 0) + 1 if saved else None, + }, ) def _defaults( @@ -638,10 +650,12 @@ class Pregel( processes = {**self.nodes} # get checkpoint from saver, or create an empty one checkpoint_config = config - checkpoint = ( - self.checkpointer.get(checkpoint_config) if self.checkpointer else None + saved = ( + self.checkpointer.get_tuple(checkpoint_config) + if self.checkpointer + else None ) - checkpoint = checkpoint or empty_checkpoint() + checkpoint = saved.checkpoint if saved else empty_checkpoint() # create channels from checkpoint with ChannelsManager( self.channels, checkpoint @@ -667,7 +681,9 @@ class Pregel( # channel updates from step N are only visible in step N+1 # channels are guaranteed to be immutable for the duration of the step, # with channel updates applied only at the transition between steps - for step in range(config["recursion_limit"] + 1): + start = saved.metadata.get("step", -1) + 1 if saved else 0 + stop = start + config["recursion_limit"] + 1 + for step in range(start, stop): next_checkpoint, next_tasks = _prepare_next_tasks( checkpoint, processes, channels, for_execution=True ) @@ -771,7 +787,9 @@ class Pregel( if self.checkpointer is not None: checkpoint = create_checkpoint(checkpoint, channels) checkpoint_config = self.checkpointer.put( - checkpoint_config, checkpoint + checkpoint_config, + checkpoint, + {"source": "loop", "step": step}, ) if stream_mode == "debug": yield map_debug_checkpoint( @@ -865,12 +883,12 @@ class Pregel( processes = {**self.nodes} # get checkpoint from saver, or create an empty one checkpoint_config = config - checkpoint = ( - await self.checkpointer.aget(checkpoint_config) + saved = ( + await self.checkpointer.aget_tuple(checkpoint_config) if self.checkpointer else None ) - checkpoint = checkpoint or empty_checkpoint() + checkpoint = saved.checkpoint if saved else empty_checkpoint() # create channels from checkpoint async with AsyncChannelsManager(self.channels, checkpoint) as channels: # map inputs to channel updates @@ -894,7 +912,9 @@ class Pregel( # channel updates from step N are only visible in step N+1, # channels are guaranteed to be immutable for the duration of the step, # channel updates being applied only at the transition between steps - for step in range(config["recursion_limit"] + 1): + start = saved.metadata.get("step", -1) + 1 if saved else 0 + stop = start + config["recursion_limit"] + 1 + for step in range(start, stop): next_checkpoint, next_tasks = _prepare_next_tasks( checkpoint, processes, channels, for_execution=True ) @@ -1008,7 +1028,9 @@ class Pregel( if self.checkpointer is not None: checkpoint = create_checkpoint(checkpoint, channels) checkpoint_config = await self.checkpointer.aput( - checkpoint_config, checkpoint + checkpoint_config, + checkpoint, + {"source": "loop", "step": step}, ) if stream_mode == "debug": yield map_debug_checkpoint( diff --git a/tests/memory_assert.py b/tests/memory_assert.py index 625b6f12e..2429cf63d 100644 --- a/tests/memory_assert.py +++ b/tests/memory_assert.py @@ -3,6 +3,7 @@ from typing import Any, Optional from langgraph.checkpoint.base import ( Checkpoint, + CheckpointMetadata, SerializerProtocol, copy_checkpoint, ) @@ -30,7 +31,12 @@ class MemorySaverAssertImmutable(MemorySaver): super().__init__(serde=serde) self.storage_for_copies = defaultdict(dict) - def put(self, config: dict, checkpoint: Checkpoint) -> None: + def put( + self, + config: dict, + checkpoint: Checkpoint, + metadata: Optional[CheckpointMetadata] = None, + ) -> None: # assert checkpoint hasn't been modified since last written thread_id = config["configurable"]["thread_id"] if saved := super().get(config): @@ -39,4 +45,4 @@ class MemorySaverAssertImmutable(MemorySaver): checkpoint ) # call super to write checkpoint - return super().put(config, checkpoint) + return super().put(config, checkpoint, metadata) diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 9180e71c9..edccfdd78 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -1202,7 +1202,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) assert ( app_w_interrupt.checkpointer.get_tuple(config).config["configurable"][ @@ -1236,7 +1236,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 1}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1320,7 +1320,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 4}, ) # test state get/update methods with interrupt_before @@ -1356,7 +1356,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) app_w_interrupt.update_state( @@ -1384,7 +1384,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 1}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1468,7 +1468,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 4}, ) # test re-invoke to continue with interrupt_before @@ -1504,7 +1504,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1836,7 +1836,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) app_w_interrupt.update_state( @@ -1862,7 +1862,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 2}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1921,7 +1921,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 5}, ) # test state get/update methods with interrupt_before @@ -1956,7 +1956,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) app_w_interrupt.update_state( @@ -1982,7 +1982,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 2}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2041,7 +2041,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 5}, ) # test w interrupt before all @@ -2064,7 +2064,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2087,7 +2087,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2126,7 +2126,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 2}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2171,7 +2171,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2210,7 +2210,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None: }, next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 2}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -3025,7 +3025,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # modify ai message @@ -3051,7 +3051,7 @@ def test_message_graph( ], next=("action",), config=next_config, - metadata={}, + metadata={"source": "update", "step": 2}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -3117,7 +3117,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 4}, ) app_w_interrupt.update_state( @@ -3153,7 +3153,7 @@ def test_message_graph( ], next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 5}, ) app_w_interrupt = workflow.compile( @@ -3199,7 +3199,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # modify ai message @@ -3228,7 +3228,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 2}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -3294,7 +3294,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 4}, ) app_w_interrupt.update_state( @@ -3330,7 +3330,7 @@ def test_message_graph( ], next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 5}, ) # add an extra message as if it came from "action" node @@ -3366,7 +3366,7 @@ def test_message_graph( ], next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 6}, ) @@ -3489,7 +3489,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value", "market": "DE"}, next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3500,7 +3500,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value slow", "market": "DE"}, next=(), config=tool_two.checkpointer.get_tuple(thread1).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3513,7 +3513,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value", "market": "US"}, next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3524,7 +3524,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value fast", "market": "US"}, next=(), config=tool_two.checkpointer.get_tuple(thread2).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) @@ -3713,7 +3713,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3724,7 +3724,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=tool_two.checkpointer.get_tuple(thread1).config, - metadata={}, + metadata={"source": "loop", "step": 3}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3737,7 +3737,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3748,7 +3748,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=tool_two.checkpointer.get_tuple(thread2).config, - metadata={}, + metadata={"source": "loop", "step": 3}, ) with SqliteSaver.from_conn_string(":memory:") as saver: @@ -3770,7 +3770,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3781,7 +3781,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=tool_two.checkpointer.get_tuple(thread1).config, - metadata={}, + metadata={"source": "loop", "step": 3}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3794,7 +3794,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3805,7 +3805,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=tool_two.checkpointer.get_tuple(thread2).config, - metadata={}, + metadata={"source": "loop", "step": 3}, ) diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 32b19b279..2a79de197 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -1280,7 +1280,7 @@ async def test_conditional_graph() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) await app_w_interrupt.aupdate_state( @@ -1308,7 +1308,7 @@ async def test_conditional_graph() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 1}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -1392,7 +1392,7 @@ async def test_conditional_graph() -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 4}, ) # test state get/update methods with interrupt_before @@ -1431,7 +1431,7 @@ async def test_conditional_graph() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) await app_w_interrupt.aupdate_state( @@ -1459,7 +1459,7 @@ async def test_conditional_graph() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 1}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -1543,7 +1543,7 @@ async def test_conditional_graph() -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 4}, ) # test re-invoke to continue with interrupt_before @@ -1582,7 +1582,7 @@ async def test_conditional_graph() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -1902,7 +1902,7 @@ async def test_conditional_graph_state() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) await app_w_interrupt.aupdate_state( @@ -1928,7 +1928,7 @@ async def test_conditional_graph_state() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 2}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -1987,7 +1987,7 @@ async def test_conditional_graph_state() -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 5}, ) # test state get/update methods with interrupt_before @@ -2024,7 +2024,7 @@ async def test_conditional_graph_state() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) await app_w_interrupt.aupdate_state( @@ -2050,7 +2050,7 @@ async def test_conditional_graph_state() -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 2}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -2109,7 +2109,7 @@ async def test_conditional_graph_state() -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "update", "step": 5}, ) @@ -2715,7 +2715,7 @@ async def test_message_graph() -> None: ], next=("action",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # modify ai message @@ -2743,7 +2743,7 @@ async def test_message_graph() -> None: ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 2}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -2796,7 +2796,7 @@ async def test_message_graph() -> None: ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "loop", "step": 4}, ) await app_w_interrupt.aupdate_state( @@ -2830,7 +2830,7 @@ async def test_message_graph() -> None: ], next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, - metadata={}, + metadata={"source": "update", "step": 5}, ) @@ -2947,7 +2947,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value", "market": "DE"}, next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -2958,7 +2958,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value slow", "market": "DE"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -2971,7 +2971,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value", "market": "US"}, next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, - metadata={}, + metadata={"source": "loop", "step": 0}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -2982,7 +2982,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None: values={"my_key": "value fast", "market": "US"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) @@ -3156,7 +3156,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -3167,7 +3167,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, - metadata={}, + metadata={"source": "loop", "step": 3}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3180,7 +3180,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3191,7 +3191,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, - metadata={}, + metadata={"source": "loop", "step": 3}, ) async with AsyncSqliteSaver.from_conn_string(":memory:") as saver: @@ -3213,7 +3213,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -3224,7 +3224,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, - metadata={}, + metadata={"source": "loop", "step": 3}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3237,7 +3237,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, - metadata={}, + metadata={"source": "loop", "step": 1}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3248,7 +3248,7 @@ async def test_branch_then() -> None: values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, - metadata={}, + metadata={"source": "loop", "step": 3}, )