diff --git a/paimon-core/src/test/java/org/apache/paimon/JavaPyE2ETest.java b/paimon-core/src/test/java/org/apache/paimon/JavaPyE2ETest.java index 524344ed54e1..8a59c4332612 100644 --- a/paimon-core/src/test/java/org/apache/paimon/JavaPyE2ETest.java +++ b/paimon-core/src/test/java/org/apache/paimon/JavaPyE2ETest.java @@ -1620,6 +1620,67 @@ public void testJavaWriteSharedShreddingMapTable() throws Exception { } } + /** Java reads shared-shredding MAP columns written by Python. */ + @Test + @EnabledIfSystemProperty(named = "run.e2e.tests", matches = "true") + public void testJavaReadSharedShreddingMapTable() throws Exception { + FileStoreTable table = + (FileStoreTable) + catalog.getTable(identifier("shared_shredding_map_python_test_parquet")); + Map> rows = new HashMap<>(); + List splits = new ArrayList<>(table.newSnapshotReader().read().dataSplits()); + try (org.apache.paimon.reader.RecordReader reader = + table.newRead().createReader(splits)) { + reader.forEachRemaining( + row -> { + int id = row.getInt(0); + for (int column = 2; column <= 3; column++) { + if (id == 3) { + assertThat(row.isNullAt(column)).isTrue(); + } else { + InternalMap required = row.getMap(column); + assertThat(required.size()).isEqualTo(id == 2 ? 0 : 1); + if (id != 2) { + InternalArray values = required.valueArray(); + long value = + column == 2 + ? values.getLong(0) + : values.getRow(0, 1).getLong(0); + assertThat(value).isEqualTo(id == 1 ? 1L : 2L); + } + } + } + if (row.isNullAt(1)) { + rows.put(id, null); + return; + } + InternalMap map = row.getMap(1); + InternalArray keys = map.keyArray(); + InternalArray values = map.valueArray(); + Map converted = new LinkedHashMap<>(); + for (int i = 0; i < map.size(); i++) { + converted.put( + keys.getString(i).toString(), + values.isNullAt(i) ? null : values.getLong(i)); + } + rows.put(id, converted); + }); + } + + assertThat(rows).containsOnlyKeys(1, 2, 3, 4); + assertThat(rows.get(1)) + .containsOnlyKeys("hot", "warm", "overflow") + .containsEntry("hot", 10L) + .containsEntry("warm", 20L) + .containsEntry("overflow", 30L); + assertThat(rows.get(2)) + .containsOnlyKeys("hot", "new") + .containsEntry("hot", null) + .containsEntry("new", 40L); + assertThat(rows.get(3)).isEmpty(); + assertThat(rows.get(4)).isNull(); + } + private Map> readMapBlobRows(FileStoreTable table) throws Exception { Map> rows = new HashMap<>(); diff --git a/paimon-python/dev/run_mixed_tests.sh b/paimon-python/dev/run_mixed_tests.sh index 620c4d989a03..54fef2a4e1c0 100755 --- a/paimon-python/dev/run_mixed_tests.sh +++ b/paimon-python/dev/run_mixed_tests.sh @@ -1033,6 +1033,21 @@ run_shared_shredding_map_test() { return 1 fi echo -e "${GREEN}✓ Python shared-shredding MAP read test completed successfully${NC}" + + echo "Running Python shared-shredding MAP write test..." + if ! python -m pytest java_py_read_write_test.py::JavaPyReadWriteTest::test_write_shared_shredding_map_for_java -v; then + echo -e "${RED}✗ Python shared-shredding MAP write test failed${NC}" + return 1 + fi + echo -e "${GREEN}✓ Python shared-shredding MAP write test completed successfully${NC}" + + cd "$PROJECT_ROOT" + echo "Running Maven test for JavaPyE2ETest.testJavaReadSharedShreddingMapTable..." + if ! mvn test -Dtest=org.apache.paimon.JavaPyE2ETest#testJavaReadSharedShreddingMapTable -pl paimon-core -q -Drun.e2e.tests=true; then + echo -e "${RED}✗ Java shared-shredding MAP read test failed${NC}" + return 1 + fi + echo -e "${GREEN}✓ Java shared-shredding MAP read test completed successfully${NC}" } # Function to run VARIANT test (Java write, Python read) diff --git a/paimon-python/pypaimon/common/options/core_options.py b/paimon-python/pypaimon/common/options/core_options.py index ce77d21f5a84..a22957388eb2 100644 --- a/paimon-python/pypaimon/common/options/core_options.py +++ b/paimon-python/pypaimon/common/options/core_options.py @@ -158,6 +158,10 @@ class CoreOptions: NESTED_SEQUENCE_FIELD = "nested-sequence-field" COUNT_LIMIT = "count-limit" MERGE_MAP_TS_FIELD = "ts-field" + MAP_STORAGE_LAYOUT = "map.storage-layout" + MAP_SHARED_SHREDDING_MAX_COLUMNS = "map.shared-shredding.max-columns" + MAP_SHARED_SHREDDING_COLUMN_PLACEMENT_POLICY = \ + "map.shared-shredding.column-placement-policy" # Basic options AUTO_CREATE: ConfigOption[bool] = ( @@ -1873,6 +1877,46 @@ def field_merge_map_ts_field(self, field_name: str) -> str: .no_default_value() ) + def map_storage_layout(self, field_name: str) -> str: + return self.options.get( + ConfigOptions.key( + f'{CoreOptions.FIELDS_PREFIX}.{field_name}.{CoreOptions.MAP_STORAGE_LAYOUT}' + ) + .string_type() + .default_value('default') + ).lower() + + def map_shared_shredding_max_columns(self, field_name: str) -> int: + value = self.options.get( + ConfigOptions.key( + f'{CoreOptions.FIELDS_PREFIX}.{field_name}.' + f'{CoreOptions.MAP_SHARED_SHREDDING_MAX_COLUMNS}' + ) + .int_type() + .default_value(256) + ) + if value <= 0: + raise ValueError( + '{} must be greater than 0'.format( + CoreOptions.MAP_SHARED_SHREDDING_MAX_COLUMNS)) + return value + + def map_shared_shredding_column_placement_policy( + self, field_name: str) -> str: + value = self.options.get( + ConfigOptions.key( + f'{CoreOptions.FIELDS_PREFIX}.{field_name}.' + f'{CoreOptions.MAP_SHARED_SHREDDING_COLUMN_PLACEMENT_POLICY}' + ) + .string_type() + .default_value('lru') + ).lower() + if value not in ('plain', 'sequential', 'lru'): + raise ValueError( + "Unsupported shared-shredding column placement policy: {}".format( + value)) + return value + @property def query_auth_enabled(self) -> bool: return self.options.get(CoreOptions.QUERY_AUTH_ENABLED) diff --git a/paimon-python/pypaimon/data/map_shared_shredding.py b/paimon-python/pypaimon/data/map_shared_shredding.py index a3d6c58bac55..13c15df75fdc 100644 --- a/paimon-python/pypaimon/data/map_shared_shredding.py +++ b/paimon-python/pypaimon/data/map_shared_shredding.py @@ -14,7 +14,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Read support for Paimon's shared-shredding MAP storage layout.""" +"""Metadata and read support for the shared-shredding MAP layout.""" import json import struct @@ -33,6 +33,7 @@ _NUM_COLUMNS = b"paimon.map.shared-shredding.num-columns" _FIELD_COLUMNS = b"paimon.map.shared-shredding.field-columns" _OVERFLOW_SET = b"paimon.map.shared-shredding.overflow-set" +_MAX_ROW_WIDTH = b"paimon.map.shared-shredding.max-row-width" _FIELD_MAPPING = "__field_mapping" _OVERFLOW = "__overflow" _PHYSICAL_COLUMN_PREFIX = "__col_" @@ -104,6 +105,40 @@ def parse_shared_shredding_selection_metadata(field: pa.Field): return name_by_id, field_to_columns, set(overflow_json), num_columns +def shared_shredding_metadata( + name_to_id, field_to_columns, overflow_set, num_columns, + max_row_width, compression): + """Build Java-compatible Arrow field metadata for one data file.""" + compression = compression.lower() + if compression not in ("none", "lz4", "zstd"): + raise ValueError( + "MAP shared-shredding only supports none/lz4/zstd compression, " + "but is {}.".format(compression)) + field_dict = json.dumps( + dict(sorted(name_to_id.items())), + separators=(",", ":"), + sort_keys=True, + ).encode("utf-8") + encoded_dict = _compress(field_dict, compression) + columns = { + str(field_id): sorted(column_ids) + for field_id, column_ids in sorted(field_to_columns.items()) + } + return { + _STORAGE_LAYOUT: b"shared-shredding", + _VERSION: b"1", + _FIELD_DICT: encoded_dict.decode("latin-1").encode("utf-8"), + _FIELD_DICT_COMPRESSION: compression.encode("utf-8"), + _FIELD_DICT_ORIGINAL_SIZE: str(len(field_dict)).encode("utf-8"), + _FIELD_COLUMNS: json.dumps( + columns, separators=(",", ":"), sort_keys=True).encode("utf-8"), + _OVERFLOW_SET: json.dumps( + sorted(overflow_set), separators=(",", ":")).encode("utf-8"), + _NUM_COLUMNS: str(num_columns).encode("utf-8"), + _MAX_ROW_WIDTH: str(max_row_width).encode("utf-8"), + } + + def map_selected_keys(description: str) -> List[str]: if not description or not description.startswith(_SELECTED_KEYS_PREFIX): raise ValueError("Invalid selected-key MAP metadata: {}".format( @@ -589,6 +624,20 @@ def _decompress(data: bytes, original_size: int, compression: str) -> bytes: return result +def _compress(data: bytes, compression: str) -> bytes: + if compression == "none": + return data + if compression == "zstd": + import zstandard as zstd + return zstd.ZstdCompressor(level=1).compress(data) + if compression == "lz4": + payload = bytes(pa.Codec("lz4_raw").compress(data)) + return struct.pack("= 0: + if column_id not in slots: + slots[column_id] = [None] * len(column) + slots[column_id][row] = items[field_id] + overflows.append([(field_id, items[field_id]) for field_id in overflow] + if overflow else None) + children = [pa.ListArray.from_arrays( + pa.array(mapping_offsets, type=pa.int32()), + pa.array(mappings, type=pa.int32()))] + empty = None + for column_id in range(self.num_columns): + value_type = self.physical_type[column_id + 1].type + if column_id in slots: + children.append(pa.array(slots.pop(column_id), type=value_type)) + else: + if empty is None: + empty = pa.array([None] * len(column), type=value_type) + children.append(empty) + children.append(pa.array(overflows, type=self.physical_type[-1].type)) + return pa.StructArray.from_arrays( + children, fields=list(self.physical_type), + mask=column.is_null() if column.null_count else None) + + def analyze(self, column): + offsets, start, end = _normalized_offsets(column) + keys = column.keys.slice(start, end - start).to_pylist() + for row, is_null in enumerate(column.is_null().to_pylist()): + if not is_null: + self._place(keys[offsets[row]:offsets[row + 1]]) + + def _place(self, keys): + field_ids = [] + for key in keys: + if key is None: + raise ValueError("Shared-shredding MAP keys cannot be null") + if not isinstance(key, str): + raise TypeError("Shared-shredding MAP keys must be strings") + field_id = self.name_to_id.setdefault(key, len(self.name_to_id)) + field_ids.append(field_id) + + mapping, overflow = self._allocate(field_ids) + for column_id, field_id in enumerate(mapping): + if field_id >= 0: + self.field_to_columns.setdefault(field_id, set()).add(column_id) + self.overflow_set.update(overflow) + self.max_row_width = max(self.max_row_width, len(field_ids)) + return field_ids, mapping, overflow + + def _allocate(self, field_ids): + if self.policy == "plain": + ordered = field_ids + return self._leading(ordered) + if self.policy == "sequential": + return self._leading(sorted(field_ids)) + return self._lru(field_ids) + + def _leading(self, field_ids): + mapping = [-1] * self.num_columns + for index, field_id in enumerate(field_ids[:self.num_columns]): + mapping[index] = field_id + return mapping, list(field_ids[self.num_columns:]) + + def _lru(self, field_ids): + mapping = [-1] * self.num_columns + next_resident = list(self._resident) + used = [False] * self.num_columns + unassigned = [] + for field_id in sorted(field_ids): + try: + column_id = self._resident.index(field_id) + except ValueError: + unassigned.append(field_id) + continue + used[column_id] = True + mapping[column_id] = field_id + + overflow = [] + for field_id in unassigned: + column_id = self._select_lru_column(used, next_resident) + if column_id < 0: + overflow.append(field_id) + continue + used[column_id] = True + mapping[column_id] = field_id + next_resident[column_id] = field_id + + touched = False + for column_id, field_id in enumerate(mapping): + if field_id >= 0: + self._last_used[column_id] = self._clock + touched = True + if touched: + self._clock += 1 + self._resident = next_resident + return mapping, overflow + + def _select_lru_column(self, used, resident): + selected = -1 + selected_last_used = None + for column_id in range(self.num_columns): + if used[column_id]: + continue + if resident[column_id] < 0: + return column_id + last_used = self._last_used[column_id] + if selected < 0 or last_used < selected_last_used: + selected = column_id + selected_last_used = last_used + return selected + + +def _bounded_batches(batch): + # One oversized row is indivisible; bound the other logical batches before + # expanding values into Python objects. The caller's input buffer is unchanged. + if batch.nbytes > _CONVERSION_BYTES and batch.num_rows > 1: + middle = batch.num_rows // 2 + yield from _bounded_batches(batch.slice(0, middle)) + yield from _bounded_batches(batch.slice(middle)) + else: + yield batch + + +def _to_python_values(column): + """Avoid Arrow 6 MAP scalars, including MAPs nested in ROW/ARRAY values.""" + data_type = column.type + if pa.types.is_struct(data_type): + children = [_to_python_values(column.field(i)) + for i in range(len(data_type))] + names = [field.name for field in data_type] + values = [dict(zip(names, (child[i] for child in children))) + for i in range(len(column))] + elif (pa.types.is_map(data_type) or pa.types.is_list(data_type) + or pa.types.is_large_list(data_type)): + offsets, start, end = _normalized_offsets(column) + if pa.types.is_map(data_type): + keys = _to_python_values(column.keys.slice(start, end - start)) + items = _to_python_values(column.items.slice(start, end - start)) + children = list(zip(keys, items)) + else: + children = _to_python_values(column.values.slice(start, end - start)) + values = [children[offsets[i]:offsets[i + 1]] for i in range(len(column))] + elif pa.types.is_fixed_size_list(data_type): + size = data_type.list_size + children = _to_python_values(column.values.slice( + column.offset * size, len(column) * size)) + values = [children[i * size:(i + 1) * size] for i in range(len(column))] + else: + return column.to_pylist() + return [None if is_null else value + for is_null, value in zip(column.is_null().to_pylist(), values)] + + +def _physical_struct_type(num_columns, item_type, logical_item_type): + mapping_type = pa.list_(pa.field( + "item", + pa.int32(), + metadata=_field_id_metadata(_array_element_id(0, 1)), + )) + fields = [pa.field( + _FIELD_MAPPING, + mapping_type, + metadata=_field_id_metadata(0), + )] + for column_id in range(num_columns): + fields.append(_field_with_ids( + _PHYSICAL_COLUMN_PREFIX + str(column_id), + item_type, + True, + logical_item_type, + column_id + 1, + )) + overflow_id = num_columns + 1 + overflow_type = _map_type_with_ids( + pa.int32(), item_type, logical_item_type, overflow_id, 0) + fields.append(pa.field( + _OVERFLOW, + overflow_type, + metadata=_field_id_metadata(overflow_id), + )) + return pa.struct(fields) + + +def _field_with_ids(name, arrow_type, nullable, logical_type, field_id): + return pa.field( + name, + _type_with_ids(arrow_type, logical_type, field_id, 0), + nullable=nullable, + metadata=_field_id_metadata(field_id), + ) + + +def _type_with_ids(arrow_type, logical_type, field_id, depth): + if isinstance(logical_type, RowType) and pa.types.is_struct(arrow_type): + arrow_by_name = {field.name: field for field in arrow_type} + return pa.struct([ + _field_with_ids( + field.name, + arrow_by_name[field.name].type, + arrow_by_name[field.name].nullable, + field.type, + field.id, + ) + for field in logical_type.fields + ]) + if isinstance(logical_type, (ArrayType, VectorType)): + child_id = _array_element_id(field_id, depth + 1) + value_field = pa.field( + "item", + _type_with_ids( + arrow_type.value_type, + logical_type.element, + field_id, + depth + 1, + ), + nullable=logical_type.element.nullable, + metadata=_field_id_metadata(child_id), + ) + if isinstance(logical_type, VectorType): + return pa.list_(value_field, logical_type.length) + return pa.list_(value_field) + if isinstance(logical_type, MapType): + return _map_type_with_ids( + arrow_type.key_type, + arrow_type.item_type, + logical_type.value, + field_id, + depth, + ) + return arrow_type + + +def _map_type_with_ids( + key_type, item_type, logical_item_type, field_id, depth): + key = pa.field( + "key", + key_type, + nullable=False, + metadata=_field_id_metadata(_map_key_id(field_id, depth + 1)), + ) + item = pa.field( + "value", + _type_with_ids( + item_type, logical_item_type, field_id, depth + 1), + nullable=logical_item_type.nullable, + metadata=_field_id_metadata(_map_value_id(field_id, depth + 1)), + ) + return pa.map_(key, item) + + +def _field_id_metadata(field_id): + value = str(field_id).encode("utf-8") + return {b"PARQUET:field_id": value, b"paimon.id": value} + + +def _array_element_id(field_id, depth): + return _FIELD_ID_BASE + field_id * _FIELD_ID_DEPTH_LIMIT + depth + + +def _map_key_id(field_id, depth): + return _FIELD_ID_BASE - field_id * _FIELD_ID_DEPTH_LIMIT - depth + + +def _map_value_id(field_id, depth): + return _FIELD_ID_BASE + field_id * _FIELD_ID_DEPTH_LIMIT + depth + + +def _contains_type(data_type, predicate): + if predicate(data_type): + return True + if isinstance(data_type, RowType): + return any(_contains_type(field.type, predicate) + for field in data_type.fields) + if isinstance(data_type, (ArrayType, VectorType, MultisetType)): + return _contains_type(data_type.element, predicate) + if isinstance(data_type, MapType): + return (_contains_type(data_type.key, predicate) + or _contains_type(data_type.value, predicate)) + return False + + +def _is_variant(data_type): + return (isinstance(data_type, AtomicType) + and data_type.type.upper() == "VARIANT") + + +def _is_blob(data_type): + return (isinstance(data_type, AtomicType) + and data_type.type.upper() == "BLOB") diff --git a/paimon-python/pypaimon/write/writer/data_vector_writer.py b/paimon-python/pypaimon/write/writer/data_vector_writer.py index 9294f67d92cb..4825b22af138 100644 --- a/paimon-python/pypaimon/write/writer/data_vector_writer.py +++ b/paimon-python/pypaimon/write/writer/data_vector_writer.py @@ -261,11 +261,13 @@ def _write_normal_data_to_file(self, data: pa.Table) -> Optional[DataFileMeta]: if data.num_rows == 0: return None + shredding_stats = {} + file_name = f"{CoreOptions.data_file_prefix(self.options)}{uuid.uuid4()}-0.{self.file_format}" file_path = self._generate_file_path(file_name) if self.file_format == CoreOptions.FILE_FORMAT_PARQUET: - self.file_io.write_parquet(file_path, data, compression=self.compression, zstd_level=self.zstd_level) + shredding_stats = self._write_parquet_data(file_path, data) elif self.file_format == CoreOptions.FILE_FORMAT_ORC: self.file_io.write_orc(file_path, data, compression=self.compression, zstd_level=self.zstd_level) elif self.file_format == CoreOptions.FILE_FORMAT_AVRO: @@ -290,7 +292,7 @@ def _write_normal_data_to_file(self, data: pa.Table) -> Optional[DataFileMeta]: min_seq, max_seq = self._append_file_sequence_range(data.num_rows) - return DataFileMeta.create( + meta = DataFileMeta.create( file_name=file_name, file_size=self.file_io.get_file_size(file_path), row_count=data.num_rows, @@ -311,6 +313,8 @@ def _write_normal_data_to_file(self, data: pa.Table) -> Optional[DataFileMeta]: file_path=file_path, write_cols=self.write_cols, ) + self._map_shared_shredding.file_completed(shredding_stats) + return meta def _validate_consistency( self, normal_meta: DataFileMeta, vector_metas: List[DataFileMeta]): diff --git a/paimon-python/pypaimon/write/writer/data_writer.py b/paimon-python/pypaimon/write/writer/data_writer.py index 24a67186c3d8..46db0578ceb8 100644 --- a/paimon-python/pypaimon/write/writer/data_writer.py +++ b/paimon-python/pypaimon/write/writer/data_writer.py @@ -29,6 +29,7 @@ from pypaimon.schema.data_types import PyarrowFieldParser from pypaimon.table.bucket_mode import BucketMode from pypaimon.table.row.generic_row import GenericRow +from pypaimon.write.map_shared_shredding_writer import MapSharedShreddingWriter from pypaimon.write.writer.mosaic_writer_options import create_mosaic_writer_options from pypaimon.write.writer.write_buffer import WriteBuffer @@ -104,6 +105,13 @@ def __init__(self, table, partition: Tuple, bucket: int, max_seq_number: int, op # Paimon field id map, used by _apply_variant_shredding; built once since # the table schema is fixed for the lifetime of this writer. self._paimon_field_id: Dict[str, int] = {pf.name: pf.id for pf in self.table.fields} + self._map_shared_shredding = MapSharedShreddingWriter( + self.table.fields, + self.options, + self.file_format, + self.changelog_file_format + if self.changelog_producer == ChangelogProducer.INPUT else None, + ) # Set by the composite writers when a flush landed its normal data file but a # later phase of the same flush failed; see their ``_close_current_writers``. @@ -259,6 +267,7 @@ def _write_data_to_file(self, data: pa.Table): extra_files = [] row_sidecar_path = None changelog_meta = None + shared_shredding_stats = {} if self._variant_shredding: data = self._apply_variant_shredding(data) @@ -270,7 +279,7 @@ def _write_data_to_file(self, data: pa.Table): # already covers. try: if self.file_format == CoreOptions.FILE_FORMAT_PARQUET: - self.file_io.write_parquet(file_path, data, compression=self.compression, zstd_level=self.zstd_level) + shared_shredding_stats = self._write_parquet_data(file_path, data) elif self.file_format == CoreOptions.FILE_FORMAT_ORC: self.file_io.write_orc(file_path, data, compression=self.compression, zstd_level=self.zstd_level) elif self.file_format == CoreOptions.FILE_FORMAT_AVRO: @@ -300,7 +309,7 @@ def _write_data_to_file(self, data: pa.Table): # min key & max key - selected_table = data.select(self.trimmed_primary_keys) + selected_table = logical_data.select(self.trimmed_primary_keys) key_columns_batch = selected_table.to_batches()[0] min_key_row_batch = key_columns_batch.slice(0, 1) max_key_row_batch = key_columns_batch.slice(key_columns_batch.num_rows - 1, 1) @@ -311,21 +320,23 @@ def _write_data_to_file(self, data: pa.Table): value_stats_enabled = self.options.metadata_stats_enabled() if value_stats_enabled: stats_fields = self.table.fields if self.table.is_primary_key_table \ - else PyarrowFieldParser.to_paimon_schema(data.schema) + else PyarrowFieldParser.to_paimon_schema(logical_data.schema) else: stats_fields = self.table.trimmed_primary_keys_fields column_stats = { - field.name: self._get_column_stats(data, field.name) + field.name: self._get_column_stats(logical_data, field.name) for field in stats_fields } key_fields = self.trimmed_primary_keys_fields - key_stats = self._collect_value_stats(data, key_fields, column_stats) + key_stats = self._collect_value_stats( + logical_data, key_fields, column_stats) if not self.options.primary_key_nullable() and not all( count == 0 for count in key_stats.null_counts): raise RuntimeError("Primary key should not be null") value_fields = stats_fields if value_stats_enabled else [] - value_stats = self._collect_value_stats(data, value_fields, column_stats) + value_stats = self._collect_value_stats( + logical_data, value_fields, column_stats) # Read the range without advancing it: the advance belongs with the # append below, so a retried flush derives the same range. @@ -368,10 +379,18 @@ def _write_data_to_file(self, data: pa.Table): raise self.sequence_generator.start = self.sequence_generator.current + self._map_shared_shredding.file_completed(shared_shredding_stats) self.committed_files.append(data_meta) if changelog_meta is not None: self.committed_changelog_files.append(changelog_meta) + def _write_parquet_data(self, path, data): + if self._map_shared_shredding.is_active(): + return self._map_shared_shredding.write_parquet( + self.file_io, path, data, self.compression, self.zstd_level) + self.file_io.write_parquet(path, data, compression=self.compression, zstd_level=self.zstd_level) + return {} + def _apply_variant_shredding(self, data: pa.Table) -> pa.Table: """Transform VARIANT columns into shredded Parquet format. @@ -414,8 +433,7 @@ def _write_changelog_file(self, data, min_key, max_key, key_stats, value_stats, try: if cl_fmt == CoreOptions.FILE_FORMAT_PARQUET: - self.file_io.write_parquet(changelog_file_path, data, compression=self.compression, - zstd_level=self.zstd_level) + self._write_parquet_data(changelog_file_path, data) elif cl_fmt == CoreOptions.FILE_FORMAT_ORC: self.file_io.write_orc(changelog_file_path, data, compression=self.compression, zstd_level=self.zstd_level) diff --git a/paimon-python/pypaimon/write/writer/dedicated_format_writer.py b/paimon-python/pypaimon/write/writer/dedicated_format_writer.py index c894182fbe0c..f79369d1cdff 100644 --- a/paimon-python/pypaimon/write/writer/dedicated_format_writer.py +++ b/paimon-python/pypaimon/write/writer/dedicated_format_writer.py @@ -703,12 +703,14 @@ def _write_normal_data_to_file(self, data: pa.Table) -> Optional[DataFileMeta]: if data.num_rows == 0: return None + shredding_stats = {} + file_name = f"{CoreOptions.data_file_prefix(self.options)}{uuid.uuid4()}-0.{self.file_format}" file_path = self._generate_file_path(file_name) # Write file based on format if self.file_format == CoreOptions.FILE_FORMAT_PARQUET: - self.file_io.write_parquet(file_path, data, compression=self.compression, zstd_level=self.zstd_level) + shredding_stats = self._write_parquet_data(file_path, data) elif self.file_format == CoreOptions.FILE_FORMAT_ORC: self.file_io.write_orc(file_path, data, compression=self.compression, zstd_level=self.zstd_level) elif self.file_format == CoreOptions.FILE_FORMAT_AVRO: @@ -728,7 +730,9 @@ def _write_normal_data_to_file(self, data: pa.Table) -> Optional[DataFileMeta]: is_external_path = self.external_path_provider is not None external_path_str = file_path if is_external_path else None - return self._create_data_file_meta(file_name, file_path, data, external_path_str) + meta = self._create_data_file_meta(file_name, file_path, data, external_path_str) + self._map_shared_shredding.file_completed(shredding_stats) + return meta def _create_data_file_meta(self, file_name: str, file_path: str, data: pa.Table, external_path: Optional[str] = None) -> DataFileMeta: