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