(
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 => {
@@ -164,8 +164,8 @@ fn binary_falsify(
return Ok(None);
}
- let lhs_falsifier = ctx.falsify(lhs)?;
- let rhs_falsifier = ctx.falsify(rhs)?;
+ let lhs_falsifier = falsify(lhs, session)?;
+ let rhs_falsifier = falsify(rhs, session)?;
or_collect(lhs_falsifier.into_iter().chain(rhs_falsifier))
}
Operator::Or => {
@@ -176,7 +176,7 @@ fn binary_falsify(
return Ok(None);
}
- match (ctx.falsify(lhs)?, ctx.falsify(rhs)?) {
+ match (falsify(lhs, session)?, falsify(rhs, session)?) {
(Some(lhs), Some(rhs)) => Some(and(lhs, rhs)),
_ => None,
}
@@ -196,7 +196,7 @@ impl StatsRewriteRule for BetweenStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ session: &VortexSession,
) -> VortexResult> {
let options = expr.as_::();
let arr = expr.child(0).clone();
@@ -205,7 +205,7 @@ impl StatsRewriteRule for BetweenStatsRewrite {
let lhs = binary(options.lower_strict.to_operator(), lower, arr.clone());
let rhs = binary(options.upper_strict.to_operator(), arr, upper);
- ctx.falsify(&and(lhs, rhs))
+ falsify(&and(lhs, rhs), session)
}
}
@@ -220,17 +220,17 @@ impl StatsRewriteRule for IsNullNullCountStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> 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<'_>,
+ _session: &VortexSession,
) -> 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())))
}
}
@@ -245,7 +245,7 @@ impl StatsRewriteRule for IsNullAllNonNullStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- _ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult > {
Ok(Some(all_non_null(expr.child(0))))
}
@@ -262,7 +262,7 @@ impl StatsRewriteRule for IsNullAllNullStatsRewrite {
fn satisfy(
&self,
expr: &BoundExpression,
- _ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult > {
Ok(Some(all_null(expr.child(0))))
}
@@ -279,17 +279,17 @@ impl StatsRewriteRule for IsNotNullNullCountStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> 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<'_>,
+ _session: &VortexSession,
) -> 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))))
}
}
@@ -304,7 +304,7 @@ impl StatsRewriteRule for IsNotNullAllNullStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- _ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult > {
Ok(Some(all_null(expr.child(0))))
}
@@ -321,7 +321,7 @@ impl StatsRewriteRule for IsNotNullAllNonNullStatsRewrite {
fn satisfy(
&self,
expr: &BoundExpression,
- _ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult > {
Ok(Some(all_non_null(expr.child(0))))
}
@@ -338,7 +338,7 @@ impl StatsRewriteRule for LikeStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult > {
let like_options = expr.as_::();
if like_options.negated || like_options.case_insensitive {
@@ -355,8 +355,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 +368,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 +393,9 @@ impl StatsRewriteRule for ListContainsNanCountStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult> {
- list_contains_falsify::(expr, ctx)
+ list_contains_falsify::(expr)
}
}
@@ -410,15 +410,14 @@ impl StatsRewriteRule for ListContainsAllNonNanStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> 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 +436,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 +450,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 +466,9 @@ impl StatsRewriteRule for DynamicComparisonNanCountStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult> {
- dynamic_comparison_falsify::(expr, ctx)
+ dynamic_comparison_falsify::(expr)
}
}
@@ -484,25 +483,24 @@ impl StatsRewriteRule for DynamicComparisonAllNonNanStatsRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> 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 +513,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 +545,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 +553,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 +566,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 +577,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 +592,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 +610,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 +623,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 +681,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,
}
}
@@ -761,7 +745,6 @@ mod tests {
use crate::scalar_fn::internal::row_count::RowCount;
use crate::stats::expr::StatFn;
use crate::stats::expr::StatOptions;
- use crate::stats::rewrite::StatsRewriteCtx;
use crate::stats::rewrite::StatsRewriteRule;
use crate::stats::session::StatsSessionExt;
@@ -898,7 +881,7 @@ mod tests {
fn falsify(
&self,
_expr: &BoundExpression,
- _ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult> {
self.0.fetch_add(1, Ordering::Relaxed);
Ok(None)
diff --git a/vortex-layout/src/layouts/zoned/aggregates/bloom_filter/scalar_fn.rs b/vortex-layout/src/layouts/zoned/aggregates/bloom_filter/scalar_fn.rs
index 847617e95f8..a6164258d9c 100644
--- a/vortex-layout/src/layouts/zoned/aggregates/bloom_filter/scalar_fn.rs
+++ b/vortex-layout/src/layouts/zoned/aggregates/bloom_filter/scalar_fn.rs
@@ -36,7 +36,6 @@ use vortex_array::scalar_fn::fns::binary::Binary;
use vortex_array::scalar_fn::fns::literal::Literal;
use vortex_array::scalar_fn::fns::operators::Operator;
use vortex_array::stats::expr::bound::stat as bound_stat;
-use vortex_array::stats::rewrite::StatsRewriteCtx;
use vortex_array::stats::rewrite::StatsRewriteRule;
use vortex_buffer::BitBufferMut;
use vortex_buffer::Buffer;
@@ -249,7 +248,7 @@ impl StatsRewriteRule for BloomEqRewrite {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ _session: &VortexSession,
) -> VortexResult > {
if *expr.as_::() != Operator::Eq {
return Ok(None);
@@ -267,7 +266,7 @@ impl StatsRewriteRule for BloomEqRewrite {
// Nulls are not stored in Bloom filters, so it is not possible to determine
// if it is present or not, so the answer is inconclusive.
- if !is_bloom_valid_dtype(&ctx.return_dtype(column)?) || literal.as_::().is_null() {
+ if !is_bloom_valid_dtype(column.dtype()) || literal.as_::().is_null() {
return Ok(None);
}
@@ -308,7 +307,6 @@ mod tests {
use vortex_array::scalar_fn::VecExecutionArgs;
use vortex_array::scalar_fn::session::ScalarFnSessionExt;
use vortex_array::stats::StatsSessionExt;
- use vortex_array::stats::rewrite::StatsRewriteCtx;
use vortex_array::stats::rewrite::StatsRewriteRule;
use vortex_array::validity::Validity;
use vortex_error::VortexResult;
@@ -377,16 +375,15 @@ mod tests {
fn bloom_rule_is_inconclusive_for_nulls() -> VortexResult<()> {
let dtype = DType::Primitive(PType::I64, Nullability::Nullable);
let session = array_session();
- let ctx = StatsRewriteCtx::new(&session);
let rule = BloomEqRewrite {
options: BloomOptions::default(),
};
let non_literal = eq(root(dtype.clone()), root(dtype.clone()));
- assert!(rule.falsify(&non_literal, &ctx)?.is_none());
+ assert!(rule.falsify(&non_literal, &session)?.is_none());
let null_literal = eq(root(dtype.clone()), lit(Scalar::null(dtype)));
- assert!(rule.falsify(&null_literal, &ctx)?.is_none());
+ assert!(rule.falsify(&null_literal, &session)?.is_none());
Ok(())
}
diff --git a/vortex-spatial/src/prune/distance.rs b/vortex-spatial/src/prune/distance.rs
index 9231b33875e..e1648b405cc 100644
--- a/vortex-spatial/src/prune/distance.rs
+++ b/vortex-spatial/src/prune/distance.rs
@@ -10,9 +10,9 @@ use vortex_array::scalar_fn::ScalarFnVTable;
use vortex_array::scalar_fn::fns::binary::Binary;
use vortex_array::scalar_fn::fns::literal::Literal;
use vortex_array::scalar_fn::fns::operators::Operator;
-use vortex_array::stats::rewrite::StatsRewriteCtx;
use vortex_array::stats::rewrite::StatsRewriteRule;
use vortex_error::VortexResult;
+use vortex_session::VortexSession;
use super::aabb_stat;
use super::geometry_and_constant;
@@ -43,7 +43,7 @@ impl StatsRewriteRule for SpatialDistancePrune {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ session: &VortexSession,
) -> VortexResult> {
// Only the ordered comparisons prune today. `== r` could prune in the future (a chunk is
// provably empty when `r` lies outside its box's [min, max] distance interval), it's just
@@ -77,10 +77,10 @@ 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 {
+ let Some(query) = query_aabb(constant, session)? else {
return Ok(None);
};
Ok(distance_prune_proof(geom, query, op, radius))
@@ -140,7 +140,6 @@ mod tests {
use vortex_array::scalar_fn::ScalarFnVTableExt;
use vortex_array::scalar_fn::fns::binary::Binary;
use vortex_array::scalar_fn::fns::operators::Operator;
- use vortex_array::stats::rewrite::StatsRewriteCtx;
use vortex_array::stats::rewrite::StatsRewriteRule;
use vortex_error::VortexResult;
@@ -174,7 +173,7 @@ mod tests {
.new_expr(operator, [distance, lit(radius.into())])
.bind(&scope)?;
- SpatialDistancePrune.falsify(&predicate, &StatsRewriteCtx::new(&session))
+ SpatialDistancePrune.falsify(&predicate, &session)
}
/// A null geometry literal (`ST_Distance(geom, NULL) <= r`) declines cleanly instead of
@@ -190,8 +189,11 @@ mod tests {
.new_expr(Operator::Lte, [distance, lit(0.5f64)])
.bind(&scope)?;
- let ctx = StatsRewriteCtx::new(&session);
- assert!(SpatialDistancePrune.falsify(&predicate, &ctx)?.is_none());
+ assert!(
+ SpatialDistancePrune
+ .falsify(&predicate, &session)?
+ .is_none()
+ );
Ok(())
}
@@ -283,8 +285,11 @@ mod tests {
let scope = point_column(vec![0.0], vec![0.0])?.dtype().clone();
let predicate = lt_eq(lit(1.0f64), lit(2.0f64)).bind(&scope)?;
- let ctx = StatsRewriteCtx::new(&session);
- assert!(SpatialDistancePrune.falsify(&predicate, &ctx)?.is_none());
+ assert!(
+ SpatialDistancePrune
+ .falsify(&predicate, &session)?
+ .is_none()
+ );
Ok(())
}
diff --git a/vortex-spatial/src/prune/intersects.rs b/vortex-spatial/src/prune/intersects.rs
index 74103003b2f..779b24d8651 100644
--- a/vortex-spatial/src/prune/intersects.rs
+++ b/vortex-spatial/src/prune/intersects.rs
@@ -6,9 +6,9 @@
use vortex_array::expr::BoundExpression;
use vortex_array::scalar_fn::ScalarFnId;
use vortex_array::scalar_fn::ScalarFnVTable;
-use vortex_array::stats::rewrite::StatsRewriteCtx;
use vortex_array::stats::rewrite::StatsRewriteRule;
use vortex_error::VortexResult;
+use vortex_session::VortexSession;
use super::aabb_stat;
use super::geometry_and_constant;
@@ -35,12 +35,12 @@ impl StatsRewriteRule for SpatialIntersectsPrune {
fn falsify(
&self,
expr: &BoundExpression,
- ctx: &StatsRewriteCtx<'_>,
+ session: &VortexSession,
) -> 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 {
+ let Some(query) = query_aabb(constant, session)? else {
return Ok(None);
};
// Disjoint iff the minimum box-to-box distance is positive. Strictly (`gt`, not `gt_eq`):
@@ -62,7 +62,6 @@ mod tests {
use vortex_array::scalar::Scalar;
use vortex_array::scalar_fn::EmptyOptions;
use vortex_array::scalar_fn::ScalarFnVTableExt;
- use vortex_array::stats::rewrite::StatsRewriteCtx;
use vortex_array::stats::rewrite::StatsRewriteRule;
use vortex_error::VortexResult;
@@ -89,7 +88,7 @@ mod tests {
let predicate = SpatialIntersects
.new_expr(EmptyOptions, operands)
.bind(&scope)?;
- SpatialIntersectsPrune.falsify(&predicate, &StatsRewriteCtx::new(&session))
+ SpatialIntersectsPrune.falsify(&predicate, &session)
}
/// Intersects is symmetric: both operand orders produce a proof.
@@ -126,8 +125,11 @@ mod tests {
.new_expr(EmptyOptions, [root(), lit(null_query)])
.bind(&scope)?;
- let ctx = StatsRewriteCtx::new(&session);
- assert!(SpatialIntersectsPrune.falsify(&predicate, &ctx)?.is_none());
+ assert!(
+ SpatialIntersectsPrune
+ .falsify(&predicate, &session)?
+ .is_none()
+ );
Ok(())
}
diff --git a/vortex-spatial/src/prune/mod.rs b/vortex-spatial/src/prune/mod.rs
index 32c2058b384..2ce211c16e3 100644
--- a/vortex-spatial/src/prune/mod.rs
+++ b/vortex-spatial/src/prune/mod.rs
@@ -37,8 +37,8 @@ use vortex_array::scalar::Scalar;
use vortex_array::scalar_fn::fns::literal::Literal;
use vortex_array::scalar_fn::fns::operators::Operator;
use vortex_array::stats::bound::stat;
-use vortex_array::stats::rewrite::StatsRewriteCtx;
use vortex_error::VortexResult;
+use vortex_session::VortexSession;
use crate::aggregate_fn::GeometryAabb;
use crate::extension::is_native_geometry;
@@ -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);
}
@@ -82,7 +81,7 @@ fn geometry_and_constant<'a>(
/// so whatever holds for the box holds for the geometry.
fn query_aabb(
constant: &Scalar,
- ctx: &StatsRewriteCtx<'_>,
+ session: &VortexSession,
) -> VortexResult >> {
// A null geometry literal has no extent to prove against, so it can never prune.
if constant.is_null() {
@@ -90,7 +89,7 @@ fn query_aabb(
}
// Decoding the constant into a concrete geometry runs through the compute stack, which needs
// an execution context.
- let mut exec = ctx.session().create_execution_ctx();
+ let mut exec = session.create_execution_ctx();
Ok(single_geometry(constant, &mut exec)?.bounding_rect())
}