use crate::event::streaming_mv::types::{AggDef, AggFunction};
use nodedb_sql::parser::preprocess::lex::{
find_ascii_case_insensitive, rfind_ascii_case_insensitive,
};
use super::super::super::result::DdlError;
fn parse_err(message: &str) -> DdlError {
DdlError {
sqlstate: "42601".to_string(),
message: message.to_string(),
}
}
pub struct ParsedStreamingMv {
pub group_by_columns: Vec<String>,
pub aggregates: Vec<AggDef>,
pub filter_expr: Option<String>,
pub source_stream: String,
}
pub fn parse_streaming_mv(query_sql: &str) -> Result<ParsedStreamingMv, DdlError> {
let query = query_sql.trim().trim_end_matches(';').trim();
let query_upper = query.to_uppercase();
let from_pos = find_ascii_case_insensitive(query, " FROM ")
.ok_or_else(|| parse_err("expected FROM clause"))?;
let after_from = query[from_pos + 6..].trim();
let source_stream = after_from
.split_whitespace()
.next()
.unwrap_or("")
.to_lowercase();
let group_by_columns = if let Some(gb_pos) = find_ascii_case_insensitive(query, " GROUP BY ") {
let gb_str = query[gb_pos + 10..].trim();
gb_str
.split(',')
.map(|s| s.trim().to_lowercase())
.filter(|s| !s.is_empty())
.collect()
} else {
Vec::new()
};
let filter_expr = if let Some(where_pos) = find_ascii_case_insensitive(query, " WHERE ") {
let end = find_ascii_case_insensitive(query, " GROUP BY ").unwrap_or(query.len());
if where_pos < end {
Some(query[where_pos + 7..end].trim().to_string())
} else {
None
}
} else {
None
};
if !query_upper.starts_with("SELECT ") {
return Err(parse_err("expected SELECT"));
}
let select_list = query[7..from_pos].trim();
let aggregates = parse_select_aggregates(select_list);
if aggregates.is_empty() {
return Err(parse_err(
"streaming MV requires at least one aggregate function (COUNT, SUM, MIN, MAX, AVG)",
));
}
Ok(ParsedStreamingMv {
group_by_columns,
aggregates,
filter_expr,
source_stream,
})
}
fn parse_select_aggregates(select_list: &str) -> Vec<AggDef> {
let mut aggregates = Vec::new();
for item in select_list.split(',') {
let item = item.trim();
if item.is_empty() || item == "*" {
continue;
}
let (expr_part, alias) = if let Some(as_pos) = rfind_ascii_case_insensitive(item, " AS ") {
(
item[..as_pos].trim(),
item[as_pos + 4..].trim().to_lowercase(),
)
} else {
(item, item.to_lowercase().replace(['(', ')', '*', ' '], "_"))
};
let expr_upper = expr_part.to_uppercase();
let func = if expr_upper.starts_with("COUNT(") {
Some(AggFunction::Count)
} else if expr_upper.starts_with("SUM(") {
Some(AggFunction::Sum)
} else if expr_upper.starts_with("MIN(") {
Some(AggFunction::Min)
} else if expr_upper.starts_with("MAX(") {
Some(AggFunction::Max)
} else if expr_upper.starts_with("AVG(") {
Some(AggFunction::Avg)
} else {
None
};
if let Some(function) = func {
let inner = expr_part
.split_once('(')
.and_then(|(_, rest)| rest.rsplit_once(')'))
.map(|(inner, _)| inner.trim().to_string())
.unwrap_or_default();
let input_expr = if inner == "*" {
String::new()
} else {
inner
};
aggregates.push(AggDef {
output_name: alias,
function,
input_expr,
});
}
}
aggregates
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_basic_streaming_mv() {
let query = "SELECT event_type, count(*) AS cnt \
FROM orders_stream \
GROUP BY event_type";
let parsed = parse_streaming_mv(query).unwrap();
assert_eq!(parsed.source_stream, "orders_stream");
assert_eq!(parsed.group_by_columns, vec!["event_type"]);
assert_eq!(parsed.aggregates.len(), 1);
assert_eq!(parsed.aggregates[0].function, AggFunction::Count);
assert_eq!(parsed.aggregates[0].output_name, "cnt");
assert!(parsed.aggregates[0].input_expr.is_empty());
}
#[test]
fn parse_multi_aggregate() {
let query = "SELECT count(*) AS cnt, sum(total) AS revenue \
FROM orders_stream \
GROUP BY event_type";
let parsed = parse_streaming_mv(query).unwrap();
assert_eq!(parsed.aggregates.len(), 2);
assert_eq!(parsed.aggregates[1].function, AggFunction::Sum);
assert_eq!(parsed.aggregates[1].input_expr, "total");
}
#[test]
fn aggregate_alias_after_unicode_expression_preserves_original_offsets() {
let parsed =
parse_streaming_mv("SELECT sum(ffff) AS total FROM orders_stream GROUP BY event_type")
.expect("streaming aggregate should parse");
assert_eq!(parsed.aggregates.len(), 1);
assert_eq!(parsed.aggregates[0].input_expr, "ffff");
assert_eq!(parsed.aggregates[0].output_name, "total");
}
#[test]
fn parse_with_where() {
let query = "SELECT count(*) AS cnt \
FROM orders_stream \
WHERE event_type = 'INSERT' \
GROUP BY collection";
let parsed = parse_streaming_mv(query).unwrap();
assert!(parsed.filter_expr.is_some());
assert!(parsed.filter_expr.unwrap().contains("event_type"));
}
#[test]
fn requires_at_least_one_aggregate() {
let query = "SELECT status FROM orders_stream GROUP BY status";
assert!(parse_streaming_mv(query).is_err());
}
#[test]
fn requires_from_clause() {
let query = "SELECT count(*) AS cnt GROUP BY status";
assert!(parse_streaming_mv(query).is_err());
}
}