From 82444d4d5a3612a0f40c2efdf97f3b16bdae1e19 Mon Sep 17 00:00:00 2001 From: Prince Kumar Date: Thu, 1 Oct 2026 13:43:08 +0000 Subject: [PATCH 1/5] feat(storage): multi-stream support in AsyncMultiRangeDownloader --- packages/google-cloud-storage/MRD_DESIGN.md | 117 +++++ .../asyncio/async_multi_range_downloader.py | 417 +++++++++++++++++- .../test_async_multi_range_downloader.py | 394 +++++++++++++++++ 3 files changed, 906 insertions(+), 22 deletions(-) create mode 100644 packages/google-cloud-storage/MRD_DESIGN.md diff --git a/packages/google-cloud-storage/MRD_DESIGN.md b/packages/google-cloud-storage/MRD_DESIGN.md new file mode 100644 index 000000000000..ec60c6089baa --- /dev/null +++ b/packages/google-cloud-storage/MRD_DESIGN.md @@ -0,0 +1,117 @@ +# Design Doc: Multi-Stream Scaling and Unfinalized Object Support in AsyncMultiRangeDownloader (MRD) + +## Background +The Google Cloud Storage (GCS) Python SDK provides `AsyncMultiRangeDownloader` (MRD) to execute parallel range reads over a single bidi-gRPC (`BidiReadObject`) connection. Higher-level storage layers like `gcsfs` use MRD to service concurrent file reads, machine learning data loaders, and columnar analytics queries (Parquet/ORC). + +In high-throughput environments like Google Cloud Storage RAPID (low-latency, zonal storage class), network bandwidth often exceeds what a single bidi-gRPC stream can handle. A single gRPC connection caps around 1.0 to 1.2 GB/s due to single-TCP and single-HTTP/2 flow control limits. Workloads requiring multi-gigabyte/sec throughput are throttled unless requests are multiplexed across multiple streams. Additionally, objects actively being appended (unfinalized objects) frequently fail or get rejected in range readers because their exact size is changing and full-object CRC32c checksums are unavailable. + +## Objective +1. Scale throughput beyond single-stream limits to achieve over 2.0 GB/s on high-bandwidth storage classes like RAPID. +2. Provide dynamic, automatic connection scaling with load-based balancing. +3. Support seamless reading of unfinalized / actively appended objects without checksum errors. +4. Keep the public API clean, backward-compatible, and easy to configure for downstream consumers like `gcsfs`. + +## Overview +We enhance `AsyncMultiRangeDownloader` with an internal stream pool (`_StreamPool`) and stream manager (`_ManagedStream`). Instead of dispatching all range requests into one multiplexer, requests are assigned to the least-loaded stream. + +When stream load exceeds a target threshold, the pool automatically spins up additional bidi-gRPC streams in the background up to a configurable maximum. Stream selection is non-blocking: incoming range reads are dispatched immediately to the stream with the lowest active load, allowing HTTP/2 flow control to manage wire-level rate limiting without internal client stalls. For unfinalized objects, MRD skips the full-object CRC32c verification while preserving per-chunk CRC32c validation. + +## Detailed Design + +### 1. Configuration (`MRDStreamConfig`) & Zero-Overhead Rollback +Callers configure multi-stream behavior via `stream_config=MRDStreamConfig(...)` on `AsyncMultiRangeDownloader.create_mrd`: + +* `min_connections` (int, default 1): Initial number of pre-warmed streams opened during initialization. +* `max_connections` (int, default 8): Upper limit on total concurrent bidi-gRPC streams. +* `target_io_depth` (int, default 8): Target concurrent in-flight requests per stream before triggering scale-up. +* `target_bytes` (int, default 8 MB): Target in-flight bytes per stream before triggering scale-up. + +**Single-Connection Rollback**: If `stream_config` is `None` or `max_connections <= 1`, `_StreamPool` is not created. Calls to `download_ranges` execute a pure direct single-stream path without any `io-ranges` or `io-bytes` tracking, locking, or pool overhead. + +### 2. Load Balancing Metric and Dynamic Scaling +Each worker tracks its active `pending_ranges` and `pending_bytes`. Load is computed as: + +Load = 0.5 * (pending_ranges / target_io_depth) + 0.5 * (pending_bytes / target_bytes) + +* **Stream Selection (Non-Blocking)**: Incoming range requests are assigned to the worker with the lowest Load. Stream selection is instantaneous and non-blocking. +* **Scale-Up Trigger**: If the best available stream has Load >= 1.0, and total connections < `max_connections`, a background task opens an additional stream. New streams reuse the existing `read_handle` and routing token for sub-millisecond connection setup without repeating object lookups. +* **Proportional Scale-Up**: Background stream creation is triggered in direct proportion to total pool load (`desired_connections = ceil(total_load)`). During large concurrent bursts, multiple background streams are launched in parallel up to `max_connections`, preventing connection bottlenecks while immediately serving requests on existing streams. +* **Flow Control**: Backpressure is delegated to HTTP/2 and TCP window flow control (`WINDOW_UPDATE` frames) and caller-level task management, avoiding internal thread/coroutine blocking within the client library. + +### 3. Unfinalized Object Handling +For objects where `is_finalized` is False: +* **Size Ratcheting**: Persisted size is dynamically updated: persisted_size = max(persisted_size, offset + length). +* **Checksum Management**: Full-object CRC32c validation is bypassed since unfinalized objects do not have a finalized full-object checksum. Per-chunk CRC32c validation remains active to protect against transport data corruption. +* **Handle Propagation**: Newly refreshed `read_handle` tokens received in read responses are shared across the pool so all subsequent stream openings connect to the latest object state. + +## Empirical Benchmark & Saturation Data + +All benchmarks below were executed directly against an actual 20 GiB file (`gs://princer-rapid-bucket/bench/bench_20GB.bin`) in Google Cloud Storage. + +### Single-Stream Saturation Heatmap (I/O Size vs I/O Depth) + +| I/O Size | Depth 1 | Depth 2 | Depth 4 | Depth 8 | Depth 16 | Depth 32 | +| :--- | :---: | :---: | :---: | :---: | :---: | :---: | +| **128 KB** | 131.7 MB/s | 306.5 MB/s | 550.2 MB/s | 747.5 MB/s | 761.3 MB/s | 774.9 MB/s | +| **256 KB** | 252.6 MB/s | 488.5 MB/s | 901.7 MB/s | 957.7 MB/s | 980.6 MB/s | 1,017.9 MB/s | +| **512 KB** | 332.0 MB/s | 579.7 MB/s | 883.5 MB/s | 1,000.2 MB/s | 1,024.1 MB/s | 1,032.7 MB/s | +| **1 MB** | 372.1 MB/s | 625.9 MB/s | 926.5 MB/s | 983.0 MB/s | 1,004.8 MB/s | 990.5 MB/s | +| **2 MB** | 406.0 MB/s | 705.4 MB/s | 975.6 MB/s | 1,019.2 MB/s | 974.3 MB/s | 1,037.3 MB/s | +| **4 MB** | 510.3 MB/s | 913.9 MB/s | 956.8 MB/s | 953.1 MB/s | 1,064.0 MB/s | 1,127.4 MB/s | +| **8 MB** | 630.2 MB/s | 947.1 MB/s | 983.5 MB/s | 1,074.0 MB/s | 1,098.1 MB/s | 1,057.1 MB/s | +| **16 MB** | 700.6 MB/s | 974.8 MB/s | 1,033.7 MB/s | 936.0 MB/s | 951.1 MB/s | 902.2 MB/s | + +### Saturation Depth per I/O Size + +| I/O Size | Depth for ~950 MB/s | Depth for >= 1.0 GB/s | In-Flight Memory at Saturation | +| :--- | :---: | :---: | :---: | +| **16 MB** | Depth 2 (975 MB/s) | Depth 4 (1,034 MB/s) | 32 MB – 64 MB | +| **8 MB** | Depth 2 (947 MB/s) | Depth 8 (1,074 MB/s) | 16 MB – 64 MB | +| **4 MB** | Depth 4 (957 MB/s) | Depth 16 (1,064 MB/s) | 16 MB – 64 MB | +| **2 MB** | Depth 4 (976 MB/s) | Depth 8 (1,019 MB/s) | 8 MB – 16 MB | +| **1 MB** | Depth 6–8 (950–983 MB/s) | Depth 16 (1,005 MB/s) | 6 MB – 16 MB | +| **512 KB**| Depth 6–8 (950–1,000 MB/s) | Depth 8 (1,000 MB/s) | 3 MB – 4 MB | +| **256 KB**| Depth 8 (958 MB/s) | Depth 32 (1,018 MB/s) | 2 MB – 8 MB | +| **128 KB**| Capped at 775 MB/s | N/A | High RPC framing overhead | + +### Full 20 GiB End-to-End Download Comparison + +| Configuration | I/O Size | Concurrency Depth | Elapsed Time | Sustained Throughput | Notes | +| :--- | :---: | :---: | :---: | :---: | :--- | +| **Single Stream (1)** | 1 MB | 8 | 16.27 s | **1,258.41 MB/s (1.32 GB/s)** | Peak single stream performance | +| **Single Stream (1)** | 8 MB | 8 | 18.47 s | **1,108.62 MB/s (1.16 GB/s)** | Standard large read | +| **Single Stream (1)** | 4 MB | 8 | 20.04 s | **1,021.94 MB/s (1.07 GB/s)** | Conservative read | +| **Multi-Stream (4–8)** | 8 MB | 32 | **11.24 s** | **1,821.62 MB/s (1.91 GB/s)** | **Full 20 GiB transferred in 11.2s** | +| **Multi-Stream (4–8)** | 8 MB | 32 (20GB load) | **10.09 s** | **2,030.60 MB/s (2.13 GB/s)** | **Peak aggregate throughput** | + +## Alternative Solutions + +| Alternative | Description | Pros | Cons / Reason Not Chosen | +| :--- | :--- | :--- | :--- | +| **External Pooling in gcsfs** | Manage multiple MRD instances in gcsfs layer. | No changes to storage package. | High overhead: duplicate open handshakes, no shared routing tokens, high memory duplication, cannot load balance sub-ranges. | +| **Static Fixed Multi-Stream** | Always open N streams upfront. | Simple implementation. | Wastes socket resources on small reads; lacks dynamic adaptation. | +| **Round-Robin Routing** | Distribute requests uniformly across streams. | Simple routing logic. | Leads to straggler stalls when range sizes differ. | +| **Dynamic Load-Aware Pool (Chosen)** | Internal stream pool with heuristic load metric and shared handles. | Zero overhead on small files, optimal throughput on heavy workloads, non-blocking dispatch. | Slight increase in internal state management. | + +## Impact Checklist +* **Backward Compatibility**: Fully preserved. Existing `AsyncMultiRangeDownloader.create_mrd` signatures and behaviors are unchanged. +* **Security & Auth**: Reuses credentials and channel state from `AsyncGrpcClient`. No new permissions or tokens required. +* **Resource Usage**: Streams scale down and close cleanly when `mrd.close()` is called. In-flight requests are distributed non-blockingly. +* **Dependencies**: No external library additions. Relies entirely on `asyncio` and `grpc.aio`. + +## Test Plan +1. **Unit Testing**: + * Verify stream load metric calculation. + * Verify scale-up triggers and least-loaded stream selection. + * Verify parameter passthrough via `MRDStreamConfig`. + * Verify dynamic persisted size ratcheting and checksum bypass on unfinalized objects. + * Ensure 100% of existing unit tests pass without regression. +2. **Integration / Live Bucket Testing**: + * Verified against live zonal bucket `gs://princer-rapid-bucket/`. + * Verified full download of 20 GiB object (`bench/bench_20GB.bin`). + * Benchmarked throughput across 48 combinations of I/O size and depth. + * Confirmed clean socket teardown and zero resource leakage. + +## Rollout / Rollback Plan +* **Rollout**: Ship as part of `google-cloud-storage` release. By default, `min_connections=1` ensures standard behavior unless callers opt in or heavy concurrent workloads trigger scaling. `gcsfs` can pass `MRDStreamConfig` to leverage high-bandwidth multi-stream reads. +* **Rollback**: Callers can set `max_connections=1` to immediately restrict execution to a single connection. The code can be reverted without schema or protocol migrations. 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..a885721bdd96 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 @@ -15,9 +15,12 @@ from __future__ import annotations import asyncio +import inspect import logging +import math +from dataclasses import dataclass from io import BytesIO -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple from google.api_core import exceptions from google.api_core.retry_async import AsyncRetry @@ -78,6 +81,165 @@ def _is_read_retryable(exc): ) +@dataclass +class MRDStreamConfig: + """Configuration parameters for multi-stream download scaling and load balancing. + + :type min_connections: int + :param min_connections: Minimum number of streams opened at initialization. Defaults to 1. + + :type max_connections: int + :param max_connections: Maximum concurrent streams allowed. Defaults to 8. + + :type target_io_depth: int + :param target_io_depth: Target outstanding requests per connection before scaling up. Defaults to 8. + + :type target_bytes: int + :param target_bytes: Target outstanding bytes per connection before scaling up. Defaults to 8 MB. + """ + + min_connections: int = 1 + max_connections: int = 8 + target_io_depth: int = 8 + target_bytes: int = 8 * 1024 * 1024 + + +class _ManagedStream: + """Wraps an active stream and its multiplexer with local load counters.""" + + def __init__( + self, + stream: _AsyncReadObjectStream, + multiplexer: _StreamMultiplexer, + ): + self.stream = stream + self.multiplexer = multiplexer + self.pending_ranges: int = 0 + self.pending_bytes: int = 0 + + def calculate_load(self, target_io_depth: int, target_bytes: int) -> float: + """Returns normalized load metric combining request count and bytes.""" + u_req = self.pending_ranges / target_io_depth if target_io_depth > 0 else 0.0 + u_bytes = self.pending_bytes / target_bytes if target_bytes > 0 else 0.0 + return 0.5 * u_req + 0.5 * u_bytes + + def record_request(self, range_count: int, byte_count: int) -> None: + self.pending_ranges += range_count + self.pending_bytes += byte_count + + def record_completion(self, range_count: int, byte_count: int) -> None: + self.pending_ranges = max(0, self.pending_ranges - range_count) + self.pending_bytes = max(0, self.pending_bytes - byte_count) + + +class _StreamPool: + """Manages a pool of _ManagedStream instances with dynamic scaling and least-loaded dispatch.""" + + def __init__( + self, + stream_factory: Callable[[], Awaitable[_ManagedStream]], + min_connections: int = 1, + max_connections: int = 8, + target_io_depth: int = 8, + target_bytes: int = 8 * 1024 * 1024, + ): + self._stream_factory = stream_factory + self.min_connections = max(1, min_connections) + self.max_connections = max(self.min_connections, max_connections) + self.target_io_depth = target_io_depth + self.target_bytes = target_bytes + + self.workers: List[_ManagedStream] = [] + self._lock = asyncio.Lock() + self._pending_scale_ups: int = 0 + self._closed = False + self._background_tasks: set[asyncio.Task] = set() + + async def add_worker(self, worker: _ManagedStream) -> None: + async with self._lock: + self.workers.append(worker) + + async def acquire_stream(self, range_count: int, req_bytes: int) -> _ManagedStream: + """Finds least loaded stream and triggers background scale-up proportional to load.""" + async with self._lock: + if self._closed: + raise ValueError("Pool is closed") + if not self.workers: + raise ValueError("No workers available in stream pool") + best = min( + self.workers, + key=lambda w: w.calculate_load(self.target_io_depth, self.target_bytes), + ) + best.record_request(range_count, req_bytes) + + # Trigger background scale-ups proportional to total load across all workers + total_load = sum( + w.calculate_load(self.target_io_depth, self.target_bytes) + for w in self.workers + ) + desired_workers = math.ceil(total_load) + planned_workers = len(self.workers) + self._pending_scale_ups + needed_scale_ups = max( + 0, min(desired_workers, self.max_connections) - planned_workers + ) + + for _ in range(needed_scale_ups): + self._pending_scale_ups += 1 + task = asyncio.create_task(self._scale_up()) + self._background_tasks.add(task) + task.add_done_callback(self._background_tasks.discard) + + return best + + async def _scale_up(self) -> None: + try: + new_worker = await self._stream_factory() + async with self._lock: + if self._closed: + try: + await new_worker.multiplexer.close() + except Exception: + pass + try: + res = new_worker.stream.close() + if inspect.isawaitable(res): + await res + except Exception: + pass + return + self.workers.append(new_worker) + except Exception as e: + logger.warning(f"Failed to scale up MRD stream: {e}") + finally: + async with self._lock: + self._pending_scale_ups = max(0, self._pending_scale_ups - 1) + + def release_stream( + self, worker: _ManagedStream, range_count: int, req_bytes: int + ) -> None: + worker.record_completion(range_count, req_bytes) + + async def close(self) -> None: + async with self._lock: + self._closed = True + self._pending_scale_ups = 0 + workers = list(self.workers) + self.workers.clear() + for task in list(self._background_tasks): + task.cancel() + for w in workers: + try: + await w.multiplexer.close() + except Exception: + pass + try: + res = w.stream.close() + if inspect.isawaitable(res): + await res + except Exception: + pass + + class AsyncMultiRangeDownloader: """Provides an interface for downloading multiple ranges of a GCS ``Object`` concurrently. @@ -120,6 +282,11 @@ async def create_mrd( read_handle: Optional[_storage_v2.BidiReadHandle] = None, retry_policy: Optional[AsyncRetry] = None, metadata: Optional[List[Tuple[str, str]]] = None, + min_connections: int = 1, + max_connections: int = 8, + target_io_depth: int = 8, + target_bytes: int = 8 * 1024 * 1024, + stream_config: Optional[MRDStreamConfig] = None, **kwargs, ) -> AsyncMultiRangeDownloader: """Initializes a MultiRangeDownloader and opens the underlying bidi-gRPC @@ -148,6 +315,21 @@ async def create_mrd( :type metadata: List[Tuple[str, str]] :param metadata: (Optional) The metadata to be sent with the ``open`` request. + :type min_connections: int + :param min_connections: (Optional) Minimum number of streams to open initially. Defaults to 1. + + :type max_connections: int + :param max_connections: (Optional) Maximum number of concurrent streams allowed. Defaults to 8. + + :type target_io_depth: int + :param target_io_depth: (Optional) Desired requests outstanding per connection before scale-up. Defaults to 8. + + :type target_bytes: int + :param target_bytes: (Optional) Desired bytes outstanding per connection. Defaults to 8 MB. + + :type stream_config: Optional[MRDStreamConfig] + :param stream_config: (Optional) Configuration dataclass grouping all multi-stream parameters. + :rtype: :class:`~google.cloud.storage.asyncio.async_multi_range_downloader.AsyncMultiRangeDownloader` :returns: An initialized AsyncMultiRangeDownloader instance for reading. """ @@ -157,6 +339,11 @@ async def create_mrd( object_name, generation=generation, read_handle=read_handle, + min_connections=min_connections, + max_connections=max_connections, + target_io_depth=target_io_depth, + target_bytes=target_bytes, + stream_config=stream_config, **kwargs, ) await mrd.open(retry_policy=retry_policy, metadata=metadata) @@ -169,6 +356,11 @@ def __init__( object_name: str, generation: Optional[int] = None, read_handle: Optional[_storage_v2.BidiReadHandle] = None, + min_connections: int = 1, + max_connections: int = 8, + target_io_depth: int = 8, + target_bytes: int = 8 * 1024 * 1024, + stream_config: Optional[MRDStreamConfig] = None, **kwargs, ) -> None: """Constructor for AsyncMultiRangeDownloader, clients are not adviced to @@ -189,6 +381,21 @@ def __init__( :type read_handle: _storage_v2.BidiReadHandle :param read_handle: (Optional) An existing read handle. + + :type min_connections: int + :param min_connections: (Optional) Minimum number of streams to open initially. Defaults to 1. + + :type max_connections: int + :param max_connections: (Optional) Maximum number of concurrent streams allowed. Defaults to 8. + + :type target_io_depth: int + :param target_io_depth: (Optional) Desired requests outstanding per connection before scale-up. Defaults to 8. + + :type target_bytes: int + :param target_bytes: (Optional) Desired bytes outstanding per connection. Defaults to 8 MB. + + :type stream_config: Optional[MRDStreamConfig] + :param stream_config: (Optional) Configuration dataclass grouping all multi-stream parameters. """ if "generation_number" in kwargs: if generation is not None: @@ -202,6 +409,12 @@ def __init__( ) generation = kwargs.pop("generation_number") + if stream_config is not None: + min_connections = stream_config.min_connections + max_connections = stream_config.max_connections + target_io_depth = stream_config.target_io_depth + target_bytes = stream_config.target_bytes + self.client = client self.bucket_name = bucket_name self.object_name = object_name @@ -216,6 +429,15 @@ def __init__( self.is_finalized: bool = False self.full_obj_server_crc32c: Optional[int] = None + self.stream_config = stream_config + self.min_connections = min_connections + self.max_connections = max_connections + self.target_io_depth = target_io_depth + self.target_bytes = target_bytes + self._pool: Optional[_StreamPool] = None + self._primary_worker: Optional[_ManagedStream] = None + self._metadata: Optional[List[Tuple[str, str]]] = None + async def __aenter__(self): """Opens the underlying bidi-gRPC connection to read from the object.""" await self.open() @@ -308,6 +530,8 @@ async def _do_open(): 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 getattr(self.read_obj_str, "routing_token", None): + self._routing_token = self.read_obj_str.routing_token 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 @@ -315,15 +539,61 @@ async def _do_open(): self._is_stream_open = True + self._metadata = list(metadata) if metadata else [] await retry_policy(_do_open)() self._multiplexer = _StreamMultiplexer(self.read_obj_str) + self._primary_worker = _ManagedStream(self.read_obj_str, self._multiplexer) + + if self.stream_config is not None and self.stream_config.max_connections > 1: + self._pool = _StreamPool( + stream_factory=self._create_new_stream_worker, + min_connections=self.min_connections, + max_connections=self.max_connections, + target_io_depth=self.target_io_depth, + target_bytes=self.target_bytes, + ) + await self._pool.add_worker(self._primary_worker) + + for _ in range(self.min_connections - 1): + worker = await self._create_new_stream_worker() + await self._pool.add_worker(worker) + else: + self._pool = None + + async def _create_new_stream_worker(self) -> _ManagedStream: + """Opens an additional stream worker using current routing and read_handle.""" + current_metadata = list(self._metadata) if self._metadata else [] + if self._routing_token: + current_metadata.append( + ("x-goog-request-params", f"routing_token={self._routing_token}") + ) + + stream = _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, + ) + await stream.open(metadata=current_metadata if current_metadata else None) - def _create_stream_factory(self, state, metadata): + if stream.generation_number: + self.generation = stream.generation_number + if stream.read_handle: + self.read_handle = stream.read_handle + if getattr(stream, "routing_token", None): + self._routing_token = stream.routing_token + + mux = _StreamMultiplexer(stream) + return _ManagedStream(stream, mux) + + def _create_stream_factory(self, state, metadata, worker=None): """Create a factory that opens a new stream with current routing state.""" + target_worker = worker or self._primary_worker async def factory(): - current_handle = state.get("read_handle") - current_token = state.get("routing_token") + current_handle = state.get("read_handle") or self.read_handle + current_token = state.get("routing_token") or self._routing_token stream = _AsyncReadObjectStream( client=self.client.grpc_client, @@ -348,23 +618,29 @@ async def factory(): self.generation = stream.generation_number if stream.read_handle: self.read_handle = stream.read_handle + if getattr(stream, "routing_token", None): + self._routing_token = stream.routing_token self.is_finalized = stream.is_finalized self.full_obj_server_crc32c = stream.full_obj_server_crc32c self.read_obj_str = stream + if target_worker is not None: + target_worker.stream = stream + if target_worker is None or target_worker == self._primary_worker: + self.read_obj_str = stream self._is_stream_open = True return stream return factory - async def download_ranges( + async def _download_ranges_on_worker( self, + worker: _ManagedStream, read_ranges: List[Tuple[int, int, BytesIO]], - lock: asyncio.Lock = None, - retry_policy: Optional[AsyncRetry] = None, - metadata: Optional[List[Tuple[str, str]]] = None, - enable_checksum: bool = True, + retry_policy: AsyncRetry, + metadata: Optional[List[Tuple[str, str]]], + enable_checksum: bool, ) -> None: """Downloads multiple byte ranges from the object into the buffers provided by user with automatic retries. @@ -425,10 +701,13 @@ async def download_ranges( # 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 + is_full_object_read = self.is_finalized and ( + (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, @@ -442,11 +721,13 @@ async def download_ranges( "read_handle": self.read_handle, "routing_token": None, "enable_checksum": enable_checksum, - "full_obj_server_crc32c": self.full_obj_server_crc32c, + "full_obj_server_crc32c": self.full_obj_server_crc32c + if self.is_finalized + else None, } read_ids = set(download_states.keys()) - queue = self._multiplexer.register(read_ids) + queue = worker.multiplexer.register(read_ids) try: attempt_count = 0 @@ -473,14 +754,16 @@ async def generator(): broken_gen = ( last_broken_generation if attempt_count > 1 - else self._multiplexer.stream_generation + else worker.multiplexer.stream_generation ) - stream_factory = self._create_stream_factory(state, metadata) - await self._multiplexer.reopen_stream( + stream_factory = self._create_stream_factory( + state, metadata, worker=worker + ) + await worker.multiplexer.reopen_stream( broken_gen, stream_factory ) - stream_generation = self._multiplexer.stream_generation + stream_generation = worker.multiplexer.stream_generation # Send Requests pending_read_ids = {r.read_id for r in requests} @@ -489,7 +772,7 @@ async def generator(): ): batch = requests[i : i + _MAX_READ_RANGES_PER_BIDI_READ_REQUEST] try: - await self._multiplexer.send( + await worker.multiplexer.send( _storage_v2.BidiReadObjectRequest(read_ranges=batch) ) except Exception: @@ -542,6 +825,88 @@ async def generator(): finally: if self._multiplexer is not None: self._multiplexer.unregister(read_ids) + if worker.multiplexer is not None: + worker.multiplexer.unregister(read_ids) + + async def download_ranges( + self, + read_ranges: List[Tuple[int, int, BytesIO]], + lock: asyncio.Lock = None, + retry_policy: Optional[AsyncRetry] = None, + metadata: Optional[List[Tuple[str, str]]] = None, + enable_checksum: bool = True, + ) -> None: + """Downloads multiple byte ranges from the object into the buffers + provided by user with automatic retries across a managed multi-stream pool. + + :type read_ranges: List[Tuple[int, int, "BytesIO"]] + :param read_ranges: A list of tuples, where each tuple represents a + combination of byte_range and writeable buffer in format - + (`start_byte`, `bytes_to_read`, `writeable_buffer`). Buffer has + to be provided by the user, and user has to make sure appropriate + memory is available in the application to avoid out-of-memory crash. + + Special cases: + if the value of `bytes_to_read` is 0, it'll be interpreted as + download all contents until the end of the file from `start_byte`. + Examples: + * (0, 0, buffer) : downloads 0 to end , i.e. entire object. + * (100, 0, buffer) : downloads from 100 to end. + + :type lock: asyncio.Lock + :param lock: (Deprecated) This parameter is deprecated and has no effect. + + :type retry_policy: :class:`~google.api_core.retry_async.AsyncRetry` + :param retry_policy: (Optional) The retry policy to use for the operation. + + :type metadata: List[Tuple[str, str]] + :param metadata: (Optional) The metadata to be sent with the request. + + :type enable_checksum: bool + :param enable_checksum: (Optional) If True, checksums are verified for downloaded data. Defaults to True. + + :raises ValueError: if the underlying bidi-GRPC stream is not open. + :raises ValueError: if the length of read_ranges is more than 1000. + :raises DataCorruption: if a checksum mismatch is detected while reading data. + + """ + + if len(read_ranges) > 1000: + raise ValueError( + "Invalid input - length of read_ranges cannot be more than 1000" + ) + + if enable_checksum: + raise_if_no_fast_crc32c() + + 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) + + # Fallback for manually mocked tests that set mrd._multiplexer without calling open() + if self._primary_worker is None and self.read_obj_str and self._multiplexer: + self._primary_worker = _ManagedStream(self.read_obj_str, self._multiplexer) + + if self._pool is not None: + total_bytes = sum(length for _, length, _ in read_ranges) + total_ranges = len(read_ranges) + worker = await self._pool.acquire_stream(total_ranges, total_bytes) + try: + await self._download_ranges_on_worker( + worker, read_ranges, retry_policy, metadata, enable_checksum + ) + finally: + self._pool.release_stream(worker, total_ranges, total_bytes) + else: + await self._download_ranges_on_worker( + self._primary_worker, + read_ranges, + retry_policy, + metadata, + enable_checksum, + ) async def close(self): """ @@ -550,16 +915,24 @@ async def close(self): if not self._is_stream_open: raise ValueError("Underlying bidi-gRPC stream is not open") + if self._pool: + await self._pool.close() + self._pool = None + 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): + if getattr(self.read_obj_str, "is_stream_open", True): + res = self.read_obj_str.close() + if inspect.isawaitable(res): + await res + except (ValueError, asyncio.CancelledError, exceptions.GoogleAPICallError): pass self.read_obj_str = None + self._primary_worker = None self._is_stream_open = False @property diff --git a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py index afe1dfd221d9..d2c393aea31d 100644 --- a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py +++ b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py @@ -865,3 +865,397 @@ def test_is_read_retryable_predicate(self): assert _is_read_retryable(exceptions.NotFound("not found")) is False assert _is_read_retryable(exceptions.PermissionDenied("denied")) is False assert _is_read_retryable(exceptions.InvalidArgument("invalid")) is False + + def test_managed_stream_load(self): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + _ManagedStream, + ) + + mock_stream = mock.MagicMock() + mock_mux = mock.MagicMock() + worker = _ManagedStream(mock_stream, mock_mux) + + # Initially zero load + assert worker.calculate_load(target_io_depth=10, target_bytes=1000) == 0.0 + + # Add 5 ranges and 500 bytes -> 5/10 = 0.5, 500/1000 = 0.5 -> load = 0.5*0.5 + 0.5*0.5 = 0.5 + worker.record_request(5, 500) + assert worker.pending_ranges == 5 + assert worker.pending_bytes == 500 + assert worker.calculate_load(target_io_depth=10, target_bytes=1000) == 0.5 + + # Release + worker.record_completion(3, 300) + assert worker.pending_ranges == 2 + assert worker.pending_bytes == 200 + assert worker.calculate_load(target_io_depth=10, target_bytes=1000) == 0.2 + + @pytest.mark.asyncio + async def test_stream_pool_scale_up_and_dispatch(self): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + _ManagedStream, + _StreamPool, + ) + + created_workers = [] + + async def stream_factory(): + w = _ManagedStream(mock.MagicMock(), mock.MagicMock()) + created_workers.append(w) + return w + + pool = _StreamPool( + stream_factory=stream_factory, + min_connections=1, + max_connections=2, + target_io_depth=2, + target_bytes=100, + ) + initial_worker = await stream_factory() + await pool.add_worker(initial_worker) + + # 1. Acquire with small load (1 range, 10 bytes) -> load = 0.5*(1/2) + 0.5*(10/100) = 0.3 < 1.0 + w1 = await pool.acquire_stream(1, 10) + assert w1 == initial_worker + assert len(pool.workers) == 1 + + # 2. Add enough load to exceed target load >= 1.0 (e.g. 3 ranges, 90 bytes) + # Total on w1: 4 ranges (hits target and triggers background scale up) + w1_again = await pool.acquire_stream(3, 90) + assert w1_again == initial_worker + + # Scale-up task was scheduled; allow event loop to run background task + await asyncio.sleep(0.01) + assert len(pool.workers) == 2 + w2 = pool.workers[1] + assert w2 != initial_worker + + # 3. Next acquire selects w2 because w2 has 0 load while w1 has high load + w_next = await pool.acquire_stream(1, 10) + assert w_next == w2 + + # 4. Release worker capacity and close + pool.release_stream(w1, 4, 100) + pool.release_stream(w2, 1, 10) + await pool.close() + assert len(pool.workers) == 0 + + @pytest.mark.asyncio + async def test_stream_pool_proportional_scale_up(self): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + _ManagedStream, + _StreamPool, + ) + + created_workers = [] + + async def stream_factory(): + w = _ManagedStream(mock.MagicMock(), mock.MagicMock()) + created_workers.append(w) + return w + + pool = _StreamPool( + stream_factory=stream_factory, + min_connections=1, + max_connections=5, + target_io_depth=2, + target_bytes=100, + ) + initial_worker = await stream_factory() + await pool.add_worker(initial_worker) + + # Huge burst: 6 ranges, 300 bytes -> load = 0.5*(6/2) + 0.5*(300/100) = 3.0 + # Desired connections = ceil(3.0) = 3. + # Should launch 2 scale-ups concurrently in the background. + w1 = await pool.acquire_stream(6, 300) + assert w1 == initial_worker + assert pool._pending_scale_ups == 2 + assert len(pool._background_tasks) == 2 + + # Allow background scale-up tasks to finish + await asyncio.sleep(0.01) + assert len(pool.workers) == 3 + assert pool._pending_scale_ups == 0 + + # Another huge burst exceeding max_connections (5): + # 10 ranges, 500 bytes -> desired workers >= 5 + # Remaining headroom to max_connections is 5 - 3 = 2. + await pool.acquire_stream(10, 500) + assert pool._pending_scale_ups == 2 + + await asyncio.sleep(0.01) + assert len(pool.workers) == 5 + assert pool._pending_scale_ups == 0 + + await pool.close() + assert len(pool.workers) == 0 + + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" + ) + @pytest.mark.asyncio + async def test_create_mrd_with_stream_config(self, mock_cls_stream): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + MRDStreamConfig, + ) + + mock_client = mock.MagicMock() + mock_client.grpc_client = mock.AsyncMock() + + s1 = mock.MagicMock() + s1.open = AsyncMock() + s1.generation_number = 1 + s1.persisted_size = 100 + s1.read_handle = b"h1" + s1.object_metadata = mock.Mock() + + s2 = mock.MagicMock() + s2.open = AsyncMock() + s2.generation_number = 1 + s2.persisted_size = 100 + s2.read_handle = b"h2" + s2.object_metadata = mock.Mock() + + mock_cls_stream.side_effect = [s1, s2] + + config = MRDStreamConfig( + min_connections=2, + max_connections=4, + target_io_depth=8, + target_bytes=2 * 1024 * 1024, + ) + + mrd = await AsyncMultiRangeDownloader.create_mrd( + mock_client, "b", "o", stream_config=config + ) + + assert mrd.min_connections == 2 + assert mrd.max_connections == 4 + assert mrd.target_io_depth == 8 + assert mrd.target_bytes == 2 * 1024 * 1024 + # Verified that 2 streams were opened initially for min_connections=2 + assert len(mrd._pool.workers) == 2 + await mrd.close() + + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" + ) + @pytest.mark.asyncio + async def test_mrd_download_ranges_triggers_pool_scaling( + self, mock_cls_stream + ): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + MRDStreamConfig, + ) + + mock_client = mock.MagicMock() + mock_client.grpc_client = mock.AsyncMock() + + s1 = mock.MagicMock() + s1.open = AsyncMock() + s1.generation_number = 1 + s1.persisted_size = 1000 + s1.read_handle = b"h1" + s1.object_metadata = mock.Mock() + + s2 = mock.MagicMock() + s2.open = AsyncMock() + s2.generation_number = 1 + s2.persisted_size = 1000 + s2.read_handle = b"h2" + s2.object_metadata = mock.Mock() + + mock_cls_stream.side_effect = [s1, s2] + + config = MRDStreamConfig( + min_connections=1, + max_connections=3, + target_io_depth=1, + target_bytes=50, + ) + mrd = await AsyncMultiRangeDownloader.create_mrd( + mock_client, "b", "o", stream_config=config + ) + assert len(mrd._pool.workers) == 1 + + with mock.patch.object( + mrd, "_download_ranges_on_worker", new=AsyncMock() + ): + # Download ranges with load > 1.0 (2 ranges, 100 bytes) + await mrd.download_ranges([(0, 50, BytesIO()), (50, 50, BytesIO())]) + await asyncio.sleep(0.01) + # Pool dynamically scaled up to 2 workers + assert len(mrd._pool.workers) == 2 + + await mrd.close() + + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._ReadResumptionStrategy" + ) + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._BidiStreamRetryManager" + ) + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" + ) + @pytest.mark.asyncio + async def test_unfinalized_object_download( + self, + mock_cls_async_read_object_stream, + mock_retry_manager_cls, + mock_strategy_cls, + ): + mock_client = mock.MagicMock() + mock_client.grpc_client = mock.AsyncMock() + + mock_stream = mock_cls_async_read_object_stream.return_value + mock_stream.open = AsyncMock() + mock_stream.generation_number = 123 + mock_stream.persisted_size = 50 + mock_stream.read_handle = b"handle" + mock_stream.is_finalized = False + mock_stream.full_obj_server_crc32c = None + + mrd = await AsyncMultiRangeDownloader.create_mrd(mock_client, "b", "o") + assert mrd.is_finalized is False + assert mrd.persisted_size == 50 + + mock_retry_manager = mock_retry_manager_cls.return_value + mock_retry_manager.execute = AsyncMock() + + # Download a range extending past initial persisted_size (offset 50, length 100 -> end 150) + buf = BytesIO() + await mrd.download_ranges([(50, 100, buf)]) + + # Verify persisted_size is retained without error (no ratcheting) + assert mrd.persisted_size == 50 + await mrd.close() + + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" + ) + @pytest.mark.asyncio + async def test_create_mrd_single_stream_bypass(self, mock_cls_stream): + mock_client = mock.MagicMock() + mock_client.grpc_client = mock.AsyncMock() + + s1 = mock.MagicMock() + s1.open = AsyncMock() + s1.generation_number = 1 + s1.persisted_size = 100 + s1.read_handle = b"h1" + s1.object_metadata = mock.Mock() + mock_cls_stream.return_value = s1 + + # Default create_mrd without stream_config -> single-stream bypass + mrd = await AsyncMultiRangeDownloader.create_mrd(mock_client, "b", "o") + assert mrd.stream_config is None + assert mrd._pool is None + + # download_ranges should use _primary_worker directly without pool + with mock.patch.object( + mrd, "_download_ranges_on_worker", new=AsyncMock() + ) as mock_dl: + await mrd.download_ranges([(0, 50, BytesIO())]) + assert mock_dl.call_count == 1 + assert mock_dl.call_args[0][0] == mrd._primary_worker + + await mrd.close() + + @pytest.mark.asyncio + async def test_stream_pool_closed_and_cancellation(self): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + _ManagedStream, + _StreamPool, + ) + + scale_up_started = asyncio.Event() + scale_up_finish = asyncio.Event() + created_workers = [] + + async def slow_factory(): + scale_up_started.set() + await scale_up_finish.wait() + mock_stream = mock.MagicMock() + mock_stream.close = AsyncMock() + mock_mux = mock.MagicMock() + mock_mux.close = AsyncMock() + w = _ManagedStream(mock_stream, mock_mux) + created_workers.append(w) + return w + + pool = _StreamPool( + stream_factory=slow_factory, + min_connections=1, + max_connections=2, + target_io_depth=1, + target_bytes=10, + ) + initial_w = _ManagedStream(mock.MagicMock(), mock.MagicMock()) + await pool.add_worker(initial_w) + + # Trigger scale up + await pool.acquire_stream(2, 20) + await scale_up_started.wait() + assert len(pool._background_tasks) == 1 + + # Close pool while scale up task is pending + await pool.close() + assert pool._closed is True + assert len(pool.workers) == 0 + + # Background task should be cancelled + with pytest.raises(asyncio.CancelledError): + await next( + iter(pool._background_tasks) + ) if pool._background_tasks else asyncio.sleep(0) + + # acquire_stream on closed pool must raise ValueError + with pytest.raises(ValueError, match="Pool is closed"): + await pool.acquire_stream(1, 10) + + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" + ) + @pytest.mark.asyncio + async def test_routing_token_preservation_and_propagation( + self, mock_cls_async_read_object_stream + ): + mock_client = mock.MagicMock() + mock_client.grpc_client = mock.AsyncMock() + + s1 = mock.MagicMock() + s1.open = AsyncMock() + s1.generation_number = 100 + s1.routing_token = "token-abc" + s1.read_handle = b"h1" + s1.persisted_size = 1000 + s1.is_finalized = True + s1.full_obj_server_crc32c = 12345 + mock_cls_async_read_object_stream.return_value = s1 + + mrd = await AsyncMultiRangeDownloader.create_mrd(mock_client, "b", "o") + # Ensure routing token from initial stream was recorded + assert mrd._routing_token == "token-abc" + + # Now when a new stream worker is opened via _create_new_stream_worker, + # it should include routing_token in the metadata + s2 = mock.MagicMock() + s2.open = AsyncMock() + s2.generation_number = 100 + s2.routing_token = "token-xyz" + s2.read_handle = b"h2" + s2.persisted_size = 1000 + mock_cls_async_read_object_stream.return_value = s2 + + w2 = await mrd._create_new_stream_worker() + assert w2 is not None + assert s2.open.call_count == 1 + call_kwargs = s2.open.call_args[1] + assert ("x-goog-request-params", "routing_token=token-abc") in call_kwargs[ + "metadata" + ] + # And after s2 opened, the token should update to s2's routing_token + assert mrd._routing_token == "token-xyz" + + await mrd.close() From b8f7493ea3e8231f2e77846d8a54f30b3f266429 Mon Sep 17 00:00:00 2001 From: Prince Kumar Date: Thu, 1 Oct 2026 17:22:10 +0000 Subject: [PATCH 2/5] fixing linting and removing design --- packages/google-cloud-storage/MRD_DESIGN.md | 117 ------------------ .../asyncio/async_multi_range_downloader.py | 1 - .../test_async_multi_range_downloader.py | 11 +- 3 files changed, 3 insertions(+), 126 deletions(-) delete mode 100644 packages/google-cloud-storage/MRD_DESIGN.md diff --git a/packages/google-cloud-storage/MRD_DESIGN.md b/packages/google-cloud-storage/MRD_DESIGN.md deleted file mode 100644 index ec60c6089baa..000000000000 --- a/packages/google-cloud-storage/MRD_DESIGN.md +++ /dev/null @@ -1,117 +0,0 @@ -# Design Doc: Multi-Stream Scaling and Unfinalized Object Support in AsyncMultiRangeDownloader (MRD) - -## Background -The Google Cloud Storage (GCS) Python SDK provides `AsyncMultiRangeDownloader` (MRD) to execute parallel range reads over a single bidi-gRPC (`BidiReadObject`) connection. Higher-level storage layers like `gcsfs` use MRD to service concurrent file reads, machine learning data loaders, and columnar analytics queries (Parquet/ORC). - -In high-throughput environments like Google Cloud Storage RAPID (low-latency, zonal storage class), network bandwidth often exceeds what a single bidi-gRPC stream can handle. A single gRPC connection caps around 1.0 to 1.2 GB/s due to single-TCP and single-HTTP/2 flow control limits. Workloads requiring multi-gigabyte/sec throughput are throttled unless requests are multiplexed across multiple streams. Additionally, objects actively being appended (unfinalized objects) frequently fail or get rejected in range readers because their exact size is changing and full-object CRC32c checksums are unavailable. - -## Objective -1. Scale throughput beyond single-stream limits to achieve over 2.0 GB/s on high-bandwidth storage classes like RAPID. -2. Provide dynamic, automatic connection scaling with load-based balancing. -3. Support seamless reading of unfinalized / actively appended objects without checksum errors. -4. Keep the public API clean, backward-compatible, and easy to configure for downstream consumers like `gcsfs`. - -## Overview -We enhance `AsyncMultiRangeDownloader` with an internal stream pool (`_StreamPool`) and stream manager (`_ManagedStream`). Instead of dispatching all range requests into one multiplexer, requests are assigned to the least-loaded stream. - -When stream load exceeds a target threshold, the pool automatically spins up additional bidi-gRPC streams in the background up to a configurable maximum. Stream selection is non-blocking: incoming range reads are dispatched immediately to the stream with the lowest active load, allowing HTTP/2 flow control to manage wire-level rate limiting without internal client stalls. For unfinalized objects, MRD skips the full-object CRC32c verification while preserving per-chunk CRC32c validation. - -## Detailed Design - -### 1. Configuration (`MRDStreamConfig`) & Zero-Overhead Rollback -Callers configure multi-stream behavior via `stream_config=MRDStreamConfig(...)` on `AsyncMultiRangeDownloader.create_mrd`: - -* `min_connections` (int, default 1): Initial number of pre-warmed streams opened during initialization. -* `max_connections` (int, default 8): Upper limit on total concurrent bidi-gRPC streams. -* `target_io_depth` (int, default 8): Target concurrent in-flight requests per stream before triggering scale-up. -* `target_bytes` (int, default 8 MB): Target in-flight bytes per stream before triggering scale-up. - -**Single-Connection Rollback**: If `stream_config` is `None` or `max_connections <= 1`, `_StreamPool` is not created. Calls to `download_ranges` execute a pure direct single-stream path without any `io-ranges` or `io-bytes` tracking, locking, or pool overhead. - -### 2. Load Balancing Metric and Dynamic Scaling -Each worker tracks its active `pending_ranges` and `pending_bytes`. Load is computed as: - -Load = 0.5 * (pending_ranges / target_io_depth) + 0.5 * (pending_bytes / target_bytes) - -* **Stream Selection (Non-Blocking)**: Incoming range requests are assigned to the worker with the lowest Load. Stream selection is instantaneous and non-blocking. -* **Scale-Up Trigger**: If the best available stream has Load >= 1.0, and total connections < `max_connections`, a background task opens an additional stream. New streams reuse the existing `read_handle` and routing token for sub-millisecond connection setup without repeating object lookups. -* **Proportional Scale-Up**: Background stream creation is triggered in direct proportion to total pool load (`desired_connections = ceil(total_load)`). During large concurrent bursts, multiple background streams are launched in parallel up to `max_connections`, preventing connection bottlenecks while immediately serving requests on existing streams. -* **Flow Control**: Backpressure is delegated to HTTP/2 and TCP window flow control (`WINDOW_UPDATE` frames) and caller-level task management, avoiding internal thread/coroutine blocking within the client library. - -### 3. Unfinalized Object Handling -For objects where `is_finalized` is False: -* **Size Ratcheting**: Persisted size is dynamically updated: persisted_size = max(persisted_size, offset + length). -* **Checksum Management**: Full-object CRC32c validation is bypassed since unfinalized objects do not have a finalized full-object checksum. Per-chunk CRC32c validation remains active to protect against transport data corruption. -* **Handle Propagation**: Newly refreshed `read_handle` tokens received in read responses are shared across the pool so all subsequent stream openings connect to the latest object state. - -## Empirical Benchmark & Saturation Data - -All benchmarks below were executed directly against an actual 20 GiB file (`gs://princer-rapid-bucket/bench/bench_20GB.bin`) in Google Cloud Storage. - -### Single-Stream Saturation Heatmap (I/O Size vs I/O Depth) - -| I/O Size | Depth 1 | Depth 2 | Depth 4 | Depth 8 | Depth 16 | Depth 32 | -| :--- | :---: | :---: | :---: | :---: | :---: | :---: | -| **128 KB** | 131.7 MB/s | 306.5 MB/s | 550.2 MB/s | 747.5 MB/s | 761.3 MB/s | 774.9 MB/s | -| **256 KB** | 252.6 MB/s | 488.5 MB/s | 901.7 MB/s | 957.7 MB/s | 980.6 MB/s | 1,017.9 MB/s | -| **512 KB** | 332.0 MB/s | 579.7 MB/s | 883.5 MB/s | 1,000.2 MB/s | 1,024.1 MB/s | 1,032.7 MB/s | -| **1 MB** | 372.1 MB/s | 625.9 MB/s | 926.5 MB/s | 983.0 MB/s | 1,004.8 MB/s | 990.5 MB/s | -| **2 MB** | 406.0 MB/s | 705.4 MB/s | 975.6 MB/s | 1,019.2 MB/s | 974.3 MB/s | 1,037.3 MB/s | -| **4 MB** | 510.3 MB/s | 913.9 MB/s | 956.8 MB/s | 953.1 MB/s | 1,064.0 MB/s | 1,127.4 MB/s | -| **8 MB** | 630.2 MB/s | 947.1 MB/s | 983.5 MB/s | 1,074.0 MB/s | 1,098.1 MB/s | 1,057.1 MB/s | -| **16 MB** | 700.6 MB/s | 974.8 MB/s | 1,033.7 MB/s | 936.0 MB/s | 951.1 MB/s | 902.2 MB/s | - -### Saturation Depth per I/O Size - -| I/O Size | Depth for ~950 MB/s | Depth for >= 1.0 GB/s | In-Flight Memory at Saturation | -| :--- | :---: | :---: | :---: | -| **16 MB** | Depth 2 (975 MB/s) | Depth 4 (1,034 MB/s) | 32 MB – 64 MB | -| **8 MB** | Depth 2 (947 MB/s) | Depth 8 (1,074 MB/s) | 16 MB – 64 MB | -| **4 MB** | Depth 4 (957 MB/s) | Depth 16 (1,064 MB/s) | 16 MB – 64 MB | -| **2 MB** | Depth 4 (976 MB/s) | Depth 8 (1,019 MB/s) | 8 MB – 16 MB | -| **1 MB** | Depth 6–8 (950–983 MB/s) | Depth 16 (1,005 MB/s) | 6 MB – 16 MB | -| **512 KB**| Depth 6–8 (950–1,000 MB/s) | Depth 8 (1,000 MB/s) | 3 MB – 4 MB | -| **256 KB**| Depth 8 (958 MB/s) | Depth 32 (1,018 MB/s) | 2 MB – 8 MB | -| **128 KB**| Capped at 775 MB/s | N/A | High RPC framing overhead | - -### Full 20 GiB End-to-End Download Comparison - -| Configuration | I/O Size | Concurrency Depth | Elapsed Time | Sustained Throughput | Notes | -| :--- | :---: | :---: | :---: | :---: | :--- | -| **Single Stream (1)** | 1 MB | 8 | 16.27 s | **1,258.41 MB/s (1.32 GB/s)** | Peak single stream performance | -| **Single Stream (1)** | 8 MB | 8 | 18.47 s | **1,108.62 MB/s (1.16 GB/s)** | Standard large read | -| **Single Stream (1)** | 4 MB | 8 | 20.04 s | **1,021.94 MB/s (1.07 GB/s)** | Conservative read | -| **Multi-Stream (4–8)** | 8 MB | 32 | **11.24 s** | **1,821.62 MB/s (1.91 GB/s)** | **Full 20 GiB transferred in 11.2s** | -| **Multi-Stream (4–8)** | 8 MB | 32 (20GB load) | **10.09 s** | **2,030.60 MB/s (2.13 GB/s)** | **Peak aggregate throughput** | - -## Alternative Solutions - -| Alternative | Description | Pros | Cons / Reason Not Chosen | -| :--- | :--- | :--- | :--- | -| **External Pooling in gcsfs** | Manage multiple MRD instances in gcsfs layer. | No changes to storage package. | High overhead: duplicate open handshakes, no shared routing tokens, high memory duplication, cannot load balance sub-ranges. | -| **Static Fixed Multi-Stream** | Always open N streams upfront. | Simple implementation. | Wastes socket resources on small reads; lacks dynamic adaptation. | -| **Round-Robin Routing** | Distribute requests uniformly across streams. | Simple routing logic. | Leads to straggler stalls when range sizes differ. | -| **Dynamic Load-Aware Pool (Chosen)** | Internal stream pool with heuristic load metric and shared handles. | Zero overhead on small files, optimal throughput on heavy workloads, non-blocking dispatch. | Slight increase in internal state management. | - -## Impact Checklist -* **Backward Compatibility**: Fully preserved. Existing `AsyncMultiRangeDownloader.create_mrd` signatures and behaviors are unchanged. -* **Security & Auth**: Reuses credentials and channel state from `AsyncGrpcClient`. No new permissions or tokens required. -* **Resource Usage**: Streams scale down and close cleanly when `mrd.close()` is called. In-flight requests are distributed non-blockingly. -* **Dependencies**: No external library additions. Relies entirely on `asyncio` and `grpc.aio`. - -## Test Plan -1. **Unit Testing**: - * Verify stream load metric calculation. - * Verify scale-up triggers and least-loaded stream selection. - * Verify parameter passthrough via `MRDStreamConfig`. - * Verify dynamic persisted size ratcheting and checksum bypass on unfinalized objects. - * Ensure 100% of existing unit tests pass without regression. -2. **Integration / Live Bucket Testing**: - * Verified against live zonal bucket `gs://princer-rapid-bucket/`. - * Verified full download of 20 GiB object (`bench/bench_20GB.bin`). - * Benchmarked throughput across 48 combinations of I/O size and depth. - * Confirmed clean socket teardown and zero resource leakage. - -## Rollout / Rollback Plan -* **Rollout**: Ship as part of `google-cloud-storage` release. By default, `min_connections=1` ensures standard behavior unless callers opt in or heavy concurrent workloads trigger scaling. `gcsfs` can pass `MRDStreamConfig` to leverage high-bandwidth multi-stream reads. -* **Rollback**: Callers can set `max_connections=1` to immediately restrict execution to a single connection. The code can be reverted without schema or protocol migrations. 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 a885721bdd96..3f16108fa2f1 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 @@ -24,7 +24,6 @@ from google.api_core import exceptions 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.asyncio._stream_multiplexer import ( diff --git a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py index d2c393aea31d..2d9be7a80975 100644 --- a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py +++ b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py @@ -20,14 +20,13 @@ import google_crc32c import pytest from google.api_core import exceptions -from google.rpc import error_details_pb2, status_pb2 - from google.cloud import _storage_v2 from google.cloud.storage.asyncio import async_read_object_stream from google.cloud.storage.asyncio.async_multi_range_downloader import ( AsyncMultiRangeDownloader, ) from google.cloud.storage.exceptions import DataCorruption +from google.rpc import error_details_pb2, status_pb2 _TEST_BUCKET_NAME = "test-bucket" _TEST_OBJECT_NAME = "test-object" @@ -1041,9 +1040,7 @@ async def test_create_mrd_with_stream_config(self, mock_cls_stream): "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" ) @pytest.mark.asyncio - async def test_mrd_download_ranges_triggers_pool_scaling( - self, mock_cls_stream - ): + async def test_mrd_download_ranges_triggers_pool_scaling(self, mock_cls_stream): from google.cloud.storage.asyncio.async_multi_range_downloader import ( MRDStreamConfig, ) @@ -1078,9 +1075,7 @@ async def test_mrd_download_ranges_triggers_pool_scaling( ) assert len(mrd._pool.workers) == 1 - with mock.patch.object( - mrd, "_download_ranges_on_worker", new=AsyncMock() - ): + with mock.patch.object(mrd, "_download_ranges_on_worker", new=AsyncMock()): # Download ranges with load > 1.0 (2 ranges, 100 bytes) await mrd.download_ranges([(0, 50, BytesIO()), (50, 50, BytesIO())]) await asyncio.sleep(0.01) From 92c553a36f21ccebe893cfaf7dd147acef929cf8 Mon Sep 17 00:00:00 2001 From: Prince Kumar Date: Thu, 1 Oct 2026 17:47:05 +0000 Subject: [PATCH 3/5] keeping feature opt-in --- .../asyncio/async_multi_range_downloader.py | 65 +++++-------------- .../test_async_multi_range_downloader.py | 3 +- 2 files changed, 17 insertions(+), 51 deletions(-) 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 3f16108fa2f1..fbba42212e6e 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 @@ -24,6 +24,7 @@ from google.api_core import exceptions 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.asyncio._stream_multiplexer import ( @@ -281,10 +282,6 @@ async def create_mrd( read_handle: Optional[_storage_v2.BidiReadHandle] = None, retry_policy: Optional[AsyncRetry] = None, metadata: Optional[List[Tuple[str, str]]] = None, - min_connections: int = 1, - max_connections: int = 8, - target_io_depth: int = 8, - target_bytes: int = 8 * 1024 * 1024, stream_config: Optional[MRDStreamConfig] = None, **kwargs, ) -> AsyncMultiRangeDownloader: @@ -314,20 +311,10 @@ async def create_mrd( :type metadata: List[Tuple[str, str]] :param metadata: (Optional) The metadata to be sent with the ``open`` request. - :type min_connections: int - :param min_connections: (Optional) Minimum number of streams to open initially. Defaults to 1. - - :type max_connections: int - :param max_connections: (Optional) Maximum number of concurrent streams allowed. Defaults to 8. - - :type target_io_depth: int - :param target_io_depth: (Optional) Desired requests outstanding per connection before scale-up. Defaults to 8. - - :type target_bytes: int - :param target_bytes: (Optional) Desired bytes outstanding per connection. Defaults to 8 MB. - :type stream_config: Optional[MRDStreamConfig] - :param stream_config: (Optional) Configuration dataclass grouping all multi-stream parameters. + :param stream_config: (Optional) Configuration dataclass grouping all multi-stream + parameters. If None, multi-stream is disabled and a single + stream is used. :rtype: :class:`~google.cloud.storage.asyncio.async_multi_range_downloader.AsyncMultiRangeDownloader` :returns: An initialized AsyncMultiRangeDownloader instance for reading. @@ -338,10 +325,6 @@ async def create_mrd( object_name, generation=generation, read_handle=read_handle, - min_connections=min_connections, - max_connections=max_connections, - target_io_depth=target_io_depth, - target_bytes=target_bytes, stream_config=stream_config, **kwargs, ) @@ -355,15 +338,11 @@ def __init__( object_name: str, generation: Optional[int] = None, read_handle: Optional[_storage_v2.BidiReadHandle] = None, - min_connections: int = 1, - max_connections: int = 8, - target_io_depth: int = 8, - target_bytes: int = 8 * 1024 * 1024, stream_config: Optional[MRDStreamConfig] = None, **kwargs, ) -> None: - """Constructor for AsyncMultiRangeDownloader, clients are not adviced to - use it directly. Instead it's adviced to use the classmethod `create_mrd`. + """Constructor for AsyncMultiRangeDownloader, clients are not advised to + use it directly. Instead it's advised to use the classmethod `create_mrd`. :type client: :class:`~google.cloud.storage.asyncio.async_grpc_client.AsyncGrpcClient` :param client: The asynchronous client to use for making API requests. @@ -381,20 +360,10 @@ def __init__( :type read_handle: _storage_v2.BidiReadHandle :param read_handle: (Optional) An existing read handle. - :type min_connections: int - :param min_connections: (Optional) Minimum number of streams to open initially. Defaults to 1. - - :type max_connections: int - :param max_connections: (Optional) Maximum number of concurrent streams allowed. Defaults to 8. - - :type target_io_depth: int - :param target_io_depth: (Optional) Desired requests outstanding per connection before scale-up. Defaults to 8. - - :type target_bytes: int - :param target_bytes: (Optional) Desired bytes outstanding per connection. Defaults to 8 MB. - :type stream_config: Optional[MRDStreamConfig] - :param stream_config: (Optional) Configuration dataclass grouping all multi-stream parameters. + :param stream_config: (Optional) Configuration dataclass grouping all multi-stream + parameters. If None, multi-stream is disabled and a single + stream is used. """ if "generation_number" in kwargs: if generation is not None: @@ -408,12 +377,6 @@ def __init__( ) generation = kwargs.pop("generation_number") - if stream_config is not None: - min_connections = stream_config.min_connections - max_connections = stream_config.max_connections - target_io_depth = stream_config.target_io_depth - target_bytes = stream_config.target_bytes - self.client = client self.bucket_name = bucket_name self.object_name = object_name @@ -429,10 +392,12 @@ def __init__( self.full_obj_server_crc32c: Optional[int] = None self.stream_config = stream_config - self.min_connections = min_connections - self.max_connections = max_connections - self.target_io_depth = target_io_depth - self.target_bytes = target_bytes + self.min_connections = stream_config.min_connections if stream_config else 1 + self.max_connections = stream_config.max_connections if stream_config else 1 + self.target_io_depth = stream_config.target_io_depth if stream_config else 8 + self.target_bytes = ( + stream_config.target_bytes if stream_config else 8 * 1024 * 1024 + ) self._pool: Optional[_StreamPool] = None self._primary_worker: Optional[_ManagedStream] = None self._metadata: Optional[List[Tuple[str, str]]] = None diff --git a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py index 2d9be7a80975..0480447fd780 100644 --- a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py +++ b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py @@ -20,13 +20,14 @@ import google_crc32c import pytest from google.api_core import exceptions +from google.rpc import error_details_pb2, status_pb2 + from google.cloud import _storage_v2 from google.cloud.storage.asyncio import async_read_object_stream from google.cloud.storage.asyncio.async_multi_range_downloader import ( AsyncMultiRangeDownloader, ) from google.cloud.storage.exceptions import DataCorruption -from google.rpc import error_details_pb2, status_pb2 _TEST_BUCKET_NAME = "test-bucket" _TEST_OBJECT_NAME = "test-object" From b7fa23322457fb186bfd2ce5753c7073dfd0bb68 Mon Sep 17 00:00:00 2001 From: Prince Kumar Date: Thu, 1 Oct 2026 18:11:25 +0000 Subject: [PATCH 4/5] refactor(storage): clean up stream teardown, eliminate redundant state, and revert artificial diffs --- .../asyncio/async_multi_range_downloader.py | 74 +++++++------------ .../test_async_multi_range_downloader.py | 16 ++-- 2 files changed, 31 insertions(+), 59 deletions(-) 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 fbba42212e6e..c588c49899e4 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 @@ -131,6 +131,19 @@ def record_completion(self, range_count: int, byte_count: int) -> None: self.pending_ranges = max(0, self.pending_ranges - range_count) self.pending_bytes = max(0, self.pending_bytes - byte_count) + async def close(self) -> None: + """Closes the underlying multiplexer and stream cleanly.""" + try: + await self.multiplexer.close() + except Exception: + pass + try: + res = self.stream.close() + if inspect.isawaitable(res): + await res + except Exception: + pass + class _StreamPool: """Manages a pool of _ManagedStream instances with dynamic scaling and least-loaded dispatch.""" @@ -196,16 +209,7 @@ async def _scale_up(self) -> None: new_worker = await self._stream_factory() async with self._lock: if self._closed: - try: - await new_worker.multiplexer.close() - except Exception: - pass - try: - res = new_worker.stream.close() - if inspect.isawaitable(res): - await res - except Exception: - pass + await new_worker.close() return self.workers.append(new_worker) except Exception as e: @@ -228,16 +232,7 @@ async def close(self) -> None: for task in list(self._background_tasks): task.cancel() for w in workers: - try: - await w.multiplexer.close() - except Exception: - pass - try: - res = w.stream.close() - if inspect.isawaitable(res): - await res - except Exception: - pass + await w.close() class AsyncMultiRangeDownloader: @@ -392,12 +387,6 @@ def __init__( self.full_obj_server_crc32c: Optional[int] = None self.stream_config = stream_config - self.min_connections = stream_config.min_connections if stream_config else 1 - self.max_connections = stream_config.max_connections if stream_config else 1 - self.target_io_depth = stream_config.target_io_depth if stream_config else 8 - self.target_bytes = ( - stream_config.target_bytes if stream_config else 8 * 1024 * 1024 - ) self._pool: Optional[_StreamPool] = None self._primary_worker: Optional[_ManagedStream] = None self._metadata: Optional[List[Tuple[str, str]]] = None @@ -494,8 +483,6 @@ async def _do_open(): 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 getattr(self.read_obj_str, "routing_token", None): - self._routing_token = self.read_obj_str.routing_token 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 @@ -511,14 +498,14 @@ async def _do_open(): if self.stream_config is not None and self.stream_config.max_connections > 1: self._pool = _StreamPool( stream_factory=self._create_new_stream_worker, - min_connections=self.min_connections, - max_connections=self.max_connections, - target_io_depth=self.target_io_depth, - target_bytes=self.target_bytes, + min_connections=self.stream_config.min_connections, + max_connections=self.stream_config.max_connections, + target_io_depth=self.stream_config.target_io_depth, + target_bytes=self.stream_config.target_bytes, ) await self._pool.add_worker(self._primary_worker) - for _ in range(self.min_connections - 1): + for _ in range(self.stream_config.min_connections - 1): worker = await self._create_new_stream_worker() await self._pool.add_worker(worker) else: @@ -545,8 +532,6 @@ async def _create_new_stream_worker(self) -> _ManagedStream: self.generation = stream.generation_number if stream.read_handle: self.read_handle = stream.read_handle - if getattr(stream, "routing_token", None): - self._routing_token = stream.routing_token mux = _StreamMultiplexer(stream) return _ManagedStream(stream, mux) @@ -582,8 +567,6 @@ async def factory(): self.generation = stream.generation_number if stream.read_handle: self.read_handle = stream.read_handle - if getattr(stream, "routing_token", None): - self._routing_token = stream.routing_token self.is_finalized = stream.is_finalized self.full_obj_server_crc32c = stream.full_obj_server_crc32c @@ -665,13 +648,10 @@ async def _download_ranges_on_worker( # 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 = self.is_finalized and ( - (offset == 0 and length == 0) - or ( - self.persisted_size is not None - and offset == 0 - and length == self.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, @@ -685,9 +665,7 @@ async def _download_ranges_on_worker( "read_handle": self.read_handle, "routing_token": None, "enable_checksum": enable_checksum, - "full_obj_server_crc32c": self.full_obj_server_crc32c - if self.is_finalized - else None, + "full_obj_server_crc32c": self.full_obj_server_crc32c, } read_ids = set(download_states.keys()) @@ -787,8 +765,6 @@ async def generator(): 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) if worker.multiplexer is not None: worker.multiplexer.unregister(read_ids) diff --git a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py index 0480447fd780..ce4d87dfc093 100644 --- a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py +++ b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py @@ -1029,10 +1029,10 @@ async def test_create_mrd_with_stream_config(self, mock_cls_stream): mock_client, "b", "o", stream_config=config ) - assert mrd.min_connections == 2 - assert mrd.max_connections == 4 - assert mrd.target_io_depth == 8 - assert mrd.target_bytes == 2 * 1024 * 1024 + assert mrd.stream_config.min_connections == 2 + assert mrd.stream_config.max_connections == 4 + assert mrd.stream_config.target_io_depth == 8 + assert mrd.stream_config.target_bytes == 2 * 1024 * 1024 # Verified that 2 streams were opened initially for min_connections=2 assert len(mrd._pool.workers) == 2 await mrd.close() @@ -1223,7 +1223,6 @@ async def test_routing_token_preservation_and_propagation( s1 = mock.MagicMock() s1.open = AsyncMock() s1.generation_number = 100 - s1.routing_token = "token-abc" s1.read_handle = b"h1" s1.persisted_size = 1000 s1.is_finalized = True @@ -1231,15 +1230,14 @@ async def test_routing_token_preservation_and_propagation( mock_cls_async_read_object_stream.return_value = s1 mrd = await AsyncMultiRangeDownloader.create_mrd(mock_client, "b", "o") - # Ensure routing token from initial stream was recorded - assert mrd._routing_token == "token-abc" + # Simulate a redirect having set _routing_token + mrd._routing_token = "token-abc" # Now when a new stream worker is opened via _create_new_stream_worker, # it should include routing_token in the metadata s2 = mock.MagicMock() s2.open = AsyncMock() s2.generation_number = 100 - s2.routing_token = "token-xyz" s2.read_handle = b"h2" s2.persisted_size = 1000 mock_cls_async_read_object_stream.return_value = s2 @@ -1251,7 +1249,5 @@ async def test_routing_token_preservation_and_propagation( assert ("x-goog-request-params", "routing_token=token-abc") in call_kwargs[ "metadata" ] - # And after s2 opened, the token should update to s2's routing_token - assert mrd._routing_token == "token-xyz" await mrd.close() From b183f7fc1ad052cbdb27ebdeacbc6020231e2bc3 Mon Sep 17 00:00:00 2001 From: Prince Kumar Date: Thu, 1 Oct 2026 19:08:39 +0000 Subject: [PATCH 5/5] minor refactoring --- .../asyncio/async_multi_range_downloader.py | 140 ++++---- .../test_async_multi_range_downloader.py | 299 ++++++++---------- 2 files changed, 218 insertions(+), 221 deletions(-) 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 c588c49899e4..e7802018672a 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 @@ -104,7 +104,7 @@ class MRDStreamConfig: target_bytes: int = 8 * 1024 * 1024 -class _ManagedStream: +class _PooledStream: """Wraps an active stream and its multiplexer with local load counters.""" def __init__( @@ -146,11 +146,11 @@ async def close(self) -> None: class _StreamPool: - """Manages a pool of _ManagedStream instances with dynamic scaling and least-loaded dispatch.""" + """Manages a pool of _PooledStream instances with dynamic scaling and least-loaded dispatch.""" def __init__( self, - stream_factory: Callable[[], Awaitable[_ManagedStream]], + stream_factory: Callable[[], Awaitable[_PooledStream]], min_connections: int = 1, max_connections: int = 8, target_io_depth: int = 8, @@ -162,38 +162,38 @@ def __init__( self.target_io_depth = target_io_depth self.target_bytes = target_bytes - self.workers: List[_ManagedStream] = [] + self.streams: List[_PooledStream] = [] self._lock = asyncio.Lock() self._pending_scale_ups: int = 0 self._closed = False self._background_tasks: set[asyncio.Task] = set() - async def add_worker(self, worker: _ManagedStream) -> None: + async def add_stream(self, pooled_stream: _PooledStream) -> None: async with self._lock: - self.workers.append(worker) + self.streams.append(pooled_stream) - async def acquire_stream(self, range_count: int, req_bytes: int) -> _ManagedStream: + async def acquire_stream(self, range_count: int, req_bytes: int) -> _PooledStream: """Finds least loaded stream and triggers background scale-up proportional to load.""" async with self._lock: if self._closed: raise ValueError("Pool is closed") - if not self.workers: - raise ValueError("No workers available in stream pool") + if not self.streams: + raise ValueError("No streams available in stream pool") best = min( - self.workers, - key=lambda w: w.calculate_load(self.target_io_depth, self.target_bytes), + self.streams, + key=lambda s: s.calculate_load(self.target_io_depth, self.target_bytes), ) best.record_request(range_count, req_bytes) - # Trigger background scale-ups proportional to total load across all workers + # Trigger background scale-ups proportional to total load across all streams total_load = sum( - w.calculate_load(self.target_io_depth, self.target_bytes) - for w in self.workers + s.calculate_load(self.target_io_depth, self.target_bytes) + for s in self.streams ) - desired_workers = math.ceil(total_load) - planned_workers = len(self.workers) + self._pending_scale_ups + desired_streams = math.ceil(total_load) + planned_streams = len(self.streams) + self._pending_scale_ups needed_scale_ups = max( - 0, min(desired_workers, self.max_connections) - planned_workers + 0, min(desired_streams, self.max_connections) - planned_streams ) for _ in range(needed_scale_ups): @@ -206,12 +206,12 @@ async def acquire_stream(self, range_count: int, req_bytes: int) -> _ManagedStre async def _scale_up(self) -> None: try: - new_worker = await self._stream_factory() + new_stream = await self._stream_factory() async with self._lock: if self._closed: - await new_worker.close() + await new_stream.close() return - self.workers.append(new_worker) + self.streams.append(new_stream) except Exception as e: logger.warning(f"Failed to scale up MRD stream: {e}") finally: @@ -219,20 +219,20 @@ async def _scale_up(self) -> None: self._pending_scale_ups = max(0, self._pending_scale_ups - 1) def release_stream( - self, worker: _ManagedStream, range_count: int, req_bytes: int + self, pooled_stream: _PooledStream, range_count: int, req_bytes: int ) -> None: - worker.record_completion(range_count, req_bytes) + pooled_stream.record_completion(range_count, req_bytes) async def close(self) -> None: async with self._lock: self._closed = True self._pending_scale_ups = 0 - workers = list(self.workers) - self.workers.clear() + streams = list(self.streams) + self.streams.clear() for task in list(self._background_tasks): task.cancel() - for w in workers: - await w.close() + for s in streams: + await s.close() class AsyncMultiRangeDownloader: @@ -388,7 +388,7 @@ def __init__( self.stream_config = stream_config self._pool: Optional[_StreamPool] = None - self._primary_worker: Optional[_ManagedStream] = None + self._primary_stream: Optional[_PooledStream] = None self._metadata: Optional[List[Tuple[str, str]]] = None async def __aenter__(self): @@ -493,26 +493,44 @@ async def _do_open(): self._metadata = list(metadata) if metadata else [] await retry_policy(_do_open)() self._multiplexer = _StreamMultiplexer(self.read_obj_str) - self._primary_worker = _ManagedStream(self.read_obj_str, self._multiplexer) + self._primary_stream = _PooledStream(self.read_obj_str, self._multiplexer) if self.stream_config is not None and self.stream_config.max_connections > 1: self._pool = _StreamPool( - stream_factory=self._create_new_stream_worker, + stream_factory=self._create_pooled_stream, min_connections=self.stream_config.min_connections, max_connections=self.stream_config.max_connections, target_io_depth=self.stream_config.target_io_depth, target_bytes=self.stream_config.target_bytes, ) - await self._pool.add_worker(self._primary_worker) - - for _ in range(self.stream_config.min_connections - 1): - worker = await self._create_new_stream_worker() - await self._pool.add_worker(worker) + await self._pool.add_stream(self._primary_stream) + + if self.stream_config.min_connections > 1: + extra_streams = await asyncio.gather( + *( + self._create_pooled_stream() + for _ in range(self.stream_config.min_connections - 1) + ), + return_exceptions=True, + ) + first_exc = next( + (s for s in extra_streams if isinstance(s, Exception)), None + ) + if first_exc is not None: + for s in extra_streams: + if not isinstance(s, Exception): + await s.close() + await self._pool.close() + self._pool = None + raise first_exc + + for stream in extra_streams: + await self._pool.add_stream(stream) else: self._pool = None - async def _create_new_stream_worker(self) -> _ManagedStream: - """Opens an additional stream worker using current routing and read_handle.""" + async def _create_pooled_stream(self) -> _PooledStream: + """Opens an additional pooled stream using current routing and read_handle.""" current_metadata = list(self._metadata) if self._metadata else [] if self._routing_token: current_metadata.append( @@ -534,11 +552,11 @@ async def _create_new_stream_worker(self) -> _ManagedStream: self.read_handle = stream.read_handle mux = _StreamMultiplexer(stream) - return _ManagedStream(stream, mux) + return _PooledStream(stream, mux) - def _create_stream_factory(self, state, metadata, worker=None): + def _create_stream_factory(self, state, metadata, pooled_stream=None): """Create a factory that opens a new stream with current routing state.""" - target_worker = worker or self._primary_worker + target_stream = pooled_stream or self._primary_stream async def factory(): current_handle = state.get("read_handle") or self.read_handle @@ -571,9 +589,9 @@ async def factory(): self.full_obj_server_crc32c = stream.full_obj_server_crc32c self.read_obj_str = stream - if target_worker is not None: - target_worker.stream = stream - if target_worker is None or target_worker == self._primary_worker: + if target_stream is not None: + target_stream.stream = stream + if target_stream is None or target_stream == self._primary_stream: self.read_obj_str = stream self._is_stream_open = True @@ -581,9 +599,9 @@ async def factory(): return factory - async def _download_ranges_on_worker( + async def _download_ranges_on_stream( self, - worker: _ManagedStream, + pooled_stream: _PooledStream, read_ranges: List[Tuple[int, int, BytesIO]], retry_policy: AsyncRetry, metadata: Optional[List[Tuple[str, str]]], @@ -669,7 +687,7 @@ async def _download_ranges_on_worker( } read_ids = set(download_states.keys()) - queue = worker.multiplexer.register(read_ids) + queue = pooled_stream.multiplexer.register(read_ids) try: attempt_count = 0 @@ -696,16 +714,16 @@ async def generator(): broken_gen = ( last_broken_generation if attempt_count > 1 - else worker.multiplexer.stream_generation + else pooled_stream.multiplexer.stream_generation ) stream_factory = self._create_stream_factory( - state, metadata, worker=worker + state, metadata, pooled_stream=pooled_stream ) - await worker.multiplexer.reopen_stream( + await pooled_stream.multiplexer.reopen_stream( broken_gen, stream_factory ) - stream_generation = worker.multiplexer.stream_generation + stream_generation = pooled_stream.multiplexer.stream_generation # Send Requests pending_read_ids = {r.read_id for r in requests} @@ -714,7 +732,7 @@ async def generator(): ): batch = requests[i : i + _MAX_READ_RANGES_PER_BIDI_READ_REQUEST] try: - await worker.multiplexer.send( + await pooled_stream.multiplexer.send( _storage_v2.BidiReadObjectRequest(read_ranges=batch) ) except Exception: @@ -765,8 +783,8 @@ async def generator(): if initial_state.get("read_handle"): self.read_handle = initial_state["read_handle"] finally: - if worker.multiplexer is not None: - worker.multiplexer.unregister(read_ids) + if pooled_stream.multiplexer is not None: + pooled_stream.multiplexer.unregister(read_ids) async def download_ranges( self, @@ -826,22 +844,22 @@ async def download_ranges( retry_policy = AsyncRetry(predicate=_is_read_retryable) # Fallback for manually mocked tests that set mrd._multiplexer without calling open() - if self._primary_worker is None and self.read_obj_str and self._multiplexer: - self._primary_worker = _ManagedStream(self.read_obj_str, self._multiplexer) + if self._primary_stream is None and self.read_obj_str and self._multiplexer: + self._primary_stream = _PooledStream(self.read_obj_str, self._multiplexer) if self._pool is not None: total_bytes = sum(length for _, length, _ in read_ranges) total_ranges = len(read_ranges) - worker = await self._pool.acquire_stream(total_ranges, total_bytes) + pooled_stream = await self._pool.acquire_stream(total_ranges, total_bytes) try: - await self._download_ranges_on_worker( - worker, read_ranges, retry_policy, metadata, enable_checksum + await self._download_ranges_on_stream( + pooled_stream, read_ranges, retry_policy, metadata, enable_checksum ) finally: - self._pool.release_stream(worker, total_ranges, total_bytes) + self._pool.release_stream(pooled_stream, total_ranges, total_bytes) else: - await self._download_ranges_on_worker( - self._primary_worker, + await self._download_ranges_on_stream( + self._primary_stream, read_ranges, retry_policy, metadata, @@ -872,7 +890,7 @@ async def close(self): except (ValueError, asyncio.CancelledError, exceptions.GoogleAPICallError): pass self.read_obj_str = None - self._primary_worker = None + self._primary_stream = None self._is_stream_open = False @property diff --git a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py index ce4d87dfc093..7cc3c04006da 100644 --- a/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py +++ b/packages/google-cloud-storage/tests/unit/asyncio/test_async_multi_range_downloader.py @@ -37,6 +37,42 @@ class TestAsyncMultiRangeDownloader: + @staticmethod + def _make_mock_client(): + mock_client = mock.MagicMock() + mock_client.grpc_client = mock.AsyncMock() + return mock_client + + @staticmethod + def _make_mock_stream( + generation=_TEST_GENERATION_NUMBER, + persisted_size=_TEST_OBJECT_SIZE, + read_handle=_TEST_READ_HANDLE, + is_finalized=True, + open_side_effect=None, + ): + s = mock.MagicMock() + s.open = AsyncMock(side_effect=open_side_effect) + s.close = AsyncMock() + s.generation_number = generation + s.persisted_size = persisted_size + s.read_handle = read_handle + s.is_finalized = is_finalized + s.full_obj_server_crc32c = 12345 if is_finalized else None + s.object_metadata = mock.Mock() + return s + + @staticmethod + def _make_dummy_pooled_stream(): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + _PooledStream, + ) + + return _PooledStream( + mock.MagicMock(close=AsyncMock()), + mock.MagicMock(close=AsyncMock()), + ) + def create_read_ranges(self, num_ranges): ranges = [] for i in range(num_ranges): @@ -866,43 +902,32 @@ def test_is_read_retryable_predicate(self): assert _is_read_retryable(exceptions.PermissionDenied("denied")) is False assert _is_read_retryable(exceptions.InvalidArgument("invalid")) is False - def test_managed_stream_load(self): - from google.cloud.storage.asyncio.async_multi_range_downloader import ( - _ManagedStream, - ) - - mock_stream = mock.MagicMock() - mock_mux = mock.MagicMock() - worker = _ManagedStream(mock_stream, mock_mux) + def test_pooled_stream_load(self): + stream = self._make_dummy_pooled_stream() # Initially zero load - assert worker.calculate_load(target_io_depth=10, target_bytes=1000) == 0.0 + assert stream.calculate_load(target_io_depth=10, target_bytes=1000) == 0.0 # Add 5 ranges and 500 bytes -> 5/10 = 0.5, 500/1000 = 0.5 -> load = 0.5*0.5 + 0.5*0.5 = 0.5 - worker.record_request(5, 500) - assert worker.pending_ranges == 5 - assert worker.pending_bytes == 500 - assert worker.calculate_load(target_io_depth=10, target_bytes=1000) == 0.5 + stream.record_request(5, 500) + assert stream.pending_ranges == 5 + assert stream.pending_bytes == 500 + assert stream.calculate_load(target_io_depth=10, target_bytes=1000) == 0.5 # Release - worker.record_completion(3, 300) - assert worker.pending_ranges == 2 - assert worker.pending_bytes == 200 - assert worker.calculate_load(target_io_depth=10, target_bytes=1000) == 0.2 + stream.record_completion(3, 300) + assert stream.pending_ranges == 2 + assert stream.pending_bytes == 200 + assert stream.calculate_load(target_io_depth=10, target_bytes=1000) == 0.2 @pytest.mark.asyncio async def test_stream_pool_scale_up_and_dispatch(self): from google.cloud.storage.asyncio.async_multi_range_downloader import ( - _ManagedStream, _StreamPool, ) - created_workers = [] - async def stream_factory(): - w = _ManagedStream(mock.MagicMock(), mock.MagicMock()) - created_workers.append(w) - return w + return self._make_dummy_pooled_stream() pool = _StreamPool( stream_factory=stream_factory, @@ -911,48 +936,41 @@ async def stream_factory(): target_io_depth=2, target_bytes=100, ) - initial_worker = await stream_factory() - await pool.add_worker(initial_worker) + initial_stream = await stream_factory() + await pool.add_stream(initial_stream) - # 1. Acquire with small load (1 range, 10 bytes) -> load = 0.5*(1/2) + 0.5*(10/100) = 0.3 < 1.0 - w1 = await pool.acquire_stream(1, 10) - assert w1 == initial_worker - assert len(pool.workers) == 1 + # 1. Acquire with small load (1 range, 10 bytes) -> load < 1.0, stays on initial stream + s1 = await pool.acquire_stream(1, 10) + assert s1 == initial_stream + assert len(pool.streams) == 1 - # 2. Add enough load to exceed target load >= 1.0 (e.g. 3 ranges, 90 bytes) - # Total on w1: 4 ranges (hits target and triggers background scale up) - w1_again = await pool.acquire_stream(3, 90) - assert w1_again == initial_worker + # 2. Add enough load to exceed target load >= 1.0 -> triggers background scale up + s1_again = await pool.acquire_stream(3, 90) + assert s1_again == initial_stream - # Scale-up task was scheduled; allow event loop to run background task + # Allow event loop to run background scale-up task await asyncio.sleep(0.01) - assert len(pool.workers) == 2 - w2 = pool.workers[1] - assert w2 != initial_worker + assert len(pool.streams) == 2 + s2 = pool.streams[1] + assert s2 != initial_stream - # 3. Next acquire selects w2 because w2 has 0 load while w1 has high load - w_next = await pool.acquire_stream(1, 10) - assert w_next == w2 + # 3. Next acquire selects idle stream s2 + assert await pool.acquire_stream(1, 10) == s2 - # 4. Release worker capacity and close - pool.release_stream(w1, 4, 100) - pool.release_stream(w2, 1, 10) + # 4. Release stream capacity and close + pool.release_stream(s1, 4, 100) + pool.release_stream(s2, 1, 10) await pool.close() - assert len(pool.workers) == 0 + assert len(pool.streams) == 0 @pytest.mark.asyncio async def test_stream_pool_proportional_scale_up(self): from google.cloud.storage.asyncio.async_multi_range_downloader import ( - _ManagedStream, _StreamPool, ) - created_workers = [] - async def stream_factory(): - w = _ManagedStream(mock.MagicMock(), mock.MagicMock()) - created_workers.append(w) - return w + return self._make_dummy_pooled_stream() pool = _StreamPool( stream_factory=stream_factory, @@ -961,34 +979,30 @@ async def stream_factory(): target_io_depth=2, target_bytes=100, ) - initial_worker = await stream_factory() - await pool.add_worker(initial_worker) - - # Huge burst: 6 ranges, 300 bytes -> load = 0.5*(6/2) + 0.5*(300/100) = 3.0 - # Desired connections = ceil(3.0) = 3. - # Should launch 2 scale-ups concurrently in the background. - w1 = await pool.acquire_stream(6, 300) - assert w1 == initial_worker + initial_stream = await stream_factory() + await pool.add_stream(initial_stream) + + # Huge burst: 6 ranges, 300 bytes -> load = 3.0 -> triggers 2 concurrent background scale-ups + s1 = await pool.acquire_stream(6, 300) + assert s1 == initial_stream assert pool._pending_scale_ups == 2 assert len(pool._background_tasks) == 2 # Allow background scale-up tasks to finish await asyncio.sleep(0.01) - assert len(pool.workers) == 3 + assert len(pool.streams) == 3 assert pool._pending_scale_ups == 0 - # Another huge burst exceeding max_connections (5): - # 10 ranges, 500 bytes -> desired workers >= 5 - # Remaining headroom to max_connections is 5 - 3 = 2. + # Burst exceeding max_connections -> capped by remaining headroom (5 - 3 = 2) await pool.acquire_stream(10, 500) assert pool._pending_scale_ups == 2 await asyncio.sleep(0.01) - assert len(pool.workers) == 5 + assert len(pool.streams) == 5 assert pool._pending_scale_ups == 0 await pool.close() - assert len(pool.workers) == 0 + assert len(pool.streams) == 0 @mock.patch( "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" @@ -999,24 +1013,11 @@ async def test_create_mrd_with_stream_config(self, mock_cls_stream): MRDStreamConfig, ) - mock_client = mock.MagicMock() - mock_client.grpc_client = mock.AsyncMock() - - s1 = mock.MagicMock() - s1.open = AsyncMock() - s1.generation_number = 1 - s1.persisted_size = 100 - s1.read_handle = b"h1" - s1.object_metadata = mock.Mock() - - s2 = mock.MagicMock() - s2.open = AsyncMock() - s2.generation_number = 1 - s2.persisted_size = 100 - s2.read_handle = b"h2" - s2.object_metadata = mock.Mock() - - mock_cls_stream.side_effect = [s1, s2] + mock_client = self._make_mock_client() + mock_cls_stream.side_effect = [ + self._make_mock_stream(read_handle=b"h1"), + self._make_mock_stream(read_handle=b"h2"), + ] config = MRDStreamConfig( min_connections=2, @@ -1033,37 +1034,48 @@ async def test_create_mrd_with_stream_config(self, mock_cls_stream): assert mrd.stream_config.max_connections == 4 assert mrd.stream_config.target_io_depth == 8 assert mrd.stream_config.target_bytes == 2 * 1024 * 1024 - # Verified that 2 streams were opened initially for min_connections=2 - assert len(mrd._pool.workers) == 2 + # Verified that 2 streams were opened concurrently for min_connections=2 + assert len(mrd._pool.streams) == 2 await mrd.close() @mock.patch( "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" ) @pytest.mark.asyncio - async def test_mrd_download_ranges_triggers_pool_scaling(self, mock_cls_stream): + async def test_mrd_open_min_connections_failure_cleanup(self, mock_cls_stream): from google.cloud.storage.asyncio.async_multi_range_downloader import ( MRDStreamConfig, ) - mock_client = mock.MagicMock() - mock_client.grpc_client = mock.AsyncMock() + mock_client = self._make_mock_client() + mock_cls_stream.side_effect = [ + self._make_mock_stream(), + self._make_mock_stream( + open_side_effect=exceptions.ServiceUnavailable("Stream failed") + ), + ] - s1 = mock.MagicMock() - s1.open = AsyncMock() - s1.generation_number = 1 - s1.persisted_size = 1000 - s1.read_handle = b"h1" - s1.object_metadata = mock.Mock() + config = MRDStreamConfig(min_connections=2, max_connections=4) - s2 = mock.MagicMock() - s2.open = AsyncMock() - s2.generation_number = 1 - s2.persisted_size = 1000 - s2.read_handle = b"h2" - s2.object_metadata = mock.Mock() + with pytest.raises(exceptions.ServiceUnavailable, match="Stream failed"): + await AsyncMultiRangeDownloader.create_mrd( + mock_client, "b", "o", stream_config=config + ) - mock_cls_stream.side_effect = [s1, s2] + @mock.patch( + "google.cloud.storage.asyncio.async_multi_range_downloader._AsyncReadObjectStream" + ) + @pytest.mark.asyncio + async def test_mrd_download_ranges_triggers_pool_scaling(self, mock_cls_stream): + from google.cloud.storage.asyncio.async_multi_range_downloader import ( + MRDStreamConfig, + ) + + mock_client = self._make_mock_client() + mock_cls_stream.side_effect = [ + self._make_mock_stream(), + self._make_mock_stream(), + ] config = MRDStreamConfig( min_connections=1, @@ -1074,14 +1086,14 @@ async def test_mrd_download_ranges_triggers_pool_scaling(self, mock_cls_stream): mrd = await AsyncMultiRangeDownloader.create_mrd( mock_client, "b", "o", stream_config=config ) - assert len(mrd._pool.workers) == 1 + assert len(mrd._pool.streams) == 1 - with mock.patch.object(mrd, "_download_ranges_on_worker", new=AsyncMock()): + with mock.patch.object(mrd, "_download_ranges_on_stream", new=AsyncMock()): # Download ranges with load > 1.0 (2 ranges, 100 bytes) await mrd.download_ranges([(0, 50, BytesIO()), (50, 50, BytesIO())]) await asyncio.sleep(0.01) - # Pool dynamically scaled up to 2 workers - assert len(mrd._pool.workers) == 2 + # Pool dynamically scaled up to 2 streams + assert len(mrd._pool.streams) == 2 await mrd.close() @@ -1101,16 +1113,14 @@ async def test_unfinalized_object_download( mock_retry_manager_cls, mock_strategy_cls, ): - mock_client = mock.MagicMock() - mock_client.grpc_client = mock.AsyncMock() - - mock_stream = mock_cls_async_read_object_stream.return_value - mock_stream.open = AsyncMock() - mock_stream.generation_number = 123 - mock_stream.persisted_size = 50 - mock_stream.read_handle = b"handle" - mock_stream.is_finalized = False - mock_stream.full_obj_server_crc32c = None + mock_client = self._make_mock_client() + mock_stream = self._make_mock_stream( + generation=123, + persisted_size=50, + read_handle=b"handle", + is_finalized=False, + ) + mock_cls_async_read_object_stream.return_value = mock_stream mrd = await AsyncMultiRangeDownloader.create_mrd(mock_client, "b", "o") assert mrd.is_finalized is False @@ -1132,53 +1142,37 @@ async def test_unfinalized_object_download( ) @pytest.mark.asyncio async def test_create_mrd_single_stream_bypass(self, mock_cls_stream): - mock_client = mock.MagicMock() - mock_client.grpc_client = mock.AsyncMock() - - s1 = mock.MagicMock() - s1.open = AsyncMock() - s1.generation_number = 1 - s1.persisted_size = 100 - s1.read_handle = b"h1" - s1.object_metadata = mock.Mock() - mock_cls_stream.return_value = s1 + mock_client = self._make_mock_client() + mock_cls_stream.return_value = self._make_mock_stream() # Default create_mrd without stream_config -> single-stream bypass mrd = await AsyncMultiRangeDownloader.create_mrd(mock_client, "b", "o") assert mrd.stream_config is None assert mrd._pool is None - # download_ranges should use _primary_worker directly without pool + # download_ranges should use _primary_stream directly without pool with mock.patch.object( - mrd, "_download_ranges_on_worker", new=AsyncMock() + mrd, "_download_ranges_on_stream", new=AsyncMock() ) as mock_dl: await mrd.download_ranges([(0, 50, BytesIO())]) assert mock_dl.call_count == 1 - assert mock_dl.call_args[0][0] == mrd._primary_worker + assert mock_dl.call_args[0][0] == mrd._primary_stream await mrd.close() @pytest.mark.asyncio async def test_stream_pool_closed_and_cancellation(self): from google.cloud.storage.asyncio.async_multi_range_downloader import ( - _ManagedStream, _StreamPool, ) scale_up_started = asyncio.Event() scale_up_finish = asyncio.Event() - created_workers = [] async def slow_factory(): scale_up_started.set() await scale_up_finish.wait() - mock_stream = mock.MagicMock() - mock_stream.close = AsyncMock() - mock_mux = mock.MagicMock() - mock_mux.close = AsyncMock() - w = _ManagedStream(mock_stream, mock_mux) - created_workers.append(w) - return w + return self._make_dummy_pooled_stream() pool = _StreamPool( stream_factory=slow_factory, @@ -1187,8 +1181,8 @@ async def slow_factory(): target_io_depth=1, target_bytes=10, ) - initial_w = _ManagedStream(mock.MagicMock(), mock.MagicMock()) - await pool.add_worker(initial_w) + initial_s = self._make_dummy_pooled_stream() + await pool.add_stream(initial_s) # Trigger scale up await pool.acquire_stream(2, 20) @@ -1198,7 +1192,7 @@ async def slow_factory(): # Close pool while scale up task is pending await pool.close() assert pool._closed is True - assert len(pool.workers) == 0 + assert len(pool.streams) == 0 # Background task should be cancelled with pytest.raises(asyncio.CancelledError): @@ -1217,34 +1211,19 @@ async def slow_factory(): async def test_routing_token_preservation_and_propagation( self, mock_cls_async_read_object_stream ): - mock_client = mock.MagicMock() - mock_client.grpc_client = mock.AsyncMock() - - s1 = mock.MagicMock() - s1.open = AsyncMock() - s1.generation_number = 100 - s1.read_handle = b"h1" - s1.persisted_size = 1000 - s1.is_finalized = True - s1.full_obj_server_crc32c = 12345 + mock_client = self._make_mock_client() + s1 = self._make_mock_stream(read_handle=b"h1") mock_cls_async_read_object_stream.return_value = s1 mrd = await AsyncMultiRangeDownloader.create_mrd(mock_client, "b", "o") - # Simulate a redirect having set _routing_token mrd._routing_token = "token-abc" - # Now when a new stream worker is opened via _create_new_stream_worker, - # it should include routing_token in the metadata - s2 = mock.MagicMock() - s2.open = AsyncMock() - s2.generation_number = 100 - s2.read_handle = b"h2" - s2.persisted_size = 1000 + s2 = self._make_mock_stream(read_handle=b"h2") mock_cls_async_read_object_stream.return_value = s2 - w2 = await mrd._create_new_stream_worker() - assert w2 is not None - assert s2.open.call_count == 1 + s2_pooled = await mrd._create_pooled_stream() + assert s2_pooled is not None + s2.open.assert_called_once() call_kwargs = s2.open.call_args[1] assert ("x-goog-request-params", "routing_token=token-abc") in call_kwargs[ "metadata"