diff --git a/src/encoding.rs b/src/encoding.rs index 60b55204a..344d2dc86 100644 --- a/src/encoding.rs +++ b/src/encoding.rs @@ -199,7 +199,7 @@ pub struct DecodeContext { /// customized. The recursion limit can be ignored by building the Prost /// crate with the `no-recursion-limit` feature. #[cfg(not(feature = "no-recursion-limit"))] - recurse_count: u32, + pub recurse_count: u32, } #[cfg(not(feature = "no-recursion-limit"))] diff --git a/src/message.rs b/src/message.rs index a190f6b47..76a56bed3 100644 --- a/src/message.rs +++ b/src/message.rs @@ -116,6 +116,20 @@ pub trait Message: Debug + Send + Sync { Self::merge(&mut message, &mut buf).map(|_| message) } + /// Decodes an instance of the message from a buffer. + /// + /// The entire buffer will be consumed. + /// + /// This is an alternative to `decode`, which allows users to specify a `DecodeContext`. + fn decode_with_context(mut buf: B, ctx: DecodeContext) -> Result + where + B: Buf, + Self: Default, + { + let mut message = Self::default(); + Self::merge_with_context(&mut message, &mut buf, ctx).map(|_| message) + } + /// Decodes a length-delimited instance of the message from the buffer. fn decode_length_delimited(buf: B) -> Result where @@ -130,12 +144,22 @@ pub trait Message: Debug + Send + Sync { /// Decodes an instance of the message from a buffer, and merges it into `self`. /// /// The entire buffer will be consumed. - fn merge(&mut self, mut buf: B) -> Result<(), DecodeError> + fn merge(&mut self, buf: B) -> Result<(), DecodeError> + where + B: Buf, + Self: Sized, + { + self.merge_with_context(buf, DecodeContext::default()) + } + + /// Decodes an instance of the message from a buffer, and merges it into `self`. + /// + /// The entire buffer will be consumed. + fn merge_with_context(&mut self, mut buf: B, ctx: DecodeContext) -> Result<(), DecodeError> where B: Buf, Self: Sized, { - let ctx = DecodeContext::default(); while buf.has_remaining() { let (tag, wire_type) = decode_key(&mut buf)?; self.merge_field(tag, wire_type, &mut buf, ctx.clone())?; diff --git a/tests/src/lib.rs b/tests/src/lib.rs index a9818d682..55b38db36 100644 --- a/tests/src/lib.rs +++ b/tests/src/lib.rs @@ -405,6 +405,29 @@ mod tests { assert!(build_and_roundtrip(101).is_err()); } + #[test] + fn test_deep_nesting_with_custom_recursion_limit() { + fn build_and_roundtrip(depth: usize) -> Result<(), prost::DecodeError> { + use crate::nesting::A; + use prost::encoding::DecodeContext; + + let mut a = Box::new(A::default()); + for _ in 0..depth { + let mut next = Box::new(A::default()); + next.a = Some(a); + a = next; + } + + let mut buf = Vec::new(); + a.encode(&mut buf).unwrap(); + A::decode_with_context(&*buf, DecodeContext { recurse_count: 200 }).map(|_| ()) + } + + assert!(build_and_roundtrip(100).is_ok()); + assert!(build_and_roundtrip(200).is_ok()); + assert!(build_and_roundtrip(201).is_err()); + } + #[test] fn test_deep_nesting_oneof() { fn build_and_roundtrip(depth: usize) -> Result<(), prost::DecodeError> {