Skip to main content

iceberg/arrow/
record_batch_partition_splitter.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use std::collections::HashMap;
19use std::sync::Arc;
20
21use arrow_array::{ArrayRef, BooleanArray, RecordBatch, StructArray};
22use arrow_buffer::BooleanBufferBuilder;
23use arrow_select::filter::filter_record_batch;
24
25use super::arrow_struct_to_literal;
26use super::partition_value_calculator::PartitionValueCalculator;
27use crate::Result;
28use crate::error::invalid_data;
29use crate::spec::{Literal, PartitionKey, PartitionSpecRef, SchemaRef, StructType};
30
31/// Column name for the projected partition values struct
32pub const PROJECTED_PARTITION_VALUE_COLUMN: &str = "_partition";
33
34/// The splitter used to split the record batch into multiple record batches by the partition spec.
35/// 1. It will project and transform the input record batch based on the partition spec, get the partitioned record batch.
36/// 2. Split the input record batch into multiple record batches based on the partitioned record batch.
37///
38/// # Partition Value Modes
39///
40/// The splitter supports two modes for obtaining partition values:
41/// - **Computed mode** (`calculator` is `Some`): Computes partition values from source columns using transforms
42/// - **Pre-computed mode** (`calculator` is `None`): Expects a `_partition` column in the input batch
43pub struct RecordBatchPartitionSplitter {
44    schema: SchemaRef,
45    partition_spec: PartitionSpecRef,
46    calculator: Option<PartitionValueCalculator>,
47    partition_type: StructType,
48}
49
50impl RecordBatchPartitionSplitter {
51    /// Create a new RecordBatchPartitionSplitter.
52    ///
53    /// # Arguments
54    ///
55    /// * `iceberg_schema` - The Iceberg schema reference
56    /// * `partition_spec` - The partition specification reference
57    /// * `calculator` - Optional calculator for computing partition values from source columns.
58    ///   - `Some(calculator)`: Compute partition values from source columns using transforms
59    ///   - `None`: Expect a pre-computed `_partition` column in the input batch
60    ///
61    /// # Returns
62    ///
63    /// Returns a new `RecordBatchPartitionSplitter` instance or an error if initialization fails.
64    pub fn try_new(
65        iceberg_schema: SchemaRef,
66        partition_spec: PartitionSpecRef,
67        calculator: Option<PartitionValueCalculator>,
68    ) -> Result<Self> {
69        let partition_type = partition_spec.partition_type(&iceberg_schema)?;
70
71        Ok(Self {
72            schema: iceberg_schema,
73            partition_spec,
74            calculator,
75            partition_type,
76        })
77    }
78
79    /// Create a new RecordBatchPartitionSplitter with computed partition values.
80    ///
81    /// This is a convenience method that creates a calculator and initializes the splitter
82    /// to compute partition values from source columns.
83    ///
84    /// # Arguments
85    ///
86    /// * `iceberg_schema` - The Iceberg schema reference
87    /// * `partition_spec` - The partition specification reference
88    ///
89    /// # Returns
90    ///
91    /// Returns a new `RecordBatchPartitionSplitter` instance or an error if initialization fails.
92    pub fn try_new_with_computed_values(
93        iceberg_schema: SchemaRef,
94        partition_spec: PartitionSpecRef,
95    ) -> Result<Self> {
96        let calculator = PartitionValueCalculator::try_new(&partition_spec, &iceberg_schema)?;
97        Self::try_new(iceberg_schema, partition_spec, Some(calculator))
98    }
99
100    /// Create a new RecordBatchPartitionSplitter expecting pre-computed partition values.
101    ///
102    /// This is a convenience method that initializes the splitter to expect a `_partition`
103    /// column in the input batches.
104    ///
105    /// # Arguments
106    ///
107    /// * `iceberg_schema` - The Iceberg schema reference
108    /// * `partition_spec` - The partition specification reference
109    ///
110    /// # Returns
111    ///
112    /// Returns a new `RecordBatchPartitionSplitter` instance or an error if initialization fails.
113    pub fn try_new_with_precomputed_values(
114        iceberg_schema: SchemaRef,
115        partition_spec: PartitionSpecRef,
116    ) -> Result<Self> {
117        Self::try_new(iceberg_schema, partition_spec, None)
118    }
119
120    /// Split the record batch into multiple record batches based on the partition spec.
121    pub fn split(&self, batch: &RecordBatch) -> Result<Vec<(PartitionKey, RecordBatch)>> {
122        let partition_structs = if let Some(calculator) = &self.calculator {
123            // Compute partition values from source columns using calculator
124            let partition_array = calculator.calculate(batch)?;
125            let struct_array = arrow_struct_to_literal(&partition_array, &self.partition_type)?;
126
127            struct_array
128                .into_iter()
129                .map(|s| {
130                    if let Some(Literal::Struct(s)) = s {
131                        Ok(s)
132                    } else {
133                        Err(invalid_data!(
134                            "Partition value is not a struct literal or is null"
135                        ))
136                    }
137                })
138                .collect::<Result<Vec<_>>>()?
139        } else {
140            // Extract partition values from pre-computed partition column
141            let partition_column = batch
142                .column_by_name(PROJECTED_PARTITION_VALUE_COLUMN)
143                .ok_or_else(|| {
144                    invalid_data!(
145                        "Partition column '{PROJECTED_PARTITION_VALUE_COLUMN}' not found in batch"
146                    )
147                })?;
148
149            let partition_struct_array = partition_column
150                .as_any()
151                .downcast_ref::<StructArray>()
152                .ok_or_else(|| invalid_data!("Partition column is not a StructArray"))?;
153
154            let arrow_struct_array = Arc::new(partition_struct_array.clone()) as ArrayRef;
155            let struct_array = arrow_struct_to_literal(&arrow_struct_array, &self.partition_type)?;
156
157            struct_array
158                .into_iter()
159                .map(|s| {
160                    if let Some(Literal::Struct(s)) = s {
161                        Ok(s)
162                    } else {
163                        Err(invalid_data!(
164                            "Partition value is not a struct literal or is null"
165                        ))
166                    }
167                })
168                .collect::<Result<Vec<_>>>()?
169        };
170
171        // Group the batch by row value.
172        let mut group_ids = HashMap::new();
173        partition_structs
174            .into_iter()
175            .enumerate()
176            .for_each(|(row_id, row)| {
177                group_ids.entry(row).or_insert(vec![]).push(row_id);
178            });
179
180        // Partition the batch with same partition partition_values
181        let mut partition_batches = Vec::with_capacity(group_ids.len());
182        for (row, row_ids) in group_ids.into_iter() {
183            // generate the bool filter array from column_ids
184            let filter_array: BooleanArray = {
185                let mut builder = BooleanBufferBuilder::new(batch.num_rows());
186                builder.append_n(batch.num_rows(), false);
187                for row_id in row_ids {
188                    builder.set_bit(row_id, true);
189                }
190                BooleanArray::new(builder.finish(), None)
191            };
192
193            // Create PartitionKey from the partition struct
194            let partition_key = PartitionKey::new(
195                self.partition_spec.as_ref().clone(),
196                self.schema.clone(),
197                row,
198            );
199
200            // filter the RecordBatch
201            partition_batches.push((partition_key, filter_record_batch(batch, &filter_array)?));
202        }
203
204        Ok(partition_batches)
205    }
206}
207
208#[cfg(test)]
209mod tests {
210    use std::sync::Arc;
211
212    use arrow_array::{Float64Array, Int32Array, RecordBatch, StringArray};
213    use arrow_schema::DataType;
214    use parquet::arrow::PARQUET_FIELD_ID_META_KEY;
215
216    use super::*;
217    use crate::arrow::schema_to_arrow_schema;
218    use crate::spec::{
219        NestedField, PartitionSpecBuilder, PrimitiveLiteral, Schema, Struct, Transform, Type,
220        UnboundPartitionField,
221    };
222
223    #[test]
224    fn test_record_batch_partition_split() {
225        let schema = Arc::new(
226            Schema::builder()
227                .with_fields(vec![
228                    NestedField::required(
229                        1,
230                        "id",
231                        Type::Primitive(crate::spec::PrimitiveType::Int),
232                    )
233                    .into(),
234                    NestedField::required(
235                        2,
236                        "name",
237                        Type::Primitive(crate::spec::PrimitiveType::String),
238                    )
239                    .into(),
240                ])
241                .build()
242                .unwrap(),
243        );
244        let partition_spec = Arc::new(
245            PartitionSpecBuilder::new(schema.clone())
246                .with_spec_id(1)
247                .add_unbound_field(
248                    UnboundPartitionField::builder()
249                        .source_ids(vec![1])
250                        .name("id_bucket".to_string())
251                        .transform(Transform::Identity)
252                        .build()
253                        .unwrap(),
254                )
255                .unwrap()
256                .build()
257                .unwrap(),
258        );
259        let partition_splitter = RecordBatchPartitionSplitter::try_new_with_computed_values(
260            schema.clone(),
261            partition_spec,
262        )
263        .expect("Failed to create splitter");
264
265        let arrow_schema = Arc::new(schema_to_arrow_schema(&schema).unwrap());
266        let id_array = Int32Array::from(vec![1, 2, 1, 3, 2, 3, 1]);
267        let data_array = StringArray::from(vec!["a", "b", "c", "d", "e", "f", "g"]);
268        let batch = RecordBatch::try_new(arrow_schema.clone(), vec![
269            Arc::new(id_array),
270            Arc::new(data_array),
271        ])
272        .expect("Failed to create RecordBatch");
273
274        let mut partitioned_batches = partition_splitter
275            .split(&batch)
276            .expect("Failed to split RecordBatch");
277        partitioned_batches.sort_by_key(|(partition_key, _)| {
278            if let PrimitiveLiteral::Int(i) = partition_key.data().fields()[0]
279                .as_ref()
280                .unwrap()
281                .as_primitive_literal()
282                .unwrap()
283            {
284                i
285            } else {
286                panic!("The partition value is not a int");
287            }
288        });
289        assert_eq!(partitioned_batches.len(), 3);
290        {
291            // check the first partition
292            let expected_id_array = Int32Array::from(vec![1, 1, 1]);
293            let expected_data_array = StringArray::from(vec!["a", "c", "g"]);
294            let expected_batch = RecordBatch::try_new(arrow_schema.clone(), vec![
295                Arc::new(expected_id_array),
296                Arc::new(expected_data_array),
297            ])
298            .expect("Failed to create expected RecordBatch");
299            assert_eq!(partitioned_batches[0].1, expected_batch);
300        }
301        {
302            // check the second partition
303            let expected_id_array = Int32Array::from(vec![2, 2]);
304            let expected_data_array = StringArray::from(vec!["b", "e"]);
305            let expected_batch = RecordBatch::try_new(arrow_schema.clone(), vec![
306                Arc::new(expected_id_array),
307                Arc::new(expected_data_array),
308            ])
309            .expect("Failed to create expected RecordBatch");
310            assert_eq!(partitioned_batches[1].1, expected_batch);
311        }
312        {
313            // check the third partition
314            let expected_id_array = Int32Array::from(vec![3, 3]);
315            let expected_data_array = StringArray::from(vec!["d", "f"]);
316            let expected_batch = RecordBatch::try_new(arrow_schema.clone(), vec![
317                Arc::new(expected_id_array),
318                Arc::new(expected_data_array),
319            ])
320            .expect("Failed to create expected RecordBatch");
321            assert_eq!(partitioned_batches[2].1, expected_batch);
322        }
323
324        let partition_values = partitioned_batches
325            .iter()
326            .map(|(partition_key, _)| partition_key.data().clone())
327            .collect::<Vec<_>>();
328        // check partition value is struct(1), struct(2), struct(3)
329        assert_eq!(partition_values, vec![
330            Struct::from_iter(vec![Some(Literal::int(1))]),
331            Struct::from_iter(vec![Some(Literal::int(2))]),
332            Struct::from_iter(vec![Some(Literal::int(3))]),
333        ]);
334    }
335
336    #[test]
337    fn test_record_batch_partition_split_with_partition_column() {
338        use arrow_array::StructArray;
339        use arrow_schema::{Field, Schema as ArrowSchema};
340
341        let schema = Arc::new(
342            Schema::builder()
343                .with_fields(vec![
344                    NestedField::required(
345                        1,
346                        "id",
347                        Type::Primitive(crate::spec::PrimitiveType::Int),
348                    )
349                    .into(),
350                    NestedField::required(
351                        2,
352                        "name",
353                        Type::Primitive(crate::spec::PrimitiveType::String),
354                    )
355                    .into(),
356                ])
357                .build()
358                .unwrap(),
359        );
360        let partition_spec = Arc::new(
361            PartitionSpecBuilder::new(schema.clone())
362                .with_spec_id(1)
363                .add_unbound_field(
364                    UnboundPartitionField::builder()
365                        .source_ids(vec![1])
366                        .name("id_bucket".to_string())
367                        .transform(Transform::Identity)
368                        .build()
369                        .unwrap(),
370                )
371                .unwrap()
372                .build()
373                .unwrap(),
374        );
375
376        // Create input schema with _partition column
377        // Note: partition field IDs start from 1000 by default
378        let partition_field = Field::new("id_bucket", DataType::Int32, false).with_metadata(
379            HashMap::from([(PARQUET_FIELD_ID_META_KEY.to_string(), "1000".to_string())]),
380        );
381        let partition_struct_field = Field::new(
382            PROJECTED_PARTITION_VALUE_COLUMN,
383            DataType::Struct(vec![partition_field.clone()].into()),
384            false,
385        );
386
387        let input_schema = Arc::new(ArrowSchema::new(vec![
388            Field::new("id", DataType::Int32, false),
389            Field::new("name", DataType::Utf8, false),
390            partition_struct_field,
391        ]));
392
393        // Create splitter expecting pre-computed partition column
394        let partition_splitter = RecordBatchPartitionSplitter::try_new_with_precomputed_values(
395            schema.clone(),
396            partition_spec,
397        )
398        .expect("Failed to create splitter");
399
400        // Create test data with pre-computed partition column
401        let id_array = Int32Array::from(vec![1, 2, 1, 3, 2, 3, 1]);
402        let data_array = StringArray::from(vec!["a", "b", "c", "d", "e", "f", "g"]);
403
404        // Create partition column (same values as id for Identity transform)
405        let partition_values = Int32Array::from(vec![1, 2, 1, 3, 2, 3, 1]);
406        let partition_struct = StructArray::from(vec![(
407            Arc::new(partition_field),
408            Arc::new(partition_values) as ArrayRef,
409        )]);
410
411        let batch = RecordBatch::try_new(input_schema.clone(), vec![
412            Arc::new(id_array),
413            Arc::new(data_array),
414            Arc::new(partition_struct),
415        ])
416        .expect("Failed to create RecordBatch");
417
418        // Split using the pre-computed partition column
419        let mut partitioned_batches = partition_splitter
420            .split(&batch)
421            .expect("Failed to split RecordBatch");
422
423        partitioned_batches.sort_by_key(|(partition_key, _)| {
424            if let PrimitiveLiteral::Int(i) = partition_key.data().fields()[0]
425                .as_ref()
426                .unwrap()
427                .as_primitive_literal()
428                .unwrap()
429            {
430                i
431            } else {
432                panic!("The partition value is not a int");
433            }
434        });
435
436        assert_eq!(partitioned_batches.len(), 3);
437
438        // Helper to extract id and name values from a batch
439        let extract_values = |batch: &RecordBatch| -> (Vec<i32>, Vec<String>) {
440            let id_col = batch
441                .column(0)
442                .as_any()
443                .downcast_ref::<Int32Array>()
444                .unwrap();
445            let name_col = batch
446                .column(1)
447                .as_any()
448                .downcast_ref::<StringArray>()
449                .unwrap();
450            (
451                id_col.values().to_vec(),
452                name_col.iter().map(|s| s.unwrap().to_string()).collect(),
453            )
454        };
455
456        // Verify partition 1: id=1, names=["a", "c", "g"]
457        let (key, batch) = &partitioned_batches[0];
458        assert_eq!(key.data(), &Struct::from_iter(vec![Some(Literal::int(1))]));
459        let (ids, names) = extract_values(batch);
460        assert_eq!(ids, vec![1, 1, 1]);
461        assert_eq!(names, vec!["a", "c", "g"]);
462
463        // Verify partition 2: id=2, names=["b", "e"]
464        let (key, batch) = &partitioned_batches[1];
465        assert_eq!(key.data(), &Struct::from_iter(vec![Some(Literal::int(2))]));
466        let (ids, names) = extract_values(batch);
467        assert_eq!(ids, vec![2, 2]);
468        assert_eq!(names, vec!["b", "e"]);
469
470        // Verify partition 3: id=3, names=["d", "f"]
471        let (key, batch) = &partitioned_batches[2];
472        assert_eq!(key.data(), &Struct::from_iter(vec![Some(Literal::int(3))]));
473        let (ids, names) = extract_values(batch);
474        assert_eq!(ids, vec![3, 3]);
475        assert_eq!(names, vec!["d", "f"]);
476    }
477
478    #[test]
479    fn test_record_batch_partition_split_signed_zero() {
480        let schema = Arc::new(
481            Schema::builder()
482                .with_fields(vec![
483                    NestedField::required(
484                        1,
485                        "id",
486                        Type::Primitive(crate::spec::PrimitiveType::Int),
487                    )
488                    .into(),
489                    NestedField::required(
490                        2,
491                        "d",
492                        Type::Primitive(crate::spec::PrimitiveType::Double),
493                    )
494                    .into(),
495                ])
496                .build()
497                .unwrap(),
498        );
499        let partition_spec = Arc::new(
500            PartitionSpecBuilder::new(schema.clone())
501                .with_spec_id(1)
502                .add_unbound_field(
503                    UnboundPartitionField::builder()
504                        .source_ids(vec![2])
505                        .name("d".to_string())
506                        .transform(Transform::Identity)
507                        .build()
508                        .unwrap(),
509                )
510                .unwrap()
511                .build()
512                .unwrap(),
513        );
514        let partition_splitter = RecordBatchPartitionSplitter::try_new_with_computed_values(
515            schema.clone(),
516            partition_spec,
517        )
518        .expect("Failed to create splitter");
519
520        let arrow_schema = Arc::new(schema_to_arrow_schema(&schema).unwrap());
521        let batch = RecordBatch::try_new(arrow_schema, vec![
522            Arc::new(Int32Array::from(vec![1, 2, 3])),
523            Arc::new(Float64Array::from(vec![-0.0, 0.0, -0.0])),
524        ])
525        .expect("Failed to create RecordBatch");
526
527        let mut partitions: Vec<(Struct, Vec<i32>)> = partition_splitter
528            .split(&batch)
529            .expect("Failed to split RecordBatch")
530            .into_iter()
531            .map(|(key, batch)| {
532                let ids = batch
533                    .column(0)
534                    .as_any()
535                    .downcast_ref::<Int32Array>()
536                    .unwrap()
537                    .values()
538                    .to_vec();
539                (key.data().clone(), ids)
540            })
541            .collect();
542        partitions.sort_by_key(|(_, ids)| ids[0]);
543
544        // -0.0 and 0.0 are different partition values, as in iceberg-java.
545        assert_eq!(partitions, vec![
546            (Struct::from_iter([Some(Literal::double(-0.0))]), vec![1, 3]),
547            (Struct::from_iter([Some(Literal::double(0.0))]), vec![2]),
548        ]);
549    }
550}