pub const K_NO_KEY: usize = usize::MAX;
const FREE_FLAG: Key = 1 << 31;
const NO_FREE: Key = FREE_FLAG - 1;
type Key = u32;
#[derive(Debug, Clone)]
struct Node<T> {
value: T,
key: Key,
}
#[derive(Clone)]
pub struct IndexedHeap<T, C> {
elements: Vec<Node<T>>,
pos: Vec<Key>,
free_head: Key,
comp: C,
}
impl<T, C> IndexedHeap<T, C>
where
C: Fn(&T, &T) -> bool,
{
pub fn new(comp: C) -> Self {
Self {
elements: Vec::new(),
pos: Vec::new(),
free_head: NO_FREE,
comp,
}
}
pub fn reserve(&mut self, capacity: usize) {
self.elements.reserve(capacity);
self.pos.reserve(capacity);
}
#[inline]
pub fn insert(&mut self, value: T) -> usize {
let index = self.elements.len();
debug_assert!(index < NO_FREE as usize, "heap larger than the key space");
let index = index as Key;
let key = if self.free_head != NO_FREE {
let key = self.free_head;
self.free_head = self.pos[key as usize] & !FREE_FLAG;
self.pos[key as usize] = index;
key
} else {
self.pos.push(index);
(self.pos.len() - 1) as Key
};
self.elements.push(Node { value, key });
self.sift_up(index as usize);
key as usize
}
#[inline]
pub fn update(&mut self, key: usize, value: T) {
let index = self.pos[key] as usize;
assert!(
index < self.elements.len(),
"update on a key not in the heap"
);
unsafe { self.elements.get_unchecked_mut(index).value = value };
if index > 0 {
let parent = Self::parent(index);
let outranks_parent = unsafe {
(self.comp)(
&self.elements.get_unchecked(index).value,
&self.elements.get_unchecked(parent).value,
)
};
if outranks_parent {
self.sift_up(index);
return;
}
}
self.sift_down(index);
}
#[inline]
pub fn pop(&mut self) -> Option<T> {
let last = self.elements.pop()?;
if self.elements.is_empty() {
self.release(last.key);
return Some(last.value);
}
let moved_key = last.key;
let root = std::mem::replace(unsafe { self.elements.get_unchecked_mut(0) }, last);
unsafe { *self.pos.get_unchecked_mut(moved_key as usize) = 0 };
self.release(root.key);
self.sift_down(0);
Some(root.value)
}
#[inline]
pub fn top(&self) -> Option<&T> {
self.elements.first().map(|node| &node.value)
}
#[inline]
pub fn get(&self, key: usize) -> Option<&T> {
let index = self.index_of(key)?;
Some(&self.elements[index].value)
}
#[inline]
pub fn is_empty(&self) -> bool {
self.elements.is_empty()
}
#[inline]
pub fn len(&self) -> usize {
self.elements.len()
}
pub fn clear(&mut self) {
while let Some(node) = self.elements.pop() {
self.release(node.key);
}
}
#[inline]
pub fn comparator(&self) -> &C {
&self.comp
}
#[inline(always)]
fn index_of(&self, key: usize) -> Option<usize> {
match self.pos.get(key) {
Some(&index) if index & FREE_FLAG == 0 => Some(index as usize),
_ => None,
}
}
#[inline(always)]
fn parent(i: usize) -> usize {
debug_assert!(i > 0, "the root has no parent");
(i - 1) / 2
}
#[inline(always)]
fn release(&mut self, key: Key) {
unsafe { *self.pos.get_unchecked_mut(key as usize) = FREE_FLAG | self.free_head };
self.free_head = key;
}
fn sift_up(&mut self, index: usize) {
if index == 0 {
return;
}
let Self {
elements,
pos,
comp,
..
} = self;
let mut hole = unsafe { Hole::new(elements, pos.as_mut_slice(), index) };
while hole.index > 0 {
let parent = (hole.index - 1) / 2;
let outranks = unsafe { comp(hole.value(), hole.element(parent)) };
if !outranks {
break;
}
unsafe { hole.move_to(parent) };
}
}
fn sift_down(&mut self, index: usize) {
let len = self.elements.len();
if index >= len {
return;
}
let Self {
elements,
pos,
comp,
..
} = self;
let mut hole = unsafe { Hole::new(elements, pos.as_mut_slice(), index) };
loop {
let left = 2 * hole.index + 1;
if left >= len {
break;
}
let right = left + 1;
let best = unsafe {
if right < len && comp(hole.element(right), hole.element(left)) {
right
} else {
left
}
};
if !unsafe { comp(hole.element(best), hole.value()) } {
break;
}
unsafe { hole.move_to(best) };
}
}
#[cfg(test)]
fn is_consistent(&self) -> bool {
for i in 1..self.elements.len() {
if (self.comp)(
&self.elements[i].value,
&self.elements[Self::parent(i)].value,
) {
return false;
}
}
for (index, node) in self.elements.iter().enumerate() {
if self.pos[node.key as usize] != index as Key {
return false;
}
}
true
}
}
struct Hole<'a, T> {
elements: &'a mut [Node<T>],
pos: &'a mut [Key],
node: std::mem::ManuallyDrop<Node<T>>,
index: usize,
}
impl<'a, T> Hole<'a, T> {
#[inline(always)]
unsafe fn new(elements: &'a mut Vec<Node<T>>, pos: &'a mut [Key], index: usize) -> Self {
debug_assert!(index < elements.len());
let node = unsafe { std::ptr::read(elements.as_ptr().add(index)) };
Self {
elements: elements.as_mut_slice(),
pos,
node: std::mem::ManuallyDrop::new(node),
index,
}
}
#[inline(always)]
fn value(&self) -> &T {
&self.node.value
}
#[inline(always)]
unsafe fn element(&self, index: usize) -> &T {
debug_assert!(index < self.elements.len() && index != self.index);
unsafe { &self.elements.get_unchecked(index).value }
}
#[inline(always)]
unsafe fn move_to(&mut self, target: usize) {
debug_assert!(target < self.elements.len() && target != self.index);
unsafe {
let ptr = self.elements.as_mut_ptr();
std::ptr::copy_nonoverlapping(ptr.add(target), ptr.add(self.index), 1);
let key = (*ptr.add(self.index)).key as usize;
*self.pos.get_unchecked_mut(key) = self.index as Key;
}
self.index = target;
}
}
impl<T> Drop for Hole<'_, T> {
#[inline(always)]
fn drop(&mut self) {
unsafe {
let node = std::mem::ManuallyDrop::take(&mut self.node);
let key = node.key as usize;
std::ptr::write(self.elements.as_mut_ptr().add(self.index), node);
*self.pos.get_unchecked_mut(key) = self.index as Key;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::Cell;
use std::rc::Rc;
fn min_heap() -> IndexedHeap<i32, fn(&i32, &i32) -> bool> {
IndexedHeap::new(|a: &i32, b: &i32| a < b)
}
#[test]
fn pops_in_priority_order() {
let mut heap = min_heap();
for value in [10, 5, 20, 1, 7] {
heap.insert(value);
}
assert_eq!(heap.len(), 5);
assert_eq!(heap.top(), Some(&1));
let mut popped = Vec::new();
while let Some(value) = heap.pop() {
popped.push(value);
}
assert_eq!(popped, vec![1, 5, 7, 10, 20]);
assert!(heap.is_empty());
assert_eq!(heap.pop(), None);
}
#[test]
fn update_moves_an_element_both_ways() {
let mut heap = min_heap();
let a = heap.insert(10);
let b = heap.insert(5);
let c = heap.insert(20);
heap.update(c, 2);
assert_eq!(heap.top(), Some(&2));
assert!(heap.is_consistent());
heap.update(c, 50);
assert_eq!(heap.top(), Some(&5));
assert!(heap.is_consistent());
assert_eq!(heap.get(a), Some(&10));
assert_eq!(heap.get(b), Some(&5));
assert_eq!(heap.get(c), Some(&50));
}
#[test]
fn update_of_the_root_does_not_underflow() {
let mut heap = min_heap();
let root = heap.insert(5);
heap.update(root, 1);
assert_eq!(heap.top(), Some(&1));
heap.update(root, 100);
assert_eq!(heap.top(), Some(&100));
}
#[test]
fn keys_are_recycled_rather_than_growing_without_bound() {
let mut heap = min_heap();
for round in 0..8 {
let keys: Vec<_> = (0..4).map(|i| heap.insert(round * 4 + i)).collect();
assert!(keys.iter().all(|&key| key < 4));
for _ in 0..4 {
heap.pop();
}
}
assert_eq!(heap.pos.len(), 4);
}
#[test]
fn clear_releases_keys_and_keeps_capacity() {
let mut heap = min_heap();
let key = heap.insert(10);
heap.insert(20);
heap.clear();
assert!(heap.is_empty());
assert_eq!(heap.get(key), None);
assert_eq!(heap.pos.len(), 2, "clear must not allocate new keys");
heap.insert(30);
assert_eq!(heap.top(), Some(&30));
assert_eq!(heap.pos.len(), 2);
}
#[test]
fn get_returns_none_for_a_popped_key() {
let mut heap = min_heap();
let key = heap.insert(1);
assert_eq!(heap.get(key), Some(&1));
heap.pop();
assert_eq!(heap.get(key), None);
assert_eq!(heap.get(K_NO_KEY), None);
assert_eq!(heap.get(12345), None);
}
#[test]
#[should_panic(expected = "update on a key not in the heap")]
fn update_of_a_popped_key_panics() {
let mut heap = min_heap();
let key = heap.insert(1);
heap.pop();
heap.update(key, 2);
}
#[test]
fn works_as_a_max_heap() {
let mut heap = IndexedHeap::new(|a: &i32, b: &i32| a > b);
for value in [3, 9, 4, 1] {
heap.insert(value);
}
assert_eq!(heap.pop(), Some(9));
assert_eq!(heap.pop(), Some(4));
assert_eq!(heap.pop(), Some(3));
assert_eq!(heap.pop(), Some(1));
}
#[test]
fn every_value_is_dropped_exactly_once() {
struct Counted(#[allow(dead_code)] usize, Rc<Cell<usize>>);
impl Drop for Counted {
fn drop(&mut self) {
self.1.set(self.1.get() + 1);
}
}
let drops = Rc::new(Cell::new(0));
{
let mut heap = IndexedHeap::new(|a: &Counted, b: &Counted| a.0 < b.0);
for i in [4, 1, 3, 2] {
heap.insert(Counted(i, Rc::clone(&drops)));
}
drop(heap.pop());
drop(heap.pop());
assert_eq!(drops.get(), 2);
heap.insert(Counted(9, Rc::clone(&drops)));
assert_eq!(drops.get(), 2);
let key = heap.insert(Counted(8, Rc::clone(&drops)));
heap.update(key, Counted(7, Rc::clone(&drops)));
assert_eq!(drops.get(), 3);
heap.clear();
assert_eq!(drops.get(), 7);
heap.insert(Counted(0, Rc::clone(&drops)));
}
assert_eq!(drops.get(), 8, "the heap must drop what it still holds");
}
#[test]
fn matches_a_reference_implementation_under_random_operations() {
let mut state = 0x2545_F491_4F6C_DD1Du64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
let mut heap = min_heap();
let mut reference: Vec<(usize, i32)> = Vec::new();
for step in 0..4000 {
let roll = next() % 100;
if roll < 45 || reference.is_empty() {
let value = (next() % 1000) as i32;
let key = heap.insert(value);
reference.push((key, value));
} else if roll < 70 {
let victim = (next() as usize) % reference.len();
let value = (next() % 1000) as i32;
let key = reference[victim].0;
heap.update(key, value);
reference[victim].1 = value;
} else {
let expected = reference.iter().map(|&(_, v)| v).min().unwrap();
assert_eq!(heap.top(), Some(&expected), "step {step}");
let popped = heap.pop().unwrap();
assert_eq!(popped, expected, "step {step}");
let victim = reference
.iter()
.position(|&(_, v)| v == popped)
.expect("popped value must be in the reference");
reference.swap_remove(victim);
}
assert!(heap.is_consistent(), "step {step}");
assert_eq!(heap.len(), reference.len(), "step {step}");
}
}
}