This is an automated email from the ASF dual-hosted git repository.

JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git


The following commit(s) were added to refs/heads/master by this push:
     new 20377469cb [python] Fix Parquet VARIANT reads with multi-chunk nested 
data (#8848)
20377469cb is described below

commit 20377469cbf7deaf7b7bc1be3643a0e7d64181fe
Author: XiaoHongbo <[email protected]>
AuthorDate: Sun Jul 26 15:11:42 2026 +0800

    [python] Fix Parquet VARIANT reads with multi-chunk nested data (#8848)
---
 paimon-python/pypaimon/data/variant_shredding.py   |   2 +-
 .../pypaimon/read/reader/format_pyarrow_reader.py  | 419 ++++++++++++++++----
 paimon-python/pypaimon/read/split_read.py          |   6 +-
 .../tests/format_pyarrow_variant_row_group_test.py | 434 +++++++++++++++++++++
 4 files changed, 785 insertions(+), 76 deletions(-)

diff --git a/paimon-python/pypaimon/data/variant_shredding.py 
b/paimon-python/pypaimon/data/variant_shredding.py
index f14946ba33..72a4508dd3 100644
--- a/paimon-python/pypaimon/data/variant_shredding.py
+++ b/paimon-python/pypaimon/data/variant_shredding.py
@@ -182,7 +182,7 @@ def _parse_typed_value_field(schema: VariantSchema, 
tv_type: pa.DataType) -> Var
     if pa.types.is_struct(tv_type):
         object_fields: List[ObjectField] = []
         for j in range(tv_type.num_fields):
-            sub_f = tv_type.field(j)
+            sub_f = tv_type[j]
             if pa.types.is_struct(sub_f.type):
                 sub_schema = build_variant_schema(sub_f.type)
             else:
diff --git a/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py 
b/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
index efaf8e0546..956ca39f46 100644
--- a/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
+++ b/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
@@ -15,7 +15,7 @@
 # specific language governing permissions and limitations
 # under the License.
 
-from typing import Any, Dict, List, Optional
+from typing import Any, Dict, List, Optional, Set
 
 import pyarrow as pa
 import pyarrow.dataset as ds
@@ -30,7 +30,15 @@ from pypaimon.data.variant_shredding import (
     is_shredded_variant,
 )
 from pypaimon.read.reader.iface.record_batch_reader import RecordBatchReader
-from pypaimon.schema.data_types import DataField, PyarrowFieldParser
+from pypaimon.schema.data_types import (
+    ArrayType,
+    AtomicType,
+    DataField,
+    MapType,
+    MultisetType,
+    PyarrowFieldParser,
+    RowType,
+)
 from pypaimon.table.special_fields import SpecialFields
 
 
@@ -48,7 +56,9 @@ class FormatPyArrowReader(RecordBatchReader):
                  read_fields: List[DataField],
                  push_down_predicate: Any, batch_size: int = 1024,
                  options: CoreOptions = None,
-                 nested_name_paths: Optional[List[List[str]]] = None):
+                 nested_name_paths: Optional[List[List[str]]] = None,
+                 predicate_field_names: Optional[Set[str]] = None):
+        self._predicate_field_names = predicate_field_names or set()
         file_path_for_pyarrow = file_io.to_filesystem_path(file_path)
         self.dataset = ds.dataset(file_path_for_pyarrow, format=file_format, 
filesystem=file_io.filesystem)
         self._file_format = file_format
@@ -62,6 +72,7 @@ class FormatPyArrowReader(RecordBatchReader):
         self._nested_name_paths = nested_name_paths
         has_nested_path = bool(
             nested_name_paths and any(len(p) > 1 for p in nested_name_paths))
+        self._has_nested_path = has_nested_path
 
         file_schema = self.dataset.schema
         if has_nested_path:
@@ -77,96 +88,212 @@ class FormatPyArrowReader(RecordBatchReader):
             self.existing_fields = [f.name for f in read_fields if f.name in 
file_schema_names]
             self.missing_fields = [f.name for f in read_fields if f.name not 
in file_schema_names]
 
-        self._shredded_schemas: Dict[str, VariantSchema] = {}
-        if options is None or options.variant_shredding_enabled():
-            top_level_names = set(file_schema.names)
-            for name in self.existing_fields:
-                if name not in top_level_names:
-                    continue
-                field_type = file_schema.field(name).type
-                if is_shredded_variant(field_type):
-                    self._shredded_schemas[name] = 
build_variant_schema(field_type)
+        self._variant_shredding_enabled = (
+            options is None or options.variant_shredding_enabled())
+        self._variant_schema_cache: Dict[pa.DataType, VariantSchema] = {}
 
-        if has_nested_path:
+        self._bounded_variant_read = (
+            self._file_format == 'parquet' and self._has_projected_variant())
+        if has_nested_path and not self._bounded_variant_read:
             existing_set = set(self.existing_fields)
             columns_dict = {}
             for f, path in zip(read_fields, nested_name_paths):
                 if f.name in existing_set:
                     columns_dict[f.name] = ds.field(*path)
-            self.reader = self.dataset.scanner(
-                columns=columns_dict,
-                filter=push_down_predicate,
-                batch_size=batch_size
-            ).to_reader()
+            self._scan_columns = columns_dict
+        elif has_nested_path:
+            self._scan_columns = None
         else:
             # Only pass existing fields to PyArrow scanner to avoid errors
-            self.reader = self.dataset.scanner(
-                columns=self.existing_fields,
-                filter=push_down_predicate,
-                batch_size=batch_size
-            ).to_reader()
+            self._scan_columns = self.existing_fields
+        self._scan_filter = push_down_predicate
+        self._scan_batch_size = batch_size
 
         self._output_schema = (
             PyarrowFieldParser.from_paimon_schema(read_fields) if read_fields 
else None
         )
 
-    def read_arrow_batch(self) -> Optional[RecordBatch]:
+        # Read projected VARIANT columns in bounded batches.
+        self._parquet_file = None
+        if self._bounded_variant_read:
+            import pyarrow.parquet as pq
+            # ParquetFile(filesystem=...) is unavailable in PyArrow 6.
+            self._parquet_file = pq.ParquetFile(
+                file_io.filesystem.open_input_file(file_path_for_pyarrow))
+        if self._parquet_file is not None:
+            self._raw_batches = self._iter_row_group_batches()
+        else:
+            reader = self.dataset.scanner(
+                columns=self._scan_columns,
+                filter=self._scan_filter,
+                batch_size=self._scan_batch_size,
+            ).to_reader()
+            self._raw_batches = self._iter_reader_batches(reader)
+
+    def _has_projected_variant(self) -> bool:
+        return any(
+            f.name in self.existing_fields
+            and _contains_variant(f.type)
+            for f in self.read_fields)
+
+    @staticmethod
+    def _iter_reader_batches(reader):
+        while True:
+            try:
+                yield reader.read_next_batch()
+            except StopIteration:
+                return
+
+    def _iter_row_group_batches(self):
+        columns = self._row_group_read_columns()
+        for row_group in self._surviving_row_group_ids():
+            for batch in self._parquet_file.iter_batches(
+                    row_groups=[row_group],
+                    columns=columns,
+                    batch_size=self._scan_batch_size):
+                if self._has_nested_path:
+                    batches = [batch]
+                    if self._scan_filter is not None:
+                        table = ds.dataset(
+                            pa.Table.from_batches([batch])
+                        ).scanner(filter=self._scan_filter).to_table()
+                        batches = table.to_batches()
+                    for filtered in batches:
+                        out = self._select_nested_fields(filtered)
+                        if out.num_rows:
+                            yield out
+                    continue
+                if self._scan_filter is None:
+                    yield self._select_existing_fields(batch)
+                    continue
+                table = ds.dataset(
+                    pa.Table.from_batches([batch])
+                ).scanner(filter=self._scan_filter).to_table()
+                if self.existing_fields:
+                    table = table.select(self.existing_fields)
+                for out in table.to_batches():
+                    if out.num_rows:
+                        yield out
+
+    def _row_group_read_columns(self):
+        if self._has_nested_path:
+            existing = set(self.existing_fields)
+            columns = []
+            for field, path in zip(self.read_fields, self._nested_name_paths):
+                if field.name in existing and path[0] not in columns:
+                    columns.append(path[0])
+        else:
+            columns = list(self.existing_fields)
+        if self._scan_filter is not None:
+            file_names = set(self.dataset.schema.names)
+            for name in self._predicate_field_names:
+                if name in file_names and name not in columns:
+                    columns.append(name)
+        return columns
+
+    def _select_existing_fields(self, batch):
+        columns = []
+        fields = []
+        for name in self.existing_fields:
+            index = batch.schema.get_field_index(name)
+            if index < 0:
+                raise KeyError("Field not found in batch: {}".format(name))
+            columns.append(batch.column(index))
+            fields.append(batch.schema.field(index))
+        return pa.RecordBatch.from_arrays(columns, schema=pa.schema(fields))
+
+    def _select_nested_fields(self, batch):
+        columns = []
+        names = []
+        existing = set(self.existing_fields)
+        for field, path in zip(self.read_fields, self._nested_name_paths):
+            if field.name not in existing:
+                continue
+            index = batch.schema.get_field_index(path[0])
+            if index < 0:
+                raise KeyError("Field not found in batch: {}".format(path[0]))
+            column = batch.column(index)
+            for name in path[1:]:
+                index = column.type.get_field_index(name)
+                column = column.flatten()[index]
+            columns.append(column)
+            names.append(field.name)
+        return pa.RecordBatch.from_arrays(columns, names=names)
+
+    def _surviving_row_group_ids(self):
+        total = self._parquet_file.num_row_groups
+        if self._scan_filter is None:
+            return range(total)
         try:
-            batch = self.reader.read_next_batch()
-
-            if self._file_format == 'orc' and self._output_schema is not None:
-                batch = self._cast_orc_time_columns(batch)
-
-            if self._shredded_schemas:
-                batch = self._assemble_shredded_variants(batch)
-
-            if not self.missing_fields:
-                return batch
-
-            def _type_for_missing(name: str) -> pa.DataType:
-                if self._output_schema is not None:
-                    idx = self._output_schema.get_field_index(name)
-                    if idx >= 0:
-                        return self._output_schema.field(idx).type
-                return pa.null()
-
-            missing_columns = [
-                pa.nulls(batch.num_rows, type=_type_for_missing(name))
-                for name in self.missing_fields
-            ]
-
-            # Reconstruct the batch with all fields in the correct order
-            all_columns = []
-            out_fields = []
-            for field_name in self._read_field_names:
-                if field_name in self.existing_fields:
-                    # Get the column from the existing batch
-                    column_idx = self.existing_fields.index(field_name)
-                    all_columns.append(batch.column(column_idx))
-                    out_fields.append(batch.schema.field(column_idx))
-                else:
-                    # Get the column from missing fields
-                    column_idx = self.missing_fields.index(field_name)
-                    col_type = _type_for_missing(field_name)
-                    all_columns.append(missing_columns[column_idx])
-                    nullable = not SpecialFields.is_system_field(field_name)
-                    out_fields.append(pa.field(field_name, col_type, 
nullable=nullable))
-            # Create a new RecordBatch with all columns
-            return pa.RecordBatch.from_arrays(all_columns, 
schema=pa.schema(out_fields))
-
-        except StopIteration:
+            ids = set()
+            for fragment in self.dataset.get_fragments(
+                    filter=self._scan_filter):
+                for row_group in fragment.split_by_row_group(
+                        self._scan_filter):
+                    ids.update(info.id for info in row_group.row_groups)
+            return sorted(ids)
+        except Exception:
+            return range(total)
+
+    def read_arrow_batch(self) -> Optional[RecordBatch]:
+        batch = next(self._raw_batches, None)
+        if batch is None:
             return None
+        return self._post_process_batch(batch)
+
+    def _post_process_batch(self, batch: RecordBatch) -> RecordBatch:
+        if self._file_format == 'orc' and self._output_schema is not None:
+            batch = self._cast_orc_time_columns(batch)
+
+        if self._variant_shredding_enabled:
+            batch = self._assemble_shredded_variants(batch)
+
+        if not self.missing_fields:
+            return batch
+
+        def _type_for_missing(name: str) -> pa.DataType:
+            if self._output_schema is not None:
+                idx = self._output_schema.get_field_index(name)
+                if idx >= 0:
+                    return self._output_schema.field(idx).type
+            return pa.null()
+
+        missing_columns = [
+            pa.nulls(batch.num_rows, type=_type_for_missing(name))
+            for name in self.missing_fields
+        ]
+
+        all_columns = []
+        out_fields = []
+        for field_name in self._read_field_names:
+            if field_name in self.existing_fields:
+                column_idx = self.existing_fields.index(field_name)
+                all_columns.append(batch.column(column_idx))
+                out_fields.append(batch.schema.field(column_idx))
+            else:
+                column_idx = self.missing_fields.index(field_name)
+                col_type = _type_for_missing(field_name)
+                all_columns.append(missing_columns[column_idx])
+                nullable = not SpecialFields.is_system_field(field_name)
+                out_fields.append(
+                    pa.field(field_name, col_type, nullable=nullable))
+        return pa.RecordBatch.from_arrays(
+            all_columns, schema=pa.schema(out_fields))
 
     def _assemble_shredded_variants(self, batch: pa.RecordBatch) -> 
pa.RecordBatch:
-        """Replace shredded VARIANT columns with standard struct<value, 
metadata>."""
         changed = False
         columns = list(batch.columns)
         fields = list(batch.schema)
+        logical_types = {field.name: field.type for field in self.read_fields}
 
         for i, f in enumerate(fields):
-            if f.name in self._shredded_schemas:
-                schema = self._shredded_schemas[f.name]
-                new_col = assemble_shredded_column(columns[i], schema)
+            logical_type = logical_types.get(f.name)
+            if logical_type is not None:
+                new_col, column_changed = _assemble_variant_column(
+                    columns[i], logical_type, self._variant_schema_cache)
+            else:
+                new_col, column_changed = columns[i], False
+            if column_changed:
                 columns[i] = new_col
                 fields[i] = pa.field(f.name, new_col.type, nullable=f.nullable)
                 changed = True
@@ -197,8 +324,12 @@ class FormatPyArrowReader(RecordBatchReader):
         return batch
 
     def close(self):
-        if self.reader is not None:
-            self.reader = None
+        self._raw_batches = None
+        if self._parquet_file is not None:
+            close = getattr(self._parquet_file, 'close', None)
+            if close is not None:
+                close()
+            self._parquet_file = None
 
 
 def _path_exists_in_arrow_schema(schema: pa.Schema, path: List[str]) -> bool:
@@ -214,5 +345,145 @@ def _path_exists_in_arrow_schema(schema: pa.Schema, path: 
List[str]) -> bool:
         idx = current_type.get_field_index(name)
         if idx < 0:
             return False
-        current_type = current_type.field(idx).type
+        current_type = current_type[idx].type
     return True
+
+
+def _contains_variant(data_type) -> bool:
+    if isinstance(data_type, AtomicType):
+        return data_type.type.upper() == 'VARIANT'
+    if isinstance(data_type, (ArrayType, MultisetType)):
+        return _contains_variant(data_type.element)
+    if isinstance(data_type, MapType):
+        return (_contains_variant(data_type.key)
+                or _contains_variant(data_type.value))
+    if isinstance(data_type, RowType):
+        return any(_contains_variant(field.type) for field in data_type.fields)
+    return False
+
+
+def _assemble_variant_column(column, data_type, schema_cache):
+    if isinstance(data_type, AtomicType):
+        if (data_type.type.upper() != 'VARIANT'
+                or not is_shredded_variant(column.type)):
+            return column, False
+        schema = schema_cache.get(column.type)
+        if schema is None:
+            schema = build_variant_schema(column.type)
+            schema_cache[column.type] = schema
+        return assemble_shredded_column(column, schema), True
+
+    if isinstance(data_type, RowType) and pa.types.is_struct(column.type):
+        logical_fields = {field.name: field.type for field in data_type.fields}
+        columns = []
+        fields = []
+        changed = False
+        for index, arrow_field in enumerate(column.type):
+            child = column.field(index)
+            logical_type = logical_fields.get(arrow_field.name)
+            if logical_type is not None:
+                child, child_changed = _assemble_variant_column(
+                    child, logical_type, schema_cache)
+                changed = changed or child_changed
+            columns.append(child)
+            fields.append(pa.field(
+                arrow_field.name,
+                child.type,
+                nullable=arrow_field.nullable,
+                metadata=arrow_field.metadata,
+            ))
+        if changed:
+            mask = column.is_null() if column.null_count else None
+            return pa.StructArray.from_arrays(
+                columns, fields=fields, mask=mask), True
+        return column, False
+
+    if (isinstance(data_type, (ArrayType, MultisetType))
+            and (pa.types.is_list(column.type)
+                 or pa.types.is_large_list(column.type))):
+        offsets, start, end = _normalized_offsets(column)
+        values = column.values.slice(start, end - start)
+        values, changed = _assemble_variant_column(
+            values, data_type.element, schema_cache)
+        if not changed:
+            return column, False
+        if pa.types.is_large_list(column.type):
+            result = pa.LargeListArray.from_arrays(offsets, values)
+            list_type = pa.large_list(pa.field(
+                column.type.value_field.name,
+                values.type,
+                nullable=column.type.value_field.nullable,
+                metadata=column.type.value_field.metadata,
+            ))
+        else:
+            result = pa.ListArray.from_arrays(offsets, values)
+            list_type = pa.list_(pa.field(
+                column.type.value_field.name,
+                values.type,
+                nullable=column.type.value_field.nullable,
+                metadata=column.type.value_field.metadata,
+            ))
+        return pa.Array.from_buffers(
+            list_type,
+            len(result),
+            result.buffers()[:2],
+            null_count=result.null_count,
+            children=[values],
+        ), True
+
+    if isinstance(data_type, MapType) and pa.types.is_map(column.type):
+        offsets, start, end = _normalized_offsets(column)
+        keys = column.keys.slice(start, end - start)
+        items = column.items.slice(start, end - start)
+        keys, key_changed = _assemble_variant_column(
+            keys, data_type.key, schema_cache)
+        items, item_changed = _assemble_variant_column(
+            items, data_type.value, schema_cache)
+        if not key_changed and not item_changed:
+            return column, False
+        result = pa.MapArray.from_arrays(offsets, keys, items)
+        map_type = pa.map_(
+            pa.field(
+                column.type.key_field.name,
+                keys.type,
+                nullable=False,
+                metadata=column.type.key_field.metadata,
+            ),
+            pa.field(
+                column.type.item_field.name,
+                items.type,
+                nullable=column.type.item_field.nullable,
+                metadata=column.type.item_field.metadata,
+            ),
+            keys_sorted=getattr(column.type, 'keys_sorted', False),
+        )
+        entries = pa.StructArray.from_arrays(
+            [keys, items], fields=[map_type.key_field, map_type.item_field])
+        return pa.Array.from_buffers(
+            map_type,
+            len(result),
+            result.buffers()[:2],
+            null_count=result.null_count,
+            children=[entries],
+        ), True
+
+    return column, False
+
+
+def _normalized_offsets(column):
+    offsets_array = getattr(column, 'offsets', None)
+    if offsets_array is None:
+        offsets_array = pa.Array.from_buffers(
+            pa.int32(),
+            len(column) + 1,
+            [None, column.buffers()[1]],
+            offset=column.offset,
+        )
+    raw_offsets = offsets_array.to_pylist()
+    start = raw_offsets[0]
+    end = raw_offsets[-1]
+    offsets = [value - start for value in raw_offsets]
+    for index, is_null in enumerate(column.is_null().to_pylist()):
+        if is_null:
+            offsets[index] = None
+    return pa.array(offsets, type=offsets_array.type), start, end
diff --git a/paimon-python/pypaimon/read/split_read.py 
b/paimon-python/pypaimon/read/split_read.py
index 6ced11f0c1..71e6b623f4 100644
--- a/paimon-python/pypaimon/read/split_read.py
+++ b/paimon-python/pypaimon/read/split_read.py
@@ -330,11 +330,15 @@ class SplitRead(ABC):
                 [nested_path_by_name[f.name] for f in ordered_read_fields]
                 if has_nested else None
             )
+            predicate_fields = (
+                predicate_field_names(self.push_down_predicate)
+                if self.push_down_predicate else set())
             format_reader = FormatPyArrowReader(
                 self.table.file_io, file_format, file_path,
                 ordered_read_fields, read_arrow_predicate, 
batch_size=batch_size,
                 options=self.table.options,
-                nested_name_paths=ordered_nested_paths)
+                nested_name_paths=ordered_nested_paths,
+                predicate_field_names=predicate_fields)
         elif file_format == CoreOptions.FILE_FORMAT_ROW:
             if has_nested:
                 raise NotImplementedError(
diff --git 
a/paimon-python/pypaimon/tests/format_pyarrow_variant_row_group_test.py 
b/paimon-python/pypaimon/tests/format_pyarrow_variant_row_group_test.py
new file mode 100644
index 0000000000..bb0f06bc6e
--- /dev/null
+++ b/paimon-python/pypaimon/tests/format_pyarrow_variant_row_group_test.py
@@ -0,0 +1,434 @@
+# 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 inspect
+import json
+import os
+import shutil
+import tempfile
+import unittest
+
+import pyarrow as pa
+import pyarrow.dataset as ds
+import pyarrow.fs as pafs
+import pyarrow.parquet as pq
+
+from pypaimon.data.generic_variant import GenericVariant
+from pypaimon.data.variant_shredding import (
+    parse_shredding_schema_option,
+    shredding_schema_to_arrow_type,
+    shred_variant_column,
+)
+from pypaimon.read.reader.format_pyarrow_reader import FormatPyArrowReader
+from pypaimon.schema.data_types import (
+    ArrayType,
+    AtomicType,
+    DataField,
+    MapType,
+    RowType,
+)
+
+_VARIANT_TYPE = pa.struct([
+    pa.field("value", pa.binary(), nullable=False),
+    pa.field("metadata", pa.binary(), nullable=False),
+])
+
+
+class _LocalFileIO:
+    filesystem = pafs.LocalFileSystem()
+
+    def to_filesystem_path(self, path):
+        return path
+
+
+def _drain(reader):
+    rows = 0
+    columns = None
+    content_keys = set()
+    while True:
+        batch = reader.read_arrow_batch()
+        if batch is None:
+            break
+        rows += batch.num_rows
+        columns = batch.schema.names
+        if "content_key" in columns:
+            content_keys |= set(
+                batch.column(columns.index("content_key")).to_pylist())
+    return rows, columns, content_keys
+
+
+class VariantRowGroupReaderTest(unittest.TestCase):
+
+    def setUp(self):
+        self.n = 2000
+        content_key = [
+            "robot_pose_raw" if i % 2 == 0 else "imu_raw"
+            for i in range(self.n)
+        ]
+        payload = [
+            {"value": b"v%d" % i, "metadata": b"m"}
+            for i in range(self.n)
+        ]
+        table = pa.table({
+            "content_key": pa.array(content_key),
+            "payload": pa.array(payload, type=_VARIANT_TYPE),
+        })
+        self.tmp = tempfile.mkdtemp()
+        self.path = os.path.join(self.tmp, "topics.parquet")
+        pq.write_table(table, self.path, row_group_size=1000)
+        self.assertEqual(2, pq.ParquetFile(self.path).num_row_groups)
+        self.read_fields = [
+            DataField(0, "content_key", AtomicType("STRING")),
+            DataField(1, "payload", AtomicType("VARIANT")),
+        ]
+
+    def _reader(self, read_fields, predicate=None, predicate_field_names=None):
+        return FormatPyArrowReader(
+            _LocalFileIO(), "parquet", self.path, read_fields,
+            predicate, batch_size=256,
+            predicate_field_names=predicate_field_names)
+
+    def _large_dictionary_payload(self):
+        if "store_schema" not in inspect.signature(pq.write_table).parameters:
+            self.skipTest("PyArrow does not support store_schema")
+        row_count = 15_340
+        values = pa.DictionaryArray.from_arrays(
+            pa.array([0] * row_count, type=pa.int32()),
+            pa.array([b"x" * 140_000], type=pa.binary()))
+        payload = pa.StructArray.from_arrays(
+            [values, pa.array([b"m"] * row_count)],
+            names=["value", "metadata"])
+        return row_count, payload
+
+    def _large_dictionary_variant(self):
+        row_count, payload = self._large_dictionary_payload()
+        path = os.path.join(self.tmp, "large-dictionary-variant.parquet")
+        pq.write_table(
+            pa.table({"payload": payload}), path,
+            use_dictionary=True, compression="zstd", store_schema=False)
+        reader = FormatPyArrowReader(
+            _LocalFileIO(), "parquet", path,
+            [DataField(0, "payload", AtomicType("VARIANT"))],
+            None, batch_size=128)
+        return pq.ParquetFile(path).num_row_groups, _drain(reader)
+
+    def _shredded_variant_payload(self):
+        schema = json.dumps({
+            "type": "ROW",
+            "fields": [{
+                "id": 0,
+                "name": "v",
+                "type": {
+                    "type": "ROW",
+                    "fields": [{
+                        "id": 1,
+                        "name": "age",
+                        "type": "BIGINT",
+                    }],
+                },
+            }],
+        })
+        obj_fields = parse_shredding_schema_option(schema)["v"]
+        target_type = shredding_schema_to_arrow_type(obj_fields)
+        variants = GenericVariant.to_arrow_array([
+            GenericVariant.from_python({"age": 1, "extra": "x"}),
+            None,
+            GenericVariant.from_python({"age": 3, "extra": "z"}),
+        ])
+        return shred_variant_column(variants, obj_fields, target_type)
+
+    def tearDown(self):
+        shutil.rmtree(self.tmp, ignore_errors=True)
+
+    def test_filter_only_column_not_in_projection(self):
+        reader = self._reader(
+            [DataField(1, "payload", AtomicType("VARIANT"))],
+            predicate=ds.field("content_key") == "robot_pose_raw",
+            predicate_field_names={"content_key"})
+        rows, columns, _ = _drain(reader)
+        self.assertEqual(self.n // 2, rows)
+        self.assertEqual(["payload"], columns)
+
+    def test_reads_all_rows_across_row_groups(self):
+        rows, columns, _ = _drain(self._reader(self.read_fields))
+        self.assertEqual(self.n, rows)
+        self.assertEqual(["content_key", "payload"], columns)
+
+    def test_reads_large_dictionary_variant_in_single_row_group(self):
+        row_groups, (rows, columns, _) = self._large_dictionary_variant()
+        self.assertEqual(1, row_groups)
+        self.assertEqual(15_340, rows)
+        self.assertEqual(["payload"], columns)
+
+    def test_predicate_returns_only_matching_rows(self):
+        predicate = ds.field("content_key") == "robot_pose_raw"
+        rows, _, content_keys = _drain(
+            self._reader(self.read_fields, predicate))
+        self.assertEqual(self.n // 2, rows)
+        self.assertEqual({"robot_pose_raw"}, content_keys)
+
+    def test_projection_returns_only_requested_columns(self):
+        rows, columns, _ = _drain(
+            self._reader([DataField(1, "payload", AtomicType("VARIANT"))]))
+        self.assertEqual(self.n, rows)
+        self.assertEqual(["payload"], columns)
+
+    def test_projection_preserves_requested_order(self):
+        reader = self._reader([
+            DataField(1, "payload", AtomicType("VARIANT")),
+            DataField(0, "content_key", AtomicType("STRING")),
+        ])
+        batch = reader.read_arrow_batch()
+        self.assertEqual(["payload", "content_key"], batch.schema.names)
+        self.assertEqual(
+            {"value": b"v0", "metadata": b"m"}, batch.column(0)[0].as_py())
+        self.assertEqual("robot_pose_raw", batch.column(1)[0].as_py())
+
+    def test_dotted_top_level_name_does_not_match_nested_path(self):
+        path = os.path.join(self.tmp, "dotted-name.parquet")
+        pq.write_table(pa.table({
+            "a": pa.array(
+                [{"b": "nested"}],
+                type=pa.struct([pa.field("b", pa.string())])),
+            "a.b": pa.array(["top"]),
+            "payload": pa.array(
+                [{"value": b"v", "metadata": b"m"}], type=_VARIANT_TYPE),
+        }), path)
+        reader = FormatPyArrowReader(
+            _LocalFileIO(), "parquet", path,
+            [
+                DataField(1, "a.b", AtomicType("STRING")),
+                DataField(2, "payload", AtomicType("VARIANT")),
+            ],
+            None, batch_size=128)
+        batch = reader.read_arrow_batch()
+        self.assertEqual(["a.b", "payload"], batch.schema.names)
+        self.assertEqual("top", batch.column(0)[0].as_py())
+        self.assertEqual(
+            {"value": b"v", "metadata": b"m"}, batch.column(1)[0].as_py())
+
+    def test_reads_variant_nested_in_container_types(self):
+        row_count, payload = self._large_dictionary_payload()
+        offsets = pa.array(range(row_count + 1), type=pa.int32())
+        variant_type = AtomicType("VARIANT")
+        cases = [
+            (
+                "row",
+                pa.StructArray.from_arrays([payload], names=["v"]),
+                RowType(True, [DataField(1, "v", variant_type)]),
+            ),
+            (
+                "array",
+                pa.ListArray.from_arrays(offsets, payload),
+                ArrayType(True, variant_type),
+            ),
+            (
+                "map",
+                pa.MapArray.from_arrays(
+                    offsets, pa.array(["k"] * row_count), payload),
+                MapType(True, AtomicType("STRING", False), variant_type),
+            ),
+        ]
+        for name, column, data_type in cases:
+            with self.subTest(name=name):
+                path = os.path.join(self.tmp, "nested-{}.parquet".format(name))
+                pq.write_table(
+                    pa.table({name: column}), path,
+                    use_dictionary=True, compression="zstd",
+                    store_schema=False)
+                reader = FormatPyArrowReader(
+                    _LocalFileIO(), "parquet", path,
+                    [DataField(0, name, data_type)], None, batch_size=128)
+                rows, columns, _ = _drain(reader)
+                self.assertEqual(row_count, rows)
+                self.assertEqual([name], columns)
+
+    def test_nested_variant_paths_use_bounded_reader(self):
+        path = os.path.join(self.tmp, "small-nested-variant.parquet")
+        payload = pa.array([
+            {"value": b"v0", "metadata": b"m"},
+            {"value": b"v1", "metadata": b"m"},
+            {"value": b"v2", "metadata": b"m"},
+        ], type=_VARIANT_TYPE)
+        offsets = pa.array([0, 1, 2, 3], type=pa.int32())
+        pq.write_table(pa.table({
+            "row": pa.StructArray.from_arrays(
+                [payload], names=["v"],
+                mask=pa.array([False, False, True])),
+            "array": pa.ListArray.from_arrays(offsets, payload),
+            "map": pa.MapArray.from_arrays(
+                offsets, pa.array(["k", "k", "k"]), payload),
+            "kind": pa.array(["keep", "drop", "keep"]),
+        }), path)
+
+        variant_type = AtomicType("VARIANT")
+        cases = [
+            ("row", RowType(True, [DataField(1, "v", variant_type)])),
+            ("array", ArrayType(True, variant_type)),
+            ("map", MapType(
+                True, AtomicType("STRING", False), variant_type)),
+        ]
+        for name, data_type in cases:
+            reader = FormatPyArrowReader(
+                _LocalFileIO(), "parquet", path,
+                [DataField(0, name, data_type)], None, batch_size=128)
+            self.assertIsNotNone(reader._parquet_file)
+            self.assertEqual(3, _drain(reader)[0])
+
+        nested_reader = FormatPyArrowReader(
+            _LocalFileIO(), "parquet", path,
+            [DataField(1, "row_v", AtomicType("VARIANT"))],
+            ds.field("kind") == "keep", batch_size=128,
+            nested_name_paths=[["row", "v"]],
+            predicate_field_names={"kind"})
+        self.assertIsNotNone(nested_reader._parquet_file)
+        batch = nested_reader.read_arrow_batch()
+        self.assertEqual(["row_v"], batch.schema.names)
+        self.assertEqual([
+            {"value": b"v0", "metadata": b"m"}, None,
+        ], batch.column(0).to_pylist())
+
+    def test_reads_nested_variant_projection(self):
+        row_count, payload = self._large_dictionary_payload()
+        path = os.path.join(self.tmp, "nested-projection.parquet")
+        pq.write_table(
+            pa.table({
+                "row": pa.StructArray.from_arrays([payload], names=["v"]),
+                "kind": pa.array(
+                    ["keep" if i % 2 == 0 else "drop"
+                     for i in range(row_count)]),
+            }),
+            path, use_dictionary=True, compression="zstd", store_schema=False)
+        reader = FormatPyArrowReader(
+            _LocalFileIO(), "parquet", path,
+            [DataField(0, "row_v", AtomicType("VARIANT"))],
+            ds.field("kind") == "keep", batch_size=128,
+            nested_name_paths=[["row", "v"]],
+            predicate_field_names={"kind"})
+        rows, columns, _ = _drain(reader)
+        self.assertEqual((row_count + 1) // 2, rows)
+        self.assertEqual(["row_v"], columns)
+
+    def test_assembles_shredded_variant_in_nested_types(self):
+        shredded = self._shredded_variant_payload()
+        values = shredded.to_pylist()
+        path = os.path.join(self.tmp, "nested-shredded-variant.parquet")
+        pq.write_table(pa.table({
+            "row": pa.array(
+                [{"v": values[0]}, {"v": values[1]}, None, {"v": values[2]}],
+                type=pa.struct([pa.field("v", shredded.type)])),
+            "array": pa.array(
+                [[values[0]], [values[1]], None, [values[2]]],
+                type=pa.list_(pa.field("item", shredded.type))),
+            "map": pa.array(
+                [[("k", values[0])], [("k", values[1])], None,
+                 [("k", values[2])]],
+                type=pa.map_(
+                    pa.field("key", pa.string(), nullable=False),
+                    pa.field("value", shredded.type))),
+            "kind": pa.array(["keep", "drop", "keep", "keep"]),
+        }), path)
+
+        variant_type = AtomicType("VARIANT")
+        data_types = {
+            "row": RowType(True, [DataField(1, "v", variant_type)]),
+            "array": ArrayType(True, variant_type),
+            "map": MapType(
+                True, AtomicType("STRING", False), variant_type),
+        }
+        outputs = {}
+        for name, data_type in data_types.items():
+            reader = FormatPyArrowReader(
+                _LocalFileIO(), "parquet", path,
+                [DataField(0, name, data_type)], None, batch_size=128)
+            outputs[name] = reader.read_arrow_batch().column(0)
+
+        self.assertEqual(_VARIANT_TYPE, outputs["row"].type[0].type)
+        self.assertEqual(_VARIANT_TYPE, outputs["array"].type.value_type)
+        self.assertEqual(_VARIANT_TYPE, outputs["map"].type.item_type)
+
+        def decode(value):
+            if value is None:
+                return None
+            return GenericVariant.from_arrow_struct(value).to_python()
+
+        row_values = outputs["row"].to_pylist()
+        self.assertEqual({"age": 1, "extra": "x"}, decode(row_values[0]["v"]))
+        self.assertIsNone(row_values[1]["v"])
+        self.assertIsNone(row_values[2])
+
+        array_values = outputs["array"].to_pylist()
+        self.assertEqual({"age": 1, "extra": "x"}, decode(array_values[0][0]))
+        self.assertIsNone(array_values[1][0])
+        self.assertIsNone(array_values[2])
+
+        map_values = outputs["map"].to_pylist()
+        self.assertEqual(
+            {"age": 1, "extra": "x"}, decode(map_values[0][0][1]))
+        self.assertIsNone(map_values[1][0][1])
+        self.assertIsNone(map_values[2])
+
+        nested_reader = FormatPyArrowReader(
+            _LocalFileIO(), "parquet", path,
+            [DataField(1, "row_v", AtomicType("VARIANT"))],
+            ds.field("kind") == "keep", batch_size=128,
+            nested_name_paths=[["row", "v"]],
+            predicate_field_names={"kind"})
+        nested = nested_reader.read_arrow_batch().column(0)
+        self.assertEqual(_VARIANT_TYPE, nested.type)
+        nested_values = nested.to_pylist()
+        self.assertEqual({"age": 1, "extra": "x"}, decode(nested_values[0]))
+        self.assertIsNone(nested_values[1])
+        self.assertEqual({"age": 3, "extra": "z"}, decode(nested_values[2]))
+
+    def test_single_row_group_scalar_read_uses_fast_path(self):
+        rows, columns, _ = _drain(
+            self._reader([DataField(0, "content_key", AtomicType("STRING"))]))
+        self.assertEqual(self.n, rows)
+        self.assertEqual(["content_key"], columns)
+
+    def test_row_group_pruning_by_statistics(self):
+        rows_per_group = 1000
+        content_key = []
+        for group in range(8):
+            content_key += (
+                ["match" if group == 3 else "other"] * rows_per_group)
+        total = len(content_key)
+        payload = [
+            {"value": b"v%d" % i, "metadata": b"m"}
+            for i in range(total)
+        ]
+        path = os.path.join(self.tmp, "clustered.parquet")
+        pq.write_table(
+            pa.table({
+                "content_key": pa.array(content_key),
+                "payload": pa.array(payload, type=_VARIANT_TYPE),
+            }),
+            path, row_group_size=rows_per_group)
+        self.assertEqual(8, pq.ParquetFile(path).num_row_groups)
+
+        reader = FormatPyArrowReader(
+            _LocalFileIO(), "parquet", path, self.read_fields,
+            ds.field("content_key") == "match", batch_size=512)
+        self.assertEqual([3], list(reader._surviving_row_group_ids()))
+        rows, _, keys = _drain(reader)
+        self.assertEqual(rows_per_group, rows)
+        self.assertEqual({"match"}, keys)
+
+
+if __name__ == "__main__":
+    unittest.main()


Reply via email to