1use std::borrow::Borrow;
25use std::collections::hash_map::RandomState;
26use std::collections::{self, BTreeSet};
27use std::fmt::{Debug, Error, Formatter};
28use std::hash::{BuildHasher, Hash};
29use std::iter::{FromIterator, FusedIterator, Sum};
30use std::ops::{Add, Deref, Mul};
31
32use archery::{SharedPointer, SharedPointerKind};
33use equivalent::Equivalent;
34
35use crate::nodes::hamt::{hash_key, Drain as NodeDrain, HashValue, Iter as NodeIter, Node};
36use crate::ordset::GenericOrdSet;
37use crate::shared_ptr::DefaultSharedPtr;
38use crate::GenericVector;
39
40#[macro_export]
55macro_rules! hashset {
56 () => { $crate::hashset::HashSet::new() };
57
58 ( $($x:expr),* ) => {{
59 let mut l = $crate::hashset::HashSet::new();
60 $(
61 l.insert($x);
62 )*
63 l
64 }};
65
66 ( $($x:expr ,)* ) => {{
67 let mut l = $crate::hashset::HashSet::new();
68 $(
69 l.insert($x);
70 )*
71 l
72 }};
73}
74
75pub type HashSet<A> = GenericHashSet<A, RandomState, DefaultSharedPtr>;
81
82pub struct GenericHashSet<A, S, P: SharedPointerKind> {
101 hasher: S,
102 root: Option<SharedPointer<Node<Value<A>, P>, P>>,
103 size: usize,
104}
105
106#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Debug)]
107struct Value<A>(A);
108
109impl<A> Deref for Value<A> {
110 type Target = A;
111 fn deref(&self) -> &Self::Target {
112 &self.0
113 }
114}
115
116impl<A> HashValue for Value<A>
119where
120 A: Hash + Eq,
121{
122 type Key = A;
123
124 fn extract_key(&self) -> &Self::Key {
125 &self.0
126 }
127
128 fn ptr_eq(&self, _other: &Self) -> bool {
129 false
130 }
131}
132
133impl<A, S, P> GenericHashSet<A, S, P>
134where
135 A: Hash + Eq + Clone,
136 S: BuildHasher + Default + Clone,
137 P: SharedPointerKind,
138{
139 #[inline]
151 #[must_use]
152 pub fn unit(a: A) -> Self {
153 GenericHashSet::new().update(a)
154 }
155}
156
157impl<A, S, P: SharedPointerKind> GenericHashSet<A, S, P> {
158 #[must_use]
160 pub fn new() -> Self
161 where
162 S: Default,
163 {
164 Self::default()
165 }
166
167 #[inline]
184 #[must_use]
185 pub fn is_empty(&self) -> bool {
186 self.len() == 0
187 }
188
189 #[inline]
201 #[must_use]
202 pub fn len(&self) -> usize {
203 self.size
204 }
205
206 pub fn ptr_eq(&self, other: &Self) -> bool {
216 match (&self.root, &other.root) {
217 (Some(a), Some(b)) => SharedPointer::ptr_eq(a, b),
218 (None, None) => true,
219 _ => false,
220 }
221 }
222
223 #[inline]
225 #[must_use]
226 pub fn with_hasher(hasher: S) -> Self {
227 GenericHashSet {
228 size: 0,
229 root: None,
230 hasher,
231 }
232 }
233
234 #[must_use]
238 pub fn hasher(&self) -> &S {
239 &self.hasher
240 }
241
242 #[inline]
244 #[must_use]
245 pub fn new_from<A2>(&self) -> GenericHashSet<A2, S, P>
246 where
247 A2: Hash + Eq + Clone,
248 S: Clone,
249 {
250 GenericHashSet {
251 size: 0,
252 root: None,
253 hasher: self.hasher.clone(),
254 }
255 }
256
257 pub fn clear(&mut self) {
274 self.root = None;
275 self.size = 0;
276 }
277
278 #[must_use]
286 pub fn iter(&self) -> Iter<'_, A, P> {
287 Iter {
288 it: NodeIter::new(self.root.as_deref(), self.size),
289 }
290 }
291}
292
293impl<A, S, P> GenericHashSet<A, S, P>
294where
295 A: Hash + Eq,
296 S: BuildHasher,
297 P: SharedPointerKind,
298{
299 fn test_eq<S2: BuildHasher, P2: SharedPointerKind>(
300 &self,
301 other: &GenericHashSet<A, S2, P2>,
302 ) -> bool {
303 if self.len() != other.len() {
304 return false;
305 }
306 let mut seen = collections::HashSet::new();
307 for value in self.iter() {
308 if !other.contains(value) {
309 return false;
310 }
311 seen.insert(value);
312 }
313 for value in other.iter() {
314 if !seen.contains(&value) {
315 return false;
316 }
317 }
318 true
319 }
320
321 #[must_use]
325 pub fn contains<Q>(&self, value: &Q) -> bool
326 where
327 Q: Hash + Equivalent<A> + ?Sized,
328 {
329 if let Some(root) = &self.root {
330 root.get(hash_key(&self.hasher, value), 0, value).is_some()
331 } else {
332 false
333 }
334 }
335
336 #[must_use]
341 pub fn is_subset<RS>(&self, other: RS) -> bool
342 where
343 RS: Borrow<Self>,
344 {
345 let o = other.borrow();
346 self.iter().all(|a| o.contains(a))
347 }
348
349 #[must_use]
355 pub fn is_proper_subset<RS>(&self, other: RS) -> bool
356 where
357 RS: Borrow<Self>,
358 {
359 self.len() != other.borrow().len() && self.is_subset(other)
360 }
361}
362
363impl<A, S, P> GenericHashSet<A, S, P>
364where
365 A: Hash + Eq + Clone,
366 S: BuildHasher + Clone,
367 P: SharedPointerKind,
368{
369 #[inline]
373 pub fn insert(&mut self, a: A) -> Option<A> {
374 let hash = hash_key(&self.hasher, &a);
375 let root = SharedPointer::make_mut(self.root.get_or_insert_with(Default::default));
376 match root.insert(hash, 0, Value(a)) {
377 None => {
378 self.size += 1;
379 None
380 }
381 Some(Value(old_value)) => Some(old_value),
382 }
383 }
384
385 pub fn remove<Q>(&mut self, value: &Q) -> Option<A>
389 where
390 Q: Hash + Equivalent<A> + ?Sized,
391 {
392 let root = SharedPointer::make_mut(self.root.get_or_insert_with(Default::default));
393 let result = root.remove(hash_key(&self.hasher, value), 0, value);
394 if result.is_some() {
395 self.size -= 1;
396 }
397 result.map(|v| v.0)
398 }
399
400 #[must_use]
418 pub fn update(&self, a: A) -> Self {
419 let mut out = self.clone();
420 out.insert(a);
421 out
422 }
423
424 #[must_use]
429 pub fn without<Q>(&self, value: &Q) -> Self
430 where
431 Q: Hash + Equivalent<A> + ?Sized,
432 {
433 let mut out = self.clone();
434 out.remove(value);
435 out
436 }
437
438 pub fn retain<F>(&mut self, mut f: F)
458 where
459 F: FnMut(&A) -> bool,
460 {
461 let Some(root) = &mut self.root else {
462 return;
463 };
464 let old_root = root.clone();
465 let root = SharedPointer::make_mut(root);
466 for (value, hash) in NodeIter::new(Some(&old_root), self.size) {
467 if !f(value) && root.remove(hash, 0, &**value).is_some() {
468 self.size -= 1;
469 }
470 }
471 }
472
473 #[must_use]
488 pub fn union(self, other: Self) -> Self {
489 let (mut to_mutate, to_consume) = if self.len() >= other.len() {
490 (self, other)
491 } else {
492 (other, self)
493 };
494 for value in to_consume {
495 to_mutate.insert(value);
496 }
497 to_mutate
498 }
499
500 #[must_use]
504 pub fn unions<I>(i: I) -> Self
505 where
506 I: IntoIterator<Item = Self>,
507 S: Default,
508 {
509 i.into_iter().fold(Self::default(), Self::union)
510 }
511
512 #[deprecated(
532 since = "2.0.1",
533 note = "to avoid conflicting behaviors between std and imbl, the `difference` alias for `symmetric_difference` will be removed."
534 )]
535 #[must_use]
536 pub fn difference(self, other: Self) -> Self {
537 self.symmetric_difference(other)
538 }
539
540 #[must_use]
555 pub fn symmetric_difference(mut self, other: Self) -> Self {
556 for value in other {
557 if self.remove(&value).is_none() {
558 self.insert(value);
559 }
560 }
561 self
562 }
563
564 #[must_use]
580 pub fn relative_complement(mut self, other: Self) -> Self {
581 for value in other {
582 let _ = self.remove(&value);
583 }
584 self
585 }
586
587 #[must_use]
602 pub fn intersection(self, other: Self) -> Self {
603 let mut out = self.new_from();
604 for value in other {
605 if self.contains(&value) {
606 out.insert(value);
607 }
608 }
609 out
610 }
611}
612
613impl<A, S, P: SharedPointerKind> Clone for GenericHashSet<A, S, P>
616where
617 A: Clone,
618 S: Clone,
619 P: SharedPointerKind,
620{
621 #[inline]
625 fn clone(&self) -> Self {
626 GenericHashSet {
627 hasher: self.hasher.clone(),
628 root: self.root.clone(),
629 size: self.size,
630 }
631 }
632}
633
634impl<A, S1, P1, S2, P2> PartialEq<GenericHashSet<A, S2, P2>> for GenericHashSet<A, S1, P1>
635where
636 A: Hash + Eq,
637 S1: BuildHasher,
638 S2: BuildHasher,
639 P1: SharedPointerKind,
640 P2: SharedPointerKind,
641{
642 fn eq(&self, other: &GenericHashSet<A, S2, P2>) -> bool {
643 self.test_eq(other)
644 }
645}
646
647impl<A, S, P> Eq for GenericHashSet<A, S, P>
648where
649 A: Hash + Eq,
650 S: BuildHasher,
651 P: SharedPointerKind,
652{
653}
654
655impl<A, S, P> Default for GenericHashSet<A, S, P>
656where
657 S: Default,
658 P: SharedPointerKind,
659{
660 fn default() -> Self {
661 GenericHashSet {
662 hasher: Default::default(),
663 root: None,
664 size: 0,
665 }
666 }
667}
668
669impl<A, S, P> Add for GenericHashSet<A, S, P>
670where
671 A: Hash + Eq + Clone,
672 S: BuildHasher + Clone,
673 P: SharedPointerKind,
674{
675 type Output = GenericHashSet<A, S, P>;
676
677 fn add(self, other: Self) -> Self::Output {
678 self.union(other)
679 }
680}
681
682impl<A, S, P> Mul for GenericHashSet<A, S, P>
683where
684 A: Hash + Eq + Clone,
685 S: BuildHasher + Clone,
686 P: SharedPointerKind,
687{
688 type Output = GenericHashSet<A, S, P>;
689
690 fn mul(self, other: Self) -> Self::Output {
691 self.intersection(other)
692 }
693}
694
695impl<A, S, P> Add for &GenericHashSet<A, S, P>
696where
697 A: Hash + Eq + Clone,
698 S: BuildHasher + Clone,
699 P: SharedPointerKind,
700{
701 type Output = GenericHashSet<A, S, P>;
702
703 fn add(self, other: Self) -> Self::Output {
704 self.clone().union(other.clone())
705 }
706}
707
708impl<A, S, P> Mul for &GenericHashSet<A, S, P>
709where
710 A: Hash + Eq + Clone,
711 S: BuildHasher + Clone,
712 P: SharedPointerKind,
713{
714 type Output = GenericHashSet<A, S, P>;
715
716 fn mul(self, other: Self) -> Self::Output {
717 self.clone().intersection(other.clone())
718 }
719}
720
721impl<A, S, P: SharedPointerKind> Sum for GenericHashSet<A, S, P>
722where
723 A: Hash + Eq + Clone,
724 S: BuildHasher + Default + Clone,
725 P: SharedPointerKind,
726{
727 fn sum<I>(it: I) -> Self
728 where
729 I: Iterator<Item = Self>,
730 {
731 it.fold(Self::default(), |a, b| a + b)
732 }
733}
734
735impl<A, S, R, P: SharedPointerKind> Extend<R> for GenericHashSet<A, S, P>
736where
737 A: Hash + Eq + Clone + From<R>,
738 S: BuildHasher + Clone,
739{
740 fn extend<I>(&mut self, iter: I)
741 where
742 I: IntoIterator<Item = R>,
743 {
744 for value in iter {
745 self.insert(From::from(value));
746 }
747 }
748}
749
750impl<A, S, P> Debug for GenericHashSet<A, S, P>
751where
752 A: Hash + Eq + Debug,
753 S: BuildHasher,
754 P: SharedPointerKind,
755{
756 fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), Error> {
757 f.debug_set().entries(self.iter()).finish()
758 }
759}
760
761pub struct Iter<'a, A, P: SharedPointerKind> {
765 it: NodeIter<'a, Value<A>, P>,
766}
767
768impl<'a, A, P: SharedPointerKind> Clone for Iter<'a, A, P> {
770 fn clone(&self) -> Self {
771 Iter {
772 it: self.it.clone(),
773 }
774 }
775}
776
777impl<'a, A, P> Iterator for Iter<'a, A, P>
778where
779 A: 'a,
780 P: SharedPointerKind,
781{
782 type Item = &'a A;
783
784 fn next(&mut self) -> Option<Self::Item> {
785 self.it.next().map(|(v, _)| &v.0)
786 }
787
788 fn size_hint(&self) -> (usize, Option<usize>) {
789 self.it.size_hint()
790 }
791}
792
793impl<'a, A, P: SharedPointerKind> ExactSizeIterator for Iter<'a, A, P> {}
794
795impl<'a, A, P: SharedPointerKind> FusedIterator for Iter<'a, A, P> {}
796
797pub struct ConsumingIter<A, P>
799where
800 A: Hash + Eq + Clone,
801 P: SharedPointerKind,
802{
803 it: NodeDrain<Value<A>, P>,
804}
805
806impl<A, P> Clone for ConsumingIter<A, P>
807where
808 A: Hash + Eq + Clone,
809 P: SharedPointerKind,
810{
811 fn clone(&self) -> Self {
812 Self {
813 it: self.it.clone(),
814 }
815 }
816}
817
818impl<A, P> Iterator for ConsumingIter<A, P>
819where
820 A: Hash + Eq + Clone,
821 P: SharedPointerKind,
822{
823 type Item = A;
824
825 fn next(&mut self) -> Option<Self::Item> {
826 self.it.next().map(|(v, _)| v.0)
827 }
828
829 fn size_hint(&self) -> (usize, Option<usize>) {
830 self.it.size_hint()
831 }
832}
833
834impl<A, P> ExactSizeIterator for ConsumingIter<A, P>
835where
836 A: Hash + Eq + Clone,
837 P: SharedPointerKind,
838{
839}
840
841impl<A, P> FusedIterator for ConsumingIter<A, P>
842where
843 A: Hash + Eq + Clone,
844 P: SharedPointerKind,
845{
846}
847
848impl<A, RA, S, P> FromIterator<RA> for GenericHashSet<A, S, P>
851where
852 A: Hash + Eq + Clone + From<RA>,
853 S: BuildHasher + Default + Clone,
854 P: SharedPointerKind,
855{
856 fn from_iter<T>(i: T) -> Self
857 where
858 T: IntoIterator<Item = RA>,
859 {
860 let mut set = Self::default();
861 for value in i {
862 set.insert(From::from(value));
863 }
864 set
865 }
866}
867
868impl<'a, A, S, P> IntoIterator for &'a GenericHashSet<A, S, P>
869where
870 A: Hash + Eq,
871 S: BuildHasher,
872 P: SharedPointerKind,
873{
874 type Item = &'a A;
875 type IntoIter = Iter<'a, A, P>;
876
877 fn into_iter(self) -> Self::IntoIter {
878 self.iter()
879 }
880}
881
882impl<A, S, P> IntoIterator for GenericHashSet<A, S, P>
883where
884 A: Hash + Eq + Clone,
885 S: BuildHasher,
886 P: SharedPointerKind,
887{
888 type Item = A;
889 type IntoIter = ConsumingIter<Self::Item, P>;
890
891 fn into_iter(self) -> Self::IntoIter {
892 ConsumingIter {
893 it: NodeDrain::new(self.root, self.size),
894 }
895 }
896}
897
898impl<A, OA, SA, SB, P1, P2> From<&GenericHashSet<&A, SA, P1>> for GenericHashSet<OA, SB, P2>
901where
902 A: ToOwned<Owned = OA> + Hash + Equivalent<A> + ?Sized,
903 OA: Hash + Eq + Clone,
904 SA: BuildHasher,
905 SB: BuildHasher + Default + Clone,
906 P1: SharedPointerKind,
907 P2: SharedPointerKind,
908{
909 fn from(set: &GenericHashSet<&A, SA, P1>) -> Self {
910 set.iter().map(|a| (*a).to_owned()).collect()
911 }
912}
913
914impl<A, S, const N: usize, P> From<[A; N]> for GenericHashSet<A, S, P>
915where
916 A: Hash + Eq + Clone,
917 S: BuildHasher + Default + Clone,
918 P: SharedPointerKind,
919{
920 fn from(arr: [A; N]) -> Self {
921 IntoIterator::into_iter(arr).collect()
922 }
923}
924
925impl<'a, A, S, P> From<&'a [A]> for GenericHashSet<A, S, P>
926where
927 A: Hash + Eq + Clone,
928 S: BuildHasher + Default + Clone,
929 P: SharedPointerKind,
930{
931 fn from(slice: &'a [A]) -> Self {
932 slice.iter().cloned().collect()
933 }
934}
935
936impl<A, S, P> From<Vec<A>> for GenericHashSet<A, S, P>
937where
938 A: Hash + Eq + Clone,
939 S: BuildHasher + Default + Clone,
940 P: SharedPointerKind,
941{
942 fn from(vec: Vec<A>) -> Self {
943 vec.into_iter().collect()
944 }
945}
946
947impl<A, S, P> From<&Vec<A>> for GenericHashSet<A, S, P>
948where
949 A: Hash + Eq + Clone,
950 S: BuildHasher + Default + Clone,
951 P: SharedPointerKind,
952{
953 fn from(vec: &Vec<A>) -> Self {
954 vec.iter().cloned().collect()
955 }
956}
957
958impl<A, S, P1, P2> From<GenericVector<A, P2>> for GenericHashSet<A, S, P1>
959where
960 A: Hash + Eq + Clone,
961 S: BuildHasher + Default + Clone,
962 P1: SharedPointerKind,
963 P2: SharedPointerKind,
964{
965 fn from(vector: GenericVector<A, P2>) -> Self {
966 vector.into_iter().collect()
967 }
968}
969
970impl<A, S, P1, P2> From<&GenericVector<A, P2>> for GenericHashSet<A, S, P1>
971where
972 A: Hash + Eq + Clone,
973 S: BuildHasher + Default + Clone,
974 P1: SharedPointerKind,
975 P2: SharedPointerKind,
976{
977 fn from(vector: &GenericVector<A, P2>) -> Self {
978 vector.iter().cloned().collect()
979 }
980}
981
982impl<A, S, P> From<collections::HashSet<A>> for GenericHashSet<A, S, P>
983where
984 A: Eq + Hash + Clone,
985 S: BuildHasher + Default + Clone,
986 P: SharedPointerKind,
987{
988 fn from(hash_set: collections::HashSet<A>) -> Self {
989 hash_set.into_iter().collect()
990 }
991}
992
993impl<A, S, P> From<&collections::HashSet<A>> for GenericHashSet<A, S, P>
994where
995 A: Eq + Hash + Clone,
996 S: BuildHasher + Default + Clone,
997 P: SharedPointerKind,
998{
999 fn from(hash_set: &collections::HashSet<A>) -> Self {
1000 hash_set.iter().cloned().collect()
1001 }
1002}
1003
1004impl<A, S, P> From<&BTreeSet<A>> for GenericHashSet<A, S, P>
1005where
1006 A: Hash + Eq + Clone,
1007 S: BuildHasher + Default + Clone,
1008 P: SharedPointerKind,
1009{
1010 fn from(btree_set: &BTreeSet<A>) -> Self {
1011 btree_set.iter().cloned().collect()
1012 }
1013}
1014
1015impl<A, S, P1, P2> From<GenericOrdSet<A, P2>> for GenericHashSet<A, S, P1>
1016where
1017 A: Ord + Hash + Eq + Clone,
1018 S: BuildHasher + Default + Clone,
1019 P1: SharedPointerKind,
1020 P2: SharedPointerKind,
1021{
1022 fn from(ordset: GenericOrdSet<A, P2>) -> Self {
1023 ordset.into_iter().collect()
1024 }
1025}
1026
1027impl<A, S, P1, P2> From<&GenericOrdSet<A, P2>> for GenericHashSet<A, S, P1>
1028where
1029 A: Ord + Hash + Eq + Clone,
1030 S: BuildHasher + Default + Clone,
1031 P1: SharedPointerKind,
1032 P2: SharedPointerKind,
1033{
1034 fn from(ordset: &GenericOrdSet<A, P2>) -> Self {
1035 ordset.into_iter().cloned().collect()
1036 }
1037}
1038
1039#[cfg(any(test, feature = "proptest"))]
1041#[doc(hidden)]
1042pub mod proptest {
1043 #[deprecated(
1044 since = "14.3.0",
1045 note = "proptest strategies have moved to imbl::proptest"
1046 )]
1047 pub use crate::proptest::hash_set;
1048}
1049
1050#[cfg(test)]
1051mod test {
1052 use super::proptest::*;
1053 use super::*;
1054 use crate::test::LolHasher;
1055 use ::proptest::num::i16;
1056 use ::proptest::proptest;
1057 use static_assertions::{assert_impl_all, assert_not_impl_any};
1058 use std::hash::BuildHasherDefault;
1059
1060 assert_impl_all!(HashSet<i32>: Send, Sync);
1061 assert_not_impl_any!(HashSet<*const i32>: Send, Sync);
1062 assert_covariant!(HashSet<T> in T);
1063
1064 #[test]
1065 fn insert_failing() {
1066 let mut set: GenericHashSet<i16, BuildHasherDefault<LolHasher>, DefaultSharedPtr> =
1067 Default::default();
1068 set.insert(14658);
1069 assert_eq!(1, set.len());
1070 set.insert(-19198);
1071 assert_eq!(2, set.len());
1072 }
1073
1074 #[test]
1075 fn match_strings_with_string_slices() {
1076 let mut set: HashSet<String> = From::from(&hashset!["foo", "bar"]);
1077 set = set.without("bar");
1078 assert!(!set.contains("bar"));
1079 set.remove("foo");
1080 assert!(!set.contains("foo"));
1081 }
1082
1083 #[test]
1084 fn macro_allows_trailing_comma() {
1085 let set1 = hashset! {"foo", "bar"};
1086 let set2 = hashset! {
1087 "foo",
1088 "bar",
1089 };
1090 assert_eq!(set1, set2);
1091 }
1092
1093 #[test]
1094 fn issue_60_drain_iterator_memory_corruption() {
1095 use crate::test::MetroHashBuilder;
1096 for i in 0..1000 {
1097 let mut lhs = vec![0, 1, 2];
1098 lhs.sort_unstable();
1099
1100 let hasher = MetroHashBuilder::new(i);
1101 let mut iset: GenericHashSet<_, MetroHashBuilder, DefaultSharedPtr> =
1102 GenericHashSet::with_hasher(hasher);
1103 for &i in &lhs {
1104 iset.insert(i);
1105 }
1106
1107 let mut rhs: Vec<_> = iset.clone().into_iter().collect();
1108 rhs.sort_unstable();
1109
1110 if lhs != rhs {
1111 println!("iteration: {}", i);
1112 println!("seed: {}", hasher.seed());
1113 println!("lhs: {}: {:?}", lhs.len(), &lhs);
1114 println!("rhs: {}: {:?}", rhs.len(), &rhs);
1115 panic!();
1116 }
1117 }
1118 }
1119
1120 proptest! {
1121 #[test]
1122 fn proptest_a_set(ref s in hash_set(".*", 10..100)) {
1123 assert!(s.len() < 100);
1124 assert!(s.len() >= 10);
1125 }
1126 }
1127}