Skip to content

Commit b183f7f

Browse files
committed
minor refactoring
1 parent b7fa233 commit b183f7f

2 files changed

Lines changed: 218 additions & 221 deletions

File tree

‎packages/google-cloud-storage/google/cloud/storage/asyncio/async_multi_range_downloader.py‎

Lines changed: 79 additions & 61 deletions
Original file line numberDiff line numberDiff line change
@@ -104,7 +104,7 @@ class MRDStreamConfig:
104104
target_bytes: int = 8 * 1024 * 1024
105105

106106

107-
class _ManagedStream:
107+
class _PooledStream:
108108
"""Wraps an active stream and its multiplexer with local load counters."""
109109

110110
def __init__(
@@ -146,11 +146,11 @@ async def close(self) -> None:
146146

147147

148148
class _StreamPool:
149-
"""Manages a pool of _ManagedStream instances with dynamic scaling and least-loaded dispatch."""
149+
"""Manages a pool of _PooledStream instances with dynamic scaling and least-loaded dispatch."""
150150

151151
def __init__(
152152
self,
153-
stream_factory: Callable[[], Awaitable[_ManagedStream]],
153+
stream_factory: Callable[[], Awaitable[_PooledStream]],
154154
min_connections: int = 1,
155155
max_connections: int = 8,
156156
target_io_depth: int = 8,
@@ -162,38 +162,38 @@ def __init__(
162162
self.target_io_depth = target_io_depth
163163
self.target_bytes = target_bytes
164164

165-
self.workers: List[_ManagedStream] = []
165+
self.streams: List[_PooledStream] = []
166166
self._lock = asyncio.Lock()
167167
self._pending_scale_ups: int = 0
168168
self._closed = False
169169
self._background_tasks: set[asyncio.Task] = set()
170170

171-
async def add_worker(self, worker: _ManagedStream) -> None:
171+
async def add_stream(self, pooled_stream: _PooledStream) -> None:
172172
async with self._lock:
173-
self.workers.append(worker)
173+
self.streams.append(pooled_stream)
174174

175-
async def acquire_stream(self, range_count: int, req_bytes: int) -> _ManagedStream:
175+
async def acquire_stream(self, range_count: int, req_bytes: int) -> _PooledStream:
176176
"""Finds least loaded stream and triggers background scale-up proportional to load."""
177177
async with self._lock:
178178
if self._closed:
179179
raise ValueError("Pool is closed")
180-
if not self.workers:
181-
raise ValueError("No workers available in stream pool")
180+
if not self.streams:
181+
raise ValueError("No streams available in stream pool")
182182
best = min(
183-
self.workers,
184-
key=lambda w: w.calculate_load(self.target_io_depth, self.target_bytes),
183+
self.streams,
184+
key=lambda s: s.calculate_load(self.target_io_depth, self.target_bytes),
185185
)
186186
best.record_request(range_count, req_bytes)
187187

188-
# Trigger background scale-ups proportional to total load across all workers
188+
# Trigger background scale-ups proportional to total load across all streams
189189
total_load = sum(
190-
w.calculate_load(self.target_io_depth, self.target_bytes)
191-
for w in self.workers
190+
s.calculate_load(self.target_io_depth, self.target_bytes)
191+
for s in self.streams
192192
)
193-
desired_workers = math.ceil(total_load)
194-
planned_workers = len(self.workers) + self._pending_scale_ups
193+
desired_streams = math.ceil(total_load)
194+
planned_streams = len(self.streams) + self._pending_scale_ups
195195
needed_scale_ups = max(
196-
0, min(desired_workers, self.max_connections) - planned_workers
196+
0, min(desired_streams, self.max_connections) - planned_streams
197197
)
198198

199199
for _ in range(needed_scale_ups):
@@ -206,33 +206,33 @@ async def acquire_stream(self, range_count: int, req_bytes: int) -> _ManagedStre
206206

207207
async def _scale_up(self) -> None:
208208
try:
209-
new_worker = await self._stream_factory()
209+
new_stream = await self._stream_factory()
210210
async with self._lock:
211211
if self._closed:
212-
await new_worker.close()
212+
await new_stream.close()
213213
return
214-
self.workers.append(new_worker)
214+
self.streams.append(new_stream)
215215
except Exception as e:
216216
logger.warning(f"Failed to scale up MRD stream: {e}")
217217
finally:
218218
async with self._lock:
219219
self._pending_scale_ups = max(0, self._pending_scale_ups - 1)
220220

221221
def release_stream(
222-
self, worker: _ManagedStream, range_count: int, req_bytes: int
222+
self, pooled_stream: _PooledStream, range_count: int, req_bytes: int
223223
) -> None:
224-
worker.record_completion(range_count, req_bytes)
224+
pooled_stream.record_completion(range_count, req_bytes)
225225

226226
async def close(self) -> None:
227227
async with self._lock:
228228
self._closed = True
229229
self._pending_scale_ups = 0
230-
workers = list(self.workers)
231-
self.workers.clear()
230+
streams = list(self.streams)
231+
self.streams.clear()
232232
for task in list(self._background_tasks):
233233
task.cancel()
234-
for w in workers:
235-
await w.close()
234+
for s in streams:
235+
await s.close()
236236

237237

238238
class AsyncMultiRangeDownloader:
@@ -388,7 +388,7 @@ def __init__(
388388

389389
self.stream_config = stream_config
390390
self._pool: Optional[_StreamPool] = None
391-
self._primary_worker: Optional[_ManagedStream] = None
391+
self._primary_stream: Optional[_PooledStream] = None
392392
self._metadata: Optional[List[Tuple[str, str]]] = None
393393

394394
async def __aenter__(self):
@@ -493,26 +493,44 @@ async def _do_open():
493493
self._metadata = list(metadata) if metadata else []
494494
await retry_policy(_do_open)()
495495
self._multiplexer = _StreamMultiplexer(self.read_obj_str)
496-
self._primary_worker = _ManagedStream(self.read_obj_str, self._multiplexer)
496+
self._primary_stream = _PooledStream(self.read_obj_str, self._multiplexer)
497497

498498
if self.stream_config is not None and self.stream_config.max_connections > 1:
499499
self._pool = _StreamPool(
500-
stream_factory=self._create_new_stream_worker,
500+
stream_factory=self._create_pooled_stream,
501501
min_connections=self.stream_config.min_connections,
502502
max_connections=self.stream_config.max_connections,
503503
target_io_depth=self.stream_config.target_io_depth,
504504
target_bytes=self.stream_config.target_bytes,
505505
)
506-
await self._pool.add_worker(self._primary_worker)
507-
508-
for _ in range(self.stream_config.min_connections - 1):
509-
worker = await self._create_new_stream_worker()
510-
await self._pool.add_worker(worker)
506+
await self._pool.add_stream(self._primary_stream)
507+
508+
if self.stream_config.min_connections > 1:
509+
extra_streams = await asyncio.gather(
510+
*(
511+
self._create_pooled_stream()
512+
for _ in range(self.stream_config.min_connections - 1)
513+
),
514+
return_exceptions=True,
515+
)
516+
first_exc = next(
517+
(s for s in extra_streams if isinstance(s, Exception)), None
518+
)
519+
if first_exc is not None:
520+
for s in extra_streams:
521+
if not isinstance(s, Exception):
522+
await s.close()
523+
await self._pool.close()
524+
self._pool = None
525+
raise first_exc
526+
527+
for stream in extra_streams:
528+
await self._pool.add_stream(stream)
511529
else:
512530
self._pool = None
513531

514-
async def _create_new_stream_worker(self) -> _ManagedStream:
515-
"""Opens an additional stream worker using current routing and read_handle."""
532+
async def _create_pooled_stream(self) -> _PooledStream:
533+
"""Opens an additional pooled stream using current routing and read_handle."""
516534
current_metadata = list(self._metadata) if self._metadata else []
517535
if self._routing_token:
518536
current_metadata.append(
@@ -534,11 +552,11 @@ async def _create_new_stream_worker(self) -> _ManagedStream:
534552
self.read_handle = stream.read_handle
535553

536554
mux = _StreamMultiplexer(stream)
537-
return _ManagedStream(stream, mux)
555+
return _PooledStream(stream, mux)
538556

539-
def _create_stream_factory(self, state, metadata, worker=None):
557+
def _create_stream_factory(self, state, metadata, pooled_stream=None):
540558
"""Create a factory that opens a new stream with current routing state."""
541-
target_worker = worker or self._primary_worker
559+
target_stream = pooled_stream or self._primary_stream
542560

543561
async def factory():
544562
current_handle = state.get("read_handle") or self.read_handle
@@ -571,19 +589,19 @@ async def factory():
571589
self.full_obj_server_crc32c = stream.full_obj_server_crc32c
572590

573591
self.read_obj_str = stream
574-
if target_worker is not None:
575-
target_worker.stream = stream
576-
if target_worker is None or target_worker == self._primary_worker:
592+
if target_stream is not None:
593+
target_stream.stream = stream
594+
if target_stream is None or target_stream == self._primary_stream:
577595
self.read_obj_str = stream
578596
self._is_stream_open = True
579597

580598
return stream
581599

582600
return factory
583601

584-
async def _download_ranges_on_worker(
602+
async def _download_ranges_on_stream(
585603
self,
586-
worker: _ManagedStream,
604+
pooled_stream: _PooledStream,
587605
read_ranges: List[Tuple[int, int, BytesIO]],
588606
retry_policy: AsyncRetry,
589607
metadata: Optional[List[Tuple[str, str]]],
@@ -669,7 +687,7 @@ async def _download_ranges_on_worker(
669687
}
670688

671689
read_ids = set(download_states.keys())
672-
queue = worker.multiplexer.register(read_ids)
690+
queue = pooled_stream.multiplexer.register(read_ids)
673691

674692
try:
675693
attempt_count = 0
@@ -696,16 +714,16 @@ async def generator():
696714
broken_gen = (
697715
last_broken_generation
698716
if attempt_count > 1
699-
else worker.multiplexer.stream_generation
717+
else pooled_stream.multiplexer.stream_generation
700718
)
701719
stream_factory = self._create_stream_factory(
702-
state, metadata, worker=worker
720+
state, metadata, pooled_stream=pooled_stream
703721
)
704-
await worker.multiplexer.reopen_stream(
722+
await pooled_stream.multiplexer.reopen_stream(
705723
broken_gen, stream_factory
706724
)
707725

708-
stream_generation = worker.multiplexer.stream_generation
726+
stream_generation = pooled_stream.multiplexer.stream_generation
709727

710728
# Send Requests
711729
pending_read_ids = {r.read_id for r in requests}
@@ -714,7 +732,7 @@ async def generator():
714732
):
715733
batch = requests[i : i + _MAX_READ_RANGES_PER_BIDI_READ_REQUEST]
716734
try:
717-
await worker.multiplexer.send(
735+
await pooled_stream.multiplexer.send(
718736
_storage_v2.BidiReadObjectRequest(read_ranges=batch)
719737
)
720738
except Exception:
@@ -765,8 +783,8 @@ async def generator():
765783
if initial_state.get("read_handle"):
766784
self.read_handle = initial_state["read_handle"]
767785
finally:
768-
if worker.multiplexer is not None:
769-
worker.multiplexer.unregister(read_ids)
786+
if pooled_stream.multiplexer is not None:
787+
pooled_stream.multiplexer.unregister(read_ids)
770788

771789
async def download_ranges(
772790
self,
@@ -826,22 +844,22 @@ async def download_ranges(
826844
retry_policy = AsyncRetry(predicate=_is_read_retryable)
827845

828846
# Fallback for manually mocked tests that set mrd._multiplexer without calling open()
829-
if self._primary_worker is None and self.read_obj_str and self._multiplexer:
830-
self._primary_worker = _ManagedStream(self.read_obj_str, self._multiplexer)
847+
if self._primary_stream is None and self.read_obj_str and self._multiplexer:
848+
self._primary_stream = _PooledStream(self.read_obj_str, self._multiplexer)
831849

832850
if self._pool is not None:
833851
total_bytes = sum(length for _, length, _ in read_ranges)
834852
total_ranges = len(read_ranges)
835-
worker = await self._pool.acquire_stream(total_ranges, total_bytes)
853+
pooled_stream = await self._pool.acquire_stream(total_ranges, total_bytes)
836854
try:
837-
await self._download_ranges_on_worker(
838-
worker, read_ranges, retry_policy, metadata, enable_checksum
855+
await self._download_ranges_on_stream(
856+
pooled_stream, read_ranges, retry_policy, metadata, enable_checksum
839857
)
840858
finally:
841-
self._pool.release_stream(worker, total_ranges, total_bytes)
859+
self._pool.release_stream(pooled_stream, total_ranges, total_bytes)
842860
else:
843-
await self._download_ranges_on_worker(
844-
self._primary_worker,
861+
await self._download_ranges_on_stream(
862+
self._primary_stream,
845863
read_ranges,
846864
retry_policy,
847865
metadata,
@@ -872,7 +890,7 @@ async def close(self):
872890
except (ValueError, asyncio.CancelledError, exceptions.GoogleAPICallError):
873891
pass
874892
self.read_obj_str = None
875-
self._primary_worker = None
893+
self._primary_stream = None
876894
self._is_stream_open = False
877895

878896
@property

0 commit comments

Comments
 (0)