use std::{
fmt,
fmt::{Display, Formatter},
};
use super::MessageError;
use crate::decode_budget::{DecodeBudget, decode_with_max_items};
pub use crate::proto::envelope::*;
#[macro_export]
macro_rules! wrap_in_envelope_body {
($($e:expr),+) => {{
use $crate::message::MessageExt;
let mut envelope_body = $crate::message::EnvelopeBody::new();
$(
envelope_body.push_part($e.to_encoded_bytes());
)*
envelope_body
}}
}
impl EnvelopeBody {
pub fn new() -> Self {
Self {
parts: Default::default(),
}
}
pub fn len(&self) -> usize {
self.parts.len()
}
pub fn total_size(&self) -> usize {
self.parts.iter().fold(0, |acc: usize, b| acc.saturating_add(b.len()))
}
pub fn is_empty(&self) -> bool {
self.parts.is_empty()
}
pub fn take_part(&mut self, index: usize) -> Option<Vec<u8>> {
Some(index)
.filter(|i| self.parts.len() > *i)
.map(|i| self.parts.remove(i))
}
pub fn push_part(&mut self, part: Vec<u8>) {
self.parts.push(part)
}
pub fn into_inner(self) -> Vec<Vec<u8>> {
self.parts
}
#[deprecated(note = "use the _with_max_items variant, which bounds what decoding can allocate")]
pub fn decode_part<T>(&self, index: usize) -> Result<Option<T>, MessageError>
where T: prost::Message + Default {
match self.parts.get(index) {
Some(part) => T::decode(part.as_slice()).map(Some).map_err(Into::into),
None => Ok(None),
}
}
pub fn decode_part_with_max_items<T>(&self, index: usize, max_items: usize) -> Result<Option<T>, MessageError>
where T: prost::Message + Default + DecodeBudget {
match self.parts.get(index) {
Some(part) => decode_with_max_items(part, max_items).map(Some).map_err(Into::into),
None => Ok(None),
}
}
}
impl Display for EnvelopeBody {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{} byte(s), {} part(s)", self.total_size(), self.len())
}
}
#[cfg(test)]
mod test {
use prost::Message;
use super::*;
fn body_of_empty_parts(count: usize) -> EnvelopeBody {
EnvelopeBody {
parts: vec![Vec::new(); count],
}
}
#[test]
fn decode_part_with_max_items_rejects_an_over_budget_part_before_decoding_it() {
let flood = body_of_empty_parts(2_000).encode_to_vec();
let at_budget = body_of_empty_parts(1_000).encode_to_vec();
let body = EnvelopeBody {
parts: vec![flood, at_budget],
};
assert_eq!(
EnvelopeBody::decode(body.parts.first().unwrap().as_slice())
.unwrap()
.len(),
2_000
);
let err = body.decode_part_with_max_items::<EnvelopeBody>(0, 1_000).unwrap_err();
assert!(err.to_string().contains("decode budget"), "unexpected error: {err}");
let decoded = body
.decode_part_with_max_items::<EnvelopeBody>(1, 1_000)
.unwrap()
.unwrap();
assert_eq!(decoded.len(), 1_000);
assert!(
body.decode_part_with_max_items::<EnvelopeBody>(2, 1_000)
.unwrap()
.is_none()
);
}
}