diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index d8af3aeca..e5a3cce55 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -5,7 +5,6 @@ from typing import Any, Optional from langchain_core.runnables import RunnableConfig from psycopg import Capabilities, Connection, Cursor, Pipeline -from psycopg.errors import UndefinedTable from psycopg.rows import DictRow, dict_row from psycopg.types.json import Jsonb from psycopg_pool import ConnectionPool @@ -76,16 +75,15 @@ class PostgresSaver(BasePostgresSaver): the first time checkpointer is used. """ with self._cursor() as cur: - try: - row = cur.execute( - "SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1" - ).fetchone() - if row is None: - version = -1 - else: - version = row["v"] - except UndefinedTable: + cur.execute(self.MIGRATIONS[0]) + results = cur.execute( + "SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1" + ) + row = results.fetchone() + if row is None: version = -1 + else: + version = row["v"] for v, migration in zip( range(version + 1, len(self.MIGRATIONS)), self.MIGRATIONS[version + 1 :], diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index 440cb452e..4c0f5295c 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -5,7 +5,6 @@ from typing import Any, Optional from langchain_core.runnables import RunnableConfig from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities -from psycopg.errors import UndefinedTable from psycopg.rows import DictRow, dict_row from psycopg.types.json import Jsonb from psycopg_pool import AsyncConnectionPool @@ -81,17 +80,15 @@ class AsyncPostgresSaver(BasePostgresSaver): the first time checkpointer is used. """ async with self._cursor() as cur: - try: - results = await cur.execute( - "SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1" - ) - row = await results.fetchone() - if row is None: - version = -1 - else: - version = row["v"] - except UndefinedTable: + await cur.execute(self.MIGRATIONS[0]) + results = await cur.execute( + "SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1" + ) + row = await results.fetchone() + if row is None: version = -1 + else: + version = row["v"] for v, migration in zip( range(version + 1, len(self.MIGRATIONS)), self.MIGRATIONS[version + 1 :], diff --git a/libs/checkpoint-postgres/tests/conftest.py b/libs/checkpoint-postgres/tests/conftest.py index ab59dbc6b..b44977ebd 100644 --- a/libs/checkpoint-postgres/tests/conftest.py +++ b/libs/checkpoint-postgres/tests/conftest.py @@ -7,6 +7,7 @@ from psycopg.rows import DictRow, dict_row from tests.embed_test_utils import CharacterEmbeddings +DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5441/" DEFAULT_URI = "postgres://postgres:postgres@localhost:5441/postgres?sslmode=disable" diff --git a/libs/checkpoint-postgres/tests/test_async.py b/libs/checkpoint-postgres/tests/test_async.py index 73c376fd2..d4d0eb8fa 100644 --- a/libs/checkpoint-postgres/tests/test_async.py +++ b/libs/checkpoint-postgres/tests/test_async.py @@ -1,7 +1,14 @@ +# type: ignore + +from contextlib import asynccontextmanager from typing import Any +from uuid import uuid4 import pytest from langchain_core.runnables import RunnableConfig +from psycopg import AsyncConnection +from psycopg.rows import dict_row +from psycopg_pool import AsyncConnectionPool from langgraph.checkpoint.base import ( Checkpoint, @@ -10,104 +17,212 @@ from langgraph.checkpoint.base import ( empty_checkpoint, ) from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver -from tests.conftest import DEFAULT_URI +from tests.conftest import DEFAULT_POSTGRES_URI -class TestAsyncPostgresSaver: - @pytest.fixture(autouse=True) - async def setup(self) -> None: - # objects for test setup - self.config_1: RunnableConfig = { - "configurable": { - "thread_id": "thread-1", - # for backwards compatibility testing - "thread_ts": "1", - "checkpoint_ns": "", - } - } - self.config_2: RunnableConfig = { - "configurable": { - "thread_id": "thread-2", - "checkpoint_id": "2", - "checkpoint_ns": "", - } - } - self.config_3: RunnableConfig = { - "configurable": { - "thread_id": "thread-2", - "checkpoint_id": "2-inner", - "checkpoint_ns": "inner", - } - } +@asynccontextmanager +async def _pool_saver(): + """Fixture for pool mode testing.""" + 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, "row_factory": dict_row}, + ) 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}") - self.chkpnt_1: Checkpoint = empty_checkpoint() - self.chkpnt_2: Checkpoint = create_checkpoint(self.chkpnt_1, {}, 1) - self.chkpnt_3: Checkpoint = empty_checkpoint() - self.metadata_1: CheckpointMetadata = { - "source": "input", - "step": 2, - "writes": {}, - "score": 1, +@asynccontextmanager +async def _pipe_saver(): + """Fixture for pipeline mode testing.""" + 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: + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI + database, + autocommit=True, + prepare_threshold=0, + row_factory=dict_row, + ) as conn: + async with conn.pipeline() as pipe: + checkpointer = AsyncPostgresSaver(conn, pipe=pipe) + await checkpointer.setup() + async with conn.pipeline() as pipe: + checkpointer = AsyncPostgresSaver(conn, 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 _base_saver(): + """Fixture for regular connection mode testing.""" + 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: + async with await AsyncConnection.connect( + DEFAULT_POSTGRES_URI + database, + autocommit=True, + prepare_threshold=0, + row_factory=dict_row, + ) as conn: + checkpointer = AsyncPostgresSaver(conn) + 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 _saver(name: str): + if name == "base": + async with _base_saver() as saver: + yield saver + elif name == "pool": + async with _pool_saver() as saver: + yield saver + elif name == "pipe": + async with _pipe_saver() as saver: + yield saver + + +@pytest.fixture +def test_data(): + """Fixture providing test data for checkpoint tests.""" + config_1: RunnableConfig = { + "configurable": { + "thread_id": "thread-1", + # for backwards compatibility testing + "thread_ts": "1", + "checkpoint_ns": "", } - self.metadata_2: CheckpointMetadata = { - "source": "loop", + } + config_2: RunnableConfig = { + "configurable": { + "thread_id": "thread-2", + "checkpoint_id": "2", + "checkpoint_ns": "", + } + } + config_3: RunnableConfig = { + "configurable": { + "thread_id": "thread-2", + "checkpoint_id": "2-inner", + "checkpoint_ns": "inner", + } + } + + chkpnt_1: Checkpoint = empty_checkpoint() + chkpnt_2: Checkpoint = create_checkpoint(chkpnt_1, {}, 1) + chkpnt_3: Checkpoint = empty_checkpoint() + + metadata_1: CheckpointMetadata = { + "source": "input", + "step": 2, + "writes": {}, + "score": 1, + } + metadata_2: CheckpointMetadata = { + "source": "loop", + "step": 1, + "writes": {"foo": "bar"}, + "score": None, + } + metadata_3: CheckpointMetadata = {} + + return { + "configs": [config_1, config_2, config_3], + "checkpoints": [chkpnt_1, chkpnt_2, chkpnt_3], + "metadata": [metadata_1, metadata_2, metadata_3], + } + + +@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"]) +async def test_asearch(request, saver_name: str, test_data) -> None: + async with _saver(saver_name) as saver: + configs = test_data["configs"] + checkpoints = test_data["checkpoints"] + metadata = test_data["metadata"] + + await saver.aput(configs[0], checkpoints[0], metadata[0], {}) + await saver.aput(configs[1], checkpoints[1], metadata[1], {}) + await saver.aput(configs[2], checkpoints[2], metadata[2], {}) + + # call method / assertions + query_1 = {"source": "input"} # search by 1 key + query_2 = { "step": 1, "writes": {"foo": "bar"}, - "score": None, - } - self.metadata_3: CheckpointMetadata = {} - async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver: - await saver.setup() + } # search by multiple keys + query_3: dict[str, Any] = {} # search by no keys, return all checkpoints + query_4 = {"source": "update", "step": 1} # no match - async def test_asearch(self) -> None: - async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver: - await saver.aput(self.config_1, self.chkpnt_1, self.metadata_1, {}) - await saver.aput(self.config_2, self.chkpnt_2, self.metadata_2, {}) - await saver.aput(self.config_3, self.chkpnt_3, self.metadata_3, {}) + search_results_1 = [c async for c in saver.alist(None, filter=query_1)] + assert len(search_results_1) == 1 + assert search_results_1[0].metadata == metadata[0] - # call method / assertions - query_1 = {"source": "input"} # search by 1 key - query_2 = { - "step": 1, - "writes": {"foo": "bar"}, - } # search by multiple keys - query_3: dict[str, Any] = {} # search by no keys, return all checkpoints - query_4 = {"source": "update", "step": 1} # no match + search_results_2 = [c async for c in saver.alist(None, filter=query_2)] + assert len(search_results_2) == 1 + assert search_results_2[0].metadata == metadata[1] - search_results_1 = [c async for c in saver.alist(None, filter=query_1)] - assert len(search_results_1) == 1 - assert search_results_1[0].metadata == self.metadata_1 + search_results_3 = [c async for c in saver.alist(None, filter=query_3)] + assert len(search_results_3) == 3 - search_results_2 = [c async for c in saver.alist(None, filter=query_2)] - assert len(search_results_2) == 1 - assert search_results_2[0].metadata == self.metadata_2 + search_results_4 = [c async for c in saver.alist(None, filter=query_4)] + assert len(search_results_4) == 0 - search_results_3 = [c async for c in saver.alist(None, filter=query_3)] - assert len(search_results_3) == 3 + # search by config (defaults to checkpoints across all namespaces) + search_results_5 = [ + c async for c in saver.alist({"configurable": {"thread_id": "thread-2"}}) + ] + assert len(search_results_5) == 2 + assert { + search_results_5[0].config["configurable"]["checkpoint_ns"], + search_results_5[1].config["configurable"]["checkpoint_ns"], + } == {"", "inner"} - search_results_4 = [c async for c in saver.alist(None, filter=query_4)] - assert len(search_results_4) == 0 - # search by config (defaults to checkpoints across all namespaces) - search_results_5 = [ - c - async for c in saver.alist({"configurable": {"thread_id": "thread-2"}}) - ] - assert len(search_results_5) == 2 - assert { - search_results_5[0].config["configurable"]["checkpoint_ns"], - search_results_5[1].config["configurable"]["checkpoint_ns"], - } == {"", "inner"} - - # TODO: test before and limit params - - async def test_null_chars(self) -> None: - async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver: - config = await saver.aput( - self.config_1, self.chkpnt_1, {"my_key": "\x00abc"}, {} - ) - assert (await saver.aget_tuple(config)).metadata["my_key"] == "abc" # type: ignore - assert [c async for c in saver.alist(None, filter={"my_key": "abc"})][ - 0 - ].metadata["my_key"] == "abc" +@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"]) +async def test_null_chars(request, saver_name: str, test_data) -> None: + async with _saver(saver_name) as saver: + config = await saver.aput( + test_data["configs"][0], + test_data["checkpoints"][0], + {"my_key": "\x00abc"}, + {}, + ) + assert (await saver.aget_tuple(config)).metadata["my_key"] == "abc" # type: ignore + assert [c async for c in saver.alist(None, filter={"my_key": "abc"})][ + 0 + ].metadata["my_key"] == "abc" diff --git a/libs/checkpoint-postgres/tests/test_sync.py b/libs/checkpoint-postgres/tests/test_sync.py index 052e699b3..fbf4c1c88 100644 --- a/libs/checkpoint-postgres/tests/test_sync.py +++ b/libs/checkpoint-postgres/tests/test_sync.py @@ -1,7 +1,14 @@ +# type: ignore + +from contextlib import contextmanager from typing import Any +from uuid import uuid4 import pytest from langchain_core.runnables import RunnableConfig +from psycopg import Connection +from psycopg.rows import dict_row +from psycopg_pool import ConnectionPool from langgraph.checkpoint.base import ( Checkpoint, @@ -10,103 +17,199 @@ from langgraph.checkpoint.base import ( empty_checkpoint, ) from langgraph.checkpoint.postgres import PostgresSaver -from tests.conftest import DEFAULT_URI +from tests.conftest import DEFAULT_POSTGRES_URI -class TestPostgresSaver: - @pytest.fixture(autouse=True) - def setup(self) -> None: - # objects for test setup - self.config_1: RunnableConfig = { - "configurable": { - "thread_id": "thread-1", - # for backwards compatibility testing - "thread_ts": "1", - "checkpoint_ns": "", - } - } - self.config_2: RunnableConfig = { - "configurable": { - "thread_id": "thread-2", - "checkpoint_id": "2", - "checkpoint_ns": "", - } - } - self.config_3: RunnableConfig = { - "configurable": { - "thread_id": "thread-2", - "checkpoint_id": "2-inner", - "checkpoint_ns": "inner", - } - } +@contextmanager +def _pool_saver(): + """Fixture for pool mode testing.""" + 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, "row_factory": dict_row}, + ) 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}") - self.chkpnt_1: Checkpoint = empty_checkpoint() - self.chkpnt_2: Checkpoint = create_checkpoint(self.chkpnt_1, {}, 1) - self.chkpnt_3: Checkpoint = empty_checkpoint() - self.metadata_1: CheckpointMetadata = { - "source": "input", - "step": 2, - "writes": {}, - "score": 1, +@contextmanager +def _pipe_saver(): + """Fixture for pipeline mode testing.""" + 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: + with Connection.connect( + DEFAULT_POSTGRES_URI + database, + autocommit=True, + prepare_threshold=0, + row_factory=dict_row, + ) as conn: + with conn.pipeline() as pipe: + checkpointer = PostgresSaver(conn, pipe=pipe) + checkpointer.setup() + with conn.pipeline() as pipe: + checkpointer = PostgresSaver(conn, 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 _base_saver(): + """Fixture for regular connection mode testing.""" + 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: + with Connection.connect( + DEFAULT_POSTGRES_URI + database, + autocommit=True, + prepare_threshold=0, + row_factory=dict_row, + ) as conn: + checkpointer = PostgresSaver(conn) + 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 _saver(name: str): + if name == "base": + with _base_saver() as saver: + yield saver + elif name == "pool": + with _pool_saver() as saver: + yield saver + elif name == "pipe": + with _pipe_saver() as saver: + yield saver + + +@pytest.fixture +def test_data(): + """Fixture providing test data for checkpoint tests.""" + config_1: RunnableConfig = { + "configurable": { + "thread_id": "thread-1", + # for backwards compatibility testing + "thread_ts": "1", + "checkpoint_ns": "", } - self.metadata_2: CheckpointMetadata = { - "source": "loop", + } + config_2: RunnableConfig = { + "configurable": { + "thread_id": "thread-2", + "checkpoint_id": "2", + "checkpoint_ns": "", + } + } + config_3: RunnableConfig = { + "configurable": { + "thread_id": "thread-2", + "checkpoint_id": "2-inner", + "checkpoint_ns": "inner", + } + } + + chkpnt_1: Checkpoint = empty_checkpoint() + chkpnt_2: Checkpoint = create_checkpoint(chkpnt_1, {}, 1) + chkpnt_3: Checkpoint = empty_checkpoint() + + metadata_1: CheckpointMetadata = { + "source": "input", + "step": 2, + "writes": {}, + "score": 1, + } + metadata_2: CheckpointMetadata = { + "source": "loop", + "step": 1, + "writes": {"foo": "bar"}, + "score": None, + } + metadata_3: CheckpointMetadata = {} + + return { + "configs": [config_1, config_2, config_3], + "checkpoints": [chkpnt_1, chkpnt_2, chkpnt_3], + "metadata": [metadata_1, metadata_2, metadata_3], + } + + +@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"]) +def test_search(saver_name: str, test_data) -> None: + with _saver(saver_name) as saver: + configs = test_data["configs"] + checkpoints = test_data["checkpoints"] + metadata = test_data["metadata"] + + saver.put(configs[0], checkpoints[0], metadata[0], {}) + saver.put(configs[1], checkpoints[1], metadata[1], {}) + saver.put(configs[2], checkpoints[2], metadata[2], {}) + + # call method / assertions + query_1 = {"source": "input"} # search by 1 key + query_2 = { "step": 1, "writes": {"foo": "bar"}, - "score": None, - } - self.metadata_3: CheckpointMetadata = {} - with PostgresSaver.from_conn_string(DEFAULT_URI) as saver: - saver.setup() + } # search by multiple keys + query_3: dict[str, Any] = {} # search by no keys, return all checkpoints + query_4 = {"source": "update", "step": 1} # no match - def test_search(self) -> None: - with PostgresSaver.from_conn_string(DEFAULT_URI) as saver: - # save checkpoints - saver.put(self.config_1, self.chkpnt_1, self.metadata_1, {}) - saver.put(self.config_2, self.chkpnt_2, self.metadata_2, {}) - saver.put(self.config_3, self.chkpnt_3, self.metadata_3, {}) + search_results_1 = list(saver.list(None, filter=query_1)) + assert len(search_results_1) == 1 + assert search_results_1[0].metadata == metadata[0] - # call method / assertions - query_1 = {"source": "input"} # search by 1 key - query_2 = { - "step": 1, - "writes": {"foo": "bar"}, - } # search by multiple keys - query_3: dict[str, Any] = {} # search by no keys, return all checkpoints - query_4 = {"source": "update", "step": 1} # no match + search_results_2 = list(saver.list(None, filter=query_2)) + assert len(search_results_2) == 1 + assert search_results_2[0].metadata == metadata[1] - search_results_1 = list(saver.list(None, filter=query_1)) - assert len(search_results_1) == 1 - assert search_results_1[0].metadata == self.metadata_1 + search_results_3 = list(saver.list(None, filter=query_3)) + assert len(search_results_3) == 3 - search_results_2 = list(saver.list(None, filter=query_2)) - assert len(search_results_2) == 1 - assert search_results_2[0].metadata == self.metadata_2 + search_results_4 = list(saver.list(None, filter=query_4)) + assert len(search_results_4) == 0 - search_results_3 = list(saver.list(None, filter=query_3)) - assert len(search_results_3) == 3 + # search by config (defaults to checkpoints across all namespaces) + search_results_5 = list(saver.list({"configurable": {"thread_id": "thread-2"}})) + assert len(search_results_5) == 2 + assert { + search_results_5[0].config["configurable"]["checkpoint_ns"], + search_results_5[1].config["configurable"]["checkpoint_ns"], + } == {"", "inner"} - search_results_4 = list(saver.list(None, filter=query_4)) - assert len(search_results_4) == 0 - # search by config (defaults to checkpoints across all namespaces) - search_results_5 = list( - saver.list({"configurable": {"thread_id": "thread-2"}}) - ) - assert len(search_results_5) == 2 - assert { - search_results_5[0].config["configurable"]["checkpoint_ns"], - search_results_5[1].config["configurable"]["checkpoint_ns"], - } == {"", "inner"} - - # TODO: test before and limit params - - def test_null_chars(self) -> None: - with PostgresSaver.from_conn_string(DEFAULT_URI) as saver: - config = saver.put(self.config_1, self.chkpnt_1, {"my_key": "\x00abc"}, {}) - assert saver.get_tuple(config).metadata["my_key"] == "abc" # type: ignore - assert ( - list(saver.list(None, filter={"my_key": "abc"}))[0].metadata["my_key"] # type: ignore - == "abc" - ) +@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"]) +def test_null_chars(saver_name: str, test_data) -> None: + with _saver(saver_name) as saver: + config = saver.put( + test_data["configs"][0], + test_data["checkpoints"][0], + {"my_key": "\x00abc"}, + {}, + ) + assert saver.get_tuple(config).metadata["my_key"] == "abc" # type: ignore + assert ( + list(saver.list(None, filter={"my_key": "abc"}))[0].metadata["my_key"] + == "abc" + )