From 54f8e9eba575ecd0230339bd73fd2ee72893a8f0 Mon Sep 17 00:00:00 2001 From: ArrayRecord Team Date: Thu, 6 Aug 2026 15:32:05 -0700 Subject: [PATCH] No public description PiperOrigin-RevId: 960538077 --- python/array_record_data_source.py | 69 +++++++++++++++++++++++++ python/array_record_data_source_test.py | 2 + 2 files changed, 71 insertions(+) diff --git a/python/array_record_data_source.py b/python/array_record_data_source.py index 0eb7db6..7a73c4a 100644 --- a/python/array_record_data_source.py +++ b/python/array_record_data_source.py @@ -441,6 +441,7 @@ def __init__( ) self._read_instructions = _get_read_instructions(paths) self._paths = [ri.filename for ri in self._read_instructions] +<<<<<<< dest: 9be4115518a4 - efmo: Refactor missing candidate rem... self._reader_pool_size = ( reader_pool_size or _get_flag_value(_ARRAY_RECORD_READER_POOL_SIZE) or 1 # pyrefly: ignore[bad-argument-type] ) @@ -453,6 +454,14 @@ def __init__( for ri in self._read_instructions ] +||||||| parent of source: f245e4a325e1 - oncall-updater: Extending oncall for... + # We open readers lazily when we need to read from them. + self._readers = [None] * len(self._read_instructions) +======= + # We open readers lazily when we need to read from them. + self._readers = [None] * len(self._read_instructions) + self._lock = threading.Lock() +>>>>>>> source: d88ec7a6be72 - huangxia: Add thread synchronization... self._num_records = sum( map(lambda x: x.num_records, self._read_instructions) ) @@ -467,8 +476,21 @@ def __enter__(self): def __exit__(self, exc_type, exc_value, traceback): logging.debug("__exit__ for ArrayRecordDataSource is called.") +<<<<<<< dest: 9be4115518a4 - efmo: Refactor missing candidate rem... for pool in self._shard_pools: pool.close_all() +||||||| parent of source: f245e4a325e1 - oncall-updater: Extending oncall for... + for reader in self._readers: + if reader: + reader.close() + self._readers = [None] * len(self._read_instructions) +======= + with self._lock: + for reader in self._readers: + if reader: + reader.close() + self._readers = [None] * len(self._read_instructions) +>>>>>>> source: d88ec7a6be72 - huangxia: Add thread synchronization... def __len__(self) -> int: return self._num_records @@ -508,6 +530,7 @@ def _split_keys_per_reader( positions_and_indices[reader_idx] = [(position, idx)] return positions_and_indices +<<<<<<< dest: 9be4115518a4 - efmo: Refactor missing candidate rem... def _read_record(self, reader: Any, position: int) -> bytes: """Helper to read a record using the best available method.""" if hasattr(reader, "read_record"): @@ -515,6 +538,37 @@ def _read_record(self, reader: Any, position: int) -> bytes: if hasattr(reader, "read"): return reader.read([position])[0] return reader[position] +||||||| parent of source: f245e4a325e1 - oncall-updater: Extending oncall for... + def _ensure_reader_exists(self, reader_idx: int) -> None: + """Threadsafe method to create corresponding reader if it doesn't exist.""" + if self._readers[reader_idx] is not None: + return + filename = self._read_instructions[reader_idx].filename + reader = _create_reader(filename, self._reader_options_string) + _check_group_size(filename, reader) + self._readers[reader_idx] = reader +======= + def _ensure_reader_exists(self, reader_idx: int) -> None: + """Threadsafe method to create corresponding reader if it doesn't exist.""" + if not (0 <= reader_idx < len(self._readers)): + raise IndexError( + f"reader_idx {reader_idx} out of range [0, {len(self._readers)})" + ) + if self._readers[reader_idx] is not None: + return + with self._lock: + if self._readers[reader_idx] is not None: + return + filename = self._read_instructions[reader_idx].filename + reader = _create_reader(filename, self._reader_options_string) + try: + _check_group_size(filename, reader) + except Exception: + if hasattr(reader, "close"): + reader.close() + raise + self._readers[reader_idx] = reader +>>>>>>> source: d88ec7a6be72 - huangxia: Add thread synchronization... def __getitem__(self, record_key: SupportsIndex) -> bytes: pool_idx, position = self._reader_idx_and_position(record_key) @@ -565,7 +619,15 @@ def read_records( def __getstate__(self): logging.debug("__getstate__ for ArrayRecordDataSource is called.") state = self.__dict__.copy() +<<<<<<< dest: 9be4115518a4 - efmo: Refactor missing candidate rem... state.pop("_shard_pools", None) +||||||| parent of source: f245e4a325e1 - oncall-updater: Extending oncall for... + del state["_readers"] +======= + del state["_readers"] + if "_lock" in state: + del state["_lock"] +>>>>>>> source: d88ec7a6be72 - huangxia: Add thread synchronization... return state def __setstate__(self, state): @@ -573,6 +635,7 @@ def __setstate__(self, state): self.__dict__.update(state) # We open readers lazily when we need to read from them. Thus, we don't # need to re-open the same files as before pickling. +<<<<<<< dest: 9be4115518a4 - efmo: Refactor missing candidate rem... self._shard_pools = [ _BoundedReaderPool( ri.filename, @@ -581,6 +644,12 @@ def __setstate__(self, state): ) for ri in self._read_instructions ] +||||||| parent of source: f245e4a325e1 - oncall-updater: Extending oncall for... + self._readers = [None] * len(self._read_instructions) +======= + self._readers = [None] * len(self._read_instructions) + self._lock = threading.Lock() +>>>>>>> source: d88ec7a6be72 - huangxia: Add thread synchronization... def __repr__(self) -> str: """Storing a hash of paths since paths can be a very long list.""" diff --git a/python/array_record_data_source_test.py b/python/array_record_data_source_test.py index b4de730..5b27d60 100644 --- a/python/array_record_data_source_test.py +++ b/python/array_record_data_source_test.py @@ -18,6 +18,8 @@ import os import pathlib import pickle +import threading +import time from unittest import mock from absl import flags