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#[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
44mod 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 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 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 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 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 pub(super) fn make_mutable<T: HashConsable>(
104 x: &mut HashConsed<T>,
105 ) -> impl DerefMut<Target = T> {
106 pub enum HashConsedMutRef<'a, T: HashConsable> {
111 Unique(&'a mut HashConsed<T>),
115 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 if Arc::strong_count(arc) != 2 {
126 return Self::new_not_unique(x);
127 }
128 {
129 let mut write_guard = INTERNED.write().unwrap();
131 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 return Self::new_not_unique(x);
140 }
141 }
144 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 **x = HashConsed::from_arc(x.0.0.clone());
181 }
182 HashConsedMutRef::NotUnique(x, new_val) => {
183 let new_val = unsafe { ManuallyDrop::take(new_val) };
185 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 pub fn new(inner: T) -> Self {
204 intern_table::intern(inner)
205 }
206 pub fn from_arc(inner: Arc<T>) -> Self {
208 intern_table::intern(inner)
209 }
210
211 pub fn as_mut(&mut self) -> impl DerefMut<Target = T> {
216 intern_table::make_mutable(self)
217 }
218 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 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}
248impl<'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
259mod 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 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 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 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}