use rudb_common::{Error, LogicalType, Result, Value};
use crate::quantile::Column;
const COMPRESSION: f64 = 100.0;
const MOST_PACKED: usize = 200;
const MOST_WAITING: usize = 800;
#[derive(Debug, Clone, Copy)]
struct Centroid {
mean: f64,
weight: f64,
}
impl Centroid {
fn add(&mut self, other: Self) {
if self.weight == 0.0 {
*self = other;
} else {
self.weight += other.weight;
self.mean += other.weight * (other.mean - self.mean) / self.weight;
}
}
}
#[derive(Debug, Clone)]
struct TDigest {
min: f64,
max: f64,
packed_weight: f64,
waiting_weight: f64,
packed: Vec<Centroid>,
waiting: Vec<f64>,
}
impl Default for TDigest {
fn default() -> Self {
Self {
min: f64::MAX,
max: f64::MIN,
packed_weight: 0.0,
waiting_weight: 0.0,
packed: Vec::new(),
waiting: Vec::new(),
}
}
}
impl TDigest {
fn add(&mut self, x: f64) {
if x.is_nan() {
return;
}
self.waiting.push(x);
self.waiting_weight += 1.0;
if self.dirty() {
self.process();
}
}
fn dirty(&self) -> bool {
self.packed.len() > MOST_PACKED || self.waiting.len() > MOST_WAITING
}
fn merge(&mut self, other: &Self) {
if !other.packed.is_empty() {
self.packed_weight += other.packed_weight;
let merged = merge_packed(&other.packed, &self.packed);
self.packed = merged;
self.bounds();
}
self.waiting.extend_from_slice(&other.waiting);
self.waiting_weight += other.waiting_weight;
if self.dirty() {
self.process();
}
}
fn bounds(&mut self) {
if let (Some(first), Some(last)) = (self.packed.first(), self.packed.last()) {
if first.mean < self.min {
self.min = first.mean;
}
if self.max < last.mean {
self.max = last.mean;
}
}
}
fn process(&mut self) {
self.waiting.sort_unstable_by(f64::total_cmp);
let merged = merge_sorted(&self.waiting, &self.packed);
let Some((&first, rest)) = merged.split_first() else {
return;
};
self.packed_weight += self.waiting_weight;
self.waiting_weight = 0.0;
self.packed.clear();
self.packed.push(first);
let total = self.packed_weight;
let mut so_far = first.weight;
let mut limit = total * integrated_q(1.0);
for ¢roid in rest {
let projected = so_far + centroid.weight;
if projected <= limit {
so_far = projected;
if let Some(last) = self.packed.last_mut() {
last.add(centroid);
}
} else {
let k1 = integrated_location(so_far / total);
limit = total * integrated_q(k1 + 1.0);
so_far += centroid.weight;
self.packed.push(centroid);
}
}
self.waiting.clear();
self.bounds();
}
fn cumulative(&self) -> Vec<f64> {
let mut cumulative = Vec::with_capacity(self.packed.len() + 1);
let mut previous = 0.0;
for centroid in &self.packed {
cumulative.push(previous + centroid.weight / 2.0);
previous += centroid.weight;
}
cumulative.push(previous);
cumulative
}
fn quantile(&self, q: f64, cumulative: &[f64]) -> f64 {
if !(0.0..=1.0).contains(&q) || self.packed.is_empty() {
return f64::NAN;
}
let packed = &self.packed;
if packed.len() == 1 {
return packed[0].mean;
}
let n = packed.len();
let index = q * self.packed_weight;
if index <= packed[0].weight / 2.0 {
return self.min + 2.0 * index / packed[0].weight * (packed[0].mean - self.min);
}
let at = cumulative.partition_point(|&c| c < index);
if at + 1 != cumulative.len() {
let z1 = index - cumulative[at - 1];
let z2 = cumulative[at] - index;
return weighted_average(packed[at - 1].mean, z2, packed[at].mean, z1);
}
let z1 = index - self.packed_weight - packed[n - 1].weight / 2.0;
let z2 = packed[n - 1].weight / 2.0 - z1;
weighted_average(packed[n - 1].mean, z1, self.max, z2)
}
}
fn merge_sorted(waiting: &[f64], packed: &[Centroid]) -> Vec<Centroid> {
let mut merged = Vec::with_capacity(waiting.len() + packed.len());
let (mut a, mut b) = (0, 0);
while a < waiting.len() && b < packed.len() {
if packed[b].mean < waiting[a] {
merged.push(packed[b]);
b += 1;
} else {
merged.push(Centroid { mean: waiting[a], weight: 1.0 });
a += 1;
}
}
merged.extend(waiting[a..].iter().map(|&mean| Centroid { mean, weight: 1.0 }));
merged.extend_from_slice(&packed[b..]);
merged
}
fn merge_packed(theirs: &[Centroid], mine: &[Centroid]) -> Vec<Centroid> {
let runs = [theirs, mine];
let mut at = [0, 0];
let mut top = usize::from(!mine.is_empty() && theirs[0].mean > mine[0].mean);
let mut merged = Vec::with_capacity(theirs.len() + mine.len());
loop {
merged.push(runs[top][at[top]]);
at[top] += 1;
let other = 1 - top;
let other_left = at[other] < runs[other].len();
if at[top] == runs[top].len() {
if !other_left {
return merged;
}
top = other;
} else if other_left && runs[other][at[other]].mean <= runs[top][at[top]].mean {
top = other;
}
}
}
fn integrated_location(q: f64) -> f64 {
COMPRESSION * ((2.0 * q - 1.0).asin() + std::f64::consts::PI / 2.0) / std::f64::consts::PI
}
fn integrated_q(k: f64) -> f64 {
let pi = std::f64::consts::PI;
((k.min(COMPRESSION) * pi / COMPRESSION - pi / 2.0).sin() + 1.0) / 2.0
}
fn weighted_average(x1: f64, w1: f64, x2: f64, w2: f64) -> f64 {
let (x1, w1, x2, w2) = if x1 <= x2 { (x1, w1, x2, w2) } else { (x2, w2, x1, w1) };
let x = (x1 * w1 + x2 * w2) / (w1 + w2);
x1.max(x.min(x2))
}
#[derive(Debug, Clone, Default)]
pub(crate) struct Digest {
digest: Option<TDigest>,
count: u64,
}
impl Digest {
pub(crate) fn push(&mut self, value: &Value) {
if let Some(x) = encode(value) {
self.push_number(x);
}
}
pub(crate) fn push_column(&mut self, column: Column<'_>, row: usize) {
match column {
Column::Reals(reals) => self.push_number(reals[row]),
#[expect(
clippy::cast_precision_loss,
reason = "the digest holds a double, as the pin's does"
)]
Column::Wholes(_, numbers) => self.push_number(numbers.at(row) as f64),
Column::Flags(_) => self.push(&column.value(row)),
}
}
fn push_number(&mut self, x: f64) {
if x.is_finite() {
self.digest.get_or_insert_default().add(x);
self.count += 1;
}
}
pub(crate) fn combine(&mut self, other: &Self) {
if let Some(theirs) = &other.digest
&& other.count > 0
{
self.digest.get_or_insert_default().merge(theirs);
self.count += other.count;
}
}
pub(crate) fn finish(&self, fraction: Option<&Value>, returns: &LogicalType) -> Result<Value> {
let Some(held) = self.digest.as_ref().filter(|_| self.count > 0) else {
return Ok(Value::Null);
};
let mut digest = held.clone();
digest.process();
let cumulative = digest.cumulative();
let share = |q: &Value| -> Result<f64> {
match *q {
Value::Float(q) => Ok(f64::from(q)),
#[expect(
clippy::cast_possible_truncation,
reason = "the pin holds the fraction as a FLOAT"
)]
Value::Double(q) => Ok(f64::from(q as f32)),
_ => Err(Error::internal("an approx_quantile fraction that is not a FLOAT")),
}
};
match (fraction, returns) {
(Some(Value::List { values, .. }), LogicalType::List(element)) => {
let values = values
.iter()
.map(|q| decode(digest.quantile(share(q)?, &cumulative), element))
.collect::<Result<Vec<_>>>()?;
Ok(Value::List { element: (**element).clone(), values })
}
(Some(q), returns) => decode(digest.quantile(share(q)?, &cumulative), returns),
(None, _) => Err(Error::internal("an approx_quantile without its fraction")),
}
}
}
#[expect(clippy::cast_precision_loss, reason = "the digest holds a double, as the pin's does")]
fn encode(value: &Value) -> Option<f64> {
Some(match *value {
Value::TinyInt(n) => f64::from(n),
Value::SmallInt(n) => f64::from(n),
Value::Integer(n) | Value::Date(n) => f64::from(n),
Value::BigInt(n) | Value::Time(n) | Value::Timestamp(n) | Value::TimestampTz(n) => n as f64,
Value::HugeInt(n) => n as f64,
Value::Float(x) => f64::from(x),
Value::Double(x) => x,
Value::Decimal { unscaled, .. } => unscaled as f64,
_ => return None,
})
}
#[expect(
clippy::cast_possible_truncation,
reason = "every whole number is clamped to its type before it narrows"
)]
fn decode(x: f64, ty: &LogicalType) -> Result<Value> {
let whole = |min: i128, max: i128| (x.round_ties_even() as i128).clamp(min, max);
let small = |min: i64, max: i64| whole(i128::from(min), i128::from(max)) as i64;
Ok(match *ty {
LogicalType::TinyInt => Value::TinyInt(small(i8::MIN.into(), i8::MAX.into()) as i8),
LogicalType::SmallInt => Value::SmallInt(small(i16::MIN.into(), i16::MAX.into()) as i16),
LogicalType::Integer => Value::Integer(small(i32::MIN.into(), i32::MAX.into()) as i32),
LogicalType::Date => Value::Date(small(i32::MIN.into(), i32::MAX.into()) as i32),
LogicalType::BigInt => Value::BigInt(small(i64::MIN, i64::MAX)),
LogicalType::Time => Value::Time(small(i64::MIN, i64::MAX)),
LogicalType::Timestamp => Value::Timestamp(small(i64::MIN, i64::MAX)),
LogicalType::TimestampTz => Value::TimestampTz(small(i64::MIN, i64::MAX)),
LogicalType::HugeInt => Value::HugeInt(whole(i128::MIN, i128::MAX)),
LogicalType::Float => {
Value::Float(x.clamp(f64::from(f32::MIN), f64::from(f32::MAX)) as f32)
}
LogicalType::Double => Value::Double(x),
LogicalType::Decimal { width, scale } => {
let unscaled = match width {
..=4 => whole(i16::MIN.into(), i16::MAX.into()),
5..=9 => whole(i32::MIN.into(), i32::MAX.into()),
10..=18 => whole(i64::MIN.into(), i64::MAX.into()),
_ => whole(i128::MIN, i128::MAX),
};
Value::Decimal { unscaled, width, scale }
}
_ => return Err(Error::internal(format!("an approx_quantile answer of type {ty}"))),
})
}
#[cfg(test)]
mod tests {
use super::*;
fn over(values: impl IntoIterator<Item = f64>) -> Digest {
let mut digest = Digest::default();
for x in values {
digest.push(&Value::Double(x));
}
digest
}
fn at(digest: &Digest, q: f32) -> Value {
digest.finish(Some(&Value::Float(q)), &LogicalType::Double).unwrap()
}
#[test]
fn a_digest_answers_what_the_pin_answers() {
let digest = over((0..1000).map(f64::from));
assert_eq!(at(&digest, 0.5), Value::Double(499.5));
let thirds = over((0..1000).map(|x| f64::from(x) / 3.0));
assert_eq!(at(&thirds, 0.1), Value::Double(33.16666716337204));
assert_eq!(at(&thirds, 0.33), Value::Double(109.83333770434064));
assert_eq!(at(&thirds, 0.9), Value::Double(299.8333253860475));
let many = over((0..100_000).map(f64::from));
assert_eq!(at(&many, 0.123), Value::Double(12299.500339746475));
assert_eq!(at(&over([f64::INFINITY, 1.0, 2.0]), 0.5), Value::Double(1.5));
assert_eq!(at(&Digest::default(), 0.5), Value::Null);
}
#[test]
fn the_top_of_values_that_are_all_negative_is_the_largest_of_them() {
assert_eq!(at(&over([-5.0, -3.0]), 1.0), Value::Double(-3.0));
}
#[test]
fn an_answer_rounds_to_even_and_stays_in_its_type() {
assert_eq!(decode(4.5, &LogicalType::Date).unwrap(), Value::Date(4));
assert_eq!(decode(499.5, &LogicalType::BigInt).unwrap(), Value::BigInt(500));
assert_eq!(decode(300.0, &LogicalType::TinyInt).unwrap(), Value::TinyInt(127));
assert_eq!(decode(9.3e18, &LogicalType::BigInt).unwrap(), Value::BigInt(i64::MAX));
let decimal = LogicalType::Decimal { width: 4, scale: 1 };
assert_eq!(
decode(49.5, &decimal).unwrap(),
Value::Decimal { unscaled: 50, width: 4, scale: 1 }
);
}
#[test]
fn merged_digests_hold_every_value() {
let mut left = over((0..5000).map(f64::from));
let right = over((5000..10_000).map(f64::from));
left.combine(&right);
left.combine(&Digest::default());
assert_eq!(left.count, 10_000);
let Value::Double(middle) = at(&left, 0.5) else { panic!("a DOUBLE answer") };
assert!((4900.0..5100.0).contains(&middle), "{middle}");
let ties = merge_packed(
&[Centroid { mean: 1.0, weight: 1.0 }, Centroid { mean: 2.0, weight: 1.0 }],
&[Centroid { mean: 1.0, weight: 2.0 }],
);
let weights: Vec<f64> = ties.iter().map(|c| c.weight).collect();
assert_eq!(weights, [1.0, 2.0, 1.0]);
}
}