langgraph: more checkpointer tests (#1263)

* langgraph: more checkpointer tests

* more tests

* lint

* update tests

---------

Co-authored-by: Nuno Campos <nuno@langchain.dev>
This commit is contained in:
Vadym Barda
2024-08-08 00:45:46 +00:00
committed by GitHub
co-authored by Nuno Campos
parent f9fa35ed82
commit a2f4d57bf2
4 changed files with 55 additions and 97 deletions
+21 -9
View File
@@ -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"}}
+25 -71
View File
@@ -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",
+8 -2
View File
@@ -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)
+1 -15
View File
@@ -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]: