use std::sync::Arc;
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};
pub const NOWHERE: usize = usize::MAX;
#[derive(Debug, Clone)]
pub struct Accumulator {
state: State,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Return {
TinyInt,
SmallInt,
Integer,
BigInt,
HugeInt,
UTinyInt,
USmallInt,
UInteger,
UBigInt,
UHugeInt,
Float,
Double,
Decimal(u8),
Other,
}
impl Return {
fn new(ty: &LogicalType) -> Self {
match ty {
LogicalType::TinyInt => Self::TinyInt,
LogicalType::SmallInt => Self::SmallInt,
LogicalType::Integer => Self::Integer,
LogicalType::BigInt => Self::BigInt,
LogicalType::HugeInt => Self::HugeInt,
LogicalType::UTinyInt => Self::UTinyInt,
LogicalType::USmallInt => Self::USmallInt,
LogicalType::UInteger => Self::UInteger,
LogicalType::UBigInt => Self::UBigInt,
LogicalType::UHugeInt => Self::UHugeInt,
LogicalType::Float => Self::Float,
LogicalType::Double => Self::Double,
LogicalType::Decimal { width, .. } => Self::Decimal(*width),
_ => Self::Other,
}
}
fn logical(self) -> LogicalType {
match self {
Self::TinyInt => LogicalType::TinyInt,
Self::SmallInt => LogicalType::SmallInt,
Self::Integer => LogicalType::Integer,
Self::BigInt => LogicalType::BigInt,
Self::HugeInt => LogicalType::HugeInt,
Self::UTinyInt => LogicalType::UTinyInt,
Self::USmallInt => LogicalType::USmallInt,
Self::UInteger => LogicalType::UInteger,
Self::UBigInt => LogicalType::UBigInt,
Self::UHugeInt => LogicalType::UHugeInt,
Self::Float => LogicalType::Float,
Self::Double => LogicalType::Double,
Self::Decimal(width) => LogicalType::Decimal { width, scale: 0 },
Self::Other => LogicalType::Null,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Kind {
CountStar,
Count,
Sum,
Avg,
Min,
Max,
}
#[derive(Debug, Clone)]
enum State {
Counted { count: i64, star: bool },
Whole { total: i128, seen: bool, returns: Return },
Real { total: f64, seen: i64, kind: Kind, returns: Return },
Mean { total: i128, seen: i64, exact: bool, returns: Return },
Scaled { total: i128, scale: u8, seen: bool, returns: Return },
Extreme { held: Option<Box<Extremum>>, least: bool },
}
#[derive(Debug, Clone)]
enum Extremum {
Held(Value),
Ranked { dictionary: Arc<Vector>, code: u32, rank: u32 },
}
impl Extremum {
fn value(&self) -> Result<Value> {
match self {
Self::Held(value) => Ok(value.clone()),
Self::Ranked { dictionary, code, .. } => dictionary.try_value_at(*code as usize),
}
}
fn settle(&mut self) -> Result<&mut Value> {
if let Self::Ranked { .. } = self {
let value = self.value()?;
*self = Self::Held(value);
}
match self {
Self::Held(value) => Ok(value),
Self::Ranked { .. } => {
Err(Error::internal("a settled extreme kept its rank".to_string()))
}
}
}
fn offer(&mut self, dictionary: &Arc<Vector>, code: u32, rank: u32, least: bool) -> Result<()> {
if let Self::Ranked { dictionary: mine, code: held_code, rank: held_rank } = self {
if Arc::ptr_eq(mine, dictionary) {
if if least { rank < *held_rank } else { rank > *held_rank } {
*held_code = code;
*held_rank = rank;
}
return Ok(());
}
}
let candidate = dictionary.try_value_at(code as usize)?;
let ordering = order(&candidate, self.settle()?)?;
if if least { ordering.is_lt() } else { ordering.is_gt() } {
*self = Self::Held(candidate);
}
Ok(())
}
}
impl Accumulator {
#[must_use]
pub fn counted(&self) -> Option<i64> {
match self.state {
State::Counted { count, .. } => Some(count),
_ => None,
}
}
#[must_use]
pub fn exact_sum(total: i128, seen: bool, returns: &LogicalType) -> Self {
Self { state: State::Whole { total, seen, returns: Return::new(returns) } }
}
#[must_use]
pub fn exact_avg(total: i128, seen: i64, returns: &LogicalType) -> Self {
Self { state: State::Mean { total, seen, exact: true, returns: Return::new(returns) } }
}
fn kind(&self) -> Kind {
match self.state {
State::Counted { star, .. } => {
if star {
Kind::CountStar
} else {
Kind::Count
}
}
State::Whole { .. } | State::Scaled { .. } => Kind::Sum,
State::Real { kind, .. } => kind,
State::Mean { .. } => Kind::Avg,
State::Extreme { least, .. } => {
if least {
Kind::Min
} else {
Kind::Max
}
}
}
}
fn returns(&self) -> Return {
match self.state {
State::Counted { .. } => Return::BigInt,
State::Whole { returns, .. }
| State::Real { returns, .. }
| State::Mean { returns, .. }
| State::Scaled { returns, .. } => returns,
State::Extreme { .. } => Return::Other,
}
}
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 scale = match returns {
LogicalType::Decimal { scale, .. } => *scale,
_ => 0,
};
let returns = Return::new(returns);
let state = match kind {
Kind::CountStar | Kind::Count => {
State::Counted { count: 0, star: kind == Kind::CountStar }
}
Kind::Avg => State::Mean { total: 0, seen: 0, exact: true, returns },
Kind::Min | Kind::Max => State::Extreme { held: None, least: kind == Kind::Min },
Kind::Sum => match returns {
Return::Decimal(_) => State::Scaled { total: 0, scale, seen: false, returns },
Return::Float | Return::Double => {
State::Real { total: 0.0, seen: 0, kind, returns }
}
_ => State::Whole { total: 0, seen: false, returns },
},
};
Ok(Self { 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 { total, seen, exact, .. } => {
match integral(value)
.filter(|_| *exact)
.and_then(|number| total.checked_add(number))
{
Some(sum) => *total = sum,
None => {
let real = if *exact { exactly(*total) } else { mean_real(*total) };
*total = mean_bits(real + approximate_or_error(value)?);
*exact = false;
}
}
*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, least } => {
let replace = match held {
None => true,
Some(current) => {
let ordering = order(value, current.settle()?)?;
if *least { ordering.is_lt() } else { ordering.is_gt() }
}
};
if replace {
*held = Some(Box::new(Extremum::Held(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.try_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 { total, exact, .. }, ty) => {
let from = if *exact { exactly(*total) } else { mean_real(*total) };
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 { total, seen, .. }, Contribution::Whole(sum)) => {
let Some(sum) = total.checked_add(sum) else { return Ok(false) };
*total = sum;
*seen += i64::try_from(live).map_err(|_| overlong())?;
}
(
State::Mean { total, seen, exact, .. },
Contribution::Real { total: carried, seen: added },
) => {
*total = mean_bits(carried);
*exact = false;
*seen += added;
}
(State::Extreme { held, .. }, Contribution::Extreme(Some(index))) => {
let candidate = input.try_value_at(index)?;
let replace = match held {
None => true,
Some(current) => {
let ordering = order(&candidate, current.settle()?)?;
if least { ordering.is_lt() } else { ordering.is_gt() }
}
};
if replace {
*held = Some(Box::new(Extremum::Held(candidate)));
}
}
(State::Extreme { .. }, Contribution::Extreme(None)) => {}
_ => return Ok(false),
}
Ok(true)
}
pub fn combine(&mut self, other: &Self) -> Result<()> {
match (&mut self.state, &other.state) {
(State::Counted { count, star }, State::Counted { count: added, star: same })
if star == same =>
{
*count += added;
}
(State::Whole { total, seen, .. }, State::Whole { total: added, seen: any, .. }) => {
*total = total.checked_add(*added).ok_or_else(overflowed)?;
*seen |= any;
}
(
State::Scaled { total, scale, seen, .. },
State::Scaled { total: added, scale: same, seen: any, .. },
) if scale == same => {
*total = total.checked_add(*added).ok_or_else(overflowed)?;
*seen |= any;
}
(State::Real { total, seen, .. }, State::Real { total: added, seen: more, .. }) => {
*total += added;
*seen += more;
}
(
State::Mean { total, seen, exact, .. },
State::Mean { total: added, seen: more, exact: whole, .. },
) => {
let both = if *exact && *whole { total.checked_add(*added) } else { None };
match both {
Some(sum) => *total = sum,
None => {
let here = if *exact { exactly(*total) } else { mean_real(*total) };
let there = if *whole { exactly(*added) } else { mean_real(*added) };
*total = mean_bits(here + there);
*exact = false;
}
}
*seen += more;
}
(State::Extreme { held, least }, State::Extreme { held: candidate, least: same })
if least == same =>
{
if let Some(candidate) = candidate {
match (held.as_deref_mut(), candidate.as_ref()) {
(None, _) => *held = Some(candidate.clone()),
(
Some(Extremum::Ranked { dictionary, rank, .. }),
Extremum::Ranked { dictionary: theirs, rank: other, .. },
) if Arc::ptr_eq(dictionary, theirs) => {
if if *least { other < rank } else { other > rank } {
*held = Some(candidate.clone());
}
}
(Some(current), _) => {
let value = candidate.value()?;
let ordering = order(&value, current.settle()?)?;
if if *least { ordering.is_lt() } else { ordering.is_gt() } {
*held = Some(candidate.clone());
}
}
}
}
}
(here, there) => {
return Err(Error::internal(format!(
"combining a {here:?} aggregate state with a {there:?} one, which are not the \
same aggregate over the same type"
)));
}
}
Ok(())
}
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);
}
let returns = self.returns().logical();
fit(*total, &returns).ok_or_else(|| {
Error::out_of_range(format!("a sum of {total} does not fit in {}", 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 self.returns() == Return::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 { total, seen, exact, .. } => {
if *seen == 0 {
return Ok(Value::Null);
}
let total = if *exact { exactly(*total) } else { mean_real(*total) };
#[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 self.returns() == Return::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() {
Return::Decimal(width) => width,
_ => rudb_common::MAX_DECIMAL_WIDTH,
};
Ok(Value::Decimal { unscaled: *total, width, scale: *scale })
}
State::Extreme { held, .. } => held.as_deref().map_or(Ok(Value::Null), Extremum::value),
}
}
pub fn finish_offset(&self, offset: i64, rows: i64) -> Result<Value> {
let Value::HugeInt(total) = self.finish()? else {
return Ok(Value::Null);
};
let added = i128::from(offset).checked_mul(i128::from(rows)).ok_or_else(overflowed)?;
Ok(Value::HugeInt(total.checked_add(added).ok_or_else(overflowed)?))
}
}
pub fn update_scattered(
states: &mut [Accumulator],
slots: &[usize],
stride: usize,
offset: usize,
input: Option<&Vector>,
rows: usize,
) -> Result<()> {
if states.is_empty() {
return Ok(());
}
if slots.len() < rows {
return Err(Error::internal(format!(
"an aggregate handed {rows} rows and {} slots",
slots.len()
)));
}
let Some(first) = states.get(offset) else {
return Err(Error::internal(format!(
"an aggregate at {offset} of {} accumulators",
states.len()
)));
};
let kind = first.kind();
let extreme = match first.state {
State::Extreme { least, .. } => Some(least),
_ => None,
};
let into = Where { slots, stride, offset };
if kind == Kind::CountStar {
for row in 0..rows {
let Some(index) = into.index(row) else { continue };
if let State::Counted { count, .. } = &mut states[index].state {
*count += 1;
}
}
return Ok(());
}
let Some(input) = input else {
return Err(Error::internal("an aggregate over 0 arguments".to_string()));
};
if input.len() < rows {
return Err(Error::internal(format!(
"an aggregate handed {rows} rows and a vector of {}",
input.len()
)));
}
let nulls = nulls_of(input);
if matches!(nulls, Validity::AllInvalid) {
return Ok(());
}
if let Some(least) = extreme {
if ranked_extremes(states, into, input, rows, &nulls, least)? {
return Ok(());
}
}
if extreme.is_some()
&& input.logical_type() == &LogicalType::Varchar
&& matches!(input.form(), Form::Flat | Form::Dictionary | Form::StringView | Form::Rle)
{
for row in 0..rows {
if !nulls.is_valid(row) {
continue;
}
let Some(index) = into.index(row) else { continue };
let bytes = input.try_bytes_at(row)?.ok_or_else(|| {
Error::internal("a valid varchar row had no borrowed text".to_string())
})?;
let State::Extreme { held, least } = &mut states[index].state else {
return Err(Error::internal("a string extreme into another state".to_string()));
};
match held {
Some(current) => {
let Value::Varchar(previous) = current.settle()? else {
return Err(Error::internal(
"a varchar extreme held another type".to_string(),
));
};
let better = if *least {
bytes < previous.as_bytes()
} else {
bytes > previous.as_bytes()
};
if better {
previous.clear();
previous.push_str(utf8(bytes)?);
}
}
None => {
let text = Value::Varchar(utf8(bytes)?.to_owned());
*held = Some(Box::new(Extremum::Held(text)));
}
}
}
return Ok(());
}
let Some(first) = states.get(offset) else {
return Err(Error::internal(format!(
"an aggregate at {offset} of {} accumulators",
states.len()
)));
};
let feed = feed_of(first, input.logical_type());
if let Some(feed) = feed {
if spread(states, into, input, rows, &nulls, feed)? {
return Ok(());
}
}
fallback::record(Kernel::Aggregate, input.form(), input.form());
for row in 0..rows {
let Some(index) = into.index(row) else { continue };
let value = input.try_value_at(row)?;
states[index].update(std::slice::from_ref(&value))?;
}
Ok(())
}
fn ranked_extremes(
states: &mut [Accumulator],
into: Where<'_>,
input: &Vector,
rows: usize,
nulls: &Validity,
least: bool,
) -> Result<bool> {
let Some((codes, dictionary)) = input.shared_dictionary_parts() else { return Ok(false) };
let Some(ranks) = dictionary.code_ranks() else { return Ok(false) };
if codes.len() < rows {
return Ok(false);
}
for (row, &code) in codes.iter().enumerate().take(rows) {
if !nulls.is_valid(row) {
continue;
}
let Some(index) = into.index(row) else { continue };
let Some(&rank) = ranks.get(code as usize) else { return Ok(false) };
let State::Extreme { held, .. } = &mut states[index].state else {
return Err(Error::internal("a ranked extreme into another state".to_string()));
};
match held {
Some(current) => current.offer(dictionary, code, rank, least)?,
None => {
let kept = Extremum::Ranked { dictionary: dictionary.clone(), code, rank };
*held = Some(Box::new(kept));
}
}
}
Ok(true)
}
fn utf8(bytes: &[u8]) -> Result<&str> {
std::str::from_utf8(bytes)
.map_err(|error| Error::conversion(format!("invalid UTF-8 in VARCHAR: {error}")))
}
#[derive(Clone, Copy)]
struct Where<'w> {
slots: &'w [usize],
stride: usize,
offset: usize,
}
impl Where<'_> {
fn index(self, row: usize) -> Option<usize> {
let slot = self.slots[row];
(slot != NOWHERE).then(|| slot * self.stride + self.offset)
}
}
#[derive(Clone, Copy)]
enum Feed {
Counted,
Whole,
Real { scale: u8 },
Extreme(bool),
}
fn feed_of(first: &Accumulator, ty: &LogicalType) -> Option<Feed> {
match (&first.state, ty) {
(State::Counted { .. }, _) => Some(Feed::Counted),
(State::Whole { .. }, _) => Some(Feed::Whole),
(State::Scaled { scale, .. }, LogicalType::Decimal { scale: held, .. })
if held == scale =>
{
Some(Feed::Whole)
}
(State::Scaled { .. }, _) => None,
(State::Mean { .. }, ty) if ty.is_integer() => Some(Feed::Whole),
(State::Mean { .. } | State::Real { .. }, ty) => {
Some(Feed::Real { scale: decimal_scale(ty) })
}
(State::Extreme { .. }, ty) if ty.is_integer() => {
Some(Feed::Extreme(first.kind() == Kind::Min))
}
(State::Extreme { .. }, _) => None,
}
}
fn spread(
states: &mut [Accumulator],
into: Where<'_>,
input: &Vector,
rows: usize,
nulls: &Validity,
feed: Feed,
) -> Result<bool> {
if matches!(feed, Feed::Counted) {
for row in 0..rows {
if !nulls.is_valid(row) {
continue;
}
let Some(index) = into.index(row) else { continue };
if let State::Counted { count, .. } = &mut states[index].state {
*count += 1;
}
}
return Ok(true);
}
match input.form() {
Form::Flat => {
let Some(data) = input.data() else { return Ok(false) };
if data.len() < rows {
return Ok(false);
}
let run = Run { input, data, rows, nulls };
scatter(states, into, &run, identity, feed)
}
Form::Dictionary | Form::Rle => {
let Some((codes, values)) = input.positions() else { return Ok(false) };
if codes.len() < rows {
return Ok(false);
}
let Some(data) = values.data() else { return Ok(false) };
let run = Run { input, data, rows, nulls };
scatter(states, into, &run, |index| codes[index] as usize, feed)
}
_ => Ok(false),
}
}
struct Run<'r> {
input: &'r Vector,
data: &'r Data,
rows: usize,
nulls: &'r Validity,
}
fn scatter<M: Fn(usize) -> usize>(
states: &mut [Accumulator],
into: Where<'_>,
run: &Run<'_>,
at: M,
feed: Feed,
) -> Result<bool> {
match feed {
Feed::Counted => Ok(true),
Feed::Whole => whole_into(states, into, run, at),
Feed::Real { scale } => real_into(states, into, run, at, scale),
Feed::Extreme(least) => extreme_into(states, into, run, at, least),
}
}
fn whole_into<M: Fn(usize) -> usize>(
states: &mut [Accumulator],
into: Where<'_>,
run: &Run<'_>,
at: M,
) -> Result<bool> {
macro_rules! each {
($(($variant:ident, $native:ty, $zero:expr)),+ $(,)?) => {
match run.data {
$(Data::$variant(values) => {
for row in 0..run.rows {
if !run.nulls.is_valid(row) {
continue;
}
let Some(index) = into.index(row) else { continue };
fold_whole(&mut states[index], i128::from(values[at(row)]))?;
}
})+
_ => return Ok(false),
}
};
}
rudb_vector::for_each_layout!(narrow, each);
Ok(true)
}
fn fold_whole(into: &mut Accumulator, number: i128) -> Result<()> {
match &mut into.state {
State::Whole { total, seen, .. } | State::Scaled { total, seen, .. } => {
*total = total.checked_add(number).ok_or_else(overflowed)?;
*seen = true;
}
State::Mean { total, seen, exact, .. } => {
match total.checked_add(number).filter(|_| *exact) {
Some(sum) => *total = sum,
None => {
let real = if *exact { exactly(*total) } else { mean_real(*total) };
*total = mean_bits(real + exactly(number));
*exact = false;
}
}
*seen += 1;
}
other => {
return Err(Error::internal(format!("an exact total into {other:?}")));
}
}
Ok(())
}
#[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_into<M: Fn(usize) -> usize>(
states: &mut [Accumulator],
into: Where<'_>,
run: &Run<'_>,
at: M,
scale: u8,
) -> Result<bool> {
let factor = pow10(scale) as f64;
let scaled = scale != 0;
macro_rules! each {
($(($variant:ident, $native:ty, $zero:expr)),+ $(,)?) => {
match run.data {
$(Data::$variant(values) => each!(@run values, |number| number as f64),)+
Data::Float32(values) => each!(@run values, f64::from),
Data::Float64(values) => each!(@run values, |number: f64| number),
_ => return Ok(false),
}
};
(@run $values:expr, $convert:expr) => {{
let values = $values;
let convert = $convert;
for row in 0..run.rows {
if !run.nulls.is_valid(row) {
continue;
}
let Some(index) = into.index(row) else { continue };
let number = convert(values[at(row)]);
fold_real(&mut states[index], if scaled { number / factor } else { number });
}
}};
}
rudb_vector::for_each_layout!(integer, each);
Ok(true)
}
fn fold_real(into: &mut Accumulator, number: f64) {
match &mut into.state {
State::Real { total, seen, .. } => {
*total += number;
*seen += 1;
}
State::Mean { total, seen, exact, .. } => {
let real = if *exact { exactly(*total) } else { mean_real(*total) };
*total = mean_bits(real + number);
*exact = false;
*seen += 1;
}
_ => {}
}
}
fn extreme_into<M: Fn(usize) -> usize>(
states: &mut [Accumulator],
into: Where<'_>,
run: &Run<'_>,
at: M,
least: bool,
) -> Result<bool> {
macro_rules! each {
($(($variant:ident, $native:ty, $zero:expr)),+ $(,)?) => {
match run.data {
$(Data::$variant(values) => {
for row in 0..run.rows {
if !run.nulls.is_valid(row) {
continue;
}
let Some(index) = into.index(row) else { continue };
let number = i128::from(values[at(row)]);
let State::Extreme { held, .. } = &mut states[index].state else {
return Err(Error::internal("an extreme into a total".to_string()));
};
let replace = match held {
None => true,
Some(current) => {
let current: &Value = current.settle()?;
let mark = integral(current).ok_or_else(|| not_narrow(current))?;
if least { number < mark } else { number > mark }
}
};
if replace {
let value = run.input.try_value_at(row)?;
*held = Some(Box::new(Extremum::Held(value)));
}
}
})+
_ => return Ok(false),
}
};
}
rudb_vector::for_each_layout!(narrow, each);
Ok(true)
}
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 mean_bits(total: f64) -> i128 {
i128::from(total.to_bits())
}
fn mean_real(bits: i128) -> f64 {
f64::from_bits(bits as u64)
}
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::<true, _>(data, identity, rows, nulls, want)
}
Form::Dictionary | Form::Rle => {
let (codes, values) = input.positions()?;
let codes = codes.get(..rows)?;
let Some(data) = values.data() else {
return match want {
Want::Extreme(least) => {
extreme_bytes(values, codes, nulls, least).map(Contribution::Extreme)
}
_ => None,
};
};
if let (Want::Whole, Validity::AllValid) = (want, nulls) {
if let Some(total) = tally(data, codes) {
return Some(Contribution::Whole(total));
}
}
collect::<false, _>(data, |index| codes[index] as usize, rows, nulls, want)
}
Form::BitPacked => {
let packed = input.packed_parts()?;
match want {
Want::Whole => {
let mut total = 0_i128;
for row in 0..rows {
if nulls.is_valid(row) {
total += packed.base() + i128::from(packed.code(row));
}
}
Some(Contribution::Whole(total))
}
Want::Real { from, .. } => {
let mut total = from;
let mut seen = 0_i64;
for row in 0..rows {
if nulls.is_valid(row) {
total += (packed.base() + i128::from(packed.code(row))) as f64;
seen += 1;
}
}
Some(Contribution::Real { total, seen })
}
Want::Extreme(least) => {
let mut found: Option<(usize, u64)> = None;
for row in 0..rows {
if !nulls.is_valid(row) {
continue;
}
let code = packed.code(row);
if found.is_none_or(
|(_, held)| {
if least { code < held } else { code > held }
},
) {
found = Some((row, code));
}
}
Some(Contribution::Extreme(found.map(|(row, _)| row)))
}
}
}
_ => None,
}
}
fn extreme_bytes(
values: &Vector,
codes: &[u32],
nulls: &Validity,
least: bool,
) -> Option<Option<usize>> {
if !matches!(values.logical_type(), LogicalType::Varchar | LogicalType::Blob) {
return None;
}
if let Some(ranks) = values.code_ranks() {
return extreme_ranked(ranks, codes, nulls, least);
}
let mut winner: Option<(usize, u32)> = None;
let mut best: Vec<u8> = Vec::new();
for (row, &code) in codes.iter().enumerate() {
if !nulls.is_valid(row) {
continue;
}
if let Some((_, held)) = winner {
if held == code {
continue;
}
}
let candidate = values.try_bytes_at(code as usize).ok()??;
let ahead = match winner {
None => true,
Some(_) => {
let ordering = candidate.cmp(best.as_slice());
if least { ordering.is_lt() } else { ordering.is_gt() }
}
};
if ahead {
best.clear();
best.extend_from_slice(candidate);
winner = Some((row, code));
}
}
Some(winner.map(|(row, _)| row))
}
fn extreme_ranked(
ranks: &[u32],
codes: &[u32],
nulls: &Validity,
least: bool,
) -> Option<Option<usize>> {
let mut winner: Option<(usize, u32)> = None;
for (row, &code) in codes.iter().enumerate() {
if !nulls.is_valid(row) {
continue;
}
let &rank = ranks.get(code as usize)?;
if winner.is_none_or(|(_, held)| if least { rank < held } else { rank > held }) {
winner = Some((row, rank));
}
}
Some(winner.map(|(row, _)| row))
}
const TALLY_LIMIT: usize = 256;
fn tally(data: &Data, codes: &[u32]) -> Option<i128> {
match data.len() {
0..=8 => tally_into::<8>(data, codes),
9..=32 => tally_into::<32>(data, codes),
33..=128 => tally_into::<128>(data, codes),
129..=TALLY_LIMIT => tally_into::<TALLY_LIMIT>(data, codes),
_ => None,
}
}
fn tally_into<const SLOTS: usize>(data: &Data, codes: &[u32]) -> Option<i128> {
macro_rules! padded {
($(($variant:ident, $native:ty, $zero:expr)),+ $(,)?) => {
match data {
$(Data::$variant(values) => padded!(@run values, $zero),)+
_ => return None,
}
};
(@run $values:expr, $zero:expr) => {{
let values = $values.as_slice();
if values.len() > SLOTS {
return None;
}
let mut table = [$zero; SLOTS];
table[..values.len()].copy_from_slice(values);
let mut total: i128 = 0;
for &code in codes {
total += i128::from(table[code as usize & (SLOTS - 1)]);
}
total
}};
}
Some(rudb_vector::for_each_layout!(narrow, padded))
}
fn straight<const DIRECT: bool, T>(values: &[T], rows: usize) -> Option<&[T]> {
if DIRECT { values.get(..rows) } else { None }
}
fn collect<const DIRECT: bool, M: Fn(usize) -> usize>(
data: &Data,
at: M,
rows: usize,
nulls: &Validity,
want: Want,
) -> Option<Contribution> {
match want {
Want::Whole => whole_sum::<DIRECT, M>(data, at, rows, nulls).map(Contribution::Whole),
Want::Real { scale, from } => real_sum::<DIRECT, M>(data, at, rows, nulls, scale, from),
Want::Extreme(least) => {
extreme::<DIRECT, M>(data, at, rows, nulls, least).map(Contribution::Extreme)
}
}
}
fn whole_sum<const DIRECT: bool, 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.as_slice();
let run = straight::<DIRECT, _>(values, rows);
let mut total: i128 = 0;
match nulls {
Validity::AllValid => match run {
Some(run) => {
for &value in run {
total += i128::from(value);
}
}
None => {
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 = match run {
Some(run) => i128::from(run[index]),
None => 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<const DIRECT: bool, 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.as_slice();
let convert = $convert;
let run = straight::<DIRECT, _>(values, rows);
let mut total = from;
let mut seen: i64 = 0;
match nulls {
Validity::AllValid => {
match run {
Some(run) => {
for &value in run {
let number = convert(value);
total += if scaled { number / factor } else { number };
}
}
None => {
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 = match run {
Some(run) => convert(run[index]),
None => 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<const DIRECT: bool, 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.as_slice();
let run = straight::<DIRECT, _>(values, rows);
let mut held = usize::MAX;
let mut mark: i128 = 0;
match nulls {
Validity::AllValid => match run {
Some(run) if !run.is_empty() => {
mark = i128::from(run[0]);
held = 0;
for (index, &value) in run.iter().enumerate().skip(1) {
let number = i128::from(value);
let win = if least { number < mark } else { number > mark };
if win {
mark = number;
held = index;
}
}
}
Some(_) => {}
None => {
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 = match run {
Some(run) => i128::from(run[index]),
None => 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::*;
#[test]
fn a_grouped_accumulator_does_not_carry_a_full_logical_type() {
assert!(size_of::<Accumulator>() <= 32, "{} bytes", size_of::<Accumulator>());
}
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 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 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.clone()).expect("in range");
let note = format!("{name} over {ty}, dictionary, one null in {nulls}");
agrees(name, &returns, &[coded, second.clone()], ¬e);
let ends: Vec<u32> = (1..=8).map(|run| (run * 13).min(97)).collect();
let runs = Vector::runs(ends, first.slice(0, 8).expect("eight values"))
.expect("one value for each run");
let note = format!("{name} over {ty}, runs, one null in {nulls}");
agrees(name, &returns, &[runs, second], ¬e);
}
}
}
}
const STRIDE: usize = 3;
const OFFSET: usize = 1;
fn deal(rows: usize, groups: usize) -> Vec<usize> {
(0..rows).map(|row| if row % 11 == 5 { NOWHERE } else { (row * 7 + 3) % groups }).collect()
}
fn group_at_a_time(
name: &str,
returns: &LogicalType,
batches: &[(Vector, Vec<usize>)],
groups: usize,
reads: bool,
) -> Result<Vec<Value>> {
let mut states = Vec::new();
for _ in 0..groups {
states.push(Accumulator::new(name, returns)?);
}
for (batch, slots) in batches {
for (row, &slot) in slots.iter().enumerate() {
if slot == NOWHERE {
continue;
}
if reads {
let value = batch.value_at(row);
states[slot].update(std::slice::from_ref(&value))?;
} else {
states[slot].update(&[])?;
}
}
}
states.iter().map(Accumulator::finish).collect()
}
fn group_at_once(
name: &str,
returns: &LogicalType,
batches: &[(Vector, Vec<usize>)],
groups: usize,
reads: bool,
) -> Result<Vec<Value>> {
let mut states = Vec::new();
for _ in 0..groups * STRIDE {
states.push(Accumulator::new(name, returns)?);
}
for (batch, slots) in batches {
let input = reads.then_some(batch);
update_scattered(&mut states, slots, STRIDE, OFFSET, input, slots.len())?;
}
(0..groups).map(|group| states[group * STRIDE + OFFSET].finish()).collect()
}
#[test]
fn every_aggregate_scattered_into_groups_agrees_with_one_accumulator_per_group() {
let mut rng = Rng(0x5eed_ca11_ab1e_0061);
let groups = 5;
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);
let reads = name != "count_star";
for nulls in [0_usize, 4, 1] {
let first = flat(ty, 97, nulls, &mut rng);
let second = flat(ty, 64, nulls, &mut rng);
let codes: Vec<u32> = (0..97).map(|index| (index % 13) as u32).collect();
let coded = Vector::dictionary(codes, first.clone()).expect("codes in range");
for (shape, batches) in [
("flat", vec![first.clone(), second.clone()]),
("dictionary", vec![coded, second.clone()]),
] {
let dealt: Vec<(Vector, Vec<usize>)> = batches
.into_iter()
.map(|batch| {
let slots = deal(batch.len(), groups);
(batch, slots)
})
.collect();
let note = format!("{name} over {ty}, {shape}, one null in {nulls}");
let slow = group_at_a_time(name, &returns, &dealt, groups, reads);
let fast = group_at_once(name, &returns, &dealt, groups, reads);
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:?}"
),
}
}
}
}
}
}
#[test]
fn a_row_that_belongs_to_no_group_is_counted_by_nobody() {
let column = Vector::from_values(
LogicalType::Integer,
&[Value::Integer(1), Value::Integer(2), Value::Integer(3), Value::Integer(4)],
)
.expect("a vector of integers");
let mut states = vec![Accumulator::new("sum", &LogicalType::HugeInt).expect("known"); 2];
let slots = [0, NOWHERE, 1, NOWHERE];
update_scattered(&mut states, &slots, 1, 0, Some(&column), 4).expect("folds them in");
assert_eq!(states[0].finish().expect("finishes"), Value::HugeInt(1));
assert_eq!(states[1].finish().expect("finishes"), Value::HugeInt(3));
}
#[test]
fn the_shapes_a_grouped_query_is_made_of_stay_off_the_row_at_a_time_path() {
let numbers = Vector::from_values(
LogicalType::Integer,
&[Value::Integer(1), Value::Integer(2), Value::Integer(3)],
)
.expect("a vector of integers");
let words = Vector::from_values(
LogicalType::Varchar,
&[Value::Varchar("a".into()), Value::Varchar("b".into()), Value::Varchar("c".into())],
)
.expect("a vector of strings");
let slots = [0_usize, 1, 0];
for (name, returns, column) in [
("count_star", LogicalType::BigInt, None),
("count", LogicalType::BigInt, Some(&numbers)),
("sum", LogicalType::HugeInt, Some(&numbers)),
("avg", LogicalType::Double, Some(&numbers)),
("min", LogicalType::Integer, Some(&numbers)),
("max", LogicalType::Integer, Some(&numbers)),
] {
fallback::reset();
let mut states = vec![Accumulator::new(name, &returns).expect("known"); 2];
update_scattered(&mut states, &slots, 1, 0, column, 3).expect("folds them in");
assert_eq!(
fallback::count(Kernel::Aggregate, Form::Flat, Form::Flat),
0,
"{name} over an integer column took the row at a time path"
);
}
fallback::reset();
let mut states = vec![Accumulator::new("min", &LogicalType::Varchar).expect("known"); 2];
update_scattered(&mut states, &slots, 1, 0, Some(&words), 3).expect("folds them in");
assert_eq!(states[0].finish().expect("finishes"), Value::Varchar("a".into()));
assert_eq!(fallback::count(Kernel::Aggregate, Form::Flat, Form::Flat), 0);
}
#[test]
fn a_sum_over_runs_takes_the_same_loop_a_dictionary_takes() {
fallback::reset();
let values = Vector::from_values(
LogicalType::Integer,
&[Value::Integer(5), Value::Null, Value::Integer(7)],
)
.expect("a vector of integers");
let runs = Vector::runs(vec![4, 6, 10], values).expect("one value for each run");
assert_eq!(runs.form(), Form::Rle);
let mut summing =
Accumulator::new("sum", &LogicalType::HugeInt).expect("a known aggregate");
summing.update_run(std::slice::from_ref(&runs), 10).expect("sums");
assert_eq!(summing.finish().expect("finishes"), Value::HugeInt(48));
assert_eq!(fallback::count(Kernel::Aggregate, Form::Rle, Form::Rle), 0);
assert_eq!(fallback::count(Kernel::Aggregate, Form::Flat, Form::Flat), 0);
fallback::reset();
}
#[test]
fn a_sum_of_numbers_stays_off_the_row_at_a_time_path_and_a_sum_of_strings_does_not() {
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)
);
}
#[derive(Debug)]
struct Filed(Vec<Vec<u8>>);
impl rudb_vector::TextSource for Filed {
fn len(&self) -> usize {
self.0.len()
}
fn bytes_at(&self, index: usize) -> Result<Option<&[u8]>> {
Ok(self.0.get(index).map(Vec::as_slice))
}
fn footprint(&self) -> usize {
self.0.iter().map(Vec::len).sum()
}
}
#[test]
fn an_extreme_over_a_dictionary_that_keeps_its_bytes_in_a_file_is_decided_on_the_bytes() {
let source = Arc::new(Filed(vec![b"pear".to_vec(), b"apple".to_vec(), b"plum".to_vec()]));
let values = Vector::external_text(LogicalType::Varchar, source).expect("three values");
let coded = Vector::dictionary(vec![0, 2, 1, 2, 0], values).expect("codes are in range");
let batch = std::slice::from_ref(&coded);
assert_eq!(
a_vector_at_a_time("min", &LogicalType::Varchar, batch).expect("finds one"),
Value::Varchar("apple".into())
);
assert_eq!(
a_vector_at_a_time("max", &LogicalType::Varchar, batch).expect("finds one"),
Value::Varchar("plum".into())
);
let smallest = gather(&coded, 5, &Validity::AllValid, Want::Extreme(true));
assert!(matches!(smallest, Some(Contribution::Extreme(Some(2)))));
let largest = gather(&coded, 5, &Validity::AllValid, Want::Extreme(false));
assert!(matches!(largest, Some(Contribution::Extreme(Some(1)))));
}
#[test]
fn a_run_shorter_than_the_vector_totals_only_the_rows_it_was_asked_for() {
let rows: Vec<Value> = (1..=10).map(Value::Integer).collect();
let vector = Vector::from_values(LogicalType::Integer, &rows).expect("a vector");
let batch = std::slice::from_ref(&vector);
for count in 0..=10usize {
let mut accumulator =
Accumulator::new("sum", &LogicalType::HugeInt).expect("a known one");
accumulator.update_run(batch, count).expect("totals");
let wanted = (count * (count + 1) / 2) as i128;
let got = accumulator.finish().expect("a total");
if count == 0 {
assert_eq!(got, Value::Null, "no rows is no total");
} else {
assert_eq!(got, Value::HugeInt(wanted), "the first {count} rows");
}
}
}
#[test]
fn a_dictionary_read_through_a_padded_copy_totals_what_the_same_rows_total_laid_out_flat() {
for distinct in [1usize, 2, 7, 255, TALLY_LIMIT, TALLY_LIMIT + 1, TALLY_LIMIT * 3] {
let entries: Vec<Value> =
(0..distinct).map(|slot| Value::Integer(slot as i32 * 7 - 11)).collect();
let values = Vector::from_values(LogicalType::Integer, &entries).expect("a dictionary");
let codes: Vec<u32> = (0..1500u32).map(|row| row * 13 % distinct as u32).collect();
let flat: Vec<Value> =
codes.iter().map(|&code| entries[code as usize].clone()).collect();
let coded = Vector::dictionary(codes, values).expect("codes are in range");
let mut counted = Accumulator::new("sum", &LogicalType::HugeInt).expect("a known one");
counted.update_run(std::slice::from_ref(&coded), 1500).expect("totals");
let laid_out = Vector::from_values(LogicalType::Integer, &flat).expect("a vector");
let mut gathered = Accumulator::new("sum", &LogicalType::HugeInt).expect("a known one");
gathered.update_run(std::slice::from_ref(&laid_out), 1500).expect("totals");
assert_eq!(
counted.finish().expect("a total"),
gathered.finish().expect("a total"),
"a dictionary of {distinct} entries"
);
}
}
#[test]
fn a_padded_dictionary_stops_at_the_rows_it_was_asked_for() {
let entries = [Value::Integer(1), Value::Integer(100)];
let values = Vector::from_values(LogicalType::Integer, &entries).expect("a dictionary");
let coded = Vector::dictionary(vec![0, 0, 0, 1, 1], values).expect("codes are in range");
let mut accumulator = Accumulator::new("sum", &LogicalType::HugeInt).expect("a known one");
accumulator.update_run(std::slice::from_ref(&coded), 3).expect("totals");
assert_eq!(
accumulator.finish().expect("a total"),
Value::HugeInt(3),
"only the three ones"
);
}
#[test]
fn a_total_of_hugeints_goes_the_row_at_a_time_way_and_still_overflows() {
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();
}
}