From d06c59340433ba1d675bf20e4a080c42da732455 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Mon, 10 Oct 2022 09:48:51 +0530 Subject: [PATCH 01/19] feat: Inline Begin transction for RW transactions --- google/cloud/spanner_v1/pool.py | 1 - google/cloud/spanner_v1/session.py | 4 +- google/cloud/spanner_v1/transaction.py | 42 +- tests/unit/test_pool.py | 6 +- tests/unit/test_session.py | 31 +- tests/unit/test_spanner.py | 610 +++++++++++++++++++++++++ tests/unit/test_transaction.py | 27 -- 7 files changed, 646 insertions(+), 75 deletions(-) create mode 100644 tests/unit/test_spanner.py diff --git a/google/cloud/spanner_v1/pool.py b/google/cloud/spanner_v1/pool.py index 56a78ef672..9c76837255 100644 --- a/google/cloud/spanner_v1/pool.py +++ b/google/cloud/spanner_v1/pool.py @@ -515,7 +515,6 @@ def begin_pending_transactions(self): """Begin all transactions for sessions added to the pool.""" while not self._pending_sessions.empty(): session = self._pending_sessions.get() - session._transaction.begin() super(TransactionPingingPool, self).put(session) diff --git a/google/cloud/spanner_v1/session.py b/google/cloud/spanner_v1/session.py index 1ab6a93626..86a9bed0e8 100644 --- a/google/cloud/spanner_v1/session.py +++ b/google/cloud/spanner_v1/session.py @@ -352,9 +352,7 @@ def run_in_transaction(self, func, *args, **kw): txn.transaction_tag = transaction_tag else: txn = self._transaction - if txn._transaction_id is None: - txn.begin() - + try: attempts += 1 return_value = func(txn, *args, **kw) diff --git a/google/cloud/spanner_v1/transaction.py b/google/cloud/spanner_v1/transaction.py index d776b12469..a8328061b0 100644 --- a/google/cloud/spanner_v1/transaction.py +++ b/google/cloud/spanner_v1/transaction.py @@ -61,9 +61,7 @@ def _check_state(self): :raises: :exc:`ValueError` if the object's state is invalid for making API requests. """ - if self._transaction_id is None: - raise ValueError("Transaction is not begun") - + if self.committed is not None: raise ValueError("Transaction is already committed") @@ -78,7 +76,11 @@ def _make_txn_selector(self): :returns: a selector configured for read-write transaction semantics. """ self._check_state() - return TransactionSelector(id=self._transaction_id) + + if self._transaction_id is None: + return TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + else: + return TransactionSelector(id=self._transaction_id) def begin(self): """Begin a transaction on the database. @@ -111,15 +113,17 @@ def begin(self): def rollback(self): """Roll back a transaction on the database.""" self._check_state() - database = self._session._database - api = database.spanner_api - metadata = _metadata_with_prefix(database.name) - with trace_call("CloudSpanner.Rollback", self._session): - api.rollback( - session=self._session.name, - transaction_id=self._transaction_id, - metadata=metadata, - ) + + if self._transaction_id is not None: + database = self._session._database + api = database.spanner_api + metadata = _metadata_with_prefix(database.name) + with trace_call("CloudSpanner.Rollback", self._session): + api.rollback( + session=self._session.name, + transaction_id=self._transaction_id, + metadata=metadata, + ) self.rolled_back = True del self._session._transaction @@ -142,6 +146,8 @@ def commit(self, return_commit_stats=False, request_options=None): :raises ValueError: if there are no mutations to commit. """ self._check_state() + if self._transaction_id is None: + self.begin() database = self._session._database api = database.spanner_api @@ -302,6 +308,10 @@ def execute_update( response = api.execute_sql( request=request, metadata=metadata, retry=retry, timeout=timeout ) + + if self._transaction_id is None and response.metadata.transaction is not None: + self._transaction_id = response.metadata.transaction.id + return response.stats.row_count_exact def batch_update(self, statements, request_options=None): @@ -378,11 +388,15 @@ def batch_update(self, statements, request_options=None): row_counts = [ result_set.stats.row_count_exact for result_set in response.result_sets ] + + for result_set in response.result_sets: + if self._transaction_id is None and result_set.metadata.transaction is not None: + self._transaction_id = result_set.metadata.transaction.id + return response.status, row_counts def __enter__(self): """Begin ``with`` block.""" - self.begin() return self def __exit__(self, exc_type, exc_val, exc_tb): diff --git a/tests/unit/test_pool.py b/tests/unit/test_pool.py index 593420187d..ee59e34183 100644 --- a/tests/unit/test_pool.py +++ b/tests/unit/test_pool.py @@ -656,7 +656,7 @@ def test_bind(self): for session in SESSIONS: session.create.assert_not_called() txn = session._transaction - txn.begin.assert_called_once_with() + txn.begin.assert_not_called() self.assertTrue(pool._pending_sessions.empty()) @@ -685,7 +685,7 @@ def test_bind_w_timestamp_race(self): for session in SESSIONS: session.create.assert_not_called() txn = session._transaction - txn.begin.assert_called_once_with() + txn.begin.assert_not_called() self.assertTrue(pool._pending_sessions.empty()) @@ -771,7 +771,7 @@ def test_begin_pending_transactions_non_empty(self): pool.begin_pending_transactions() # no raise for txn in TRANSACTIONS: - txn.begin.assert_called_once_with() + txn.begin.assert_not_called() self.assertTrue(pending.empty()) diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 0f297654bb..97195734aa 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -725,17 +725,6 @@ def unit_of_work(txn, *args, **kw): self.assertEqual(args, ()) self.assertEqual(kw, {}) - expected_options = TransactionOptions(read_write=TransactionOptions.ReadWrite()) - gax_api.begin_transaction.assert_called_once_with( - session=self.SESSION_NAME, - options=expected_options, - metadata=[("google-cloud-resource-prefix", database.name)], - ) - gax_api.rollback.assert_called_once_with( - session=self.SESSION_NAME, - transaction_id=TRANSACTION_ID, - metadata=[("google-cloud-resource-prefix", database.name)], - ) def test_run_in_transaction_callback_raises_non_abort_rpc_error(self): from google.api_core.exceptions import Cancelled @@ -780,12 +769,6 @@ def unit_of_work(txn, *args, **kw): self.assertEqual(args, ()) self.assertEqual(kw, {}) - expected_options = TransactionOptions(read_write=TransactionOptions.ReadWrite()) - gax_api.begin_transaction.assert_called_once_with( - session=self.SESSION_NAME, - options=expected_options, - metadata=[("google-cloud-resource-prefix", database.name)], - ) gax_api.rollback.assert_not_called() def test_run_in_transaction_w_args_w_kwargs_wo_abort(self): @@ -1141,16 +1124,10 @@ def unit_of_work(txn, *args, **kw): self.assertEqual(kw, {}) expected_options = TransactionOptions(read_write=TransactionOptions.ReadWrite()) - self.assertEqual( - gax_api.begin_transaction.call_args_list, - [ - mock.call( - session=self.SESSION_NAME, - options=expected_options, - metadata=[("google-cloud-resource-prefix", database.name)], - ) - ] - * 2, + gax_api.begin_transaction.assert_called_once_with( + session=self.SESSION_NAME, + options=expected_options, + metadata=[("google-cloud-resource-prefix", database.name)], ) request = CommitRequest( session=self.SESSION_NAME, diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py new file mode 100644 index 0000000000..1ec87e9125 --- /dev/null +++ b/tests/unit/test_spanner.py @@ -0,0 +1,610 @@ +# Copyright 2016 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 dataclasses import fields +from google.protobuf.struct_pb2 import Struct +from google.cloud.spanner_v1 import ( + PartialResultSet, + ResultSetMetadata, + ResultSetStats, + ResultSet, + RequestOptions, + Type, + TypeCode, + ExecuteSqlRequest, + ReadRequest, + StructType, + TransactionOptions, + TransactionSelector, + ExecuteBatchDmlRequest, + ExecuteBatchDmlResponse, + param_types +) +from google.cloud.spanner_v1.types import transaction as transaction_type +from google.cloud.spanner_v1.keyset import KeySet + +from google.cloud.spanner_v1._helpers import ( + _make_value_pb, + _merge_query_options, +) + +import mock + +from google.api_core.retry import Retry +from google.api_core import gapic_v1 + +from tests._helpers import OpenTelemetryBase, StatusCode + +TABLE_NAME = "citizens" +COLUMNS = ["email", "first_name", "last_name", "age"] +VALUES = [ + ["phred@exammple.com", "Phred", "Phlyntstone", 32], + ["bharney@example.com", "Bharney", "Rhubble", 31], +] +DML_QUERY = """\ +INSERT INTO citizens(first_name, last_name, age) +VALUES ("Phred", "Phlyntstone", 32) +""" +DML_QUERY_WITH_PARAM = """ +INSERT INTO citizens(first_name, last_name, age) +VALUES ("Phred", "Phlyntstone", @age) +""" +SQL_QUERY = """\ +SELECT first_name, last_name, age FROM citizens ORDER BY age""" +SQL_QUERY_WITH_PARAM = """ +SELECT first_name, last_name, email FROM citizens WHERE age <= @max_age""" +PARAMS = {"age": 30} +PARAM_TYPES = {"age": Type(code=TypeCode.INT64)} + + +class TestTransaction(OpenTelemetryBase): + + PROJECT_ID = "project-id" + INSTANCE_ID = "instance-id" + INSTANCE_NAME = "projects/" + PROJECT_ID + "/instances/" + INSTANCE_ID + DATABASE_ID = "database-id" + DATABASE_NAME = INSTANCE_NAME + "/databases/" + DATABASE_ID + SESSION_ID = "session-id" + SESSION_NAME = DATABASE_NAME + "/sessions/" + SESSION_ID + TRANSACTION_ID = b"DEADBEEF" + TRANSACTION_TAG = "transaction-tag" + + BASE_ATTRIBUTES = { + "db.type": "spanner", + "db.url": "spanner.googleapis.com", + "db.instance": "testing", + "net.host.name": "spanner.googleapis.com", + } + + def _getTargetClass(self): + from google.cloud.spanner_v1.transaction import Transaction + + return Transaction + + def _make_one(self, session, *args, **kwargs): + transaction = self._getTargetClass()(session, *args, **kwargs) + session._transaction = transaction + return transaction + + def _make_spanner_api(self): + from google.cloud.spanner_v1 import SpannerClient + + return mock.create_autospec(SpannerClient, instance=True) + + def _execute_update_helper( + self, + transaction, + database, + count=0, + query_options=None, + request_options=None, + retry=gapic_v1.method.DEFAULT, + timeout=gapic_v1.method.DEFAULT, + begin=True + ): + MODE = 2 # PROFILE + stats_pb = ResultSetStats(row_count_exact=1) + + api = database.spanner_api = self._make_spanner_api() + + if begin is True: + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(transaction=transaction_pb) + api.execute_sql.return_value = ResultSet(stats=stats_pb, metadata=metadata_pb) + else: + api.execute_sql.return_value = ResultSet(stats=stats_pb) + + transaction.transaction_tag = self.TRANSACTION_TAG + transaction._execute_sql_count = count + + if request_options is None: + request_options = RequestOptions() + elif type(request_options) == dict: + request_options = RequestOptions(request_options) + + row_count = transaction.execute_update( + DML_QUERY_WITH_PARAM, + PARAMS, + PARAM_TYPES, + query_mode=MODE, + query_options=query_options, + request_options=request_options, + retry=retry, + timeout=timeout, + ) + + self.assertEqual(row_count, count + 1) + + if begin is True: + expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + else: + expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) + + expected_params = Struct( + fields={key: _make_value_pb(value) for (key, value) in PARAMS.items()} + ) + + expected_query_options = database._instance._client._query_options + if query_options: + expected_query_options = _merge_query_options( + expected_query_options, query_options + ) + expected_request_options = request_options + expected_request_options.transaction_tag = self.TRANSACTION_TAG + + expected_request = ExecuteSqlRequest( + session=self.SESSION_NAME, + sql=DML_QUERY_WITH_PARAM, + transaction=expected_transaction, + params=expected_params, + param_types=PARAM_TYPES, + query_mode=MODE, + query_options=expected_query_options, + request_options=request_options, + seqno=count, + ) + api.execute_sql.assert_called_once_with( + request=expected_request, + retry=retry, + timeout=timeout, + metadata=[("google-cloud-resource-prefix", database.name)], + ) + + self.assertEqual(transaction._execute_sql_count, count + 1) + + def _execute_sql_helper( + self, + transaction, + database, + count=0, + partition=None, + sql_count=0, + query_options=None, + request_options=None, + timeout=gapic_v1.method.DEFAULT, + retry=gapic_v1.method.DEFAULT, + begin=True + ): + + + VALUES = [["bharney", "rhubbyl", 31], ["phred", "phlyntstone", 32]] + VALUE_PBS = [[_make_value_pb(item) for item in row] for row in VALUES] + MODE = 2 # PROFILE + struct_type_pb = StructType( + fields=[ + StructType.Field(name="first_name", type_=Type(code=TypeCode.STRING)), + StructType.Field(name="last_name", type_=Type(code=TypeCode.STRING)), + StructType.Field(name="age", type_=Type(code=TypeCode.INT64)), + ] + ) + if begin is True: + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) + else: + metadata_pb = ResultSetMetadata(row_type=struct_type_pb) + stats_pb = ResultSetStats( + query_stats=Struct(fields={"rows_returned": _make_value_pb(2)}) + ) + result_sets = [ + PartialResultSet(metadata=metadata_pb), + PartialResultSet(stats=stats_pb), + ] + for i in range(len(result_sets)): + result_sets[i].values.extend(VALUE_PBS[i]) + iterator = _MockIterator(*result_sets) + api = database.spanner_api = self._make_spanner_api() + api.execute_streaming_sql.return_value = iterator + transaction._execute_sql_count = sql_count + transaction._read_request_count = count + + if request_options is None: + request_options = RequestOptions() + elif type(request_options) == dict: + request_options = RequestOptions(request_options) + + result_set = transaction.execute_sql( + SQL_QUERY_WITH_PARAM, + PARAMS, + PARAM_TYPES, + query_mode=MODE, + query_options=query_options, + request_options=request_options, + partition=partition, + retry=retry, + timeout=timeout, + ) + + self.assertEqual(transaction._read_request_count, count + 1) + + self.assertEqual(list(result_set), VALUES) + self.assertEqual(result_set.metadata, metadata_pb) + self.assertEqual(result_set.stats, stats_pb) + + if begin is True: + expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + else: + expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) + + expected_params = Struct( + fields={key: _make_value_pb(value) for (key, value) in PARAMS.items()} + ) + + expected_query_options = database._instance._client._query_options + if query_options: + expected_query_options = _merge_query_options( + expected_query_options, query_options + ) + expected_request_options = request_options + + expected_request = ExecuteSqlRequest( + session=self.SESSION_NAME, + sql=SQL_QUERY_WITH_PARAM, + transaction=expected_transaction, + params=expected_params, + param_types=PARAM_TYPES, + query_mode=MODE, + query_options=expected_query_options, + request_options=expected_request_options, + partition_token=partition, + seqno=sql_count, + ) + api.execute_streaming_sql.assert_called_once_with( + request=expected_request, + metadata=[("google-cloud-resource-prefix", database.name)], + timeout=timeout, + retry=retry, + ) + + self.assertEqual(transaction._execute_sql_count, sql_count + 1) + + def _read_helper( + self, + transaction, + database, + count=0, + partition=None, + timeout=gapic_v1.method.DEFAULT, + retry=gapic_v1.method.DEFAULT, + request_options=None, + begin=True + ): + VALUES = [["bharney", 31], ["phred", 32]] + VALUE_PBS = [[_make_value_pb(item) for item in row] for row in VALUES] + struct_type_pb = StructType( + fields=[ + StructType.Field(name="name", type_=Type(code=TypeCode.STRING)), + StructType.Field(name="age", type_=Type(code=TypeCode.INT64)), + ] + ) + + if begin is True: + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) + else: + metadata_pb = ResultSetMetadata(row_type=struct_type_pb) + + stats_pb = ResultSetStats( + query_stats=Struct(fields={"rows_returned": _make_value_pb(2)}) + ) + result_sets = [ + PartialResultSet(metadata=metadata_pb), + PartialResultSet(stats=stats_pb), + ] + for i in range(len(result_sets)): + result_sets[i].values.extend(VALUE_PBS[i]) + KEYS = [["bharney@example.com"], ["phred@example.com"]] + keyset = KeySet(keys=KEYS) + INDEX = "email-address-index" + LIMIT = 20 + api = database.spanner_api = self._make_spanner_api() + api.streaming_read.return_value = _MockIterator(*result_sets) + transaction._read_request_count = count + + if request_options is None: + request_options = RequestOptions() + elif type(request_options) == dict: + request_options = RequestOptions(request_options) + if partition is not None: # 'limit' and 'partition' incompatible + result_set = transaction.read( + TABLE_NAME, + COLUMNS, + keyset, + index=INDEX, + partition=partition, + retry=retry, + timeout=timeout, + request_options=request_options, + ) + else: + result_set = transaction.read( + TABLE_NAME, + COLUMNS, + keyset, + index=INDEX, + limit=LIMIT, + retry=retry, + timeout=timeout, + request_options=request_options, + ) + + self.assertEqual(transaction._read_request_count, count + 1) + + self.assertIs(result_set._source, transaction) + + self.assertEqual(list(result_set), VALUES) + self.assertEqual(result_set.metadata, metadata_pb) + self.assertEqual(result_set.stats, stats_pb) + + if begin is True: + expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + else: + expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) + + if partition is not None: + expected_limit = 0 + else: + expected_limit = LIMIT + + # Transaction tag is ignored for read request. + expected_request_options = request_options + expected_request_options.transaction_tag = None + + expected_request = ReadRequest( + session=self.SESSION_NAME, + table=TABLE_NAME, + columns=COLUMNS, + key_set=keyset._to_pb(), + transaction=expected_transaction, + index=INDEX, + limit=expected_limit, + partition_token=partition, + request_options=expected_request_options, + ) + api.streaming_read.assert_called_once_with( + request=expected_request, + metadata=[("google-cloud-resource-prefix", database.name)], + retry=retry, + timeout=timeout, + ) + + def _batch_update_helper(self, transaction, database, error_after=None, count=0, request_options=None, begin=True): + from google.rpc.status_pb2 import Status + insert_dml = "INSERT INTO table(pkey, desc) VALUES (%pkey, %desc)" + insert_params = {"pkey": 12345, "desc": "DESCRIPTION"} + insert_param_types = {"pkey": param_types.INT64, "desc": param_types.STRING} + update_dml = 'UPDATE table SET desc = desc + "-amended"' + delete_dml = "DELETE FROM table WHERE desc IS NULL" + + dml_statements = [ + (insert_dml, insert_params, insert_param_types), + update_dml, + delete_dml, + ] + + stats_pbs = [ + ResultSetStats(row_count_exact=1), + ResultSetStats(row_count_exact=2), + ResultSetStats(row_count_exact=3), + ] + if error_after is not None: + stats_pbs = stats_pbs[:error_after] + expected_status = Status(code=400) + else: + expected_status = Status(code=200) + expected_row_counts = [stats.row_count_exact for stats in stats_pbs] + if begin is True: + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(transaction=transaction_pb) + result_sets_pb = [ResultSet(stats=stats_pb, metadata= metadata_pb) for stats_pb in stats_pbs] + + else: + result_sets_pb = [ResultSet(stats=stats_pb) for stats_pb in stats_pbs] + + response = ExecuteBatchDmlResponse( + status=expected_status, + result_sets=result_sets_pb, + ) + + api = database.spanner_api = self._make_spanner_api() + api.execute_batch_dml.return_value = response + transaction.transaction_tag = self.TRANSACTION_TAG + transaction._execute_sql_count = count + + if request_options is None: + request_options = RequestOptions() + elif type(request_options) == dict: + request_options = RequestOptions(request_options) + + status, row_counts = transaction.batch_update( + dml_statements, request_options=request_options + ) + + self.assertEqual(status, expected_status) + self.assertEqual(row_counts, expected_row_counts) + + if begin is True: + expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + else: + expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) + + expected_insert_params = Struct( + fields={ + key: _make_value_pb(value) for (key, value) in insert_params.items() + } + ) + expected_statements = [ + ExecuteBatchDmlRequest.Statement( + sql=insert_dml, + params=expected_insert_params, + param_types=insert_param_types, + ), + ExecuteBatchDmlRequest.Statement(sql=update_dml), + ExecuteBatchDmlRequest.Statement(sql=delete_dml), + ] + expected_request_options = request_options + expected_request_options.transaction_tag = self.TRANSACTION_TAG + + expected_request = ExecuteBatchDmlRequest( + session=self.SESSION_NAME, + transaction=expected_transaction, + statements=expected_statements, + seqno=count, + request_options=expected_request_options, + ) + api.execute_batch_dml.assert_called_once_with( + request=expected_request, + metadata=[("google-cloud-resource-prefix", database.name)], + ) + + self.assertEqual(transaction._execute_sql_count, count + 1) + + def test_insert(self, transaction): + from google.cloud.spanner_v1 import Mutation + + transaction.insert(TABLE_NAME, columns=COLUMNS, values=VALUES) + + self.assertEqual(len(transaction._mutations), 1) + mutation = transaction._mutations[0] + self.assertIsInstance(mutation, Mutation) + write = mutation.insert + self.assertIsInstance(write, Mutation.Write) + self.assertEqual(write.table, TABLE_NAME) + self.assertEqual(write.columns, COLUMNS) + self._compare_values(write.values, VALUES) + + def test_transaction_should_include_begin_with_first_update(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._execute_update_helper(transaction=transaction, database=database) + + def test_transaction_should_include_begin_with_first_query(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._execute_sql_helper(transaction=transaction, database=database) + + def test_transaction_should_include_begin_with_first_read(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._read_helper(transaction=transaction, database=database) + + def test_transaction_should_include_begin_with_first_batch_update(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._batch_update_helper(transaction=transaction, database=database) + + def test_transaction_should_use_transaction_id_returned_by_first_query(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._execute_sql_helper(transaction=transaction, database=database) + self._execute_update_helper(transaction=transaction, database=database, begin=False) + + def test_transaction_should_use_transaction_id_returned_by_first_update(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._execute_update_helper(transaction=transaction, database=database, begin=True) + self._execute_sql_helper(transaction=transaction, database=database, begin=False) + + def test_transaction_should_use_transaction_id_returned_by_first_read(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._read_helper(transaction=transaction, database=database, begin=True) + self._batch_update_helper(transaction=transaction, database=database, begin=False) + + def test_transaction_should_use_transaction_id_returned_by_first_batch_update(self): + database = _Database() + session = _Session(database) + transaction = self._make_one(session) + self._batch_update_helper(transaction=transaction, database=database, begin=True) + self._read_helper(transaction=transaction, database=database, begin=False) + +class _Client(object): + def __init__(self): + from google.cloud.spanner_v1 import ExecuteSqlRequest + + self._query_options = ExecuteSqlRequest.QueryOptions(optimizer_version="1") + + +class _Instance(object): + def __init__(self): + self._client = _Client() + + +class _Database(object): + def __init__(self): + self.name = "testing" + self._instance = _Instance() + + +class _Session(object): + + _transaction = None + + def __init__(self, database=None, name=TestTransaction.SESSION_NAME): + self._database = database + self.name = name + +class _MockIterator(object): + def __init__(self, *values, **kw): + self._iter_values = iter(values) + self._fail_after = kw.pop("fail_after", False) + self._error = kw.pop("error", Exception) + + def __iter__(self): + return self + + def __next__(self): + try: + return next(self._iter_values) + except StopIteration: + if self._fail_after: + raise self._error + raise + + next = __next__ diff --git a/tests/unit/test_transaction.py b/tests/unit/test_transaction.py index d4d9c99c02..f9e471b8f1 100644 --- a/tests/unit/test_transaction.py +++ b/tests/unit/test_transaction.py @@ -91,12 +91,6 @@ def test_ctor_defaults(self): self.assertTrue(transaction._multi_use) self.assertEqual(transaction._execute_sql_count, 0) - def test__check_state_not_begun(self): - session = _Session() - transaction = self._make_one(session) - with self.assertRaises(ValueError): - transaction._check_state() - def test__check_state_already_committed(self): session = _Session() transaction = self._make_one(session) @@ -194,14 +188,6 @@ def test_begin_ok(self): "CloudSpanner.BeginTransaction", attributes=TestTransaction.BASE_ATTRIBUTES ) - def test_rollback_not_begun(self): - session = _Session() - transaction = self._make_one(session) - with self.assertRaises(ValueError): - transaction.rollback() - - self.assertNoSpans() - def test_rollback_already_committed(self): session = _Session() transaction = self._make_one(session) @@ -267,14 +253,6 @@ def test_rollback_ok(self): "CloudSpanner.Rollback", attributes=TestTransaction.BASE_ATTRIBUTES ) - def test_commit_not_begun(self): - session = _Session() - transaction = self._make_one(session) - with self.assertRaises(ValueError): - transaction.commit() - - self.assertNoSpans() - def test_commit_already_committed(self): session = _Session() transaction = self._make_one(session) @@ -840,11 +818,6 @@ def test_context_mgr_failure(self): self.assertEqual(api._committed, None) - session_id, txn_id, metadata = api._rolled_back - self.assertEqual(session_id, session.name) - self.assertEqual(txn_id, self.TRANSACTION_ID) - self.assertEqual(metadata, [("google-cloud-resource-prefix", database.name)]) - class _Client(object): def __init__(self): From 84923b5e77308ba05e1b515484762ac5c17155f8 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Tue, 8 Nov 2022 14:28:23 +0530 Subject: [PATCH 02/19] ILB with lock for execute update and batch update --- google/cloud/spanner_v1/snapshot.py | 1 + google/cloud/spanner_v1/transaction.py | 110 +++++- tests/unit/test_spanner.py | 446 +++++++++++++++---------- 3 files changed, 370 insertions(+), 187 deletions(-) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index a55c3994c4..7359737f2e 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -106,6 +106,7 @@ class _SnapshotBase(_SessionWrapper): _transaction_id = None _read_request_count = 0 _execute_sql_count = 0 + _inline_begin_started = False def _make_txn_selector(self): """Helper for :meth:`read` / :meth:`execute_sql`. diff --git a/google/cloud/spanner_v1/transaction.py b/google/cloud/spanner_v1/transaction.py index a8328061b0..a3c2b11c57 100644 --- a/google/cloud/spanner_v1/transaction.py +++ b/google/cloud/spanner_v1/transaction.py @@ -13,7 +13,8 @@ # limitations under the License. """Spanner read-write transaction support.""" - +import functools +import threading from google.protobuf.struct_pb2 import Struct from google.cloud.spanner_v1._helpers import ( @@ -48,6 +49,7 @@ class Transaction(_SnapshotBase, _BatchBase): commit_stats = None _multi_use = True _execute_sql_count = 0 + _lock = threading.Lock() def __init__(self, session): if session._transaction is not None: @@ -82,6 +84,16 @@ def _make_txn_selector(self): else: return TransactionSelector(id=self._transaction_id) + def _execute_request( + self, method, request, trace_name=None, session=None, attributes=None + ): + transaction = self._make_txn_selector() + request.transaction = transaction + with trace_call(trace_name, session, attributes): + response = method(request=request) + + return response + def begin(self): """Begin a transaction on the database. @@ -270,7 +282,7 @@ def execute_update( params_pb = self._make_params_pb(params, param_types) database = self._session._database metadata = _metadata_with_prefix(database.name) - transaction = self._make_txn_selector() + api = database.spanner_api seqno, self._execute_sql_count = ( @@ -294,7 +306,6 @@ def execute_update( request = ExecuteSqlRequest( session=self._session.name, sql=dml, - transaction=transaction, params=params_pb, param_types=param_types, query_mode=query_mode, @@ -302,15 +313,45 @@ def execute_update( seqno=seqno, request_options=request_options, ) - with trace_call( - "CloudSpanner.ReadWriteTransaction", self._session, trace_attributes - ): - response = api.execute_sql( - request=request, metadata=metadata, retry=retry, timeout=timeout - ) - if self._transaction_id is None and response.metadata.transaction is not None: - self._transaction_id = response.metadata.transaction.id + method = functools.partial( + api.execute_sql, + request=request, + metadata=metadata, + retry=retry, + timeout=timeout, + ) + + if self._transaction_id is None: + with self._lock: + if self._inline_begin_started is False: + response = self._execute_request( + method, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes + ) + + if self._transaction_id is None and response.metadata.transaction is not None: + self._transaction_id = response.metadata.transaction.id + self._inline_begin_started = True + else: + response = self._execute_request( + method, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes + ) + else: + response = self._execute_request( + method, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes + ) return response.stats.row_count_exact @@ -378,21 +419,54 @@ def batch_update(self, statements, request_options=None): } request = ExecuteBatchDmlRequest( session=self._session.name, - transaction=transaction, statements=parsed, seqno=seqno, request_options=request_options, ) - with trace_call("CloudSpanner.DMLTransaction", self._session, trace_attributes): - response = api.execute_batch_dml(request=request, metadata=metadata) + + method = functools.partial( + api.execute_batch_dml, + request=request, + metadata=metadata, + ) + + if self._transaction_id is None: + with self._lock: + if self._inline_begin_started is False: + response = self._execute_request( + method, + request, + "CloudSpanner.DMLTransaction", + self._session, + trace_attributes + ) + + for result_set in response.result_sets: + if self._transaction_id is None and result_set.metadata.transaction is not None: + self._transaction_id = result_set.metadata.transaction.id + + self._inline_begin_started = True + else: + response = self._execute_request( + method, + request, + "CloudSpanner.DMLTransaction", + self._session, + trace_attributes + ) + else: + response = self._execute_request( + method, + request, + "CloudSpanner.DMLTransaction", + self._session, + trace_attributes + ) + row_counts = [ result_set.stats.row_count_exact for result_set in response.result_sets ] - for result_set in response.result_sets: - if self._transaction_id is None and result_set.metadata.transaction is not None: - self._transaction_id = result_set.metadata.transaction.id - return response.status, row_counts def __enter__(self): diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index 1ec87e9125..c46b27cad1 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -14,6 +14,7 @@ from dataclasses import fields +import threading from google.protobuf.struct_pb2 import Struct from google.cloud.spanner_v1 import ( PartialResultSet, @@ -67,7 +68,25 @@ SELECT first_name, last_name, email FROM citizens WHERE age <= @max_age""" PARAMS = {"age": 30} PARAM_TYPES = {"age": Type(code=TypeCode.INT64)} - +KEYS = [["bharney@example.com"], ["phred@example.com"]] +KEYSET = KeySet(keys=KEYS) +INDEX = "email-address-index" +LIMIT = 20 +MODE = 2 +RETRY=gapic_v1.method.DEFAULT +TIMEOUT=gapic_v1.method.DEFAULT +REQUEST_OPTIONS = RequestOptions() +insert_dml = "INSERT INTO table(pkey, desc) VALUES (%pkey, %desc)" +insert_params = {"pkey": 12345, "desc": "DESCRIPTION"} +insert_param_types = {"pkey": param_types.INT64, "desc": param_types.STRING} +update_dml = 'UPDATE table SET desc = desc + "-amended"' +delete_dml = "DELETE FROM table WHERE desc IS NULL" + +dml_statements = [ + (insert_dml, insert_params, insert_param_types), + update_dml, + delete_dml, +] class TestTransaction(OpenTelemetryBase): @@ -106,49 +125,34 @@ def _make_spanner_api(self): def _execute_update_helper( self, transaction, - database, + api, count=0, query_options=None, - request_options=None, - retry=gapic_v1.method.DEFAULT, - timeout=gapic_v1.method.DEFAULT, - begin=True ): - MODE = 2 # PROFILE stats_pb = ResultSetStats(row_count_exact=1) - - api = database.spanner_api = self._make_spanner_api() - - if begin is True: - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID - ) - metadata_pb = ResultSetMetadata(transaction=transaction_pb) - api.execute_sql.return_value = ResultSet(stats=stats_pb, metadata=metadata_pb) - else: - api.execute_sql.return_value = ResultSet(stats=stats_pb) + + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(transaction=transaction_pb) + api.execute_sql.return_value = ResultSet(stats=stats_pb, metadata=metadata_pb) transaction.transaction_tag = self.TRANSACTION_TAG transaction._execute_sql_count = count - if request_options is None: - request_options = RequestOptions() - elif type(request_options) == dict: - request_options = RequestOptions(request_options) - row_count = transaction.execute_update( DML_QUERY_WITH_PARAM, PARAMS, PARAM_TYPES, query_mode=MODE, query_options=query_options, - request_options=request_options, - retry=retry, - timeout=timeout, + request_options=REQUEST_OPTIONS, + retry=RETRY, + timeout=TIMEOUT, ) - self.assertEqual(row_count, count + 1) + def _execute_update_expected_request(self, database, query_options=None, begin=True, count=0): if begin is True: expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) else: @@ -163,7 +167,7 @@ def _execute_update_helper( expected_query_options = _merge_query_options( expected_query_options, query_options ) - expected_request_options = request_options + expected_request_options = REQUEST_OPTIONS expected_request_options.transaction_tag = self.TRANSACTION_TAG expected_request = ExecuteSqlRequest( @@ -174,36 +178,24 @@ def _execute_update_helper( param_types=PARAM_TYPES, query_mode=MODE, query_options=expected_query_options, - request_options=request_options, + request_options=expected_request_options, seqno=count, ) - api.execute_sql.assert_called_once_with( - request=expected_request, - retry=retry, - timeout=timeout, - metadata=[("google-cloud-resource-prefix", database.name)], - ) - - self.assertEqual(transaction._execute_sql_count, count + 1) + + return expected_request def _execute_sql_helper( self, transaction, database, + api, count=0, partition=None, sql_count=0, query_options=None, - request_options=None, - timeout=gapic_v1.method.DEFAULT, - retry=gapic_v1.method.DEFAULT, - begin=True ): - - VALUES = [["bharney", "rhubbyl", 31], ["phred", "phlyntstone", 32]] VALUE_PBS = [[_make_value_pb(item) for item in row] for row in VALUES] - MODE = 2 # PROFILE struct_type_pb = StructType( fields=[ StructType.Field(name="first_name", type_=Type(code=TypeCode.STRING)), @@ -211,13 +203,10 @@ def _execute_sql_helper( StructType.Field(name="age", type_=Type(code=TypeCode.INT64)), ] ) - if begin is True: - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID - ) - metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) - else: - metadata_pb = ResultSetMetadata(row_type=struct_type_pb) + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) stats_pb = ResultSetStats( query_stats=Struct(fields={"rows_returned": _make_value_pb(2)}) ) @@ -228,26 +217,20 @@ def _execute_sql_helper( for i in range(len(result_sets)): result_sets[i].values.extend(VALUE_PBS[i]) iterator = _MockIterator(*result_sets) - api = database.spanner_api = self._make_spanner_api() api.execute_streaming_sql.return_value = iterator transaction._execute_sql_count = sql_count transaction._read_request_count = count - if request_options is None: - request_options = RequestOptions() - elif type(request_options) == dict: - request_options = RequestOptions(request_options) - result_set = transaction.execute_sql( SQL_QUERY_WITH_PARAM, PARAMS, PARAM_TYPES, query_mode=MODE, query_options=query_options, - request_options=request_options, + request_options=REQUEST_OPTIONS, partition=partition, - retry=retry, - timeout=timeout, + retry=RETRY, + timeout=TIMEOUT, ) self.assertEqual(transaction._read_request_count, count + 1) @@ -255,7 +238,9 @@ def _execute_sql_helper( self.assertEqual(list(result_set), VALUES) self.assertEqual(result_set.metadata, metadata_pb) self.assertEqual(result_set.stats, stats_pb) + self.assertEqual(transaction._execute_sql_count, sql_count + 1) + def _execute_sql_expected_request(self, database, partition=None, query_options=None, begin=True, sql_count=0): if begin is True: expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) else: @@ -270,7 +255,6 @@ def _execute_sql_helper( expected_query_options = _merge_query_options( expected_query_options, query_options ) - expected_request_options = request_options expected_request = ExecuteSqlRequest( session=self.SESSION_NAME, @@ -280,29 +264,19 @@ def _execute_sql_helper( param_types=PARAM_TYPES, query_mode=MODE, query_options=expected_query_options, - request_options=expected_request_options, + request_options=REQUEST_OPTIONS, partition_token=partition, seqno=sql_count, ) - api.execute_streaming_sql.assert_called_once_with( - request=expected_request, - metadata=[("google-cloud-resource-prefix", database.name)], - timeout=timeout, - retry=retry, - ) - - self.assertEqual(transaction._execute_sql_count, sql_count + 1) + + return expected_request def _read_helper( self, transaction, - database, + api, count=0, partition=None, - timeout=gapic_v1.method.DEFAULT, - retry=gapic_v1.method.DEFAULT, - request_options=None, - begin=True ): VALUES = [["bharney", 31], ["phred", 32]] VALUE_PBS = [[_make_value_pb(item) for item in row] for row in VALUES] @@ -313,14 +287,11 @@ def _read_helper( ] ) - if begin is True: - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID - ) - metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) - else: - metadata_pb = ResultSetMetadata(row_type=struct_type_pb) - + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) + stats_pb = ResultSetStats( query_stats=Struct(fields={"rows_returned": _make_value_pb(2)}) ) @@ -330,39 +301,31 @@ def _read_helper( ] for i in range(len(result_sets)): result_sets[i].values.extend(VALUE_PBS[i]) - KEYS = [["bharney@example.com"], ["phred@example.com"]] - keyset = KeySet(keys=KEYS) - INDEX = "email-address-index" - LIMIT = 20 - api = database.spanner_api = self._make_spanner_api() + api.streaming_read.return_value = _MockIterator(*result_sets) transaction._read_request_count = count - if request_options is None: - request_options = RequestOptions() - elif type(request_options) == dict: - request_options = RequestOptions(request_options) if partition is not None: # 'limit' and 'partition' incompatible result_set = transaction.read( TABLE_NAME, COLUMNS, - keyset, + KEYSET, index=INDEX, partition=partition, - retry=retry, - timeout=timeout, - request_options=request_options, + retry=RETRY, + timeout=TIMEOUT, + request_options=REQUEST_OPTIONS, ) else: result_set = transaction.read( TABLE_NAME, COLUMNS, - keyset, + KEYSET, index=INDEX, limit=LIMIT, - retry=retry, - timeout=timeout, - request_options=request_options, + retry=RETRY, + timeout=TIMEOUT, + request_options=REQUEST_OPTIONS, ) self.assertEqual(transaction._read_request_count, count + 1) @@ -373,6 +336,8 @@ def _read_helper( self.assertEqual(result_set.metadata, metadata_pb) self.assertEqual(result_set.stats, stats_pb) + def _read_helper_expected_request(self, partition=None, begin=True, count=0): + if begin is True: expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) else: @@ -384,40 +349,25 @@ def _read_helper( expected_limit = LIMIT # Transaction tag is ignored for read request. - expected_request_options = request_options + expected_request_options = REQUEST_OPTIONS expected_request_options.transaction_tag = None expected_request = ReadRequest( session=self.SESSION_NAME, table=TABLE_NAME, columns=COLUMNS, - key_set=keyset._to_pb(), + key_set=KEYSET._to_pb(), transaction=expected_transaction, index=INDEX, limit=expected_limit, partition_token=partition, request_options=expected_request_options, ) - api.streaming_read.assert_called_once_with( - request=expected_request, - metadata=[("google-cloud-resource-prefix", database.name)], - retry=retry, - timeout=timeout, - ) + + return expected_request - def _batch_update_helper(self, transaction, database, error_after=None, count=0, request_options=None, begin=True): + def _batch_update_helper(self, transaction, database, api, error_after=None, count=0,): from google.rpc.status_pb2 import Status - insert_dml = "INSERT INTO table(pkey, desc) VALUES (%pkey, %desc)" - insert_params = {"pkey": 12345, "desc": "DESCRIPTION"} - insert_param_types = {"pkey": param_types.INT64, "desc": param_types.STRING} - update_dml = 'UPDATE table SET desc = desc + "-amended"' - delete_dml = "DELETE FROM table WHERE desc IS NULL" - - dml_statements = [ - (insert_dml, insert_params, insert_param_types), - update_dml, - delete_dml, - ] stats_pbs = [ ResultSetStats(row_count_exact=1), @@ -430,38 +380,30 @@ def _batch_update_helper(self, transaction, database, error_after=None, count=0, else: expected_status = Status(code=200) expected_row_counts = [stats.row_count_exact for stats in stats_pbs] - if begin is True: - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID - ) - metadata_pb = ResultSetMetadata(transaction=transaction_pb) - result_sets_pb = [ResultSet(stats=stats_pb, metadata= metadata_pb) for stats_pb in stats_pbs] - - else: - result_sets_pb = [ResultSet(stats=stats_pb) for stats_pb in stats_pbs] + transaction_pb = transaction_type.Transaction( + id = self.TRANSACTION_ID + ) + metadata_pb = ResultSetMetadata(transaction=transaction_pb) + result_sets_pb = [ResultSet(stats=stats_pb, metadata= metadata_pb) for stats_pb in stats_pbs] response = ExecuteBatchDmlResponse( status=expected_status, result_sets=result_sets_pb, ) - api = database.spanner_api = self._make_spanner_api() api.execute_batch_dml.return_value = response transaction.transaction_tag = self.TRANSACTION_TAG transaction._execute_sql_count = count - if request_options is None: - request_options = RequestOptions() - elif type(request_options) == dict: - request_options = RequestOptions(request_options) - status, row_counts = transaction.batch_update( - dml_statements, request_options=request_options + dml_statements, request_options=REQUEST_OPTIONS ) self.assertEqual(status, expected_status) self.assertEqual(row_counts, expected_row_counts) + self.assertEqual(transaction._execute_sql_count, count + 1) + def _batch_update_expected_request(self, begin=True, count=0): if begin is True: expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) else: @@ -481,7 +423,8 @@ def _batch_update_helper(self, transaction, database, error_after=None, count=0, ExecuteBatchDmlRequest.Statement(sql=update_dml), ExecuteBatchDmlRequest.Statement(sql=delete_dml), ] - expected_request_options = request_options + + expected_request_options = REQUEST_OPTIONS expected_request_options.transaction_tag = self.TRANSACTION_TAG expected_request = ExecuteBatchDmlRequest( @@ -491,78 +434,243 @@ def _batch_update_helper(self, transaction, database, error_after=None, count=0, seqno=count, request_options=expected_request_options, ) - api.execute_batch_dml.assert_called_once_with( - request=expected_request, - metadata=[("google-cloud-resource-prefix", database.name)], - ) - - self.assertEqual(transaction._execute_sql_count, count + 1) - - def test_insert(self, transaction): - from google.cloud.spanner_v1 import Mutation - - transaction.insert(TABLE_NAME, columns=COLUMNS, values=VALUES) - - self.assertEqual(len(transaction._mutations), 1) - mutation = transaction._mutations[0] - self.assertIsInstance(mutation, Mutation) - write = mutation.insert - self.assertIsInstance(write, Mutation.Write) - self.assertEqual(write.table, TABLE_NAME) - self.assertEqual(write.columns, COLUMNS) - self._compare_values(write.values, VALUES) + + return expected_request def test_transaction_should_include_begin_with_first_update(self): database = _Database() session = _Session(database) + api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._execute_update_helper(transaction=transaction, database=database) + self._execute_update_helper(transaction=transaction, api=api) + + api.execute_sql.assert_called_once_with( + request=self._execute_update_expected_request(database=database), + retry=RETRY, + timeout=TIMEOUT, + metadata=[("google-cloud-resource-prefix", database.name)], + ) def test_transaction_should_include_begin_with_first_query(self): database = _Database() session = _Session(database) + api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._execute_sql_helper(transaction=transaction, database=database) + self._execute_sql_helper(transaction=transaction, database=database, api=api) + + api.execute_streaming_sql.assert_called_once_with( + request=self._execute_sql_expected_request(database=database), + metadata=[("google-cloud-resource-prefix", database.name)], + timeout=TIMEOUT, + retry=RETRY, + ) def test_transaction_should_include_begin_with_first_read(self): database = _Database() session = _Session(database) + api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._read_helper(transaction=transaction, database=database) + self._read_helper(transaction=transaction, api=api) + + api.streaming_read.assert_called_once_with( + request=self._read_helper_expected_request(), + metadata=[("google-cloud-resource-prefix", database.name)], + retry=RETRY, + timeout=TIMEOUT, + ) def test_transaction_should_include_begin_with_first_batch_update(self): database = _Database() session = _Session(database) + api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._batch_update_helper(transaction=transaction, database=database) + self._batch_update_helper(transaction=transaction, database=database, api=api) + api.execute_batch_dml.assert_called_once_with( + request=self._batch_update_expected_request(), + metadata=[("google-cloud-resource-prefix", database.name)], + ) def test_transaction_should_use_transaction_id_returned_by_first_query(self): database = _Database() session = _Session(database) + api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._execute_sql_helper(transaction=transaction, database=database) - self._execute_update_helper(transaction=transaction, database=database, begin=False) + self._execute_sql_helper(transaction=transaction, database=database, api=api) + api.execute_streaming_sql.assert_called_once_with( + request=self._execute_sql_expected_request(database=database), + retry=gapic_v1.method.DEFAULT, + timeout=gapic_v1.method.DEFAULT, + metadata=[("google-cloud-resource-prefix", database.name)], + ) + + self._execute_update_helper(transaction=transaction, api=api) + api.execute_sql.assert_called_once_with( + request=self._execute_update_expected_request(database=database, begin=False), + retry=gapic_v1.method.DEFAULT, + timeout=gapic_v1.method.DEFAULT, + metadata=[("google-cloud-resource-prefix", database.name)], + ) def test_transaction_should_use_transaction_id_returned_by_first_update(self): database = _Database() session = _Session(database) + api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._execute_update_helper(transaction=transaction, database=database, begin=True) - self._execute_sql_helper(transaction=transaction, database=database, begin=False) + self._execute_update_helper(transaction=transaction, api=api) + api.execute_sql.assert_called_once_with( + request=self._execute_update_expected_request(database=database), + retry=gapic_v1.method.DEFAULT, + timeout=gapic_v1.method.DEFAULT, + metadata=[("google-cloud-resource-prefix", database.name)], + ) + + self._execute_sql_helper(transaction=transaction, database=database, api=api) + api.execute_streaming_sql.assert_called_once_with( + request=self._execute_sql_expected_request(database=database, begin=False), + retry=gapic_v1.method.DEFAULT, + timeout=gapic_v1.method.DEFAULT, + metadata=[("google-cloud-resource-prefix", database.name)], + ) def test_transaction_should_use_transaction_id_returned_by_first_read(self): database = _Database() session = _Session(database) + api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._read_helper(transaction=transaction, database=database, begin=True) - self._batch_update_helper(transaction=transaction, database=database, begin=False) + self._read_helper(transaction=transaction, api=api) + api.streaming_read.assert_called_once_with( + request=self._read_helper_expected_request(), + metadata=[("google-cloud-resource-prefix", database.name)], + retry=RETRY, + timeout=TIMEOUT, + ) + + self._batch_update_helper(transaction=transaction, database=database, api=api) + api.execute_batch_dml.assert_called_once_with( + request=self._batch_update_expected_request(begin=False), + metadata=[("google-cloud-resource-prefix", database.name)], + ) def test_transaction_should_use_transaction_id_returned_by_first_batch_update(self): database = _Database() + api = database.spanner_api = self._make_spanner_api() + session = _Session(database) + transaction = self._make_one(session) + self._batch_update_helper(transaction=transaction, database=database, api=api) + api.execute_batch_dml.assert_called_once_with( + request=self._batch_update_expected_request(), + metadata=[("google-cloud-resource-prefix", database.name)], + ) + self._read_helper(transaction=transaction, api=api) + api.streaming_read.assert_called_once_with( + request=self._read_helper_expected_request(begin=False), + metadata=[("google-cloud-resource-prefix", database.name)], + retry=RETRY, + timeout=TIMEOUT, + ) + + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_execute_update(self): + database = _Database() + api = database.spanner_api = self._make_spanner_api() + session = _Session(database) + transaction = self._make_one(session) + threads = [] + threads.append(threading.Thread(target=self._execute_update_helper, kwargs={'transaction' : transaction, 'api' : api})) + threads.append(threading.Thread(target=self._execute_update_helper, kwargs={'transaction' : transaction, 'api' : api})) + for thread in threads: + thread.start() + + for thread in threads: + thread.join() + + self._batch_update_helper(transaction=transaction, database=database, api=api) + + api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database), + retry=RETRY, + timeout=TIMEOUT, + metadata=[("google-cloud-resource-prefix", database.name)]) + + api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database, begin=False), + retry=RETRY, + timeout=TIMEOUT, + metadata=[("google-cloud-resource-prefix", database.name)]) + + api.execute_batch_dml.assert_any_call( + request=self._batch_update_expected_request(begin=False), + metadata=[("google-cloud-resource-prefix", database.name)]) + + self.assertEqual(api.execute_sql.call_count, 2) + self.assertEqual(api.execute_batch_dml.call_count, 1) + + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_batch_update(self): + database = _Database() + api = database.spanner_api = self._make_spanner_api() session = _Session(database) transaction = self._make_one(session) - self._batch_update_helper(transaction=transaction, database=database, begin=True) - self._read_helper(transaction=transaction, database=database, begin=False) + threads = [] + threads.append(threading.Thread(target=self._batch_update_helper, kwargs={'transaction' : transaction, 'database' : database, 'api' : api})) + threads.append(threading.Thread(target=self._batch_update_helper, kwargs={'transaction' : transaction, 'database' : database, 'api' : api})) + for thread in threads: + thread.start() + + for thread in threads: + thread.join() + + self._execute_update_helper(transaction=transaction, api=api) + + api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database, begin=False), + retry=RETRY, + timeout=TIMEOUT, + metadata=[("google-cloud-resource-prefix", database.name)]) + + api.execute_batch_dml.assert_any_call( + request=self._batch_update_expected_request(), + metadata=[("google-cloud-resource-prefix", database.name)]) + + api.execute_batch_dml.assert_any_call( + request=self._batch_update_expected_request(begin=False), + metadata=[("google-cloud-resource-prefix", database.name)]) + + self.assertEqual(api.execute_sql.call_count, 1) + self.assertEqual(api.execute_batch_dml.call_count, 2) + + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_read(self): + database = _Database() + api = database.spanner_api = self._make_spanner_api() + session = _Session(database) + transaction = self._make_one(session) + threads = [] + threads.append(threading.Thread(target=self._read_helper, kwargs={'transaction' : transaction, 'api' : api})) + threads.append(threading.Thread(target=self._read_helper, kwargs={'transaction' : transaction, 'api' : api})) + for thread in threads: + thread.start() + + for thread in threads: + thread.join() + + self._execute_update_helper(transaction=transaction, api=api) + + api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database, begin=False), + retry=RETRY, + timeout=TIMEOUT, + metadata=[("google-cloud-resource-prefix", database.name)]) + + api.streaming_read.assert_called_once_with( + request=self._read_helper_expected_request(), + metadata=[("google-cloud-resource-prefix", database.name)], + retry=RETRY, + timeout=TIMEOUT, + ) + + api.streaming_read.assert_called_once_with( + request=self._read_helper_expected_request(begin=False), + metadata=[("google-cloud-resource-prefix", database.name)], + retry=RETRY, + timeout=TIMEOUT, + ) + + self.assertEqual(api.execute_sql.call_count, 1) + self.assertEqual(api.streaming_read.call_count, 2) class _Client(object): def __init__(self): From 42535490b8ac506bd15462af615ee070e1dc556e Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 9 Nov 2022 18:14:14 +0530 Subject: [PATCH 03/19] Added lock for execute sql and read method --- google/cloud/spanner_v1/database.py | 3 +- google/cloud/spanner_v1/snapshot.py | 101 +++++++++++++++----- google/cloud/spanner_v1/transaction.py | 63 ++++-------- tests/unit/test_snapshot.py | 127 ++++++++++++++++++------- tests/unit/test_spanner.py | 68 +++++++++++-- 5 files changed, 247 insertions(+), 115 deletions(-) diff --git a/google/cloud/spanner_v1/database.py b/google/cloud/spanner_v1/database.py index 7d2384beed..5d89ce958d 100644 --- a/google/cloud/spanner_v1/database.py +++ b/google/cloud/spanner_v1/database.py @@ -564,7 +564,6 @@ def execute_pdml(): request = ExecuteSqlRequest( session=session.name, sql=dml, - transaction=txn_selector, params=params_pb, param_types=param_types, query_options=query_options, @@ -575,7 +574,7 @@ def execute_pdml(): metadata=metadata, ) - iterator = _restart_on_unavailable(method, request) + iterator = _restart_on_unavailable(self = None, method = method, request= request, isPdml=True, transactionSelector=txn_selector) result_set = StreamedResultSet(iterator) list(result_set) # consume all partials diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index 7359737f2e..fedb129bcb 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -15,7 +15,7 @@ """Model a set of read-only queries to a database as a snapshot.""" import functools - +import threading from google.protobuf.struct_pb2 import Struct from google.cloud.spanner_v1 import ExecuteSqlRequest from google.cloud.spanner_v1 import ReadRequest @@ -43,7 +43,7 @@ def _restart_on_unavailable( - method, request, trace_name=None, session=None, attributes=None + self, method, request, trace_name=None, session=None, attributes=None, isPdml = False, transactionSelector = None ): """Restart iteration after :exc:`.ServiceUnavailable`. @@ -53,8 +53,15 @@ def _restart_on_unavailable( :type request: proto :param request: request proto to call the method with """ + resume_token = b"" item_buffer = [] + if isPdml is True: + transaction = transactionSelector + else: + transaction = self._make_txn_selector() + + request.transaction = transaction with trace_call(trace_name, session, attributes): iterator = method(request=request) while True: @@ -68,6 +75,11 @@ def _restart_on_unavailable( del item_buffer[:] with trace_call(trace_name, session, attributes): request.resume_token = resume_token + if isPdml is True: + transaction = transactionSelector + else: + transaction = self._make_txn_selector() + request.transaction = transaction iterator = method(request=request) continue except InternalServerError as exc: @@ -80,6 +92,11 @@ def _restart_on_unavailable( del item_buffer[:] with trace_call(trace_name, session, attributes): request.resume_token = resume_token + if isPdml is True: + transaction = transactionSelector + else: + transaction = self._make_txn_selector() + request.transaction = transaction iterator = method(request=request) continue @@ -106,7 +123,7 @@ class _SnapshotBase(_SessionWrapper): _transaction_id = None _read_request_count = 0 _execute_sql_count = 0 - _inline_begin_started = False + _lock = threading.Lock() def _make_txn_selector(self): """Helper for :meth:`read` / :meth:`execute_sql`. @@ -181,13 +198,12 @@ def read( if self._read_request_count > 0: if not self._multi_use: raise ValueError("Cannot re-use single-use snapshot.") - if self._transaction_id is None: + if self._transaction_id is None and self._read_only: raise ValueError("Transaction ID pending.") database = self._session._database api = database.spanner_api metadata = _metadata_with_prefix(database.name) - transaction = self._make_txn_selector() if request_options is None: request_options = RequestOptions() @@ -205,7 +221,6 @@ def read( table=table, columns=columns, key_set=keyset._to_pb(), - transaction=transaction, index=index, limit=limit, partition_token=partition, @@ -220,16 +235,33 @@ def read( ) trace_attributes = {"table_id": table, "columns": columns} - iterator = _restart_on_unavailable( - restart, - request, - "CloudSpanner.ReadOnlyTransaction", - self._session, - trace_attributes, - ) - + + if self._transaction_id is None: + with self._lock: + iterator = _restart_on_unavailable( + self, + restart, + request, + "CloudSpanner.ReadOnlyTransaction", + self._session, + trace_attributes, + ) + self._read_request_count += 1 + if self._multi_use: + return StreamedResultSet(iterator, source=self) + else: + return StreamedResultSet(iterator) + else: + iterator = _restart_on_unavailable( + self, + restart, + request, + "CloudSpanner.ReadOnlyTransaction", + self._session, + trace_attributes, + ) + self._read_request_count += 1 - if self._multi_use: return StreamedResultSet(iterator, source=self) else: @@ -302,7 +334,7 @@ def execute_sql( if self._read_request_count > 0: if not self._multi_use: raise ValueError("Cannot re-use single-use snapshot.") - if self._transaction_id is None: + if self._transaction_id is None and self._read_only: raise ValueError("Transaction ID pending.") if params is not None: @@ -316,7 +348,7 @@ def execute_sql( database = self._session._database metadata = _metadata_with_prefix(database.name) - transaction = self._make_txn_selector() + api = database.spanner_api # Query-level options have higher precedence than client-level and @@ -337,7 +369,6 @@ def execute_sql( request = ExecuteSqlRequest( session=self._session.name, sql=sql, - transaction=transaction, params=params_pb, param_types=param_types, query_mode=query_mode, @@ -355,17 +386,35 @@ def execute_sql( ) trace_attributes = {"db.statement": sql} - iterator = _restart_on_unavailable( - restart, - request, - "CloudSpanner.ReadWriteTransaction", - self._session, - trace_attributes, - ) + + if self._transaction_id is None: + with self._lock: + iterator = _restart_on_unavailable( + self, + restart, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes, + ) + self._read_request_count += 1 + self._execute_sql_count += 1 + if self._multi_use: + return StreamedResultSet(iterator, source=self) + else: + return StreamedResultSet(iterator) + else: + iterator = _restart_on_unavailable( + self, + restart, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes, + ) self._read_request_count += 1 self._execute_sql_count += 1 - if self._multi_use: return StreamedResultSet(iterator, source=self) else: diff --git a/google/cloud/spanner_v1/transaction.py b/google/cloud/spanner_v1/transaction.py index a3c2b11c57..39f6cc1024 100644 --- a/google/cloud/spanner_v1/transaction.py +++ b/google/cloud/spanner_v1/transaction.py @@ -324,26 +324,15 @@ def execute_update( if self._transaction_id is None: with self._lock: - if self._inline_begin_started is False: - response = self._execute_request( - method, - request, - "CloudSpanner.ReadWriteTransaction", - self._session, - trace_attributes - ) - - if self._transaction_id is None and response.metadata.transaction is not None: - self._transaction_id = response.metadata.transaction.id - self._inline_begin_started = True - else: - response = self._execute_request( - method, - request, - "CloudSpanner.ReadWriteTransaction", - self._session, - trace_attributes - ) + response = self._execute_request( + method, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes + ) + if self._transaction_id is None and response.metadata.transaction is not None: + self._transaction_id = response.metadata.transaction.id else: response = self._execute_request( method, @@ -399,7 +388,6 @@ def batch_update(self, statements, request_options=None): database = self._session._database metadata = _metadata_with_prefix(database.name) - transaction = self._make_txn_selector() api = database.spanner_api seqno, self._execute_sql_count = ( @@ -432,28 +420,17 @@ def batch_update(self, statements, request_options=None): if self._transaction_id is None: with self._lock: - if self._inline_begin_started is False: - response = self._execute_request( - method, - request, - "CloudSpanner.DMLTransaction", - self._session, - trace_attributes - ) - - for result_set in response.result_sets: - if self._transaction_id is None and result_set.metadata.transaction is not None: - self._transaction_id = result_set.metadata.transaction.id - - self._inline_begin_started = True - else: - response = self._execute_request( - method, - request, - "CloudSpanner.DMLTransaction", - self._session, - trace_attributes - ) + response = self._execute_request( + method, + request, + "CloudSpanner.DMLTransaction", + self._session, + trace_attributes + ) + + for result_set in response.result_sets: + if self._transaction_id is None and result_set.metadata.transaction is not None: + self._transaction_id = result_set.metadata.transaction.id else: response = self._execute_request( method, diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index 5b515f1bbb..711ff7f479 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -46,15 +46,32 @@ "db.instance": "testing", "net.host.name": "spanner.googleapis.com", } + +class Test_SnapshotBase(OpenTelemetryBase): + + PROJECT_ID = "project-id" + INSTANCE_ID = "instance-id" + INSTANCE_NAME = "projects/" + PROJECT_ID + "/instances/" + INSTANCE_ID + DATABASE_ID = "database-id" + DATABASE_NAME = INSTANCE_NAME + "/databases/" + DATABASE_ID + SESSION_ID = "session-id" + SESSION_NAME = DATABASE_NAME + "/sessions/" + SESSION_ID + def _getTargetClass(self): + from google.cloud.spanner_v1.snapshot import _SnapshotBase + + return _SnapshotBase + + def _make_one(self, session): + return self._getTargetClass()(session) -class Test_restart_on_unavailable(OpenTelemetryBase): def _call_fut( - self, restart, request, span_name=None, session=None, attributes=None + self, derived, restart, request, span_name=None, session=None, attributes=None ): from google.cloud.spanner_v1.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1.snapshot import Snapshot - return _restart_on_unavailable(restart, request, span_name, session, attributes) + return _restart_on_unavailable(derived, restart, request, span_name, session, attributes) def _make_item(self, value, resume_token=b""): return mock.Mock( @@ -65,7 +82,11 @@ def test_iteration_w_empty_raw(self): raw = _MockIterator() request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], return_value=raw) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), []) restart.assert_called_once_with(request=request) self.assertNoSpans() @@ -75,7 +96,11 @@ def test_iteration_w_non_empty_raw(self): raw = _MockIterator(*ITEMS) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], return_value=raw) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(ITEMS)) restart.assert_called_once_with(request=request) self.assertNoSpans() @@ -90,7 +115,11 @@ def test_iteration_w_raw_w_resume_tken(self): raw = _MockIterator(*ITEMS) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], return_value=raw) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(ITEMS)) restart.assert_called_once_with(request=request) self.assertNoSpans() @@ -107,7 +136,11 @@ def test_iteration_w_raw_raising_unavailable_no_token(self): after = _MockIterator(*ITEMS) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(ITEMS)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, b"") @@ -130,7 +163,11 @@ def test_iteration_w_raw_raising_retryable_internal_error_no_token(self): after = _MockIterator(*ITEMS) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(ITEMS)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, b"") @@ -148,7 +185,11 @@ def test_iteration_w_raw_raising_non_retryable_internal_error_no_token(self): after = _MockIterator(*ITEMS) request = mock.Mock(spec=["resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) with self.assertRaises(InternalServerError): list(resumable) restart.assert_called_once_with(request=request) @@ -166,7 +207,11 @@ def test_iteration_w_raw_raising_unavailable(self): after = _MockIterator(*LAST) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(FIRST + LAST)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, RESUME_TOKEN) @@ -188,7 +233,11 @@ def test_iteration_w_raw_raising_retryable_internal_error(self): after = _MockIterator(*LAST) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(FIRST + LAST)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, RESUME_TOKEN) @@ -206,7 +255,11 @@ def test_iteration_w_raw_raising_non_retryable_internal_error(self): after = _MockIterator(*LAST) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) with self.assertRaises(InternalServerError): list(resumable) restart.assert_called_once_with(request=request) @@ -223,7 +276,11 @@ def test_iteration_w_raw_raising_unavailable_after_token(self): after = _MockIterator(*SECOND) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(FIRST + SECOND)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, RESUME_TOKEN) @@ -244,7 +301,11 @@ def test_iteration_w_raw_raising_retryable_internal_error_after_token(self): after = _MockIterator(*SECOND) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(FIRST + SECOND)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, RESUME_TOKEN) @@ -261,7 +322,11 @@ def test_iteration_w_raw_raising_non_retryable_internal_error_after_token(self): after = _MockIterator(*SECOND) request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) - resumable = self._call_fut(restart, request) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request) with self.assertRaises(InternalServerError): list(resumable) restart.assert_called_once_with(request=request) @@ -273,8 +338,12 @@ def test_iteration_w_span_creation(self): raw = _MockIterator() request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], return_value=raw) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) resumable = self._call_fut( - restart, request, name, _Session(_Database()), extra_atts + derived, restart, request, name, _Session(_Database()), extra_atts ) self.assertEqual(list(resumable), []) self.assertSpanAttributes(name, attributes=dict(BASE_ATTRIBUTES, test_att=1)) @@ -293,7 +362,11 @@ def test_iteration_w_multiple_span_creation(self): request = mock.Mock(test="test", spec=["test", "resume_token"]) restart = mock.Mock(spec=[], side_effect=[before, after]) name = "TestSpan" - resumable = self._call_fut(restart, request, name, _Session(_Database())) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + resumable = self._call_fut(derived, restart, request, name, _Session(_Database())) self.assertEqual(list(resumable), list(FIRST + LAST)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, RESUME_TOKEN) @@ -312,25 +385,6 @@ def test_iteration_w_multiple_span_creation(self): }, ) - -class Test_SnapshotBase(OpenTelemetryBase): - - PROJECT_ID = "project-id" - INSTANCE_ID = "instance-id" - INSTANCE_NAME = "projects/" + PROJECT_ID + "/instances/" + INSTANCE_ID - DATABASE_ID = "database-id" - DATABASE_NAME = INSTANCE_NAME + "/databases/" + DATABASE_ID - SESSION_ID = "session-id" - SESSION_NAME = DATABASE_NAME + "/sessions/" + SESSION_ID - - def _getTargetClass(self): - from google.cloud.spanner_v1.snapshot import _SnapshotBase - - return _SnapshotBase - - def _make_one(self, session): - return self._getTargetClass()(session) - def _makeDerived(self, session): class _Derived(self._getTargetClass()): @@ -876,7 +930,6 @@ def _partition_read_helper( derived._multi_use = multi_use if w_txn: derived._transaction_id = TXN_ID - tokens = list( derived.partition_read( TABLE_NAME, diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index c46b27cad1..7b4a42691d 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -187,7 +187,6 @@ def _execute_update_expected_request(self, database, query_options=None, begin=T def _execute_sql_helper( self, transaction, - database, api, count=0, partition=None, @@ -256,6 +255,8 @@ def _execute_sql_expected_request(self, database, partition=None, query_options= expected_query_options, query_options ) + expected_request_options = REQUEST_OPTIONS + expected_request_options.transaction_tag = None expected_request = ExecuteSqlRequest( session=self.SESSION_NAME, sql=SQL_QUERY_WITH_PARAM, @@ -264,7 +265,7 @@ def _execute_sql_expected_request(self, database, partition=None, query_options= param_types=PARAM_TYPES, query_mode=MODE, query_options=expected_query_options, - request_options=REQUEST_OPTIONS, + request_options=expected_request_options, partition_token=partition, seqno=sql_count, ) @@ -456,7 +457,7 @@ def test_transaction_should_include_begin_with_first_query(self): session = _Session(database) api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._execute_sql_helper(transaction=transaction, database=database, api=api) + self._execute_sql_helper(transaction=transaction, api=api) api.execute_streaming_sql.assert_called_once_with( request=self._execute_sql_expected_request(database=database), @@ -495,7 +496,7 @@ def test_transaction_should_use_transaction_id_returned_by_first_query(self): session = _Session(database) api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._execute_sql_helper(transaction=transaction, database=database, api=api) + self._execute_sql_helper(transaction=transaction, api=api) api.execute_streaming_sql.assert_called_once_with( request=self._execute_sql_expected_request(database=database), retry=gapic_v1.method.DEFAULT, @@ -524,7 +525,7 @@ def test_transaction_should_use_transaction_id_returned_by_first_update(self): metadata=[("google-cloud-resource-prefix", database.name)], ) - self._execute_sql_helper(transaction=transaction, database=database, api=api) + self._execute_sql_helper(transaction=transaction, api=api) api.execute_streaming_sql.assert_called_once_with( request=self._execute_sql_expected_request(database=database, begin=False), retry=gapic_v1.method.DEFAULT, @@ -650,28 +651,81 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ self._execute_update_helper(transaction=transaction, api=api) + begin_read_write = 0 + for call in api.mock_calls: + if "read_write" in call.__str__(): + begin_read_write+=1 + + self.assertEqual(begin_read_write, 1) api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database, begin=False), retry=RETRY, timeout=TIMEOUT, metadata=[("google-cloud-resource-prefix", database.name)]) - api.streaming_read.assert_called_once_with( + api.streaming_read.assert_any_call( request=self._read_helper_expected_request(), metadata=[("google-cloud-resource-prefix", database.name)], retry=RETRY, timeout=TIMEOUT, ) - api.streaming_read.assert_called_once_with( + api.streaming_read.assert_any_call( request=self._read_helper_expected_request(begin=False), metadata=[("google-cloud-resource-prefix", database.name)], retry=RETRY, timeout=TIMEOUT, ) + self.assertEqual(api.execute_sql.call_count, 1) self.assertEqual(api.streaming_read.call_count, 2) + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_query(self): + database = _Database() + api = database.spanner_api = self._make_spanner_api() + session = _Session(database) + transaction = self._make_one(session) + threads = [] + threads.append(threading.Thread(target=self._execute_sql_helper, kwargs={'transaction' : transaction, 'api' : api})) + threads.append(threading.Thread(target=self._execute_sql_helper, kwargs={'transaction' : transaction, 'api' : api})) + for thread in threads: + thread.start() + + for thread in threads: + thread.join() + + self._execute_update_helper(transaction=transaction, api=api) + + begin_read_write = 0 + for call in api.mock_calls: + if "read_write" in call.__str__(): + begin_read_write+=1 + + self.assertEqual(begin_read_write, 1) + api.execute_sql.assert_any_call( + request=self._execute_update_expected_request(database, begin=False), + retry=RETRY, + timeout=TIMEOUT, + metadata=[("google-cloud-resource-prefix", database.name)]) + + req = self._execute_sql_expected_request(database) + api.execute_streaming_sql.assert_any_call( + request=req, + metadata=[("google-cloud-resource-prefix", database.name)], + retry=RETRY, + timeout=TIMEOUT, + ) + + api.execute_streaming_sql.assert_any_call( + request=self._execute_sql_expected_request(database, begin=False), + metadata=[("google-cloud-resource-prefix", database.name)], + retry=RETRY, + timeout=TIMEOUT, + ) + + self.assertEqual(api.execute_sql.call_count, 1) + self.assertEqual(api.execute_streaming_sql.call_count, 2) + class _Client(object): def __init__(self): from google.cloud.spanner_v1 import ExecuteSqlRequest From a64d88236bda2eab096736db63ca116e8bd16dd1 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 23 Nov 2022 11:18:15 +0530 Subject: [PATCH 04/19] fix: lint fix and testcases --- google/cloud/spanner_v1/database.py | 8 +- google/cloud/spanner_v1/session.py | 2 +- google/cloud/spanner_v1/snapshot.py | 17 +- google/cloud/spanner_v1/transaction.py | 64 ++++---- tests/system/test_session_api.py | 3 +- tests/unit/test_session.py | 1 - tests/unit/test_snapshot.py | 11 +- tests/unit/test_spanner.py | 215 +++++++++++++++++-------- 8 files changed, 213 insertions(+), 108 deletions(-) diff --git a/google/cloud/spanner_v1/database.py b/google/cloud/spanner_v1/database.py index 5d89ce958d..24be4b2bf1 100644 --- a/google/cloud/spanner_v1/database.py +++ b/google/cloud/spanner_v1/database.py @@ -574,7 +574,13 @@ def execute_pdml(): metadata=metadata, ) - iterator = _restart_on_unavailable(self = None, method = method, request= request, isPdml=True, transactionSelector=txn_selector) + iterator = _restart_on_unavailable( + self=None, + method=method, + request=request, + isPdml=True, + transactionSelector=txn_selector, + ) result_set = StreamedResultSet(iterator) list(result_set) # consume all partials diff --git a/google/cloud/spanner_v1/session.py b/google/cloud/spanner_v1/session.py index 86a9bed0e8..06f41b9495 100644 --- a/google/cloud/spanner_v1/session.py +++ b/google/cloud/spanner_v1/session.py @@ -352,7 +352,7 @@ def run_in_transaction(self, func, *args, **kw): txn.transaction_tag = transaction_tag else: txn = self._transaction - + try: attempts += 1 return_value = func(txn, *args, **kw) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index fedb129bcb..9d2b219d01 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -43,7 +43,14 @@ def _restart_on_unavailable( - self, method, request, trace_name=None, session=None, attributes=None, isPdml = False, transactionSelector = None + self, + method, + request, + trace_name=None, + session=None, + attributes=None, + isPdml=False, + transactionSelector=None, ): """Restart iteration after :exc:`.ServiceUnavailable`. @@ -53,7 +60,7 @@ def _restart_on_unavailable( :type request: proto :param request: request proto to call the method with """ - + resume_token = b"" item_buffer = [] if isPdml is True: @@ -235,7 +242,7 @@ def read( ) trace_attributes = {"table_id": table, "columns": columns} - + if self._transaction_id is None: with self._lock: iterator = _restart_on_unavailable( @@ -260,7 +267,7 @@ def read( self._session, trace_attributes, ) - + self._read_request_count += 1 if self._multi_use: return StreamedResultSet(iterator, source=self) @@ -348,7 +355,7 @@ def execute_sql( database = self._session._database metadata = _metadata_with_prefix(database.name) - + api = database.spanner_api # Query-level options have higher precedence than client-level and diff --git a/google/cloud/spanner_v1/transaction.py b/google/cloud/spanner_v1/transaction.py index 39f6cc1024..873b988440 100644 --- a/google/cloud/spanner_v1/transaction.py +++ b/google/cloud/spanner_v1/transaction.py @@ -63,7 +63,7 @@ def _check_state(self): :raises: :exc:`ValueError` if the object's state is invalid for making API requests. """ - + if self.committed is not None: raise ValueError("Transaction is already committed") @@ -78,9 +78,11 @@ def _make_txn_selector(self): :returns: a selector configured for read-write transaction semantics. """ self._check_state() - + if self._transaction_id is None: - return TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + return TransactionSelector( + begin=TransactionOptions(read_write=TransactionOptions.ReadWrite()) + ) else: return TransactionSelector(id=self._transaction_id) @@ -91,7 +93,7 @@ def _execute_request( request.transaction = transaction with trace_call(trace_name, session, attributes): response = method(request=request) - + return response def begin(self): @@ -282,7 +284,7 @@ def execute_update( params_pb = self._make_params_pb(params, param_types) database = self._session._database metadata = _metadata_with_prefix(database.name) - + api = database.spanner_api seqno, self._execute_sql_count = ( @@ -325,21 +327,24 @@ def execute_update( if self._transaction_id is None: with self._lock: response = self._execute_request( - method, - request, - "CloudSpanner.ReadWriteTransaction", - self._session, - trace_attributes + method, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes, ) - if self._transaction_id is None and response.metadata.transaction is not None: + if ( + self._transaction_id is None + and response.metadata.transaction is not None + ): self._transaction_id = response.metadata.transaction.id else: response = self._execute_request( - method, - request, - "CloudSpanner.ReadWriteTransaction", - self._session, - trace_attributes + method, + request, + "CloudSpanner.ReadWriteTransaction", + self._session, + trace_attributes, ) return response.stats.row_count_exact @@ -421,23 +426,26 @@ def batch_update(self, statements, request_options=None): if self._transaction_id is None: with self._lock: response = self._execute_request( - method, - request, - "CloudSpanner.DMLTransaction", - self._session, - trace_attributes + method, + request, + "CloudSpanner.DMLTransaction", + self._session, + trace_attributes, ) - + for result_set in response.result_sets: - if self._transaction_id is None and result_set.metadata.transaction is not None: + if ( + self._transaction_id is None + and result_set.metadata.transaction is not None + ): self._transaction_id = result_set.metadata.transaction.id else: response = self._execute_request( - method, - request, - "CloudSpanner.DMLTransaction", - self._session, - trace_attributes + method, + request, + "CloudSpanner.DMLTransaction", + self._session, + trace_attributes, ) row_counts = [ diff --git a/tests/system/test_session_api.py b/tests/system/test_session_api.py index aedcbcaa55..4872bcae8e 100644 --- a/tests/system/test_session_api.py +++ b/tests/system/test_session_api.py @@ -1088,11 +1088,10 @@ def unit_of_work(transaction): session.run_in_transaction(unit_of_work) span_list = ot_exporter.get_finished_spans() - assert len(span_list) == 6 + assert len(span_list) == 5 expected_span_names = [ "CloudSpanner.CreateSession", "CloudSpanner.Commit", - "CloudSpanner.BeginTransaction", "CloudSpanner.DMLTransaction", "CloudSpanner.Commit", "Test Span", diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 97195734aa..aa919071d8 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -725,7 +725,6 @@ def unit_of_work(txn, *args, **kw): self.assertEqual(args, ()) self.assertEqual(kw, {}) - def test_run_in_transaction_callback_raises_non_abort_rpc_error(self): from google.api_core.exceptions import Cancelled from google.cloud.spanner_v1 import ( diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index 711ff7f479..7a93651430 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -46,7 +46,8 @@ "db.instance": "testing", "net.host.name": "spanner.googleapis.com", } - + + class Test_SnapshotBase(OpenTelemetryBase): PROJECT_ID = "project-id" @@ -71,7 +72,9 @@ def _call_fut( from google.cloud.spanner_v1.snapshot import _restart_on_unavailable from google.cloud.spanner_v1.snapshot import Snapshot - return _restart_on_unavailable(derived, restart, request, span_name, session, attributes) + return _restart_on_unavailable( + derived, restart, request, span_name, session, attributes + ) def _make_item(self, value, resume_token=b""): return mock.Mock( @@ -366,7 +369,9 @@ def test_iteration_w_multiple_span_creation(self): database.spanner_api = self._make_spanner_api() session = _Session(database) derived = self._makeDerived(session) - resumable = self._call_fut(derived, restart, request, name, _Session(_Database())) + resumable = self._call_fut( + derived, restart, request, name, _Session(_Database()) + ) self.assertEqual(list(resumable), list(FIRST + LAST)) self.assertEqual(len(restart.mock_calls), 2) self.assertEqual(request.resume_token, RESUME_TOKEN) diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index 7b4a42691d..339ae9cc94 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -31,7 +31,7 @@ TransactionSelector, ExecuteBatchDmlRequest, ExecuteBatchDmlResponse, - param_types + param_types, ) from google.cloud.spanner_v1.types import transaction as transaction_type from google.cloud.spanner_v1.keyset import KeySet @@ -73,8 +73,8 @@ INDEX = "email-address-index" LIMIT = 20 MODE = 2 -RETRY=gapic_v1.method.DEFAULT -TIMEOUT=gapic_v1.method.DEFAULT +RETRY = gapic_v1.method.DEFAULT +TIMEOUT = gapic_v1.method.DEFAULT REQUEST_OPTIONS = RequestOptions() insert_dml = "INSERT INTO table(pkey, desc) VALUES (%pkey, %desc)" insert_params = {"pkey": 12345, "desc": "DESCRIPTION"} @@ -88,6 +88,7 @@ delete_dml, ] + class TestTransaction(OpenTelemetryBase): PROJECT_ID = "project-id" @@ -121,7 +122,7 @@ def _make_spanner_api(self): from google.cloud.spanner_v1 import SpannerClient return mock.create_autospec(SpannerClient, instance=True) - + def _execute_update_helper( self, transaction, @@ -130,13 +131,11 @@ def _execute_update_helper( query_options=None, ): stats_pb = ResultSetStats(row_count_exact=1) - - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID - ) + + transaction_pb = transaction_type.Transaction(id=self.TRANSACTION_ID) metadata_pb = ResultSetMetadata(transaction=transaction_pb) api.execute_sql.return_value = ResultSet(stats=stats_pb, metadata=metadata_pb) - + transaction.transaction_tag = self.TRANSACTION_TAG transaction._execute_sql_count = count @@ -152,9 +151,13 @@ def _execute_update_helper( ) self.assertEqual(row_count, count + 1) - def _execute_update_expected_request(self, database, query_options=None, begin=True, count=0): + def _execute_update_expected_request( + self, database, query_options=None, begin=True, count=0 + ): if begin is True: - expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + expected_transaction = TransactionSelector( + begin=TransactionOptions(read_write=TransactionOptions.ReadWrite()) + ) else: expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) @@ -181,7 +184,7 @@ def _execute_update_expected_request(self, database, query_options=None, begin=T request_options=expected_request_options, seqno=count, ) - + return expected_request def _execute_sql_helper( @@ -202,10 +205,10 @@ def _execute_sql_helper( StructType.Field(name="age", type_=Type(code=TypeCode.INT64)), ] ) - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID + transaction_pb = transaction_type.Transaction(id=self.TRANSACTION_ID) + metadata_pb = ResultSetMetadata( + row_type=struct_type_pb, transaction=transaction_pb ) - metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) stats_pb = ResultSetStats( query_stats=Struct(fields={"rows_returned": _make_value_pb(2)}) ) @@ -219,7 +222,7 @@ def _execute_sql_helper( api.execute_streaming_sql.return_value = iterator transaction._execute_sql_count = sql_count transaction._read_request_count = count - + result_set = transaction.execute_sql( SQL_QUERY_WITH_PARAM, PARAMS, @@ -239,9 +242,13 @@ def _execute_sql_helper( self.assertEqual(result_set.stats, stats_pb) self.assertEqual(transaction._execute_sql_count, sql_count + 1) - def _execute_sql_expected_request(self, database, partition=None, query_options=None, begin=True, sql_count=0): + def _execute_sql_expected_request( + self, database, partition=None, query_options=None, begin=True, sql_count=0 + ): if begin is True: - expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + expected_transaction = TransactionSelector( + begin=TransactionOptions(read_write=TransactionOptions.ReadWrite()) + ) else: expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) @@ -269,7 +276,7 @@ def _execute_sql_expected_request(self, database, partition=None, query_options= partition_token=partition, seqno=sql_count, ) - + return expected_request def _read_helper( @@ -288,11 +295,11 @@ def _read_helper( ] ) - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID + transaction_pb = transaction_type.Transaction(id=self.TRANSACTION_ID) + metadata_pb = ResultSetMetadata( + row_type=struct_type_pb, transaction=transaction_pb ) - metadata_pb = ResultSetMetadata(row_type=struct_type_pb,transaction=transaction_pb) - + stats_pb = ResultSetStats( query_stats=Struct(fields={"rows_returned": _make_value_pb(2)}) ) @@ -302,10 +309,10 @@ def _read_helper( ] for i in range(len(result_sets)): result_sets[i].values.extend(VALUE_PBS[i]) - + api.streaming_read.return_value = _MockIterator(*result_sets) transaction._read_request_count = count - + if partition is not None: # 'limit' and 'partition' incompatible result_set = transaction.read( TABLE_NAME, @@ -338,9 +345,11 @@ def _read_helper( self.assertEqual(result_set.stats, stats_pb) def _read_helper_expected_request(self, partition=None, begin=True, count=0): - + if begin is True: - expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + expected_transaction = TransactionSelector( + begin=TransactionOptions(read_write=TransactionOptions.ReadWrite()) + ) else: expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) @@ -364,10 +373,17 @@ def _read_helper_expected_request(self, partition=None, begin=True, count=0): partition_token=partition, request_options=expected_request_options, ) - + return expected_request - def _batch_update_helper(self, transaction, database, api, error_after=None, count=0,): + def _batch_update_helper( + self, + transaction, + database, + api, + error_after=None, + count=0, + ): from google.rpc.status_pb2 import Status stats_pbs = [ @@ -381,11 +397,11 @@ def _batch_update_helper(self, transaction, database, api, error_after=None, cou else: expected_status = Status(code=200) expected_row_counts = [stats.row_count_exact for stats in stats_pbs] - transaction_pb = transaction_type.Transaction( - id = self.TRANSACTION_ID - ) + transaction_pb = transaction_type.Transaction(id=self.TRANSACTION_ID) metadata_pb = ResultSetMetadata(transaction=transaction_pb) - result_sets_pb = [ResultSet(stats=stats_pb, metadata= metadata_pb) for stats_pb in stats_pbs] + result_sets_pb = [ + ResultSet(stats=stats_pb, metadata=metadata_pb) for stats_pb in stats_pbs + ] response = ExecuteBatchDmlResponse( status=expected_status, @@ -406,10 +422,12 @@ def _batch_update_helper(self, transaction, database, api, error_after=None, cou def _batch_update_expected_request(self, begin=True, count=0): if begin is True: - expected_transaction = TransactionSelector(begin=TransactionOptions(read_write=TransactionOptions.ReadWrite())) + expected_transaction = TransactionSelector( + begin=TransactionOptions(read_write=TransactionOptions.ReadWrite()) + ) else: expected_transaction = TransactionSelector(id=self.TRANSACTION_ID) - + expected_insert_params = Struct( fields={ key: _make_value_pb(value) for (key, value) in insert_params.items() @@ -435,7 +453,7 @@ def _batch_update_expected_request(self, begin=True, count=0): seqno=count, request_options=expected_request_options, ) - + return expected_request def test_transaction_should_include_begin_with_first_update(self): @@ -506,7 +524,9 @@ def test_transaction_should_use_transaction_id_returned_by_first_query(self): self._execute_update_helper(transaction=transaction, api=api) api.execute_sql.assert_called_once_with( - request=self._execute_update_expected_request(database=database, begin=False), + request=self._execute_update_expected_request( + database=database, begin=False + ), retry=gapic_v1.method.DEFAULT, timeout=gapic_v1.method.DEFAULT, metadata=[("google-cloud-resource-prefix", database.name)], @@ -551,7 +571,7 @@ def test_transaction_should_use_transaction_id_returned_by_first_read(self): request=self._batch_update_expected_request(begin=False), metadata=[("google-cloud-resource-prefix", database.name)], ) - + def test_transaction_should_use_transaction_id_returned_by_first_batch_update(self): database = _Database() api = database.spanner_api = self._make_spanner_api() @@ -570,14 +590,26 @@ def test_transaction_should_use_transaction_id_returned_by_first_batch_update(se timeout=TIMEOUT, ) - def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_execute_update(self): + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_execute_update( + self, + ): database = _Database() api = database.spanner_api = self._make_spanner_api() session = _Session(database) transaction = self._make_one(session) threads = [] - threads.append(threading.Thread(target=self._execute_update_helper, kwargs={'transaction' : transaction, 'api' : api})) - threads.append(threading.Thread(target=self._execute_update_helper, kwargs={'transaction' : transaction, 'api' : api})) + threads.append( + threading.Thread( + target=self._execute_update_helper, + kwargs={"transaction": transaction, "api": api}, + ) + ) + threads.append( + threading.Thread( + target=self._execute_update_helper, + kwargs={"transaction": transaction, "api": api}, + ) + ) for thread in threads: thread.start() @@ -586,31 +618,48 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ self._batch_update_helper(transaction=transaction, database=database, api=api) - api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database), + api.execute_sql.assert_any_call( + request=self._execute_update_expected_request(database), retry=RETRY, timeout=TIMEOUT, - metadata=[("google-cloud-resource-prefix", database.name)]) + metadata=[("google-cloud-resource-prefix", database.name)], + ) - api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database, begin=False), + api.execute_sql.assert_any_call( + request=self._execute_update_expected_request(database, begin=False), retry=RETRY, timeout=TIMEOUT, - metadata=[("google-cloud-resource-prefix", database.name)]) + metadata=[("google-cloud-resource-prefix", database.name)], + ) api.execute_batch_dml.assert_any_call( request=self._batch_update_expected_request(begin=False), - metadata=[("google-cloud-resource-prefix", database.name)]) + metadata=[("google-cloud-resource-prefix", database.name)], + ) self.assertEqual(api.execute_sql.call_count, 2) self.assertEqual(api.execute_batch_dml.call_count, 1) - def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_batch_update(self): + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_batch_update( + self, + ): database = _Database() api = database.spanner_api = self._make_spanner_api() session = _Session(database) transaction = self._make_one(session) threads = [] - threads.append(threading.Thread(target=self._batch_update_helper, kwargs={'transaction' : transaction, 'database' : database, 'api' : api})) - threads.append(threading.Thread(target=self._batch_update_helper, kwargs={'transaction' : transaction, 'database' : database, 'api' : api})) + threads.append( + threading.Thread( + target=self._batch_update_helper, + kwargs={"transaction": transaction, "database": database, "api": api}, + ) + ) + threads.append( + threading.Thread( + target=self._batch_update_helper, + kwargs={"transaction": transaction, "database": database, "api": api}, + ) + ) for thread in threads: thread.start() @@ -619,30 +668,46 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ self._execute_update_helper(transaction=transaction, api=api) - api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database, begin=False), + api.execute_sql.assert_any_call( + request=self._execute_update_expected_request(database, begin=False), retry=RETRY, timeout=TIMEOUT, - metadata=[("google-cloud-resource-prefix", database.name)]) + metadata=[("google-cloud-resource-prefix", database.name)], + ) api.execute_batch_dml.assert_any_call( request=self._batch_update_expected_request(), - metadata=[("google-cloud-resource-prefix", database.name)]) + metadata=[("google-cloud-resource-prefix", database.name)], + ) api.execute_batch_dml.assert_any_call( request=self._batch_update_expected_request(begin=False), - metadata=[("google-cloud-resource-prefix", database.name)]) + metadata=[("google-cloud-resource-prefix", database.name)], + ) self.assertEqual(api.execute_sql.call_count, 1) self.assertEqual(api.execute_batch_dml.call_count, 2) - def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_read(self): + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_read( + self, + ): database = _Database() api = database.spanner_api = self._make_spanner_api() session = _Session(database) transaction = self._make_one(session) threads = [] - threads.append(threading.Thread(target=self._read_helper, kwargs={'transaction' : transaction, 'api' : api})) - threads.append(threading.Thread(target=self._read_helper, kwargs={'transaction' : transaction, 'api' : api})) + threads.append( + threading.Thread( + target=self._read_helper, + kwargs={"transaction": transaction, "api": api}, + ) + ) + threads.append( + threading.Thread( + target=self._read_helper, + kwargs={"transaction": transaction, "api": api}, + ) + ) for thread in threads: thread.start() @@ -654,13 +719,15 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ begin_read_write = 0 for call in api.mock_calls: if "read_write" in call.__str__(): - begin_read_write+=1 + begin_read_write += 1 self.assertEqual(begin_read_write, 1) - api.execute_sql.assert_any_call(request=self._execute_update_expected_request(database, begin=False), + api.execute_sql.assert_any_call( + request=self._execute_update_expected_request(database, begin=False), retry=RETRY, timeout=TIMEOUT, - metadata=[("google-cloud-resource-prefix", database.name)]) + metadata=[("google-cloud-resource-prefix", database.name)], + ) api.streaming_read.assert_any_call( request=self._read_helper_expected_request(), @@ -676,18 +743,29 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ timeout=TIMEOUT, ) - self.assertEqual(api.execute_sql.call_count, 1) self.assertEqual(api.streaming_read.call_count, 2) - def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_query(self): + def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_query( + self, + ): database = _Database() api = database.spanner_api = self._make_spanner_api() session = _Session(database) transaction = self._make_one(session) threads = [] - threads.append(threading.Thread(target=self._execute_sql_helper, kwargs={'transaction' : transaction, 'api' : api})) - threads.append(threading.Thread(target=self._execute_sql_helper, kwargs={'transaction' : transaction, 'api' : api})) + threads.append( + threading.Thread( + target=self._execute_sql_helper, + kwargs={"transaction": transaction, "api": api}, + ) + ) + threads.append( + threading.Thread( + target=self._execute_sql_helper, + kwargs={"transaction": transaction, "api": api}, + ) + ) for thread in threads: thread.start() @@ -699,15 +777,16 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ begin_read_write = 0 for call in api.mock_calls: if "read_write" in call.__str__(): - begin_read_write+=1 + begin_read_write += 1 self.assertEqual(begin_read_write, 1) api.execute_sql.assert_any_call( request=self._execute_update_expected_request(database, begin=False), retry=RETRY, timeout=TIMEOUT, - metadata=[("google-cloud-resource-prefix", database.name)]) - + metadata=[("google-cloud-resource-prefix", database.name)], + ) + req = self._execute_sql_expected_request(database) api.execute_streaming_sql.assert_any_call( request=req, @@ -722,10 +801,11 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ retry=RETRY, timeout=TIMEOUT, ) - + self.assertEqual(api.execute_sql.call_count, 1) self.assertEqual(api.execute_streaming_sql.call_count, 2) + class _Client(object): def __init__(self): from google.cloud.spanner_v1 import ExecuteSqlRequest @@ -752,6 +832,7 @@ def __init__(self, database=None, name=TestTransaction.SESSION_NAME): self._database = database self.name = name + class _MockIterator(object): def __init__(self, *values, **kw): self._iter_values = iter(values) From f7963a798e8d93d5d58e966c6109e6e73e0c3b49 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 23 Nov 2022 17:56:10 +0530 Subject: [PATCH 05/19] fix: lint --- tests/unit/test_session.py | 2 -- tests/unit/test_snapshot.py | 1 - tests/unit/test_spanner.py | 4 +--- 3 files changed, 1 insertion(+), 6 deletions(-) diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index aa919071d8..d36326ac49 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -683,7 +683,6 @@ def test_transaction_w_existing_txn(self): def test_run_in_transaction_callback_raises_non_gax_error(self): from google.cloud.spanner_v1 import ( Transaction as TransactionPB, - TransactionOptions, ) from google.cloud.spanner_v1.transaction import Transaction @@ -729,7 +728,6 @@ def test_run_in_transaction_callback_raises_non_abort_rpc_error(self): from google.api_core.exceptions import Cancelled from google.cloud.spanner_v1 import ( Transaction as TransactionPB, - TransactionOptions, ) from google.cloud.spanner_v1.transaction import Transaction diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index 7a93651430..dac47f5241 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -70,7 +70,6 @@ def _call_fut( self, derived, restart, request, span_name=None, session=None, attributes=None ): from google.cloud.spanner_v1.snapshot import _restart_on_unavailable - from google.cloud.spanner_v1.snapshot import Snapshot return _restart_on_unavailable( derived, restart, request, span_name, session, attributes diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index 339ae9cc94..4dcdf8190a 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -13,7 +13,6 @@ # limitations under the License. -from dataclasses import fields import threading from google.protobuf.struct_pb2 import Struct from google.cloud.spanner_v1 import ( @@ -43,10 +42,9 @@ import mock -from google.api_core.retry import Retry from google.api_core import gapic_v1 -from tests._helpers import OpenTelemetryBase, StatusCode +from tests._helpers import OpenTelemetryBase TABLE_NAME = "citizens" COLUMNS = ["email", "first_name", "last_name", "age"] From 4d0c0d4630eb05d64317c5320ca561ab69b76f3c Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 30 Nov 2022 11:11:43 +0530 Subject: [PATCH 06/19] fix: Set transction id along with resume token --- google/cloud/spanner_v1/snapshot.py | 2 ++ tests/unit/test_snapshot.py | 39 ++++++++++++++++++++++++++--- 2 files changed, 38 insertions(+), 3 deletions(-) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index 9d2b219d01..bbad11d2ac 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -77,6 +77,8 @@ def _restart_on_unavailable( item_buffer.append(item) if item.resume_token: resume_token = item.resume_token + if self._transaction_id is None and item.metadata is not None and item.metadata.transaction is not None and item.metadata.transaction.id is not None: + self._transaction_id = item.metadata.transaction.id break except ServiceUnavailable: del item_buffer[:] diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index dac47f5241..410ef72f13 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -75,11 +75,11 @@ def _call_fut( derived, restart, request, span_name, session, attributes ) - def _make_item(self, value, resume_token=b""): + def _make_item(self, value, resume_token=b"", metadata=None): return mock.Mock( - value=value, resume_token=resume_token, spec=["value", "resume_token"] + value=value, resume_token=resume_token, metadata=metadata, spec=["value", "resume_token", "metadata"] ) - + def test_iteration_w_empty_raw(self): raw = _MockIterator() request = mock.Mock(test="test", spec=["test", "resume_token"]) @@ -288,6 +288,39 @@ def test_iteration_w_raw_raising_unavailable_after_token(self): self.assertEqual(request.resume_token, RESUME_TOKEN) self.assertNoSpans() + def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): + from google.api_core.exceptions import ServiceUnavailable + + from google.cloud.spanner_v1 import (ResultSetMetadata) + from google.cloud.spanner_v1 import ( + Transaction as TransactionPB, + ReadRequest, + ) + + transaction_pb = TransactionPB(id=TXN_ID) + metadata_pb = ResultSetMetadata(transaction=transaction_pb) + FIRST = (self._make_item(0), self._make_item(1, resume_token=RESUME_TOKEN, metadata=metadata_pb)) + SECOND = (self._make_item(2), self._make_item(3)) + before = _MockIterator( + *FIRST, fail_after=True, error=ServiceUnavailable("testing") + ) + after = _MockIterator(*SECOND) + request = ReadRequest(transaction=None) + restart = mock.Mock(spec=[], side_effect=[before, after]) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + derived._multi_use = True + + resumable = self._call_fut(derived, restart, request) + + self.assertEqual(list(resumable), list(FIRST + SECOND)) + self.assertEqual(len(restart.mock_calls), 2) + + self.assertEqual(request.resume_token, RESUME_TOKEN) + self.assertNoSpans() + def test_iteration_w_raw_raising_retryable_internal_error_after_token(self): from google.api_core.exceptions import InternalServerError From 8d814ad017ee60b1c9259bd42d9145f120cd087d Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 30 Nov 2022 11:28:00 +0530 Subject: [PATCH 07/19] fix: lint --- google/cloud/spanner_v1/snapshot.py | 7 ++++++- tests/unit/test_snapshot.py | 20 +++++++++++++------- 2 files changed, 19 insertions(+), 8 deletions(-) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index bbad11d2ac..59ce7ddf48 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -77,7 +77,12 @@ def _restart_on_unavailable( item_buffer.append(item) if item.resume_token: resume_token = item.resume_token - if self._transaction_id is None and item.metadata is not None and item.metadata.transaction is not None and item.metadata.transaction.id is not None: + if ( + self._transaction_id is None + and item.metadata is not None + and item.metadata.transaction is not None + and item.metadata.transaction.id is not None + ): self._transaction_id = item.metadata.transaction.id break except ServiceUnavailable: diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index 410ef72f13..f03c326d12 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -77,9 +77,12 @@ def _call_fut( def _make_item(self, value, resume_token=b"", metadata=None): return mock.Mock( - value=value, resume_token=resume_token, metadata=metadata, spec=["value", "resume_token", "metadata"] + value=value, + resume_token=resume_token, + metadata=metadata, + spec=["value", "resume_token", "metadata"], ) - + def test_iteration_w_empty_raw(self): raw = _MockIterator() request = mock.Mock(test="test", spec=["test", "resume_token"]) @@ -291,7 +294,7 @@ def test_iteration_w_raw_raising_unavailable_after_token(self): def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): from google.api_core.exceptions import ServiceUnavailable - from google.cloud.spanner_v1 import (ResultSetMetadata) + from google.cloud.spanner_v1 import ResultSetMetadata from google.cloud.spanner_v1 import ( Transaction as TransactionPB, ReadRequest, @@ -299,7 +302,10 @@ def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): transaction_pb = TransactionPB(id=TXN_ID) metadata_pb = ResultSetMetadata(transaction=transaction_pb) - FIRST = (self._make_item(0), self._make_item(1, resume_token=RESUME_TOKEN, metadata=metadata_pb)) + FIRST = ( + self._make_item(0), + self._make_item(1, resume_token=RESUME_TOKEN, metadata=metadata_pb), + ) SECOND = (self._make_item(2), self._make_item(3)) before = _MockIterator( *FIRST, fail_after=True, error=ServiceUnavailable("testing") @@ -312,12 +318,12 @@ def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): session = _Session(database) derived = self._makeDerived(session) derived._multi_use = True - + resumable = self._call_fut(derived, restart, request) - + self.assertEqual(list(resumable), list(FIRST + SECOND)) self.assertEqual(len(restart.mock_calls), 2) - + self.assertEqual(request.resume_token, RESUME_TOKEN) self.assertNoSpans() From 59d7c1b93071a9920eba9682d500115bbc701f83 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 30 Nov 2022 13:21:00 +0530 Subject: [PATCH 08/19] fix: test cases --- google/cloud/spanner_v1/snapshot.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index 59ce7ddf48..75f30a290d 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -78,7 +78,8 @@ def _restart_on_unavailable( if item.resume_token: resume_token = item.resume_token if ( - self._transaction_id is None + self is not None + and self._transaction_id is None and item.metadata is not None and item.metadata.transaction is not None and item.metadata.transaction.id is not None From b42b66fbff3994312f5f00de27d55a270d9d172a Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 30 Nov 2022 14:53:54 +0530 Subject: [PATCH 09/19] fix: few more test case for restart on unavailable --- tests/unit/test_snapshot.py | 67 ++++++++++++++++++++++++++++++++++++- 1 file changed, 66 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index f03c326d12..ae22bc5167 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -291,6 +291,67 @@ def test_iteration_w_raw_raising_unavailable_after_token(self): self.assertEqual(request.resume_token, RESUME_TOKEN) self.assertNoSpans() + def test_iteration_w_raw_w_multiuse(self): + from google.cloud.spanner_v1 import ( + ReadRequest, + ) + + FIRST = ( + self._make_item(0), + self._make_item(1), + ) + before = _MockIterator(*FIRST) + request = ReadRequest(transaction=None) + restart = mock.Mock(spec=[], return_value=before) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + derived._multi_use = True + resumable = self._call_fut(derived, restart, request) + self.assertEqual(list(resumable), list(FIRST)) + self.assertEqual(len(restart.mock_calls), 1) + begin = 0 + for call in restart.mock_calls: + if "begin" in call.__str__(): + begin += 1 + + self.assertEqual(begin, 1) + self.assertNoSpans() + + def test_iteration_w_raw_raising_unavailable_w_multiuse(self): + from google.api_core.exceptions import ServiceUnavailable + from google.cloud.spanner_v1 import ( + ReadRequest, + ) + + FIRST = ( + self._make_item(0), + self._make_item(1), + ) + SECOND = (self._make_item(2), self._make_item(3)) + before = _MockIterator( + *FIRST, fail_after=True, error=ServiceUnavailable("testing") + ) + after = _MockIterator(*SECOND) + request = ReadRequest(transaction=None) + restart = mock.Mock(spec=[], side_effect=[before, after]) + database = _Database() + database.spanner_api = self._make_spanner_api() + session = _Session(database) + derived = self._makeDerived(session) + derived._multi_use = True + resumable = self._call_fut(derived, restart, request) + self.assertEqual(list(resumable), list(SECOND)) + self.assertEqual(len(restart.mock_calls), 2) + begin = 0 + for call in restart.mock_calls: + if "begin" in call.__str__(): + begin += 1 + + self.assertEqual(begin, 2) + self.assertNoSpans() + def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): from google.api_core.exceptions import ServiceUnavailable @@ -323,8 +384,12 @@ def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): self.assertEqual(list(resumable), list(FIRST + SECOND)) self.assertEqual(len(restart.mock_calls), 2) - + transaction_id_selector = 0 self.assertEqual(request.resume_token, RESUME_TOKEN) + for call in restart.mock_calls: + if 'id: "DEAFBEAD"' in call.__str__(): + transaction_id_selector += 1 + self.assertEqual(transaction_id_selector, 2) self.assertNoSpans() def test_iteration_w_raw_raising_retryable_internal_error_after_token(self): From 4814a280edd508b8300d1956bc5ef2b9c6653763 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Tue, 6 Dec 2022 12:00:31 +0530 Subject: [PATCH 10/19] test: Batch update error test case --- tests/unit/test_spanner.py | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index 4dcdf8190a..5bd6b0b91e 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -507,6 +507,26 @@ def test_transaction_should_include_begin_with_first_batch_update(self): metadata=[("google-cloud-resource-prefix", database.name)], ) + def test_transaction_should_use_transaction_id_if_error_with_first_batch_update(self): + database = _Database() + session = _Session(database) + api = database.spanner_api = self._make_spanner_api() + transaction = self._make_one(session) + self._batch_update_helper(transaction=transaction, database=database, api=api, error_after=2) + api.execute_batch_dml.assert_called_once_with( + request=self._batch_update_expected_request(begin=True), + metadata=[("google-cloud-resource-prefix", database.name)], + ) + self._execute_update_helper(transaction=transaction, api=api) + api.execute_sql.assert_called_once_with( + request=self._execute_update_expected_request( + database=database, begin=False + ), + retry=gapic_v1.method.DEFAULT, + timeout=gapic_v1.method.DEFAULT, + metadata=[("google-cloud-resource-prefix", database.name)], + ) + def test_transaction_should_use_transaction_id_returned_by_first_query(self): database = _Database() session = _Session(database) From 9f6ce71aec91b3537bb6aea4d1432340719e97b4 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Tue, 6 Dec 2022 12:25:27 +0530 Subject: [PATCH 11/19] fix: lint --- tests/unit/test_spanner.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index 5bd6b0b91e..a6952d91e3 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -507,12 +507,16 @@ def test_transaction_should_include_begin_with_first_batch_update(self): metadata=[("google-cloud-resource-prefix", database.name)], ) - def test_transaction_should_use_transaction_id_if_error_with_first_batch_update(self): + def test_transaction_should_use_transaction_id_if_error_with_first_batch_update( + self, + ): database = _Database() session = _Session(database) api = database.spanner_api = self._make_spanner_api() transaction = self._make_one(session) - self._batch_update_helper(transaction=transaction, database=database, api=api, error_after=2) + self._batch_update_helper( + transaction=transaction, database=database, api=api, error_after=2 + ) api.execute_batch_dml.assert_called_once_with( request=self._batch_update_expected_request(begin=True), metadata=[("google-cloud-resource-prefix", database.name)], From 3a002c4d1a56cfa6bb1363198db9d16c49238ccd Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 7 Dec 2022 19:24:33 +0530 Subject: [PATCH 12/19] fix: Code review comments --- google/cloud/spanner_v1/database.py | 4 +- google/cloud/spanner_v1/snapshot.py | 55 ++++++++++++----------- google/cloud/spanner_v1/transaction.py | 22 ++++++++-- tests/unit/test_session.py | 6 +++ tests/unit/test_snapshot.py | 61 ++++++++++++++++++++------ tests/unit/test_spanner.py | 2 +- tests/unit/test_transaction.py | 22 ++++++++++ 7 files changed, 127 insertions(+), 45 deletions(-) diff --git a/google/cloud/spanner_v1/database.py b/google/cloud/spanner_v1/database.py index 24be4b2bf1..346f181075 100644 --- a/google/cloud/spanner_v1/database.py +++ b/google/cloud/spanner_v1/database.py @@ -575,11 +575,9 @@ def execute_pdml(): ) iterator = _restart_on_unavailable( - self=None, method=method, request=request, - isPdml=True, - transactionSelector=txn_selector, + transaction_selector=txn_selector, ) result_set = StreamedResultSet(iterator) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index 75f30a290d..0ee307b131 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -27,6 +27,7 @@ from google.api_core.exceptions import InternalServerError from google.api_core.exceptions import ServiceUnavailable +from google.api_core.exceptions import InvalidArgument from google.api_core import gapic_v1 from google.cloud.spanner_v1._helpers import _make_value_pb from google.cloud.spanner_v1._helpers import _merge_query_options @@ -43,14 +44,13 @@ def _restart_on_unavailable( - self, method, request, trace_name=None, session=None, attributes=None, - isPdml=False, - transactionSelector=None, + transaction=None, + transaction_selector=None, ): """Restart iteration after :exc:`.ServiceUnavailable`. @@ -63,12 +63,15 @@ def _restart_on_unavailable( resume_token = b"" item_buffer = [] - if isPdml is True: - transaction = transactionSelector - else: - transaction = self._make_txn_selector() - request.transaction = transaction + if transaction is not None: + transaction_selector = transaction._make_txn_selector() + elif transaction_selector is None: + raise InvalidArgument( + "Either transaction or transaction_selector should be set" + ) + + request.transaction = transaction_selector with trace_call(trace_name, session, attributes): iterator = method(request=request) while True: @@ -77,24 +80,23 @@ def _restart_on_unavailable( item_buffer.append(item) if item.resume_token: resume_token = item.resume_token + # Setting the transaction id because the transaction begin was inlined for first rpc. if ( - self is not None - and self._transaction_id is None + transaction is not None + and transaction._transaction_id is None and item.metadata is not None and item.metadata.transaction is not None and item.metadata.transaction.id is not None ): - self._transaction_id = item.metadata.transaction.id + transaction._transaction_id = item.metadata.transaction.id break except ServiceUnavailable: del item_buffer[:] with trace_call(trace_name, session, attributes): request.resume_token = resume_token - if isPdml is True: - transaction = transactionSelector - else: - transaction = self._make_txn_selector() - request.transaction = transaction + if transaction is not None: + transaction_selector = transaction._make_txn_selector() + request.transaction = transaction_selector iterator = method(request=request) continue except InternalServerError as exc: @@ -107,11 +109,9 @@ def _restart_on_unavailable( del item_buffer[:] with trace_call(trace_name, session, attributes): request.resume_token = resume_token - if isPdml is True: - transaction = transactionSelector - else: - transaction = self._make_txn_selector() - request.transaction = transaction + if transaction is not None: + transaction_selector = transaction._make_txn_selector() + request.transaction = transaction_selector iterator = method(request=request) continue @@ -252,14 +252,15 @@ def read( trace_attributes = {"table_id": table, "columns": columns} if self._transaction_id is None: + # lock is added to handle the inline begin for first rpc with self._lock: iterator = _restart_on_unavailable( - self, restart, request, "CloudSpanner.ReadOnlyTransaction", self._session, trace_attributes, + transaction=self, ) self._read_request_count += 1 if self._multi_use: @@ -268,15 +269,16 @@ def read( return StreamedResultSet(iterator) else: iterator = _restart_on_unavailable( - self, restart, request, "CloudSpanner.ReadOnlyTransaction", self._session, trace_attributes, + transaction=self, ) self._read_request_count += 1 + if self._multi_use: return StreamedResultSet(iterator, source=self) else: @@ -403,33 +405,36 @@ def execute_sql( trace_attributes = {"db.statement": sql} if self._transaction_id is None: + # lock is added to handle the inline begin for first rpc with self._lock: iterator = _restart_on_unavailable( - self, restart, request, "CloudSpanner.ReadWriteTransaction", self._session, trace_attributes, + transaction=self, ) self._read_request_count += 1 self._execute_sql_count += 1 + if self._multi_use: return StreamedResultSet(iterator, source=self) else: return StreamedResultSet(iterator) else: iterator = _restart_on_unavailable( - self, restart, request, "CloudSpanner.ReadWriteTransaction", self._session, trace_attributes, + transaction=self, ) self._read_request_count += 1 self._execute_sql_count += 1 + if self._multi_use: return StreamedResultSet(iterator, source=self) else: diff --git a/google/cloud/spanner_v1/transaction.py b/google/cloud/spanner_v1/transaction.py index 873b988440..b9697f6817 100644 --- a/google/cloud/spanner_v1/transaction.py +++ b/google/cloud/spanner_v1/transaction.py @@ -89,6 +89,14 @@ def _make_txn_selector(self): def _execute_request( self, method, request, trace_name=None, session=None, attributes=None ): + """Helper method to execute request after fetching transaction selector. + + :type method: callable + :param method: function returning iterator + + :type request: proto + :param request: request proto to call the method with + """ transaction = self._make_txn_selector() request.transaction = transaction with trace_call(trace_name, session, attributes): @@ -160,8 +168,10 @@ def commit(self, return_commit_stats=False, request_options=None): :raises ValueError: if there are no mutations to commit. """ self._check_state() - if self._transaction_id is None: + if self._transaction_id is None and len(self._mutations) > 0: self.begin() + elif self._transaction_id is None and len(self._mutations) is 0: + raise ValueError("Transaction is not begun") database = self._session._database api = database.spanner_api @@ -284,7 +294,6 @@ def execute_update( params_pb = self._make_params_pb(params, param_types) database = self._session._database metadata = _metadata_with_prefix(database.name) - api = database.spanner_api seqno, self._execute_sql_count = ( @@ -325,6 +334,7 @@ def execute_update( ) if self._transaction_id is None: + # lock is added to handle the inline begin for first rpc with self._lock: response = self._execute_request( method, @@ -333,8 +343,11 @@ def execute_update( self._session, trace_attributes, ) + # Setting the transaction id because the transaction begin was inlined for first rpc. if ( self._transaction_id is None + and response is not None + and response.metadata is not None and response.metadata.transaction is not None ): self._transaction_id = response.metadata.transaction.id @@ -424,6 +437,7 @@ def batch_update(self, statements, request_options=None): ) if self._transaction_id is None: + # lock is added to handle the inline begin for first rpc with self._lock: response = self._execute_request( method, @@ -432,13 +446,15 @@ def batch_update(self, statements, request_options=None): self._session, trace_attributes, ) - + # Setting the transaction id because the transaction begin was inlined for first rpc. for result_set in response.result_sets: if ( self._transaction_id is None + and result_set.metadata is not None and result_set.metadata.transaction is not None ): self._transaction_id = result_set.metadata.transaction.id + break else: response = self._execute_request( method, diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index d36326ac49..763e8fa7f2 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -723,6 +723,10 @@ def unit_of_work(txn, *args, **kw): self.assertTrue(txn.rolled_back) self.assertEqual(args, ()) self.assertEqual(kw, {}) + # Transaction only has mutation operations. + # Exception was raised before commit, hence transaction did not begin. + # Therefore rollback was not called. + gax_api.rollback.assert_not_called() def test_run_in_transaction_callback_raises_non_abort_rpc_error(self): from google.api_core.exceptions import Cancelled @@ -1121,6 +1125,8 @@ def unit_of_work(txn, *args, **kw): self.assertEqual(kw, {}) expected_options = TransactionOptions(read_write=TransactionOptions.ReadWrite()) + + # First call was aborted before commit operation, Therefore no begin rpc was made during first attempt. gax_api.begin_transaction.assert_called_once_with( session=self.SESSION_NAME, options=expected_options, diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index ae22bc5167..74cf5f613a 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -48,23 +48,39 @@ } -class Test_SnapshotBase(OpenTelemetryBase): - - PROJECT_ID = "project-id" - INSTANCE_ID = "instance-id" - INSTANCE_NAME = "projects/" + PROJECT_ID + "/instances/" + INSTANCE_ID - DATABASE_ID = "database-id" - DATABASE_NAME = INSTANCE_NAME + "/databases/" + DATABASE_ID - SESSION_ID = "session-id" - SESSION_NAME = DATABASE_NAME + "/sessions/" + SESSION_ID - +class Test_restart_on_unavailable(OpenTelemetryBase): def _getTargetClass(self): from google.cloud.spanner_v1.snapshot import _SnapshotBase return _SnapshotBase - def _make_one(self, session): - return self._getTargetClass()(session) + def _makeDerived(self, session): + class _Derived(self._getTargetClass()): + + _transaction_id = None + _multi_use = False + + def _make_txn_selector(self): + from google.cloud.spanner_v1 import ( + TransactionOptions, + TransactionSelector, + ) + + if self._transaction_id: + return TransactionSelector(id=self._transaction_id) + options = TransactionOptions( + read_only=TransactionOptions.ReadOnly(strong=True) + ) + if self._multi_use: + return TransactionSelector(begin=options) + return TransactionSelector(single_use=options) + + return _Derived(session) + + def _make_spanner_api(self): + from google.cloud.spanner_v1 import SpannerClient + + return mock.create_autospec(SpannerClient, instance=True) def _call_fut( self, derived, restart, request, span_name=None, session=None, attributes=None @@ -72,7 +88,7 @@ def _call_fut( from google.cloud.spanner_v1.snapshot import _restart_on_unavailable return _restart_on_unavailable( - derived, restart, request, span_name, session, attributes + restart, request, span_name, session, attributes, transaction=derived ) def _make_item(self, value, resume_token=b"", metadata=None): @@ -493,6 +509,25 @@ def test_iteration_w_multiple_span_creation(self): }, ) + +class Test_SnapshotBase(OpenTelemetryBase): + + PROJECT_ID = "project-id" + INSTANCE_ID = "instance-id" + INSTANCE_NAME = "projects/" + PROJECT_ID + "/instances/" + INSTANCE_ID + DATABASE_ID = "database-id" + DATABASE_NAME = INSTANCE_NAME + "/databases/" + DATABASE_ID + SESSION_ID = "session-id" + SESSION_NAME = DATABASE_NAME + "/sessions/" + SESSION_ID + + def _getTargetClass(self): + from google.cloud.spanner_v1.snapshot import _SnapshotBase + + return _SnapshotBase + + def _make_one(self, session): + return self._getTargetClass()(session) + def _makeDerived(self, session): class _Derived(self._getTargetClass()): diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index a6952d91e3..4f447bcd98 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -1,4 +1,4 @@ -# Copyright 2016 Google LLC All rights reserved. +# Copyright 2022 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. diff --git a/tests/unit/test_transaction.py b/tests/unit/test_transaction.py index f9e471b8f1..2bf489f34f 100644 --- a/tests/unit/test_transaction.py +++ b/tests/unit/test_transaction.py @@ -188,6 +188,20 @@ def test_begin_ok(self): "CloudSpanner.BeginTransaction", attributes=TestTransaction.BASE_ATTRIBUTES ) + def test_rollback_not_begun(self): + database = _Database() + api = database.spanner_api = self._make_spanner_api() + session = _Session(database) + transaction = self._make_one(session) + + transaction.rollback() + self.assertTrue(transaction.rolled_back) + + # Since there was no transaction to be rolled back, rollbacl rpc is not called. + api.rollback.assert_not_called() + + self.assertNoSpans() + def test_rollback_already_committed(self): session = _Session() transaction = self._make_one(session) @@ -253,6 +267,14 @@ def test_rollback_ok(self): "CloudSpanner.Rollback", attributes=TestTransaction.BASE_ATTRIBUTES ) + def test_commit_not_begun(self): + session = _Session() + transaction = self._make_one(session) + with self.assertRaises(ValueError): + transaction.commit() + + self.assertNoSpans() + def test_commit_already_committed(self): session = _Session() transaction = self._make_one(session) From 1640cbd1ea7806058e045e774447355c9ec50d26 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Thu, 8 Dec 2022 15:02:18 +0530 Subject: [PATCH 13/19] fix: test cases + lint --- google/cloud/spanner_v1/transaction.py | 2 +- tests/system/test_session_api.py | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/google/cloud/spanner_v1/transaction.py b/google/cloud/spanner_v1/transaction.py index b9697f6817..ce34054ab9 100644 --- a/google/cloud/spanner_v1/transaction.py +++ b/google/cloud/spanner_v1/transaction.py @@ -170,7 +170,7 @@ def commit(self, return_commit_stats=False, request_options=None): self._check_state() if self._transaction_id is None and len(self._mutations) > 0: self.begin() - elif self._transaction_id is None and len(self._mutations) is 0: + elif self._transaction_id is None and len(self._mutations) == 0: raise ValueError("Transaction is not begun") database = self._session._database diff --git a/tests/system/test_session_api.py b/tests/system/test_session_api.py index 4872bcae8e..c9c5c8a959 100644 --- a/tests/system/test_session_api.py +++ b/tests/system/test_session_api.py @@ -1027,6 +1027,7 @@ def test_transaction_batch_update_wo_statements(sessions_database, sessions_to_d sessions_to_delete.append(session) with session.transaction() as transaction: + transaction.begin() with pytest.raises(exceptions.InvalidArgument): transaction.batch_update([]) From 6f102f94e2281caa113f574c3ba0930fb9dc97af Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Fri, 9 Dec 2022 15:06:49 +0530 Subject: [PATCH 14/19] fix: code review comments --- google/cloud/spanner_v1/snapshot.py | 6 ++++++ tests/unit/test_transaction.py | 2 +- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index 0ee307b131..04b577df61 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -59,6 +59,12 @@ def _restart_on_unavailable( :type request: proto :param request: request proto to call the method with + + :type transaction: :class:`google.cloud.spanner_v1.snapshot._SnapshotBase` + :param transaction: Snapshot or Transaction class object based on the type of transaction + + :type transaction_selector: :class:`transaction_pb2.TransactionSelector` + :param transaction_selector: Transaction selector object to be used in request if transaction is not passed """ resume_token = b"" diff --git a/tests/unit/test_transaction.py b/tests/unit/test_transaction.py index 2bf489f34f..5fb69b4979 100644 --- a/tests/unit/test_transaction.py +++ b/tests/unit/test_transaction.py @@ -835,9 +835,9 @@ def test_context_mgr_failure(self): raise Exception("bail out") self.assertEqual(transaction.committed, None) + # Rollback rpc will not be called as there is no transaction id to be rolled back, rolled_back flag will be marked as true. self.assertTrue(transaction.rolled_back) self.assertEqual(len(transaction._mutations), 1) - self.assertEqual(api._committed, None) From 2db734de2451f2ca619ee0f60ca1538929daaba5 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Fri, 9 Dec 2022 15:53:24 +0530 Subject: [PATCH 15/19] fix: deprecate transactionpingingpool msg --- google/cloud/spanner_v1/pool.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/google/cloud/spanner_v1/pool.py b/google/cloud/spanner_v1/pool.py index 9c76837255..8d92503172 100644 --- a/google/cloud/spanner_v1/pool.py +++ b/google/cloud/spanner_v1/pool.py @@ -19,7 +19,7 @@ from google.cloud.exceptions import NotFound from google.cloud.spanner_v1._helpers import _metadata_with_prefix - +from warnings import warn _NOW = datetime.datetime.utcnow # unit tests may replace @@ -447,6 +447,9 @@ def ping(self): class TransactionPingingPool(PingingPool): """Concrete session pool implementation: + Deprecated: TransactionPingingPool no longer begins a transaction for each of its sessions at startup. + Hence the TransactionPingingPool is same as :class:`PingingPool` , and maybe removed in the future. + In addition to the features of :class:`PingingPool`, this class creates and begins a transaction for each of its sessions at startup. @@ -473,6 +476,12 @@ class TransactionPingingPool(PingingPool): """ def __init__(self, size=10, default_timeout=10, ping_interval=3000, labels=None): + """This throws a deprecation warning on initialization.""" + warn( + f"{self.__class__.__name__} is deprecated.", + DeprecationWarning, + stacklevel=2, + ) self._pending_sessions = queue.Queue() super(TransactionPingingPool, self).__init__( From 7474a8b5cb3e0a6a923aef74386f501087027f24 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Mon, 12 Dec 2022 14:54:55 +0530 Subject: [PATCH 16/19] fix: review comments Co-authored-by: larkee <31196561+larkee@users.noreply.github.com> --- google/cloud/spanner_v1/pool.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/google/cloud/spanner_v1/pool.py b/google/cloud/spanner_v1/pool.py index 8d92503172..dd5698eecc 100644 --- a/google/cloud/spanner_v1/pool.py +++ b/google/cloud/spanner_v1/pool.py @@ -448,7 +448,7 @@ def ping(self): class TransactionPingingPool(PingingPool): """Concrete session pool implementation: Deprecated: TransactionPingingPool no longer begins a transaction for each of its sessions at startup. - Hence the TransactionPingingPool is same as :class:`PingingPool` , and maybe removed in the future. + Hence the TransactionPingingPool is same as :class:`PingingPool` and maybe removed in the future. In addition to the features of :class:`PingingPool`, this class From 1070ecda23436d7ba6c34a99531a1ffa81859fc9 Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Mon, 12 Dec 2022 14:59:47 +0530 Subject: [PATCH 17/19] fix: Apply suggestions from code review Co-authored-by: larkee <31196561+larkee@users.noreply.github.com> --- google/cloud/spanner_v1/pool.py | 1 + tests/unit/test_session.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/google/cloud/spanner_v1/pool.py b/google/cloud/spanner_v1/pool.py index dd5698eecc..a6d613bf94 100644 --- a/google/cloud/spanner_v1/pool.py +++ b/google/cloud/spanner_v1/pool.py @@ -447,6 +447,7 @@ def ping(self): class TransactionPingingPool(PingingPool): """Concrete session pool implementation: + Deprecated: TransactionPingingPool no longer begins a transaction for each of its sessions at startup. Hence the TransactionPingingPool is same as :class:`PingingPool` and maybe removed in the future. diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 763e8fa7f2..08c48ba748 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -1126,7 +1126,7 @@ def unit_of_work(txn, *args, **kw): expected_options = TransactionOptions(read_write=TransactionOptions.ReadWrite()) - # First call was aborted before commit operation, Therefore no begin rpc was made during first attempt. + # First call was aborted before commit operation, therefore no begin rpc was made during first attempt. gax_api.begin_transaction.assert_called_once_with( session=self.SESSION_NAME, options=expected_options, From d4ba6f5a9875d67c09e41076b0597874c9af915e Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Mon, 12 Dec 2022 16:27:39 +0530 Subject: [PATCH 18/19] fix: review comments --- google/cloud/spanner_v1/snapshot.py | 21 ++++++++--------- tests/unit/test_session.py | 3 ++- tests/unit/test_snapshot.py | 35 ++++++++++++++++------------- tests/unit/test_spanner.py | 18 +++++++-------- 4 files changed, 40 insertions(+), 37 deletions(-) diff --git a/google/cloud/spanner_v1/snapshot.py b/google/cloud/spanner_v1/snapshot.py index 04b577df61..f1fff8b533 100644 --- a/google/cloud/spanner_v1/snapshot.py +++ b/google/cloud/spanner_v1/snapshot.py @@ -64,7 +64,8 @@ def _restart_on_unavailable( :param transaction: Snapshot or Transaction class object based on the type of transaction :type transaction_selector: :class:`transaction_pb2.TransactionSelector` - :param transaction_selector: Transaction selector object to be used in request if transaction is not passed + :param transaction_selector: Transaction selector object to be used in request if transaction is not passed, + if both transaction_selector and transaction are passed, then transaction is given priority. """ resume_token = b"" @@ -84,17 +85,17 @@ def _restart_on_unavailable( try: for item in iterator: item_buffer.append(item) + # Setting the transaction id because the transaction begin was inlined for first rpc. + if ( + transaction is not None + and transaction._transaction_id is None + and item.metadata is not None + and item.metadata.transaction is not None + and item.metadata.transaction.id is not None + ): + transaction._transaction_id = item.metadata.transaction.id if item.resume_token: resume_token = item.resume_token - # Setting the transaction id because the transaction begin was inlined for first rpc. - if ( - transaction is not None - and transaction._transaction_id is None - and item.metadata is not None - and item.metadata.transaction is not None - and item.metadata.transaction.id is not None - ): - transaction._transaction_id = item.metadata.transaction.id break except ServiceUnavailable: del item_buffer[:] diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 08c48ba748..5b331eebf2 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -725,8 +725,9 @@ def unit_of_work(txn, *args, **kw): self.assertEqual(kw, {}) # Transaction only has mutation operations. # Exception was raised before commit, hence transaction did not begin. - # Therefore rollback was not called. + # Therefore rollback and begin transaction was not called. gax_api.rollback.assert_not_called() + gax_api.begin_transaction.assert_not_called() def test_run_in_transaction_callback_raises_non_abort_rpc_error(self): from google.api_core.exceptions import Cancelled diff --git a/tests/unit/test_snapshot.py b/tests/unit/test_snapshot.py index 74cf5f613a..c3ea162f11 100644 --- a/tests/unit/test_snapshot.py +++ b/tests/unit/test_snapshot.py @@ -327,12 +327,10 @@ def test_iteration_w_raw_w_multiuse(self): resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(FIRST)) self.assertEqual(len(restart.mock_calls), 1) - begin = 0 - for call in restart.mock_calls: - if "begin" in call.__str__(): - begin += 1 - - self.assertEqual(begin, 1) + begin_count = sum( + [1 for args in restart.call_args_list if "begin" in args.kwargs.__str__()] + ) + self.assertEqual(begin_count, 1) self.assertNoSpans() def test_iteration_w_raw_raising_unavailable_w_multiuse(self): @@ -360,12 +358,12 @@ def test_iteration_w_raw_raising_unavailable_w_multiuse(self): resumable = self._call_fut(derived, restart, request) self.assertEqual(list(resumable), list(SECOND)) self.assertEqual(len(restart.mock_calls), 2) - begin = 0 - for call in restart.mock_calls: - if "begin" in call.__str__(): - begin += 1 + begin_count = sum( + [1 for args in restart.call_args_list if "begin" in args.kwargs.__str__()] + ) - self.assertEqual(begin, 2) + # Since the transaction id was not set before the Unavailable error, the statement will be retried with inline begin. + self.assertEqual(begin_count, 2) self.assertNoSpans() def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): @@ -400,12 +398,17 @@ def test_iteration_w_raw_raising_unavailable_after_token_w_multiuse(self): self.assertEqual(list(resumable), list(FIRST + SECOND)) self.assertEqual(len(restart.mock_calls), 2) - transaction_id_selector = 0 self.assertEqual(request.resume_token, RESUME_TOKEN) - for call in restart.mock_calls: - if 'id: "DEAFBEAD"' in call.__str__(): - transaction_id_selector += 1 - self.assertEqual(transaction_id_selector, 2) + transaction_id_selector_count = sum( + [ + 1 + for args in restart.call_args_list + if 'id: "DEAFBEAD"' in args.kwargs.__str__() + ] + ) + + # Statement will be retried with Transaction id. + self.assertEqual(transaction_id_selector_count, 2) self.assertNoSpans() def test_iteration_w_raw_raising_retryable_internal_error_after_token(self): diff --git a/tests/unit/test_spanner.py b/tests/unit/test_spanner.py index 4f447bcd98..a7c41c5f4f 100644 --- a/tests/unit/test_spanner.py +++ b/tests/unit/test_spanner.py @@ -738,12 +738,11 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ self._execute_update_helper(transaction=transaction, api=api) - begin_read_write = 0 - for call in api.mock_calls: - if "read_write" in call.__str__(): - begin_read_write += 1 + begin_read_write_count = sum( + [1 for call in api.mock_calls if "read_write" in call.kwargs.__str__()] + ) - self.assertEqual(begin_read_write, 1) + self.assertEqual(begin_read_write_count, 1) api.execute_sql.assert_any_call( request=self._execute_update_expected_request(database, begin=False), retry=RETRY, @@ -796,12 +795,11 @@ def test_transaction_for_concurrent_statement_should_begin_one_transaction_with_ self._execute_update_helper(transaction=transaction, api=api) - begin_read_write = 0 - for call in api.mock_calls: - if "read_write" in call.__str__(): - begin_read_write += 1 + begin_read_write_count = sum( + [1 for call in api.mock_calls if "read_write" in call.kwargs.__str__()] + ) - self.assertEqual(begin_read_write, 1) + self.assertEqual(begin_read_write_count, 1) api.execute_sql.assert_any_call( request=self._execute_update_expected_request(database, begin=False), retry=RETRY, From a565281fc6f190c5f2f424d81f6d4ec1e97681ad Mon Sep 17 00:00:00 2001 From: surbhigarg92 Date: Wed, 14 Dec 2022 15:07:39 +0530 Subject: [PATCH 19/19] fix: review comment Update tests/unit/test_session.py Co-authored-by: larkee <31196561+larkee@users.noreply.github.com> --- tests/unit/test_session.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 48a8bb38e8..edad4ce777 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -800,7 +800,7 @@ def unit_of_work(txn, *args, **kw): self.assertEqual(kw, {}) # Transaction only has mutation operations. # Exception was raised before commit, hence transaction did not begin. - # Therefore rollback and begin transaction was not called. + # Therefore rollback and begin transaction were not called. gax_api.rollback.assert_not_called() gax_api.begin_transaction.assert_not_called()