pub const DEFAULT_MAX_DECODE_ITEMS: usize = 16_384;
const MAX_DEPTH: usize = 128;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DecodeBudgetExceeded {
pub items: usize,
pub max: usize,
}
#[derive(Debug)]
pub struct Budget {
items: usize,
max: usize,
depth: usize,
}
impl Budget {
pub fn new(max: usize) -> Self {
Self {
items: 0,
max,
depth: 0,
}
}
pub fn items(&self) -> usize {
self.items
}
fn exceeded(&self) -> DecodeBudgetExceeded {
DecodeBudgetExceeded {
items: self.items,
max: self.max,
}
}
pub fn charge_items(&mut self, count: usize) -> Result<(), DecodeBudgetExceeded> {
self.items = self.items.saturating_add(count);
if self.items > self.max {
return Err(self.exceeded());
}
Ok(())
}
pub fn charge_message<T: DecodeBudget + ?Sized>(&mut self, contents: &[u8]) -> Result<(), DecodeBudgetExceeded> {
self.items = self.items.saturating_add(1);
if self.items > self.max || self.depth >= MAX_DEPTH {
return Err(self.exceeded());
}
self.depth = self.depth.saturating_add(1);
let result = T::count_messages(contents, self);
self.depth = self.depth.saturating_sub(1);
result
}
pub fn charge_map_entry<V: DecodeBudget + ?Sized>(&mut self, contents: &[u8]) -> Result<(), DecodeBudgetExceeded> {
self.items = self.items.saturating_add(1);
if self.items > self.max {
return Err(self.exceeded());
}
walk_len_fields(contents, self, |tag, value, budget| {
if tag == 2 {
budget.charge_message::<V>(value)
} else {
Ok(())
}
})
}
}
pub trait DecodeBudget {
fn count_messages(buf: &[u8], budget: &mut Budget) -> Result<(), DecodeBudgetExceeded>;
fn count_oneof_field(_tag: u32, _contents: &[u8], _budget: &mut Budget) -> Result<(), DecodeBudgetExceeded> {
Ok(())
}
}
macro_rules! impl_no_embedded_messages {
($($ty:ty),* $(,)?) => {
$(
impl DecodeBudget for $ty {
fn count_messages(_buf: &[u8], _budget: &mut Budget) -> Result<(), DecodeBudgetExceeded> {
Ok(())
}
}
)*
};
}
impl_no_embedded_messages!((), bool, u32, u64, i32, i64, f32, f64, String, Vec<u8>, bytes::Bytes);
pub fn check_decode_budget<T: DecodeBudget + ?Sized>(
buf: &[u8],
max_items: usize,
) -> Result<usize, DecodeBudgetExceeded> {
let mut budget = Budget::new(max_items);
T::count_messages(buf, &mut budget)?;
Ok(budget.items())
}
pub fn decode_with_max_items<T>(buf: &[u8], max_items: usize) -> Result<T, prost::DecodeError>
where T: prost::Message + Default + DecodeBudget {
check_decode_budget::<T>(buf, max_items).map_err(|err| {
prost::DecodeError::new(format!(
"message exceeds the decode budget ({} embedded items, at most {} allowed)",
err.items, err.max
))
})?;
T::decode(buf)
}
pub fn walk_len_fields<F>(buf: &[u8], budget: &mut Budget, f: F) -> Result<(), DecodeBudgetExceeded>
where F: FnMut(u32, &[u8], &mut Budget) -> Result<(), DecodeBudgetExceeded> {
walk_fields(buf, budget, &[], &[], f)
}
pub fn walk_fields<F>(
buf: &[u8],
budget: &mut Budget,
map_tags: &[u32],
repeated_scalar_tags: &[u32],
mut f: F,
) -> Result<(), DecodeBudgetExceeded>
where
F: FnMut(u32, &[u8], &mut Budget) -> Result<(), DecodeBudgetExceeded>,
{
let mut pos = 0usize;
while pos < buf.len() {
let Some(key) = read_varint(buf, &mut pos) else {
return Ok(());
};
let Ok(tag) = u32::try_from(key >> 3) else {
return Ok(());
};
if tag == 0 {
return Ok(());
}
let wire_type = key & 0x7;
if wire_type > 5 {
return Ok(());
}
let wire_type = if map_tags.contains(&tag) { 2 } else { wire_type };
if matches!(wire_type, 0 | 1 | 5) && repeated_scalar_tags.contains(&tag) {
budget.charge_items(1)?;
}
let skip = match wire_type {
0 => {
if read_varint(buf, &mut pos).is_none() {
return Ok(());
}
0
},
1 => 8,
2 => {
let Some(len) = read_varint(buf, &mut pos).and_then(|len| usize::try_from(len).ok()) else {
return Ok(());
};
let Some(contents) = pos.checked_add(len).and_then(|end| buf.get(pos..end)) else {
return Ok(());
};
f(tag, contents, budget)?;
len
},
3 => {
if !skip_group(buf, &mut pos, tag, 0) {
return Ok(());
}
0
},
5 => 4,
_ => return Ok(()),
};
match pos.checked_add(skip).filter(|end| *end <= buf.len()) {
Some(end) => pos = end,
None => return Ok(()),
}
}
Ok(())
}
fn skip_group(buf: &[u8], pos: &mut usize, group_tag: u32, depth: usize) -> bool {
if depth >= MAX_DEPTH {
return false;
}
loop {
let Some(key) = read_varint(buf, pos) else {
return false;
};
let Ok(tag) = u32::try_from(key >> 3) else {
return false;
};
if tag == 0 {
return false;
}
let skip = match key & 0x7 {
0 => {
if read_varint(buf, pos).is_none() {
return false;
}
0
},
1 => 8,
2 => match read_varint(buf, pos).and_then(|len| usize::try_from(len).ok()) {
Some(len) => len,
None => return false,
},
3 => {
if !skip_group(buf, pos, tag, depth.saturating_add(1)) {
return false;
}
0
},
4 => return tag == group_tag,
5 => 4,
_ => return false,
};
match pos.checked_add(skip).filter(|end| *end <= buf.len()) {
Some(end) => *pos = end,
None => return false,
}
}
}
fn read_varint(buf: &[u8], pos: &mut usize) -> Option<u64> {
let mut value = 0u64;
for i in 0..10u32 {
let byte = *buf.get(*pos)?;
*pos = pos.checked_add(1)?;
value |= u64::from(byte & 0x7f).checked_shl(i.checked_mul(7)?)?;
if byte & 0x80 == 0 {
return Some(value);
}
}
None
}
#[cfg(test)]
mod test {
use std::collections::{BTreeMap, HashMap};
use prost::Message;
use tari_comms_rpc_macros::DecodeBudget;
use super::*;
#[derive(Clone, PartialEq, Message, DecodeBudget)]
struct Leaf {
#[prost(uint32, tag = "1")]
value: u32,
}
#[derive(Clone, PartialEq, Message, DecodeBudget)]
struct Inner {
#[prost(bytes = "vec", tag = "1")]
data: Vec<u8>,
#[prost(message, repeated, tag = "2")]
leaves: Vec<Leaf>,
}
#[derive(Clone, PartialEq, prost::Oneof, DecodeBudget)]
enum Choice {
#[prost(message, tag = "3")]
Message(Inner),
#[prost(bytes, tag = "4")]
Bytes(Vec<u8>),
#[prost(message, tag = "7")]
Boxed(Box<Leaf>),
}
#[derive(Clone, PartialEq, Message, DecodeBudget)]
struct Outer {
#[prost(message, repeated, tag = "1")]
items: Vec<Inner>,
#[prost(message, optional, boxed, tag = "2")]
single: Option<Box<Inner>>,
#[prost(oneof = "Choice", tags = "3, 4, 7")]
choice: Option<Choice>,
#[prost(map = "string, message", tag = "5")]
by_name: HashMap<String, Leaf>,
#[prost(bytes = "vec", tag = "6")]
blob: Vec<u8>,
#[prost(string, tag = "8")]
text: String,
#[prost(uint64, repeated, tag = "9")]
numbers: Vec<u64>,
#[prost(bytes = "vec", repeated, tag = "10")]
hashes: Vec<Vec<u8>>,
#[prost(string, repeated, tag = "11")]
names: Vec<String>,
#[prost(map = "int32, bytes", tag = "12")]
blobs: HashMap<i32, Vec<u8>>,
#[prost(fixed32, repeated, tag = "13")]
fixed: Vec<u32>,
#[prost(btree_map = "string, message", tag = "14")]
tree: BTreeMap<String, Leaf>,
}
#[derive(Clone, PartialEq, Message, DecodeBudget)]
struct Node {
#[prost(message, optional, boxed, tag = "1")]
child: Option<Box<Node>>,
}
fn count<T: DecodeBudget>(msg: &impl Message, max: usize) -> Result<usize, DecodeBudgetExceeded> {
check_decode_budget::<T>(&msg.encode_to_vec(), max)
}
fn fake_messages(len: usize) -> Vec<u8> {
[0x0a, 0x00].repeat(len / 2)
}
fn leaf() -> Leaf {
Leaf { value: 1 }
}
fn inner(leaves: usize) -> Inner {
Inner {
data: fake_messages(64),
leaves: vec![leaf(); leaves],
}
}
#[test]
fn it_counts_every_message_instance() {
let msg = Outer {
items: vec![inner(2); 3],
single: Some(Box::new(inner(4))),
choice: Some(Choice::Message(inner(5))),
by_name: [("a".to_string(), leaf()), ("b".to_string(), leaf())]
.into_iter()
.collect(),
blob: fake_messages(1_000),
text: "\n\0\n\0".to_string(),
numbers: vec![10; 100],
hashes: vec![fake_messages(64); 3],
names: vec![String::from_utf8(fake_messages(64)).unwrap(); 2],
blobs: [(1, fake_messages(64)), (2, fake_messages(64))].into_iter().collect(),
fixed: vec![1; 10],
tree: [("t".to_string(), leaf())].into_iter().collect(),
};
assert_eq!(
count::<Outer>(&msg, DEFAULT_MAX_DECODE_ITEMS).unwrap(),
9 + 5 + 6 + 4 + 100 + 3 + 2 + 2 + 40 + 2
);
let boxed_variant = Outer {
choice: Some(Choice::Boxed(Box::new(leaf()))),
..Default::default()
};
assert_eq!(count::<Outer>(&boxed_variant, DEFAULT_MAX_DECODE_ITEMS).unwrap(), 1);
}
#[test]
fn it_never_enters_bytes_or_strings() {
let msg = Outer {
blob: fake_messages(1_000_000),
choice: Some(Choice::Bytes(fake_messages(1_000_000))),
text: String::from_utf8(fake_messages(1_000_000)).unwrap(),
single: Some(Box::new(Inner {
data: fake_messages(1_000_000),
leaves: vec![],
})),
..Default::default()
};
assert_eq!(count::<Outer>(&msg, 1).unwrap(), 1);
}
fn len_field(tag: u32, contents: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(contents.len().saturating_add(16));
prost::encoding::encode_key(tag, prost::encoding::WireType::LengthDelimited, &mut buf);
prost::encoding::encode_varint(contents.len() as u64, &mut buf);
buf.extend_from_slice(contents);
buf
}
fn empty_elements(tag: u32, count: usize) -> Vec<u8> {
len_field(tag, &[]).repeat(count)
}
const MAX_REQUEST_SIZE: usize = 6 * 1024 * 1024;
#[cfg(feature = "rpc")]
const _: () = assert!(MAX_REQUEST_SIZE == crate::protocol::rpc::RPC_MAX_REQUEST_SIZE);
const FLOOD_ELEMENTS: usize = MAX_REQUEST_SIZE / 2 - 8;
#[test]
fn it_rejects_a_flat_flood() {
let flood = empty_elements(1, FLOOD_ELEMENTS);
assert!(flood.len() <= MAX_REQUEST_SIZE);
let err = check_decode_budget::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
assert_eq!(err, DecodeBudgetExceeded {
items: DEFAULT_MAX_DECODE_ITEMS + 1,
max: DEFAULT_MAX_DECODE_ITEMS
});
}
#[test]
fn it_rejects_a_flood_of_bytes_or_string_elements() {
check_decode_budget::<Outer>(&empty_elements(10, FLOOD_ELEMENTS), DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
check_decode_budget::<Outer>(&empty_elements(11, FLOOD_ELEMENTS), DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
}
#[test]
fn a_map_field_is_charged_whatever_its_wire_type() {
use prost::encoding::{WireType, encode_key, encode_varint};
let mut entry = len_field(1, b"a");
entry.extend_from_slice(&len_field(2, &leaf().encode_to_vec()));
let tail = empty_elements(1, 10);
for wire_type in [
WireType::Varint,
WireType::SixtyFourBit,
WireType::LengthDelimited,
WireType::StartGroup,
WireType::EndGroup,
WireType::ThirtyTwoBit,
] {
let mut bytes = Vec::new();
encode_key(5, wire_type, &mut bytes);
encode_varint(entry.len() as u64, &mut bytes);
bytes.extend_from_slice(&entry);
bytes.extend_from_slice(&tail);
let decoded = Outer::decode(bytes.as_slice()).unwrap();
assert_eq!(decoded.by_name.len(), 1, "{wire_type:?}");
assert_eq!(decoded.items.len(), 10, "{wire_type:?}");
assert_eq!(
check_decode_budget::<Outer>(&bytes, DEFAULT_MAX_DECODE_ITEMS).unwrap(),
12,
"{wire_type:?}"
);
}
let mut flood = Vec::new();
for _ in 0..100_000 {
encode_key(5, WireType::Varint, &mut flood);
encode_varint(entry.len() as u64, &mut flood);
flood.extend_from_slice(&entry);
}
check_decode_budget::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
}
#[test]
fn it_rejects_a_flood_of_scalar_map_entries() {
let mut flood = Vec::new();
for key in 0..100_000u64 {
let mut entry = Vec::new();
prost::encoding::encode_key(1, prost::encoding::WireType::Varint, &mut entry);
prost::encoding::encode_varint(key, &mut entry);
flood.extend_from_slice(&len_field(12, &entry));
}
check_decode_budget::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
let msg = Outer {
blobs: (0..100).map(|key| (key, vec![])).collect(),
..Default::default()
};
assert_eq!(count::<Outer>(&msg, DEFAULT_MAX_DECODE_ITEMS).unwrap(), 100);
}
#[test]
fn it_rejects_a_flood_of_packed_scalars() {
let flood = len_field(9, &vec![1u8; 2 * FLOOD_ELEMENTS]);
check_decode_budget::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
let msg = Outer {
numbers: vec![u64::MAX; 100],
..Default::default()
};
assert_eq!(count::<Outer>(&msg, DEFAULT_MAX_DECODE_ITEMS).unwrap(), 1_000);
}
#[test]
fn unpacked_repeated_scalars_are_charged_one_per_element() {
use prost::encoding::{WireType, encode_key, encode_varint};
let mut bytes = Vec::new();
for _ in 0..100 {
encode_key(9, WireType::Varint, &mut bytes);
encode_varint(7, &mut bytes);
}
for _ in 0..50 {
encode_key(13, WireType::ThirtyTwoBit, &mut bytes);
bytes.extend_from_slice(&7u32.to_le_bytes());
}
bytes.extend_from_slice(&len_field(9, &[7u8; 10]));
let decoded = Outer::decode(bytes.as_slice()).unwrap();
assert_eq!(decoded.numbers.len(), 110);
assert_eq!(decoded.fixed.len(), 50);
assert_eq!(
check_decode_budget::<Outer>(&bytes, DEFAULT_MAX_DECODE_ITEMS).unwrap(),
160
);
let mut element = Vec::new();
encode_key(9, WireType::Varint, &mut element);
encode_varint(0, &mut element);
let flood = element.repeat(FLOOD_ELEMENTS);
assert_eq!(
Outer::decode(flood.get(..2 * 1_000).unwrap()).unwrap().numbers.len(),
1_000
);
check_decode_budget::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
}
#[test]
fn a_btree_map_field_is_charged_like_a_map() {
use prost::encoding::{WireType, encode_key, encode_varint};
let msg = Outer {
tree: [("a".to_string(), leaf()), ("b".to_string(), leaf())]
.into_iter()
.collect(),
..Default::default()
};
assert_eq!(count::<Outer>(&msg, DEFAULT_MAX_DECODE_ITEMS).unwrap(), 4);
let mut entry = len_field(1, b"a");
entry.extend_from_slice(&len_field(2, &leaf().encode_to_vec()));
let mut bytes = Vec::new();
encode_key(14, WireType::Varint, &mut bytes);
encode_varint(entry.len() as u64, &mut bytes);
bytes.extend_from_slice(&entry);
bytes.extend_from_slice(&empty_elements(1, 10));
let decoded = Outer::decode(bytes.as_slice()).unwrap();
assert_eq!(decoded.tree.len(), 1);
assert_eq!(decoded.items.len(), 10);
assert_eq!(
check_decode_budget::<Outer>(&bytes, DEFAULT_MAX_DECODE_ITEMS).unwrap(),
12
);
check_decode_budget::<Outer>(&len_field(14, &[]).repeat(FLOOD_ELEMENTS), DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
}
fn group_prefix() -> Vec<u8> {
let mut buf = Vec::new();
prost::encoding::encode_key(1000, prost::encoding::WireType::StartGroup, &mut buf);
prost::encoding::encode_key(1, prost::encoding::WireType::Varint, &mut buf);
prost::encoding::encode_varint(7, &mut buf);
prost::encoding::encode_key(1000, prost::encoding::WireType::EndGroup, &mut buf);
buf
}
fn with_group_prefix(msg: &impl Message) -> Vec<u8> {
let mut buf = group_prefix();
buf.extend_from_slice(&msg.encode_to_vec());
buf
}
#[test]
fn a_group_does_not_hide_a_flood() {
let mut flood = group_prefix();
flood.extend_from_slice(&empty_elements(1, FLOOD_ELEMENTS));
check_decode_budget::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
let mut inner = group_prefix();
inner.extend_from_slice(&empty_elements(2, FLOOD_ELEMENTS));
check_decode_budget::<Outer>(&len_field(2, &inner), DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
let small = Outer {
items: vec![Inner::default(); 10],
..Default::default()
};
let bytes = with_group_prefix(&small);
assert_eq!(Outer::decode(bytes.as_slice()).unwrap(), small);
assert_eq!(
check_decode_budget::<Outer>(&bytes, DEFAULT_MAX_DECODE_ITEMS).unwrap(),
10
);
}
#[test]
fn it_rejects_a_nested_flood() {
let flood = Outer {
single: Some(Box::new(inner(100_000))),
..Default::default()
};
count::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
let flood = Outer {
choice: Some(Choice::Message(inner(100_000))),
..Default::default()
};
count::<Outer>(&flood, DEFAULT_MAX_DECODE_ITEMS).unwrap_err();
}
#[test]
fn it_rejects_one_over_the_cap() {
let msg = Outer {
items: vec![Inner::default(); 100],
..Default::default()
};
assert_eq!(count::<Outer>(&msg, 100).unwrap(), 100);
let err = count::<Outer>(&msg, 99).unwrap_err();
assert_eq!(err, DecodeBudgetExceeded { items: 100, max: 99 });
}
#[test]
fn malformed_wire_stops_the_count_without_rejecting() {
let malformed: &[&[u8]] = &[
&[0x00, 0x01],
&[0x0e, 0x01],
&[0x0f, 0x01],
&[0x7c],
&[0x7b, 0x08, 0x01],
&[0x7b, 0x74],
&[0x80],
&[0x08, 0x80],
&[0x0a, 0x7f, 0x01],
&[0x79, 0x01],
&[0x7d, 0x01],
];
for bytes in malformed {
assert_eq!(check_decode_budget::<Outer>(bytes, 0).unwrap(), 0);
assert!(Outer::decode(*bytes).is_err(), "prost accepted {bytes:02x?}");
let mut after_valid_field = Outer {
items: vec![Inner::default()],
..Default::default()
}
.encode_to_vec();
after_valid_field.extend_from_slice(bytes);
assert!(Outer::decode(after_valid_field.as_slice()).is_err());
}
let mut bytes = Outer {
items: vec![Inner::default(); 2],
..Default::default()
}
.encode_to_vec();
bytes.extend_from_slice(&[0x0f, 0x0a, 0x00]);
assert_eq!(check_decode_budget::<Outer>(&bytes, 10).unwrap(), 2);
assert!(Outer::decode(bytes.as_slice()).is_err());
}
#[test]
fn it_bounds_the_depth_of_recursive_types() {
let mut node = Node::default();
for _ in 0..MAX_DEPTH {
node = Node {
child: Some(Box::new(node)),
};
}
assert_eq!(count::<Node>(&node, usize::MAX).unwrap(), MAX_DEPTH);
let node = Node {
child: Some(Box::new(node)),
};
count::<Node>(&node, usize::MAX).unwrap_err();
}
#[test]
fn scalar_payloads_count_nothing() {
assert_eq!(count::<u64>(&u64::MAX, 0).unwrap(), 0);
assert_eq!(count::<String>(&"\n\0\n\0".to_string(), 0).unwrap(), 0);
assert_eq!(count::<Vec<u8>>(&fake_messages(1_000), 0).unwrap(), 0);
assert_eq!(count::<()>(&(), 0).unwrap(), 0);
}
}