diff --git a/src/opaque.rs b/src/opaque.rs index 1623766..d98eacb 100644 --- a/src/opaque.rs +++ b/src/opaque.rs @@ -34,9 +34,14 @@ pub struct RegisterFirstMessage { impl TryFrom<&[u8]> for RegisterFirstMessage { type Error = ProtocolError; fn try_from(first_message_bytes: &[u8]) -> Result { + let checked_slice = check_slice_size( + first_message_bytes, + Grp::ElemLen::to_usize(), + "first_message_bytes", + )?; // Check that the message is actually containing an element of the // correct subgroup - let arr = GenericArray::from_slice(first_message_bytes); + let arr = GenericArray::from_slice(checked_slice); let alpha = Grp::from_element_slice(arr)?; Ok(Self { alpha }) } @@ -142,14 +147,24 @@ pub struct LoginFirstMessage { impl TryFrom<&[u8]> for LoginFirstMessage { type Error = ProtocolError; fn try_from(first_message_bytes: &[u8]) -> Result { + let min_expected_len = ::ElemLen::to_usize(); + let checked_slice = (if first_message_bytes.len() <= min_expected_len { + Err(InternalPakeError::SizeError { + name: "first_message_bytes", + len: min_expected_len, + actual_len: first_message_bytes.len(), + }) + } else { + Ok(first_message_bytes) + })?; // Check that the message is actually containing an element of the // correct subgroup let elem_len = ::ElemLen::to_usize(); - let arr = GenericArray::from_slice(&first_message_bytes[..elem_len]); + let arr = GenericArray::from_slice(&checked_slice[..elem_len]); let alpha = CS::Group::from_element_slice(arr)?; let ke1_message = >::KE1Message::try_from( - first_message_bytes[elem_len..].to_vec(), + checked_slice[elem_len..].to_vec(), )?; Ok(Self { alpha, ke1_message }) } @@ -275,12 +290,23 @@ pub struct ClientRegistration { impl TryFrom<&[u8]> for ClientRegistration { type Error = ProtocolError; fn try_from(bytes: &[u8]) -> Result { + let min_expected_len = ::ScalarLen::to_usize(); + let checked_slice = (if bytes.len() <= min_expected_len { + Err(InternalPakeError::SizeError { + name: "client_registration_bytes", + len: min_expected_len, + actual_len: bytes.len(), + }) + } else { + Ok(bytes) + })?; + // Check that the message is actually containing an element of the // correct subgroup - let scalar_len = ::ScalarLen::to_usize(); - let blinding_factor_bytes = GenericArray::from_slice(&bytes[..scalar_len]); + let scalar_len = min_expected_len; + let blinding_factor_bytes = GenericArray::from_slice(&checked_slice[..scalar_len]); let blinding_factor = CS::Group::from_scalar_slice(blinding_factor_bytes)?; - let password = bytes[scalar_len..].to_vec(); + let password = checked_slice[scalar_len..].to_vec(); Ok(Self { blinding_factor, password, @@ -629,11 +655,23 @@ impl TryFrom<&[u8]> for ClientLogin { type Error = ProtocolError; fn try_from(bytes: &[u8]) -> Result { let scalar_len = ::ScalarLen::to_usize(); - let blinding_factor_bytes = GenericArray::from_slice(&bytes[..scalar_len]); - let blinding_factor = CS::Group::from_scalar_slice(blinding_factor_bytes)?; let ke1_state_size = >::ke1_state_size(); + + let min_expected_len = scalar_len + ke1_state_size; + let checked_slice = (if bytes.len() <= min_expected_len { + Err(InternalPakeError::SizeError { + name: "client_login_bytes", + len: min_expected_len, + actual_len: bytes.len(), + }) + } else { + Ok(bytes) + })?; + + let blinding_factor_bytes = GenericArray::from_slice(&checked_slice[..scalar_len]); + let blinding_factor = CS::Group::from_scalar_slice(blinding_factor_bytes)?; let ke1_state = >::KE1State::try_from( - bytes[scalar_len..scalar_len + ke1_state_size].to_vec(), + checked_slice[scalar_len..scalar_len + ke1_state_size].to_vec(), )?; let password = bytes[scalar_len + ke1_state_size..].to_vec(); Ok(Self { diff --git a/src/tests/serialization.rs b/src/tests/serialization.rs index fad8057..47ead1e 100644 --- a/src/tests/serialization.rs +++ b/src/tests/serialization.rs @@ -17,6 +17,7 @@ use crate::{ use curve25519_dalek::ristretto::RistrettoPoint; use generic_array::typenum::Unsigned; +use proptest::{collection::vec, prelude::*}; use rand_core::{OsRng, RngCore}; use sha2::{Digest, Sha256}; @@ -171,3 +172,57 @@ fn login_first_message_roundtrip() { let reg_bytes = reg.to_bytes(); assert_eq!(reg_bytes, ke1m); } + +proptest! { + + #[test] + fn test_nocrash_register_first_message(bytes in vec(any::(), 0..200)) { + RegisterFirstMessage::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_register_second_message(bytes in vec(any::(), 0..200)) { + RegisterSecondMessage::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_register_third_message(bytes in vec(any::(), 0..200)) { + RegisterThirdMessage::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_login_first_message(bytes in vec(any::(), 0..500)) { + LoginFirstMessage::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_login_second_message(bytes in vec(any::(), 0..500)) { + LoginSecondMessage::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_login_third_message(bytes in vec(any::(), 0..500)) { + LoginThirdMessage::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_client_registration(bytes in vec(any::(), 0..700)) { + ClientRegistration::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_server_registration(bytes in vec(any::(), 0..700)) { + ServerRegistration::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_client_login(bytes in vec(any::(), 0..700)) { + ClientLogin::::try_from(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_server_login(bytes in vec(any::(), 0..700)) { + ServerLogin::::try_from(&bytes[..]).map_or(true, |_| true); + } + +}