diff --git a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py index 7ee88c123..da353c2b4 100644 --- a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py @@ -97,7 +97,8 @@ class InMemorySaver( self.stack.enter_context(self.blobs) # type: ignore[arg-type] def __enter__(self) -> InMemorySaver: - return self.stack.__enter__() + self.stack.__enter__() + return self def __exit__( self, @@ -108,7 +109,8 @@ class InMemorySaver( return self.stack.__exit__(exc_type, exc_value, traceback) async def __aenter__(self) -> InMemorySaver: - return self.stack.__enter__() + self.stack.__enter__() + return self async def __aexit__( self, diff --git a/libs/checkpoint/tests/test_memory.py b/libs/checkpoint/tests/test_memory.py index e3fc4aa58..7d85f4ed5 100644 --- a/libs/checkpoint/tests/test_memory.py +++ b/libs/checkpoint/tests/test_memory.py @@ -188,7 +188,14 @@ class TestMemorySaver: assert len(search_results_4) == 0 -def test_memory_saver() -> None: +async def test_memory_saver() -> None: from langgraph.checkpoint.memory import InMemorySaver - assert isinstance(InMemorySaver(), InMemorySaver) + memory_saver = InMemorySaver() + assert isinstance(memory_saver, InMemorySaver) + + async with memory_saver as async_memory_saver: + assert async_memory_saver is memory_saver + + with memory_saver as sync_memory_saver: + assert sync_memory_saver is memory_saver