Skip to main content

iceberg/arrow/
partition_value_calculator.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
18//! Partition value calculation for Iceberg tables.
19//!
20//! This module provides utilities for calculating partition values from record batches
21//! based on a partition specification.
22
23use std::sync::Arc;
24
25use arrow_array::{ArrayRef, RecordBatch, StructArray};
26use arrow_schema::DataType;
27
28use super::record_batch_projector::RecordBatchProjector;
29use super::type_to_arrow_type;
30use crate::Result;
31use crate::error::invalid_data;
32use crate::spec::{PartitionSpec, Schema, StructType, Type};
33use crate::transform::{BoxedTransformFunction, create_transform_function};
34
35/// Calculator for partition values in Iceberg tables.
36///
37/// This struct handles the projection of source columns and application of
38/// partition transforms to compute partition values for a given record batch.
39#[derive(Debug)]
40pub struct PartitionValueCalculator {
41    projector: RecordBatchProjector,
42    transform_functions: Vec<BoxedTransformFunction>,
43    partition_type: StructType,
44    partition_arrow_type: DataType,
45}
46
47impl PartitionValueCalculator {
48    /// Create a new PartitionValueCalculator.
49    ///
50    /// # Arguments
51    ///
52    /// * `partition_spec` - The partition specification
53    /// * `table_schema` - The Iceberg table schema
54    ///
55    /// # Returns
56    ///
57    /// Returns a new `PartitionValueCalculator` instance or an error if initialization fails.
58    ///
59    /// # Errors
60    ///
61    /// Returns an error if:
62    /// - The partition spec is unpartitioned
63    /// - Transform function creation fails
64    /// - Projector initialization fails
65    pub fn try_new(partition_spec: &PartitionSpec, table_schema: &Schema) -> Result<Self> {
66        if partition_spec.is_unpartitioned() {
67            return Err(invalid_data!(
68                "Cannot create partition calculator for unpartitioned table"
69            ));
70        }
71
72        // Create transform functions for each partition field
73        let transform_functions: Vec<BoxedTransformFunction> = partition_spec
74            .fields()
75            .iter()
76            .map(|pf| create_transform_function(&pf.transform))
77            .collect::<Result<Vec<_>>>()?;
78
79        // Extract source field IDs for projection
80        let source_field_ids: Vec<i32> = partition_spec
81            .fields()
82            .iter()
83            .map(|pf| pf.source_id)
84            .collect();
85
86        // Create projector for extracting source columns
87        let projector = RecordBatchProjector::from_iceberg_schema(
88            Arc::new(table_schema.clone()),
89            &source_field_ids,
90        )?;
91
92        // Get partition type information
93        let partition_type = partition_spec.partition_type(table_schema)?;
94        let partition_arrow_type = type_to_arrow_type(&Type::Struct(partition_type.clone()))?;
95
96        Ok(Self {
97            projector,
98            transform_functions,
99            partition_type,
100            partition_arrow_type,
101        })
102    }
103
104    /// Get the partition type as an Iceberg StructType.
105    pub fn partition_type(&self) -> &StructType {
106        &self.partition_type
107    }
108
109    /// Get the partition type as an Arrow DataType.
110    pub fn partition_arrow_type(&self) -> &DataType {
111        &self.partition_arrow_type
112    }
113
114    /// Calculate partition values for a record batch.
115    ///
116    /// This method:
117    /// 1. Projects the source columns from the batch
118    /// 2. Applies partition transforms to each source column
119    /// 3. Constructs a StructArray containing the partition values
120    ///
121    /// # Arguments
122    ///
123    /// * `batch` - The record batch to calculate partition values for
124    ///
125    /// # Returns
126    ///
127    /// Returns an ArrayRef containing a StructArray of partition values, or an error if calculation fails.
128    ///
129    /// # Errors
130    ///
131    /// Returns an error if:
132    /// - Column projection fails
133    /// - Transform application fails
134    /// - StructArray construction fails
135    pub fn calculate(&self, batch: &RecordBatch) -> Result<ArrayRef> {
136        // Project source columns from the batch
137        let source_columns = self.projector.project_column(batch.columns())?;
138
139        // Get expected struct fields for the result
140        let expected_struct_fields = match &self.partition_arrow_type {
141            DataType::Struct(fields) => fields.clone(),
142            _ => {
143                return Err(invalid_data!("Expected partition type must be a struct"));
144            }
145        };
146
147        // Apply transforms to each source column
148        let mut partition_values = Vec::with_capacity(self.transform_functions.len());
149        for (source_column, transform_fn) in source_columns.iter().zip(&self.transform_functions) {
150            let partition_value = transform_fn.transform(source_column.clone())?;
151            partition_values.push(partition_value);
152        }
153
154        // Construct the StructArray
155        let struct_array = StructArray::try_new(expected_struct_fields, partition_values, None)
156            .map_err(|e| invalid_data!("Failed to create partition struct array: {e}"))?;
157
158        Ok(Arc::new(struct_array))
159    }
160}
161
162#[cfg(test)]
163mod tests {
164    use std::sync::Arc;
165
166    use arrow_array::{Int32Array, RecordBatch, StringArray};
167    use arrow_schema::{Field, Schema as ArrowSchema};
168
169    use super::*;
170    use crate::spec::{NestedField, PartitionSpecBuilder, PrimitiveType, Transform};
171
172    #[test]
173    fn test_partition_calculator_identity_transform() {
174        let table_schema = Schema::builder()
175            .with_schema_id(0)
176            .with_fields(vec![
177                NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
178                NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)).into(),
179            ])
180            .build()
181            .unwrap();
182
183        let partition_spec = PartitionSpecBuilder::new(Arc::new(table_schema.clone()))
184            .add_partition_field("id", "id_partition", Transform::Identity)
185            .unwrap()
186            .build()
187            .unwrap();
188
189        let calculator = PartitionValueCalculator::try_new(&partition_spec, &table_schema).unwrap();
190
191        // Verify partition type
192        assert_eq!(calculator.partition_type().fields().len(), 1);
193        assert_eq!(calculator.partition_type().fields()[0].name, "id_partition");
194
195        // Create test batch
196        let arrow_schema = Arc::new(ArrowSchema::new(vec![
197            Field::new("id", DataType::Int32, false),
198            Field::new("name", DataType::Utf8, false),
199        ]));
200
201        let batch = RecordBatch::try_new(arrow_schema, vec![
202            Arc::new(Int32Array::from(vec![10, 20, 30])),
203            Arc::new(StringArray::from(vec!["a", "b", "c"])),
204        ])
205        .unwrap();
206
207        // Calculate partition values
208        let result = calculator.calculate(&batch).unwrap();
209        let struct_array = result.as_any().downcast_ref::<StructArray>().unwrap();
210
211        let id_partition = struct_array
212            .column_by_name("id_partition")
213            .unwrap()
214            .as_any()
215            .downcast_ref::<Int32Array>()
216            .unwrap();
217
218        assert_eq!(id_partition.value(0), 10);
219        assert_eq!(id_partition.value(1), 20);
220        assert_eq!(id_partition.value(2), 30);
221    }
222
223    #[test]
224    fn test_partition_calculator_unpartitioned_error() {
225        let table_schema = Schema::builder()
226            .with_schema_id(0)
227            .with_fields(vec![
228                NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
229            ])
230            .build()
231            .unwrap();
232
233        let partition_spec = PartitionSpecBuilder::new(Arc::new(table_schema.clone()))
234            .build()
235            .unwrap();
236
237        let result = PartitionValueCalculator::try_new(&partition_spec, &table_schema);
238        assert!(result.is_err());
239        assert!(
240            result
241                .unwrap_err()
242                .to_string()
243                .contains("unpartitioned table")
244        );
245    }
246}