From 98f1821897cd2800e5bffb2a70541056145e99cc Mon Sep 17 00:00:00 2001 From: Kevin Lewi Date: Tue, 15 Jun 2021 17:31:12 -0700 Subject: [PATCH] Adding identity element checks and ensuring non-zero scalar selection --- src/errors.rs | 2 ++ src/group.rs | 45 ++++++++++++++++++--------- src/keypair.rs | 3 +- src/messages.rs | 30 ++++++++++++++++++ src/opaque.rs | 2 +- src/oprf.rs | 3 +- src/serialization/tests.rs | 63 +++++++++++++++++++++++++++++++++++--- src/tests/full_test.rs | 24 +++++++++++++-- 8 files changed, 148 insertions(+), 24 deletions(-) diff --git a/src/errors.rs b/src/errors.rs index 1578770..e422250 100644 --- a/src/errors.rs +++ b/src/errors.rs @@ -73,6 +73,8 @@ pub enum PakeError { InvalidLoginError, /// Error with serializing / deserializing protocol messages SerializationError, + /// Identity group element was encountered during deserialization, which is invalid + IdentityGroupElementError, } // This is meant to express future(ly) non-trivial ways of converting the diff --git a/src/group.rs b/src/group.rs index 34f8e2c..2d02a11 100644 --- a/src/group.rs +++ b/src/group.rs @@ -12,6 +12,7 @@ use curve25519_dalek::{ constants::RISTRETTO_BASEPOINT_POINT, ristretto::{CompressedRistretto, RistrettoPoint}, scalar::Scalar, + traits::Identity, }; use generic_array::{ typenum::{U32, U64}, @@ -35,7 +36,7 @@ pub trait Group: Copy + Sized + for<'a> Mul<&'a ::Scalar, Output scalar_bits: &GenericArray, ) -> Result; /// picks a scalar at random - fn random_scalar(rng: &mut R) -> Self::Scalar; + fn random_nonzero_scalar(rng: &mut R) -> Self::Scalar; /// Serializes a scalar to bytes fn scalar_as_bytes(scalar: &Self::Scalar) -> &GenericArray; /// The multiplicative inverse of this scalar @@ -64,6 +65,9 @@ pub trait Group: Copy + Sized + for<'a> Mul<&'a ::Scalar, Output /// Multiply the point by a scalar, represented as a slice fn mult_by_slice(&self, scalar: &GenericArray) -> Self; + + /// Returns if the group element is equal to the identity (1) + fn is_identity(&self) -> bool; } /// The implementation of such a subgroup for Ristretto @@ -77,20 +81,28 @@ impl Group for RistrettoPoint { bits.copy_from_slice(scalar_bits); Ok(Scalar::from_bytes_mod_order(bits)) } - fn random_scalar(rng: &mut R) -> Self::Scalar { - #[cfg(not(test))] - { - let mut scalar_bytes = [0u8; 64]; - rng.fill_bytes(&mut scalar_bytes); - Scalar::from_bytes_mod_order_wide(&scalar_bytes) - } + fn random_nonzero_scalar(rng: &mut R) -> Self::Scalar { + loop { + let scalar = { + #[cfg(not(test))] + { + let mut scalar_bytes = [0u8; 64]; + rng.fill_bytes(&mut scalar_bytes); + Scalar::from_bytes_mod_order_wide(&scalar_bytes) + } - // Tests need an exact conversion from bytes to scalar, sampling only 32 bytes from rng - #[cfg(test)] - { - let mut scalar_bytes = [0u8; 32]; - rng.fill_bytes(&mut scalar_bytes); - Scalar::from_bytes_mod_order(scalar_bytes) + // Tests need an exact conversion from bytes to scalar, sampling only 32 bytes from rng + #[cfg(test)] + { + let mut scalar_bytes = [0u8; 32]; + rng.fill_bytes(&mut scalar_bytes); + Scalar::from_bytes_mod_order(scalar_bytes) + } + }; + + if scalar != Scalar::zero() { + break scalar; + } } } fn scalar_as_bytes(scalar: &Self::Scalar) -> &GenericArray { @@ -134,4 +146,9 @@ impl Group for RistrettoPoint { let arr: [u8; 32] = scalar.as_slice().try_into().expect("Wrong length"); self * Scalar::from_bits(arr) } + + /// Returns if the group element is equal to the identity (1) + fn is_identity(&self) -> bool { + self == &Self::identity() + } } diff --git a/src/keypair.rs b/src/keypair.rs index d22a983..093cf83 100644 --- a/src/keypair.rs +++ b/src/keypair.rs @@ -79,7 +79,7 @@ impl KeyPair { /// Generating a random key pair given a cryptographic rng pub(crate) fn generate_random(rng: &mut R) -> Self { - let sk = G::random_scalar(rng); + let sk = G::random_nonzero_scalar(rng); let sk_bytes = G::scalar_as_bytes(&sk); let pk = G::base_point().mult_by_slice(sk_bytes); Self { @@ -174,6 +174,7 @@ impl Key { GenericArray::clone_from_slice(&self.0[..]) } + #[allow(clippy::unnecessary_wraps)] fn from_arr(key_bytes: &GenericArray) -> Result { Ok(Key(key_bytes.to_vec())) } diff --git a/src/messages.rs b/src/messages.rs index dbdfcae..32c4956 100644 --- a/src/messages.rs +++ b/src/messages.rs @@ -28,6 +28,14 @@ pub struct RegistrationRequest { pub(crate) alpha: CS::Group, } +impl RegistrationRequest { + /// Only used for testing purposes + #[cfg(test)] + pub fn get_alpha_for_testing(&self) -> CS::Group { + self.alpha + } +} + // Cannot be derived because it would require for CS to be Clone. impl Clone for RegistrationRequest { fn clone(&self) -> Self { @@ -49,6 +57,11 @@ impl RegistrationRequest { // correct subgroup let arr = GenericArray::from_slice(checked_slice); let alpha = CS::Group::from_element_slice(arr)?; + + // Throw an error if the identity group element is encountered + if alpha.is_identity() { + return Err(PakeError::IdentityGroupElementError.into()); + } Ok(Self { alpha }) } } @@ -91,6 +104,13 @@ impl RegistrationResponse { // correct subgroup let arr = GenericArray::from_slice(&checked_slice[..elem_len]); let beta = CS::Group::from_element_slice(arr)?; + + // Throw an error if the identity group element is encountered + if beta.is_identity() { + return Err(PakeError::IdentityGroupElementError.into()); + } + + // Ensure that public key is valid let server_s_pk = KeyPair::::check_public_key(PublicKey::from_bytes( &checked_slice[elem_len..], )?)?; @@ -188,6 +208,11 @@ impl CredentialRequest { let arr = GenericArray::from_slice(&checked_slice[..elem_len]); let alpha = CS::Group::from_element_slice(arr)?; + // Throw an error if the identity group element is encountered + if alpha.is_identity() { + return Err(PakeError::IdentityGroupElementError.into()); + } + let ke1_message = >::KE1Message::from_bytes::( &checked_slice[elem_len..], @@ -258,6 +283,11 @@ impl CredentialResponse { let arr = GenericArray::from_slice(beta_bytes); let beta = CS::Group::from_element_slice(arr)?; + // Throw an error if the identity group element is encountered + if beta.is_identity() { + return Err(PakeError::IdentityGroupElementError.into()); + } + let unchecked_server_s_pk = PublicKey::from_bytes(&checked_slice[elem_len..elem_len + key_len])?; let server_s_pk = KeyPair::::check_public_key(unchecked_server_s_pk)?; diff --git a/src/opaque.rs b/src/opaque.rs index 35caaa0..32930af 100644 --- a/src/opaque.rs +++ b/src/opaque.rs @@ -383,7 +383,7 @@ impl ServerRegistration { server_s_pk: &PublicKey, ) -> Result, ProtocolError> { // RFC: generate oprf_key (salt) and v_u = g^oprf_key - let oprf_key = CS::Group::random_scalar(rng); + let oprf_key = CS::Group::random_nonzero_scalar(rng); // Compute beta = alpha^oprf_key let beta = oprf::evaluate::(message.alpha, &oprf_key); diff --git a/src/oprf.rs b/src/oprf.rs index e9f7347..bc3ec4b 100644 --- a/src/oprf.rs +++ b/src/oprf.rs @@ -30,7 +30,8 @@ pub(crate) fn blind( input: &[u8], blinding_factor_rng: &mut R, ) -> Result<(Token, G), InternalPakeError> { - let blind = G::random_scalar(blinding_factor_rng); + // Choose a random scalar that must be non-zero + let blind = G::random_nonzero_scalar(blinding_factor_rng); 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; diff --git a/src/serialization/tests.rs b/src/serialization/tests.rs index 2a4e932..78052f4 100644 --- a/src/serialization/tests.rs +++ b/src/serialization/tests.rs @@ -6,6 +6,7 @@ use crate::{ ciphersuite::CipherSuite, envelope::{Envelope, InnerEnvelopeMode}, + errors::*, group::Group, key_exchange::{ traits::{FromBytes, KeyExchange, ToBytes}, @@ -16,7 +17,7 @@ use crate::{ *, }; -use curve25519_dalek::ristretto::RistrettoPoint; +use curve25519_dalek::{ristretto::RistrettoPoint, traits::Identity}; use generic_array::typenum::Unsigned; use generic_bytes::SizedBytes; use proptest::{collection::vec, prelude::*}; @@ -53,7 +54,7 @@ fn random_ristretto_point() -> RistrettoPoint { fn client_registration_roundtrip() { let pw = b"hunter2"; let mut rng = OsRng; - let sc = ::random_scalar(&mut rng); + let sc = ::random_nonzero_scalar(&mut rng); // serialization order: scalar, password let bytes: Vec = [&sc.as_bytes()[..], &pw[..]].concat(); @@ -67,7 +68,7 @@ fn server_registration_roundtrip() { // If we don't have envelope and client_pk, the server registration just // contains the prf key let mut rng = OsRng; - let oprf_key = ::random_scalar(&mut rng); + let oprf_key = ::random_nonzero_scalar(&mut rng); let mut oprf_bytes: Vec = vec![]; oprf_bytes.extend_from_slice(oprf_key.as_bytes()); let reg = ServerRegistration::::deserialize(&oprf_bytes[..]).unwrap(); @@ -106,6 +107,17 @@ fn registration_request_roundtrip() { let r1 = RegistrationRequest::::deserialize(input.as_slice()).unwrap(); let r1_bytes = r1.serialize(); assert_eq!(input, r1_bytes); + + // Assert that identity group element is rejected + let identity = RistrettoPoint::identity(); + let identity_bytes = identity.to_arr().to_vec(); + + assert!( + match RegistrationRequest::::deserialize(identity_bytes.as_slice()) { + Err(ProtocolError::VerificationError(PakeError::IdentityGroupElementError)) => true, + _ => false, + } + ); } #[test] @@ -123,6 +135,17 @@ fn registration_response_roundtrip() { let r2 = RegistrationResponse::::deserialize(input.as_slice()).unwrap(); let r2_bytes = r2.serialize(); assert_eq!(input, r2_bytes); + + // Assert that identity group element is rejected + let identity = RistrettoPoint::identity(); + let identity_bytes = identity.to_arr().to_vec(); + + assert!(match RegistrationResponse::::deserialize( + &[identity_bytes, pubkey_bytes.to_vec()].concat() + ) { + Err(ProtocolError::VerificationError(PakeError::IdentityGroupElementError)) => true, + _ => false, + }); } #[test] @@ -183,6 +206,17 @@ fn credential_request_roundtrip() { let l1 = CredentialRequest::::deserialize(input.as_slice()).unwrap(); let l1_bytes = l1.serialize(); assert_eq!(input, l1_bytes); + + // Assert that identity group element is rejected + let identity = RistrettoPoint::identity(); + let identity_bytes = identity.to_arr().to_vec(); + + assert!(match CredentialRequest::::deserialize( + &[identity_bytes, ke1m.to_vec()].concat() + ) { + Err(ProtocolError::VerificationError(PakeError::IdentityGroupElementError)) => true, + _ => false, + }); } #[test] @@ -226,15 +260,34 @@ fn credential_response_roundtrip() { ] .concat(); + let serialized_envelope = envelope.serialize(); + let mut input = Vec::new(); input.extend_from_slice(pt_bytes.as_slice()); input.extend_from_slice(&pubkey_bytes.as_slice()); - input.extend_from_slice(&envelope.serialize()); + input.extend_from_slice(&serialized_envelope); input.extend_from_slice(&ke2m[..]); let l2 = CredentialResponse::::deserialize(&input).unwrap(); let l2_bytes = l2.serialize(); assert_eq!(input, l2_bytes); + + // Assert that identity group element is rejected + let identity = RistrettoPoint::identity(); + let identity_bytes = identity.to_arr().to_vec(); + + assert!(match CredentialRequest::::deserialize( + &[ + identity_bytes, + pubkey_bytes.to_vec(), + serialized_envelope, + ke2m.to_vec() + ] + .concat() + ) { + Err(ProtocolError::VerificationError(PakeError::IdentityGroupElementError)) => true, + _ => false, + }); } #[test] @@ -254,7 +307,7 @@ fn login_third_message_roundtrip() { fn client_login_roundtrip() { let pw = b"hunter2"; let mut rng = OsRng; - let sc = ::random_scalar(&mut rng); + let sc = ::random_nonzero_scalar(&mut rng); let client_e_kp = Default::generate_random_keypair(&mut rng); let mut client_nonce = vec![0u8; NonceLen::to_usize()]; diff --git a/src/tests/full_test.rs b/src/tests/full_test.rs index 08862e2..3c9142e 100644 --- a/src/tests/full_test.rs +++ b/src/tests/full_test.rs @@ -16,7 +16,7 @@ use crate::{ tests::mock_rng::CycleRng, *, }; -use curve25519_dalek::ristretto::RistrettoPoint; +use curve25519_dalek::{ristretto::RistrettoPoint, traits::Identity}; use generic_array::typenum::Unsigned; use generic_bytes::SizedBytes; use rand::{rngs::OsRng, RngCore}; @@ -285,7 +285,7 @@ fn generate_parameters() -> TestVectorParameters { let mut server_nonce = vec![0u8; NonceLen::to_usize()]; rng.fill_bytes(&mut server_nonce); - let blinding_factor = CS::Group::random_scalar(&mut rng); + let blinding_factor = CS::Group::random_nonzero_scalar(&mut rng); let blinding_factor_bytes = CS::Group::scalar_as_bytes(&blinding_factor).clone(); let info1 = b"info1"; @@ -1050,3 +1050,23 @@ fn test_zeroize_server_login_finish() -> Result<(), ProtocolError> { Ok(()) } + +#[test] +fn test_scalar_always_nonzero() -> Result<(), ProtocolError> { + // Start out with a bunch of zeros to force resampling of scalar + let mut client_registration_rng = CycleRng::new([vec![0u8; 128], vec![1u8; 128]].concat()); + let client_registration_start_result = + ClientRegistration::::start( + &mut client_registration_rng, + STR_PASSWORD.as_bytes(), + )?; + + assert_ne!( + RistrettoPoint::identity(), + client_registration_start_result + .message + .get_alpha_for_testing() + ); + + Ok(()) +}