diff --git a/vortex-array/src/stats/rewrite.rs b/vortex-array/src/stats/rewrite.rs index ddf74ee5dab..012c29659c5 100644 --- a/vortex-array/src/stats/rewrite.rs +++ b/vortex-array/src/stats/rewrite.rs @@ -101,11 +101,6 @@ impl<'a> StatsRewriteCtx<'a> { self.session } - /// Return the dtype of `expr` within this rewrite scope. - pub fn return_dtype(&self, expr: &BoundExpression) -> VortexResult { - Ok(expr.dtype().clone()) - } - /// Rewrite `expr` into a stats-backed falsifier. pub fn falsify(&self, expr: &BoundExpression) -> VortexResult> { self.ensure_predicate(expr)?; @@ -118,8 +113,9 @@ impl<'a> StatsRewriteCtx<'a> { rewrite(expr, self, StatsRewriteRule::satisfy) } + #[inline] fn ensure_predicate(&self, expr: &BoundExpression) -> VortexResult<()> { - let dtype = self.return_dtype(expr)?; + let dtype = expr.dtype(); vortex_ensure!( matches!(dtype, DType::Bool(_)), "Stats rewrites require a boolean predicate, got {dtype}", diff --git a/vortex-array/src/stats/rewrite/builtins.rs b/vortex-array/src/stats/rewrite/builtins.rs index 3cb5fdb06df..83b26fd6695 100644 --- a/vortex-array/src/stats/rewrite/builtins.rs +++ b/vortex-array/src/stats/rewrite/builtins.rs @@ -120,43 +120,42 @@ fn binary_falsify( Ok(match operator { Operator::Eq => { - let left = min(lhs, ctx).zip(max(rhs, ctx)).map(|(a, b)| gt(a, b)); - let right = min(rhs, ctx).zip(max(lhs, ctx)).map(|(a, b)| gt(a, b)); + let left = min(lhs).zip(max(rhs)).map(|(a, b)| gt(a, b)); + let right = min(rhs).zip(max(lhs)).map(|(a, b)| gt(a, b)); or_collect(left.into_iter().chain(right)) - .map(|value_predicate| with_non_nan_guards::

(ctx, [lhs, rhs], value_predicate)) + .map(|value_predicate| with_non_nan_guards::

([lhs, rhs], value_predicate)) .transpose()? .flatten() } - Operator::NotEq => min(lhs, ctx) - .zip(max(rhs, ctx)) - .zip(max(lhs, ctx).zip(min(rhs, ctx))) + Operator::NotEq => min(lhs) + .zip(max(rhs)) + .zip(max(lhs).zip(min(rhs))) .map(|((min_lhs, max_rhs), (max_lhs, min_rhs))| { with_non_nan_guards::

( - ctx, [lhs, rhs], and(eq(min_lhs, max_rhs), eq(max_lhs, min_rhs)), ) }) .transpose()? .flatten(), - Operator::Gt => max(lhs, ctx) - .zip(min(rhs, ctx)) - .map(|(a, b)| with_non_nan_guards::

(ctx, [lhs, rhs], lt_eq(a, b))) + Operator::Gt => max(lhs) + .zip(min(rhs)) + .map(|(a, b)| with_non_nan_guards::

([lhs, rhs], lt_eq(a, b))) .transpose()? .flatten(), - Operator::Gte => max(lhs, ctx) - .zip(min(rhs, ctx)) - .map(|(a, b)| with_non_nan_guards::

(ctx, [lhs, rhs], lt(a, b))) + Operator::Gte => max(lhs) + .zip(min(rhs)) + .map(|(a, b)| with_non_nan_guards::

([lhs, rhs], lt(a, b))) .transpose()? .flatten(), - Operator::Lt => min(lhs, ctx) - .zip(max(rhs, ctx)) - .map(|(a, b)| with_non_nan_guards::

(ctx, [lhs, rhs], gt_eq(a, b))) + Operator::Lt => min(lhs) + .zip(max(rhs)) + .map(|(a, b)| with_non_nan_guards::

([lhs, rhs], gt_eq(a, b))) .transpose()? .flatten(), - Operator::Lte => min(lhs, ctx) - .zip(max(rhs, ctx)) - .map(|(a, b)| with_non_nan_guards::

(ctx, [lhs, rhs], gt(a, b))) + Operator::Lte => min(lhs) + .zip(max(rhs)) + .map(|(a, b)| with_non_nan_guards::

([lhs, rhs], gt(a, b))) .transpose()? .flatten(), Operator::And => { @@ -220,17 +219,17 @@ impl StatsRewriteRule for IsNullNullCountStatsRewrite { fn falsify( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - Ok(null_count(expr.child(0), ctx).map(|null_count| eq(null_count, lit(0u64)))) + Ok(null_count(expr.child(0)).map(|null_count| eq(null_count, lit(0u64)))) } fn satisfy( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - Ok(null_count(expr.child(0), ctx).map(|null_count| eq(null_count, row_count()))) + Ok(null_count(expr.child(0)).map(|null_count| eq(null_count, row_count()))) } } @@ -279,17 +278,17 @@ impl StatsRewriteRule for IsNotNullNullCountStatsRewrite { fn falsify( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - Ok(null_count(expr.child(0), ctx).map(|null_count| eq(null_count, row_count()))) + Ok(null_count(expr.child(0)).map(|null_count| eq(null_count, row_count()))) } fn satisfy( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - Ok(null_count(expr.child(0), ctx).map(|null_count| eq(null_count, lit(0u64)))) + Ok(null_count(expr.child(0)).map(|null_count| eq(null_count, lit(0u64)))) } } @@ -338,7 +337,7 @@ impl StatsRewriteRule for LikeStatsRewrite { fn falsify( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { let like_options = expr.as_::(); if like_options.negated || like_options.case_insensitive { @@ -355,8 +354,8 @@ impl StatsRewriteRule for LikeStatsRewrite { let source = expr.child(0); Ok(match LikeVariant::from_str(pattern) { Some(LikeVariant::Exact(text)) => { - min(source, ctx) - .zip(max(source, ctx)) + min(source) + .zip(max(source)) .map(|(source_min, source_max)| { or( gt(source_min, lit(text.as_ref())), @@ -368,8 +367,8 @@ impl StatsRewriteRule for LikeStatsRewrite { let Some(successor) = prefix.to_string().increment().ok() else { return Ok(None); }; - min(source, ctx) - .zip(max(source, ctx)) + min(source) + .zip(max(source)) .map(|(source_min, source_max)| { or( gt_eq(source_min, lit(successor)), @@ -393,9 +392,9 @@ impl StatsRewriteRule for ListContainsNanCountStatsRewrite { fn falsify( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - list_contains_falsify::(expr, ctx) + list_contains_falsify::(expr) } } @@ -410,15 +409,14 @@ impl StatsRewriteRule for ListContainsAllNonNanStatsRewrite { fn falsify( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - list_contains_falsify::(expr, ctx) + list_contains_falsify::(expr) } } fn list_contains_falsify( expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { let list = expr.child(0); let needle = expr.child(1); @@ -437,10 +435,10 @@ fn list_contains_falsify( return Ok(P::EMIT_UNGUARDED_REWRITES.then(|| lit(true))); } - let Some(value_max) = max(needle, ctx) else { + let Some(value_max) = max(needle) else { return Ok(None); }; - let Some(value_min) = min(needle, ctx) else { + let Some(value_min) = min(needle) else { return Ok(None); }; @@ -451,7 +449,7 @@ fn list_contains_falsify( ) })); value_predicate - .map(|value_predicate| with_non_nan_guards::

(ctx, [needle], value_predicate)) + .map(|value_predicate| with_non_nan_guards::

([needle], value_predicate)) .transpose() .map(Option::flatten) } @@ -467,9 +465,9 @@ impl StatsRewriteRule for DynamicComparisonNanCountStatsRewrite { fn falsify( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - dynamic_comparison_falsify::(expr, ctx) + dynamic_comparison_falsify::(expr) } } @@ -484,25 +482,24 @@ impl StatsRewriteRule for DynamicComparisonAllNonNanStatsRewrite { fn falsify( &self, expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, + _ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - dynamic_comparison_falsify::(expr, ctx) + dynamic_comparison_falsify::(expr) } } fn dynamic_comparison_falsify( expr: &BoundExpression, - ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { let dynamic = expr.as_::(); let lhs = expr.child(0); let Some((operator, lhs_stat)) = (match dynamic.operator { CompareOperator::Eq | CompareOperator::NotEq => None, - CompareOperator::Gt => max(lhs, ctx).map(|lhs_stat| (CompareOperator::Lte, lhs_stat)), - CompareOperator::Gte => max(lhs, ctx).map(|lhs_stat| (CompareOperator::Lt, lhs_stat)), - CompareOperator::Lt => min(lhs, ctx).map(|lhs_stat| (CompareOperator::Gte, lhs_stat)), - CompareOperator::Lte => min(lhs, ctx).map(|lhs_stat| (CompareOperator::Gt, lhs_stat)), + CompareOperator::Gt => max(lhs).map(|lhs_stat| (CompareOperator::Lte, lhs_stat)), + CompareOperator::Gte => max(lhs).map(|lhs_stat| (CompareOperator::Lt, lhs_stat)), + CompareOperator::Lt => min(lhs).map(|lhs_stat| (CompareOperator::Gte, lhs_stat)), + CompareOperator::Lte => min(lhs).map(|lhs_stat| (CompareOperator::Gt, lhs_stat)), }) else { return Ok(None); }; @@ -515,19 +512,19 @@ fn dynamic_comparison_falsify( }, lhs_stat, ); - with_non_nan_guards::

(ctx, [lhs], value_predicate) + with_non_nan_guards::

([lhs], value_predicate) } -fn min(expr: &BoundExpression, ctx: &StatsRewriteCtx<'_>) -> Option { - stat_expr(expr, Stat::Min, ctx) +fn min(expr: &BoundExpression) -> Option { + stat_expr(expr, Stat::Min) } -fn max(expr: &BoundExpression, ctx: &StatsRewriteCtx<'_>) -> Option { - stat_expr(expr, Stat::Max, ctx) +fn max(expr: &BoundExpression) -> Option { + stat_expr(expr, Stat::Max) } -fn null_count(expr: &BoundExpression, ctx: &StatsRewriteCtx<'_>) -> Option { - stat_expr(expr, Stat::NullCount, ctx) +fn null_count(expr: &BoundExpression) -> Option { + stat_expr(expr, Stat::NullCount) } fn all_null(expr: &BoundExpression) -> BoundExpression { @@ -547,7 +544,7 @@ enum NanCheck { trait NonNanProof { const EMIT_UNGUARDED_REWRITES: bool; - fn check(ctx: &StatsRewriteCtx<'_>, expr: &BoundExpression) -> VortexResult; + fn check(expr: &BoundExpression) -> VortexResult; } struct NanCountProof; @@ -555,12 +552,10 @@ struct NanCountProof; impl NonNanProof for NanCountProof { const EMIT_UNGUARDED_REWRITES: bool = true; - fn check(ctx: &StatsRewriteCtx<'_>, expr: &BoundExpression) -> VortexResult { - non_nan_check(ctx, expr, |expr| { - match stat_expr(expr, Stat::NaNCount, ctx) { - Some(nan_count) => NanCheck::Check(eq(nan_count, lit(0u64))), - None => NanCheck::Unavailable, - } + fn check(expr: &BoundExpression) -> VortexResult { + non_nan_check(expr, |expr| match stat_expr(expr, Stat::NaNCount) { + Some(nan_count) => NanCheck::Check(eq(nan_count, lit(0u64))), + None => NanCheck::Unavailable, }) } } @@ -570,8 +565,8 @@ struct AllNonNanProof; impl NonNanProof for AllNonNanProof { const EMIT_UNGUARDED_REWRITES: bool = false; - fn check(ctx: &StatsRewriteCtx<'_>, expr: &BoundExpression) -> VortexResult { - non_nan_check(ctx, expr, |expr| { + fn check(expr: &BoundExpression) -> VortexResult { + non_nan_check(expr, |expr| { NanCheck::Check(stat_fn(expr.clone(), AllNonNan.bind(AggregateEmptyOptions))) }) } @@ -581,7 +576,6 @@ impl NonNanProof for AllNonNanProof { // candidate value is known to be non-NaN. Cast result dtypes are not enough: a cast // from float to non-float still needs a proof about the float source values. fn non_nan_check( - ctx: &StatsRewriteCtx<'_>, expr: &BoundExpression, proof: impl FnOnce(&BoundExpression) -> NanCheck, ) -> VortexResult { @@ -597,14 +591,14 @@ fn non_nan_check( } if expr.is::() { - if !has_nans(&ctx.return_dtype(expr.child(0))?) { + if !has_nans(expr.child(0).dtype()) { return Ok(NanCheck::NotNeeded); } - return non_nan_check(ctx, expr.child(0), proof); + return non_nan_check(expr.child(0), proof); } - if !has_nans(&ctx.return_dtype(expr)?) { + if !has_nans(expr.dtype()) { return Ok(NanCheck::NotNeeded); } @@ -615,11 +609,7 @@ fn has_nans(dtype: &DType) -> bool { dtype.is_float() } -fn stat_expr( - expr: &BoundExpression, - stat: Stat, - ctx: &StatsRewriteCtx<'_>, -) -> Option { +fn stat_expr(expr: &BoundExpression, stat: Stat) -> Option { if let Some(literal) = literal_stat(expr, stat) { return Some(literal); } @@ -632,28 +622,26 @@ fn stat_expr( } if let Some(dtype) = expr.as_opt::() { - return cast_stat(expr.child(0), dtype, stat, ctx); + return cast_stat(expr.child(0), dtype, stat); } let aggregate_fn = stat.aggregate_fn()?; // The aggregate may not support the expression's dtype, e.g. min/max over structs, // even when the predicate itself is well-typed. Such stats cannot be lowered later, // so do not reference them in the rewrite. - let input_dtype = ctx.return_dtype(expr).ok()?; aggregate_fn - .return_dtype(&input_dtype) + .return_dtype(expr.dtype()) .is_some() .then(|| stat_fn(expr.clone(), aggregate_fn)) } fn with_non_nan_guards<'a, P: NonNanProof>( - ctx: &StatsRewriteCtx<'_>, exprs: impl IntoIterator, value_predicate: BoundExpression, ) -> VortexResult> { let mut nan_checks = Vec::new(); for expr in exprs { - match P::check(ctx, expr)? { + match P::check(expr)? { NanCheck::NotNeeded => {} NanCheck::Check(check) => nan_checks.push(check), NanCheck::Unavailable => return Ok(None), @@ -692,15 +680,10 @@ fn literal_stat(expr: &BoundExpression, stat: Stat) -> Option { } } -fn cast_stat( - expr: &BoundExpression, - dtype: &DType, - stat: Stat, - ctx: &StatsRewriteCtx<'_>, -) -> Option { +fn cast_stat(expr: &BoundExpression, dtype: &DType, stat: Stat) -> Option { match stat { - Stat::Min | Stat::Max => stat_expr(expr, stat, ctx).map(|stat| cast(stat, dtype.clone())), - Stat::NaNCount | Stat::Sum | Stat::UncompressedSizeInBytes => stat_expr(expr, stat, ctx), + Stat::Min | Stat::Max => stat_expr(expr, stat).map(|stat| cast(stat, dtype.clone())), + Stat::NaNCount | Stat::Sum | Stat::UncompressedSizeInBytes => stat_expr(expr, stat), Stat::NullCount | Stat::IsConstant | Stat::IsSorted | Stat::IsStrictSorted => None, } } diff --git a/vortex-spatial/src/prune/distance.rs b/vortex-spatial/src/prune/distance.rs index 9231b33875e..f669eb197a3 100644 --- a/vortex-spatial/src/prune/distance.rs +++ b/vortex-spatial/src/prune/distance.rs @@ -77,7 +77,7 @@ impl StatsRewriteRule for SpatialDistancePrune { return Ok(None); } - let Some((geom, constant)) = geometry_and_constant(distance, ctx)? else { + let Some((geom, constant)) = geometry_and_constant(distance)? else { return Ok(None); }; let Some(query) = query_aabb(constant, ctx)? else { diff --git a/vortex-spatial/src/prune/intersects.rs b/vortex-spatial/src/prune/intersects.rs index 74103003b2f..e161ba399bb 100644 --- a/vortex-spatial/src/prune/intersects.rs +++ b/vortex-spatial/src/prune/intersects.rs @@ -37,7 +37,7 @@ impl StatsRewriteRule for SpatialIntersectsPrune { expr: &BoundExpression, ctx: &StatsRewriteCtx<'_>, ) -> VortexResult> { - let Some((geom, constant)) = geometry_and_constant(expr, ctx)? else { + let Some((geom, constant)) = geometry_and_constant(expr)? else { return Ok(None); }; let Some(query) = query_aabb(constant, ctx)? else { diff --git a/vortex-spatial/src/prune/mod.rs b/vortex-spatial/src/prune/mod.rs index 32c2058b384..d39adbdade7 100644 --- a/vortex-spatial/src/prune/mod.rs +++ b/vortex-spatial/src/prune/mod.rs @@ -51,10 +51,9 @@ use crate::extension::single_geometry; /// shape (in either operand order), or the column's dtype carries no [`GeometryAabb`] statistic. /// An asymmetric predicate (e.g. a future contains) must recover which operand is the column /// itself instead of calling this. -fn geometry_and_constant<'a>( - expr: &'a BoundExpression, - ctx: &StatsRewriteCtx<'_>, -) -> VortexResult> { +fn geometry_and_constant( + expr: &BoundExpression, +) -> VortexResult> { // The predicate is symmetric, so the column (scope root) and the constant may be on either // side. let (lhs, rhs) = (expr.child(0), expr.child(1)); @@ -68,7 +67,7 @@ fn geometry_and_constant<'a>( // A `GeometryAabb` stat reference only binds for dtypes it supports; anything else (e.g. a // WKB column) must fall through to the scan. - if !is_native_geometry(&ctx.return_dtype(geom)?) { + if !is_native_geometry(geom.dtype()) { return Ok(None); }