From c9570b505b20e20a2e3ed1c18396b7ea18e9d21c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Knut=20Olav=20L=C3=B8ite?= Date: Thu, 17 Sep 2026 09:18:45 +0200 Subject: [PATCH] perf(spanner): optimize request ID header generation and retry closures Optimize request ID generation and RPC retry hot paths across Database, Transaction, Snapshot, and Batch: - Use `int.from_bytes(os.urandom(8), "big")` for process-level random ID generation and pre-format the static process prefix once at import time. - Cache per-database request ID prefixes directly in `instance.__dict__` via a non-data descriptor (`_CachedPrefixDescriptor`), lazily computing on first access and invalidating when `_channel_id` changes. - Inline `Database.with_error_augmentation` and construct gRPC metadata lists with unpacking (`[*prior_metadata, ...]`) instead of intermediate list copies and `.append()`. - Replace thread-safe `AtomicCounter` instances with local integer counters (`attempt = 0`) inside sequential retry closures, and pass active OpenTelemetry spans directly to avoid repeated `contextvars` lookups. - Eliminate transient `functools.partial` allocations and redundant `*args, **kwargs` unpacking inside retry closures. - Make `wrap_with_request_id` idempotent on repeated wrapping and replace stale request IDs when retrying failed calls. --- .../google/cloud/spanner_v1/_async/batch.py | 22 +- .../cloud/spanner_v1/_async/database.py | 89 +++++--- .../cloud/spanner_v1/_async/snapshot.py | 37 ++-- .../cloud/spanner_v1/_async/transaction.py | 92 ++++----- .../google/cloud/spanner_v1/_helpers.py | 62 +++--- .../google/cloud/spanner_v1/batch.py | 17 +- .../google/cloud/spanner_v1/database.py | 84 +++++--- .../google/cloud/spanner_v1/exceptions.py | 13 +- .../cloud/spanner_v1/request_id_header.py | 59 +++--- .../google/cloud/spanner_v1/snapshot.py | 36 ++-- .../google/cloud/spanner_v1/transaction.py | 91 ++++---- .../tests/unit/_async/test_database.py | 108 ++++++++++ .../tests/unit/_async/test_transaction.py | 73 ++++++- .../tests/unit/test__helpers.py | 194 ++++++++++++++++++ .../tests/unit/test_database.py | 107 ++++++++++ .../tests/unit/test_exceptions.py | 34 +++ .../tests/unit/test_transaction.py | 62 +++++- 17 files changed, 900 insertions(+), 280 deletions(-) diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/batch.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/batch.py index 966519265d50..92de7adaead0 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/batch.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/batch.py @@ -15,7 +15,6 @@ """Context manager for Cloud Spanner batched writes.""" __CROSS_SYNC_OUTPUT__ = "google.cloud.spanner_v1.batch" -import functools import time from typing import List, Optional @@ -25,7 +24,6 @@ from google.cloud.aio._cross_sync import CrossSync from google.cloud.spanner_v1._async._helpers import _retry, _retry_on_aborted_exception from google.cloud.spanner_v1._helpers import ( - AtomicCounter, _check_rst_stream_error, _make_list_value_pb, _make_list_value_pbs, @@ -341,13 +339,11 @@ async def wrapped_method(): metadata, span, ) - commit_method = functools.partial( - api.commit, - request=commit_request, - metadata=call_metadata, - ) with error_augmenter: - return await commit_method() + return await api.commit( + request=commit_request, + metadata=call_metadata, + ) response = await _retry_on_aborted_exception( wrapped_method, @@ -478,27 +474,27 @@ async def batch_write( ) as span, MetricsCapture(self._resource_info), ): - attempt = AtomicCounter(0) + attempt = 0 nth_request = getattr(database, "_next_nth_request", 0) def wrapped_method(): + nonlocal attempt + attempt += 1 batch_write_request = BatchWriteRequest( session=session.name, mutation_groups=mutation_groups, request_options=request_options, exclude_txn_from_change_streams=exclude_txn_from_change_streams, ) - batch_write_method = functools.partial( - api.batch_write, + return api.batch_write( request=batch_write_request, metadata=database.metadata_with_request_id( nth_request, - attempt.increment(), + attempt, metadata, span, ), ) - return batch_write_method() response = await _retry( wrapped_method, diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py index c12769172c23..3be33183dc6a 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py @@ -58,11 +58,14 @@ _merge_query_options, _metadata_with_leader_aware_routing, _metadata_with_prefix, - _metadata_with_request_id, - _metadata_with_request_id_and_req_id, ) from google.cloud.spanner_v1.keyset import KeySet from google.cloud.spanner_v1.merged_result_set import MergedResultSet +from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, + _CachedPrefixDescriptor, +) from google.cloud.spanner_v1.services.spanner.async_client import ( SpannerAsyncClient as SpannerClient, ) @@ -180,6 +183,8 @@ class Database(object): __transport_lock = threading.Lock() __transports_to_channel_id = dict() + _channel_id_val = 0 + _req_id_prefix = _CachedPrefixDescriptor() def __init__( self, @@ -550,23 +555,17 @@ def spanner_api(self): return self._spanner_api - def metadata_with_request_id( - self, nth_request, nth_attempt, prior_metadata=[], span=None - ): - if span is None: - span = get_current_span() + @property + def _channel_id(self): + return self._channel_id_val - return _metadata_with_request_id( - self._nth_client_id, - self._channel_id, - nth_request, - nth_attempt, - prior_metadata, - span, - ) + @_channel_id.setter + def _channel_id(self, value): + self._channel_id_val = value + self.__dict__.pop("_req_id_prefix", None) def metadata_and_request_id( - self, nth_request, nth_attempt, prior_metadata=[], span=None + self, nth_request, nth_attempt, prior_metadata=None, span=None ): """Return metadata and request ID string. @@ -585,17 +584,38 @@ def metadata_and_request_id( if span is None: span = get_current_span() - return _metadata_with_request_id_and_req_id( - self._nth_client_id, - self._channel_id, - nth_request, - nth_attempt, - prior_metadata, - span, + req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}" + metadata = ( + [*prior_metadata, (REQ_ID_HEADER_KEY, req_id)] + if prior_metadata + else [(REQ_ID_HEADER_KEY, req_id)] ) + if span is not None and span.is_recording(): + span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) + + return metadata, req_id + + def metadata_with_request_id( + self, nth_request, nth_attempt, prior_metadata=None, span=None + ): + if span is None: + span = get_current_span() + + req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}" + metadata = ( + [*prior_metadata, (REQ_ID_HEADER_KEY, req_id)] + if prior_metadata + else [(REQ_ID_HEADER_KEY, req_id)] + ) + + if span is not None and span.is_recording(): + span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) + + return metadata + def with_error_augmentation( - self, nth_request, nth_attempt, prior_metadata=[], span=None + self, nth_request, nth_attempt, prior_metadata=None, span=None ): """Context manager for gRPC calls with error augmentation. @@ -614,16 +634,17 @@ def with_error_augmentation( if span is None: span = get_current_span() - metadata, request_id = _metadata_with_request_id_and_req_id( - self._nth_client_id, - self._channel_id, - nth_request, - nth_attempt, - prior_metadata, - span, + req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}" + metadata = ( + [*prior_metadata, (REQ_ID_HEADER_KEY, req_id)] + if prior_metadata + else [(REQ_ID_HEADER_KEY, req_id)] ) - return metadata, _augment_errors_with_request_id(request_id) + if span is not None and span.is_recording(): + span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) + + return metadata, _augment_errors_with_request_id(req_id) def __eq__(self, other): if not isinstance(other, self.__class__): @@ -975,13 +996,13 @@ async def execute_pdml(): @property def _next_nth_request(self): if self._instance and self._instance._client: - return self._instance._client._next_nth_request + return getattr(self._instance._client, "_next_nth_request", 1) return 1 @property def _nth_client_id(self): if self._instance and self._instance._client: - return self._instance._client._nth_client_id + return getattr(self._instance._client, "_nth_client_id", 0) return 0 def session(self, labels=None, database_role=None): 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 5e64e106db99..fd0e44889709 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 @@ -31,7 +31,6 @@ from google.cloud.spanner_v1._async._helpers import _retry from google.cloud.spanner_v1._async.streamed import StreamedResultSet from google.cloud.spanner_v1._helpers import ( - AtomicCounter, _augment_error_with_request_id, _check_rst_stream_error, _make_value_pb, @@ -760,20 +759,20 @@ async def partition_read( MetricsCapture(self._resource_info), ): nth_request = getattr(database, "_next_nth_request", 0) - attempt = AtomicCounter() + attempt = 0 async def attempt_tracking_method(): + nonlocal attempt + attempt += 1 all_metadata = database.metadata_with_request_id( - nth_request, attempt.increment(), metadata, span + nth_request, attempt, metadata, span ) - partition_read_method = functools.partial( - api.partition_read, + return await api.partition_read( request=partition_read_request, metadata=all_metadata, retry=retry, timeout=timeout, ) - return await partition_read_method() response = await _retry( attempt_tracking_method, @@ -843,20 +842,20 @@ async def partition_query( MetricsCapture(self._resource_info), ): nth_request = getattr(database, "_next_nth_request", 0) - attempt = AtomicCounter() + attempt = 0 async def attempt_tracking_method(): + nonlocal attempt + attempt += 1 all_metadata = database.metadata_with_request_id( - nth_request, attempt.increment(), metadata, span + nth_request, attempt, metadata, span ) - partition_query_method = functools.partial( - api.partition_query, + return await api.partition_query( request=partition_query_request, metadata=all_metadata, retry=retry, timeout=timeout, ) - return await partition_query_method() response = await _retry( attempt_tracking_method, @@ -917,22 +916,22 @@ async def _begin_transaction( MetricsCapture(self._resource_info), ): nth_request = getattr(database, "_next_nth_request", 0) - attempt = AtomicCounter() + attempt = 0 async def wrapped_method(): + nonlocal attempt + attempt += 1 begin_transaction_request = BeginTransactionRequest( **begin_request_kwargs ) call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.increment(), metadata, span - ) - begin_transaction_method = functools.partial( - api.begin_transaction, - request=begin_transaction_request, - metadata=call_metadata, + nth_request, attempt, metadata, span ) with error_augmenter: - return await begin_transaction_method() + return await api.begin_transaction( + request=begin_transaction_request, + metadata=call_metadata, + ) async def before_next_retry(nth_retry, delay_in_seconds): add_span_event( diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/transaction.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/transaction.py index a551c6be2bfd..d8b79d8948a8 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/transaction.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/transaction.py @@ -29,7 +29,6 @@ from google.cloud.spanner_v1._async.batch import _BatchBase from google.cloud.spanner_v1._async.snapshot import _SnapshotBase from google.cloud.spanner_v1._helpers import ( - AtomicCounter, _check_rst_stream_error, _make_value_pb, _merge_client_context, @@ -154,10 +153,10 @@ async def _execute_request( session._database, "observability_options", None ), metadata=metadata, - ), + ) as span, MetricsCapture(self._resource_info), ): - method = functools.partial(method, request=request) + method = functools.partial(method, request=request, span=span) response = await _retry( method, allowed_exceptions={InternalServerError: _check_rst_stream_error}, @@ -200,25 +199,24 @@ async def rollback(self) -> None: ) as span, MetricsCapture(self._resource_info), ): - attempt = AtomicCounter(0) + attempt = 0 nth_request = database._next_nth_request - def wrapped_method(*args, **kwargs): - attempt.increment() + async def wrapped_method(): + nonlocal attempt + attempt += 1 call_metadata, error_augmenter = database.with_error_augmentation( nth_request, - attempt.value, + attempt, metadata, span, ) - rollback_method = functools.partial( - api.rollback, - session=session.name, - transaction_id=self._transaction_id, - metadata=call_metadata, - ) with error_augmenter: - return rollback_method(*args, **kwargs) + return await api.rollback( + session=session.name, + transaction_id=self._transaction_id, + metadata=call_metadata, + ) await _retry( wrapped_method, @@ -234,7 +232,6 @@ async def _reset_and_begin(self): self._execute_sql_request_count = 0 await self.begin() - @CrossSync.convert @CrossSync.convert async def commit( self, return_commit_stats=False, request_options=None, max_commit_delay=None @@ -324,11 +321,12 @@ async def commit( add_span_event(span, "Starting Commit") - attempt = AtomicCounter(0) + attempt = 0 nth_request = database._next_nth_request - async def wrapped_method(*args, **kwargs): - attempt.increment() + async def wrapped_method(): + nonlocal attempt + attempt += 1 commit_request_args = { "mutations": mutations, **common_commit_request_args, @@ -340,17 +338,15 @@ async def wrapped_method(*args, **kwargs): call_metadata, error_augmenter = database.with_error_augmentation( nth_request, - attempt.value, + attempt, metadata, span, ) - commit_method = functools.partial( - api.commit, - request=CommitRequest(**commit_request_args), - metadata=call_metadata, - ) with error_augmenter: - return await commit_method(*args, **kwargs) + return await api.commit( + request=CommitRequest(**commit_request_args), + metadata=call_metadata, + ) commit_retry_event_name = "Transaction Commit Attempt Failed. Retrying" @@ -557,22 +553,21 @@ async def execute_update( ) nth_request = database._next_nth_request - attempt = AtomicCounter(0) + attempt = 0 - async def wrapped_method(*args, **kwargs): - attempt.increment() + async def wrapped_method(request=execute_sql_request, span=None): + nonlocal attempt + attempt += 1 call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.value, metadata - ) - execute_sql_method = functools.partial( - api.execute_sql, - request=execute_sql_request, - metadata=call_metadata, - retry=retry, - timeout=timeout, + nth_request, attempt, metadata, span ) with error_augmenter: - return await execute_sql_method(*args, **kwargs) + return await api.execute_sql( + request=request, + metadata=call_metadata, + retry=retry, + timeout=timeout, + ) result_set_pb: ResultSet = await self._execute_request( wrapped_method, @@ -713,22 +708,21 @@ async def batch_update( ) nth_request = database._next_nth_request - attempt = AtomicCounter(0) + attempt = 0 - async def wrapped_method(*args, **kwargs): - attempt.increment() + async def wrapped_method(request=execute_batch_dml_request, span=None): + nonlocal attempt + attempt += 1 call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.value, metadata - ) - execute_batch_dml_method = functools.partial( - api.execute_batch_dml, - request=execute_batch_dml_request, - metadata=call_metadata, - retry=retry, - timeout=timeout, + nth_request, attempt, metadata, span ) with error_augmenter: - return await execute_batch_dml_method(*args, **kwargs) + return await api.execute_batch_dml( + request=request, + metadata=call_metadata, + retry=retry, + timeout=timeout, + ) response_pb: ExecuteBatchDmlResponse = await self._execute_request( wrapped_method, 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 a9740883afa4..0eb66af46fdb 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_helpers.py @@ -24,7 +24,6 @@ import threading import time import uuid -from contextlib import contextmanager from google.api_core import datetime_helpers from google.api_core.exceptions import Aborted @@ -898,6 +897,8 @@ def _get_retry_delay(cause, attempts, default_retry_delay=None): class AtomicCounter: + __slots__ = ("__lock", "__value") + def __init__(self, start_value=0): self.__lock = threading.Lock() self.__value = start_value @@ -939,35 +940,51 @@ def reset(self): self.__value = 0 -def _metadata_with_request_id(*args, **kwargs): +def _metadata_with_request_id( + client_id, channel_id, nth_request, attempt, other_metadata=None, span=None +): """Return metadata with request ID header. This function returns only the metadata list (not a tuple), maintaining backward compatibility with existing code. Args: - *args: Arguments to pass to with_request_id - **kwargs: Keyword arguments to pass to with_request_id + client_id: The client identifier. + channel_id: The channel identifier. + nth_request: The request sequence number. + attempt: The attempt sequence number. + other_metadata: Prior metadata list. + span: Optional trace span. Returns: list: gRPC metadata with request ID header """ - return with_request_id_metadata_only(*args, **kwargs) + return with_request_id_metadata_only( + client_id, channel_id, nth_request, attempt, other_metadata, span + ) -def _metadata_with_request_id_and_req_id(*args, **kwargs): +def _metadata_with_request_id_and_req_id( + client_id, channel_id, nth_request, attempt, other_metadata=None, span=None +): """Return both metadata and request ID string. This is used when we need to augment errors with the request ID. Args: - *args: Arguments to pass to with_request_id - **kwargs: Keyword arguments to pass to with_request_id + client_id: The client identifier. + channel_id: The channel identifier. + nth_request: The request sequence number. + attempt: The attempt sequence number. + other_metadata: Prior metadata list. + span: Optional trace span. Returns: tuple: (metadata, request_id) """ - return with_request_id(*args, **kwargs) + return with_request_id( + client_id, channel_id, nth_request, attempt, other_metadata, span + ) def _augment_error_with_request_id(error, request_id=None): @@ -983,22 +1000,21 @@ def _augment_error_with_request_id(error, request_id=None): return wrap_with_request_id(error, request_id) -@contextmanager -def _augment_errors_with_request_id(request_id): - """Context manager to augment exceptions with request ID. +class _augment_errors_with_request_id: + """Context manager to augment exceptions with request ID.""" - Args: - request_id (str): The request ID to include in exceptions + __slots__ = ("_request_id",) - Yields: - None - """ - try: - yield - except Exception as exc: - augmented = _augment_error_with_request_id(exc, request_id) - # Use exception chaining to preserve the original exception - raise augmented from exc + def __init__(self, request_id): + self._request_id = request_id + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + if exc_val is not None: + _augment_error_with_request_id(exc_val, self._request_id) + return False def _merge_Transaction_Options( diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/batch.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/batch.py index c0ff1cc1a613..974afab9f0d4 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/batch.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/batch.py @@ -17,7 +17,6 @@ """Context manager for Cloud Spanner batched writes.""" -import functools import time from typing import List, Optional @@ -25,7 +24,6 @@ from google.cloud._helpers import _datetime_to_pb_timestamp from google.cloud.spanner_v1._helpers import ( - AtomicCounter, _check_rst_stream_error, _make_list_value_pb, _make_list_value_pbs, @@ -294,11 +292,8 @@ def wrapped_method(): call_metadata, error_augmenter = database.with_error_augmentation( getattr(database, "_next_nth_request", 0), 1, metadata, span ) - commit_method = functools.partial( - api.commit, request=commit_request, metadata=call_metadata - ) with error_augmenter: - return commit_method() + return api.commit(request=commit_request, metadata=call_metadata) response = _retry_on_aborted_exception( wrapped_method, @@ -414,24 +409,24 @@ def batch_write(self, request_options=None, exclude_txn_from_change_streams=Fals ) as span, MetricsCapture(self._resource_info), ): - attempt = AtomicCounter(0) + attempt = 0 nth_request = getattr(database, "_next_nth_request", 0) def wrapped_method(): + nonlocal attempt + attempt += 1 batch_write_request = BatchWriteRequest( session=session.name, mutation_groups=mutation_groups, request_options=request_options, exclude_txn_from_change_streams=exclude_txn_from_change_streams, ) - batch_write_method = functools.partial( - api.batch_write, + return api.batch_write( request=batch_write_request, metadata=database.metadata_with_request_id( - nth_request, attempt.increment(), metadata, span + nth_request, attempt, metadata, span ), ) - return batch_write_method() response = _retry( wrapped_method, diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py index 0508e0ae23c8..f2b57d4996a6 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/database.py @@ -49,8 +49,6 @@ _merge_query_options, _metadata_with_leader_aware_routing, _metadata_with_prefix, - _metadata_with_request_id, - _metadata_with_request_id_and_req_id, ) from google.cloud.spanner_v1._opentelemetry_tracing import ( add_span_event, @@ -66,6 +64,11 @@ from google.cloud.spanner_v1.merged_result_set import MergedResultSet from google.cloud.spanner_v1.metrics.metrics_capture import MetricsCapture from google.cloud.spanner_v1.pool import BurstyPool +from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, + _CachedPrefixDescriptor, +) from google.cloud.spanner_v1.services.spanner.client import ( SpannerClient as SpannerClient, ) @@ -152,6 +155,8 @@ class Database(object): _spanner_api: SpannerClient = None __transport_lock = threading.Lock() __transports_to_channel_id = dict() + _channel_id_val = 0 + _req_id_prefix = _CachedPrefixDescriptor() def __init__( self, @@ -474,22 +479,17 @@ def spanner_api(self): self._channel_id = channel_id return self._spanner_api - def metadata_with_request_id( - self, nth_request, nth_attempt, prior_metadata=[], span=None - ): - if span is None: - span = get_current_span() - return _metadata_with_request_id( - self._nth_client_id, - self._channel_id, - nth_request, - nth_attempt, - prior_metadata, - span, - ) + @property + def _channel_id(self): + return self._channel_id_val + + @_channel_id.setter + def _channel_id(self, value): + self._channel_id_val = value + self.__dict__.pop("_req_id_prefix", None) def metadata_and_request_id( - self, nth_request, nth_attempt, prior_metadata=[], span=None + self, nth_request, nth_attempt, prior_metadata=None, span=None ): """Return metadata and request ID string. @@ -506,17 +506,33 @@ def metadata_and_request_id( tuple: (metadata_list, request_id_string)""" if span is None: span = get_current_span() - return _metadata_with_request_id_and_req_id( - self._nth_client_id, - self._channel_id, - nth_request, - nth_attempt, - prior_metadata, - span, + req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}" + metadata = ( + [*prior_metadata, (REQ_ID_HEADER_KEY, req_id)] + if prior_metadata + else [(REQ_ID_HEADER_KEY, req_id)] + ) + if span is not None and span.is_recording(): + span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) + return (metadata, req_id) + + def metadata_with_request_id( + self, nth_request, nth_attempt, prior_metadata=None, span=None + ): + if span is None: + span = get_current_span() + req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}" + metadata = ( + [*prior_metadata, (REQ_ID_HEADER_KEY, req_id)] + if prior_metadata + else [(REQ_ID_HEADER_KEY, req_id)] ) + if span is not None and span.is_recording(): + span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) + return metadata def with_error_augmentation( - self, nth_request, nth_attempt, prior_metadata=[], span=None + self, nth_request, nth_attempt, prior_metadata=None, span=None ): """Context manager for gRPC calls with error augmentation. @@ -533,15 +549,15 @@ def with_error_augmentation( tuple: (metadata_list, context_manager)""" if span is None: span = get_current_span() - metadata, request_id = _metadata_with_request_id_and_req_id( - self._nth_client_id, - self._channel_id, - nth_request, - nth_attempt, - prior_metadata, - span, + req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}" + metadata = ( + [*prior_metadata, (REQ_ID_HEADER_KEY, req_id)] + if prior_metadata + else [(REQ_ID_HEADER_KEY, req_id)] ) - return (metadata, _augment_errors_with_request_id(request_id)) + if span is not None and span.is_recording(): + span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) + return (metadata, _augment_errors_with_request_id(req_id)) def __eq__(self, other): if not isinstance(other, self.__class__): @@ -855,13 +871,13 @@ def execute_pdml(): @property def _next_nth_request(self): if self._instance and self._instance._client: - return self._instance._client._next_nth_request + return getattr(self._instance._client, "_next_nth_request", 1) return 1 @property def _nth_client_id(self): if self._instance and self._instance._client: - return self._instance._client._nth_client_id + return getattr(self._instance._client, "_nth_client_id", 0) return 0 def session(self, labels=None, database_role=None): diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/exceptions.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/exceptions.py index 361079b4f2b0..2a7c13fe2ce8 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/exceptions.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/exceptions.py @@ -34,9 +34,18 @@ def wrap_with_request_id(error, request_id=None): the original error unchanged. """ if isinstance(error, GoogleAPICallError) and request_id: + old_req_id = getattr(error, "request_id", None) + if old_req_id == request_id: + return error # Add request_id as an attribute for programmatic access error.request_id = request_id # Modify the message to include request_id so it appears in logs - if hasattr(error, "message") and error.message: - error.message = f"{error.message}, request_id = {request_id}" + message = getattr(error, "message", None) + if isinstance(message, str) and message: + if old_req_id and f", request_id = {old_req_id}" in message: + error.message = message.replace( + f", request_id = {old_req_id}", f", request_id = {request_id}" + ) + else: + error.message = f"{message}, request_id = {request_id}" return error diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/request_id_header.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/request_id_header.py index 1a5da534e962..ed8a8b31dd21 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/request_id_header.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/request_id_header.py @@ -19,54 +19,67 @@ def generate_rand_uint64(): - b = os.urandom(8) - return ( - b[7] & 0xFF - | (b[6] & 0xFF) << 8 - | (b[5] & 0xFF) << 16 - | (b[4] & 0xFF) << 24 - | (b[3] & 0xFF) << 32 - | (b[2] & 0xFF) << 36 - | (b[1] & 0xFF) << 48 - | (b[0] & 0xFF) << 56 - ) + return int.from_bytes(os.urandom(8), "big") REQ_RAND_PROCESS_ID = generate_rand_uint64() +_REQ_ID_PROCESS_PREFIX = f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}." X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR = "x_goog_spanner_request_id" +class _CachedPrefixDescriptor: + """Non-data descriptor that computes and caches _req_id_prefix in instance __dict__.""" + + def __get__(self, instance, owner=None): + if instance is None: + return self + prefix = ( + f"{_REQ_ID_PROCESS_PREFIX}{instance._nth_client_id}.{instance._channel_id}." + ) + instance.__dict__["_req_id_prefix"] = prefix + return prefix + + def with_request_id( - client_id, channel_id, nth_request, attempt, other_metadata=[], span=None + client_id, channel_id, nth_request, attempt, other_metadata=None, span=None ): - req_id = build_request_id(client_id, channel_id, nth_request, attempt) - all_metadata = (other_metadata or []).copy() - all_metadata.append((REQ_ID_HEADER_KEY, req_id)) + req_id = f"{_REQ_ID_PROCESS_PREFIX}{client_id}.{channel_id}.{nth_request}.{attempt}" + all_metadata = ( + [*other_metadata, (REQ_ID_HEADER_KEY, req_id)] + if other_metadata + else [(REQ_ID_HEADER_KEY, req_id)] + ) - if span: + if span is not None and span.is_recording(): span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) return all_metadata, req_id def with_request_id_metadata_only( - client_id, channel_id, nth_request, attempt, other_metadata=[], span=None + client_id, channel_id, nth_request, attempt, other_metadata=None, span=None ): """Return metadata with request ID header, discarding the request ID value.""" - all_metadata, _ = with_request_id( - client_id, channel_id, nth_request, attempt, other_metadata, span + req_id = f"{_REQ_ID_PROCESS_PREFIX}{client_id}.{channel_id}.{nth_request}.{attempt}" + all_metadata = ( + [*other_metadata, (REQ_ID_HEADER_KEY, req_id)] + if other_metadata + else [(REQ_ID_HEADER_KEY, req_id)] ) + + if span is not None and span.is_recording(): + span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id) + return all_metadata def build_request_id(client_id, channel_id, nth_request, attempt): - return f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.{client_id}.{channel_id}.{nth_request}.{attempt}" + return f"{_REQ_ID_PROCESS_PREFIX}{client_id}.{channel_id}.{nth_request}.{attempt}" def parse_request_id(request_id_str): - splits = request_id_str.split(".") - version, rand_process_id, client_id, channel_id, nth_request, nth_attempt = list( - map(lambda v: int(v), splits) + version, rand_process_id, client_id, channel_id, nth_request, nth_attempt = map( + int, request_id_str.split(".") ) return ( version, 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 ce2bd07306d8..123a0879c8d6 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/snapshot.py @@ -31,7 +31,6 @@ from google.cloud.aio._cross_sync import CrossSync from google.cloud.spanner_v1._helpers import ( - AtomicCounter, _augment_error_with_request_id, _check_rst_stream_error, _make_value_pb, @@ -680,20 +679,20 @@ def partition_read( MetricsCapture(self._resource_info), ): nth_request = getattr(database, "_next_nth_request", 0) - attempt = AtomicCounter() + attempt = 0 def attempt_tracking_method(): + nonlocal attempt + attempt += 1 all_metadata = database.metadata_with_request_id( - nth_request, attempt.increment(), metadata, span + nth_request, attempt, metadata, span ) - partition_read_method = functools.partial( - api.partition_read, + return api.partition_read( request=partition_read_request, metadata=all_metadata, retry=retry, timeout=timeout, ) - return partition_read_method() response = _retry( attempt_tracking_method, @@ -755,20 +754,20 @@ def partition_query( MetricsCapture(self._resource_info), ): nth_request = getattr(database, "_next_nth_request", 0) - attempt = AtomicCounter() + attempt = 0 def attempt_tracking_method(): + nonlocal attempt + attempt += 1 all_metadata = database.metadata_with_request_id( - nth_request, attempt.increment(), metadata, span + nth_request, attempt, metadata, span ) - partition_query_method = functools.partial( - api.partition_query, + return api.partition_query( request=partition_query_request, metadata=all_metadata, retry=retry, timeout=timeout, ) - return partition_query_method() response = _retry( attempt_tracking_method, @@ -820,22 +819,21 @@ def _begin_transaction( MetricsCapture(self._resource_info), ): nth_request = getattr(database, "_next_nth_request", 0) - attempt = AtomicCounter() + attempt = 0 def wrapped_method(): + nonlocal attempt + attempt += 1 begin_transaction_request = BeginTransactionRequest( **begin_request_kwargs ) call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.increment(), metadata, span - ) - begin_transaction_method = functools.partial( - api.begin_transaction, - request=begin_transaction_request, - metadata=call_metadata, + nth_request, attempt, metadata, span ) with error_augmenter: - return begin_transaction_method() + return api.begin_transaction( + request=begin_transaction_request, metadata=call_metadata + ) def before_next_retry(nth_retry, delay_in_seconds): add_span_event( diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/transaction.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/transaction.py index 43c58400d1ad..63dad3b41aea 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/transaction.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/transaction.py @@ -26,7 +26,6 @@ from google.protobuf.struct_pb2 import Struct from google.cloud.spanner_v1._helpers import ( - AtomicCounter, _check_rst_stream_error, _make_value_pb, _merge_client_context, @@ -124,10 +123,10 @@ def _execute_request( session._database, "observability_options", None ), metadata=metadata, - ), + ) as span, MetricsCapture(self._resource_info), ): - method = functools.partial(method, request=request) + method = functools.partial(method, request=request, span=span) response = _retry( method, allowed_exceptions={InternalServerError: _check_rst_stream_error}, @@ -163,22 +162,21 @@ def rollback(self) -> None: ) as span, MetricsCapture(self._resource_info), ): - attempt = AtomicCounter(0) + attempt = 0 nth_request = database._next_nth_request - def wrapped_method(*args, **kwargs): - attempt.increment() + def wrapped_method(): + nonlocal attempt + attempt += 1 call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.value, metadata, span - ) - rollback_method = functools.partial( - api.rollback, - session=session.name, - transaction_id=self._transaction_id, - metadata=call_metadata, + nth_request, attempt, metadata, span ) with error_augmenter: - return rollback_method(*args, **kwargs) + return api.rollback( + session=session.name, + transaction_id=self._transaction_id, + metadata=call_metadata, + ) _retry( wrapped_method, @@ -266,11 +264,12 @@ def commit( "request_options": request_options, } add_span_event(span, "Starting Commit") - attempt = AtomicCounter(0) + attempt = 0 nth_request = database._next_nth_request - def wrapped_method(*args, **kwargs): - attempt.increment() + def wrapped_method(): + nonlocal attempt + attempt += 1 commit_request_args = { "mutations": mutations, **common_commit_request_args, @@ -279,15 +278,13 @@ def wrapped_method(*args, **kwargs): if is_multiplexed and self._precommit_token is not None: commit_request_args["precommit_token"] = self._precommit_token call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.value, metadata, span - ) - commit_method = functools.partial( - api.commit, - request=CommitRequest(**commit_request_args), - metadata=call_metadata, + nth_request, attempt, metadata, span ) with error_augmenter: - return commit_method(*args, **kwargs) + return api.commit( + request=CommitRequest(**commit_request_args), + metadata=call_metadata, + ) commit_retry_event_name = "Transaction Commit Attempt Failed. Retrying" @@ -459,22 +456,21 @@ def execute_update( last_statement=last_statement, ) nth_request = database._next_nth_request - attempt = AtomicCounter(0) + attempt = 0 - def wrapped_method(*args, **kwargs): - attempt.increment() + def wrapped_method(request=execute_sql_request, span=None): + nonlocal attempt + attempt += 1 call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.value, metadata - ) - execute_sql_method = functools.partial( - api.execute_sql, - request=execute_sql_request, - metadata=call_metadata, - retry=retry, - timeout=timeout, + nth_request, attempt, metadata, span ) with error_augmenter: - return execute_sql_method(*args, **kwargs) + return api.execute_sql( + request=request, + metadata=call_metadata, + retry=retry, + timeout=timeout, + ) result_set_pb: ResultSet = self._execute_request( wrapped_method, @@ -595,22 +591,21 @@ def batch_update( last_statements=last_statement, ) nth_request = database._next_nth_request - attempt = AtomicCounter(0) + attempt = 0 - def wrapped_method(*args, **kwargs): - attempt.increment() + def wrapped_method(request=execute_batch_dml_request, span=None): + nonlocal attempt + attempt += 1 call_metadata, error_augmenter = database.with_error_augmentation( - nth_request, attempt.value, metadata - ) - execute_batch_dml_method = functools.partial( - api.execute_batch_dml, - request=execute_batch_dml_request, - metadata=call_metadata, - retry=retry, - timeout=timeout, + nth_request, attempt, metadata, span ) with error_augmenter: - return execute_batch_dml_method(*args, **kwargs) + return api.execute_batch_dml( + request=request, + metadata=call_metadata, + retry=retry, + timeout=timeout, + ) response_pb: ExecuteBatchDmlResponse = self._execute_request( wrapped_method, 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..d952f399bedb 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_database.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_database.py @@ -161,6 +161,114 @@ async def test_close(self): database._sessions_manager.close.assert_called_once() + @CrossSync.pytest + async def test_req_id_prefix_and_caching(self): + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + ) + + instance = _Instance(self.INSTANCE_NAME) + database = await self._make_one(self.DATABASE_ID, instance) + database._channel_id = 42 + self.assertNotIn("_req_id_prefix", database.__dict__) + + expected_prefix = ( + f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.42." + ) + self.assertEqual(database._req_id_prefix, expected_prefix) + self.assertEqual(database.__dict__["_req_id_prefix"], expected_prefix) + + # Invalidate when channel_id changes + database._channel_id = 43 + self.assertNotIn("_req_id_prefix", database.__dict__) + expected_prefix_2 = ( + f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.43." + ) + self.assertEqual(database._req_id_prefix, expected_prefix_2) + self.assertEqual(database.__dict__["_req_id_prefix"], expected_prefix_2) + + @CrossSync.pytest + async def test_database_request_id_methods(self): + from unittest import mock + + from google.api_core.exceptions import GoogleAPICallError + + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, + ) + + instance = _Instance(self.INSTANCE_NAME) + database = await self._make_one(self.DATABASE_ID, instance) + database._channel_id = 1 + + mock_rec_span = mock.Mock() + mock_rec_span.is_recording.return_value = True + mock_non_rec_span = mock.Mock() + mock_non_rec_span.is_recording.return_value = False + + # Test metadata_with_request_id + meta = database.metadata_with_request_id( + 5, 2, [("foo", "bar")], span=mock_rec_span + ) + expected_id = f"{database._req_id_prefix}5.2" + self.assertEqual(meta, [("foo", "bar"), (REQ_ID_HEADER_KEY, expected_id)]) + mock_rec_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, expected_id + ) + + meta_no_prior = database.metadata_with_request_id( + 5, 3, prior_metadata=None, span=mock_non_rec_span + ) + self.assertEqual( + meta_no_prior, [(REQ_ID_HEADER_KEY, f"{database._req_id_prefix}5.3")] + ) + mock_non_rec_span.set_attribute.assert_not_called() + + # Test metadata_and_request_id + mock_rec_span.reset_mock() + meta2, req_id = database.metadata_and_request_id(6, 1, span=mock_rec_span) + expected_id_2 = f"{database._req_id_prefix}6.1" + self.assertEqual(req_id, expected_id_2) + self.assertEqual(meta2, [(REQ_ID_HEADER_KEY, expected_id_2)]) + mock_rec_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, expected_id_2 + ) + + meta2_prior, _ = database.metadata_and_request_id( + 6, 2, prior_metadata=[("k", "v")], span=mock_non_rec_span + ) + self.assertEqual( + meta2_prior, + [("k", "v"), (REQ_ID_HEADER_KEY, f"{database._req_id_prefix}6.2")], + ) + mock_non_rec_span.set_attribute.assert_not_called() + + # Test with_error_augmentation + mock_rec_span.reset_mock() + meta3, error_aug = database.with_error_augmentation( + 7, 1, prior_metadata=[("a", "b")], span=mock_rec_span + ) + expected_id_3 = f"{database._req_id_prefix}7.1" + self.assertEqual(meta3, [("a", "b"), (REQ_ID_HEADER_KEY, expected_id_3)]) + mock_rec_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, expected_id_3 + ) + + err = GoogleAPICallError("mock error") + with self.assertRaises(GoogleAPICallError): + with error_aug: + raise err + self.assertEqual(err.request_id, expected_id_3) + + meta3_no_prior, _ = database.with_error_augmentation( + 7, 2, prior_metadata=None, span=mock_non_rec_span + ) + self.assertEqual( + meta3_no_prior, [(REQ_ID_HEADER_KEY, f"{database._req_id_prefix}7.2")] + ) + @CrossSync.pytest async def test_sessions_manager_close(self): from google.cloud.spanner_v1._async.database_sessions_manager import ( diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_transaction.py b/packages/google-cloud-spanner/tests/unit/_async/test_transaction.py index 9e71af613ae4..1d348031345d 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_transaction.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_transaction.py @@ -237,6 +237,28 @@ async def test_rollback_w_other_error(self, mock_region): ), ) + @mock.patch( + "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", + return_value="global", + ) + @CrossSync.pytest + async def test_rollback_w_grpc_error(self, mock_region): + from google.api_core.exceptions import Unknown + + database = _Database() + database.spanner_api = self._make_spanner_api() + err = Unknown("grpc error") + database.spanner_api.rollback.side_effect = err + session = _Session(database) + transaction = self._make_one(session) + transaction._transaction_id = TRANSACTION_ID + + with pytest.raises(Unknown): + await transaction.rollback() + + req_id = f"1.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.{database._channel_id}.1.1" + self.assertEqual(getattr(err, "request_id", None), req_id) + @mock.patch( "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", return_value="global", @@ -280,6 +302,47 @@ async def test_rollback_ok(self, mock_region): ), ) + @mock.patch( + "google.cloud.spanner_v1._async._helpers.asyncio.sleep", + new_callable=mock.AsyncMock, + ) + @mock.patch( + "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", + return_value="global", + ) + @CrossSync.pytest + async def test_rollback_w_retry(self, mock_region, mock_sleep): + from google.api_core.exceptions import InternalServerError + from google.protobuf.empty_pb2 import Empty + + empty_pb = Empty() + database = _Database() + metadata_calls = [] + + async def mock_rollback(session=None, transaction_id=None, metadata=None): + metadata_calls.append(metadata) + if len(metadata_calls) == 1: + raise InternalServerError("RST_STREAM") + return empty_pb + + api = database.spanner_api = _FauxSpannerAPI() + api.rollback = mock_rollback + + session = _Session(database) + transaction = self._make_one(session) + transaction._transaction_id = TRANSACTION_ID + + await transaction.rollback() + + self.assertTrue(transaction.rolled_back) + self.assertEqual(len(metadata_calls), 2) + + req_id_1 = f"1.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.{database._channel_id}.1.1" + req_id_2 = f"1.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.{database._channel_id}.1.2" + + self.assertIn(("x-goog-spanner-request-id", req_id_1), metadata_calls[0]) + self.assertIn(("x-goog-spanner-request-id", req_id_2), metadata_calls[1]) + @CrossSync.pytest async def test_commit_not_begun(self): database = _Database() @@ -884,7 +947,9 @@ async def _execute_update_helper( ) expected_attributes = self._build_span_attributes( - database, **{"db.statement": DML_QUERY_WITH_PARAM} + database, + x_goog_spanner_request_id=f"1.{REQ_RAND_PROCESS_ID}.{_Client.NTH_CLIENT.value}.1.1.1", + **{"db.statement": DML_QUERY_WITH_PARAM}, ) if request_options.request_tag: expected_attributes["request.tag"] = request_options.request_tag @@ -1566,15 +1631,15 @@ class _FauxSpannerAPI(object): def __init__(self, **kwargs): self.__dict__.update(**kwargs) - def begin_transaction(self, session=None, options=None, metadata=None): + async def begin_transaction(self, session=None, options=None, metadata=None): self._begun = (session, options, metadata) return self._begin_transaction_response - def rollback(self, session=None, transaction_id=None, metadata=None): + async def rollback(self, session=None, transaction_id=None, metadata=None): self._rolled_back = (session, transaction_id, metadata) return self._rollback_response - def commit( + async def commit( self, request=None, metadata=None, diff --git a/packages/google-cloud-spanner/tests/unit/test__helpers.py b/packages/google-cloud-spanner/tests/unit/test__helpers.py index 3776bbd26141..c416e606b853 100644 --- a/packages/google-cloud-spanner/tests/unit/test__helpers.py +++ b/packages/google-cloud-spanner/tests/unit/test__helpers.py @@ -2049,3 +2049,197 @@ def test_create_spanner_omni_transport_interceptors_and_credentials_fallback(sel self.assertIsInstance( mock_factory.call_args[1]["credentials"], AnonymousCredentials ) + + +class TestRequestIdHelpers(unittest.TestCase): + def test_build_request_id(self): + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + build_request_id, + ) + + req_id = build_request_id(1, 2, 3, 4) + expected = f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.1.2.3.4" + self.assertEqual(req_id, expected) + + def test_with_request_id(self): + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + with_request_id, + ) + + expected_id = f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.10.20.30.1" + metadata, req_id = with_request_id(10, 20, 30, 1, [("prior", "val")]) + self.assertEqual(req_id, expected_id) + self.assertEqual(metadata, [("prior", "val"), (REQ_ID_HEADER_KEY, expected_id)]) + + def test_with_request_id_metadata_only(self): + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + with_request_id_metadata_only, + ) + + expected_id = f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.10.20.30.1" + metadata = with_request_id_metadata_only(10, 20, 30, 1) + self.assertEqual(metadata, [(REQ_ID_HEADER_KEY, expected_id)]) + + def test_with_request_id_span_recording(self): + from unittest import mock + + from google.cloud.spanner_v1.request_id_header import ( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, + with_request_id, + ) + + mock_span = mock.Mock() + mock_span.is_recording.return_value = True + + _, req_id = with_request_id(1, 1, 1, 1, span=mock_span) + mock_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id + ) + + def test_with_request_id_span_non_recording(self): + from unittest import mock + + from google.cloud.spanner_v1.request_id_header import with_request_id + + mock_span = mock.Mock() + mock_span.is_recording.return_value = False + + with_request_id(1, 1, 1, 1, span=mock_span) + mock_span.set_attribute.assert_not_called() + + def test_augment_errors_with_request_id_success(self): + from google.cloud.spanner_v1._helpers import _augment_errors_with_request_id + + with _augment_errors_with_request_id("test-req-id"): + val = 42 + self.assertEqual(val, 42) + + def test_augment_errors_with_request_id_google_api_call_error(self): + from google.api_core.exceptions import GoogleAPICallError + + from google.cloud.spanner_v1._helpers import _augment_errors_with_request_id + + err = GoogleAPICallError("something went wrong") + with self.assertRaises(GoogleAPICallError) as ctx: + with _augment_errors_with_request_id("test-req-id"): + raise err + + raised = ctx.exception + self.assertIs(raised, err) + self.assertEqual(getattr(raised, "request_id", None), "test-req-id") + self.assertIn("request_id = test-req-id", raised.message) + + def test_augment_errors_with_request_id_non_api_error(self): + from google.cloud.spanner_v1._helpers import _augment_errors_with_request_id + + err = ValueError("regular error") + with self.assertRaises(ValueError) as ctx: + with _augment_errors_with_request_id("test-req-id"): + raise err + + raised = ctx.exception + self.assertIs(raised, err) + self.assertFalse(hasattr(raised, "request_id")) + + def test_atomic_counter_slots(self): + from google.cloud.spanner_v1._helpers import AtomicCounter + + counter = AtomicCounter() + self.assertFalse(hasattr(counter, "__dict__")) + self.assertEqual(counter.value, 0) + self.assertEqual(counter.increment(), 1) + self.assertEqual(counter.value, 1) + counter += 2 + self.assertEqual(counter.value, 3) + counter.reset() + self.assertEqual(counter.value, 0) + + def test_helpers_metadata_with_request_id_wrappers(self): + from google.cloud.spanner_v1._helpers import ( + _metadata_with_request_id, + _metadata_with_request_id_and_req_id, + ) + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + ) + + expected_id = f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.1.2.3.4" + meta = _metadata_with_request_id(1, 2, 3, 4, [("key", "val")]) + self.assertEqual(meta, [("key", "val"), (REQ_ID_HEADER_KEY, expected_id)]) + + meta_tuple, req_id = _metadata_with_request_id_and_req_id(1, 2, 3, 4) + self.assertEqual(req_id, expected_id) + self.assertEqual(meta_tuple, [(REQ_ID_HEADER_KEY, expected_id)]) + + def test_parse_request_id(self): + from google.cloud.spanner_v1.request_id_header import parse_request_id + + parsed = parse_request_id("1.12345.2.3.4.5") + self.assertEqual(parsed, (1, 12345, 2, 3, 4, 5)) + + def test_with_request_id_metadata_only_span_recording(self): + from unittest import mock + + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, + with_request_id_metadata_only, + ) + + mock_span = mock.Mock() + mock_span.is_recording.return_value = True + + expected_id = f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.1.2.3.4" + meta = with_request_id_metadata_only( + 1, 2, 3, 4, other_metadata=[("k", "v")], span=mock_span + ) + self.assertEqual(meta, [("k", "v"), (REQ_ID_HEADER_KEY, expected_id)]) + mock_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, expected_id + ) + + def test_with_request_id_metadata_only_span_non_recording(self): + from unittest import mock + + from google.cloud.spanner_v1.request_id_header import ( + with_request_id_metadata_only, + ) + + mock_span = mock.Mock() + mock_span.is_recording.return_value = False + + with_request_id_metadata_only(1, 2, 3, 4, span=mock_span) + mock_span.set_attribute.assert_not_called() + + def test_cached_prefix_descriptor(self): + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + _CachedPrefixDescriptor, + ) + + class DummyDB: + _nth_client_id = 5 + _channel_id = 10 + _req_id_prefix = _CachedPrefixDescriptor() + + # Class access returns descriptor instance + self.assertIsInstance(DummyDB._req_id_prefix, _CachedPrefixDescriptor) + + # Instance access computes and caches in __dict__ + db = DummyDB() + expected = f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.5.10." + self.assertEqual(db._req_id_prefix, expected) + self.assertEqual(db.__dict__["_req_id_prefix"], expected) diff --git a/packages/google-cloud-spanner/tests/unit/test_database.py b/packages/google-cloud-spanner/tests/unit/test_database.py index 21c382b45438..0bf856446f52 100644 --- a/packages/google-cloud-spanner/tests/unit/test_database.py +++ b/packages/google-cloud-spanner/tests/unit/test_database.py @@ -163,6 +163,113 @@ def test_ctor_w_database_role(self): self.assertIs(database._instance, instance) self.assertIs(database.database_role, self.DATABASE_ROLE) + def test_req_id_prefix_and_caching(self): + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_VERSION, + REQ_RAND_PROCESS_ID, + ) + + instance = _Instance(self.INSTANCE_NAME) + database = self._make_one(self.DATABASE_ID, instance) + database._channel_id = 42 + self.assertNotIn("_req_id_prefix", database.__dict__) + + expected_prefix = ( + f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.42." + ) + self.assertEqual(database._req_id_prefix, expected_prefix) + self.assertEqual(database.__dict__["_req_id_prefix"], expected_prefix) + + # Invalidate when channel_id changes + database._channel_id = 43 + self.assertNotIn("_req_id_prefix", database.__dict__) + expected_prefix_2 = ( + f"{REQ_ID_VERSION}.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.43." + ) + self.assertEqual(database._req_id_prefix, expected_prefix_2) + self.assertEqual(database.__dict__["_req_id_prefix"], expected_prefix_2) + + def test_database_request_id_methods(self): + from unittest import mock + + from google.api_core.exceptions import GoogleAPICallError + + from google.cloud.spanner_v1.request_id_header import ( + REQ_ID_HEADER_KEY, + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, + ) + + instance = _Instance(self.INSTANCE_NAME) + database = self._make_one(self.DATABASE_ID, instance) + database._channel_id = 1 + + mock_rec_span = mock.Mock() + mock_rec_span.is_recording.return_value = True + mock_non_rec_span = mock.Mock() + mock_non_rec_span.is_recording.return_value = False + + # Test metadata_with_request_id + meta = database.metadata_with_request_id( + 5, 2, [("foo", "bar")], span=mock_rec_span + ) + expected_id = f"{database._req_id_prefix}5.2" + self.assertEqual(meta, [("foo", "bar"), (REQ_ID_HEADER_KEY, expected_id)]) + mock_rec_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, expected_id + ) + + meta_no_prior = database.metadata_with_request_id( + 5, 3, prior_metadata=None, span=mock_non_rec_span + ) + self.assertEqual( + meta_no_prior, [(REQ_ID_HEADER_KEY, f"{database._req_id_prefix}5.3")] + ) + mock_non_rec_span.set_attribute.assert_not_called() + + # Test metadata_and_request_id + mock_rec_span.reset_mock() + meta2, req_id = database.metadata_and_request_id(6, 1, span=mock_rec_span) + expected_id_2 = f"{database._req_id_prefix}6.1" + self.assertEqual(req_id, expected_id_2) + self.assertEqual(meta2, [(REQ_ID_HEADER_KEY, expected_id_2)]) + mock_rec_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, expected_id_2 + ) + + meta2_prior, _ = database.metadata_and_request_id( + 6, 2, prior_metadata=[("k", "v")], span=mock_non_rec_span + ) + self.assertEqual( + meta2_prior, + [("k", "v"), (REQ_ID_HEADER_KEY, f"{database._req_id_prefix}6.2")], + ) + mock_non_rec_span.set_attribute.assert_not_called() + + # Test with_error_augmentation + mock_rec_span.reset_mock() + meta3, error_aug = database.with_error_augmentation( + 7, 1, prior_metadata=[("a", "b")], span=mock_rec_span + ) + expected_id_3 = f"{database._req_id_prefix}7.1" + self.assertEqual(meta3, [("a", "b"), (REQ_ID_HEADER_KEY, expected_id_3)]) + mock_rec_span.set_attribute.assert_called_once_with( + X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, expected_id_3 + ) + + err = GoogleAPICallError("mock error") + with self.assertRaises(GoogleAPICallError): + with error_aug: + raise err + self.assertEqual(err.request_id, expected_id_3) + + meta3_no_prior, _ = database.with_error_augmentation( + 7, 2, prior_metadata=None, span=mock_non_rec_span + ) + self.assertEqual( + meta3_no_prior, [(REQ_ID_HEADER_KEY, f"{database._req_id_prefix}7.2")] + ) + mock_non_rec_span.set_attribute.assert_not_called() + def test_ctor_w_route_to_leader_disbled(self): client = _Client(route_to_leader_enabled=False) instance = _Instance(self.INSTANCE_NAME, client=client) diff --git a/packages/google-cloud-spanner/tests/unit/test_exceptions.py b/packages/google-cloud-spanner/tests/unit/test_exceptions.py index 7f0cfb3c455b..3ae42911cb80 100644 --- a/packages/google-cloud-spanner/tests/unit/test_exceptions.py +++ b/packages/google-cloud-spanner/tests/unit/test_exceptions.py @@ -61,6 +61,40 @@ def test_wrap_with_request_id_with_non_google_api_error(self): self.assertIs(result, error) self.assertFalse(hasattr(result, "request_id")) + def test_wrap_with_request_id_idempotent_and_retry(self): + """Test that re-wrapping does not duplicate request_id suffix and updates on retry.""" + error = Aborted("Transaction aborted") + req_id_1 = "1.12345.1.0.1.1" + req_id_2 = "1.12345.1.0.1.2" + + # First wrap + wrap_with_request_id(error, req_id_1) + self.assertEqual(error.request_id, req_id_1) + self.assertEqual(error.message, f"Transaction aborted, request_id = {req_id_1}") + + # Wrapping again with the same request_id should be a no-op + wrap_with_request_id(error, req_id_1) + self.assertEqual(error.request_id, req_id_1) + self.assertEqual(error.message, f"Transaction aborted, request_id = {req_id_1}") + + # Wrapping with a new request_id on retry should replace the old request_id in message + wrap_with_request_id(error, req_id_2) + self.assertEqual(error.request_id, req_id_2) + self.assertEqual(error.message, f"Transaction aborted, request_id = {req_id_2}") + + # Wrapping error with empty message sets request_id attribute without altering empty message + empty_err = Aborted("") + wrap_with_request_id(empty_err, req_id_1) + self.assertEqual(empty_err.request_id, req_id_1) + self.assertEqual(empty_err.message, "") + + # Wrapping error with non-string message (e.g. int/mock) sets request_id without raising TypeError + non_str_err = Aborted("msg") + non_str_err.message = 12345 + wrap_with_request_id(non_str_err, req_id_1) + self.assertEqual(non_str_err.request_id, req_id_1) + self.assertEqual(non_str_err.message, 12345) + if __name__ == "__main__": unittest.main() diff --git a/packages/google-cloud-spanner/tests/unit/test_transaction.py b/packages/google-cloud-spanner/tests/unit/test_transaction.py index aa9e9ae14995..fd568db63fd7 100644 --- a/packages/google-cloud-spanner/tests/unit/test_transaction.py +++ b/packages/google-cloud-spanner/tests/unit/test_transaction.py @@ -221,6 +221,27 @@ def test_rollback_w_other_error(self, mock_region): ), ) + @mock.patch( + "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", + return_value="global", + ) + def test_rollback_w_grpc_error(self, mock_region): + from google.api_core.exceptions import Unknown + + database = _Database() + database.spanner_api = self._make_spanner_api() + err = Unknown("grpc error") + database.spanner_api.rollback.side_effect = err + session = _Session(database) + transaction = self._make_one(session) + transaction._transaction_id = TRANSACTION_ID + + with self.assertRaises(Unknown): + transaction.rollback() + + req_id = f"1.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.{database._channel_id}.1.1" + self.assertEqual(getattr(err, "request_id", None), req_id) + @mock.patch( "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", return_value="global", @@ -263,6 +284,43 @@ def test_rollback_ok(self, mock_region): ), ) + @mock.patch("google.cloud.spanner_v1._helpers.time.sleep") + @mock.patch( + "google.cloud.spanner_v1._opentelemetry_tracing._get_cloud_region", + return_value="global", + ) + def test_rollback_w_retry(self, mock_region, mock_sleep): + from google.api_core.exceptions import InternalServerError + from google.protobuf.empty_pb2 import Empty + + empty_pb = Empty() + database = _Database() + metadata_calls = [] + + def mock_rollback(session=None, transaction_id=None, metadata=None): + metadata_calls.append(metadata) + if len(metadata_calls) == 1: + raise InternalServerError("RST_STREAM") + return empty_pb + + api = database.spanner_api = _FauxSpannerAPI() + api.rollback = mock_rollback + + session = _Session(database) + transaction = self._make_one(session) + transaction._transaction_id = TRANSACTION_ID + + transaction.rollback() + + self.assertTrue(transaction.rolled_back) + self.assertEqual(len(metadata_calls), 2) + + req_id_1 = f"1.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.{database._channel_id}.1.1" + req_id_2 = f"1.{REQ_RAND_PROCESS_ID}.{database._nth_client_id}.{database._channel_id}.1.2" + + self.assertIn(("x-goog-spanner-request-id", req_id_1), metadata_calls[0]) + self.assertIn(("x-goog-spanner-request-id", req_id_2), metadata_calls[1]) + def test_commit_not_begun(self): database = _Database() database.spanner_api = self._make_spanner_api() @@ -842,7 +900,9 @@ def _execute_update_helper( ) expected_attributes = self._build_span_attributes( - database, **{"db.statement": DML_QUERY_WITH_PARAM} + database, + x_goog_spanner_request_id=f"1.{REQ_RAND_PROCESS_ID}.{_Client.NTH_CLIENT.value}.1.1.1", + **{"db.statement": DML_QUERY_WITH_PARAM}, ) if request_options.request_tag: expected_attributes["request.tag"] = request_options.request_tag