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
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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"
)
Expand All @@ -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
Comment thread
sakthivelmanii marked this conversation as resolved.

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(
Expand Down
3 changes: 2 additions & 1 deletion packages/sqlalchemy-spanner/noxfile.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down
6 changes: 5 additions & 1 deletion packages/sqlalchemy-spanner/tests/test_suite_20.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
141 changes: 141 additions & 0 deletions packages/sqlalchemy-spanner/tests/unit/test_timeout.py
Original file line number Diff line number Diff line change
@@ -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)
Loading