use crate::error::Error;
use crate::req::SearchCriteria;
#[derive(Debug, Clone)]
pub struct SortedView<T> {
items: Vec<T>,
cumulative_weights: Vec<u64>,
total_weight: u64,
}
impl<T> SortedView<T>
where
T: Clone + Ord,
{
pub(super) fn new(mut weighted_items: Vec<(T, u64)>) -> Self {
if weighted_items.is_empty() {
return Self {
items: vec![],
cumulative_weights: vec![],
total_weight: 0,
};
}
weighted_items.sort_unstable_by(|a, b| a.0.cmp(&b.0));
let mut items: Vec<T> = Vec::with_capacity(weighted_items.len());
let mut cumulative_weights = Vec::with_capacity(weighted_items.len());
let mut cumulative_weight = 0u64;
for (item, weight) in weighted_items {
if let Some(last) = items.last() {
if last == &item {
cumulative_weight += weight;
let last_idx = cumulative_weights.len() - 1;
cumulative_weights[last_idx] = cumulative_weight;
continue;
}
}
cumulative_weight += weight;
items.push(item);
cumulative_weights.push(cumulative_weight);
}
Self {
items,
cumulative_weights,
total_weight: cumulative_weight,
}
}
pub fn is_empty(&self) -> bool {
self.items.is_empty()
}
pub fn len(&self) -> usize {
self.items.len()
}
pub fn total_weight(&self) -> u64 {
self.total_weight
}
pub fn rank(&self, item: &T, criteria: SearchCriteria) -> Result<f64, Error> {
if self.is_empty() {
return Err(Error::invalid_argument("sketch is empty"));
}
match criteria {
SearchCriteria::Inclusive => {
let pos = self.items.partition_point(|x| x <= item);
if pos == 0 {
Ok(0.0)
} else {
Ok(self.cumulative_weights[pos - 1] as f64 / self.total_weight as f64)
}
}
SearchCriteria::Exclusive => {
let pos = self.items.partition_point(|x| x < item);
if pos == 0 {
Ok(0.0)
} else {
Ok(self.cumulative_weights[pos - 1] as f64 / self.total_weight as f64)
}
}
}
}
pub fn quantile(&self, rank: f64, criteria: SearchCriteria) -> Result<T, Error> {
if self.is_empty() {
return Err(Error::invalid_argument("sketch is empty"));
}
if !(0.0..=1.0).contains(&rank) {
return Err(Error::invalid_argument(format!(
"rank {rank} must be in [0, 1]"
)));
}
if rank == 0.0 {
match criteria {
SearchCriteria::Inclusive => return Ok(self.items[0].clone()),
SearchCriteria::Exclusive => return Ok(self.items[0].clone()),
}
}
if rank == 1.0 {
return Ok(self.items[self.items.len() - 1].clone());
}
let target_weight = match criteria {
SearchCriteria::Inclusive => (rank * self.total_weight as f64).ceil() as u64,
SearchCriteria::Exclusive => (rank * self.total_weight as f64) as u64,
};
let index = match criteria {
SearchCriteria::Inclusive => {
self.cumulative_weights
.partition_point(|&w| w < target_weight)
}
SearchCriteria::Exclusive => {
self.cumulative_weights
.partition_point(|&w| w <= target_weight)
}
};
if index >= self.items.len() {
return Ok(self.items[self.items.len() - 1].clone());
}
Ok(self.items[index].clone())
}
pub fn pmf(&self, split_points: &[T], criteria: SearchCriteria) -> Result<Vec<f64>, Error> {
if self.is_empty() {
return Err(Error::invalid_argument("sketch is empty"));
}
self.validate_split_points(split_points)?;
let mut result = Vec::with_capacity(split_points.len() + 1);
let mut prev_rank = 0.0;
for split_point in split_points {
let rank = self.rank(split_point, criteria)?;
result.push(rank - prev_rank);
prev_rank = rank;
}
result.push(1.0 - prev_rank);
Ok(result)
}
pub fn cdf(&self, split_points: &[T], criteria: SearchCriteria) -> Result<Vec<f64>, Error> {
if self.is_empty() {
return Err(Error::invalid_argument("sketch is empty"));
}
self.validate_split_points(split_points)?;
let mut result = Vec::with_capacity(split_points.len() + 1);
let mut cumulative = 0.0;
let pmf = self.pmf(split_points, criteria)?;
for mass in pmf {
cumulative += mass;
result.push(cumulative);
}
Ok(result)
}
fn validate_split_points(&self, split_points: &[T]) -> Result<(), Error> {
if split_points.windows(2).any(|pair| pair[0] >= pair[1]) {
return Err(Error::invalid_argument(
"Split points must be unique and monotonically increasing".to_string(),
));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use googletest::assert_that;
use googletest::prelude::all;
use googletest::prelude::anything;
use googletest::prelude::err;
use googletest::prelude::ge;
use googletest::prelude::le;
use googletest::prelude::near;
use super::*;
fn create_test_view() -> SortedView<i32> {
let weighted_items = vec![(1, 1), (3, 1), (5, 1), (7, 1), (9, 1)];
SortedView::new(weighted_items)
}
#[test]
fn test_sorted_view_creation() {
let view = create_test_view();
assert_eq!(view.len(), 5);
assert_eq!(view.total_weight(), 5);
assert!(!view.is_empty());
}
#[test]
fn test_rank_queries() -> Result<(), Error> {
let view = create_test_view();
assert_that!(view.rank(&1, SearchCriteria::Inclusive)?, near(0.2, 1e-10));
assert_that!(view.rank(&1, SearchCriteria::Exclusive)?, near(0.0, 1e-10));
assert_that!(view.rank(&2, SearchCriteria::Inclusive)?, near(0.2, 1e-10));
assert_that!(view.rank(&6, SearchCriteria::Inclusive)?, near(0.6, 1e-10));
assert_that!(view.rank(&0, SearchCriteria::Inclusive)?, near(0.0, 1e-10));
assert_that!(view.rank(&10, SearchCriteria::Inclusive)?, near(1.0, 1e-10));
Ok(())
}
#[test]
fn test_quantile_queries() -> Result<(), Error> {
let view = create_test_view();
assert_eq!(view.quantile(0.0, SearchCriteria::Inclusive)?, 1);
assert_eq!(view.quantile(1.0, SearchCriteria::Inclusive)?, 9);
let median = view.quantile(0.5, SearchCriteria::Inclusive)?;
assert_that!(median, all!(ge(3), le(7)));
let q25 = view.quantile(0.25, SearchCriteria::Inclusive)?;
let q75 = view.quantile(0.75, SearchCriteria::Inclusive)?;
assert_that!(q25, le(median));
assert_that!(median, le(q75));
Ok(())
}
#[test]
fn test_pmf() -> Result<(), Error> {
let view = create_test_view();
let split_points = vec![3, 7];
let pmf = view.pmf(&split_points, SearchCriteria::Inclusive)?;
assert_eq!(pmf.len(), 3);
let sum: f64 = pmf.iter().sum();
assert_that!(sum, near(1.0, 1e-10));
Ok(())
}
#[test]
fn test_cdf() -> Result<(), Error> {
let view = create_test_view();
let split_points = vec![3, 7];
let cdf = view.cdf(&split_points, SearchCriteria::Inclusive)?;
assert_eq!(cdf.len(), 3);
for i in 1..cdf.len() {
assert_that!(cdf[i], ge(cdf[i - 1]));
}
assert_that!(cdf[cdf.len() - 1], near(1.0, 1e-10));
Ok(())
}
#[test]
fn test_empty_view() {
let view: SortedView<i32> = SortedView::new(vec![]);
assert!(view.is_empty());
assert_eq!(view.len(), 0);
assert_eq!(view.total_weight(), 0);
assert_that!(view.rank(&5, SearchCriteria::Inclusive), err(anything()));
assert_that!(
view.quantile(0.5, SearchCriteria::Inclusive),
err(anything())
);
}
}