use super::field_mappings::{FieldMappings, FieldType};
use crate::parser::{Aggregation, GroupBy, StatsNode};
use serde_json::{json, Value as JsonValue};
const UNLIMITED_BUCKET_SIZE: usize = 10000;
const BUCKET_LIMIT_AGG_NAME: &str = "_tql_bucket_limit";
fn is_single_value_metric(os_type: &str) -> bool {
matches!(
os_type,
"value_count" | "cardinality" | "sum" | "min" | "max" | "avg"
)
}
fn agg_alias(agg: &Aggregation, index: usize) -> String {
agg.alias
.clone()
.unwrap_or_else(|| format!("{}_{}", agg.function.to_lowercase(), index))
}
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" | "standard_deviation" => Some("extended_stats"),
"percentile" | "percentiles" | "p" | "pct" => Some("percentiles"),
"percentile_rank" | "percentile_ranks" | "pct_rank" | "pct_ranks" => {
Some("percentile_ranks")
}
"values" | "unique" | "distinct" => 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(agg, 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" | "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" | "distinct" => Ok(
json!({ "terms": { "field": field, "size": UNLIMITED_BUCKET_SIZE } }),
),
_ => {
Ok(json!({ os_type: { "field": field } }))
}
}
}
fn resolve_aggregation_modifier(
aggregations: &[Aggregation],
group_levels: usize,
) -> Result<Option<(String, &'static str, usize)>, String> {
for (i, agg) in aggregations.iter().enumerate() {
let Some(modifier) = agg.modifier.as_deref() else {
continue;
};
let limit = agg.limit.unwrap_or(10);
if group_levels > 1 {
return Err(format!(
"'{modifier} {limit}' on an aggregation cannot be pushed down alongside \
{group_levels} group-by fields: OpenSearch ranks the buckets of one `terms` \
aggregation, while the in-memory engine ranks the flattened cross product of \
all group-by levels. Group by a single field, or move the limit onto the \
group-by field (`by <field> {modifier} {limit}`), which ranks by doc_count."
));
}
if limit == 0 {
return Err(format!(
"'{modifier} 0' on an aggregation asks for zero buckets, which OpenSearch cannot \
express: `terms` requires size > 0 and `bucket_sort` requires a positive size."
));
}
let func = agg.function.to_lowercase();
let os_type = map_aggregation(&func)
.ok_or_else(|| format!("Unknown aggregation function: {}", func))?;
if !is_single_value_metric(os_type) {
return Err(format!(
"'{modifier} {limit}' cannot rank buckets by '{func}': OpenSearch cannot order a \
`terms` aggregation by a value that is not a single-value metric. Rank by count, \
sum, min, max, avg or cardinality, or move the limit onto the group-by field \
(`by <field> {modifier} {limit}`)."
));
}
let direction = if modifier == "top" { "desc" } else { "asc" };
return Ok(Some((agg_alias(agg, i), direction, limit)));
}
Ok(None)
}
fn build_grouped_aggregations(
aggregations: &[Aggregation],
group_by: &[GroupBy],
field_mappings: Option<&FieldMappings>,
) -> Result<JsonValue, String> {
let inner_aggs = build_simple_aggregations(aggregations)?;
let modifier = resolve_aggregation_modifier(aggregations, group_by.len())?;
if group_by.len() > 1 {
if let Some(gb) = group_by.iter().find(|gb| gb.bucket_size == Some(0)) {
return Err(format!(
"'top 0' on group-by field '{}' asks for zero buckets at that level, which \
OpenSearch cannot express: `terms` requires size > 0.",
gb.field
));
}
}
let mut current_aggs = inner_aggs;
let last_level = group_by.len() - 1;
for (depth, gb) in group_by.iter().rev().enumerate() {
let level = last_level - depth;
let is_outermost = level == 0;
let resolved_field = resolve_aggregation_field(&gb.field, field_mappings);
let mut terms_body = serde_json::Map::new();
terms_body.insert("field".to_string(), json!(resolved_field));
match (is_outermost, modifier.as_ref()) {
(true, Some((alias, direction, limit))) => {
terms_body.insert("size".to_string(), json!(limit));
terms_body.insert("order".to_string(), json!({ alias.as_str(): direction }));
}
_ => match gb.bucket_size {
Some(n) if n > 0 => {
terms_body.insert("size".to_string(), json!(n));
terms_body.insert("order".to_string(), json!({ "_count": "desc" }));
}
_ => {
terms_body.insert("size".to_string(), json!(UNLIMITED_BUCKET_SIZE));
}
},
}
let mut terms_agg = json!({ "terms": JsonValue::Object(terms_body) });
let mut sub_aggs = current_aggs
.as_object()
.cloned()
.unwrap_or_else(serde_json::Map::new);
if is_outermost {
if let (Some(_), Some(n)) = (modifier.as_ref(), gb.bucket_size) {
if n > 0 {
sub_aggs.insert(
BUCKET_LIMIT_AGG_NAME.to_string(),
json!({
"bucket_sort": {
"sort": [{ "_count": { "order": "desc" } }],
"size": n
}
}),
);
}
}
}
if !sub_aggs.is_empty() {
terms_agg["aggs"] = JsonValue::Object(sub_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;
#[test]
fn the_values_aggregation_uses_the_named_bucket_size() {
for spelling in ["values", "unique", "distinct"] {
let stats = parse_stats(&format!("| stats {spelling}(user.name)"));
let dsl = translate_stats(&stats, None).expect("translation failed");
let size = dsl
.pointer("/aggs")
.and_then(|aggs| aggs.as_object())
.and_then(|aggs| aggs.values().next())
.and_then(|agg| agg.pointer("/terms/size"))
.unwrap_or_else(|| panic!("`{spelling}` emitted no terms bucket: {dsl}"));
assert_eq!(
*size,
serde_json::json!(UNLIMITED_BUCKET_SIZE),
"`{spelling}` disagrees with the named size"
);
}
}
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());
}
fn terms_body(query: &str) -> JsonValue {
let stats = parse_stats(query);
let dsl = translate_stats(&stats, None).unwrap_or_else(|e| panic!("{query:?}: {e}"));
dsl["aggs"]["group_by_department"]["terms"].clone()
}
fn outer(query: &str) -> JsonValue {
let stats = parse_stats(query);
let dsl = translate_stats(&stats, None).unwrap_or_else(|e| panic!("{query:?}: {e}"));
dsl["aggs"]["group_by_department"].clone()
}
#[test]
fn aggregation_modifier_ranks_by_the_aggregate() {
for query in [
"| stats sum(salary) top 3 by department",
"| stats sum(salary, top 3) by department",
] {
assert_eq!(
terms_body(query),
json!({"field": "department", "size": 3, "order": {"sum_0": "desc"}}),
"{query:?}"
);
}
}
#[test]
fn bottom_reverses_the_direction() {
for query in [
"| stats sum(salary) bottom 3 by department",
"| stats sum(salary, bottom 3) by department",
] {
assert_eq!(
terms_body(query),
json!({"field": "department", "size": 3, "order": {"sum_0": "asc"}}),
"{query:?}"
);
}
}
#[test]
fn group_by_bucket_limit_ranks_by_doc_count() {
assert_eq!(
terms_body("| stats count() by department top 3"),
json!({"field": "department", "size": 3, "order": {"_count": "desc"}})
);
}
#[test]
fn unmodified_group_by_asks_for_every_bucket() {
let body = terms_body("| stats sum(salary) by department");
assert_eq!(body["size"], UNLIMITED_BUCKET_SIZE);
assert!(body.get("order").is_none(), "{body}");
}
#[test]
fn both_modifiers_emit_two_passes_not_one() {
let agg = outer("| stats sum(salary) top 4 by department top 2");
assert_eq!(
agg["terms"],
json!({"field": "department", "size": 4, "order": {"sum_0": "desc"}})
);
assert_eq!(
agg["aggs"][BUCKET_LIMIT_AGG_NAME],
json!({"bucket_sort": {"sort": [{"_count": {"order": "desc"}}], "size": 2}})
);
assert_eq!(agg["aggs"]["sum_0"], json!({"sum": {"field": "salary"}}));
}
#[test]
fn bucket_sort_is_absent_when_only_one_modifier_is_present() {
for query in [
"| stats sum(salary) top 3 by department",
"| stats count() by department top 3",
"| stats sum(salary) by department",
] {
assert!(
outer(query)["aggs"].get(BUCKET_LIMIT_AGG_NAME).is_none(),
"{query:?} emitted a second pass it does not need"
);
}
}
#[test]
fn top_zero_on_a_group_by_field_is_a_no_op() {
let stats = parse_stats("| stats count() by role top 0");
let dsl = translate_stats(&stats, None).unwrap();
let body = &dsl["aggs"]["group_by_role"]["terms"];
assert_eq!(body["size"], UNLIMITED_BUCKET_SIZE);
assert!(body.get("order").is_none(), "{body}");
}
#[test]
fn multilevel_bucket_limits_nest_per_level() {
let stats = parse_stats("| stats count() by department top 3, role top 2");
let dsl = translate_stats(&stats, None).unwrap();
let outer_agg = &dsl["aggs"]["group_by_department"];
assert_eq!(
outer_agg["terms"],
json!({"field": "department", "size": 3, "order": {"_count": "desc"}})
);
assert_eq!(
outer_agg["aggs"]["group_by_role"]["terms"],
json!({"field": "role", "size": 2, "order": {"_count": "desc"}})
);
}
#[test]
fn multilevel_without_limits_asks_for_every_bucket_at_every_level() {
let stats = parse_stats("| stats count() by department, role");
let dsl = translate_stats(&stats, None).unwrap();
let outer_agg = &dsl["aggs"]["group_by_department"];
assert_eq!(outer_agg["terms"]["size"], UNLIMITED_BUCKET_SIZE);
assert_eq!(
outer_agg["aggs"]["group_by_role"]["terms"]["size"],
UNLIMITED_BUCKET_SIZE
);
}
#[test]
fn the_order_path_follows_the_aggregation_position() {
let agg = outer("| stats count(), sum(salary) top 3 by department");
let order = agg["terms"]["order"].as_object().expect("an order");
let (key, _) = order.iter().next().expect("one order key");
assert_eq!(key, "sum_1");
assert!(agg["aggs"].get(key).is_some(), "{}", agg["aggs"]);
}
#[test]
fn an_uppercase_function_name_still_names_the_aggregation_it_emits() {
let agg = outer("| stats SUM(salary) TOP 3 by department");
assert_eq!(
agg["terms"],
json!({"field": "department", "size": 3, "order": {"sum_0": "desc"}})
);
assert!(agg["aggs"].get("sum_0").is_some(), "{}", agg["aggs"]);
}
#[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() }));
}
}