Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
198 changes: 198 additions & 0 deletions crates/asap-aware-mapping/src/frequency_rewrite.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,198 @@
//! Narrow frequency recognition over the resolved scalar/operator graph.
use asap_types::ir::non_asap::any_measure_filtered;
use asap_types::{
ir::{ExprSemantics, NonASAPOp, Operator, OperatorNode, ProjectItem, ScalarExpr},
pre_asap::{
AggIntent, ArithmeticOpKind, CompareOpKind, DataType, GroupKeys, Reduction, ScalarValue,
},
};
use std::rc::Rc;

// Only follow projections. Crossing a filter/limit would change the population.
fn expand(
mut expr: ScalarExpr,
mut node: Rc<OperatorNode>,
) -> Option<(ScalarExpr, Rc<OperatorNode>)> {
while let Some(NonASAPOp::Project { cols, child, .. }) = node.non_asap() {
expr = substitute(&expr, cols)?;
node = Rc::clone(child);
}
Some((expr, node))
}
fn substitute(expr: &ScalarExpr, cols: &[ProjectItem]) -> Option<ScalarExpr> {
Some(match expr {
ScalarExpr::Column(id) => cols.get(*id)?.expr.clone(),
ScalarExpr::Cast {
expr,
to,
try_cast: false,
} if *to == DataType::Float64 => ScalarExpr::Cast {
expr: Box::new(substitute(expr, cols)?),
to: to.clone(),
try_cast: false,
},
ScalarExpr::Arithmetic {
op,
left,
right,
semantics,
} => ScalarExpr::Arithmetic {
op: op.clone(),
left: Box::new(substitute(left, cols)?),
right: Box::new(substitute(right, cols)?),
semantics: *semantics,
},
ScalarExpr::FunctionCall { name, args } => ScalarExpr::FunctionCall {
name: name.clone(),
args: args
.iter()
.map(|arg| substitute(arg, cols))
.collect::<Option<_>>()?,
},
_ => return None,
})
}
fn uncast(expr: &ScalarExpr) -> &ScalarExpr {
match expr {
ScalarExpr::Cast {
expr,
to: DataType::Float64,
try_cast: false,
} => uncast(expr),
_ => expr,
}
}

pub(super) fn frequency_l2_rewrite(root: &Rc<OperatorNode>) -> Option<Rc<OperatorNode>> {
let NonASAPOp::Project {
cols,
child,
qualifier,
} = root.non_asap()?
else {
return None;
};
let [item] = cols.as_slice() else {
return None;
};
let (expr, outer) = expand(item.expr.clone(), Rc::clone(child))?;
let ScalarExpr::FunctionCall { name, args } = &expr else {
return None;
};
if !name.eq_ignore_ascii_case("sqrt") {
return None;
}
let [arg] = args.as_slice() else {
return None;
};
if !matches!(uncast(arg), ScalarExpr::Column(0)) {
return None;
}
let NonASAPOp::Aggregate {
reduction: Reduction::Reduce(keys),
measures,
filters,
having: None,
child,
..
} = outer.non_asap()?
else {
return None;
};
if keys.is_without() || !keys.keys().is_empty() || any_measure_filtered(filters) {
return None;
}
let [AggIntent::Sum { col: Some(col) }] = measures.as_slice() else {
return None;
};
let (product, inner) = expand(ScalarExpr::Column(*col), Rc::clone(child))?;
let ScalarExpr::Arithmetic {
op: ArithmeticOpKind::Mul,
left,
right,
semantics: ExprSemantics::Sql,
} = &product
else {
return None;
};
// SQL Int64 multiplication can overflow. Admit only products already typed Float64.
if product.scalar_type(&inner.schema).ok()?.0 != DataType::Float64 {
return None;
}
let NonASAPOp::Aggregate {
reduction: Reduction::Reduce(keys),
measures,
filters,
having: None,
child,
..
} = inner.non_asap()?
else {
return None;
};
let [key] = keys.keys() else {
return None;
};
if keys.is_without() || any_measure_filtered(filters) {
return None;
}
let [AggIntent::Count { accuracy }] = measures.as_slice() else {
return None;
};
if !matches!(
(uncast(left), uncast(right)),
(ScalarExpr::Column(1), ScalarExpr::Column(1))
) {
return None;
}
let field = child.schema.fields.get(*key)?;
// COUNT(*) GROUP BY NULL creates a real group; the frequency intent skips it.
if field.nullable
|| !matches!(
field.plain_dtype()?,
DataType::Bool | DataType::Int64 | DataType::Utf8
)
{
return None;
}
let aggregate = OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Aggregate {
reduction: Reduction::Reduce(GroupKeys::none()),
measures: vec![AggIntent::FrequencyL2 {
col: Some(*key),
accuracy: accuracy.clone(),
}],
output_names: vec!["frequency_l2".into()],
filters: vec![],
having: None,
child: Rc::clone(child),
}))
.ok()?;
// L2 is positive for any nonempty unit-update population. Restore SQL SUM's
// NULL on an empty relation without introducing another count computation.
let rewritten = OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Project {
cols: vec![ProjectItem {
alias: Some(root.schema.fields.first()?.name.clone()),
expr: ScalarExpr::Case {
operand: None,
branches: vec![(
ScalarExpr::Compare {
left: Box::new(ScalarExpr::Column(0)),
op: CompareOpKind::Eq,
right: Box::new(ScalarExpr::Literal(ScalarValue::Float64(0.0))),
semantics: ExprSemantics::Sql,
},
ScalarExpr::Cast {
expr: Box::new(ScalarExpr::Literal(ScalarValue::Null)),
to: DataType::Float64,
try_cast: false,
},
)],
else_expr: Some(Box::new(ScalarExpr::Column(0))),
},
}],
qualifier: qualifier.clone(),
child: aggregate,
}))
.ok()?;
(root.schema == rewritten.schema).then_some(rewritten)
}
1 change: 1 addition & 0 deletions crates/asap-aware-mapping/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,7 @@ pub mod empirical_resources;
pub mod erp;
pub mod exact_composition;
pub mod explanation;
mod frequency_rewrite;
mod function_rules;
pub mod grouping;
pub mod pane_sharing;
Expand Down
17 changes: 17 additions & 0 deletions crates/asap-aware-mapping/src/rewrite.rs
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,14 @@
//! that reshaping — see "Non-goals" below for why it does not also decide
//! whether the reshaping is worth it.
//!
//! ## SQL frequency recognition
//!
//! A floating-point `SQRT(SUM(c*c))` over grouped unit counts exposes a
//! frequency L2 alternative through the same semantic strategy. Projection
//! lineage, predicates, accuracy and empty-input NULL are preserved. Integer
//! products and nullable grouping keys are excluded because overflow and NULL
//! groups have observable SQL behavior. See `frequency_rewrite` for the rule.
//!
//! ## Scope
//!
//! Ordinary `by(...)` averages use a schema-preserving projection. Temporal
Expand Down Expand Up @@ -399,9 +407,18 @@ impl ReplacementStrategy for SemanticEquivalentRewriteStrategy {
fn matches(&self, target: &TargetSubDAG<'_>) -> bool {
avg_rewrite_target(target.root).is_some()
|| composed_aggregate_rewrite(target.root).is_some()
|| crate::frequency_rewrite::frequency_l2_rewrite(target.root).is_some()
}

fn replacements(&self, target: &TargetSubDAG<'_>) -> Vec<ReplacementSubDAG> {
if let Some(rewritten) = crate::frequency_rewrite::frequency_l2_rewrite(target.root) {
return vec![ReplacementSubDAG {
strategy: "SemanticEquivalentRewriteStrategy",
replacement: Replacement::SubDAG(rewritten),
provenance: crate::replacement::ReplacementProvenance::LogicalRewrite,
rationale: "recognize a floating SQL frequency L2 product while preserving empty-input NULL and the original exact candidate".into(),
}];
}
if let Some(rewritten) = composed_aggregate_rewrite(target.root) {
return vec![ReplacementSubDAG {
strategy: "SemanticEquivalentRewriteStrategy",
Expand Down
24 changes: 23 additions & 1 deletion crates/asap-physical-operators/src/expressions/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,16 @@ pub(super) fn evaluate(
.collect::<Result<Vec<_>, _>>()?;
return Ok(Value::Float64(promql_function(name, &values)?));
}
if name.eq_ignore_ascii_case("sqrt") {
return match evaluate(&args[0], row, schema)? {
Value::Null => Ok(Value::Null),
Value::Float64(value) => Ok(Value::Float64(value.sqrt())),
Value::Int64(value) => Ok(Value::Float64((value as f64).sqrt())),
_ => Err(Error::Invalid(
"SQL sqrt requires a numeric argument".into(),
)),
};
}
if name == "promql_drop_metric_name" {
let Value::Utf8(encoded) = evaluate(&args[0], row, schema)? else {
return Err(Error::Invalid("series identity must be Utf8".into()));
Expand Down Expand Up @@ -645,7 +655,19 @@ fn validate(expr: &ScalarExpr, schema: &planner_types::pre_asap::Schema) -> Resu
Ok(())
}
ScalarExpr::FunctionCall { name, args } => {
if name != "promql_drop_metric_name"
if name.eq_ignore_ascii_case("sqrt") {
if args.len() != 1
|| !matches!(
args[0]
.scalar_type(schema)
.map_err(|e| Error::Invalid(e.to_string()))?
.0,
DataType::Int64 | DataType::Float64 | DataType::Null
)
{
return Err(invalid());
}
} else if name != "promql_drop_metric_name"
&& planner_types::pre_asap::scalar_type_rules::promql_function_arity(name).is_none()
&& name != "asap_struct_field"
&& name != "asap_element_access"
Expand Down
30 changes: 30 additions & 0 deletions crates/asap-physical-operators/tests/physical_semantics.rs
Original file line number Diff line number Diff line change
Expand Up @@ -831,3 +831,33 @@ fn exact_frequency_grouping_and_entropy_bits() {
}
assert!(unary(input, vec![], operator).is_empty());
}

// SQL SQRT propagates NULL and accepts numeric inputs with a floating result.
#[test]
fn sql_sqrt_executes_numeric_and_null_arguments() {
for (dtype, value, expected) in [
(DataType::Int64, Value::Int64(9), 3.0),
(DataType::Float64, Value::Float64(2.25), 1.5),
] {
let input = schema(&[("v", dtype, true)]);
let expression = QueryExpr::FunctionCall {
name: "sqrt".into(),
args: vec![QueryExpr::Column(0)],
};
let compiled = CompiledExpression::compile(&expression, &input).unwrap();
assert!(matches!(compiled.evaluate(&[value]).unwrap(), Value::Float64(v) if v == expected));
assert!(matches!(
compiled.evaluate(&[Value::Null]).unwrap(),
Value::Null
));
}
let input = schema(&[("v", DataType::Float64, false)]);
let expression = QueryExpr::FunctionCall {
name: "sqrt".into(),
args: vec![QueryExpr::Column(0)],
};
let compiled = CompiledExpression::compile(&expression, &input).unwrap();
assert!(
matches!(compiled.evaluate(&[Value::Float64(-1.0)]).unwrap(), Value::Float64(v) if v.is_nan())
);
}
Loading
Loading