From ceac66ceb06e19bea8794c0fcfdb6ade602889d0 Mon Sep 17 00:00:00 2001 From: Sakthivel Subramanian Date: Mon, 28 Sep 2026 10:24:45 +0000 Subject: [PATCH 1/2] feat(sqlalchemy-spanner): wire timeout execution option through to DBAPI Connection.timeout - Reset connection timeout to None on pool checkin in reset_connection. - Support timeout execution option in SpannerExecutionContext.pre_exec. - Use scoped save-and-restore in SpannerExecutionContext to restore connection timeout in post_exec and handle_dbapi_exception. - Add unit tests covering timeout lifecycle and exception restoration. --- .../sqlalchemy_spanner/sqlalchemy_spanner.py | 33 ++++ packages/sqlalchemy-spanner/noxfile.py | 3 +- .../sqlalchemy-spanner/tests/test_suite_20.py | 6 +- .../tests/unit/test_timeout.py | 141 ++++++++++++++++++ 4 files changed, 181 insertions(+), 2 deletions(-) create mode 100644 packages/sqlalchemy-spanner/tests/unit/test_timeout.py diff --git a/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py b/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py index ee8e72eb5665..d0834c137ecc 100644 --- a/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py +++ b/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py @@ -70,6 +70,8 @@ def reset_connection(dbapi_conn, connection_record, reset_state=None): dbapi_conn.staleness = None dbapi_conn.read_only = False + if hasattr(dbapi_conn, "timeout"): + dbapi_conn.timeout = None # register a method to get a single value of a JSON object @@ -200,7 +202,12 @@ def wrapper(self, connection, *args, **kwargs): return wrapper +_UNSET = object() + + class SpannerExecutionContext(DefaultExecutionContext): + _previous_timeout = _UNSET + def pre_exec(self): """ Apply execution options to the DB API connection before @@ -228,6 +235,12 @@ def pre_exec(self): if request_tag: self.cursor.request_tag = request_tag + if "timeout" in self.execution_options: + conn = getattr(self._dbapi_connection, "connection", self._dbapi_connection) + if conn is not None and hasattr(conn, "timeout"): + self._previous_timeout = conn.timeout + conn.timeout = self.execution_options["timeout"] + ignore_transaction_warnings = self.execution_options.get( "ignore_transaction_warnings" ) @@ -238,6 +251,26 @@ def pre_exec(self): ignore_transaction_warnings ) + def _restore_connection_timeout(self): + if self._previous_timeout is not _UNSET: + try: + conn = getattr( + self._dbapi_connection, "connection", self._dbapi_connection + ) + if conn is not None and hasattr(conn, "timeout"): + conn.timeout = self._previous_timeout + except Exception: + pass + self._previous_timeout = _UNSET + + def post_exec(self): + super(SpannerExecutionContext, self).post_exec() + self._restore_connection_timeout() + + def handle_dbapi_exception(self, e): + self._restore_connection_timeout() + super(SpannerExecutionContext, self).handle_dbapi_exception(e) + def fire_sequence(self, seq, type_): """Builds a statement for fetching next value of the sequence.""" return self._execute_scalar( diff --git a/packages/sqlalchemy-spanner/noxfile.py b/packages/sqlalchemy-spanner/noxfile.py index 56884d95296d..d083e8eddcde 100644 --- a/packages/sqlalchemy-spanner/noxfile.py +++ b/packages/sqlalchemy-spanner/noxfile.py @@ -120,10 +120,11 @@ class = StreamHandler SQLALCHEMY_14_DEPENDENCIES = [ "sqlalchemy>=1.4,<2.0", + "alembic<1.20", ] SQLALCHEMY_20_DEPENDENCIES = [ - "sqlalchemy>=2.0", + "sqlalchemy>=2.0,<2.1", ] UNIT_TEST_PYTHON_VERSIONS = ["3.10", "3.11", "3.12", "3.13", "3.14", "3.15"] diff --git a/packages/sqlalchemy-spanner/tests/test_suite_20.py b/packages/sqlalchemy-spanner/tests/test_suite_20.py index 6c975004a8c4..87e8c2ef120b 100644 --- a/packages/sqlalchemy-spanner/tests/test_suite_20.py +++ b/packages/sqlalchemy-spanner/tests/test_suite_20.py @@ -78,7 +78,11 @@ LongNameBlowoutTest as _LongNameBlowoutTest, ) from sqlalchemy.testing.suite.test_ddl import TableDDLTest as _TableDDLTest -from sqlalchemy.testing.suite.test_deprecations import * # noqa: F401, F403 + +try: + from sqlalchemy.testing.suite.test_deprecations import * # noqa: F401, F403 +except ImportError: + pass from sqlalchemy.testing.suite.test_dialect import * # noqa: F401, F403 from sqlalchemy.testing.suite.test_dialect import ( DifficultParametersTest as _DifficultParametersTest, diff --git a/packages/sqlalchemy-spanner/tests/unit/test_timeout.py b/packages/sqlalchemy-spanner/tests/unit/test_timeout.py new file mode 100644 index 000000000000..838fc9257a46 --- /dev/null +++ b/packages/sqlalchemy-spanner/tests/unit/test_timeout.py @@ -0,0 +1,141 @@ +# Copyright 2026 Google LLC All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from unittest import mock + +from sqlalchemy.testing import eq_ +from sqlalchemy.testing.plugin.plugin_base import fixtures + +from google.cloud import spanner_dbapi +from google.cloud.sqlalchemy_spanner.sqlalchemy_spanner import ( + _UNSET, + SpannerDialect, + SpannerExecutionContext, + reset_connection, +) + + +class SqlAlchemyTimeoutTest(fixtures.TestBase): + def test_reset_connection_clears_timeout(self): + dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) + dbapi_conn.timeout = 30.0 + dbapi_conn.inside_transaction = False + + reset_connection(dbapi_conn, None) + + eq_(dbapi_conn.timeout, None) + + def test_pre_exec_sets_timeout(self): + context = SpannerExecutionContext() + context.execution_options = {"timeout": 45.0} + + dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) + dbapi_conn.timeout = None + context._dbapi_connection = mock.MagicMock() + context._dbapi_connection.connection = dbapi_conn + + context.pre_exec() + + eq_(dbapi_conn.timeout, 45.0) + + def test_query_without_timeout_does_not_alter_connection_timeout(self): + dialect = SpannerDialect() + dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) + dbapi_conn.timeout = 25.0 + + context = SpannerExecutionContext() + context.dialect = dialect + context.execution_options = {} + context._dbapi_connection = mock.MagicMock() + context._dbapi_connection.connection = dbapi_conn + + context.pre_exec() + eq_(dbapi_conn.timeout, 25.0) + + context.post_exec() + eq_(dbapi_conn.timeout, 25.0) + + def test_statement_level_timeout_restores_previous_timeout_in_post_exec(self): + dialect = SpannerDialect() + dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) + dbapi_conn.timeout = 25.0 + + context = SpannerExecutionContext() + context.dialect = dialect + context.execution_options = {"timeout": 5.0} + context._dbapi_connection = mock.MagicMock() + context._dbapi_connection.connection = dbapi_conn + + context.pre_exec() + eq_(dbapi_conn.timeout, 5.0) + + context.post_exec() + eq_(dbapi_conn.timeout, 25.0) + + def test_statement_level_timeout_restores_previous_timeout_on_dbapi_exception(self): + dialect = SpannerDialect() + dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) + dbapi_conn.timeout = 25.0 + + context = SpannerExecutionContext() + context.dialect = dialect + context.execution_options = {"timeout": 5.0} + context._dbapi_connection = mock.MagicMock() + context._dbapi_connection.connection = dbapi_conn + + context.pre_exec() + eq_(dbapi_conn.timeout, 5.0) + + context.handle_dbapi_exception(Exception("Statement timeout/error")) + eq_(dbapi_conn.timeout, 25.0) + + def test_statement_level_timeout_none_overrides_and_restores(self): + dialect = SpannerDialect() + dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) + dbapi_conn.timeout = 25.0 + + context = SpannerExecutionContext() + context.dialect = dialect + context.execution_options = {"timeout": None} + context._dbapi_connection = mock.MagicMock() + context._dbapi_connection.connection = dbapi_conn + + context.pre_exec() + eq_(dbapi_conn.timeout, None) + + context.post_exec() + eq_(dbapi_conn.timeout, 25.0) + + def test_statement_level_timeout_swallows_exception_on_broken_conn(self): + dialect = SpannerDialect() + dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) + dbapi_conn.timeout = 25.0 + + context = SpannerExecutionContext() + context.dialect = dialect + context.execution_options = {"timeout": 5.0} + context._dbapi_connection = mock.MagicMock() + context._dbapi_connection.connection = dbapi_conn + + context.pre_exec() + eq_(dbapi_conn.timeout, 5.0) + + # Simulate broken connection raising on timeout assignment + type(dbapi_conn).timeout = mock.PropertyMock( + side_effect=Exception("Connection broken") + ) + + # Should not raise exception + context.handle_dbapi_exception(Exception("Original DBAPI error")) + eq_(context._previous_timeout, _UNSET) From 3d4a8953f894be547026c311c2264d852893154c Mon Sep 17 00:00:00 2001 From: Sakthivel Subramanian Date: Sun, 4 Oct 2026 18:59:54 +0000 Subject: [PATCH 2/2] Addressed comments --- .../sqlalchemy_spanner/sqlalchemy_spanner.py | 35 +------ .../tests/unit/test_timeout.py | 93 +++---------------- 2 files changed, 14 insertions(+), 114 deletions(-) diff --git a/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py b/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py index d0834c137ecc..43433690d41f 100644 --- a/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py +++ b/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py @@ -70,8 +70,6 @@ def reset_connection(dbapi_conn, connection_record, reset_state=None): dbapi_conn.staleness = None dbapi_conn.read_only = False - if hasattr(dbapi_conn, "timeout"): - dbapi_conn.timeout = None # register a method to get a single value of a JSON object @@ -202,12 +200,7 @@ def wrapper(self, connection, *args, **kwargs): return wrapper -_UNSET = object() - - class SpannerExecutionContext(DefaultExecutionContext): - _previous_timeout = _UNSET - def pre_exec(self): """ Apply execution options to the DB API connection before @@ -235,11 +228,9 @@ def pre_exec(self): if request_tag: self.cursor.request_tag = request_tag - if "timeout" in self.execution_options: - conn = getattr(self._dbapi_connection, "connection", self._dbapi_connection) - if conn is not None and hasattr(conn, "timeout"): - self._previous_timeout = conn.timeout - conn.timeout = self.execution_options["timeout"] + timeout = self.execution_options.get("timeout") + if timeout is not None: + self.cursor.timeout = timeout ignore_transaction_warnings = self.execution_options.get( "ignore_transaction_warnings" @@ -251,26 +242,6 @@ def pre_exec(self): ignore_transaction_warnings ) - def _restore_connection_timeout(self): - if self._previous_timeout is not _UNSET: - try: - conn = getattr( - self._dbapi_connection, "connection", self._dbapi_connection - ) - if conn is not None and hasattr(conn, "timeout"): - conn.timeout = self._previous_timeout - except Exception: - pass - self._previous_timeout = _UNSET - - def post_exec(self): - super(SpannerExecutionContext, self).post_exec() - self._restore_connection_timeout() - - def handle_dbapi_exception(self, e): - self._restore_connection_timeout() - super(SpannerExecutionContext, self).handle_dbapi_exception(e) - def fire_sequence(self, seq, type_): """Builds a statement for fetching next value of the sequence.""" return self._execute_scalar( diff --git a/packages/sqlalchemy-spanner/tests/unit/test_timeout.py b/packages/sqlalchemy-spanner/tests/unit/test_timeout.py index 838fc9257a46..984e892b3ae6 100644 --- a/packages/sqlalchemy-spanner/tests/unit/test_timeout.py +++ b/packages/sqlalchemy-spanner/tests/unit/test_timeout.py @@ -19,7 +19,6 @@ from google.cloud import spanner_dbapi from google.cloud.sqlalchemy_spanner.sqlalchemy_spanner import ( - _UNSET, SpannerDialect, SpannerExecutionContext, reset_connection, @@ -27,115 +26,45 @@ class SqlAlchemyTimeoutTest(fixtures.TestBase): - def test_reset_connection_clears_timeout(self): + def test_reset_connection_preserves_default_timeout(self): dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) dbapi_conn.timeout = 30.0 dbapi_conn.inside_transaction = False reset_connection(dbapi_conn, None) - eq_(dbapi_conn.timeout, None) + eq_(dbapi_conn.timeout, 30.0) - def test_pre_exec_sets_timeout(self): + def test_pre_exec_sets_cursor_timeout(self): context = SpannerExecutionContext() context.execution_options = {"timeout": 45.0} + context.cursor = mock.MagicMock() + context.cursor.timeout = None - dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) - dbapi_conn.timeout = None - context._dbapi_connection = mock.MagicMock() - context._dbapi_connection.connection = dbapi_conn - - context.pre_exec() - - eq_(dbapi_conn.timeout, 45.0) - - def test_query_without_timeout_does_not_alter_connection_timeout(self): - dialect = SpannerDialect() dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) dbapi_conn.timeout = 25.0 - - context = SpannerExecutionContext() - context.dialect = dialect - context.execution_options = {} context._dbapi_connection = mock.MagicMock() context._dbapi_connection.connection = dbapi_conn context.pre_exec() - eq_(dbapi_conn.timeout, 25.0) - context.post_exec() + eq_(context.cursor.timeout, 45.0) eq_(dbapi_conn.timeout, 25.0) - def test_statement_level_timeout_restores_previous_timeout_in_post_exec(self): + def test_query_without_timeout_does_not_alter_cursor_or_connection_timeout(self): dialect = SpannerDialect() dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) dbapi_conn.timeout = 25.0 context = SpannerExecutionContext() context.dialect = dialect - context.execution_options = {"timeout": 5.0} - context._dbapi_connection = mock.MagicMock() - context._dbapi_connection.connection = dbapi_conn - - context.pre_exec() - eq_(dbapi_conn.timeout, 5.0) - - context.post_exec() - eq_(dbapi_conn.timeout, 25.0) - - def test_statement_level_timeout_restores_previous_timeout_on_dbapi_exception(self): - dialect = SpannerDialect() - dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) - dbapi_conn.timeout = 25.0 - - context = SpannerExecutionContext() - context.dialect = dialect - context.execution_options = {"timeout": 5.0} - context._dbapi_connection = mock.MagicMock() - context._dbapi_connection.connection = dbapi_conn - - context.pre_exec() - eq_(dbapi_conn.timeout, 5.0) - - context.handle_dbapi_exception(Exception("Statement timeout/error")) - eq_(dbapi_conn.timeout, 25.0) - - def test_statement_level_timeout_none_overrides_and_restores(self): - dialect = SpannerDialect() - dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) - dbapi_conn.timeout = 25.0 - - context = SpannerExecutionContext() - context.dialect = dialect - context.execution_options = {"timeout": None} + context.execution_options = {} + context.cursor = mock.MagicMock() + context.cursor.timeout = None context._dbapi_connection = mock.MagicMock() context._dbapi_connection.connection = dbapi_conn context.pre_exec() - eq_(dbapi_conn.timeout, None) - context.post_exec() + eq_(context.cursor.timeout, None) eq_(dbapi_conn.timeout, 25.0) - - def test_statement_level_timeout_swallows_exception_on_broken_conn(self): - dialect = SpannerDialect() - dbapi_conn = mock.MagicMock(spec=spanner_dbapi.Connection) - dbapi_conn.timeout = 25.0 - - context = SpannerExecutionContext() - context.dialect = dialect - context.execution_options = {"timeout": 5.0} - context._dbapi_connection = mock.MagicMock() - context._dbapi_connection.connection = dbapi_conn - - context.pre_exec() - eq_(dbapi_conn.timeout, 5.0) - - # Simulate broken connection raising on timeout assignment - type(dbapi_conn).timeout = mock.PropertyMock( - side_effect=Exception("Connection broken") - ) - - # Should not raise exception - context.handle_dbapi_exception(Exception("Original DBAPI error")) - eq_(context._previous_timeout, _UNSET)