This commit is contained in:
vbarda
2024-11-18 11:12:21 -05:00
parent 5abbb79e1b
commit f0505155a2
3 changed files with 27 additions and 4 deletions
@@ -363,6 +363,16 @@ class PostgresSaver(BasePostgresSaver):
),
)
def _check_pipeline_support(self, conn: Connection[DictRow]) -> None:
if self.supports_pipeline is not None:
return
try:
with conn.pipeline():
self.supports_pipeline = True
except NotSupportedError:
self.supports_pipeline = False
@contextmanager
def _cursor(self, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]:
"""Create a database cursor as a context manager.
@@ -384,14 +394,15 @@ class PostgresSaver(BasePostgresSaver):
if pipeline:
self.pipe.sync()
elif pipeline:
self._check_pipeline_support(conn)
# a connection not in pipeline mode can only be used by one
# thread/coroutine at a time, so we acquire a lock
try:
if self.supports_pipeline:
with self.lock, conn.pipeline(), conn.cursor(
binary=True, row_factory=dict_row
) as cur:
yield cur
except NotSupportedError:
else:
# Use connection's transaction context manager when pipeline mode not supported
with self.lock, conn.transaction(), conn.cursor(
binary=True, row_factory=dict_row
@@ -319,6 +319,16 @@ class AsyncPostgresSaver(BasePostgresSaver):
async with self._cursor(pipeline=True) as cur:
await cur.executemany(query, params)
async def _check_pipeline_support(self, conn: AsyncConnection[DictRow]) -> None:
if self.supports_pipeline is not None:
return
try:
async with conn.pipeline():
self.supports_pipeline = True
except NotSupportedError:
self.supports_pipeline = False
@asynccontextmanager
async def _cursor(
self, *, pipeline: bool = False
@@ -342,14 +352,15 @@ class AsyncPostgresSaver(BasePostgresSaver):
if pipeline:
await self.pipe.sync()
elif pipeline:
await self._check_pipeline_support(conn)
# a connection not in pipeline mode can only be used by one
# thread/coroutine at a time, so we acquire a lock
try:
if self.supports_pipeline:
async with self.lock, conn.pipeline(), conn.cursor(
binary=True, row_factory=dict_row
) as cur:
yield cur
except NotSupportedError:
else:
# 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
@@ -133,6 +133,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
INSERT_CHECKPOINT_WRITES_SQL = INSERT_CHECKPOINT_WRITES_SQL
jsonplus_serde = JsonPlusSerializer()
supports_pipeline: Optional[bool] = None
def _load_checkpoint(
self,