charon_lib/utils/
dedup.rs1use indexmap::IndexMap as SeqHashMap;
7use rustc_hash::FxHashMap;
8use serde::{Deserialize, Serialize};
9use serde_state::{DeserializeState, SerializeState};
10use std::any::type_name;
11use std::cell::RefCell;
12use std::hash::Hash;
13
14use crate::utils::type_map::{Mappable, Mapper, TypeMap};
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
19pub struct DedupId(u32);
20
21pub trait Dedup: Mappable + Clone + Eq + Hash {}
24impl<T> Dedup for T where T: Mappable + Clone + Eq + Hash {}
25
26pub trait DedupSerializerState: Sized {
29 fn record_serialized<T: Dedup>(&self, value: &T) -> Option<Result<DedupId, DedupId>>;
33 fn record_deserialized<T: Dedup>(&self, id: DedupId, value: T);
35 fn get_deserialized<T: Dedup>(&self, id: DedupId) -> Option<T>;
37}
38
39impl DedupSerializerState for () {
41 fn record_serialized<T: Dedup>(&self, _value: &T) -> Option<Result<DedupId, DedupId>> {
42 None
43 }
44 fn record_deserialized<T: Dedup>(&self, _id: DedupId, _value: T) {}
45 fn get_deserialized<T: Dedup>(&self, _id: DedupId) -> Option<T> {
46 None
47 }
48}
49
50struct SerializeTableMapper;
51impl Mapper for SerializeTableMapper {
52 type Value<T: Mappable> = FxHashMap<T, DedupId>;
53}
54struct DeserializeTableMapper;
55impl Mapper for DeserializeTableMapper {
56 type Value<T: Mappable> = SeqHashMap<DedupId, T>;
57}
58
59#[derive(Default)]
61pub struct DedupSerializer {
62 ser: RefCell<TypeMap<SerializeTableMapper>>,
64 de: RefCell<TypeMap<DeserializeTableMapper>>,
66}
67
68impl DedupSerializerState for DedupSerializer {
69 fn record_serialized<T: Dedup>(&self, value: &T) -> Option<Result<DedupId, DedupId>> {
70 let mut ser = self.ser.borrow_mut();
71 let table = ser.or_default::<T>();
72 Some(match table.get(value) {
73 Some(&id) => Err(id),
74 None => {
75 let id = DedupId(table.len().try_into().unwrap());
76 table.insert(value.clone(), id);
77 Ok(id)
78 }
79 })
80 }
81 fn record_deserialized<T: Dedup>(&self, id: DedupId, value: T) {
82 self.de.borrow_mut().or_default::<T>().insert(id, value);
83 }
84 fn get_deserialized<T: Dedup>(&self, id: DedupId) -> Option<T> {
85 self.de
86 .borrow()
87 .get::<T>()
88 .and_then(|table| table.get(&id))
89 .cloned()
90 }
91}
92
93#[derive(Serialize, Deserialize, SerializeState, DeserializeState)]
96#[serde_state(state_implements = DedupSerializerState)]
97pub enum SerDedup<T> {
98 Value(#[serde_state(stateless)] DedupId, T),
101 #[serde_state(stateless)]
104 Deduplicated(DedupId),
105 Untagged(T),
107}
108
109pub fn serialize_dedup<T, R, State, S>(
112 value: &T,
113 repr: R,
114 state: &State,
115 serializer: S,
116) -> Result<S::Ok, S::Error>
117where
118 T: Dedup,
119 R: SerializeState<State>,
120 State: DedupSerializerState,
121 S: serde::Serializer,
122{
123 let repr = match state.record_serialized(value) {
124 Some(Ok(id)) => SerDedup::Value(id, repr),
125 Some(Err(id)) => SerDedup::Deduplicated(id),
126 None => SerDedup::Untagged(repr),
127 };
128 repr.serialize_state(state, serializer)
129}
130
131pub fn deserialize_dedup<'de, T, R, State, D>(
134 state: &State,
135 deserializer: D,
136 build: impl FnOnce(R) -> T,
137) -> Result<T, D::Error>
138where
139 T: Dedup,
140 R: DeserializeState<'de, State>,
141 State: DedupSerializerState,
142 D: serde::Deserializer<'de>,
143{
144 use serde::de::Error;
145 Ok(
146 match SerDedup::<R>::deserialize_state(state, deserializer)? {
147 SerDedup::Value(id, repr) => {
148 let value = build(repr);
149 state.record_deserialized(id, value.clone());
150 value
151 }
152 SerDedup::Deduplicated(id) => state.get_deserialized(id).ok_or_else(|| {
153 let msg = format!(
154 "can't deserialize deduplicated value of type {}; \
155 were you careful with managing the deduplication state?",
156 type_name::<T>()
157 );
158 D::Error::custom(msg)
159 })?,
160 SerDedup::Untagged(repr) => build(repr),
161 },
162 )
163}
164
165pub fn stateless_deserialize_error<T>() -> String {
168 format!(
169 "trying to deserialize a deduplicated value using serde's `{ty}::deserialize` method. \
170 This won't work, use serde_state's \
171 `{ty}::deserialize_state(&DedupSerializer::default(), _)` instead",
172 ty = type_name::<T>(),
173 )
174}