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]: 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 9c6847d09..82bd0c414 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -558,7 +558,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 +758,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 +774,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 +927,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 +1262,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 +1334,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: @@ -7550,8 +7541,7 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: def test_nested_graph_interrupts( checkpointer_fct: Callable[[], BaseCheckpointSaver], ) -> None: - try: - checkpointer = checkpointer_fct() + with checkpointer_fct() as checkpointer: class InnerState(TypedDict): my_key: str @@ -8729,9 +8719,6 @@ def test_nested_graph_interrupts( parent_config=None, ), ] - finally: - if hasattr(checkpointer, "__exit__"): - checkpointer.__exit__(None, None, None) @pytest.mark.parametrize( @@ -8746,7 +8733,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] @@ -8861,9 +8848,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 @@ -8879,7 +8863,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 @@ -8965,9 +8949,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 9c4e16015..23065b3d2 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -2,7 +2,8 @@ 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, @@ -66,6 +67,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 +250,7 @@ async def test_step_timeout_on_stream_hang() -> None: [ MemorySaverAssertImmutable(), AsyncSqliteSaver.from_conn_string(":memory:"), - None, + NoneContextManager(), ], ids=[ "memory", @@ -247,7 +261,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 +325,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 +332,7 @@ async def test_cancel_graph_astream( [ MemorySaverAssertImmutable(), AsyncSqliteSaver.from_conn_string(":memory:"), - None, + NoneContextManager(), ], ids=[ "memory", @@ -332,7 +343,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 +413,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 +687,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 +888,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 +904,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 +1066,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 +1408,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 +1480,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: @@ -6143,8 +6142,7 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None: async def test_nested_graph_interrupts( checkpointer_fct: Callable[[], BaseCheckpointSaver], ) -> None: - try: - checkpointer = checkpointer_fct() + async with checkpointer_fct() as checkpointer: class InnerState(TypedDict): my_key: str @@ -7335,9 +7333,6 @@ async def test_nested_graph_interrupts( parent_config=None, ), ] - finally: - if hasattr(checkpointer, "__aexit__"): - await checkpointer.__aexit__(None, None, None) @pytest.mark.parametrize( @@ -7354,7 +7349,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] @@ -7471,9 +7466,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 @@ -7491,7 +7483,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 @@ -7582,9 +7574,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: