diff --git a/packages/google-cloud-storage/google/cloud/storage/_bucket_metadata_cache.py b/packages/google-cloud-storage/google/cloud/storage/_bucket_metadata_cache.py index a8bc1759294e..c932a25fb29a 100644 --- a/packages/google-cloud-storage/google/cloud/storage/_bucket_metadata_cache.py +++ b/packages/google-cloud-storage/google/cloud/storage/_bucket_metadata_cache.py @@ -14,6 +14,7 @@ """In-memory LRU cache for bucket metadata supporting App-centric Observability (ACO).""" +import asyncio import logging import threading @@ -58,11 +59,18 @@ def get_or_queue_fetch(self, bucket_name): # bypass starting duplicate fetches. return None else: - # fire a background thread and get bucket metadata. + # fire a background fetch and get bucket metadata. self._inflight_fetches.add(bucket_name) - threading.Thread( - target=self._fetch_background, args=(bucket_name,), daemon=True - ).start() + if getattr(self._client, "_is_async_grpc_client", False) is True: + try: + loop = asyncio.get_running_loop() + loop.create_task(self._fetch_background_async(bucket_name)) + except RuntimeError: + self._inflight_fetches.discard(bucket_name) + else: + threading.Thread( + target=self._fetch_background, args=(bucket_name,), daemon=True + ).start() return None def check_and_evict(self, bucket_name): @@ -73,11 +81,20 @@ def check_and_evict(self, bucket_name): if bucket_name in self._inflight_checks: return self._inflight_checks.add(bucket_name) - threading.Thread( - target=self._verify_existence_background, - args=(bucket_name,), - daemon=True, - ).start() + if getattr(self._client, "_is_async_grpc_client", False) is True: + try: + loop = asyncio.get_running_loop() + loop.create_task( + self._verify_existence_background_async(bucket_name) + ) + except RuntimeError: + self._inflight_checks.discard(bucket_name) + else: + threading.Thread( + target=self._verify_existence_background, + args=(bucket_name,), + daemon=True, + ).start() def _verify_existence_background(self, bucket_name): try: @@ -92,6 +109,24 @@ def _verify_existence_background(self, bucket_name): with self._lock: self._inflight_checks.discard(bucket_name) + async def _verify_existence_background_async(self, bucket_name): + try: + from google.cloud import _storage_v2 as storage_v2 + + request = storage_v2.GetBucketRequest( + name=f"projects/_/buckets/{bucket_name}" + ) + await self._client.grpc_client.get_bucket(request=request, timeout=10.0) + except (NotFound, api_exceptions.NotFound): + self.evict(bucket_name) + except Exception as e: + logger.debug( + f"Async background verification for bucket existence failed for {bucket_name}: {e}" + ) + finally: + with self._lock: + self._inflight_checks.discard(bucket_name) + def _fetch_background(self, bucket_name): """Asynchronously fetch bucket metadata and update the cache.""" try: @@ -112,26 +147,72 @@ def _fetch_background(self, bucket_name): with self._lock: self._inflight_fetches.discard(bucket_name) - def update_from_bucket(self, bucket): - """Update cache from a Bucket instance.""" - if not bucket or not bucket.name: + async def _fetch_background_async(self, bucket_name): + """Asynchronously fetch bucket metadata via gRPC and update the cache.""" + try: + from google.cloud import _storage_v2 as storage_v2 + + request = storage_v2.GetBucketRequest( + name=f"projects/_/buckets/{bucket_name}" + ) + bucket = await self._client.grpc_client.get_bucket( + request=request, timeout=10.0 + ) + self.update_from_bucket(bucket, bucket_name=bucket_name) + except (NotFound, api_exceptions.NotFound): + self.evict(bucket_name) + except api_exceptions.Forbidden: + self.update_cache( + bucket_name, f"projects/_/buckets/{bucket_name}", "global" + ) + except Exception as e: + logger.debug( + f"Async background fetch for bucket metadata failed for {bucket_name}: {e}" + ) + finally: + with self._lock: + self._inflight_fetches.discard(bucket_name) + + def update_from_bucket(self, bucket, bucket_name=None): + """Update cache from a Bucket instance or storage_v2.Bucket proto.""" + if not bucket: + return + name = bucket_name or getattr(bucket, "name", None) + if not name or not isinstance(name, str): return + if name.startswith("projects/") and "/buckets/" in name: + name = name.split("/buckets/", 1)[1] project_number = getattr(bucket, "project_number", None) - location = getattr(bucket, "location", None) or "global" - location = location.lower() - location_type = getattr(bucket, "location_type", None) or "region" - location_type = location_type.lower() + if not project_number: + proj_attr = getattr(bucket, "project", None) + if isinstance(proj_attr, (str, int)): + proj_str = str(proj_attr) + if proj_str.startswith("projects/"): + project_number = proj_str.split("projects/", 1)[1] + elif proj_str: + project_number = proj_str + + loc_attr = getattr(bucket, "location", None) + location = ( + loc_attr.lower() if isinstance(loc_attr, str) and loc_attr else "global" + ) + loc_type_attr = getattr(bucket, "location_type", None) + location_type = ( + loc_type_attr.lower() + if isinstance(loc_type_attr, str) and loc_type_attr + else "region" + ) if location_type in ("multi-region", "dual-region"): location = "global" - if project_number: - destination_id = f"projects/{project_number}/buckets/{bucket.name}" + if project_number and str(project_number) != "_": + destination_id = f"projects/{project_number}/buckets/{name}" else: - destination_id = f"projects/_/buckets/{bucket.name}" + destination_id = f"projects/_/buckets/{name}" - self.update_cache(bucket.name, destination_id, location) + self.update_cache(name, destination_id, location) def update_cache(self, bucket_name, destination_id, location): """Thread-safely update or insert a cache entry with bounded size.""" diff --git a/packages/google-cloud-storage/google/cloud/storage/_helpers.py b/packages/google-cloud-storage/google/cloud/storage/_helpers.py index 4e9d435347d8..6cd72f722827 100644 --- a/packages/google-cloud-storage/google/cloud/storage/_helpers.py +++ b/packages/google-cloud-storage/google/cloud/storage/_helpers.py @@ -19,6 +19,7 @@ import base64 import datetime +import inspect import logging import os import secrets @@ -32,18 +33,15 @@ from google.auth import environment_vars from google.cloud.exceptions import NotFound -from google.cloud.storage._opentelemetry_tracing import ( - _is_bucket_metadata_disabled, -) -from google.cloud.storage._opentelemetry_tracing import ( - create_trace_span as _base_create_trace_span, -) +from google.cloud.storage import _opentelemetry_tracing from google.cloud.storage.constants import _DEFAULT_TIMEOUT from google.cloud.storage.retry import ( DEFAULT_RETRY, DEFAULT_RETRY_IF_METAGENERATION_SPECIFIED, ) +_base_create_trace_span = None + _logger = logging.getLogger(__name__) STORAGE_EMULATOR_ENV_VAR = "STORAGE_EMULATOR_HOST" # Despite name, includes scheme. @@ -149,61 +147,121 @@ def _validate_name(name): return name -@contextmanager -def create_trace_span_helper(client, bucket_name, name, attributes=None, **kwargs): - span_attrs = dict(attributes) if attributes else {} - - if ( - bucket_name - and isinstance(bucket_name, str) - and client - and hasattr(client, "_bucket_metadata_cache") - and client._bucket_metadata_cache - and not _is_bucket_metadata_disabled() - ): - try: - if name in ( - "Storage.Client.getBucket", - "Storage.Client.lookupBucket", - "Storage.Bucket.reload", - "Storage.Bucket.exists", - ): - cached = client._bucket_metadata_cache.get(bucket_name) - else: - cached = client._bucket_metadata_cache.get_or_queue_fetch(bucket_name) - - if cached and isinstance(cached, tuple) and len(cached) == 2: - dest_id, loc = cached - span_attrs.update( - { - "gcp.resource.destination.id": dest_id, - "gcp.resource.destination.location": loc, - } - ) - except Exception as e: - _logger.debug(f"Failed cache lookup in create_trace_span_helper: {e}") +class _TraceSpanHelperContext: + """Context manager supporting both sync and async tracing span creation with bucket metadata.""" - if "client" not in kwargs and client: - kwargs["client"] = client + def __init__(self, client, bucket_name, name, attributes=None, **kwargs): + self.client = client + self.bucket_name = bucket_name + self.name = name + self.attributes = attributes + self.kwargs = kwargs + self._base_cm = None + + def _prepare_base_cm(self): + span_attrs = dict(self.attributes) if self.attributes else {} + client = self.client + bucket_name = self.bucket_name + name = self.name + + if ( + bucket_name + and isinstance(bucket_name, str) + and client + and hasattr(client, "_bucket_metadata_cache") + and client._bucket_metadata_cache + and _opentelemetry_tracing._is_otel_traces_enabled() + and not _opentelemetry_tracing._is_bucket_metadata_disabled() + ): + try: + if name in ( + "Storage.Client.getBucket", + "Storage.Client.lookupBucket", + "Storage.Bucket.reload", + "Storage.Bucket.exists", + ): + cached = client._bucket_metadata_cache.get(bucket_name) + else: + cached = client._bucket_metadata_cache.get_or_queue_fetch( + bucket_name + ) - with _base_create_trace_span(name, attributes=span_attrs, **kwargs) as span: - try: - yield span - except (NotFound, api_exceptions.NotFound): - if ( - bucket_name - and isinstance(bucket_name, str) - and client - and hasattr(client, "_bucket_metadata_cache") - and client._bucket_metadata_cache - ): - try: - client._bucket_metadata_cache.check_and_evict(bucket_name) - except Exception as e: - _logger.debug( - f"Failed cache eviction on 404 in create_trace_span_helper: {e}" + if cached and isinstance(cached, tuple) and len(cached) == 2: + dest_id, loc = cached + span_attrs.update( + { + "gcp.resource.destination.id": dest_id, + "gcp.resource.destination.location": loc, + } ) - raise + except Exception as e: + _logger.debug(f"Failed cache lookup in create_trace_span_helper: {e}") + + kwargs = dict(self.kwargs) + if "client" not in kwargs and client: + kwargs["client"] = client + + create_span_fn = ( + _base_create_trace_span + if _base_create_trace_span is not None + else _opentelemetry_tracing.create_trace_span + ) + self._base_cm = create_span_fn(name, attributes=span_attrs, **kwargs) + return self._base_cm + + def _handle_not_found(self): + if ( + self.bucket_name + and isinstance(self.bucket_name, str) + and self.client + and hasattr(self.client, "_bucket_metadata_cache") + and self.client._bucket_metadata_cache + and _opentelemetry_tracing._is_otel_traces_enabled() + ): + try: + self.client._bucket_metadata_cache.check_and_evict(self.bucket_name) + except Exception as e: + _logger.debug( + f"Failed cache eviction on 404 in create_trace_span_helper: {e}" + ) + + def __enter__(self): + self._prepare_base_cm() + return self._base_cm.__enter__() + + def __exit__(self, exc_type, exc_val, exc_tb): + if exc_val is not None and isinstance( + exc_val, (NotFound, api_exceptions.NotFound) + ): + self._handle_not_found() + if self._base_cm is not None: + return self._base_cm.__exit__(exc_type, exc_val, exc_tb) + return False + + async def __aenter__(self): + self._prepare_base_cm() + if hasattr(self._base_cm, "__aenter__"): + res = self._base_cm.__aenter__() + return await res if inspect.isawaitable(res) else res + return self._base_cm.__enter__() + + async def __aexit__(self, exc_type, exc_val, exc_tb): + if exc_val is not None and isinstance( + exc_val, (NotFound, api_exceptions.NotFound) + ): + self._handle_not_found() + if self._base_cm is not None: + if hasattr(self._base_cm, "__aexit__"): + res = self._base_cm.__aexit__(exc_type, exc_val, exc_tb) + return await res if inspect.isawaitable(res) else res + return self._base_cm.__exit__(exc_type, exc_val, exc_tb) + return False + + +def create_trace_span_helper(client, bucket_name, name, attributes=None, **kwargs): + return _TraceSpanHelperContext( + client, bucket_name, name, attributes=attributes, **kwargs + ) class _PropertyMixin(object): diff --git a/packages/google-cloud-storage/google/cloud/storage/_opentelemetry_tracing.py b/packages/google-cloud-storage/google/cloud/storage/_opentelemetry_tracing.py index 1d9e4b88270b..aead5b5e82f3 100644 --- a/packages/google-cloud-storage/google/cloud/storage/_opentelemetry_tracing.py +++ b/packages/google-cloud-storage/google/cloud/storage/_opentelemetry_tracing.py @@ -16,7 +16,6 @@ import logging import os -from contextlib import contextmanager from urllib.parse import urlparse from google.api_core import exceptions as api_exceptions @@ -44,6 +43,16 @@ def _is_bucket_metadata_disabled() -> bool: enable_otel_traces = _parse_bool_env( ENABLE_OTEL_TRACES_ENV_VAR, _DEFAULT_ENABLE_OTEL_TRACES_VALUE ) + + +def _is_otel_traces_enabled() -> bool: + if not HAS_OPENTELEMETRY or not enable_otel_traces: + return False + return _parse_bool_env( + ENABLE_OTEL_TRACES_ENV_VAR, _DEFAULT_ENABLE_OTEL_TRACES_VALUE + ) + + logger = logging.getLogger(__name__) @@ -75,30 +84,88 @@ def _is_bucket_metadata_disabled() -> bool: } -@contextmanager -def create_trace_span(name, attributes=None, client=None, api_request=None, retry=None): +class _TraceSpanContext: + """Context manager supporting both sync and async tracing spans.""" + + def __init__( + self, + name, + attributes=None, + client=None, + api_request=None, + retry=None, + rpc_system="http", + ): + self.name = name + self.attributes = attributes + self.client = client + self.api_request = api_request + self.retry = retry + self.rpc_system = rpc_system + self._span_cm = None + self._span = None + + def __enter__(self): + if not _is_otel_traces_enabled(): + return None + + tracer = trace.get_tracer(__name__) + final_attributes = _get_final_attributes( + self.attributes, + self.client, + self.api_request, + self.retry, + rpc_system=self.rpc_system, + ) + self._span_cm = tracer.start_as_current_span( + name=self.name, kind=trace.SpanKind.CLIENT, attributes=final_attributes + ) + self._span = self._span_cm.__enter__() + return self._span + + def __exit__(self, exc_type, exc_val, exc_tb): + if self._span_cm is not None: + if exc_val is not None and isinstance( + exc_val, api_exceptions.GoogleAPICallError + ): + self._span.set_status(trace.Status(trace.StatusCode.ERROR)) + self._span.record_exception(exc_val) + return self._span_cm.__exit__(exc_type, exc_val, exc_tb) + return False + + async def __aenter__(self): + return self.__enter__() + + async def __aexit__(self, exc_type, exc_val, exc_tb): + return self.__exit__(exc_type, exc_val, exc_tb) + + +def create_trace_span( + name, + attributes=None, + client=None, + api_request=None, + retry=None, + rpc_system="http", +): """Creates a context manager for a new span and set it as the current span - in the configured tracer. If no configuration exists yields None.""" - if not HAS_OPENTELEMETRY or not enable_otel_traces: - yield None - return - - tracer = trace.get_tracer(__name__) - final_attributes = _get_final_attributes(attributes, client, api_request, retry) - # Yield new span. - with tracer.start_as_current_span( - name=name, kind=trace.SpanKind.CLIENT, attributes=final_attributes - ) as span: - try: - yield span - except api_exceptions.GoogleAPICallError as error: - span.set_status(trace.Status(trace.StatusCode.ERROR)) - span.record_exception(error) - raise - - -def _get_final_attributes(attributes=None, client=None, api_request=None, retry=None): + in the configured tracer. Supports both sync and async context managers. + If no configuration exists yields None.""" + return _TraceSpanContext( + name=name, + attributes=attributes, + client=client, + api_request=api_request, + retry=retry, + rpc_system=rpc_system, + ) + + +def _get_final_attributes( + attributes=None, client=None, api_request=None, retry=None, rpc_system="http" +): collected_attr = _default_attributes.copy() + collected_attr["rpc.system"] = rpc_system collected_attr.update(_cloud_trace_adoption_attrs) if api_request: collected_attr.update(_set_api_request_attr(api_request, client)) diff --git a/packages/google-cloud-storage/google/cloud/storage/asyncio/_utils.py b/packages/google-cloud-storage/google/cloud/storage/asyncio/_utils.py index 53e8503c7f8b..e2c066f0608e 100644 --- a/packages/google-cloud-storage/google/cloud/storage/asyncio/_utils.py +++ b/packages/google-cloud-storage/google/cloud/storage/asyncio/_utils.py @@ -15,6 +15,8 @@ import google_crc32c from google.api_core import exceptions +from google.cloud.storage import _opentelemetry_tracing + def raise_if_no_fast_crc32c(): """Check if the C-accelerated version of google-crc32c is available. @@ -38,3 +40,35 @@ def update_write_handle_if_exists(obj, response): """Update the write_handle attribute of an object if it exists in the response.""" if hasattr(response, "write_handle") and response.write_handle is not None: obj.write_handle = response.write_handle + + +def inject_traceparent_to_metadata(metadata=None): + """Inject W3C traceparent (and tracestate) into gRPC metadata tuple or list.""" + meta_list = list(metadata) if metadata else [] + if not _opentelemetry_tracing._is_otel_traces_enabled(): + return ( + tuple(meta_list) + if isinstance(metadata, tuple) or metadata is None + else meta_list + ) + + try: + from opentelemetry.trace.propagation.tracecontext import ( + TraceContextTextMapPropagator, + ) + + carrier = {} + TraceContextTextMapPropagator().inject(carrier) + existing_keys = {k.lower() for k, _ in meta_list} + for key, val in carrier.items(): + if key.lower() not in existing_keys: + meta_list.append((key, val)) + except Exception: + pass + + return ( + tuple(meta_list) + if isinstance(metadata, tuple) or metadata is None + else meta_list + ) + diff --git a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_appendable_object_writer.py b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_appendable_object_writer.py index 7ef8ef099e6a..73c13914b1e5 100644 --- a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_appendable_object_writer.py +++ b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_appendable_object_writer.py @@ -25,6 +25,7 @@ from google.cloud._storage_v2.types import BidiWriteObjectRedirectedError from google.cloud._storage_v2.types.storage import BidiWriteObjectRequest from google.cloud.storage import Blob +from google.cloud.storage._helpers import create_trace_span_helper from google.cloud.storage.asyncio.async_grpc_client import ( AsyncGrpcClient, ) @@ -336,55 +337,64 @@ async def open( if self._is_stream_open: raise ValueError("Underlying bidi-gRPC stream is already open") - retry_policy = self._merge_retry_policy(retry_policy) + async with create_trace_span_helper( + self.client, + self.bucket_name, + "Storage.AsyncAppendableObjectWriter.open", + rpc_system="grpc", + ): + retry_policy = self._merge_retry_policy(retry_policy) + + async def _do_open(): + current_metadata = list(metadata) if metadata else [] + + # Cleanup stream from previous failed attempt, if any. + if self.write_obj_stream: + if self.write_obj_stream.is_stream_open: + try: + await self.write_obj_stream.close() + except Exception as e: + logger.warning( + f"Error closing previous write stream during open retry. Got exception: {e}" + ) + self.write_obj_stream = None + self._is_stream_open = False - async def _do_open(): - current_metadata = list(metadata) if metadata else [] + self.write_obj_stream = _AsyncWriteObjectStream( + client=self.client.grpc_client, + bucket_name=self.bucket_name, + object_name=self.object_name, + blob=self.blob, + generation_number=self.generation, + write_handle=self.write_handle, + routing_token=self._routing_token, + ) - # Cleanup stream from previous failed attempt, if any. - if self.write_obj_stream: - if self.write_obj_stream.is_stream_open: - try: - await self.write_obj_stream.close() - except Exception as e: - logger.warning( - f"Error closing previous write stream during open retry. Got exception: {e}" + if self._routing_token: + current_metadata.append( + ( + "x-goog-request-params", + f"routing_token={self._routing_token}", ) - self.write_obj_stream = None - self._is_stream_open = False - - self.write_obj_stream = _AsyncWriteObjectStream( - client=self.client.grpc_client, - bucket_name=self.bucket_name, - object_name=self.object_name, - blob=self.blob, - generation_number=self.generation, - write_handle=self.write_handle, - routing_token=self._routing_token, - ) + ) - if self._routing_token: - current_metadata.append( - ("x-goog-request-params", f"routing_token={self._routing_token}") + await self.write_obj_stream.open( + metadata=current_metadata if current_metadata else None ) - await self.write_obj_stream.open( - metadata=current_metadata if current_metadata else None - ) - - if self.write_obj_stream.generation_number: - self.generation = self.write_obj_stream.generation_number - if self.write_obj_stream.write_handle: - self.write_handle = self.write_obj_stream.write_handle - if self.write_obj_stream.persisted_size is not None: - self.persisted_size = self.write_obj_stream.persisted_size - # set offset while opening - self.offset = self.persisted_size + if self.write_obj_stream.generation_number: + self.generation = self.write_obj_stream.generation_number + if self.write_obj_stream.write_handle: + self.write_handle = self.write_obj_stream.write_handle + if self.write_obj_stream.persisted_size is not None: + self.persisted_size = self.write_obj_stream.persisted_size + # set offset while opening + self.offset = self.persisted_size - self._is_stream_open = True - self._routing_token = None + self._is_stream_open = True + self._routing_token = None - await retry_policy(_do_open)() + await retry_policy(_do_open)() async def append( self, @@ -423,100 +433,109 @@ async def append( logger.debug("No data provided to append; returning without action.") return - if retry_policy is None: - retry_policy = AsyncRetry(predicate=_is_write_retryable) - - strategy = _WriteResumptionStrategy() - buffer = io.BytesIO(data) - attempt_count = 0 - - def send_and_recv_generator( - requests: List[BidiWriteObjectRequest], - state: Dict[str, _WriteState], - metadata: Optional[List[Tuple[str, str]]] = None, + async with create_trace_span_helper( + self.client, + self.bucket_name, + "Storage.AsyncAppendableObjectWriter.append", + attributes={"gcp.storage.chunk.size": len(data)}, + rpc_system="grpc", ): - async def generator(): - nonlocal attempt_count - nonlocal requests - attempt_count += 1 - resp = None - write_state = state["write_state"] - # If this is a retry or redirect, we must re-open the stream - if attempt_count > 1 or write_state.routing_token: - logger.info( - f"Re-opening the stream with attempt_count: {attempt_count}" - ) - - current_metadata = list(metadata) if metadata else [] - if write_state.routing_token: - current_metadata.append( - ( - "x-goog-request-params", - f"routing_token={write_state.routing_token}", - ) + if retry_policy is None: + retry_policy = AsyncRetry(predicate=_is_write_retryable) + + strategy = _WriteResumptionStrategy() + buffer = io.BytesIO(data) + attempt_count = 0 + + def send_and_recv_generator( + requests: List[BidiWriteObjectRequest], + state: Dict[str, _WriteState], + metadata: Optional[List[Tuple[str, str]]] = None, + ): + async def generator(): + nonlocal attempt_count + nonlocal requests + attempt_count += 1 + resp = None + write_state = state["write_state"] + # If this is a retry or redirect, we must re-open the stream + if attempt_count > 1 or write_state.routing_token: + logger.info( + f"Re-opening the stream with attempt_count: {attempt_count}" ) - self._routing_token = write_state.routing_token - - self._is_stream_open = False - await self.open(metadata=current_metadata) - write_state.persisted_size = self.persisted_size - write_state.write_handle = self.write_handle - write_state.routing_token = None - - write_state.user_buffer.seek(write_state.persisted_size) - write_state.bytes_sent = write_state.persisted_size - write_state.bytes_since_last_flush = 0 - self.bytes_appended_since_last_flush = 0 - - requests = strategy.generate_requests(state) - - for chunk_req in requests: - await self.write_obj_stream.send(chunk_req) - if chunk_req.flush: - self._flush_count += 1 - - resp = None - if chunk_req.state_lookup: - # TODO: if there's error, it'll raise error - # and will be handled by `recover_state_on_failure` - resp = await self.write_obj_stream.recv() - - if resp: - if resp.persisted_size is not None: - self.persisted_size = resp.persisted_size - state["write_state"].persisted_size = resp.persisted_size - self.offset = self.persisted_size - if resp.write_handle: - self.write_handle = resp.write_handle - state["write_state"].write_handle = resp.write_handle - - yield resp - - return generator() - - # State initialization - write_state = _WriteState( - _MAX_CHUNK_SIZE_BYTES, - buffer, - self.flush_interval, - enable_checksum=enable_checksum, - ) - write_state.write_handle = self.write_handle - write_state.persisted_size = self.persisted_size - # offset is set during `open()` call. - write_state.bytes_sent = self.offset or 0 - write_state.bytes_since_last_flush = self.bytes_appended_since_last_flush - - retry_manager = _BidiStreamRetryManager( - _WriteResumptionStrategy(), - lambda r, s: send_and_recv_generator(r, s, metadata), - ) - await retry_manager.execute({"write_state": write_state}, retry_policy) + current_metadata = list(metadata) if metadata else [] + if write_state.routing_token: + current_metadata.append( + ( + "x-goog-request-params", + f"routing_token={write_state.routing_token}", + ) + ) + self._routing_token = write_state.routing_token + + self._is_stream_open = False + await self.open(metadata=current_metadata) + + write_state.persisted_size = self.persisted_size + write_state.write_handle = self.write_handle + write_state.routing_token = None + + write_state.user_buffer.seek(write_state.persisted_size) + write_state.bytes_sent = write_state.persisted_size + write_state.bytes_since_last_flush = 0 + self.bytes_appended_since_last_flush = 0 + + requests = strategy.generate_requests(state) + + for chunk_req in requests: + await self.write_obj_stream.send(chunk_req) + if chunk_req.flush: + self._flush_count += 1 + + resp = None + if chunk_req.state_lookup: + # TODO: if there's error, it'll raise error + # and will be handled by `recover_state_on_failure` + resp = await self.write_obj_stream.recv() + + if resp: + if resp.persisted_size is not None: + self.persisted_size = resp.persisted_size + state["write_state"].persisted_size = ( + resp.persisted_size + ) + self.offset = self.persisted_size + if resp.write_handle: + self.write_handle = resp.write_handle + state["write_state"].write_handle = resp.write_handle + + yield resp + + return generator() + + # State initialization + write_state = _WriteState( + _MAX_CHUNK_SIZE_BYTES, + buffer, + self.flush_interval, + enable_checksum=enable_checksum, + ) + write_state.write_handle = self.write_handle + write_state.persisted_size = self.persisted_size + # offset is set during `open()` call. + write_state.bytes_sent = self.offset or 0 + write_state.bytes_since_last_flush = self.bytes_appended_since_last_flush + + retry_manager = _BidiStreamRetryManager( + _WriteResumptionStrategy(), + lambda r, s: send_and_recv_generator(r, s, metadata), + ) + await retry_manager.execute({"write_state": write_state}, retry_policy) - # Sync local markers - self.bytes_appended_since_last_flush = write_state.bytes_since_last_flush - self.offset = write_state.bytes_sent + # Sync local markers + self.bytes_appended_since_last_flush = write_state.bytes_since_last_flush + self.offset = write_state.bytes_sent async def simple_flush(self) -> None: """Flushes the data to the server. @@ -549,17 +568,23 @@ async def flush(self) -> int: if not self._is_stream_open: raise ValueError("Stream is not open. Call open() before flush().") - await self.write_obj_stream.send( - _storage_v2.BidiWriteObjectRequest( - flush=True, - state_lookup=True, + async with create_trace_span_helper( + self.client, + self.bucket_name, + "Storage.AsyncAppendableObjectWriter.flush", + rpc_system="grpc", + ): + await self.write_obj_stream.send( + _storage_v2.BidiWriteObjectRequest( + flush=True, + state_lookup=True, + ) ) - ) - response = await self.write_obj_stream.recv() - self.persisted_size = response.persisted_size - self.offset = self.persisted_size - self.bytes_appended_since_last_flush = 0 - return self.persisted_size + response = await self.write_obj_stream.recv() + self.persisted_size = response.persisted_size + self.offset = self.persisted_size + self.bytes_appended_since_last_flush = 0 + return self.persisted_size async def close( self, @@ -612,43 +637,49 @@ async def close( "full_object_checksum can only be provided when finalize_on_close is True." ) - if finalize_on_close: - return await self.finalize( - full_object_checksum=full_object_checksum, - retry_policy=retry_policy, - ) + async with create_trace_span_helper( + self.client, + self.bucket_name, + "Storage.AsyncAppendableObjectWriter.close", + rpc_system="grpc", + ): + if finalize_on_close: + return await self.finalize( + full_object_checksum=full_object_checksum, + retry_policy=retry_policy, + ) - retry_policy = self._merge_retry_policy(retry_policy) + retry_policy = self._merge_retry_policy(retry_policy) - attempt_count = 0 - expected_offset = self.offset + attempt_count = 0 + expected_offset = self.offset - async def _do_close(): - nonlocal attempt_count - attempt_count += 1 + async def _do_close(): + nonlocal attempt_count + attempt_count += 1 - if attempt_count > 1: - logger.info( - f"Re-opening the stream for close retry attempt: {attempt_count}" - ) - self._is_stream_open = False - await self.open() - if ( - self.offset is not None - and expected_offset is not None - and self.offset != expected_offset - ): - raise exceptions.InternalServerError( - f"Unrecoverable data loss during reconnect. Expected offset {expected_offset}, got {self.offset}" + if attempt_count > 1: + logger.info( + f"Re-opening the stream for close retry attempt: {attempt_count}" ) + self._is_stream_open = False + await self.open() + if ( + self.offset is not None + and expected_offset is not None + and self.offset != expected_offset + ): + raise exceptions.InternalServerError( + f"Unrecoverable data loss during reconnect. Expected offset {expected_offset}, got {self.offset}" + ) - await self.write_obj_stream.close() - return self.persisted_size + await self.write_obj_stream.close() + return self.persisted_size - try: - return await retry_policy(_do_close)() - finally: - self._is_stream_open = False + try: + return await retry_policy(_do_close)() + finally: + self._is_stream_open = False async def finalize( self, diff --git a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_grpc_client.py b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_grpc_client.py index 18ee348401c7..5d75141e5b72 100644 --- a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_grpc_client.py +++ b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_grpc_client.py @@ -22,6 +22,9 @@ DEFAULT_CLIENT_INFO, ) from google.cloud.storage import __version__ +from google.cloud.storage._bucket_metadata_cache import BucketMetadataCache +from google.cloud.storage._helpers import create_trace_span_helper +from google.cloud.storage.asyncio import _utils _DEFAULT_HOST = "storage.googleapis.com" @@ -73,6 +76,8 @@ def __init__( *, attempt_direct_path=True, ): + self._is_async_grpc_client = True + self._bucket_metadata_cache = BucketMetadataCache(self) if isinstance(credentials, auth_credentials.AnonymousCredentials): if client_options is None or client_options.api_endpoint is None: raise ValueError( @@ -212,24 +217,31 @@ async def delete_object( :param retry: (Optional) Designation of what errors, if any, should be retried. """ _validate_metadata(metadata) - # The gRPC API requires the bucket name to be in the format "projects/_/buckets/bucket_name" - bucket_path = f"projects/_/buckets/{bucket_name}" - request = storage_v2.DeleteObjectRequest( - bucket=bucket_path, - object=object_name, - generation=generation, - if_generation_match=if_generation_match, - if_generation_not_match=if_generation_not_match, - if_metageneration_match=if_metageneration_match, - if_metageneration_not_match=if_metageneration_not_match, - **kwargs, - ) - await self._grpc_client.delete_object( - request=request, - metadata=metadata, - timeout=timeout, - retry=retry, - ) + async with create_trace_span_helper( + self, + bucket_name, + "Storage.AsyncGrpcClient.deleteObject", + rpc_system="grpc", + ): + # The gRPC API requires the bucket name to be in the format "projects/_/buckets/bucket_name" + bucket_path = f"projects/_/buckets/{bucket_name}" + request = storage_v2.DeleteObjectRequest( + bucket=bucket_path, + object=object_name, + generation=generation, + if_generation_match=if_generation_match, + if_generation_not_match=if_generation_not_match, + if_metageneration_match=if_metageneration_match, + if_metageneration_not_match=if_metageneration_not_match, + **kwargs, + ) + final_metadata = _utils.inject_traceparent_to_metadata(metadata) + await self._grpc_client.delete_object( + request=request, + metadata=final_metadata, + timeout=timeout, + retry=retry, + ) async def get_object( self, @@ -292,24 +304,31 @@ async def get_object( :returns: The object metadata resource. """ _validate_metadata(metadata) - bucket_path = f"projects/_/buckets/{bucket_name}" - - request = storage_v2.GetObjectRequest( - bucket=bucket_path, - object=object_name, - generation=generation, - if_generation_match=if_generation_match, - if_generation_not_match=if_generation_not_match, - if_metageneration_match=if_metageneration_match, - if_metageneration_not_match=if_metageneration_not_match, - soft_deleted=soft_deleted or False, - **kwargs, - ) + async with create_trace_span_helper( + self, + bucket_name, + "Storage.AsyncGrpcClient.getObject", + rpc_system="grpc", + ): + bucket_path = f"projects/_/buckets/{bucket_name}" + + request = storage_v2.GetObjectRequest( + bucket=bucket_path, + object=object_name, + generation=generation, + if_generation_match=if_generation_match, + if_generation_not_match=if_generation_not_match, + if_metageneration_match=if_metageneration_match, + if_metageneration_not_match=if_metageneration_not_match, + soft_deleted=soft_deleted or False, + **kwargs, + ) - # Calls the underlying GAPIC StorageAsyncClient.get_object method - return await self._grpc_client.get_object( - request=request, - metadata=metadata, - timeout=timeout, - retry=retry, - ) + final_metadata = _utils.inject_traceparent_to_metadata(metadata) + # Calls the underlying GAPIC StorageAsyncClient.get_object method + return await self._grpc_client.get_object( + request=request, + metadata=final_metadata, + timeout=timeout, + retry=retry, + ) diff --git a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_multi_range_downloader.py b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_multi_range_downloader.py index 32757bf754b7..ac3ac712ecd8 100644 --- a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_multi_range_downloader.py +++ b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_multi_range_downloader.py @@ -23,7 +23,10 @@ from google.api_core.retry_async import AsyncRetry from google.cloud import _storage_v2 -from google.cloud.storage._helpers import generate_random_56_bit_integer +from google.cloud.storage._helpers import ( + create_trace_span_helper, + generate_random_56_bit_integer, +) from google.cloud.storage.asyncio._stream_multiplexer import ( _StreamEnd, _StreamError, @@ -244,79 +247,90 @@ async def open( if self._is_stream_open: raise ValueError("Underlying bidi-gRPC stream is already open") - if retry_policy is None: + async with create_trace_span_helper( + self.client, + self.bucket_name, + "Storage.AsyncMultiRangeDownloader.open", + rpc_system="grpc", + ): + if retry_policy is None: - def on_error_wrapper(exc): - self._open_retries += 1 - self._on_open_error(exc) + def on_error_wrapper(exc): + self._open_retries += 1 + self._on_open_error(exc) - retry_policy = AsyncRetry( - predicate=_is_open_retryable, on_error=on_error_wrapper - ) - else: - original_on_error = retry_policy._on_error - - def combined_on_error(exc): - self._open_retries += 1 - self._on_open_error(exc) - if original_on_error: - original_on_error(exc) - - retry_policy = AsyncRetry( - predicate=_is_open_retryable, - initial=retry_policy._initial, - maximum=retry_policy._maximum, - multiplier=retry_policy._multiplier, - deadline=retry_policy._deadline, - on_error=combined_on_error, - ) + retry_policy = AsyncRetry( + predicate=_is_open_retryable, on_error=on_error_wrapper + ) + else: + original_on_error = retry_policy._on_error + + def combined_on_error(exc): + self._open_retries += 1 + self._on_open_error(exc) + if original_on_error: + original_on_error(exc) + + retry_policy = AsyncRetry( + predicate=_is_open_retryable, + initial=retry_policy._initial, + maximum=retry_policy._maximum, + multiplier=retry_policy._multiplier, + deadline=retry_policy._deadline, + on_error=combined_on_error, + ) - async def _do_open(): - current_metadata = list(metadata) if metadata else [] + async def _do_open(): + current_metadata = list(metadata) if metadata else [] - # Cleanup stream from previous failed attempt, if any. - if self.read_obj_str: - if self.read_obj_str.is_stream_open: - try: - await self.read_obj_str.close() - except exceptions.GoogleAPICallError as e: - logger.warning( - f"Failed to close existing stream during resumption: {e}" - ) - self.read_obj_str = None - self._is_stream_open = False + # Cleanup stream from previous failed attempt, if any. + if self.read_obj_str: + if self.read_obj_str.is_stream_open: + try: + await self.read_obj_str.close() + except exceptions.GoogleAPICallError as e: + logger.warning( + f"Failed to close existing stream during resumption: {e}" + ) + self.read_obj_str = None + self._is_stream_open = False + + self.read_obj_str = _AsyncReadObjectStream( + client=self.client.grpc_client, + bucket_name=self.bucket_name, + object_name=self.object_name, + generation_number=self.generation, + read_handle=self.read_handle, + ) - self.read_obj_str = _AsyncReadObjectStream( - client=self.client.grpc_client, - bucket_name=self.bucket_name, - object_name=self.object_name, - generation_number=self.generation, - read_handle=self.read_handle, - ) + if self._routing_token: + current_metadata.append( + ( + "x-goog-request-params", + f"routing_token={self._routing_token}", + ) + ) + self._routing_token = None - if self._routing_token: - current_metadata.append( - ("x-goog-request-params", f"routing_token={self._routing_token}") + await self.read_obj_str.open( + metadata=current_metadata if current_metadata else None ) - self._routing_token = None - - await self.read_obj_str.open( - metadata=current_metadata if current_metadata else None - ) - if self.read_obj_str.generation_number: - self.generation = self.read_obj_str.generation_number - if self.read_obj_str.read_handle: - self.read_handle = self.read_obj_str.read_handle - if self.read_obj_str.persisted_size is not None: - self.persisted_size = self.read_obj_str.persisted_size - self.is_finalized = self.read_obj_str.is_finalized - self.full_obj_server_crc32c = self.read_obj_str.full_obj_server_crc32c + if self.read_obj_str.generation_number: + self.generation = self.read_obj_str.generation_number + if self.read_obj_str.read_handle: + self.read_handle = self.read_obj_str.read_handle + if self.read_obj_str.persisted_size is not None: + self.persisted_size = self.read_obj_str.persisted_size + self.is_finalized = self.read_obj_str.is_finalized + self.full_obj_server_crc32c = ( + self.read_obj_str.full_obj_server_crc32c + ) - self._is_stream_open = True + self._is_stream_open = True - await retry_policy(_do_open)() - self._multiplexer = _StreamMultiplexer(self.read_obj_str) + await retry_policy(_do_open)() + self._multiplexer = _StreamMultiplexer(self.read_obj_str) def _create_stream_factory(self, state, metadata): """Create a factory that opens a new stream with current routing state.""" @@ -412,136 +426,149 @@ async def download_ranges( if not self._is_stream_open: raise ValueError("Underlying bidi-gRPC stream is not open") - if retry_policy is None: - retry_policy = AsyncRetry(predicate=_is_read_retryable) - - # Initialize Global State for Retry Strategy - download_states = {} - for read_range in read_ranges: - read_id = generate_random_56_bit_integer() - # Unpack tuple into self-documenting variable names to improve readability. - offset, length, user_buffer = read_range - - # Heuristic to detect full object reads: - # - Implicit full object read: start offset is 0 and length is 0 (read all). - # - Explicit full object read: start offset is 0 and length matches the exact persisted size. - is_full_object_read = (offset == 0 and length == 0) or ( - self.persisted_size is not None - and offset == 0 - and length == self.persisted_size - ) - download_states[read_id] = _DownloadState( - initial_offset=offset, - initial_length=length, - user_buffer=user_buffer, - is_full_object_read=is_full_object_read, - ) - - initial_state = { - "download_states": download_states, - "read_handle": self.read_handle, - "routing_token": None, - "enable_checksum": enable_checksum, - "full_obj_server_crc32c": self.full_obj_server_crc32c, - } - - read_ids = set(download_states.keys()) - queue = self._multiplexer.register(read_ids) - - try: - attempt_count = 0 - last_broken_generation = None - - def send_and_recv_via_multiplexer( - requests: List[_storage_v2.ReadRange], - state: Dict[str, Any], - ): - async def generator(): - nonlocal attempt_count, last_broken_generation - attempt_count += 1 - - if attempt_count > 1: - logger.info( - f"Resuming download (attempt {attempt_count}) for {len(requests)} ranges." - ) + async with create_trace_span_helper( + self.client, + self.bucket_name, + "Storage.AsyncMultiRangeDownloader.downloadRanges", + attributes={"gcp.storage.range.count": len(read_ranges)}, + rpc_system="grpc", + ): + if retry_policy is None: + retry_policy = AsyncRetry(predicate=_is_read_retryable) + + # Initialize Global State for Retry Strategy + download_states = {} + for read_range in read_ranges: + read_id = generate_random_56_bit_integer() + # Unpack tuple into self-documenting variable names to improve readability. + offset, length, user_buffer = read_range + + # Heuristic to detect full object reads: + # - Implicit full object read: start offset is 0 and length is 0 (read all). + # - Explicit full object read: start offset is 0 and length matches the exact persisted size. + is_full_object_read = (offset == 0 and length == 0) or ( + self.persisted_size is not None + and offset == 0 + and length == self.persisted_size + ) + download_states[read_id] = _DownloadState( + initial_offset=offset, + initial_length=length, + user_buffer=user_buffer, + is_full_object_read=is_full_object_read, + ) - # Reopen stream if needed - should_reopen = ( - attempt_count > 1 and last_broken_generation is not None - ) or (attempt_count == 1 and metadata is not None) - if should_reopen: - broken_gen = ( - last_broken_generation - if attempt_count > 1 - else self._multiplexer.stream_generation - ) - stream_factory = self._create_stream_factory(state, metadata) - await self._multiplexer.reopen_stream( - broken_gen, stream_factory - ) + initial_state = { + "download_states": download_states, + "read_handle": self.read_handle, + "routing_token": None, + "enable_checksum": enable_checksum, + "full_obj_server_crc32c": self.full_obj_server_crc32c, + } - stream_generation = self._multiplexer.stream_generation + read_ids = set(download_states.keys()) + queue = self._multiplexer.register(read_ids) - # Send Requests - pending_read_ids = {r.read_id for r in requests} - for i in range( - 0, len(requests), _MAX_READ_RANGES_PER_BIDI_READ_REQUEST - ): - batch = requests[i : i + _MAX_READ_RANGES_PER_BIDI_READ_REQUEST] - try: - await self._multiplexer.send( - _storage_v2.BidiReadObjectRequest(read_ranges=batch) + try: + attempt_count = 0 + last_broken_generation = None + + def send_and_recv_via_multiplexer( + requests: List[_storage_v2.ReadRange], + state: Dict[str, Any], + ): + async def generator(): + nonlocal attempt_count, last_broken_generation + attempt_count += 1 + + if attempt_count > 1: + logger.info( + f"Resuming download (attempt {attempt_count}) for {len(requests)} ranges." ) - except Exception: - last_broken_generation = stream_generation - raise - # Receive Responses - while pending_read_ids: - item = await queue.get() + # Reopen stream if needed + should_reopen = ( + attempt_count > 1 and last_broken_generation is not None + ) or (attempt_count == 1 and metadata is not None) + if should_reopen: + broken_gen = ( + last_broken_generation + if attempt_count > 1 + else self._multiplexer.stream_generation + ) + stream_factory = self._create_stream_factory( + state, metadata + ) + await self._multiplexer.reopen_stream( + broken_gen, stream_factory + ) - if isinstance(item, _StreamEnd): - if pending_read_ids: - last_broken_generation = stream_generation - raise exceptions.ServiceUnavailable( - "Stream ended with pending read_ids" - ) - break - - if isinstance(item, _StreamError): - if item.generation < stream_generation: - continue # stale error, skip - last_broken_generation = item.generation - raise item.exception - - # Track completion - if item.object_data_ranges: - for data_range in item.object_data_ranges: - if data_range.range_end: - pending_read_ids.discard( - data_range.read_range.read_id + stream_generation = self._multiplexer.stream_generation + + # Send Requests + pending_read_ids = {r.read_id for r in requests} + for i in range( + 0, len(requests), _MAX_READ_RANGES_PER_BIDI_READ_REQUEST + ): + batch = requests[ + i : i + _MAX_READ_RANGES_PER_BIDI_READ_REQUEST + ] + try: + await self._multiplexer.send( + _storage_v2.BidiReadObjectRequest( + read_ranges=batch ) - yield item + ) + except Exception: + last_broken_generation = stream_generation + raise - return generator() + # Receive Responses + while pending_read_ids: + item = await queue.get() - strategy = _ReadResumptionStrategy() - retry_manager = _BidiStreamRetryManager( - strategy, send_and_recv_via_multiplexer - ) + if isinstance(item, _StreamEnd): + if pending_read_ids: + last_broken_generation = stream_generation + raise exceptions.ServiceUnavailable( + "Stream ended with pending read_ids" + ) + break + + if isinstance(item, _StreamError): + if item.generation < stream_generation: + continue # stale error, skip + last_broken_generation = item.generation + raise item.exception + + # Track completion + if item.object_data_ranges: + for data_range in item.object_data_ranges: + if data_range.range_end: + pending_read_ids.discard( + data_range.read_range.read_id + ) + yield item + + return generator() + + strategy = _ReadResumptionStrategy() + retry_manager = _BidiStreamRetryManager( + strategy, send_and_recv_via_multiplexer + ) - try: - await retry_manager.execute(initial_state, retry_policy) - except DataCorruption: - if self.is_stream_open: - await self.close() - raise - - if initial_state.get("read_handle"): - self.read_handle = initial_state["read_handle"] - finally: - if self._multiplexer is not None: - self._multiplexer.unregister(read_ids) + try: + await retry_manager.execute(initial_state, retry_policy) + except DataCorruption: + if self.is_stream_open: + await self.close() + raise + + if initial_state.get("read_handle"): + self.read_handle = initial_state["read_handle"] + finally: + if self._multiplexer is not None: + self._multiplexer.unregister(read_ids) async def close(self): """ @@ -550,17 +577,23 @@ async def close(self): if not self._is_stream_open: raise ValueError("Underlying bidi-gRPC stream is not open") - if self._multiplexer: - await self._multiplexer.close() - self._multiplexer = None + async with create_trace_span_helper( + self.client, + self.bucket_name, + "Storage.AsyncMultiRangeDownloader.close", + rpc_system="grpc", + ): + if self._multiplexer: + await self._multiplexer.close() + self._multiplexer = None - if self.read_obj_str: - try: - await self.read_obj_str.close() - except (asyncio.CancelledError, exceptions.GoogleAPICallError): - pass - self.read_obj_str = None - self._is_stream_open = False + if self.read_obj_str: + try: + await self.read_obj_str.close() + except (asyncio.CancelledError, exceptions.GoogleAPICallError): + pass + self.read_obj_str = None + self._is_stream_open = False @property def is_stream_open(self) -> bool: diff --git a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_read_object_stream.py b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_read_object_stream.py index d3c0780673c7..c1a719bbba8e 100644 --- a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_read_object_stream.py +++ b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_read_object_stream.py @@ -19,6 +19,7 @@ from google.api_core.bidi_async import AsyncBidiRpc from google.cloud import _storage_v2 +from google.cloud.storage.asyncio import _utils from google.cloud.storage.asyncio.async_abstract_object_stream import ( _AsyncAbstractObjectStream, ) @@ -124,6 +125,7 @@ async def open(self, metadata: Optional[List[Tuple[str, str]]] = None) -> None: current_metadata = other_metadata current_metadata.append(("x-goog-request-params", "&".join(request_params))) + current_metadata = _utils.inject_traceparent_to_metadata(current_metadata) self.socket_like_rpc = AsyncBidiRpc( self.rpc, diff --git a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_write_object_stream.py b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_write_object_stream.py index 48f786c7654d..369f3d280004 100644 --- a/packages/google-cloud-storage/google/cloud/storage/asyncio/async_write_object_stream.py +++ b/packages/google-cloud-storage/google/cloud/storage/asyncio/async_write_object_stream.py @@ -154,6 +154,7 @@ async def open(self, metadata: Optional[List[Tuple[str, str]]] = None) -> None: final_metadata.append((key, value)) final_metadata.append(("x-goog-request-params", "&".join(request_param_values))) + final_metadata = _utils.inject_traceparent_to_metadata(final_metadata) self.socket_like_rpc = AsyncBidiRpc( self.rpc, diff --git a/packages/google-cloud-storage/tests/unit/asyncio/test_zonal_observability.py b/packages/google-cloud-storage/tests/unit/asyncio/test_zonal_observability.py new file mode 100644 index 000000000000..50ad23ec3ada --- /dev/null +++ b/packages/google-cloud-storage/tests/unit/asyncio/test_zonal_observability.py @@ -0,0 +1,278 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for OpenTelemetry tracing in Zonal Buckets (Rapid Storage) async gRPC classes.""" + +import importlib +from io import BytesIO +from unittest import mock + +import pytest + +from google.cloud import _storage_v2 as storage_v2 +from google.cloud.storage import _opentelemetry_tracing +from google.cloud.storage.asyncio import _utils, async_grpc_client +from google.cloud.storage.asyncio.async_appendable_object_writer import ( + AsyncAppendableObjectWriter, +) +from google.cloud.storage.asyncio.async_multi_range_downloader import ( + AsyncMultiRangeDownloader, +) + + +@pytest.fixture +def exporter(monkeypatch): + """Set up OpenTelemetry InMemorySpanExporter and enable tracing.""" + try: + from opentelemetry import trace as trace_api + from opentelemetry.sdk.trace import TracerProvider, export + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + except ImportError: + pytest.skip("OpenTelemetry SDK packages are required for this test suite.") + + monkeypatch.setenv("ENABLE_GCS_PYTHON_CLIENT_OTEL_TRACES", "true") + monkeypatch.delenv("DISABLE_GCS_PYTHON_CLIENT_OTEL_BUCKET_METADATA", raising=False) + importlib.reload(_opentelemetry_tracing) + + if hasattr(trace_api, "_TRACER_PROVIDER_SET_ONCE"): + trace_api._TRACER_PROVIDER_SET_ONCE._done = False + trace_api._TRACER_PROVIDER = None + + tracer_provider = TracerProvider() + memory_exporter = InMemorySpanExporter() + span_processor = export.SimpleSpanProcessor(memory_exporter) + tracer_provider.add_span_processor(span_processor) + trace_api.set_tracer_provider(tracer_provider) + + yield memory_exporter + memory_exporter.clear() + monkeypatch.setenv("ENABLE_GCS_PYTHON_CLIENT_OTEL_TRACES", "false") + importlib.reload(_opentelemetry_tracing) + + +@pytest.fixture +def mock_client(): + """Create an AsyncGrpcClient with a mocked underlying GAPIC client.""" + with mock.patch("google.cloud._storage_v2.StorageAsyncClient"): + client = async_grpc_client.AsyncGrpcClient( + credentials=mock.Mock(), + ) + client._grpc_client = mock.AsyncMock() + # Pre-populate ACO bucket metadata cache for zonal bucket verification + client._bucket_metadata_cache.update_cache( + "my-zonal-bucket", + "projects/123456789/buckets/my-zonal-bucket", + "us-east1-a", + ) + return client + + +def test_inject_traceparent_to_metadata_when_disabled(monkeypatch): + monkeypatch.setattr(_opentelemetry_tracing, "enable_otel_traces", False) + + orig = (("x-goog-request-params", "bucket=foo"),) + result = _utils.inject_traceparent_to_metadata(orig) + assert result == orig + + +def test_inject_traceparent_to_metadata_when_enabled(exporter): + with _opentelemetry_tracing.create_trace_span( + "Test.ParentSpan", rpc_system="grpc" + ): + orig = (("x-goog-request-params", "bucket=foo"),) + result = _utils.inject_traceparent_to_metadata(orig) + keys = [k for k, _ in result] + assert "traceparent" in keys + traceparent_val = dict(result)["traceparent"] + assert traceparent_val.startswith("00-") + + +@pytest.mark.asyncio +async def test_async_grpc_client_get_and_delete_object_spans(exporter, mock_client): + mock_client._grpc_client.get_object.return_value = storage_v2.Object( + name="obj1", bucket="projects/_/buckets/my-zonal-bucket" + ) + mock_client._grpc_client.delete_object.return_value = None + + await mock_client.get_object("my-zonal-bucket", "obj1") + await mock_client.delete_object("my-zonal-bucket", "obj1") + + spans = exporter.get_finished_spans() + assert len(spans) == 2 + + get_span, del_span = spans[0], spans[1] + assert get_span.name == "Storage.AsyncGrpcClient.getObject" + assert get_span.attributes["rpc.system"] == "grpc" + assert get_span.attributes["gcp.client.service"] == "storage" + assert ( + get_span.attributes["gcp.resource.destination.id"] + == "projects/123456789/buckets/my-zonal-bucket" + ) + assert get_span.attributes["gcp.resource.destination.location"] == "us-east1-a" + + assert del_span.name == "Storage.AsyncGrpcClient.deleteObject" + assert del_span.attributes["rpc.system"] == "grpc" + + # Verify traceparent was injected into GAPIC call metadata + get_kwargs = mock_client._grpc_client.get_object.call_args.kwargs + metadata_keys = [k for k, _ in get_kwargs["metadata"]] + assert "traceparent" in metadata_keys + + +@pytest.mark.asyncio +async def test_async_appendable_object_writer_spans(exporter, mock_client): + writer = AsyncAppendableObjectWriter( + client=mock_client, + bucket_name="my-zonal-bucket", + object_name="append-obj", + ) + + with mock.patch( + "google.cloud.storage.asyncio.async_appendable_object_writer._AsyncWriteObjectStream" + ) as mock_stream_cls: + mock_stream = mock.AsyncMock() + mock_stream.generation_number = 1001 + mock_stream.write_handle = storage_v2.BidiWriteHandle(handle=b"handle-1") + mock_stream.persisted_size = 0 + mock_stream.recv.return_value = storage_v2.BidiWriteObjectResponse( + persisted_size=11 + ) + mock_stream_cls.return_value = mock_stream + + await writer.open() + await writer.append(b"hello world") + await writer.flush() + await writer.close() + + spans = exporter.get_finished_spans() + span_names = [s.name for s in spans] + assert span_names == [ + "Storage.AsyncAppendableObjectWriter.open", + "Storage.AsyncAppendableObjectWriter.append", + "Storage.AsyncAppendableObjectWriter.flush", + "Storage.AsyncAppendableObjectWriter.close", + ] + + append_span = spans[1] + assert append_span.attributes["rpc.system"] == "grpc" + assert append_span.attributes["gcp.storage.chunk.size"] == 11 + assert ( + append_span.attributes["gcp.resource.destination.id"] + == "projects/123456789/buckets/my-zonal-bucket" + ) + assert append_span.attributes["gcp.resource.destination.location"] == "us-east1-a" + + +@pytest.mark.asyncio +async def test_async_multi_range_downloader_spans(exporter, mock_client): + mrd = AsyncMultiRangeDownloader( + client=mock_client, + bucket_name="my-zonal-bucket", + object_name="read-obj", + ) + + with mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" + ) as mock_stream_cls, mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._StreamMultiplexer" + ) as mock_mux_cls, mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._BidiStreamRetryManager" + ) as mock_retry_mgr_cls: + mock_stream = mock.AsyncMock() + mock_stream.generation_number = 2002 + mock_stream.read_handle = storage_v2.BidiReadHandle(handle=b"rhandle-1") + mock_stream.persisted_size = 1024 + mock_stream.is_finalized = True + mock_stream.full_obj_server_crc32c = None + mock_stream_cls.return_value = mock_stream + + mock_mux = mock.AsyncMock() + mock_mux.register = mock.Mock(return_value=mock.AsyncMock()) + mock_mux.unregister = mock.Mock() + mock_mux_cls.return_value = mock_mux + + mock_retry_mgr = mock.AsyncMock() + mock_retry_mgr_cls.return_value = mock_retry_mgr + + await mrd.open() + buf1, buf2 = BytesIO(), BytesIO() + await mrd.download_ranges([(0, 100, buf1), (200, 300, buf2)], enable_checksum=False) + await mrd.close() + + spans = exporter.get_finished_spans() + span_names = [s.name for s in spans] + assert span_names == [ + "Storage.AsyncMultiRangeDownloader.open", + "Storage.AsyncMultiRangeDownloader.downloadRanges", + "Storage.AsyncMultiRangeDownloader.close", + ] + + download_span = spans[1] + assert download_span.attributes["rpc.system"] == "grpc" + assert download_span.attributes["gcp.storage.range.count"] == 2 + assert ( + download_span.attributes["gcp.resource.destination.id"] + == "projects/123456789/buckets/my-zonal-bucket" + ) + assert download_span.attributes["gcp.resource.destination.location"] == "us-east1-a" + + +@pytest.mark.asyncio +async def test_bucket_metadata_cache_async_grpc_fetch(mock_client): + cache = mock_client._bucket_metadata_cache + cache.clear() + + mock_client._grpc_client.get_bucket.return_value = storage_v2.Bucket( + name="projects/_/buckets/new-zonal-bucket", + project="projects/987654321", + location="US-WEST1-B", + location_type="zone", + ) + + # First call misses cache and schedules _fetch_background_async on event loop + res = cache.get_or_queue_fetch("new-zonal-bucket") + assert res is None + + # Allow scheduled task on event loop to complete + import asyncio + + await asyncio.sleep(0.01) + + cached = cache.get("new-zonal-bucket") + assert cached == ( + "projects/987654321/buckets/new-zonal-bucket", + "us-west1-b", + ) + + +@pytest.mark.asyncio +async def test_create_trace_span_helper_with_sync_mock_context_manager(mock_client): + """Verify fallback when _base_create_trace_span returns a synchronous context manager.""" + from contextlib import contextmanager + from google.cloud.storage import _helpers + + fake_span = object() + + @contextmanager + def sync_cm(*args, **kwargs): + yield fake_span + + with mock.patch.object(_helpers, "_base_create_trace_span", side_effect=sync_cm): + async with _helpers.create_trace_span_helper( + mock_client, "my-zonal-bucket", "Test.SyncFallback" + ) as span: + assert span is fake_span +