diff --git a/examples/digital_locker.rs b/examples/digital_locker.rs index 8107238..4a2fe4f 100644 --- a/examples/digital_locker.rs +++ b/examples/digital_locker.rs @@ -145,7 +145,7 @@ fn open_locker( ClientLoginStartParameters::default(), ) .unwrap(); - let credential_request_bytes = client_login_start_result.message.serialize(); + let credential_request_bytes = client_login_start_result.message.serialize().unwrap(); // Client sends credential_request_bytes to server @@ -160,7 +160,7 @@ fn open_locker( ServerLoginStartParameters::default(), ) .unwrap(); - let credential_response_bytes = server_login_start_result.message.serialize(); + let credential_response_bytes = server_login_start_result.message.serialize().unwrap(); // Server sends credential_response_bytes to client @@ -174,7 +174,7 @@ fn open_locker( return Err(String::from("Incorrect password, please try again.")); } let client_login_finish_result = result.unwrap(); - let credential_finalization_bytes = client_login_finish_result.message.serialize(); + let credential_finalization_bytes = client_login_finish_result.message.serialize().unwrap(); // Client sends credential_finalization_bytes to server diff --git a/examples/simple_login.rs b/examples/simple_login.rs index 44f9572..b291e02 100644 --- a/examples/simple_login.rs +++ b/examples/simple_login.rs @@ -99,7 +99,7 @@ fn account_login( ClientLoginStartParameters::default(), ) .unwrap(); - let credential_request_bytes = client_login_start_result.message.serialize(); + let credential_request_bytes = client_login_start_result.message.serialize().unwrap(); // Client sends credential_request_bytes to server @@ -113,7 +113,7 @@ fn account_login( ServerLoginStartParameters::default(), ) .unwrap(); - let credential_response_bytes = server_login_start_result.message.serialize(); + let credential_response_bytes = server_login_start_result.message.serialize().unwrap(); // Server sends credential_response_bytes to client @@ -127,7 +127,7 @@ fn account_login( return false; } let client_login_finish_result = result.unwrap(); - let credential_finalization_bytes = client_login_finish_result.message.serialize(); + let credential_finalization_bytes = client_login_finish_result.message.serialize().unwrap(); // Client sends credential_finalization_bytes to server diff --git a/src/envelope.rs b/src/envelope.rs index 50a909b..86d5f7c 100644 --- a/src/envelope.rs +++ b/src/envelope.rs @@ -177,8 +177,8 @@ impl Envelope { client_s_sk: &[u8], server_s_pk: &[u8], optional_ids: Option<(Vec, Vec)>, - ) -> Result<(Self, GenericArray::OutputSize>), InternalPakeError> { - let aad = construct_aad(server_s_pk, &optional_ids); + ) -> Result<(Self, GenericArray::OutputSize>), ProtocolError> { + let aad = construct_aad(server_s_pk, &optional_ids)?; Self::seal_raw(rng, key, &client_s_sk, &aad, mode_from_ids(&optional_ids)) } @@ -190,7 +190,7 @@ impl Envelope { plaintext: &[u8], aad: &[u8], mode: InnerEnvelopeMode, - ) -> Result<(Self, GenericArray::OutputSize>), InternalPakeError> { + ) -> Result<(Self, GenericArray::OutputSize>), ProtocolError> { let mut nonce = vec![0u8; NONCE_LEN]; rng.fill_bytes(&mut nonce); @@ -239,18 +239,18 @@ impl Envelope { key: &[u8], server_s_pk: &[u8], optional_ids: &Option<(Vec, Vec)>, - ) -> Result, InternalPakeError> { + ) -> Result, ProtocolError> { // First, check that mode matches if self.inner_envelope.mode != mode_from_ids(optional_ids) { - return Err(InternalPakeError::IncompatibleEnvelopeModeError); + return Err(InternalPakeError::IncompatibleEnvelopeModeError.into()); } - let aad = construct_aad(server_s_pk, optional_ids); + let aad = construct_aad(server_s_pk, optional_ids)?; let opened = self.open_raw(key, &aad)?; if opened.plaintext.len() != ::Len::to_usize() { // Plaintext should consist of a single key - return Err(InternalPakeError::UnexpectedEnvelopeContentsError); + return Err(InternalPakeError::UnexpectedEnvelopeContentsError.into()); } Ok(OpenedEnvelope { @@ -325,12 +325,15 @@ impl Drop for Envelope { // Helper functions -fn construct_aad(server_s_pk: &[u8], optional_ids: &Option<(Vec, Vec)>) -> Vec { - let ids = optional_ids - .iter() - .flat_map(|(l, r)| [serialize(l, 2), serialize(r, 2)].concat()) - .collect(); - [server_s_pk.to_vec(), ids].concat() +fn construct_aad( + server_s_pk: &[u8], + optional_ids: &Option<(Vec, Vec)>, +) -> Result, ProtocolError> { + let ids = match optional_ids { + Some((l, r)) => [serialize(l, 2)?, serialize(r, 2)?].concat(), + None => vec![], + }; + Ok([server_s_pk.to_vec(), ids].concat()) } pub(crate) fn mode_from_ids(optional_ids: &Option<(Vec, Vec)>) -> InnerEnvelopeMode { diff --git a/src/key_exchange/traits.rs b/src/key_exchange/traits.rs index e1df71a..e9b356b 100644 --- a/src/key_exchange/traits.rs +++ b/src/key_exchange/traits.rs @@ -61,11 +61,11 @@ pub trait KeyExchange { } pub trait ToBytes { - fn to_bytes(&self) -> Vec; + fn to_bytes(&self) -> Result, ProtocolError>; } pub trait ToBytesWithPointers { - fn to_bytes(&self) -> Vec; + fn to_bytes(&self) -> Result, ProtocolError>; // Only used for tests to grab raw pointers to data #[cfg(test)] diff --git a/src/key_exchange/tripledh.rs b/src/key_exchange/tripledh.rs index a3326de..b4ed3c4 100644 --- a/src/key_exchange/tripledh.rs +++ b/src/key_exchange/tripledh.rs @@ -89,9 +89,9 @@ impl KeyExchange for TripleDH { let mut transcript_hasher = D::new() .chain(STR_3DH) - .chain(&serialize(&id_u, 2)) + .chain(&serialize(&id_u, 2)?) .chain(&serialized_credential_request[..]) - .chain(&serialize(&id_s, 2)) + .chain(&serialize(&id_s, 2)?) .chain(&l2_bytes[..]) .chain(&server_nonce[..]) .chain(&server_e_kp.public().to_arr()); @@ -119,7 +119,7 @@ impl KeyExchange for TripleDH { .map(|(&x1, &x2)| x1 ^ x2) .collect(); - transcript_hasher.update(&serialize(&ciphertext, 2)); + transcript_hasher.update(&serialize(&ciphertext, 2)?); let mut mac_hasher = Hmac::::new_varkey(&km2).map_err(|_| InternalPakeError::HmacError)?; @@ -157,9 +157,9 @@ impl KeyExchange for TripleDH { ) -> Result<(Vec, Vec, Self::KE3Message), ProtocolError> { let mut transcript_hasher = D::new() .chain(STR_3DH) - .chain(&serialize(&id_u, 2)) + .chain(&serialize(&id_u, 2)?) .chain(&serialized_credential_request) - .chain(&serialize(&id_s, 2)) + .chain(&serialize(&id_s, 2)?) .chain(&l2_component[..]) .chain(&ke2_message.to_bytes_without_info_or_mac()); @@ -175,7 +175,7 @@ impl KeyExchange for TripleDH { &transcript_hasher.clone().finalize(), )?; - transcript_hasher.update(&serialize(&ke2_message.e_info[..], 2)); + transcript_hasher.update(&serialize(&ke2_message.e_info[..], 2)?); let mut server_mac = Hmac::::new_varkey(&km2).map_err(|_| InternalPakeError::HmacError)?; @@ -269,9 +269,9 @@ impl TryFrom<&[u8]> for Ke1State { } impl ToBytesWithPointers for Ke1State { - fn to_bytes(&self) -> Vec { + fn to_bytes(&self) -> Result, ProtocolError> { let output: Vec = [&self.client_e_sk.to_arr(), &self.client_nonce[..]].concat(); - output + Ok(output) } #[cfg(test)] @@ -287,13 +287,13 @@ impl ToBytesWithPointers for Ke1State { } impl ToBytes for Ke1Message { - fn to_bytes(&self) -> Vec { - [ + fn to_bytes(&self) -> Result, ProtocolError> { + Ok([ &self.client_nonce[..], - &serialize(&self.info, 2), + &serialize(&self.info, 2)?, &self.client_e_pk.to_arr(), ] - .concat() + .concat()) } } @@ -339,13 +339,13 @@ impl> Drop for Ke2State { } impl> ToBytesWithPointers for Ke2State { - fn to_bytes(&self) -> Vec { - [ + fn to_bytes(&self) -> Result, ProtocolError> { + Ok([ &self.km3[..], &self.hashed_transcript[..], &self.session_key[..], ] - .concat() + .concat()) } #[cfg(test)] @@ -384,13 +384,13 @@ impl> TryFrom<&[u8]> for Ke2State { } impl> ToBytes for Ke2Message { - fn to_bytes(&self) -> Vec { - [ + fn to_bytes(&self) -> Result, ProtocolError> { + Ok([ &self.to_bytes_without_info_or_mac(), - &serialize(&self.e_info, 2), + &serialize(&self.e_info, 2)?, &self.mac[..], ] - .concat() + .concat()) } } @@ -449,8 +449,8 @@ pub struct Ke3Message> { } impl> ToBytes for Ke3Message { - fn to_bytes(&self) -> Vec { - self.mac.to_vec() + fn to_bytes(&self) -> Result, ProtocolError> { + Ok(self.mac.to_vec()) } } @@ -546,9 +546,9 @@ fn hkdf_expand_label_extracted( let mut opaque_label: Vec = Vec::new(); opaque_label.extend_from_slice(&STR_OPAQUE); opaque_label.extend_from_slice(&label); - hkdf_label.extend_from_slice(&serialize(&opaque_label, 1)); + hkdf_label.extend_from_slice(&serialize(&opaque_label, 1)?); - hkdf_label.extend_from_slice(&serialize(&context, 1)); + hkdf_label.extend_from_slice(&serialize(&context, 1)?); hkdf.expand(&hkdf_label, &mut okm) .map_err(|_| InternalPakeError::HkdfError)?; diff --git a/src/map_to_curve.rs b/src/map_to_curve.rs index 3875b88..a39be1a 100644 --- a/src/map_to_curve.rs +++ b/src/map_to_curve.rs @@ -6,7 +6,7 @@ //! Defines the GroupWithMapToCurve trait to specify how to map a password to a //! curve point -use crate::errors::InternalPakeError; +use crate::errors::{InternalPakeError, ProtocolError}; use crate::group::Group; use crate::hash::Hash; use crate::serialization::i2osp; @@ -22,12 +22,12 @@ pub trait GroupWithMapToCurve: Group { const SUITE_ID: usize; /// transforms a password and domain separation tag (DST) into a curve point - fn map_to_curve(msg: &[u8], dst: &[u8]) -> Result; + fn map_to_curve(msg: &[u8], dst: &[u8]) -> Result; /// Generates the contextString parameter as defined in /// - fn get_context_string(mode: u8) -> Vec { - [i2osp(mode as usize, 1), i2osp(Self::SUITE_ID, 2)].concat() + fn get_context_string(mode: u8) -> Result, ProtocolError> { + Ok([i2osp(mode as usize, 1)?, i2osp(Self::SUITE_ID, 2)?].concat()) } } @@ -36,7 +36,7 @@ impl GroupWithMapToCurve for RistrettoPoint { // Implements the hash_to_ristretto255() function from // https://www.ietf.org/archive/id/draft-irtf-cfrg-hash-to-curve-10.txt - fn map_to_curve(msg: &[u8], dst: &[u8]) -> Result { + fn map_to_curve(msg: &[u8], dst: &[u8]) -> Result { let uniform_bytes = expand_message_xmd::(msg, dst, ::OutputSize::to_usize())?; Ok(::hash_to_curve( @@ -65,24 +65,24 @@ pub(crate) fn expand_message_xmd( msg: &[u8], dst: &[u8], len_in_bytes: usize, -) -> Result, InternalPakeError> { +) -> Result, ProtocolError> { let b_in_bytes = ::OutputSize::to_usize(); let r_in_bytes = ::BlockSize::to_usize(); let ell = div_ceil(len_in_bytes, b_in_bytes); if ell > 255 { - return Err(InternalPakeError::HashToCurveError); + return Err(InternalPakeError::HashToCurveError.into()); } - let dst_prime = [dst, &i2osp(dst.len(), 1)].concat(); - let z_pad = i2osp(0, r_in_bytes); - let l_i_b_str = i2osp(len_in_bytes, 2); - let msg_prime = [&z_pad, msg, &l_i_b_str, &i2osp(0, 1), &dst_prime].concat(); + let dst_prime = [dst, &i2osp(dst.len(), 1)?].concat(); + let z_pad = i2osp(0, r_in_bytes)?; + let l_i_b_str = i2osp(len_in_bytes, 2)?; + let msg_prime = [&z_pad, msg, &l_i_b_str, &i2osp(0, 1)?, &dst_prime].concat(); let mut b: Vec> = vec![H::digest(&msg_prime).to_vec()]; // b[0] let mut h = H::new(); h.update(&b[0]); - h.update(&i2osp(1, 1)); + h.update(&i2osp(1, 1)?); h.update(&dst_prime); b.push(h.finalize_reset().to_vec()); // b[1] @@ -91,7 +91,7 @@ pub(crate) fn expand_message_xmd( for i in 2..(ell + 1) { h.update(xor(&b[0], &b[i - 1])?); - h.update(&i2osp(i, 1)); + h.update(&i2osp(i, 1)?); h.update(&dst_prime); b.push(h.finalize_reset().to_vec()); // b[i] uniform_bytes.extend_from_slice(&b[i]); diff --git a/src/messages.rs b/src/messages.rs index 692ec4a..083fad3 100644 --- a/src/messages.rs +++ b/src/messages.rs @@ -131,11 +131,11 @@ pub struct CredentialRequest { impl CredentialRequest { /// Serialization into bytes - pub fn serialize(&self) -> Vec { + pub fn serialize(&self) -> Result, ProtocolError> { let mut credential_request: Vec = Vec::new(); credential_request.extend_from_slice(&self.alpha.to_arr()); - credential_request.extend_from_slice(&self.ke1_message.to_bytes()); - credential_request + credential_request.extend_from_slice(&self.ke1_message.to_bytes()?); + Ok(credential_request) } /// Deserialization from bytes @@ -172,12 +172,12 @@ pub struct CredentialResponse { impl CredentialResponse { /// Serialization into bytes - pub fn serialize(&self) -> Vec { - [ + pub fn serialize(&self) -> Result, ProtocolError> { + Ok([ Self::serialize_without_ke(&self.beta, &self.server_s_pk, &self.envelope), - self.ke2_message.to_bytes(), + self.ke2_message.to_bytes()?, ] - .concat() + .concat()) } pub(crate) fn serialize_without_ke( @@ -237,7 +237,7 @@ pub struct CredentialFinalization { impl CredentialFinalization { /// Serialization into bytes - pub fn serialize(&self) -> Vec { + pub fn serialize(&self) -> Result, ProtocolError> { self.ke3_message.to_bytes() } diff --git a/src/opaque.rs b/src/opaque.rs index 35c6f6d..efb960d 100644 --- a/src/opaque.rs +++ b/src/opaque.rs @@ -398,15 +398,15 @@ pub struct ClientLogin { impl ClientLogin { /// Serialization into bytes - pub fn serialize(&self) -> Vec { + pub fn serialize(&self) -> Result, ProtocolError> { let output: Vec = [ &CS::Group::scalar_as_bytes(&self.token.blind)[..], - &serialize(&self.serialized_credential_request, 2), - &serialize(&self.ke1_state.to_bytes(), 2), + &serialize(&self.serialized_credential_request, 2)?, + &serialize(&self.ke1_state.to_bytes()?, 2)?, &self.token.data, ] .concat(); - output + Ok(output) } /// Deserialization from bytes @@ -541,7 +541,7 @@ impl ClientLogin { let (ke1_state, ke1_message) = CS::KeyExchange::generate_ke1(info, rng)?; let credential_request = CredentialRequest { alpha, ke1_message }; - let serialized_credential_request = credential_request.serialize(); + let serialized_credential_request = credential_request.serialize()?; Ok(ClientLoginStartResult { message: credential_request, @@ -605,8 +605,10 @@ impl ClientLogin { .envelope .open(&password_derived_key, &server_s_pk_bytes, &optional_ids) .map_err(|e| match e { - InternalPakeError::SealOpenHmacError => PakeError::InvalidLoginError, - err => PakeError::from(err), + ProtocolError::VerificationError(PakeError::CryptoError( + InternalPakeError::SealOpenHmacError, + )) => ProtocolError::VerificationError(PakeError::InvalidLoginError), + err => err, })?; let client_s_sk = Key::from_bytes(&opened_envelope.client_s_sk)?; @@ -693,7 +695,7 @@ pub struct ServerLoginFinishResult { impl ServerLogin { /// Serialization into bytes - pub fn serialize(&self) -> Vec { + pub fn serialize(&self) -> Result, ProtocolError> { self.ke2_state.to_bytes() } @@ -780,7 +782,7 @@ impl ServerLogin { Some((id_u, id_s)) => (id_u, id_s), }; - let l1_bytes = &l1.serialize(); + let l1_bytes = &l1.serialize()?; let beta = oprf::evaluate(l1.alpha, &password_file.oprf_key); let server_s_pk = KeyPair::::public_from_private(&server_s_sk); @@ -945,7 +947,7 @@ impl Drop for ServerLogin { fn get_password_derived_key, D: Hash>( token: &oprf::Token, beta: G, -) -> Result, InternalPakeError> { - let oprf_output = oprf::finalize::(&token.data, &token.blind, beta); - SH::hash(oprf_output) +) -> Result, ProtocolError> { + let oprf_output = oprf::finalize::(&token.data, &token.blind, beta)?; + SH::hash(oprf_output).map_err(ProtocolError::from) } diff --git a/src/oprf.rs b/src/oprf.rs index ad1e883..c701329 100644 --- a/src/oprf.rs +++ b/src/oprf.rs @@ -4,7 +4,7 @@ // LICENSE file in the root directory of this source tree. use crate::{ - errors::InternalPakeError, group::Group, hash::Hash, map_to_curve::GroupWithMapToCurve, + errors::ProtocolError, group::Group, hash::Hash, map_to_curve::GroupWithMapToCurve, serialization::serialize, }; use digest::Digest; @@ -28,9 +28,9 @@ static MODE_BASE: u8 = 0x00; pub(crate) fn blind( input: &[u8], blinding_factor_rng: &mut R, -) -> Result<(Token, G), InternalPakeError> { +) -> Result<(Token, G), ProtocolError> { let blind = G::random_scalar(blinding_factor_rng); - let dst = [STR_VOPRF, &G::get_context_string(MODE_BASE)].concat(); + let dst = [STR_VOPRF, &G::get_context_string(MODE_BASE)?].concat(); let mapped_point = G::map_to_curve::(input, &dst)?; let blind_token = mapped_point * &blind; Ok(( @@ -54,7 +54,7 @@ pub(crate) fn finalize( input: &[u8], blind: &G::Scalar, evaluated_element: G, -) -> GenericArray::OutputSize> { +) -> Result::OutputSize>, ProtocolError> { let unblinded_element = evaluated_element * &G::scalar_invert(blind); finalize_after_unblind::(input, unblinded_element) } @@ -62,15 +62,15 @@ pub(crate) fn finalize( fn finalize_after_unblind( input: &[u8], unblinded_element: G, -) -> GenericArray::OutputSize> { - let finalize_dst = [STR_VOPRF_FINALIZE, &G::get_context_string(MODE_BASE)].concat(); +) -> Result::OutputSize>, ProtocolError> { + let finalize_dst = [STR_VOPRF_FINALIZE, &G::get_context_string(MODE_BASE)?].concat(); let hash_input = [ - serialize(input, 2), - serialize(&unblinded_element.to_arr().to_vec(), 2), - serialize(&finalize_dst, 2), + serialize(input, 2)?, + serialize(&unblinded_element.to_arr().to_vec(), 2)?, + serialize(&finalize_dst, 2)?, ] .concat(); - ::digest(&hash_input) + Ok(::digest(&hash_input)) } //////////////////////// @@ -83,7 +83,7 @@ fn finalize_after_unblind( pub fn blind_shim( input: &[u8], blinding_factor_rng: &mut R, -) -> Result<(Token, G), InternalPakeError> { +) -> Result<(Token, G), ProtocolError> { blind::(input, blinding_factor_rng) } @@ -100,8 +100,8 @@ pub fn evaluate_shim(point: G, oprf_key: &G::Scalar) -> G { pub fn finalize_shim( token: &Token, point: G, -) -> Result::OutputSize>, InternalPakeError> { - Ok(finalize::(&token.data, &token.blind, point)) +) -> Result::OutputSize>, ProtocolError> { + finalize::(&token.data, &token.blind, point) } /////////// @@ -119,30 +119,34 @@ mod tests { use sha2::Sha512; fn prf(input: &[u8], oprf_key: &[u8; 32]) -> GenericArray::OutputSize> { - let dst = [STR_VOPRF, &RistrettoPoint::get_context_string(MODE_BASE)].concat(); + let dst = [ + STR_VOPRF, + &RistrettoPoint::get_context_string(MODE_BASE).unwrap(), + ] + .concat(); let point = RistrettoPoint::map_to_curve::(input, &dst).unwrap(); let scalar = RistrettoPoint::from_scalar_slice(GenericArray::from_slice(&oprf_key[..])).unwrap(); let res = point * scalar; - finalize_after_unblind::(&input, res) + finalize_after_unblind::(&input, res).unwrap() } #[test] - fn oprf_retrieval() -> Result<(), InternalPakeError> { + fn oprf_retrieval() { let input = b"hunter2"; let mut rng = OsRng; - let (token, alpha) = blind::<_, RistrettoPoint, Sha512>(&input[..], &mut rng)?; + let (token, alpha) = blind::<_, RistrettoPoint, Sha512>(&input[..], &mut rng).unwrap(); let oprf_key_bytes = arr![ u8; 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, ]; - let oprf_key = RistrettoPoint::from_scalar_slice(&oprf_key_bytes)?; + let oprf_key = RistrettoPoint::from_scalar_slice(&oprf_key_bytes).unwrap(); let beta = evaluate::(alpha, &oprf_key); - let res = finalize::(&token.data, &token.blind, beta); + let res = + finalize::(&token.data, &token.blind, beta).unwrap(); let res2 = prf(&input[..], &oprf_key.as_bytes()); assert_eq!(res, res2); - Ok(()) } #[test] @@ -151,11 +155,16 @@ mod tests { let mut input = vec![0u8; 64]; rng.fill_bytes(&mut input); let (token, alpha) = blind::<_, RistrettoPoint, sha2::Sha512>(&input, &mut rng).unwrap(); - let res = finalize::(&token.data, &token.blind, alpha); + let res = + finalize::(&token.data, &token.blind, alpha).unwrap(); - let dst = [STR_VOPRF, &RistrettoPoint::get_context_string(MODE_BASE)].concat(); + let dst = [ + STR_VOPRF, + &RistrettoPoint::get_context_string(MODE_BASE).unwrap(), + ] + .concat(); let point = RistrettoPoint::map_to_curve::(&input, &dst).unwrap(); - let res2 = finalize_after_unblind::(&input, point); + let res2 = finalize_after_unblind::(&input, point).unwrap(); assert_eq!(res, res2); } diff --git a/src/serialization/mod.rs b/src/serialization/mod.rs index e05462b..a0c75be 100644 --- a/src/serialization/mod.rs +++ b/src/serialization/mod.rs @@ -6,17 +6,24 @@ use crate::errors::PakeError; // Corresponds to the I2OSP() function from RFC8017 -pub(crate) fn i2osp(input: usize, length: usize) -> Vec { - if length <= std::mem::size_of::() { - return (&input.to_be_bytes()[std::mem::size_of::() - length..]).to_vec(); +pub(crate) fn i2osp(input: usize, length: usize) -> Result, PakeError> { + let sizeof_usize = std::mem::size_of::(); + + // Check if input >= 256^length + if (sizeof_usize as u32 - input.leading_zeros() / 8) > length as u32 { + return Err(PakeError::SerializationError); + } + + if length <= sizeof_usize { + return Ok((&input.to_be_bytes()[sizeof_usize - length..]).to_vec()); } let mut output = vec![0u8; length]; output.splice( - length - std::mem::size_of::()..length, + length - sizeof_usize..length, input.to_be_bytes().iter().cloned(), ); - output + Ok(output) } // Corresponds to the OS2IP() function from RFC8017 @@ -31,8 +38,8 @@ pub(crate) fn os2ip(input: &[u8]) -> Result { } // Computes I2OSP(len(input), max_bytes) || input -pub(crate) fn serialize(input: &[u8], max_bytes: usize) -> Vec { - [&i2osp(input.len(), max_bytes), input].concat() +pub(crate) fn serialize(input: &[u8], max_bytes: usize) -> Result, PakeError> { + Ok([&i2osp(input.len(), max_bytes)?, input].concat()) } // Tokenizes an input of the format I2OSP(len(input), max_bytes) || input, outputting @@ -55,3 +62,22 @@ pub(crate) fn tokenize(input: &[u8], size_bytes: usize) -> Result<(Vec, Vec< #[cfg(test)] mod tests; + +#[cfg(test)] +mod unit_tests { + use super::*; + + // Test the error condition for I2OSP + #[test] + fn test_i2osp_err_check() { + assert!(i2osp(0, 1).is_ok()); + + assert!(i2osp(255, 1).is_ok()); + assert!(i2osp(256, 1).is_err()); + assert!(i2osp(257, 1).is_err()); + + assert!(i2osp(256 * 256 - 1, 2).is_ok()); + assert!(i2osp(256 * 256, 2).is_err()); + assert!(i2osp(256 * 256 + 1, 2).is_err()); + } +} diff --git a/src/serialization/tests.rs b/src/serialization/tests.rs index 596639c..bef03bc 100644 --- a/src/serialization/tests.rs +++ b/src/serialization/tests.rs @@ -172,7 +172,7 @@ fn credential_request_roundtrip() { let ke1m: Vec = [ &client_nonce[..], - &serialize(&info.to_vec(), 2), + &serialize(&info.to_vec(), 2).unwrap(), &client_e_kp.public(), ] .concat(); @@ -182,7 +182,7 @@ fn credential_request_roundtrip() { input.extend_from_slice(&ke1m[..]); let l1 = CredentialRequest::::deserialize(input.as_slice()).unwrap(); - let l1_bytes = l1.serialize(); + let l1_bytes = l1.serialize().unwrap(); assert_eq!(input, l1_bytes); } @@ -222,7 +222,7 @@ fn credential_response_roundtrip() { let ke2m: Vec = [ &server_nonce[..], &server_e_kp.public(), - &serialize(&e_info.to_vec(), 2), + &serialize(&e_info.to_vec(), 2).unwrap(), &mac[..], ] .concat(); @@ -234,7 +234,7 @@ fn credential_response_roundtrip() { input.extend_from_slice(&ke2m[..]); let l2 = CredentialResponse::::deserialize(&input).unwrap(); - let l2_bytes = l2.serialize(); + let l2_bytes = l2.serialize().unwrap(); assert_eq!(input, l2_bytes); } @@ -247,7 +247,7 @@ fn login_third_message_roundtrip() { let input: Vec = [&mac[..]].concat(); let l3 = CredentialFinalization::::deserialize(&input).unwrap(); - let l3_bytes = l3.serialize(); + let l3_bytes = l3.serialize().unwrap(); assert_eq!(input, l3_bytes); } @@ -267,13 +267,13 @@ fn client_login_roundtrip() { // serialization order: scalar, credential_request, ke1_state, password let bytes: Vec = [ &sc.as_bytes()[..], - &serialize(&serialized_credential_request, 2), - &serialize(&l1_data, 2), + &serialize(&serialized_credential_request, 2).unwrap(), + &serialize(&l1_data, 2).unwrap(), &pw[..], ] .concat(); let reg = ClientLogin::::deserialize(&bytes[..]).unwrap(); - let reg_bytes = reg.serialize(); + let reg_bytes = reg.serialize().unwrap(); assert_eq!(reg_bytes, bytes); } @@ -290,14 +290,14 @@ fn ke1_message_roundtrip() { let ke1m: Vec = [ &client_nonce[..], - &serialize(&info.to_vec(), 2), + &serialize(&info.to_vec(), 2).unwrap(), &client_e_kp.public(), ] .concat(); let reg = >::KE1Message::try_from(&ke1m[..]) .unwrap(); - let reg_bytes = reg.to_bytes(); + let reg_bytes = reg.to_bytes().unwrap(); assert_eq!(reg_bytes, ke1m); } @@ -316,7 +316,7 @@ fn ke2_message_roundtrip() { let ke2m: Vec = [ &server_nonce[..], &server_e_kp.public(), - &serialize(&e_info.to_vec(), 2), + &serialize(&e_info.to_vec(), 2).unwrap(), &mac[..], ] .concat(); @@ -324,7 +324,7 @@ fn ke2_message_roundtrip() { let reg = >::KE2Message::try_from(&ke2m[..]) .unwrap(); - let reg_bytes = reg.to_bytes(); + let reg_bytes = reg.to_bytes().unwrap(); assert_eq!(reg_bytes, ke2m); } @@ -339,7 +339,7 @@ fn ke3_message_roundtrip() { let reg = >::KE3Message::try_from(&ke3m[..]) .unwrap(); - let reg_bytes = reg.to_bytes(); + let reg_bytes = reg.to_bytes().unwrap(); assert_eq!(reg_bytes, ke3m); } @@ -347,7 +347,7 @@ proptest! { #[test] fn test_i2osp_os2ip(bytes in vec(any::(), 0..std::mem::size_of::())) { - assert_eq!(i2osp(os2ip(&bytes)?, bytes.len()), bytes); + assert_eq!(i2osp(os2ip(&bytes)?, bytes.len())?, bytes); } #[test] diff --git a/src/tests/full_test.rs b/src/tests/full_test.rs index 85fc6f3..02ffd6d 100644 --- a/src/tests/full_test.rs +++ b/src/tests/full_test.rs @@ -358,8 +358,16 @@ fn generate_parameters() -> TestVectorParameters { ClientLoginStartParameters::WithInfo(info1.to_vec()), ) .unwrap(); - let credential_request_bytes = client_login_start_result.message.serialize().to_vec(); - let client_login_state = client_login_start_result.state.serialize().to_vec(); + let credential_request_bytes = client_login_start_result + .message + .serialize() + .unwrap() + .to_vec(); + let client_login_state = client_login_start_result + .state + .serialize() + .unwrap() + .to_vec(); let mut server_e_sk_and_nonce_rng = CycleRng::new( [ @@ -380,8 +388,16 @@ fn generate_parameters() -> TestVectorParameters { ), ) .unwrap(); - let credential_response_bytes = server_login_start_result.message.serialize().to_vec(); - let server_login_state = server_login_start_result.state.serialize().to_vec(); + let credential_response_bytes = server_login_start_result + .message + .serialize() + .unwrap() + .to_vec(); + let server_login_state = server_login_start_result + .state + .serialize() + .unwrap() + .to_vec(); let client_login_finish_result = client_login_start_result .state @@ -390,7 +406,7 @@ fn generate_parameters() -> TestVectorParameters { ClientLoginFinishParameters::WithIdentifiers(id_u.to_vec(), id_s.to_vec()), ) .unwrap(); - let credential_finalization_bytes = client_login_finish_result.message.serialize(); + let credential_finalization_bytes = client_login_finish_result.message.serialize().unwrap(); TestVectorParameters { client_s_pk: client_s_kp.public().to_arr().to_vec(), @@ -535,11 +551,11 @@ fn test_credential_request() -> Result<(), ProtocolError> { )?; assert_eq!( hex::encode(¶meters.credential_request), - hex::encode(client_login_start_result.message.serialize()) + hex::encode(client_login_start_result.message.serialize().unwrap()) ); assert_eq!( hex::encode(¶meters.client_login_state), - hex::encode(client_login_start_result.state.serialize()) + hex::encode(client_login_start_result.state.serialize()?) ); Ok(()) } @@ -569,11 +585,11 @@ fn test_credential_response() -> Result<(), ProtocolError> { ); assert_eq!( hex::encode(¶meters.credential_response), - hex::encode(server_login_start_result.message.serialize()) + hex::encode(server_login_start_result.message.serialize().unwrap()) ); assert_eq!( hex::encode(¶meters.server_login_state), - hex::encode(server_login_start_result.state.serialize()) + hex::encode(server_login_start_result.state.serialize().unwrap()) ); Ok(()) } @@ -606,7 +622,7 @@ fn test_credential_finalization() -> Result<(), ProtocolError> { ); assert_eq!( hex::encode(¶meters.credential_finalization), - hex::encode(client_login_finish_result.message.serialize()) + hex::encode(client_login_finish_result.message.serialize().unwrap()) ); assert_eq!( hex::encode(¶meters.export_key), diff --git a/src/tests/opaque_test_vectors.rs b/src/tests/opaque_test_vectors.rs index 0e9d0a3..ff68c9f 100644 --- a/src/tests/opaque_test_vectors.rs +++ b/src/tests/opaque_test_vectors.rs @@ -457,7 +457,7 @@ fn test_ke1() -> Result<(), ProtocolError> { )?; assert_eq!( hex::encode(¶meters.KE1), - hex::encode(client_login_start_result.message.serialize()) + hex::encode(client_login_start_result.message.serialize()?) ); } Ok(()) @@ -492,7 +492,7 @@ fn test_ke2() -> Result<(), ProtocolError> { ); assert_eq!( hex::encode(¶meters.KE2), - hex::encode(server_login_start_result.message.serialize()) + hex::encode(server_login_start_result.message.serialize()?) ); } Ok(()) @@ -536,7 +536,7 @@ fn test_ke3() -> Result<(), ProtocolError> { ); assert_eq!( hex::encode(¶meters.KE3), - hex::encode(client_login_finish_result.message.serialize()) + hex::encode(client_login_finish_result.message.serialize()?) ); assert_eq!( hex::encode(¶meters.export_key), diff --git a/src/tests/voprf_test_vectors.rs b/src/tests/voprf_test_vectors.rs index 7ad5274..e8368f6 100644 --- a/src/tests/voprf_test_vectors.rs +++ b/src/tests/voprf_test_vectors.rs @@ -63,7 +63,7 @@ fn populate_test_vectors(values: &Value) -> VOPRFTestVectorParameters { // Tests input -> blind, blinded_element #[test] -fn test_blind() -> Result<(), PakeError> { +fn test_blind() -> Result<(), ProtocolError> { for tv in OPRF_RISTRETTO255_SHA512 { let parameters = populate_test_vectors(&serde_json::from_str(tv).unwrap()); let mut rng = CycleRng::new(parameters.blind.to_vec()); @@ -106,7 +106,7 @@ fn test_evaluate() -> Result<(), PakeError> { // Tests input, blind, evaluation_element -> output #[test] -fn test_finalize() -> Result<(), PakeError> { +fn test_finalize() -> Result<(), ProtocolError> { for tv in OPRF_RISTRETTO255_SHA512 { let parameters = populate_test_vectors(&serde_json::from_str(tv).unwrap()); @@ -116,7 +116,7 @@ fn test_finalize() -> Result<(), PakeError> { RistrettoPoint::from_element_slice(GenericArray::from_slice( ¶meters.evaluation_element, ))?, - ); + )?; assert_eq!(¶meters.output, &output.to_vec()); }