Skip to main content

charon_lib/utils/
hash_cons.rs

1use derive_generic_visitor::{Drive, DriveMut, DriveTwo, Visit, VisitMut, VisitTwo};
2use std::hash::Hash;
3use std::ops::{ControlFlow, Deref, DerefMut};
4use std::sync::Arc;
5
6use crate::utils::hash_by_addr::HashByAddr;
7use crate::utils::type_map::Mappable;
8
9/// Hash-consed data structure: a reference-counted wrapper that guarantees that two equal
10/// value will be stored at the same address. This makes it possible to use the pointer address
11/// as a hash value.
12// Warning: a `derive` should not introduce a way to create a new `HashConsed` value without
13// going through the interning table.
14#[derive(PartialEq, Eq, Hash)]
15pub struct HashConsed<T>(HashByAddr<Arc<T>>);
16
17impl<T> Clone for HashConsed<T> {
18    fn clone(&self) -> Self {
19        Self(self.0.clone())
20    }
21}
22
23impl<T> HashConsed<T> {
24    pub fn inner(&self) -> &T {
25        self.0.0.as_ref()
26    }
27}
28
29impl<T: PartialOrd> PartialOrd for HashConsed<T> {
30    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
31        self.inner().partial_cmp(other.inner())
32    }
33}
34
35impl<T: Ord> Ord for HashConsed<T> {
36    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
37        self.inner().cmp(other.inner())
38    }
39}
40
41pub trait HashConsable: Hash + PartialEq + Eq + Clone + Mappable {}
42impl<T> HashConsable for T where T: Hash + PartialEq + Eq + Clone + Mappable {}
43
44// Private module that contains the static we'll use as interning map. A value of type
45// `HashCons` MUST NOT be created in any other way than this table, else hashing and euqality
46// on it will be broken. Note that this likely means that if a crate uses charon both as a
47// direct dependency and as a dylib, then the static will be duplicated, causing hashing and
48// equality on `HashCons` to be broken.
49mod intern_table {
50    use rustc_hash::FxBuildHasher;
51    use std::borrow::Borrow;
52    use std::mem::ManuallyDrop;
53    use std::ops::DerefMut;
54    use std::sync::{Arc, LazyLock, RwLock};
55
56    use super::{HashConsable, HashConsed};
57    use crate::utils::hash_by_addr::HashByAddr;
58    use crate::utils::type_map::{Mappable, Mapper, TypeMap};
59
60    type SeqHashSet<T> = indexmap::IndexSet<T, FxBuildHasher>;
61
62    // This is a static mutable `SeqHashSet<Arc<T>>` that records for each `T` value a unique
63    // `Arc<T>` that contains the same value. Values inside the set are hashed/compared
64    // as is normal for `T`.
65    // Once we've gotten an `Arc` out of the set however, we're sure that "T-equality"
66    // implies address-equality, hence the `HashByAddr` wrapper preserves correct equality
67    // and hashing behavior.
68    struct InternMapper;
69    impl Mapper for InternMapper {
70        type Value<T: Mappable> = SeqHashSet<Arc<T>>;
71    }
72    static INTERNED: LazyLock<RwLock<TypeMap<InternMapper>>> = LazyLock::new(Default::default);
73
74    // The excessive generality is to make it work for both `U = T` and `U = Arc<T>`.
75    pub(super) fn intern<T: HashConsable, U>(inner: U) -> HashConsed<T>
76    where
77        Arc<T>: Borrow<U>,
78        U: Into<Arc<T>> + std::hash::Hash,
79        U: indexmap::Equivalent<Arc<T>>,
80    {
81        // Fast read-only check.
82        let arc = if let read_guard = INTERNED.read().unwrap()
83            && let Some(set) = read_guard.get::<T>()
84            && let Some(arc) = set.get(&inner)
85        {
86            arc.clone()
87        } else {
88            // Concurrent access is possible right here, so we have to check everything again.
89            let mut write_guard = INTERNED.write().unwrap();
90            let set: &mut SeqHashSet<Arc<T>> = write_guard.or_default::<T>();
91            if let Some(arc) = set.get(&inner) {
92                arc.clone()
93            } else {
94                let arc: Arc<T> = inner.into();
95                set.insert(arc.clone());
96                arc
97            }
98        };
99        HashConsed(HashByAddr(arc))
100    }
101
102    /// The returned value must not be leaked as this would break the hash-consing invariant.
103    pub(super) fn make_mutable<T: HashConsable>(
104        x: &mut HashConsed<T>,
105    ) -> impl DerefMut<Target = T> {
106        /// A reference to a `HashConsed` for the purposes of mutating the contained value. Avoids
107        /// clones when possible.
108        ///
109        /// This value must not be leaked as that would invalidate the hash-consing invariant.
110        pub enum HashConsedMutRef<'a, T: HashConsable> {
111            /// The contained arc is known to have strong_count 1, so can be mutated directly.
112            /// The hash-consing invariant is broken: the table does not know about this value and we
113            /// must restore the invariant at the end.
114            Unique(&'a mut HashConsed<T>),
115            /// The value was shared, so we simply made a clone.
116            NotUnique(&'a mut HashConsed<T>, ManuallyDrop<T>),
117        }
118
119        impl<'a, T: HashConsable> HashConsedMutRef<'a, T> {
120            pub fn new(x: &'a mut HashConsed<T>) -> Self {
121                let arc = &mut x.0.0;
122                // Every value has at least two pointers: the current value and the one stored in the
123                // global map. If there are exactly two, we may mutate directly by discarding the one in
124                // the global map temporarily.
125                if Arc::strong_count(arc) != 2 {
126                    return Self::new_not_unique(x);
127                }
128                {
129                    // Take the write guard just long enough to drop the other `Arc` to this value.
130                    let mut write_guard = INTERNED.write().unwrap();
131                    // Check the count again, it could have changed concurrently.
132                    if Arc::strong_count(arc) != 2 {
133                        return Self::new_not_unique(x);
134                    }
135                    if let Some(other_arc) = write_guard.or_default::<T>().swap_take(&*arc) {
136                        drop(other_arc);
137                    } else {
138                        // Nothing was removed, early return.
139                        return Self::new_not_unique(x);
140                    }
141                    // The Arc was removed from the map; `x` is invalid as interning the same value would
142                    // result in a different pointer. NO MORE EARLY RETURN until we fix that.
143                }
144                // If we are still the sole owner, we can now mutate in-place.
145                if Arc::get_mut(arc).is_some() {
146                    Self::Unique(x)
147                } else {
148                    Self::new_not_unique(x)
149                }
150            }
151            fn new_not_unique(x: &'a mut HashConsed<T>) -> Self {
152                Self::NotUnique(x, ManuallyDrop::new(x.inner().clone()))
153            }
154        }
155
156        impl<'a, T: HashConsable> std::ops::Deref for HashConsedMutRef<'a, T> {
157            type Target = T;
158            fn deref(&self) -> &Self::Target {
159                match self {
160                    HashConsedMutRef::Unique(x) => x,
161                    HashConsedMutRef::NotUnique(_, val) => val,
162                }
163            }
164        }
165        impl<'a, T: HashConsable> std::ops::DerefMut for HashConsedMutRef<'a, T> {
166            fn deref_mut(&mut self) -> &mut Self::Target {
167                match self {
168                    HashConsedMutRef::Unique(x) => Arc::get_mut(&mut x.0.0).unwrap(),
169                    HashConsedMutRef::NotUnique(_, val) => val,
170                }
171            }
172        }
173
174        impl<'a, T: HashConsable> Drop for HashConsedMutRef<'a, T> {
175            fn drop(&mut self) {
176                match self {
177                    HashConsedMutRef::Unique(x) => {
178                        // Re-establish the interning invariant. If the same value was added to the map in the
179                        // meantime, we'll get a pointer to that.
180                        **x = HashConsed::from_arc(x.0.0.clone());
181                    }
182                    HashConsedMutRef::NotUnique(x, new_val) => {
183                        // SAFETY: we won't touch it again.
184                        let new_val = unsafe { ManuallyDrop::take(new_val) };
185                        // Re-intern the new value if it changed.
186                        if new_val != *x.inner() {
187                            **x = HashConsed::new(new_val);
188                        }
189                    }
190                }
191            }
192        }
193        HashConsedMutRef::new(x)
194    }
195}
196
197impl<T> HashConsed<T>
198where
199    T: HashConsable,
200{
201    /// Deduplicate the values by hashing them. This deduplication is crucial for the hashing
202    /// function to be correct. This is the only function allowed to create `Self` values.
203    pub fn new(inner: T) -> Self {
204        intern_table::intern(inner)
205    }
206    /// Rarely used: in case we already have an `Arc`, may avoid an allocation.
207    pub fn from_arc(inner: Arc<T>) -> Self {
208        intern_table::intern(inner)
209    }
210
211    /// Get a reference to the pointed-to value that can be mutated. Avoids cloning/allocation if
212    /// this is the sole pointer to that value.
213    ///
214    /// The returned value must not be leaked as this would break the hash-consing invariant.
215    pub fn as_mut(&mut self) -> impl DerefMut<Target = T> {
216        intern_table::make_mutable(self)
217    }
218    /// Clones if needed to get mutable access to the inner value.
219    pub fn with_inner_mut<R>(&mut self, f: impl FnOnce(&mut T) -> R) -> R {
220        f(&mut self.as_mut())
221    }
222}
223
224impl<T> Deref for HashConsed<T> {
225    type Target = T;
226    fn deref(&self) -> &Self::Target {
227        self.inner()
228    }
229}
230
231impl<T: std::fmt::Debug> std::fmt::Debug for HashConsed<T> {
232    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
233        // Hide the `HashByAddr` wrapper.
234        f.debug_tuple("HashConsed").field(self.inner()).finish()
235    }
236}
237
238impl<'s, T, V: Visit<'s, T>> Drive<'s, V> for HashConsed<T> {
239    fn drive_inner(&'s self, v: &mut V) -> ControlFlow<V::Break> {
240        v.visit(self.inner())
241    }
242}
243impl<'s, T, V: VisitTwo<'s, T>> DriveTwo<'s, V> for HashConsed<T> {
244    fn drive_two_inner(&'s self, other: &'s Self, v: &mut V) -> ControlFlow<V::Break> {
245        v.visit(self.inner(), other.inner())
246    }
247}
248/// Note: this explores the inner value mutably by cloning and re-hashing afterwards.
249impl<'s, T, V> DriveMut<'s, V> for HashConsed<T>
250where
251    T: HashConsable,
252    V: for<'a> VisitMut<'a, T>,
253{
254    fn drive_inner_mut(&'s mut self, v: &mut V) -> ControlFlow<V::Break> {
255        self.with_inner_mut(|inner| v.visit(inner))
256    }
257}
258
259/// `HashCons` values are deduplicated in the serialized output: see [`crate::utils::dedup`].
260mod serialize {
261    use serde::{Deserialize, Serialize};
262    use serde_state::{DeserializeState, SerializeState};
263
264    use super::{HashConsable, HashConsed};
265    use crate::utils::dedup::*;
266
267    impl<T> Serialize for HashConsed<T>
268    where
269        T: Serialize + HashConsable,
270    {
271        fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
272        where
273            S: serde::Serializer,
274        {
275            SerDedup::Untagged(self.inner()).serialize(serializer)
276        }
277    }
278    /// Options for the state are `()` to serialize values normally and `DedupSerializer`
279    /// to deduplicate identical values in the serialized output.
280    impl<T, State> SerializeState<State> for HashConsed<T>
281    where
282        T: SerializeState<State> + HashConsable,
283        State: DedupSerializerState,
284    {
285        fn serialize_state<S>(&self, state: &State, serializer: S) -> Result<S::Ok, S::Error>
286        where
287            S: serde::Serializer,
288        {
289            serialize_dedup(self, self.inner(), state, serializer)
290        }
291    }
292
293    impl<'de, T> Deserialize<'de> for HashConsed<T>
294    where
295        T: Deserialize<'de> + HashConsable,
296    {
297        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
298        where
299            D: serde::Deserializer<'de>,
300        {
301            use serde::de::Error;
302            let repr: SerDedup<T> = SerDedup::deserialize(deserializer)?;
303            match repr {
304                SerDedup::Value { .. } | SerDedup::Deduplicated { .. } => {
305                    Err(D::Error::custom(stateless_deserialize_error::<T>()))
306                }
307                SerDedup::Untagged(val) => Ok(HashConsed::new(val)),
308            }
309        }
310    }
311    impl<'de, T, State> DeserializeState<'de, State> for HashConsed<T>
312    where
313        T: DeserializeState<'de, State> + HashConsable,
314        State: DedupSerializerState,
315    {
316        fn deserialize_state<D>(state: &State, deserializer: D) -> Result<Self, D::Error>
317        where
318            D: serde::Deserializer<'de>,
319        {
320            deserialize_dedup(state, deserializer, HashConsed::new)
321        }
322    }
323}
324
325#[test]
326fn test_hash_cons() {
327    let x = HashConsed::new(42u32);
328    let y = HashConsed::new(42u32);
329    assert_eq!(x, y);
330    // Test a serialization round-trip.
331    let z = serde_json::from_value(serde_json::to_value(x.clone()).unwrap()).unwrap();
332    assert_eq!(x, z);
333}
334
335#[test]
336fn test_hash_cons_concurrent() {
337    use itertools::Itertools;
338    let handles = (0..10)
339        .map(|_| std::thread::spawn(|| std::hint::black_box(HashConsed::new(42u32))))
340        .collect_vec();
341    let values = handles.into_iter().map(|h| h.join().unwrap()).collect_vec();
342    assert!(values.iter().all_equal())
343}
344
345#[test]
346fn test_hash_cons_dedup() {
347    use crate::utils::dedup::DedupSerializer;
348    use serde_state::{DeserializeState, SerializeState};
349    type Ty = HashConsed<TyKind>;
350    #[derive(Debug, Clone, PartialEq, Eq, Hash)]
351    #[derive(SerializeState, DeserializeState)]
352    #[serde_state(state = DedupSerializer)]
353    enum TyKind {
354        Bool,
355        Pair(Ty, Ty),
356    }
357
358    // Build a value with some redundancy.
359    let bool1 = HashConsed::new(TyKind::Bool);
360    let bool2 = HashConsed::new(TyKind::Bool);
361    let pair = HashConsed::new(TyKind::Pair(bool1.clone(), bool2));
362    let triple = HashConsed::new(TyKind::Pair(bool1, pair));
363
364    let state = DedupSerializer::default();
365    let json_val = triple
366        .serialize_state(&state, serde_json::value::Serializer)
367        .unwrap();
368    let state = DedupSerializer::default();
369    let round_tripped = Ty::deserialize_state(&state, json_val).unwrap();
370
371    assert_eq!(triple, round_tripped);
372}