diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_session.py b/packages/google-cloud-spanner/tests/unit/_async/test_session.py index 98758b6b904b..b20dda7e2435 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_session.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_session.py @@ -1219,9 +1219,13 @@ async def unit_of_work(transaction): pass await session.create() - await session.run_in_transaction(unit_of_work) + with mock.patch( + "google.cloud.spanner_v1._async._helpers.asyncio.sleep" + ) as sleep_mock: + await session.run_in_transaction(unit_of_work) self.assertEqual(begin_transaction.call_count, 2) + sleep_mock.assert_called_once() begin_transaction.assert_called_with( request=BeginTransactionRequest( @@ -1261,10 +1265,14 @@ async def unit_of_work(transaction): pass await session.create() - await session.run_in_transaction(unit_of_work) + with mock.patch( + "google.cloud.spanner_v1._async._helpers.asyncio.sleep" + ) as sleep_mock: + await session.run_in_transaction(unit_of_work) # Verify retried BeginTransaction API call. self.assertEqual(begin_transaction.call_count, 2) + sleep_mock.assert_called_once() begin_transaction.assert_called_with( request=BeginTransactionRequest( @@ -1308,10 +1316,14 @@ async def unit_of_work(transaction): pass await session.create() - await session.run_in_transaction(unit_of_work) + with mock.patch( + "google.cloud.spanner_v1._async._helpers.asyncio.sleep" + ) as sleep_mock: + await session.run_in_transaction(unit_of_work) # Verify retried BeginTransaction API call. self.assertEqual(begin_transaction.call_count, 2) + sleep_mock.assert_called_once() begin_transaction.assert_called_with( request=BeginTransactionRequest( @@ -1812,7 +1824,7 @@ def _time(): return _results[0] return 1.0 - with mock.patch("time.time", _time): + with mock.patch("google.cloud.spanner_v1._async._helpers.time.time", _time): with mock.patch( "google.cloud.spanner_v1._async._helpers.asyncio.sleep", new_callable=mock.AsyncMock, @@ -1893,7 +1905,7 @@ def _time(): return 1.0 with ( - mock.patch("time.time", _time), + mock.patch("google.cloud.spanner_v1._async._helpers.time.time", _time), mock.patch( "google.cloud.spanner_v1._helpers.random.random", return_value=0 ), @@ -2817,11 +2829,15 @@ def _time_func(): return 3 # check if current time > deadline - with mock.patch("time.time", _time_func): + with mock.patch( + "google.cloud.spanner_v1._async._helpers.time.time", _time_func + ): with pytest.raises(Exception): _delay_until_retry(exc_mock, 2, 1, default_retry_delay=0) - with mock.patch("time.time", _time_func): + with mock.patch( + "google.cloud.spanner_v1._async._helpers.time.time", _time_func + ): with mock.patch( "google.cloud.spanner_v1._helpers._get_retry_delay" ) as get_retry_delay_mock: diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py b/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py index da90f929f615..39926148d9de 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py @@ -1367,9 +1367,14 @@ async def test_begin_transaction_retry(self): TransactionPB(id=TXN_ID), ] - tid = await snapshot._begin_transaction() + with mock.patch( + "google.cloud.spanner_v1._async._helpers.asyncio.sleep" + ) as sleep_mock: + tid = await snapshot._begin_transaction() + self.assertEqual(tid, TXN_ID) self.assertEqual(api.begin_transaction.call_count, 2) + sleep_mock.assert_called_once_with(2) async def test_update_for_transaction_pb_w_precommit_token(self): from google.cloud.spanner_v1.types import MultiplexedSessionPrecommitToken diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_transaction_extra.py b/packages/google-cloud-spanner/tests/unit/_async/test_transaction_extra.py index 4e21378cb6b2..69ce9c4d2a68 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_transaction_extra.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_transaction_extra.py @@ -135,8 +135,13 @@ async def test_commit_retry_and_precommit_token(self): final_resp, ] - await txn.commit() + with mock.patch( + "google.cloud.spanner_v1._async._helpers.asyncio.sleep" + ) as sleep_mock: + await txn.commit() + self.assertEqual(self.db.spanner_api.commit.call_count, 3) + sleep_mock.assert_called_once_with(2) async def test_execute_update_request_options_dict(self): # coverage for line 503 diff --git a/packages/google-cloud-spanner/tests/unit/omni/test_login_client.py b/packages/google-cloud-spanner/tests/unit/omni/test_login_client.py index 3d080cd67859..fe33bfea5efc 100644 --- a/packages/google-cloud-spanner/tests/unit/omni/test_login_client.py +++ b/packages/google-cloud-spanner/tests/unit/omni/test_login_client.py @@ -41,9 +41,9 @@ def test_login_successful_flow(self): params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) @@ -248,9 +248,9 @@ def test_login_missing_access_token_in_final_response(self): password = "test_password" params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) diff --git a/packages/google-cloud-spanner/tests/unit/omni/test_opaque.py b/packages/google-cloud-spanner/tests/unit/omni/test_opaque.py index e2816ff0c22a..17e93eb90721 100644 --- a/packages/google-cloud-spanner/tests/unit/omni/test_opaque.py +++ b/packages/google-cloud-spanner/tests/unit/omni/test_opaque.py @@ -486,9 +486,9 @@ def test_oprf_evaluate(self): def test_authenticator_validation(self): valid_params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) @@ -512,9 +512,9 @@ def test_authenticator_validation(self): def test_authenticator_state_errors(self): params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) @@ -541,9 +541,9 @@ def test_authenticator_state_errors(self): def test_user_authenticator_clear_and_del(self): params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) @@ -584,9 +584,9 @@ def test_full_opaque_handshake_simulation(self): params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) @@ -708,9 +708,9 @@ def test_full_opaque_handshake_simulation(self): def test_final_request_invalid_masked_response_length(self): params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) @@ -964,9 +964,9 @@ def test_final_request_edge_cases(self): params = authentication_pb2.HashParameters( argon2_id_parameters=authentication_pb2.HashParameters.Argon2IdParameters( - iteration_count=3, - memory_usage=64 * 1024, - parallelism=4, + iteration_count=1, + memory_usage=8, + parallelism=1, hash_size=32, ) ) diff --git a/packages/google-cloud-spanner/tests/unit/test_batch.py b/packages/google-cloud-spanner/tests/unit/test_batch.py index 933c47bfa92d..247bdb58ab6c 100644 --- a/packages/google-cloud-spanner/tests/unit/test_batch.py +++ b/packages/google-cloud-spanner/tests/unit/test_batch.py @@ -345,12 +345,25 @@ def test_aborted_exception_on_commit_with_retries(self, mock_region): batch.insert(TABLE_NAME, COLUMNS, VALUES) # Assertion: Ensure that calling batch.commit() raises Aborted - with self.assertRaises(Aborted) as context: - batch.commit(timeout_secs=1.0, default_retry_delay=0) + delay_call_count = 0 + + def fake_delay(exc, *args, **kwargs): + nonlocal delay_call_count + delay_call_count += 1 + if delay_call_count >= 2: + raise exc + + with mock.patch( + "google.cloud.spanner_v1._helpers._delay_until_retry", + side_effect=fake_delay, + ): + with self.assertRaises(Aborted) as context: + batch.commit(timeout_secs=1.0, default_retry_delay=0) # Verify exception includes request_id attribute self.assertIn("409 Transaction was aborted", str(context.exception)) self.assertTrue(hasattr(context.exception, "request_id")) + self.assertEqual(delay_call_count, 2) self.assertGreater( api.commit.call_count, 1, "commit should be called more than once" ) diff --git a/packages/google-cloud-spanner/tests/unit/test_client.py b/packages/google-cloud-spanner/tests/unit/test_client.py index ac279e031230..241b4efa02c3 100644 --- a/packages/google-cloud-spanner/tests/unit/test_client.py +++ b/packages/google-cloud-spanner/tests/unit/test_client.py @@ -59,6 +59,17 @@ def _get_target_class(self): def _make_one(self, *args, **kwargs): return self._get_target_class()(*args, **kwargs) + @staticmethod + def _make_instance_admin_api(): + from google.cloud.spanner_admin_instance_v1 import InstanceAdminClient + from google.cloud.spanner_admin_instance_v1.services.instance_admin.transports.base import ( + InstanceAdminTransport, + ) + + mock_transport = mock.create_autospec(InstanceAdminTransport, instance=True) + mock_transport._wrapped_methods = {} + return InstanceAdminClient(transport=mock_transport) + def _constructor_test_helper( self, expected_scopes, @@ -627,15 +638,14 @@ def test_project_name_property(self): def test_list_instance_configs(self): from google.cloud.spanner_admin_instance_v1 import ( - InstanceAdminClient, - ListInstanceConfigsRequest, - ListInstanceConfigsResponse, + InstanceConfig as InstanceConfigPB, ) from google.cloud.spanner_admin_instance_v1 import ( - InstanceConfig as InstanceConfigPB, + ListInstanceConfigsRequest, + ListInstanceConfigsResponse, ) - api = InstanceAdminClient(credentials=AnonymousCredentials()) + api = self._make_instance_admin_api() credentials = build_scoped_credentials() client = self._make_one(project=self.PROJECT, credentials=credentials) client._instance_admin_api = api @@ -676,16 +686,15 @@ def test_list_instance_configs(self): def test_list_instance_configs_w_options(self): from google.cloud.spanner_admin_instance_v1 import ( - InstanceAdminClient, - ListInstanceConfigsRequest, - ListInstanceConfigsResponse, + InstanceConfig as InstanceConfigPB, ) from google.cloud.spanner_admin_instance_v1 import ( - InstanceConfig as InstanceConfigPB, + ListInstanceConfigsRequest, + ListInstanceConfigsResponse, ) credentials = build_scoped_credentials() - api = InstanceAdminClient(credentials=credentials) + api = self._make_instance_admin_api() client = self._make_one(project=self.PROJECT, credentials=credentials) client._instance_admin_api = api @@ -754,15 +763,16 @@ def test_instance_factory_explicit(self): self.assertIs(instance._client, client) def test_list_instances(self): - from google.cloud.spanner_admin_instance_v1 import Instance as InstancePB from google.cloud.spanner_admin_instance_v1 import ( - InstanceAdminClient, + Instance as InstancePB, + ) + from google.cloud.spanner_admin_instance_v1 import ( ListInstancesRequest, ListInstancesResponse, ) + api = self._make_instance_admin_api() credentials = build_scoped_credentials() - api = InstanceAdminClient(credentials=credentials) client = self._make_one(project=self.PROJECT, credentials=credentials) client._instance_admin_api = api @@ -806,13 +816,12 @@ def test_list_instances(self): def test_list_instances_w_options(self): from google.cloud.spanner_admin_instance_v1 import ( - InstanceAdminClient, ListInstancesRequest, ListInstancesResponse, ) + api = self._make_instance_admin_api() credentials = build_scoped_credentials() - api = InstanceAdminClient(credentials=credentials) client = self._make_one(project=self.PROJECT, credentials=credentials) client._instance_admin_api = api diff --git a/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py b/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py index 47fa96e5a468..f7b1fd50b88a 100644 --- a/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py +++ b/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py @@ -11,14 +11,13 @@ # 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. +import threading from datetime import timedelta from os import environ -from time import sleep, time -from typing import Callable from unittest import TestCase from google.api_core.exceptions import BadRequest, FailedPrecondition -from mock import MagicMock, Mock, patch +from mock import DEFAULT, MagicMock, Mock, patch from google.cloud.spanner_v1.database_sessions_manager import ( DatabaseSessionsManager, @@ -27,12 +26,6 @@ from tests._builders import build_database -# Shorten polling and refresh intervals for testing. -@patch.multiple( - DatabaseSessionsManager, - _MAINTENANCE_THREAD_POLLING_INTERVAL=timedelta(seconds=1), - _MAINTENANCE_THREAD_REFRESH_INTERVAL=timedelta(seconds=2), -) class TestDatabaseSessionManager(TestCase): @classmethod def setUpClass(cls): @@ -64,7 +57,8 @@ def tearDown(self): if thread and thread.is_alive(): manager._multiplexed_session_terminate_event.set() - self._assert_true_with_timeout(lambda: not thread.is_alive()) + thread.join(timeout=10) + self.assertFalse(thread.is_alive()) def test_read_only_pooled(self): manager = self._manager @@ -180,19 +174,31 @@ def test_read_write_multiplexed(self): # Verify create_session was called. manager._database.spanner_api.create_session.assert_called_once() + @patch.multiple( + DatabaseSessionsManager, + _MAINTENANCE_THREAD_POLLING_INTERVAL=timedelta(milliseconds=5), + _MAINTENANCE_THREAD_REFRESH_INTERVAL=timedelta(milliseconds=10), + ) def test_multiplexed_maintenance(self): manager = self._manager self._enable_multiplexed_sessions() + rotated = threading.Event() + + def on_create_session(*args, **kwargs): + if manager._database.spanner_api.create_session.call_count > 1: + rotated.set() + return DEFAULT + + manager._database.spanner_api.create_session.side_effect = on_create_session + # Maintenance thread is started. session_1 = manager.get_session(TransactionType.READ_ONLY) self.assertTrue(session_1.is_multiplexed) self.assertTrue(manager._multiplexed_session_thread.is_alive()) - # Wait for maintenance thread to execute. - self._assert_true_with_timeout( - lambda: manager._database.spanner_api.create_session.call_count > 1 - ) + # Wait for maintenance thread to execute without polling. + self.assertTrue(rotated.wait(timeout=10.0)) # Verify that maintenance thread created new multiplexed session. session_2 = manager.get_session(TransactionType.READ_ONLY) @@ -222,24 +228,32 @@ def test_concurrent_get_multiplexed_session_no_deadlock(self): # Mock maintenance thread creation to avoid spawning background tasks manager._build_maintenance_thread = Mock(return_value=Mock()) - # Mock _build_multiplexed_session to include a suspension point - async def slow_build(): - await asyncio.sleep(0.5) - return Mock() - - manager._build_multiplexed_session = slow_build - # Enable multiplexed sessions in environment for verification environ[DatabaseSessionsManager._ENV_VAR_MULTIPLEXED] = "true" async def run_concurrent(): + entered_slow_build = asyncio.Event() + release_slow_build = asyncio.Event() + + # Mock _build_multiplexed_session to include a suspension point + async def slow_build(): + entered_slow_build.set() + await release_slow_build.wait() + return Mock() + + manager._build_multiplexed_session = slow_build + # Trigger Coroutine 1 task1 = asyncio.create_task(manager._get_multiplexed_session()) - await asyncio.sleep(0.1) # Allow Coroutine 1 to acquire lock and suspend + # Wait until Coroutine 1 has acquired the lock and suspended + await entered_slow_build.wait() - # Trigger Coroutine 2 - this would previously block the main thread + # Trigger Coroutine 2 while Coroutine 1 holds the lock task2 = asyncio.create_task(manager._get_multiplexed_session()) + # Release Coroutine 1 to complete + release_slow_build.set() + await asyncio.gather(task1, task2) try: @@ -645,22 +659,6 @@ def test_rotate_multiplexed_session_build_failure(self): self.assertIs(manager._multiplexed_session, current_session) current_session.delete.assert_not_called() - def _assert_true_with_timeout(self, condition: Callable) -> None: - """Asserts that the given condition is met within a timeout period. - - :type condition: Callable - :param condition: A callable that returns a boolean indicating whether the condition is met. - """ - - sleep_seconds = 0.1 - timeout_seconds = 10 - - start_time = time() - while not condition() and time() - start_time < timeout_seconds: - sleep(sleep_seconds) - - self.assertTrue(condition()) - @staticmethod def _disable_multiplexed_sessions() -> None: """Sets environment variables to disable multiplexed sessions for all transactions types.""" diff --git a/packages/google-cloud-spanner/tests/unit/test_instance.py b/packages/google-cloud-spanner/tests/unit/test_instance.py index b089777cc853..9f52d70b1518 100644 --- a/packages/google-cloud-spanner/tests/unit/test_instance.py +++ b/packages/google-cloud-spanner/tests/unit/test_instance.py @@ -15,7 +15,6 @@ import unittest import mock -from google.auth.credentials import AnonymousCredentials from google.cloud.spanner_v1 import DefaultTransactionOptions @@ -52,6 +51,17 @@ def _getTargetClass(self): def _make_one(self, *args, **kwargs): return self._getTargetClass()(*args, **kwargs) + @staticmethod + def _make_database_admin_api(): + from google.cloud.spanner_admin_database_v1 import DatabaseAdminClient + from google.cloud.spanner_admin_database_v1.services.database_admin.transports.base import ( + DatabaseAdminTransport, + ) + + mock_transport = mock.create_autospec(DatabaseAdminTransport, instance=True) + mock_transport._wrapped_methods = {} + return DatabaseAdminClient(transport=mock_transport) + def test_constructor_defaults(self): from google.cloud.spanner_v1.instance import DEFAULT_NODE_COUNT @@ -586,14 +596,15 @@ def test_database_factory_explicit(self): self.assertIs(database._proto_descriptors, proto_descriptors) def test_list_databases(self): - from google.cloud.spanner_admin_database_v1 import Database as DatabasePB from google.cloud.spanner_admin_database_v1 import ( - DatabaseAdminClient, + Database as DatabasePB, + ) + from google.cloud.spanner_admin_database_v1 import ( ListDatabasesRequest, ListDatabasesResponse, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) @@ -629,12 +640,11 @@ def test_list_databases(self): def test_list_databases_w_options(self): from google.cloud.spanner_admin_database_v1 import ( - DatabaseAdminClient, ListDatabasesRequest, ListDatabasesResponse, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) @@ -711,14 +721,15 @@ def test_backup_factory_explicit(self): self.assertEqual(backup._encryption_config, encryption_config) def test_list_backups_defaults(self): - from google.cloud.spanner_admin_database_v1 import Backup as BackupPB from google.cloud.spanner_admin_database_v1 import ( - DatabaseAdminClient, + Backup as BackupPB, + ) + from google.cloud.spanner_admin_database_v1 import ( ListBackupsRequest, ListBackupsResponse, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) @@ -752,14 +763,15 @@ def test_list_backups_defaults(self): ) def test_list_backups_w_options(self): - from google.cloud.spanner_admin_database_v1 import Backup as BackupPB from google.cloud.spanner_admin_database_v1 import ( - DatabaseAdminClient, + Backup as BackupPB, + ) + from google.cloud.spanner_admin_database_v1 import ( ListBackupsRequest, ListBackupsResponse, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) @@ -801,12 +813,11 @@ def test_list_backup_operations_defaults(self): from google.cloud.spanner_admin_database_v1 import ( CreateBackupMetadata, - DatabaseAdminClient, ListBackupOperationsRequest, ListBackupOperationsResponse, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) @@ -849,12 +860,11 @@ def test_list_backup_operations_w_options(self): from google.cloud.spanner_admin_database_v1 import ( CreateBackupMetadata, - DatabaseAdminClient, ListBackupOperationsRequest, ListBackupOperationsResponse, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) @@ -899,13 +909,12 @@ def test_list_database_operations_defaults(self): from google.cloud.spanner_admin_database_v1 import ( CreateDatabaseMetadata, - DatabaseAdminClient, ListDatabaseOperationsRequest, ListDatabaseOperationsResponse, OptimizeRestoredDatabaseMetadata, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) @@ -955,7 +964,6 @@ def test_list_database_operations_w_options(self): from google.protobuf.any_pb2 import Any from google.cloud.spanner_admin_database_v1 import ( - DatabaseAdminClient, ListDatabaseOperationsRequest, ListDatabaseOperationsResponse, RestoreDatabaseMetadata, @@ -963,7 +971,7 @@ def test_list_database_operations_w_options(self): UpdateDatabaseDdlMetadata, ) - api = DatabaseAdminClient(credentials=AnonymousCredentials()) + api = self._make_database_admin_api() client = _Client(self.PROJECT) client.database_admin_api = api instance = self._make_one(self.INSTANCE_ID, client) diff --git a/packages/google-cloud-spanner/tests/unit/test_metrics_concurrency.py b/packages/google-cloud-spanner/tests/unit/test_metrics_concurrency.py index 83d88b0ab023..1ad20eef2740 100644 --- a/packages/google-cloud-spanner/tests/unit/test_metrics_concurrency.py +++ b/packages/google-cloud-spanner/tests/unit/test_metrics_concurrency.py @@ -13,7 +13,6 @@ # limitations under the License. import threading -import time import unittest from google.cloud.spanner_v1.metrics.metrics_capture import MetricsCapture @@ -34,6 +33,7 @@ def test_concurrent_tracers(self): factory.enabled = True errors = [] + barrier = threading.Barrier(10) def worker(idx): try: @@ -43,14 +43,15 @@ def worker(idx): tracer = SpannerMetricsTracerFactory.get_current_tracer() if tracer is None: errors.append(f"Thread {idx}: Tracer is None inside Capture") + barrier.abort() return # Set a unique attribute for this thread project_name = f"project-{idx}" tracer.set_project(project_name) - # Simulate some work - time.sleep(0.01) + # Synchronize all threads so they are all concurrently inside MetricsCapture + barrier.wait(timeout=10.0) # Verify verify we still have OUR tracer current_tracer = SpannerMetricsTracerFactory.get_current_tracer() @@ -67,7 +68,10 @@ def worker(idx): if interceptor_tracer is not tracer: errors.append(f"Thread {idx}: Interceptor tracer mismatch") + except threading.BrokenBarrierError: + pass except Exception as e: + barrier.abort() errors.append(f"Thread {idx}: Exception {e}") threads = [] @@ -79,6 +83,7 @@ def worker(idx): for t in threads: t.join() + self.assertFalse(barrier.broken, "Barrier timed out or was broken") self.assertEqual(errors, [], f"Concurrency errors found: {errors}") def test_context_var_cleanup(self): diff --git a/packages/google-cloud-spanner/tests/unit/test_session.py b/packages/google-cloud-spanner/tests/unit/test_session.py index 17704cb59b2f..0bb064d20a9b 100644 --- a/packages/google-cloud-spanner/tests/unit/test_session.py +++ b/packages/google-cloud-spanner/tests/unit/test_session.py @@ -1151,9 +1151,11 @@ def unit_of_work(transaction): list(transaction.read(TABLE_NAME, COLUMNS, KEYSET)) session.create() - session.run_in_transaction(unit_of_work) + with mock.patch("google.cloud.spanner_v1._helpers.time.sleep") as sleep_mock: + session.run_in_transaction(unit_of_work) self.assertEqual(begin_transaction.call_count, 2) + sleep_mock.assert_called_once() begin_transaction.assert_called_with( request=BeginTransactionRequest( @@ -1191,10 +1193,12 @@ def unit_of_work(transaction): list(transaction.read(TABLE_NAME, COLUMNS, KEYSET)) session.create() - session.run_in_transaction(unit_of_work) + with mock.patch("google.cloud.spanner_v1._helpers.time.sleep") as sleep_mock: + session.run_in_transaction(unit_of_work) # Verify retried BeginTransaction API call. self.assertEqual(begin_transaction.call_count, 2) + sleep_mock.assert_called_once() begin_transaction.assert_called_with( request=BeginTransactionRequest( @@ -1236,10 +1240,12 @@ def unit_of_work(transaction): list(transaction.read(TABLE_NAME, COLUMNS, KEYSET)) session.create() - session.run_in_transaction(unit_of_work) + with mock.patch("google.cloud.spanner_v1._helpers.time.sleep") as sleep_mock: + session.run_in_transaction(unit_of_work) # Verify retried BeginTransaction API call. self.assertEqual(begin_transaction.call_count, 2) + sleep_mock.assert_called_once() begin_transaction.assert_called_with( request=BeginTransactionRequest( @@ -1524,7 +1530,7 @@ def unit_of_work(txn, *args, **kw): called_with.append((txn, args, kw)) txn.insert(TABLE_NAME, COLUMNS, VALUES) - with mock.patch("time.sleep") as sleep_mock: + with mock.patch("google.cloud.spanner_v1._helpers.time.sleep") as sleep_mock: session.run_in_transaction(unit_of_work, "abc", some_arg="def") sleep_mock.assert_called_once_with(RETRY_SECONDS + RETRY_NANOS / 1.0e9) @@ -1639,7 +1645,7 @@ def unit_of_work(txn, *args, **kw): raise _make_rpc_error(Aborted, trailing_metadata) txn.insert(TABLE_NAME, COLUMNS, VALUES) - with mock.patch("time.sleep") as sleep_mock: + with mock.patch("google.cloud.spanner_v1._helpers.time.sleep") as sleep_mock: session.run_in_transaction(unit_of_work) sleep_mock.assert_called_once_with(RETRY_SECONDS + RETRY_NANOS / 1.0e9) @@ -1726,8 +1732,10 @@ def _time(): return _results[0] return 1.0 - with mock.patch("time.time", _time): - with mock.patch("time.sleep") as sleep_mock: + with mock.patch("google.cloud.spanner_v1._helpers.time.time", _time): + with mock.patch( + "google.cloud.spanner_v1._helpers.time.sleep" + ) as sleep_mock: # Exception has request_id attribute added with self.assertRaises(Aborted) as context: session.run_in_transaction(unit_of_work, "abc", timeout_secs=1) @@ -1803,11 +1811,11 @@ def _time(): return 1.0 with ( - mock.patch("time.time", _time), + mock.patch("google.cloud.spanner_v1._helpers.time.time", _time), mock.patch( "google.cloud.spanner_v1._helpers.random.random", return_value=0 ), - mock.patch("time.sleep") as sleep_mock, + mock.patch("google.cloud.spanner_v1._helpers.time.sleep") as sleep_mock, ): # Exception has request_id attribute added with self.assertRaises(Aborted) as context: @@ -2219,7 +2227,7 @@ def unit_of_work(txn, *args, **kw): called_with.append((txn, args, kw)) txn.insert(TABLE_NAME, COLUMNS, VALUES) - with mock.patch("time.sleep") as sleep_mock: + with mock.patch("google.cloud.spanner_v1._helpers.time.sleep") as sleep_mock: session.run_in_transaction( unit_of_work, "abc", @@ -2669,15 +2677,17 @@ def _time_func(): return 3 # check if current time > deadline - with mock.patch("time.time", _time_func): + with mock.patch("google.cloud.spanner_v1._helpers.time.time", _time_func): with self.assertRaises(Exception): _delay_until_retry(exc_mock, 2, 1, default_retry_delay=0) - with mock.patch("time.time", _time_func): + with mock.patch("google.cloud.spanner_v1._helpers.time.time", _time_func): with mock.patch( "google.cloud.spanner_v1._helpers._get_retry_delay" ) as get_retry_delay_mock: - with mock.patch("time.sleep") as sleep_mock: + with mock.patch( + "google.cloud.spanner_v1._helpers.time.sleep" + ) as sleep_mock: get_retry_delay_mock.return_value = None _delay_until_retry(exc_mock, 6, 1) diff --git a/packages/google-cloud-spanner/tests/unit/test_snapshot.py b/packages/google-cloud-spanner/tests/unit/test_snapshot.py index bd17f9cca606..bfc52ee4461d 100644 --- a/packages/google-cloud-spanner/tests/unit/test_snapshot.py +++ b/packages/google-cloud-spanner/tests/unit/test_snapshot.py @@ -1084,11 +1084,12 @@ def test_begin_precommit_token(self, mock_region): self._execute_begin(derived) + @mock.patch("google.cloud.spanner_v1._helpers.time.sleep") @mock.patch( "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", return_value="global", ) - def test_begin_retry_for_internal_server_error(self, mock_region): + def test_begin_retry_for_internal_server_error(self, mock_region, mock_sleep): derived = _build_snapshot_derived(multi_use=True) begin_transaction = derived._session._database.spanner_api.begin_transaction @@ -1098,6 +1099,7 @@ def test_begin_retry_for_internal_server_error(self, mock_region): ] self._execute_begin(derived, attempts=2) + mock_sleep.assert_called_once_with(2) expected_statuses = [ ( @@ -1108,11 +1110,12 @@ def test_begin_retry_for_internal_server_error(self, mock_region): actual_statuses = self.finished_spans_events_statuses() self.assertEqual(expected_statuses, actual_statuses) + @mock.patch("google.cloud.spanner_v1._helpers.time.sleep") @mock.patch( "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", return_value="global", ) - def test_begin_retry_for_aborted(self, mock_region): + def test_begin_retry_for_aborted(self, mock_region, mock_sleep): derived = _build_snapshot_derived(multi_use=True) begin_transaction = derived._session._database.spanner_api.begin_transaction @@ -1122,6 +1125,7 @@ def test_begin_retry_for_aborted(self, mock_region): ] self._execute_begin(derived, attempts=2) + mock_sleep.assert_called_once_with(2) expected_statuses = [ ( @@ -2057,7 +2061,8 @@ def test_partition_read_other_error(self, mock_region): ), ) - def test_partition_read_w_retry(self): + @mock.patch("google.cloud.spanner_v1._helpers.time.sleep") + def test_partition_read_w_retry(self, mock_sleep): from google.cloud.spanner_v1 import Partition, PartitionResponse, Transaction from google.cloud.spanner_v1.keyset import KeySet @@ -2087,6 +2092,7 @@ def test_partition_read_w_retry(self): list(derived.partition_read(TABLE_NAME, COLUMNS, keyset)) self.assertEqual(api.partition_read.call_count, 2) + mock_sleep.assert_called_once_with(2) @mock.patch( "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region",