use std::{
cmp::Ordering,
collections::btree_map::{self, BTreeMap, Entry},
hash::Hash,
iter::FusedIterator,
num::NonZeroUsize,
ops::{BitAnd, BitOr, BitXor, RangeBounds, Sub},
};
#[derive(Clone, Debug, Default, Eq)]
pub struct NumericalMultiset<T> {
value_to_multiplicity: BTreeMap<T, NonZeroUsize>,
len: usize,
}
impl<T> NumericalMultiset<T> {
#[must_use = "Only effect is to produce a result"]
pub fn new() -> Self {
Self {
value_to_multiplicity: BTreeMap::new(),
len: 0,
}
}
pub fn clear(&mut self) {
self.value_to_multiplicity.clear();
self.len = 0;
}
#[must_use = "Only effect is to produce a result"]
pub fn len(&self) -> usize {
self.len
}
#[must_use = "Only effect is to produce a result"]
pub fn num_values(&self) -> usize {
self.value_to_multiplicity.len()
}
#[must_use = "Only effect is to produce a result"]
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[must_use = "Only effect is to produce a result"]
pub fn into_values(
self,
) -> impl DoubleEndedIterator<Item = T> + ExactSizeIterator + FusedIterator {
self.value_to_multiplicity.into_keys()
}
fn reset_len(&mut self) {
self.len = self.value_to_multiplicity.values().map(|x| x.get()).sum();
}
}
impl<T: Copy> NumericalMultiset<T> {
#[must_use = "Only effect is to produce a result"]
pub fn iter(&self) -> Iter<'_, T> {
self.into_iter()
}
#[must_use = "Only effect is to produce a result"]
pub fn values(
&self,
) -> impl DoubleEndedIterator<Item = T> + ExactSizeIterator + FusedIterator + Clone {
self.value_to_multiplicity.keys().copied()
}
}
impl<T: Ord> NumericalMultiset<T> {
#[inline]
#[must_use = "Only effect is to produce a result"]
pub fn contains(&self, value: T) -> bool {
self.value_to_multiplicity.contains_key(&value)
}
#[inline]
#[must_use = "Only effect is to produce a result"]
pub fn multiplicity(&self, value: T) -> Option<NonZeroUsize> {
self.value_to_multiplicity.get(&value).copied()
}
#[must_use = "Only effect is to produce a result"]
pub fn is_disjoint(&self, other: &Self) -> bool {
let mut iter1 = self.value_to_multiplicity.keys().peekable();
let mut iter2 = other.value_to_multiplicity.keys().peekable();
'joint_iter: loop {
match (iter1.peek(), iter2.peek()) {
(Some(value1), Some(value2)) => {
match value1.cmp(value2) {
Ordering::Less => {
let _ = iter1.next();
continue 'joint_iter;
}
Ordering::Greater => {
let _ = iter2.next();
continue 'joint_iter;
}
Ordering::Equal => return false,
}
}
(Some(_), None) | (None, Some(_)) | (None, None) => return true,
}
}
}
#[must_use = "Only effect is to produce a result"]
pub fn is_subset(&self, other: &Self) -> bool {
let mut other_iter = other.value_to_multiplicity.iter().peekable();
for (value, &multiplicity) in self.value_to_multiplicity.iter() {
'other_iter: loop {
match other_iter.peek() {
Some((other_value, other_multiplicity)) => match value.cmp(other_value) {
Ordering::Less => return false,
Ordering::Greater => {
let _ = other_iter.next();
continue 'other_iter;
}
Ordering::Equal => {
if **other_multiplicity < multiplicity {
return false;
}
let _ = other_iter.next();
break 'other_iter;
}
},
None => return false,
}
}
}
true
}
#[must_use = "Only effect is to produce a result"]
pub fn is_superset(&self, other: &Self) -> bool {
other.is_subset(self)
}
#[inline]
#[must_use = "Invalid removal should be handled"]
pub fn pop_all_first(&mut self) -> Option<(T, NonZeroUsize)> {
self.value_to_multiplicity
.pop_first()
.inspect(|(_value, count)| self.len -= count.get())
}
#[inline]
#[must_use = "Invalid removal should be handled"]
pub fn pop_all_last(&mut self) -> Option<(T, NonZeroUsize)> {
self.value_to_multiplicity
.pop_last()
.inspect(|(_value, count)| self.len -= count.get())
}
#[inline]
pub fn insert(&mut self, value: T) -> Option<NonZeroUsize> {
self.insert_multiple(value, NonZeroUsize::new(1).unwrap())
}
#[inline]
pub fn insert_multiple(&mut self, value: T, count: NonZeroUsize) -> Option<NonZeroUsize> {
let result = match self.value_to_multiplicity.entry(value) {
Entry::Vacant(v) => {
v.insert(count);
None
}
Entry::Occupied(mut o) => {
let old_count = *o.get();
*o.get_mut() = old_count
.checked_add(count.get())
.expect("Multiplicity counter has overflown");
Some(old_count)
}
};
self.len += count.get();
result
}
#[inline]
pub fn replace_all(&mut self, value: T, count: NonZeroUsize) -> Option<NonZeroUsize> {
let result = match self.value_to_multiplicity.entry(value) {
Entry::Vacant(v) => {
v.insert(count);
None
}
Entry::Occupied(mut o) => {
let old_count = *o.get();
*o.get_mut() = count;
self.len -= old_count.get();
Some(old_count)
}
};
self.len += count.get();
result
}
#[inline]
#[must_use = "Invalid removal should be handled"]
pub fn remove(&mut self, value: T) -> Option<NonZeroUsize> {
match self.value_to_multiplicity.entry(value) {
Entry::Vacant(_) => None,
Entry::Occupied(mut o) => {
let old_multiplicity = *o.get();
self.len -= 1;
match NonZeroUsize::new(old_multiplicity.get() - 1) {
Some(new_multiplicity) => {
*o.get_mut() = new_multiplicity;
}
None => {
o.remove_entry();
}
}
Some(old_multiplicity)
}
}
}
#[inline]
#[must_use = "Invalid removal should be handled"]
pub fn remove_all(&mut self, value: T) -> Option<NonZeroUsize> {
let result = self.value_to_multiplicity.remove(&value);
self.len -= result.map_or(0, |nz| nz.get());
result
}
pub fn split_off(&mut self, value: T) -> Self {
let mut result = Self {
value_to_multiplicity: self.value_to_multiplicity.split_off(&value),
len: 0,
};
self.reset_len();
result.reset_len();
result
}
}
impl<T: Copy + Ord> NumericalMultiset<T> {
#[must_use = "Only effect is to produce a result"]
pub fn range<R>(
&self,
range: R,
) -> impl DoubleEndedIterator<Item = (T, NonZeroUsize)> + FusedIterator
where
R: RangeBounds<T>,
{
self.value_to_multiplicity
.range(range)
.map(|(&k, &v)| (k, v))
}
#[must_use = "Only effect is to produce a result"]
pub fn difference<'a>(
&'a self,
other: &'a Self,
) -> impl Iterator<Item = (T, NonZeroUsize)> + Clone + 'a {
let mut iter = self.iter();
let mut other_iter = other.iter().peekable();
std::iter::from_fn(move || {
let (mut value, mut multiplicity) = iter.next()?;
'other_iter: loop {
match other_iter.peek() {
Some((other_value, other_multiplicity)) => match value.cmp(other_value) {
Ordering::Less => return Some((value, multiplicity)),
Ordering::Greater => {
let _ = other_iter.next();
continue 'other_iter;
}
Ordering::Equal => {
if multiplicity > *other_multiplicity {
let difference_multiplicity = NonZeroUsize::new(
multiplicity.get() - other_multiplicity.get(),
)
.expect("Checked above that this is fine");
let _ = other_iter.next();
return Some((value, difference_multiplicity));
} else {
let _ = other_iter.next();
(value, multiplicity) = iter.next()?;
continue 'other_iter;
}
}
},
None => return Some((value, multiplicity)),
}
}
})
}
#[must_use = "Only effect is to produce a result"]
pub fn symmetric_difference<'a>(
&'a self,
other: &'a Self,
) -> impl Iterator<Item = (T, NonZeroUsize)> + Clone + 'a {
let mut iter1 = self.iter().peekable();
let mut iter2 = other.iter().peekable();
std::iter::from_fn(move || {
'joint_iter: loop {
match (iter1.peek(), iter2.peek()) {
(Some((value1, multiplicity1)), Some((value2, multiplicity2))) => {
match value1.cmp(value2) {
Ordering::Less => return iter1.next(),
Ordering::Greater => return iter2.next(),
Ordering::Equal => {
if multiplicity1 != multiplicity2 {
let value12 = *value1;
let difference_multiplicity = NonZeroUsize::new(
multiplicity1.get().abs_diff(multiplicity2.get()),
)
.expect("Checked above that this is fine");
let _ = (iter1.next(), iter2.next());
return Some((value12, difference_multiplicity));
} else {
let _ = (iter1.next(), iter2.next());
continue 'joint_iter;
}
}
}
}
(Some(_), None) => return iter1.next(),
(None, Some(_)) => return iter2.next(),
(None, None) => return None,
}
}
})
}
#[must_use = "Only effect is to produce a result"]
pub fn intersection<'a>(
&'a self,
other: &'a Self,
) -> impl Iterator<Item = (T, NonZeroUsize)> + Clone + 'a {
let mut iter1 = self.iter().peekable();
let mut iter2 = other.iter().peekable();
std::iter::from_fn(move || {
'joint_iter: loop {
match (iter1.peek(), iter2.peek()) {
(Some((value1, multiplicity1)), Some((value2, multiplicity2))) => {
match value1.cmp(value2) {
Ordering::Less => {
let _ = iter1.next();
continue 'joint_iter;
}
Ordering::Greater => {
let _ = iter2.next();
continue 'joint_iter;
}
Ordering::Equal => {
let value12 = *value1;
let multiplicity12 = *multiplicity1.min(multiplicity2);
let _ = (iter1.next(), iter2.next());
return Some((value12, multiplicity12));
}
}
}
(Some(_), None) | (None, Some(_)) | (None, None) => return None,
}
}
})
}
#[must_use = "Only effect is to produce a result"]
pub fn union<'a>(
&'a self,
other: &'a Self,
) -> impl Iterator<Item = (T, NonZeroUsize)> + Clone + 'a {
let mut iter1 = self.iter().peekable();
let mut iter2 = other.iter().peekable();
std::iter::from_fn(move || match (iter1.peek(), iter2.peek()) {
(Some((value1, multiplicity1)), Some((value2, multiplicity2))) => {
match value1.cmp(value2) {
Ordering::Less => iter1.next(),
Ordering::Greater => iter2.next(),
Ordering::Equal => {
let value12 = *value1;
let multiplicity12 = *multiplicity1.max(multiplicity2);
let _ = (iter1.next(), iter2.next());
Some((value12, multiplicity12))
}
}
}
(Some(_), None) => iter1.next(),
(None, Some(_)) => iter2.next(),
(None, None) => None,
})
}
#[inline]
#[must_use = "Only effect is to produce a result"]
pub fn first(&self) -> Option<(T, NonZeroUsize)> {
self.value_to_multiplicity
.first_key_value()
.map(|(&k, &v)| (k, v))
}
#[inline]
#[must_use = "Only effect is to produce a result"]
pub fn last(&self) -> Option<(T, NonZeroUsize)> {
self.value_to_multiplicity
.last_key_value()
.map(|(&k, &v)| (k, v))
}
#[inline]
#[must_use = "Invalid removal should be handled"]
pub fn pop_first(&mut self) -> Option<T> {
let mut occupied = self.value_to_multiplicity.first_entry()?;
let old_multiplicity = *occupied.get();
let value = *occupied.key();
match NonZeroUsize::new(old_multiplicity.get() - 1) {
Some(new_multiplicity) => {
*occupied.get_mut() = new_multiplicity;
}
None => {
occupied.remove_entry();
}
}
self.len -= 1;
Some(value)
}
#[inline]
#[must_use = "Invalid removal should be handled"]
pub fn pop_last(&mut self) -> Option<T> {
let mut occupied = self.value_to_multiplicity.last_entry()?;
let old_multiplicity = *occupied.get();
let value = *occupied.key();
match NonZeroUsize::new(old_multiplicity.get() - 1) {
Some(new_multiplicity) => {
*occupied.get_mut() = new_multiplicity;
}
None => {
occupied.remove_entry();
}
}
self.len -= 1;
Some(value)
}
pub fn retain(&mut self, mut f: impl FnMut(T, &mut NonZeroUsize) -> bool) {
self.value_to_multiplicity.retain(|&k, v| f(k, v));
self.reset_len();
}
pub fn append(&mut self, other: &mut Self) {
if self.is_empty() {
std::mem::swap(self, other);
return;
}
for (value, multiplicity) in other.iter() {
self.insert_multiple(value, multiplicity);
}
other.clear();
}
}
impl<T: Copy + Ord> BitAnd<&NumericalMultiset<T>> for &NumericalMultiset<T> {
type Output = NumericalMultiset<T>;
fn bitand(self, rhs: &NumericalMultiset<T>) -> Self::Output {
self.intersection(rhs).collect()
}
}
impl<T: Copy + Ord> BitOr<&NumericalMultiset<T>> for &NumericalMultiset<T> {
type Output = NumericalMultiset<T>;
fn bitor(self, rhs: &NumericalMultiset<T>) -> Self::Output {
self.union(rhs).collect()
}
}
impl<T: Copy + Ord> BitXor<&NumericalMultiset<T>> for &NumericalMultiset<T> {
type Output = NumericalMultiset<T>;
fn bitxor(self, rhs: &NumericalMultiset<T>) -> Self::Output {
self.symmetric_difference(rhs).collect()
}
}
impl<T: Ord> Extend<T> for NumericalMultiset<T> {
fn extend<I: IntoIterator<Item = T>>(&mut self, iter: I) {
for element in iter {
self.insert(element);
}
}
}
impl<T: Ord> Extend<(T, NonZeroUsize)> for NumericalMultiset<T> {
fn extend<I: IntoIterator<Item = (T, NonZeroUsize)>>(&mut self, iter: I) {
for (value, count) in iter {
self.insert_multiple(value, count);
}
}
}
impl<T: Ord> FromIterator<T> for NumericalMultiset<T> {
fn from_iter<I: IntoIterator<Item = T>>(iter: I) -> Self {
let mut result = Self::new();
result.extend(iter);
result
}
}
impl<T: Ord> FromIterator<(T, NonZeroUsize)> for NumericalMultiset<T> {
fn from_iter<I: IntoIterator<Item = (T, NonZeroUsize)>>(iter: I) -> Self {
let mut result = Self::new();
result.extend(iter);
result
}
}
impl<T: Hash> Hash for NumericalMultiset<T> {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.value_to_multiplicity.hash(state)
}
}
impl<'a, T: Copy> IntoIterator for &'a NumericalMultiset<T> {
type Item = (T, NonZeroUsize);
type IntoIter = Iter<'a, T>;
fn into_iter(self) -> Self::IntoIter {
Iter(self.value_to_multiplicity.iter())
}
}
#[derive(Clone, Debug, Default)]
pub struct Iter<'a, T: Copy>(btree_map::Iter<'a, T, NonZeroUsize>);
impl<T: Copy> DoubleEndedIterator for Iter<'_, T> {
#[inline]
fn next_back(&mut self) -> Option<Self::Item> {
self.0.next_back().map(|(&k, &v)| (k, v))
}
#[inline]
fn nth_back(&mut self, n: usize) -> Option<Self::Item> {
self.0.nth_back(n).map(|(&k, &v)| (k, v))
}
}
impl<T: Copy> ExactSizeIterator for Iter<'_, T> {
fn len(&self) -> usize {
self.0.len()
}
}
impl<T: Copy> FusedIterator for Iter<'_, T> {}
impl<T: Copy> Iterator for Iter<'_, T> {
type Item = (T, NonZeroUsize);
#[inline]
fn next(&mut self) -> Option<Self::Item> {
self.0.next().map(|(&k, &v)| (k, v))
}
fn size_hint(&self) -> (usize, Option<usize>) {
self.0.size_hint()
}
fn count(self) -> usize
where
Self: Sized,
{
self.0.count()
}
fn last(mut self) -> Option<Self::Item>
where
Self: Sized,
{
self.0.next_back().map(|(&k, &v)| (k, v))
}
#[inline]
fn nth(&mut self, n: usize) -> Option<Self::Item> {
self.0.nth(n).map(|(&k, &v)| (k, v))
}
fn is_sorted(self) -> bool
where
Self: Sized,
Self::Item: PartialOrd,
{
true
}
}
impl<T> IntoIterator for NumericalMultiset<T> {
type Item = (T, NonZeroUsize);
type IntoIter = IntoIter<T>;
fn into_iter(self) -> Self::IntoIter {
IntoIter(self.value_to_multiplicity.into_iter())
}
}
#[derive(Debug, Default)]
pub struct IntoIter<T>(btree_map::IntoIter<T, NonZeroUsize>);
impl<T> DoubleEndedIterator for IntoIter<T> {
#[inline]
fn next_back(&mut self) -> Option<Self::Item> {
self.0.next_back()
}
#[inline]
fn nth_back(&mut self, n: usize) -> Option<Self::Item> {
self.0.nth_back(n)
}
}
impl<T> ExactSizeIterator for IntoIter<T> {
fn len(&self) -> usize {
self.0.len()
}
}
impl<T> FusedIterator for IntoIter<T> {}
impl<T> Iterator for IntoIter<T> {
type Item = (T, NonZeroUsize);
#[inline]
fn next(&mut self) -> Option<Self::Item> {
self.0.next()
}
fn size_hint(&self) -> (usize, Option<usize>) {
self.0.size_hint()
}
fn count(self) -> usize
where
Self: Sized,
{
self.0.count()
}
fn last(mut self) -> Option<Self::Item>
where
Self: Sized,
{
self.0.next_back()
}
#[inline]
fn nth(&mut self, n: usize) -> Option<Self::Item> {
self.0.nth(n)
}
fn is_sorted(self) -> bool
where
Self: Sized,
Self::Item: PartialOrd,
{
true
}
}
impl<T: Ord> Ord for NumericalMultiset<T> {
fn cmp(&self, other: &Self) -> Ordering {
self.value_to_multiplicity.cmp(&other.value_to_multiplicity)
}
}
impl<T: PartialEq> PartialEq for NumericalMultiset<T> {
fn eq(&self, other: &Self) -> bool {
self.len == other.len && self.value_to_multiplicity == other.value_to_multiplicity
}
}
impl<T: PartialOrd> PartialOrd for NumericalMultiset<T> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
self.value_to_multiplicity
.partial_cmp(&other.value_to_multiplicity)
}
}
impl<T: Copy + Ord> Sub<&NumericalMultiset<T>> for &NumericalMultiset<T> {
type Output = NumericalMultiset<T>;
fn sub(self, rhs: &NumericalMultiset<T>) -> Self::Output {
self.difference(rhs).collect()
}
}
#[cfg(test)]
mod test {
use super::*;
use proptest::{prelude::*, sample::SizeRange};
use std::{
cmp::Ordering,
collections::HashSet,
fmt::Debug,
hash::{BuildHasher, RandomState},
ops::Range,
};
const ONE: NonZeroUsize = NonZeroUsize::MIN;
fn check_equal_iterable<V, It1, It2>(it1: It1, it2: It2)
where
It1: IntoIterator<Item = V>,
It2: IntoIterator<Item = V>,
V: Debug + PartialEq,
{
assert_eq!(
it1.into_iter().collect::<Vec<_>>(),
it2.into_iter().collect::<Vec<_>>(),
);
}
fn check_empty_iterable<It>(it: It)
where
It: IntoIterator,
It::Item: Debug + PartialEq,
{
check_equal_iterable(it, std::iter::empty());
}
fn check_any_mset_pair(mset1: &NumericalMultiset<i32>, mset2: &NumericalMultiset<i32>) {
let intersection = mset1 & mset2;
for (val, mul) in &intersection {
assert_eq!(
mul,
mset1
.multiplicity(val)
.unwrap()
.min(mset2.multiplicity(val).unwrap()),
);
}
for val1 in mset1.values() {
assert!(intersection.contains(val1) || !mset2.contains(val1));
}
for val2 in mset2.values() {
assert!(intersection.contains(val2) || !mset1.contains(val2));
}
check_equal_iterable(mset1.intersection(mset2), &intersection);
let union = mset1 | mset2;
for (val, mul) in &union {
assert_eq!(
mul.get(),
mset1
.multiplicity(val)
.map_or(0, |nz| nz.get())
.max(mset2.multiplicity(val).map_or(0, |nz| nz.get()))
);
}
for val in mset1.values().chain(mset2.values()) {
assert!(union.contains(val));
}
check_equal_iterable(mset1.union(mset2), &union);
let difference = mset1 - mset2;
for (val, mul) in &difference {
assert_eq!(
mul.get(),
mset1
.multiplicity(val)
.unwrap()
.get()
.checked_sub(mset2.multiplicity(val).map_or(0, |nz| nz.get()))
.unwrap()
);
}
for (val, mul1) in mset1 {
assert!(difference.contains(val) || mset2.multiplicity(val).unwrap() >= mul1);
}
check_equal_iterable(mset1.difference(mset2), difference);
let symmetric_difference = mset1 ^ mset2;
for (val, mul) in &symmetric_difference {
assert_eq!(
mul.get(),
mset1
.multiplicity(val)
.map_or(0, |nz| nz.get())
.abs_diff(mset2.multiplicity(val).map_or(0, |nz| nz.get()))
);
}
for (val1, mul1) in mset1 {
assert!(
symmetric_difference.contains(val1) || mset2.multiplicity(val1).unwrap() >= mul1
);
}
for (val2, mul2) in mset2 {
assert!(
symmetric_difference.contains(val2) || mset1.multiplicity(val2).unwrap() >= mul2
);
}
check_equal_iterable(mset1.symmetric_difference(mset2), symmetric_difference);
assert_eq!(mset1.is_disjoint(mset2), intersection.is_empty(),);
if mset1.is_subset(mset2) {
for (val, mul1) in mset1 {
assert!(mset2.multiplicity(val).unwrap() >= mul1);
}
} else {
assert!(mset1
.iter()
.any(|(val1, mul1)| { mset2.multiplicity(val1).is_none_or(|mul2| mul2 < mul1) }))
}
assert_eq!(mset2.is_superset(mset1), mset1.is_subset(mset2));
let mut combined = mset1.clone();
let mut appended = mset2.clone();
combined.append(&mut appended);
assert_eq!(
combined,
mset1
.iter()
.chain(mset2.iter())
.collect::<NumericalMultiset<_>>()
);
assert!(appended.is_empty());
let mut extended_by_tuples = mset1.clone();
extended_by_tuples.extend(mset2.iter());
assert_eq!(extended_by_tuples, combined);
let mut extended_by_values = mset1.clone();
extended_by_values.extend(
mset2
.iter()
.flat_map(|(val, mul)| std::iter::repeat_n(val, mul.get())),
);
assert_eq!(
extended_by_values, combined,
"{mset1:?} + {mset2:?} != {extended_by_values:?}"
);
assert_eq!(
mset1.cmp(mset2),
mset1
.value_to_multiplicity
.cmp(&mset2.value_to_multiplicity)
);
assert_eq!(mset1.partial_cmp(mset2), Some(mset1.cmp(mset2)));
}
fn check_any_mset(mset: &NumericalMultiset<i32>, contents: &[(i32, NonZeroUsize)]) {
let sorted_contents = contents
.iter()
.map(|(v, m)| (*v, m.get()))
.collect::<BTreeMap<i32, usize>>();
check_equal_iterable(
mset.iter().map(|(val, mul)| (val, mul.get())),
sorted_contents.iter().map(|(&k, &v)| (k, v)),
);
check_equal_iterable(mset, mset.iter());
check_equal_iterable(mset.clone(), mset.iter());
check_equal_iterable(mset.range(..), mset.iter());
check_equal_iterable(mset.values(), sorted_contents.keys().copied());
check_equal_iterable(mset.clone().into_values(), mset.values());
assert_eq!(mset.len(), sorted_contents.values().sum());
assert_eq!(mset.num_values(), contents.len());
assert_eq!(mset.is_empty(), contents.is_empty());
for (&val, &mul) in &sorted_contents {
assert!(mset.contains(val));
assert_eq!(mset.multiplicity(val).unwrap().get(), mul);
}
assert_eq!(
mset.first().map(|(val, mul)| (val, mul.get())),
sorted_contents.first_key_value().map(|(&k, &v)| (k, v)),
);
assert_eq!(
mset.last().map(|(val, mul)| (val, mul.get())),
sorted_contents.last_key_value().map(|(&k, &v)| (k, v)),
);
#[allow(clippy::eq_op)]
{
assert_eq!(mset, mset);
}
assert_eq!(*mset, mset.clone());
assert_eq!(mset.cmp(mset), Ordering::Equal);
assert_eq!(mset.partial_cmp(mset), Some(mset.cmp(mset)));
let state = RandomState::new();
assert_eq!(
state.hash_one(mset),
state.hash_one(&mset.value_to_multiplicity),
);
let mut mutable = mset.clone();
if let Some((first, first_mul)) = mset.first() {
assert_eq!(mutable.pop_all_first(), Some((first, first_mul)));
assert_eq!(mutable.len(), mset.len() - first_mul.get());
assert_eq!(mutable.num_values(), mset.num_values() - 1);
assert!(!mutable.contains(first));
assert_eq!(mutable.multiplicity(first), None);
assert_ne!(mutable, *mset);
assert_eq!(mutable.insert_multiple(first, first_mul), None);
assert_eq!(mutable, *mset);
assert_eq!(mutable.pop_first(), Some(first));
assert_eq!(mutable.len(), mset.len() - 1);
let new_first_mul = NonZeroUsize::new(first_mul.get() - 1);
let first_is_single = new_first_mul.is_none();
assert_eq!(
mutable.num_values(),
mset.num_values() - first_is_single as usize
);
assert_eq!(mutable.contains(first), !first_is_single);
assert_eq!(mutable.multiplicity(first), new_first_mul);
assert_ne!(mutable, *mset);
assert_eq!(mutable.insert(first), new_first_mul);
assert_eq!(mutable, *mset);
let (last, last_mul) = mset.last().unwrap();
assert_eq!(mutable.pop_all_last(), Some((last, last_mul)));
assert_eq!(mutable.len(), mset.len() - last_mul.get());
assert_eq!(mutable.num_values(), mset.num_values() - 1);
assert!(!mutable.contains(last));
assert_eq!(mutable.multiplicity(last), None);
assert_ne!(mutable, *mset);
assert_eq!(mutable.insert_multiple(last, last_mul), None);
assert_eq!(mutable, *mset);
assert_eq!(mutable.pop_last(), Some(last));
assert_eq!(mutable.len(), mset.len() - 1);
let new_last_mul = NonZeroUsize::new(last_mul.get() - 1);
let last_is_single = new_last_mul.is_none();
assert_eq!(
mutable.num_values(),
mset.num_values() - last_is_single as usize
);
assert_eq!(mutable.contains(last), !last_is_single);
assert_eq!(mutable.multiplicity(last), new_last_mul);
assert_ne!(mutable, *mset);
assert_eq!(mutable.insert(last), new_last_mul);
assert_eq!(mutable, *mset);
} else {
assert!(mset.is_empty());
assert_eq!(mutable.pop_first(), None);
assert!(mutable.is_empty());
assert_eq!(mutable.pop_all_first(), None);
assert!(mutable.is_empty());
assert_eq!(mutable.pop_last(), None);
assert!(mutable.is_empty());
assert_eq!(mutable.pop_all_last(), None);
assert!(mutable.is_empty());
}
let mut retain_all = mset.clone();
retain_all.retain(|_, _| true);
assert_eq!(retain_all, *mset);
let mut retain_nothing = mset.clone();
retain_nothing.retain(|_, _| false);
assert!(retain_nothing.is_empty());
}
fn check_empty_mset(empty: &NumericalMultiset<i32>) {
check_any_mset(empty, &[]);
assert_eq!(empty.len(), 0);
assert_eq!(empty.num_values(), 0);
assert!(empty.is_empty());
assert_eq!(empty.first(), None);
assert_eq!(empty.last(), None);
check_empty_iterable(empty.iter());
check_empty_iterable(empty.values());
check_empty_iterable(empty.clone());
check_empty_iterable(empty.clone().into_values());
let mut mutable = empty.clone();
assert_eq!(mutable.pop_first(), None);
assert_eq!(mutable.pop_last(), None);
assert_eq!(mutable.pop_all_first(), None);
assert_eq!(mutable.pop_all_last(), None);
}
fn check_clear_outcome(mut mset: NumericalMultiset<i32>) {
mset.clear();
check_empty_mset(&mset);
}
#[test]
fn empty() {
check_empty_mset(&NumericalMultiset::default());
let mset = NumericalMultiset::<i32>::new();
check_empty_mset(&mset);
check_clear_outcome(mset);
}
fn max_multiplicity() -> usize {
SizeRange::default().end_excl()
}
fn multiplicity() -> impl Strategy<Value = NonZeroUsize> {
prop_oneof![Just(1), Just(2), 3..max_multiplicity()]
.prop_map(|m| NonZeroUsize::new(m).unwrap())
}
fn mset_contents() -> impl Strategy<Value = Vec<(i32, NonZeroUsize)>> {
any::<HashSet<i32>>().prop_flat_map(|values| {
prop::collection::vec(multiplicity(), values.len()).prop_map(move |multiplicities| {
values.iter().copied().zip(multiplicities).collect()
})
})
}
proptest! {
#[test]
fn single(contents in mset_contents()) {
for mset in [
contents.iter().copied().collect(),
contents.iter().flat_map(|(v, m)| {
std::iter::repeat_n(*v, m.get())
}).collect(),
] {
check_any_mset(&mset, &contents);
check_any_mset_pair(&mset, &mset);
let empty = NumericalMultiset::default();
check_any_mset_pair(&mset, &empty);
check_any_mset_pair(&empty, &mset);
check_clear_outcome(mset);
}
}
}
fn mset() -> impl Strategy<Value = NumericalMultiset<i32>> {
mset_contents().prop_map(NumericalMultiset::from_iter)
}
fn mset_and_value() -> impl Strategy<Value = (NumericalMultiset<i32>, i32)> {
mset().prop_flat_map(|mset| {
if mset.is_empty() {
(Just(mset), any::<i32>()).boxed()
} else {
let inner_value = prop::sample::select(mset.values().collect::<Vec<_>>());
let value = prop_oneof![inner_value, any::<i32>()];
(Just(mset), value).boxed()
}
})
}
proptest! {
#[test]
fn with_value((initial, value) in mset_and_value()) {
if let Some(&mul) = initial.value_to_multiplicity.get(&value) {
assert!(initial.contains(value));
assert_eq!(initial.multiplicity(value), Some(mul));
{
let mut mset = initial.clone();
assert_eq!(mset.insert(value), Some(mul));
let mut expected = initial.clone();
let Entry::Occupied(mut entry) = expected.value_to_multiplicity.entry(value) else {
unreachable!();
};
*entry.get_mut() = mul.checked_add(1).unwrap();
expected.len += 1;
assert_eq!(mset, expected);
}
{
let mut mset = initial.clone();
assert_eq!(mset.remove(value), Some(mul));
let mut expected = initial.clone();
if mul == ONE {
expected.value_to_multiplicity.remove(&value);
} else {
let Entry::Occupied(mut entry) = expected.value_to_multiplicity.entry(value) else {
unreachable!();
};
*entry.get_mut() = NonZeroUsize::new(mul.get() - 1).unwrap();
}
expected.len -= 1;
assert_eq!(mset, expected);
}
{
let mut mset = initial.clone();
assert_eq!(mset.remove_all(value), Some(mul));
let mut expected = initial.clone();
expected.value_to_multiplicity.remove(&value);
expected.len -= mul.get();
assert_eq!(mset, expected);
}
} else {
assert!(!initial.contains(value));
assert_eq!(initial.multiplicity(value), None);
{
let mut mset = initial.clone();
assert_eq!(mset.insert(value), None);
let mut expected = initial.clone();
expected.value_to_multiplicity.insert(value, ONE);
expected.len += 1;
assert_eq!(mset, expected);
}
{
let mut mset = initial.clone();
assert_eq!(mset.remove(value), None);
assert_eq!(mset, initial);
assert_eq!(mset.remove_all(value), None);
assert_eq!(mset, initial);
}
}
{
let mut mset = initial.clone();
let ge = mset.split_off(value);
let lt = mset;
let mut expected_lt = initial.clone();
let ge_value_to_multiplicity = expected_lt.value_to_multiplicity.split_off(&value);
expected_lt.reset_len();
assert_eq!(lt, expected_lt);
let expected_ge = NumericalMultiset::from_iter(ge_value_to_multiplicity);
assert_eq!(ge, expected_ge);
}
}
#[test]
fn with_value_and_multiplicity((initial, value) in mset_and_value(),
new_mul in multiplicity()) {
if let Some(&initial_mul) = initial.value_to_multiplicity.get(&value) {
{
let mut mset = initial.clone();
assert_eq!(mset.insert_multiple(value, new_mul), Some(initial_mul));
let mut expected = initial.clone();
let Entry::Occupied(mut entry) = expected.value_to_multiplicity.entry(value) else {
unreachable!();
};
*entry.get_mut() = initial_mul.checked_add(new_mul.get()).unwrap();
expected.len += new_mul.get();
assert_eq!(mset, expected);
}
{
let mut mset = initial.clone();
assert_eq!(mset.replace_all(value, new_mul), Some(initial_mul));
let mut expected = initial.clone();
let Entry::Occupied(mut entry) = expected.value_to_multiplicity.entry(value) else {
unreachable!();
};
*entry.get_mut() = new_mul;
expected.len = expected.len - initial_mul.get() + new_mul.get();
assert_eq!(mset, expected);
}
} else {
let mut inserted = initial.clone();
assert_eq!(inserted.insert_multiple(value, new_mul), None);
let mut expected = initial.clone();
expected.value_to_multiplicity.insert(value, new_mul);
expected.len += new_mul.get();
assert_eq!(inserted, expected);
let mut replaced = initial.clone();
assert_eq!(replaced.replace_all(value, new_mul), None);
assert_eq!(replaced, expected);
}
{
let f = |v, m: &mut NonZeroUsize| {
if v <= value && *m <= new_mul {
*m = m.checked_add(42).unwrap();
true
} else {
false
}
};
let mut retained = initial.clone();
retained.retain(f);
let mut expected = initial.clone();
expected.value_to_multiplicity.retain(|&v, m| f(v, m));
expected.reset_len();
assert_eq!(retained, expected);
}
}
}
fn mset_and_value_range() -> impl Strategy<Value = (NumericalMultiset<i32>, Range<i32>)> {
let pair_to_range = |values: [i32; 2]| values[0]..values[1];
mset().prop_flat_map(move |mset| {
if mset.is_empty() {
(Just(mset), any::<[i32; 2]>().prop_map(pair_to_range)).boxed()
} else {
let inner_value = || prop::sample::select(mset.values().collect::<Vec<_>>());
let value = || prop_oneof![3 => inner_value(), 2 => any::<i32>()];
let range = [value(), value()].prop_map(pair_to_range);
(Just(mset), range).boxed()
}
})
}
proptest! {
#[test]
fn range((mset, range) in mset_and_value_range()) {
match std::panic::catch_unwind(|| {
mset.range(range.clone()).collect::<Vec<_>>()
}) {
Ok(output) => check_equal_iterable(output, mset.value_to_multiplicity.range(range).map(|(&v, &m)| (v, m))),
Err(_panicked) => assert!(range.start > range.end),
}
}
}
fn mset_pair() -> impl Strategy<Value = (NumericalMultiset<i32>, NumericalMultiset<i32>)> {
mset().prop_flat_map(|mset1| {
if mset1.is_empty() {
(Just(mset1), mset()).boxed()
} else {
let related = prop::sample::subsequence(
mset1.iter().collect::<Vec<_>>(),
0..mset1.num_values(),
)
.prop_flat_map(move |subseq| {
subseq
.into_iter()
.map(|(v, m)| {
let m = m.get();
let multiplicity = match m {
1 => prop_oneof![Just(1), 2..max_multiplicity()].boxed(),
2 => prop_oneof![Just(1), Just(2), 3..max_multiplicity()].boxed(),
_ if m + 1 < max_multiplicity() => {
prop_oneof![Just(1), 2..m, Just(m), (m + 1)..max_multiplicity()]
.boxed()
}
_ => prop_oneof![Just(1), 2..max_multiplicity()].boxed(),
}
.prop_map(|m| NonZeroUsize::new(m).unwrap());
(Just(v), multiplicity)
})
.collect::<Vec<_>>()
})
.prop_map(|elems| elems.into_iter().collect());
let related_pair = (Just(mset1.clone()), related, any::<bool>()).prop_map(
|(mset1, mset2, flip)| {
if flip {
(mset2, mset1)
} else {
(mset1, mset2)
}
},
);
prop_oneof![
1 => (Just(mset1), mset()),
4 => related_pair,
]
.boxed()
}
})
}
proptest! {
#[test]
fn pair((mset1, mset2) in mset_pair()) {
check_any_mset_pair(&mset1, &mset2);
}
}
}