From ceac66ceb06e19bea8794c0fcfdb6ade602889d0 Mon Sep 17 00:00:00 2001 From: Sakthivel Subramanian Date: Mon, 28 Sep 2026 10:24:45 +0000 Subject: [PATCH] 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)