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#[derive(Copy, Clone, Debug, Serialize, Deserialize)]
23pub struct RollCompensation(pub u64);
24
25#[derive(Clone, Debug, Serialize, Deserialize)]
27pub struct RollUpdate {
28 pub roll_purchases: u64,
30 pub roll_sales: u64,
32 }
34
35impl RollUpdate {
36 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 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 pub fn is_nil(&self) -> bool {
74 self.roll_purchases == 0 && self.roll_sales == 0
75 }
76}
77
78pub struct RollUpdateSerializer {
80 u64_serializer: U64VarIntSerializer,
81}
82
83impl RollUpdateSerializer {
84 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 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
119pub struct RollUpdateDeserializer {
121 u64_deserializer: U64VarIntDeserializer,
122}
123
124impl RollUpdateDeserializer {
125 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 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#[derive(Clone, Debug, Serialize, Deserialize, Default)]
181pub struct RollUpdates(pub PreHashMap<Address, RollUpdate>);
182
183impl RollUpdates {
184 pub fn get_involved_addresses(&self) -> PreHashSet<Address> {
186 self.0.keys().copied().collect()
187 }
188
189 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 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 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 #[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 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#[derive(Clone, Debug, Serialize, Deserialize, Default)]
249pub struct RollCounts(pub BTreeMap<Address, u64>);
250
251impl RollCounts {
252 pub fn new() -> Self {
254 RollCounts(BTreeMap::new())
255 }
256
257 pub fn len(&self) -> usize {
259 self.0.len()
260 }
261
262 pub fn is_empty(&self) -> bool {
264 self.0.is_empty()
265 }
266
267 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 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 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 #[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 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}