diff --git a/der/src/reader/slice.rs b/der/src/reader/slice.rs index 4e5a7c4d5..3340cb380 100644 --- a/der/src/reader/slice.rs +++ b/der/src/reader/slice.rs @@ -99,8 +99,20 @@ impl<'a> Reader<'a> for SliceReader<'a> { let resumption = self.position.split_nested(len)?; let ret = f(self); + let finished = self.is_finished(); + let decoded = self.position.current(); + let remaining = self.remaining_len(); + self.bytes = bytes; self.position.resume_nested(resumption); + + if ret.is_ok() && !finished { + self.failed = true; + return Err(self + .error(ErrorKind::TrailingData { decoded, remaining }) + .into()); + }; + ret } @@ -153,10 +165,10 @@ impl<'a> Reader<'a> for SliceReader<'a> { } #[cfg(test)] -#[allow(clippy::unwrap_used, clippy::panic)] +#[allow(clippy::unwrap_used, clippy::panic, reason = "tests")] mod tests { use super::SliceReader; - use crate::{Decode, ErrorKind, Length, Reader}; + use crate::{Decode, Error, ErrorKind, Length, Reader}; use hex_literal::hex; // INTEGER: 42 @@ -217,4 +229,26 @@ mod tests { err.kind() ); } + + #[test] + fn nested_trailing_data() { + let der = hex!("0102"); + let mut reader = SliceReader::new(&der).unwrap(); + + let err: Error = reader + .read_nested(2u8.into(), |reader| { + reader.read_slice(1u8.into())?; + Ok(()) + }) + .unwrap_err(); + + assert_eq!(Length::ONE, err.position().unwrap()); + assert_eq!( + ErrorKind::TrailingData { + decoded: 1u8.into(), + remaining: 1u8.into(), + }, + err.kind() + ); + } }