1use super::*;
19use crate::error::invalid_data;
20use crate::spec::VariantType;
21
22struct PruneColumn {
23 selected: HashSet<i32>,
24 select_full_types: bool,
25}
26
27pub 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 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 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 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 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 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}