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: