Skip to main content

charon_lib/utils/
dedup.rs

1//! Deduplication of repeated values in the serialized output.
2//!
3//! Note that the deduplication scheme is order-dependent: it relies on the fact that
4//! serialization and deserialization traverse the value in the same order.
5
6use 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/// Identifies a deduplicated value amongst the values of its type within a single serialized
17/// output. Ids are allocated in the order in which we serialize the values.
18#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
19#[derive(Serialize, Deserialize)]
20pub struct DedupId(u32);
21
22/// A value that we deduplicate in the serialized output. We identify values by equality, hence
23/// the bounds.
24pub trait Dedup: Mappable + Clone + Eq + Hash {}
25impl<T> Dedup for T where T: Mappable + Clone + Eq + Hash {}
26
27/// The state threaded through (de)serialization to deduplicate values. Use `()` to serialize
28/// values normally and [`DedupSerializer`] to deduplicate them.
29pub trait DedupSerializerState: Sized {
30    /// Record that we're serializing this value. Returns `None` if we're not deduplicating
31    /// values, `Some(Ok(id))` the first time we meet a given value (it must then be serialized
32    /// in full), and `Some(Err(id))` afterwards (only the id must be serialized).
33    fn record_serialized<T: Dedup>(&self, value: &T) -> Option<Result<DedupId, DedupId>>;
34    /// Record that we deserialized the value with this id.
35    fn record_deserialized<T: Dedup>(&self, id: DedupId, value: T);
36    /// Find the previously-deserialized value with that id.
37    fn get_deserialized<T: Dedup>(&self, id: DedupId) -> Option<T>;
38}
39
40/// Don't deduplicate anything.
41impl 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/// Deduplicate the values of each type, in one table per type.
61#[derive(Default)]
62pub struct DedupSerializer {
63    // Table used for serialization: the values we've already emitted, with the id we gave them.
64    ser: RefCell<TypeMap<SerializeTableMapper>>,
65    // Table used for deserialization: the values we've read so far, by id.
66    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/// How we represent a deduplicated value in the serialized output. `T` is the serialized form of
95/// the value.
96#[derive(Serialize, Deserialize, SerializeState, DeserializeState)]
97#[serde_state(state_implements = DedupSerializerState)]
98pub enum SerDedup<T> {
99    /// A value represented normally, accompanied by its id. This is emitted the first time we
100    /// serialize a given value: subsequent times will use `SerDedup::Deduplicated` instead.
101    Value(#[serde_state(stateless)] DedupId, T),
102    /// A value represented by its id. The actual value must have been emitted as a
103    /// `SerDedup::Value` with that same id earlier.
104    #[serde_state(stateless)]
105    Deduplicated(DedupId),
106    /// A plain value without an id, emitted when we're not deduplicating.
107    Untagged(T),
108}
109
110/// Serialize `value`, deduplicating it if the state says so. `repr` is the serialized form of
111/// `value`, only used the first time we meet it.
112pub 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
132/// Deserialize a value that may have been deduplicated. `build` reconstructs the value from its
133/// serialized form.
134pub 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
166/// The error we report when a deduplicated value is deserialized with serde's stateless
167/// `Deserialize` impl, which can't resolve the ids.
168pub 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}