diff --git a/asap-planner-rs/src/optimizer/aqe_extractor.rs b/asap-planner-rs/src/optimizer/aqe_extractor.rs index 3e4a2d48..15ed2064 100644 --- a/asap-planner-rs/src/optimizer/aqe_extractor.rs +++ b/asap-planner-rs/src/optimizer/aqe_extractor.rs @@ -2,6 +2,8 @@ use std::collections::HashMap; use asap_types::query_requirements::{build_query_requirements_promql, QueryRequirements}; use asap_types::PromQLSchema; +use promql_parser::parser::token::{self, TokenType}; +use promql_parser::parser::{AggregateExpr, Call, Expr, Function}; use promql_utilities::data_model::KeyByLabelNames; use promql_utilities::query_logics::enums::Statistic; @@ -138,10 +140,11 @@ pub(super) fn gcd(a: u64, b: u64) -> u64 { /// into their arms. Scalar arms (e.g. the `100` in `rate(x[5m]) * 100`) are /// dropped — they contribute no AQE. Only arithmetic operators are split; /// comparison and set operators are left as-is (treated as opaque leaves). +/// An `avg` leaf becomes its sum and count leaves (see `rewrite_avg`). fn decompose_to_leaves(query: &str) -> Vec { let (lhs, rhs) = match parse_binary_arms(query) { Some(arms) => arms, - None => return vec![query.to_string()], + None => return rewrite_avg(query), }; let mut leaves = Vec::new(); @@ -153,6 +156,42 @@ fn decompose_to_leaves(query: &str) -> Vec { leaves } +/// Rewrite an `avg` leaf into the sum and count leaves it is computed from: +/// `avg by (l) (x)` → `sum by (l) (x)`, `count by (l) (x)`, and +/// `avg_over_time(x[5m])` → `sum_over_time(x[5m])`, `count_over_time(x[5m])`. +/// Any other query is returned unchanged. +fn rewrite_avg(query: &str) -> Vec { + match promql_parser::parser::parse(query) { + Ok(Expr::Aggregate(agg)) if agg.op.id() == token::T_AVG => [token::T_SUM, token::T_COUNT] + .into_iter() + .map(|op| { + Expr::Aggregate(AggregateExpr { + op: TokenType::new(op), + ..agg.clone() + }) + .to_string() + }) + .collect(), + // sum_over_time and count_over_time share avg_over_time's signature. + Ok(Expr::Call(call)) if call.func.name == "avg_over_time" => { + ["sum_over_time", "count_over_time"] + .into_iter() + .map(|name| { + Expr::Call(Call { + func: Function { + name, + ..call.func.clone() + }, + args: call.args.clone(), + }) + .to_string() + }) + .collect() + } + _ => vec![query.to_string()], + } +} + /// Try to extract `QueryRequirements` from a single leaf PromQL query string. /// Returns `None` if the query cannot be parsed or does not match any pattern. fn extract_requirements( @@ -301,6 +340,74 @@ mod tests { )); } + fn sorted_statistics_and_queries( + items: &[OptimizerItem], + ) -> Vec<(Vec, Vec)> { + let mut out: Vec<_> = items + .iter() + .map(|i| (i.requirements.statistics.clone(), i.query_strings.clone())) + .collect(); + out.sort_by_key(|(_, q)| q.clone()); + out + } + + // A multi-statistic [Sum, Count] item has no single-sketch candidate, so + // avg must reach the optimizer as separate sum and count items. + #[test] + fn spatial_avg_becomes_sum_and_count_items() { + let rqes = vec![rqe("avg by (job) (metric)", 60_000)]; + let items = extract_aqes(&rqes, &empty_schema(), 15_000).unwrap(); + assert_eq!( + sorted_statistics_and_queries(&items), + vec![ + ( + vec![Statistic::Count], + vec!["count by (job) (metric)".into()] + ), + (vec![Statistic::Sum], vec!["sum by (job) (metric)".into()]), + ] + ); + for item in &items { + assert_eq!( + item.requirements.grouping_labels, + KeyByLabelNames::new(vec!["job".into()]) + ); + } + } + + #[test] + fn avg_over_time_becomes_sum_and_count_over_time_items() { + let rqes = vec![rqe("avg_over_time(metric[5m])", 60_000)]; + let items = extract_aqes(&rqes, &empty_schema(), 15_000).unwrap(); + assert_eq!( + sorted_statistics_and_queries(&items), + vec![ + ( + vec![Statistic::Count], + vec!["count_over_time(metric[5m])".into()] + ), + ( + vec![Statistic::Sum], + vec!["sum_over_time(metric[5m])".into()] + ), + ] + ); + for item in &items { + assert_eq!(item.requirements.data_range_ms, 300_000); + } + } + + #[test] + fn avg_arm_of_binary_query_is_rewritten() { + let rqes = vec![rqe( + "avg_over_time(metric_a[5m]) / sum_over_time(metric_b[5m])", + 60_000, + )]; + let items = extract_aqes(&rqes, &empty_schema(), 15_000).unwrap(); + assert_eq!(items.len(), 3); + assert!(items.iter().all(|i| i.requirements.statistics.len() == 1)); + } + #[test] fn spatial_only_query_sets_range_to_scrape_interval() { let rqes = vec![rqe("sum(metric)", 60_000)];