Skip to content

Commit 5a0f7a5

Browse files
committed
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.
1 parent 9a1fb24 commit 5a0f7a5

17 files changed

Lines changed: 854 additions & 277 deletions

packages/google-cloud-spanner/google/cloud/spanner_v1/_async/batch.py

Lines changed: 9 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,6 @@
1515
"""Context manager for Cloud Spanner batched writes."""
1616

1717
__CROSS_SYNC_OUTPUT__ = "google.cloud.spanner_v1.batch"
18-
import functools
1918
import time
2019
from typing import List, Optional
2120

@@ -25,7 +24,6 @@
2524
from google.cloud.aio._cross_sync import CrossSync
2625
from google.cloud.spanner_v1._async._helpers import _retry, _retry_on_aborted_exception
2726
from google.cloud.spanner_v1._helpers import (
28-
AtomicCounter,
2927
_check_rst_stream_error,
3028
_make_list_value_pb,
3129
_make_list_value_pbs,
@@ -341,13 +339,11 @@ async def wrapped_method():
341339
metadata,
342340
span,
343341
)
344-
commit_method = functools.partial(
345-
api.commit,
346-
request=commit_request,
347-
metadata=call_metadata,
348-
)
349342
with error_augmenter:
350-
return await commit_method()
343+
return await api.commit(
344+
request=commit_request,
345+
metadata=call_metadata,
346+
)
351347

352348
response = await _retry_on_aborted_exception(
353349
wrapped_method,
@@ -478,27 +474,27 @@ async def batch_write(
478474
) as span,
479475
MetricsCapture(self._resource_info),
480476
):
481-
attempt = AtomicCounter(0)
477+
attempt = 0
482478
nth_request = getattr(database, "_next_nth_request", 0)
483479

484480
def wrapped_method():
481+
nonlocal attempt
482+
attempt += 1
485483
batch_write_request = BatchWriteRequest(
486484
session=session.name,
487485
mutation_groups=mutation_groups,
488486
request_options=request_options,
489487
exclude_txn_from_change_streams=exclude_txn_from_change_streams,
490488
)
491-
batch_write_method = functools.partial(
492-
api.batch_write,
489+
return api.batch_write(
493490
request=batch_write_request,
494491
metadata=database.metadata_with_request_id(
495492
nth_request,
496-
attempt.increment(),
493+
attempt,
497494
metadata,
498495
span,
499496
),
500497
)
501-
return batch_write_method()
502498

503499
response = await _retry(
504500
wrapped_method,

packages/google-cloud-spanner/google/cloud/spanner_v1/_async/database.py

Lines changed: 55 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -58,11 +58,14 @@
5858
_merge_query_options,
5959
_metadata_with_leader_aware_routing,
6060
_metadata_with_prefix,
61-
_metadata_with_request_id,
62-
_metadata_with_request_id_and_req_id,
6361
)
6462
from google.cloud.spanner_v1.keyset import KeySet
6563
from google.cloud.spanner_v1.merged_result_set import MergedResultSet
64+
from google.cloud.spanner_v1.request_id_header import (
65+
REQ_ID_HEADER_KEY,
66+
X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR,
67+
_CachedPrefixDescriptor,
68+
)
6669
from google.cloud.spanner_v1.services.spanner.async_client import (
6770
SpannerAsyncClient as SpannerClient,
6871
)
@@ -180,6 +183,8 @@ class Database(object):
180183

181184
__transport_lock = threading.Lock()
182185
__transports_to_channel_id = dict()
186+
_channel_id_val = 0
187+
_req_id_prefix = _CachedPrefixDescriptor()
183188

184189
def __init__(
185190
self,
@@ -550,23 +555,17 @@ def spanner_api(self):
550555

551556
return self._spanner_api
552557

553-
def metadata_with_request_id(
554-
self, nth_request, nth_attempt, prior_metadata=[], span=None
555-
):
556-
if span is None:
557-
span = get_current_span()
558+
@property
559+
def _channel_id(self):
560+
return self._channel_id_val
558561

559-
return _metadata_with_request_id(
560-
self._nth_client_id,
561-
self._channel_id,
562-
nth_request,
563-
nth_attempt,
564-
prior_metadata,
565-
span,
566-
)
562+
@_channel_id.setter
563+
def _channel_id(self, value):
564+
self._channel_id_val = value
565+
self.__dict__.pop("_req_id_prefix", None)
567566

568567
def metadata_and_request_id(
569-
self, nth_request, nth_attempt, prior_metadata=[], span=None
568+
self, nth_request, nth_attempt, prior_metadata=None, span=None
570569
):
571570
"""Return metadata and request ID string.
572571
@@ -585,17 +584,38 @@ def metadata_and_request_id(
585584
if span is None:
586585
span = get_current_span()
587586

588-
return _metadata_with_request_id_and_req_id(
589-
self._nth_client_id,
590-
self._channel_id,
591-
nth_request,
592-
nth_attempt,
593-
prior_metadata,
594-
span,
587+
req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}"
588+
metadata = (
589+
[*prior_metadata, (REQ_ID_HEADER_KEY, req_id)]
590+
if prior_metadata
591+
else [(REQ_ID_HEADER_KEY, req_id)]
595592
)
596593

594+
if span is not None and span.is_recording():
595+
span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id)
596+
597+
return metadata, req_id
598+
599+
def metadata_with_request_id(
600+
self, nth_request, nth_attempt, prior_metadata=None, span=None
601+
):
602+
if span is None:
603+
span = get_current_span()
604+
605+
req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}"
606+
metadata = (
607+
[*prior_metadata, (REQ_ID_HEADER_KEY, req_id)]
608+
if prior_metadata
609+
else [(REQ_ID_HEADER_KEY, req_id)]
610+
)
611+
612+
if span is not None and span.is_recording():
613+
span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id)
614+
615+
return metadata
616+
597617
def with_error_augmentation(
598-
self, nth_request, nth_attempt, prior_metadata=[], span=None
618+
self, nth_request, nth_attempt, prior_metadata=None, span=None
599619
):
600620
"""Context manager for gRPC calls with error augmentation.
601621
@@ -614,16 +634,17 @@ def with_error_augmentation(
614634
if span is None:
615635
span = get_current_span()
616636

617-
metadata, request_id = _metadata_with_request_id_and_req_id(
618-
self._nth_client_id,
619-
self._channel_id,
620-
nth_request,
621-
nth_attempt,
622-
prior_metadata,
623-
span,
637+
req_id = f"{self._req_id_prefix}{nth_request}.{nth_attempt}"
638+
metadata = (
639+
[*prior_metadata, (REQ_ID_HEADER_KEY, req_id)]
640+
if prior_metadata
641+
else [(REQ_ID_HEADER_KEY, req_id)]
624642
)
625643

626-
return metadata, _augment_errors_with_request_id(request_id)
644+
if span is not None and span.is_recording():
645+
span.set_attribute(X_GOOG_SPANNER_REQUEST_ID_SPAN_ATTR, req_id)
646+
647+
return metadata, _augment_errors_with_request_id(req_id)
627648

628649
def __eq__(self, other):
629650
if not isinstance(other, self.__class__):
@@ -975,13 +996,13 @@ async def execute_pdml():
975996
@property
976997
def _next_nth_request(self):
977998
if self._instance and self._instance._client:
978-
return self._instance._client._next_nth_request
999+
return getattr(self._instance._client, "_next_nth_request", 1)
9791000
return 1
9801001

9811002
@property
9821003
def _nth_client_id(self):
9831004
if self._instance and self._instance._client:
984-
return self._instance._client._nth_client_id
1005+
return getattr(self._instance._client, "_nth_client_id", 0)
9851006
return 0
9861007

9871008
def session(self, labels=None, database_role=None):

packages/google-cloud-spanner/google/cloud/spanner_v1/_async/snapshot.py

Lines changed: 18 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,6 @@
3131
from google.cloud.spanner_v1._async._helpers import _retry
3232
from google.cloud.spanner_v1._async.streamed import StreamedResultSet
3333
from google.cloud.spanner_v1._helpers import (
34-
AtomicCounter,
3534
_augment_error_with_request_id,
3635
_check_rst_stream_error,
3736
_make_value_pb,
@@ -760,20 +759,20 @@ async def partition_read(
760759
MetricsCapture(self._resource_info),
761760
):
762761
nth_request = getattr(database, "_next_nth_request", 0)
763-
attempt = AtomicCounter()
762+
attempt = 0
764763

765764
async def attempt_tracking_method():
765+
nonlocal attempt
766+
attempt += 1
766767
all_metadata = database.metadata_with_request_id(
767-
nth_request, attempt.increment(), metadata, span
768+
nth_request, attempt, metadata, span
768769
)
769-
partition_read_method = functools.partial(
770-
api.partition_read,
770+
return await api.partition_read(
771771
request=partition_read_request,
772772
metadata=all_metadata,
773773
retry=retry,
774774
timeout=timeout,
775775
)
776-
return await partition_read_method()
777776

778777
response = await _retry(
779778
attempt_tracking_method,
@@ -843,20 +842,20 @@ async def partition_query(
843842
MetricsCapture(self._resource_info),
844843
):
845844
nth_request = getattr(database, "_next_nth_request", 0)
846-
attempt = AtomicCounter()
845+
attempt = 0
847846

848847
async def attempt_tracking_method():
848+
nonlocal attempt
849+
attempt += 1
849850
all_metadata = database.metadata_with_request_id(
850-
nth_request, attempt.increment(), metadata, span
851+
nth_request, attempt, metadata, span
851852
)
852-
partition_query_method = functools.partial(
853-
api.partition_query,
853+
return await api.partition_query(
854854
request=partition_query_request,
855855
metadata=all_metadata,
856856
retry=retry,
857857
timeout=timeout,
858858
)
859-
return await partition_query_method()
860859

861860
response = await _retry(
862861
attempt_tracking_method,
@@ -917,22 +916,22 @@ async def _begin_transaction(
917916
MetricsCapture(self._resource_info),
918917
):
919918
nth_request = getattr(database, "_next_nth_request", 0)
920-
attempt = AtomicCounter()
919+
attempt = 0
921920

922921
async def wrapped_method():
922+
nonlocal attempt
923+
attempt += 1
923924
begin_transaction_request = BeginTransactionRequest(
924925
**begin_request_kwargs
925926
)
926927
call_metadata, error_augmenter = database.with_error_augmentation(
927-
nth_request, attempt.increment(), metadata, span
928-
)
929-
begin_transaction_method = functools.partial(
930-
api.begin_transaction,
931-
request=begin_transaction_request,
932-
metadata=call_metadata,
928+
nth_request, attempt, metadata, span
933929
)
934930
with error_augmenter:
935-
return await begin_transaction_method()
931+
return await api.begin_transaction(
932+
request=begin_transaction_request,
933+
metadata=call_metadata,
934+
)
936935

937936
async def before_next_retry(nth_retry, delay_in_seconds):
938937
add_span_event(

0 commit comments

Comments
 (0)