From 0b6b5db87a4232646e306cbc8a97ec7c8a7c2c4b Mon Sep 17 00:00:00 2001 From: Sangkyoon Nam Date: Mon, 5 Oct 2026 21:23:43 +0900 Subject: [PATCH] fix: end each psycopg2 persister call's transaction Co-Authored-By: Claude Fable 5.1 --- burr/integrations/persisters/b_psycopg2.py | 111 +++++++++++------- .../persisters/test_postgresql.py | 96 +++++++++++++++ 2 files changed, 162 insertions(+), 45 deletions(-) diff --git a/burr/integrations/persisters/b_psycopg2.py b/burr/integrations/persisters/b_psycopg2.py index 26425f805..e6f8ef8f2 100644 --- a/burr/integrations/persisters/b_psycopg2.py +++ b/burr/integrations/persisters/b_psycopg2.py @@ -22,6 +22,7 @@ except ImportError as e: base.require_plugin(e, "postgresql") +import contextlib import json import logging from typing import Literal, Optional @@ -51,7 +52,9 @@ class PostgreSQLPersister(persistence.BaseStatePersister): p = PostgreSQLPersister.from_values("postgres", "postgres", "my_password", "localhost", 54320, table_name="burr_state") - + ``is_initialized``, ``list_app_ids``, ``load`` and ``save`` each end their transaction before + returning (commit on success, rollback on failure). Give the persister a connection of its own + rather than one you keep your own transaction open on. """ PARTITION_KEY_DEFAULT = "" @@ -107,6 +110,25 @@ def set_serde_kwargs(self, serde_kwargs: dict): """Sets the serde_kwargs for the persister.""" self.serde_kwargs = serde_kwargs + @contextlib.contextmanager + def _transaction(self): + """Yields a cursor and ends the transaction when the block exits: commit on success, + rollback on an exception. psycopg2 opens a transaction on the first statement and keeps it + open, so without this a failed insert leaves the connection in an aborted transaction and + a read leaves it idle in transaction. This doesn't use ``with self.connection`` because + psycopg2 refuses to re-enter it from a caller's own ``with connection`` block. + """ + try: + yield self.connection.cursor() + except BaseException: + try: + self.connection.rollback() + except Exception: + # Usually the connection is already gone; the original error is the useful one. + logger.debug("Rollback after a failed call also failed", exc_info=True) + raise + self.connection.commit() + def create_table(self, table_name: str): """Helper function to create the table where things are stored.""" cursor = self.connection.cursor() @@ -142,24 +164,24 @@ def is_initialized(self) -> bool: """ if self._initialized: return True - cursor = self.connection.cursor() - cursor.execute( - "SELECT EXISTS (SELECT FROM information_schema.tables WHERE table_name = %s)", - (self.table_name,), - ) - self._initialized = cursor.fetchone()[0] + with self._transaction() as cursor: + cursor.execute( + "SELECT EXISTS (SELECT FROM information_schema.tables WHERE table_name = %s)", + (self.table_name,), + ) + self._initialized = cursor.fetchone()[0] return self._initialized def list_app_ids(self, partition_key: str, **kwargs) -> list[str]: """Lists the app_ids for a given partition_key.""" - cursor = self.connection.cursor() - cursor.execute( - f"SELECT DISTINCT app_id, created_at FROM {self.table_name} " - "WHERE partition_key = %s " - "ORDER BY created_at DESC", - (partition_key,), - ) - app_ids = [row[0] for row in cursor.fetchall()] + with self._transaction() as cursor: + cursor.execute( + f"SELECT DISTINCT app_id, created_at FROM {self.table_name} " + "WHERE partition_key = %s " + "ORDER BY created_at DESC", + (partition_key,), + ) + app_ids = [row[0] for row in cursor.fetchall()] return app_ids def load( @@ -178,29 +200,29 @@ def load( if partition_key is None: partition_key = self.PARTITION_KEY_DEFAULT logger.debug("Loading %s, %s, %s", partition_key, app_id, sequence_id) - cursor = self.connection.cursor() - if app_id is None: - # get latest for all app_ids - cursor.execute( - f"SELECT position, state, sequence_id, app_id, created_at, status FROM {self.table_name} " - f"WHERE partition_key = %s " - f"ORDER BY CREATED_AT DESC LIMIT 1", - (partition_key,), - ) - elif sequence_id is None: - cursor.execute( - f"SELECT position, state, sequence_id, app_id, created_at, status FROM {self.table_name} " - f"WHERE partition_key = %s AND app_id = %s " - f"ORDER BY sequence_id DESC LIMIT 1", - (partition_key, app_id), - ) - else: - cursor.execute( - f"SELECT position, state, sequence_id, app_id, created_at, status FROM {self.table_name} " - f"WHERE partition_key = %s AND app_id = %s AND sequence_id = %s ", - (partition_key, app_id, sequence_id), - ) - row = cursor.fetchone() + with self._transaction() as cursor: + if app_id is None: + # get latest for all app_ids + cursor.execute( + f"SELECT position, state, sequence_id, app_id, created_at, status FROM {self.table_name} " + f"WHERE partition_key = %s " + f"ORDER BY CREATED_AT DESC LIMIT 1", + (partition_key,), + ) + elif sequence_id is None: + cursor.execute( + f"SELECT position, state, sequence_id, app_id, created_at, status FROM {self.table_name} " + f"WHERE partition_key = %s AND app_id = %s " + f"ORDER BY sequence_id DESC LIMIT 1", + (partition_key, app_id), + ) + else: + cursor.execute( + f"SELECT position, state, sequence_id, app_id, created_at, status FROM {self.table_name} " + f"WHERE partition_key = %s AND app_id = %s AND sequence_id = %s ", + (partition_key, app_id, sequence_id), + ) + row = cursor.fetchone() if row is None: return None _state = state.State.deserialize(row[1], **self.serde_kwargs) @@ -250,14 +272,13 @@ def save( state, status, ) - cursor = self.connection.cursor() json_state = json.dumps(state.serialize(**self.serde_kwargs)) - cursor.execute( - f"INSERT INTO {self.table_name} (partition_key, app_id, sequence_id, position, state, status) " - "VALUES (%s, %s, %s, %s, %s, %s)", - (partition_key, app_id, sequence_id, position, json_state, status), - ) - self.connection.commit() + with self._transaction() as cursor: + cursor.execute( + f"INSERT INTO {self.table_name} (partition_key, app_id, sequence_id, position, state, status) " + "VALUES (%s, %s, %s, %s, %s, %s)", + (partition_key, app_id, sequence_id, position, json_state, status), + ) def cleanup(self): """Closes the connection to the database.""" diff --git a/tests/integrations/persisters/test_postgresql.py b/tests/integrations/persisters/test_postgresql.py index 0e8cd700a..52543acc5 100644 --- a/tests/integrations/persisters/test_postgresql.py +++ b/tests/integrations/persisters/test_postgresql.py @@ -18,6 +18,8 @@ import os import pickle +import psycopg2 +import psycopg2.extensions import pytest from burr.core import state @@ -57,6 +59,100 @@ def test_list_app_ids(postgresql_persister): assert "app_id2" in app_ids +def test_failed_save_does_not_poison_the_connection(postgresql_persister): + """A rejected insert is rolled back, so later calls on the same persister still work.""" + postgresql_persister.save("txn", "app", 1, "pos", state.State({"a": 1}), "completed") + with pytest.raises(psycopg2.errors.UniqueViolation): + postgresql_persister.save("txn", "app", 1, "pos", state.State({"a": 1}), "completed") + assert postgresql_persister.load("txn", "app", 1)["state"].get_all() == {"a": 1} + assert postgresql_persister.list_app_ids("txn") == ["app"] + postgresql_persister.save("txn", "app", 2, "pos", state.State({"a": 2}), "completed") + assert postgresql_persister.load("txn", "app")["sequence_id"] == 2 + + +def _is_idle(persister): + return ( + persister.connection.get_transaction_status() == psycopg2.extensions.TRANSACTION_STATUS_IDLE + ) + + +def test_reads_leave_no_transaction_open(postgresql_persister): + postgresql_persister.save("txn-read", "app", 1, "pos", state.State({"a": 1}), "completed") + postgresql_persister.load("txn-read", "app") + assert _is_idle(postgresql_persister) + postgresql_persister.load("txn-read", "missing") + assert _is_idle(postgresql_persister) + postgresql_persister.list_app_ids("txn-read") + assert _is_idle(postgresql_persister) + fresh = PostgreSQLPersister.from_values( + db_name="postgres", + user="postgres", + password="postgres", + host="localhost", + port=5432, + table_name="testtable", + ) + try: + assert fresh.is_initialized() + assert _is_idle(fresh) + finally: + fresh.cleanup() + + +def test_interrupted_call_rolls_back(postgresql_persister): + """A KeyboardInterrupt mid-statement must not leave the transaction open.""" + real = postgresql_persister.connection + + class InterruptingCursor: + def __init__(self): + self.cursor = real.cursor() + + def execute(self, *args, **kwargs): + self.cursor.execute(*args, **kwargs) + raise KeyboardInterrupt() + + class ConnectionProxy: + def cursor(self): + return InterruptingCursor() + + def __getattr__(self, name): + return getattr(real, name) + + interrupted = PostgreSQLPersister(ConnectionProxy(), table_name="testtable") + with pytest.raises(KeyboardInterrupt): + interrupted.load("txn-interrupt", "app") + assert real.get_transaction_status() == psycopg2.extensions.TRANSACTION_STATUS_IDLE + + +def test_persister_works_inside_a_callers_connection_block(postgresql_persister): + """Callers who wrap the persister in their own ``with connection`` block keep working.""" + with postgresql_persister.connection: + postgresql_persister.save("txn-nested", "app", 1, "pos", state.State({"a": 1}), "completed") + assert postgresql_persister.load("txn-nested", "app", 1)["state"].get_all() == {"a": 1} + + +def test_save_after_read_is_stamped_after_another_connections_save(postgresql_persister): + """created_at comes from the transaction start, so a read must not keep one open + across another connection's save.""" + other = PostgreSQLPersister.from_values( + db_name="postgres", + user="postgres", + password="postgres", + host="localhost", + port=5432, + table_name="testtable", + ) + try: + postgresql_persister.load("txn-order", "missing") + other.save("txn-order", "app-b", 1, "pos", state.State({}), "completed") + postgresql_persister.save("txn-order", "app-a", 1, "pos", state.State({}), "completed") + created_a = postgresql_persister.load("txn-order", "app-a", 1)["created_at"] + created_b = postgresql_persister.load("txn-order", "app-b", 1)["created_at"] + assert created_a >= created_b + finally: + other.cleanup() + + def test_load_nonexistent_key(postgresql_persister): state_data = postgresql_persister.load("pk", "nonexistent_key") assert state_data is None