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