massa_models/
rolls.rs

1use crate::{
2    address::Address,
3    error::ModelsError,
4    error::ModelsResult as Result,
5    prehash::{PreHashMap, PreHashSet},
6};
7use massa_serialization::{
8    Deserializer, SerializeError, Serializer, U64VarIntDeserializer, U64VarIntSerializer,
9};
10use nom::{
11    error::{context, ContextError, ParseError},
12    sequence::tuple,
13    IResult, Parser,
14};
15use serde::{Deserialize, Serialize};
16use std::collections::hash_map;
17use std::ops::Bound::Included;
18
19use std::collections::{btree_map, BTreeMap};
20
21/// just a `u64` to keep track of the roll sells and buys during a cycle
22#[derive(Copy, Clone, Debug, Serialize, Deserialize)]
23pub struct RollCompensation(pub u64);
24
25/// roll sales and purchases
26#[derive(Clone, Debug, Serialize, Deserialize)]
27pub struct RollUpdate {
28    /// roll purchases
29    pub roll_purchases: u64,
30    /// roll sales
31    pub roll_sales: u64,
32    // Here is space for registering any denunciations/resets
33}
34
35impl RollUpdate {
36    /// chain two roll updates, compensate and return compensation count
37    fn chain(&mut self, change: &Self) -> Result<RollCompensation> {
38        let compensation_other = std::cmp::min(change.roll_purchases, change.roll_sales);
39        self.roll_purchases = self
40            .roll_purchases
41            .checked_add(change.roll_purchases - compensation_other)
42            .ok_or_else(|| {
43                ModelsError::InvalidRollUpdate(
44                    "roll_purchases overflow in RollUpdate::chain".into(),
45                )
46            })?;
47        self.roll_sales = self
48            .roll_sales
49            .checked_add(change.roll_sales - compensation_other)
50            .ok_or_else(|| {
51                ModelsError::InvalidRollUpdate("roll_sales overflow in RollUpdate::chain".into())
52            })?;
53
54        let compensation_self = self.compensate().0;
55
56        let compensation_total = compensation_other
57            .checked_add(compensation_self)
58            .ok_or_else(|| {
59                ModelsError::InvalidRollUpdate("compensation overflow in RollUpdate::chain".into())
60            })?;
61        Ok(RollCompensation(compensation_total))
62    }
63
64    /// compensate a roll update, return compensation count
65    pub fn compensate(&mut self) -> RollCompensation {
66        let compensation = std::cmp::min(self.roll_purchases, self.roll_sales);
67        self.roll_purchases -= compensation;
68        self.roll_sales -= compensation;
69        RollCompensation(compensation)
70    }
71
72    /// true if the update has no effect
73    pub fn is_nil(&self) -> bool {
74        self.roll_purchases == 0 && self.roll_sales == 0
75    }
76}
77
78/// Serializer for `RollUpdate`
79pub struct RollUpdateSerializer {
80    u64_serializer: U64VarIntSerializer,
81}
82
83impl RollUpdateSerializer {
84    /// Creates a new `RollUpdateSerializer`
85    pub fn new() -> Self {
86        RollUpdateSerializer {
87            u64_serializer: U64VarIntSerializer::new(),
88        }
89    }
90}
91
92impl Default for RollUpdateSerializer {
93    fn default() -> Self {
94        Self::new()
95    }
96}
97
98impl Serializer<RollUpdate> for RollUpdateSerializer {
99    /// ## Example:
100    /// ```rust
101    /// use massa_models::rolls::{RollUpdate, RollUpdateSerializer};
102    /// use massa_serialization::Serializer;
103    ///
104    /// let roll_update = RollUpdate {
105    ///   roll_purchases: 1,
106    ///   roll_sales: 2,
107    /// };
108    /// let mut buffer = vec![];
109    /// RollUpdateSerializer::new().serialize(&roll_update, &mut buffer).unwrap();
110    /// ```
111    fn serialize(&self, value: &RollUpdate, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
112        self.u64_serializer
113            .serialize(&value.roll_purchases, buffer)?;
114        self.u64_serializer.serialize(&value.roll_sales, buffer)?;
115        Ok(())
116    }
117}
118
119/// Deserializer for `RollUpdate`
120pub struct RollUpdateDeserializer {
121    u64_deserializer: U64VarIntDeserializer,
122}
123
124impl RollUpdateDeserializer {
125    /// Creates a new `RollUpdateDeserializer`
126    pub fn new() -> Self {
127        RollUpdateDeserializer {
128            u64_deserializer: U64VarIntDeserializer::new(Included(0), Included(u64::MAX)),
129        }
130    }
131}
132
133impl Default for RollUpdateDeserializer {
134    fn default() -> Self {
135        Self::new()
136    }
137}
138
139impl Deserializer<RollUpdate> for RollUpdateDeserializer {
140    /// ## Example:
141    /// ```rust
142    /// use massa_models::rolls::{RollUpdate, RollUpdateDeserializer, RollUpdateSerializer};
143    /// use massa_serialization::{Serializer, Deserializer, DeserializeError};
144    ///
145    /// let roll_update = RollUpdate {
146    ///   roll_purchases: 1,
147    ///   roll_sales: 2,
148    /// };
149    /// let mut buffer = vec![];
150    /// RollUpdateSerializer::new().serialize(&roll_update, &mut buffer).unwrap();
151    /// let (rest, roll_update_deserialized) = RollUpdateDeserializer::new().deserialize::<DeserializeError>(&buffer).unwrap();
152    /// assert_eq!(rest.len(), 0);
153    /// assert_eq!(roll_update.roll_purchases, roll_update_deserialized.roll_purchases);
154    /// assert_eq!(roll_update.roll_sales, roll_update_deserialized.roll_sales);
155    /// ```
156    fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
157        &self,
158        buffer: &'a [u8],
159    ) -> IResult<&'a [u8], RollUpdate, E> {
160        context(
161            "Failed RollUpdate deserialization",
162            tuple((
163                context("Failed roll_purchases deserialization", |input| {
164                    self.u64_deserializer.deserialize(input)
165                }),
166                context("Failed roll_sales deserialization", |input| {
167                    self.u64_deserializer.deserialize(input)
168                }),
169            )),
170        )
171        .map(|(roll_purchases, roll_sales)| RollUpdate {
172            roll_purchases,
173            roll_sales,
174        })
175        .parse(buffer)
176    }
177}
178
179/// maps addresses to roll updates
180#[derive(Clone, Debug, Serialize, Deserialize, Default)]
181pub struct RollUpdates(pub PreHashMap<Address, RollUpdate>);
182
183impl RollUpdates {
184    /// the addresses impacted by the updates
185    pub fn get_involved_addresses(&self) -> PreHashSet<Address> {
186        self.0.keys().copied().collect()
187    }
188
189    /// chains with another `RollUpdates`, compensates and returns compensations
190    pub fn chain(
191        &mut self,
192        updates: &RollUpdates,
193    ) -> Result<PreHashMap<Address, RollCompensation>> {
194        let mut res = PreHashMap::default();
195        for (addr, update) in updates.0.iter() {
196            res.insert(*addr, self.apply(addr, update)?);
197            // remove if nil
198            if let hash_map::Entry::Occupied(occ) = self.0.entry(*addr) {
199                if occ.get().is_nil() {
200                    occ.remove();
201                }
202            }
203        }
204        Ok(res)
205    }
206
207    /// applies a `RollUpdate`, compensates and returns compensation
208    pub fn apply(&mut self, addr: &Address, update: &RollUpdate) -> Result<RollCompensation> {
209        if update.is_nil() {
210            return Ok(RollCompensation(0));
211        }
212        match self.0.entry(*addr) {
213            hash_map::Entry::Occupied(mut occ) => occ.get_mut().chain(update),
214            hash_map::Entry::Vacant(vac) => {
215                let mut compensated_update = update.clone();
216                let compensation = compensated_update.compensate();
217                vac.insert(compensated_update);
218                Ok(compensation)
219            }
220        }
221    }
222
223    /// get the roll update for a subset of addresses
224    #[must_use]
225    pub fn clone_subset(&self, addrs: &PreHashSet<Address>) -> Self {
226        Self(
227            addrs
228                .iter()
229                .filter_map(|addr| self.0.get(addr).map(|v| (*addr, v.clone())))
230                .collect(),
231        )
232    }
233
234    /// merge another roll updates into self, overwriting existing data
235    /// addresses that are in not other are removed from self
236    pub fn sync_from(&mut self, addrs: &PreHashSet<Address>, mut other: RollUpdates) {
237        for addr in addrs.iter() {
238            if let Some(new_val) = other.0.remove(addr) {
239                self.0.insert(*addr, new_val);
240            } else {
241                self.0.remove(addr);
242            }
243        }
244    }
245}
246
247/// counts the roll for each address
248#[derive(Clone, Debug, Serialize, Deserialize, Default)]
249pub struct RollCounts(pub BTreeMap<Address, u64>);
250
251impl RollCounts {
252    /// Makes a new, empty `RollCounts`.
253    pub fn new() -> Self {
254        RollCounts(BTreeMap::new())
255    }
256
257    /// Returns the number of elements in the `RollCounts`.
258    pub fn len(&self) -> usize {
259        self.0.len()
260    }
261
262    /// Returns true if the `RollCounts` contains no elements.
263    pub fn is_empty(&self) -> bool {
264        self.0.is_empty()
265    }
266
267    /// applies `RollUpdates` to self with compensations
268    pub fn apply_updates(&mut self, updates: &RollUpdates) -> Result<()> {
269        for (addr, update) in updates.0.iter() {
270            match self.0.entry(*addr) {
271                btree_map::Entry::Occupied(mut occ) => {
272                    let cur_val = *occ.get();
273                    if update.roll_purchases >= update.roll_sales {
274                        *occ.get_mut() = cur_val
275                            .checked_add(update.roll_purchases - update.roll_sales)
276                            .ok_or_else(|| {
277                                ModelsError::InvalidRollUpdate(
278                                    "overflow while incrementing roll count".into(),
279                                )
280                            })?;
281                    } else {
282                        *occ.get_mut() = cur_val
283                            .checked_sub(update.roll_sales - update.roll_purchases)
284                            .ok_or_else(|| {
285                                ModelsError::InvalidRollUpdate(
286                                    "underflow while decrementing roll count".into(),
287                                )
288                            })?;
289                    }
290                    if *occ.get() == 0 {
291                        // remove if 0
292                        occ.remove();
293                    }
294                }
295                btree_map::Entry::Vacant(vac) => {
296                    if update.roll_purchases >= update.roll_sales {
297                        if update.roll_purchases > update.roll_sales {
298                            // ignore if 0
299                            vac.insert(update.roll_purchases - update.roll_sales);
300                        }
301                    } else {
302                        return Err(ModelsError::InvalidRollUpdate(
303                            "underflow while decrementing roll count".into(),
304                        ));
305                    }
306                }
307            }
308        }
309        Ok(())
310    }
311
312    /// get roll counts for a subset of addresses.
313    #[must_use]
314    pub fn clone_subset(&self, addrs: &PreHashSet<Address>) -> Self {
315        Self(
316            addrs
317                .iter()
318                .filter_map(|addr| self.0.get(addr).map(|v| (*addr, *v)))
319                .collect(),
320        )
321    }
322
323    /// merge another roll counts into self, overwriting existing data
324    /// addresses that are in not other are removed from self
325    pub fn sync_from(&mut self, addrs: &PreHashSet<Address>, mut other: RollCounts) {
326        for addr in addrs.iter() {
327            if let Some(new_val) = other.0.remove(addr) {
328                self.0.insert(*addr, new_val);
329            } else {
330                self.0.remove(addr);
331            }
332        }
333    }
334}