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
25 changes: 16 additions & 9 deletions crates/executor/src/expressions/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -127,15 +127,22 @@ 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.eq_ignore_ascii_case("sqrt") || name.eq_ignore_ascii_case("ln") {
let value = match evaluate(&args[0], row, schema)? {
Value::Null => return Ok(Value::Null),
Value::Float64(value) => value,
Value::Int64(value) => value as f64,
_ => {
return Err(Error::Invalid(
"SQL math function requires a numeric argument".into(),
))
}
};
return Ok(Value::Float64(if name.eq_ignore_ascii_case("sqrt") {
value.sqrt()
} else {
value.ln()
}));
}
if name == "promql_drop_metric_name" {
let Value::Utf8(encoded) = evaluate(&args[0], row, schema)? else {
Expand Down Expand Up @@ -656,7 +663,7 @@ fn validate(expr: &ScalarExpr, schema: &planner_types::ir::schema::Schema) -> Re
Ok(())
}
ScalarExpr::FunctionCall { name, args } => {
if name.eq_ignore_ascii_case("sqrt") {
if name.eq_ignore_ascii_case("sqrt") || name.eq_ignore_ascii_case("ln") {
if args.len() != 1
|| !matches!(
args[0]
Expand Down
35 changes: 35 additions & 0 deletions crates/executor/src/operators/aggregate/mod.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,20 @@
use super::*;
impl Operator {
/// A SQL SUM window over the complete, unordered input relation.
pub fn sql_window_sum(input: SchemaRef, column: usize, name: String) -> Result<Self, Error> {
let (dtype, _) = plain(&input, column)?;
if !matches!(dtype, DataType::Int64 | DataType::Float64) {
return Err(invalid("SQL window SUM requires a numeric column"));
}
let mut output = (*input).clone();
output.fields.push(result_field(&name, dtype.clone(), true));
Ok(Self {
kind: Kind::SQLWindowSum { column },
inputs: vec![input],
output: Arc::new(output),
})
}

pub fn aggregate(
input: SchemaRef,
groups: Vec<usize>,
Expand Down Expand Up @@ -171,6 +186,26 @@ pub(super) fn execute<'a>(
Ok(futures::stream::once(async move {
let (rows, _memory) = collect_rows(input, &context).await?;
let result = match &operator.kind {
Kind::SQLWindowSum { column } => {
let mut work = Cooperative::new(&context);
let total = reduce_one(
&rows,
&Reduction::Sum(*column),
&operator.inputs[0],
&mut work,
&context,
)
.await?;
let mut workspace = Workspace::new(&context)?;
let mut result = Vec::with_capacity(rows.len());
for mut row in rows {
work.checkpoint().await?;
row.push(total.clone());
workspace.grow(row_bytes(&row))?;
result.push(row);
}
result
}
Kind::Window {
intent,
coordinate,
Expand Down
7 changes: 6 additions & 1 deletion crates/executor/src/operators/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,9 @@ enum Kind {
groups: Vec<usize>,
window: Option<(i64, i64)>,
},
SQLWindowSum {
column: usize,
},
Aggregate {
groups: Vec<usize>,
measures: Vec<Reduction>,
Expand Down Expand Up @@ -281,6 +284,7 @@ impl PhysicalOperator<Batch, SchemaRef> for Operator {
| Kind::SeriesBinary { .. }
| Kind::SeriesHistogramQuantile { .. }
| Kind::SeriesRelabel { .. }
| Kind::SQLWindowSum { .. }
| Kind::Aggregate { .. }
| Kind::Window { .. }
| Kind::Join { .. }
Expand Down Expand Up @@ -333,6 +337,7 @@ impl PhysicalOperator<Batch, SchemaRef> for Operator {
Kind::Filter(_) => "Filter",
Kind::Limit { .. } => "Limit",
Kind::Sort { .. } => "Sort",
Kind::SQLWindowSum { .. } => "SQLWindowSum",
Kind::Aggregate { .. } => "Aggregate",
Kind::Window { .. } => "WindowAggregate",
Kind::SemiJoin { .. } => "SemiJoin",
Expand Down Expand Up @@ -384,7 +389,7 @@ impl PhysicalOperator<Batch, SchemaRef> for Operator {
Kind::Filter(_) => filter::execute(self, inputs, context),
Kind::Limit { .. } => limit::execute(self, inputs, context),
Kind::Sort { .. } => sort::execute(self, inputs, context),
Kind::Window { .. } | Kind::Aggregate { .. } => {
Kind::SQLWindowSum { .. } | Kind::Window { .. } | Kind::Aggregate { .. } => {
aggregate::execute(self, inputs, context)
}
Kind::Join { .. } | Kind::SemiJoin { .. } => joins::execute(self, inputs, context),
Expand Down
9 changes: 9 additions & 0 deletions crates/executor/src/operators/unchecked.rs
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,15 @@ impl TryFrom<UncheckedOperator> for Operator {
groups,
window,
} => Operator::window(input(0)?, *intent, coordinate, value, groups, window)?,
Kind::SQLWindowSum { column } => {
let name = output
.fields
.last()
.ok_or_else(|| invalid("SQL window SUM output missing"))?
.name
.clone();
Operator::sql_window_sum(input(0)?, column, name)?
}
Kind::Aggregate { groups, measures } => {
if groups.len() + measures.len() != output.fields.len() {
return Err(invalid("aggregate width mismatch"));
Expand Down
33 changes: 33 additions & 0 deletions crates/executor/src/physical_planner/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -820,6 +820,39 @@ fn bind_operation(node: &PhysicalASAPDAGNode, inputs: &[SchemaRef]) -> Result<Op
*offset as u64,
groups(input, partition_by)?,
),
NonASAPOpKind::SQLWindowFunc {
func: planner_types::ir::operator::WindowFuncKind::Sum,
args,
partition_by,
order_by,
frame: Some(frame),
output_name,
} => {
use planner_types::ir::{
operator::{WindowFrameBound, WindowFrameOffset},
scalar::ScalarValue,
};
let [WireScalarExpr::Column(column)] = args.as_slice() else {
return Err(invalid("SQL window SUM requires one column"));
};
if partition_by.is_without()
|| !partition_by.keys().is_empty()
|| !order_by.is_empty()
|| !matches!(
frame.start_bound,
WindowFrameBound::Preceding(WindowFrameOffset::Scalar(ScalarValue::Null))
)
|| !matches!(
frame.end_bound,
WindowFrameBound::Following(WindowFrameOffset::Scalar(ScalarValue::Null))
)
{
return Err(invalid(
"native SQL window SUM requires the complete unordered relation",
));
}
Operator::sql_window_sum(input.clone(), *column, output_name.clone())
}
NonASAPOpKind::Aggregate {
reduction,
measures,
Expand Down
16 changes: 16 additions & 0 deletions crates/executor/tests/blocking_resources.rs
Original file line number Diff line number Diff line change
Expand Up @@ -291,3 +291,19 @@ fn frequency_dictionary_enforces_memory_budget() {
assert_eq!(run.retained_bytes(), 0);
}
}

// A complete SUM window accounts for its expanded output and frees memory on failure.
#[test]
fn complete_window_sum_enforces_workspace_budget() {
let sources = source(64);
let run = context(12_000);
let inputs = sources.execute(&[0], run.clone()).unwrap();
let operator = Operator::sql_window_sum(schema(1), 0, "total".into()).unwrap();
let mut output = operator.start(inputs, run.clone()).unwrap();
assert!(matches!(
block_on(output.next()),
Some(Err(Error::MemoryLimit))
));
drop(output);
assert_eq!(run.retained_bytes(), 0);
}
56 changes: 56 additions & 0 deletions crates/executor/tests/physical_semantics.rs
Original file line number Diff line number Diff line change
Expand Up @@ -945,3 +945,59 @@ fn sql_sqrt_executes_numeric_and_null_arguments() {
matches!(compiled.evaluate(&[Value::Float64(-1.0)]).unwrap(), Value::Float64(v) if v.is_nan())
);
}

// A complete SQL SUM window keeps every row and appends one nullable total, including recovery.
#[test]
fn complete_sql_sum_window_preserves_rows_and_nulls() {
let input = schema(&[("v", DataType::Int64, true)]);
let operator = Operator::sql_window_sum(input.clone(), 0, "total".into()).unwrap();
let operator: Operator =
serde_json::from_slice(&serde_json::to_vec(&operator).unwrap()).unwrap();
for (rows, expected) in [
(vec![], None),
(vec![vec![Value::Null]], None),
(
vec![
vec![Value::Int64(1)],
vec![Value::Null],
vec![Value::Int64(3)],
],
Some(4),
),
] {
let original = rows.clone();
let actual = unary(input.clone(), vec![rows], operator.clone());
assert_eq!(actual.len(), original.len());
for (row, original) in actual.iter().zip(original) {
assert_eq!(row[0].key().unwrap(), original[0].key().unwrap());
match (&row[1], expected) {
(Value::Null, None) => {}
(Value::Int64(value), Some(expected)) => assert_eq!(*value, expected),
other => panic!("wrong complete-window sum: {other:?}"),
}
}
}
}

// SQL LN preserves nullable numeric signatures and natural-log units.
#[test]
fn sql_ln_executes_numeric_and_null_arguments() {
for (dtype, value) in [
(DataType::Int64, Value::Int64(2)),
(DataType::Float64, Value::Float64(2.0)),
] {
let input = schema(&[("v", dtype, true)]);
let expression = ScalarExpr::FunctionCall {
name: "ln".into(),
args: vec![ScalarExpr::Column(0)],
};
let compiled = CompiledExpression::compile(&expression, &input).unwrap();
assert!(
matches!(compiled.evaluate(&[value]).unwrap(), Value::Float64(v) if v == std::f64::consts::LN_2)
);
assert!(matches!(
compiled.evaluate(&[Value::Null]).unwrap(),
Value::Null
));
}
}
Loading