@@ -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
148148class _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
238238class 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