use std::cmp::Ordering;
use minarrow::Numeric;
use vec64::Vec64;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HeapOrder {
Max,
Min,
}
#[inline(always)]
fn cmp_val<T: Numeric + PartialOrd>(order: HeapOrder, a: T, b: T) -> Ordering {
let base = a.partial_cmp(&b).unwrap_or(Ordering::Equal);
match order {
HeapOrder::Max => base,
HeapOrder::Min => base.reverse(),
}
}
#[derive(Debug, Clone)]
pub struct BinaryHeap64<T: Numeric + PartialOrd> {
data: Vec<T>,
order: HeapOrder,
nan_start: usize,
}
impl<T: Numeric + PartialOrd> BinaryHeap64<T> {
#[inline]
pub fn new() -> Self {
Self { data: Vec::new(), order: HeapOrder::Min, nan_start: 0 }
}
#[inline]
pub fn new_min() -> Self {
Self::new()
}
#[inline]
pub fn new_max() -> Self {
Self { data: Vec::new(), order: HeapOrder::Max, nan_start: 0 }
}
#[inline]
pub fn new_min_cap(capacity: usize) -> Self {
Self { data: Vec::with_capacity(capacity), order: HeapOrder::Min, nan_start: 0 }
}
#[inline]
pub fn new_max_cap(capacity: usize) -> Self {
Self { data: Vec::with_capacity(capacity), order: HeapOrder::Max, nan_start: 0 }
}
#[inline]
pub fn len(&self) -> usize {
self.data.len()
}
#[inline]
pub fn real_len(&self) -> usize {
self.nan_start
}
#[inline]
pub fn nan_count(&self) -> usize {
self.data.len() - self.nan_start
}
#[inline]
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
#[inline]
pub fn peek(&self) -> Option<T> {
if self.nan_start == 0 { None } else { Some(self.data[0]) }
}
#[inline]
pub fn push(&mut self, val: T) {
if val.partial_cmp(&val).is_none() {
self.data.push(val);
return;
}
self.data.push(val);
let new_idx = self.data.len() - 1;
if new_idx != self.nan_start {
self.data.swap(self.nan_start, new_idx);
}
self.nan_start += 1;
self.sift_up(self.nan_start - 1);
}
#[inline]
pub fn pop(&mut self) -> Option<T> {
if self.nan_start == 0 {
return None;
}
self.nan_start -= 1;
self.data.swap(0, self.nan_start);
let val = if self.nan_count() == 0 {
self.data.pop().unwrap()
} else {
let last = self.data.len() - 1;
self.data.swap(self.nan_start, last);
self.data.pop().unwrap()
};
if self.nan_start > 0 {
self.sift_down(0, self.nan_start);
}
Some(val)
}
#[inline]
pub fn push_pop(&mut self, val: T) -> T {
if val.partial_cmp(&val).is_none() {
return val;
}
if self.nan_start == 0 {
return val;
}
let order = self.order;
if cmp_val(order, val, self.data[0]) != Ordering::Greater {
return val;
}
let root = self.data[0];
self.data[0] = val;
self.sift_down(0, self.nan_start);
root
}
#[inline]
pub fn into_vec(self) -> Vec<T> {
self.data
}
pub fn into_sorted_vec64(mut self) -> Vec64<T> {
if self.nan_start <= 1 {
return self.data.into();
}
let mut end = self.nan_start;
while end > 1 {
end -= 1;
self.data.swap(0, end);
self.sift_down(0, end);
}
self.data.into()
}
#[inline]
pub fn as_slice(&self) -> &[T] {
&self.data
}
#[inline]
pub fn order(&self) -> HeapOrder {
self.order
}
pub fn from_vec64(data: Vec64<T>, order: HeapOrder) -> Self {
let v: Vec<T> = data.into_iter().collect();
Self::from_vec(v, order)
}
pub fn from_vec(mut data: Vec<T>, order: HeapOrder) -> Self {
let mut nan_start = data.len();
let mut i = 0;
while i < nan_start {
if data[i].partial_cmp(&data[i]).is_none() {
nan_start -= 1;
data.swap(i, nan_start);
} else {
i += 1;
}
}
let mut heap = Self { data, order, nan_start };
if nan_start > 1 {
for i in (0..nan_start / 2).rev() {
heap.sift_down(i, nan_start);
}
}
heap
}
#[inline]
fn sift_up(&mut self, mut idx: usize) {
let order = self.order;
while idx > 0 {
let parent = (idx - 1) / 2;
if cmp_val(order, self.data[idx], self.data[parent]) != Ordering::Greater {
break;
}
self.data.swap(idx, parent);
idx = parent;
}
}
#[inline]
fn sift_down(&mut self, mut idx: usize, end: usize) {
let order = self.order;
loop {
let left = 2 * idx + 1;
if left >= end {
break;
}
let right = left + 1;
let child = if right < end && cmp_val(order, self.data[right], self.data[left]) == Ordering::Greater {
right
} else {
left
};
if cmp_val(order, self.data[child], self.data[idx]) != Ordering::Greater {
break;
}
self.data.swap(idx, child);
idx = child;
}
}
}
impl<T: Numeric + PartialOrd> Default for BinaryHeap64<T> {
fn default() -> Self {
Self::new()
}
}
impl<T: Numeric + PartialOrd> PartialEq for BinaryHeap64<T> {
fn eq(&self, other: &Self) -> bool {
if self.order != other.order || self.data.len() != other.data.len()
|| self.nan_count() != other.nan_count()
{
return false;
}
let mut a: Vec<T> = self.data[..self.nan_start].to_vec();
let mut b: Vec<T> = other.data[..other.nan_start].to_vec();
a.sort_by(|x, y| x.partial_cmp(y).unwrap_or(Ordering::Equal));
b.sort_by(|x, y| x.partial_cmp(y).unwrap_or(Ordering::Equal));
a == b
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn min_heap_f64() {
let mut h = BinaryHeap64::<f64>::new();
h.push(3.0);
h.push(1.0);
h.push(4.0);
h.push(1.0);
h.push(5.0);
assert_eq!(h.len(), 5);
assert_eq!(h.real_len(), 5);
assert_eq!(h.nan_count(), 0);
assert_eq!(h.peek(), Some(1.0));
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(3.0));
assert_eq!(h.pop(), Some(4.0));
assert_eq!(h.pop(), Some(5.0));
assert_eq!(h.pop(), None);
}
#[test]
fn max_heap_f64() {
let mut h = BinaryHeap64::<f64>::new_max();
h.push(3.0);
h.push(1.0);
h.push(4.0);
h.push(1.0);
h.push(5.0);
assert_eq!(h.peek(), Some(5.0));
assert_eq!(h.pop(), Some(5.0));
assert_eq!(h.pop(), Some(4.0));
assert_eq!(h.pop(), Some(3.0));
}
#[test]
fn min_heap_i64() {
let mut h = BinaryHeap64::<i64>::new();
h.push(3);
h.push(1);
h.push(4);
h.push(-2);
h.push(5);
assert_eq!(h.peek(), Some(-2));
assert_eq!(h.pop(), Some(-2));
assert_eq!(h.pop(), Some(1));
assert_eq!(h.pop(), Some(3));
assert_eq!(h.nan_count(), 0);
}
#[test]
fn max_heap_u32() {
let mut h = BinaryHeap64::<u32>::new_max();
h.push(10);
h.push(3);
h.push(7);
assert_eq!(h.peek(), Some(10));
assert_eq!(h.pop(), Some(10));
assert_eq!(h.pop(), Some(7));
}
#[test]
fn topk_via_min_heap() {
let mut h = BinaryHeap64::<f64>::new_min_cap(3);
for v in [10.0, 2.0, 8.0, 5.0, 1.0, 9.0, 3.0, 7.0] {
if h.real_len() < 3 {
h.push(v);
} else if v > h.peek().unwrap() {
h.pop();
h.push(v);
}
}
let sorted = h.into_sorted_vec64();
assert_eq!(&*sorted, &[10.0, 9.0, 8.0]);
}
#[test]
fn bottomk_via_max_heap() {
let mut h = BinaryHeap64::<f64>::new_max_cap(3);
for v in [10.0, 2.0, 8.0, 5.0, 1.0, 9.0, 3.0, 7.0] {
if h.real_len() < 3 {
h.push(v);
} else if v < h.peek().unwrap() {
h.pop();
h.push(v);
}
}
let sorted = h.into_sorted_vec64();
assert_eq!(&*sorted, &[1.0, 2.0, 3.0]);
}
#[test]
fn inf_ordering() {
let mut h = BinaryHeap64::<f64>::new_max();
h.push(f64::NEG_INFINITY);
h.push(1.0);
h.push(f64::INFINITY);
h.push(0.0);
assert_eq!(h.pop(), Some(f64::INFINITY));
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(0.0));
assert_eq!(h.pop(), Some(f64::NEG_INFINITY));
}
#[test]
fn nan_excluded_from_heap() {
let mut h = BinaryHeap64::<f64>::new_max();
h.push(f64::NAN);
h.push(1.0);
h.push(f64::INFINITY);
h.push(0.0);
assert_eq!(h.len(), 4);
assert_eq!(h.real_len(), 3);
assert_eq!(h.nan_count(), 1);
assert_eq!(h.peek(), Some(f64::INFINITY));
assert_eq!(h.pop(), Some(f64::INFINITY));
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(0.0));
assert_eq!(h.pop(), None);
}
#[test]
fn nan_at_tail_of_sorted() {
let mut h = BinaryHeap64::<f64>::new_max();
h.push(3.0);
h.push(f64::NAN);
h.push(1.0);
h.push(f64::NAN);
h.push(2.0);
let sorted = h.into_sorted_vec64();
assert_eq!(sorted[0], 1.0);
assert_eq!(sorted[1], 2.0);
assert_eq!(sorted[2], 3.0);
assert!(sorted[3].is_nan());
assert!(sorted[4].is_nan());
}
#[test]
fn nan_at_tail_min_heap() {
let mut h = BinaryHeap64::<f64>::new();
h.push(3.0);
h.push(f64::NAN);
h.push(1.0);
h.push(2.0);
let sorted = h.into_sorted_vec64();
assert_eq!(sorted[0], 3.0);
assert_eq!(sorted[1], 2.0);
assert_eq!(sorted[2], 1.0);
assert!(sorted[3].is_nan());
}
#[test]
fn nan_never_displaces_via_push_pop() {
let mut h = BinaryHeap64::<f64>::new_min_cap(2);
h.push(1.0);
h.push(2.0);
let returned = h.push_pop(f64::NAN);
assert!(returned.is_nan());
assert_eq!(h.peek(), Some(1.0));
assert_eq!(h.real_len(), 2);
}
#[test]
fn all_nan() {
let mut h = BinaryHeap64::<f64>::new();
h.push(f64::NAN);
h.push(f64::NAN);
h.push(f64::NAN);
assert_eq!(h.len(), 3);
assert_eq!(h.real_len(), 0);
assert_eq!(h.peek(), None);
assert_eq!(h.pop(), None);
}
#[test]
fn negative_zero() {
let mut h = BinaryHeap64::<f64>::new();
h.push(-0.0);
h.push(0.0);
assert_eq!(h.real_len(), 2);
h.pop();
h.pop();
assert_eq!(h.real_len(), 0);
}
#[test]
fn from_vec64_with_nan() {
let data: Vec64<f64> = vec![3.0, f64::NAN, 1.0, 4.0, f64::NAN].into();
let mut h = BinaryHeap64::from_vec64(data, HeapOrder::Min);
assert_eq!(h.real_len(), 3);
assert_eq!(h.nan_count(), 2);
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(3.0));
assert_eq!(h.pop(), Some(4.0));
assert_eq!(h.pop(), None);
}
#[test]
fn from_vec64_integers() {
let data: Vec64<i32> = vec![3, 1, 4, 1, 5, 9, 2, 6].into();
let mut h = BinaryHeap64::from_vec64(data, HeapOrder::Max);
assert_eq!(h.pop(), Some(9));
assert_eq!(h.pop(), Some(6));
assert_eq!(h.pop(), Some(5));
}
#[test]
fn into_sorted_min() {
let mut h = BinaryHeap64::<f64>::new();
for v in [3.0, 1.0, 4.0, 1.0, 5.0, 9.0] {
h.push(v);
}
let sorted = h.into_sorted_vec64();
assert_eq!(&*sorted, &[9.0, 5.0, 4.0, 3.0, 1.0, 1.0]);
}
#[test]
fn into_sorted_max() {
let mut h = BinaryHeap64::<f64>::new_max();
for v in [3.0, 1.0, 4.0, 1.0, 5.0, 9.0] {
h.push(v);
}
let sorted = h.into_sorted_vec64();
assert_eq!(&*sorted, &[1.0, 1.0, 3.0, 4.0, 5.0, 9.0]);
}
#[test]
fn into_sorted_integers() {
let mut h = BinaryHeap64::<i32>::new();
for v in [3, 1, 4, 1, 5] {
h.push(v);
}
let sorted = h.into_sorted_vec64();
assert_eq!(&*sorted, &[5, 4, 3, 1, 1]);
}
#[test]
fn min_heap_f32() {
let mut h = BinaryHeap64::<f32>::new();
h.push(3.0);
h.push(1.0);
h.push(4.0);
h.push(1.0);
h.push(5.0);
assert_eq!(h.peek(), Some(1.0));
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(3.0));
}
#[test]
fn nan_f32() {
let mut h = BinaryHeap64::<f32>::new_max();
h.push(f32::NAN);
h.push(2.0);
h.push(1.0);
assert_eq!(h.real_len(), 2);
assert_eq!(h.nan_count(), 1);
assert_eq!(h.peek(), Some(2.0));
let sorted = h.into_sorted_vec64();
assert_eq!(sorted[0], 1.0);
assert_eq!(sorted[1], 2.0);
assert!(sorted[2].is_nan());
}
#[test]
fn inf_f32() {
let mut h = BinaryHeap64::<f32>::new();
h.push(f32::INFINITY);
h.push(1.0);
h.push(f32::NEG_INFINITY);
assert_eq!(h.pop(), Some(f32::NEG_INFINITY));
assert_eq!(h.pop(), Some(1.0));
assert_eq!(h.pop(), Some(f32::INFINITY));
}
#[test]
fn empty_operations() {
let mut h = BinaryHeap64::<f64>::new();
assert_eq!(h.peek(), None);
assert_eq!(h.pop(), None);
assert_eq!(h.push_pop(5.0), 5.0);
assert!(h.is_empty());
}
#[test]
fn single_element() {
let mut h = BinaryHeap64::<f64>::new();
h.push(42.0);
assert_eq!(h.peek(), Some(42.0));
assert_eq!(h.pop(), Some(42.0));
assert_eq!(h.len(), 0);
}
}