use std::ops::AddAssign;
use std::ops::DivAssign;
use super::KeyedVec;
use super::StorageKey;
use crate::containers::HashSet;
use crate::pumpkin_assert_moderate;
#[derive(Debug, Clone)]
pub struct KeyValueHeap<Key, Value> {
values: Vec<Value>,
map_key_to_position: KeyedVec<Key, usize>,
map_position_to_key: Vec<Key>,
end_position: usize,
}
impl<Key: StorageKey, Value> Default for KeyValueHeap<Key, Value> {
fn default() -> Self {
Self {
values: Default::default(),
map_key_to_position: Default::default(),
map_position_to_key: Default::default(),
end_position: Default::default(),
}
}
}
impl<Key, Value> KeyValueHeap<Key, Value> {
pub(crate) const fn new() -> Self {
Self {
values: Vec::new(),
map_key_to_position: KeyedVec::new(),
map_position_to_key: Vec::new(),
end_position: 0,
}
}
}
impl<Key, Value> KeyValueHeap<Key, Value>
where
Key: StorageKey + Copy,
Value: AddAssign<Value> + DivAssign<Value> + PartialOrd + Default + Copy,
{
pub(crate) fn keys(&self) -> impl Iterator<Item = Key> + '_ {
self.map_position_to_key[..self.end_position]
.iter()
.copied()
}
pub(crate) fn peek_max(&self) -> Option<(&Key, &Value)> {
if self.has_no_nonremoved_elements() {
None
} else {
Some((
&self.map_position_to_key[0],
&self.values[self.map_key_to_position[&self.map_position_to_key[0]]],
))
}
}
pub fn get_value(&self, key: Key) -> &Value {
pumpkin_assert_moderate!(
key.index() < self.map_key_to_position.len(),
"Attempted to get key with index {} for a map with length {}",
key.index(),
self.map_key_to_position.len()
);
&self.values[self.map_key_to_position[key]]
}
pub fn pop_max(&mut self) -> Option<Key> {
if !self.has_no_nonremoved_elements() {
let best_key = self.map_position_to_key[0];
pumpkin_assert_moderate!(0 == self.map_key_to_position[best_key]);
self.delete_key(best_key);
Some(best_key)
} else {
None
}
}
pub fn increment(&mut self, key: Key, increment: Value) {
let position = self.map_key_to_position[key];
self.values[position] += increment;
if self.is_key_present(key) {
self.sift_up(position);
}
}
pub fn restore_key(&mut self, key: Key) {
if !self.is_key_present(key) {
let position = self.map_key_to_position[key];
pumpkin_assert_moderate!(position >= self.end_position);
self.swap_positions(position, self.end_position);
self.end_position += 1;
self.sift_up(self.end_position - 1);
}
}
pub fn delete_key(&mut self, key: Key) {
if self.is_key_present(key) {
let position = self.map_key_to_position[key];
self.swap_positions(position, self.end_position - 1);
self.end_position -= 1;
if position < self.end_position {
self.sift_down(position);
}
}
}
pub fn len(&self) -> usize {
self.values.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn num_nonremoved_elements(&self) -> usize {
self.end_position
}
pub(crate) fn has_no_nonremoved_elements(&self) -> bool {
self.num_nonremoved_elements() == 0
}
pub fn is_key_present(&self, key: Key) -> bool {
key.index() < self.map_key_to_position.len()
&& self.map_key_to_position[key] < self.end_position
}
pub fn grow(&mut self, key: Key, value: Value) {
let last_index = self.values.len();
self.values.push(value);
let _ = self.map_key_to_position.push(last_index);
self.map_position_to_key.push(key);
pumpkin_assert_moderate!(
self.map_position_to_key[last_index].index() == key.index()
&& self.map_key_to_position[key] == last_index
);
self.swap_positions(self.end_position, last_index);
self.end_position += 1;
self.sift_up(self.end_position - 1);
}
pub fn clear(&mut self) {
self.values.clear();
self.map_key_to_position.clear();
self.map_position_to_key.clear();
self.end_position = 0;
}
pub fn divide_values(&mut self, divisor: Value) {
for value in self.values.iter_mut() {
*value /= divisor;
}
}
fn swap_positions(&mut self, a: usize, b: usize) {
let key_i = self.map_position_to_key[a];
pumpkin_assert_moderate!(self.map_key_to_position[key_i] == a);
let key_j = self.map_position_to_key[b];
pumpkin_assert_moderate!(self.map_key_to_position[key_j] == b);
self.values.swap(a, b);
self.map_position_to_key.swap(a, b);
self.map_key_to_position.swap(key_i.index(), key_j.index());
pumpkin_assert_moderate!(
self.map_key_to_position[key_i] == b && self.map_key_to_position[key_j] == a
);
pumpkin_assert_moderate!(
self.map_key_to_position
.iter()
.collect::<HashSet<&usize>>()
.len()
== self.map_key_to_position.len()
)
}
fn sift_up(&mut self, position: usize) {
if position > 0 {
let parent_position = KeyValueHeap::<Key, Value>::get_parent_position(position);
if self.values[parent_position] < self.values[position] {
self.swap_positions(parent_position, position);
self.sift_up(parent_position);
}
}
}
fn sift_down(&mut self, position: usize) {
pumpkin_assert_moderate!(position < self.end_position);
if !self.is_heap_locally(position) {
let largest_child_position = self.get_largest_child_position(position);
self.swap_positions(largest_child_position, position);
self.sift_down(largest_child_position);
}
}
fn is_heap_locally(&self, position: usize) -> bool {
let left_child_position = KeyValueHeap::<Key, Value>::get_left_child_position(position);
let right_child_position = KeyValueHeap::<Key, Value>::get_right_child_position(position);
if self.is_leaf(position) {
return true;
}
if right_child_position >= self.end_position {
return self.values[position] >= self.values[left_child_position];
}
self.values[position] >= self.values[left_child_position]
&& self.values[position] >= self.values[right_child_position]
}
fn is_leaf(&self, position: usize) -> bool {
KeyValueHeap::<Key, Value>::get_left_child_position(position) >= self.end_position
}
fn get_largest_child_position(&self, position: usize) -> usize {
pumpkin_assert_moderate!(!self.is_leaf(position));
let left_child_position = KeyValueHeap::<Key, Value>::get_left_child_position(position);
let right_child_position = KeyValueHeap::<Key, Value>::get_right_child_position(position);
if right_child_position < self.end_position
&& self.values[right_child_position] > self.values[left_child_position]
{
right_child_position
} else {
left_child_position
}
}
fn get_parent_position(child_position: usize) -> usize {
pumpkin_assert_moderate!(child_position > 0, "Root has no parent.");
(child_position - 1) / 2
}
fn get_left_child_position(position: usize) -> usize {
2 * position + 1
}
fn get_right_child_position(position: usize) -> usize {
2 * position + 2
}
}
#[cfg(test)]
mod test {
use super::KeyValueHeap;
#[test]
fn failing_test_case() {
let mut heap: KeyValueHeap<usize, u32> = KeyValueHeap::default();
heap.grow(0, 7);
heap.grow(1, 5);
assert_eq!(heap.pop_max().unwrap(), 0);
heap.grow(2, 7);
heap.grow(3, 6);
assert_eq!(heap.pop_max().unwrap(), 2);
assert_eq!(heap.pop_max().unwrap(), 3);
}
#[test]
fn failing_test_case2() {
let mut heap: KeyValueHeap<usize, u32> = KeyValueHeap::default();
heap.grow(0, 5);
heap.grow(1, 7);
heap.grow(2, 6);
assert_eq!(heap.pop_max().unwrap(), 1);
assert_eq!(heap.pop_max().unwrap(), 2);
}
fn heap_sort_test_helper(numbers: Vec<usize>) {
let mut sorted_numbers = numbers.clone();
sorted_numbers.sort();
sorted_numbers.reverse();
let mut heap: KeyValueHeap<usize, usize> = KeyValueHeap::default();
for n in numbers.iter().enumerate() {
heap.grow(n.0, *n.1);
}
let mut heap_sorted_vector: Vec<usize> = vec![];
while let Some(index) = heap.pop_max() {
heap_sorted_vector.push(numbers[index]);
}
assert_eq!(heap_sorted_vector, sorted_numbers);
}
#[test]
fn trivial() {
let mut heap: KeyValueHeap<usize, usize> = KeyValueHeap::default();
heap.grow(0, 5);
assert_eq!(heap.pop_max(), Some(0));
assert!(heap.has_no_nonremoved_elements());
assert_eq!(heap.pop_max(), None);
}
#[test]
fn trivial_sort() {
heap_sort_test_helper(vec![5]);
}
#[test]
fn simple() {
heap_sort_test_helper(vec![5, 10]);
}
#[test]
fn random1() {
heap_sort_test_helper(vec![5, 10, 3]);
}
#[test]
fn random2() {
heap_sort_test_helper(vec![3, 10, 5]);
}
#[test]
fn random3() {
heap_sort_test_helper(vec![1, 2, 3, 4]);
}
#[test]
fn duplicates() {
heap_sort_test_helper(vec![2, 2, 1, 1, 3, 3, 3]);
}
}