Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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.

Expand All @@ -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.

Expand All @@ -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__):
Expand Down Expand Up @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
Loading
Loading