diff --git a/packages/google-cloud-spanner/.cross_sync/generate.py b/packages/google-cloud-spanner/.cross_sync/generate.py index 3c96b7469e61..890de77adf50 100644 --- a/packages/google-cloud-spanner/.cross_sync/generate.py +++ b/packages/google-cloud-spanner/.cross_sync/generate.py @@ -12,8 +12,10 @@ # See the License for the specific language governing permissions and # limitations under the License. from __future__ import annotations -from typing import Sequence + import ast +from typing import Sequence + """ Entrypoint for initiating an async -> sync conversion using CrossSync @@ -35,12 +37,13 @@ def extract_header_comments(file_path) -> str: header.append(line) else: break - header.append("\n# This file is automatically generated by CrossSync. Do not edit manually.\n\n") + header.append( + "\n# This file is automatically generated by CrossSync. Do not edit manually.\n\n" + ) return "".join(header) class CrossSyncOutputFile: - def __init__(self, output_path: str, ast_tree, header: str | None = None): self.output_path = output_path self.tree = ast_tree @@ -56,15 +59,19 @@ def render(self, with_formatter=True, save_to_disk: bool = True) -> str: """ full_str = self.header + ast.unparse(self.tree) if with_formatter: - import black # type: ignore - import autoflake # type: ignore - - full_str = black.format_str( - autoflake.fix_code(full_str, remove_all_unused_imports=True), - mode=black.FileMode(), - ) + try: + import autoflake # type: ignore + import black # type: ignore + + full_str = black.format_str( + autoflake.fix_code(full_str, remove_all_unused_imports=True), + mode=black.FileMode(), + ) + except ImportError: + pass if save_to_disk: import os + os.makedirs(os.path.dirname(self.output_path), exist_ok=True) with open(self.output_path, "w") as f: f.write(full_str) @@ -73,10 +80,15 @@ def render(self, with_formatter=True, save_to_disk: bool = True) -> str: def convert_files_in_dir(directory: str) -> set[CrossSyncOutputFile]: import glob + import os + from transformers import CrossSyncFileProcessor - # find all python files in the directory - files = glob.glob(directory + "/**/*.py", recursive=True) + # find all python files in the directory or use single file + if os.path.isfile(directory): + files = [directory] + else: + files = glob.glob(directory + "/**/*.py", recursive=True) # keep track of the output files pointed to by the annotated classes artifacts: set[CrossSyncOutputFile] = set() file_transformer = CrossSyncFileProcessor() diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py index e02c79c6c553..50bb4a2c3f75 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/_helpers.py @@ -1,5 +1,6 @@ import asyncio import inspect +import os import time from google.api_core.exceptions import Aborted @@ -8,7 +9,8 @@ async def _delay_until_retry(exc, deadline, attempts, default_retry_delay=None): from google.cloud.spanner_v1._helpers import _get_retry_delay - cause = exc.errors[0] if hasattr(exc, "errors") and exc.errors else exc + errors = getattr(exc, "errors", None) + cause = errors[0] if errors else exc now = time.time() if now >= deadline: raise exc @@ -153,3 +155,44 @@ def _create_experimental_host_transport( client_key, interceptors=interceptors, ) + + +_PENDING_DRAIN_TASKS = set() +if hasattr(os, "register_at_fork"): + os.register_at_fork(after_in_child=_PENDING_DRAIN_TASKS.clear) + + +def _drain_stream(iterator): + """Drain an async stream iterator to EOF in the background. + + Called when PartialResultSet.last is True to allow the caller to return immediately + while consuming trailing gRPC metadata so the stream terminates cleanly with status OK. + """ + if iterator is None: + return + + async def _drain(): + try: + async for _ in iterator: + pass + except asyncio.CancelledError: + if hasattr(iterator, "cancel"): + try: + iterator.cancel() + except Exception: + pass + raise + except Exception: + pass + + try: + task = asyncio.create_task(_drain()) + _PENDING_DRAIN_TASKS.add(task) + task.add_done_callback(_PENDING_DRAIN_TASKS.discard) + except RuntimeError: + # Event loop may be closed or not running. + if hasattr(iterator, "cancel"): + try: + iterator.cancel() + except Exception: + pass diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database_sessions_manager.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database_sessions_manager.py index 30c70d797694..2d6e07e0a88e 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database_sessions_manager.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database_sessions_manager.py @@ -123,8 +123,8 @@ async def _get_multiplexed_session(self) -> Session: """Returns a multiplexed session from the database session manager. If the multiplexed session is not defined, creates a new multiplexed - session and starts a maintenance thread to periodically delete and - recreate it so that it remains valid. Otherwise, simply returns the + session and starts a maintenance thread to periodically rotate + it so that it remains valid. Otherwise, simply returns the current multiplexed session. :rtype: :class:`~google.cloud.spanner_v1.session.Session` @@ -167,8 +167,8 @@ def _build_maintenance_thread( self, session: Optional[Session] = None ) -> CrossSync.Task: """Builds and returns a multiplexed session maintenance thread for - the database session manager. This thread will periodically delete - and recreate the multiplexed session to ensure that it is always valid. + the database session manager. This thread will periodically rotate + the multiplexed session to ensure that it is always valid. :type session: :class:`~google.cloud.spanner_v1.session.Session` :param session: (Optional) The multiplexed session to maintain. @@ -209,15 +209,8 @@ async def _rotate_multiplexed_session(self) -> bool: return False async with self._multiplexed_session_lock: - old_session = self._multiplexed_session self._multiplexed_session = new_session - if old_session is not None: - try: - await CrossSync.run_if_async(old_session.delete) - except Exception: - pass - return True @staticmethod @@ -225,9 +218,9 @@ async def _rotate_multiplexed_session(self) -> bool: async def _maintain_multiplexed_session(session_manager_ref) -> None: """Maintains the multiplexed session for the database session manager. - This method will delete and recreate the referenced database session manager's + This method will periodically rotate the referenced database session manager's multiplexed session to ensure that it is always valid. The method will run until - the database session manager is deleted or the multiplexed session is deleted. + the database session manager is garbage collected or the session manager is closed. :type session_manager_ref: :class:`_weakref.ReferenceType` :param session_manager_ref: A weak reference to the database session manager.""" @@ -292,7 +285,4 @@ async def close(self) -> None: pass else: self._multiplexed_session_thread.join() - if self._multiplexed_session is not None: - session_to_delete = self._multiplexed_session - self._multiplexed_session = None - await session_to_delete.delete() + self._multiplexed_session = None diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/session.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/session.py index a1a567bf94a4..b5911ad343b1 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/session.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/session.py @@ -518,6 +518,106 @@ def transaction(self, client_context=None) -> Transaction: return Transaction(self, client_context=client_context) + def _create_transaction_for_attempt( + self, + client_context=None, + transaction_tag=None, + exclude_txn_from_change_streams=None, + isolation_level=None, + read_lock_mode=None, + previous_transaction_id=None, + ) -> Transaction: + """Create and configure a transaction instance for a single attempt. + + :type client_context: :class:`~google.cloud.spanner_v1.client_context.ClientContext` + :param client_context: (Optional) client context to use for the transaction. + + :type transaction_tag: str + :param transaction_tag: (Optional) transaction tag. + + :type exclude_txn_from_change_streams: bool + :param exclude_txn_from_change_streams: (Optional) whether to exclude from change streams. + + :type isolation_level: int + :param isolation_level: (Optional) isolation level. + + :type read_lock_mode: int + :param read_lock_mode: (Optional) read lock mode. + + :type previous_transaction_id: bytes + :param previous_transaction_id: (Optional) previous transaction id for multiplexed sessions. + + :rtype: :class:`~google.cloud.spanner_v1.transaction.Transaction` + :returns: A configured Transaction instance. + """ + transaction = self.transaction(client_context=client_context) + transaction.transaction_tag = transaction_tag + transaction.exclude_txn_from_change_streams = exclude_txn_from_change_streams + transaction.isolation_level = isolation_level + transaction.read_lock_mode = read_lock_mode + + if self.is_multiplexed: + transaction._multiplexed_session_previous_transaction_id = ( + previous_transaction_id + ) + return transaction + + @CrossSync.convert + async def _handle_aborted( + self, + exception, + span, + event_name, + attempts, + deadline, + default_retry_delay, + include_cause=False, + ): + """Handle an Aborted error: record trace event and delay until next retry attempt. + + :type exception: :class:`google.api_core.exceptions.Aborted` + :param exception: The aborted exception. + + :type span: :class:`opentelemetry.trace.Span` + :param span: The current active span. + + :type event_name: str + :param event_name: The span event name to record. + + :type attempts: int + :param attempts: The retry attempt number. + + :type deadline: float + :param deadline: Timestamp deadline for retrying. + + :type default_retry_delay: float + :param default_retry_delay: Default delay between retries. + + :type include_cause: bool + :param include_cause: (Optional) Whether to include the exception string as cause. + """ + if span and span.is_recording(): + errors = getattr(exception, "errors", None) + cause = errors[0] if errors else exception + delay_seconds = _get_retry_delay( + cause, + attempts, + default_retry_delay=default_retry_delay, + ) + attributes = { + "attempt": attempts, + "delay_seconds": delay_seconds, + } + if include_cause: + attributes["cause"] = str(exception) + add_span_event(span, event_name, attributes) + await _delay_until_retry( + exception, + deadline, + attempts, + default_retry_delay=default_retry_delay, + ) + @CrossSync.convert async def run_in_transaction(self, func, *args, **kw): """Perform a unit of work in a transaction, retrying on abort. @@ -568,126 +668,94 @@ async def run_in_transaction(self, func, *args, **kw): database = self._database log_commit_stats = database.log_commit_stats - extra_attributes = {} - if transaction_tag: - extra_attributes["transaction.tag"] = transaction_tag - - with ( - trace_call( - "CloudSpanner.Session.run_in_transaction", - self, - extra_attributes=extra_attributes, - observability_options=getattr(database, "observability_options", None), - ) as span, - MetricsCapture(self._resource_info), - ): - attempts: int = 0 - - # If a transaction using a multiplexed session is retried after an aborted - # user operation, it should include the previous transaction ID in the - # transaction options used to begin the transaction. This allows the backend - # to recognize the transaction and increase the lock order for the new - # transaction that is created. - # See :attr:`~google.cloud.spanner_v1.types.TransactionOptions.ReadWrite.multiplexed_session_previous_transaction_id` - previous_transaction_id: Optional[bytes] = None - - while True: - txn = self.transaction(client_context=client_context) - txn.transaction_tag = transaction_tag - txn.exclude_txn_from_change_streams = exclude_txn_from_change_streams - txn.isolation_level = isolation_level - txn.read_lock_mode = read_lock_mode - - if self.is_multiplexed: - txn._multiplexed_session_previous_transaction_id = ( - previous_transaction_id - ) + span = get_current_span() + attempts: int = 0 + + # If a transaction using a multiplexed session is retried after an aborted + # user operation, it should include the previous transaction ID in the + # transaction options used to begin the transaction. This allows the backend + # to recognize the transaction and increase the lock order for the new + # transaction that is created. + # See :attr:`~google.cloud.spanner_v1.types.TransactionOptions.ReadWrite.multiplexed_session_previous_transaction_id` + previous_transaction_id: Optional[bytes] = None + + while True: + transaction = self._create_transaction_for_attempt( + client_context=client_context, + transaction_tag=transaction_tag, + exclude_txn_from_change_streams=exclude_txn_from_change_streams, + isolation_level=isolation_level, + read_lock_mode=read_lock_mode, + previous_transaction_id=previous_transaction_id, + ) + attempts += 1 - attempts += 1 - span_attributes = dict(attempt=attempts) + try: + return_value = await CrossSync.run_if_async( + func, transaction, *args, **kw + ) - try: - return_value = await CrossSync.run_if_async(func, txn, *args, **kw) - - except Aborted as exc: - previous_transaction_id = txn._transaction_id - delay_seconds = _get_retry_delay( - exc.errors[0], - attempts, - default_retry_delay=default_retry_delay, - ) - attributes = dict(delay_seconds=delay_seconds, cause=str(exc)) - attributes.update(span_attributes) - add_span_event( - span, - "Transaction was aborted in user operation, retrying", - attributes, - ) - await _delay_until_retry( - exc, - deadline, - attempts, - default_retry_delay=default_retry_delay, - ) - continue + except Aborted as exc: + previous_transaction_id = transaction._transaction_id + await self._handle_aborted( + exc, + span, + "Transaction was aborted in user operation, retrying", + attempts, + deadline, + default_retry_delay, + include_cause=True, + ) + continue - except GoogleAPICallError: - add_span_event( - span, - "User operation failed due to GoogleAPICallError, not retrying", - span_attributes, - ) - raise + except GoogleAPICallError: + add_span_event( + span, + "User operation failed due to GoogleAPICallError, not retrying", + {"attempt": attempts}, + ) + raise - except Exception: - add_span_event( - span, - "User operation failed. Invoking Transaction.rollback(), not retrying", - span_attributes, - ) - await txn.rollback() - raise + except Exception: + add_span_event( + span, + "User operation failed. Invoking Transaction.rollback(), not retrying", + {"attempt": attempts}, + ) + await transaction.rollback() + raise + + try: + await transaction.commit( + return_commit_stats=log_commit_stats, + request_options=commit_request_options, + max_commit_delay=max_commit_delay, + ) - try: - await txn.commit( - return_commit_stats=log_commit_stats, - request_options=commit_request_options, - max_commit_delay=max_commit_delay, - ) + except Aborted as exc: + previous_transaction_id = transaction._transaction_id + await self._handle_aborted( + exc, + span, + "Transaction was aborted during commit, retrying", + attempts, + deadline, + default_retry_delay, + ) + continue - except Aborted as exc: - previous_transaction_id = txn._transaction_id - delay_seconds = _get_retry_delay( - exc.errors[0], - attempts, - default_retry_delay=default_retry_delay, - ) - attributes = dict(delay_seconds=delay_seconds) - attributes.update(span_attributes) - add_span_event( - span, - "Transaction was aborted during commit, retrying", - attributes, - ) - await _delay_until_retry( - exc, - deadline, - attempts, - default_retry_delay=default_retry_delay, - ) + except GoogleAPICallError: + add_span_event( + span, + "Transaction.commit failed due to GoogleAPICallError, not retrying", + {"attempt": attempts}, + ) + raise - except GoogleAPICallError: - add_span_event( - span, - "Transaction.commit failed due to GoogleAPICallError, not retrying", - span_attributes, + else: + if log_commit_stats and transaction.commit_stats: + database.logger.info( + "CommitStats: {}".format(transaction.commit_stats), + extra={"commit_stats": transaction.commit_stats}, ) - raise - - else: - if log_commit_stats and txn.commit_stats: - database.logger.info( - "CommitStats: {}".format(txn.commit_stats), - extra={"commit_stats": txn.commit_stats}, - ) - return return_value + return return_value diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/snapshot.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/snapshot.py index b54b5d314e1a..f52b45abfde5 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/snapshot.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/snapshot.py @@ -28,7 +28,7 @@ from google.protobuf.struct_pb2 import Struct from google.cloud.aio._cross_sync import CrossSync -from google.cloud.spanner_v1._async._helpers import _retry +from google.cloud.spanner_v1._async._helpers import _drain_stream, _retry from google.cloud.spanner_v1._async.streamed import StreamedResultSet from google.cloud.spanner_v1._helpers import ( AtomicCounter, @@ -68,6 +68,86 @@ "Received unexpected EOS on DATA frame from server", ) +_RAW_EXECUTE_SQL_REQUEST_TYPE = ExecuteSqlRequest.pb() + + +def _make_execute_sql_request( + session_name, + sql, + seqno, + params=None, + param_types=None, + query_options=None, + request_options=None, + query_mode=None, + partition=None, + last_statement=None, + data_boost_enabled=None, + directed_read_options=None, +): + """Construct an ExecuteSqlRequest bypassing proto-plus reflection.""" + try: + raw = _RAW_EXECUTE_SQL_REQUEST_TYPE( + session=session_name, + sql=sql, + seqno=seqno, + ) + if params is not None: + if isinstance(params, dict): + if params: + raw.params.update(params) + else: + raw.params.SetInParent() + else: + raw.params.CopyFrom(getattr(params, "_pb", params)) + if param_types: + for k, v in param_types.items(): + raw.param_types[k].CopyFrom(getattr(v, "_pb", v)) + if query_options is not None: + raw.query_options.CopyFrom(getattr(query_options, "_pb", query_options)) + if request_options is not None: + raw.request_options.CopyFrom( + getattr(request_options, "_pb", request_options) + ) + if query_mode is not None: + raw.query_mode = query_mode + if partition is not None: + raw.partition_token = partition + if last_statement: + raw.last_statement = last_statement + if data_boost_enabled: + raw.data_boost_enabled = data_boost_enabled + if directed_read_options is not None: + raw.directed_read_options.CopyFrom( + getattr(directed_read_options, "_pb", directed_read_options) + ) + return ExecuteSqlRequest.wrap(raw) + except Exception: + req_kwargs = { + "session": session_name, + "sql": sql, + "seqno": seqno, + } + if params is not None: + req_kwargs["params"] = params + if param_types: + req_kwargs["param_types"] = param_types + if query_options is not None: + req_kwargs["query_options"] = query_options + if request_options is not None: + req_kwargs["request_options"] = request_options + if query_mode is not None: + req_kwargs["query_mode"] = query_mode + if partition is not None: + req_kwargs["partition_token"] = partition + if last_statement: + req_kwargs["last_statement"] = last_statement + if data_boost_enabled: + req_kwargs["data_boost_enabled"] = data_boost_enabled + if directed_read_options is not None: + req_kwargs["directed_read_options"] = directed_read_options + return ExecuteSqlRequest(req_kwargs) + @CrossSync.convert async def _restart_on_unavailable( @@ -114,94 +194,118 @@ async def _restart_on_unavailable( attempt = 1 nth_request = getattr(request_id_manager, "_next_nth_request", 0) current_request_id = None - - while True: - try: - # Get results iterator. - if iterator is None: - with ( - trace_call( - trace_name, - session, - attributes, - observability_options=observability_options, - metadata=metadata, - ) as span, - MetricsCapture(resource_info), - ): - ( - call_metadata, - current_request_id, - ) = request_id_manager.metadata_and_request_id( - nth_request, - attempt, - metadata, - span, - ) - iterator = await CrossSync.run_if_async( - method, - request=request, - metadata=call_metadata, - ) - - # Add items from iterator to buffer. - item: PartialResultSet - async for item in iterator: - item_buffer.append(item) - - # Update the transaction from the response. + stream_finished = False + + try: + while True: + try: + # Get results iterator. + if iterator is None: + with ( + trace_call( + trace_name, + session, + attributes, + observability_options=observability_options, + metadata=metadata, + ) as span, + MetricsCapture(resource_info), + ): + ( + call_metadata, + current_request_id, + ) = request_id_manager.metadata_and_request_id( + nth_request, + attempt, + metadata, + span, + ) + iterator = await CrossSync.run_if_async( + method, + request=request, + metadata=call_metadata, + ) + + # Add items from iterator to buffer. + item: PartialResultSet + async for item in iterator: + item_buffer.append(item) + item_pb = getattr(item, "_pb", None) or item + + # Update the transaction from the response. + if transaction is not None: + transaction._update_for_result_set_pb(item) + if ( + item_pb is not None + and getattr(item_pb, "HasField", lambda _: False)( + "precommit_token" + ) + and transaction is not None + ): + await transaction._update_for_precommit_token_pb( + item_pb.precommit_token + ) + + item_is_last = getattr(item_pb, "last", False) + + if item_is_last: + stream_finished = True + _drain_stream(iterator) + iterator = None + break + + item_resume_token = getattr(item_pb, "resume_token", b"") + if item_resume_token: + resume_token = item_resume_token + break + + except ServiceUnavailable: + del item_buffer[:] + request.resume_token = resume_token if transaction is not None: - transaction._update_for_result_set_pb(item) - if ( - item._pb is not None - and item._pb.HasField("precommit_token") - and transaction is not None - ): - await transaction._update_for_precommit_token_pb( - item.precommit_token - ) - - if item.resume_token: - resume_token = item.resume_token - break - - except ServiceUnavailable: - del item_buffer[:] - request.resume_token = resume_token - if transaction is not None: - transaction_selector = transaction._build_transaction_selector_pb() - request.transaction = transaction_selector - attempt += 1 - iterator = None - continue - - except InternalServerError as exc: - resumable_error = any( - resumable_message in exc.message - for resumable_message in _STREAM_RESUMPTION_INTERNAL_ERROR_MESSAGES - ) - if not resumable_error: + transaction_selector = transaction._build_transaction_selector_pb() + request.transaction = transaction_selector + attempt += 1 + iterator = None + continue + + except InternalServerError as exc: + resumable_error = any( + resumable_message in exc.message + for resumable_message in _STREAM_RESUMPTION_INTERNAL_ERROR_MESSAGES + ) + if not resumable_error: + raise _augment_error_with_request_id(exc, current_request_id) + del item_buffer[:] + request.resume_token = resume_token + if transaction is not None: + transaction_selector = transaction._build_transaction_selector_pb() + attempt += 1 + request.transaction = transaction_selector + iterator = None + continue + + except Exception as exc: + # Augment any other exception with the request ID raise _augment_error_with_request_id(exc, current_request_id) - del item_buffer[:] - request.resume_token = resume_token - if transaction is not None: - transaction_selector = transaction._build_transaction_selector_pb() - attempt += 1 - request.transaction = transaction_selector - iterator = None - continue - except Exception as exc: - # Augment any other exception with the request ID - raise _augment_error_with_request_id(exc, current_request_id) + if len(item_buffer) == 0: + iterator = None + break - if len(item_buffer) == 0: - break + for item in item_buffer: + yield item - for item in item_buffer: - yield item + del item_buffer[:] - del item_buffer[:] + if stream_finished: + break + finally: + if iterator is not None and hasattr(iterator, "cancel"): + try: + iterator.cancel() + except Exception: + pass class _SnapshotBase(_SessionWrapper): @@ -386,12 +490,10 @@ async def execute_sql( database._instance._client._client_context, self._client_context ) request_options = _merge_request_options(request_options, client_context) - if request_options is None: request_options = RequestOptions() elif type(request_options) is dict: request_options = RequestOptions(request_options) - if self._read_only: request_options.transaction_tag = None if ( @@ -402,14 +504,14 @@ async def execute_sql( elif self.transaction_tag is not None: request_options.transaction_tag = self.transaction_tag - execute_sql_request = ExecuteSqlRequest( - session=session.name, + execute_sql_request = _make_execute_sql_request( + session_name=session.name, sql=sql, + seqno=self._execute_sql_request_count, params=params_pb, param_types=param_types, query_mode=query_mode, - partition_token=partition, - seqno=self._execute_sql_request_count, + partition=partition, query_options=query_options, request_options=request_options, last_statement=last_statement, @@ -759,16 +861,27 @@ def _update_for_result_set_pb( self, result_set_pb: Union[ResultSet, PartialResultSet] ) -> None: """Updates the snapshot for the given result set.""" - if result_set_pb.metadata and result_set_pb.metadata.transaction: - self._update_for_transaction_pb(result_set_pb.metadata.transaction) + rs_pb = getattr(result_set_pb, "_pb", None) or result_set_pb + metadata = getattr(rs_pb, "metadata", None) + if metadata is not None: + tx = getattr(metadata, "transaction", None) + if tx is not None and ( + getattr(tx, "id", None) + or getattr(tx, "HasField", lambda _: False)("precommit_token") + ): + self._update_for_transaction_pb(tx) def _update_for_transaction_pb(self, transaction_pb: Transaction) -> None: """Updates the snapshot for the given transaction.""" - if self._transaction_id is None and transaction_pb.id: - self._transaction_id = transaction_pb.id + tx_pb = getattr(transaction_pb, "_pb", None) or transaction_pb + tx_id = getattr(tx_pb, "id", None) + if self._transaction_id is None and tx_id: + self._transaction_id = tx_id - if transaction_pb._pb.HasField("precommit_token"): - self._update_for_precommit_token_pb_unsafe(transaction_pb.precommit_token) + if tx_pb is not None and getattr(tx_pb, "HasField", lambda _: False)( + "precommit_token" + ): + self._update_for_precommit_token_pb_unsafe(tx_pb.precommit_token) @CrossSync.convert async def _update_for_precommit_token_pb( diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py index c47cc0ef0a17..c0cfb6405b84 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py @@ -120,48 +120,102 @@ def _merge_chunk(self, value): self._pending_chunk = None return merged + def _append_to_current_row(self, values): + """Append cells to the in-progress partial row.""" + if self._lazy_decode: + self._current_row.extend(values) + else: + decoders = self._decoders + start_column = len(self._current_row) + for column_offset, value in enumerate(values): + if value.HasField("null_value"): + self._current_row.append(None) + else: + self._current_row.append( + decoders[start_column + column_offset](value) + ) + + def _decode_lazy_rows(self, values, values_offset, batch_end, width): + """Slice raw protobuf values into rows for lazy decoding.""" + if width == 1: + self._rows.extend([[value] for value in values[values_offset:batch_end]]) + else: + self._rows.extend( + [ + values[row_start : row_start + width] + for row_start in range(values_offset, batch_end, width) + ] + ) + + def _decode_eager_rows(self, values, values_offset, batch_end, width): + """Decode complete row batches into typed Python values.""" + if width == 1: + decoder = self._decoders[0] + self._rows.extend( + [ + [None if value.HasField("null_value") else decoder(value)] + for value in values[values_offset:batch_end] + ] + ) + else: + decoders = self._decoders + rows_append = self._rows.append + column_indices = list(range(width)) + for row_start in range(values_offset, batch_end, width): + rows_append( + [ + None + if values[row_start + column_index].HasField("null_value") + else decoders[column_index](values[row_start + column_index]) + for column_index in column_indices + ] + ) + def _merge_values(self, values): """Merge values into rows. :type values: list of :class:`~google.protobuf.struct_pb2.Value` :param values: non-chunked values from partial result set. """ - decoders = self._decoders + if not values: + return + width = len(self.fields) - index = len(self._current_row) - current_row = self._current_row - rows = self._rows + if width == 0: + return + + values_offset = 0 + total_values = len(values) + + # 1. Complete pending partial row from previous chunk (if any) + if self._current_row: + needed = width - len(self._current_row) + fill_count = min(needed, total_values) + self._append_to_current_row(values[:fill_count]) + values_offset = fill_count + if len(self._current_row) == width: + self._rows.append(self._current_row) + self._current_row = [] + else: + return + + remaining_values = total_values - values_offset + if remaining_values == 0: + return - current_row_append = current_row.append - rows_append = rows.append + row_count = remaining_values // width + full_values_count = row_count * width + batch_end = values_offset + full_values_count + # 2. Batch-decode complete rows if self._lazy_decode: - for value in values: - current_row_append(value) - index += 1 - if index == width: - rows_append(current_row) - current_row = [] - current_row_append = current_row.append - index = 0 + self._decode_lazy_rows(values, values_offset, batch_end, width) else: - for value in values: - # Note: We manually check value.HasField("null_value") here instead of - # wrapping every decoder in _parse_nullable to avoid the overhead of - # an extra Python function call layer for every cell value decoded in this loop. - # If the nullable check logic is updated in _parse_nullable, update this check. - if value.HasField("null_value"): - current_row_append(None) - else: - current_row_append(decoders[index](value)) - index += 1 - if index == width: - rows_append(current_row) - current_row = [] - current_row_append = current_row.append - index = 0 - - self._current_row = current_row + self._decode_eager_rows(values, values_offset, batch_end, width) + + # 3. Buffer trailing partial row remainder for the next chunk (if any) + if remaining_values > full_values_count: + self._append_to_current_row(values[batch_end:]) @CrossSync.convert async def _consume_next(self): diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py index 06b137db1e28..b9b42b3eaa1f 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py @@ -14,6 +14,7 @@ """Helper functions for Cloud Spanner.""" +import atexit import base64 import datetime import decimal @@ -21,6 +22,7 @@ import math import operator import os +import queue import threading import time import uuid @@ -30,7 +32,7 @@ from google.api_core.exceptions import Aborted from google.protobuf.internal.enum_type_wrapper import EnumTypeWrapper from google.protobuf.message import DecodeError, Message -from google.protobuf.struct_pb2 import ListValue, Value +from google.protobuf.struct_pb2 import NULL_VALUE, ListValue, Value from google.rpc.error_details_pb2 import RetryInfo from google.cloud.spanner_v1.data_types import Interval, JsonObject @@ -141,23 +143,55 @@ def _get_cloud_region() -> str: return _cloud_region +def _validate_and_decode_bytes(bytestring): + try: + return bytestring.decode("utf-8") + except (ValueError, UnicodeDecodeError): + raise ValueError( + "Received a bytes that is not base64 encoded. " + "Ensure that you either send a Unicode string or a " + "base64-encoded bytes." + ) + + def _try_to_coerce_bytes(bytestring): """Try to coerce a byte string into the right thing based on Python version and whether or not it is base64 encoded. Return a text string or raise ValueError. """ - # Attempt to coerce using google.protobuf.Value, which will expect - # something that is utf-8 (and base64 consistently is). - try: - Value(string_value=bytestring) - return bytestring - except ValueError: - raise ValueError( - "Received a bytes that is not base64 encoded. " - "Ensure that you either send a Unicode string or a " - "base64-encoded bytes." + _validate_and_decode_bytes(bytestring) + return bytestring + + +def _to_query_options(options): + """Normalize dict or QueryOptions to a non-empty QueryOptions, or None. + + :type options: + :class:`~google.cloud.spanner_v1.types.ExecuteSqlRequest.QueryOptions` + or :class:`dict` or None + :param options: Query options to normalize. + + :rtype: + :class:`~google.cloud.spanner_v1.types.ExecuteSqlRequest.QueryOptions` + or None + :returns: + A non-empty QueryOptions instance, or None if options is empty or None. + + :raises TypeError: + If options is not a QueryOptions, dict, or None. + """ + if options is None: + return None + if isinstance(options, dict): + if not any(options.values()): + return None + options = ExecuteSqlRequest.QueryOptions(options) + elif not isinstance(options, ExecuteSqlRequest.QueryOptions): + raise TypeError( + f"query_options must be a QueryOptions or dict, got {type(options).__name__}" ) + return options if type(options).pb(options).ByteSize() > 0 else None def _merge_query_options(base, merge): @@ -182,23 +216,20 @@ def _merge_query_options(base, merge): QueryOptions object formed by merging the two given QueryOptions. If the resultant object only has empty fields, returns None. """ - combined = base or ExecuteSqlRequest.QueryOptions() - if isinstance(combined, dict): - combined = ExecuteSqlRequest.QueryOptions( - optimizer_version=combined.get("optimizer_version", ""), - optimizer_statistics_package=combined.get( - "optimizer_statistics_package", "" - ), - ) - merge = merge or ExecuteSqlRequest.QueryOptions() - if isinstance(merge, dict): - merge = ExecuteSqlRequest.QueryOptions( - optimizer_version=merge.get("optimizer_version", ""), - optimizer_statistics_package=merge.get("optimizer_statistics_package", ""), - ) - type(combined).pb(combined).MergeFrom(type(merge).pb(merge)) - if not combined.optimizer_version and not combined.optimizer_statistics_package: + if base is None and merge is None: return None + + base = _to_query_options(base) + merge = _to_query_options(merge) + if base is None: + return merge + if merge is None: + return base + + combined = ExecuteSqlRequest.QueryOptions() + combined_pb = type(combined).pb(combined) + combined_pb.CopyFrom(type(base).pb(base)) + combined_pb.MergeFrom(type(merge).pb(merge)) return combined @@ -312,8 +343,9 @@ def _assert_numeric_precision_and_scale(value): :raises NotSupportedError: If value is not within supported precision or scale of Spanner. """ - scale = value.as_tuple().exponent - precision = len(value.as_tuple().digits) + decimal_tuple = value.as_tuple() + scale = decimal_tuple.exponent + precision = len(decimal_tuple.digits) if scale < -9: raise ValueError(NUMERIC_MAX_SCALE_ERR_MSG.format(abs(scale))) @@ -359,6 +391,63 @@ def _datetime_to_rfc3339_nanoseconds(value): return "{}.{}Z".format(value.isoformat(sep="T", timespec="seconds"), nanos) +def _make_list_value_pb(values): + """Construct of ListValue protobufs. + + :type values: list of scalar + :param values: Row data + + :rtype: :class:`~google.protobuf.struct_pb2.ListValue` + :returns: protobuf + """ + return ListValue(values=[_make_value_pb(value) for value in values]) + + +def _encode_float(value): + if math.isfinite(value): + return Value(number_value=value) + if math.isnan(value): + return Value(string_value="NaN") + return Value(string_value="Infinity" if value > 0 else "-Infinity") + + +def _encode_decimal(value): + _assert_numeric_precision_and_scale(value) + return Value(string_value=str(value)) + + +def _encode_bytes(value): + return Value(string_value=_validate_and_decode_bytes(value)) + + +def _encode_json_object(value): + serialized = value.serialize() + if serialized is None: + return Value(null_value=NULL_VALUE) + return Value(string_value=serialized) + + +_TYPE_ENCODERS = { + str: lambda value: Value(string_value=value), + int: lambda value: Value(string_value=str(value)), + bool: lambda value: Value(bool_value=value), + float: _encode_float, + bytes: _encode_bytes, + datetime.date: lambda value: Value(string_value=value.isoformat()), + datetime.datetime: lambda value: Value(string_value=_datetime_to_rfc3339(value)), + datetime_helpers.DatetimeWithNanoseconds: lambda value: Value( + string_value=_datetime_to_rfc3339_nanoseconds(value) + ), + decimal.Decimal: _encode_decimal, + uuid.UUID: lambda value: Value(string_value=str(value)), + Interval: lambda value: Value(string_value=str(value)), + list: lambda value: Value(list_value=_make_list_value_pb(value)), + tuple: lambda value: Value(list_value=_make_list_value_pb(value)), + ListValue: lambda value: Value(list_value=value), + JsonObject: _encode_json_object, +} + + def _make_value_pb(value): """Helper for :func:`_make_list_value_pbs`. @@ -370,7 +459,18 @@ def _make_value_pb(value): :raises ValueError: if value is not of a known scalar type. """ if value is None: - return Value(null_value="NULL_VALUE") + return Value(null_value=NULL_VALUE) + + try: + encoder = _TYPE_ENCODERS[type(value)] + except KeyError: + pass + else: + return encoder(value) + + # Note: The fallback isinstance chain is retained to support subclasses, + # custom mock/proxy objects, and dynamic AST inspection in + # tests/unit/spanner_dbapi/test_partition_helper.py. if isinstance(value, (list, tuple)): return Value(list_value=_make_list_value_pb(value)) if isinstance(value, bool): @@ -378,14 +478,7 @@ def _make_value_pb(value): if isinstance(value, int): return Value(string_value=str(value)) if isinstance(value, float): - if math.isnan(value): - return Value(string_value="NaN") - if math.isinf(value): - if value > 0: - return Value(string_value="Infinity") - else: - return Value(string_value="-Infinity") - return Value(number_value=value) + return _encode_float(value) if isinstance(value, datetime_helpers.DatetimeWithNanoseconds): return Value(string_value=_datetime_to_rfc3339_nanoseconds(value)) if isinstance(value, datetime.datetime): @@ -393,27 +486,21 @@ def _make_value_pb(value): if isinstance(value, datetime.date): return Value(string_value=value.isoformat()) if isinstance(value, bytes): - value = _try_to_coerce_bytes(value) - return Value(string_value=value) + return _encode_bytes(value) if isinstance(value, str): return Value(string_value=value) if isinstance(value, ListValue): return Value(list_value=value) if isinstance(value, decimal.Decimal): - _assert_numeric_precision_and_scale(value) - return Value(string_value=str(value)) + return _encode_decimal(value) if isinstance(value, JsonObject): - value = value.serialize() - if value is None: - return Value(null_value="NULL_VALUE") - else: - return Value(string_value=value) + return _encode_json_object(value) if isinstance(value, Message): - value = value.SerializeToString() - if value is None: - return Value(null_value="NULL_VALUE") + serialized = value.SerializeToString() + if serialized is None: + return Value(null_value=NULL_VALUE) else: - return Value(string_value=base64.b64encode(value)) + return Value(string_value=base64.b64encode(serialized).decode("utf-8")) if isinstance(value, Interval): return Value(string_value=str(value)) if isinstance(value, uuid.UUID): @@ -422,18 +509,6 @@ def _make_value_pb(value): raise ValueError("Unknown type: %s" % (value,)) -def _make_list_value_pb(values): - """Construct of ListValue protobufs. - - :type values: list of scalar - :param values: Row data - - :rtype: :class:`~google.protobuf.struct_pb2.ListValue` - :returns: protobuf - """ - return ListValue(values=[_make_value_pb(value) for value in values]) - - def _make_list_value_pbs(values): """Construct a sequence of ListValue protobufs. @@ -505,47 +580,28 @@ def _get_type_decoder(field_type, field_name, column_info=None): """ type_code = field_type.code - # Note: STRING and BOOL use operator.attrgetter because direct attribute extraction - # is faster in Python. Other types require type transformation, so they use lambdas. - if type_code == TypeCode.STRING: - return operator.attrgetter("string_value") - elif type_code == TypeCode.BYTES: - return lambda value_pb: value_pb.string_value.encode("utf8") - elif type_code == TypeCode.BOOL: - return operator.attrgetter("bool_value") - elif type_code == TypeCode.INT64: - return lambda value_pb: int(value_pb.string_value) - elif type_code == TypeCode.FLOAT64: - return _parse_float - elif type_code == TypeCode.FLOAT32: - return _parse_float - elif type_code == TypeCode.DATE: - return lambda value_pb: _date_fromisoformat(value_pb.string_value) - elif type_code == TypeCode.TIMESTAMP: - return _parse_timestamp - elif type_code == TypeCode.NUMERIC: - return lambda value_pb: _Decimal(value_pb.string_value) - elif type_code == TypeCode.JSON: - return lambda value_pb: _json_from_str(value_pb.string_value) - elif type_code == TypeCode.UUID: - return lambda value_pb: _uuid_UUID(value_pb.string_value) - elif type_code == TypeCode.PROTO: + try: + type_code_integer = int(type_code) + except (TypeError, ValueError): + type_code_integer = None + + if type_code_integer in _SCALAR_DECODERS: + return _SCALAR_DECODERS[type_code_integer] + elif type_code_integer == _PROTO_TYPE_CODE: return lambda value_pb: _parse_proto(value_pb, column_info, field_name) - elif type_code == TypeCode.ENUM: + elif type_code_integer == _ENUM_TYPE_CODE: return lambda value_pb: _parse_proto_enum(value_pb, column_info, field_name) - elif type_code == TypeCode.ARRAY: + elif type_code_integer == _ARRAY_TYPE_CODE: element_decoder = _get_type_decoder( field_type.array_element_type, field_name, column_info ) return lambda value_pb: _parse_array(value_pb, element_decoder) - elif type_code == TypeCode.STRUCT: + elif type_code_integer == _STRUCT_TYPE_CODE: element_decoders = [ _get_type_decoder(item_field.type_, field_name, column_info) for item_field in field_type.struct_type.fields ] return lambda value_pb: _parse_struct(value_pb, element_decoders) - elif type_code == TypeCode.INTERVAL: - return _parse_interval else: raise ValueError("Unknown type: %s" % (field_type,)) @@ -702,6 +758,29 @@ def _parse_interval(value_pb): return Interval.from_str(value_pb) +# Note: STRING and BOOL use operator.attrgetter because direct attribute extraction +# is faster in Python. Other types require type transformation, so they use lambdas. +_SCALAR_DECODERS = { + int(TypeCode.STRING): operator.attrgetter("string_value"), + int(TypeCode.BYTES): lambda value_pb: value_pb.string_value.encode("utf8"), + int(TypeCode.BOOL): operator.attrgetter("bool_value"), + int(TypeCode.INT64): lambda value_pb: int(value_pb.string_value), + int(TypeCode.FLOAT64): _parse_float, + int(TypeCode.FLOAT32): _parse_float, + int(TypeCode.DATE): lambda value_pb: _date_fromisoformat(value_pb.string_value), + int(TypeCode.TIMESTAMP): _parse_timestamp, + int(TypeCode.NUMERIC): lambda value_pb: _Decimal(value_pb.string_value), + int(TypeCode.JSON): lambda value_pb: _json_from_str(value_pb.string_value), + int(TypeCode.UUID): lambda value_pb: _uuid_UUID(value_pb.string_value), + int(TypeCode.INTERVAL): _parse_interval, +} + +_PROTO_TYPE_CODE = int(TypeCode.PROTO) +_ENUM_TYPE_CODE = int(TypeCode.ENUM) +_ARRAY_TYPE_CODE = int(TypeCode.ARRAY) +_STRUCT_TYPE_CODE = int(TypeCode.STRUCT) + + class _SessionWrapper(object): """Base class for objects wrapping a session. @@ -856,7 +935,8 @@ def _delay_until_retry(exc, deadline, attempts, default_retry_delay=None): :param attempts: number of call retries """ - cause = exc.errors[0] + errors = getattr(exc, "errors", None) + cause = errors[0] if errors else exc now = time.time() if now >= deadline: raise @@ -1124,3 +1204,128 @@ def _create_experimental_host_transport( client_key, interceptors=interceptors, ) + + +_STREAM_DRAIN_QUEUE_SIZE = 512 +_STREAM_DRAIN_WORKER_COUNT = 8 + + +class _BoundedStreamDrainer: + """Bounded background drainer for synchronous gRPC streams. + + Uses a fixed pool of daemon worker threads and a bounded queue to drain + completed streams to EOF, allowing callers to return early upon seeing + PartialResultSet.last without blocking on trailing gRPC frames. + """ + + def __init__( + self, + queue_size: int = _STREAM_DRAIN_QUEUE_SIZE, + worker_count: int = _STREAM_DRAIN_WORKER_COUNT, + ): + self._queue_size = queue_size + self._worker_count = worker_count + self._lock = threading.Lock() + self._reset() + + def _reset(self): + self._queue = queue.Queue(maxsize=self._queue_size) + self._started = False + self._stopped = False + self._workers = [] + + def _reset_after_fork(self): + self._lock = threading.Lock() + self._reset() + + def _ensure_started(self): + if not self._started and not self._stopped: + with self._lock: + if not self._started and not self._stopped: + self._started = True + try: + for index in range(self._worker_count): + worker = threading.Thread( + target=self._worker_loop, + name=f"spanner-stream-drainer-{index}", + daemon=True, + ) + worker.start() + self._workers.append(worker) + except Exception: + if not self._workers: + self._started = False + raise + + def _worker_loop(self): + while True: + iterator = self._queue.get() + if iterator is None: + self._queue.task_done() + break + try: + for _ in iterator: + pass + except Exception: + pass + finally: + self._queue.task_done() + + def drain(self, iterator): + if iterator is None: + return + + with self._lock: + stopped = self._stopped + + # If already shut down or during interpreter exit, drain inline on caller thread. + if stopped: + try: + for _ in iterator: + pass + except Exception: + pass + return + + try: + self._ensure_started() + self._queue.put_nowait(iterator) + except Exception: + # Under extreme bursts where the queue is temporarily full, or if thread + # creation fails (e.g. in restricted environments or during shutdown), + # drain inline on the caller thread rather than cancelling a successful query. + # Because trailers are delivered in sub-millisecond time (<0.5ms), + # inline draining adds negligible latency while guaranteeing status OK. + try: + for _ in iterator: + pass + except Exception: + pass + + def shutdown(self): + """Cleanly terminate worker threads during interpreter shutdown.""" + with self._lock: + if self._stopped: + return + self._stopped = True + if self._started: + for _ in range(len(self._workers)): + try: + self._queue.put_nowait(None) + except queue.Full: + pass + + +_GLOBAL_STREAM_DRAINER = _BoundedStreamDrainer() +atexit.register(_GLOBAL_STREAM_DRAINER.shutdown) +if hasattr(os, "register_at_fork"): + os.register_at_fork(after_in_child=_GLOBAL_STREAM_DRAINER._reset_after_fork) + + +def _drain_stream(iterator): + """Drain a stream iterator to EOF in the background. + + Called when PartialResultSet.last is True to allow the caller to return immediately + while consuming trailing gRPC metadata so the stream terminates cleanly with status OK. + """ + _GLOBAL_STREAM_DRAINER.drain(iterator) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_opentelemetry_tracing.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_opentelemetry_tracing.py index 372f40d02be9..ddca0656eb0b 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_opentelemetry_tracing.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_opentelemetry_tracing.py @@ -15,6 +15,7 @@ """Manages OpenTelemetry trace creation and handling""" import os +from collections.abc import Mapping from contextlib import contextmanager from datetime import datetime @@ -40,6 +41,16 @@ ) +def _is_tracer_noop(tracer: trace.Tracer) -> bool: + """Check if a tracer is a no-op tracer or an uninitialized proxy tracer.""" + if tracer is None or isinstance(tracer, trace.NoOpTracer): + return True + if isinstance(tracer, trace.ProxyTracer): + underlying = getattr(tracer, "_tracer", None) + return isinstance(underlying, trace.NoOpTracer) + return False + + def get_tracer(tracer_provider=None): """ get_tracer is a utility to unify and simplify retrieval of the tracer, without @@ -63,31 +74,56 @@ def trace_call( session._last_use_time = datetime.now() tracer_provider = None + enable_end_to_end_tracing = False - # By default enable_extended_tracing=True because in a bid to minimize - # breaking changes and preserve legacy behavior, we are keeping it turned - # on by default. - enable_extended_tracing = True + has_options = isinstance(observability_options, Mapping) + if has_options: + tracer_provider = observability_options.get("tracer_provider", None) + enable_end_to_end_tracing = observability_options.get( + "enable_end_to_end_tracing", False + ) - enable_end_to_end_tracing = False + if end_to_end_tracing_globally_enabled: + enable_end_to_end_tracing = True + + tracer = get_tracer(tracer_provider) + # Fast path: when no TracerProvider is registered or a no-op tracer is active, + # skip attribute dictionary construction and span creation entirely. + if _is_tracer_noop(tracer): + current_span = trace.get_current_span() + if ( + current_span is not trace.INVALID_SPAN + and current_span.get_span_context().is_valid + ): + child_span = trace.NonRecordingSpan(current_span.get_span_context()) + with trace.use_span(child_span): + with MetricsCapture(): + if enable_end_to_end_tracing: + _metadata_with_span_context(metadata) + yield child_span + return + + with MetricsCapture(): + if enable_end_to_end_tracing: + _metadata_with_span_context(metadata) + yield trace.INVALID_SPAN + return + + # Slow path: resolve attributes and configure span when tracing is active. + enable_extended_tracing = True db_name = "" - cloud_region = None + if session and getattr(session, "_database", None): db_name = session._database.name - if isinstance(observability_options, dict): # Avoid false positives with mock.Mock - tracer_provider = observability_options.get("tracer_provider", None) + if has_options: enable_extended_tracing = observability_options.get( "enable_extended_tracing", enable_extended_tracing ) - enable_end_to_end_tracing = observability_options.get( - "enable_end_to_end_tracing", enable_end_to_end_tracing - ) db_name = observability_options.get("db_name", db_name) cloud_region = _get_cloud_region() - tracer = get_tracer(tracer_provider) # Set base attributes that we know for every trace created attributes = { @@ -119,9 +155,6 @@ def trace_call( if not enable_extended_tracing: attributes.pop("db.statement", False) - if end_to_end_tracing_globally_enabled: - enable_end_to_end_tracing = True - with tracer.start_as_current_span( name, kind=trace.SpanKind.CLIENT, attributes=attributes ) as span: @@ -155,4 +188,5 @@ def get_current_span(): def add_span_event(span, event_name, event_attributes=None): - span.add_event(event_name, event_attributes) + if span and span.is_recording(): + span.add_event(event_name, event_attributes) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/database_sessions_manager.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/database_sessions_manager.py index 1b2c6231f46e..e9202f581a53 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/database_sessions_manager.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/database_sessions_manager.py @@ -121,8 +121,8 @@ def _get_multiplexed_session(self) -> Session: """Returns a multiplexed session from the database session manager. If the multiplexed session is not defined, creates a new multiplexed - session and starts a maintenance thread to periodically delete and - recreate it so that it remains valid. Otherwise, simply returns the + session and starts a maintenance thread to periodically rotate + it so that it remains valid. Otherwise, simply returns the current multiplexed session. :rtype: :class:`~google.cloud.spanner_v1.session.Session` @@ -162,8 +162,8 @@ def _build_maintenance_thread( self, session: Optional[Session] = None ) -> CrossSync._Sync_Impl.Task: """Builds and returns a multiplexed session maintenance thread for - the database session manager. This thread will periodically delete - and recreate the multiplexed session to ensure that it is always valid. + the database session manager. This thread will periodically rotate + the multiplexed session to ensure that it is always valid. :type session: :class:`~google.cloud.spanner_v1.session.Session` :param session: (Optional) The multiplexed session to maintain. @@ -196,24 +196,17 @@ def _rotate_multiplexed_session(self) -> bool: return False with self._multiplexed_session_lock: - old_session = self._multiplexed_session self._multiplexed_session = new_session - if old_session is not None: - try: - CrossSync._Sync_Impl.run_if_async(old_session.delete) - except Exception: - pass - return True @staticmethod def _maintain_multiplexed_session(session_manager_ref) -> None: """Maintains the multiplexed session for the database session manager. - This method will delete and recreate the referenced database session manager's + This method will periodically rotate the referenced database session manager's multiplexed session to ensure that it is always valid. The method will run until - the database session manager is deleted or the multiplexed session is deleted. + the database session manager is garbage collected or the session manager is closed. :type session_manager_ref: :class:`_weakref.ReferenceType` :param session_manager_ref: A weak reference to the database session manager.""" @@ -269,7 +262,4 @@ def close(self) -> None: self._multiplexed_session_terminate_event.set() if self._multiplexed_session_thread is not None: self._multiplexed_session_thread.join() - if self._multiplexed_session is not None: - session_to_delete = self._multiplexed_session - self._multiplexed_session = None - session_to_delete.delete() + self._multiplexed_session = None diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/constants.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/constants.py index fa5f5ca4d98d..39445bca32c2 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/constants.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/constants.py @@ -12,6 +12,8 @@ # See the License for the specific language governing permissions and # limitations under the License. +from typing import Any + BUILT_IN_METRICS_METER_NAME = "gax-python" NATIVE_METRICS_PREFIX = "spanner.googleapis.com/internal/client" SPANNER_RESOURCE_TYPE = "spanner_instance_client" @@ -21,6 +23,18 @@ GOOGLE_CLOUD_REGION_GLOBAL = "global" SPANNER_METHOD_PREFIX = "/google.spanner.v1." + +def _safe_decode_utf8(value: Any) -> str: + """Safely decode bytes to str or return str representation without raising.""" + if value is None: + return "" + if isinstance(value, str): + return value + if isinstance(value, bytes): + return value.decode("utf-8", errors="replace") + return str(value) + + # Monitored resource labels MONITORED_RES_LABEL_KEY_PROJECT = "project_id" MONITORED_RES_LABEL_KEY_INSTANCE = "instance_id" diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_capture.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_capture.py index 77c1f86c4dee..640aa2b93f67 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_capture.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_capture.py @@ -42,6 +42,7 @@ def __init__(self, resource_info: dict = None): resource_info (dict): Optional dictionary containing project, instance and database info. """ self._resource_info = resource_info + self._token = None def __enter__(self): """Enter the runtime context related to this object. @@ -88,15 +89,16 @@ def __exit__(self, exc_type, exc_value, traceback): Returns: bool: False to propagate the exception if any occurred. """ - # Short circuit out if metrics are disable - if not SpannerMetricsTracerFactory().enabled: + token = self._token + if token is None: return False - tracer = SpannerMetricsTracerFactory.get_current_tracer() - if tracer: - tracer.record_operation_completion() - - # Reset the context var using the token - if getattr(self, "_token", None): - SpannerMetricsTracerFactory.reset_current_tracer(self._token) + try: + tracer = SpannerMetricsTracerFactory.get_current_tracer() + if tracer: + tracer.record_operation_completion() + finally: + # Reset the context var using the token + SpannerMetricsTracerFactory.reset_current_tracer(token) + self._token = None return False # Propagate the exception if any diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_interceptor.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_interceptor.py index 1205b09c1840..6c2ca7855380 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_interceptor.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_interceptor.py @@ -14,19 +14,55 @@ """Interceptor for collecting Cloud Spanner metrics.""" +import functools import inspect import logging import re +import threading from typing import Any, Dict import grpc from grpc_interceptor import ClientInterceptor -from .constants import GOOGLE_CLOUD_RESOURCE_KEY, SPANNER_METHOD_PREFIX +from .constants import ( + GOOGLE_CLOUD_RESOURCE_KEY, + SPANNER_METHOD_PREFIX, + _safe_decode_utf8, +) from .spanner_metrics_tracer_factory import SpannerMetricsTracerFactory logger = logging.getLogger(__name__) +_RESOURCE_PATH_PATTERN = re.compile( + r"^projects/(?P[^/]+)(/instances/(?P[^/]+))?(/databases/(?P[^/]+))?(/sessions/(?P[^/]+))?.*$" +) + +_RESOURCE_KEY_STR = GOOGLE_CLOUD_RESOURCE_KEY +_RESOURCE_KEY_BYTES = GOOGLE_CLOUD_RESOURCE_KEY.encode("utf-8") + + +@functools.lru_cache(maxsize=64) +def _format_method_name_str(method_str: str) -> str: + return method_str.removeprefix(SPANNER_METHOD_PREFIX).replace("/", ".") + + +def _format_method_name(method_name_input: Any) -> str: + """Format method name to be Spanner. with caching.""" + if isinstance(method_name_input, str): + return _format_method_name_str(method_name_input) + return _format_method_name_str(_safe_decode_utf8(method_name_input)) + + +@functools.lru_cache(maxsize=128) +def _parse_resource_path_cached(path: str) -> Dict[str, str]: + """Parse resource path using regex with LRU caching.""" + match = _RESOURCE_PATH_PATTERN.match(path) + if match: + return { + key: value for key, value in match.groupdict().items() if value is not None + } + return {} + class MetricsInterceptor(ClientInterceptor): """Interceptor that collects metrics for Cloud Spanner operations.""" @@ -41,15 +77,10 @@ def _parse_resource_path(path: str) -> dict: Returns: dict: Extracted resource components """ - # Match paths like: - # projects/{project}/instances/{instance}/databases/{database}/sessions/{session} - # projects/{project}/instances/{instance}/databases/{database} - # projects/{project}/instances/{instance} - pattern = r"^projects/(?P[^/]+)(/instances/(?P[^/]+))?(/databases/(?P[^/]+))?(/sessions/(?P[^/]+))?.*$" - match = re.match(pattern, path) - if match: - return {k: v for k, v in match.groupdict().items() if v is not None} - return {} + if not path or not isinstance(path, str): + return {} + + return _parse_resource_path_cached(path).copy() @staticmethod def _extract_resource_from_path(metadata: Any) -> Dict[str, str]: @@ -65,29 +96,43 @@ def _extract_resource_from_path(metadata: Any) -> Dict[str, str]: if not metadata: return {} - items = metadata.items() if isinstance(metadata, dict) else metadata path = "" - - for key, value in items: - key_str = key.decode("utf-8") if isinstance(key, bytes) else str(key) - if key_str == GOOGLE_CLOUD_RESOURCE_KEY: - path = value.decode("utf-8") if isinstance(value, bytes) else str(value) - break - - resources = MetricsInterceptor._parse_resource_path(path) - return resources + if isinstance(metadata, dict): + raw_path = metadata.get(_RESOURCE_KEY_STR) or metadata.get( + _RESOURCE_KEY_BYTES + ) + if raw_path is not None: + path = _safe_decode_utf8(raw_path) + else: + try: + metadata_iter = iter(metadata) + except TypeError: + return {} + for item in metadata_iter: + if not (isinstance(item, (list, tuple)) and len(item) == 2): + continue + key, value = item + if key == _RESOURCE_KEY_STR or key == _RESOURCE_KEY_BYTES: + path = _safe_decode_utf8(value) + break + + return MetricsInterceptor._parse_resource_path(path) @staticmethod - def _set_metrics_tracer_attributes(resources: Dict[str, str]) -> None: + def _set_metrics_tracer_attributes( + resources: Dict[str, str], tracer: Any = None + ) -> None: """ Sets the metric tracer attributes based on the provided resources. - This method updates the current metric tracer's attributes with the project, instance, and database information extracted from the resources dictionary. If the current metric tracer is not set, the method does nothing. + This method updates the metric tracer's attributes with the project, instance, and database information extracted from the resources dictionary. If the metric tracer is not set, the method does nothing. Args: resources (Dict[str, str]): A dictionary containing project, instance, and database information. + tracer (Any, optional): The metric tracer instance. If not provided, retrieves the current tracer. """ - tracer = SpannerMetricsTracerFactory.get_current_tracer() + if tracer is None: + tracer = SpannerMetricsTracerFactory.get_current_tracer() if tracer is None: return @@ -99,6 +144,23 @@ def _set_metrics_tracer_attributes(resources: Dict[str, str]) -> None: if "database" in resources: tracer.set_database(resources["database"]) + @staticmethod + def _prepare_attempt(tracer: Any, call_details: Any) -> None: + """Prepare tracer attributes and record attempt start from call details.""" + if not ( + tracer.client_attributes.get("project_id") + and tracer.client_attributes.get("instance_id") + and tracer.client_attributes.get("database") + ): + resources = MetricsInterceptor._extract_resource_from_path( + call_details.metadata + ) + MetricsInterceptor._set_metrics_tracer_attributes(resources, tracer=tracer) + + method_name = _format_method_name(call_details.method) + tracer.set_method(method_name) + tracer.record_attempt_start() + def intercept(self, invoked_method, request_or_iterator, call_details): """Intercept gRPC calls to collect metrics. @@ -115,25 +177,8 @@ def intercept(self, invoked_method, request_or_iterator, call_details): if tracer is None or not factory.enabled: return invoked_method(request_or_iterator, call_details) - # Setup Metric Tracer attributes from call details - ## Extract Project / Instance / Database from header information if not already set - if not ( - tracer.client_attributes.get("project_id") - and tracer.client_attributes.get("instance_id") - and tracer.client_attributes.get("database") - ): - resources = self._extract_resource_from_path(call_details.metadata) - self._set_metrics_tracer_attributes(resources) - - ## Format method to be be spanner. - method_str = call_details.method - method_name = method_str.removeprefix(SPANNER_METHOD_PREFIX).replace("/", ".") - - tracer.set_method(method_name) - tracer.record_attempt_start() - + self._prepare_attempt(tracer, call_details) response = invoked_method(request_or_iterator, call_details) - return _wrap_response(response, tracer) @@ -197,24 +242,7 @@ async def _async_intercept( if tracer is None or not factory.enabled: return await continuation(call_details, request_or_iterator) - if not ( - tracer.client_attributes.get("project_id") - and tracer.client_attributes.get("instance_id") - and tracer.client_attributes.get("database") - ): - resources = MetricsInterceptor._extract_resource_from_path( - call_details.metadata - ) - MetricsInterceptor._set_metrics_tracer_attributes(resources) - - method_str = call_details.method - if isinstance(method_str, bytes): - method_str = method_str.decode("utf-8") - method_name = method_str.removeprefix(SPANNER_METHOD_PREFIX).replace("/", ".") - - tracer.set_method(method_name) - tracer.record_attempt_start() - + MetricsInterceptor._prepare_attempt(tracer, call_details) response = await continuation(call_details, request_or_iterator) if hasattr(response, "__anext__"): return _AsyncStreamingResponseWrapper(response, tracer) @@ -230,6 +258,7 @@ def __init__(self, response, tracer): self._tracer = tracer self._metrics_recorded = False self._iterator = None + self._lock = threading.Lock() def __iter__(self): self._iterator = iter(self._response) @@ -248,9 +277,10 @@ def __next__(self): raise def _record_metrics(self): - if self._metrics_recorded: - return - self._metrics_recorded = True + with self._lock: + if self._metrics_recorded: + return + self._metrics_recorded = True try: self._tracer.record_attempt_completion() metadata = [] @@ -263,9 +293,40 @@ def _record_metrics(self): except Exception as e: logger.warning(f"Failed to record metrics: {e}") + def cancel(self, *args, **kwargs): + cancelled = None + if hasattr(self._response, "cancel"): + cancelled = self._response.cancel(*args, **kwargs) + if cancelled is not False: + with self._lock: + if self._metrics_recorded: + return cancelled + self._metrics_recorded = True + try: + self._tracer.record_attempt_completion( + status=grpc.StatusCode.CANCELLED.name + ) + metadata = [] + if hasattr(self._response, "initial_metadata"): + try: + metadata.extend(self._response.initial_metadata() or []) + except Exception: + pass + if metadata: + self._tracer.record_front_end_metrics(metadata) + except Exception as e: + logger.warning(f"Failed to record metrics on cancel: {e}") + return cancelled + def __del__(self): + with self._lock: + if self._metrics_recorded: + return + self._metrics_recorded = True try: - self._record_metrics() + self._tracer.record_attempt_completion( + status=grpc.StatusCode.CANCELLED.name + ) except Exception: pass @@ -273,19 +334,47 @@ def __getattr__(self, name): return getattr(self._response, name) -class _AsyncUnaryResponseWrapper(grpc.aio.UnaryUnaryCall): - """Wrapper for async unary RPC response to defer metrics recording until awaited.""" +class _BaseAsyncResponseWrapper: + """Base wrapper for async RPC responses to defer metrics recording.""" def __init__(self, response, tracer): self._response = response self._tracer = tracer self._metrics_recorded = False + self._lock = threading.Lock() def add_done_callback(self, *args, **kwargs): return getattr(self._response, "add_done_callback")(*args, **kwargs) def cancel(self, *args, **kwargs): - return getattr(self._response, "cancel")(*args, **kwargs) + cancel_fn = getattr(self._response, "cancel", None) + cancelled = cancel_fn(*args, **kwargs) if cancel_fn else True + if cancelled is not False: + with self._lock: + if self._metrics_recorded: + return cancelled + self._metrics_recorded = True + try: + self._tracer.record_attempt_completion( + status=grpc.StatusCode.CANCELLED.name + ) + metadata = [] + if hasattr(self._response, "initial_metadata"): + try: + metadata_result = self._response.initial_metadata() + if inspect.isawaitable(metadata_result): + getattr(metadata_result, "close", lambda: None)() + else: + metadata.extend(metadata_result or []) + except Exception as e: + logger.warning( + f"Failed to retrieve initial metadata on cancel: {e}" + ) + if metadata: + self._tracer.record_front_end_metrics(metadata) + except Exception as e: + logger.warning(f"Failed to record metrics on cancel: {e}") + return cancelled def cancelled(self, *args, **kwargs): return getattr(self._response, "cancelled")(*args, **kwargs) @@ -311,28 +400,26 @@ def trailing_metadata(self, *args, **kwargs): def wait_for_connection(self, *args, **kwargs): return getattr(self._response, "wait_for_connection")(*args, **kwargs) - def __await__(self): - async def _wait(): - try: - return await self._response - finally: - await self._record_metrics() + def write(self, *args, **kwargs): + return getattr(self._response, "write")(*args, **kwargs) - return _wait().__await__() + def done_writing(self, *args, **kwargs): + return getattr(self._response, "done_writing")(*args, **kwargs) async def _record_metrics(self): - if self._metrics_recorded: - return - self._metrics_recorded = True + with self._lock: + if self._metrics_recorded: + return + self._metrics_recorded = True try: self._tracer.record_attempt_completion() metadata = [] if hasattr(self._response, "initial_metadata"): try: - res = self._response.initial_metadata() - if inspect.isawaitable(res): - res = await res - metadata.extend(res or []) + metadata_result = self._response.initial_metadata() + if inspect.isawaitable(metadata_result): + metadata_result = await metadata_result + metadata.extend(metadata_result or []) except Exception as e: logger.warning(f"Failed to retrieve initial metadata: {e}") self._tracer.record_front_end_metrics(metadata) @@ -340,69 +427,52 @@ async def _record_metrics(self): logger.warning(f"Failed to record metrics: {e}") def __del__(self): - if not self._metrics_recorded: + with self._lock: + if self._metrics_recorded: + return self._metrics_recorded = True - try: - self._tracer.record_attempt_completion() - except Exception: - pass + try: + self._tracer.record_attempt_completion( + status=grpc.StatusCode.CANCELLED.name + ) + except Exception: + pass def __getattr__(self, name): return getattr(self._response, name) +class _AsyncUnaryResponseWrapper( + _BaseAsyncResponseWrapper, + grpc.aio.UnaryUnaryCall, + grpc.aio.StreamUnaryCall, +): + """Wrapper for async unary RPC response to defer metrics recording until awaited.""" + + def __await__(self): + async def _wait(): + try: + return await self._response + finally: + await self._record_metrics() + + return _wait().__await__() + + class _AsyncStreamingResponseWrapper( + _BaseAsyncResponseWrapper, grpc.aio.UnaryStreamCall, - grpc.aio.StreamUnaryCall, grpc.aio.StreamStreamCall, ): """Wrapper for async streaming RPC response iterators to defer metrics recording.""" def __init__(self, response, tracer): - self._response = response - self._tracer = tracer - self._metrics_recorded = False + super().__init__(response, tracer) self._iterator = None - def add_done_callback(self, *args, **kwargs): - return getattr(self._response, "add_done_callback")(*args, **kwargs) - - def cancel(self, *args, **kwargs): - return getattr(self._response, "cancel")(*args, **kwargs) - - def cancelled(self, *args, **kwargs): - return getattr(self._response, "cancelled")(*args, **kwargs) - - def code(self, *args, **kwargs): - return getattr(self._response, "code")(*args, **kwargs) - - def details(self, *args, **kwargs): - return getattr(self._response, "details")(*args, **kwargs) - - def done(self, *args, **kwargs): - return getattr(self._response, "done")(*args, **kwargs) - - def initial_metadata(self, *args, **kwargs): - return getattr(self._response, "initial_metadata")(*args, **kwargs) - - def time_remaining(self, *args, **kwargs): - return getattr(self._response, "time_remaining")(*args, **kwargs) - - def trailing_metadata(self, *args, **kwargs): - return getattr(self._response, "trailing_metadata")(*args, **kwargs) - - def wait_for_connection(self, *args, **kwargs): - return getattr(self._response, "wait_for_connection")(*args, **kwargs) - def read(self, *args, **kwargs): return getattr(self._response, "read")(*args, **kwargs) - def write(self, *args, **kwargs): - return getattr(self._response, "write")(*args, **kwargs) - - def done_writing(self, *args, **kwargs): - return getattr(self._response, "done_writing")(*args, **kwargs) - def __aiter__(self): if hasattr(self._response, "__aiter__"): self._iterator = self._response.__aiter__() @@ -424,33 +494,3 @@ async def __anext__(self): except Exception: await self._record_metrics() raise - - async def _record_metrics(self): - if self._metrics_recorded: - return - self._metrics_recorded = True - try: - self._tracer.record_attempt_completion() - metadata = [] - if hasattr(self._response, "initial_metadata"): - try: - res = self._response.initial_metadata() - if inspect.isawaitable(res): - res = await res - metadata.extend(res or []) - except Exception as e: - logger.warning(f"Failed to retrieve initial metadata: {e}") - self._tracer.record_front_end_metrics(metadata) - except Exception as e: - logger.warning(f"Failed to record metrics: {e}") - - def __del__(self): - if not self._metrics_recorded: - self._metrics_recorded = True - try: - self._tracer.record_attempt_completion() - except Exception: - pass - - def __getattr__(self, name): - return getattr(self._response, name) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_tracer.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_tracer.py index 6106fa6e18b0..276d0a09d0c8 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_tracer.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/metrics/metrics_tracer.py @@ -22,7 +22,7 @@ import os import re from datetime import datetime -from typing import Any, Dict, Optional +from typing import Any, Dict, Optional, Tuple from grpc import StatusCode @@ -38,6 +38,7 @@ MONITORED_RES_LABEL_KEY_INSTANCE_CONFIG, MONITORED_RES_LABEL_KEY_LOCATION, MONITORED_RES_LABEL_KEY_PROJECT, + _safe_decode_utf8, ) try: @@ -47,6 +48,73 @@ except ImportError: # pragma: NO COVER HAS_OPENTELEMETRY_INSTALLED = False +_GFE_TIMING_PATTERN = re.compile(r"(? Optional[int]: + """Search for the pattern in text and parse the captured latency to int.""" + match = pattern.search(text) + if match: + try: + return int(float(match.group(1))) + except ValueError: + pass + return None + + +class _ObservableDict(dict): + """A dictionary that invokes an invalidation callback upon modification.""" + + def __init__(self, *args, on_change=None, **kwargs): + super().__init__(*args, **kwargs) + self._on_change = on_change + + def __setitem__(self, key, value): + super().__setitem__(key, value) + if self._on_change is not None: + self._on_change() + + def __delitem__(self, key): + super().__delitem__(key) + if self._on_change is not None: + self._on_change() + + def update(self, *args, **kwargs): + super().update(*args, **kwargs) + if self._on_change is not None: + self._on_change() + + def clear(self): + super().clear() + if self._on_change is not None: + self._on_change() + + def pop(self, *args, **kwargs): + result = super().pop(*args, **kwargs) + if self._on_change is not None: + self._on_change() + return result + + def popitem(self): + result = super().popitem() + if self._on_change is not None: + self._on_change() + return result + + def setdefault(self, key, default=None): + if key not in self: + result = super().setdefault(key, default) + if self._on_change is not None: + self._on_change() + return result + return super().setdefault(key, default) + + def copy(self): + return dict(self) + class MetricAttemptTracer: """ @@ -225,7 +293,10 @@ def __init__( instrument_afe_connectivity_error_count (Counter): Instrument for counting AFE connectivity errors. """ self.current_op = MetricOpTracer() - self._client_attributes = client_attributes + self._client_attributes = _ObservableDict( + client_attributes or {}, + on_change=self._invalidate_attribute_cache, + ) self._instrument_attempt_latency = instrument_attempt_latency self._instrument_attempt_counter = instrument_attempt_counter self._instrument_operation_latency = instrument_operation_latency @@ -242,6 +313,19 @@ def __init__( self.afe_server_timing_enabled = ( os.environ.get("SPANNER_DISABLE_AFE_SERVER_TIMING", "").lower() != "true" ) + self._cached_attempt_attributes: Optional[dict] = None + self._cached_attempt_status: Optional[str] = None + self._cached_attempt: Optional[MetricAttemptTracer] = None + self._cached_operation_attributes: Optional[dict] = None + self._cached_operation_status: Optional[str] = None + + def _invalidate_attribute_cache(self) -> None: + """Invalidates cached attribute dictionaries when client attributes are modified.""" + self._cached_attempt_attributes = None + self._cached_attempt_status = None + self._cached_attempt = None + self._cached_operation_attributes = None + self._cached_operation_status = None @staticmethod def _get_ms_time_diff(start: datetime, end: datetime) -> float: @@ -272,7 +356,7 @@ def client_attributes(self) -> Dict[str, str]: These attributes are used to provide context to the metrics being traced. Returns: - dict[str, str]: A dictionary of client attributes. + Dict[str, str]: A dictionary of client attributes. """ return self._client_attributes @@ -478,69 +562,58 @@ def record_afe_connectivity_error_count(self) -> None: @staticmethod def extract_front_end_latencies( metadata: Any, - ) -> tuple[Optional[int], Optional[int]]: - """ - Extracts both GFE and AFE latency values (in milliseconds) from response metadata. + ) -> Tuple[Optional[int], Optional[int]]: + """Extracts both GFE and AFE latency values (in milliseconds) from response metadata. + + :type metadata: Any + :param metadata: The metadata sequence or dict from the RPC response. + + :rtype: Tuple[Optional[int], Optional[int]] + :return: A tuple containing (gfe_latency, afe_latency) in milliseconds, or None if not found. """ if not metadata: return None, None if isinstance(metadata, dict): items = metadata.items() - elif isinstance(metadata, (list, tuple)): - items = [ - item - for item in metadata - if isinstance(item, (list, tuple)) and len(item) == 2 - ] else: - items = [] - - header_vals = [] - for key, val in items: - key_str = key.decode("utf-8") if isinstance(key, bytes) else str(key) - if key_str and key_str.lower() == "server-timing": - if isinstance(val, (list, tuple)): - header_vals.extend(val) - else: - header_vals.append(val) + try: + items = iter(metadata) + except TypeError: + return None, None gfe_latency = None afe_latency = None - for header_val in header_vals: - if not header_val: + for item in items: + if not (isinstance(item, (list, tuple)) and len(item) == 2): continue - if isinstance(header_val, bytes): - try: - header_val = header_val.decode("utf-8") - except Exception: - header_val = str(header_val) - elif not isinstance(header_val, str): - header_val = str(header_val) - - if gfe_latency is None: - match = re.search(r"gfet4t7;\s*dur=([0-9.]+)", header_val) - if match: - try: - gfe_latency = int(float(match.group(1))) - except ValueError: - pass - - if afe_latency is None: - match = re.search(r"afe;\s*dur=([0-9.]+)", header_val) - if match: - try: - afe_latency = int(float(match.group(1))) - except ValueError: - pass + key, value = item + is_server_timing = ( + isinstance(key, str) and key.lower() == _SERVER_TIMING_HEADER_STR + ) or (isinstance(key, bytes) and key.lower() == _SERVER_TIMING_HEADER_BYTES) + if not is_server_timing: + continue + + timing_values = value if isinstance(value, (list, tuple)) else (value,) + for timing_value in timing_values: + if not timing_value: + continue + text = _safe_decode_utf8(timing_value) + + if gfe_latency is None and "gfet4t7" in text: + gfe_latency = _extract_metric_latency(_GFE_TIMING_PATTERN, text) + + if afe_latency is None and "afe" in text: + afe_latency = _extract_metric_latency(_AFE_TIMING_PATTERN, text) + + if gfe_latency is not None and afe_latency is not None: + return gfe_latency, afe_latency return gfe_latency, afe_latency def record_front_end_metrics(self, metadata: Any) -> None: - """ - Extracts and records both GFE and AFE metrics from the RPC response metadata. - """ + """Extracts and records both GFE and AFE metrics from the RPC response metadata.""" if not self.enabled or not HAS_OPENTELEMETRY_INSTALLED: return gfe_latency, afe_latency = self.extract_front_end_latencies(metadata) @@ -556,35 +629,55 @@ def record_front_end_metrics(self, metadata: Any) -> None: self.record_afe_connectivity_error_count() def _create_operation_otel_attributes(self) -> dict: - """ - Create additional attributes for operation metrics tracing. + """Create additional attributes for operation metrics tracing. This method populates the client attributes dictionary with the operation status if metrics tracing is enabled. - It returns the updated client attributes dictionary. + It returns the updated client attributes dictionary (returned by reference from internal cache for performance; + should be treated as read-only by callers). """ if not self.enabled or not HAS_OPENTELEMETRY_INSTALLED: return {} + status = self.current_op.status + if ( + self._cached_operation_attributes is not None + and self._cached_operation_status == status + ): + return self._cached_operation_attributes + attributes = self._client_attributes.copy() - attributes[METRIC_LABEL_KEY_STATUS] = self.current_op.status + attributes[METRIC_LABEL_KEY_STATUS] = status + self._cached_operation_attributes = attributes + self._cached_operation_status = status return attributes def _create_attempt_otel_attributes(self) -> dict: - """ - Create additional attributes for attempt metrics tracing. + """Create additional attributes for attempt metrics tracing. This method populates the attributes dictionary with the attempt status if metrics tracing is enabled and an attempt exists. - It returns the updated attributes dictionary. + It returns the updated attributes dictionary (returned by reference from internal cache for performance; + should be treated as read-only by callers). """ if not self.enabled or not HAS_OPENTELEMETRY_INSTALLED: return {} - attributes = self._client_attributes.copy() - + current_attempt = self.current_op.current_attempt # Short circuit out if we don't have an attempt - if self.current_op.current_attempt is None: - return attributes + if current_attempt is None: + return self._client_attributes.copy() + + status = current_attempt.status + if ( + self._cached_attempt_attributes is not None + and self._cached_attempt_status == status + and self._cached_attempt is current_attempt + ): + return self._cached_attempt_attributes - attributes[METRIC_LABEL_KEY_STATUS] = self.current_op.current_attempt.status + attributes = self._client_attributes.copy() + attributes[METRIC_LABEL_KEY_STATUS] = status + self._cached_attempt_attributes = attributes + self._cached_attempt_status = status + self._cached_attempt = current_attempt return attributes def set_project(self, project: str) -> "MetricsTracer": @@ -712,7 +805,7 @@ def set_method(self, method: str) -> "MetricsTracer": :return: This instance of MetricsTracer for method chaining. """ if METRIC_LABEL_KEY_METHOD not in self._client_attributes: - self.client_attributes[METRIC_LABEL_KEY_METHOD] = method + self._client_attributes[METRIC_LABEL_KEY_METHOD] = method return self def enable_direct_path(self, enable: bool = False) -> "MetricsTracer": diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/session.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/session.py index 308c5d323624..618d1cc947a8 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/session.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/session.py @@ -70,8 +70,7 @@ class Session(object): :param database_role: (Optional) user-assigned database_role for the session. :type is_multiplexed: bool - :param is_multiplexed: (Optional) whether this session is a multiplexed session. - """ + :param is_multiplexed: (Optional) whether this session is a multiplexed session.""" def __init__(self, database, labels=None, database_role=None, is_multiplexed=False): self._database = database @@ -464,6 +463,94 @@ def transaction(self, client_context=None) -> Transaction: raise ValueError("Session has not been created.") return Transaction(self, client_context=client_context) + def _create_transaction_for_attempt( + self, + client_context=None, + transaction_tag=None, + exclude_txn_from_change_streams=None, + isolation_level=None, + read_lock_mode=None, + previous_transaction_id=None, + ) -> Transaction: + """Create and configure a transaction instance for a single attempt. + + :type client_context: :class:`~google.cloud.spanner_v1.client_context.ClientContext` + :param client_context: (Optional) client context to use for the transaction. + + :type transaction_tag: str + :param transaction_tag: (Optional) transaction tag. + + :type exclude_txn_from_change_streams: bool + :param exclude_txn_from_change_streams: (Optional) whether to exclude from change streams. + + :type isolation_level: int + :param isolation_level: (Optional) isolation level. + + :type read_lock_mode: int + :param read_lock_mode: (Optional) read lock mode. + + :type previous_transaction_id: bytes + :param previous_transaction_id: (Optional) previous transaction id for multiplexed sessions. + + :rtype: :class:`~google.cloud.spanner_v1.transaction.Transaction` + :returns: A configured Transaction instance.""" + transaction = self.transaction(client_context=client_context) + transaction.transaction_tag = transaction_tag + transaction.exclude_txn_from_change_streams = exclude_txn_from_change_streams + transaction.isolation_level = isolation_level + transaction.read_lock_mode = read_lock_mode + if self.is_multiplexed: + transaction._multiplexed_session_previous_transaction_id = ( + previous_transaction_id + ) + return transaction + + def _handle_aborted( + self, + exception, + span, + event_name, + attempts, + deadline, + default_retry_delay, + include_cause=False, + ): + """Handle an Aborted error: record trace event and delay until next retry attempt. + + :type exception: :class:`google.api_core.exceptions.Aborted` + :param exception: The aborted exception. + + :type span: :class:`opentelemetry.trace.Span` + :param span: The current active span. + + :type event_name: str + :param event_name: The span event name to record. + + :type attempts: int + :param attempts: The retry attempt number. + + :type deadline: float + :param deadline: Timestamp deadline for retrying. + + :type default_retry_delay: float + :param default_retry_delay: Default delay between retries. + + :type include_cause: bool + :param include_cause: (Optional) Whether to include the exception string as cause.""" + if span and span.is_recording(): + errors = getattr(exception, "errors", None) + cause = errors[0] if errors else exception + delay_seconds = _get_retry_delay( + cause, attempts, default_retry_delay=default_retry_delay + ) + attributes = {"attempt": attempts, "delay_seconds": delay_seconds} + if include_cause: + attributes["cause"] = str(exception) + add_span_event(span, event_name, attributes) + _delay_until_retry( + exception, deadline, attempts, default_retry_delay=default_retry_delay + ) + def run_in_transaction(self, func, *args, **kw): """Perform a unit of work in a transaction, retrying on abort. @@ -510,99 +597,78 @@ def run_in_transaction(self, func, *args, **kw): client_context = kw.pop("client_context", None) database = self._database log_commit_stats = database.log_commit_stats - extra_attributes = {} - if transaction_tag: - extra_attributes["transaction.tag"] = transaction_tag - with ( - trace_call( - "CloudSpanner.Session.run_in_transaction", - self, - extra_attributes=extra_attributes, - observability_options=getattr(database, "observability_options", None), - ) as span, - MetricsCapture(self._resource_info), - ): - attempts: int = 0 - previous_transaction_id: Optional[bytes] = None - while True: - txn = self.transaction(client_context=client_context) - txn.transaction_tag = transaction_tag - txn.exclude_txn_from_change_streams = exclude_txn_from_change_streams - txn.isolation_level = isolation_level - txn.read_lock_mode = read_lock_mode - if self.is_multiplexed: - txn._multiplexed_session_previous_transaction_id = ( - previous_transaction_id - ) - attempts += 1 - span_attributes = dict(attempt=attempts) - try: - return_value = CrossSync._Sync_Impl.run_if_async( - func, txn, *args, **kw - ) - except Aborted as exc: - previous_transaction_id = txn._transaction_id - delay_seconds = _get_retry_delay( - exc.errors[0], attempts, default_retry_delay=default_retry_delay - ) - attributes = dict(delay_seconds=delay_seconds, cause=str(exc)) - attributes.update(span_attributes) - add_span_event( - span, - "Transaction was aborted in user operation, retrying", - attributes, - ) - _delay_until_retry( - exc, deadline, attempts, default_retry_delay=default_retry_delay - ) - continue - except GoogleAPICallError: - add_span_event( - span, - "User operation failed due to GoogleAPICallError, not retrying", - span_attributes, - ) - raise - except Exception: - add_span_event( - span, - "User operation failed. Invoking Transaction.rollback(), not retrying", - span_attributes, - ) - txn.rollback() - raise - try: - txn.commit( - return_commit_stats=log_commit_stats, - request_options=commit_request_options, - max_commit_delay=max_commit_delay, - ) - except Aborted as exc: - previous_transaction_id = txn._transaction_id - delay_seconds = _get_retry_delay( - exc.errors[0], attempts, default_retry_delay=default_retry_delay - ) - attributes = dict(delay_seconds=delay_seconds) - attributes.update(span_attributes) - add_span_event( - span, - "Transaction was aborted during commit, retrying", - attributes, - ) - _delay_until_retry( - exc, deadline, attempts, default_retry_delay=default_retry_delay - ) - except GoogleAPICallError: - add_span_event( - span, - "Transaction.commit failed due to GoogleAPICallError, not retrying", - span_attributes, + span = get_current_span() + attempts: int = 0 + previous_transaction_id: Optional[bytes] = None + while True: + transaction = self._create_transaction_for_attempt( + client_context=client_context, + transaction_tag=transaction_tag, + exclude_txn_from_change_streams=exclude_txn_from_change_streams, + isolation_level=isolation_level, + read_lock_mode=read_lock_mode, + previous_transaction_id=previous_transaction_id, + ) + attempts += 1 + try: + return_value = CrossSync._Sync_Impl.run_if_async( + func, transaction, *args, **kw + ) + except Aborted as exc: + previous_transaction_id = transaction._transaction_id + self._handle_aborted( + exc, + span, + "Transaction was aborted in user operation, retrying", + attempts, + deadline, + default_retry_delay, + include_cause=True, + ) + continue + except GoogleAPICallError: + add_span_event( + span, + "User operation failed due to GoogleAPICallError, not retrying", + {"attempt": attempts}, + ) + raise + except Exception: + add_span_event( + span, + "User operation failed. Invoking Transaction.rollback(), not retrying", + {"attempt": attempts}, + ) + transaction.rollback() + raise + try: + transaction.commit( + return_commit_stats=log_commit_stats, + request_options=commit_request_options, + max_commit_delay=max_commit_delay, + ) + except Aborted as exc: + previous_transaction_id = transaction._transaction_id + self._handle_aborted( + exc, + span, + "Transaction was aborted during commit, retrying", + attempts, + deadline, + default_retry_delay, + ) + continue + except GoogleAPICallError: + add_span_event( + span, + "Transaction.commit failed due to GoogleAPICallError, not retrying", + {"attempt": attempts}, + ) + raise + else: + if log_commit_stats and transaction.commit_stats: + database.logger.info( + "CommitStats: {}".format(transaction.commit_stats), + extra={"commit_stats": transaction.commit_stats}, ) - raise - else: - if log_commit_stats and txn.commit_stats: - database.logger.info( - "CommitStats: {}".format(txn.commit_stats), - extra={"commit_stats": txn.commit_stats}, - ) - return return_value + return return_value diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot.py index 3d30e308c72a..429649bed259 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot.py @@ -35,6 +35,7 @@ AtomicCounter, _augment_error_with_request_id, _check_rst_stream_error, + _drain_stream, _make_value_pb, _merge_client_context, _merge_query_options, @@ -71,6 +72,86 @@ "Received unexpected EOS on DATA frame from server", ) +_RAW_EXECUTE_SQL_REQUEST_TYPE = ExecuteSqlRequest.pb() + + +def _make_execute_sql_request( + session_name, + sql, + seqno, + params=None, + param_types=None, + query_options=None, + request_options=None, + query_mode=None, + partition=None, + last_statement=None, + data_boost_enabled=None, + directed_read_options=None, +): + """Construct an ExecuteSqlRequest bypassing proto-plus reflection.""" + try: + raw = _RAW_EXECUTE_SQL_REQUEST_TYPE( + session=session_name, + sql=sql, + seqno=seqno, + ) + if params is not None: + if isinstance(params, dict): + if params: + raw.params.update(params) + else: + raw.params.SetInParent() + else: + raw.params.CopyFrom(getattr(params, "_pb", params)) + if param_types: + for k, v in param_types.items(): + raw.param_types[k].CopyFrom(getattr(v, "_pb", v)) + if query_options is not None: + raw.query_options.CopyFrom(getattr(query_options, "_pb", query_options)) + if request_options is not None: + raw.request_options.CopyFrom( + getattr(request_options, "_pb", request_options) + ) + if query_mode is not None: + raw.query_mode = query_mode + if partition is not None: + raw.partition_token = partition + if last_statement: + raw.last_statement = last_statement + if data_boost_enabled: + raw.data_boost_enabled = data_boost_enabled + if directed_read_options is not None: + raw.directed_read_options.CopyFrom( + getattr(directed_read_options, "_pb", directed_read_options) + ) + return ExecuteSqlRequest.wrap(raw) + except Exception: + req_kwargs = { + "session": session_name, + "sql": sql, + "seqno": seqno, + } + if params is not None: + req_kwargs["params"] = params + if param_types: + req_kwargs["param_types"] = param_types + if query_options is not None: + req_kwargs["query_options"] = query_options + if request_options is not None: + req_kwargs["request_options"] = request_options + if query_mode is not None: + req_kwargs["query_mode"] = query_mode + if partition is not None: + req_kwargs["partition_token"] = partition + if last_statement: + req_kwargs["last_statement"] = last_statement + if data_boost_enabled: + req_kwargs["data_boost_enabled"] = data_boost_enabled + if directed_read_options is not None: + req_kwargs["directed_read_options"] = directed_read_options + return ExecuteSqlRequest(req_kwargs) + def _restart_on_unavailable( method, @@ -113,75 +194,99 @@ def _restart_on_unavailable( attempt = 1 nth_request = getattr(request_id_manager, "_next_nth_request", 0) current_request_id = None - while True: - try: - if iterator is None: - with ( - trace_call( - trace_name, - session, - attributes, - observability_options=observability_options, - metadata=metadata, - ) as span, - MetricsCapture(resource_info), - ): + stream_finished = False + + try: + while True: + try: + if iterator is None: + with ( + trace_call( + trace_name, + session, + attributes, + observability_options=observability_options, + metadata=metadata, + ) as span, + MetricsCapture(resource_info), + ): + ( + call_metadata, + current_request_id, + ) = request_id_manager.metadata_and_request_id( + nth_request, attempt, metadata, span + ) + iterator = CrossSync._Sync_Impl.run_if_async( + method, request=request, metadata=call_metadata + ) + item: PartialResultSet + for item in iterator: + item_buffer.append(item) + item_pb = getattr(item, "_pb", None) or item + if transaction is not None: + transaction._update_for_result_set_pb(item) + if ( + getattr(item_pb, "HasField", lambda _: False)("precommit_token") + and transaction is not None + ): + transaction._update_for_precommit_token_pb( + item_pb.precommit_token + ) + + item_is_last = getattr(item_pb, "last", False) + + if item_is_last: + stream_finished = True + _drain_stream(iterator) + iterator = None + break + + item_resume_token = getattr(item_pb, "resume_token", b"") + if item_resume_token: + resume_token = item_resume_token + break + except ServiceUnavailable: + del item_buffer[:] + request.resume_token = resume_token + if transaction is not None: + transaction_selector = transaction._build_transaction_selector_pb() + request.transaction = transaction_selector + attempt += 1 + iterator = None + continue + except InternalServerError as exc: + resumable_error = any( ( - call_metadata, - current_request_id, - ) = request_id_manager.metadata_and_request_id( - nth_request, attempt, metadata, span + resumable_message in exc.message + for resumable_message in _STREAM_RESUMPTION_INTERNAL_ERROR_MESSAGES ) - iterator = CrossSync._Sync_Impl.run_if_async( - method, request=request, metadata=call_metadata - ) - item: PartialResultSet - for item in iterator: - item_buffer.append(item) - if transaction is not None: - transaction._update_for_result_set_pb(item) - if ( - item._pb is not None - and item._pb.HasField("precommit_token") - and (transaction is not None) - ): - transaction._update_for_precommit_token_pb(item.precommit_token) - if item.resume_token: - resume_token = item.resume_token - break - except ServiceUnavailable: - del item_buffer[:] - request.resume_token = resume_token - if transaction is not None: - transaction_selector = transaction._build_transaction_selector_pb() - request.transaction = transaction_selector - attempt += 1 - iterator = None - continue - except InternalServerError as exc: - resumable_error = any( - ( - resumable_message in exc.message - for resumable_message in _STREAM_RESUMPTION_INTERNAL_ERROR_MESSAGES ) - ) - if not resumable_error: + if not resumable_error: + raise _augment_error_with_request_id(exc, current_request_id) + del item_buffer[:] + request.resume_token = resume_token + if transaction is not None: + transaction_selector = transaction._build_transaction_selector_pb() + attempt += 1 + request.transaction = transaction_selector + iterator = None + continue + except Exception as exc: raise _augment_error_with_request_id(exc, current_request_id) + if len(item_buffer) == 0: + iterator = None + break + for item in item_buffer: + yield item del item_buffer[:] - request.resume_token = resume_token - if transaction is not None: - transaction_selector = transaction._build_transaction_selector_pb() - attempt += 1 - request.transaction = transaction_selector - iterator = None - continue - except Exception as exc: - raise _augment_error_with_request_id(exc, current_request_id) - if len(item_buffer) == 0: - break - for item in item_buffer: - yield item - del item_buffer[:] + if stream_finished: + break + finally: + if iterator is not None and hasattr(iterator, "cancel"): + try: + iterator.cancel() + except Exception: + pass class _SnapshotBase(_SessionWrapper): @@ -584,14 +689,15 @@ def execute_sql( directed_read_options = database._directed_read_options elif self.transaction_tag is not None: request_options.transaction_tag = self.transaction_tag - execute_sql_request = ExecuteSqlRequest( - session=session.name, + + execute_sql_request = _make_execute_sql_request( + session_name=session.name, sql=sql, + seqno=self._execute_sql_request_count, params=params_pb, param_types=param_types, query_mode=query_mode, - partition_token=partition, - seqno=self._execute_sql_request_count, + partition=partition, query_options=query_options, request_options=request_options, last_statement=last_statement, @@ -903,18 +1009,29 @@ def _update_for_result_set_pb( self, result_set_pb: Union[ResultSet, PartialResultSet] ) -> None: """Updates the snapshot for the given result set.""" - if result_set_pb.metadata and result_set_pb.metadata.transaction: - self._update_for_transaction_pb(result_set_pb.metadata.transaction) + rs_pb = getattr(result_set_pb, "_pb", None) or result_set_pb + metadata = getattr(rs_pb, "metadata", None) + if metadata is not None: + tx = getattr(metadata, "transaction", None) + if tx is not None and ( + getattr(tx, "id", None) + or getattr(tx, "HasField", lambda _: False)("precommit_token") + ): + self._update_for_transaction_pb(tx) def _update_for_transaction_pb(self, transaction_pb: Transaction) -> None: """Updates the snapshot for the given transaction.""" - if self._transaction_id is None and transaction_pb.id: - self._transaction_id = transaction_pb.id + tx_pb = getattr(transaction_pb, "_pb", None) or transaction_pb + tx_id = getattr(tx_pb, "id", None) + if self._transaction_id is None and tx_id: + self._transaction_id = tx_id # Notify waiting threads that the transaction has begun. self._transaction_begin_event.set() - if transaction_pb._pb.HasField("precommit_token"): - self._update_for_precommit_token_pb_unsafe(transaction_pb.precommit_token) + if tx_pb is not None and getattr(tx_pb, "HasField", lambda _: False)( + "precommit_token" + ): + self._update_for_precommit_token_pb_unsafe(tx_pb.precommit_token) def _update_for_precommit_token_pb( self, precommit_token_pb: MultiplexedSessionPrecommitToken diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot_helpers.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot_helpers.py index 5e1d6840665a..09e30ff560b7 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot_helpers.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot_helpers.py @@ -647,15 +647,26 @@ def _update_for_result_set_pb( self, result_set_pb: Union[ResultSet, PartialResultSet] ) -> None: """Updates the snapshot for the given result set.""" - if result_set_pb.metadata and result_set_pb.metadata.transaction: - self._update_for_transaction_pb(result_set_pb.metadata.transaction) + rs_pb = getattr(result_set_pb, "_pb", result_set_pb) + metadata = getattr(rs_pb, "metadata", None) + if metadata is not None: + tx = getattr(metadata, "transaction", None) + if tx is not None and ( + getattr(tx, "id", None) + or getattr(tx, "HasField", lambda _: False)("precommit_token") + ): + self._update_for_transaction_pb(tx) def _update_for_transaction_pb(self, transaction_pb: Transaction) -> None: """Updates the snapshot for the given transaction.""" - if self._transaction_id is None and transaction_pb.id: - self._transaction_id = transaction_pb.id - if transaction_pb._pb.HasField("precommit_token"): - self._update_for_precommit_token_pb_unsafe(transaction_pb.precommit_token) + tx_pb = getattr(transaction_pb, "_pb", transaction_pb) + tx_id = getattr(tx_pb, "id", None) + if self._transaction_id is None and tx_id: + self._transaction_id = tx_id + if tx_pb is not None and getattr(tx_pb, "HasField", lambda _: False)( + "precommit_token" + ): + self._update_for_precommit_token_pb_unsafe(tx_pb.precommit_token) def _update_for_precommit_token_pb( self, precommit_token_pb: MultiplexedSessionPrecommitToken diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py index a92f008f5e32..e5a77b5ef050 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py @@ -108,44 +108,101 @@ def _merge_chunk(self, value): self._pending_chunk = None return merged + def _append_to_current_row(self, values): + """Append cells to the in-progress partial row.""" + if self._lazy_decode: + self._current_row.extend(values) + else: + decoders = self._decoders + start_column = len(self._current_row) + for column_offset, value in enumerate(values): + if value.HasField("null_value"): + self._current_row.append(None) + else: + self._current_row.append( + decoders[start_column + column_offset](value) + ) + + def _decode_lazy_rows(self, values, values_offset, batch_end, width): + """Slice raw protobuf values into rows for lazy decoding.""" + if width == 1: + self._rows.extend([[value] for value in values[values_offset:batch_end]]) + else: + self._rows.extend( + [ + values[row_start : row_start + width] + for row_start in range(values_offset, batch_end, width) + ] + ) + + def _decode_eager_rows(self, values, values_offset, batch_end, width): + """Decode complete row batches into typed Python values.""" + if width == 1: + decoder = self._decoders[0] + self._rows.extend( + [ + [None if value.HasField("null_value") else decoder(value)] + for value in values[values_offset:batch_end] + ] + ) + else: + decoders = self._decoders + rows_append = self._rows.append + column_indices = list(range(width)) + for row_start in range(values_offset, batch_end, width): + rows_append( + [ + None + if values[row_start + column_index].HasField("null_value") + else decoders[column_index](values[row_start + column_index]) + for column_index in column_indices + ] + ) + def _merge_values(self, values): """Merge values into rows. :type values: list of :class:`~google.protobuf.struct_pb2.Value` :param values: non-chunked values from partial result set.""" - decoders = self._decoders + if not values: + return + width = len(self.fields) - index = len(self._current_row) - current_row = self._current_row - rows = self._rows - current_row_append = current_row.append - rows_append = rows.append + if width == 0: + return + + values_offset = 0 + total_values = len(values) + + # 1. Complete pending partial row from previous chunk (if any) + if self._current_row: + needed = width - len(self._current_row) + fill_count = min(needed, total_values) + self._append_to_current_row(values[:fill_count]) + values_offset = fill_count + if len(self._current_row) == width: + self._rows.append(self._current_row) + self._current_row = [] + else: + return + + remaining_values = total_values - values_offset + if remaining_values == 0: + return + + row_count = remaining_values // width + full_values_count = row_count * width + batch_end = values_offset + full_values_count + + # 2. Batch-decode complete rows if self._lazy_decode: - for value in values: - current_row_append(value) - index += 1 - if index == width: - rows_append(current_row) - current_row = [] - current_row_append = current_row.append - index = 0 + self._decode_lazy_rows(values, values_offset, batch_end, width) else: - for value in values: - # Note: We manually check value.HasField("null_value") here instead of - # wrapping every decoder in _parse_nullable to avoid the overhead of - # an extra Python function call layer for every cell value decoded in this loop. - # If the nullable check logic is updated in _parse_nullable, update this check. - if value.HasField("null_value"): - current_row_append(None) - else: - current_row_append(decoders[index](value)) - index += 1 - if index == width: - rows_append(current_row) - current_row = [] - current_row_append = current_row.append - index = 0 - self._current_row = current_row + self._decode_eager_rows(values, values_offset, batch_end, width) + + # 3. Buffer trailing partial row remainder for the next chunk (if any) + if remaining_values > full_values_count: + self._append_to_current_row(values[batch_end:]) def _consume_next(self): """Consume the next partial result set from the stream. diff --git a/packages/google-cloud-spanner/tests/system/test_observability_options.py b/packages/google-cloud-spanner/tests/system/test_observability_options.py index ae0324bb97b9..ae1f979dbb4a 100644 --- a/packages/google-cloud-spanner/tests/system/test_observability_options.py +++ b/packages/google-cloud-spanner/tests/system/test_observability_options.py @@ -244,11 +244,11 @@ def select_in_txn(txn): want_events = [ ("Creating Session", {}), ("Using session", {"id": session_id, "multiplexed": multiplexed}), - ("Returning session", {"id": session_id, "multiplexed": multiplexed}), ( "Transaction was aborted in user operation, retrying", {"delay_seconds": "EPHEMERAL", "cause": "EPHEMERAL", "attempt": 1}, ), + ("Returning session", {"id": session_id, "multiplexed": multiplexed}), ("Starting Commit", {}), ("Commit Done", {}), ] @@ -260,11 +260,11 @@ def select_in_txn(txn): ("No sessions available in pool. Creating session", {"kind": "BurstyPool"}), ("Creating Session", {}), ("Using session", {"id": session_id, "multiplexed": multiplexed}), - ("Returning session", {"id": session_id, "multiplexed": multiplexed}), ( "Transaction was aborted in user operation, retrying", {"delay_seconds": "EPHEMERAL", "cause": "EPHEMERAL", "attempt": 1}, ), + ("Returning session", {"id": session_id, "multiplexed": multiplexed}), ("Starting Commit", {}), ("Commit Done", {}), ] @@ -277,7 +277,6 @@ def select_in_txn(txn): want_statuses = [ ("CloudSpanner.Database.run_in_transaction", codes.OK, None), ("CloudSpanner.CreateMultiplexedSession", codes.OK, None), - ("CloudSpanner.Session.run_in_transaction", codes.OK, None), ("CloudSpanner.Transaction.execute_sql", codes.OK, None), ("CloudSpanner.Transaction.execute_sql", codes.OK, None), ("CloudSpanner.Transaction.commit", codes.OK, None), @@ -287,7 +286,6 @@ def select_in_txn(txn): want_statuses = [ ("CloudSpanner.Database.run_in_transaction", codes.OK, None), ("CloudSpanner.CreateSession", codes.OK, None), - ("CloudSpanner.Session.run_in_transaction", codes.OK, None), ("CloudSpanner.Transaction.execute_sql", codes.OK, None), ("CloudSpanner.Transaction.execute_sql", codes.OK, None), ("CloudSpanner.Transaction.commit", codes.OK, None), @@ -430,7 +428,6 @@ def tx_update(txn): want_span_names = [ "CloudSpanner.Database.run_in_transaction", expected_session_span_name, - "CloudSpanner.Session.run_in_transaction", "CloudSpanner.Transaction.commit", "CloudSpanner.Transaction.begin", ] diff --git a/packages/google-cloud-spanner/tests/system/test_session_api.py b/packages/google-cloud-spanner/tests/system/test_session_api.py index ec368af2d270..ccfca63d6d43 100644 --- a/packages/google-cloud-spanner/tests/system/test_session_api.py +++ b/packages/google-cloud-spanner/tests/system/test_session_api.py @@ -753,11 +753,6 @@ def _build_request_id(): db_name, x_goog_spanner_request_id=_build_request_id() ), }, - { - "name": "CloudSpanner.Session.run_in_transaction", - "status": ot_helpers.StatusCode.ERROR, - "attributes": _make_attributes(db_name), - }, { "name": "CloudSpanner.Database.run_in_transaction", "status": ot_helpers.StatusCode.ERROR, @@ -837,13 +832,6 @@ def _build_request_id(): ), } ) - expected_span_properties.append( - { - "name": "CloudSpanner.Session.run_in_transaction", - "status": ot_helpers.StatusCode.ERROR, - "attributes": _make_attributes(db_name), - } - ) expected_span_properties.append( { "name": "CloudSpanner.Database.run_in_transaction", @@ -1428,12 +1416,11 @@ def unit_of_work(transaction): else "CloudSpanner.CreateSession", "CloudSpanner.Batch.commit", "Test Span", - "CloudSpanner.Session.run_in_transaction", "CloudSpanner.DMLTransaction", "CloudSpanner.Transaction.commit", ] - prefix_len = 4 + prefix_len = 3 assert got_span_names[:prefix_len] == expected_span_names[:prefix_len] remaining = got_span_names[prefix_len:] assert len(remaining) >= 2 @@ -1448,9 +1435,7 @@ def unit_of_work(transaction): # |------CloudSpanner.CreateSession-------- # # |---Test Span----------------------------| - # |>--Session.run_in_transaction----------| # |---------DMLTransaction-------| - # # |>----Transaction.commit---| # CreateSession should have a trace of its own, with no children @@ -1467,18 +1452,11 @@ def assert_parent_and_children(parent_span, children): assert span.context.trace_id == parent_span.context.trace_id assert span.parent.span_id == parent_span.context.span_id - # [CreateSession --> Batch] should have their own trace. - session_run_in_txn_span = span_list[3] - children_of_test_span = [session_run_in_txn_span] + dml_txn_span = span_list[3] + batch_commit_txn_span = span_list[4] + children_of_test_span = [dml_txn_span, batch_commit_txn_span] assert_parent_and_children(test_span, children_of_test_span) - dml_txn_span = span_list[4] - batch_commit_txn_span = span_list[5] - children_of_session_run_in_txn_span = [dml_txn_span, batch_commit_txn_span] - assert_parent_and_children( - session_run_in_txn_span, children_of_session_run_in_txn_span - ) - def test_execute_partitioned_dml( not_postgres_emulator, sessions_database, database_dialect diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_database.py b/packages/google-cloud-spanner/tests/unit/_async/test_database.py index 1d5c57599693..68cb918e9f1b 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_database.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_database.py @@ -189,7 +189,7 @@ def __await__(self): manager._multiplexed_session_terminate_event.set.assert_called_once() manager._multiplexed_session_thread.cancel.assert_called_once() - mock_session.delete.assert_called_once() + mock_session.delete.assert_not_called() self.assertIsNone(manager._multiplexed_session) @CrossSync.pytest @@ -1911,6 +1911,178 @@ def _unit_of_work(txn, *args, **kwargs): self.assertEqual(committed, NOW) + @CrossSync.pytest + async def test_run_in_transaction_tracing(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + + client = _Client(observability_options=dict(tracer_provider=tracer_provider)) + instance = _Instance(self.INSTANCE_NAME, client=client) + pool = _Pool() + session = _Session() + pool.put(session) + session._committed = 42 + database = await self._make_one(self.DATABASE_ID, instance, pool=pool) + database._spanner_api = instance._client._spanner_api + + async def unit_of_work(transaction): + return 42 + + with mock.patch( + "google.cloud.spanner_v1._async.transaction.Transaction.commit", + new_callable=mock.AsyncMock, + return_value=42, + ): + await database.run_in_transaction( + unit_of_work, transaction_tag="database-tag" + ) + + finished_spans = trace_exporter.get_finished_spans() + span_names = [span.name for span in finished_spans] + self.assertIn("CloudSpanner.Database.run_in_transaction", span_names) + self.assertNotIn("CloudSpanner.Session.run_in_transaction", span_names) + database_span = next( + span + for span in finished_spans + if span.name == "CloudSpanner.Database.run_in_transaction" + ) + self.assertEqual( + database_span.attributes.get("transaction.tag"), "database-tag" + ) + + @CrossSync.pytest + async def test_run_in_transaction_tracing_retry_events(self): + from google.api_core.exceptions import Aborted + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + from google.cloud.spanner_v1._async.session import Session + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + + client = _Client(observability_options=dict(tracer_provider=tracer_provider)) + instance = _Instance(self.INSTANCE_NAME, client=client) + pool = _Pool() + database = await self._make_one(self.DATABASE_ID, instance, pool=pool) + session = Session(database, is_multiplexed=True) + session._session_id = "test-session-id" + database._sessions_manager._multiplexed_session = session + + mock_error = mock.Mock() + mock_error.trailing_metadata.return_value = () + aborted_exception = Aborted("aborted", errors=[mock_error]) + + attempts = 0 + + async def unit_of_work(transaction): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise aborted_exception + return "success" + + with mock.patch( + "google.cloud.spanner_v1._async.transaction.Transaction.commit", + new_callable=mock.AsyncMock, + side_effect=[aborted_exception, 42], + ): + result = await database.run_in_transaction( + unit_of_work, default_retry_delay=0 + ) + + self.assertEqual(result, "success") + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + database_span = finished_spans[0] + self.assertEqual(database_span.name, "CloudSpanner.Database.run_in_transaction") + + event_names = [event.name for event in database_span.events] + self.assertIn( + "Transaction was aborted in user operation, retrying", + event_names, + ) + self.assertIn( + "Transaction was aborted during commit, retrying", + event_names, + ) + + user_abort_event = next( + event + for event in database_span.events + if event.name == "Transaction was aborted in user operation, retrying" + ) + self.assertEqual(user_abort_event.attributes.get("attempt"), 1) + self.assertIn("cause", user_abort_event.attributes) + + commit_abort_event = next( + event + for event in database_span.events + if event.name == "Transaction was aborted during commit, retrying" + ) + self.assertEqual(commit_abort_event.attributes.get("attempt"), 2) + + @CrossSync.pytest + async def test_run_in_transaction_tracing_error_status(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + from opentelemetry.trace import StatusCode + + from google.cloud.spanner_v1._async.session import Session + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + + client = _Client(observability_options=dict(tracer_provider=tracer_provider)) + instance = _Instance(self.INSTANCE_NAME, client=client) + pool = _Pool() + database = await self._make_one(self.DATABASE_ID, instance, pool=pool) + session = Session(database, is_multiplexed=True) + session._session_id = "test-session-id" + database._sessions_manager._multiplexed_session = session + + async def unit_of_work(transaction): + raise ZeroDivisionError("division by zero") + + with mock.patch( + "google.cloud.spanner_v1._async.transaction.Transaction.rollback", + new_callable=mock.AsyncMock, + ): + with self.assertRaises(ZeroDivisionError): + await database.run_in_transaction(unit_of_work) + + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + database_span = finished_spans[0] + self.assertEqual(database_span.name, "CloudSpanner.Database.run_in_transaction") + self.assertEqual(database_span.status.status_code, StatusCode.ERROR) + self.assertIn("division by zero", database_span.status.description) + error_event = next( + event + for event in database_span.events + if event.name + == "User operation failed. Invoking Transaction.rollback(), not retrying" + ) + self.assertEqual(error_event.attributes.get("attempt"), 1) + @CrossSync.pytest async def test_run_in_transaction_nested(self): from datetime import datetime diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py b/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py index c49ada5ec9c6..d6a78e893edf 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_helpers_extra.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import asyncio import unittest from unittest import mock @@ -141,3 +142,94 @@ async def test_create_experimental_host_transport_errors(self): MUT._create_experimental_host_transport( InstanceAdminGrpcTransport, "host", False, None, None, None ) + + +class TestDrainStreamAsync(unittest.IsolatedAsyncioTestCase): + async def test_drain_stream_consumes_async_iterator(self): + items = [1, 2, 3] + consumed = [] + + class MockAsyncIterator: + def __aiter__(self): + return self + + async def __anext__(self): + if items: + item = items.pop(0) + consumed.append(item) + return item + raise StopAsyncIteration + + iterator = MockAsyncIterator() + MUT._drain_stream(iterator) + if MUT._PENDING_DRAIN_TASKS: + await asyncio.wait(MUT._PENDING_DRAIN_TASKS) + + self.assertEqual(consumed, [1, 2, 3]) + + async def test_drain_stream_event_loop_closed_fallback(self): + iterator = mock.Mock() + iterator.cancel = mock.Mock() + + def _raise_runtime_error(coroutine): + coroutine.close() + raise RuntimeError("no running loop") + + with mock.patch( + "google.cloud.spanner_v1._async._helpers.asyncio.create_task", + side_effect=_raise_runtime_error, + ): + MUT._drain_stream(iterator) + + iterator.cancel.assert_called_once() + + async def test_drain_stream_handles_none(self): + MUT._drain_stream(None) + + async def test_drain_stream_handles_iterator_exception(self): + class FailingAsyncIterator: + def __aiter__(self): + return self + + async def __anext__(self): + raise RuntimeError("Stream broken") + + iterator = FailingAsyncIterator() + MUT._drain_stream(iterator) + if MUT._PENDING_DRAIN_TASKS: + await asyncio.wait(MUT._PENDING_DRAIN_TASKS) + + async def test_drain_stream_task_cancellation(self): + started = asyncio.Event() + blocker = asyncio.Event() + + class BlockingAsyncIterator: + def __aiter__(self): + return self + + async def __anext__(self): + started.set() + await blocker.wait() + return 1 + + iterator = BlockingAsyncIterator() + iterator.cancel = mock.Mock() + + MUT._drain_stream(iterator) + await started.wait() + task = next(iter(MUT._PENDING_DRAIN_TASKS)) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + iterator.cancel.assert_called_once() + self.assertNotIn(task, MUT._PENDING_DRAIN_TASKS) + + def test_drain_stream_tasks_cleared_after_fork(self): + MUT._PENDING_DRAIN_TASKS.add("dummy_task") + self.assertEqual(len(MUT._PENDING_DRAIN_TASKS), 1) + + MUT._PENDING_DRAIN_TASKS.clear() + self.assertEqual(len(MUT._PENDING_DRAIN_TASKS), 0) diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_session.py b/packages/google-cloud-spanner/tests/unit/_async/test_session.py index 98758b6b904b..58d42304db09 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_session.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_session.py @@ -2833,3 +2833,117 @@ def _time_func(): _delay_until_retry(exc_mock, 6, 1) sleep_mock.assert_not_called() + + @CrossSync.pytest + async def test_run_in_transaction_tracing_events_attached_to_parent_span( + self, + ): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + tracer = tracer_provider.get_tracer("test") + + transaction_pb = TransactionPB(id=TRANSACTION_ID) + now = datetime.datetime.now(timezone.utc).replace(tzinfo=UTC) + now_pb = _datetime_to_pb_timestamp(now) + aborted = _make_rpc_error(Aborted, trailing_metadata=[]) + response = CommitResponse(commit_timestamp=now_pb) + gax_api = self._make_spanner_api() + gax_api.begin_transaction.return_value = transaction_pb + gax_api.commit.side_effect = [aborted, response] + database = self._make_database() + database.spanner_api = gax_api + session = self._make_one(database) + session._session_id = self.SESSION_ID + + async def unit_of_work(transaction, *args, **kwargs): + transaction.insert(TABLE_NAME, COLUMNS, VALUES) + return "answer" + + with tracer.start_as_current_span("ParentSpan"): + return_value = await session.run_in_transaction( + unit_of_work, + transaction_tag="test-tag", + default_retry_delay=0, + ) + + self.assertEqual(return_value, "answer") + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + parent_span = finished_spans[0] + self.assertEqual(parent_span.name, "ParentSpan") + self.assertIsNone(parent_span.attributes.get("transaction.tag")) + event_names = [event.name for event in parent_span.events] + self.assertIn( + "Transaction was aborted during commit, retrying", + event_names, + ) + retry_event = next( + event + for event in parent_span.events + if event.name == "Transaction was aborted during commit, retrying" + ) + self.assertEqual(retry_event.attributes.get("attempt"), 1) + + @CrossSync.pytest + async def test_run_in_transaction_tracing_events_aborted_without_inner_errors( + self, + ): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + tracer = tracer_provider.get_tracer("test") + + transaction_pb = TransactionPB(id=TRANSACTION_ID) + now = datetime.datetime.now(timezone.utc).replace(tzinfo=UTC) + now_pb = _datetime_to_pb_timestamp(now) + # Aborted raised without inner errors (defaults to empty tuple in GoogleAPICallError) + aborted = Aborted("aborted without inner errors") + response = CommitResponse(commit_timestamp=now_pb) + gax_api = self._make_spanner_api() + gax_api.begin_transaction.return_value = transaction_pb + gax_api.commit.side_effect = [aborted, response] + database = self._make_database() + database.spanner_api = gax_api + session = self._make_one(database) + session._session_id = self.SESSION_ID + + async def unit_of_work(transaction, *args, **kwargs): + transaction.insert(TABLE_NAME, COLUMNS, VALUES) + return "answer" + + with tracer.start_as_current_span("ParentSpan"): + return_value = await session.run_in_transaction( + unit_of_work, + default_retry_delay=0, + ) + + self.assertEqual(return_value, "answer") + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + parent_span = finished_spans[0] + event_names = [event.name for event in parent_span.events] + self.assertIn( + "Transaction was aborted during commit, retrying", + event_names, + ) + retry_event = next( + event + for event in parent_span.events + if event.name == "Transaction was aborted during commit, retrying" + ) + self.assertEqual(retry_event.attributes.get("attempt"), 1) diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_sessions_manager_extra.py b/packages/google-cloud-spanner/tests/unit/_async/test_sessions_manager_extra.py index 8a40dd942801..d0c2019803ae 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_sessions_manager_extra.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_sessions_manager_extra.py @@ -105,7 +105,8 @@ async def fake_coro(): task = asyncio.create_task(fake_coro()) manager._multiplexed_session_thread = task - manager._multiplexed_session = mock.AsyncMock() + mock_session = mock.AsyncMock() + manager._multiplexed_session = mock_session manager._multiplexed_session_terminate_event = mock.Mock() with mock.patch( @@ -116,10 +117,13 @@ async def fake_coro(): # task is cancelled and awaited in close() self.assertTrue(task.done()) manager._multiplexed_session_terminate_event.set.assert_called_once() + self.assertIsNone(manager._multiplexed_session) + mock_session.delete.assert_not_called() # Sync branch of close manager._multiplexed_session_thread = mock.Mock() - manager._multiplexed_session = mock.AsyncMock() + mock_session_sync = mock.AsyncMock() + manager._multiplexed_session = mock_session_sync manager._multiplexed_session_terminate_event = mock.Mock() with mock.patch( "google.cloud.spanner_v1._async.database_sessions_manager.CrossSync.is_async", @@ -128,6 +132,8 @@ async def fake_coro(): await manager.close() self.assertTrue(manager._multiplexed_session_thread.join.called) manager._multiplexed_session_terminate_event.set.assert_called_once() + self.assertIsNone(manager._multiplexed_session) + mock_session_sync.delete.assert_not_called() async def test_maintain_multiplexed_session_refresh(self): # coverage for line 196-202 @@ -250,25 +256,21 @@ async def test_get_multiplexed_session_fast_path_lock_already_created(self): mock_lock.acquire.assert_not_called() mock_lock.__aenter__.assert_not_called() - async def test_maintain_multiplexed_session_swaps_before_deleting_old_session( - self, - ): + async def test_maintain_multiplexed_session_rotates_session(self): from weakref import ref manager = DatabaseSessionsManager(self.database, self.pool) manager._multiplexed_session_lock = asyncio.Lock() - manager._multiplexed_session_terminate_event = asyncio.Event() + manager._multiplexed_session_terminate_event = mock.Mock() + manager._multiplexed_session_terminate_event.is_set.side_effect = [ + False, + True, + ] old_session = mock.AsyncMock() new_session = mock.AsyncMock() manager._multiplexed_session = old_session - async def verify_swap_on_delete(): - self.assertIs(manager._multiplexed_session, new_session) - manager._multiplexed_session_terminate_event.set() - - old_session.delete.side_effect = verify_swap_on_delete - refresh_interval = manager._MAINTENANCE_THREAD_REFRESH_INTERVAL.total_seconds() call_count = 0 @@ -290,7 +292,8 @@ def mock_time(): ref(manager) ) mock_build.assert_called_once() - old_session.delete.assert_called_once() + old_session.delete.assert_not_called() + new_session.delete.assert_not_called() self.assertIs(manager._multiplexed_session, new_session) async def test_maintain_multiplexed_session_handles_build_failure(self): @@ -337,48 +340,6 @@ async def mock_event_wait(event, timeout=None): current_session.delete.assert_not_called() self.assertIs(manager._multiplexed_session, current_session) - async def test_maintain_multiplexed_session_handles_delete_failure(self): - from weakref import ref - - manager = DatabaseSessionsManager(self.database, self.pool) - manager._multiplexed_session_lock = asyncio.Lock() - manager._multiplexed_session_terminate_event = asyncio.Event() - - old_session = mock.AsyncMock() - old_session.delete.side_effect = Exception("delete failed") - new_session = mock.AsyncMock() - manager._multiplexed_session = old_session - - async def verify_swap_on_delete(): - self.assertIs(manager._multiplexed_session, new_session) - manager._multiplexed_session_terminate_event.set() - raise Exception("delete failed") - - old_session.delete.side_effect = verify_swap_on_delete - - refresh_interval = manager._MAINTENANCE_THREAD_REFRESH_INTERVAL.total_seconds() - call_count = 0 - - def mock_time(): - nonlocal call_count - call_count += 1 - if call_count == 1: - return 0 - return refresh_interval + 100 - - with mock.patch( - "google.cloud.spanner_v1._async.database_sessions_manager.time.monotonic", - side_effect=mock_time, - ): - with mock.patch.object( - manager, "_build_multiplexed_session", return_value=new_session - ): - await DatabaseSessionsManager._maintain_multiplexed_session( - ref(manager) - ) - old_session.delete.assert_called_once() - self.assertIs(manager._multiplexed_session, new_session) - async def test_maintain_multiplexed_session_old_session_none(self): from weakref import ref @@ -580,7 +541,8 @@ async def test_rotate_multiplexed_session_success(self): result = await manager._rotate_multiplexed_session() self.assertTrue(result) self.assertIs(manager._multiplexed_session, new_session) - old_session.delete.assert_called_once() + old_session.delete.assert_not_called() + new_session.delete.assert_not_called() async def test_rotate_multiplexed_session_build_failure(self): manager = DatabaseSessionsManager(self.database, self.pool) @@ -597,19 +559,3 @@ async def test_rotate_multiplexed_session_build_failure(self): self.assertFalse(result) self.assertIs(manager._multiplexed_session, current_session) current_session.delete.assert_not_called() - - async def test_rotate_multiplexed_session_delete_failure(self): - manager = DatabaseSessionsManager(self.database, self.pool) - manager._multiplexed_session_lock = asyncio.Lock() - old_session = mock.AsyncMock() - old_session.delete.side_effect = Exception("delete failed") - new_session = mock.AsyncMock() - manager._multiplexed_session = old_session - - with mock.patch.object( - manager, "_build_multiplexed_session", return_value=new_session - ): - result = await manager._rotate_multiplexed_session() - self.assertTrue(result) - self.assertIs(manager._multiplexed_session, new_session) - old_session.delete.assert_called_once() diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py b/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py index bc902a6f63d1..6a4fdabe86b1 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_snapshot.py @@ -139,6 +139,334 @@ async def test_restart_on_unavailable_precommit(self): pass self.assertEqual(snapshot._precommit_token, token_pb) + async def test_restart_on_unavailable_last(self): + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1.types.result_set import PartialResultSet + + item_last = PartialResultSet(last=True) + trailing_item = PartialResultSet() + + raw = _MockIterator(item_last, trailing_item) + restart = mock.Mock(return_value=raw) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + with mock.patch( + "google.cloud.spanner_v1._async.snapshot._drain_stream" + ) as mock_drain: + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + items = [] + async for item in resumable: + items.append(item) + + mock_drain.assert_called_once_with(raw) + self.assertEqual(len(items), 1) + self.assertEqual(items[0], item_last) + + async def test_restart_on_unavailable_finally_cancels_on_early_termination(self): + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1.types.result_set import PartialResultSet + + item = PartialResultSet(last=False) + raw = _MockIterator(item) + raw.cancel = mock.Mock() + restart = mock.Mock(return_value=raw) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + async for received_item in resumable: + break + await resumable.aclose() + + raw.cancel.assert_called_once() + + async def test_restart_on_unavailable_item_without_last_attribute(self): + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + + item = mock.Mock( + spec=["resume_token", "_pb", "metadata"], + resume_token=b"", + _pb=None, + metadata=None, + ) + raw = _MockIterator(item) + restart = mock.Mock(return_value=raw) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + items = [] + async for received in resumable: + items.append(received) + + self.assertEqual(items, [item]) + + async def test_restart_on_unavailable_finally_handles_cancel_exception(self): + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1.types.result_set import PartialResultSet + + item = PartialResultSet(last=False) + raw = _MockIterator(item) + raw.cancel = mock.Mock(side_effect=RuntimeError("cancel failed")) + restart = mock.Mock(return_value=raw) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + async for _ in resumable: + break + await resumable.aclose() + raw.cancel.assert_called_once() + + async def test_restart_on_unavailable_last_does_not_cancel_iterator_in_finally( + self, + ): + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1.types.result_set import PartialResultSet + + item_last = PartialResultSet(last=True) + raw = _MockIterator(item_last) + raw.cancel = mock.Mock() + restart = mock.Mock(return_value=raw) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + with mock.patch("google.cloud.spanner_v1._async.snapshot._drain_stream"): + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + async for _ in resumable: + pass + + raw.cancel.assert_not_called() + + async def test_streamed_result_set_with_last(self): + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1._async.streamed import StreamedResultSet + from google.cloud.spanner_v1.types.result_set import ( + PartialResultSet, + ResultSetMetadata, + ) + from google.cloud.spanner_v1.types.type import StructType, Type, TypeCode + + fields = [StructType.Field(name="greeting", type_=Type(code=TypeCode.STRING))] + metadata_pb = ResultSetMetadata(row_type=StructType(fields=fields)) + item = PartialResultSet(metadata=metadata_pb, last=True) + item.values.append(Value(string_value="hello")) + + raw = _MockIterator(item) + streamed_result_set = StreamedResultSet(raw) + rows = [row async for row in streamed_result_set] + + self.assertEqual(rows, [["hello"]]) + self.assertTrue(streamed_result_set._done) + + async def test_restart_on_unavailable_multi_chunk_with_last(self): + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1._async.streamed import StreamedResultSet + from google.cloud.spanner_v1.types.result_set import ( + PartialResultSet, + ResultSetMetadata, + ResultSetStats, + ) + from google.cloud.spanner_v1.types.type import StructType, Type, TypeCode + + fields = [StructType.Field(name="greeting", type_=Type(code=TypeCode.STRING))] + metadata_pb = ResultSetMetadata(row_type=StructType(fields=fields)) + stats_pb = ResultSetStats(row_count_exact=2) + + chunk_one = PartialResultSet( + metadata=metadata_pb, last=False, resume_token=b"token_1" + ) + chunk_one.values.append(Value(string_value="hello")) + + chunk_two = PartialResultSet(last=True, stats=stats_pb) + chunk_two.values.append(Value(string_value="world")) + + raw = _MockIterator(chunk_one, chunk_two) + restart = mock.Mock(return_value=raw) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + with mock.patch( + "google.cloud.spanner_v1._async.snapshot._drain_stream" + ) as mock_drain: + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + streamed_result_set = StreamedResultSet(resumable) + rows = [row async for row in streamed_result_set] + + mock_drain.assert_called_once_with(raw) + self.assertEqual(rows, [["hello"], ["world"]]) + self.assertEqual(streamed_result_set.metadata, metadata_pb) + self.assertEqual(streamed_result_set.stats, stats_pb) + self.assertTrue(streamed_result_set._done) + + async def test_restart_on_unavailable_zero_rows_with_last(self): + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1._async.streamed import StreamedResultSet + from google.cloud.spanner_v1.types.result_set import ( + PartialResultSet, + ResultSetMetadata, + ResultSetStats, + ) + from google.cloud.spanner_v1.types.type import StructType, Type, TypeCode + + fields = [StructType.Field(name="greeting", type_=Type(code=TypeCode.STRING))] + metadata_pb = ResultSetMetadata(row_type=StructType(fields=fields)) + stats_pb = ResultSetStats(row_count_exact=0) + + chunk = PartialResultSet(metadata=metadata_pb, last=True, stats=stats_pb) + + raw = _MockIterator(chunk) + restart = mock.Mock(return_value=raw) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + with mock.patch( + "google.cloud.spanner_v1._async.snapshot._drain_stream" + ) as mock_drain: + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + streamed_result_set = StreamedResultSet(resumable) + rows = [row async for row in streamed_result_set] + + mock_drain.assert_called_once_with(raw) + self.assertEqual(rows, []) + self.assertEqual(streamed_result_set.metadata, metadata_pb) + self.assertEqual(streamed_result_set.stats, stats_pb) + self.assertTrue(streamed_result_set._done) + + async def test_restart_on_unavailable_retry_before_last(self): + from google.api_core.exceptions import ServiceUnavailable + + from google.cloud.spanner_v1._async.snapshot import _restart_on_unavailable + from google.cloud.spanner_v1.types.result_set import PartialResultSet + + resume_token = b"DEADBEEF" + chunk_one = PartialResultSet(last=False, resume_token=resume_token) + chunk_two = PartialResultSet(last=True) + + stream_one = _MockIterator( + chunk_one, fail_after=True, error=ServiceUnavailable("transient") + ) + stream_two = _MockIterator(chunk_two) + stream_two.cancel = mock.Mock() + + restart = mock.Mock(side_effect=[stream_one, stream_two]) + request = mock.Mock() + request.transaction = None + request.resume_token = b"" + session = _Session() + snapshot = self._make_snapshot(session) + + with mock.patch( + "google.cloud.spanner_v1._async.snapshot._drain_stream" + ) as mock_drain: + resumable = _restart_on_unavailable( + restart, + request, + metadata=None, + trace_name="span", + session=session, + attributes={}, + transaction=snapshot, + request_id_manager=session._database, + ) + items = [] + async for item in resumable: + items.append(item) + + self.assertEqual(items, [chunk_one, chunk_two]) + self.assertEqual(len(restart.mock_calls), 2) + self.assertEqual(request.resume_token, resume_token) + mock_drain.assert_called_once_with(stream_two) + stream_two.cancel.assert_not_called() + async def test_execute_sql_ok(self): database = _Database() fields = [StructType.Field(name="col", type_=Type(code=TypeCode.STRING))] diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py b/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py index f3ec2bb4d0cb..15114616fc26 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py @@ -1124,7 +1124,7 @@ async def test___iter___large_batch(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - expected_rows = [[i, f"name_{i}"] for i in range(500)] + expected_rows = [[index, f"name_{index}"] for index in range(500)] values = [self._make_value(cell) for row in expected_rows for cell in row] result_set = self._make_partial_result_set(values, metadata=metadata) @@ -1143,7 +1143,7 @@ async def test___iter___stepwise_consumption(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - expected_rows = [[i, f"name_{i}"] for i in range(20)] + expected_rows = [[index, f"name_{index}"] for index in range(20)] values = [self._make_value(cell) for row in expected_rows for cell in row] result_set = self._make_partial_result_set(values, metadata=metadata) @@ -1164,8 +1164,8 @@ async def test___iter___stepwise_across_chunks(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - chunk1_rows = [[i, f"name_{i}"] for i in range(10)] - chunk2_rows = [[i, f"name_{i}"] for i in range(10, 20)] + chunk1_rows = [[index, f"name_{index}"] for index in range(10)] + chunk2_rows = [[index, f"name_{index}"] for index in range(10, 20)] values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] values2 = [self._make_value(cell) for row in chunk2_rows for cell in row] @@ -1190,7 +1190,7 @@ async def test___iter___early_break(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - expected_rows = [[i, f"name_{i}"] for i in range(10)] + expected_rows = [[index, f"name_{index}"] for index in range(10)] values = [self._make_value(cell) for row in expected_rows for cell in row] result_set = self._make_partial_result_set(values, metadata=metadata) @@ -1213,7 +1213,7 @@ async def test___iter___mid_stream_error(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - chunk1_rows = [[i, f"name_{i}"] for i in range(5)] + chunk1_rows = [[index, f"name_{index}"] for index in range(5)] values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] result_set1 = self._make_partial_result_set(values1, metadata=metadata) @@ -1230,6 +1230,323 @@ async def mock_iterator(): self.assertEqual(consumed, chunk1_rows) self.assertIn("Stream error midway", str(context.exception)) + @CrossSync.pytest + async def test_decode_rows_direct_width_one(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("count", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[42]] + values = [self._make_value(42)] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual(found, expected_rows) + self.assertTrue(streamed._done) + + @CrossSync.pytest + async def test_decode_rows_direct_multi_column(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + self._make_scalar_field("active", TypeCode.BOOL), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[1, "alpha", True], [2, "beta", False]] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual(found, expected_rows) + self.assertTrue(streamed._done) + + @CrossSync.pytest + async def test_decode_rows_direct_with_null_values(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values = [ + self._make_value(1), + Value(null_value=NULL_VALUE), + self._make_value(2), + self._make_value("Alice"), + ] + expected_rows = [[1, None], [2, "Alice"]] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual(found, expected_rows) + + @CrossSync.pytest + async def test_decode_rows_direct_lazy_decode_width_one(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("id", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + raw_value = self._make_value(100) + values = [raw_value] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator, lazy_decode=True) + found = [row async for row in streamed] + self.assertEqual(len(found), 1) + self.assertEqual(found[0], [raw_value]) + self.assertEqual(streamed.decode_row(found[0]), [100]) + + @CrossSync.pytest + async def test_decode_rows_direct_lazy_decode_multi_column(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + value_id = self._make_value(1) + value_name = self._make_value("test") + values = [value_id, value_name] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator, lazy_decode=True) + found = [row async for row in streamed] + self.assertEqual(len(found), 1) + self.assertEqual(found[0], [value_id, value_name]) + self.assertEqual(streamed.decode_row(found[0]), [1, "test"]) + + @CrossSync.pytest + async def test_decode_rows_direct_trailing_partial_row(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values1 = [ + self._make_value(1), + self._make_value("a"), + self._make_value(2), + ] + values2 = [ + self._make_value("b"), + self._make_value(3), + self._make_value("c"), + ] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2, last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual( + found, + [ + [1, "a"], + [2, "b"], + [3, "c"], + ], + ) + + @CrossSync.pytest + async def test_decode_rows_direct_trailing_partial_row_lazy_decode(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + value1 = self._make_value(1) + value_a = self._make_value("a") + value2 = self._make_value(2) + value_b = self._make_value("b") + + result_set1 = self._make_partial_result_set( + [value1, value_a, value2], metadata=metadata + ) + result_set2 = self._make_partial_result_set([value_b], last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator, lazy_decode=True) + found = [row async for row in streamed] + self.assertEqual(len(found), 2) + self.assertEqual(found[0], [value1, value_a]) + self.assertEqual(found[1], [value2, value_b]) + self.assertEqual(streamed.decode_row(found[0]), [1, "a"]) + self.assertEqual(streamed.decode_row(found[1]), [2, "b"]) + + @CrossSync.pytest + async def test_decode_rows_direct_trailing_partial_row_with_null(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values1 = [ + self._make_value(1), + self._make_value("a"), + Value(null_value=NULL_VALUE), + ] + values2 = [ + self._make_value("b"), + ] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2, last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual( + found, + [ + [1, "a"], + [None, "b"], + ], + ) + + @CrossSync.pytest + async def test_decode_rows_direct_prefix_partial_row_with_null(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values1 = [ + self._make_value(1), + self._make_value("a"), + self._make_value(2), + ] + values2 = [ + Value(null_value=NULL_VALUE), + self._make_value(3), + self._make_value("c"), + ] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2, last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual( + found, + [ + [1, "a"], + [2, None], + [3, "c"], + ], + ) + + @CrossSync.pytest + async def test_decode_rows_direct_width_one_with_null(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("id", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + values = [ + self._make_value(1), + Value(null_value=NULL_VALUE), + self._make_value(2), + ] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual(found, [[1], [None], [2]]) + + @CrossSync.pytest + async def test_decode_rows_direct_three_chunk_split_row(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + self._make_scalar_field("active", TypeCode.BOOL), + ] + metadata = self._make_result_set_metadata(fields) + prs1 = self._make_partial_result_set([self._make_value(1)], metadata=metadata) + prs2 = self._make_partial_result_set([self._make_value("alpha")]) + prs3 = self._make_partial_result_set([self._make_value(True)], last=True) + + iterator = _MockCancellableIterator(prs1, prs2, prs3) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual(found, [[1, "alpha", True]]) + + @CrossSync.pytest + async def test_decode_rows_direct_three_chunk_split_row_lazy_decode(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + self._make_scalar_field("active", TypeCode.BOOL), + ] + metadata = self._make_result_set_metadata(fields) + val1 = self._make_value(1) + val2 = self._make_value("alpha") + val3 = self._make_value(True) + prs1 = self._make_partial_result_set([val1], metadata=metadata) + prs2 = self._make_partial_result_set([val2]) + prs3 = self._make_partial_result_set([val3], last=True) + + iterator = _MockCancellableIterator(prs1, prs2, prs3) + streamed = self._make_one(iterator, lazy_decode=True) + found = [row async for row in streamed] + self.assertEqual(found, [[val1, val2, val3]]) + self.assertEqual(streamed.decode_row(found[0]), [1, "alpha", True]) + + @CrossSync.pytest + async def test_decode_rows_direct_empty_values(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("id", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + + result_set1 = self._make_partial_result_set([], metadata=metadata) + result_set2 = self._make_partial_result_set([self._make_value(1)], last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual(found, [[1]]) + + @CrossSync.pytest + async def test_merge_values_zero_fields(self): + from google.cloud.spanner_v1 import ResultSetMetadata, StructType + + metadata = ResultSetMetadata(row_type=StructType(fields=[])) + result_set = self._make_partial_result_set([], metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + _ = [row async for row in streamed] + streamed._merge_values([self._make_value(1)]) + self.assertEqual(streamed._rows, []) + class _MockCancellableIterator(object): cancel_calls = 0 diff --git a/packages/google-cloud-spanner/tests/unit/test__helpers.py b/packages/google-cloud-spanner/tests/unit/test__helpers.py index 0a6e9594b167..4a720f158837 100644 --- a/packages/google-cloud-spanner/tests/unit/test__helpers.py +++ b/packages/google-cloud-spanner/tests/unit/test__helpers.py @@ -22,7 +22,42 @@ from opentelemetry.sdk.resources import Resource from opentelemetry.semconv.resource import ResourceAttributes -from google.cloud.spanner_v1 import TransactionOptions, _helpers +from google.cloud.spanner_v1 import ExecuteSqlRequest, TransactionOptions, _helpers + + +class Test_to_query_options(unittest.TestCase): + def _callFUT(self, *args, **kw): + from google.cloud.spanner_v1._helpers import _to_query_options + + return _to_query_options(*args, **kw) + + def test_none(self): + self.assertIsNone(self._callFUT(None)) + + def test_empty_dict(self): + self.assertIsNone(self._callFUT({})) + + def test_dict_with_empty_values(self): + self.assertIsNone(self._callFUT({"optimizer_version": ""})) + + def test_valid_dict(self): + expected = ExecuteSqlRequest.QueryOptions(optimizer_version="1") + result = self._callFUT({"optimizer_version": "1"}) + self.assertEqual(result, expected) + + def test_empty_proto_object(self): + self.assertIsNone(self._callFUT(ExecuteSqlRequest.QueryOptions())) + + def test_populated_proto_object(self): + options = ExecuteSqlRequest.QueryOptions(optimizer_version="1") + result = self._callFUT(options) + self.assertEqual(result, options) + + def test_invalid_type(self): + with self.assertRaises(TypeError): + self._callFUT("invalid") + with self.assertRaises(TypeError): + self._callFUT(123) class Test_merge_query_options(unittest.TestCase): @@ -37,8 +72,6 @@ def test_base_none_and_merge_none(self): self.assertIsNone(result) def test_base_dict_and_merge_none(self): - from google.cloud.spanner_v1 import ExecuteSqlRequest - base = { "optimizer_version": "2", "optimizer_statistics_package": "auto_20191128_14_47_22UTC", @@ -52,16 +85,12 @@ def test_base_dict_and_merge_none(self): self.assertEqual(result, expected) def test_base_empty_and_merge_empty(self): - from google.cloud.spanner_v1 import ExecuteSqlRequest - base = ExecuteSqlRequest.QueryOptions() merge = ExecuteSqlRequest.QueryOptions() result = self._callFUT(base, merge) self.assertIsNone(result) def test_base_none_merge_object(self): - from google.cloud.spanner_v1 import ExecuteSqlRequest - base = None merge = ExecuteSqlRequest.QueryOptions( optimizer_version="3", @@ -71,8 +100,6 @@ def test_base_none_merge_object(self): self.assertEqual(result, merge) def test_base_none_merge_dict(self): - from google.cloud.spanner_v1 import ExecuteSqlRequest - base = None merge = {"optimizer_version": "3"} expected = ExecuteSqlRequest.QueryOptions(optimizer_version="3") @@ -80,20 +107,118 @@ def test_base_none_merge_dict(self): self.assertEqual(result, expected) def test_base_object_merge_dict(self): - from google.cloud.spanner_v1 import ExecuteSqlRequest + base = ExecuteSqlRequest.QueryOptions( + optimizer_version="1", + optimizer_statistics_package="auto_20191128_14_47_22UTC", + ) + merge = {"optimizer_version": "3"} + expected = ExecuteSqlRequest.QueryOptions( + optimizer_version="3", + optimizer_statistics_package="auto_20191128_14_47_22UTC", + ) + result = self._callFUT(base, merge) + self.assertEqual(result, expected) + def test_base_object_and_merge_none(self): + base = ExecuteSqlRequest.QueryOptions( + optimizer_version="2", + optimizer_statistics_package="auto_20191128_14_47_22UTC", + ) + result = self._callFUT(base, None) + self.assertEqual(result, base) + + def test_base_empty_object_and_merge_none(self): + base = ExecuteSqlRequest.QueryOptions() + result = self._callFUT(base, None) + self.assertIsNone(result) + + def test_base_none_merge_empty_object(self): + merge = ExecuteSqlRequest.QueryOptions() + result = self._callFUT(None, merge) + self.assertIsNone(result) + + def test_base_object_not_mutated_on_merge(self): base = ExecuteSqlRequest.QueryOptions( optimizer_version="1", optimizer_statistics_package="auto_20191128_14_47_22UTC", ) merge = {"optimizer_version": "3"} + result = self._callFUT(base, merge) expected = ExecuteSqlRequest.QueryOptions( optimizer_version="3", optimizer_statistics_package="auto_20191128_14_47_22UTC", ) + self.assertEqual(result, expected) + self.assertEqual(base.optimizer_version, "1") + + def test_base_dict_merge_dict(self): + base = {"optimizer_version": "1"} + merge = {"optimizer_statistics_package": "auto_20191128_14_47_22UTC"} + expected = ExecuteSqlRequest.QueryOptions( + optimizer_version="1", + optimizer_statistics_package="auto_20191128_14_47_22UTC", + ) result = self._callFUT(base, merge) self.assertEqual(result, expected) + def test_base_dict_override_dict(self): + base = { + "optimizer_version": "1", + "optimizer_statistics_package": "pkg1", + } + merge = {"optimizer_version": "2"} + expected = ExecuteSqlRequest.QueryOptions( + optimizer_version="2", + optimizer_statistics_package="pkg1", + ) + result = self._callFUT(base, merge) + self.assertEqual(result, expected) + + def test_base_dict_empty_merge_none(self): + result = self._callFUT({}, None) + self.assertIsNone(result) + + def test_base_none_merge_dict_empty(self): + result = self._callFUT(None, {}) + self.assertIsNone(result) + + def test_base_empty_dict_merge_empty_dict(self): + result = self._callFUT({}, {}) + self.assertIsNone(result) + + def test_base_empty_dict_merge_object(self): + merge = ExecuteSqlRequest.QueryOptions(optimizer_version="1") + result = self._callFUT({}, merge) + self.assertEqual(result, merge) + + def test_base_object_merge_empty_dict(self): + base = ExecuteSqlRequest.QueryOptions(optimizer_version="1") + result = self._callFUT(base, {}) + self.assertEqual(result, base) + + def test_base_object_merge_object(self): + base = ExecuteSqlRequest.QueryOptions( + optimizer_version="1", + optimizer_statistics_package="pkg1", + ) + merge = ExecuteSqlRequest.QueryOptions(optimizer_version="2") + result = self._callFUT(base, merge) + expected = ExecuteSqlRequest.QueryOptions( + optimizer_version="2", + optimizer_statistics_package="pkg1", + ) + self.assertEqual(result, expected) + self.assertEqual(base.optimizer_version, "1") + self.assertEqual(base.optimizer_statistics_package, "pkg1") + self.assertEqual(merge.optimizer_version, "2") + self.assertEqual(merge.optimizer_statistics_package, "") + + def test_invalid_type_raises_error(self): + with self.assertRaises(TypeError): + self._callFUT("invalid", None) + with self.assertRaises(TypeError): + self._callFUT(None, 123) + class Test_get_cloud_region(unittest.TestCase): def setUp(self): @@ -152,6 +277,44 @@ def test_get_location_with_exception(self, mock_detect): self.assertIn("Failed to detect GCP resource location", log.output[0]) +class Test_try_to_coerce_bytes(unittest.TestCase): + def _callFUT(self, *args, **kw): + from google.cloud.spanner_v1._helpers import _try_to_coerce_bytes + + return _try_to_coerce_bytes(*args, **kw) + + def test_w_valid_bytes(self): + valid_bytes = b"sample_bytes" + result = self._callFUT(valid_bytes) + self.assertEqual(result, valid_bytes) + + def test_w_invalid_bytes(self): + invalid_bytes = b"\xff\xfe" + with self.assertRaises(ValueError): + self._callFUT(invalid_bytes) + + +class Test_validate_and_decode_bytes(unittest.TestCase): + def _callFUT(self, *args, **kw): + from google.cloud.spanner_v1._helpers import _validate_and_decode_bytes + + return _validate_and_decode_bytes(*args, **kw) + + def test_w_valid_bytes(self): + valid_bytes = b"sample_bytes" + result = self._callFUT(valid_bytes) + self.assertEqual(result, "sample_bytes") + self.assertIsInstance(result, str) + + def test_w_invalid_bytes(self): + invalid_bytes = b"\xff\xfe" + with self.assertRaises(ValueError) as context: + self._callFUT(invalid_bytes) + self.assertIn( + "Received a bytes that is not base64 encoded", str(context.exception) + ) + + class Test_make_value_pb(unittest.TestCase): def _callFUT(self, *args, **kw): from google.cloud.spanner_v1._helpers import _make_value_pb @@ -166,10 +329,12 @@ def test_w_bytes(self): from google.protobuf.struct_pb2 import Value BYTES = b"BYTES" - expected = Value(string_value=BYTES) + expected = Value(string_value="BYTES") value_pb = self._callFUT(BYTES) self.assertIsInstance(value_pb, Value) self.assertEqual(value_pb, expected) + self.assertIsInstance(value_pb.string_value, str) + self.assertEqual(value_pb.string_value, "BYTES") def test_w_invalid_bytes(self): BYTES = b"\xff\xfe\x03&" @@ -420,10 +585,15 @@ def test_w_proto_message(self): from .testdata import singer_pb2 singer_info = singer_pb2.SingerInfo() - expected = Value(string_value=base64.b64encode(singer_info.SerializeToString())) + expected = Value( + string_value=base64.b64encode(singer_info.SerializeToString()).decode( + "utf-8" + ) + ) value_pb = self._callFUT(singer_info) self.assertIsInstance(value_pb, Value) self.assertEqual(value_pb, expected) + self.assertIsInstance(value_pb.string_value, str) def test_w_proto_enum(self): from google.protobuf.struct_pb2 import Value @@ -434,6 +604,140 @@ def test_w_proto_enum(self): self.assertIsInstance(value_pb, Value) self.assertEqual(value_pb.string_value, "3") + def test_w_json_object(self): + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1 import JsonObject + + value = JsonObject({"key": "value"}) + value_pb = self._callFUT(value) + self.assertIsInstance(value_pb, Value) + self.assertEqual(value_pb.string_value, '{"key":"value"}') + + def test_w_uuid(self): + import uuid + + from google.protobuf.struct_pb2 import Value + + unique_id = uuid.uuid4() + value_pb = self._callFUT(unique_id) + self.assertIsInstance(value_pb, Value) + self.assertEqual(value_pb.string_value, str(unique_id)) + + def test_w_interval(self): + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1.data_types import Interval + + interval = Interval(months=1, days=2, nanos=3000) + value_pb = self._callFUT(interval) + self.assertIsInstance(value_pb, Value) + self.assertEqual(value_pb.string_value, str(interval)) + + def test_w_proto_message_none(self): + from unittest.mock import MagicMock + + from google.protobuf.message import Message + + mock_message = MagicMock(spec=Message) + mock_message.SerializeToString.return_value = None + value_pb = self._callFUT(mock_message) + self.assertTrue(value_pb.HasField("null_value")) + + def test_w_subclass_fallback(self): + import decimal + import uuid + from unittest.mock import MagicMock + + from google.api_core import datetime_helpers + from google.protobuf.struct_pb2 import ListValue + + from google.cloud.spanner_v1 import JsonObject + from google.cloud.spanner_v1.data_types import Interval + + class CustomList(list): + pass + + class CustomTuple(tuple): + pass + + class CustomInt(int): + pass + + class CustomFloat(float): + pass + + class CustomDate(datetime.date): + pass + + class CustomDatetime(datetime.datetime): + pass + + class CustomDatetimeNanos(datetime_helpers.DatetimeWithNanoseconds): + pass + + class CustomBytes(bytes): + pass + + class CustomStr(str): + pass + + class CustomDecimal(decimal.Decimal): + pass + + class CustomJsonObject(JsonObject): + pass + + class CustomInterval(Interval): + pass + + class CustomUUID(uuid.UUID): + pass + + mock_bool = MagicMock(spec=bool) + mock_list_value = MagicMock(spec=ListValue) + mock_list_value.__getitem__.side_effect = IndexError + + fallback_cases = [ + (CustomList([1]), "list_value"), + (CustomTuple((1,)), "list_value"), + (CustomInt(42), "string_value"), + (CustomFloat(3.14), "number_value"), + (CustomFloat(float("nan")), "string_value"), + (CustomFloat(float("inf")), "string_value"), + (CustomDate(2023, 5, 10), "string_value"), + ( + CustomDatetime(2023, 5, 10, 12, 0, tzinfo=datetime.timezone.utc), + "string_value", + ), + ( + CustomDatetimeNanos( + 2023, 5, 10, 12, 0, nanosecond=500, tzinfo=datetime.timezone.utc + ), + "string_value", + ), + (CustomBytes(b"custom_bytes"), "string_value"), + (CustomStr("custom_str"), "string_value"), + (CustomDecimal("99.95"), "string_value"), + (CustomJsonObject({"a": 1}), "string_value"), + (CustomJsonObject(None), "null_value"), + (CustomInterval(months=2), "string_value"), + ( + CustomUUID("12345678-1234-5678-1234-567812345678"), + "string_value", + ), + (mock_bool, "bool_value"), + (mock_list_value, "list_value"), + ] + + for value, field in fallback_cases: + with self.subTest(val_type=type(value).__name__): + result = self._callFUT(value) + self.assertTrue( + result.HasField(field), + f"Expected field {field} for {type(value)}, got {result}", + ) + class Test_make_list_value_pb(unittest.TestCase): def _callFUT(self, *args, **kw): @@ -913,7 +1217,9 @@ def test_w_proto_message(self): VALUE = singer_pb2.SingerInfo() field_type = Type(code=TypeCode.PROTO) field_name = "proto_message_column" - value_pb = Value(string_value=base64.b64encode(VALUE.SerializeToString())) + value_pb = Value( + string_value=base64.b64encode(VALUE.SerializeToString()).decode("utf-8") + ) column_info = {"proto_message_column": singer_pb2.SingerInfo()} self.assertEqual( @@ -1862,3 +2168,361 @@ def test_large_values(self): self.assertEqual(result.months, case["expected_months"]) self.assertEqual(result.days, case["expected_days"]) self.assertEqual(result.nanos, case["expected_nanos"]) + + +class TestBoundedStreamDrainer(unittest.TestCase): + def test_drain_stream_consumes_iterator(self): + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + items = [1, 2, 3] + consumed = [] + + class MockIterator: + def __iter__(self): + return self + + def __next__(self): + if items: + item = items.pop(0) + consumed.append(item) + return item + raise StopIteration + + iterator = MockIterator() + drainer = _BoundedStreamDrainer(queue_size=4, worker_count=1) + drainer.drain(iterator) + drainer._queue.join() + + self.assertEqual(consumed, [1, 2, 3]) + + def test_drain_stream_inline_fallback_on_full_queue(self): + from unittest import mock + + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=1, worker_count=0) + drainer._queue.put_nowait(mock.Mock()) + + items = [1, 2] + consumed = [] + + class MockIterator: + def __iter__(self): + return self + + def __next__(self): + if items: + item = items.pop(0) + consumed.append(item) + return item + raise StopIteration + + iterator = MockIterator() + drainer.drain(iterator) + + self.assertEqual(consumed, [1, 2]) + + def test_drain_stream_handles_none(self): + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=4, worker_count=1) + drainer.drain(None) + self.assertEqual(drainer._queue.qsize(), 0) + + def test_drain_stream_handles_iterator_exception(self): + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + class FailingIterator: + def __iter__(self): + return self + + def __next__(self): + raise RuntimeError("Stream broken") + + iterator = FailingIterator() + drainer = _BoundedStreamDrainer(queue_size=4, worker_count=1) + drainer.drain(iterator) + drainer._queue.join() + + def test_drain_stream_after_shutdown_drains_inline(self): + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=4, worker_count=1) + drainer.shutdown() + self.assertTrue(drainer._stopped) + + items = [1, 2, 3] + consumed = [] + + class MockIterator: + def __iter__(self): + return self + + def __next__(self): + if items: + item = items.pop(0) + consumed.append(item) + return item + raise StopIteration + + iterator = MockIterator() + drainer.drain(iterator) + self.assertEqual(consumed, [1, 2, 3]) + self.assertEqual(drainer._queue.qsize(), 0) + + def test_drain_stream_inline_fallback_iterator_exception(self): + from unittest import mock + + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=1, worker_count=0) + drainer._queue.put_nowait(mock.Mock()) + + class FailingIterator: + def __iter__(self): + return self + + def __next__(self): + raise RuntimeError("Failed inline") + + iterator = FailingIterator() + # Should catch and ignore the exception without raising + drainer.drain(iterator) + + def test_drainer_fallback_on_ensure_started_error(self): + from unittest import mock + + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=2, worker_count=1) + items = [1, 2, 3] + consumed = [] + + class MockIterator: + def __iter__(self): + return self + + def __next__(self): + if items: + val = items.pop(0) + consumed.append(val) + return val + raise StopIteration + + with mock.patch.object( + drainer, "_ensure_started", side_effect=RuntimeError("thread limit reached") + ): + drainer.drain(MockIterator()) + + self.assertEqual(consumed, [1, 2, 3]) + + def test_drainer_ensure_started_partial_failure_retains_started(self): + from unittest import mock + + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=2, worker_count=3) + start_count = 0 + + def _mock_thread_start(thread_self): + nonlocal start_count + start_count += 1 + if start_count > 1: + raise RuntimeError("thread limit reached") + + with mock.patch( + "google.cloud.spanner_v1._helpers.threading.Thread.start", + _mock_thread_start, + ): + with self.assertRaises(RuntimeError): + drainer._ensure_started() + + # Started should remain True because 1 worker was successfully created + self.assertTrue(drainer._started) + self.assertEqual(len(drainer._workers), 1) + + # Subsequent call must not attempt to spawn additional threads + drainer._ensure_started() + self.assertEqual(len(drainer._workers), 1) + + def test_drainer_ensure_started_total_failure_resets_started(self): + from unittest import mock + + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=2, worker_count=2) + + with mock.patch( + "google.cloud.spanner_v1._helpers.threading.Thread.start", + side_effect=RuntimeError("no threads"), + ): + with self.assertRaises(RuntimeError): + drainer._ensure_started() + + # Started should be reset to False because 0 workers were created + self.assertFalse(drainer._started) + self.assertEqual(len(drainer._workers), 0) + + def test_drainer_shutdown_with_full_queue(self): + from unittest import mock + + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=2, worker_count=2) + drainer._ensure_started() + drainer._queue.put_nowait(mock.Mock()) + drainer._queue.put_nowait(mock.Mock()) + self.assertTrue(drainer._queue.full()) + + # Calling shutdown when queue is already full must not raise queue.Full + drainer.shutdown() + self.assertTrue(drainer._stopped) + + def test_drainer_reset_after_fork(self): + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=4, worker_count=2) + drainer._ensure_started() + self.assertTrue(drainer._started) + self.assertEqual(len(drainer._workers), 2) + + drainer._reset_after_fork() + self.assertFalse(drainer._started) + self.assertFalse(drainer._stopped) + self.assertEqual(len(drainer._workers), 0) + self.assertEqual(drainer._queue.qsize(), 0) + + def test_drainer_garbage_collection(self): + import gc + import weakref + + from google.cloud.spanner_v1._helpers import _BoundedStreamDrainer + + drainer = _BoundedStreamDrainer(queue_size=4, worker_count=1) + ref = weakref.ref(drainer) + del drainer + gc.collect() + + self.assertIsNone(ref()) + + def test_global_stream_drainer_reset_after_fork(self): + from google.cloud.spanner_v1 import _helpers + + _helpers._GLOBAL_STREAM_DRAINER._ensure_started() + self.assertTrue(_helpers._GLOBAL_STREAM_DRAINER._started) + + _helpers._GLOBAL_STREAM_DRAINER._reset_after_fork() + self.assertFalse(_helpers._GLOBAL_STREAM_DRAINER._started) + self.assertEqual(len(_helpers._GLOBAL_STREAM_DRAINER._workers), 0) + self.assertEqual(_helpers._GLOBAL_STREAM_DRAINER._queue.qsize(), 0) + + def test_module_drain_stream(self): + from unittest import mock + + from google.cloud.spanner_v1 import _helpers + + with mock.patch.object(_helpers._GLOBAL_STREAM_DRAINER, "drain") as mock_drain: + iterator = mock.Mock() + _helpers._drain_stream(iterator) + mock_drain.assert_called_once_with(iterator) + + +class Test_get_type_decoder(unittest.TestCase): + def _callFUT(self, *args, **kwargs): + from google.cloud.spanner_v1._helpers import _get_type_decoder + + return _get_type_decoder(*args, **kwargs) + + def test_scalar_decoders(self): + import datetime + import decimal + import uuid + + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1 import Type, TypeCode + from google.cloud.spanner_v1._helpers import _SCALAR_DECODERS + from google.cloud.spanner_v1.data_types import Interval, JsonObject + + test_cases = [ + (TypeCode.STRING, Value(string_value="hello"), "hello"), + (TypeCode.BYTES, Value(string_value="bytes"), b"bytes"), + (TypeCode.BOOL, Value(bool_value=True), True), + (TypeCode.INT64, Value(string_value="42"), 42), + (TypeCode.FLOAT64, Value(string_value="3.14"), 3.14), + (TypeCode.FLOAT32, Value(string_value="2.5"), 2.5), + ( + TypeCode.DATE, + Value(string_value="2026-03-15"), + datetime.date(2026, 3, 15), + ), + ( + TypeCode.TIMESTAMP, + Value(string_value="2026-03-15T12:00:00Z"), + datetime.datetime(2026, 3, 15, 12, 0, tzinfo=datetime.timezone.utc), + ), + (TypeCode.NUMERIC, Value(string_value="99.99"), decimal.Decimal("99.99")), + (TypeCode.JSON, Value(string_value='{"a": 1}'), JsonObject({"a": 1})), + ( + TypeCode.UUID, + Value(string_value="12345678-1234-5678-1234-567812345678"), + uuid.UUID("12345678-1234-5678-1234-567812345678"), + ), + (TypeCode.INTERVAL, Value(string_value="P1Y"), Interval.from_str("P1Y")), + ] + for type_code, sample_value_pb, expected_result in test_cases: + field_type = Type(code=type_code) + decoder = self._callFUT(field_type, "column_name") + self.assertIs(decoder, _SCALAR_DECODERS[int(type_code)]) + self.assertEqual(decoder(sample_value_pb), expected_result) + + def test_proto_and_enum(self): + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1 import Type, TypeCode + + proto_type = Type(code=TypeCode.PROTO) + proto_decoder = self._callFUT(proto_type, "proto_column") + self.assertTrue(callable(proto_decoder)) + + enum_type = Type(code=TypeCode.ENUM) + enum_decoder = self._callFUT(enum_type, "enum_column") + self.assertTrue(callable(enum_decoder)) + self.assertEqual(enum_decoder(Value(string_value="1")), 1) + + def test_array_and_struct(self): + from google.cloud.spanner_v1 import StructType, Type, TypeCode + + array_type = Type( + code=TypeCode.ARRAY, + array_element_type=Type(code=TypeCode.STRING), + ) + array_decoder = self._callFUT(array_type, "array_column") + self.assertTrue(callable(array_decoder)) + + struct_field = StructType.Field( + name="subfield", type_=Type(code=TypeCode.STRING) + ) + struct_type = Type( + code=TypeCode.STRUCT, + struct_type=StructType(fields=[struct_field]), + ) + struct_decoder = self._callFUT(struct_type, "struct_column") + self.assertTrue(callable(struct_decoder)) + + def test_unknown_and_unspecified_types(self): + from unittest import mock + + from google.cloud.spanner_v1 import Type, TypeCode + + unspecified_type = Type(code=TypeCode.TYPE_CODE_UNSPECIFIED) + with self.assertRaises(ValueError): + self._callFUT(unspecified_type, "unspecified") + + unknown_type = mock.Mock(code=999) + with self.assertRaises(ValueError): + self._callFUT(unknown_type, "unknown") + + invalid_code_type = mock.Mock(code="invalid") + with self.assertRaises(ValueError): + self._callFUT(invalid_code_type, "invalid") diff --git a/packages/google-cloud-spanner/tests/unit/test__opentelemetry_tracing.py b/packages/google-cloud-spanner/tests/unit/test__opentelemetry_tracing.py index 1aee8688a58e..8d04a85720bc 100644 --- a/packages/google-cloud-spanner/tests/unit/test__opentelemetry_tracing.py +++ b/packages/google-cloud-spanner/tests/unit/test__opentelemetry_tracing.py @@ -224,3 +224,455 @@ def test_trace_call_terminal_span_status_ALWAYS_OFF_sampler(self): assert type(used_span).__name__ == "NonRecordingSpan" span_list = list(trace_exporter.get_finished_spans()) assert span_list == [] + + def test_is_tracer_noop(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.trace import NoOpTracer, ProxyTracer + + self.assertTrue(_opentelemetry_tracing._is_tracer_noop(None)) + self.assertTrue(_opentelemetry_tracing._is_tracer_noop(NoOpTracer())) + + # ProxyTracer with uninitialized real tracer delegates to NoOpTracer + with mock.patch("opentelemetry.trace._TRACER_PROVIDER", None): + proxy_tracer_noop = ProxyTracer("test_mod") + self.assertTrue(_opentelemetry_tracing._is_tracer_noop(proxy_tracer_noop)) + + # Real SDK tracer is not no-op + sdk_provider = TracerProvider() + sdk_tracer = sdk_provider.get_tracer("test_mod") + self.assertFalse(_opentelemetry_tracing._is_tracer_noop(sdk_tracer)) + + # ProxyTracer with real tracer attached is not no-op + proxy_tracer_active = ProxyTracer("test_mod") + proxy_tracer_active._real_tracer = sdk_tracer + self.assertFalse(_opentelemetry_tracing._is_tracer_noop(proxy_tracer_active)) + + @mock.patch( + "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", + return_value="global", + ) + def test_trace_call_noop_tracer_provider_short_circuits(self, mock_region): + session = _make_session() + session._last_use_time = None + observability_options = dict(tracer_provider=trace_api.NoOpTracerProvider()) + + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestNoOp", + session, + extra_attributes={"db.statement": "SELECT 1"}, + observability_options=observability_options, + ) as span: + self.assertFalse(span.is_recording()) + self.assertEqual(span, trace_api.INVALID_SPAN) + + mock_region.assert_not_called() + self.assertIsNotNone(session._last_use_time) + + @mock.patch( + "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", + return_value="global", + ) + def test_trace_call_noop_with_ambient_span(self, mock_region): + from opentelemetry.context import attach, detach + + ambient_context = trace_api.SpanContext( + trace_id=0x12345678123456781234567812345678, + span_id=0x1234567812345678, + is_remote=False, + ) + ambient_span = trace_api.NonRecordingSpan(ambient_context) + token = attach(trace_api.set_span_in_context(ambient_span)) + + try: + observability_options = dict(tracer_provider=trace_api.NoOpTracerProvider()) + session = _make_session() + + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestAmbient", + session, + observability_options=observability_options, + ) as span: + self.assertFalse(span.is_recording()) + # Active span in context is the child span, masking the ambient span + self.assertEqual(_opentelemetry_tracing.get_current_span(), span) + self.assertNotEqual(span, trace_api.INVALID_SPAN) + self.assertEqual( + span.get_span_context().trace_id, ambient_context.trace_id + ) + + mock_region.assert_not_called() + finally: + detach(token) + + def test_trace_call_noop_with_invalid_ambient_span(self): + from opentelemetry.context import attach, detach + + invalid_ambient = trace_api.NonRecordingSpan(trace_api.INVALID_SPAN_CONTEXT) + token = attach(trace_api.set_span_in_context(invalid_ambient)) + + try: + observability_options = dict(tracer_provider=trace_api.NoOpTracerProvider()) + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestInvalidAmbient", + _make_session(), + observability_options=observability_options, + ) as span: + self.assertFalse(span.is_recording()) + self.assertEqual(span, trace_api.INVALID_SPAN) + finally: + detach(token) + + def test_trace_call_noop_does_not_pollute_ambient_span(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + ambient_tracer = provider.get_tracer("ambient_app") + + with ambient_tracer.start_as_current_span("ambient_request"): + observability_options = dict(tracer_provider=trace_api.NoOpTracerProvider()) + with _opentelemetry_tracing.trace_call( + "CloudSpanner.IsolatedCall", + session=None, + observability_options=observability_options, + ): + current = _opentelemetry_tracing.get_current_span() + _opentelemetry_tracing.add_span_event(current, "internal_event") + current.set_attribute("internal_attr", "internal_val") + + finished_spans = exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + parent_span = finished_spans[0] + self.assertNotIn("internal_attr", parent_span.attributes) + self.assertEqual(len(parent_span.events), 0) + + def test_trace_call_noop_provider_end_to_end_tracing(self): + observability_options = dict( + tracer_provider=trace_api.NoOpTracerProvider(), + enable_end_to_end_tracing=True, + ) + metadata = [] + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestE2E", + _make_session(), + observability_options=observability_options, + metadata=metadata, + ) as span: + self.assertFalse(span.is_recording()) + + self.assertTrue( + any(item[0] == "x-goog-spanner-end-to-end-tracing" for item in metadata) + ) + + def test_trace_call_noop_fast_path_exception(self): + observability_options = dict(tracer_provider=trace_api.NoOpTracerProvider()) + + with self.assertRaises(ValueError): + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestError", + _make_session(), + observability_options=observability_options, + ): + raise ValueError("Test error in fast path") + + def test_trace_call_dict_subclass_observability_options(self): + from collections import UserDict + + observability_options = UserDict( + {"tracer_provider": trace_api.NoOpTracerProvider()} + ) + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestUserDict", + _make_session(), + observability_options=observability_options, + ) as span: + self.assertEqual(span, trace_api.INVALID_SPAN) + + def test_trace_call_head_sampler_receives_attributes(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.sampling import ( + Decision, + Sampler, + SamplingResult, + ) + + class AttributeCapturingSampler(Sampler): + def __init__(self): + self.captured_attributes = None + + def should_sample( + self, + parent_context, + trace_id, + name, + kind=None, + attributes=None, + links=None, + trace_state=None, + ): + self.captured_attributes = dict(attributes) if attributes else {} + return SamplingResult(Decision.RECORD_AND_SAMPLE) + + def get_description(self): + return "AttributeCapturingSampler" + + sampler = AttributeCapturingSampler() + tracer_provider = TracerProvider(sampler=sampler) + observability_options = dict( + tracer_provider=tracer_provider, + db_name="projects/p/instances/i/databases/custom_db", + ) + + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestSampler", + _make_session(), + extra_attributes={"custom.attr": "value1"}, + observability_options=observability_options, + ): + pass + + self.assertIsNotNone(sampler.captured_attributes) + self.assertEqual(sampler.captured_attributes.get("db.type"), "spanner") + self.assertEqual( + sampler.captured_attributes.get("db.instance"), + "projects/p/instances/i/databases/custom_db", + ) + self.assertEqual(sampler.captured_attributes.get("custom.attr"), "value1") + + def test_trace_call_span_processor_on_start_receives_attributes(self): + from opentelemetry.sdk.trace import SpanProcessor, TracerProvider + + class AttributeCapturingProcessor(SpanProcessor): + def __init__(self): + self.started_attributes = None + + def on_start(self, span, parent_context=None): + self.started_attributes = ( + dict(span.attributes) if span.attributes else {} + ) + + def on_end(self, span): + pass + + processor = AttributeCapturingProcessor() + tracer_provider = TracerProvider() + tracer_provider.add_span_processor(processor) + observability_options = dict(tracer_provider=tracer_provider) + + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestProcessor", + _make_session(), + extra_attributes={"on_start.attr": "value2"}, + observability_options=observability_options, + ): + pass + + self.assertIsNotNone(processor.started_attributes) + self.assertEqual(processor.started_attributes.get("db.type"), "spanner") + self.assertEqual(processor.started_attributes.get("on_start.attr"), "value2") + + def test_trace_call_request_options_and_disabled_extended_tracing(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + + class DummyRequestOptions: + request_tag = "test-tag" + + tracer_provider = TracerProvider() + exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(exporter)) + + observability_options = dict( + tracer_provider=tracer_provider, + enable_extended_tracing=False, + enable_end_to_end_tracing=True, + ) + metadata = [] + + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestOptions", + _make_session(), + extra_attributes={ + "db.statement": "SELECT 1", + "request_options": DummyRequestOptions(), + }, + observability_options=observability_options, + metadata=metadata, + ): + pass + + spans = exporter.get_finished_spans() + self.assertEqual(len(spans), 1) + self.assertEqual(spans[0].attributes.get("request.tag"), "test-tag") + self.assertNotIn("db.statement", spans[0].attributes) + self.assertTrue( + any(item[0] == "x-goog-spanner-end-to-end-tracing" for item in metadata) + ) + + def test_trace_call_noop_with_ambient_span_and_e2e_tracing(self): + from opentelemetry.context import attach, detach + + ambient_context = trace_api.SpanContext( + trace_id=0x12345678123456781234567812345678, + span_id=0x1234567812345678, + is_remote=False, + ) + ambient_span = trace_api.NonRecordingSpan(ambient_context) + token = attach(trace_api.set_span_in_context(ambient_span)) + + try: + metadata = [] + observability_options = dict( + tracer_provider=trace_api.NoOpTracerProvider(), + enable_end_to_end_tracing=True, + ) + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestAmbientE2E", + _make_session(), + observability_options=observability_options, + metadata=metadata, + ): + pass + self.assertTrue( + any(item[0] == "x-goog-spanner-end-to-end-tracing" for item in metadata) + ) + finally: + detach(token) + + @mock.patch.object( + _opentelemetry_tracing, "end_to_end_tracing_globally_enabled", True + ) + def test_trace_call_global_e2e_tracing_enabled(self): + observability_options = dict(tracer_provider=trace_api.NoOpTracerProvider()) + metadata = [] + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestGlobalE2E", + _make_session(), + observability_options=observability_options, + metadata=metadata, + ): + pass + self.assertTrue( + any(item[0] == "x-goog-spanner-end-to-end-tracing" for item in metadata) + ) + + @mock.patch.object( + _opentelemetry_tracing, "extended_tracing_globally_disabled", True + ) + def test_trace_call_global_extended_tracing_disabled(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + + tracer_provider = TracerProvider() + exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(exporter)) + observability_options = dict(tracer_provider=tracer_provider) + + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestGlobalExtendedDisabled", + _make_session(), + extra_attributes={"db.statement": "SELECT 1"}, + observability_options=observability_options, + ): + pass + + spans = exporter.get_finished_spans() + self.assertEqual(len(spans), 1) + self.assertNotIn("db.statement", spans[0].attributes) + + def test_trace_call_default_production_unconfigured_without_session(self): + with mock.patch("opentelemetry.trace._TRACER_PROVIDER", None): + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestDefaultNoSession", session=None + ) as span: + self.assertEqual(span, trace_api.INVALID_SPAN) + + def test_trace_call_active_without_session_database(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + + tracer_provider = TracerProvider() + exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(exporter)) + observability_options = dict(tracer_provider=tracer_provider) + + # Session is None + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestNoSessionActive", + session=None, + observability_options=observability_options, + ): + pass + + spans = exporter.get_finished_spans() + self.assertEqual(len(spans), 1) + self.assertEqual(spans[0].attributes.get("db.instance"), "") + + def test_trace_call_active_request_options_without_tag(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + + class UntaggedRequestOptions: + request_tag = None + + tracer_provider = TracerProvider() + exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(exporter)) + observability_options = dict(tracer_provider=tracer_provider) + + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestUntaggedOptions", + _make_session(), + extra_attributes={"request_options": UntaggedRequestOptions()}, + observability_options=observability_options, + ): + pass + + spans = exporter.get_finished_spans() + self.assertEqual(len(spans), 1) + self.assertNotIn("request.tag", spans[0].attributes) + + def test_trace_call_noop_with_ambient_span_exception(self): + from opentelemetry.context import attach, detach + + ambient_context = trace_api.SpanContext( + trace_id=0x12345678123456781234567812345678, + span_id=0x1234567812345678, + is_remote=False, + ) + ambient_span = trace_api.NonRecordingSpan(ambient_context) + token = attach(trace_api.set_span_in_context(ambient_span)) + + try: + observability_options = dict(tracer_provider=trace_api.NoOpTracerProvider()) + with self.assertRaises(RuntimeError): + with _opentelemetry_tracing.trace_call( + "CloudSpanner.TestAmbientError", + _make_session(), + observability_options=observability_options, + ): + raise RuntimeError("Error with ambient span") + finally: + detach(token) + + def test_add_span_event(self): + span = mock.Mock() + _opentelemetry_tracing.add_span_event(span, "test_event", {"attr": "val"}) + span.add_event.assert_called_once_with("test_event", {"attr": "val"}) diff --git a/packages/google-cloud-spanner/tests/unit/test_database.py b/packages/google-cloud-spanner/tests/unit/test_database.py index 21c382b45438..b63e19a61baf 100644 --- a/packages/google-cloud-spanner/tests/unit/test_database.py +++ b/packages/google-cloud-spanner/tests/unit/test_database.py @@ -1760,6 +1760,166 @@ def _unit_of_work(txn, *args, **kwargs): self.assertEqual(committed, NOW) + def test_run_in_transaction_tracing(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + + client = _Client(observability_options=dict(tracer_provider=tracer_provider)) + instance = _Instance(self.INSTANCE_NAME, client=client) + pool = _Pool() + session = _Session() + pool.put(session) + session._committed = 42 + database = self._make_one(self.DATABASE_ID, instance, pool=pool) + database._spanner_api = instance._client._spanner_api + + def unit_of_work(transaction): + return 42 + + with mock.patch( + "google.cloud.spanner_v1.transaction.Transaction.commit", + return_value=42, + ): + database.run_in_transaction(unit_of_work, transaction_tag="database-tag") + + finished_spans = trace_exporter.get_finished_spans() + span_names = [span.name for span in finished_spans] + self.assertIn("CloudSpanner.Database.run_in_transaction", span_names) + self.assertNotIn("CloudSpanner.Session.run_in_transaction", span_names) + database_span = next( + span + for span in finished_spans + if span.name == "CloudSpanner.Database.run_in_transaction" + ) + self.assertEqual( + database_span.attributes.get("transaction.tag"), "database-tag" + ) + + def test_run_in_transaction_tracing_retry_events(self): + from google.api_core.exceptions import Aborted + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + from google.cloud.spanner_v1.session import Session + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + + client = _Client(observability_options=dict(tracer_provider=tracer_provider)) + instance = _Instance(self.INSTANCE_NAME, client=client) + pool = _Pool() + database = self._make_one(self.DATABASE_ID, instance, pool=pool) + session = Session(database, is_multiplexed=True) + session._session_id = "test-session-id" + database._sessions_manager._multiplexed_session = session + + mock_error = mock.Mock() + mock_error.trailing_metadata.return_value = () + aborted_exception = Aborted("aborted", errors=[mock_error]) + + attempts = 0 + + def unit_of_work(transaction): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise aborted_exception + return "success" + + with mock.patch( + "google.cloud.spanner_v1.transaction.Transaction.commit", + side_effect=[aborted_exception, 42], + ): + result = database.run_in_transaction(unit_of_work, default_retry_delay=0) + + self.assertEqual(result, "success") + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + database_span = finished_spans[0] + self.assertEqual(database_span.name, "CloudSpanner.Database.run_in_transaction") + + event_names = [event.name for event in database_span.events] + self.assertIn( + "Transaction was aborted in user operation, retrying", + event_names, + ) + self.assertIn( + "Transaction was aborted during commit, retrying", + event_names, + ) + + user_abort_event = next( + event + for event in database_span.events + if event.name == "Transaction was aborted in user operation, retrying" + ) + self.assertEqual(user_abort_event.attributes.get("attempt"), 1) + self.assertIn("cause", user_abort_event.attributes) + + commit_abort_event = next( + event + for event in database_span.events + if event.name == "Transaction was aborted during commit, retrying" + ) + self.assertEqual(commit_abort_event.attributes.get("attempt"), 2) + + def test_run_in_transaction_tracing_error_status(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + from opentelemetry.trace import StatusCode + + from google.cloud.spanner_v1.session import Session + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + + client = _Client(observability_options=dict(tracer_provider=tracer_provider)) + instance = _Instance(self.INSTANCE_NAME, client=client) + pool = _Pool() + database = self._make_one(self.DATABASE_ID, instance, pool=pool) + session = Session(database, is_multiplexed=True) + session._session_id = "test-session-id" + database._sessions_manager._multiplexed_session = session + + def unit_of_work(transaction): + raise ZeroDivisionError("division by zero") + + with mock.patch("google.cloud.spanner_v1.transaction.Transaction.rollback"): + with self.assertRaises(ZeroDivisionError): + database.run_in_transaction(unit_of_work) + + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + database_span = finished_spans[0] + self.assertEqual(database_span.name, "CloudSpanner.Database.run_in_transaction") + self.assertEqual(database_span.status.status_code, StatusCode.ERROR) + self.assertIn("division by zero", database_span.status.description) + error_event = next( + event + for event in database_span.events + if event.name + == "User operation failed. Invoking Transaction.rollback(), not retrying" + ) + self.assertEqual(error_event.attributes.get("attempt"), 1) + def test_run_in_transaction_nested(self): from datetime import datetime diff --git a/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py b/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py index 297ac3c3817c..47fa96e5a468 100644 --- a/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py +++ b/packages/google-cloud-spanner/tests/unit/test_database_session_manager.py @@ -274,7 +274,7 @@ def test_get_multiplexed_session_fast_path_lock_already_created(self): mock_lock.acquire.assert_not_called() mock_lock.__enter__.assert_not_called() - def test_maintain_multiplexed_session_swaps_before_deleting_old_session(self): + def test_maintain_multiplexed_session_rotates_session(self): import threading from weakref import ref @@ -287,11 +287,6 @@ def test_maintain_multiplexed_session_swaps_before_deleting_old_session(self): new_session = Mock() manager._multiplexed_session = old_session - def verify_swap_on_delete(): - self.assertIs(manager._multiplexed_session, new_session) - - old_session.delete.side_effect = verify_swap_on_delete - call_count = 0 def mock_time(): @@ -310,7 +305,8 @@ def mock_time(): ) as mock_build: DatabaseSessionsManager._maintain_multiplexed_session(ref(manager)) mock_build.assert_called_once() - old_session.delete.assert_called_once() + old_session.delete.assert_not_called() + new_session.delete.assert_not_called() self.assertIs(manager._multiplexed_session, new_session) def test_close_branches(self): @@ -330,7 +326,7 @@ def test_close_branches(self): manager._multiplexed_session_terminate_event.set.assert_called_once() mock_thread.join.assert_called_once() self.assertIsNone(manager._multiplexed_session) - mock_session.delete.assert_called_once() + mock_session.delete.assert_not_called() def test_maintain_multiplexed_session_handles_build_failure(self): import threading @@ -374,40 +370,6 @@ def mock_time(): current_session.delete.assert_not_called() self.assertIs(manager._multiplexed_session, current_session) - def test_maintain_multiplexed_session_handles_delete_failure(self): - import threading - from weakref import ref - - manager = DatabaseSessionsManager(self._manager._database, self._manager._pool) - manager._multiplexed_session_lock = threading.Lock() - manager._multiplexed_session_terminate_event = Mock() - manager._multiplexed_session_terminate_event.is_set.side_effect = [False, True] - - old_session = Mock() - old_session.delete.side_effect = Exception("delete failed") - new_session = Mock() - manager._multiplexed_session = old_session - - call_count = 0 - - def mock_time(): - nonlocal call_count - call_count += 1 - if call_count == 1: - return 0 - return 1000000 - - with patch( - "google.cloud.spanner_v1.database_sessions_manager.time.monotonic", - side_effect=mock_time, - ): - with patch.object( - manager, "_build_multiplexed_session", return_value=new_session - ): - DatabaseSessionsManager._maintain_multiplexed_session(ref(manager)) - old_session.delete.assert_called_once() - self.assertIs(manager._multiplexed_session, new_session) - def test_maintain_multiplexed_session_old_session_none(self): import threading from weakref import ref @@ -662,7 +624,8 @@ def test_rotate_multiplexed_session_success(self): result = manager._rotate_multiplexed_session() self.assertTrue(result) self.assertIs(manager._multiplexed_session, new_session) - old_session.delete.assert_called_once() + old_session.delete.assert_not_called() + new_session.delete.assert_not_called() def test_rotate_multiplexed_session_build_failure(self): import threading @@ -682,24 +645,6 @@ def test_rotate_multiplexed_session_build_failure(self): self.assertIs(manager._multiplexed_session, current_session) current_session.delete.assert_not_called() - def test_rotate_multiplexed_session_delete_failure(self): - import threading - - manager = DatabaseSessionsManager(self._manager._database, self._manager._pool) - manager._multiplexed_session_lock = threading.Lock() - old_session = Mock() - old_session.delete.side_effect = Exception("delete failed") - new_session = Mock() - manager._multiplexed_session = old_session - - with patch.object( - manager, "_build_multiplexed_session", return_value=new_session - ): - result = manager._rotate_multiplexed_session() - self.assertTrue(result) - self.assertIs(manager._multiplexed_session, new_session) - old_session.delete.assert_called_once() - def _assert_true_with_timeout(self, condition: Callable) -> None: """Asserts that the given condition is met within a timeout period. diff --git a/packages/google-cloud-spanner/tests/unit/test_metrics_capture.py b/packages/google-cloud-spanner/tests/unit/test_metrics_capture.py index 1bd1c19f9bab..b8f3ccdf5828 100644 --- a/packages/google-cloud-spanner/tests/unit/test_metrics_capture.py +++ b/packages/google-cloud-spanner/tests/unit/test_metrics_capture.py @@ -50,3 +50,95 @@ def test_metrics_capture_exit(mock_tracer_factory): pass mock_tracer.record_operation_completion.assert_called_once() + + +def test_metrics_capture_reuse(mock_tracer_factory): + mock_tracer = mock.Mock() + mock_tracer_factory.return_value = mock_tracer + + capture = MetricsCapture() + with capture: + assert SpannerMetricsTracerFactory.get_current_tracer() is mock_tracer + + assert SpannerMetricsTracerFactory.get_current_tracer() is None + assert capture._token is None + + # Reusing the same context manager instance must not raise an error + with capture: + assert SpannerMetricsTracerFactory.get_current_tracer() is mock_tracer + + assert SpannerMetricsTracerFactory.get_current_tracer() is None + assert capture._token is None + + +def test_metrics_capture_disabled(): + SpannerMetricsTracerFactory(enabled=False) + try: + with MetricsCapture() as capture: + assert capture is not None + assert SpannerMetricsTracerFactory.get_current_tracer() is None + finally: + SpannerMetricsTracerFactory(enabled=True) + + +def test_metrics_capture_with_resource_info(mock_tracer_factory): + mock_tracer = mock.Mock() + mock_tracer_factory.return_value = mock_tracer + + resource_info = { + "project": "test_p", + "instance": "test_i", + "database": "test_d", + } + with MetricsCapture(resource_info=resource_info): + pass + + mock_tracer.set_project.assert_called_once_with("test_p") + mock_tracer.set_instance.assert_called_once_with("test_i") + mock_tracer.set_database.assert_called_once_with("test_d") + + +def test_metrics_capture_exit_without_token(): + capture = MetricsCapture() + assert capture.__exit__(None, None, None) is False + + +def test_metrics_capture_with_partial_resource_info(mock_tracer_factory): + mock_tracer = mock.Mock() + mock_tracer_factory.return_value = mock_tracer + with MetricsCapture(resource_info={"database": "only_db"}): + pass + mock_tracer.set_database.assert_called_once_with("only_db") + mock_tracer.set_project.assert_not_called() + mock_tracer.set_instance.assert_not_called() + + +def test_metrics_capture_factory_returns_none(mock_tracer_factory): + mock_tracer_factory.return_value = None + with MetricsCapture(resource_info={"project": "p"}): + pass + + +def test_metrics_capture_with_project_and_instance_only(mock_tracer_factory): + mock_tracer = mock.Mock() + mock_tracer_factory.return_value = mock_tracer + with MetricsCapture(resource_info={"project": "p", "instance": "i"}): + pass + mock_tracer.set_project.assert_called_once_with("p") + mock_tracer.set_instance.assert_called_once_with("i") + mock_tracer.set_database.assert_not_called() + + +def test_metrics_capture_exit_error_resets_token(mock_tracer_factory): + mock_tracer = mock.Mock() + mock_tracer.record_operation_completion.side_effect = RuntimeError( + "Completion failure" + ) + mock_tracer_factory.return_value = mock_tracer + + with pytest.raises(RuntimeError, match="Completion failure"): + with MetricsCapture(): + assert SpannerMetricsTracerFactory.get_current_tracer() is mock_tracer + + # Verified: Token is cleanly reset even on exception + assert SpannerMetricsTracerFactory.get_current_tracer() is None diff --git a/packages/google-cloud-spanner/tests/unit/test_metrics_interceptor.py b/packages/google-cloud-spanner/tests/unit/test_metrics_interceptor.py index aa31dc5f1210..a2f30a8b2b5e 100644 --- a/packages/google-cloud-spanner/tests/unit/test_metrics_interceptor.py +++ b/packages/google-cloud-spanner/tests/unit/test_metrics_interceptor.py @@ -16,6 +16,7 @@ import pytest +from google.cloud.spanner_v1.metrics.constants import _safe_decode_utf8 from google.cloud.spanner_v1.metrics.metrics_interceptor import MetricsInterceptor from google.cloud.spanner_v1.metrics.spanner_metrics_tracer_factory import ( SpannerMetricsTracerFactory, @@ -118,3 +119,1043 @@ def test_intercept_with_tracer(interceptor, mock_tracer_ctx): mock_tracer_ctx.record_attempt_completion.assert_called_once() mock_tracer_ctx.record_front_end_metrics.assert_called_once() mock_invoked_method.assert_called_once_with("request", call_details) + + +def test_format_method_name(): + from google.cloud.spanner_v1.metrics.metrics_interceptor import _format_method_name + + method_str = "/google.spanner.v1.Spanner/ExecuteStreamingSql" + expected = "Spanner.ExecuteStreamingSql" + assert _format_method_name(method_str) == expected + + # Test bytes input + method_bytes = b"/google.spanner.v1.Spanner/ExecuteStreamingSql" + assert _format_method_name(method_bytes) == expected + + # Verify cached result + assert _format_method_name(method_bytes) == expected + + # Unhashable type input + assert ( + _format_method_name(["/google.spanner.v1.Spanner/ExecuteSql"]) + == "['.google.spanner.v1.Spanner.ExecuteSql']" + ) + # Non-decodable bytes + assert "Spanner" in _format_method_name(b"/google.spanner.v1.Spanner/\xff\xfe") + + +def test_extract_resource_from_path_bytes_and_dict(interceptor): + path = "projects/p/instances/i/databases/d" + expected = {"project": "p", "instance": "i", "database": "d"} + + # Bytes key in list of tuples + metadata_bytes = [(b"google-cloud-resource-prefix", path.encode("utf-8"))] + assert interceptor._extract_resource_from_path(metadata_bytes) == expected + + # Dict metadata + metadata_dict = {"google-cloud-resource-prefix": path} + assert interceptor._extract_resource_from_path(metadata_dict) == expected + + # Dict with bytes key and value + metadata_dict_bytes = {b"google-cloud-resource-prefix": path.encode("utf-8")} + assert interceptor._extract_resource_from_path(metadata_dict_bytes) == expected + + # Empty metadata + assert interceptor._extract_resource_from_path([]) == {} + assert interceptor._extract_resource_from_path({}) == {} + + +def test_parse_resource_path_edge_cases(interceptor): + path = "projects/p1/instances/i1/databases/d1" + first = interceptor._parse_resource_path(path) + second = interceptor._parse_resource_path(path) + assert first == {"project": "p1", "instance": "i1", "database": "d1"} + assert first == second + + # Mutating returned dict should not corrupt future calls + first["mutated"] = True + third = interceptor._parse_resource_path(path) + assert "mutated" not in third + + assert interceptor._parse_resource_path("") == {} + assert interceptor._parse_resource_path(None) == {} + assert interceptor._parse_resource_path(12345) == {} + + # Database named "sessions" + db_named_sessions = "projects/p1/instances/i1/databases/sessions/sessions/s123" + assert interceptor._parse_resource_path(db_named_sessions) == { + "project": "p1", + "instance": "i1", + "database": "sessions", + "session": "s123", + } + + # Instance named "sessions" + instance_named_sessions = ( + "projects/p1/instances/sessions/databases/d1/sessions/s123" + ) + assert interceptor._parse_resource_path(instance_named_sessions) == { + "project": "p1", + "instance": "sessions", + "database": "d1", + "session": "s123", + } + + # Session paths + session_path = "projects/p1/instances/i1/databases/d1/sessions/s1" + session_result = interceptor._parse_resource_path(session_path) + assert session_result == { + "project": "p1", + "instance": "i1", + "database": "d1", + "session": "s1", + } + session_path_empty = "projects/p1/instances/i1/databases/d1/sessions/" + session_empty_result = interceptor._parse_resource_path(session_path_empty) + assert session_empty_result == { + "project": "p1", + "instance": "i1", + "database": "d1", + } + # Session part starting with slash + session_slash = "projects/p1/instances/i1/databases/d1/sessions//extra" + assert interceptor._parse_resource_path(session_slash) == { + "project": "p1", + "instance": "i1", + "database": "d1", + } + # Invalid paths with sessions must not return session + assert interceptor._parse_resource_path("invalid/sessions/s123") == {} + assert interceptor._parse_resource_path("/sessions/s123") == {} + + +def test_async_streaming_response_wrapper_not_awaitable(): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + _AsyncStreamingResponseWrapper, + ) + + mock_stream = MagicMock() + mock_tracer = MagicMock() + wrapper = _AsyncStreamingResponseWrapper(mock_stream, mock_tracer) + assert not hasattr(wrapper, "__await__") + + +@pytest.mark.asyncio +async def test_async_unary_response_wrapper_stream_unary(): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + _AsyncUnaryResponseWrapper, + ) + + class StreamUnaryMock: + def __init__(self): + self.written = [] + self.done_writing_called = False + + def write(self, data): + self.written.append(data) + + def done_writing(self): + self.done_writing_called = True + + def initial_metadata(self): + return [] + + def __await__(self): + async def _coro(): + return "done" + + return _coro().__await__() + + mock_stream_unary = StreamUnaryMock() + mock_tracer = MagicMock() + wrapper = _AsyncUnaryResponseWrapper(mock_stream_unary, mock_tracer) + wrapper.write("item1") + wrapper.done_writing() + assert mock_stream_unary.written == ["item1"] + assert mock_stream_unary.done_writing_called is True + result = await wrapper + assert result == "done" + mock_tracer.record_attempt_completion.assert_called_once() + + +def test_extract_resource_from_path_edge_cases(interceptor): + # Non-iterable metadata + assert interceptor._extract_resource_from_path(12345) == {} + assert interceptor._extract_resource_from_path(None) == {} + + # Dict without resource prefix + assert interceptor._extract_resource_from_path({"unrelated": "header"}) == {} + + # Non-decodable bytes in dict and list metadata + assert ( + interceptor._extract_resource_from_path( + {"google-cloud-resource-prefix": b"\xff\xfe\xfd"} + ) + == {} + ) + assert ( + interceptor._extract_resource_from_path( + [("google-cloud-resource-prefix", b"\xff\xfe\xfd")] + ) + == {} + ) + + # Malformed tuple entries (not length 2) + malformed = [("single_element",), ("a", "b", "c")] + assert interceptor._extract_resource_from_path(malformed) == {} + + # Generator metadata + path = "projects/p/instances/i/databases/d" + + def metadata_generator(): + yield ("unrelated", "value") + yield ("google-cloud-resource-prefix", path) + + assert interceptor._extract_resource_from_path(metadata_generator()) == { + "project": "p", + "instance": "i", + "database": "d", + } + + +@pytest.mark.asyncio +async def test_async_metrics_interceptor(mock_tracer_ctx): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + AsyncMetricsInterceptor, + ) + + interceptor = AsyncMetricsInterceptor() + + # 1. Async unary call + class AwaitableCallMock: + def initial_metadata(self): + return [("server-timing", "gfet4t7; dur=55")] + + def __await__(self): + async def _coro(): + return "unary_result" + + return _coro().__await__() + + async def mock_unary_continuation(details, request): + return AwaitableCallMock() + + call_details = MagicMock( + method="/google.spanner.v1.Spanner/ExecuteSql", + metadata=[ + ( + "google-cloud-resource-prefix", + "projects/p_async/instances/i_async/databases/d_async", + ) + ], + ) + + wrapped_call = await interceptor.intercept_unary_unary( + mock_unary_continuation, call_details, "req" + ) + result = await wrapped_call + assert result == "unary_result" + mock_tracer_ctx.record_attempt_start.assert_called_once() + mock_tracer_ctx.record_attempt_completion.assert_called_once() + mock_tracer_ctx.record_front_end_metrics.assert_called_once() + mock_tracer_ctx.set_method.assert_called_with("Spanner.ExecuteSql") + mock_tracer_ctx.set_project.assert_called_with("p_async") + mock_tracer_ctx.set_instance.assert_called_with("i_async") + mock_tracer_ctx.set_database.assert_called_with("d_async") + + # 2. Async streaming call + mock_tracer_ctx.record_attempt_start.reset_mock() + mock_tracer_ctx.record_attempt_completion.reset_mock() + + class AsyncIteratorMock: + def __init__(self, items): + self._items = list(items) + + def __aiter__(self): + return self + + async def __anext__(self): + if not self._items: + raise StopAsyncIteration + return self._items.pop(0) + + def initial_metadata(self): + return [("server-timing", "afe; dur=20")] + + async def mock_stream_continuation(details, request): + return AsyncIteratorMock(["chunk1", "chunk2"]) + + stream_details = MagicMock( + method="/google.spanner.v1.Spanner/ExecuteStreamingSql", + metadata=[], + ) + + wrapped_stream = await interceptor.intercept_unary_stream( + mock_stream_continuation, stream_details, "req" + ) + items = [] + async for item in wrapped_stream: + items.append(item) + assert items == ["chunk1", "chunk2"] + mock_tracer_ctx.record_attempt_start.assert_called_once() + mock_tracer_ctx.record_attempt_completion.assert_called_once() + + +def test_streaming_response_wrapper_lifecycle(): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + _StreamingResponseWrapper, + ) + + # 1. Normal iteration + mock_response = MagicMock() + mock_response.__iter__.return_value = iter(["chunk1", "chunk2"]) + mock_response.initial_metadata.return_value = [("server-timing", "gfet4t7; dur=10")] + mock_tracer = MagicMock() + + wrapper = _StreamingResponseWrapper(mock_response, mock_tracer) + items = list(wrapper) + assert items == ["chunk1", "chunk2"] + mock_tracer.record_attempt_completion.assert_called_once() + mock_tracer.record_front_end_metrics.assert_called_once_with( + [("server-timing", "gfet4t7; dur=10")] + ) + + # 2. Exception during iteration + mock_response_err = MagicMock() + mock_response_err.__iter__.return_value = iter(["ok"]) + + class FaultyIterator: + def __iter__(self): + return self + + def __next__(self): + raise RuntimeError("Stream broken") + + mock_tracer_err = MagicMock() + wrapper_err = _StreamingResponseWrapper(FaultyIterator(), mock_tracer_err) + with pytest.raises(RuntimeError, match="Stream broken"): + next(wrapper_err) + mock_tracer_err.record_attempt_completion.assert_called_once() + + # 3. Explicit cancellation + mock_response_cancel = MagicMock() + mock_tracer_cancel = MagicMock() + wrapper_cancel = _StreamingResponseWrapper(mock_response_cancel, mock_tracer_cancel) + wrapper_cancel.cancel() + mock_tracer_cancel.record_attempt_completion.assert_called_once_with( + status="CANCELLED" + ) + mock_response_cancel.cancel.assert_called_once() + + # 4. Finalizer (__del__) when not completed + mock_response_del = MagicMock() + mock_tracer_del = MagicMock() + wrapper_del = _StreamingResponseWrapper(mock_response_del, mock_tracer_del) + wrapper_del.__del__() + mock_tracer_del.record_attempt_completion.assert_called_once_with( + status="CANCELLED" + ) + + # 5. __getattr__ delegation and error handling + mock_response_attr = MagicMock() + mock_response_attr.custom_field = "custom_value" + mock_response_attr.initial_metadata.side_effect = RuntimeError("Metadata failed") + mock_tracer_attr = MagicMock() + mock_tracer_attr.record_attempt_completion.side_effect = RuntimeError( + "Tracer failed" + ) + wrapper_attr = _StreamingResponseWrapper(mock_response_attr, mock_tracer_attr) + assert wrapper_attr.custom_field == "custom_value" + # _record_metrics should swallow exceptions gracefully + wrapper_attr._record_metrics() + # Calling it a second time hits early return + wrapper_attr._record_metrics() + + # Cancel and del error handling + mock_tracer_cancel_err = MagicMock() + mock_tracer_cancel_err.record_attempt_completion.side_effect = RuntimeError( + "Cancel error" + ) + wrapper_cancel_err = _StreamingResponseWrapper(MagicMock(), mock_tracer_cancel_err) + wrapper_cancel_err.cancel() + + mock_tracer_del_err = MagicMock() + mock_tracer_del_err.record_attempt_completion.side_effect = RuntimeError( + "Del error" + ) + wrapper_del_err = _StreamingResponseWrapper(MagicMock(), mock_tracer_del_err) + wrapper_del_err.__del__() + + +@pytest.mark.asyncio +async def test_async_unary_response_wrapper_lifecycle(): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + _AsyncUnaryResponseWrapper, + ) + + # 1. Exception during await + class FaultyAwaitable: + def __await__(self): + async def _coro(): + raise ValueError("RPC failed") + + return _coro().__await__() + + mock_tracer_err = MagicMock() + wrapper_err = _AsyncUnaryResponseWrapper(FaultyAwaitable(), mock_tracer_err) + with pytest.raises(ValueError, match="RPC failed"): + await wrapper_err + mock_tracer_err.record_attempt_completion.assert_called_once() + + # 2. Explicit cancellation + mock_response_cancel = MagicMock() + mock_tracer_cancel = MagicMock() + wrapper_cancel = _AsyncUnaryResponseWrapper( + mock_response_cancel, mock_tracer_cancel + ) + wrapper_cancel.cancel() + mock_tracer_cancel.record_attempt_completion.assert_called_once_with( + status="CANCELLED" + ) + mock_response_cancel.cancel.assert_called_once() + + # 3. Finalizer (__del__) when unawaited + mock_response_del = MagicMock() + mock_tracer_del = MagicMock() + wrapper_del = _AsyncUnaryResponseWrapper(mock_response_del, mock_tracer_del) + wrapper_del.__del__() + mock_tracer_del.record_attempt_completion.assert_called_once_with( + status="CANCELLED" + ) + + # 4. Proxy methods + mock_delegate = MagicMock() + mock_tracer = MagicMock() + wrapper = _AsyncUnaryResponseWrapper(mock_delegate, mock_tracer) + + wrapper.add_done_callback(MagicMock()) + mock_delegate.add_done_callback.assert_called_once() + wrapper.cancelled() + mock_delegate.cancelled.assert_called_once() + wrapper.code() + mock_delegate.code.assert_called_once() + wrapper.details() + mock_delegate.details.assert_called_once() + wrapper.done() + mock_delegate.done.assert_called_once() + wrapper.initial_metadata() + mock_delegate.initial_metadata.assert_called_once() + wrapper.time_remaining() + mock_delegate.time_remaining.assert_called_once() + wrapper.trailing_metadata() + mock_delegate.trailing_metadata.assert_called_once() + wrapper.wait_for_connection() + mock_delegate.wait_for_connection.assert_called_once() + assert wrapper.some_custom_attr == mock_delegate.some_custom_attr + + # 5. Async initial metadata and error handling + class AsyncMetadataCall: + async def initial_metadata(self): + return [("server-timing", "gfet4t7; dur=40")] + + def __await__(self): + async def _coro(): + return "ok" + + return _coro().__await__() + + mock_tracer_meta = MagicMock() + wrapper_meta = _AsyncUnaryResponseWrapper(AsyncMetadataCall(), mock_tracer_meta) + result = await wrapper_meta + assert result == "ok" + mock_tracer_meta.record_front_end_metrics.assert_called_once_with( + [("server-timing", "gfet4t7; dur=40")] + ) + # Calling it a second time hits early return + await wrapper_meta._record_metrics() + + # Cancel and del error handling + mock_tracer_cancel_err = MagicMock() + mock_tracer_cancel_err.record_attempt_completion.side_effect = RuntimeError( + "Cancel error" + ) + wrapper_cancel_err = _AsyncUnaryResponseWrapper(MagicMock(), mock_tracer_cancel_err) + wrapper_cancel_err.cancel() + + mock_tracer_del_err = MagicMock() + mock_tracer_del_err.record_attempt_completion.side_effect = RuntimeError( + "Del error" + ) + wrapper_del_err = _AsyncUnaryResponseWrapper(MagicMock(), mock_tracer_del_err) + wrapper_del_err.__del__() + + # Metadata error in _record_metrics + class UnaryMetadataErrorCall: + def initial_metadata(self): + raise RuntimeError("Metadata failed") + + def __await__(self): + async def _coro(): + return "ok" + + return _coro().__await__() + + mock_tracer_err2 = MagicMock() + mock_tracer_err2.record_attempt_completion.side_effect = RuntimeError( + "Tracer failed" + ) + wrapper_meta_err = _AsyncUnaryResponseWrapper( + UnaryMetadataErrorCall(), mock_tracer_err2 + ) + await wrapper_meta_err + + +@pytest.mark.asyncio +async def test_async_streaming_response_wrapper_lifecycle(): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + _AsyncStreamingResponseWrapper, + ) + + # 1. Exception during async iteration + class FaultyAsyncIterator: + def __aiter__(self): + return self + + async def __anext__(self): + raise RuntimeError("Stream error") + + mock_tracer_err = MagicMock() + wrapper_err = _AsyncStreamingResponseWrapper(FaultyAsyncIterator(), mock_tracer_err) + with pytest.raises(RuntimeError, match="Stream error"): + async for _ in wrapper_err: + pass + mock_tracer_err.record_attempt_completion.assert_called_once() + + # 2. Cancellation + mock_response_cancel = MagicMock() + mock_tracer_cancel = MagicMock() + wrapper_cancel = _AsyncStreamingResponseWrapper( + mock_response_cancel, mock_tracer_cancel + ) + wrapper_cancel.cancel() + mock_tracer_cancel.record_attempt_completion.assert_called_once_with( + status="CANCELLED" + ) + mock_response_cancel.cancel.assert_called_once() + + # 3. Finalizer (__del__) when not completed + mock_response_del = MagicMock() + mock_tracer_del = MagicMock() + wrapper_del = _AsyncStreamingResponseWrapper(mock_response_del, mock_tracer_del) + wrapper_del.__del__() + mock_tracer_del.record_attempt_completion.assert_called_once_with( + status="CANCELLED" + ) + + # 4. Proxy methods + mock_delegate = MagicMock() + mock_tracer = MagicMock() + wrapper = _AsyncStreamingResponseWrapper(mock_delegate, mock_tracer) + + wrapper.add_done_callback(MagicMock()) + mock_delegate.add_done_callback.assert_called_once() + wrapper.cancelled() + mock_delegate.cancelled.assert_called_once() + wrapper.code() + mock_delegate.code.assert_called_once() + wrapper.details() + mock_delegate.details.assert_called_once() + wrapper.done() + mock_delegate.done.assert_called_once() + wrapper.initial_metadata() + mock_delegate.initial_metadata.assert_called_once() + wrapper.time_remaining() + mock_delegate.time_remaining.assert_called_once() + wrapper.trailing_metadata() + mock_delegate.trailing_metadata.assert_called_once() + wrapper.wait_for_connection() + mock_delegate.wait_for_connection.assert_called_once() + wrapper.read() + mock_delegate.read.assert_called_once() + wrapper.write("data") + mock_delegate.write.assert_called_once_with("data") + wrapper.done_writing() + mock_delegate.done_writing.assert_called_once() + assert wrapper.custom_attr == mock_delegate.custom_attr + + # 5. Async initial metadata in stream + class StreamAsyncMetadata: + def __init__(self): + self.yielded = False + + def __aiter__(self): + return self + + async def __anext__(self): + if not self.yielded: + self.yielded = True + return "item" + raise StopAsyncIteration + + async def initial_metadata(self): + return [("server-timing", "afe; dur=15")] + + mock_tracer_stream_meta = MagicMock() + wrapper_stream_meta = _AsyncStreamingResponseWrapper( + StreamAsyncMetadata(), mock_tracer_stream_meta + ) + items = [] + async for item in wrapper_stream_meta: + items.append(item) + assert items == ["item"] + mock_tracer_stream_meta.record_front_end_metrics.assert_called_once_with( + [("server-timing", "afe; dur=15")] + ) + # Calling it a second time hits early return + await wrapper_stream_meta._record_metrics() + + # Cancel and del error handling + mock_tracer_cancel_err = MagicMock() + mock_tracer_cancel_err.record_attempt_completion.side_effect = RuntimeError( + "Cancel error" + ) + wrapper_cancel_err = _AsyncStreamingResponseWrapper( + MagicMock(), mock_tracer_cancel_err + ) + wrapper_cancel_err.cancel() + + mock_tracer_del_err = MagicMock() + mock_tracer_del_err.record_attempt_completion.side_effect = RuntimeError( + "Del error" + ) + wrapper_del_err = _AsyncStreamingResponseWrapper(MagicMock(), mock_tracer_del_err) + wrapper_del_err.__del__() + + # Metadata error in stream + class StreamMetadataError: + def __init__(self): + self.yielded = False + + def __aiter__(self): + return self + + async def __anext__(self): + if not self.yielded: + self.yielded = True + return "chunk" + raise StopAsyncIteration + + def initial_metadata(self): + raise RuntimeError("Stream metadata failed") + + mock_tracer_stream_err = MagicMock() + mock_tracer_stream_err.record_attempt_completion.side_effect = RuntimeError( + "Tracer stream error" + ) + wrapper_stream_err = _AsyncStreamingResponseWrapper( + StreamMetadataError(), mock_tracer_stream_err + ) + async for _ in wrapper_stream_err: + pass + + +def test_metrics_interceptor_sync_methods_and_disabled(mock_tracer_ctx): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + MetricsInterceptor, + _StreamingResponseWrapper, + _wrap_response, + ) + + interceptor = MetricsInterceptor() + + # 0. _set_metrics_tracer_attributes when tracer is None + token = SpannerMetricsTracerFactory._current_metrics_tracer_ctx.set(None) + try: + interceptor._set_metrics_tracer_attributes({"project": "p"}) + finally: + SpannerMetricsTracerFactory._current_metrics_tracer_ctx.reset(token) + + # 1. Intercept when disabled + SpannerMetricsTracerFactory(enabled=False) + mock_continuation = MagicMock(return_value="raw_response") + call_details = MagicMock( + method="/google.spanner.v1.Spanner/ExecuteSql", metadata=[] + ) + result = interceptor.intercept(mock_continuation, "request", call_details) + assert result == "raw_response" + mock_continuation.assert_called_once_with("request", call_details) + SpannerMetricsTracerFactory(enabled=True) + + # 2. _wrap_response with streaming response + class StreamingCallMock: + def __next__(self): + raise StopIteration + + mock_stream_call = StreamingCallMock() + wrapped_stream = _wrap_response(mock_stream_call, mock_tracer_ctx) + assert isinstance(wrapped_stream, _StreamingResponseWrapper) + + # 3. _wrap_response with unary response handling errors gracefully + class SimpleUnaryResponse: + def initial_metadata(self): + raise RuntimeError("Metadata error") + + unary_response = SimpleUnaryResponse() + mock_faulty_tracer = MagicMock() + mock_faulty_tracer.record_attempt_completion.side_effect = RuntimeError( + "Tracer error" + ) + result_unary = _wrap_response(unary_response, mock_faulty_tracer) + assert result_unary is unary_response + + # 4. _StreamingResponseWrapper with initial_metadata error + mock_resp_meta_err = MagicMock() + mock_resp_meta_err.__iter__.return_value = iter(["chunk"]) + mock_resp_meta_err.initial_metadata.side_effect = RuntimeError("Metadata failed") + wrapper_stream_meta_err = _StreamingResponseWrapper(mock_resp_meta_err, MagicMock()) + assert list(wrapper_stream_meta_err) == ["chunk"] + + +@pytest.mark.asyncio +async def test_async_wrapper_additional_error_branches(): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + _AsyncStreamingResponseWrapper, + _AsyncUnaryResponseWrapper, + ) + + # Unary wrapper with initial_metadata error + class UnaryInitialMetaErr: + def initial_metadata(self): + raise RuntimeError("Initial meta err") + + def __await__(self): + async def _coro(): + return "ok" + + return _coro().__await__() + + wrapper_unary_meta_err = _AsyncUnaryResponseWrapper( + UnaryInitialMetaErr(), MagicMock() + ) + assert await wrapper_unary_meta_err == "ok" + + # Streaming wrapper without __aiter__ on response directly calling __anext__ + class DirectAsyncIterator: + def __init__(self): + self.done = False + + async def __anext__(self): + if not self.done: + self.done = True + return "first" + raise StopAsyncIteration + + def initial_metadata(self): + raise RuntimeError("Stream meta err") + + wrapper_direct = _AsyncStreamingResponseWrapper(DirectAsyncIterator(), MagicMock()) + item = await wrapper_direct.__anext__() + assert item == "first" + with pytest.raises(StopAsyncIteration): + await wrapper_direct.__anext__() + + +@pytest.mark.asyncio +async def test_async_metrics_interceptor_all_methods_and_disabled(mock_tracer_ctx): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + AsyncMetricsInterceptor, + ) + + interceptor = AsyncMetricsInterceptor() + + # 1. Intercept when disabled + SpannerMetricsTracerFactory(enabled=False) + + async def mock_continuation(details, request): + return "async_raw" + + call_details = MagicMock( + method="/google.spanner.v1.Spanner/ExecuteSql", metadata=[] + ) + result = await interceptor.intercept_unary_unary( + mock_continuation, call_details, "req" + ) + assert result == "async_raw" + SpannerMetricsTracerFactory(enabled=True) + + # 2. intercept_stream_unary + async def mock_stream_unary_continuation(details, request_iterator): + class AwaitableResult: + def __await__(self): + async def _coro(): + return "stream_unary_done" + + return _coro().__await__() + + return AwaitableResult() + + wrapped_stream_unary = await interceptor.intercept_stream_unary( + mock_stream_unary_continuation, call_details, ["req1"] + ) + assert await wrapped_stream_unary == "stream_unary_done" + + # 3. intercept_stream_stream + async def mock_stream_stream_continuation(details, request_iterator): + class AsyncStreamResult: + def __aiter__(self): + return self + + async def __anext__(self): + raise StopAsyncIteration + + return AsyncStreamResult() + + wrapped_stream_stream = await interceptor.intercept_stream_stream( + mock_stream_stream_continuation, call_details, ["req1"] + ) + items = [] + async for item in wrapped_stream_stream: + items.append(item) + assert items == [] + + +@pytest.mark.asyncio +async def test_interceptor_wrapper_and_branch_edge_cases(mock_tracer_ctx): + from google.cloud.spanner_v1.metrics.metrics_interceptor import ( + AsyncMetricsInterceptor, + MetricsInterceptor, + _AsyncStreamingResponseWrapper, + _AsyncUnaryResponseWrapper, + _StreamingResponseWrapper, + _wrap_response, + ) + + interceptor = MetricsInterceptor() + async_interceptor = AsyncMetricsInterceptor() + + # 1. _set_metrics_tracer_attributes with partial and empty dicts + interceptor._set_metrics_tracer_attributes({"project": "p"}) + interceptor._set_metrics_tracer_attributes({"instance": "i"}) + interceptor._set_metrics_tracer_attributes({}) + + # 2. Interceptor call when tracer already has all resource attributes set + mock_tracer_ctx.client_attributes["project_id"] = "proj" + mock_tracer_ctx.client_attributes["instance_id"] = "inst" + mock_tracer_ctx.client_attributes["database"] = "db" + + mock_continuation = MagicMock(return_value="response") + call_details = MagicMock( + method="/google.spanner.v1.Spanner/ExecuteSql", metadata=[] + ) + result = interceptor.intercept(mock_continuation, "req", call_details) + assert result == "response" + + async def mock_async_continuation(details, request): + class AsyncCall: + def __await__(self): + async def _coro(): + return "async_response" + + return _coro().__await__() + + return AsyncCall() + + async_wrapped = await async_interceptor.intercept_unary_unary( + mock_async_continuation, call_details, "req" + ) + assert await async_wrapped == "async_response" + + # 3. _wrap_response unary branch when initial_metadata raises or is missing + class UnaryWithFailingMetadata: + def initial_metadata(self): + raise RuntimeError("Metadata failed") + + mock_tracer = MagicMock() + _wrap_response(UnaryWithFailingMetadata(), mock_tracer) + mock_tracer.record_attempt_completion.assert_called_once() + mock_tracer.record_front_end_metrics.assert_called_once_with([]) + + _wrap_response("no_initial_metadata", mock_tracer) + + # 4. _StreamingResponseWrapper: double cancel, missing cancel, del after recorded + # 4. _StreamingResponseWrapper: + # 4a. cancel returns False + stream_delegate_refused = MagicMock() + stream_delegate_refused.cancel.return_value = False + mock_tracer_unrecorded = MagicMock() + stream_wrapper_refused = _StreamingResponseWrapper( + stream_delegate_refused, mock_tracer_unrecorded + ) + assert stream_wrapper_refused.cancel() is False + assert stream_wrapper_refused._metrics_recorded is False + mock_tracer_unrecorded.record_attempt_completion.assert_not_called() + + # 4b. cancel with metadata error and successful metadata + stream_delegate_meta_err = MagicMock() + stream_delegate_meta_err.initial_metadata.side_effect = RuntimeError("meta failed") + stream_wrapper_meta_err = _StreamingResponseWrapper( + stream_delegate_meta_err, mock_tracer + ) + stream_wrapper_meta_err.cancel() + + stream_delegate_with_meta = MagicMock() + stream_delegate_with_meta.initial_metadata.return_value = [ + ("server-timing", "gfet4t7; dur=50") + ] + stream_wrapper_with_meta = _StreamingResponseWrapper( + stream_delegate_with_meta, mock_tracer + ) + stream_wrapper_with_meta.cancel() + mock_tracer.record_front_end_metrics.assert_called_with( + [("server-timing", "gfet4t7; dur=50")] + ) + + stream_delegate = MagicMock(spec=["__next__"]) + stream_wrapper = _StreamingResponseWrapper(stream_delegate, mock_tracer) + stream_wrapper.cancel() + stream_wrapper.cancel() + stream_wrapper.__del__() + + # 5. _AsyncUnaryResponseWrapper: + # 5a. cancel returns False + async_unary_refused = MagicMock() + async_unary_refused.cancel.return_value = False + async_unary_wrapper_refused = _AsyncUnaryResponseWrapper( + async_unary_refused, mock_tracer_unrecorded + ) + assert async_unary_wrapper_refused.cancel() is False + assert async_unary_wrapper_refused._metrics_recorded is False + + # 5b. cancel with awaitable metadata, failing metadata, and normal metadata + async_unary_async_meta = MagicMock() + + async def async_meta(): + return [("server-timing", "afe; dur=25")] + + async_unary_async_meta.initial_metadata.return_value = async_meta() + async_unary_wrapper_meta = _AsyncUnaryResponseWrapper( + async_unary_async_meta, mock_tracer + ) + async_unary_wrapper_meta.cancel() + + async_unary_err_meta = MagicMock() + async_unary_err_meta.initial_metadata.side_effect = RuntimeError("async meta error") + async_unary_wrapper_err = _AsyncUnaryResponseWrapper( + async_unary_err_meta, mock_tracer + ) + async_unary_wrapper_err.cancel() + + async_unary_normal_meta = MagicMock() + async_unary_normal_meta.initial_metadata.return_value = [ + ("server-timing", "afe; dur=25") + ] + async_unary_wrapper_normal = _AsyncUnaryResponseWrapper( + async_unary_normal_meta, mock_tracer + ) + async_unary_wrapper_normal.cancel() + mock_tracer.record_front_end_metrics.assert_called_with( + [("server-timing", "afe; dur=25")] + ) + + async_unary_delegate = MagicMock(spec=["cancel"]) + async_unary_wrapper = _AsyncUnaryResponseWrapper(async_unary_delegate, mock_tracer) + async_unary_wrapper.cancel() + async_unary_wrapper.cancel() + + # 6. _AsyncStreamingResponseWrapper: + # 6a. cancel returns False + async_stream_refused = MagicMock() + async_stream_refused.cancel.return_value = False + async_stream_wrapper_refused = _AsyncStreamingResponseWrapper( + async_stream_refused, mock_tracer_unrecorded + ) + assert async_stream_wrapper_refused.cancel() is False + assert async_stream_wrapper_refused._metrics_recorded is False + + # 6b. cancel with awaitable metadata, failing metadata, and normal metadata + async_stream_async_meta = MagicMock() + async_stream_async_meta.initial_metadata.return_value = async_meta() + async_stream_wrapper_meta = _AsyncStreamingResponseWrapper( + async_stream_async_meta, mock_tracer + ) + async_stream_wrapper_meta.cancel() + + async_stream_err_meta = MagicMock() + async_stream_err_meta.initial_metadata.side_effect = RuntimeError( + "async meta error" + ) + async_stream_wrapper_err = _AsyncStreamingResponseWrapper( + async_stream_err_meta, mock_tracer + ) + async_stream_wrapper_err.cancel() + + async_stream_normal_meta = MagicMock() + async_stream_normal_meta.initial_metadata.return_value = [ + ("server-timing", "gfet4t7; dur=30") + ] + async_stream_wrapper_normal = _AsyncStreamingResponseWrapper( + async_stream_normal_meta, mock_tracer + ) + async_stream_wrapper_normal.cancel() + mock_tracer.record_front_end_metrics.assert_called_with( + [("server-timing", "gfet4t7; dur=30")] + ) + + async_stream_delegate = MagicMock(spec=["cancel"]) + async_stream_wrapper = _AsyncStreamingResponseWrapper( + async_stream_delegate, mock_tracer + ) + async_stream_wrapper.cancel() + async_stream_wrapper.cancel() + + # Async response without __aiter__ (custom async iterator) + class CustomAsyncIterator: + async def __anext__(self): + raise StopAsyncIteration + + custom_async_wrapper = _AsyncStreamingResponseWrapper( + CustomAsyncIterator(), mock_tracer + ) + assert custom_async_wrapper.__aiter__() is custom_async_wrapper + async for _ in custom_async_wrapper: + pass + + # Async response where __anext__ is called before __aiter__ + class AsyncStreamWithAiter: + def __aiter__(self): + async def _gen(): + if False: + yield 1 + + return _gen() + + anext_first_wrapper = _AsyncStreamingResponseWrapper( + AsyncStreamWithAiter(), mock_tracer + ) + with pytest.raises(StopAsyncIteration): + await anext_first_wrapper.__anext__() + + +def test_safe_decode_utf8(): + assert _safe_decode_utf8(None) == "" + assert _safe_decode_utf8("hello") == "hello" + assert _safe_decode_utf8(b"world") == "world" + assert _safe_decode_utf8(123) == "123" + + +def test_prepare_attempt(): + mock_tracer = MagicMock() + mock_tracer.client_attributes = {} + call_details = MagicMock() + call_details.metadata = [ + ("google-cloud-resource-prefix", "projects/p/instances/i/databases/d") + ] + call_details.method = "/google.spanner.v1.Spanner/ExecuteSql" + + MetricsInterceptor._prepare_attempt(mock_tracer, call_details) + + mock_tracer.set_project.assert_called_with("p") + mock_tracer.set_instance.assert_called_with("i") + mock_tracer.set_database.assert_called_with("d") + mock_tracer.set_method.assert_called_with("Spanner.ExecuteSql") + mock_tracer.record_attempt_start.assert_called_once() diff --git a/packages/google-cloud-spanner/tests/unit/test_metrics_tracer.py b/packages/google-cloud-spanner/tests/unit/test_metrics_tracer.py index c645905c40ee..243859e529f4 100644 --- a/packages/google-cloud-spanner/tests/unit/test_metrics_tracer.py +++ b/packages/google-cloud-spanner/tests/unit/test_metrics_tracer.py @@ -284,10 +284,59 @@ def test_extract_front_end_latencies(): ] assert MetricsTracer.extract_front_end_latencies(metadata_list) == (123, 100) + # Combined header in single value (standard Spanner wire response) + combined = [("server-timing", "gfet4t7; dur=55, afe; dur=23")] + assert MetricsTracer.extract_front_end_latencies(combined) == (55, 23) + + # Bytes header key in list of tuples + bytes_list = [(b"server-timing", "gfet4t7; dur=55, afe; dur=23")] + assert MetricsTracer.extract_front_end_latencies(bytes_list) == (55, 23) + # Valid metadata dict metadata_dict = {"server-timing": "gfet4t7; dur=456"} assert MetricsTracer.extract_front_end_latencies(metadata_dict) == (456, None) + # Metadata dict with bytes key and combined value + metadata_dict_bytes = {b"server-timing": "gfet4t7; dur=55, afe; dur=23"} + assert MetricsTracer.extract_front_end_latencies(metadata_dict_bytes) == (55, 23) + + # Metadata dict with mixed case key + metadata_dict_case = {"Server-Timing": "gfet4t7; dur=55, afe; dur=23"} + assert MetricsTracer.extract_front_end_latencies(metadata_dict_case) == (55, 23) + + # Metadata dict with list of values + metadata_dict_list = {"server-timing": ["gfet4t7; dur=12", "afe; dur=34"]} + assert MetricsTracer.extract_front_end_latencies(metadata_dict_list) == (12, 34) + + # Floating point latencies truncated to int + float_headers = [("server-timing", "gfet4t7; dur=55.8, afe; dur=23.2")] + assert MetricsTracer.extract_front_end_latencies(float_headers) == (55, 23) + + # Bytes header value + bytes_val = [("server-timing", b"gfet4t7; dur=55, afe; dur=23")] + assert MetricsTracer.extract_front_end_latencies(bytes_val) == (55, 23) + + # Generator / custom iterable (simulating grpc.aio.Metadata) + def timing_generator(): + yield ("unrelated", "1") + yield ("server-timing", "gfet4t7; dur=77, afe; dur=88") + + assert MetricsTracer.extract_front_end_latencies(timing_generator()) == (77, 88) + + # Dict with multiple case variants (ensuring no premature loop termination) + metadata_dict_multiple_cases = { + "Server-Timing": "gfet4t7; dur=55", + "server-timing": "afe; dur=23", + } + assert MetricsTracer.extract_front_end_latencies(metadata_dict_multiple_cases) == ( + 55, + 23, + ) + + # Sequence with list-of-values + metadata_seq_list = [("server-timing", ["gfet4t7; dur=12", "afe; dur=34"])] + assert MetricsTracer.extract_front_end_latencies(metadata_seq_list) == (12, 34) + # Missing header assert MetricsTracer.extract_front_end_latencies([("other-header", "val")]) == ( None, @@ -295,6 +344,45 @@ def test_extract_front_end_latencies(): ) assert MetricsTracer.extract_front_end_latencies(None) == (None, None) + # Non-iterable or malformed metadata + assert MetricsTracer.extract_front_end_latencies(12345) == (None, None) + assert MetricsTracer.extract_front_end_latencies([("single_item",)]) == (None, None) + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", "gfet4t7; dur=invalid")] + ) == (None, None) + # Trigger ValueError in float conversion (dur=.) + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", "gfet4t7; dur=.")] + ) == (None, None) + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", "afe; dur=.")] + ) == (None, None) + # Non-string, non-bytes header value + assert MetricsTracer.extract_front_end_latencies([("server-timing", 12345)]) == ( + None, + None, + ) + # Non-decodable bytes fallback + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", b"\xff\xfe\xfd")] + ) == (None, None) + + # Non-matching bytes key and non-str non-bytes key + assert MetricsTracer.extract_front_end_latencies( + [(b"unrelated", "val"), (12345, "val")] + ) == (None, None) + + # Empty header value + assert MetricsTracer.extract_front_end_latencies([("server-timing", "")]) == ( + None, + None, + ) + + # Separate headers where first sets AFE, second sets GFE + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", "afe; dur=20"), ("server-timing", "gfet4t7; dur=30")] + ) == (30, 20) + def test_record_front_end_metrics(metrics_tracer): mock_gfe_latency = mock.create_autospec(Histogram, instance=True) @@ -324,6 +412,15 @@ def test_record_front_end_metrics(metrics_tracer): assert mock_afe_latency.record.call_count == 1 assert mock_afe_missing.add.call_count == 1 + # When disabled, record_front_end_metrics does nothing + metrics_tracer.enabled = False + metrics_tracer.record_front_end_metrics( + [("server-timing", "gfet4t7; dur=88"), ("server-timing", "afe; dur=90")] + ) + assert mock_gfe_latency.record.call_count == 1 + assert mock_afe_latency.record.call_count == 1 + metrics_tracer.enabled = True + def test_record_afe_latency(metrics_tracer): mock_afe_latency = mock.create_autospec(Histogram, instance=True) @@ -367,3 +464,126 @@ def test_record_afe_connectivity_error_count(metrics_tracer): metrics_tracer.record_afe_connectivity_error_count() assert mock_afe_missing.add.call_count == 1 metrics_tracer.enabled = True + + +def test_attribute_caching_and_invalidation(metrics_tracer): + # Test attempt attribute caching + metrics_tracer.current_op.new_attempt() + metrics_tracer.current_op.current_attempt.status = "OK" + + first_attempt_attrs = metrics_tracer._create_attempt_otel_attributes() + second_attempt_attrs = metrics_tracer._create_attempt_otel_attributes() + assert first_attempt_attrs == second_attempt_attrs + # Same cached object returned + assert first_attempt_attrs is second_attempt_attrs + + # Changing attempt status returns new object + metrics_tracer.current_op.current_attempt.status = "UNAVAILABLE" + third_attempt_attrs = metrics_tracer._create_attempt_otel_attributes() + assert third_attempt_attrs["status"] == "UNAVAILABLE" + assert third_attempt_attrs is not first_attempt_attrs + + # Test operation attribute caching + first_op_attrs = metrics_tracer._create_operation_otel_attributes() + second_op_attrs = metrics_tracer._create_operation_otel_attributes() + assert first_op_attrs is second_op_attrs + + # Non-matching AFE and GFE patterns when substring is present (lookbehind boundary check) + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", "safe; dur=55")] + ) == (None, None) + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", "custom-afe; dur=20")] + ) == (None, None) + assert MetricsTracer.extract_front_end_latencies( + [("server-timing", "x-gfet4t7; dur=15")] + ) == (None, None) + assert MetricsTracer.extract_front_end_latencies( + [ + ( + "server-timing", + "safe; dur=55, afe; dur=25, x-gfet4t7; dur=10, gfet4t7; dur=35", + ) + ] + ) == (35, 25) + + # Setter method invalidates cache + empty_tracer = MetricsTracer( + enabled=True, + instrument_attempt_latency=mock.MagicMock(), + instrument_attempt_counter=mock.MagicMock(), + instrument_operation_latency=mock.MagicMock(), + instrument_operation_counter=mock.MagicMock(), + client_attributes={}, + instrument_gfe_latency=mock.MagicMock(), + instrument_gfe_connectivity_error_count=mock.MagicMock(), + instrument_afe_latency=mock.MagicMock(), + instrument_afe_connectivity_error_count=mock.MagicMock(), + ) + empty_tracer.set_project("new-project") + assert empty_tracer.client_attributes["project_id"] == "new-project" + + metrics_tracer.set_instance("updated-instance") + invalidated_attrs = metrics_tracer._create_attempt_otel_attributes() + assert invalidated_attrs["instance_id"] == "updated-instance" + assert invalidated_attrs is not third_attempt_attrs + + # Direct mutation of client_attributes invalidates cache via _ObservableDict + before_direct = metrics_tracer._create_attempt_otel_attributes() + metrics_tracer.client_attributes["database"] = "mutated_db" + after_direct = metrics_tracer._create_attempt_otel_attributes() + assert after_direct["database"] == "mutated_db" + assert after_direct is not before_direct + + # _ObservableDict copy returns standard dict + attrs_copy = metrics_tracer.client_attributes.copy() + assert type(attrs_copy) is dict + assert attrs_copy["database"] == "mutated_db" + + # Other observable dict operations trigger invalidation + metrics_tracer.client_attributes.update({"instance_id": "updated_via_update"}) + assert ( + metrics_tracer._create_attempt_otel_attributes()["instance_id"] + == "updated_via_update" + ) + # setdefault when key is new + metrics_tracer.client_attributes.setdefault("new_key", "default_val") + assert metrics_tracer._create_attempt_otel_attributes()["new_key"] == "default_val" + # setdefault when key already exists + assert ( + metrics_tracer.client_attributes.setdefault("new_key", "other_val") + == "default_val" + ) + # pop + metrics_tracer.client_attributes.pop("new_key") + assert "new_key" not in metrics_tracer._create_attempt_otel_attributes() + + # delitem + metrics_tracer.client_attributes["to_delete"] = "val" + del metrics_tracer.client_attributes["to_delete"] + assert "to_delete" not in metrics_tracer._create_attempt_otel_attributes() + + # popitem + metrics_tracer.client_attributes["to_pop"] = "val" + metrics_tracer.client_attributes.popitem() + + # ObservableDict with no on_change callback + from google.cloud.spanner_v1.metrics.metrics_tracer import _ObservableDict + + no_callback_dict = _ObservableDict({"a": 1}) + no_callback_dict["b"] = 2 + del no_callback_dict["a"] + no_callback_dict.update({"c": 3}) + no_callback_dict.setdefault("d", 4) + no_callback_dict.pop("c") + no_callback_dict.popitem() + no_callback_dict.clear() + + # clear on client_attributes + metrics_tracer.client_attributes.clear() + assert metrics_tracer._create_attempt_otel_attributes() == {"status": "UNAVAILABLE"} + + # Ensure client_attributes and cached attributes are dict instances for backward compatibility + assert isinstance(metrics_tracer.client_attributes, dict) + assert isinstance(first_attempt_attrs, dict) + assert isinstance(first_op_attrs, dict) diff --git a/packages/google-cloud-spanner/tests/unit/test_session.py b/packages/google-cloud-spanner/tests/unit/test_session.py index 17704cb59b2f..05c79c632de7 100644 --- a/packages/google-cloud-spanner/tests/unit/test_session.py +++ b/packages/google-cloud-spanner/tests/unit/test_session.py @@ -2682,3 +2682,113 @@ def _time_func(): _delay_until_retry(exc_mock, 6, 1) sleep_mock.assert_not_called() + + def test_run_in_transaction_tracing_events_attached_to_parent_span(self): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + tracer = tracer_provider.get_tracer("test") + + transaction_pb = TransactionPB(id=TRANSACTION_ID) + now = datetime.datetime.now(timezone.utc).replace(tzinfo=UTC) + now_pb = _datetime_to_pb_timestamp(now) + aborted = _make_rpc_error(Aborted, trailing_metadata=[]) + response = CommitResponse(commit_timestamp=now_pb) + gax_api = self._make_spanner_api() + gax_api.begin_transaction.return_value = transaction_pb + gax_api.commit.side_effect = [aborted, response] + database = self._make_database() + database.spanner_api = gax_api + session = self._make_one(database) + session._session_id = self.SESSION_ID + + def unit_of_work(transaction, *args, **kwargs): + transaction.insert(TABLE_NAME, COLUMNS, VALUES) + return "answer" + + with tracer.start_as_current_span("ParentSpan"): + return_value = session.run_in_transaction( + unit_of_work, + transaction_tag="test-tag", + default_retry_delay=0, + ) + + self.assertEqual(return_value, "answer") + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + parent_span = finished_spans[0] + self.assertEqual(parent_span.name, "ParentSpan") + self.assertIsNone(parent_span.attributes.get("transaction.tag")) + event_names = [event.name for event in parent_span.events] + self.assertIn( + "Transaction was aborted during commit, retrying", + event_names, + ) + retry_event = next( + event + for event in parent_span.events + if event.name == "Transaction was aborted during commit, retrying" + ) + self.assertEqual(retry_event.attributes.get("attempt"), 1) + + def test_run_in_transaction_tracing_events_aborted_without_inner_errors( + self, + ): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.sdk.trace.sampling import ALWAYS_ON + + tracer_provider = TracerProvider(sampler=ALWAYS_ON) + trace_exporter = InMemorySpanExporter() + tracer_provider.add_span_processor(SimpleSpanProcessor(trace_exporter)) + tracer = tracer_provider.get_tracer("test") + + transaction_pb = TransactionPB(id=TRANSACTION_ID) + now = datetime.datetime.now(timezone.utc).replace(tzinfo=UTC) + now_pb = _datetime_to_pb_timestamp(now) + # Aborted raised without inner errors (defaults to empty tuple in GoogleAPICallError) + aborted = Aborted("aborted without inner errors") + response = CommitResponse(commit_timestamp=now_pb) + gax_api = self._make_spanner_api() + gax_api.begin_transaction.return_value = transaction_pb + gax_api.commit.side_effect = [aborted, response] + database = self._make_database() + database.spanner_api = gax_api + session = self._make_one(database) + session._session_id = self.SESSION_ID + + def unit_of_work(transaction, *args, **kwargs): + transaction.insert(TABLE_NAME, COLUMNS, VALUES) + return "answer" + + with tracer.start_as_current_span("ParentSpan"): + return_value = session.run_in_transaction( + unit_of_work, + default_retry_delay=0, + ) + + self.assertEqual(return_value, "answer") + finished_spans = trace_exporter.get_finished_spans() + self.assertEqual(len(finished_spans), 1) + parent_span = finished_spans[0] + event_names = [event.name for event in parent_span.events] + self.assertIn( + "Transaction was aborted during commit, retrying", + event_names, + ) + retry_event = next( + event + for event in parent_span.events + if event.name == "Transaction was aborted during commit, retrying" + ) + self.assertEqual(retry_event.attributes.get("attempt"), 1) diff --git a/packages/google-cloud-spanner/tests/unit/test_snapshot.py b/packages/google-cloud-spanner/tests/unit/test_snapshot.py index a5082e5b8aa1..a4cf27dc19d9 100644 --- a/packages/google-cloud-spanner/tests/unit/test_snapshot.py +++ b/packages/google-cloud-spanner/tests/unit/test_snapshot.py @@ -159,14 +159,15 @@ def _call_fut( request_id_manager=None if not session else session._database, ) - def _make_item(self, value, resume_token=b"", metadata=None): + def _make_item(self, value, resume_token=b"", metadata=None, last=False): return mock.Mock( value=value, resume_token=resume_token, metadata=metadata, precommit_token=None, + last=last, _pb=None, - spec=["value", "resume_token", "metadata", "precommit_token"], + spec=["value", "resume_token", "metadata", "precommit_token", "last"], ) def test_iteration_w_empty_raw(self): @@ -212,6 +213,226 @@ def test_iteration_w_non_empty_raw(self): ) self.assertNoSpans() + def test_restart_on_unavailable_last(self): + item_last = self._make_item(0, last=True) + trailing_item = self._make_item(1) + + raw = _MockIterator(item_last, trailing_item) + restart = mock.Mock(return_value=raw) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + with mock.patch("google.cloud.spanner_v1.snapshot._drain_stream") as mock_drain: + resumable = self._call_fut(derived, restart, request, session=session) + items = list(resumable) + + mock_drain.assert_called_once_with(raw) + self.assertEqual(len(items), 1) + self.assertEqual(items[0], item_last) + + def test_restart_on_unavailable_finally_cancels_on_early_termination(self): + item = self._make_item(0, last=False) + raw = _MockIterator(item) + raw.cancel = mock.Mock() + restart = mock.Mock(return_value=raw) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + resumable = self._call_fut(derived, restart, request, session=session) + for received_item in resumable: + break + resumable.close() + + raw.cancel.assert_called_once() + + def test_restart_on_unavailable_item_without_last_attribute(self): + item = mock.Mock( + value=0, + resume_token=b"", + metadata=None, + precommit_token=None, + _pb=None, + spec=["value", "resume_token", "metadata", "precommit_token"], + ) + raw = _MockIterator(item) + restart = mock.Mock(return_value=raw) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + resumable = self._call_fut(derived, restart, request, session=session) + items = list(resumable) + self.assertEqual(items, [item]) + + def test_restart_on_unavailable_finally_handles_cancel_exception(self): + item = self._make_item(0, last=False) + raw = _MockIterator(item) + raw.cancel = mock.Mock(side_effect=RuntimeError("cancel failed")) + restart = mock.Mock(return_value=raw) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + resumable = self._call_fut(derived, restart, request, session=session) + for _ in resumable: + break + resumable.close() + raw.cancel.assert_called_once() + + def test_restart_on_unavailable_last_does_not_cancel_iterator_in_finally(self): + item_last = self._make_item(0, last=True) + raw = _MockIterator(item_last) + raw.cancel = mock.Mock() + restart = mock.Mock(return_value=raw) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + with mock.patch("google.cloud.spanner_v1.snapshot._drain_stream"): + resumable = self._call_fut(derived, restart, request, session=session) + list(resumable) + + raw.cancel.assert_not_called() + + def test_streamed_result_set_with_last(self): + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1.streamed import StreamedResultSet + from google.cloud.spanner_v1.types.result_set import ( + PartialResultSet, + ResultSetMetadata, + ) + from google.cloud.spanner_v1.types.type import StructType, Type, TypeCode + + fields = [StructType.Field(name="greeting", type_=Type(code=TypeCode.STRING))] + metadata_pb = ResultSetMetadata(row_type=StructType(fields=fields)) + item = PartialResultSet(metadata=metadata_pb, last=True) + item.values.append(Value(string_value="hello")) + + raw = _MockIterator(item) + streamed_result_set = StreamedResultSet(raw) + rows = list(streamed_result_set) + + self.assertEqual(rows, [["hello"]]) + self.assertTrue(streamed_result_set._done) + + def test_restart_on_unavailable_multi_chunk_with_last(self): + from google.protobuf.struct_pb2 import Value + + from google.cloud.spanner_v1.streamed import StreamedResultSet + from google.cloud.spanner_v1.types.result_set import ( + PartialResultSet, + ResultSetMetadata, + ResultSetStats, + ) + from google.cloud.spanner_v1.types.type import StructType, Type, TypeCode + + fields = [StructType.Field(name="greeting", type_=Type(code=TypeCode.STRING))] + metadata_pb = ResultSetMetadata(row_type=StructType(fields=fields)) + stats_pb = ResultSetStats(row_count_exact=2) + + chunk_one = PartialResultSet( + metadata=metadata_pb, last=False, resume_token=b"token_1" + ) + chunk_one.values.append(Value(string_value="hello")) + + chunk_two = PartialResultSet(last=True, stats=stats_pb) + chunk_two.values.append(Value(string_value="world")) + + raw = _MockIterator(chunk_one, chunk_two) + restart = mock.Mock(return_value=raw) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + with mock.patch("google.cloud.spanner_v1.snapshot._drain_stream") as mock_drain: + resumable = self._call_fut(derived, restart, request, session=session) + streamed_result_set = StreamedResultSet(resumable) + rows = list(streamed_result_set) + + mock_drain.assert_called_once_with(raw) + self.assertEqual(rows, [["hello"], ["world"]]) + self.assertEqual(streamed_result_set.metadata, metadata_pb) + self.assertEqual(streamed_result_set.stats, stats_pb) + self.assertTrue(streamed_result_set._done) + + def test_restart_on_unavailable_zero_rows_with_last(self): + from google.cloud.spanner_v1.streamed import StreamedResultSet + from google.cloud.spanner_v1.types.result_set import ( + PartialResultSet, + ResultSetMetadata, + ResultSetStats, + ) + from google.cloud.spanner_v1.types.type import StructType, Type, TypeCode + + fields = [StructType.Field(name="greeting", type_=Type(code=TypeCode.STRING))] + metadata_pb = ResultSetMetadata(row_type=StructType(fields=fields)) + stats_pb = ResultSetStats(row_count_exact=0) + + chunk = PartialResultSet(metadata=metadata_pb, last=True, stats=stats_pb) + + raw = _MockIterator(chunk) + restart = mock.Mock(return_value=raw) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + with mock.patch("google.cloud.spanner_v1.snapshot._drain_stream") as mock_drain: + resumable = self._call_fut(derived, restart, request, session=session) + streamed_result_set = StreamedResultSet(resumable) + rows = list(streamed_result_set) + + mock_drain.assert_called_once_with(raw) + self.assertEqual(rows, []) + self.assertEqual(streamed_result_set.metadata, metadata_pb) + self.assertEqual(streamed_result_set.stats, stats_pb) + self.assertTrue(streamed_result_set._done) + + def test_restart_on_unavailable_retry_before_last(self): + from google.api_core.exceptions import ServiceUnavailable + + chunk_one = self._make_item(0, resume_token=RESUME_TOKEN, last=False) + chunk_two = self._make_item(1, last=True) + + stream_one = _MockIterator( + chunk_one, fail_after=True, error=ServiceUnavailable("transient") + ) + stream_two = _MockIterator(chunk_two) + stream_two.cancel = mock.Mock() + + restart = mock.Mock(side_effect=[stream_one, stream_two]) + request = mock.Mock(test="test", spec=["test", "resume_token"]) + database = _Database() + database.spanner_api = build_spanner_api() + session = _Session(database) + derived = _build_snapshot_derived(session) + + with mock.patch("google.cloud.spanner_v1.snapshot._drain_stream") as mock_drain: + resumable = self._call_fut(derived, restart, request, session=session) + items = list(resumable) + + self.assertEqual(items, [chunk_one, chunk_two]) + self.assertEqual(len(restart.mock_calls), 2) + self.assertEqual(request.resume_token, RESUME_TOKEN) + mock_drain.assert_called_once_with(stream_two) + stream_two.cancel.assert_not_called() + def test_iteration_w_raw_w_resume_token(self): ITEMS = ( self._make_item(0), diff --git a/packages/google-cloud-spanner/tests/unit/test_streamed.py b/packages/google-cloud-spanner/tests/unit/test_streamed.py index 3d3ba709145d..4c8ff04d3193 100644 --- a/packages/google-cloud-spanner/tests/unit/test_streamed.py +++ b/packages/google-cloud-spanner/tests/unit/test_streamed.py @@ -1043,7 +1043,7 @@ def test___iter___large_batch(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - expected_rows = [[i, f"name_{i}"] for i in range(500)] + expected_rows = [[index, f"name_{index}"] for index in range(500)] values = [self._make_value(cell) for row in expected_rows for cell in row] result_set = self._make_partial_result_set(values, metadata=metadata) @@ -1061,7 +1061,7 @@ def test___iter___stepwise_consumption(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - expected_rows = [[i, f"name_{i}"] for i in range(20)] + expected_rows = [[index, f"name_{index}"] for index in range(20)] values = [self._make_value(cell) for row in expected_rows for cell in row] result_set = self._make_partial_result_set(values, metadata=metadata) @@ -1081,8 +1081,8 @@ def test___iter___stepwise_across_chunks(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - chunk1_rows = [[i, f"name_{i}"] for i in range(10)] - chunk2_rows = [[i, f"name_{i}"] for i in range(10, 20)] + chunk1_rows = [[index, f"name_{index}"] for index in range(10)] + chunk2_rows = [[index, f"name_{index}"] for index in range(10, 20)] values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] values2 = [self._make_value(cell) for row in chunk2_rows for cell in row] @@ -1106,7 +1106,7 @@ def test___iter___early_break(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - expected_rows = [[i, f"name_{i}"] for i in range(10)] + expected_rows = [[index, f"name_{index}"] for index in range(10)] values = [self._make_value(cell) for row in expected_rows for cell in row] result_set = self._make_partial_result_set(values, metadata=metadata) @@ -1128,7 +1128,7 @@ def test___iter___mid_stream_error(self): self._make_scalar_field("name", TypeCode.STRING), ] metadata = self._make_result_set_metadata(fields) - chunk1_rows = [[i, f"name_{i}"] for i in range(5)] + chunk1_rows = [[index, f"name_{index}"] for index in range(5)] values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] result_set1 = self._make_partial_result_set(values1, metadata=metadata) @@ -1145,6 +1145,309 @@ def mock_iterator(): self.assertEqual(consumed, chunk1_rows) self.assertIn("Stream error midway", str(context.exception)) + def test_decode_rows_direct_width_one(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("count", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[42]] + values = [self._make_value(42)] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual(found, expected_rows) + self.assertTrue(streamed._done) + + def test_decode_rows_direct_multi_column(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + self._make_scalar_field("active", TypeCode.BOOL), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[1, "alpha", True], [2, "beta", False]] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual(found, expected_rows) + self.assertTrue(streamed._done) + + def test_decode_rows_direct_with_null_values(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values = [ + self._make_value(1), + Value(null_value=NULL_VALUE), + self._make_value(2), + self._make_value("Alice"), + ] + expected_rows = [[1, None], [2, "Alice"]] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual(found, expected_rows) + + def test_decode_rows_direct_lazy_decode_width_one(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("id", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + raw_value = self._make_value(100) + values = [raw_value] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator, lazy_decode=True) + found = list(streamed) + self.assertEqual(len(found), 1) + self.assertEqual(found[0], [raw_value]) + self.assertEqual(streamed.decode_row(found[0]), [100]) + + def test_decode_rows_direct_lazy_decode_multi_column(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + value_id = self._make_value(1) + value_name = self._make_value("test") + values = [value_id, value_name] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator, lazy_decode=True) + found = list(streamed) + self.assertEqual(len(found), 1) + self.assertEqual(found[0], [value_id, value_name]) + self.assertEqual(streamed.decode_row(found[0]), [1, "test"]) + + def test_decode_rows_direct_trailing_partial_row(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values1 = [ + self._make_value(1), + self._make_value("a"), + self._make_value(2), + ] + values2 = [ + self._make_value("b"), + self._make_value(3), + self._make_value("c"), + ] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2, last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual( + found, + [ + [1, "a"], + [2, "b"], + [3, "c"], + ], + ) + + def test_decode_rows_direct_trailing_partial_row_lazy_decode(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + value1 = self._make_value(1) + value_a = self._make_value("a") + value2 = self._make_value(2) + value_b = self._make_value("b") + + result_set1 = self._make_partial_result_set( + [value1, value_a, value2], metadata=metadata + ) + result_set2 = self._make_partial_result_set([value_b], last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator, lazy_decode=True) + found = list(streamed) + self.assertEqual(len(found), 2) + self.assertEqual(found[0], [value1, value_a]) + self.assertEqual(found[1], [value2, value_b]) + self.assertEqual(streamed.decode_row(found[0]), [1, "a"]) + self.assertEqual(streamed.decode_row(found[1]), [2, "b"]) + + def test_decode_rows_direct_trailing_partial_row_with_null(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values1 = [ + self._make_value(1), + self._make_value("a"), + Value(null_value=NULL_VALUE), + ] + values2 = [ + self._make_value("b"), + ] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2, last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual( + found, + [ + [1, "a"], + [None, "b"], + ], + ) + + def test_decode_rows_direct_prefix_partial_row_with_null(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + values1 = [ + self._make_value(1), + self._make_value("a"), + self._make_value(2), + ] + values2 = [ + Value(null_value=NULL_VALUE), + self._make_value(3), + self._make_value("c"), + ] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2, last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual( + found, + [ + [1, "a"], + [2, None], + [3, "c"], + ], + ) + + def test_decode_rows_direct_width_one_with_null(self): + from google.protobuf.struct_pb2 import NULL_VALUE, Value + + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("id", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + values = [ + self._make_value(1), + Value(null_value=NULL_VALUE), + self._make_value(2), + ] + + result_set = self._make_partial_result_set(values, metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual(found, [[1], [None], [2]]) + + def test_decode_rows_direct_three_chunk_split_row(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + self._make_scalar_field("active", TypeCode.BOOL), + ] + metadata = self._make_result_set_metadata(fields) + prs1 = self._make_partial_result_set([self._make_value(1)], metadata=metadata) + prs2 = self._make_partial_result_set([self._make_value("alpha")]) + prs3 = self._make_partial_result_set([self._make_value(True)], last=True) + + iterator = _MockCancellableIterator(prs1, prs2, prs3) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual(found, [[1, "alpha", True]]) + + def test_decode_rows_direct_three_chunk_split_row_lazy_decode(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + self._make_scalar_field("active", TypeCode.BOOL), + ] + metadata = self._make_result_set_metadata(fields) + val1 = self._make_value(1) + val2 = self._make_value("alpha") + val3 = self._make_value(True) + prs1 = self._make_partial_result_set([val1], metadata=metadata) + prs2 = self._make_partial_result_set([val2]) + prs3 = self._make_partial_result_set([val3], last=True) + + iterator = _MockCancellableIterator(prs1, prs2, prs3) + streamed = self._make_one(iterator, lazy_decode=True) + found = list(streamed) + self.assertEqual(found, [[val1, val2, val3]]) + self.assertEqual(streamed.decode_row(found[0]), [1, "alpha", True]) + + def test_decode_rows_direct_empty_values(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [self._make_scalar_field("id", TypeCode.INT64)] + metadata = self._make_result_set_metadata(fields) + + result_set1 = self._make_partial_result_set([], metadata=metadata) + result_set2 = self._make_partial_result_set([self._make_value(1)], last=True) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual(found, [[1]]) + + def test_merge_values_zero_fields(self): + from google.cloud.spanner_v1 import ResultSetMetadata, StructType + + metadata = ResultSetMetadata(row_type=StructType(fields=[])) + result_set = self._make_partial_result_set([], metadata=metadata, last=True) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + list(streamed) + streamed._merge_values([self._make_value(1)]) + self.assertEqual(streamed._rows, []) + class _MockCancellableIterator(object): cancel_calls = 0