1use std::borrow::Borrow;
14use std::collections::HashSet;
15use std::hash::{Hash, Hasher};
16use std::mem::size_of_val;
17
18use num::{PrimInt, Unsigned};
19use wyhash::WyHash;
20
21use crate::mphf::{Mphf, MphfError, DEFAULT_GAMMA};
22
23#[derive(Default)]
25#[cfg_attr(feature = "rkyv_derive", derive(rkyv::Archive, rkyv::Deserialize, rkyv::Serialize))]
26#[cfg_attr(feature = "rkyv_derive", archive_attr(derive(rkyv::CheckBytes)))]
27#[cfg_attr(feature = "serde", derive(serde::Serialize))]
28#[cfg_attr(
29 feature = "serde",
30 serde(bound(serialize = "K: serde::Serialize, ST: serde::Serialize"))
31)]
32pub struct Set<K, const B: usize = 32, const S: usize = 8, ST = u8, H = WyHash>
33where
34 ST: PrimInt + Unsigned,
35 H: Hasher + Default,
36{
37 pub(crate) mphf: Mphf<B, S, ST, H>,
39 keys: Box<[K]>,
41}
42
43#[cfg(feature = "serde")]
44#[derive(serde::Deserialize)]
45#[serde(bound(deserialize = "K: serde::Deserialize<'de>, ST: serde::Deserialize<'de>"))]
46struct SetUnchecked<K, const B: usize = 32, const S: usize = 8, ST = u8, H = WyHash>
47where
48 ST: PrimInt + Unsigned,
49 H: Hasher + Default,
50{
51 pub(crate) mphf: Mphf<B, S, ST, H>,
52 keys: Box<[K]>,
53}
54
55#[cfg(feature = "serde")]
56impl<'de, K, const B: usize, const S: usize, ST, H> serde::Deserialize<'de> for Set<K, B, S, ST, H>
57where
58 K: serde::Deserialize<'de> + Hash,
59 ST: serde::Deserialize<'de> + PrimInt + Unsigned,
60 H: Hasher + Default,
61{
62 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
63 where
64 D: serde::Deserializer<'de>,
65 {
66 use crate::ValidateKeyResult;
67 use serde::de::Error;
68
69 let this = SetUnchecked::deserialize(deserializer)?;
70
71 this.mphf.validate_keys(&this.keys).map_err(|e| match e {
72 ValidateKeyResult::InvalidKeyCount => {
73 Error::custom("key count should equal the number of set bits in the MPHF")
74 }
75 ValidateKeyResult::IncorrectKeyOrder => Error::custom("keys should correspond to MPHF index"),
76 })?;
77
78 Ok(Self { mphf: this.mphf, keys: this.keys })
79 }
80}
81
82impl<K, const B: usize, const S: usize, ST, H> Set<K, B, S, ST, H>
83where
84 K: Eq + Hash,
85 ST: PrimInt + Unsigned,
86 H: Hasher + Default,
87{
88 pub fn from_iter_with_params<I>(iter: I, gamma: f32) -> Result<Self, MphfError>
98 where
99 I: IntoIterator<Item = K>,
100 {
101 let mut keys: Vec<K> = iter.into_iter().collect();
102
103 let mphf = Mphf::from_slice(&keys, gamma)?;
104
105 for i in 0..keys.len() {
107 loop {
108 let idx: usize = mphf.get(&keys[i]).unwrap();
109 if idx == i {
110 break;
111 }
112 keys.swap(i, idx);
113 }
114 }
115
116 Ok(Set { mphf, keys: keys.into_boxed_slice() })
117 }
118
119 #[inline]
130 pub fn contains<Q>(&self, key: &Q) -> bool
131 where
132 K: Borrow<Q> + PartialEq<Q>,
133 Q: Hash + Eq + ?Sized,
134 {
135 self.mphf
137 .get(key)
138 .map(|idx| unsafe { self.keys.get_unchecked(idx) == key })
139 .unwrap_or_default()
140 }
141
142 #[inline]
152 pub fn len(&self) -> usize {
153 self.keys.len()
154 }
155
156 #[inline]
168 pub fn is_empty(&self) -> bool {
169 self.keys.is_empty()
170 }
171
172 #[inline]
184 pub fn iter(&self) -> impl Iterator<Item = &K> {
185 self.keys.iter()
186 }
187
188 #[inline]
198 pub fn size(&self) -> usize {
199 size_of_val(self) + self.mphf.size() + size_of_val(self.keys.as_ref())
200 }
201}
202
203impl<K> TryFrom<HashSet<K>> for Set<K>
205where
206 K: Eq + Hash,
207{
208 type Error = MphfError;
209
210 #[inline]
211 fn try_from(value: HashSet<K>) -> Result<Self, Self::Error> {
212 Set::from_iter_with_params(value, DEFAULT_GAMMA)
213 }
214}
215
216#[cfg(feature = "rkyv_derive")]
218impl<K, const B: usize, const S: usize, ST, H> ArchivedSet<K, B, S, ST, H>
219where
220 K: Eq + Hash + rkyv::Archive,
221 K::Archived: PartialEq<K>,
222 ST: PrimInt + Unsigned + rkyv::Archive<Archived = ST>,
223 H: Hasher + Default,
224{
225 #[inline]
239 pub fn contains<Q>(&self, key: &Q) -> bool
240 where
241 K: Borrow<Q>,
242 <K as rkyv::Archive>::Archived: PartialEq<Q>,
243 Q: ?Sized + Hash + Eq,
244 {
245 self.mphf
247 .get(key)
248 .map(|idx| unsafe { self.keys.get_unchecked(idx) == key })
249 .unwrap_or_default()
250 }
251}
252
253#[cfg(test)]
254mod tests {
255 use super::*;
256 use paste::paste;
257 use proptest::prelude::*;
258 use rand::{Rng, SeedableRng};
259 use rand_chacha::ChaCha8Rng;
260
261 fn gen_set(items_num: usize) -> HashSet<u64> {
262 let mut rng = ChaCha8Rng::seed_from_u64(123);
263
264 (0..items_num).map(|_| rng.gen::<u64>()).collect()
265 }
266
267 #[test]
268 fn test_set_with_hashset() {
269 let original_set = gen_set(1000);
271
272 let set = Set::try_from(original_set.clone()).unwrap();
274
275 assert_eq!(set.len(), original_set.len());
277
278 assert_eq!(set.is_empty(), original_set.is_empty());
280
281 for key in &original_set {
283 assert!(set.contains(key));
284 }
285
286 for &k in set.iter() {
288 assert!(original_set.contains(&k));
289 }
290
291 assert_eq!(set.size(), 8540);
293 }
294
295 #[test]
297 fn test_contains_borrow() {
298 let set = Set::try_from(HashSet::from(["a".to_string(), "b".to_string()])).unwrap();
299
300 assert!(set.contains("a"));
301 assert!(set.contains("b"));
302 assert!(!set.contains("c"));
303 }
304
305 #[cfg(feature = "rkyv_derive")]
306 #[test]
307 fn test_rkyv() {
308 let original_set = gen_set(1000);
310 let set = Set::try_from(original_set.clone()).unwrap();
311 let rkyv_bytes = rkyv::to_bytes::<_, 1024>(&set).unwrap();
312
313 let rkyv_set = rkyv::check_archived_root::<Set<u64>>(&rkyv_bytes).unwrap();
314
315 for k in original_set.iter() {
317 assert!(rkyv_set.contains(k));
318 }
319 }
320
321 #[cfg(feature = "rkyv_derive")]
322 #[test]
323 fn test_rkyv_contains_borrow() {
324 let set = Set::try_from(HashSet::from(["a".to_string(), "b".to_string()])).unwrap();
325 let rkyv_bytes = rkyv::to_bytes::<_, 1024>(&set).unwrap();
326 let rkyv_set = rkyv::check_archived_root::<Set<String>>(&rkyv_bytes).unwrap();
327
328 assert!(rkyv_set.contains("a"));
329 assert!(rkyv_set.contains("b"));
330 assert!(!rkyv_set.contains("c"));
331 }
332
333 #[cfg(feature = "serde")]
334 #[test]
335 fn test_serde() {
336 let original_set = gen_set(1000);
338 let set = Set::try_from(original_set.clone()).unwrap();
339
340 let bytes = rmp_serde::to_vec(&set).unwrap();
341 let de: Set<u64> = rmp_serde::from_slice(&bytes).unwrap();
342
343 assert_eq!(de.len(), original_set.len());
344
345 for k in original_set.iter() {
347 assert!(de.contains(k));
348 }
349 }
350
351 macro_rules! proptest_set_model {
352 ($(($b:expr, $s:expr, $gamma:expr)),* $(,)?) => {
353 $(
354 paste! {
355 proptest! {
356 #[test]
357 fn [<proptest_set_model_ $b _ $s _ $gamma>](model: HashSet<u64>, arbitrary: HashSet<u64>) {
358 let entropy_set: Set<u64, $b, $s> = Set::from_iter_with_params(
359 model.clone(),
360 $gamma as f32 / 100.0
361 ).unwrap();
362
363 assert_eq!(entropy_set.len(), model.len());
365 assert_eq!(entropy_set.is_empty(), model.is_empty());
366
367 for elm in &model {
369 assert!(entropy_set.contains(&elm));
370 }
371
372 for elm in arbitrary {
374 assert_eq!(
375 model.contains(&elm),
376 entropy_set.contains(&elm),
377 );
378 }
379 }
380 }
381 }
382 )*
383 };
384 }
385
386 proptest_set_model!(
387 (2, 8, 100),
389 (4, 8, 100),
390 (7, 8, 100),
391 (8, 8, 100),
392 (15, 8, 100),
393 (16, 8, 100),
394 (23, 8, 100),
395 (24, 8, 100),
396 (31, 8, 100),
397 (32, 8, 100),
398 (33, 8, 100),
399 (48, 8, 100),
400 (53, 8, 100),
401 (61, 8, 100),
402 (63, 8, 100),
403 (64, 8, 100),
404 (32, 7, 100),
405 (32, 5, 100),
406 (32, 4, 100),
407 (32, 3, 100),
408 (32, 1, 100),
409 (32, 0, 100),
410 (32, 8, 200),
411 (32, 6, 200),
412 );
413
414 proptest! {
415 #[test]
416 fn test_set_contains(model: HashSet<u64>, arbitrary: HashSet<u64>) {
417 let entropy_set = Set::try_from(model.clone()).unwrap();
418
419 for elm in &model {
420 assert!(entropy_set.contains(elm));
421 }
422
423 for elm in arbitrary {
424 assert_eq!(
425 model.contains(&elm),
426 entropy_set.contains(&elm),
427 );
428 }
429 }
430 }
431}