1use std::cmp::Reverse;
21use std::collections::{HashMap, HashSet};
22
23use crate::error::invalid_data;
24use crate::spec::{
25 NestedField, NestedFieldRef, PartitionField, PartitionSpec, Schema, StructType, Transform, Type,
26};
27use crate::{Error, ErrorKind, Result};
28
29pub fn compute_unified_partition_type<'a>(
51 partition_specs: impl Iterator<Item = &'a PartitionSpec>,
52 schema: &Schema,
53) -> Result<StructType> {
54 let mut specs: Vec<&PartitionSpec> = partition_specs.collect();
55 specs.sort_by_key(|s| Reverse(s.spec_id()));
56
57 let active_field_ids = all_active_field_ids(specs.iter().copied(), schema);
58
59 let mut field_map: HashMap<i32, &PartitionField> = HashMap::new();
60 let mut type_map: HashMap<i32, Type> = HashMap::new();
61 let mut name_map: HashMap<i32, String> = HashMap::new();
62
63 for spec in &specs {
64 for field in spec.fields() {
65 let field_id = field.field_id;
66
67 if matches!(field.transform, Transform::Unknown) {
72 return Err(invalid_data!(
73 "Partition field '{}' uses an unknown transform whose result type \
74 cannot be determined",
75 field.name
76 ));
77 }
78
79 if !active_field_ids.contains(&field_id) {
80 continue;
81 }
82
83 let source_field = match schema.field_by_id(field.source_id) {
84 Some(f) => f,
85 None => continue,
86 };
87
88 match field_map.get(&field_id) {
89 None => {
90 let res_type = field.transform.result_type(&source_field.field_type)?;
91 field_map.insert(field_id, field);
92 type_map.insert(field_id, res_type);
93 name_map.insert(field_id, field.name.clone());
94 }
95 Some(existing) => {
96 if !equivalent_ignoring_names(field, existing) {
99 return Err(invalid_data!(
100 "Conflicting partition fields for field id {field_id}: \
101 '{}' and '{}'",
102 field.name,
103 existing.name
104 ));
105 }
106
107 if is_void_transform(existing) && !is_void_transform(field) {
111 let res_type = field.transform.result_type(&source_field.field_type)?;
112 field_map.insert(field_id, field);
113 type_map.insert(field_id, res_type);
114 }
115 }
116 }
117 }
118 }
119
120 let mut field_ids: Vec<i32> = field_map.keys().copied().collect();
121 field_ids.sort();
122
123 let struct_fields = field_ids
124 .into_iter()
125 .map(|fid| -> Result<NestedFieldRef> {
126 let name = name_map.get(&fid).ok_or_else(|| {
127 Error::new(
128 ErrorKind::Unexpected,
129 format!("Missing name for partition field {fid}"),
130 )
131 })?;
132 let ty = type_map.remove(&fid).ok_or_else(|| {
133 Error::new(
134 ErrorKind::Unexpected,
135 format!("Missing type for partition field {fid}"),
136 )
137 })?;
138 Ok(NestedField::optional(fid, name, ty).into())
139 })
140 .collect::<Result<Vec<_>>>()?;
141
142 Ok(StructType::new(struct_fields))
143}
144
145fn is_void_transform(field: &PartitionField) -> bool {
146 matches!(field.transform, Transform::Void)
147}
148
149fn equivalent_ignoring_names(field: &PartitionField, other: &PartitionField) -> bool {
153 field.field_id == other.field_id
154 && field.source_id == other.source_id
155 && compatible_transforms(&field.transform, &other.transform)
156}
157
158fn compatible_transforms(t1: &Transform, t2: &Transform) -> bool {
161 t1 == t2 || matches!(t1, Transform::Void) || matches!(t2, Transform::Void)
162}
163
164fn all_active_field_ids<'a>(
165 partition_specs: impl Iterator<Item = &'a PartitionSpec>,
166 schema: &Schema,
167) -> HashSet<i32> {
168 partition_specs
169 .flat_map(|spec| spec.fields().iter())
170 .filter(|field| schema.field_by_id(field.source_id).is_some())
171 .map(|field| field.field_id)
172 .collect()
173}
174
175#[cfg(test)]
176mod tests {
177 use std::sync::Arc;
178
179 use super::*;
180 use crate::spec::{
181 NestedField, PrimitiveType, Transform, Type, UnboundPartitionField, UnboundPartitionSpec,
182 };
183
184 fn test_schema() -> Schema {
185 Schema::builder()
186 .with_fields(vec![
187 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
188 NestedField::required(2, "data", Type::Primitive(PrimitiveType::String)).into(),
189 NestedField::required(3, "ts", Type::Primitive(PrimitiveType::Timestamp)).into(),
190 NestedField::required(4, "category", Type::Primitive(PrimitiveType::String)).into(),
191 ])
192 .build()
193 .unwrap()
194 }
195
196 fn build_spec(
197 schema: &Schema,
198 spec_id: i32,
199 fields: Vec<(i32, &str, Transform)>,
200 ) -> PartitionSpec {
201 let mut builder = UnboundPartitionSpec::builder().with_spec_id(spec_id);
202 for (source_id, name, transform) in fields {
203 builder = builder
204 .add_partition_field(
205 UnboundPartitionField::builder()
206 .source_ids(vec![source_id])
207 .name(name)
208 .transform(transform)
209 .build()
210 .unwrap(),
211 )
212 .unwrap();
213 }
214 builder.build().bind(schema.clone()).unwrap()
215 }
216
217 #[test]
218 fn test_single_spec_identity() {
219 let schema = test_schema();
220 let spec = build_spec(&schema, 0, vec![(4, "category", Transform::Identity)]);
221
222 let result = compute_unified_partition_type([&spec].into_iter(), &schema).unwrap();
223 assert_eq!(result.fields().len(), 1);
224 assert_eq!(result.fields()[0].name, "category");
225 assert_eq!(
226 *result.fields()[0].field_type,
227 Type::Primitive(PrimitiveType::String)
228 );
229 }
230
231 #[test]
232 fn test_single_spec_with_year_transform() {
233 let schema = test_schema();
234 let spec = build_spec(&schema, 0, vec![(3, "ts_year", Transform::Year)]);
235
236 let result = compute_unified_partition_type([&spec].into_iter(), &schema).unwrap();
237 assert_eq!(result.fields().len(), 1);
238 assert_eq!(result.fields()[0].name, "ts_year");
239 assert_eq!(
240 *result.fields()[0].field_type,
241 Type::Primitive(PrimitiveType::Int)
242 );
243 }
244
245 #[test]
246 fn test_unpartitioned() {
247 let schema = test_schema();
248 let spec = PartitionSpec::unpartition_spec();
249 let result = compute_unified_partition_type([&spec].into_iter(), &schema).unwrap();
250 assert!(result.fields().is_empty());
251 }
252
253 #[test]
254 fn test_multiple_fields_sorted_by_id() {
255 let schema = test_schema();
256 let spec = build_spec(&schema, 0, vec![
257 (3, "ts_year", Transform::Year),
258 (4, "category", Transform::Identity),
259 ]);
260
261 let result = compute_unified_partition_type([&spec].into_iter(), &schema).unwrap();
262 assert_eq!(result.fields().len(), 2);
263 assert!(result.fields()[0].id < result.fields()[1].id);
264 }
265
266 #[test]
267 fn test_newer_name_takes_precedence() {
268 let schema = test_schema();
269
270 let spec_v0 = PartitionSpec::builder(Arc::new(schema.clone()))
272 .with_spec_id(0)
273 .add_unbound_field(
274 UnboundPartitionField::builder()
275 .source_ids(vec![4])
276 .field_id(1000)
277 .name("cat_old".to_string())
278 .transform(Transform::Identity)
279 .build()
280 .unwrap(),
281 )
282 .unwrap()
283 .build()
284 .unwrap();
285
286 let spec_v1 = PartitionSpec::builder(Arc::new(schema.clone()))
288 .with_spec_id(1)
289 .add_unbound_field(
290 UnboundPartitionField::builder()
291 .source_ids(vec![4])
292 .field_id(1000)
293 .name("cat_new".to_string())
294 .transform(Transform::Identity)
295 .build()
296 .unwrap(),
297 )
298 .unwrap()
299 .build()
300 .unwrap();
301
302 let result =
303 compute_unified_partition_type([&spec_v0, &spec_v1].into_iter(), &schema).unwrap();
304 assert_eq!(result.fields().len(), 1);
305 assert_eq!(result.fields()[0].name, "cat_new");
306 }
307
308 #[test]
309 fn test_void_replaced_by_older_non_void() {
310 let schema = test_schema();
311
312 let spec_v0 = PartitionSpec::builder(Arc::new(schema.clone()))
314 .with_spec_id(0)
315 .add_unbound_field(
316 UnboundPartitionField::builder()
317 .source_ids(vec![4])
318 .field_id(1000)
319 .name("category".to_string())
320 .transform(Transform::Identity)
321 .build()
322 .unwrap(),
323 )
324 .unwrap()
325 .build()
326 .unwrap();
327
328 let spec_v1 = PartitionSpec::builder(Arc::new(schema.clone()))
330 .with_spec_id(1)
331 .add_unbound_field(
332 UnboundPartitionField::builder()
333 .source_ids(vec![4])
334 .field_id(1000)
335 .name("category_v2".to_string())
336 .transform(Transform::Void)
337 .build()
338 .unwrap(),
339 )
340 .unwrap()
341 .build()
342 .unwrap();
343
344 let result =
345 compute_unified_partition_type([&spec_v0, &spec_v1].into_iter(), &schema).unwrap();
346
347 assert_eq!(result.fields().len(), 1);
348 assert_eq!(result.fields()[0].name, "category_v2");
350 assert_eq!(
352 *result.fields()[0].field_type,
353 Type::Primitive(PrimitiveType::String)
354 );
355 }
356
357 #[test]
358 fn test_dropped_source_column_skipped() {
359 let schema = Schema::builder()
361 .with_fields(vec![
362 NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)).into(),
363 NestedField::required(2, "data", Type::Primitive(PrimitiveType::String)).into(),
364 ])
365 .build()
366 .unwrap();
367
368 let spec = serde_json::from_value::<PartitionSpec>(serde_json::json!({
371 "spec-id": 0,
372 "fields": [{
373 "source-id": 4,
374 "field-id": 1000,
375 "name": "category",
376 "transform": "identity"
377 }]
378 }))
379 .unwrap();
380
381 let result = compute_unified_partition_type([&spec].into_iter(), &schema).unwrap();
382 assert!(result.fields().is_empty());
383 }
384
385 #[test]
386 fn test_evolution_adds_new_field() {
387 let schema = test_schema();
388
389 let spec_v0 = build_spec(&schema, 0, vec![(4, "category", Transform::Identity)]);
391
392 let spec_v1 = PartitionSpec::builder(Arc::new(schema.clone()))
394 .with_spec_id(1)
395 .add_unbound_field(
396 UnboundPartitionField::builder()
397 .source_ids(vec![4])
398 .field_id(spec_v0.fields()[0].field_id)
399 .name("category".to_string())
400 .transform(Transform::Identity)
401 .build()
402 .unwrap(),
403 )
404 .unwrap()
405 .add_partition_field("ts", "ts_year", Transform::Year)
406 .unwrap()
407 .build()
408 .unwrap();
409
410 let result =
411 compute_unified_partition_type([&spec_v0, &spec_v1].into_iter(), &schema).unwrap();
412 assert_eq!(result.fields().len(), 2);
413 }
414
415 #[test]
416 fn test_unknown_transform_errors() {
417 let schema = test_schema();
418
419 let spec = serde_json::from_value::<PartitionSpec>(serde_json::json!({
422 "spec-id": 0,
423 "fields": [{
424 "source-id": 4,
425 "field-id": 1000,
426 "name": "category",
427 "transform": "unknown"
428 }]
429 }))
430 .unwrap();
431
432 let err = compute_unified_partition_type([&spec].into_iter(), &schema).unwrap_err();
433 assert_eq!(err.kind(), ErrorKind::DataInvalid);
434 }
435
436 #[test]
437 fn test_conflicting_partition_fields_error() {
438 let schema = test_schema();
439
440 let spec_v0 = serde_json::from_value::<PartitionSpec>(serde_json::json!({
442 "spec-id": 0,
443 "fields": [{
444 "source-id": 4,
445 "field-id": 1000,
446 "name": "category",
447 "transform": "identity"
448 }]
449 }))
450 .unwrap();
451
452 let spec_v1 = serde_json::from_value::<PartitionSpec>(serde_json::json!({
455 "spec-id": 1,
456 "fields": [{
457 "source-id": 3,
458 "field-id": 1000,
459 "name": "ts_year",
460 "transform": "year"
461 }]
462 }))
463 .unwrap();
464
465 let err =
466 compute_unified_partition_type([&spec_v0, &spec_v1].into_iter(), &schema).unwrap_err();
467 assert_eq!(err.kind(), ErrorKind::DataInvalid);
468 }
469}