Skip to content

Commit 67f0bde

Browse files
committed
Avoid expanding unreferenced struct plan fields
Signed-off-by: Joe Isaacs <joe.isaacs@live.co.uk>
1 parent dd71e39 commit 67f0bde

2 files changed

Lines changed: 76 additions & 15 deletions

File tree

vortex-layout/src/plan/plans/pack.rs

Lines changed: 8 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -352,8 +352,7 @@ impl PlanParentReduceRule<Pack> for ExpressionPackRule {
352352
.get(&ExactBoundExpr(expression.clone()))
353353
.vortex_expect("Bound expression missing free-field annotations")
354354
.clone();
355-
let expanded_root = expanded_struct_root(child.dtype(), fields)?;
356-
let expanded = expand_struct_root(expression.clone(), &expanded_root, fields)?;
355+
let expanded = expand_struct_root(expression.clone(), fields)?;
357356
let partitioned =
358357
partition_bound(expanded.clone(), make_bound_free_field_annotator(fields))?;
359358

@@ -469,14 +468,13 @@ fn expanded_struct_root(
469468

470469
fn expand_struct_root(
471470
expression: BoundExpression,
472-
expanded_root: &BoundExpression,
473471
fields: &StructFields,
474472
) -> VortexResult<BoundExpression> {
475473
Ok(expression
476474
.transform_down(|node| {
477475
if node.is_root() {
478476
return Ok(Transformed {
479-
value: expanded_root.clone(),
477+
value: expanded_struct_root(node.dtype(), fields)?,
480478
changed: true,
481479
order: TraversalOrder::Skip,
482480
});
@@ -493,28 +491,23 @@ fn expand_struct_root(
493491
return Ok(Transformed::no(node));
494492
}
495493

496-
if let Some(field_name) = scalar_fn.as_opt::<GetItem>() {
497-
let index = fields.find(field_name).ok_or_else(|| {
498-
vortex_err!("Field {field_name} not found while expanding struct root")
499-
})?;
494+
if scalar_fn.is::<GetItem>() {
500495
return Ok(Transformed {
501-
value: expanded_root.children()[index].clone(),
502-
changed: true,
496+
value: node,
497+
changed: false,
503498
order: TraversalOrder::Skip,
504499
});
505500
}
506501

507502
if let Some(selection) = scalar_fn.as_opt::<Select>() {
508503
let names = selection.normalize_to_included_fields(fields.names())?;
504+
let root = node.children()[0].clone();
509505
let children = names
510506
.iter()
511507
.map(|name| {
512-
let index = fields
513-
.find(name)
514-
.vortex_expect("normalized selection fields must exist in the root");
515-
expanded_root.children()[index].clone()
508+
BoundExpression::try_new(GetItem.bind(name.clone()), [root.clone()])
516509
})
517-
.collect();
510+
.collect::<VortexResult<Vec<_>>>()?;
518511
return Ok(Transformed {
519512
value: bound_pack(names, children)?,
520513
changed: true,

vortex-layout/src/plan/tests.rs

Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,8 @@ use vortex_array::expr::is_null;
3030
use vortex_array::expr::lit;
3131
use vortex_array::expr::pack;
3232
use vortex_array::expr::root;
33+
use vortex_array::expr::select;
34+
use vortex_array::expr::select_exclude;
3335
use vortex_error::VortexResult;
3436
use vortex_error::vortex_err;
3537
use vortex_io::runtime::single::block_on;
@@ -279,6 +281,72 @@ fn struct_plan_appends_validity_when_nullable() -> VortexResult<()> {
279281
Ok(())
280282
}
281283

284+
#[test]
285+
fn struct_select_preserves_struct_output() -> VortexResult<()> {
286+
let field_dtype = primitive(PType::I32, Nullability::NonNullable);
287+
let layout = StructLayout::new(
288+
3,
289+
DType::Struct(
290+
StructFields::from_iter([("a", field_dtype.clone()), ("b", field_dtype.clone())]),
291+
Nullability::NonNullable,
292+
),
293+
vec![
294+
flat(3, field_dtype.clone(), 0),
295+
flat(3, field_dtype.clone(), 1),
296+
],
297+
)
298+
.into_layout();
299+
300+
let plan = make_eval(select(["a"], root()), make_plan(layout)?)?.into_plan();
301+
let optimized = optimize(plan)?;
302+
303+
assert_eq!(
304+
optimized.dtype(),
305+
&DType::Struct(
306+
StructFields::from_iter([("a", field_dtype)]),
307+
Nullability::NonNullable,
308+
)
309+
);
310+
insta::assert_snapshot!(optimized.display_tree(), @r"
311+
root: vortex.plan.eval({a=i32}, rows=3) expr=pack(a: $)
312+
child: vortex.plan.segment_scan(i32, rows=3)
313+
");
314+
Ok(())
315+
}
316+
317+
#[test]
318+
fn struct_select_exclude_preserves_struct_output() -> VortexResult<()> {
319+
let field_dtype = primitive(PType::I32, Nullability::NonNullable);
320+
let layout = StructLayout::new(
321+
3,
322+
DType::Struct(
323+
StructFields::from_iter([("a", field_dtype.clone()), ("b", field_dtype.clone())]),
324+
Nullability::NonNullable,
325+
),
326+
vec![
327+
flat(3, field_dtype.clone(), 0),
328+
flat(3, field_dtype.clone(), 1),
329+
],
330+
)
331+
.into_layout();
332+
333+
let plan = make_eval(select_exclude(["b"], root()), make_plan(layout)?)?.into_plan();
334+
let optimized = optimize(plan)?;
335+
336+
assert_eq!(
337+
optimized.dtype(),
338+
&DType::Struct(
339+
StructFields::from_iter([("a", field_dtype)]),
340+
Nullability::NonNullable,
341+
)
342+
);
343+
insta::assert_snapshot!(optimized.display_tree(), @r"
344+
root: vortex.plan.eval({a=i32}, rows=3) expr=pack(a: $)
345+
child: vortex.plan.segment_scan(i32, rows=3)
346+
");
347+
Ok(())
348+
}
349+
282350
#[test]
283351
fn with_children_rejects_mismatched_arity() -> VortexResult<()> {
284352
let dtype = primitive(PType::I32, Nullability::NonNullable);

0 commit comments

Comments
 (0)