Skip to main content

iceberg/spec/schema/
prune_columns.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 super::*;
19use crate::error::invalid_data;
20use crate::spec::VariantType;
21
22struct PruneColumn {
23    selected: HashSet<i32>,
24    select_full_types: bool,
25}
26
27/// Visit a schema and returns only the fields selected by id set
28pub fn prune_columns(
29    schema: &Schema,
30    selected: impl IntoIterator<Item = i32>,
31    select_full_types: bool,
32) -> Result<Type> {
33    let mut visitor = PruneColumn::new(HashSet::from_iter(selected), select_full_types);
34    let result = visit_schema(schema, &mut visitor);
35
36    match result {
37        Ok(s) => {
38            if let Some(struct_type) = s {
39                Ok(struct_type)
40            } else {
41                Ok(Type::Struct(StructType::default()))
42            }
43        }
44        Err(e) => Err(e),
45    }
46}
47
48impl PruneColumn {
49    fn new(selected: HashSet<i32>, select_full_types: bool) -> Self {
50        Self {
51            selected,
52            select_full_types,
53        }
54    }
55
56    fn project_selected_struct(projected_field: Option<Type>) -> Result<StructType> {
57        match projected_field {
58            // If the field is a StructType, return it as such
59            Some(Type::Struct(s)) => Ok(s),
60            Some(_) => Err(Error::new(
61                ErrorKind::Unexpected,
62                "Projected field with struct type must be struct".to_string(),
63            )),
64            // If projected_field is None or not a StructType, return an empty StructType
65            None => Ok(StructType::default()),
66        }
67    }
68    fn project_list(list: &ListType, element_result: Type) -> Result<ListType> {
69        if *list.element_field.field_type == element_result {
70            return Ok(list.clone());
71        }
72        Ok(ListType {
73            element_field: Arc::new(NestedField {
74                id: list.element_field.id,
75                name: list.element_field.name.clone(),
76                required: list.element_field.required,
77                field_type: Box::new(element_result),
78                doc: list.element_field.doc.clone(),
79                initial_default: list.element_field.initial_default.clone(),
80                write_default: list.element_field.write_default.clone(),
81            }),
82        })
83    }
84    fn project_map(map: &MapType, value_result: Type) -> Result<MapType> {
85        if *map.value_field.field_type == value_result {
86            return Ok(map.clone());
87        }
88        Ok(MapType {
89            key_field: map.key_field.clone(),
90            value_field: Arc::new(NestedField {
91                id: map.value_field.id,
92                name: map.value_field.name.clone(),
93                required: map.value_field.required,
94                field_type: Box::new(value_result),
95                doc: map.value_field.doc.clone(),
96                initial_default: map.value_field.initial_default.clone(),
97                write_default: map.value_field.write_default.clone(),
98            }),
99        })
100    }
101}
102
103impl SchemaVisitor for PruneColumn {
104    type T = Option<Type>;
105
106    fn schema(&mut self, _schema: &Schema, value: Option<Type>) -> Result<Option<Type>> {
107        Ok(Some(value.unwrap()))
108    }
109
110    fn field(&mut self, field: &NestedFieldRef, value: Option<Type>) -> Result<Option<Type>> {
111        if self.selected.contains(&field.id) {
112            if self.select_full_types {
113                Ok(Some(*field.field_type.clone()))
114            } else if field.field_type.is_struct() {
115                Ok(Some(Type::Struct(PruneColumn::project_selected_struct(
116                    value,
117                )?)))
118            } else if !field.field_type.is_nested() {
119                Ok(Some(*field.field_type.clone()))
120            } else {
121                Err(invalid_data!(
122                    "Can't project list or map field directly when not selecting full type."
123                )
124                .with_context("field_id", field.id.to_string())
125                .with_context("field_type", field.field_type.to_string()))
126            }
127        } else {
128            Ok(value)
129        }
130    }
131
132    fn r#struct(
133        &mut self,
134        r#struct: &StructType,
135        results: Vec<Option<Type>>,
136    ) -> Result<Option<Type>> {
137        let fields = r#struct.fields();
138        let mut selected_field = Vec::with_capacity(fields.len());
139        let mut same_type = true;
140
141        for (field, projected_type) in zip_eq(fields.iter(), results.iter()) {
142            if let Some(projected_type) = projected_type {
143                if *field.field_type == *projected_type {
144                    selected_field.push(field.clone());
145                } else {
146                    same_type = false;
147                    let new_field = NestedField {
148                        id: field.id,
149                        name: field.name.clone(),
150                        required: field.required,
151                        field_type: Box::new(projected_type.clone()),
152                        doc: field.doc.clone(),
153                        initial_default: field.initial_default.clone(),
154                        write_default: field.write_default.clone(),
155                    };
156                    selected_field.push(Arc::new(new_field));
157                }
158            }
159        }
160
161        if !selected_field.is_empty() {
162            if selected_field.len() == fields.len() && same_type {
163                return Ok(Some(Type::Struct(r#struct.clone())));
164            } else {
165                return Ok(Some(Type::Struct(StructType::new(selected_field))));
166            }
167        }
168        Ok(None)
169    }
170
171    fn list(&mut self, list: &ListType, value: Option<Type>) -> Result<Option<Type>> {
172        if self.selected.contains(&list.element_field.id) {
173            if self.select_full_types {
174                Ok(Some(Type::List(list.clone())))
175            } else if list.element_field.field_type.is_struct() {
176                let projected_struct = PruneColumn::project_selected_struct(value).unwrap();
177                Ok(Some(Type::List(PruneColumn::project_list(
178                    list,
179                    Type::Struct(projected_struct),
180                )?)))
181            } else if list.element_field.field_type.is_primitive() {
182                Ok(Some(Type::List(list.clone())))
183            } else {
184                Err(invalid_data!(
185                    "Cannot explicitly project List or Map types, List element {} of type {} was selected",
186                    list.element_field.id,
187                    list.element_field.field_type
188                ))
189            }
190        } else if let Some(result) = value {
191            Ok(Some(Type::List(PruneColumn::project_list(list, result)?)))
192        } else {
193            Ok(None)
194        }
195    }
196
197    fn map(
198        &mut self,
199        map: &MapType,
200        _key_value: Option<Type>,
201        value: Option<Type>,
202    ) -> Result<Option<Type>> {
203        if self.selected.contains(&map.value_field.id) {
204            if self.select_full_types {
205                Ok(Some(Type::Map(map.clone())))
206            } else if map.value_field.field_type.is_struct() {
207                let projected_struct =
208                    PruneColumn::project_selected_struct(Some(value.unwrap())).unwrap();
209                Ok(Some(Type::Map(PruneColumn::project_map(
210                    map,
211                    Type::Struct(projected_struct),
212                )?)))
213            } else if map.value_field.field_type.is_primitive() {
214                Ok(Some(Type::Map(map.clone())))
215            } else {
216                Err(invalid_data!(
217                    "Cannot explicitly project List or Map types, Map value {} of type {} was selected",
218                    map.value_field.id,
219                    map.value_field.field_type
220                ))
221            }
222        } else if let Some(value_result) = value {
223            Ok(Some(Type::Map(PruneColumn::project_map(
224                map,
225                value_result,
226            )?)))
227        } else if self.selected.contains(&map.key_field.id) {
228            Ok(Some(Type::Map(map.clone())))
229        } else {
230            Ok(None)
231        }
232    }
233
234    fn primitive(&mut self, _p: &PrimitiveType) -> Result<Option<Type>> {
235        Ok(None)
236    }
237
238    fn variant(&mut self, _v: &VariantType) -> Result<Self::T> {
239        Ok(None)
240    }
241}
242
243#[cfg(test)]
244mod tests {
245    use Type::Primitive;
246
247    use super::*;
248    use crate::spec::schema::tests::table_schema_nested;
249
250    #[test]
251    fn test_schema_prune_columns_string() {
252        let expected_type = Type::from(
253            Schema::builder()
254                .with_fields(vec![
255                    NestedField::optional(1, "foo", Primitive(PrimitiveType::String)).into(),
256                ])
257                .build()
258                .unwrap()
259                .as_struct()
260                .clone(),
261        );
262        let schema = table_schema_nested();
263        let selected: HashSet<i32> = HashSet::from([1]);
264        let result = prune_columns(&schema, selected, false);
265        assert!(result.is_ok());
266        assert_eq!(result.unwrap(), expected_type);
267    }
268
269    #[test]
270    fn test_schema_prune_columns_string_full() {
271        let expected_type = Type::from(
272            Schema::builder()
273                .with_fields(vec![
274                    NestedField::optional(1, "foo", Primitive(PrimitiveType::String)).into(),
275                ])
276                .build()
277                .unwrap()
278                .as_struct()
279                .clone(),
280        );
281        let schema = table_schema_nested();
282        let selected: HashSet<i32> = HashSet::from([1]);
283        let result = prune_columns(&schema, selected, true);
284        assert!(result.is_ok());
285        assert_eq!(result.unwrap(), expected_type);
286    }
287
288    #[test]
289    fn test_schema_prune_columns_list() {
290        let expected_type = Type::from(
291            Schema::builder()
292                .with_fields(vec![
293                    NestedField::required(
294                        4,
295                        "qux",
296                        Type::List(ListType {
297                            element_field: NestedField::list_element(
298                                5,
299                                Primitive(PrimitiveType::String),
300                                true,
301                            )
302                            .into(),
303                        }),
304                    )
305                    .into(),
306                ])
307                .build()
308                .unwrap()
309                .as_struct()
310                .clone(),
311        );
312        let schema = table_schema_nested();
313        let selected: HashSet<i32> = HashSet::from([5]);
314        let result = prune_columns(&schema, selected, false);
315        assert!(result.is_ok());
316        assert_eq!(result.unwrap(), expected_type);
317    }
318
319    #[test]
320    fn test_prune_columns_list_itself() {
321        let schema = table_schema_nested();
322        let selected: HashSet<i32> = HashSet::from([4]);
323        let result = prune_columns(&schema, selected, false);
324        assert!(result.is_err());
325    }
326
327    #[test]
328    fn test_schema_prune_columns_list_full() {
329        let expected_type = Type::from(
330            Schema::builder()
331                .with_fields(vec![
332                    NestedField::required(
333                        4,
334                        "qux",
335                        Type::List(ListType {
336                            element_field: NestedField::list_element(
337                                5,
338                                Primitive(PrimitiveType::String),
339                                true,
340                            )
341                            .into(),
342                        }),
343                    )
344                    .into(),
345                ])
346                .build()
347                .unwrap()
348                .as_struct()
349                .clone(),
350        );
351        let schema = table_schema_nested();
352        let selected: HashSet<i32> = HashSet::from([5]);
353        let result = prune_columns(&schema, selected, true);
354        assert!(result.is_ok());
355        assert_eq!(result.unwrap(), expected_type);
356    }
357
358    #[test]
359    fn test_prune_columns_map() {
360        let expected_type = Type::from(
361            Schema::builder()
362                .with_fields(vec![
363                    NestedField::required(
364                        6,
365                        "quux",
366                        Type::Map(MapType {
367                            key_field: NestedField::map_key_element(
368                                7,
369                                Primitive(PrimitiveType::String),
370                            )
371                            .into(),
372                            value_field: NestedField::map_value_element(
373                                8,
374                                Type::Map(MapType {
375                                    key_field: NestedField::map_key_element(
376                                        9,
377                                        Primitive(PrimitiveType::String),
378                                    )
379                                    .into(),
380                                    value_field: NestedField::map_value_element(
381                                        10,
382                                        Primitive(PrimitiveType::Int),
383                                        true,
384                                    )
385                                    .into(),
386                                }),
387                                true,
388                            )
389                            .into(),
390                        }),
391                    )
392                    .into(),
393                ])
394                .build()
395                .unwrap()
396                .as_struct()
397                .clone(),
398        );
399        let schema = table_schema_nested();
400        let selected: HashSet<i32> = HashSet::from([9]);
401        let result = prune_columns(&schema, selected, false);
402        assert!(result.is_ok());
403        assert_eq!(result.unwrap(), expected_type);
404    }
405
406    #[test]
407    fn test_prune_columns_map_itself() {
408        let schema = table_schema_nested();
409        let selected: HashSet<i32> = HashSet::from([6]);
410        let result = prune_columns(&schema, selected, false);
411        assert!(result.is_err());
412    }
413
414    #[test]
415    fn test_prune_columns_map_full() {
416        let expected_type = Type::from(
417            Schema::builder()
418                .with_fields(vec![
419                    NestedField::required(
420                        6,
421                        "quux",
422                        Type::Map(MapType {
423                            key_field: NestedField::map_key_element(
424                                7,
425                                Primitive(PrimitiveType::String),
426                            )
427                            .into(),
428                            value_field: NestedField::map_value_element(
429                                8,
430                                Type::Map(MapType {
431                                    key_field: NestedField::map_key_element(
432                                        9,
433                                        Primitive(PrimitiveType::String),
434                                    )
435                                    .into(),
436                                    value_field: NestedField::map_value_element(
437                                        10,
438                                        Primitive(PrimitiveType::Int),
439                                        true,
440                                    )
441                                    .into(),
442                                }),
443                                true,
444                            )
445                            .into(),
446                        }),
447                    )
448                    .into(),
449                ])
450                .build()
451                .unwrap()
452                .as_struct()
453                .clone(),
454        );
455        let schema = table_schema_nested();
456        let selected: HashSet<i32> = HashSet::from([9]);
457        let result = prune_columns(&schema, selected, true);
458        assert!(result.is_ok());
459        assert_eq!(result.unwrap(), expected_type);
460    }
461
462    #[test]
463    fn test_prune_columns_map_key() {
464        let expected_type = Type::from(
465            Schema::builder()
466                .with_fields(vec![
467                    NestedField::required(
468                        6,
469                        "quux",
470                        Type::Map(MapType {
471                            key_field: NestedField::map_key_element(
472                                7,
473                                Primitive(PrimitiveType::String),
474                            )
475                            .into(),
476                            value_field: NestedField::map_value_element(
477                                8,
478                                Type::Map(MapType {
479                                    key_field: NestedField::map_key_element(
480                                        9,
481                                        Primitive(PrimitiveType::String),
482                                    )
483                                    .into(),
484                                    value_field: NestedField::map_value_element(
485                                        10,
486                                        Primitive(PrimitiveType::Int),
487                                        true,
488                                    )
489                                    .into(),
490                                }),
491                                true,
492                            )
493                            .into(),
494                        }),
495                    )
496                    .into(),
497                ])
498                .build()
499                .unwrap()
500                .as_struct()
501                .clone(),
502        );
503        let schema = table_schema_nested();
504        let selected: HashSet<i32> = HashSet::from([10]);
505        let result = prune_columns(&schema, selected, false);
506        assert!(result.is_ok());
507        assert_eq!(result.unwrap(), expected_type);
508    }
509
510    #[test]
511    fn test_prune_columns_struct() {
512        let expected_type = Type::from(
513            Schema::builder()
514                .with_fields(vec![
515                    NestedField::optional(
516                        15,
517                        "person",
518                        Type::Struct(StructType::new(vec![
519                            NestedField::optional(16, "name", Primitive(PrimitiveType::String))
520                                .into(),
521                        ])),
522                    )
523                    .into(),
524                ])
525                .build()
526                .unwrap()
527                .as_struct()
528                .clone(),
529        );
530        let schema = table_schema_nested();
531        let selected: HashSet<i32> = HashSet::from([16]);
532        let result = prune_columns(&schema, selected, false);
533        assert!(result.is_ok());
534        assert_eq!(result.unwrap(), expected_type);
535    }
536
537    #[test]
538    fn test_prune_columns_struct_full() {
539        let expected_type = Type::from(
540            Schema::builder()
541                .with_fields(vec![
542                    NestedField::optional(
543                        15,
544                        "person",
545                        Type::Struct(StructType::new(vec![
546                            NestedField::optional(16, "name", Primitive(PrimitiveType::String))
547                                .into(),
548                        ])),
549                    )
550                    .into(),
551                ])
552                .build()
553                .unwrap()
554                .as_struct()
555                .clone(),
556        );
557        let schema = table_schema_nested();
558        let selected: HashSet<i32> = HashSet::from([16]);
559        let result = prune_columns(&schema, selected, true);
560        assert!(result.is_ok());
561        assert_eq!(result.unwrap(), expected_type);
562    }
563
564    #[test]
565    fn test_prune_columns_empty_struct() {
566        let schema_with_empty_struct_field = Schema::builder()
567            .with_fields(vec![
568                NestedField::optional(15, "person", Type::Struct(StructType::new(vec![]))).into(),
569            ])
570            .build()
571            .unwrap();
572        let expected_type = Type::from(
573            Schema::builder()
574                .with_fields(vec![
575                    NestedField::optional(15, "person", Type::Struct(StructType::new(vec![])))
576                        .into(),
577                ])
578                .build()
579                .unwrap()
580                .as_struct()
581                .clone(),
582        );
583        let selected: HashSet<i32> = HashSet::from([15]);
584        let result = prune_columns(&schema_with_empty_struct_field, selected, false);
585        assert!(result.is_ok());
586        assert_eq!(result.unwrap(), expected_type);
587    }
588
589    #[test]
590    fn test_prune_columns_empty_struct_full() {
591        let schema_with_empty_struct_field = Schema::builder()
592            .with_fields(vec![
593                NestedField::optional(15, "person", Type::Struct(StructType::new(vec![]))).into(),
594            ])
595            .build()
596            .unwrap();
597        let expected_type = Type::from(
598            Schema::builder()
599                .with_fields(vec![
600                    NestedField::optional(15, "person", Type::Struct(StructType::new(vec![])))
601                        .into(),
602                ])
603                .build()
604                .unwrap()
605                .as_struct()
606                .clone(),
607        );
608        let selected: HashSet<i32> = HashSet::from([15]);
609        let result = prune_columns(&schema_with_empty_struct_field, selected, true);
610        assert!(result.is_ok());
611        assert_eq!(result.unwrap(), expected_type);
612    }
613
614    #[test]
615    fn test_prune_columns_struct_in_map() {
616        let schema_with_struct_in_map_field = Schema::builder()
617            .with_schema_id(1)
618            .with_fields(vec![
619                NestedField::required(
620                    6,
621                    "id_to_person",
622                    Type::Map(MapType {
623                        key_field: NestedField::map_key_element(7, Primitive(PrimitiveType::Int))
624                            .into(),
625                        value_field: NestedField::map_value_element(
626                            8,
627                            Type::Struct(StructType::new(vec![
628                                NestedField::optional(10, "name", Primitive(PrimitiveType::String))
629                                    .into(),
630                                NestedField::required(11, "age", Primitive(PrimitiveType::Int))
631                                    .into(),
632                            ])),
633                            true,
634                        )
635                        .into(),
636                    }),
637                )
638                .into(),
639            ])
640            .build()
641            .unwrap();
642        let expected_type = Type::from(
643            Schema::builder()
644                .with_fields(vec![
645                    NestedField::required(
646                        6,
647                        "id_to_person",
648                        Type::Map(MapType {
649                            key_field: NestedField::map_key_element(
650                                7,
651                                Primitive(PrimitiveType::Int),
652                            )
653                            .into(),
654                            value_field: NestedField::map_value_element(
655                                8,
656                                Type::Struct(StructType::new(vec![
657                                    NestedField::required(11, "age", Primitive(PrimitiveType::Int))
658                                        .into(),
659                                ])),
660                                true,
661                            )
662                            .into(),
663                        }),
664                    )
665                    .into(),
666                ])
667                .build()
668                .unwrap()
669                .as_struct()
670                .clone(),
671        );
672        let selected: HashSet<i32> = HashSet::from([11]);
673        let result = prune_columns(&schema_with_struct_in_map_field, selected, false);
674        assert!(result.is_ok());
675        assert_eq!(result.unwrap(), expected_type);
676    }
677    #[test]
678    fn test_prune_columns_struct_in_map_full() {
679        let schema = Schema::builder()
680            .with_schema_id(1)
681            .with_fields(vec![
682                NestedField::required(
683                    6,
684                    "id_to_person",
685                    Type::Map(MapType {
686                        key_field: NestedField::map_key_element(7, Primitive(PrimitiveType::Int))
687                            .into(),
688                        value_field: NestedField::map_value_element(
689                            8,
690                            Type::Struct(StructType::new(vec![
691                                NestedField::optional(10, "name", Primitive(PrimitiveType::String))
692                                    .into(),
693                                NestedField::required(11, "age", Primitive(PrimitiveType::Int))
694                                    .into(),
695                            ])),
696                            true,
697                        )
698                        .into(),
699                    }),
700                )
701                .into(),
702            ])
703            .build()
704            .unwrap();
705        let expected_type = Type::from(
706            Schema::builder()
707                .with_fields(vec![
708                    NestedField::required(
709                        6,
710                        "id_to_person",
711                        Type::Map(MapType {
712                            key_field: NestedField::map_key_element(
713                                7,
714                                Primitive(PrimitiveType::Int),
715                            )
716                            .into(),
717                            value_field: NestedField::map_value_element(
718                                8,
719                                Type::Struct(StructType::new(vec![
720                                    NestedField::required(11, "age", Primitive(PrimitiveType::Int))
721                                        .into(),
722                                ])),
723                                true,
724                            )
725                            .into(),
726                        }),
727                    )
728                    .into(),
729                ])
730                .build()
731                .unwrap()
732                .as_struct()
733                .clone(),
734        );
735        let selected: HashSet<i32> = HashSet::from([11]);
736        let result = prune_columns(&schema, selected, true);
737        assert!(result.is_ok());
738        assert_eq!(result.unwrap(), expected_type);
739    }
740
741    #[test]
742    fn test_prune_columns_select_original_schema() {
743        let schema = table_schema_nested();
744        let selected: HashSet<i32> = (0..schema.highest_field_id() + 1).collect();
745        let result = prune_columns(&schema, selected, true);
746        assert!(result.is_ok());
747        assert_eq!(result.unwrap(), Type::Struct(schema.as_struct().clone()));
748    }
749
750    #[test]
751    fn test_prune_columns_variant() {
752        // foo (String, id=1) + v (Variant, id=2).
753        let schema = Schema::builder()
754            .with_fields(vec![
755                NestedField::optional(1, "foo", Primitive(PrimitiveType::String)).into(),
756                NestedField::optional(2, "v", Type::Variant(VariantType)).into(),
757            ])
758            .build()
759            .unwrap();
760
761        // A variant is a leaf (like a primitive): selecting it keeps it, the same way
762        // for select_full_types true and false.
763        let only_variant = Type::Struct(StructType::new(vec![
764            NestedField::optional(2, "v", Type::Variant(VariantType)).into(),
765        ]));
766        for full in [false, true] {
767            let result = prune_columns(&schema, HashSet::from([2]), full).unwrap();
768            assert_eq!(result, only_variant, "select_full_types={full}");
769        }
770
771        // Selecting a sibling prunes the variant out.
772        let only_foo = Type::Struct(StructType::new(vec![
773            NestedField::optional(1, "foo", Primitive(PrimitiveType::String)).into(),
774        ]));
775        let result = prune_columns(&schema, HashSet::from([1]), false).unwrap();
776        assert_eq!(result, only_foo);
777    }
778}