From 0a5220aa0715a0756d748e68c991c5167ef239ca Mon Sep 17 00:00:00 2001 From: vbarda Date: Thu, 14 Nov 2024 18:55:02 -0500 Subject: [PATCH] code review --- .../langgraph/checkpoint/postgres/__init__.py | 76 +++++++++++------- .../langgraph/checkpoint/postgres/aio.py | 79 ++++++++++++------- 2 files changed, 99 insertions(+), 56 deletions(-) diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index 7085e107a..260d12863 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -1,10 +1,10 @@ import threading -from contextlib import contextmanager, nullcontext +from contextlib import contextmanager from typing import Any, Iterator, Optional, Sequence, Union from langchain_core.runnables import RunnableConfig from psycopg import Connection, Cursor, Pipeline -from psycopg.errors import UndefinedTable +from psycopg.errors import NotSupportedError, UndefinedTable from psycopg.rows import DictRow, dict_row from psycopg.types.json import Jsonb from psycopg_pool import ConnectionPool @@ -308,29 +308,27 @@ class PostgresSaver(BasePostgresSaver): } } - with self._cursor() as cur: - # Use connection's transaction context manager when not in pipeline mode - with cur.connection.transaction() if self.pipe is None else nullcontext(): - cur.executemany( - self.UPSERT_CHECKPOINT_BLOBS_SQL, - self._dump_blobs( - thread_id, - checkpoint_ns, - copy.pop("channel_values"), # type: ignore[misc] - new_versions, - ), - ) - cur.execute( - self.UPSERT_CHECKPOINTS_SQL, - ( - thread_id, - checkpoint_ns, - checkpoint["id"], - checkpoint_id, - Jsonb(self._dump_checkpoint(copy)), - self._dump_metadata(metadata), - ), - ) + with self._cursor(pipeline=True) as cur: + cur.executemany( + self.UPSERT_CHECKPOINT_BLOBS_SQL, + self._dump_blobs( + thread_id, + checkpoint_ns, + copy.pop("channel_values"), # type: ignore[misc] + new_versions, + ), + ) + cur.execute( + self.UPSERT_CHECKPOINTS_SQL, + ( + thread_id, + checkpoint_ns, + checkpoint["id"], + checkpoint_id, + Jsonb(self._dump_checkpoint(copy)), + self._dump_metadata(metadata), + ), + ) return next_config def put_writes( @@ -353,7 +351,7 @@ class PostgresSaver(BasePostgresSaver): if all(w[0] in WRITES_IDX_MAP for w in writes) else self.INSERT_CHECKPOINT_WRITES_SQL ) - with self._cursor() as cur: + with self._cursor(pipeline=True) as cur: cur.executemany( query, self._dump_writes( @@ -366,7 +364,14 @@ class PostgresSaver(BasePostgresSaver): ) @contextmanager - def _cursor(self) -> Iterator[Cursor[DictRow]]: + def _cursor(self, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]: + """Create a database cursor as a context manager. + + Args: + pipeline (bool): whether to use pipeline for the DB operations inside the context manager. + Will be applied regardless of whether the PostgresSaver instance was initialized with a pipeline. + If pipeline mode is not supported, will fall back to using transaction context manager. + """ with _get_connection(self.conn) as conn: if self.pipe: # a connection in pipeline mode can be used concurrently @@ -376,7 +381,22 @@ class PostgresSaver(BasePostgresSaver): with conn.cursor(binary=True, row_factory=dict_row) as cur: yield cur finally: - self.pipe.sync() + if pipeline: + self.pipe.sync() + elif pipeline: + # a connection not in pipeline mode can only be used by one + # thread/coroutine at a time, so we acquire a lock + try: + with self.lock, conn.pipeline(), conn.cursor( + binary=True, row_factory=dict_row + ) as cur: + yield cur + except NotSupportedError: + # Use connection's transaction context manager when pipeline mode not supported + with self.lock, conn.transaction(), conn.cursor( + binary=True, row_factory=dict_row + ) as cur: + yield cur else: with self.lock, conn.cursor(binary=True, row_factory=dict_row) as cur: yield cur diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index eac629605..43429ae4d 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -1,10 +1,10 @@ import asyncio -from contextlib import asynccontextmanager, nullcontext +from contextlib import asynccontextmanager from typing import Any, AsyncIterator, Iterator, Optional, Sequence, Union from langchain_core.runnables import RunnableConfig from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline -from psycopg.errors import UndefinedTable +from psycopg.errors import NotSupportedError, UndefinedTable from psycopg.rows import DictRow, dict_row from psycopg.types.json import Jsonb from psycopg_pool import AsyncConnectionPool @@ -264,29 +264,28 @@ class AsyncPostgresSaver(BasePostgresSaver): } } - async with self._cursor() as cur: - async with cur.connection.transaction() if self.pipe is None else nullcontext(): - await cur.executemany( - self.UPSERT_CHECKPOINT_BLOBS_SQL, - await asyncio.to_thread( - self._dump_blobs, - thread_id, - checkpoint_ns, - copy.pop("channel_values"), # type: ignore[misc] - new_versions, - ), - ) - await cur.execute( - self.UPSERT_CHECKPOINTS_SQL, - ( - thread_id, - checkpoint_ns, - checkpoint["id"], - checkpoint_id, - Jsonb(self._dump_checkpoint(copy)), - self._dump_metadata(metadata), - ), - ) + async with self._cursor(pipeline=True) as cur: + await cur.executemany( + self.UPSERT_CHECKPOINT_BLOBS_SQL, + await asyncio.to_thread( + self._dump_blobs, + thread_id, + checkpoint_ns, + copy.pop("channel_values"), # type: ignore[misc] + new_versions, + ), + ) + await cur.execute( + self.UPSERT_CHECKPOINTS_SQL, + ( + thread_id, + checkpoint_ns, + checkpoint["id"], + checkpoint_id, + Jsonb(self._dump_checkpoint(copy)), + self._dump_metadata(metadata), + ), + ) return next_config async def aput_writes( @@ -317,11 +316,20 @@ class AsyncPostgresSaver(BasePostgresSaver): task_id, writes, ) - async with self._cursor() as cur: + async with self._cursor(pipeline=True) as cur: await cur.executemany(query, params) @asynccontextmanager - async def _cursor(self) -> AsyncIterator[AsyncCursor[DictRow]]: + async def _cursor( + self, *, pipeline: bool = False + ) -> AsyncIterator[AsyncCursor[DictRow]]: + """Create a database cursor as a context manager. + + Args: + pipeline (bool): whether to use pipeline for the DB operations inside the context manager. + Will be applied regardless of whether the AsyncPostgresSaver instance was initialized with a pipeline. + If pipeline mode is not supported, will fall back to using transaction context manager. + """ async with _get_connection(self.conn) as conn: if self.pipe: # a connection in pipeline mode can be used concurrently @@ -331,7 +339,22 @@ class AsyncPostgresSaver(BasePostgresSaver): async with conn.cursor(binary=True, row_factory=dict_row) as cur: yield cur finally: - await self.pipe.sync() + if pipeline: + await self.pipe.sync() + elif pipeline: + # a connection not in pipeline mode can only be used by one + # thread/coroutine at a time, so we acquire a lock + try: + async with self.lock, conn.pipeline(), conn.cursor( + binary=True, row_factory=dict_row + ) as cur: + yield cur + except NotSupportedError: + # Use connection's transaction context manager when pipeline mode not supported + async with self.lock, conn.transaction(), conn.cursor( + binary=True, row_factory=dict_row + ) as cur: + yield cur else: async with self.lock, conn.cursor( binary=True, row_factory=dict_row