1use std::borrow::Borrow;
15use std::collections::HashMap;
16use std::hash::{Hash, Hasher};
17use std::mem::size_of_val;
18
19use bitpacking::{BitPacker, BitPacker1x};
20use num::{PrimInt, Unsigned};
21use wyhash::WyHash;
22
23use crate::mphf::{Mphf, DEFAULT_GAMMA};
24
25#[derive(Default)]
35#[cfg_attr(feature = "rkyv_derive", derive(rkyv::Archive, rkyv::Deserialize, rkyv::Serialize))]
36#[cfg_attr(feature = "rkyv_derive", archive_attr(derive(rkyv::CheckBytes)))]
37#[cfg_attr(feature = "serde", derive(serde::Serialize))]
38#[cfg_attr(
39 feature = "serde",
40 serde(bound(serialize = "K: serde::Serialize, ST: serde::Serialize"))
41)]
42pub struct MapWithDictBitpacked<K, const B: usize = 32, const S: usize = 8, ST = u8, H = WyHash>
43where
44 ST: PrimInt + Unsigned,
45 H: Hasher + Default,
46{
47 mphf: Mphf<B, S, ST, H>,
49 keys: Box<[K]>,
51 values_index: Box<[usize]>,
53 #[cfg_attr(feature = "serde", serde(with = "serde_bytes"))]
55 values_dict: Box<[u8]>,
56}
57
58#[cfg(feature = "serde")]
59#[derive(serde::Deserialize)]
60#[serde(bound(deserialize = "K: serde::Deserialize<'de>, ST: serde::Deserialize<'de>"))]
61struct MapWithDictBitpackedUnchecked<K, const B: usize = 32, const S: usize = 8, ST = u8, H = WyHash>
62where
63 ST: PrimInt + Unsigned,
64 H: Hasher + Default,
65{
66 mphf: Mphf<B, S, ST, H>,
67 keys: Box<[K]>,
68 values_index: Box<[usize]>,
69 #[serde(with = "serde_bytes")]
70 values_dict: Box<[u8]>,
71}
72
73#[cfg(feature = "serde")]
77impl<'de, K, const B: usize, const S: usize, ST, H> serde::Deserialize<'de> for MapWithDictBitpacked<K, B, S, ST, H>
78where
79 K: serde::Deserialize<'de> + Hash,
80 ST: serde::Deserialize<'de> + PrimInt + Unsigned,
81 H: Hasher + Default,
82{
83 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
84 where
85 D: serde::Deserializer<'de>,
86 {
87 use crate::{ValidateKeyResult, ValidateValueResult};
88 use serde::de::Error;
89
90 let this = MapWithDictBitpackedUnchecked::deserialize(deserializer)?;
91
92 this.mphf.validate_keys(&this.keys).map_err(|e| match e {
93 ValidateKeyResult::InvalidKeyCount => {
94 Error::custom("key count should equal the number of set bits in the MPHF")
95 }
96 ValidateKeyResult::IncorrectKeyOrder => Error::custom("keys should correspond to MPHF index"),
97 })?;
98
99 this.mphf
100 .validate_values(&this.keys, &this.values_index, &this.values_dict)
101 .map_err(|e| match e {
102 ValidateValueResult::KeyValueLenMismatch => Error::custom("key count should equal value count"),
103 ValidateValueResult::InvalidValueIndex => Error::custom("value index is out of bounds of value dict"),
104 })?;
105
106 if this.values_index.iter().any(|&i| this.values_dict[i] > 32) {
107 return Err(Error::custom("value index out of num_bits bounds"));
108 }
109
110 Ok(Self {
111 mphf: this.mphf,
112 keys: this.keys,
113 values_index: this.values_index,
114 values_dict: this.values_dict,
115 })
116 }
117}
118
119#[derive(Debug)]
121pub enum Error {
122 MphfError(crate::mphf::MphfError),
124 NotEqualValuesLengths,
126}
127
128impl<K, const B: usize, const S: usize, ST, H> MapWithDictBitpacked<K, B, S, ST, H>
129where
130 K: Hash + PartialEq + Clone,
131 ST: PrimInt + Unsigned,
132 H: Hasher + Default,
133{
134 pub fn from_iter_with_params<I>(iter: I, gamma: f32) -> Result<Self, Error>
136 where
137 I: IntoIterator<Item = (K, Vec<u32>)>,
138 {
139 let mut keys = vec![];
140 let mut offsets_cache = HashMap::new();
141 let mut values_index = vec![];
142 let mut values_dict = vec![];
143
144 let mut iter = iter.into_iter().peekable();
145 let v_len = iter.peek().map_or(0, |(_, v)| v.len());
146
147 for (k, v) in iter {
148 keys.push(k.clone());
149
150 if v.len() != v_len {
151 return Err(Error::NotEqualValuesLengths);
152 }
153
154 if let Some(&offset) = offsets_cache.get(&v) {
155 values_index.push(offset);
157 } else {
158 let offset = values_dict.len();
160 offsets_cache.insert(v.clone(), offset);
161 values_index.push(offset);
162
163 pack_values(&v, &mut values_dict);
165 }
166 }
167
168 values_dict.resize(values_dict.len() + 4 * VALUES_BLOCK_LEN, 0);
170
171 let mphf = Mphf::from_slice(&keys, gamma).map_err(Error::MphfError)?;
172
173 for i in 0..keys.len() {
175 loop {
176 let idx = mphf.get(&keys[i]).unwrap();
177 if idx == i {
178 break;
179 }
180 keys.swap(i, idx);
181 values_index.swap(i, idx);
182 }
183 }
184
185 Ok(MapWithDictBitpacked {
186 mphf,
187 keys: keys.into_boxed_slice(),
188 values_index: values_index.into_boxed_slice(),
189 values_dict: values_dict.into_boxed_slice(),
190 })
191 }
192
193 #[inline]
212 pub fn get_values<Q>(&self, key: &Q, values: &mut [u32]) -> bool
213 where
214 K: Borrow<Q> + PartialEq<Q>,
215 Q: Hash + Eq + ?Sized,
216 {
217 let idx = match self.mphf.get(key) {
218 Some(idx) => idx,
219 None => return false,
220 };
221
222 unsafe {
224 if self.keys.get_unchecked(idx) != key {
225 return false;
226 }
227
228 let value_idx = *self.values_index.get_unchecked(idx);
230 let dict = self.values_dict.get_unchecked(value_idx..);
231 unpack_values(dict, values);
232 }
233
234 true
235 }
236
237 #[inline]
247 pub fn len(&self) -> usize {
248 self.keys.len()
249 }
250
251 #[inline]
263 pub fn is_empty(&self) -> bool {
264 self.keys.is_empty()
265 }
266
267 #[inline]
278 pub fn contains_key<Q>(&self, key: &Q) -> bool
279 where
280 K: Borrow<Q> + PartialEq<Q>,
281 Q: Hash + Eq + ?Sized,
282 {
283 if let Some(idx) = self.mphf.get(key) {
284 unsafe { self.keys.get_unchecked(idx) == key }
286 } else {
287 false
288 }
289 }
290
291 #[inline]
303 pub fn iter(&self, n: usize) -> impl Iterator<Item = (&K, Vec<u32>)> {
304 self.keys().zip(self.values_index.iter()).map(move |(key, &value_idx)| {
305 let mut values = vec![0; n];
306 let dict = unsafe { self.values_dict.get_unchecked(value_idx..) };
308 unpack_values(dict, &mut values);
309 (key, values)
310 })
311 }
312
313 #[inline]
325 pub fn keys(&self) -> impl Iterator<Item = &K> {
326 self.keys.iter()
327 }
328
329 #[inline]
341 pub fn values(&self, n: usize) -> impl Iterator<Item = Vec<u32>> + '_ {
342 self.values_index.iter().map(move |&value_idx| {
343 let mut values = vec![0; n];
344 let dict = unsafe { self.values_dict.get_unchecked(value_idx..) };
346 unpack_values(dict, &mut values);
347 values
348 })
349 }
350
351 pub fn size(&self) -> usize {
361 size_of_val(self)
362 + self.mphf.size()
363 + size_of_val(self.keys.as_ref())
364 + size_of_val(self.values_index.as_ref())
365 + size_of_val(self.values_dict.as_ref())
366 }
367}
368
369impl<K> TryFrom<HashMap<K, Vec<u32>>> for MapWithDictBitpacked<K>
371where
372 K: PartialEq + Hash + Clone,
373{
374 type Error = Error;
375
376 #[inline]
377 fn try_from(value: HashMap<K, Vec<u32>>) -> Result<Self, Self::Error> {
378 MapWithDictBitpacked::from_iter_with_params(value, DEFAULT_GAMMA)
379 }
380}
381
382const VALUES_BLOCK_LEN: usize = BitPacker1x::BLOCK_LEN;
384
385fn pack_values(values: &[u32], dict: &mut Vec<u8>) {
388 let bitpacker = BitPacker1x::new();
390
391 for block in values.chunks(VALUES_BLOCK_LEN) {
392 let mut values_block = [0u32; VALUES_BLOCK_LEN];
393 let mut values_packed_block = [0u8; 4 * VALUES_BLOCK_LEN];
394
395 values_block[..block.len()].copy_from_slice(block);
396
397 let num_bits = bitpacker.num_bits(&values_block);
399
400 bitpacker.compress(&values_block, &mut values_packed_block, num_bits);
402
403 let size = (block.len() * (num_bits as usize)).div_ceil(8);
405 dict.push(num_bits);
406 dict.extend_from_slice(&values_packed_block[..size]);
407 }
408}
409
410fn unpack_values(dict: &[u8], res: &mut [u32]) {
413 let bitpacker = BitPacker1x::new();
414 let mut dict = dict;
415 for block in res.chunks_mut(VALUES_BLOCK_LEN) {
416 let mut values_block = [0u32; VALUES_BLOCK_LEN];
417
418 let num_bits = dict[0];
420 dict = &dict[1..];
421
422 let size = (block.len() * (num_bits as usize)).div_ceil(8);
424 bitpacker.decompress(dict, &mut values_block, num_bits);
425 dict = &dict[size..];
426
427 block.copy_from_slice(&values_block[..block.len()]);
428 }
429}
430
431#[cfg(feature = "rkyv_derive")]
433impl<K, const B: usize, const S: usize, ST, H> ArchivedMapWithDictBitpacked<K, B, S, ST, H>
434where
435 K: PartialEq + Hash + rkyv::Archive,
436 K::Archived: PartialEq<K>,
437 ST: PrimInt + Unsigned + rkyv::Archive<Archived = ST>,
438 H: Hasher + Default,
439{
440 #[inline]
457 pub fn get_values(&self, key: &K, values: &mut [u32]) -> bool {
458 let idx = match self.mphf.get(key) {
459 Some(idx) => idx,
460 None => return false,
461 };
462
463 unsafe {
465 if self.keys.get_unchecked(idx) != key {
466 return false;
467 }
468
469 let value_idx = *self.values_index.get_unchecked(idx) as usize;
471 let dict = self.values_dict.get_unchecked(value_idx..);
472 unpack_values(dict, values);
473 }
474
475 true
476 }
477}
478
479#[cfg(test)]
480mod tests {
481 use super::*;
482 use paste::paste;
483 use proptest::prelude::*;
484 use rand::{Rng, SeedableRng};
485 use rand_chacha::ChaCha8Rng;
486 use std::collections::{hash_map::RandomState, HashSet};
487 use test_case::test_case;
488
489 #[test_case(
490 &[] => Vec::<u8>::new();
491 "empty values"
492 )]
493 #[test_case(
494 &[0] => vec![0];
495 "single 0-bit value"
496 )]
497 #[test_case(
498 &[0; 10] => vec![0];
499 "10 0-bit value"
500 )]
501 #[test_case(
502 &[0; 77] => vec![0, 0, 0];
503 "77 0-bit values (3 blocks)"
504 )]
505 #[test_case(
506 &[1] => vec![1, 1];
507 "single 1-bit value"
508 )]
509 #[test_case(
510 &[1; 10] => vec![1, 0b11111111, 0b00000011];
511 "10 1-bit value"
512 )]
513 #[test_case(
514 &[1; 32] => vec![1, 0b11111111, 0b11111111, 0b11111111, 0b11111111];
515 "32 1-bit value"
516 )]
517 #[test_case(
518 &[1; 33] => vec![1, 0b11111111, 0b11111111, 0b11111111, 0b11111111, 1, 0b00000001];
519 "33 1-bit value"
520 )]
521 #[test_case(
522 &[1, 2, 3, 4, 5, 6, 7, 8, 9, 10] => vec![4, 0b0010_0001, 0b0100_0011, 0b0110_0101, 0b1000_0111, 0b1010_1001];
523 "10 4-bit value"
524 )]
525 fn test_pack_unpack(values: &[u32]) -> Vec<u8> {
526 let mut dict = vec![];
527 pack_values(values, &mut dict);
528
529 let mut padded_dict = dict.clone();
530 padded_dict.resize(dict.len() + 4 * VALUES_BLOCK_LEN, 0);
531
532 let mut unpacked_values = vec![0; values.len()];
533 unpack_values(&padded_dict, &mut unpacked_values);
534
535 assert_eq!(values, unpacked_values);
536
537 dict
538 }
539
540 #[test]
541 fn test_pack_unpack_random() {
542 let max_n = 200;
543 let mut rng = ChaCha8Rng::seed_from_u64(123);
544 let mut dict = vec![];
545 let mut values = vec![];
546 let mut unpacked_values = vec![];
547
548 for n in 1..=max_n {
549 for num_bits in 0..=32 {
550 values.clear();
551 values.extend((0..n).map(|_| rng.gen::<u32>() & ((1u32 << (num_bits % 32)) - 1)));
552 dict.clear();
553
554 pack_values(&values, &mut dict);
555 assert!(!dict.is_empty());
556
557 dict.resize(dict.len() + 4 * VALUES_BLOCK_LEN, 0);
558 unpacked_values.resize(n, 0);
559 unpack_values(&dict, &mut unpacked_values);
560
561 assert_eq!(values, unpacked_values);
562 }
563 }
564 }
565
566 fn gen_map(items_num: usize, values_num: usize) -> HashMap<u64, Vec<u32>> {
567 let mut rng = ChaCha8Rng::seed_from_u64(123);
568
569 (0..items_num)
570 .map(|_| {
571 let key = rng.gen::<u64>();
572 let value = (0..values_num).map(|_| rng.gen_range(1..=10)).collect();
573 (key, value)
574 })
575 .collect()
576 }
577
578 #[test]
579 fn test_map_with_dict_bitpacked() {
580 let items_num = 1000;
581 let values_num = 10;
582 let original_map = gen_map(items_num, values_num);
583
584 let map = MapWithDictBitpacked::try_from(original_map.clone()).unwrap();
585
586 assert_eq!(map.len(), original_map.len());
588
589 assert_eq!(map.is_empty(), original_map.is_empty());
591
592 let mut values_buf = vec![0; values_num];
594 for (key, value) in &original_map {
595 assert!(map.get_values(key, &mut values_buf));
596 assert_eq!(value, &values_buf);
597 assert!(map.contains_key(key));
598 }
599
600 for (&k, v) in map.iter(values_num) {
602 assert_eq!(original_map.get(&k), Some(&v));
603 }
604
605 for k in map.keys() {
607 assert!(original_map.contains_key(k));
608 }
609
610 for v in map.values(values_num) {
612 assert!(original_map.values().any(|val| val == &v));
613 }
614
615 assert_eq!(map.size(), 22664);
617 }
618
619 #[cfg(feature = "rkyv_derive")]
620 #[test]
621 fn test_rkyv() {
622 let items_num = 1000;
624 let values_num = 10;
625 let original_map = gen_map(items_num, values_num);
626 let map = MapWithDictBitpacked::try_from(original_map.clone()).unwrap();
627 let rkyv_bytes = rkyv::to_bytes::<_, 1024>(&map).unwrap();
628
629 let rkyv_map = rkyv::check_archived_root::<MapWithDictBitpacked<u64>>(&rkyv_bytes).unwrap();
630
631 let mut values_buf = vec![0; values_num];
633 for (k, v) in original_map {
634 rkyv_map.get_values(&k, &mut values_buf);
635 assert_eq!(v, values_buf);
636 }
637 }
638
639 #[cfg(feature = "serde")]
640 #[test]
641 fn test_serde() {
642 let items_num = 1000;
644 let values_num = 10;
645 let original_map = gen_map(items_num, values_num);
646 let map = MapWithDictBitpacked::try_from(original_map.clone()).unwrap();
647
648 let bytes = rmp_serde::to_vec(&map).unwrap();
649 let de: MapWithDictBitpacked<u64> = rmp_serde::from_slice(&bytes).unwrap();
650
651 assert_eq!(de.len(), original_map.len());
652
653 let mut values_buf = vec![0; values_num];
655 for (k, v) in &original_map {
656 assert!(de.get_values(k, &mut values_buf));
657 assert_eq!(v, &values_buf);
658 }
659 }
660
661 macro_rules! proptest_map_with_dict_bitpacked_model {
662 ($(($b:expr, $s:expr, $gamma:expr, $n:expr)),* $(,)?) => {
663 $(
664 paste! {
665 proptest! {
666 #[test]
667 fn [<proptest_map_with_dict_bitpacked_model_ $b _ $s _ $n _ $gamma>](model: HashMap<u64, [u32; $n]>, arbitrary: HashSet<u64>) {
668 let entropy_map: MapWithDictBitpacked<u64, $b, $s> = MapWithDictBitpacked::from_iter_with_params(
669 model.iter().map(|(&k, v)| (k, Vec::from(v))),
670 $gamma as f32 / 100.0
671 ).unwrap();
672
673 assert_eq!(entropy_map.len(), model.len());
675 assert_eq!(entropy_map.is_empty(), model.is_empty());
676
677 assert_eq!(
679 HashSet::<_, RandomState>::from_iter(entropy_map.keys()),
680 HashSet::from_iter(model.keys())
681 );
682 assert_eq!(
683 HashSet::<_, RandomState>::from_iter(entropy_map.values($n)),
684 HashSet::from_iter(model.values().map(Vec::from))
685 );
686
687 for (k, v) in &model {
689 assert!(entropy_map.contains_key(&k));
690
691 let mut buf = [0u32; $n];
692 assert!(entropy_map.get_values(&k, &mut buf));
693 assert_eq!(&buf, v);
694 }
695
696 for k in arbitrary {
698 assert_eq!(
699 model.contains_key(&k),
700 entropy_map.contains_key(&k),
701 );
702 let mut buf = [0u32; $n];
703 let contains = entropy_map.get_values(&k, &mut buf);
704 assert_eq!(contains, model.contains_key(&k));
705 if contains {
706 assert_eq!(Some(&buf), model.get(&k));
707 }
708 }
709 }
710 }
711 }
712 )*
713 };
714 }
715
716 proptest_map_with_dict_bitpacked_model!(
717 (2, 8, 100, 10),
719 (4, 8, 100, 10),
720 (7, 8, 100, 10),
721 (8, 8, 100, 10),
722 (15, 8, 100, 10),
723 (16, 8, 100, 10),
724 (23, 8, 100, 10),
725 (24, 8, 100, 10),
726 (31, 8, 100, 10),
727 (32, 8, 100, 10),
728 (33, 8, 100, 10),
729 (48, 8, 100, 10),
730 (53, 8, 100, 10),
731 (61, 8, 100, 10),
732 (63, 8, 100, 10),
733 (64, 8, 100, 10),
734 (32, 7, 100, 10),
735 (32, 5, 100, 10),
736 (32, 4, 100, 10),
737 (32, 3, 100, 10),
738 (32, 1, 100, 10),
739 (32, 0, 100, 10),
740 (32, 8, 200, 10),
741 (32, 6, 200, 10),
742 );
743}