diff --git a/crates/ironrdp-pdu/src/mcs.rs b/crates/ironrdp-pdu/src/mcs.rs index 03bb2ebe..f5d225ce 100644 --- a/crates/ironrdp-pdu/src/mcs.rs +++ b/crates/ironrdp-pdu/src/mcs.rs @@ -225,6 +225,7 @@ impl DomainMcsPdu { } fn read_mcspdu_header(src: &mut ReadCursor<'_>, name: &'static str) -> crate::Result { + ensure_size!(name: name, in: src, size: 1); let choice = src.read_u8(); DomainMcsPdu::from_choice(choice).ok_or(crate::Error::InvalidMessage { @@ -235,6 +236,7 @@ fn read_mcspdu_header(src: &mut ReadCursor<'_>, name: &'static str) -> crate::Re } fn peek_mcspdu_header(src: &mut ReadCursor<'_>, name: &'static str) -> crate::Result { + ensure_size!(name: name, in: src, size: 1); let choice = src.peek_u8(); DomainMcsPdu::from_choice(choice).ok_or(crate::Error::InvalidMessage { @@ -580,6 +582,8 @@ impl<'de> McsPdu<'de> for SendDataRequest<'de> { } fn mcs_body_decode(src: &mut ReadCursor<'de>, tpdu_user_data_size: usize) -> Result { + let src_len_before = src.len(); + read_mcspdu_header(src, Self::MCS_NAME)?.check_expected(Self::MCS_NAME, DomainMcsPdu::SendDataRequest)?; let initiator_id = per::read_u16(src, BASE_CHANNEL_ID)?; @@ -587,9 +591,21 @@ impl<'de> McsPdu<'de> for SendDataRequest<'de> { let _data_priority_and_segmentation = src.read_u8(); - let (length, _) = per::read_length(src); - debug_assert!(usize::from(length) <= tpdu_user_data_size - 7); - let user_data = Cow::Borrowed(src.read_slice(usize::from(length))); + let (length, _) = per::read_length(src)?; + let length = usize::from(length); + + let src_len_after = src.len(); + + if length > tpdu_user_data_size.saturating_sub(src_len_before - src_len_after) { + return Err(crate::Error::InvalidMessage { + name: Self::MCS_NAME, + field: "user-data-length", + reason: "inconsistent with user data size advertised in TPDU", + }); + } + + ensure_size!(name: Self::MCS_NAME, in: src, size: length); + let user_data = Cow::Borrowed(src.read_slice(length)); Ok(Self { initiator_id, @@ -645,6 +661,8 @@ impl<'de> McsPdu<'de> for SendDataIndication<'de> { } fn mcs_body_decode(src: &mut ReadCursor<'de>, tpdu_user_data_size: usize) -> Result { + let src_len_before = src.len(); + read_mcspdu_header(src, Self::MCS_NAME)?.check_expected(Self::MCS_NAME, DomainMcsPdu::SendDataIndication)?; let initiator_id = per::read_u16(src, BASE_CHANNEL_ID)?; @@ -652,9 +670,21 @@ impl<'de> McsPdu<'de> for SendDataIndication<'de> { let _data_priority_and_segmentation = src.read_u8(); - let (length, _) = per::read_length(src); - debug_assert!(usize::from(length) <= tpdu_user_data_size - 7); - let user_data = Cow::Borrowed(src.read_slice(usize::from(length))); + let (length, _) = per::read_length(src)?; + let length = usize::from(length); + + let src_len_after = src.len(); + + if length > tpdu_user_data_size.saturating_sub(src_len_before - src_len_after) { + return Err(crate::Error::InvalidMessage { + name: Self::MCS_NAME, + field: "user-data-length", + reason: "inconsistent with user data size advertised in TPDU", + }); + } + + ensure_size!(name: Self::MCS_NAME, in: src, size: length); + let user_data = Cow::Borrowed(src.read_slice(length)); Ok(Self { initiator_id, diff --git a/crates/ironrdp-pdu/src/per.rs b/crates/ironrdp-pdu/src/per.rs index 01dc026d..03a254da 100644 --- a/crates/ironrdp-pdu/src/per.rs +++ b/crates/ironrdp-pdu/src/per.rs @@ -8,16 +8,18 @@ pub(crate) const U16_SIZE: usize = 2; const OBJECT_ID_SIZE: usize = 6; -pub(crate) fn read_length(src: &mut ReadCursor<'_>) -> (u16, usize) { +pub(crate) fn read_length(src: &mut ReadCursor<'_>) -> crate::Result<(u16, usize)> { + ensure_size!(name: "PER LENGTH", in: src, size: 1); let a = src.read_u8(); if a & 0x80 != 0 { + ensure_size!(name: "PER LENGTH", in: src, size: 1); let b = src.read_u8(); let length = ((u16::from(a) & !0x80) << 8) + u16::from(b); - (length, 2) + Ok((length, 2)) } else { - (u16::from(a), 1) + Ok((u16::from(a), 1)) } } @@ -82,13 +84,22 @@ pub(crate) fn write_padding(dst: &mut WriteCursor<'_>, padding_length: usize) { } pub(crate) fn read_u32(src: &mut ReadCursor<'_>) -> crate::Result { - let (length, _) = read_length(src); + let (length, _) = read_length(src)?; match length { 0 => Ok(0), - 1 => Ok(u32::from(src.read_u8())), - 2 => Ok(u32::from(src.read_u16_be())), - 4 => Ok(src.read_u32_be()), + 1 => { + ensure_size!(name: "PER U32", in: src, size: 1); + Ok(u32::from(src.read_u8())) + } + 2 => { + ensure_size!(name: "PER U32", in: src, size: 2); + Ok(u32::from(src.read_u16_be())) + } + 4 => { + ensure_size!(name: "PER U32", in: src, size: 4); + Ok(src.read_u32_be()) + } _ => Err(crate::Error::Other { context: "PER", reason: "invalid length for u32", @@ -110,6 +121,7 @@ pub(crate) fn write_u32(dst: &mut WriteCursor<'_>, value: u32) { } pub(crate) fn read_u16(src: &mut ReadCursor<'_>, min: u16) -> crate::Result { + ensure_size!(name: "PER U16", in: src, size: 2); min.checked_add(src.read_u16_be()).ok_or(crate::Error::Other { context: "PER", reason: "invalid u16", @@ -129,6 +141,7 @@ pub(crate) fn write_u16(dst: &mut WriteCursor<'_>, value: u16, min: u16) -> crat } pub(crate) fn read_enum(src: &mut ReadCursor<'_>, count: u8) -> crate::Result { + ensure_size!(name: "PER ENUM", in: src, size: 1); let enumerated = src.read_u8(); if u16::from(enumerated) + 1 > u16::from(count) { @@ -146,7 +159,7 @@ pub(crate) fn write_enum(dst: &mut WriteCursor<'_>, enumerated: u8) { } pub(crate) fn read_object_id(src: &mut ReadCursor<'_>) -> crate::Result<[u8; OBJECT_ID_SIZE]> { - let (length, _) = read_length(src); + let (length, _) = read_length(src)?; if length != 5 { return Err(crate::Error::Other { @@ -155,6 +168,8 @@ pub(crate) fn read_object_id(src: &mut ReadCursor<'_>) -> crate::Result<[u8; OBJ }); } + ensure_size!(name: "PER OID", in: src, size: 5); + let first_two_tuples = src.read_u8(); let mut read_object_ids = [0u8; OBJECT_ID_SIZE]; @@ -178,9 +193,11 @@ pub(crate) fn write_object_id(dst: &mut WriteCursor<'_>, object_ids: [u8; OBJECT } } -pub(crate) fn read_octet_string(src: &mut ReadCursor<'_>, min: usize) -> Vec { - let (length, _) = read_length(src); - src.read_slice(min + usize::from(length)).to_owned() +pub(crate) fn read_octet_string(src: &mut ReadCursor<'_>, min: usize) -> crate::Result> { + let (length, _) = read_length(src)?; + let read_len = min + usize::from(length); + ensure_size!(name: "PER OCTET_STRING", in: src, size: read_len); + Ok(src.read_slice(read_len).to_owned()) } pub(crate) fn write_octet_string(dst: &mut WriteCursor<'_>, octet_string: &[u8], min: usize) { @@ -194,10 +211,12 @@ pub(crate) fn write_octet_string(dst: &mut WriteCursor<'_>, octet_string: &[u8], dst.write_slice(octet_string); } -pub(crate) fn read_numeric_string(src: &mut ReadCursor<'_>, min: u16) { - let (length, _) = read_length(src); - let length = (length + min + 1) / 2; - src.advance(usize::from(length)) +pub(crate) fn read_numeric_string(src: &mut ReadCursor<'_>, min: u16) -> crate::Result<()> { + let (length, _) = read_length(src)?; + let length = usize::from((length + min + 1) / 2); + ensure_size!(name: "PER NUMERIC_STRING", in: src, size: length); + src.advance(length); + Ok(()) } pub(crate) fn write_numeric_string(dst: &mut WriteCursor<'_>, num_str: &[u8], min: usize) { @@ -466,7 +485,7 @@ mod tests { fn read_length_is_correct_length() { let mut src = ReadCursor::new(&[0x05]); - let (length, sizeof_length) = read_length(&mut src); + let (length, sizeof_length) = read_length(&mut src).unwrap(); assert_eq!(5, length); assert_eq!(src.len(), 0); @@ -477,7 +496,7 @@ mod tests { fn read_length_is_correct_long_length() { let mut src = ReadCursor::new(&[0x80, 0x8d]); - let (length, sizeof_length) = read_length(&mut src); + let (length, sizeof_length) = read_length(&mut src).unwrap(); assert_eq!(141, length); assert_eq!(src.len(), 0); @@ -706,7 +725,7 @@ mod tests { let buf = [0x00, 0x10]; let mut src = ReadCursor::new(&buf); - read_numeric_string(&mut src, 1); + read_numeric_string(&mut src, 1).unwrap(); } #[test] @@ -728,7 +747,7 @@ mod tests { let buf = [0x00, 0x44, 0x75, 0x63, 0x61]; let mut src = ReadCursor::new(&buf); - assert_eq!(b"Duca", read_octet_string(&mut src, 4).as_slice()); + assert_eq!(b"Duca", read_octet_string(&mut src, 4).unwrap().as_slice()); } #[test] diff --git a/crates/ironrdp-testsuite-core/tests/pdu/regression.rs b/crates/ironrdp-testsuite-core/tests/pdu/regression.rs index 8b137891..ef7bc632 100644 --- a/crates/ironrdp-testsuite-core/tests/pdu/regression.rs +++ b/crates/ironrdp-testsuite-core/tests/pdu/regression.rs @@ -1 +1,11 @@ +const REGRESSION_DATA_FOLDER: &str = concat!(env!("CARGO_MANIFEST_DIR"), "/test_data/pdu/regression/"); +#[test] +pub fn check() { + for entry in std::fs::read_dir(REGRESSION_DATA_FOLDER).unwrap() { + let entry = entry.unwrap(); + println!("Check {}", entry.path().display()); + let test_case = std::fs::read(entry.path()).unwrap(); + ironrdp_fuzzing::oracles::pdu_decode(&test_case); + } +}