From 64c30508d92d12220b0dcaa13c892455411e1680 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 5 Aug 2024 16:19:56 -0400 Subject: [PATCH 1/4] checkpoint-sqlite: stop using checkpointers as context managers, use contextmanager only in from_conn_string --- .../langgraph/checkpoint/sqlite/__init__.py | 41 +++++++------------ .../langgraph/checkpoint/sqlite/aio.py | 29 ++++--------- 2 files changed, 24 insertions(+), 46 deletions(-) diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py index 78b601cb8..54c14d5ea 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py @@ -1,12 +1,10 @@ import sqlite3 import threading -from contextlib import AbstractContextManager, contextmanager +from contextlib import contextmanager from hashlib import md5 -from types import TracebackType from typing import Any, AsyncIterator, Dict, Iterator, Optional, Sequence, Tuple from langchain_core.runnables import RunnableConfig -from typing_extensions import Self from langgraph.checkpoint.base import ( BaseCheckpointSaver, @@ -32,7 +30,7 @@ _AIO_ERROR_MSG = ( ) -class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): +class SqliteSaver(BaseCheckpointSaver): """A checkpoint saver that stores checkpoints in a SQLite database. Note: @@ -82,43 +80,34 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): self.lock = threading.Lock() @classmethod - def from_conn_string(cls, conn_string: str) -> "SqliteSaver": + @contextmanager + def from_conn_string(cls, conn_string: str) -> Iterator["SqliteSaver"]: """Create a new SqliteSaver instance from a connection string. Args: conn_string (str): The SQLite connection string. - Returns: + Yields: SqliteSaver: A new SqliteSaver instance. Examples: In memory: - memory = SqliteSaver.from_conn_string(":memory:") + with SqliteSaver.from_conn_string(":memory:") as memory: + ... To disk: - memory = SqliteSaver.from_conn_string("checkpoints.sqlite") + with SqliteSaver.from_conn_string("checkpoints.sqlite") as memory: + ... """ - return SqliteSaver( - conn=sqlite3.connect( - conn_string, - # https://ricardoanderegg.com/posts/python-sqlite-thread-safety/ - check_same_thread=False, - ) - ) - - def __enter__(self) -> Self: - return self - - def __exit__( - self, - __exc_type: Optional[type[BaseException]], - __exc_value: Optional[BaseException], - __traceback: Optional[TracebackType], - ) -> Optional[bool]: - return self.conn.close() + with sqlite3.connect( + conn_string, + # https://ricardoanderegg.com/posts/python-sqlite-thread-safety/ + check_same_thread=False, + ) as conn: + yield SqliteSaver(conn) def setup(self) -> None: """Set up the checkpoint database. diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py index 69e07ad47..5f6e88854 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py @@ -1,7 +1,6 @@ import asyncio import functools -from contextlib import AbstractAsyncContextManager -from types import TracebackType +from contextlib import asynccontextmanager from typing import ( Any, AsyncIterator, @@ -15,7 +14,6 @@ from typing import ( import aiosqlite from langchain_core.runnables import RunnableConfig -from typing_extensions import Self from langgraph.checkpoint.base import ( BaseCheckpointSaver, @@ -45,7 +43,7 @@ def not_implemented_sync_method(func: T) -> T: return wrapper -class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): +class AsyncSqliteSaver(BaseCheckpointSaver): """An asynchronous checkpoint saver that stores checkpoints in a SQLite database. This class provides an asynchronous interface for saving and retrieving checkpoints @@ -135,29 +133,20 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): self.is_setup = False @classmethod - def from_conn_string(cls, conn_string: str) -> "AsyncSqliteSaver": + @asynccontextmanager + async def from_conn_string( + cls, conn_string: str + ) -> AsyncIterator["AsyncSqliteSaver"]: """Create a new AsyncSqliteSaver instance from a connection string. Args: conn_string (str): The SQLite connection string. - Returns: + Yields: AsyncSqliteSaver: A new AsyncSqliteSaver instance. """ - - return AsyncSqliteSaver(conn=aiosqlite.connect(conn_string)) - - async def __aenter__(self) -> Self: - return self - - async def __aexit__( - self, - __exc_type: Optional[type[BaseException]], - __exc_value: Optional[BaseException], - __traceback: Optional[TracebackType], - ) -> Optional[bool]: - if self.is_setup: - return await self.conn.close() + async with aiosqlite.connect(conn_string) as conn: + yield AsyncSqliteSaver(conn) @not_implemented_sync_method def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: From 55717c9bdef661c561871c1771b8bed697adcbb4 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 5 Aug 2024 16:52:41 -0400 Subject: [PATCH 2/4] update memory saver + tests --- .../checkpoint/langgraph/checkpoint/memory.py | 28 +++++++- libs/langgraph/tests/test_pregel.py | 42 +++-------- libs/langgraph/tests/test_pregel_async.py | 70 ++++++++----------- 3 files changed, 67 insertions(+), 73 deletions(-) diff --git a/libs/checkpoint/langgraph/checkpoint/memory.py b/libs/checkpoint/langgraph/checkpoint/memory.py index db8da8ad0..c724cd08d 100644 --- a/libs/checkpoint/langgraph/checkpoint/memory.py +++ b/libs/checkpoint/langgraph/checkpoint/memory.py @@ -1,6 +1,8 @@ import asyncio from collections import defaultdict +from contextlib import AbstractAsyncContextManager, AbstractContextManager from functools import partial +from types import TracebackType from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Tuple from langchain_core.runnables import RunnableConfig @@ -15,7 +17,9 @@ from langgraph.checkpoint.base import ( ) -class MemorySaver(BaseCheckpointSaver): +class MemorySaver( + BaseCheckpointSaver, AbstractContextManager, AbstractAsyncContextManager +): """An in-memory checkpoint saver. This checkpoint saver stores checkpoints in memory using a defaultdict. @@ -57,6 +61,28 @@ class MemorySaver(BaseCheckpointSaver): self.storage = defaultdict(lambda: defaultdict(dict)) self.writes = defaultdict(list) + def __enter__(self) -> "MemorySaver": + return self + + def __exit__( + self, + exc_type: Optional[type[BaseException]], + exc_value: Optional[BaseException], + traceback: Optional[TracebackType], + ) -> Optional[bool]: + return + + async def __aenter__(self) -> "MemorySaver": + return self + + async def __aexit__( + self, + __exc_type: Optional[type[BaseException]], + __exc_value: Optional[BaseException], + __traceback: Optional[TracebackType], + ) -> Optional[bool]: + return + def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: """Get a checkpoint tuple from the in-memory storage. diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 7c7d5ec33..ef7c72c42 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -8,7 +8,6 @@ from contextlib import contextmanager from typing import ( Annotated, Any, - Callable, Dict, Generator, List, @@ -558,7 +557,7 @@ def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None: def test_invoke_two_processes_in_out_interrupt( checkpointer: BaseCheckpointSaver, mocker: MockerFixture ) -> None: - try: + with checkpointer as checkpointer: add_one = mocker.Mock(side_effect=lambda x: x + 1) one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox") two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output") @@ -758,9 +757,6 @@ def test_invoke_two_processes_in_out_interrupt( assert [c for c in app.stream(None, fork_config, stream_mode="updates")] == [ {"one": {"inbox": 4}} ] - finally: - if hasattr(checkpointer, "__exit__"): - checkpointer.__exit__(None, None, None) @pytest.mark.parametrize( @@ -777,7 +773,7 @@ def test_invoke_two_processes_in_out_interrupt( def test_fork_always_re_runs_nodes( checkpointer: BaseCheckpointSaver, mocker: MockerFixture ) -> None: - try: + with checkpointer as checkpointer: add_one = mocker.Mock(side_effect=lambda _: 1) builder = StateGraph(Annotated[int, operator.add]) @@ -930,9 +926,6 @@ def test_fork_always_re_runs_nodes( {"add_one": 1}, {"add_one": 1}, ] - finally: - if hasattr(checkpointer, "__exit__"): - checkpointer.__exit__(None, None, None) def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None: @@ -1268,7 +1261,7 @@ def test_invoke_checkpoint(mocker: MockerFixture) -> None: ], ) def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None: - try: + with checkpointer as checkpointer: class State(TypedDict): value: Annotated[int, operator.add] @@ -1340,9 +1333,6 @@ def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None: two.rtn = {"value": 3} # both the pending write and the new write were applied, 1 + 2 + 3 = 6 assert graph.invoke(None, thread1) == {"value": 6} - finally: - if getattr(checkpointer, "__exit__", None): - checkpointer.__exit__(None, None, None) def test_cond_edge_after_send() -> None: @@ -7514,12 +7504,12 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: ] -@pytest.mark.repeat(10) +# @pytest.mark.repeat(10) @pytest.mark.parametrize( - "checkpointer_fct", + "checkpointer", [ - lambda: MemorySaverAssertImmutable(put_sleep=0.2), - lambda: SqliteSaver.from_conn_string(":memory:"), + MemorySaverAssertImmutable(put_sleep=0.2), + SqliteSaver.from_conn_string(":memory:"), ], ids=[ "memory", @@ -7527,10 +7517,9 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: ], ) def test_nested_graph_interrupts( - checkpointer_fct: Callable[[], BaseCheckpointSaver], + checkpointer: BaseCheckpointSaver, ) -> None: - try: - checkpointer = checkpointer_fct() + with checkpointer as checkpointer: class InnerState(TypedDict): my_key: str @@ -8708,9 +8697,6 @@ def test_nested_graph_interrupts( parent_config=None, ), ] - finally: - if hasattr(checkpointer, "__exit__"): - checkpointer.__exit__(None, None, None) @pytest.mark.parametrize( @@ -8725,7 +8711,7 @@ def test_nested_graph_interrupts( ], ) def test_nested_graph_interrupts_parallel(checkpointer: BaseCheckpointSaver) -> None: - try: + with checkpointer as checkpointer: class InnerState(TypedDict): my_key: Annotated[str, operator.add] @@ -8840,9 +8826,6 @@ def test_nested_graph_interrupts_parallel(checkpointer: BaseCheckpointSaver) -> "my_key": "got here and there and parallel and back again", }, ] - finally: - if hasattr(checkpointer, "__exit__"): - checkpointer.__exit__(None, None, None) @pytest.mark.skip @@ -8858,7 +8841,7 @@ def test_nested_graph_interrupts_parallel(checkpointer: BaseCheckpointSaver) -> ], ) def test_doubly_nested_graph_interrupts(checkpointer: BaseCheckpointSaver) -> None: - try: + with checkpointer as checkpointer: class State(TypedDict): my_key: str @@ -8944,9 +8927,6 @@ def test_doubly_nested_graph_interrupts(checkpointer: BaseCheckpointSaver) -> No "my_key": "hi my value here and there and back again", }, ] - finally: - if hasattr(checkpointer, "__exit__"): - checkpointer.__exit__(None, None, None) def test_repeat_condition(snapshot: SnapshotAssertion) -> None: diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index c646e21a2..f83d51aa2 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -2,13 +2,13 @@ import asyncio import json import operator from collections import Counter -from contextlib import asynccontextmanager, contextmanager +from contextlib import AbstractAsyncContextManager, asynccontextmanager, contextmanager +from types import TracebackType from typing import ( Annotated, Any, AsyncGenerator, AsyncIterator, - Callable, Dict, Generator, List, @@ -66,6 +66,19 @@ from tests.memory_assert import ( from tests.messages import _AnyIdAIMessage, _AnyIdHumanMessage +class NoneContextManager(AbstractAsyncContextManager): + async def __aenter__(self) -> None: + return None + + async def __aexit__( + self, + __exc_type: Optional[type[BaseException]], + __exc_value: Optional[BaseException], + __traceback: Optional[TracebackType], + ) -> Optional[bool]: + return + + async def test_checkpoint_errors() -> None: class FaultyGetCheckpointer(MemorySaver): async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: @@ -236,7 +249,7 @@ async def test_step_timeout_on_stream_hang() -> None: [ MemorySaverAssertImmutable(), AsyncSqliteSaver.from_conn_string(":memory:"), - None, + NoneContextManager(), ], ids=[ "memory", @@ -247,7 +260,7 @@ async def test_step_timeout_on_stream_hang() -> None: async def test_cancel_graph_astream( checkpointer: Optional[BaseCheckpointSaver], ) -> None: - try: + async with checkpointer as checkpointer: class State(TypedDict): value: Annotated[int, operator.add] @@ -311,9 +324,6 @@ async def test_cancel_graph_astream( "alittlewhile", ) assert state.metadata == {"source": "loop", "step": 0, "writes": None} - finally: - if getattr(checkpointer, "__aexit__", None): - await checkpointer.__aexit__(None, None, None) @pytest.mark.parametrize( @@ -321,7 +331,7 @@ async def test_cancel_graph_astream( [ MemorySaverAssertImmutable(), AsyncSqliteSaver.from_conn_string(":memory:"), - None, + NoneContextManager(), ], ids=[ "memory", @@ -332,7 +342,7 @@ async def test_cancel_graph_astream( async def test_cancel_graph_astream_events_v2( checkpointer: Optional[BaseCheckpointSaver], ) -> None: - try: + async with checkpointer as checkpointer: class State(TypedDict): value: int @@ -402,9 +412,6 @@ async def test_cancel_graph_astream_events_v2( "step": 1, "writes": {"alittlewhile": {"value": 2}}, } - finally: - if getattr(checkpointer, "__aexit__", None): - await checkpointer.__aexit__(None, None, None) async def test_node_schemas_custom_output() -> None: @@ -679,7 +686,7 @@ async def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None: async def test_invoke_two_processes_in_out_interrupt( checkpointer: BaseCheckpointSaver, mocker: MockerFixture ) -> None: - try: + async with checkpointer as checkpointer: add_one = mocker.Mock(side_effect=lambda x: x + 1) one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox") two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output") @@ -880,9 +887,6 @@ async def test_invoke_two_processes_in_out_interrupt( assert [ c async for c in app.astream(None, fork_config, stream_mode="updates") ] == [{"one": {"inbox": 4}}] - finally: - if hasattr(checkpointer, "__aexit__"): - await checkpointer.__aexit__(None, None, None) @pytest.mark.parametrize( @@ -899,7 +903,7 @@ async def test_invoke_two_processes_in_out_interrupt( async def test_fork_always_re_runs_nodes( checkpointer: BaseCheckpointSaver, mocker: MockerFixture ) -> None: - try: + async with checkpointer as checkpointer: add_one = mocker.Mock(side_effect=lambda _: 1) builder = StateGraph(Annotated[int, operator.add]) @@ -1061,9 +1065,6 @@ async def test_fork_always_re_runs_nodes( {"add_one": 1}, {"add_one": 1}, ] - finally: - if hasattr(checkpointer, "__aexit__"): - await checkpointer.__aexit__(None, None, None) async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None: @@ -1406,7 +1407,7 @@ async def test_invoke_checkpoint(mocker: MockerFixture) -> None: ], ) async def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None: - try: + async with checkpointer as checkpointer: class State(TypedDict): value: Annotated[int, operator.add] @@ -1478,9 +1479,6 @@ async def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None: two.rtn = {"value": 3} # both the pending write and the new write were applied, 1 + 2 + 3 = 6 assert await graph.ainvoke(None, thread1) == {"value": 6} - finally: - if getattr(checkpointer, "__aexit__", None): - await checkpointer.__aexit__(None, None, None) async def test_cond_edge_after_send() -> None: @@ -5991,10 +5989,10 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None: @pytest.mark.repeat(10) @pytest.mark.parametrize( - "checkpointer_fct", + "checkpointer", [ - lambda: MemorySaverAssertImmutable(put_sleep=0.2), - lambda: AsyncSqliteSaver.from_conn_string(":memory:"), + MemorySaverAssertImmutable(put_sleep=0.2), + AsyncSqliteSaver.from_conn_string(":memory:"), ], ids=[ "memory", @@ -6002,10 +6000,9 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None: ], ) async def test_nested_graph_interrupts( - checkpointer_fct: Callable[[], BaseCheckpointSaver], + checkpointer: BaseCheckpointSaver, ) -> None: - try: - checkpointer = checkpointer_fct() + async with checkpointer as checkpointer: class InnerState(TypedDict): my_key: str @@ -7196,9 +7193,6 @@ async def test_nested_graph_interrupts( parent_config=None, ), ] - finally: - if hasattr(checkpointer, "__aexit__"): - await checkpointer.__aexit__(None, None, None) @pytest.mark.parametrize( @@ -7215,7 +7209,7 @@ async def test_nested_graph_interrupts( async def test_nested_graph_interrupts_parallel( checkpointer: BaseCheckpointSaver, ) -> None: - try: + async with checkpointer as checkpointer: class InnerState(TypedDict): my_key: Annotated[str, operator.add] @@ -7332,9 +7326,6 @@ async def test_nested_graph_interrupts_parallel( "my_key": "got here and there and parallel and back again", }, ] - finally: - if hasattr(checkpointer, "__aexit__"): - await checkpointer.__aexit__(None, None, None) @pytest.mark.skip @@ -7352,7 +7343,7 @@ async def test_nested_graph_interrupts_parallel( async def test_doubly_nested_graph_interrupts( checkpointer: BaseCheckpointSaver, ) -> None: - try: + async with checkpointer as checkpointer: class State(TypedDict): my_key: str @@ -7443,9 +7434,6 @@ async def test_doubly_nested_graph_interrupts( "my_key": "hi my value here and there and back again", }, ] - finally: - if hasattr(checkpointer, "__aexit__"): - await checkpointer.__aexit__(None, None, None) async def test_checkpoint_metadata() -> None: From 9aaa73cc5579f019c3460d0fa4cb624da16328ec Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 5 Aug 2024 16:55:19 -0400 Subject: [PATCH 3/4] undo comment --- libs/langgraph/tests/test_pregel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index ef7c72c42..a8b4b9849 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7504,7 +7504,7 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: ] -# @pytest.mark.repeat(10) +@pytest.mark.repeat(10) @pytest.mark.parametrize( "checkpointer", [ From 566e6c9b60c1da8ab54641315f77f2bd45fe114b Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 5 Aug 2024 17:03:50 -0400 Subject: [PATCH 4/4] fix --- libs/langgraph/tests/test_pregel.py | 11 ++++++----- libs/langgraph/tests/test_pregel_async.py | 11 ++++++----- 2 files changed, 12 insertions(+), 10 deletions(-) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index a8b4b9849..7001f8844 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -8,6 +8,7 @@ from contextlib import contextmanager from typing import ( Annotated, Any, + Callable, Dict, Generator, List, @@ -7506,10 +7507,10 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: @pytest.mark.repeat(10) @pytest.mark.parametrize( - "checkpointer", + "checkpointer_fct", [ - MemorySaverAssertImmutable(put_sleep=0.2), - SqliteSaver.from_conn_string(":memory:"), + lambda: MemorySaverAssertImmutable(put_sleep=0.2), + lambda: SqliteSaver.from_conn_string(":memory:"), ], ids=[ "memory", @@ -7517,9 +7518,9 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: ], ) def test_nested_graph_interrupts( - checkpointer: BaseCheckpointSaver, + checkpointer_fct: Callable[[], BaseCheckpointSaver], ) -> None: - with checkpointer as checkpointer: + with checkpointer_fct() as checkpointer: class InnerState(TypedDict): my_key: str diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index f83d51aa2..1b8c3df91 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -9,6 +9,7 @@ from typing import ( Any, AsyncGenerator, AsyncIterator, + Callable, Dict, Generator, List, @@ -5989,10 +5990,10 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None: @pytest.mark.repeat(10) @pytest.mark.parametrize( - "checkpointer", + "checkpointer_fct", [ - MemorySaverAssertImmutable(put_sleep=0.2), - AsyncSqliteSaver.from_conn_string(":memory:"), + lambda: MemorySaverAssertImmutable(put_sleep=0.2), + lambda: AsyncSqliteSaver.from_conn_string(":memory:"), ], ids=[ "memory", @@ -6000,9 +6001,9 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None: ], ) async def test_nested_graph_interrupts( - checkpointer: BaseCheckpointSaver, + checkpointer_fct: Callable[[], BaseCheckpointSaver], ) -> None: - async with checkpointer as checkpointer: + async with checkpointer_fct() as checkpointer: class InnerState(TypedDict): my_key: str