use yo_common::num::{parse_f64, parse_i64};
use yo_common::{Code, Error, Result};
use yo_kv::{Db, Foreign, Keyspace};
use yo_sketch::tdigest::TDigest;
use super::args::{self, Args};
use super::table::Spec;
use crate::reply::Out;
const DEFAULT_COMPRESSION: i64 = 100;
const EXISTS: &str = "T-Digest: key already exists";
const MISSING: &str = "T-Digest: key does not exist";
const BAD_COMPRESSION: &str = "T-Digest: error parsing compression parameter";
const COMPRESSION_RANGE: &str = "T-Digest: compression parameter needs to be a positive integer";
const KEYWORD: &str = "T-Digest: wrong keyword";
const NO_MEMORY: &str = "T-Digest: allocation failed";
const NO_MEMORY_DEST: &str = "T-Digest: allocation of destination digest failed";
const BAD_VAL: &str = "T-Digest: error parsing val parameter";
const NOT_FINITE: &str = "T-Digest: val parameter needs to be a finite number";
const OVERFLOW: &str = "T-Digest: overflow detected";
const BAD_NUMKEYS: &str = "T-Digest: error parsing numkeys";
const NUMKEYS_RANGE: &str = "T-Digest: numkeys needs to be a positive integer";
const BAD_QUANTILE: &str = "T-Digest: error parsing quantile";
const QUANTILE_RANGE: &str = "T-Digest: quantile should be in [0,1]";
const BAD_CDF: &str = "T-Digest: error parsing cdf";
const BAD_VALUE: &str = "T-Digest: error parsing value";
const BAD_RANK: &str = "T-Digest: error parsing rank";
const RANK_NEGATIVE: &str = "T-Digest: rank needs to be non negative";
const BAD_LOW: &str = "T-Digest: error parsing low_cut_percentile";
const BAD_HIGH: &str = "T-Digest: error parsing high_cut_percentile";
const CUT_RANGE: &str = "T-Digest: low_cut_percentile and high_cut_percentile should be in [0,1]";
const CUT_ORDER: &str = "T-Digest: low_cut_percentile should be lower than high_cut_percentile";
const WRONG_KIND: &str = "Operation against a key holding the wrong kind of value";
#[derive(Debug)]
pub(super) struct TDigestBody {
t: TDigest,
}
impl Foreign for TDigestBody {
fn type_name(&self) -> &'static str {
"TDIS-TYPE"
}
fn encoding(&self) -> &'static str {
"raw"
}
fn memory_bytes(&self) -> usize {
self.t.memory_bytes()
}
fn is_empty(&self) -> bool {
false
}
}
pub(super) fn execute(db: &Db, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
if spec.name == "tdigest.merge" {
return merge(db, args, out);
}
let mut held = db.hold(args.get(1));
let db = &mut *held;
match spec.name {
"tdigest.create" => create(db, args, out),
"tdigest.reset" => reset(db, args, out),
"tdigest.add" => add(db, args, out),
"tdigest.min" => ends(db, args, out, false),
"tdigest.max" => ends(db, args, out, true),
"tdigest.quantile" => quantile(db, args, out),
"tdigest.cdf" => cdf(db, args, out),
"tdigest.trimmed_mean" => trimmed_mean(db, args, out),
"tdigest.rank" => rank(db, args, out, false),
"tdigest.revrank" => rank(db, args, out, true),
"tdigest.byrank" => by_rank(db, args, out, false),
"tdigest.byrevrank" => by_rank(db, args, out, true),
"tdigest.info" => info(db, args, out),
other => unreachable!("{other} is not a t digest command"),
}
}
fn create(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() != 2 && args.len() != 4 {
return Err(args::wrong_arity("tdigest.create"));
}
let key = args.get(1);
if read(db, key)?.is_some() {
return Err(Error::new(Code::Invalid, EXISTS));
}
let mut compression = DEFAULT_COMPRESSION;
if args.len() == 4 {
if !args::is(args.get(2), b"compression") && !args::is(args.get(3), b"compression") {
return Err(Error::new(Code::Invalid, KEYWORD));
}
compression = size(args.get(3))?;
}
let Some(t) = TDigest::new(compression) else {
return Err(Error::new(Code::Invalid, NO_MEMORY));
};
db.put_foreign(key, Box::new(TDigestBody { t }));
out.ok();
Ok(())
}
fn reset(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let body = digest(db, args.get(1))?;
body.t.reset();
out.ok();
Ok(())
}
fn add(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let body = digest(db, args.get(1))?;
let mut values = Vec::with_capacity(args.len() - 2);
for i in 2..args.len() {
let Some(v) = double(args.get(i)) else {
return Err(Error::new(Code::Invalid, BAD_VAL));
};
if !v.is_finite() {
return Err(Error::new(Code::Invalid, NOT_FINITE));
}
values.push(v);
}
for v in values {
if body.t.add(v, 1).is_err() {
return Err(Error::new(Code::Invalid, OVERFLOW));
}
}
out.ok();
Ok(())
}
fn merge(db: &Db, args: Args<'_>, out: &mut Out) -> Result<()> {
let dest = args.get(1);
let dest_compression = read(&mut db.hold(dest), dest)?.map(|body| body.t.compression());
let Some(numkeys) = parse_i64(args.get(2)) else {
return Err(Error::new(Code::Invalid, BAD_NUMKEYS));
};
if numkeys <= 0 {
return Err(Error::new(Code::Invalid, NUMKEYS_RANGE));
}
let sources = usize::try_from(numkeys).unwrap_or(usize::MAX);
if sources > args.len() - 3 {
return Err(args::wrong_arity("tdigest.merge"));
}
let rest = sources + 3;
let mut compression = dest_compression;
let mut override_dest = false;
if rest < args.len() {
let at = (rest..args.len()).find(|&i| args::is(args.get(i), b"compression"));
if let Some(at) = at {
if at + 1 >= args.len() {
return Err(args::wrong_arity("tdigest.merge"));
}
compression = Some(size(args.get(at + 1))?);
}
let has_override = (rest..args.len()).any(|i| args::is(args.get(i), b"override"));
if has_override {
override_dest = true;
if at.is_none() {
compression = None;
}
}
if at.is_none() && !has_override {
return Err(Error::new(Code::Invalid, KEYWORD));
}
}
let mut largest = 0;
for i in 3..rest {
let name = args.get(i);
let found = if name == dest {
dest_compression
} else {
read(&mut db.hold(name), name)?.map(|body| body.t.compression())
};
let Some(c) = found else {
return Err(Error::new(Code::Invalid, MISSING));
};
largest = largest.max(c);
}
let Some(mut into) = TDigest::new(compression.unwrap_or(largest)) else {
return Err(Error::new(Code::Invalid, NO_MEMORY_DEST));
};
if !override_dest && dest_compression.is_some() {
let from = compressed(&mut db.hold(dest), dest)?;
fold(&mut into, &from)?;
}
for i in 3..rest {
let key = args.get(i);
let from = compressed(&mut db.hold(key), key)?;
fold(&mut into, &from)?;
}
db.hold(dest)
.put_foreign(dest, Box::new(TDigestBody { t: into }));
out.ok();
Ok(())
}
fn compressed(db: &mut Keyspace, key: &[u8]) -> Result<Vec<(f64, i64)>> {
let body = digest(db, key)?;
if body.t.compress().is_err() {
return Err(Error::new(Code::Invalid, OVERFLOW));
}
Ok(body.t.centroids())
}
fn fold(into: &mut TDigest, from: &[(f64, i64)]) -> Result<()> {
if into.compress().is_err() {
return Err(Error::new(Code::Invalid, OVERFLOW));
}
for &(mean, weight) in from {
if into.add(mean, weight).is_err() {
return Err(Error::new(Code::Invalid, OVERFLOW));
}
}
Ok(())
}
fn ends(db: &mut Keyspace, args: Args<'_>, out: &mut Out, top: bool) -> Result<()> {
let body = digest(db, args.get(1))?;
let value = if body.t.size() > 0 {
if top { body.t.max() } else { body.t.min() }
} else {
f64::NAN
};
out.double(value);
Ok(())
}
fn quantile(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let body = digest(db, args.get(1))?;
let mut wanted = Vec::with_capacity(args.len() - 2);
for i in 2..args.len() {
let Some(q) = double(args.get(i)) else {
return Err(Error::new(Code::Invalid, BAD_QUANTILE));
};
if q < 0.0 || q > 1.0 {
return Err(Error::new(Code::Invalid, QUANTILE_RANGE));
}
wanted.push(q);
}
let mut values = vec![0.0; wanted.len()];
let mut at = 0;
while at < wanted.len() {
let mut end = at;
while end + 1 < wanted.len() && wanted[end] <= wanted[end + 1] {
end += 1;
}
body.t.quantiles(&wanted[at..=end], &mut values[at..=end]);
at = end + 1;
}
out.array(values.len());
for v in values {
out.double(v);
}
Ok(())
}
fn cdf(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let body = digest(db, args.get(1))?;
let mut wanted = Vec::with_capacity(args.len() - 2);
for i in 2..args.len() {
let Some(v) = double(args.get(i)) else {
return Err(Error::new(Code::Invalid, BAD_CDF));
};
wanted.push(v);
}
out.array(wanted.len());
for v in wanted {
let answer = body.t.cdf(v);
out.double(answer);
}
Ok(())
}
fn trimmed_mean(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let body = digest(db, args.get(1))?;
let Some(low) = double(args.get(2)) else {
return Err(Error::new(Code::Invalid, BAD_LOW));
};
let Some(high) = double(args.get(3)) else {
return Err(Error::new(Code::Invalid, BAD_HIGH));
};
if low < 0.0 || low > 1.0 || high < 0.0 || high > 1.0 {
return Err(Error::new(Code::Invalid, CUT_RANGE));
}
if low >= high {
return Err(Error::new(Code::Invalid, CUT_ORDER));
}
let value = body.t.trimmed_mean(low, high);
out.double(value);
Ok(())
}
fn rank(db: &mut Keyspace, args: Args<'_>, out: &mut Out, reverse: bool) -> Result<()> {
let body = digest(db, args.get(1))?;
let mut wanted = Vec::with_capacity(args.len() - 2);
for i in 2..args.len() {
let Some(v) = double(args.get(i)) else {
return Err(Error::new(Code::Invalid, BAD_VALUE));
};
wanted.push(v);
}
#[allow(clippy::cast_precision_loss)]
let size = body.t.size() as f64;
let min = body.t.min();
let max = body.t.max();
out.array(wanted.len());
for v in wanted {
let answer = if size == 0.0 {
-2.0
} else if v < min {
if reverse { size } else { -1.0 }
} else if v > max {
if reverse { -1.0 } else { size }
} else {
let at = body.t.cdf(v) * size;
let at = if reverse { at.round() } else { half_down(at) };
if reverse { (size - at).round() } else { at }
};
#[allow(clippy::cast_possible_truncation)]
out.int(answer as i64);
}
Ok(())
}
fn by_rank(db: &mut Keyspace, args: Args<'_>, out: &mut Out, reverse: bool) -> Result<()> {
let body = digest(db, args.get(1))?;
let mut wanted = Vec::with_capacity(args.len() - 2);
for i in 2..args.len() {
let Some(rank) = parse_i64(args.get(i)) else {
return Err(Error::new(Code::Invalid, BAD_RANK));
};
if rank < 0 {
return Err(Error::new(Code::Invalid, RANK_NEGATIVE));
}
wanted.push(rank);
}
#[allow(clippy::cast_precision_loss)]
let size = body.t.size() as f64;
out.array(wanted.len());
for rank in wanted {
#[allow(clippy::cast_precision_loss)]
let rank = rank as f64;
let answer = if size == 0.0 {
f64::NAN
} else if rank == 0.0 {
if reverse { body.t.max() } else { body.t.min() }
} else if rank >= size {
if reverse {
f64::NEG_INFINITY
} else {
f64::INFINITY
}
} else {
let at = if reverse { size - rank - 1.0 } else { rank };
body.t.quantile(at / size)
};
out.double(answer);
}
Ok(())
}
fn info(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let body = digest(db, args.get(1))?;
let t = &body.t;
out.map(9);
out.simple(b"Compression");
out.int(t.compression());
out.simple(b"Capacity");
out.uint(t.capacity() as u64);
out.simple(b"Merged nodes");
out.uint(t.merged_nodes() as u64);
out.simple(b"Unmerged nodes");
out.uint(t.unmerged_nodes() as u64);
out.simple(b"Merged weight");
out.int(t.merged_weight());
out.simple(b"Unmerged weight");
out.int(t.unmerged_weight());
out.simple(b"Observations");
out.int(t.size());
out.simple(b"Total compressions");
out.int(t.compressions());
out.simple(b"Memory usage");
out.uint(t.reported_bytes());
Ok(())
}
fn half_down(f: f64) -> f64 {
let whole = f.trunc();
let frac = f - whole;
if frac.abs() <= 0.5 {
return whole;
}
if whole >= 0.0 {
whole + 1.0
} else {
whole - 1.0
}
}
fn double(arg: &[u8]) -> Option<f64> {
let v = parse_f64(arg)?;
if v.is_finite() {
if v == 0.0 && !is_zero(arg) {
return None;
}
return Some(v);
}
if is_infinity(arg) { Some(v) } else { None }
}
fn is_infinity(arg: &[u8]) -> bool {
let body = match arg.first() {
Some(b'+' | b'-') => &arg[1..],
_ => arg,
};
body.eq_ignore_ascii_case(b"inf") || body.eq_ignore_ascii_case(b"infinity")
}
fn is_zero(arg: &[u8]) -> bool {
let mut digits = arg.iter().take_while(|&&c| c != b'e' && c != b'E');
digits.all(|&c| !c.is_ascii_digit() || c == b'0')
}
fn size(arg: &[u8]) -> Result<i64> {
let Some(n) = parse_i64(arg) else {
return Err(Error::new(Code::Invalid, BAD_COMPRESSION));
};
if n <= 0 {
return Err(Error::new(Code::Invalid, COMPRESSION_RANGE));
}
Ok(n)
}
fn digest<'k>(db: &'k mut Keyspace, key: &[u8]) -> Result<&'k mut TDigestBody> {
match write(db, key)? {
Some(body) => Ok(body),
None => Err(Error::new(Code::Invalid, MISSING)),
}
}
fn write<'k>(db: &'k mut Keyspace, key: &[u8]) -> Result<Option<&'k mut TDigestBody>> {
match db.foreign_mut(key)? {
Some(body) => match body.downcast_mut::<TDigestBody>() {
Some(body) => Ok(Some(body)),
None => Err(Error::new(Code::WrongType, WRONG_KIND)),
},
None => Ok(None),
}
}
fn read<'k>(db: &'k mut Keyspace, key: &[u8]) -> Result<Option<&'k TDigestBody>> {
match db.foreign(key)? {
Some(body) => match body.downcast_ref::<TDigestBody>() {
Some(body) => Ok(Some(body)),
None => Err(Error::new(Code::WrongType, WRONG_KIND)),
},
None => Ok(None),
}
}