use std::fmt::Write;
use std::hash::Hash;
use std::mem;
use ahash::HashMap;
use anyhow::{Result, bail, ensure};
use revision::revisioned;
use surrealdb_strand::Strand;
use surrealdb_types::ToSql;
use crate::err::Error;
use crate::expr::field::Selector;
use crate::expr::statements::define::DefineConfigStatement;
use crate::expr::statements::{
CreateStatement, DefineAccessStatement, DefineApiStatement, DefineFieldStatement,
DefineFunctionStatement, DefineIndexStatement, InsertStatement, RelateStatement,
UpdateStatement, UpsertStatement,
};
use crate::expr::visit::{MutVisitor, VisitMut};
use crate::expr::{Expr, Field, Fields, Function, Groups, Idiom, Part, SelectStatement};
use crate::val::{Array, Datetime, Number, Object, TryAdd as _, TryFloatDiv, TryMul, Value};
#[revisioned(revision = 1)]
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub(crate) enum Aggregation {
Count,
CountValue(usize),
NumberMax(usize),
NumberMin(usize),
Sum(usize),
Mean(usize),
StdDev(usize),
Variance(usize),
DatetimeMax(usize),
DatetimeMin(usize),
Accumulate(usize),
}
impl Aggregation {
pub(crate) fn to_stat(&self) -> AggregationStat {
match *self {
Aggregation::Count => AggregationStat::Count {
count: 0,
},
Aggregation::CountValue(arg) => AggregationStat::CountValue {
arg,
count: 0,
},
Aggregation::NumberMax(arg) => AggregationStat::NumberMax {
arg,
max: f64::NEG_INFINITY.into(),
},
Aggregation::NumberMin(arg) => AggregationStat::NumberMin {
arg,
min: f64::INFINITY.into(),
},
Aggregation::Sum(arg) => AggregationStat::Sum {
arg,
sum: 0.0.into(),
},
Aggregation::Mean(arg) => AggregationStat::Mean {
arg,
count: 0,
sum: 0.0.into(),
},
Aggregation::StdDev(arg) => AggregationStat::StdDev {
arg,
sum: 0.0.into(),
sum_of_squares: 0.0.into(),
count: 0,
},
Aggregation::Variance(arg) => AggregationStat::Variance {
arg,
sum: 0.0.into(),
sum_of_squares: 0.0.into(),
count: 0,
},
Aggregation::DatetimeMax(arg) => AggregationStat::TimeMax {
arg,
max: Datetime::MIN_UTC,
},
Aggregation::DatetimeMin(arg) => AggregationStat::TimeMin {
arg,
min: Datetime::MAX_UTC,
},
Aggregation::Accumulate(arg) => AggregationStat::Accumulate {
arg,
values: Vec::new(),
},
}
}
}
#[revisioned(revision = 1)]
#[derive(Clone, Debug, PartialEq)]
pub(crate) enum AggregationStat {
Count {
count: i64,
},
CountValue {
arg: usize,
count: i64,
},
NumberMax {
arg: usize,
max: Number,
},
NumberMin {
arg: usize,
min: Number,
},
Sum {
arg: usize,
sum: Number,
},
Mean {
arg: usize,
sum: Number,
count: i64,
},
StdDev {
arg: usize,
sum: Number,
sum_of_squares: Number,
count: i64,
},
Variance {
arg: usize,
sum: Number,
sum_of_squares: Number,
count: i64,
},
TimeMax {
arg: usize,
max: Datetime,
},
TimeMin {
arg: usize,
min: Datetime,
},
Accumulate {
arg: usize,
values: Vec<Value>,
},
}
impl AggregationStat {
pub(crate) fn get_count(aggregation_stats: &[AggregationStat]) -> Option<i64> {
aggregation_stats.iter().find_map(|x| match x {
AggregationStat::Count {
count,
}
| AggregationStat::Mean {
count,
..
}
| AggregationStat::Variance {
count,
..
}
| AggregationStat::StdDev {
count,
..
} => Some(*count),
_ => None,
})
}
}
pub fn write_aggregate_field_name(s: &mut String, idx: usize) {
write!(s, "_a{}", idx).expect("writing into a string cannot fail");
}
pub fn write_group_field_name(s: &mut String, idx: usize) {
write!(s, "_g{}", idx).expect("writing into a string cannot fail");
}
fn format_indexed_strand(prefix: &str, idx: usize) -> Strand {
Strand::from_display(format_args!("{prefix}{idx}"))
}
pub fn aggregate_field_name(idx: usize) -> Strand {
format_indexed_strand("_a", idx)
}
pub fn group_field_name(idx: usize) -> Strand {
format_indexed_strand("_g", idx)
}
pub(crate) fn add_to_aggregation_stats(
arguments: &[Value],
stats: &mut [AggregationStat],
) -> Result<()> {
for stat in stats {
match stat {
AggregationStat::Count {
count,
} => {
*count += 1;
}
AggregationStat::CountValue {
arg,
count,
} => {
*count += arguments[*arg].is_truthy() as i64;
}
AggregationStat::NumberMax {
arg,
max,
} => {
let Value::Number(ref n) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "math::max".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `number` but found `{}`",
arguments[*arg].to_sql()
),
})
};
if *max < *n {
*max = *n
}
}
AggregationStat::NumberMin {
arg,
min,
} => {
let Value::Number(ref n) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "math::min".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `number` but found `{}`",
arguments[*arg].to_sql()
),
})
};
if *min > *n {
*min = *n
}
}
AggregationStat::Sum {
arg,
sum,
} => {
let Value::Number(ref n) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "math::sum".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `number` but found `{}`",
arguments[*arg].to_sql()
),
})
};
*sum = (*sum).try_add(*n)?;
}
AggregationStat::Mean {
arg,
sum,
count,
} => {
let Value::Number(ref n) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "math::mean".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `number` but found `{}`",
arguments[*arg].to_sql()
),
})
};
*sum = (*sum).try_add(*n)?;
*count += 1;
}
AggregationStat::StdDev {
arg,
sum,
sum_of_squares,
count,
} => {
let Value::Number(ref n) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "math::stddev".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `number` but found `{}`",
arguments[*arg].to_sql()
),
})
};
*sum = (*sum).try_add(*n)?;
*sum_of_squares = (*sum_of_squares).try_add(n.try_mul(*n)?)?;
*count += 1;
}
AggregationStat::Variance {
arg,
sum,
sum_of_squares,
count,
} => {
let Value::Number(ref n) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "math::variance".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `number` but found `{}`",
arguments[*arg].to_sql()
),
})
};
*sum = (*sum).try_add(*n)?;
*sum_of_squares = (*sum_of_squares).try_add(n.try_mul(*n)?)?;
*count += 1;
}
AggregationStat::TimeMax {
arg,
max,
} => {
let Value::Datetime(ref d) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "time::max".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `datetime` but found `{}`",
arguments[*arg].to_sql()
),
})
};
if *max < *d {
*max = *d;
}
}
AggregationStat::TimeMin {
arg,
min,
} => {
let Value::Datetime(ref d) = arguments[*arg] else {
bail!(Error::InvalidFunctionArguments {
name: "time::min".to_string(),
message: format!(
"Argument 1 was the wrong type. Expected `datetime` but found `{}`",
arguments[*arg].to_sql()
),
})
};
if *min > *d {
*min = *d;
}
}
AggregationStat::Accumulate {
arg,
values,
} => {
values.push(arguments[*arg].clone());
}
}
}
Ok(())
}
pub(crate) fn create_field_document(group: &[Value], stats: &[AggregationStat]) -> Object {
let mut res = Object::default();
for (idx, a) in stats.iter().enumerate() {
let value = match a {
AggregationStat::Count {
count,
}
| AggregationStat::CountValue {
count,
..
} => Value::from(Number::from(*count)),
AggregationStat::NumberMax {
max,
..
} => (*max).into(),
AggregationStat::NumberMin {
min,
..
} => (*min).into(),
AggregationStat::Sum {
sum,
..
} => (*sum).into(),
AggregationStat::Mean {
sum,
count,
..
} => sum.try_float_div((*count).into()).unwrap_or(f64::NAN.into()).into(),
AggregationStat::StdDev {
sum,
sum_of_squares,
count,
..
} => {
let num = if *count == 0 {
Number::from(f64::NAN)
} else if *count == 1 {
Number::from(0.0)
} else {
let mean = *sum / Number::from(*count);
let variance = (*sum_of_squares - (*sum * mean)) / Number::from(*count - 1);
if variance == Number::from(0.0) {
Number::from(0.0)
} else {
variance.sqrt()
}
};
num.into()
}
AggregationStat::Variance {
sum,
sum_of_squares,
count,
..
} => {
let num = if *count == 0 {
Number::from(f64::NAN)
} else if *count == 1 {
Number::from(0.0)
} else {
let mean = *sum / Number::from(*count);
(*sum_of_squares - (*sum * mean)) / Number::from(*count - 1)
};
num.into()
}
AggregationStat::TimeMax {
max,
..
} => (*max).into(),
AggregationStat::TimeMin {
min,
..
} => (*min).into(),
AggregationStat::Accumulate {
values,
..
} => Value::Array(Array(values.clone())),
};
res.0.insert(aggregate_field_name(idx), value);
}
for (idx, g) in group.iter().enumerate() {
res.0.insert(group_field_name(idx), g.clone());
}
res
}
struct AggregateExprCollector<'a> {
support_acummulate: bool,
within_aggregate_argument: bool,
exprs_map: &'a mut HashMap<Expr, usize>,
aggregations: &'a mut Vec<Aggregation>,
groups: &'a Groups,
}
impl AggregateExprCollector<'_> {
fn push_aggregate_function<F: Fn(usize) -> Aggregation>(
&mut self,
name: &str,
args: &[Expr],
f: F,
) -> Result<()> {
ensure!(
args.len() == 1,
Error::InvalidFunctionArguments {
name: name.to_string(),
message: "Expected 1 argument".to_string()
}
);
let expr = args[0].clone();
let len = self.exprs_map.len();
let arg = *self.exprs_map.entry(expr).or_insert_with(|| len);
self.aggregations.push(f(arg));
Ok(())
}
}
impl MutVisitor for AggregateExprCollector<'_> {
type Error = anyhow::Error;
fn visit_mut_expr(&mut self, s: &mut Expr) -> Result<(), Self::Error> {
match s {
Expr::FunctionCall(f) => {
if let Function::Normal(x) = &f.receiver {
match x.as_str() {
"count" => {
if f.arguments.is_empty() {
self.aggregations.push(Aggregation::Count);
} else {
self.push_aggregate_function(
"count",
&f.arguments,
Aggregation::CountValue,
)?;
}
}
"math::max" => {
self.push_aggregate_function(
"math::max",
&f.arguments,
Aggregation::NumberMax,
)?;
}
"math::min" => {
self.push_aggregate_function(
"math::min",
&f.arguments,
Aggregation::NumberMin,
)?;
}
"math::sum" => {
self.push_aggregate_function(
"math::sum",
&f.arguments,
Aggregation::Sum,
)?;
}
"math::mean" => {
self.push_aggregate_function(
"math::mean",
&f.arguments,
Aggregation::Mean,
)?;
}
"math::stddev" => {
self.push_aggregate_function(
"math::stddev",
&f.arguments,
Aggregation::StdDev,
)?;
}
"math::variance" => {
self.push_aggregate_function(
"math::variance",
&f.arguments,
Aggregation::Variance,
)?;
}
"time::max" => {
self.push_aggregate_function(
"time::max",
&f.arguments,
Aggregation::DatetimeMax,
)?;
}
"time::min" => {
self.push_aggregate_function(
"time::min",
&f.arguments,
Aggregation::DatetimeMin,
)?;
}
_ => {
return f.visit_mut(self);
}
}
} else {
return f.visit_mut(self);
}
self.within_aggregate_argument = true;
for a in f.arguments.iter_mut() {
a.visit_mut(self)?;
}
self.within_aggregate_argument = false;
*s = Expr::Idiom(Idiom::field(aggregate_field_name(self.aggregations.len() - 1)));
Ok(())
}
Expr::Param(p) => {
if p.as_str() == "this" {
bail!(Error::Query{
message: "Found a `$this` parameter refering to the document of a group by select statement\n\
Select statements with a group by currently have no defined document to refer to".to_string()
});
}
Ok(())
}
Expr::Idiom(i) => {
if !self.within_aggregate_argument {
if let Some(group_idx) = self.groups.0.iter().position(|x| x.0 == *i) {
i.visit_mut(self)?;
*s = Expr::Idiom(Idiom::field(group_field_name(group_idx)));
} else if let Some(Part::Field(field)) = i.0.first_mut() {
if self.support_acummulate {
let field_name = mem::replace(
field,
aggregate_field_name(self.aggregations.len()),
);
let len = self.exprs_map.len();
let arg = *self
.exprs_map
.entry(Expr::Idiom(Idiom::field(field_name)))
.or_insert_with(|| len);
self.aggregations.push(Aggregation::Accumulate(arg))
} else {
bail!(Error::Query {
message: format!(
"Found idiom `{}` within the selector of a materialized aggregate view.\n\
Selection of document fields which are not used within the argument of an optimized aggregate function is currently not supported",
i.to_sql()
)
})
}
}
Ok(())
} else {
i.visit_mut(self)
}
}
x => x.visit_mut(self),
}
}
fn visit_mut_create(&mut self, s: &mut CreateStatement) -> Result<(), Self::Error> {
for e in s.what.iter_mut() {
self.visit_mut_expr(e)?;
}
self.visit_mut_expr(&mut s.timeout)?;
if let Some(d) = &mut s.data {
ParentRewritor.visit_mut_data(d)?;
}
Ok(())
}
fn visit_mut_select(&mut self, s: &mut SelectStatement) -> Result<(), Self::Error> {
for v in s.what.iter_mut() {
self.visit_mut_expr(v)?;
}
if let Some(l) = s.limit.as_mut() {
self.visit_mut_expr(&mut l.0)?;
}
self.visit_mut_expr(&mut s.version)?;
ParentRewritor.visit_mut_fields(&mut s.fields)?;
for o in s.omit.iter_mut() {
ParentRewritor.visit_mut_expr(o)?;
}
if let Some(c) = s.cond.as_mut() {
ParentRewritor.visit_mut_expr(&mut c.0)?;
}
if let Some(s) = s.split.as_mut() {
for s in s.0.iter_mut() {
ParentRewritor.visit_mut_idiom(&mut s.0)?;
}
}
if let Some(g) = s.group.as_mut() {
for g in g.0.iter_mut() {
ParentRewritor.visit_mut_idiom(&mut g.0)?;
}
}
if let Some(o) = s.order.as_mut() {
ParentRewritor.visit_mut_ordering(o)?;
}
if let Some(f) = s.fetch.as_mut() {
for f in f.iter_mut() {
ParentRewritor.visit_mut_expr(&mut f.0)?;
}
}
Ok(())
}
fn visit_mut_update(&mut self, s: &mut UpdateStatement) -> Result<(), Self::Error> {
for e in s.what.iter_mut() {
self.visit_mut_expr(e)?;
}
if let Some(e) = &mut s.data {
ParentRewritor.visit_mut_data(e)?;
}
if let Some(e) = &mut s.cond {
ParentRewritor.visit_mut_expr(&mut e.0)?;
}
self.visit_mut_expr(&mut s.timeout)?;
Ok(())
}
fn visit_mut_upsert(&mut self, s: &mut UpsertStatement) -> Result<(), Self::Error> {
for e in s.what.iter_mut() {
self.visit_mut_expr(e)?;
}
if let Some(d) = &mut s.data {
ParentRewritor.visit_mut_data(d)?;
}
if let Some(e) = &mut s.cond {
ParentRewritor.visit_mut_expr(&mut e.0)?;
}
self.visit_mut_expr(&mut s.timeout)?;
Ok(())
}
fn visit_mut_relate(&mut self, s: &mut RelateStatement) -> Result<(), Self::Error> {
self.visit_mut_expr(&mut s.through)?;
self.visit_mut_expr(&mut s.from)?;
self.visit_mut_expr(&mut s.to)?;
self.visit_mut_expr(&mut s.timeout)?;
if let Some(d) = s.data.as_mut() {
ParentRewritor.visit_mut_data(d)?;
}
if let Some(o) = s.output.as_mut() {
ParentRewritor.visit_mut_output(o)?;
}
Ok(())
}
fn visit_mut_insert(&mut self, i: &mut InsertStatement) -> Result<(), Self::Error> {
if let Some(into) = &mut i.into {
self.visit_mut_expr(into)?;
}
self.visit_mut_expr(&mut i.timeout)?;
ParentRewritor.visit_mut_data(&mut i.data)?;
if let Some(update) = i.update.as_mut() {
ParentRewritor.visit_mut_data(update)?;
}
if let Some(o) = i.output.as_mut() {
ParentRewritor.visit_mut_output(o)?;
}
Ok(())
}
fn visit_mut_define_api(
&mut self,
d: &mut DefineApiStatement,
) -> std::result::Result<(), Self::Error> {
self.visit_mut_expr(&mut d.path)?;
self.visit_mut_expr(&mut d.comment)?;
Ok(())
}
fn visit_mut_define_function(
&mut self,
d: &mut DefineFunctionStatement,
) -> std::result::Result<(), Self::Error> {
self.visit_mut_expr(&mut d.comment)?;
Ok(())
}
fn visit_mut_define_access(
&mut self,
d: &mut DefineAccessStatement,
) -> std::result::Result<(), Self::Error> {
self.visit_mut_expr(&mut d.name)?;
self.visit_mut_expr(&mut d.comment)?;
self.visit_mut_expr(&mut d.duration.grant)?;
self.visit_mut_expr(&mut d.duration.token)?;
self.visit_mut_expr(&mut d.duration.session)?;
Ok(())
}
fn visit_mut_define_index(&mut self, d: &mut DefineIndexStatement) -> Result<(), Self::Error> {
self.visit_mut_expr(&mut d.name)?;
self.visit_mut_expr(&mut d.comment)?;
Ok(())
}
fn visit_mut_define_field(&mut self, d: &mut DefineFieldStatement) -> Result<(), Self::Error> {
self.visit_mut_expr(&mut d.name)?;
self.visit_mut_expr(&mut d.what)?;
self.visit_mut_expr(&mut d.comment)?;
Ok(())
}
}
struct ParentRewritor;
impl MutVisitor for ParentRewritor {
type Error = Error;
fn visit_mut_expr(&mut self, e: &mut Expr) -> Result<(), Self::Error> {
if let Expr::Param(p) = e
&& p.as_str() == "parent"
{
return Err(Error::Query{
message: "Found a `$parent` parameter refering to the document of a GROUP select statement\n\
Select statements with a GROUP BY or GROUP ALL currently have no defined document to refer to".to_string()
});
}
e.visit_mut(self)
}
fn visit_mut_create(&mut self, s: &mut CreateStatement) -> Result<(), Self::Error> {
for e in s.what.iter_mut() {
self.visit_mut_expr(e)?;
}
self.visit_mut_expr(&mut s.timeout)?;
Ok(())
}
fn visit_mut_select(&mut self, s: &mut SelectStatement) -> Result<(), Self::Error> {
self.visit_mut_fields(&mut s.fields)?;
for v in s.what.iter_mut() {
self.visit_mut_expr(v)?;
}
if let Some(l) = s.limit.as_mut() {
self.visit_mut_expr(&mut l.0)?;
}
self.visit_mut_expr(&mut s.version)?;
Ok(())
}
fn visit_mut_update(&mut self, s: &mut UpdateStatement) -> Result<(), Self::Error> {
for e in s.what.iter_mut() {
self.visit_mut_expr(e)?;
}
self.visit_mut_expr(&mut s.timeout)?;
Ok(())
}
fn visit_mut_upsert(&mut self, s: &mut UpsertStatement) -> Result<(), Self::Error> {
for e in s.what.iter_mut() {
self.visit_mut_expr(e)?;
}
self.visit_mut_expr(&mut s.timeout)?;
Ok(())
}
fn visit_mut_relate(&mut self, s: &mut RelateStatement) -> Result<(), Self::Error> {
self.visit_mut_expr(&mut s.through)?;
self.visit_mut_expr(&mut s.from)?;
self.visit_mut_expr(&mut s.to)?;
self.visit_mut_expr(&mut s.timeout)?;
Ok(())
}
fn visit_mut_insert(&mut self, i: &mut InsertStatement) -> Result<(), Self::Error> {
if let Some(into) = &mut i.into {
self.visit_mut_expr(into)?;
}
self.visit_mut_expr(&mut i.timeout)?;
Ok(())
}
fn visit_mut_define_api(
&mut self,
d: &mut DefineApiStatement,
) -> std::result::Result<(), Self::Error> {
self.visit_mut_expr(&mut d.path)?;
self.visit_mut_expr(&mut d.comment)?;
Ok(())
}
fn visit_mut_permission(&mut self, _: &mut super::Permission) -> Result<(), Self::Error> {
Ok(())
}
fn visit_mut_define_config(
&mut self,
_: &mut DefineConfigStatement,
) -> Result<(), Self::Error> {
Ok(())
}
fn visit_mut_define_function(
&mut self,
_: &mut DefineFunctionStatement,
) -> std::result::Result<(), Self::Error> {
Ok(())
}
fn visit_mut_define_access(
&mut self,
d: &mut DefineAccessStatement,
) -> std::result::Result<(), Self::Error> {
self.visit_mut_expr(&mut d.name)?;
self.visit_mut_expr(&mut d.comment)?;
self.visit_mut_expr(&mut d.duration.grant)?;
self.visit_mut_expr(&mut d.duration.token)?;
self.visit_mut_expr(&mut d.duration.session)?;
Ok(())
}
fn visit_mut_define_index(&mut self, d: &mut DefineIndexStatement) -> Result<(), Self::Error> {
self.visit_mut_expr(&mut d.name)?;
self.visit_mut_expr(&mut d.comment)?;
Ok(())
}
fn visit_mut_define_field(&mut self, d: &mut DefineFieldStatement) -> Result<(), Self::Error> {
self.visit_mut_expr(&mut d.name)?;
self.visit_mut_expr(&mut d.what)?;
self.visit_mut_expr(&mut d.comment)?;
Ok(())
}
}
#[revisioned(revision = 1)]
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub(crate) enum AggregateFields {
Value(Expr),
Fields(Vec<(Idiom, Expr)>),
}
#[revisioned(revision = 1)]
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub(crate) struct AggregationAnalysis {
pub(crate) aggregate_arguments: Vec<Expr>,
pub(crate) aggregations: Vec<Aggregation>,
pub(crate) group_expressions: Vec<Expr>,
pub(crate) fields: AggregateFields,
}
impl AggregationAnalysis {
pub(crate) fn analyze_fields_groups(
fields: &Fields,
groups: &Groups,
materialized_view: bool,
) -> Result<Self> {
let mut aggregations = Vec::new();
let mut exprs_map = HashMap::default();
let mut group_expressions = Vec::with_capacity(groups.len());
for g in groups.0.iter() {
group_expressions.push(Expr::Idiom(g.0.clone()))
}
let mut collect = AggregateExprCollector {
support_acummulate: !materialized_view,
within_aggregate_argument: false,
exprs_map: &mut exprs_map,
aggregations: &mut aggregations,
groups,
};
let fields = match fields {
Fields::Value(field) => {
let mut expr = field.expr.clone();
collect.visit_mut_expr(&mut expr)?;
AggregateFields::Value(expr)
}
Fields::Select(fields) => {
let mut collect_fields = Vec::with_capacity(fields.len());
for f in fields.iter() {
let Field::Single(Selector {
expr,
alias,
}) = f
else {
bail!(Error::InvalidAggregationSelector {
expr: f.to_sql()
})
};
if let Some((alias, x)) = alias.as_ref().and_then(|alias| {
group_expressions
.iter()
.position(|x| {
if let Expr::Idiom(i) = x {
*i == *alias
} else {
false
}
})
.map(|x| (alias, x))
}) {
group_expressions[x] = expr.clone();
collect_fields
.push((alias.clone(), Expr::Idiom(Idiom::field(group_field_name(x)))));
} else {
let name = alias.clone().unwrap_or_else(|| expr.to_idiom());
let mut expr = expr.clone();
collect.visit_mut_expr(&mut expr)?;
collect_fields.push((name, expr))
}
}
AggregateFields::Fields(collect_fields)
}
};
let mut aggregate_arguments = Vec::with_capacity(exprs_map.len());
for (k, v) in exprs_map {
if aggregate_arguments.len() > v {
aggregate_arguments[v] = k
} else {
for _ in aggregate_arguments.len()..v {
aggregate_arguments.push(Expr::Break)
}
aggregate_arguments.push(k)
}
}
if materialized_view
&& !aggregations.iter().any(|x| {
matches!(
x,
Aggregation::Count
| Aggregation::Mean(_)
| Aggregation::StdDev(_)
| Aggregation::Variance(_)
)
}) {
aggregations.push(Aggregation::Count)
}
Ok(Self {
aggregations,
aggregate_arguments,
group_expressions,
fields,
})
}
}