use super::field_mappings::{FieldMappings, FieldType};
use crate::parser::{Aggregation, GroupBy, StatsNode};
use serde_json::{json, Value as JsonValue};
const DEFAULT_BUCKET_SIZE: usize = 5;
fn map_aggregation(func: &str) -> Option<&'static str> {
match func {
"count" => Some("value_count"),
"unique_count" | "cardinality" => Some("cardinality"),
"sum" => Some("sum"),
"min" => Some("min"),
"max" => Some("max"),
"average" | "avg" | "mean" => Some("avg"),
"median" | "med" => Some("percentiles"),
"std" | "stdev" | "standard_deviation" => Some("extended_stats"),
"percentile" | "percentiles" | "p" | "pct" => Some("percentiles"),
"percentile_rank" | "percentile_ranks" | "pct_rank" | "pct_ranks" => {
Some("percentile_ranks")
}
"values" | "unique" => Some("terms"),
_ => None,
}
}
fn resolve_aggregation_field(field: &str, field_mappings: Option<&FieldMappings>) -> String {
if let Some(mappings) = field_mappings {
if let Some(ft) = mappings.get_field_type(field) {
if *ft == FieldType::Text {
return format!("{}.keyword", field);
}
}
let keyword_field = format!("{}.keyword", field);
if mappings.get_field_type(&keyword_field).is_some() {
return keyword_field;
}
}
field.to_string()
}
pub fn translate_stats(
stats: &StatsNode,
field_mappings: Option<&FieldMappings>,
) -> Result<JsonValue, String> {
if stats.aggregations.is_empty() {
return Err("No aggregations specified in stats query".to_string());
}
let aggs_dsl = if stats.group_by.is_empty() {
build_simple_aggregations(&stats.aggregations)?
} else {
build_grouped_aggregations(&stats.aggregations, &stats.group_by, field_mappings)?
};
Ok(json!({ "aggs": aggs_dsl }))
}
fn build_simple_aggregations(aggregations: &[Aggregation]) -> Result<JsonValue, String> {
let mut aggs = serde_json::Map::new();
for (i, agg) in aggregations.iter().enumerate() {
let func = agg.function.to_lowercase();
let field = agg.field.as_deref().unwrap_or("*");
let alias = agg
.alias
.clone()
.unwrap_or_else(|| format!("{}_{}", func, i));
let agg_dsl = build_single_aggregation(&func, field, agg)?;
aggs.insert(alias, agg_dsl);
}
Ok(JsonValue::Object(aggs))
}
fn build_single_aggregation(
func: &str,
field: &str,
agg: &Aggregation,
) -> Result<JsonValue, String> {
if func == "count" && field == "*" {
return Ok(json!({ "value_count": { "field": "_id" } }));
}
let os_type =
map_aggregation(func).ok_or_else(|| format!("Unknown aggregation function: {}", func))?;
match func {
"median" | "med" => Ok(json!({ "percentiles": { "field": field, "percents": [50.0] } })),
"std" | "stdev" | "standard_deviation" => {
Ok(json!({ "extended_stats": { "field": field } }))
}
"percentile" | "percentiles" | "p" | "pct" => {
let percents = agg
.percentile_values
.as_ref()
.cloned()
.unwrap_or_else(|| vec![50.0]);
Ok(json!({ "percentiles": { "field": field, "percents": percents } }))
}
"percentile_rank" | "percentile_ranks" | "pct_rank" | "pct_ranks" => {
let values = agg.rank_values.as_ref().cloned().unwrap_or_default();
if values.is_empty() {
return Err("percentile_rank requires at least one value".to_string());
}
Ok(json!({ "percentile_ranks": { "field": field, "values": values } }))
}
"values" | "unique" => Ok(json!({ "terms": { "field": field, "size": 10000 } })),
_ => {
Ok(json!({ os_type: { "field": field } }))
}
}
}
fn build_grouped_aggregations(
aggregations: &[Aggregation],
group_by: &[GroupBy],
field_mappings: Option<&FieldMappings>,
) -> Result<JsonValue, String> {
let inner_aggs = build_simple_aggregations(aggregations)?;
let mut order_field: Option<String> = None;
let mut order_direction = "desc";
for (i, agg) in aggregations.iter().enumerate() {
if let Some(ref modifier) = agg.modifier {
let alias = agg
.alias
.clone()
.unwrap_or_else(|| format!("{}_{}", agg.function, i));
order_field = Some(alias);
order_direction = if modifier == "top" { "desc" } else { "asc" };
break;
}
}
let mut current_aggs = inner_aggs;
for (i, gb) in group_by.iter().rev().enumerate() {
let bucket_size = gb.bucket_size.unwrap_or(DEFAULT_BUCKET_SIZE);
let resolved_field = resolve_aggregation_field(&gb.field, field_mappings);
let mut terms_agg = json!({
"terms": {
"field": resolved_field,
"size": bucket_size
}
});
if i == group_by.len() - 1 {
if let Some(ref of_) = order_field {
terms_agg["terms"]["order"] = json!({ of_: order_direction });
}
}
if current_aggs.as_object().is_some_and(|m| !m.is_empty()) {
terms_agg["aggs"] = current_aggs;
}
let key = format!("group_by_{}", gb.field);
current_aggs = json!({ key: terms_agg });
}
Ok(current_aggs)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser::TqlParser;
fn parse_stats(query: &str) -> StatsNode {
let parser = TqlParser::new();
let ast = parser.parse(query).expect("parse failed");
match ast {
crate::parser::AstNode::StatsExpr(s) => s,
crate::parser::AstNode::QueryWithStats(qws) => qws.stats,
_ => panic!("expected stats AST, got {:?}", ast),
}
}
#[test]
fn test_simple_count_star() {
let stats = parse_stats("| stats count(*)");
let dsl = translate_stats(&stats, None).unwrap();
let aggs = &dsl["aggs"];
assert!(aggs
.as_object()
.unwrap()
.values()
.any(|v| { v.get("value_count").is_some() }));
}
#[test]
fn test_count_by_field() {
let stats = parse_stats("| stats count(*) by event.code");
let dsl = translate_stats(&stats, None).unwrap();
let aggs = &dsl["aggs"];
let group = aggs.get("group_by_event.code").expect("missing group_by");
assert!(group.get("terms").is_some());
assert!(group.get("aggs").is_some());
}
#[test]
fn test_multiple_group_by() {
let stats = parse_stats("| stats count(*) by host.name, event.code");
let dsl = translate_stats(&stats, None).unwrap();
let aggs = &dsl["aggs"];
let outer = aggs.get("group_by_host.name").expect("missing outer group");
assert!(outer.get("terms").is_some());
let inner_aggs = outer.get("aggs").expect("missing inner aggs");
assert!(inner_aggs.get("group_by_event.code").is_some());
}
#[test]
fn test_avg_aggregation() {
let stats = parse_stats("| stats avg(response_time)");
let dsl = translate_stats(&stats, None).unwrap();
let aggs = &dsl["aggs"];
assert!(aggs
.as_object()
.unwrap()
.values()
.any(|v| { v.get("avg").is_some() }));
}
}