use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet};
use std::{
collections::HashSet,
hash::Hash,
iter::{FromIterator, FusedIterator},
};
#[derive(Debug, Clone)]
pub struct MultiMap<K, V> {
map: FxHashMap<K, FxHashSet<V>>,
}
impl<K, V> MultiMap<K, V> {
pub fn new() -> Self {
Self::default()
}
pub fn keys(&self) -> impl ExactSizeIterator<Item = &K> + FusedIterator {
self.map.keys()
}
pub fn len(&self) -> usize {
self.keys().len()
}
#[allow(dead_code)]
pub fn iter_pairs(&self) -> impl Iterator<Item = (&K, &V)> {
self.map
.iter()
.flat_map(|(k, v_set)| v_set.iter().map(move |v| (k, v)))
}
pub fn iter(&self) -> impl Iterator<Item = (&K, &FxHashSet<V>)> {
self.map.iter()
}
pub fn values(&self) -> impl Iterator<Item = &V> {
self.map.values().flatten()
}
}
impl<K, V> MultiMap<K, V>
where
K: Eq + Hash,
V: Eq + Hash,
{
pub fn get<Q>(&self, key: &Q) -> &FxHashSet<V>
where
K: std::borrow::Borrow<Q>,
Q: Eq + Hash,
{
self.map
.get(key)
.unwrap_or(const { &HashSet::with_hasher(FxBuildHasher) })
}
pub fn insert(&mut self, key: K, value: V) -> bool {
let set = self.map.entry(key).or_default();
set.insert(value)
}
#[cfg(feature = "dev")]
pub fn remove_one<F>(&mut self, key: K, pred: F) -> Option<V>
where
F: FnMut(&V) -> bool,
{
use std::collections::hash_map::Entry;
match self.map.entry(key) {
Entry::Occupied(mut entry) => {
let set = entry.get_mut();
let removed = set.extract_if(pred).next();
if set.is_empty() {
entry.remove();
}
removed
}
Entry::Vacant(_) => None,
}
}
pub fn contains_key<Q>(&self, k: &Q) -> bool
where
K: std::borrow::Borrow<Q>,
Q: Eq + Hash,
{
!self.get(k).is_empty()
}
}
impl<K, V> Default for MultiMap<K, V> {
fn default() -> Self {
Self {
map: FxHashMap::default(),
}
}
}
impl<K, V> FromIterator<(K, V)> for MultiMap<K, V>
where
K: Eq + Hash,
V: Eq + Hash,
{
fn from_iter<I>(iter: I) -> Self
where
I: IntoIterator<Item = (K, V)>,
{
let mut result = Self::new();
for (k, v) in iter {
result.insert(k, v);
}
result
}
}
#[cfg(test)]
mod tests {
use std::collections::HashSet;
use super::*;
#[test]
fn simple_lifecycle() {
let mut multimap = MultiMap::new();
assert!(multimap.get(&23).is_empty());
assert!(multimap.insert(23, 45));
assert_eq!(multimap.get(&23), &FxHashSet::from_iter([45]));
assert!(!multimap.insert(23, 45));
assert_eq!(multimap.get(&23), &FxHashSet::from_iter([45]));
assert!(multimap.insert(23, 67));
assert_eq!(multimap.get(&23), &FxHashSet::from_iter([45, 67]));
let full_set = multimap
.iter_pairs()
.map(|(x, y)| (*x, *y))
.collect::<HashSet<(i32, i32)>>();
let expected_full_set = HashSet::from([(23, 45), (23, 67)]);
assert_eq!(full_set, expected_full_set);
}
#[cfg(feature = "dev")]
#[test]
fn removals() {
let mut multimap = MultiMap::new();
assert!(multimap.insert(23, 45));
assert!(multimap.insert(23, 67));
assert_eq!(multimap.remove_one(23, |x| *x == 45), Some(45));
assert_eq!(multimap.get(&23), &FxHashSet::from_iter([67]));
assert_eq!(multimap.remove_one(23, |x| *x == 11), None);
assert_eq!(multimap.remove_one(11, |x| *x == 67), None);
assert_eq!(multimap.len(), 1);
assert_eq!(multimap.remove_one(23, |x| *x == 67), Some(67));
assert_eq!(multimap.len(), 0);
multimap.insert(1, 2);
multimap.insert(1, 3);
assert!(multimap.remove_one(1, |_| true).is_some());
assert_eq!(multimap.get(&1).len(), 1);
}
}