diff --git a/pydal/connection.py b/pydal/connection.py index d79f2935c..6d9d1c5e2 100644 --- a/pydal/connection.py +++ b/pydal/connection.py @@ -113,7 +113,8 @@ def set_connection(self, connection: Any, run_hooks: bool = False) -> None: When ``connection`` is non-None: also issue a cursor; run the hooks if requested; run ``test_connection`` if - ``check_active_connection`` is True. + ``check_active_connection`` is True; then commit any connection + initialization work so the connection is returned in an idle state. """ setattr(THREAD_LOCAL, self._connection_uname_, connection) if connection: @@ -122,6 +123,16 @@ def set_connection(self, connection: Any, run_hooks: bool = False) -> None: self.after_connection_hook() if self.check_active_connection: self.test_connection() + if run_hooks or self.check_active_connection: + # DB-API drivers commonly disable autocommit. In that mode, + # connection hooks and even a read-only liveness query can + # start a transaction implicitly. This method runs only + # while acquiring a connection, before it is exposed to the + # caller, so this is the safe boundary at which to finalize + # that initialization work. Without it, PostgreSQL can + # report an otherwise unused connection as + # ``idle in transaction`` indefinitely. + connection.commit() else: setattr(THREAD_LOCAL, self._cursors_uname_, None) diff --git a/tests/__init__.py b/tests/__init__.py index 40de7d308..1f6550119 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -16,6 +16,7 @@ from .ast_statements import * from .ast_subselect import * from .ast_translate import * +from .connection import * from .cross_dialect import * from .driver_io import * from .tier2_units import * diff --git a/tests/connection.py b/tests/connection.py new file mode 100644 index 000000000..07cc82848 --- /dev/null +++ b/tests/connection.py @@ -0,0 +1,78 @@ +# -*- coding: utf-8 -*- + +"""Unit coverage for connection acquisition and initialization.""" + +from pydal.connection import ConnectionPool + +from ._compat import unittest + + +class RecordingCursor: + def close(self): + return None + + +class RecordingConnection: + def __init__(self, events): + self.events = events + self.in_transaction = False + + def cursor(self): + return RecordingCursor() + + def commit(self): + self.events.append("commit") + self.in_transaction = False + + +class RecordingConnectionPool(ConnectionPool): + def __init__(self, events): + super().__init__() + self.events = events + + def after_connection_hook(self): + self.events.append("hook") + self.connection.in_transaction = True + + def test_connection(self): + self.events.append("test") + self.connection.in_transaction = True + + +class TestConnectionInitialization(unittest.TestCase): + def test_initialization_transaction_is_committed_before_connection_is_returned( + self, + ): + events = [] + connection = RecordingConnection(events) + pool = RecordingConnectionPool(events) + + pool.set_connection(connection, run_hooks=True) + + self.assertEqual(events, ["hook", "test", "commit"]) + self.assertFalse(connection.in_transaction) + + def test_pooled_connection_check_is_committed_before_checkout(self): + events = [] + connection = RecordingConnection(events) + pool = RecordingConnectionPool(events) + + pool.set_connection(connection, run_hooks=False) + + self.assertEqual(events, ["test", "commit"]) + self.assertFalse(connection.in_transaction) + + def test_hooks_are_committed_when_connection_check_is_disabled(self): + events = [] + connection = RecordingConnection(events) + pool = RecordingConnectionPool(events) + pool.check_active_connection = False + + pool.set_connection(connection, run_hooks=True) + + self.assertEqual(events, ["hook", "commit"]) + self.assertFalse(connection.in_transaction) + + +if __name__ == "__main__": + unittest.main()