iceberg/arrow/
partition_value_calculator.rs1use 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#[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 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 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 let source_field_ids: Vec<i32> = partition_spec
81 .fields()
82 .iter()
83 .map(|pf| pf.source_id)
84 .collect();
85
86 let projector = RecordBatchProjector::from_iceberg_schema(
88 Arc::new(table_schema.clone()),
89 &source_field_ids,
90 )?;
91
92 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 pub fn partition_type(&self) -> &StructType {
106 &self.partition_type
107 }
108
109 pub fn partition_arrow_type(&self) -> &DataType {
111 &self.partition_arrow_type
112 }
113
114 pub fn calculate(&self, batch: &RecordBatch) -> Result<ArrayRef> {
136 let source_columns = self.projector.project_column(batch.columns())?;
138
139 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 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 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 assert_eq!(calculator.partition_type().fields().len(), 1);
193 assert_eq!(calculator.partition_type().fields()[0].name, "id_partition");
194
195 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 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}