1use 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
31pub const PROJECTED_PARTITION_VALUE_COLUMN: &str = "_partition";
33
34pub struct RecordBatchPartitionSplitter {
44 schema: SchemaRef,
45 partition_spec: PartitionSpecRef,
46 calculator: Option<PartitionValueCalculator>,
47 partition_type: StructType,
48}
49
50impl RecordBatchPartitionSplitter {
51 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 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 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 pub fn split(&self, batch: &RecordBatch) -> Result<Vec<(PartitionKey, RecordBatch)>> {
122 let partition_structs = if let Some(calculator) = &self.calculator {
123 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 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 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 let mut partition_batches = Vec::with_capacity(group_ids.len());
182 for (row, row_ids) in group_ids.into_iter() {
183 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 let partition_key = PartitionKey::new(
195 self.partition_spec.as_ref().clone(),
196 self.schema.clone(),
197 row,
198 );
199
200 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 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 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 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 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 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 let partition_splitter = RecordBatchPartitionSplitter::try_new_with_precomputed_values(
395 schema.clone(),
396 partition_spec,
397 )
398 .expect("Failed to create splitter");
399
400 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 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 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 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 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 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 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 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}