From 745f7f83dabd25372eccc4ec6aec9b2abc4d2a65 Mon Sep 17 00:00:00 2001 From: Nikhil Sinha Date: Tue, 25 Aug 2026 18:10:02 +0530 Subject: [PATCH 1/3] fix aggregate evaluation for alert --- src/alerts/alerts_utils.rs | 282 +++++++++++++++++++++++-------------- src/alerts/mod.rs | 118 +++++++++++++++- 2 files changed, 295 insertions(+), 105 deletions(-) diff --git a/src/alerts/alerts_utils.rs b/src/alerts/alerts_utils.rs index 919dd412a..4d521922e 100644 --- a/src/alerts/alerts_utils.rs +++ b/src/alerts/alerts_utils.rs @@ -19,7 +19,8 @@ use std::{collections::HashMap, fmt::Display}; use actix_web::{Either, http::header::HeaderMap}; -use arrow_array::{Array, Float64Array, Int64Array, RecordBatch}; +use arrow::{compute::cast, datatypes::DataType, util::display::array_value_to_string}; +use arrow_array::{Array, Float64Array, RecordBatch}; use datafusion::{ logical_expr::{Literal, LogicalPlan}, prelude::{Expr, lit}, @@ -30,7 +31,7 @@ use crate::{ alerts::{ AlertTrait, LogicalOperator, WhereConfigOperator, alert_structs::{AlertQueryResult, ConditionConfig, Conditions, GroupResult}, - extract_aggregate_aliases, + resolve_alert_output_layout, }, handlers::http::{ cluster::send_query_request, @@ -127,7 +128,7 @@ async fn execute_local_query( } }; - Ok(extract_group_results(records, raw_logical_plan)) + extract_group_results(records, raw_logical_plan) } /// Execute alert query remotely (Prism mode) @@ -172,52 +173,45 @@ fn convert_result_to_group_results( .as_array() .ok_or_else(|| AlertError::CustomError("Expected array in query result".to_string()))?; - let aggregate_aliases = extract_aggregate_aliases(&plan); + let layout = resolve_alert_output_layout(&plan)?; - if array_val.is_empty() || aggregate_aliases.is_empty() { + if array_val.is_empty() { return Ok(AlertQueryResult { groups: vec![], - is_simple_query: true, + is_simple_query: layout.dimension_indices.is_empty(), }); } - // take the first entry and extract the column name / alias - let (agg_condition, alias) = &aggregate_aliases[0]; - - let aggregate_key = if let Some(alias) = alias { - alias - } else { - agg_condition - }; - - // Find the aggregate column from the first row - let first_row = array_val[0] - .as_object() - .ok_or_else(|| AlertError::CustomError("Expected object in query result".to_string()))?; - - let is_simple_query = first_row.len() == 1; + let aggregate_key = &layout.measure_name; + let is_simple_query = layout.dimension_indices.is_empty(); let mut groups = Vec::new(); // Process each row as a separate group for row in array_val { if let Some(object) = row.as_object() { - let mut group_values = HashMap::new(); - let mut aggregate_value = 0.0; - - for (key, value) in object { - if key == aggregate_key { - aggregate_value = value.as_f64().ok_or_else(|| { - AlertError::CustomError(format!( - "Non-numeric value found in aggregate column '{}'", - aggregate_key - )) - })?; - } else { - // This is a GROUP BY column - group_values - .insert(key.clone(), value.to_string().trim_matches('"').to_string()); + let aggregate_value = match object.get(aggregate_key) { + Some(serde_json::Value::Null) => 0.0, + Some(value) => value.as_f64().ok_or_else(|| { + AlertError::CustomError(format!( + "Non-numeric value found in aggregate column '{aggregate_key}'" + )) + })?, + None => { + return Err(AlertError::CustomError(format!( + "Aggregate column '{aggregate_key}' missing from query result" + ))); } - } + }; + let group_values = object + .iter() + .filter(|(key, _)| *key != aggregate_key) + .map(|(key, value)| { + ( + key.clone(), + value.to_string().trim_matches('"').to_string(), + ) + }) + .collect(); groups.push(GroupResult { group_values, @@ -232,38 +226,12 @@ fn convert_result_to_group_results( }) } -/// Extract numeric value from an Arrow array at the given row index -fn extract_numeric_value(column: &dyn Array, row_index: usize) -> f64 { - if let Some(float_array) = column.as_any().downcast_ref::() { - if !float_array.is_null(row_index) { - return float_array.value(row_index); - } - } else if let Some(int_array) = column.as_any().downcast_ref::() - && !int_array.is_null(row_index) - { - return int_array.value(row_index) as f64; - } - 0.0 -} - /// Extract string value from an Arrow array at the given row index -fn extract_string_value(column: &dyn Array, row_index: usize) -> String { - use arrow_array::StringArray; - - if let Some(string_array) = column.as_any().downcast_ref::() { - if !string_array.is_null(row_index) { - return string_array.value(row_index).to_string(); - } - } else if let Some(int_array) = column.as_any().downcast_ref::() { - if !int_array.is_null(row_index) { - return int_array.value(row_index).to_string(); - } - } else if let Some(float_array) = column.as_any().downcast_ref::() - && !float_array.is_null(row_index) - { - return float_array.value(row_index).to_string(); +fn extract_string_value(column: &dyn Array, row_index: usize) -> Result { + if column.is_null(row_index) { + return Ok("null".to_string()); } - "null".to_string() + Ok(array_value_to_string(column, row_index)?) } pub fn evaluate_condition(operator: &AlertOperator, actual: f64, expected: f64) -> bool { @@ -327,51 +295,49 @@ async fn update_alert_state( } /// Extract group results from record batches, supporting both simple and GROUP BY queries -fn extract_group_results(records: Vec, plan: LogicalPlan) -> AlertQueryResult { +fn extract_group_results( + records: Vec, + plan: LogicalPlan, +) -> Result { trace!("records-\n{records:?}"); - let aggregate_aliases = extract_aggregate_aliases(&plan); + let layout = resolve_alert_output_layout(&plan)?; - // since there is going to be only one aggregate, we'll check if it is empty - if aggregate_aliases.is_empty() || records.is_empty() { - return AlertQueryResult { + if records.is_empty() { + return Ok(AlertQueryResult { groups: vec![], - is_simple_query: true, - }; + is_simple_query: layout.dimension_indices.is_empty(), + }); } - // take the first entry and extract the column name / alias - let (agg_condition, alias) = &aggregate_aliases[0]; - - let alias = if let Some(alias) = alias { - alias - } else { - agg_condition - }; - - let first_batch = &records[0]; - let schema = first_batch.schema(); - - // Determine if this is a simple query (no GROUP BY) or a grouped query - let is_simple_query = schema.fields().len() == 1; + let is_simple_query = layout.dimension_indices.is_empty(); let mut groups = Vec::new(); for batch in &records { + let schema = batch.schema(); + let measure_values = cast(batch.column(layout.measure_index), &DataType::Float64)?; + let measure_values = measure_values + .as_any() + .downcast_ref::() + .ok_or_else(|| { + AlertError::CustomError("Failed to cast alert value to Float64".into()) + })?; for row_index in 0..batch.num_rows() { let mut group_values = HashMap::new(); - let mut aggregate_value = 0.0; - - // Extract values for each column - for (col_index, field) in schema.fields().iter().enumerate() { - let column = batch.column(col_index); - if field.name().eq(alias) { - aggregate_value = extract_numeric_value(column, row_index) - } else { - // This is a GROUP BY column - let value = extract_string_value(column, row_index); - group_values.insert(field.name().clone(), value); - } + let aggregate_value = if measure_values.is_null(row_index) { + 0.0 + } else { + measure_values.value(row_index) + }; + + for dimension_index in &layout.dimension_indices { + let field = schema.field(*dimension_index); + let value = extract_string_value( + batch.column(*dimension_index).as_ref(), + row_index, + )?; + group_values.insert(field.name().clone(), value); } groups.push(GroupResult { @@ -381,10 +347,10 @@ fn extract_group_results(records: Vec, plan: LogicalPlan) -> AlertQ } } - AlertQueryResult { + Ok(AlertQueryResult { groups, is_simple_query, - } + }) } pub fn get_filter_string(where_clause: &Conditions) -> Result { @@ -713,7 +679,117 @@ impl Display for ValueType { #[cfg(test)] mod tests { use super::*; - use crate::alerts::WhereConfigOperator; + use crate::alerts::{WhereConfigOperator, resolve_alert_output_layout}; + use datafusion::prelude::SessionContext; + + const WRAPPED_AGGREGATE_QUERY: &str = r#" + SELECT + app, + ROUND(SUM(cost), 4) AS estimated_spend_usd + FROM ( + VALUES + ('Codex', CAST(1.23456 AS DOUBLE)), + ('Codex', CAST(2.0 AS DOUBLE)), + ('Claude Code', CAST(4.5 AS DOUBLE)) + ) AS usage(app, cost) + GROUP BY app + ORDER BY app + "#; + + #[tokio::test] + async fn wrapped_aggregate_is_resolved_by_lineage() { + let context = SessionContext::new(); + let dataframe = context.sql(WRAPPED_AGGREGATE_QUERY).await.unwrap(); + let plan = dataframe.logical_plan().clone(); + let layout = resolve_alert_output_layout(&plan).unwrap(); + + assert_eq!(layout.measure_index, 1); + assert_eq!(layout.measure_name, "estimated_spend_usd"); + assert_eq!(layout.dimension_indices, [0]); + + let result = extract_group_results(dataframe.collect().await.unwrap(), plan).unwrap(); + assert!(!result.is_simple_query); + assert_eq!(result.groups.len(), 2); + assert_eq!(result.groups[0].group_values["app"], "Claude Code"); + assert_eq!(result.groups[0].aggregate_value, 4.5); + assert_eq!(result.groups[1].group_values["app"], "Codex"); + assert!((result.groups[1].aggregate_value - 3.2346).abs() < f64::EPSILON); + } + + #[tokio::test] + async fn wrapped_aggregate_uses_final_alias_for_remote_results() { + let context = SessionContext::new(); + let plan = context + .state() + .create_logical_plan(WRAPPED_AGGREGATE_QUERY) + .await + .unwrap(); + let result = convert_result_to_group_results( + serde_json::json!([ + {"app": "Codex", "estimated_spend_usd": 3.2346} + ]), + plan, + ) + .unwrap(); + + assert_eq!(result.groups[0].group_values["app"], "Codex"); + assert_eq!(result.groups[0].aggregate_value, 3.2346); + } + + #[tokio::test] + async fn aggregate_lineage_survives_nested_scalar_expressions() { + let context = SessionContext::new(); + let plan = context + .state() + .create_logical_plan( + r#" + SELECT + app, + CAST( + CASE + WHEN SUM(cost) > 0 + THEN ABS(SUM(cost) * 2) + 1 + ELSE 0 + END + AS DOUBLE + ) AS score + FROM ( + VALUES ('Codex', CAST(1.0 AS DOUBLE)) + ) AS usage(app, cost) + GROUP BY app + "#, + ) + .await + .unwrap(); + + let layout = resolve_alert_output_layout(&plan).unwrap(); + assert_eq!(layout.measure_name, "score"); + assert_eq!(layout.dimension_indices, [0]); + } + + #[tokio::test] + async fn multiple_aggregate_derived_outputs_are_rejected() { + let context = SessionContext::new(); + let plan = context + .state() + .create_logical_plan( + r#" + SELECT + app, + SUM(cost) AS raw_spend, + ROUND(SUM(cost), 2) AS rounded_spend + FROM ( + VALUES ('Codex', CAST(1.0 AS DOUBLE)) + ) AS usage(app, cost) + GROUP BY app + "#, + ) + .await + .unwrap(); + + let error = resolve_alert_output_layout(&plan).unwrap_err(); + assert!(error.to_string().contains("found 2")); + } // ------------------------------------------------------------------------- // sanitize_array_elements — happy-path cases diff --git a/src/alerts/mod.rs b/src/alerts/mod.rs index b8e1a8921..5af305d35 100644 --- a/src/alerts/mod.rs +++ b/src/alerts/mod.rs @@ -21,7 +21,7 @@ use actix_web::http::header::ContentType; use arrow_schema::{ArrowError, DataType, Schema}; use async_trait::async_trait; use chrono::Utc; -use datafusion::logical_expr::{LogicalPlan, Projection}; +use datafusion::logical_expr::{LogicalPlan, Projection, utils::find_aggregate_exprs}; use datafusion::prelude::Expr; use datafusion::sql::sqlparser::parser::ParserError; use derive_more::FromStrError; @@ -854,7 +854,121 @@ pub async fn get_number_of_agg_exprs( .map_err(|err| AlertError::CustomError(format!("Failed to parse query: {err}")))?; // Check if the plan structure indicates an aggregate query - get_number_of_agg_exprs_from_plan(&logical_plan) + let aggregate_count = get_number_of_agg_exprs_from_plan(&logical_plan)?; + if aggregate_count == 1 { + resolve_alert_output_layout(&logical_plan)?; + } + Ok(aggregate_count) +} + +/// Output columns used when evaluating an alert query. +/// +/// Exactly one final output column must derive from the query's aggregate +/// expression. Every other output column is treated as a group dimension. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AlertOutputLayout { + pub measure_index: usize, + pub measure_name: String, + pub dimension_indices: Vec, +} + +/// Resolve the final aggregate-derived output without relying on function or +/// alias names. Aggregate lineage is propagated through arbitrary projection +/// expressions, including scalar functions, UDFs, casts, arithmetic, and CASE. +pub fn resolve_alert_output_layout(plan: &LogicalPlan) -> Result { + let measure_indices = aggregate_output_indices(plan); + let [measure_index] = measure_indices.as_slice() else { + return Err(AlertError::InvalidAlertQuery(format!( + "Alert query must return exactly one aggregate-derived value, found {}", + measure_indices.len() + ))); + }; + + let measure_field = plan.schema().field(*measure_index); + if !measure_field.data_type().is_numeric() { + return Err(AlertError::InvalidAlertQuery(format!( + "Alert aggregate-derived value '{}' must be numeric, found {}", + measure_field.name(), + measure_field.data_type() + ))); + } + + Ok(AlertOutputLayout { + measure_index: *measure_index, + measure_name: measure_field.name().clone(), + dimension_indices: (0..plan.schema().fields().len()) + .filter(|index| index != measure_index) + .collect(), + }) +} + +fn aggregate_output_indices(plan: &LogicalPlan) -> Vec { + let mut indices = match plan { + LogicalPlan::Aggregate(aggregate) => { + (aggregate.group_expr.len()..aggregate.schema.fields().len()).collect() + } + LogicalPlan::Projection(projection) => { + let input_indices = aggregate_output_indices(&projection.input); + projection + .expr + .iter() + .enumerate() + .filter_map(|(index, expr)| { + expression_depends_on_aggregate( + expr, + projection.input.schema(), + &input_indices, + ) + .then_some(index) + }) + .collect() + } + LogicalPlan::Window(window) => { + let mut indices = aggregate_output_indices(&window.input); + let input_len = window.input.schema().fields().len(); + for (index, expr) in window.window_expr.iter().enumerate() { + if expression_depends_on_aggregate(expr, window.input.schema(), &indices) { + indices.push(input_len + index); + } + } + indices + } + LogicalPlan::Union(union) => union + .inputs + .iter() + .flat_map(|input| aggregate_output_indices(input)) + .collect(), + _ => { + let inputs = plan.inputs(); + if let [input] = inputs.as_slice() + && input.schema().fields().len() == plan.schema().fields().len() + { + aggregate_output_indices(input) + } else { + Vec::new() + } + } + }; + indices.sort_unstable(); + indices.dedup(); + indices +} + +fn expression_depends_on_aggregate( + expr: &Expr, + input_schema: &datafusion::common::DFSchema, + aggregate_indices: &[usize], +) -> bool { + if !find_aggregate_exprs([expr]).is_empty() { + return true; + } + + expr.column_refs().iter().any(|column| { + input_schema + .maybe_index_of_column(column) + .or_else(|| input_schema.index_of_column_by_name(None, &column.name)) + .is_some_and(|index| aggregate_indices.contains(&index)) + }) } /// Extract the projection which deals with aggregation From 9f86b91f759152faeda4498ec8d186f41de1e143 Mon Sep 17 00:00:00 2001 From: Nikhil Sinha Date: Tue, 25 Aug 2026 18:15:13 +0530 Subject: [PATCH 2/3] fmt --- src/alerts/alerts_utils.rs | 13 +++---------- src/alerts/mod.rs | 8 ++------ 2 files changed, 5 insertions(+), 16 deletions(-) diff --git a/src/alerts/alerts_utils.rs b/src/alerts/alerts_utils.rs index 4d521922e..837e16938 100644 --- a/src/alerts/alerts_utils.rs +++ b/src/alerts/alerts_utils.rs @@ -205,12 +205,7 @@ fn convert_result_to_group_results( let group_values = object .iter() .filter(|(key, _)| *key != aggregate_key) - .map(|(key, value)| { - ( - key.clone(), - value.to_string().trim_matches('"').to_string(), - ) - }) + .map(|(key, value)| (key.clone(), value.to_string().trim_matches('"').to_string())) .collect(); groups.push(GroupResult { @@ -333,10 +328,8 @@ fn extract_group_results( for dimension_index in &layout.dimension_indices { let field = schema.field(*dimension_index); - let value = extract_string_value( - batch.column(*dimension_index).as_ref(), - row_index, - )?; + let value = + extract_string_value(batch.column(*dimension_index).as_ref(), row_index)?; group_values.insert(field.name().clone(), value); } diff --git a/src/alerts/mod.rs b/src/alerts/mod.rs index 5af305d35..5c0313527 100644 --- a/src/alerts/mod.rs +++ b/src/alerts/mod.rs @@ -914,12 +914,8 @@ fn aggregate_output_indices(plan: &LogicalPlan) -> Vec { .iter() .enumerate() .filter_map(|(index, expr)| { - expression_depends_on_aggregate( - expr, - projection.input.schema(), - &input_indices, - ) - .then_some(index) + expression_depends_on_aggregate(expr, projection.input.schema(), &input_indices) + .then_some(index) }) .collect() } From ad2fb785b77543879b7e1a4abb284b4ec9f227d2 Mon Sep 17 00:00:00 2001 From: Nikhil Sinha Date: Tue, 25 Aug 2026 18:34:15 +0530 Subject: [PATCH 3/3] fix coderabbit comments --- src/alerts/alerts_utils.rs | 53 +++++++++++++++++++++++++++++++++++++- src/alerts/mod.rs | 6 ++++- 2 files changed, 57 insertions(+), 2 deletions(-) diff --git a/src/alerts/alerts_utils.rs b/src/alerts/alerts_utils.rs index 837e16938..611fd458d 100644 --- a/src/alerts/alerts_utils.rs +++ b/src/alerts/alerts_utils.rs @@ -205,7 +205,13 @@ fn convert_result_to_group_results( let group_values = object .iter() .filter(|(key, _)| *key != aggregate_key) - .map(|(key, value)| (key.clone(), value.to_string().trim_matches('"').to_string())) + .map(|(key, value)| { + let rendered = match value { + serde_json::Value::String(text) => text.clone(), + other => other.to_string(), + }; + (key.clone(), rendered) + }) .collect(); groups.push(GroupResult { @@ -729,6 +735,26 @@ mod tests { assert_eq!(result.groups[0].aggregate_value, 3.2346); } + #[tokio::test] + async fn remote_string_dimensions_are_not_json_escaped() { + let context = SessionContext::new(); + let plan = context + .state() + .create_logical_plan(WRAPPED_AGGREGATE_QUERY) + .await + .unwrap(); + let app = r#"he said "hi" at C:\temp"#; + let result = convert_result_to_group_results( + serde_json::json!([ + {"app": app, "estimated_spend_usd": 3.2346} + ]), + plan, + ) + .unwrap(); + + assert_eq!(result.groups[0].group_values["app"], app); + } + #[tokio::test] async fn aggregate_lineage_survives_nested_scalar_expressions() { let context = SessionContext::new(); @@ -760,6 +786,31 @@ mod tests { assert_eq!(layout.dimension_indices, [0]); } + #[tokio::test] + async fn rollup_internal_grouping_id_is_not_a_measure() { + let context = SessionContext::new(); + let plan = context + .state() + .create_logical_plan( + r#" + SELECT + region, + app, + ROUND(SUM(cost), 2) AS spend + FROM ( + VALUES ('us-east', 'Codex', CAST(1.0 AS DOUBLE)) + ) AS usage(region, app, cost) + GROUP BY ROLLUP(region, app) + "#, + ) + .await + .unwrap(); + + let layout = resolve_alert_output_layout(&plan).unwrap(); + assert_eq!(layout.measure_name, "spend"); + assert_eq!(layout.dimension_indices, [0, 1]); + } + #[tokio::test] async fn multiple_aggregate_derived_outputs_are_rejected() { let context = SessionContext::new(); diff --git a/src/alerts/mod.rs b/src/alerts/mod.rs index 5c0313527..6e9422af1 100644 --- a/src/alerts/mod.rs +++ b/src/alerts/mod.rs @@ -902,10 +902,13 @@ pub fn resolve_alert_output_layout(plan: &LogicalPlan) -> Result Vec { let mut indices = match plan { LogicalPlan::Aggregate(aggregate) => { - (aggregate.group_expr.len()..aggregate.schema.fields().len()).collect() + let schema_len = aggregate.schema.fields().len(); + let aggregate_start = schema_len - aggregate.aggr_expr.len(); + (aggregate_start..schema_len).collect() } LogicalPlan::Projection(projection) => { let input_indices = aggregate_output_indices(&projection.input); @@ -950,6 +953,7 @@ fn aggregate_output_indices(plan: &LogicalPlan) -> Vec { indices } +/// Check whether an expression directly contains or references an aggregate value. fn expression_depends_on_aggregate( expr: &Expr, input_schema: &datafusion::common::DFSchema,