Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
111 changes: 66 additions & 45 deletions burr/integrations/persisters/b_psycopg2.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
except ImportError as e:
base.require_plugin(e, "postgresql")

import contextlib
import json
import logging
from typing import Literal, Optional
Expand Down Expand Up @@ -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 = ""
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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(
Expand All @@ -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)
Expand Down Expand Up @@ -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."""
Expand Down
96 changes: 96 additions & 0 deletions tests/integrations/persisters/test_postgresql.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
import os
import pickle

import psycopg2
import psycopg2.extensions
import pytest

from burr.core import state
Expand Down Expand Up @@ -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
Expand Down
Loading