use rudb_common::{Error, LogicalType, Result, Value};
use rudb_vector::{Data, Form, Validity, Vector};
use crate::compare::order;
use crate::fallback::{self, Kernel};
use crate::number::{fit, integral, pow10, rescale};
use crate::shape::{identity, nulls_of};
#[derive(Debug, Clone)]
pub struct Accumulator {
kind: Kind,
returns: LogicalType,
state: State,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Kind {
CountStar,
Count,
Sum,
Avg,
Min,
Max,
}
#[derive(Debug, Clone)]
enum State {
Counted(i64),
Whole { total: i128, seen: bool },
Real { total: f64, seen: i64 },
Mean { whole: i128, real: f64, seen: i64, exact: bool },
Scaled { total: i128, scale: u8, seen: bool },
Extreme(Option<Value>),
}
impl Accumulator {
pub fn new(name: &str, returns: &LogicalType) -> Result<Self> {
let kind = match name {
"count_star" => Kind::CountStar,
"count" => Kind::Count,
"sum" => Kind::Sum,
"avg" => Kind::Avg,
"min" => Kind::Min,
"max" => Kind::Max,
other => {
return Err(Error::not_implemented(format!("the {other} aggregate")));
}
};
let state = match kind {
Kind::CountStar | Kind::Count => State::Counted(0),
Kind::Avg => State::Mean { whole: 0, real: 0.0, seen: 0, exact: true },
Kind::Min | Kind::Max => State::Extreme(None),
Kind::Sum => match returns {
LogicalType::Decimal { scale, .. } => {
State::Scaled { total: 0, scale: *scale, seen: false }
}
LogicalType::Float | LogicalType::Double => State::Real { total: 0.0, seen: 0 },
_ => State::Whole { total: 0, seen: false },
},
};
Ok(Self { kind, returns: returns.clone(), state })
}
pub fn update(&mut self, args: &[Value]) -> Result<()> {
if self.kind == Kind::CountStar {
if let State::Counted(count) = &mut self.state {
*count += 1;
}
return Ok(());
}
let value = match args {
[only] => only,
_ => {
return Err(Error::internal(format!("an aggregate over {} arguments", args.len())));
}
};
if value.is_null() {
return Ok(());
}
match &mut self.state {
State::Counted(count) => *count += 1,
State::Whole { total, seen } => {
let whole = integral(value).ok_or_else(|| not_narrow(value))?;
*total = total.checked_add(whole).ok_or_else(overflowed)?;
*seen = true;
}
State::Real { total, seen } => {
*total += approximate_or_error(value)?;
*seen += 1;
}
State::Mean { whole, real, seen, exact } => {
match integral(value)
.filter(|_| *exact)
.and_then(|number| whole.checked_add(number))
{
Some(total) => *whole = total,
None => {
if *exact {
*real = exactly(*whole);
*exact = false;
}
*real += approximate_or_error(value)?;
}
}
*seen += 1;
}
State::Scaled { total, scale, seen } => {
let unscaled = at_scale(value, *scale).ok_or_else(|| not_narrow(value))?;
*total = total.checked_add(unscaled).ok_or_else(overflowed)?;
*seen = true;
}
State::Extreme(held) => {
let replace = match held {
None => true,
Some(current) => {
let ordering = order(value, current)?;
match self.kind {
Kind::Min => ordering.is_lt(),
_ => ordering.is_gt(),
}
}
};
if replace {
*held = Some(value.clone());
}
}
}
Ok(())
}
pub fn update_run(&mut self, args: &[Vector], rows: usize) -> Result<()> {
if self.kind == Kind::CountStar {
if let State::Counted(count) = &mut self.state {
*count += i64::try_from(rows).map_err(|_| overlong())?;
}
return Ok(());
}
let input = match args {
[only] => only,
_ => {
return Err(Error::internal(format!("an aggregate over {} arguments", args.len())));
}
};
if input.len() < rows {
return Err(Error::internal(format!(
"an aggregate handed {rows} rows and a vector of {}",
input.len()
)));
}
if self.folded(input, rows)? {
return Ok(());
}
fallback::record(Kernel::Aggregate, input.form(), input.form());
for row in 0..rows {
let value = input.value_at(row);
self.update(std::slice::from_ref(&value))?;
}
Ok(())
}
fn folded(&mut self, input: &Vector, rows: usize) -> Result<bool> {
let nulls = nulls_of(input);
if let State::Counted(count) = &mut self.state {
*count += i64::try_from(nulls.count_valid(rows)).map_err(|_| overlong())?;
return Ok(true);
}
let least = self.kind == Kind::Min;
let want = match (&self.state, input.logical_type()) {
(State::Whole { .. }, _) => Want::Whole,
(State::Real { total, .. }, ty) => {
Want::Real { scale: decimal_scale(ty), from: *total }
}
(State::Mean { exact: true, .. }, ty) if ty.is_integer() => Want::Whole,
(State::Mean { whole, real, exact, .. }, ty) => {
let from = if *exact { exactly(*whole) } else { *real };
Want::Real { scale: decimal_scale(ty), from }
}
(State::Scaled { scale, .. }, LogicalType::Decimal { scale: held, .. })
if held == scale =>
{
Want::Whole
}
(State::Scaled { .. }, _) => return Ok(false),
(State::Extreme(_), _) => Want::Extreme(least),
(State::Counted(_), _) => return Ok(false),
};
let Some(contribution) = gather(input, rows, &nulls, want) else {
return Ok(false);
};
let live = nulls.count_valid(rows);
match (&mut self.state, contribution) {
(
State::Whole { total, seen } | State::Scaled { total, seen, .. },
Contribution::Whole(sum),
) => {
*total = total.checked_add(sum).ok_or_else(overflowed)?;
*seen |= live > 0;
}
(State::Real { total, seen }, Contribution::Real { total: carried, seen: added }) => {
*total = carried;
*seen += added;
}
(State::Mean { whole, seen, .. }, Contribution::Whole(sum)) => {
let Some(total) = whole.checked_add(sum) else { return Ok(false) };
*whole = total;
*seen += i64::try_from(live).map_err(|_| overlong())?;
}
(
State::Mean { real, seen, exact, .. },
Contribution::Real { total: carried, seen: added },
) => {
*real = carried;
*exact = false;
*seen += added;
}
(State::Extreme(held), Contribution::Extreme(Some(index))) => {
let candidate = input.value_at(index);
let replace = match held {
None => true,
Some(current) => {
let ordering = order(&candidate, current)?;
if least { ordering.is_lt() } else { ordering.is_gt() }
}
};
if replace {
*held = Some(candidate);
}
}
(State::Extreme(_), Contribution::Extreme(None)) => {}
_ => return Ok(false),
}
Ok(true)
}
pub fn finish(&self) -> Result<Value> {
match &self.state {
State::Counted(count) => Ok(Value::BigInt(*count)),
State::Whole { total, seen } => {
if !seen {
return Ok(Value::Null);
}
fit(*total, &self.returns).ok_or_else(|| {
Error::out_of_range(format!(
"a sum of {total} does not fit in {}",
self.returns
))
})
}
State::Real { total, seen } => {
if *seen == 0 {
return Ok(Value::Null);
}
#[expect(
clippy::cast_precision_loss,
reason = "the count of rows in one group is well inside the exact range"
)]
let answer = if self.kind == Kind::Avg { total / *seen as f64 } else { *total };
if matches!(self.returns, LogicalType::Float) {
#[expect(
clippy::cast_possible_truncation,
reason = "a declared FLOAT result is a FLOAT"
)]
return Ok(Value::Float(answer as f32));
}
Ok(Value::Double(answer))
}
State::Mean { whole, real, seen, exact } => {
if *seen == 0 {
return Ok(Value::Null);
}
let total = if *exact { exactly(*whole) } else { *real };
#[expect(
clippy::cast_precision_loss,
reason = "the count of rows in one group is well inside the exact range"
)]
let answer = total / *seen as f64;
if matches!(self.returns, LogicalType::Float) {
#[expect(
clippy::cast_possible_truncation,
reason = "a declared FLOAT result is a FLOAT"
)]
return Ok(Value::Float(answer as f32));
}
Ok(Value::Double(answer))
}
State::Scaled { total, scale, seen } => {
if !seen {
return Ok(Value::Null);
}
let width = match self.returns {
LogicalType::Decimal { width, .. } => width,
_ => rudb_common::MAX_DECIMAL_WIDTH,
};
Ok(Value::Decimal { unscaled: *total, width, scale: *scale })
}
State::Extreme(held) => Ok(held.clone().unwrap_or(Value::Null)),
}
}
}
fn not_narrow(value: &Value) -> Error {
Error::not_implemented(format!("summing a {}", value.logical_type()))
}
#[expect(
clippy::cast_precision_loss,
reason = "a total past 2^53 rounding once here is the definition of a double result"
)]
fn exactly(total: i128) -> f64 {
total as f64
}
fn approximate_or_error(value: &Value) -> Result<f64> {
crate::number::approximate(value).ok_or_else(|| not_narrow(value))
}
fn at_scale(value: &Value, scale: u8) -> Option<i128> {
match *value {
Value::Decimal { unscaled, scale: held, .. } => rescale(unscaled, held, scale),
_ => integral(value).and_then(|whole| whole.checked_mul(pow10(scale))),
}
}
fn overflowed() -> Error {
Error::out_of_range("Overflow in the running total of a sum".to_string())
}
fn overlong() -> Error {
Error::out_of_range("more rows in one vector than a count can hold".to_string())
}
fn decimal_scale(ty: &LogicalType) -> u8 {
match *ty {
LogicalType::Decimal { scale, .. } => scale,
_ => 0,
}
}
#[derive(Clone, Copy)]
enum Want {
Whole,
Real { scale: u8, from: f64 },
Extreme(bool),
}
enum Contribution {
Whole(i128),
Real { total: f64, seen: i64 },
Extreme(Option<usize>),
}
fn gather(input: &Vector, rows: usize, nulls: &Validity, want: Want) -> Option<Contribution> {
match input.form() {
Form::Flat => {
let data = input.data()?;
if data.len() < rows {
return None;
}
collect(data, identity, rows, nulls, want)
}
Form::Dictionary => {
let (codes, values) = input.dictionary_parts()?;
if codes.len() < rows {
return None;
}
collect(values.data()?, |index| codes[index] as usize, rows, nulls, want)
}
_ => None,
}
}
fn collect<M: Fn(usize) -> usize>(
data: &Data,
at: M,
rows: usize,
nulls: &Validity,
want: Want,
) -> Option<Contribution> {
match want {
Want::Whole => whole_sum(data, at, rows, nulls).map(Contribution::Whole),
Want::Real { scale, from } => real_sum(data, at, rows, nulls, scale, from),
Want::Extreme(least) => extreme(data, at, rows, nulls, least).map(Contribution::Extreme),
}
}
fn whole_sum<M: Fn(usize) -> usize>(
data: &Data,
at: M,
rows: usize,
nulls: &Validity,
) -> Option<i128> {
macro_rules! summed {
($(($variant:ident, $native:ty, $zero:expr)),+ $(,)?) => {
match data {
$(Data::$variant(values) => summed!(@run values),)+
_ => return None,
}
};
(@run $values:expr) => {{
let values = $values;
let mut total: i128 = 0;
match nulls {
Validity::AllValid => {
for index in 0..rows {
total += i128::from(values[at(index)]);
}
}
Validity::AllInvalid => {}
Validity::Mask(mask) => {
for start in (0..rows).step_by(64) {
let word = mask.word(start / 64);
for index in start..(start + 64).min(rows) {
let number = i128::from(values[at(index)]);
total += if word >> (index - start) & 1 == 1 { number } else { 0 };
}
}
}
}
total
}};
}
Some(rudb_vector::for_each_layout!(narrow, summed))
}
#[expect(
clippy::cast_precision_loss,
reason = "a wide integer past 2^53 losing digits is what a double is, and this is the float path"
)]
fn real_sum<M: Fn(usize) -> usize>(
data: &Data,
at: M,
rows: usize,
nulls: &Validity,
scale: u8,
from: f64,
) -> Option<Contribution> {
let factor = pow10(scale) as f64;
let scaled = scale != 0;
let all = i64::try_from(rows).ok()?;
macro_rules! added {
($(($variant:ident, $native:ty, $zero:expr)),+ $(,)?) => {
match data {
$(Data::$variant(values) => added!(@run values, |number| number as f64),)+
Data::Float32(values) => added!(@run values, f64::from),
Data::Float64(values) => added!(@run values, |number: f64| number),
_ => return None,
}
};
(@run $values:expr, $convert:expr) => {{
let values = $values;
let convert = $convert;
let mut total = from;
let mut seen: i64 = 0;
match nulls {
Validity::AllValid => {
for index in 0..rows {
let number = convert(values[at(index)]);
total += if scaled { number / factor } else { number };
}
seen = all;
}
Validity::AllInvalid => {}
Validity::Mask(mask) => {
for start in (0..rows).step_by(64) {
let word = mask.word(start / 64);
for index in start..(start + 64).min(rows) {
if word >> (index - start) & 1 == 0 {
continue;
}
let number = convert(values[at(index)]);
total += if scaled { number / factor } else { number };
seen += 1;
}
}
}
}
(total, seen)
}};
}
let (total, seen) = rudb_vector::for_each_layout!(integer, added);
Some(Contribution::Real { total, seen })
}
fn extreme<M: Fn(usize) -> usize>(
data: &Data,
at: M,
rows: usize,
nulls: &Validity,
least: bool,
) -> Option<Option<usize>> {
macro_rules! best {
($(($variant:ident, $native:ty, $zero:expr)),+ $(,)?) => {
match data {
$(Data::$variant(values) => best!(@run values),)+
_ => return None,
}
};
(@run $values:expr) => {{
let values = $values;
let mut held = usize::MAX;
let mut mark: i128 = 0;
match nulls {
Validity::AllValid => {
if rows > 0 {
mark = i128::from(values[at(0)]);
held = 0;
for index in 1..rows {
let number = i128::from(values[at(index)]);
let win = if least { number < mark } else { number > mark };
if win {
mark = number;
held = index;
}
}
}
}
Validity::AllInvalid => {}
Validity::Mask(mask) => {
for start in (0..rows).step_by(64) {
let word = mask.word(start / 64);
for index in start..(start + 64).min(rows) {
if word >> (index - start) & 1 == 0 {
continue;
}
let number = i128::from(values[at(index)]);
let win = if least { number < mark } else { number > mark };
if held == usize::MAX || win {
mark = number;
held = index;
}
}
}
}
}
(held != usize::MAX).then_some(held)
}};
}
Some(rudb_vector::for_each_layout!(narrow, best))
}
#[cfg(test)]
mod tests {
use super::*;
fn run(name: &str, returns: &LogicalType, rows: &[Value]) -> Value {
let mut accumulator = Accumulator::new(name, returns).expect("a known aggregate");
for row in rows {
accumulator.update(std::slice::from_ref(row)).expect("accumulates");
}
accumulator.finish().expect("finishes")
}
#[test]
fn count_star_counts_rows_and_count_counts_values() {
let mut stars = Accumulator::new("count_star", &LogicalType::BigInt).expect("known");
for _ in 0..3 {
stars.update(&[]).expect("no arguments");
}
assert_eq!(stars.finish().expect("finishes"), Value::BigInt(3));
let counted = run(
"count",
&LogicalType::BigInt,
&[Value::Integer(1), Value::Null, Value::Integer(3)],
);
assert_eq!(counted, Value::BigInt(2));
}
#[test]
fn a_sum_of_nothing_is_null_and_a_count_of_nothing_is_zero() {
assert_eq!(run("sum", &LogicalType::HugeInt, &[]), Value::Null);
assert_eq!(run("sum", &LogicalType::HugeInt, &[Value::Null]), Value::Null);
assert_eq!(run("count", &LogicalType::BigInt, &[]), Value::BigInt(0));
assert_eq!(run("count_star", &LogicalType::BigInt, &[]), Value::BigInt(0));
}
#[test]
fn a_sum_of_integers_accumulates_wider_than_it_reads() {
let rows = vec![Value::Integer(i32::MAX); 4];
let total = run("sum", &LogicalType::HugeInt, &rows);
assert_eq!(total, Value::HugeInt(i128::from(i32::MAX) * 4));
}
#[test]
fn an_average_divides_by_the_rows_it_saw_rather_than_the_rows_there_were() {
let average =
run("avg", &LogicalType::Double, &[Value::Integer(1), Value::Null, Value::Integer(3)]);
assert_eq!(average, Value::Double(2.0));
}
const WIDE: [i64; 4] = [435090932899640449, 435090932899640450, 1000003, 999999999999999999];
fn wide_mean() -> f64 {
exactly(WIDE.iter().map(|&number| i128::from(number)).sum()) / 4.0
}
fn wide_values() -> Vec<Value> {
WIDE.iter().map(|&number| Value::BigInt(number)).collect()
}
#[test]
fn an_average_of_whole_numbers_adds_them_up_exactly_and_divides_once() {
let mut running = 0.0_f64;
for value in wide_values() {
running += crate::number::approximate(&value).expect("a number");
}
assert_ne!(running / 4.0, wide_mean(), "the two ways of averaging have to differ here");
assert_eq!(run("avg", &LogicalType::Double, &wide_values()), Value::Double(wide_mean()));
}
#[test]
fn the_vector_path_averages_whole_numbers_exactly_as_well() {
let _turn = fallback::TURN.lock().expect("no test panics while holding this");
let values = wide_values();
let vector = Vector::from_values(LogicalType::BigInt, &values).expect("a vector of these");
let mut accumulator = Accumulator::new("avg", &LogicalType::Double).expect("a known one");
accumulator.update_run(std::slice::from_ref(&vector), values.len()).expect("folds them in");
assert_eq!(accumulator.finish().expect("finishes"), Value::Double(wide_mean()));
}
#[test]
fn an_average_of_doubles_is_the_running_total_the_float_path_produces() {
let rows = [Value::Double(1e17), Value::Double(1.0), Value::Double(3.0)];
let mut running = 0.0_f64;
for value in &rows {
running += crate::number::approximate(value).expect("a number");
}
assert_eq!(run("avg", &LogicalType::Double, &rows), Value::Double(running / 3.0));
}
#[test]
fn min_and_max_skip_nulls_and_keep_the_value_rather_than_a_number() {
let smallest = run(
"min",
&LogicalType::Varchar,
&[Value::Varchar("b".into()), Value::Null, Value::Varchar("a".into())],
);
assert_eq!(smallest, Value::Varchar("a".into()));
let largest = run(
"max",
&LogicalType::Integer,
&[Value::Integer(1), Value::Integer(7), Value::Integer(3)],
);
assert_eq!(largest, Value::Integer(7));
}
#[test]
fn a_decimal_sums_at_its_own_scale() {
let ty = LogicalType::decimal(10, 2).expect("a legal decimal");
let total = run(
"sum",
&ty,
&[
Value::Decimal { unscaled: 250, width: 10, scale: 2 },
Value::Decimal { unscaled: 125, width: 10, scale: 2 },
],
);
assert_eq!(total, Value::Decimal { unscaled: 375, width: 10, scale: 2 });
}
#[test]
fn an_aggregate_nobody_has_written_says_which_one() {
let error = Accumulator::new("median", &LogicalType::Double)
.expect_err("median is not written yet");
assert!(error.message().contains("the median aggregate"), "{error}");
}
fn row_at_a_time(name: &str, returns: &LogicalType, batches: &[Vector]) -> Result<Value> {
let mut accumulator = Accumulator::new(name, returns)?;
for batch in batches {
for row in 0..batch.len() {
let value = batch.value_at(row);
accumulator.update(std::slice::from_ref(&value))?;
}
}
accumulator.finish()
}
fn a_vector_at_a_time(name: &str, returns: &LogicalType, batches: &[Vector]) -> Result<Value> {
let mut accumulator = Accumulator::new(name, returns)?;
for batch in batches {
accumulator.update_run(std::slice::from_ref(batch), batch.len())?;
}
accumulator.finish()
}
fn agrees(name: &str, returns: &LogicalType, batches: &[Vector], note: &str) {
let slow = row_at_a_time(name, returns, batches);
let fast = a_vector_at_a_time(name, returns, batches);
match (slow, fast) {
(Ok(slow), Ok(fast)) => assert_eq!(slow, fast, "{note}"),
(Err(slow), Err(fast)) => {
assert_eq!(slow.message(), fast.message(), "{note}");
}
(slow, fast) => {
panic!(
"{note}: one path answered and the other did not, {slow:?} against {fast:?}"
);
}
}
}
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
self.0 ^= self.0 << 13;
self.0 ^= self.0 >> 7;
self.0 ^= self.0 << 17;
self.0
}
}
fn small(rng: &mut Rng) -> i64 {
(rng.next() % 201) as i64 - 100
}
fn sample(ty: &LogicalType, rng: &mut Rng) -> Value {
let number = small(rng);
let positive = number.unsigned_abs();
match *ty {
LogicalType::TinyInt => Value::TinyInt(number as i8),
LogicalType::SmallInt => Value::SmallInt(number as i16),
LogicalType::Integer => Value::Integer(number as i32),
LogicalType::BigInt => Value::BigInt(number),
LogicalType::HugeInt => Value::HugeInt(i128::from(number)),
LogicalType::UTinyInt => Value::UTinyInt(positive as u8),
LogicalType::USmallInt => Value::USmallInt(positive as u16),
LogicalType::UInteger => Value::UInteger(positive as u32),
LogicalType::UBigInt => Value::UBigInt(positive),
LogicalType::Float => Value::Float(number as f32 / 8.0),
LogicalType::Double => Value::Double(number as f64 / 8.0),
LogicalType::Decimal { width, scale } => {
Value::Decimal { unscaled: i128::from(number) * 7, width, scale }
}
LogicalType::Varchar => Value::Varchar(format!("w{number}")),
_ => panic!("no sample for {ty}"),
}
}
fn flat(ty: &LogicalType, rows: usize, nulls: usize, rng: &mut Rng) -> Vector {
let values: Vec<Value> = (0..rows)
.map(
|index| {
if nulls > 0 && index % nulls == 0 { Value::Null } else { sample(ty, rng) }
},
)
.collect();
Vector::from_values(ty.clone(), &values).expect("a vector of this type")
}
fn returns_of(name: &str, ty: &LogicalType) -> LogicalType {
match name {
"count" | "count_star" => LogicalType::BigInt,
"avg" => LogicalType::Double,
"min" | "max" => ty.clone(),
_ => match *ty {
LogicalType::Decimal { scale, .. } => {
LogicalType::decimal(rudb_common::MAX_DECIMAL_WIDTH, scale)
.expect("the widest decimal at this scale is legal")
}
LogicalType::Float | LogicalType::Double => LogicalType::Double,
_ => LogicalType::HugeInt,
},
}
}
#[test]
fn every_aggregate_over_every_type_agrees_with_the_row_at_a_time_path() {
let _turn = fallback::TURN.lock().expect("no test panics while holding this");
let mut rng = Rng(0x5eed_ca11_ab1e_0003);
let types = [
LogicalType::TinyInt,
LogicalType::SmallInt,
LogicalType::Integer,
LogicalType::BigInt,
LogicalType::HugeInt,
LogicalType::UTinyInt,
LogicalType::USmallInt,
LogicalType::UInteger,
LogicalType::UBigInt,
LogicalType::Float,
LogicalType::Double,
LogicalType::decimal(9, 2).expect("a legal decimal"),
LogicalType::decimal(18, 4).expect("a legal decimal"),
LogicalType::decimal(30, 6).expect("a legal decimal"),
LogicalType::Varchar,
];
for ty in &types {
for name in ["count_star", "count", "sum", "avg", "min", "max"] {
let returns = returns_of(name, ty);
for nulls in [0_usize, 4, 1] {
let first = flat(ty, 97, nulls, &mut rng);
let second = flat(ty, 64, nulls, &mut rng);
let note = format!("{name} over {ty}, flat, one null in {nulls}");
agrees(name, &returns, &[first.clone(), second.clone()], ¬e);
let codes: Vec<u32> = (0..97).map(|index| (index % 13) as u32).collect();
let coded = Vector::dictionary(codes, first).expect("codes are in range");
let note = format!("{name} over {ty}, dictionary, one null in {nulls}");
agrees(name, &returns, &[coded, second], ¬e);
}
}
}
}
#[test]
fn a_sum_of_numbers_stays_off_the_row_at_a_time_path_and_a_sum_of_strings_does_not() {
let _turn = fallback::TURN.lock().expect("no test panics while holding this");
fallback::reset();
let numbers = Vector::from_values(
LogicalType::Integer,
&[Value::Integer(1), Value::Integer(2), Value::Integer(3)],
)
.expect("a vector of integers");
let mut summing =
Accumulator::new("sum", &LogicalType::HugeInt).expect("a known aggregate");
summing.update_run(std::slice::from_ref(&numbers), 3).expect("sums");
assert_eq!(summing.finish().expect("finishes"), Value::HugeInt(6));
assert_eq!(fallback::count(Kernel::Aggregate, Form::Flat, Form::Flat), 0);
let words = Vector::from_values(
LogicalType::Varchar,
&[Value::Varchar("a".into()), Value::Null, Value::Varchar("b".into())],
)
.expect("a vector of strings");
let mut counting = Accumulator::new("count", &LogicalType::BigInt).expect("a known one");
counting.update_run(std::slice::from_ref(&words), 3).expect("counts");
assert_eq!(counting.finish().expect("finishes"), Value::BigInt(2));
assert_eq!(fallback::count(Kernel::Aggregate, Form::Flat, Form::Flat), 0);
let mut wrong = Accumulator::new("sum", &LogicalType::HugeInt).expect("a known aggregate");
let error =
wrong.update_run(std::slice::from_ref(&words), 3).expect_err("cannot sum those");
assert!(error.message().contains("summing a"), "{error}");
assert_eq!(fallback::count(Kernel::Aggregate, Form::Flat, Form::Flat), 1);
fallback::reset();
}
#[test]
fn a_floating_point_sum_carries_the_running_total_into_the_next_vector() {
let first =
Vector::from_values(LogicalType::Double, &[Value::Double(1.0e16)]).expect("a vector");
let second = Vector::from_values(LogicalType::Double, &vec![Value::Double(1.0); 8])
.expect("a vector");
let batches = [first, second];
let slow = row_at_a_time("sum", &LogicalType::Double, &batches).expect("sums");
let fast = a_vector_at_a_time("sum", &LogicalType::Double, &batches).expect("sums");
assert_eq!(slow, fast);
assert_eq!(slow, Value::Double(1.0e16));
assert_ne!(1.0e16 + 8.0, 1.0e16);
}
#[test]
fn a_null_behind_a_dictionary_code_is_skipped_by_every_aggregate() {
let values = Vector::from_values(
LogicalType::Integer,
&[Value::Null, Value::Integer(5), Value::Integer(9)],
)
.expect("a vector of integers");
let coded = Vector::dictionary(vec![0, 1, 0, 2, 0], values).expect("codes are in range");
let batch = std::slice::from_ref(&coded);
assert_eq!(
a_vector_at_a_time("count", &LogicalType::BigInt, batch).expect("counts"),
Value::BigInt(2)
);
assert_eq!(
a_vector_at_a_time("sum", &LogicalType::HugeInt, batch).expect("sums"),
Value::HugeInt(14)
);
assert_eq!(
a_vector_at_a_time("min", &LogicalType::Integer, batch).expect("finds one"),
Value::Integer(5)
);
}
#[test]
fn a_total_of_hugeints_goes_the_row_at_a_time_way_and_still_overflows() {
let _turn = fallback::TURN.lock().expect("no test panics while holding this");
fallback::reset();
let rows = vec![Value::HugeInt(i128::MAX); 2];
let vector = Vector::from_values(LogicalType::HugeInt, &rows).expect("a vector");
let mut accumulator = Accumulator::new("sum", &LogicalType::HugeInt).expect("a known one");
let error =
accumulator.update_run(std::slice::from_ref(&vector), 2).expect_err("overflows");
assert!(error.message().contains("Overflow in the running total"), "{error}");
assert_eq!(fallback::count(Kernel::Aggregate, Form::Flat, Form::Flat), 1);
fallback::reset();
}
}