1use crate::error::ModelsError;
4use crate::prehash::{PreHashSet, PreHashed};
5use bitvec::prelude::BitVec;
6use massa_serialization::{
7 Deserializer, SerializeError, Serializer, U32VarIntDeserializer, U32VarIntSerializer,
8 U64VarIntDeserializer, U64VarIntSerializer,
9};
10use nom::bytes::complete::take;
11use nom::multi::{length_count, length_data};
12use nom::sequence::preceded;
13use nom::{branch::alt, Parser, ToUsize};
14use nom::{
15 error::{context, ContextError, ErrorKind, ParseError},
16 IResult,
17};
18use num::integer::div_ceil;
19use std::convert::TryInto;
20use std::marker::PhantomData;
21use std::mem::size_of;
22use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
23use std::ops::Bound;
24use Bound::Included;
25
26pub trait SerializeMinBEInt {
28 fn to_be_bytes_min(self, max_value: Self) -> Result<Vec<u8>, ModelsError>;
30}
31
32impl SerializeMinBEInt for u32 {
33 fn to_be_bytes_min(self, max_value: Self) -> Result<Vec<u8>, ModelsError> {
34 if self > max_value {
35 return Err(ModelsError::SerializeError("integer out of bounds".into()));
36 }
37 let skip_bytes = (max_value.leading_zeros() as usize) / 8;
38 Ok(self.to_be_bytes()[skip_bytes..].to_vec())
39 }
40}
41
42impl SerializeMinBEInt for u64 {
43 fn to_be_bytes_min(self, max_value: Self) -> Result<Vec<u8>, ModelsError> {
44 if self > max_value {
45 return Err(ModelsError::SerializeError("integer out of bounds".into()));
46 }
47 let skip_bytes = (max_value.leading_zeros() as usize) / 8;
48 Ok(self.to_be_bytes()[skip_bytes..].to_vec())
49 }
50}
51
52pub trait DeserializeMinBEInt: Sized {
54 fn from_be_bytes_min(buffer: &[u8], max_value: Self) -> Result<(Self, usize), ModelsError>;
57}
58
59pub const fn u32_be_bytes_min_length(max_value: u32) -> usize {
61 size_of::<u32>() - (max_value.leading_zeros() as usize) / 8
62}
63
64pub const fn u64_be_bytes_min_length(max_value: u64) -> usize {
66 size_of::<u64>() - (max_value.leading_zeros() as usize) / 8
67}
68impl DeserializeMinBEInt for u32 {
69 fn from_be_bytes_min(buffer: &[u8], max_value: Self) -> Result<(Self, usize), ModelsError> {
70 let read_bytes = u32_be_bytes_min_length(max_value);
71 let skip_bytes = size_of::<Self>() - read_bytes;
72 if buffer.len() < read_bytes {
73 return Err(ModelsError::SerializeError("unexpected buffer END".into()));
74 }
75 let mut buf = [0u8; size_of::<Self>()];
76 buf[skip_bytes..].clone_from_slice(&buffer[..read_bytes]);
77 let res = u32::from_be_bytes(buf);
78 if res > max_value {
79 return Err(ModelsError::SerializeError(
80 "integer outside of bounds".into(),
81 ));
82 }
83 Ok((res, read_bytes))
84 }
85}
86
87impl DeserializeMinBEInt for u64 {
88 fn from_be_bytes_min(buffer: &[u8], max_value: Self) -> Result<(Self, usize), ModelsError> {
89 let read_bytes = u64_be_bytes_min_length(max_value);
90 let skip_bytes = size_of::<Self>() - read_bytes;
91 if buffer.len() < read_bytes {
92 return Err(ModelsError::SerializeError("unexpected buffer END".into()));
93 }
94 let mut buf = [0u8; size_of::<Self>()];
95 buf[skip_bytes..].clone_from_slice(&buffer[..read_bytes]);
96 let res = u64::from_be_bytes(buf);
97 if res > max_value {
98 return Err(ModelsError::SerializeError(
99 "integer outside of bounds".into(),
100 ));
101 }
102 Ok((res, read_bytes))
103 }
104}
105
106pub fn array_from_slice<const ARRAY_SIZE: usize>(
108 buffer: &[u8],
109) -> Result<[u8; ARRAY_SIZE], ModelsError> {
110 if buffer.len() < ARRAY_SIZE {
111 return Err(ModelsError::BufferError(
112 "slice too small to extract array".into(),
113 ));
114 }
115 buffer[..ARRAY_SIZE].try_into().map_err(|err| {
116 ModelsError::BufferError(format!("could not extract array from slice: {}", err))
117 })
118}
119
120pub fn u8_from_slice(buffer: &[u8]) -> Result<u8, ModelsError> {
122 if buffer.is_empty() {
123 return Err(ModelsError::BufferError(
124 "could not read u8 from empty buffer".into(),
125 ));
126 }
127 Ok(buffer[0])
128}
129
130#[derive(Default, Clone)]
132pub struct IpAddrSerializer;
133
134impl IpAddrSerializer {
135 pub const fn new() -> Self {
137 Self
138 }
139}
140
141impl Serializer<IpAddr> for IpAddrSerializer {
142 fn serialize(&self, value: &IpAddr, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
154 match value {
155 IpAddr::V4(ip_v4) => {
156 buffer.push(4u8);
157 buffer.extend(ip_v4.octets());
158 }
159 IpAddr::V6(ip_v6) => {
160 buffer.push(6u8);
161 buffer.extend(ip_v6.octets());
162 }
163 };
164 Ok(())
165 }
166}
167
168#[derive(Default, Clone)]
170pub struct IpAddrDeserializer;
171
172impl IpAddrDeserializer {
173 pub const fn new() -> Self {
175 Self
176 }
177}
178
179impl Deserializer<IpAddr> for IpAddrDeserializer {
180 fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
196 &self,
197 buffer: &'a [u8],
198 ) -> IResult<&'a [u8], IpAddr, E> {
199 context(
200 "Failed IpAddr deserialization",
201 alt((
202 preceded(
203 |input| nom::bytes::complete::tag([4u8])(input),
204 |input: &'a [u8]| {
205 let (rest, addr) = take(4usize)(input)?;
206 let addr: [u8; 4] = addr.try_into().unwrap();
207 Ok((rest, IpAddr::V4(Ipv4Addr::from(addr))))
208 },
209 ),
210 preceded(
211 |input| nom::bytes::complete::tag([6u8])(input),
212 |input: &'a [u8]| {
213 let (rest, addr) = take(16usize)(input)?;
214 let addr: [u8; 16] = addr.try_into().unwrap();
216 Ok((rest, IpAddr::V6(Ipv6Addr::from(addr))))
217 },
218 ),
219 )),
220 )(buffer)
221 }
222}
223
224#[derive(Clone)]
226pub struct VecU8Serializer {
227 len_serializer: U64VarIntSerializer,
228}
229
230impl VecU8Serializer {
231 pub fn new() -> Self {
233 Self {
234 len_serializer: U64VarIntSerializer::new(),
235 }
236 }
237}
238
239impl Default for VecU8Serializer {
240 fn default() -> Self {
241 Self::new()
242 }
243}
244
245impl Serializer<Vec<u8>> for VecU8Serializer {
246 fn serialize(&self, value: &Vec<u8>, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
257 let len: u64 = value.len().try_into().map_err(|err| {
258 SerializeError::NumberTooBig(format!("too many entries data in VecU8: {}", err))
259 })?;
260 self.len_serializer.serialize(&len, buffer)?;
261 buffer.extend(value);
262 Ok(())
263 }
264}
265
266#[derive(Clone)]
268pub struct VecU8Deserializer {
269 varint_u64_deserializer: U64VarIntDeserializer,
270}
271
272impl VecU8Deserializer {
273 pub const fn new(min_length: Bound<u64>, max_length: Bound<u64>) -> Self {
275 Self {
276 varint_u64_deserializer: U64VarIntDeserializer::new(min_length, max_length),
277 }
278 }
279}
280
281impl Deserializer<Vec<u8>> for VecU8Deserializer {
282 fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
297 &self,
298 buffer: &'a [u8],
299 ) -> IResult<&'a [u8], Vec<u8>, E> {
300 context("Failed Vec<u8> deserialization", |input| {
301 length_data(|input| self.varint_u64_deserializer.deserialize(input))(input)
302 })
303 .map(|res| res.to_vec())
304 .parse(buffer)
305 }
306}
307
308#[derive(Clone)]
310pub struct VecSerializer<T, ST>
311where
312 ST: Serializer<T>,
313{
314 len_serializer: U64VarIntSerializer,
315 data_serializer: ST,
316 phantom_t: PhantomData<T>,
317}
318
319impl<T, ST> VecSerializer<T, ST>
320where
321 ST: Serializer<T>,
322{
323 pub fn new(data_serializer: ST) -> Self {
325 Self {
326 len_serializer: U64VarIntSerializer::new(),
327 data_serializer,
328 phantom_t: PhantomData,
329 }
330 }
331}
332
333impl<T, ST> Serializer<Vec<T>> for VecSerializer<T, ST>
334where
335 ST: Serializer<T>,
336{
337 fn serialize(&self, value: &Vec<T>, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
338 self.len_serializer
339 .serialize(&(value.len() as u64), buffer)?;
340 for elem in value {
341 self.data_serializer.serialize(elem, buffer)?;
342 }
343 Ok(())
344 }
345}
346
347#[derive(Clone)]
349pub struct VecDeserializer<T, ST>
350where
351 ST: Deserializer<T> + Clone,
352{
353 varint_u64_deserializer: U64VarIntDeserializer,
354 data_deserializer: ST,
355 phantom_t: PhantomData<T>,
356}
357
358impl<T, ST> VecDeserializer<T, ST>
359where
360 ST: Deserializer<T> + Clone,
361{
362 pub const fn new(
364 data_deserializer: ST,
365 min_length: Bound<u64>,
366 max_length: Bound<u64>,
367 ) -> Self {
368 Self {
369 varint_u64_deserializer: U64VarIntDeserializer::new(min_length, max_length),
370 data_deserializer,
371 phantom_t: PhantomData,
372 }
373 }
374}
375
376impl<T, ST> Deserializer<Vec<T>> for VecDeserializer<T, ST>
377where
378 ST: Deserializer<T> + Clone,
379{
380 fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
381 &self,
382 buffer: &'a [u8],
383 ) -> IResult<&'a [u8], Vec<T>, E> {
384 context("Failed Vec<_> deserialization", |input| {
385 length_count(
386 context("length", |input| {
387 self.varint_u64_deserializer.deserialize(input)
388 }),
389 context("data", |input| self.data_deserializer.deserialize(input)),
390 )(input)
391 })
392 .parse(buffer)
393 }
394}
395
396#[derive(Clone)]
398pub struct PreHashSetSerializer<T, ST>
399where
400 ST: Serializer<T>,
401{
402 len_serializer: U64VarIntSerializer,
403 data_serializer: ST,
404 phantom_t: PhantomData<T>,
405}
406
407impl<T, ST> PreHashSetSerializer<T, ST>
408where
409 ST: Serializer<T>,
410{
411 pub fn new(data_serializer: ST) -> Self {
413 Self {
414 len_serializer: U64VarIntSerializer::new(),
415 data_serializer,
416 phantom_t: PhantomData,
417 }
418 }
419}
420
421impl<T, ST> Serializer<PreHashSet<T>> for PreHashSetSerializer<T, ST>
422where
423 ST: Serializer<T>,
424 T: PreHashed,
425{
426 fn serialize(&self, value: &PreHashSet<T>, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
427 self.len_serializer
428 .serialize(&(value.len() as u64), buffer)?;
429 for elem in value {
430 self.data_serializer.serialize(elem, buffer)?;
431 }
432 Ok(())
433 }
434}
435
436#[derive(Clone)]
438pub struct PreHashSetDeserializer<T, ST>
439where
440 ST: Deserializer<T> + Clone,
441{
442 varint_u64_deserializer: U64VarIntDeserializer,
443 data_deserializer: ST,
444 phantom_t: PhantomData<T>,
445}
446
447impl<T, ST> PreHashSetDeserializer<T, ST>
448where
449 ST: Deserializer<T> + Clone,
450{
451 pub const fn new(
453 data_deserializer: ST,
454 min_length: Bound<u64>,
455 max_length: Bound<u64>,
456 ) -> Self {
457 Self {
458 varint_u64_deserializer: U64VarIntDeserializer::new(min_length, max_length),
459 data_deserializer,
460 phantom_t: PhantomData,
461 }
462 }
463}
464
465impl<T, ST> Deserializer<PreHashSet<T>> for PreHashSetDeserializer<T, ST>
466where
467 ST: Deserializer<T> + Clone,
468 T: PreHashed + std::cmp::Eq + std::hash::Hash,
469{
470 fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
471 &self,
472 buffer: &'a [u8],
473 ) -> IResult<&'a [u8], PreHashSet<T>, E> {
474 context("Failed PreHashSet<_> deserialization", |input| {
475 length_count(
476 context("length", |input| {
477 self.varint_u64_deserializer.deserialize(input)
478 }),
479 context("data", |input| self.data_deserializer.deserialize(input)),
480 )(input)
481 })
482 .map(|vec| vec.into_iter().collect())
483 .parse(buffer)
484 }
485}
486
487#[derive(Clone)]
488pub struct StringSerializer<SL, L>
490where
491 SL: Serializer<L>,
492 L: TryFrom<usize>,
493{
494 length_serializer: SL,
495 marker_l: std::marker::PhantomData<L>,
496}
497
498impl<SL, L> StringSerializer<SL, L>
499where
500 SL: Serializer<L>,
501 L: TryFrom<usize>,
502{
503 pub fn new(length_serializer: SL) -> Self {
508 Self {
509 length_serializer,
510 marker_l: std::marker::PhantomData,
511 }
512 }
513}
514
515impl<SL, L> Serializer<String> for StringSerializer<SL, L>
516where
517 SL: Serializer<L>,
518 L: TryFrom<usize>,
519{
520 fn serialize(&self, value: &String, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
521 self.length_serializer.serialize(
522 &value.len().try_into().map_err(|_| {
523 SerializeError::StringTooBig("The string is too big to be serialized".to_string())
524 })?,
525 buffer,
526 )?;
527 buffer.extend(value.as_bytes());
528 Ok(())
529 }
530}
531
532#[derive(Clone)]
534pub struct StringDeserializer<DL, L>
535where
536 DL: Deserializer<L>,
537 L: TryFrom<usize> + ToUsize,
538{
539 length_deserializer: DL,
540 marker_l: std::marker::PhantomData<L>,
541}
542
543impl<DL, L> StringDeserializer<DL, L>
544where
545 DL: Deserializer<L>,
546 L: TryFrom<usize> + ToUsize,
547{
548 pub const fn new(length_deserializer: DL) -> Self {
553 Self {
554 length_deserializer,
555 marker_l: std::marker::PhantomData,
556 }
557 }
558}
559
560impl<DL, L> Deserializer<String> for StringDeserializer<DL, L>
561where
562 DL: Deserializer<L>,
563 L: TryFrom<usize> + ToUsize,
564{
565 fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
566 &self,
567 buffer: &'a [u8],
568 ) -> IResult<&'a [u8], String, E> {
569 let (rest, res) = length_data(|input| self.length_deserializer.deserialize(input))
570 .map(|data| {
571 String::from_utf8(data.to_vec()).map_err(|_| {
572 nom::Err::Error(ParseError::from_error_kind(
573 data,
574 nom::error::ErrorKind::Verify,
575 ))
576 })
577 })
578 .parse(buffer)?;
579 Ok((rest, res?))
580 }
581}
582
583#[derive(Clone)]
584pub struct BitVecSerializer {
586 u32_serializer: U32VarIntSerializer,
587}
588
589impl BitVecSerializer {
590 pub fn new() -> BitVecSerializer {
592 BitVecSerializer {
593 u32_serializer: U32VarIntSerializer::new(),
594 }
595 }
596}
597
598impl Default for BitVecSerializer {
599 fn default() -> Self {
600 Self::new()
601 }
602}
603
604impl Serializer<BitVec<u8>> for BitVecSerializer {
605 fn serialize(&self, value: &BitVec<u8>, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
606 let n_entries: u32 = value.len().try_into().map_err(|err| {
607 SerializeError::NumberTooBig(format!(
608 "too many entries when serializing a `BitVec<u8>`: {}",
609 err
610 ))
611 })?;
612 self.u32_serializer.serialize(&n_entries, buffer)?;
613 buffer.extend(value.clone().into_vec());
614 Ok(())
615 }
616}
617
618#[derive(Clone)]
619pub struct BitVecDeserializer {
621 u32_deserializer: U32VarIntDeserializer,
622}
623
624impl BitVecDeserializer {
625 pub fn new() -> BitVecDeserializer {
627 BitVecDeserializer {
628 u32_deserializer: U32VarIntDeserializer::new(
629 Bound::Included(u32::MIN),
630 Included(u32::MAX),
631 ),
632 }
633 }
634}
635
636impl Default for BitVecDeserializer {
637 fn default() -> Self {
638 Self::new()
639 }
640}
641
642impl Deserializer<BitVec<u8>> for BitVecDeserializer {
643 fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
644 &self,
645 buffer: &'a [u8],
646 ) -> IResult<&'a [u8], BitVec<u8>, E> {
647 context("Failed rng_seed deserialization", |input| {
648 let (rest, n_entries) = self.u32_deserializer.deserialize(input)?;
649 let bits_u8_len = div_ceil(n_entries, u8::BITS) as usize;
650 if rest.len() < bits_u8_len {
651 return Err(nom::Err::Error(ParseError::from_error_kind(
652 input,
653 ErrorKind::Eof,
654 )));
655 }
656 let mut rng_seed: BitVec<u8> = BitVec::try_from_vec(rest[..bits_u8_len].to_vec())
657 .map_err(|_| nom::Err::Error(ParseError::from_error_kind(input, ErrorKind::Eof)))?;
658 rng_seed.truncate(n_entries as usize);
659 if rng_seed.len() != n_entries as usize {
660 return Err(nom::Err::Error(ParseError::from_error_kind(
661 input,
662 ErrorKind::Eof,
663 )));
664 }
665 Ok((&rest[bits_u8_len..], rng_seed))
666 })
667 .map(|elements| elements.into_iter().collect())
668 .parse(buffer)
669 }
670}
671
672#[cfg(test)]
673mod tests {
674 use super::*;
675 use massa_serialization::DeserializeError;
676 use serial_test::serial;
677 use std::ops::Bound::Included;
678 #[test]
679 #[serial]
680 fn vec_u8() {
681 let vec: Vec<u8> = vec![9, 8, 7];
682 let vec_u8_serializer = VecU8Serializer::new();
683 let vec_u8_deserializer = VecU8Deserializer::new(Included(u64::MIN), Included(u64::MAX));
684 let mut serialized = Vec::new();
685 vec_u8_serializer.serialize(&vec, &mut serialized).unwrap();
686 let (rest, new_vec) = vec_u8_deserializer
687 .deserialize::<DeserializeError>(&serialized)
688 .unwrap();
689 assert!(rest.is_empty());
690 assert_eq!(vec, new_vec);
691 }
692
693 #[test]
694 #[serial]
695 fn vec_u8_big_length() {
696 let vec: Vec<u8> = vec![9, 8, 7];
697 let len: u64 = 10;
698 let mut serialized = Vec::new();
699 U64VarIntSerializer::new()
700 .serialize(&len, &mut serialized)
701 .unwrap();
702 serialized.extend(vec);
703 let vec_u8_deserializer = VecU8Deserializer::new(Included(u64::MIN), Included(u64::MAX));
704 let _ = vec_u8_deserializer
705 .deserialize::<DeserializeError>(&serialized)
706 .expect_err("Should fail too long size");
707 }
708
709 #[test]
710 #[serial]
711 fn vec_u8_min_length() {
712 let vec: Vec<u8> = vec![9, 8, 7];
713 let len: u64 = 1;
714 let mut serialized = Vec::new();
715 U64VarIntSerializer::new()
716 .serialize(&len, &mut serialized)
717 .unwrap();
718 serialized.extend(vec);
719 let vec_u8_deserializer = VecU8Deserializer::new(Included(u64::MIN), Included(u64::MAX));
720 let (rest, res) = vec_u8_deserializer
721 .deserialize::<DeserializeError>(&serialized)
722 .unwrap();
723 assert_eq!(rest, &[8, 7]);
724 assert_eq!(res, &[9])
725 }
726
727 #[test]
728 #[serial]
729 fn test_be_min() {
730 let x32 = 70_000u32;
731 let x64 = 10_000_000_000u64;
732
733 let mut res: Vec<u8> = Vec::new();
735 res.extend(x32.to_be_bytes_min(70_001).unwrap());
736 assert_eq!(res.len(), 3);
737 res.extend(x64.to_be_bytes_min(10_000_000_001).unwrap());
738 assert_eq!(res.len(), 3 + 5);
739
740 assert!(x32.to_be_bytes_min(69_999).is_err());
742 assert!(x64.to_be_bytes_min(9_999_999_999).is_err());
743
744 let buf = res.as_slice();
746 let mut cursor = 0;
747 let (out_x32, delta) = u32::from_be_bytes_min(&buf[cursor..], 70_001).unwrap();
748 assert_eq!(out_x32, x32);
749 cursor += delta;
750 let (out_x64, delta) = u64::from_be_bytes_min(&buf[cursor..], 10_000_000_001).unwrap();
751 assert_eq!(out_x64, x64);
752 cursor += delta;
753 assert_eq!(cursor, buf.len());
754 }
755
756 #[test]
757 #[serial]
758 fn test_array_from_slice_with_zero_u64() {
759 let zero: u64 = 0;
760 let res = array_from_slice(&zero.to_be_bytes()).unwrap();
761 assert_eq!(zero, u64::from_be_bytes(res));
762 }
763}