From b6cc0228616c71586a1e9b29225a8ed4d4f32c38 Mon Sep 17 00:00:00 2001 From: William FH <13333726+hinthornw@users.noreply.github.com> Date: Tue, 2 Dec 2025 11:47:41 -0800 Subject: [PATCH] =?UTF-8?q?fix(checkpoint):=20InMemorySaver=20context=20ma?= =?UTF-8?q?nagers=20should=20return=20self=20in=E2=80=A6=20(#6529)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit h/t to @lexi-k for openening. merged here to check/fix CI Co-authored-by: lexi-k <69981673+lexi-k@users.noreply.github.com> --- .../langgraph/checkpoint/memory/__init__.py | 6 ++++-- libs/checkpoint/tests/test_memory.py | 11 +++++++++-- 2 files changed, 13 insertions(+), 4 deletions(-) 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