mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-10 03:37:51 +02:00
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:
co-authored by
Nuno Campos
parent
f9fa35ed82
commit
a2f4d57bf2
@@ -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"}}
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]:
|
||||
|
||||
Reference in New Issue
Block a user