diff --git a/docs/docs/pypaimon/blob.md b/docs/docs/pypaimon/blob.md index ddb9c40be4b2..a2a4b120fb07 100644 --- a/docs/docs/pypaimon/blob.md +++ b/docs/docs/pypaimon/blob.md @@ -136,6 +136,13 @@ Without `blob-as-descriptor=true`, blob values are materialized before `row.get_blob(...)` returns; `new_input_stream()` then reads from in-memory bytes, not from storage. +For data-evolution reads, PyPaimon applies user filters, row-level authorization +filters, and limits before materializing projected scalar BLOB payloads. A user +or authorization filter that references a BLOB value keeps that field eager. +Column masking is applied after payload materialization. Set +`read.defer-blob-resolve=false` to restore eager materialization. ARRAY and MAP +elements containing BLOB values are not deferred. + ## Lower-level: `Blob.from_bytes` When you already have raw or descriptor bytes (for example from a custom diff --git a/docs/docs/pypaimon/pytorch.md b/docs/docs/pypaimon/pytorch.md index 0f7f7bbdef6d..185b997b7b53 100644 --- a/docs/docs/pypaimon/pytorch.md +++ b/docs/docs/pypaimon/pytorch.md @@ -59,6 +59,10 @@ when it is false, it will read the full amount of data into memory. **`prefetch_concurrency`** (default: 1): When streaming is true, number of threads used for parallel prefetch within each DataLoader worker. Set to a value greater than 1 to partition splits across threads and increase read throughput. Has no effect when streaming is false. +When the read builder has a `LIMIT`, streaming reads use one DataLoader worker +and disable split prefetch fan-out. This preserves one global remaining-row +quota and avoids materializing BLOB payloads beyond the limit. + ## Shuffle PyPaimon supports streaming shuffle for PyTorch `IterableDataset`. The shuffle diff --git a/paimon-python/pypaimon/common/options/core_options.py b/paimon-python/pypaimon/common/options/core_options.py index 3b890cc0d62d..55124ba51f68 100644 --- a/paimon-python/pypaimon/common/options/core_options.py +++ b/paimon-python/pypaimon/common/options/core_options.py @@ -903,6 +903,16 @@ class CoreOptions: ) ) + READ_DEFER_BLOB_RESOLVE: ConfigOption[bool] = ( + ConfigOptions.key("read.defer-blob-resolve") + .boolean_type() + .default_value(True) + .with_description( + "Whether filtered or limited data-evolution reads should apply " + "row selection before materializing projected scalar BLOB payloads." + ) + ) + READ_BATCH_SIZE: ConfigOption[int] = ( ConfigOptions.key("read.batch-size") .int_type() @@ -1526,6 +1536,9 @@ def local_cache_block_size(self) -> MemorySize: def local_cache_whitelist(self) -> str: return self.options.get(CoreOptions.LOCAL_CACHE_WHITELIST) + def read_defer_blob_resolve(self) -> bool: + return self.options.get(CoreOptions.READ_DEFER_BLOB_RESOLVE) + def read_batch_size(self, default=None) -> int: return self.options.get(CoreOptions.READ_BATCH_SIZE, default or 1024) diff --git a/paimon-python/pypaimon/read/datasource/torch_dataset.py b/paimon-python/pypaimon/read/datasource/torch_dataset.py index 5eb3485dddd1..78dad9258618 100644 --- a/paimon-python/pypaimon/read/datasource/torch_dataset.py +++ b/paimon-python/pypaimon/read/datasource/torch_dataset.py @@ -101,6 +101,11 @@ def _row_to_dict(self, offset_row) -> dict: return row_dict def _worker_splits(self, worker_info) -> List[Split]: + if self.table_read.limit is not None: + if worker_info is None or worker_info.id == 0: + return self.splits + return [] + if worker_info is None: return self.splits @@ -164,7 +169,7 @@ def __iter__(self): worker_info = torch.utils.data.get_worker_info() splits_to_process = self._worker_splits(worker_info) - if self.prefetch_concurrency > 1: + if self.prefetch_concurrency > 1 and self.table_read.limit is None: for row in self._iter_rows(splits_to_process): yield row return @@ -288,7 +293,8 @@ def __iter__(self): worker_id = worker_info.id if worker_info is not None else 0 splits_to_process = self._worker_splits(worker_info) - if self.max_buffer_input_splits == 1: + if (self.table_read.limit is not None + or self.max_buffer_input_splits == 1): rows = self._iter_ordered_rows(splits_to_process) else: rows = self._iter_interleaved_rows(splits_to_process) diff --git a/paimon-python/pypaimon/read/reader/deferred_blob_resolve_reader.py b/paimon-python/pypaimon/read/reader/deferred_blob_resolve_reader.py new file mode 100644 index 000000000000..b52c47c051b8 --- /dev/null +++ b/paimon-python/pypaimon/read/reader/deferred_blob_resolve_reader.py @@ -0,0 +1,75 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +from typing import List, Optional + +import pyarrow as pa +from pyarrow import RecordBatch + +from pypaimon.read.reader.iface.record_batch_reader import RecordBatchReader +from pypaimon.table.row.blob import Blob + + +class DeferredBlobResolveReader(RecordBatchReader): + """Materialize projected BLOB payloads after row filtering. + + This must remain the outermost BLOB materialization layer because adopted + metadata still identifies the materialized columns as logical BLOB fields. + """ + + def __init__(self, inner: RecordBatchReader, file_io, + blob_field_names: List[str], blob_parallelism: int = 1): + self._inner = inner + self._file_io = file_io + self._blob_field_names = blob_field_names + self._blob_parallelism = max(1, blob_parallelism) + self._adopt_metadata(inner) + + def read_arrow_batch(self) -> Optional[RecordBatch]: + batch = self._inner.read_arrow_batch() + if batch is None: + return None + + columns = list(batch.columns) + fields = list(batch.schema) + changed = False + for field_name in self._blob_field_names: + column_index = batch.schema.get_field_index(field_name) + if column_index < 0: + continue + values = batch.column(column_index).to_pylist() + blobs = [Blob.from_bytes(value, self._file_io) for value in values] + payloads = self._file_io.read_blobs_concurrent( + blobs, self._blob_parallelism) + source_field = batch.schema.field(column_index) + columns[column_index] = pa.array(payloads, type=pa.large_binary()) + fields[column_index] = pa.field( + field_name, + pa.large_binary(), + nullable=source_field.nullable, + metadata=source_field.metadata, + ) + changed = True + if not changed: + return batch + return pa.RecordBatch.from_arrays( + columns, + schema=pa.schema(fields, metadata=batch.schema.metadata), + ) + + def close(self) -> None: + self._inner.close() diff --git a/paimon-python/pypaimon/read/split_read.py b/paimon-python/pypaimon/read/split_read.py index 6ced11f0c127..0f1ee79568a9 100644 --- a/paimon-python/pypaimon/read/split_read.py +++ b/paimon-python/pypaimon/read/split_read.py @@ -42,7 +42,10 @@ MergeAllBatchReader, DataEvolutionMergeReader) from pypaimon.read.reader.concat_record_reader import ConcatRecordReader +from pypaimon.read.reader.auth_masking_reader import AuthFilterReader from pypaimon.read.reader.data_file_batch_reader import DataFileBatchReader +from pypaimon.read.reader.deferred_blob_resolve_reader import \ + DeferredBlobResolveReader from pypaimon.read.reader.drop_delete_reader import DropDeleteRecordReader from pypaimon.read.reader.empty_record_reader import EmptyFileRecordReader from pypaimon.read.reader.field_bunch import BlobBunch, DataBunch, FieldBunch, VectorBunch @@ -81,6 +84,33 @@ KEY_PREFIX = "_KEY_" KEY_FIELD_ID_START = 1000000 NULL_FIELD_INDEX = -1 + + +def deferred_blob_field_names(table, read_fields: List[DataField], + predicate: Optional[Predicate], + limit: Optional[int], + has_post_filter: bool = False) -> set: + # An auth filter also selects rows; defer past it too, like a predicate/limit. + if ((predicate is None and limit is None and not has_post_filter) + or CoreOptions.blob_as_descriptor(table.options) + or not table.options.read_defer_blob_resolve()): + return set() + + inline_fields = ( + CoreOptions.blob_descriptor_fields(table.options) + | CoreOptions.blob_view_fields(table.options) + ) + predicate_fields = ( + predicate_field_names(predicate) if predicate is not None else set() + ) + return { + read_fields[index].name + for index in blob_field_indices(read_fields) + if read_fields[index].name not in inline_fields + and read_fields[index].name not in predicate_fields + } + + ROW_SIDECAR_FORMAT = CoreOptions.FILE_FORMAT_ROW _COMPRESS_EXTENSIONS = frozenset(['gz', 'bz2', 'deflate', 'snappy', 'lz4', 'zst']) @@ -283,7 +313,7 @@ def file_reader_supplier(self, file: DataFileMeta, for_merge_read: bool, if has_nested: raise NotImplementedError( "Nested-field projection is not supported on BLOB files") - blob_as_descriptor = CoreOptions.blob_as_descriptor(self.table.options) + blob_as_descriptor = self._read_blob_as_descriptor(read_file_fields) blob_parallelism = self._blob_parallelism format_reader = FormatBlobReader(self.table.file_io, file_path, read_file_fields, self.read_fields, read_arrow_predicate, blob_as_descriptor, @@ -419,6 +449,12 @@ def file_reader_supplier(self, file: DataFileMeta, for_merge_read: bool, return reader + def _read_blob_as_descriptor(self, field_names: List[str]) -> bool: + if CoreOptions.blob_as_descriptor(self.table.options): + return True + deferred_fields = getattr(self, '_deferred_blob_fields', set()) + return any(field_name in deferred_fields for field_name in field_names) + @staticmethod def _row_sidecar_file_name(file: DataFileMeta) -> Optional[str]: row_files = [ @@ -995,7 +1031,10 @@ def __init__( nested_name_paths: Optional[List[List[str]]] = None, limit: Optional[int] = None, outer_extract_name_paths: Optional[List[List[str]]] = None, - outer_flat_read_type: Optional[List[DataField]] = None): + outer_flat_read_type: Optional[List[DataField]] = None, + post_merge_filter=None, + eager_blob_fields=None, + post_filter_after_inline=False): self.row_ranges = None actual_split = split if isinstance(split, IndexedSplit): @@ -1008,6 +1047,11 @@ def __init__( ) self.outer_extract_name_paths = outer_extract_name_paths self.outer_flat_read_type = outer_flat_read_type + self._post_merge_filter = post_merge_filter + # Apply the auth filter after inline BLOB resolution, so scalar BLOBs still defer. + self._post_filter_after_inline = post_filter_after_inline + self._eager_blob_fields = set(eager_blob_fields or []) + self._deferred_blob_fields = self._deferred_blob_field_names() def _push_down_predicate(self) -> Optional[Predicate]: # Data evolution: files may have different schemas, so we don't push predicate @@ -1027,8 +1071,35 @@ def create_reader(self) -> RecordReader: prescan_reader_factory=lambda names: self._create_prescan_reader(names), blob_parallelism=blob_parallelism) + if self._post_filter_after_inline: + if self._post_merge_filter is not None: + reader = AuthFilterReader(reader, self._post_merge_filter) + if self.limit is not None: + reader = LimitedRecordBatchReader(reader, self.limit) + + if self._deferred_blob_fields: + blob_names = [ + field.name for field in self.read_fields + if field.name in self._deferred_blob_fields + ] + reader = DeferredBlobResolveReader( + reader, + self.table.file_io, + blob_names, + blob_parallelism=self._blob_parallelism, + ) + return reader + def _deferred_blob_field_names(self) -> set: + return deferred_blob_field_names( + self.table, + self.read_fields, + self.predicate_for_reader, + self.limit, + has_post_filter=self._post_merge_filter is not None, + ) - self._eager_blob_fields + def _create_raw_reader(self) -> RecordReader: """Core read logic: split_by_row_id -> suppliers -> ConcatBatchReader -> filter.""" files = self.split.files @@ -1065,6 +1136,9 @@ def _create_raw_reader(self) -> RecordReader: else: reader = merge_reader + if self._post_merge_filter is not None and not self._post_filter_after_inline: + reader = AuthFilterReader(reader, self._post_merge_filter) + if self.outer_extract_name_paths: if self.outer_flat_read_type is None: raise ValueError( @@ -1075,7 +1149,7 @@ def _create_raw_reader(self) -> RecordReader: reader = NestedLeafBatchReader( reader, self.outer_extract_name_paths, self.outer_flat_read_type) - if self.limit is not None: + if self.limit is not None and not self._post_filter_after_inline: reader = LimitedRecordBatchReader(reader, self.limit) return reader @@ -1138,17 +1212,16 @@ def _create_prescan_reader(self, field_names): if not prescan_fields: return EmptyRecordBatchReader() - # When there's a normal field predicate, don't push down limit to prescan reader - # because the outer reader will apply predicate+limit filtering, - # while prescan reader would only apply limit without normal field predicate - # TODO support limit+predicate push down + # Skip limit push-down when the outer reader also selects rows (predicate or auth + # filter): prescan's first-N rows would differ from the outer set. TODO: push down. + skip_limit = self.predicate is not None or self._post_merge_filter is not None prescan_read = DataEvolutionSplitRead( table=self.table, predicate=self.predicate, read_type=prescan_fields, split=self.split, row_tracking_enabled=False, - limit=None if self.predicate else self.limit, + limit=None if skip_limit else self.limit, ) prescan_read.row_ranges = self.row_ranges return prescan_read._create_raw_reader() @@ -1297,7 +1370,7 @@ def _create_union_reader(self, need_merge_files: List[DataFileMeta], deletion_ve [read_fields[0]] ).field(0).type, self.row_ranges, - CoreOptions.blob_as_descriptor(self.table.options), + self._read_blob_as_descriptor([read_fields[0].name]), deletion_vector=deletion_vector, batch_size=batch_size, blob_parallelism=self._blob_parallelism, @@ -1352,7 +1425,7 @@ def _create_raw_blob_file_reader( read_fields, self.read_fields, None, - CoreOptions.blob_as_descriptor(self.table.options), + self._read_blob_as_descriptor(read_fields), batch_size=self.table.options.read_batch_size(), row_indices=row_indices, blob_parallelism=blob_parallelism, diff --git a/paimon-python/pypaimon/read/table_read.py b/paimon-python/pypaimon/read/table_read.py index 6e641b50963a..c2ba44545a40 100644 --- a/paimon-python/pypaimon/read/table_read.py +++ b/paimon-python/pypaimon/read/table_read.py @@ -24,16 +24,18 @@ import pyarrow from pypaimon.common.predicate import Predicate +from pypaimon.common.predicate_json_parser import extract_referenced_fields from pypaimon.read.push_down_utils import predicate_field_names from pypaimon.read.query_auth_split import QueryAuthSplit from pypaimon.read.reader.auth_masking_reader import ( AuthFilterReader, AuthMaskingReader, ColumnProjectReader, RecordReaderToBatchAdapter, BatchToRecordReaderAdapter) from pypaimon.read.reader.iface.record_batch_reader import RecordBatchReader +from pypaimon.read.reader.limited_record_reader import LimitedRecordBatchReader from pypaimon.read.split import Split from pypaimon.read.split_read import (DataEvolutionSplitRead, MergeFileSplitRead, RawFileSplitRead, - SplitRead) + SplitRead, deferred_blob_field_names) from pypaimon.schema.data_types import DataField, PyarrowFieldParser from pypaimon.table.row.offset_row import OffsetRow @@ -111,6 +113,15 @@ def __init__( self._predicate_extra_fields = self._predicate_fields_outside_read_type() self._scan_read_type = self.read_type + self._predicate_extra_fields self._output_column_names = [f.name for f in self.read_type] + self._deferred_blob_fields = ( + deferred_blob_field_names( + self.table, + self._scan_read_type, + self.predicate, + limit, + ) + if self.table.options.data_evolution_enabled() else set() + ) self.include_row_kind = include_row_kind self.nested_name_paths = nested_name_paths self.limit = limit @@ -124,7 +135,9 @@ def _record_generator(): for split in splits: if limit is not None and count >= limit: return - reader = self.__create_reader_for_split(split) + remaining = None if limit is None else limit - count + reader = self.__create_reader_for_split( + split, limit=remaining) try: for batch in iter(reader.read_batch, None): for row in iter(batch.next, None): @@ -188,7 +201,9 @@ def to_arrow( order. Must be ``>= 1``. Note that with ``>= 2`` (or auto) and a ``limit`` set, the returned rows are an arbitrary subset of the requested size, since which splits fill the row - quota first is non-deterministic. + quota first is non-deterministic. Data-evolution reads with + deferred BLOB resolution run serially when a limit may discard + rows, so payloads are not materialized from discarded splits. blob_parallelism: number of threads for concurrent blob reads within each batch. ``None`` or ``1`` (default) reads blobs serially; ``>= 2`` uses a thread pool with ``pread`` for @@ -230,7 +245,8 @@ def _arrow_batch_generator(self, splits: List[Split], schema: pyarrow.Schema, for split in splits: if remaining is not None and remaining <= 0: break - reader = self.__create_reader_for_split(split, blob_parallelism) + reader = self.__create_reader_for_split( + split, blob_parallelism, limit=remaining) try: if isinstance(reader, RecordBatchReader): for batch in iter(reader.read_arrow_batch, None): @@ -331,7 +347,25 @@ def _should_run_parallel( overhead, no behavior change). A single split is never parallelized since there is nothing to fan out across. """ - return effective >= 2 and len(splits) >= 2 + deferred_limit_may_prune = ( + self.limit is not None + and self._deferred_blob_fields + and not self._limit_covers_all_splits(splits) + ) + return (effective >= 2 and len(splits) >= 2 + and not deferred_limit_may_prune) + + def _limit_covers_all_splits(self, splits: List[Split]) -> bool: + """Return whether split metadata proves that LIMIT cannot drop rows.""" + total_rows = 0 + for split in splits: + merged_row_count = split.merged_row_count() + if merged_row_count is None: + return False + total_rows += merged_row_count + if total_rows > self.limit: + return False + return True def _to_arrow_parallel( self, @@ -650,12 +684,33 @@ def to_torch( dataset = TorchDataset(self, splits) return dataset - def _create_split_read(self, split: Split, blob_parallelism: int = 1, read_type=None) -> SplitRead: - sr = self._build_split_read(split, read_type) + def _create_split_read(self, split: Split, blob_parallelism: int = 1, + read_type=None, limit: Optional[int] = None, + push_down_limit: bool = True, + post_merge_filter=None, + eager_blob_fields=None, + post_filter_after_inline: bool = False) -> SplitRead: + sr = self._build_split_read( + split, + read_type, + limit, + push_down_limit, + post_merge_filter, + eager_blob_fields, + post_filter_after_inline, + ) sr._blob_parallelism = blob_parallelism return sr - def _build_split_read(self, split: Split, read_type=None) -> SplitRead: + def _build_split_read(self, split: Split, read_type=None, + limit: Optional[int] = None, + push_down_limit: bool = True, + post_merge_filter=None, + eager_blob_fields=None, + post_filter_after_inline: bool = False) -> SplitRead: + effective_limit = ( + self.limit if limit is None else limit + ) if push_down_limit else None effective_read_type = read_type if read_type is not None else self.read_type scan_read_type = self._with_predicate_extra_fields(read_type) if read_type is not None else self._scan_read_type if self.table.is_primary_key_table and not split.raw_convertible: @@ -705,7 +760,7 @@ def _build_split_read(self, split: Split, read_type=None) -> SplitRead: outer_extract_name_paths=outer_extract_name_paths, outer_flat_read_type=( effective_read_type if outer_extract_name_paths else None), - limit=self.limit, + limit=effective_limit, ) elif self.table.options.data_evolution_enabled(): if self.nested_name_paths and any( @@ -726,7 +781,10 @@ def _build_split_read(self, split: Split, read_type=None) -> SplitRead: outer_extract_name_paths=outer_extract_name_paths, outer_flat_read_type=( self.read_type if outer_extract_name_paths else None), - limit=self.limit, + limit=effective_limit, + post_merge_filter=post_merge_filter, + eager_blob_fields=eager_blob_fields, + post_filter_after_inline=post_filter_after_inline, ) else: inner_read_type = scan_read_type @@ -752,7 +810,7 @@ def _build_split_read(self, split: Split, read_type=None) -> SplitRead: outer_extract_name_paths=outer_extract_name_paths, outer_flat_read_type=( effective_read_type if outer_extract_name_paths else None), - limit=self.limit, + limit=effective_limit, ) def _project_batch_to_output(self, batch: pyarrow.RecordBatch) -> pyarrow.RecordBatch: @@ -812,18 +870,27 @@ def _widen_to_top_level_for_merge(self) -> List[DataField]: widened.append(field) return widened - def __create_reader_for_split(self, split, blob_parallelism=1): + def __create_reader_for_split(self, split, blob_parallelism=1, + limit: Optional[int] = None): auth_result = None if isinstance(split, QueryAuthSplit): auth_result = split.auth_result split = split.split if auth_result is not None: - return self.__authed_reader(split, auth_result, blob_parallelism) - else: - return self._create_split_read(split, blob_parallelism=blob_parallelism).create_reader() - - def __authed_reader(self, split, auth_result, blob_parallelism=1): + return self.__authed_reader( + split, auth_result, blob_parallelism, limit) + if limit is None: + return self._create_split_read( + split, blob_parallelism=blob_parallelism).create_reader() + return self._create_split_read( + split, + blob_parallelism=blob_parallelism, + limit=limit, + ).create_reader() + + def __authed_reader(self, split, auth_result, blob_parallelism=1, + limit: Optional[int] = None): table_fields = self.table.fields read_fields = self.read_type @@ -832,9 +899,34 @@ def __authed_reader(self, split, auth_result, blob_parallelism=1): if extra_fields: effective_read_type = read_fields + extra_fields - reader = self._create_split_read( - split, blob_parallelism=blob_parallelism, - read_type=effective_read_type).create_reader() + filter_fn = auth_result.extract_row_filter() + effective_limit = self.limit if limit is None else limit + auth_fields = ( + self._auth_filter_field_names(auth_result, effective_read_type) + if filter_fn is not None else set() + ) + inline_blob_fields = ( + self.table.options.blob_descriptor_fields() + | self.table.options.blob_view_fields() + ) + embed_filter = ( + filter_fn is not None + and self.table.options.data_evolution_enabled() + ) + # If the auth filter references an inline BLOB, run it after inline resolution (in + # the split read) so it sees resolved payloads while scalar BLOBs still defer. + post_filter_after_inline = embed_filter and bool(auth_fields & inline_blob_fields) + split_read = self._create_split_read( + split, + blob_parallelism=blob_parallelism, + read_type=effective_read_type, + limit=limit, + push_down_limit=filter_fn is None or embed_filter, + post_merge_filter=filter_fn if embed_filter else None, + eager_blob_fields=auth_fields if embed_filter else None, + post_filter_after_inline=post_filter_after_inline, + ) + reader = split_read.create_reader() needs_convert_back = False if not isinstance(reader, RecordBatchReader): @@ -842,9 +934,10 @@ def __authed_reader(self, split, auth_result, blob_parallelism=1): reader = RecordReaderToBatchAdapter(reader, schema, include_row_kind=self.include_row_kind) needs_convert_back = True - filter_fn = auth_result.extract_row_filter() - if filter_fn: + if filter_fn and not embed_filter: reader = AuthFilterReader(reader, filter_fn) + if effective_limit is not None: + reader = LimitedRecordBatchReader(reader, effective_limit) if auth_result.column_masking: reader = AuthMaskingReader(reader, auth_result.column_masking, effective_read_type) @@ -858,6 +951,16 @@ def __authed_reader(self, split, auth_result, blob_parallelism=1): return reader + @staticmethod + def _auth_filter_field_names(auth_result, read_fields) -> set: + filters = getattr(auth_result, "filter", None) + if not filters: + return {field.name for field in read_fields} + names = set() + for filter_json in filters: + names.update(extract_referenced_fields(filter_json)) + return names + @staticmethod def convert_rows_to_arrow_batch(row_tuples: List[tuple], schema: pyarrow.Schema) -> pyarrow.RecordBatch: columns_data = zip(*row_tuples) diff --git a/paimon-python/pypaimon/tests/deferred_blob_resolve_test.py b/paimon-python/pypaimon/tests/deferred_blob_resolve_test.py new file mode 100644 index 000000000000..3dd60d6c4e54 --- /dev/null +++ b/paimon-python/pypaimon/tests/deferred_blob_resolve_test.py @@ -0,0 +1,414 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import json +import os +import shutil +import tempfile +import unittest +from unittest.mock import patch + +import pyarrow as pa +import pyarrow.compute as pc + +from pypaimon import CatalogFactory, Schema +from pypaimon.read.query_auth_split import QueryAuthSplit + + +_ROW_COUNT = 10 +_TABLE_OPTIONS = { + "row-tracking.enabled": "true", + "data-evolution.enabled": "true", +} + + +class _BlobCountingFileIO: + + def __init__(self, inner): + self._inner = inner + self.blobs_fetched = 0 + + def read_blobs_concurrent(self, blobs, parallelism): + self.blobs_fetched += sum(blob is not None for blob in blobs) + return self._inner.read_blobs_concurrent(blobs, parallelism) + + def __getattr__(self, name): + return getattr(self._inner, name) + + +class _RejectScoreOneAuthResult: + column_masking = None + filter = [json.dumps({ + "kind": "LEAF", + "transform": { + "name": "FIELD_REF", + "fieldRef": {"index": 2, "name": "score", "type": "INT"}, + }, + "function": "NOT_EQUAL", + "literals": [1], + })] + + @staticmethod + def get_extra_fields_for_filter(read_fields, table_fields): + return [] + + @staticmethod + def extract_row_filter(): + return lambda batch: pc.not_equal(batch.column("score"), 1) + + +class _PayloadAuthResult: + column_masking = None + + def __init__(self, expected_payload): + self._expected_payload = expected_payload + self.filter = [json.dumps({ + "kind": "LEAF", + "transform": { + "name": "FIELD_REF", + "fieldRef": { + "index": 1, + "name": "payload", + "type": "BYTES", + }, + }, + "function": "EQUAL", + "literals": [], + })] + + @staticmethod + def get_extra_fields_for_filter(read_fields, table_fields): + return [] + + def extract_row_filter(self): + return lambda batch: pc.equal( + batch.column("payload"), self._expected_payload) + + +class DeferredBlobResolveTest(unittest.TestCase): + + @classmethod + def setUpClass(cls): + cls.tempdir = tempfile.mkdtemp() + cls.catalog = CatalogFactory.create({ + "warehouse": os.path.join(cls.tempdir, "warehouse") + }) + cls.catalog.create_database("default", False) + cls.schema = pa.schema([ + ("sample_id", pa.string()), + ("payload", pa.large_binary()), + ("score", pa.int32()), + ]) + + @classmethod + def tearDownClass(cls): + shutil.rmtree(cls.tempdir, ignore_errors=True) + + def _create_table(self, name, extra_options=None, payloads=None, + partition_keys=None, sample_ids=None): + options = dict(_TABLE_OPTIONS) + options.update(extra_options or {}) + identifier = "default.%s" % name + self.catalog.create_table( + identifier, + Schema.from_pyarrow_schema( + self.schema, + partition_keys=partition_keys, + options=options, + ), + False, + ) + table = self.catalog.get_table(identifier) + write_builder = table.new_batch_write_builder() + writer = write_builder.new_write() + commit = write_builder.new_commit() + if sample_ids is None: + sample_ids = [ + "sample_%d" % index for index in range(_ROW_COUNT) + ] + writer.write_arrow(pa.table({ + "sample_id": sample_ids, + "payload": ( + payloads if payloads is not None else + [bytes([index]) * 1024 for index in range(_ROW_COUNT)] + ), + "score": list(range(_ROW_COUNT)), + }, schema=self.schema)) + commit.commit(writer.prepare_commit()) + writer.close() + commit.close() + return self.catalog.get_table(identifier) + + def _read(self, table, predicate, limit=None, blob_parallelism=None): + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + read_builder = table.new_read_builder() + if predicate is not None: + read_builder = read_builder.with_filter(predicate) + read_builder = read_builder.with_projection( + ["sample_id", "payload", "score"]) + if limit is not None: + read_builder = read_builder.with_limit(limit) + splits = read_builder.new_scan().plan().splits() + table_read = read_builder.new_read() + if blob_parallelism is None: + batch_reader = table_read.to_arrow_batch_reader(splits) + else: + batch_reader = table_read.to_arrow_batch_reader( + splits, blob_parallelism=blob_parallelism) + result = pa.Table.from_batches(batch_reader) + return result, counting_file_io + + def test_fetches_payloads_only_for_filtered_rows(self): + table = self._create_table("defer_filtered") + predicate = table.new_read_builder().new_predicate_builder().less_than( + "score", 5) + + result, counting_file_io = self._read(table, predicate) + + self.assertEqual(5, result.num_rows) + self.assertEqual(5, counting_file_io.blobs_fetched) + self.assertEqual( + [bytes([index]) * 1024 for index in range(5)], + result.column("payload").to_pylist(), + ) + + def test_applies_limit_before_fetching_payloads(self): + table = self._create_table("defer_limit") + predicate = table.new_read_builder().new_predicate_builder().less_than( + "score", 8) + + result, counting_file_io = self._read(table, predicate, limit=2) + + self.assertEqual(2, result.num_rows) + self.assertEqual(2, counting_file_io.blobs_fetched) + + def test_limit_without_predicate_defers_payloads(self): + table = self._create_table("defer_limit_only") + + result, counting_file_io = self._read(table, None, limit=2) + + self.assertEqual(2, result.num_rows) + self.assertEqual(2, counting_file_io.blobs_fetched) + + def test_limit_does_not_prefetch_payloads_across_splits(self): + table = self._create_table( + "defer_limit_splits", + extra_options={"source.split.target-size": "1b"}, + partition_keys=["sample_id"], + ) + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + splits = table.new_read_builder().new_scan().plan().splits() + read_builder = table.new_read_builder().with_limit(1) + + table_read = read_builder.new_read() + with patch.object( + table_read, + "_to_arrow_parallel", + side_effect=AssertionError("deferred LIMIT must run serially"), + ) as parallel_read: + result = table_read.to_arrow(splits, parallelism=4) + + self.assertGreater(len(splits), 1) + parallel_read.assert_not_called() + self.assertEqual(1, result.num_rows) + self.assertEqual(1, counting_file_io.blobs_fetched) + + def test_limit_covering_all_rows_preserves_parallelism(self): + table = self._create_table( + "defer_limit_all_rows", + extra_options={"source.split.target-size": "1b"}, + partition_keys=["sample_id"], + ) + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + read_builder = table.new_read_builder().with_limit(_ROW_COUNT) + splits = read_builder.new_scan().plan().splits() + table_read = read_builder.new_read() + + with patch.object( + table_read, + "_to_arrow_parallel", + wraps=table_read._to_arrow_parallel, + ) as parallel_read: + result = table_read.to_arrow(splits, parallelism=4) + + self.assertGreater(len(splits), 1) + parallel_read.assert_called_once() + self.assertEqual(_ROW_COUNT, result.num_rows) + self.assertEqual(_ROW_COUNT, counting_file_io.blobs_fetched) + + def test_iterator_passes_remaining_limit_across_splits(self): + table = self._create_table( + "defer_iterator_limit_splits", + extra_options={"source.split.target-size": "1b"}, + partition_keys=["sample_id"], + sample_ids=["a"] + ["b"] * (_ROW_COUNT - 1), + ) + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + read_builder = table.new_read_builder().with_projection( + ["sample_id", "payload", "score"] + ).with_limit(2) + splits = read_builder.new_scan().plan().splits() + + rows = list(read_builder.new_read().to_iterator(splits)) + + self.assertEqual(2, len(splits)) + self.assertEqual(2, len(rows)) + self.assertEqual(2, counting_file_io.blobs_fetched) + + def test_iterator_applies_limit_after_auth_filter(self): + table = self._create_table( + "defer_iterator_auth_limit_splits", + extra_options={"source.split.target-size": "1b"}, + partition_keys=["sample_id"], + sample_ids=["a"] + ["b"] * (_ROW_COUNT - 1), + ) + read_builder = table.new_read_builder().with_projection( + ["sample_id", "payload", "score"] + ).with_limit(2) + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + auth_result = _RejectScoreOneAuthResult() + splits = [ + QueryAuthSplit(split, auth_result) + for split in read_builder.new_scan().plan().splits() + ] + + scores = [ + row.get_field(2) + for row in read_builder.new_read().to_iterator(splits) + ] + + self.assertEqual([0, 2], scores) + self.assertEqual(2, counting_file_io.blobs_fetched) + + def test_auth_blob_filter_keeps_eager_resolution(self): + table = self._create_table("defer_auth_blob_filter") + expected_payload = bytes([3]) * 1024 + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + read_builder = table.new_read_builder().with_projection( + ["sample_id", "payload", "score"] + ).with_limit(1) + auth_result = _PayloadAuthResult(expected_payload) + splits = [ + QueryAuthSplit(split, auth_result) + for split in read_builder.new_scan().plan().splits() + ] + + result = pa.Table.from_batches( + read_builder.new_read().to_arrow_batch_reader( + splits, blob_parallelism=4) + ) + + self.assertEqual([expected_payload], result.column("payload").to_pylist()) + self.assertEqual(_ROW_COUNT, counting_file_io.blobs_fetched) + + def test_auth_only_defers_non_auth_payloads(self): + # An auth filter with no predicate/limit still defers scalar BLOBs, so payloads of + # the rows the auth filter drops are not read. + table = self._create_table("defer_auth_only") + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + read_builder = table.new_read_builder().with_projection( + ["sample_id", "payload", "score"]) + splits = [ + QueryAuthSplit(split, _RejectScoreOneAuthResult()) + for split in read_builder.new_scan().plan().splits() + ] + + result = pa.Table.from_batches( + read_builder.new_read().to_arrow_batch_reader( + splits, blob_parallelism=4)) + + self.assertEqual(_ROW_COUNT - 1, result.num_rows) + self.assertEqual(_ROW_COUNT - 1, counting_file_io.blobs_fetched) + + def test_preserves_null_payloads_after_filtering(self): + payloads = [ + None if index == 1 else bytes([index]) * 1024 + for index in range(_ROW_COUNT) + ] + table = self._create_table("defer_null", payloads=payloads) + predicate = table.new_read_builder().new_predicate_builder().less_than( + "score", 4) + + result, counting_file_io = self._read(table, predicate) + + self.assertEqual(4, result.num_rows) + self.assertEqual(3, counting_file_io.blobs_fetched) + self.assertEqual(payloads[:4], result.column("payload").to_pylist()) + + def test_can_disable_deferred_resolution(self): + table = self._create_table( + "defer_disabled", + {"read.defer-blob-resolve": "false"}, + ) + predicate = table.new_read_builder().new_predicate_builder().less_than( + "score", 5) + + result, counting_file_io = self._read( + table, predicate, blob_parallelism=4) + + self.assertEqual(5, result.num_rows) + self.assertEqual(_ROW_COUNT, counting_file_io.blobs_fetched) + + def test_blob_predicate_keeps_eager_resolution(self): + table = self._create_table("defer_blob_predicate") + expected_payload = bytes([3]) * 1024 + predicate = table.new_read_builder().new_predicate_builder().equal( + "payload", expected_payload) + + result, counting_file_io = self._read( + table, predicate, blob_parallelism=4) + + self.assertEqual(1, result.num_rows) + self.assertEqual([expected_payload], result.column("payload").to_pylist()) + self.assertEqual(_ROW_COUNT, counting_file_io.blobs_fetched) + + def test_defers_payloads_for_blob_fallback_reader(self): + table = self._create_table("defer_fallback") + update_builder = table.new_batch_write_builder() + table_update = update_builder.new_update().with_update_type(["payload"]) + updated_payload = b"updated-payload" + update_messages = table_update.update_by_arrow_with_row_id(pa.table({ + "_ROW_ID": pa.array([3], type=pa.int64()), + "payload": pa.array([updated_payload], type=pa.large_binary()), + })) + update_builder.new_commit().commit(update_messages) + + predicate = table.new_read_builder().new_predicate_builder().less_than( + "score", 5) + result, counting_file_io = self._read(table, predicate) + + self.assertEqual(5, result.num_rows) + self.assertEqual(5, counting_file_io.blobs_fetched) + payload_by_score = dict(zip( + result.column("score").to_pylist(), + result.column("payload").to_pylist(), + )) + self.assertEqual( + updated_payload, + payload_by_score[3], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/paimon-python/pypaimon/tests/torch_read_test.py b/paimon-python/pypaimon/tests/torch_read_test.py index 5f55cb2bc892..6a1d559b451f 100644 --- a/paimon-python/pypaimon/tests/torch_read_test.py +++ b/paimon-python/pypaimon/tests/torch_read_test.py @@ -29,6 +29,20 @@ from pypaimon.table.file_store_table import FileStoreTable +class _BlobCountingFileIO: + + def __init__(self, inner): + self._inner = inner + self.blobs_fetched = 0 + + def read_blobs_concurrent(self, blobs, parallelism): + self.blobs_fetched += sum(blob is not None for blob in blobs) + return self._inner.read_blobs_concurrent(blobs, parallelism) + + def __getattr__(self, name): + return getattr(self._inner, name) + + class TorchReadTest(unittest.TestCase): @classmethod def setUpClass(cls): @@ -143,6 +157,77 @@ def test_torch_streaming_prefetch_concurrency(self): self.assertEqual(sorted_user_ids, expected_user_ids) self.assertEqual(sorted_behaviors, expected_behaviors) + def test_torch_streaming_limit_is_global_across_workers(self): + schema = Schema.from_pyarrow_schema( + self.pa_schema, partition_keys=['user_id']) + self.catalog.create_table( + 'default.test_torch_streaming_global_limit', schema, False) + table = self.catalog.get_table( + 'default.test_torch_streaming_global_limit') + self._write_test_table(table) + + read_builder = table.new_read_builder() \ + .with_projection(['user_id']) \ + .with_limit(3) + splits = read_builder.new_scan().plan().splits() + dataset = read_builder.new_read().to_torch( + splits, + streaming=True, + prefetch_concurrency=4, + ) + + user_ids = self._collect_torch_user_ids(dataset, num_workers=2) + + self.assertEqual(3, len(user_ids)) + self.assertEqual(3, len(set(user_ids))) + + def test_torch_streaming_limit_does_not_prefetch_blob_payloads(self): + pa_schema = pa.schema([ + ('sample_id', pa.string()), + ('payload', pa.large_binary()), + ('score', pa.int32()), + ]) + schema = Schema.from_pyarrow_schema( + pa_schema, + partition_keys=['sample_id'], + options={ + 'row-tracking.enabled': 'true', + 'data-evolution.enabled': 'true', + 'source.split.target-size': '1b', + }, + ) + self.catalog.create_table( + 'default.test_torch_streaming_blob_limit', schema, False) + table = self.catalog.get_table( + 'default.test_torch_streaming_blob_limit') + write_builder = table.new_batch_write_builder() + writer = write_builder.new_write() + commit = write_builder.new_commit() + writer.write_arrow(pa.table({ + 'sample_id': ['sample_%d' % index for index in range(10)], + 'payload': [bytes([index]) * 1024 for index in range(10)], + 'score': list(range(10)), + }, schema=pa_schema)) + commit.commit(writer.prepare_commit()) + writer.close() + commit.close() + + counting_file_io = _BlobCountingFileIO(table.file_io) + table.file_io = counting_file_io + read_builder = table.new_read_builder() \ + .with_projection(['sample_id', 'payload', 'score']) \ + .with_limit(2) + splits = read_builder.new_scan().plan().splits() + + rows = list(read_builder.new_read().to_torch( + splits, + streaming=True, + prefetch_concurrency=4, + )) + + self.assertEqual(2, len(rows)) + self.assertEqual(2, counting_file_io.blobs_fetched) + def test_blob_torch_read(self): """Test end-to-end blob functionality using blob descriptors.""" import random