diff --git a/libs/langgraph/tests/test_interruption.py b/libs/langgraph/tests/test_interruption.py index 0c9da279a..e6682e964 100644 --- a/libs/langgraph/tests/test_interruption.py +++ b/libs/langgraph/tests/test_interruption.py @@ -1,10 +1,18 @@ from typing import TypedDict -from langgraph.checkpoint.memory import MemorySaver +import pytest +from pytest_mock import MockerFixture + from langgraph.graph import END, START, StateGraph -def test_interruption_without_state_updates(): +@pytest.mark.parametrize( + "checkpointer_name", + ["memory", "sqlite", "postgres", "postgres_pipe"], +) +def test_interruption_without_state_updates( + request: pytest.FixtureRequest, checkpointer_name: str, mocker: MockerFixture +) -> None: """Test interruption without state updates. This test confirms that interrupting doesn't require a state key having been updated in the prev step""" @@ -23,9 +31,8 @@ def test_interruption_without_state_updates(): builder.add_edge("step_2", "step_3") builder.add_edge("step_3", END) - memory = MemorySaver() - - graph = builder.compile(checkpointer=memory, interrupt_after="*") + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") + graph = builder.compile(checkpointer=checkpointer, interrupt_after="*") initial_input = {"input": "hello world"} thread = {"configurable": {"thread_id": "1"}} @@ -40,7 +47,13 @@ def test_interruption_without_state_updates(): assert graph.get_state(thread).next == () -async def test_interruption_without_state_updates_async(): +@pytest.mark.parametrize( + "checkpointer_name", + ["memory", "sqlite_aio", "postgres_aio", "postgres_aio_pipe"], +) +async def test_interruption_without_state_updates_async( + request: pytest.FixtureRequest, checkpointer_name: str, mocker: MockerFixture +): """Test interruption without state updates. This test confirms that interrupting doesn't require a state key having been updated in the prev step""" @@ -59,9 +72,8 @@ async def test_interruption_without_state_updates_async(): builder.add_edge("step_2", "step_3") builder.add_edge("step_3", END) - memory = MemorySaver() - - graph = builder.compile(checkpointer=memory, interrupt_after="*") + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") + graph = builder.compile(checkpointer=checkpointer, interrupt_after="*") initial_input = {"input": "hello world"} thread = {"configurable": {"thread_id": "1"}} diff --git a/libs/langgraph/tests/test_prebuilt.py b/libs/langgraph/tests/test_prebuilt.py index 97ef429c1..ba9c88afc 100644 --- a/libs/langgraph/tests/test_prebuilt.py +++ b/libs/langgraph/tests/test_prebuilt.py @@ -18,11 +18,8 @@ from langchain_core.tools import BaseTool from langchain_core.tools import tool as dec_tool from pydantic import BaseModel as BaseModelV2 -from langgraph.checkpoint.base import BaseCheckpointSaver from langgraph.prebuilt import ToolNode, ValidationNode, create_react_agent from langgraph.prebuilt.tool_node import InjectedState -from tests.any_str import AnyStr -from tests.memory_assert import MemorySaverAssertImmutable from tests.messages import _AnyIdHumanMessage @@ -54,18 +51,13 @@ class FakeToolCallingModel(BaseChatModel): @pytest.mark.parametrize( - "checkpointer", - [ - MemorySaverAssertImmutable(), - None, - ], - ids=[ - "memory", - "none", - ], + "checkpointer_name", + ["memory", "sqlite", "postgres", "postgres_pipe"], ) -def test_no_modifier(checkpointer: Optional[BaseCheckpointSaver]): +def test_no_modifier(request: pytest.FixtureRequest, checkpointer_name: str) -> None: + checkpointer = request.getfixturevalue("checkpointer_" + checkpointer_name) model = FakeToolCallingModel() + agent = create_react_agent(model, [], checkpointer=checkpointer) inputs = [HumanMessage("hi?")] thread = {"configurable": {"thread_id": "123"}} @@ -76,30 +68,12 @@ def test_no_modifier(checkpointer: Optional[BaseCheckpointSaver]): if checkpointer: saved = checkpointer.get_tuple(thread) assert saved is not None - assert saved.checkpoint == { - "v": 1, - "ts": AnyStr(), - "id": AnyStr(), - "channel_values": { - "messages": [ - _AnyIdHumanMessage(content="hi?"), - AIMessage(content="hi?", id="0"), - ], - "agent": "agent", - }, - "channel_versions": { - "__start__": 2, - "messages": 3, - "start:agent": 3, - "agent": 3, - }, - "versions_seen": { - "__input__": {}, - "__start__": {"__start__": 1}, - "agent": {"start:agent": 2}, - }, - "pending_sends": [], - "current_tasks": {}, + assert saved.checkpoint["channel_values"] == { + "messages": [ + _AnyIdHumanMessage(content="hi?"), + AIMessage(content="hi?", id="0"), + ], + "agent": "agent", } assert saved.metadata == { "source": "loop", @@ -110,18 +84,16 @@ def test_no_modifier(checkpointer: Optional[BaseCheckpointSaver]): @pytest.mark.parametrize( - "checkpointer", - [ - MemorySaverAssertImmutable(), - None, - ], - ids=[ - "memory", - "none", - ], + "checkpointer_name", + ["memory", "sqlite_aio", "postgres_aio", "postgres_aio_pipe"], ) -async def test_no_modifier_async(checkpointer: Optional[BaseCheckpointSaver]): +async def test_no_modifier_async( + request: pytest.FixtureRequest, checkpointer_name: str +) -> None: + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") + model = FakeToolCallingModel() + agent = create_react_agent(model, [], checkpointer=checkpointer) inputs = [HumanMessage("hi?")] thread = {"configurable": {"thread_id": "123"}} @@ -132,30 +104,12 @@ async def test_no_modifier_async(checkpointer: Optional[BaseCheckpointSaver]): if checkpointer: saved = await checkpointer.aget_tuple(thread) assert saved is not None - assert saved.checkpoint == { - "v": 1, - "ts": AnyStr(), - "id": AnyStr(), - "channel_values": { - "messages": [ - _AnyIdHumanMessage(content="hi?"), - AIMessage(content="hi?", id="0"), - ], - "agent": "agent", - }, - "channel_versions": { - "__start__": 2, - "messages": 3, - "start:agent": 3, - "agent": 3, - }, - "versions_seen": { - "__input__": {}, - "__start__": {"__start__": 1}, - "agent": {"start:agent": 2}, - }, - "pending_sends": [], - "current_tasks": {}, + assert saved.checkpoint["channel_values"] == { + "messages": [ + _AnyIdHumanMessage(content="hi?"), + AIMessage(content="hi?", id="0"), + ], + "agent": "agent", } assert saved.metadata == { "source": "loop", diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 728effc0e..71f2f8c59 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -9200,7 +9200,13 @@ def test_checkpoint_metadata() -> None: assert chkpnt_tuple.metadata["test_config_4"] == "bar" -def test_remove_message_via_state_update(): +@pytest.mark.parametrize( + "checkpointer_name", + ["memory", "sqlite", "postgres", "postgres_pipe"], +) +def test_remove_message_via_state_update( + request: pytest.FixtureRequest, checkpointer_name: str +) -> None: from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage workflow = MessageGraph() @@ -9216,7 +9222,7 @@ def test_remove_message_via_state_update(): workflow.set_entry_point("chatbot") workflow.add_edge("chatbot", END) - checkpointer = MemorySaver() + checkpointer = request.getfixturevalue("checkpointer_" + checkpointer_name) app = workflow.compile(checkpointer=checkpointer) config = {"configurable": {"thread_id": "1"}} output = app.invoke([HumanMessage(content="Hi")], config=config) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 989d2a1e5..e5c305dd8 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -2,8 +2,7 @@ import asyncio import json import operator from collections import Counter -from contextlib import AbstractAsyncContextManager, asynccontextmanager, contextmanager -from types import TracebackType +from contextlib import asynccontextmanager, contextmanager from typing import ( Annotated, Any, @@ -67,19 +66,6 @@ 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]: