From f3049081025f02c6d661584d5f35f50b4c2ba043 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 3 May 2024 15:40:30 -0700 Subject: [PATCH] Add metadata to checkpoints - not yet used in this PR --- langgraph/checkpoint/aiosqlite.py | 69 ++++++++++++++--------------- langgraph/checkpoint/base.py | 13 +++++- langgraph/checkpoint/memory.py | 32 ++++++++++---- langgraph/checkpoint/sqlite.py | 73 +++++++++++++++++-------------- langgraph/pregel/__init__.py | 15 +++++-- langgraph/pregel/types.py | 2 + tests/test_pregel.py | 39 +++++++++++++++++ tests/test_pregel_async.py | 29 ++++++++++++ 8 files changed, 188 insertions(+), 84 deletions(-) diff --git a/langgraph/checkpoint/aiosqlite.py b/langgraph/checkpoint/aiosqlite.py index 62818dc16..e0758d84d 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 AsyncIterator, Optional +from typing import Any, AsyncIterator, Optional import aiosqlite from langchain_core.runnables import RunnableConfig @@ -130,6 +130,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): thread_ts TEXT NOT NULL, parent_ts TEXT, checkpoint BLOB, + metadata BLOB, PRIMARY KEY (thread_id, thread_ts) ); """ @@ -155,7 +156,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): await self.setup() if config["configurable"].get("thread_ts"): async with self.conn.execute( - "SELECT checkpoint, parent_ts FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", + "SELECT checkpoint, parent_ts, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", ( str(config["configurable"]["thread_id"]), str(config["configurable"]["thread_ts"]), @@ -165,20 +166,19 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): return CheckpointTuple( config, self.serde.loads(value[0]), - ( - { - "configurable": { - "thread_id": config["configurable"]["thread_id"], - "thread_ts": value[1], - } + self.serde.loads(value[2]) if value[2] is not None else None, + { + "configurable": { + "thread_id": config["configurable"]["thread_id"], + "thread_ts": value[1], } - if value[1] - else None - ), + } + if value[1] + else None, ) else: async with self.conn.execute( - "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", + "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", (str(config["configurable"]["thread_id"]),), ) as cursor: if value := await cursor.fetchone(): @@ -190,16 +190,15 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): } }, self.serde.loads(value[3]), - ( - { - "configurable": { - "thread_id": value[0], - "thread_ts": value[2], - } + self.serde.loads(value[4]) if value[4] is not None else None, + { + "configurable": { + "thread_id": value[0], + "thread_ts": value[2], } - if value[2] - else None - ), + } + if value[2] + else None, ) async def alist( @@ -224,9 +223,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): """ await self.setup() query = ( - "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC" + "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC" if before is None - else "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC" + else "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC" ) if limit: query += f" LIMIT {limit}" @@ -241,24 +240,21 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): ) ), ) as cursor: - async for thread_id, thread_ts, parent_ts, value in cursor: + async for thread_id, thread_ts, parent_ts, value, metadata in cursor: yield CheckpointTuple( {"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}}, self.serde.loads(value), - ( - { - "configurable": { - "thread_id": thread_id, - "thread_ts": parent_ts, - } - } - if parent_ts - else None - ), + self.serde.loads(metadata) if metadata is not None else None, + {"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}} + if parent_ts + else None, ) async def aput( - self, config: RunnableConfig, checkpoint: Checkpoint + self, + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: Optional[dict[str, Any]] = None, ) -> RunnableConfig: """Save a checkpoint to the database asynchronously. @@ -274,12 +270,13 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): """ await self.setup() async with self.conn.execute( - "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint) VALUES (?, ?, ?, ?)", + "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint, metadata) VALUES (?, ?, ?, ?, ?)", ( str(config["configurable"]["thread_id"]), checkpoint["ts"], config["configurable"].get("thread_ts"), self.serde.dumps(checkpoint), + self.serde.dumps(metadata) if metadata is not None else None, ), ): await self.conn.commit() diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index 86d71c17f..6bfc79885 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -83,6 +83,7 @@ class CheckpointAt(StrEnum): class CheckpointTuple(NamedTuple): config: RunnableConfig checkpoint: Checkpoint + metadata: Optional[dict[str, Any]] parent_config: Optional[RunnableConfig] = None @@ -139,7 +140,12 @@ class BaseCheckpointSaver(ABC): ) -> Iterator[CheckpointTuple]: raise NotImplementedError - def put(self, config: RunnableConfig, checkpoint: Checkpoint) -> RunnableConfig: + def put( + self, + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: dict[str, Any], + ) -> RunnableConfig: raise NotImplementedError async def aget(self, config: RunnableConfig) -> Optional[Checkpoint]: @@ -159,6 +165,9 @@ class BaseCheckpointSaver(ABC): raise NotImplementedError async def aput( - self, config: RunnableConfig, checkpoint: Checkpoint + self, + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: dict[str, Any], ) -> RunnableConfig: raise NotImplementedError diff --git a/langgraph/checkpoint/memory.py b/langgraph/checkpoint/memory.py index 055f92d54..dab21794f 100644 --- a/langgraph/checkpoint/memory.py +++ b/langgraph/checkpoint/memory.py @@ -1,6 +1,6 @@ import asyncio from collections import defaultdict -from typing import AsyncIterator, Iterator, Optional +from typing import Any, AsyncIterator, Iterator, Optional from langchain_core.runnables import RunnableConfig @@ -39,7 +39,7 @@ class MemorySaver(BaseCheckpointSaver): asyncio.run(coro) # Output: 2 """ - storage: defaultdict[str, dict[str, Checkpoint]] + storage: defaultdict[str, dict[str, tuple[bytes, bytes]]] def __init__( self, @@ -66,16 +66,21 @@ class MemorySaver(BaseCheckpointSaver): """ thread_id = config["configurable"]["thread_id"] if ts := config["configurable"].get("thread_ts"): - if checkpoint := self.storage[thread_id].get(ts): + if saved := self.storage[thread_id].get(ts): + checkpoint, metadata = saved return CheckpointTuple( - config=config, checkpoint=self.serde.loads(checkpoint) + config=config, + checkpoint=self.serde.loads(checkpoint), + metadata=self.serde.loads(metadata), ) else: if checkpoints := self.storage[thread_id]: ts = max(checkpoints.keys()) + checkpoint, metadata = checkpoints[ts] return CheckpointTuple( config={"configurable": {"thread_id": thread_id, "thread_ts": ts}}, - checkpoint=self.serde.loads(checkpoints[ts]), + checkpoint=self.serde.loads(checkpoint), + metadata=self.serde.loads(metadata), ) def list( @@ -99,7 +104,7 @@ class MemorySaver(BaseCheckpointSaver): Iterator[CheckpointTuple]: An iterator of checkpoint tuples. """ thread_id = config["configurable"]["thread_id"] - for ts, checkpoint in self.storage[thread_id].items(): + for ts, (checkpoint, metadata) in self.storage[thread_id].items(): if before and ts >= before["configurable"]["thread_ts"]: continue if limit is not None and limit <= 0: @@ -108,9 +113,15 @@ class MemorySaver(BaseCheckpointSaver): yield CheckpointTuple( config={"configurable": {"thread_id": thread_id, "thread_ts": ts}}, checkpoint=self.serde.loads(checkpoint), + metadata=self.serde.loads(metadata), ) - def put(self, config: RunnableConfig, checkpoint: Checkpoint) -> RunnableConfig: + def put( + self, + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: dict[str, Any] = None, + ) -> RunnableConfig: """Save a checkpoint to the in-memory storage. This method saves a checkpoint to the in-memory storage. The checkpoint is associated @@ -124,7 +135,12 @@ class MemorySaver(BaseCheckpointSaver): RunnableConfig: The updated config containing the saved checkpoint's timestamp. """ self.storage[config["configurable"]["thread_id"]].update( - {checkpoint["ts"]: self.serde.dumps(checkpoint)} + { + checkpoint["ts"]: ( + self.serde.dumps(checkpoint), + self.serde.dumps(metadata or {}), + ) + } ) return { "configurable": { diff --git a/langgraph/checkpoint/sqlite.py b/langgraph/checkpoint/sqlite.py index e4c3be2a3..f45c28be4 100644 --- a/langgraph/checkpoint/sqlite.py +++ b/langgraph/checkpoint/sqlite.py @@ -146,6 +146,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): thread_ts TEXT NOT NULL, parent_ts TEXT, checkpoint BLOB, + metadata BLOB, PRIMARY KEY (thread_id, thread_ts) ); """ @@ -211,7 +212,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): with self.cursor(transaction=False) as cur: if config["configurable"].get("thread_ts"): cur.execute( - "SELECT checkpoint, parent_ts FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", + "SELECT checkpoint, parent_ts, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", ( str(config["configurable"]["thread_id"]), str(config["configurable"]["thread_ts"]), @@ -221,20 +222,19 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): return CheckpointTuple( config, self.serde.loads(value[0]), - ( - { - "configurable": { - "thread_id": config["configurable"]["thread_id"], - "thread_ts": value[1], - } + self.serde.loads(value[2]) if value[2] is not None else None, + { + "configurable": { + "thread_id": config["configurable"]["thread_id"], + "thread_ts": value[1], } - if value[1] - else None - ), + } + if value[1] + else None, ) else: cur.execute( - "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", + "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", (str(config["configurable"]["thread_id"]),), ) if value := cur.fetchone(): @@ -246,16 +246,15 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): } }, self.serde.loads(value[3]), - ( - { - "configurable": { - "thread_id": value[0], - "thread_ts": value[2], - } + self.serde.loads(value[4]) if value[4] is not None else None, + { + "configurable": { + "thread_id": value[0], + "thread_ts": value[2], } - if value[2] - else None - ), + } + if value[2] + else None, ) def list( @@ -289,9 +288,9 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): print(checkpoints) # Output: [CheckpointTuple(...), ...] """ query = ( - "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC" + "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC" if before is None - else "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC" + else "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC" ) if limit: query += f" LIMIT {limit}" @@ -307,23 +306,27 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): ) ), ) - for thread_id, thread_ts, parent_ts, value in cur: + for thread_id, thread_ts, parent_ts, value, metadata in cur: yield CheckpointTuple( {"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}}, self.serde.loads(value), - ( - { - "configurable": { - "thread_id": thread_id, - "thread_ts": parent_ts, - } + self.serde.loads(metadata) if metadata is not None else None, + { + "configurable": { + "thread_id": thread_id, + "thread_ts": parent_ts, } - if parent_ts - else None - ), + } + if parent_ts + else None, ) - def put(self, config: RunnableConfig, checkpoint: Checkpoint) -> RunnableConfig: + def put( + self, + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: dict[str, Any] = None, + ) -> RunnableConfig: """Save a checkpoint to the database. This method saves a checkpoint to the SQLite database. The checkpoint is associated @@ -332,6 +335,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): Args: config (RunnableConfig): The config to associate with the checkpoint. checkpoint (Checkpoint): The checkpoint to save. + metadata (Optional[dict[str, Any]]): Additional metadata to save with the checkpoint. Defaults to None. Returns: RunnableConfig: The updated config containing the saved checkpoint's timestamp. @@ -347,12 +351,13 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): """ with self.cursor() as cur: cur.execute( - "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint) VALUES (?, ?, ?, ?)", + "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint, metadata) VALUES (?, ?, ?, ?, ?)", ( str(config["configurable"]["thread_id"]), checkpoint["ts"], config["configurable"].get("thread_ts"), self.serde.dumps(checkpoint), + self.serde.dumps(metadata) if metadata is not None else None, ), ) return { diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index d83eb906b..c29f55b86 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -343,6 +343,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, + saved.metadata or {}, ) async def aget_state(self, config: RunnableConfig) -> StateSnapshot: @@ -361,6 +362,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, + saved.metadata or {}, ) def get_state_history( @@ -374,7 +376,7 @@ class Pregel( if not self.checkpointer: raise ValueError("No checkpointer set") - for config, checkpoint, parent_config in self.checkpointer.list( + for config, checkpoint, metadata, parent_config in self.checkpointer.list( config, before=before, limit=limit ): with ChannelsManager(self.channels, checkpoint) as channels: @@ -385,6 +387,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, + metadata or {}, parent_config, ) @@ -399,9 +402,12 @@ class Pregel( if not self.checkpointer: raise ValueError("No checkpointer set") - async for config, checkpoint, parent_config in self.checkpointer.alist( - config, before=before, limit=limit - ): + async for ( + config, + checkpoint, + metadata, + parent_config, + ) in self.checkpointer.alist(config, before=before, limit=limit): async with AsyncChannelsManager(self.channels, checkpoint) as channels: _, next_tasks = _prepare_next_tasks( checkpoint, self.nodes, channels, for_execution=False @@ -410,6 +416,7 @@ class Pregel( read_channels(channels, self.stream_channels_asis), tuple(name for name, _ in next_tasks), config, + metadata or {}, parent_config, ) diff --git a/langgraph/pregel/types.py b/langgraph/pregel/types.py index 7c6075c98..4359b5994 100644 --- a/langgraph/pregel/types.py +++ b/langgraph/pregel/types.py @@ -25,6 +25,8 @@ class StateSnapshot(NamedTuple): """Nodes to execute in the next step, if any""" config: RunnableConfig """Config used to fetch this snapshot""" + metadata: dict[str, Any] + """Metadata associated with this snapshot""" parent_config: Optional[RunnableConfig] = None """Config used to fetch the parent snapshot, if any""" diff --git a/tests/test_pregel.py b/tests/test_pregel.py index c8f20eda6..2ba1d1d77 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -1228,6 +1228,7 @@ def test_conditional_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert ( app_w_interrupt.checkpointer.get_tuple(config).config["configurable"][ @@ -1261,6 +1262,7 @@ def test_conditional_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1344,6 +1346,7 @@ def test_conditional_graph( }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) # test state get/update methods with interrupt_before @@ -1379,6 +1382,7 @@ def test_conditional_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) app_w_interrupt.update_state( @@ -1406,6 +1410,7 @@ def test_conditional_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1489,6 +1494,7 @@ def test_conditional_graph( }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) # test re-invoke to continue with interrupt_before @@ -1524,6 +1530,7 @@ def test_conditional_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1860,6 +1867,7 @@ def test_conditional_state_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) app_w_interrupt.update_state( @@ -1885,6 +1893,7 @@ def test_conditional_state_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -1943,6 +1952,7 @@ def test_conditional_state_graph( }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) # test state get/update methods with interrupt_before @@ -1977,6 +1987,7 @@ def test_conditional_state_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) app_w_interrupt.update_state( @@ -2002,6 +2013,7 @@ def test_conditional_state_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2060,6 +2072,7 @@ def test_conditional_state_graph( }, next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) # test w interrupt before all @@ -2082,6 +2095,7 @@ def test_conditional_state_graph( }, next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2104,6 +2118,7 @@ def test_conditional_state_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2142,6 +2157,7 @@ def test_conditional_state_graph( }, next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2186,6 +2202,7 @@ def test_conditional_state_graph( }, next=("tools",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -2224,6 +2241,7 @@ def test_conditional_state_graph( }, next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -3042,6 +3060,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) # modify ai message @@ -3067,6 +3086,7 @@ def test_message_graph( ], next=("action",), config=next_config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -3132,6 +3152,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) app_w_interrupt.update_state( @@ -3167,6 +3188,7 @@ def test_message_graph( ], next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) app_w_interrupt = workflow.compile( @@ -3212,6 +3234,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) # modify ai message @@ -3240,6 +3263,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c for c in app_w_interrupt.stream(None, config)] == [ @@ -3305,6 +3329,7 @@ def test_message_graph( ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) app_w_interrupt.update_state( @@ -3340,6 +3365,7 @@ def test_message_graph( ], next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) # add an extra message as if it came from "action" node @@ -3375,6 +3401,7 @@ def test_message_graph( ], next=("agent",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) @@ -3503,6 +3530,7 @@ def test_start_branch_then( values={"my_key": "value", "market": "DE"}, next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, + metadata={}, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3513,6 +3541,7 @@ def test_start_branch_then( values={"my_key": "value slow", "market": "DE"}, next=(), config=tool_two.checkpointer.get_tuple(thread1).config, + metadata={}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3525,6 +3554,7 @@ def test_start_branch_then( values={"my_key": "value", "market": "US"}, next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, + metadata={}, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3535,6 +3565,7 @@ def test_start_branch_then( values={"my_key": "value fast", "market": "US"}, next=(), config=tool_two.checkpointer.get_tuple(thread2).config, + metadata={}, ) @@ -3855,6 +3886,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, + metadata={}, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3865,6 +3897,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=tool_two.checkpointer.get_tuple(thread1).config, + metadata={}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3877,6 +3910,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, + metadata={}, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3887,6 +3921,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=tool_two.checkpointer.get_tuple(thread2).config, + metadata={}, ) with SqliteSaver.from_conn_string(":memory:") as saver: @@ -3909,6 +3944,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=tool_two.checkpointer.get_tuple(thread1).config, + metadata={}, ) # resume, for same result as above assert tool_two.invoke(None, thread1, debug=1) == { @@ -3919,6 +3955,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=tool_two.checkpointer.get_tuple(thread1).config, + metadata={}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3931,6 +3968,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=tool_two.checkpointer.get_tuple(thread2).config, + metadata={}, ) # resume, for same result as above assert tool_two.invoke(None, thread2, debug=1) == { @@ -3941,6 +3979,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) - values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=tool_two.checkpointer.get_tuple(thread2).config, + metadata={}, ) diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index e453460c9..79e4b3e33 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -1306,6 +1306,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) await app_w_interrupt.aupdate_state( @@ -1333,6 +1334,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -1416,6 +1418,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) # test state get/update methods with interrupt_before @@ -1454,6 +1457,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) await app_w_interrupt.aupdate_state( @@ -1481,6 +1485,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -1564,6 +1569,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) # test re-invoke to continue with interrupt_before @@ -1602,6 +1608,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -1924,6 +1931,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) await app_w_interrupt.aupdate_state( @@ -1949,6 +1957,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -2007,6 +2016,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) # test state get/update methods with interrupt_before @@ -2043,6 +2053,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) await app_w_interrupt.aupdate_state( @@ -2068,6 +2079,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None: }, next=("tools",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -2126,6 +2138,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None: }, next=(), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) @@ -2734,6 +2747,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None: ], next=("action",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, + metadata={}, ) # modify ai message @@ -2761,6 +2775,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None: ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) assert [c async for c in app_w_interrupt.astream(None, config)] == [ @@ -2813,6 +2828,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None: ], next=("action",), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) await app_w_interrupt.aupdate_state( @@ -2846,6 +2862,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None: ], next=(), config=app_w_interrupt.checkpointer.get_tuple(config).config, + metadata={}, ) @@ -2968,6 +2985,7 @@ async def test_start_branch_then( values={"my_key": "value", "market": "DE"}, next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, + metadata={}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -2978,6 +2996,7 @@ async def test_start_branch_then( values={"my_key": "value slow", "market": "DE"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, + metadata={}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -2990,6 +3009,7 @@ async def test_start_branch_then( values={"my_key": "value", "market": "US"}, next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, + metadata={}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3000,6 +3020,7 @@ async def test_start_branch_then( values={"my_key": "value fast", "market": "US"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, + metadata={}, ) @@ -3308,6 +3329,7 @@ async def test_branch_then( values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, + metadata={}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -3318,6 +3340,7 @@ async def test_branch_then( values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, + metadata={}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3330,6 +3353,7 @@ async def test_branch_then( values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, + metadata={}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3340,6 +3364,7 @@ async def test_branch_then( values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, + metadata={}, ) async with AsyncSqliteSaver.from_conn_string(":memory:") as saver: @@ -3362,6 +3387,7 @@ async def test_branch_then( values={"my_key": "value prepared", "market": "DE"}, next=("tool_two_slow",), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, + metadata={}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread1, debug=1) == { @@ -3372,6 +3398,7 @@ async def test_branch_then( values={"my_key": "value prepared slow finished", "market": "DE"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread1)).config, + metadata={}, ) thread2 = {"configurable": {"thread_id": "2"}} @@ -3384,6 +3411,7 @@ async def test_branch_then( values={"my_key": "value prepared", "market": "US"}, next=("tool_two_fast",), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, + metadata={}, ) # resume, for same result as above assert await tool_two.ainvoke(None, thread2, debug=1) == { @@ -3394,6 +3422,7 @@ async def test_branch_then( values={"my_key": "value prepared fast finished", "market": "US"}, next=(), config=(await tool_two.checkpointer.aget_tuple(thread2)).config, + metadata={}, )