From 909a4591a8adb4461bf54ec463094c4d6e16199f Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 9 May 2025 12:42:44 -0700 Subject: [PATCH] Port checkpointer/store fixtures in langgraph-prebuilt to idiomatic pattern (#4629) --- libs/langgraph/tests/conftest.py | 1 - libs/prebuilt/tests/conftest.py | 515 +++++-------------- libs/prebuilt/tests/conftest_checkpointer.py | 185 +++++++ libs/prebuilt/tests/conftest_store.py | 154 ++++++ libs/prebuilt/tests/test_react_agent.py | 148 ++---- 5 files changed, 507 insertions(+), 496 deletions(-) create mode 100644 libs/prebuilt/tests/conftest_checkpointer.py create mode 100644 libs/prebuilt/tests/conftest_store.py diff --git a/libs/langgraph/tests/conftest.py b/libs/langgraph/tests/conftest.py index 504cb5785..8ca3a8daf 100644 --- a/libs/langgraph/tests/conftest.py +++ b/libs/langgraph/tests/conftest.py @@ -39,7 +39,6 @@ from tests.conftest_store import ( pytest.register_assert_rewrite("tests.memory_assert") -DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5442/" # TODO: fix this once core is released IS_LANGCHAIN_CORE_030_OR_GREATER = version.parse(core_version) >= version.parse( "0.3.0.dev0" diff --git a/libs/prebuilt/tests/conftest.py b/libs/prebuilt/tests/conftest.py index 20f9e15f6..8ba0771c2 100644 --- a/libs/prebuilt/tests/conftest.py +++ b/libs/prebuilt/tests/conftest.py @@ -1,30 +1,36 @@ -import sys -from contextlib import asynccontextmanager -from typing import AsyncIterator, Optional -from uuid import UUID, uuid4 +from collections.abc import AsyncIterator, Iterator +from uuid import UUID import pytest from langchain_core import __version__ as core_version from packaging import version -from psycopg import AsyncConnection, Connection -from psycopg_pool import AsyncConnectionPool, ConnectionPool from pytest_mock import MockerFixture from langgraph.checkpoint.base import BaseCheckpointSaver -from langgraph.checkpoint.postgres import PostgresSaver, ShallowPostgresSaver -from langgraph.checkpoint.postgres.aio import ( - AsyncPostgresSaver, - AsyncShallowPostgresSaver, -) -from langgraph.checkpoint.sqlite import SqliteSaver -from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver from langgraph.store.base import BaseStore -from langgraph.store.memory import InMemoryStore -from langgraph.store.postgres import AsyncPostgresStore, PostgresStore +from tests.conftest_checkpointer import ( + _checkpointer_memory, + _checkpointer_postgres, + _checkpointer_postgres_aio, + _checkpointer_postgres_aio_pipe, + _checkpointer_postgres_aio_pool, + _checkpointer_postgres_pipe, + _checkpointer_postgres_pool, + _checkpointer_sqlite, + _checkpointer_sqlite_aio, +) +from tests.conftest_store import ( + _store_memory, + _store_postgres, + _store_postgres_aio, + _store_postgres_aio_pipe, + _store_postgres_aio_pool, + _store_postgres_pipe, + _store_postgres_pool, +) pytest.register_assert_rewrite("tests.memory_assert") -DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5442/" # TODO: fix this once core is released IS_LANGCHAIN_CORE_030_OR_GREATER = version.parse(core_version) >= version.parse( "0.3.0.dev0" @@ -47,375 +53,41 @@ def deterministic_uuids(mocker: MockerFixture) -> MockerFixture: # checkpointer fixtures -@pytest.fixture(scope="function") -def checkpointer_memory(): - from tests.memory_assert import MemorySaverAssertImmutable - - yield MemorySaverAssertImmutable() - - -@pytest.fixture(scope="function") -def checkpointer_sqlite(): - with SqliteSaver.from_conn_string(":memory:") as checkpointer: - yield checkpointer - - -@asynccontextmanager -async def _checkpointer_sqlite_aio(): - async with AsyncSqliteSaver.from_conn_string(":memory:") as checkpointer: - yield checkpointer - - -@pytest.fixture(scope="function") -def checkpointer_postgres(): - database = f"test_{uuid4().hex[:16]}" - # create unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - with PostgresSaver.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as checkpointer: - checkpointer.setup() - yield checkpointer - finally: - # drop unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"DROP DATABASE {database}") - - -@pytest.fixture(scope="function") -def checkpointer_postgres_shallow(): - database = f"test_{uuid4().hex[:16]}" - # create unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - with ShallowPostgresSaver.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as checkpointer: - checkpointer.setup() - yield checkpointer - finally: - # drop unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"DROP DATABASE {database}") - - -@pytest.fixture(scope="function") -def checkpointer_postgres_pipe(): - database = f"test_{uuid4().hex[:16]}" - # create unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - with PostgresSaver.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as checkpointer: - checkpointer.setup() - # setup can't run inside pipeline because of implicit transaction - with checkpointer.conn.pipeline() as pipe: - checkpointer.pipe = pipe - yield checkpointer - finally: - # drop unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"DROP DATABASE {database}") - - -@pytest.fixture(scope="function") -def checkpointer_postgres_pool(): - database = f"test_{uuid4().hex[:16]}" - # create unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - with ConnectionPool( - DEFAULT_POSTGRES_URI + database, max_size=10, kwargs={"autocommit": True} - ) as pool: - checkpointer = PostgresSaver(pool) - checkpointer.setup() - yield checkpointer - finally: - # drop unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"DROP DATABASE {database}") - - -@asynccontextmanager -async def _checkpointer_postgres_aio(): - if sys.version_info < (3, 10): - pytest.skip("Async Postgres tests require Python 3.10+") - database = f"test_{uuid4().hex[:16]}" - # create unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - async with AsyncPostgresSaver.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as checkpointer: - await checkpointer.setup() - yield checkpointer - finally: - # drop unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"DROP DATABASE {database}") - - -@asynccontextmanager -async def _checkpointer_postgres_aio_shallow(): - if sys.version_info < (3, 10): - pytest.skip("Async Postgres tests require Python 3.10+") - database = f"test_{uuid4().hex[:16]}" - # create unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - async with AsyncShallowPostgresSaver.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as checkpointer: - await checkpointer.setup() - yield checkpointer - finally: - # drop unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"DROP DATABASE {database}") - - -@asynccontextmanager -async def _checkpointer_postgres_aio_pipe(): - if sys.version_info < (3, 10): - pytest.skip("Async Postgres tests require Python 3.10+") - database = f"test_{uuid4().hex[:16]}" - # create unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - async with AsyncPostgresSaver.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as checkpointer: - await checkpointer.setup() - # setup can't run inside pipeline because of implicit transaction - async with checkpointer.conn.pipeline() as pipe: - checkpointer.pipe = pipe - yield checkpointer - finally: - # drop unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"DROP DATABASE {database}") - - -@asynccontextmanager -async def _checkpointer_postgres_aio_pool(): - if sys.version_info < (3, 10): - pytest.skip("Async Postgres tests require Python 3.10+") - database = f"test_{uuid4().hex[:16]}" - # create unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"CREATE DATABASE {database}") - try: - # yield checkpointer - async with AsyncConnectionPool( - DEFAULT_POSTGRES_URI + database, max_size=10, kwargs={"autocommit": True} - ) as pool: - checkpointer = AsyncPostgresSaver(pool) - await checkpointer.setup() - yield checkpointer - finally: - # drop unique db - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"DROP DATABASE {database}") - - -@asynccontextmanager -async def awith_checkpointer( - checkpointer_name: Optional[str], -) -> AsyncIterator[BaseCheckpointSaver]: - if checkpointer_name is None: - yield None - elif checkpointer_name == "memory": - from tests.memory_assert import MemorySaverAssertImmutable - - yield MemorySaverAssertImmutable() - elif checkpointer_name == "sqlite_aio": - async with _checkpointer_sqlite_aio() as checkpointer: - yield checkpointer - elif checkpointer_name == "postgres_aio": - async with _checkpointer_postgres_aio() as checkpointer: - yield checkpointer - elif checkpointer_name == "postgres_aio_shallow": - async with _checkpointer_postgres_aio_shallow() as checkpointer: - yield checkpointer - elif checkpointer_name == "postgres_aio_pipe": - async with _checkpointer_postgres_aio_pipe() as checkpointer: - yield checkpointer - elif checkpointer_name == "postgres_aio_pool": - async with _checkpointer_postgres_aio_pool() as checkpointer: - yield checkpointer - else: - raise NotImplementedError(f"Unknown checkpointer: {checkpointer_name}") - - -@asynccontextmanager -async def _store_postgres_aio(): - if sys.version_info < (3, 10): - pytest.skip("Async Postgres tests require Python 3.10+") - database = f"test_{uuid4().hex[:16]}" - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"CREATE DATABASE {database}") - try: - async with AsyncPostgresStore.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as store: - await store.setup() - yield store - finally: - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"DROP DATABASE {database}") - - -@asynccontextmanager -async def _store_postgres_aio_pipe(): - if sys.version_info < (3, 10): - pytest.skip("Async Postgres tests require Python 3.10+") - database = f"test_{uuid4().hex[:16]}" - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"CREATE DATABASE {database}") - try: - async with AsyncPostgresStore.from_conn_string( - DEFAULT_POSTGRES_URI + database - ) as store: - await store.setup() # Run in its own transaction - async with AsyncPostgresStore.from_conn_string( - DEFAULT_POSTGRES_URI + database, pipeline=True - ) as store: - yield store - finally: - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"DROP DATABASE {database}") - - -@asynccontextmanager -async def _store_postgres_aio_pool(): - if sys.version_info < (3, 10): - pytest.skip("Async Postgres tests require Python 3.10+") - database = f"test_{uuid4().hex[:16]}" - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"CREATE DATABASE {database}") - try: - async with AsyncPostgresStore.from_conn_string( - DEFAULT_POSTGRES_URI + database, - pool_config={"max_size": 10}, - ) as store: - await store.setup() - yield store - finally: - async with await AsyncConnection.connect( - DEFAULT_POSTGRES_URI, autocommit=True - ) as conn: - await conn.execute(f"DROP DATABASE {database}") - - -@pytest.fixture(scope="function") -def store_postgres(): - database = f"test_{uuid4().hex[:16]}" - # create unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"CREATE DATABASE {database}") - try: - # yield store - with PostgresStore.from_conn_string(DEFAULT_POSTGRES_URI + database) as store: - store.setup() - yield store - finally: - # drop unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"DROP DATABASE {database}") - - -@pytest.fixture(scope="function") -def store_postgres_pipe(): - database = f"test_{uuid4().hex[:16]}" - # create unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"CREATE DATABASE {database}") - try: - # yield store - with PostgresStore.from_conn_string(DEFAULT_POSTGRES_URI + database) as store: - store.setup() # Run in its own transaction - with PostgresStore.from_conn_string( - DEFAULT_POSTGRES_URI + database, pipeline=True - ) as store: - yield store - finally: - # drop unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"DROP DATABASE {database}") - - -@pytest.fixture(scope="function") -def store_postgres_pool(): - database = f"test_{uuid4().hex[:16]}" - # create unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"CREATE DATABASE {database}") - try: - # yield store - with PostgresStore.from_conn_string( - DEFAULT_POSTGRES_URI + database, pool_config={"max_size": 10} - ) as store: - store.setup() - yield store - finally: - # drop unique db - with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: - conn.execute(f"DROP DATABASE {database}") - - -@pytest.fixture(scope="function") -def store_in_memory(): - yield InMemoryStore() - - -@asynccontextmanager -async def awith_store(store_name: Optional[str]) -> AsyncIterator[BaseStore]: +@pytest.fixture( + scope="function", + params=["in_memory", "postgres", "postgres_pipe", "postgres_pool"], +) +def sync_store(request: pytest.FixtureRequest) -> Iterator[BaseStore]: + store_name = request.param if store_name is None: yield None elif store_name == "in_memory": - yield InMemoryStore() + with _store_memory() as store: + yield store + elif store_name == "postgres": + with _store_postgres() as store: + yield store + elif store_name == "postgres_pipe": + with _store_postgres_pipe() as store: + yield store + elif store_name == "postgres_pool": + with _store_postgres_pool() as store: + yield store + else: + raise NotImplementedError(f"Unknown store {store_name}") + + +@pytest.fixture( + scope="function", + params=["in_memory", "postgres_aio", "postgres_aio_pipe", "postgres_aio_pool"], +) +async def async_store(request: pytest.FixtureRequest) -> AsyncIterator[BaseStore]: + store_name = request.param + if store_name is None: + yield None + elif store_name == "in_memory": + with _store_memory() as store: + yield store elif store_name == "postgres_aio": async with _store_postgres_aio() as store: yield store @@ -429,20 +101,67 @@ async def awith_store(store_name: Optional[str]) -> AsyncIterator[BaseStore]: raise NotImplementedError(f"Unknown store {store_name}") -ALL_CHECKPOINTERS_SYNC = [ - "memory", - "sqlite", - "postgres", - "postgres_pipe", - "postgres_pool", - "postgres_shallow", -] +@pytest.fixture( + scope="function", + params=[ + "memory", + "sqlite", + "postgres", + "postgres_pipe", + "postgres_pool", + ], +) +def sync_checkpointer( + request: pytest.FixtureRequest, +) -> Iterator[BaseCheckpointSaver]: + checkpointer_name = request.param + if checkpointer_name == "memory": + with _checkpointer_memory() as checkpointer: + yield checkpointer + elif checkpointer_name == "sqlite": + with _checkpointer_sqlite() as checkpointer: + yield checkpointer + elif checkpointer_name == "postgres": + with _checkpointer_postgres() as checkpointer: + yield checkpointer + elif checkpointer_name == "postgres_pipe": + with _checkpointer_postgres_pipe() as checkpointer: + yield checkpointer + elif checkpointer_name == "postgres_pool": + with _checkpointer_postgres_pool() as checkpointer: + yield checkpointer + else: + raise NotImplementedError(f"Unknown checkpointer: {checkpointer_name}") -ALL_CHECKPOINTERS_ASYNC = [ - "memory", - "sqlite_aio", - "postgres_aio", - "postgres_aio_pipe", - "postgres_aio_pool", - "postgres_aio_shallow", -] + +@pytest.fixture( + scope="function", + params=[ + "memory", + "sqlite_aio", + "postgres_aio", + "postgres_aio_pipe", + "postgres_aio_pool", + ], +) +async def async_checkpointer( + request: pytest.FixtureRequest, +) -> AsyncIterator[BaseCheckpointSaver]: + checkpointer_name = request.param + if checkpointer_name == "memory": + with _checkpointer_memory() as checkpointer: + yield checkpointer + elif checkpointer_name == "sqlite_aio": + async with _checkpointer_sqlite_aio() as checkpointer: + yield checkpointer + elif checkpointer_name == "postgres_aio": + async with _checkpointer_postgres_aio() as checkpointer: + yield checkpointer + elif checkpointer_name == "postgres_aio_pipe": + async with _checkpointer_postgres_aio_pipe() as checkpointer: + yield checkpointer + elif checkpointer_name == "postgres_aio_pool": + async with _checkpointer_postgres_aio_pool() as checkpointer: + yield checkpointer + else: + raise NotImplementedError(f"Unknown checkpointer: {checkpointer_name}") diff --git a/libs/prebuilt/tests/conftest_checkpointer.py b/libs/prebuilt/tests/conftest_checkpointer.py new file mode 100644 index 000000000..b3c782001 --- /dev/null +++ b/libs/prebuilt/tests/conftest_checkpointer.py @@ -0,0 +1,185 @@ +import sys +from contextlib import asynccontextmanager, contextmanager +from uuid import uuid4 + +import pytest +from psycopg import AsyncConnection, Connection +from psycopg_pool import AsyncConnectionPool, ConnectionPool + +from langgraph.checkpoint.postgres import PostgresSaver +from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver +from langgraph.checkpoint.sqlite import SqliteSaver +from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver +from tests.memory_assert import MemorySaverAssertImmutable + +DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5442/" + + +@contextmanager +def _checkpointer_memory(): + yield MemorySaverAssertImmutable() + + +@contextmanager +def _checkpointer_sqlite(): + with SqliteSaver.from_conn_string(":memory:") as checkpointer: + yield checkpointer + + +@contextmanager +def _checkpointer_postgres(): + database = f"test_{uuid4().hex[:16]}" + # create unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"CREATE DATABASE {database}") + try: + # yield checkpointer + with PostgresSaver.from_conn_string( + DEFAULT_POSTGRES_URI + database + ) as checkpointer: + checkpointer.setup() + yield checkpointer + finally: + # drop unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"DROP DATABASE {database}") + + +@contextmanager +def _checkpointer_postgres_pipe(): + database = f"test_{uuid4().hex[:16]}" + # create unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"CREATE DATABASE {database}") + try: + # yield checkpointer + with PostgresSaver.from_conn_string( + DEFAULT_POSTGRES_URI + database + ) as checkpointer: + checkpointer.setup() + # setup can't run inside pipeline because of implicit transaction + with checkpointer.conn.pipeline() as pipe: + checkpointer.pipe = pipe + yield checkpointer + finally: + # drop unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"DROP DATABASE {database}") + + +@contextmanager +def _checkpointer_postgres_pool(): + database = f"test_{uuid4().hex[:16]}" + # create unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"CREATE DATABASE {database}") + try: + # yield checkpointer + with ConnectionPool( + DEFAULT_POSTGRES_URI + database, max_size=10, kwargs={"autocommit": True} + ) as pool: + checkpointer = PostgresSaver(pool) + checkpointer.setup() + yield checkpointer + finally: + # drop unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"DROP DATABASE {database}") + + +@asynccontextmanager +async def _checkpointer_sqlite_aio(): + async with AsyncSqliteSaver.from_conn_string(":memory:") as checkpointer: + yield checkpointer + + +@asynccontextmanager +async def _checkpointer_postgres_aio(): + if sys.version_info < (3, 10): + pytest.skip("Async Postgres tests require Python 3.10+") + database = f"test_{uuid4().hex[:16]}" + # create unique db + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"CREATE DATABASE {database}") + try: + # yield checkpointer + async with AsyncPostgresSaver.from_conn_string( + DEFAULT_POSTGRES_URI + database + ) as checkpointer: + await checkpointer.setup() + yield checkpointer + finally: + # drop unique db + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"DROP DATABASE {database}") + + +@asynccontextmanager +async def _checkpointer_postgres_aio_pipe(): + if sys.version_info < (3, 10): + pytest.skip("Async Postgres tests require Python 3.10+") + database = f"test_{uuid4().hex[:16]}" + # create unique db + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"CREATE DATABASE {database}") + try: + # yield checkpointer + async with AsyncPostgresSaver.from_conn_string( + DEFAULT_POSTGRES_URI + database + ) as checkpointer: + await checkpointer.setup() + # setup can't run inside pipeline because of implicit transaction + async with checkpointer.conn.pipeline() as pipe: + checkpointer.pipe = pipe + yield checkpointer + finally: + # drop unique db + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"DROP DATABASE {database}") + + +@asynccontextmanager +async def _checkpointer_postgres_aio_pool(): + if sys.version_info < (3, 10): + pytest.skip("Async Postgres tests require Python 3.10+") + database = f"test_{uuid4().hex[:16]}" + # create unique db + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"CREATE DATABASE {database}") + try: + # yield checkpointer + async with AsyncConnectionPool( + DEFAULT_POSTGRES_URI + database, max_size=10, kwargs={"autocommit": True} + ) as pool: + checkpointer = AsyncPostgresSaver(pool) + await checkpointer.setup() + yield checkpointer + finally: + # drop unique db + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"DROP DATABASE {database}") + + +__all__ = [ + "_checkpointer_memory", + "_checkpointer_sqlite", + "_checkpointer_postgres", + "_checkpointer_postgres_pipe", + "_checkpointer_postgres_pool", + "_checkpointer_sqlite_aio", + "_checkpointer_postgres_aio", + "_checkpointer_postgres_aio_pipe", + "_checkpointer_postgres_aio_pool", +] diff --git a/libs/prebuilt/tests/conftest_store.py b/libs/prebuilt/tests/conftest_store.py new file mode 100644 index 000000000..9d047a8f8 --- /dev/null +++ b/libs/prebuilt/tests/conftest_store.py @@ -0,0 +1,154 @@ +import sys +from contextlib import asynccontextmanager, contextmanager +from uuid import uuid4 + +import pytest +from psycopg import AsyncConnection, Connection + +from langgraph.store.memory import InMemoryStore +from langgraph.store.postgres import AsyncPostgresStore, PostgresStore + +DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5442/" + + +@contextmanager +def _store_memory(): + store = InMemoryStore() + yield store + + +@contextmanager +def _store_postgres(): + database = f"test_{uuid4().hex[:16]}" + # create unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"CREATE DATABASE {database}") + try: + # yield store + with PostgresStore.from_conn_string(DEFAULT_POSTGRES_URI + database) as store: + store.setup() + yield store + finally: + # drop unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"DROP DATABASE {database}") + + +@contextmanager +def _store_postgres_pipe(): + database = f"test_{uuid4().hex[:16]}" + # create unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"CREATE DATABASE {database}") + try: + # yield store + with PostgresStore.from_conn_string(DEFAULT_POSTGRES_URI + database) as store: + store.setup() # Run in its own transaction + with PostgresStore.from_conn_string( + DEFAULT_POSTGRES_URI + database, pipeline=True + ) as store: + yield store + finally: + # drop unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"DROP DATABASE {database}") + + +@contextmanager +def _store_postgres_pool(): + database = f"test_{uuid4().hex[:16]}" + # create unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"CREATE DATABASE {database}") + try: + # yield store + with PostgresStore.from_conn_string( + DEFAULT_POSTGRES_URI + database, pool_config={"max_size": 10} + ) as store: + store.setup() + yield store + finally: + # drop unique db + with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn: + conn.execute(f"DROP DATABASE {database}") + + +@asynccontextmanager +async def _store_postgres_aio(): + if sys.version_info < (3, 10): + pytest.skip("Async Postgres tests require Python 3.10+") + database = f"test_{uuid4().hex[:16]}" + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"CREATE DATABASE {database}") + try: + async with AsyncPostgresStore.from_conn_string( + DEFAULT_POSTGRES_URI + database + ) as store: + await store.setup() + yield store + finally: + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"DROP DATABASE {database}") + + +@asynccontextmanager +async def _store_postgres_aio_pipe(): + if sys.version_info < (3, 10): + pytest.skip("Async Postgres tests require Python 3.10+") + database = f"test_{uuid4().hex[:16]}" + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"CREATE DATABASE {database}") + try: + async with AsyncPostgresStore.from_conn_string( + DEFAULT_POSTGRES_URI + database + ) as store: + await store.setup() # Run in its own transaction + async with AsyncPostgresStore.from_conn_string( + DEFAULT_POSTGRES_URI + database, pipeline=True + ) as store: + yield store + finally: + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"DROP DATABASE {database}") + + +@asynccontextmanager +async def _store_postgres_aio_pool(): + if sys.version_info < (3, 10): + pytest.skip("Async Postgres tests require Python 3.10+") + database = f"test_{uuid4().hex[:16]}" + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"CREATE DATABASE {database}") + try: + async with AsyncPostgresStore.from_conn_string( + DEFAULT_POSTGRES_URI + database, + pool_config={"max_size": 10}, + ) as store: + await store.setup() + yield store + finally: + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI, autocommit=True + ) as conn: + await conn.execute(f"DROP DATABASE {database}") + + +__all__ = [ + "_store_memory", + "_store_postgres", + "_store_postgres_pipe", + "_store_postgres_pool", + "_store_postgres_aio", + "_store_postgres_aio_pipe", + "_store_postgres_aio_pool", +] diff --git a/libs/prebuilt/tests/test_react_agent.py b/libs/prebuilt/tests/test_react_agent.py index df092675d..25d17000b 100644 --- a/libs/prebuilt/tests/test_react_agent.py +++ b/libs/prebuilt/tests/test_react_agent.py @@ -55,12 +55,7 @@ from langgraph.store.memory import InMemoryStore from langgraph.types import Command, Interrupt, interrupt from langgraph.utils.config import get_stream_writer from tests.any_str import AnyStr -from tests.conftest import ( - ALL_CHECKPOINTERS_ASYNC, - ALL_CHECKPOINTERS_SYNC, - IS_LANGCHAIN_CORE_030_OR_GREATER, - awith_checkpointer, -) +from tests.conftest import IS_LANGCHAIN_CORE_030_OR_GREATER from tests.messages import _AnyIdHumanMessage, _AnyIdToolMessage from tests.model import FakeToolCallingModel @@ -69,20 +64,14 @@ pytestmark = pytest.mark.anyio REACT_TOOL_CALL_VERSIONS = ["v1", "v2"] -@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) -def test_no_prompt( - request: pytest.FixtureRequest, checkpointer_name: str, version: str -) -> None: - checkpointer: BaseCheckpointSaver = request.getfixturevalue( - "checkpointer_" + checkpointer_name - ) +def test_no_prompt(sync_checkpointer: BaseCheckpointSaver, version: str) -> None: model = FakeToolCallingModel() agent = create_react_agent( model, [], - checkpointer=checkpointer, + checkpointer=sync_checkpointer, version=version, ) inputs = [HumanMessage("hi?")] @@ -91,54 +80,50 @@ def test_no_prompt( expected_response = {"messages": inputs + [AIMessage(content="hi?", id="0")]} assert response == expected_response - if checkpointer: - saved = checkpointer.get_tuple(thread) - assert saved is not None - assert saved.checkpoint["channel_values"] == { - "messages": [ - _AnyIdHumanMessage(content="hi?"), - AIMessage(content="hi?", id="0"), - ], - } - assert saved.metadata == { - "parents": {}, - "source": "loop", - "writes": {"agent": {"messages": [AIMessage(content="hi?", id="0")]}}, - "step": 1, - "thread_id": "123", - } - assert saved.pending_writes == [] + saved = sync_checkpointer.get_tuple(thread) + assert saved is not None + assert saved.checkpoint["channel_values"] == { + "messages": [ + _AnyIdHumanMessage(content="hi?"), + AIMessage(content="hi?", id="0"), + ], + } + assert saved.metadata == { + "parents": {}, + "source": "loop", + "writes": {"agent": {"messages": [AIMessage(content="hi?", id="0")]}}, + "step": 1, + "thread_id": "123", + } + assert saved.pending_writes == [] -@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) -async def test_no_prompt_async(checkpointer_name: str) -> None: - async with awith_checkpointer(checkpointer_name) as checkpointer: - model = FakeToolCallingModel() +async def test_no_prompt_async(async_checkpointer: BaseCheckpointSaver) -> None: + model = FakeToolCallingModel() - agent = create_react_agent(model, [], checkpointer=checkpointer) - inputs = [HumanMessage("hi?")] - thread = {"configurable": {"thread_id": "123"}} - response = await agent.ainvoke({"messages": inputs}, thread, debug=True) - expected_response = {"messages": inputs + [AIMessage(content="hi?", id="0")]} - assert response == expected_response + agent = create_react_agent(model, [], checkpointer=async_checkpointer) + inputs = [HumanMessage("hi?")] + thread = {"configurable": {"thread_id": "123"}} + response = await agent.ainvoke({"messages": inputs}, thread, debug=True) + expected_response = {"messages": inputs + [AIMessage(content="hi?", id="0")]} + assert response == expected_response - if checkpointer: - saved = await checkpointer.aget_tuple(thread) - assert saved is not None - assert saved.checkpoint["channel_values"] == { - "messages": [ - _AnyIdHumanMessage(content="hi?"), - AIMessage(content="hi?", id="0"), - ], - } - assert saved.metadata == { - "parents": {}, - "source": "loop", - "writes": {"agent": {"messages": [AIMessage(content="hi?", id="0")]}}, - "step": 1, - "thread_id": "123", - } - assert saved.pending_writes == [] + saved = await async_checkpointer.aget_tuple(thread) + assert saved is not None + assert saved.checkpoint["channel_values"] == { + "messages": [ + _AnyIdHumanMessage(content="hi?"), + AIMessage(content="hi?", id="0"), + ], + } + assert saved.metadata == { + "parents": {}, + "source": "loop", + "writes": {"agent": {"messages": [AIMessage(content="hi?", id="0")]}}, + "step": 1, + "thread_id": "123", + } + assert saved.pending_writes == [] def test_system_message_prompt(): @@ -515,19 +500,13 @@ class CustomStatePydantic(AgentStatePydantic): not IS_LANGCHAIN_CORE_030_OR_GREATER, reason="Langchain core 0.3.0 or greater is required", ) -@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) @pytest.mark.parametrize("state_schema", [CustomState, CustomStatePydantic]) def test_react_agent_update_state( - request: pytest.FixtureRequest, - checkpointer_name: str, + sync_checkpointer: BaseCheckpointSaver, version: str, state_schema: StateSchemaType, ) -> None: - checkpointer: BaseCheckpointSaver = request.getfixturevalue( - "checkpointer_" + checkpointer_name - ) - @dec_tool def get_user_name(tool_call_id: Annotated[str, InjectedToolCallId]): """Retrieve user name""" @@ -569,7 +548,7 @@ def test_react_agent_update_state( [get_user_name], state_schema=state_schema, prompt=prompt, - checkpointer=checkpointer, + checkpointer=sync_checkpointer, version=version, ) config = {"configurable": {"thread_id": "1"}} @@ -590,21 +569,10 @@ def test_react_agent_update_state( not IS_LANGCHAIN_CORE_030_OR_GREATER, reason="Langchain core 0.3.0 or greater is required", ) -@pytest.mark.parametrize( - "checkpointer_name", - [ - checkpointer - for checkpointer in ALL_CHECKPOINTERS_SYNC - if "shallow" not in checkpointer - ], -) @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) def test_react_agent_parallel_tool_calls( - request: pytest.FixtureRequest, checkpointer_name: str, version: str + sync_checkpointer: BaseCheckpointSaver, version: str ) -> None: - checkpointer: BaseCheckpointSaver = request.getfixturevalue( - "checkpointer_" + checkpointer_name - ) human_assistance_execution_count = 0 @dec_tool @@ -635,7 +603,7 @@ def test_react_agent_parallel_tool_calls( agent = create_react_agent( model, [human_assistance, get_weather], - checkpointer=checkpointer, + checkpointer=sync_checkpointer, version=version, ) config = {"configurable": {"thread_id": "1"}} @@ -1125,14 +1093,9 @@ def test_inspect_react() -> None: @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) -@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_react_with_subgraph_tools( - request: pytest.FixtureRequest, checkpointer_name: str, version: str + sync_checkpointer: BaseCheckpointSaver, version: str ) -> None: - checkpointer: BaseCheckpointSaver = request.getfixturevalue( - "checkpointer_" + checkpointer_name - ) - class State(TypedDict): a: int b: int @@ -1183,7 +1146,7 @@ def test_react_with_subgraph_tools( agent = create_react_agent( model, tool_node, - checkpointer=checkpointer, + checkpointer=sync_checkpointer, version=version, ) result = agent.invoke( @@ -1270,14 +1233,9 @@ def test_tool_node_stream_writer() -> None: @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) -@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_tool_node_node_interrupt( - request: pytest.FixtureRequest, checkpointer_name: str, version: str + sync_checkpointer: BaseCheckpointSaver, version: str ) -> None: - checkpointer: BaseCheckpointSaver = request.getfixturevalue( - "checkpointer_" + checkpointer_name - ) - def tool_normal(some_val: int) -> str: """Tool docstring.""" return "normal" @@ -1301,7 +1259,7 @@ def test_tool_node_node_interrupt( agent = create_react_agent( model, [tool_interrupt, tool_normal], - checkpointer=checkpointer, + checkpointer=sync_checkpointer, version=version, ) result = agent.invoke({"messages": [HumanMessage("hi?")]}, config) @@ -1333,10 +1291,6 @@ def test_tool_node_node_interrupt( elif version == "v2": assert result["messages"] == expected_messages - # TODO: figure out why this is not working w/ shallow postgres checkpointer - if "shallow" in checkpointer_name: - return - state = agent.get_state(config) assert state.next == ("tools",) task = state.tasks[0]