@@ -38,6 +38,7 @@ def __getitem__(self, record_keys: Sequence[int]) -> Sequence[T]:
3838import itertools
3939import os
4040import pathlib
41+ import queue
4142import re
4243import typing
4344from typing import Any , Callable , Iterator , List , Mapping , Protocol , Sequence , SupportsIndex , Tuple , TypeVar , Union
@@ -271,7 +272,9 @@ def __init__(
271272 self ._read_instructions = _get_read_instructions (paths )
272273 self ._paths = [ri .filename for ri in self ._read_instructions ]
273274 # We open readers lazily when we need to read from them.
274- self ._readers = [None ] * len (self ._read_instructions )
275+ # To ensure thread safety while allowing concurrent reads, we maintain a
276+ # pool of readers for each shard.
277+ self ._reader_pools = [queue .LifoQueue () for _ in self ._read_instructions ]
275278 self ._num_records = sum (
276279 map (lambda x : x .num_records , self ._read_instructions )
277280 )
@@ -286,10 +289,11 @@ def __enter__(self):
286289
287290 def __exit__ (self , exc_type , exc_value , traceback ):
288291 logging .debug ("__exit__ for ArrayRecordDataSource is called." )
289- for reader in self ._readers :
290- if reader :
291- reader .close ()
292- self ._readers = [None ] * len (self ._read_instructions )
292+ for pool in self ._reader_pools :
293+ while not pool .empty ():
294+ reader = pool .get ()
295+ if reader :
296+ reader .close ()
293297
294298 def __len__ (self ) -> int :
295299 return self ._num_records
@@ -329,21 +333,29 @@ def _split_keys_per_reader(
329333 positions_and_indices [reader_idx ] = [(position , idx )]
330334 return positions_and_indices
331335
332- def _ensure_reader_exists (self , reader_idx : int ) -> None :
333- """Threadsafe method to create corresponding reader if it doesn't exist."""
334- if self ._readers [reader_idx ] is not None :
335- return
336- filename = self ._read_instructions [reader_idx ].filename
337- reader = _create_reader (filename , self ._reader_options_string )
338- _check_group_size (filename , reader )
339- self ._readers [reader_idx ] = reader
336+ def _get_reader (self , reader_idx : int ) -> Any :
337+ """Gets a reader from the pool or creates a new one."""
338+ try :
339+ return self ._reader_pools [reader_idx ].get_nowait ()
340+ except queue .Empty :
341+ filename = self ._read_instructions [reader_idx ].filename
342+ reader = _create_reader (filename , self ._reader_options_string )
343+ _check_group_size (filename , reader )
344+ return reader
345+
346+ def _release_reader (self , reader_idx : int , reader : Any ) -> None :
347+ """Returns a reader to the pool."""
348+ self ._reader_pools [reader_idx ].put (reader )
340349
341350 def __getitem__ (self , record_key : SupportsIndex ) -> bytes :
342351 reader_idx , position = self ._reader_idx_and_position (record_key )
343- self ._ensure_reader_exists (reader_idx )
344- if hasattr (self ._readers [reader_idx ], "read" ):
345- return self ._readers [reader_idx ].read ([position ])[0 ]
346- return self ._readers [reader_idx ][position ]
352+ reader = self ._get_reader (reader_idx )
353+ try :
354+ if hasattr (reader , "read" ):
355+ return reader .read ([position ])[0 ]
356+ return reader [position ]
357+ finally :
358+ self ._release_reader (reader_idx , reader )
347359
348360 def __getitems__ (
349361 self , record_keys : Sequence [SupportsIndex ]
@@ -352,14 +364,16 @@ def read_records(
352364 reader_idx : int , reader_positions_and_indices : Sequence [Tuple [int , int ]]
353365 ) -> Sequence [Tuple [Any , int ]]:
354366 """Reads records using the given reader keeping track of the indices."""
355- # Initialize readers lazily when we need to read from them.
356- self ._ensure_reader_exists (reader_idx )
357- positions , indices = list (zip (* reader_positions_and_indices ))
358- if hasattr (self ._readers [reader_idx ], "read" ):
359- records = self ._readers [reader_idx ].read (positions ) # pytype: disable=attribute-error
360- else :
361- records = [self ._readers [reader_idx ][p ] for p in positions ]
362- return list (zip (records , indices ))
367+ reader = self ._get_reader (reader_idx )
368+ try :
369+ positions , indices = list (zip (* reader_positions_and_indices ))
370+ if hasattr (reader , "read" ):
371+ records = reader .read (positions ) # pytype: disable=attribute-error
372+ else :
373+ records = [reader [p ] for p in positions ]
374+ return list (zip (records , indices ))
375+ finally :
376+ self ._release_reader (reader_idx , reader )
363377
364378 positions_and_indices = self ._split_keys_per_reader (record_keys )
365379 num_threads = _get_flag_value (_GRAIN_NUM_THREADS_FETCHING_RECORDS )
@@ -390,15 +404,23 @@ def read_records(
390404 def __getstate__ (self ):
391405 logging .debug ("__getstate__ for ArrayRecordDataSource is called." )
392406 state = self .__dict__ .copy ()
393- del state ["_readers " ]
407+ del state ["_reader_pools " ]
394408 return state
395409
396410 def __setstate__ (self , state ):
397411 logging .debug ("__setstate__ for ArrayRecordDataSource is called." )
398412 self .__dict__ .update (state )
399413 # We open readers lazily when we need to read from them. Thus, we don't
400414 # need to re-open the same files as before pickling.
401- self ._readers = [None ] * len (self ._read_instructions )
415+ self ._reader_pools = [queue .LifoQueue () for _ in self ._read_instructions ]
416+
417+ def _peek_readers (self ) -> List [Any ]:
418+ """Returns a list of readers (one per shard) or None (for testing)."""
419+ readers = []
420+ for pool in self ._reader_pools :
421+ with pool .mutex :
422+ readers .append (pool .queue [- 1 ] if pool .queue else None )
423+ return readers
402424
403425 def __repr__ (self ) -> str :
404426 """Storing a hash of paths since paths can be a very long list."""
0 commit comments