Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 69 additions & 0 deletions python/array_record_data_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
)
Expand All @@ -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)
)
Expand All @@ -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
Expand Down Expand Up @@ -508,13 +530,45 @@ 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"):
return reader.read_record(position)
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)
Expand Down Expand Up @@ -565,14 +619,23 @@ 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):
logging.debug("__setstate__ for ArrayRecordDataSource is called.")
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,
Expand All @@ -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."""
Expand Down
2 changes: 2 additions & 0 deletions python/array_record_data_source_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
import os
import pathlib
import pickle
import threading
import time
from unittest import mock

from absl import flags
Expand Down
Loading