diff --git a/Cargo.lock b/Cargo.lock index 648d584..fbeb2af 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1,5 +1,7 @@ # This file is automatically @generated by Cargo. # It is not intended for manual editing. +version = 3 + [[package]] name = "aead" version = "0.3.2" @@ -556,7 +558,7 @@ checksum = "624a8340c38c1b80fd549087862da4ba43e08858af025b236e509b6649fc13d5" [[package]] name = "opaque-ke" -version = "0.5.0" +version = "0.5.1-pre.1" dependencies = [ "anyhow", "base64", diff --git a/Cargo.toml b/Cargo.toml index a8c3597..ac6248e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "opaque-ke" -version = "0.5.0" +version = "0.5.1-pre.1" repository = "https://github.com/novifinancial/opaque-ke" keywords = ["cryptography", "crypto", "opaque", "passwords", "authentication"] description = "An implementation of the OPAQUE password-authenticated key exchange protocol" diff --git a/src/envelope.rs b/src/envelope.rs index a10a4f5..50a909b 100644 --- a/src/envelope.rs +++ b/src/envelope.rs @@ -16,6 +16,7 @@ use hkdf::Hkdf; use hmac::{Hmac, Mac, NewMac}; use rand::{CryptoRng, RngCore}; use std::convert::TryFrom; +use zeroize::Zeroize; // Constant string used as salt for HKDF computation const STR_PAD: &[u8] = b"Pad"; @@ -24,7 +25,8 @@ const STR_EXPORT_KEY: &[u8] = b"ExportKey"; const NONCE_LEN: usize = 32; -#[derive(Clone, Copy, PartialEq)] +#[derive(Clone, Copy, PartialEq, Zeroize)] +#[zeroize(drop)] pub(crate) enum InnerEnvelopeMode { Base = 1, CustomIdentifier = 2, @@ -41,6 +43,8 @@ impl TryFrom for InnerEnvelopeMode { } } +#[derive(Clone, Zeroize)] +#[zeroize(drop)] pub(crate) struct InnerEnvelope { mode: InnerEnvelopeMode, nonce: Vec, @@ -78,6 +82,15 @@ impl InnerEnvelope { bytes[NONCE_LEN + key_len..].to_vec(), )) } + + #[cfg(test)] + pub fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + vec![ + /* Cannot easily get raw pointer of enum value, otherwise would do self.mode.as_ptr() */ + (self.nonce.as_ptr(), self.nonce.len()), + (self.ciphertext.as_ptr(), self.ciphertext.len()), + ] + } } /// This struct is an instantiation of the envelope as described in @@ -90,6 +103,7 @@ impl InnerEnvelope { /// The specification update has simplified this assumption by taking /// an XOR-based approach without compromising on security, and to avoid /// the confusion around the implementation of an RKR-secure encryption. +#[derive(Clone)] pub(crate) struct Envelope { inner_envelope: InnerEnvelope, hmac: GenericArray::OutputSize>, @@ -284,6 +298,29 @@ impl Envelope { ), }) } + + #[cfg(test)] + pub fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + [ + self.inner_envelope.as_byte_ptrs(), + vec![(self.hmac.as_ptr(), self.hmac.len())], + ] + .concat() + } +} + +// This can't be derived because of the use of a phantom parameter +impl Zeroize for Envelope { + fn zeroize(&mut self) { + self.inner_envelope.zeroize(); + self.hmac.zeroize(); + } +} + +impl Drop for Envelope { + fn drop(&mut self) { + self.zeroize(); + } } // Helper functions diff --git a/src/key_exchange/traits.rs b/src/key_exchange/traits.rs index c88ad01..e1df71a 100644 --- a/src/key_exchange/traits.rs +++ b/src/key_exchange/traits.rs @@ -12,10 +12,11 @@ use crate::{ use rand::{CryptoRng, RngCore}; use std::convert::TryFrom; +use zeroize::Zeroize; pub trait KeyExchange { - type KE1State: for<'r> TryFrom<&'r [u8], Error = PakeError> + ToBytes; - type KE2State: for<'r> TryFrom<&'r [u8], Error = PakeError> + ToBytes; + type KE1State: for<'r> TryFrom<&'r [u8], Error = PakeError> + ToBytesWithPointers + Zeroize; + type KE2State: for<'r> TryFrom<&'r [u8], Error = PakeError> + ToBytesWithPointers + Zeroize; type KE1Message: for<'r> TryFrom<&'r [u8], Error = PakeError> + ToBytes; type KE2Message: for<'r> TryFrom<&'r [u8], Error = PakeError> + ToBytes; type KE3Message: for<'r> TryFrom<&'r [u8], Error = PakeError> + ToBytes; @@ -62,3 +63,11 @@ pub trait KeyExchange { pub trait ToBytes { fn to_bytes(&self) -> Vec; } + +pub trait ToBytesWithPointers { + fn to_bytes(&self) -> Vec; + + // Only used for tests to grab raw pointers to data + #[cfg(test)] + fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)>; +} diff --git a/src/key_exchange/tripledh.rs b/src/key_exchange/tripledh.rs index 708158a..2c504f5 100644 --- a/src/key_exchange/tripledh.rs +++ b/src/key_exchange/tripledh.rs @@ -11,7 +11,7 @@ use crate::{ }, group::Group, hash::Hash, - key_exchange::traits::{KeyExchange, ToBytes}, + key_exchange::traits::{KeyExchange, ToBytes, ToBytesWithPointers}, keypair::{Key, KeyPair, SizedBytesExt}, serialization::{serialize, tokenize}, }; @@ -24,6 +24,7 @@ use generic_bytes::SizedBytes; use hkdf::Hkdf; use hmac::{Hmac, Mac, NewMac}; use rand::{CryptoRng, RngCore}; +use zeroize::Zeroize; use std::convert::TryFrom; @@ -237,7 +238,8 @@ impl KeyExchange for TripleDH { } /// The client state produced after the first key exchange message -#[derive(PartialEq, Eq)] +#[derive(PartialEq, Eq, Zeroize)] +#[zeroize(drop)] pub struct Ke1State { client_e_sk: Key, client_nonce: GenericArray, @@ -267,11 +269,22 @@ impl TryFrom<&[u8]> for Ke1State { } } -impl ToBytes for Ke1State { +impl ToBytesWithPointers for Ke1State { fn to_bytes(&self) -> Vec { let output: Vec = [&self.client_e_sk.to_arr(), &self.client_nonce[..]].concat(); output } + + #[cfg(test)] + fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + vec![ + ( + self.client_e_sk.as_ptr(), + ::Len::to_usize(), + ), + (self.client_nonce.as_ptr(), NonceLen::to_usize()), + ] + } } impl ToBytes for Ke1Message { @@ -311,15 +324,22 @@ pub struct Ke2State> { session_key: GenericArray, } -/// The second key exchange message -pub struct Ke2Message> { - server_nonce: GenericArray, - server_e_pk: Key, - e_info: Vec, - mac: GenericArray, +// This can't be derived because of the use of a phantom parameter +impl> Zeroize for Ke2State { + fn zeroize(&mut self) { + self.km3.zeroize(); + self.hashed_transcript.zeroize(); + self.session_key.zeroize(); + } } -impl> ToBytes for Ke2State { +impl> Drop for Ke2State { + fn drop(&mut self) { + self.zeroize(); + } +} + +impl> ToBytesWithPointers for Ke2State { fn to_bytes(&self) -> Vec { [ &self.km3[..], @@ -328,6 +348,23 @@ impl> ToBytes for Ke2State { ] .concat() } + + #[cfg(test)] + fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + vec![ + (self.km3.as_ptr(), HashLen::to_usize()), + (self.hashed_transcript.as_ptr(), HashLen::to_usize()), + (self.session_key.as_ptr(), HashLen::to_usize()), + ] + } +} + +/// The second key exchange message +pub struct Ke2Message> { + server_nonce: GenericArray, + server_e_pk: Key, + e_info: Vec, + mac: GenericArray, } impl> TryFrom<&[u8]> for Ke2State { diff --git a/src/keypair.rs b/src/keypair.rs index edae451..27f85f1 100644 --- a/src/keypair.rs +++ b/src/keypair.rs @@ -5,8 +5,12 @@ //! Contains the keypair types that must be supplied for the OPAQUE API +#![allow(unsafe_code)] + use crate::errors::InternalPakeError; use crate::group::Group; +#[cfg(test)] +use generic_array::typenum::Unsigned; use generic_array::{typenum::U32, GenericArray}; use generic_bytes::{SizedBytes, TryFromSizedBytesError}; #[cfg(test)] @@ -31,13 +35,27 @@ pub trait SizedBytesExt: SizedBytes { impl SizedBytesExt for T where T: SizedBytes {} /// A Keypair trait with public-private verification -#[derive(Clone, Debug, PartialEq, Eq, Zeroize)] +#[derive(Clone, Debug, PartialEq, Eq)] pub struct KeyPair { pk: Key, sk: Key, _g: PhantomData, } +// This can't be derived because of the use of a phantom parameter +impl Zeroize for KeyPair { + fn zeroize(&mut self) { + self.pk.zeroize(); + self.sk.zeroize(); + } +} + +impl Drop for KeyPair { + fn drop(&mut self) { + self.zeroize(); + } +} + impl KeyPair { /// The public key component pub fn public(&self) -> &Key { @@ -100,6 +118,14 @@ impl KeyPair { let pk = Self::public_from_private(&sk); Self::new(pk, sk) } + + #[cfg(test)] + pub fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + vec![ + (self.pk.as_ptr(), ::Len::to_usize()), + (self.sk.as_ptr(), ::Len::to_usize()), + ] + } } #[cfg(test)] @@ -149,7 +175,41 @@ impl SizedBytes for Key { #[cfg(test)] mod tests { use super::*; + use crate::errors::*; use curve25519_dalek::ristretto::RistrettoPoint; + use generic_array::typenum::Unsigned; + use rand::rngs::OsRng; + use std::slice::from_raw_parts; + + #[test] + fn test_zeroize_key() -> Result<(), ProtocolError> { + let key_len = ::Len::to_usize(); + let mut key = Key(vec![1u8; key_len]); + let ptr = key.as_ptr(); + + key.zeroize(); + + let bytes = unsafe { from_raw_parts(ptr, key_len) }; + assert!(bytes.iter().all(|&x| x == 0)); + + Ok(()) + } + + #[test] + fn test_zeroize_keypair() -> Result<(), ProtocolError> { + let mut rng = OsRng; + let mut keypair = KeyPair::::generate_random(&mut rng); + let ptrs = keypair.as_byte_ptrs(); + + keypair.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) + } proptest! { #[test] diff --git a/src/opaque.rs b/src/opaque.rs index e867ef1..35c6f6d 100644 --- a/src/opaque.rs +++ b/src/opaque.rs @@ -11,7 +11,7 @@ use crate::{ errors::{utils::check_slice_size_atleast, InternalPakeError, PakeError, ProtocolError}, group::Group, hash::Hash, - key_exchange::traits::{KeyExchange, ToBytes}, + key_exchange::traits::{KeyExchange, ToBytesWithPointers}, keypair::{Key, KeyPair, SizedBytesExt}, map_to_curve::GroupWithMapToCurve, oprf, @@ -73,6 +73,14 @@ impl ClientRegistration { }, }) } + + #[cfg(test)] + pub fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + vec![ + (self.token.data.as_ptr(), self.token.data.len()), + /* cannot provide raw pointer to self.token.blind until this is exposed in curve25519_dalek::scalar::Scalar */ + ] + } } /// Optional parameters for client registration finish @@ -140,6 +148,9 @@ pub struct ClientRegistrationFinishResult { pub message: RegistrationUpload, /// The export key output by client registration pub export_key: GenericArray::OutputSize>, + /// Instance of the ClientRegistration, only used in tests for checking zeroize + #[cfg(test)] + pub state: ClientRegistration, } impl ClientRegistration { @@ -203,38 +214,12 @@ impl ClientRegistration { client_s_pk: client_static_keypair.public().clone(), }, export_key, + #[cfg(test)] + state: self, }) } } -// This can't be derived because of the use of a phantom parameter -impl Zeroize for ClientRegistration { - fn zeroize(&mut self) { - self.token.data.zeroize(); - self.token.blind.zeroize(); - } -} - -impl Drop for ClientRegistration { - fn drop(&mut self) { - self.zeroize(); - } -} - -// This can't be derived because of the use of a phantom parameter -impl Zeroize for ClientLogin { - fn zeroize(&mut self) { - self.token.data.zeroize(); - self.token.blind.zeroize(); - } -} - -impl Drop for ClientLogin { - fn drop(&mut self) { - self.zeroize(); - } -} - /// Contains the fields that are returned by a server registration start pub struct ServerRegistrationStartResult { /// The registration resposne message to send to the client @@ -295,6 +280,21 @@ impl ServerRegistration { }) } + #[cfg(test)] + pub fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + [ + match &self.envelope { + Some(env) => env.as_byte_ptrs(), + None => vec![], + }, + match &self.client_s_pk { + Some(pk) => vec![(pk.as_ptr(), pk.len())], + None => vec![], + }, + /* cannot provide raw pointer to self.oprf_key until this is exposed in curve25519_dalek::scalar::Scalar */ + ].concat() + } + /// From the client's "blinded" password, returns a response to be /// sent back to the client, as well as a ServerRegistration /// @@ -380,7 +380,7 @@ impl ServerRegistration { Ok(Self { envelope: Some(message.envelope), client_s_pk: Some(message.client_s_pk), - oprf_key: self.oprf_key, + oprf_key: self.oprf_key.clone(), }) } } @@ -440,6 +440,18 @@ impl ClientLogin { serialized_credential_request, }) } + + #[cfg(test)] + pub fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + [ + vec![ + (self.token.data.as_ptr(), self.token.data.len()), + /* cannot provide raw pointer to self.token.blind until this is exposed in curve25519_dalek::scalar::Scalar */ + ], + self.ke1_state.as_byte_ptrs(), + vec![ (self.serialized_credential_request.as_ptr(), self.serialized_credential_request.len()) ], + ].concat() + } } /// Optional parameters for client login start @@ -488,6 +500,9 @@ pub struct ClientLoginFinishResult { pub server_s_pk: Key, /// The confidential info sent by the client pub confidential_info: Vec, + /// Instance of the ClientLogin, only used in tests for checking zeroize + #[cfg(test)] + pub state: ClientLogin, } impl ClientLogin { @@ -626,6 +641,8 @@ impl ClientLogin { session_key, export_key: opened_envelope.export_key.clone(), server_s_pk: l2.server_s_pk, + #[cfg(test)] + state: self, }) } } @@ -665,9 +682,13 @@ pub struct ServerLoginStartResult { } /// Contains the fields that are returned by a server login finish -pub struct ServerLoginFinishResult { +pub struct ServerLoginFinishResult { /// The session key between client and server pub session_key: Vec, + _cs: PhantomData, + /// Instance of the ClientRegistration, only used in tests for checking zeroize + #[cfg(test)] + pub state: ServerLogin, } impl ServerLogin { @@ -728,6 +749,7 @@ impl ServerLogin { ) -> Result, ProtocolError> { let client_s_pk = password_file .client_s_pk + .clone() .ok_or(InternalPakeError::SealError)?; let (e_info, optional_ids) = match params { @@ -740,7 +762,10 @@ impl ServerLogin { } }; - let envelope = password_file.envelope.ok_or(InternalPakeError::SealError)?; + let envelope = password_file + .envelope + .clone() + .ok_or(InternalPakeError::SealError)?; if envelope.get_mode() != mode_from_ids(&optional_ids) { return Err(InternalPakeError::IncompatibleEnvelopeModeError.into()); } @@ -827,9 +852,9 @@ impl ServerLogin { /// # Ok::<(), ProtocolError>(()) /// ``` pub fn finish( - &self, + self, message: CredentialFinalization, - ) -> Result { + ) -> Result, ProtocolError> { let session_key = >::finish_ke( message.ke3_message, &self.ke2_state, @@ -841,11 +866,82 @@ impl ServerLogin { err => err, })?; - Ok(ServerLoginFinishResult { session_key }) + Ok(ServerLoginFinishResult { + session_key, + _cs: PhantomData, + #[cfg(test)] + state: self, + }) + } + + #[cfg(test)] + pub fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { + self.ke2_state.as_byte_ptrs() + } +} + +// Zeroize on drop implementations + +// This can't be derived because of the use of a phantom parameter +impl Zeroize for ClientRegistration { + fn zeroize(&mut self) { + self.token.data.zeroize(); + self.token.blind.zeroize(); + } +} + +impl Drop for ClientRegistration { + fn drop(&mut self) { + self.zeroize(); + } +} + +// This can't be derived because of the use of a phantom parameter +impl Zeroize for ServerRegistration { + fn zeroize(&mut self) { + self.envelope.zeroize(); + self.client_s_pk.zeroize(); + self.oprf_key.zeroize(); + } +} + +impl Drop for ServerRegistration { + fn drop(&mut self) { + self.zeroize(); + } +} + +// This can't be derived because of the use of a phantom parameter +impl Zeroize for ClientLogin { + fn zeroize(&mut self) { + self.token.data.zeroize(); + self.token.blind.zeroize(); + self.ke1_state.zeroize(); + self.serialized_credential_request.zeroize(); + } +} + +impl Drop for ClientLogin { + fn drop(&mut self) { + self.zeroize(); + } +} + +// This can't be derived because of the use of a phantom parameter +impl Zeroize for ServerLogin { + fn zeroize(&mut self) { + self.ke2_state.zeroize(); + } +} + +impl Drop for ServerLogin { + fn drop(&mut self) { + self.zeroize(); } } // Helper functions + fn get_password_derived_key, D: Hash>( token: &oprf::Token, beta: G, diff --git a/src/tests/full_test.rs b/src/tests/full_test.rs index e68b814..85fc6f3 100644 --- a/src/tests/full_test.rs +++ b/src/tests/full_test.rs @@ -3,6 +3,8 @@ // This source code is licensed under the MIT license found in the // LICENSE file in the root directory of this source tree. +#![allow(unsafe_code)] + use crate::{ ciphersuite::CipherSuite, errors::*, @@ -19,6 +21,8 @@ use generic_array::typenum::Unsigned; use generic_bytes::SizedBytes; use rand::{rngs::OsRng, RngCore}; use serde_json::Value; +use std::slice::from_raw_parts; +use zeroize::Zeroize; // Tests // ===== @@ -65,6 +69,8 @@ pub struct TestVectorParameters { pub session_key: Vec, } +static STR_PASSWORD: &str = "password"; + static TEST_VECTOR: &str = r#" { "client_s_pk": "6e0a6082dd29936c44b47ecb8a5fe72e4b321a0ac314b0080ca4c48afdabd215", @@ -416,7 +422,7 @@ fn generate_parameters() -> TestVectorParameters { server_registration_state, client_login_state, server_login_state, - session_key: client_login_finish_result.session_key, + session_key: client_login_finish_result.session_key.clone(), export_key: client_registration_finish_result.export_key.to_vec(), } } @@ -623,7 +629,7 @@ fn test_server_login_finish() -> Result<(), ProtocolError> { assert_eq!( hex::encode(parameters.session_key), - hex::encode(server_login_result.session_key) + hex::encode(&server_login_result.session_key) ); Ok(()) @@ -680,8 +686,8 @@ fn test_complete_flow( .finish(client_login_finish_result.message)?; assert_eq!( - hex::encode(server_login_finish_result.session_key), - hex::encode(client_login_finish_result.session_key) + hex::encode(&server_login_finish_result.session_key), + hex::encode(&client_login_finish_result.session_key) ); assert_eq!( hex::encode(client_registration_finish_result.export_key), @@ -706,3 +712,305 @@ fn test_complete_flow_success() -> Result<(), ProtocolError> { fn test_complete_flow_fail() -> Result<(), ProtocolError> { test_complete_flow(b"good password", b"bad password") } + +// Zeroize tests + +#[test] +fn test_zeroize_client_registration_start() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let client_registration_start_result = + ClientRegistration::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + )?; + + let mut state = client_registration_start_result.state; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} + +#[test] +fn test_zeroize_server_registration_start() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let mut server_rng = OsRng; + let server_kp = RistrettoSha5123dhNoSlowHash::generate_random_keypair(&mut server_rng); + let client_registration_start_result = + ClientRegistration::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + )?; + let server_registration_start_result = + ServerRegistration::::start( + &mut server_rng, + client_registration_start_result.message, + server_kp.public(), + )?; + + let mut state = server_registration_start_result.state; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} + +#[test] +fn test_zeroize_client_registration_finish() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let mut server_rng = OsRng; + let server_kp = RistrettoSha5123dhNoSlowHash::generate_random_keypair(&mut server_rng); + let client_registration_start_result = + ClientRegistration::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + )?; + let server_registration_start_result = + ServerRegistration::::start( + &mut server_rng, + client_registration_start_result.message, + server_kp.public(), + )?; + let client_registration_finish_result = client_registration_start_result.state.finish( + &mut client_rng, + server_registration_start_result.message, + ClientRegistrationFinishParameters::default(), + )?; + + let mut state = client_registration_finish_result.state; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} + +#[test] +fn test_zeroize_server_registration_finish() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let mut server_rng = OsRng; + let server_kp = RistrettoSha5123dhNoSlowHash::generate_random_keypair(&mut server_rng); + let client_registration_start_result = + ClientRegistration::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + )?; + let server_registration_start_result = + ServerRegistration::::start( + &mut server_rng, + client_registration_start_result.message, + server_kp.public(), + )?; + let client_registration_finish_result = client_registration_start_result.state.finish( + &mut client_rng, + server_registration_start_result.message, + ClientRegistrationFinishParameters::default(), + )?; + let p_file = server_registration_start_result + .state + .finish(client_registration_finish_result.message)?; + + let mut state = p_file; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} + +#[test] +fn test_zeroize_client_login_start() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let client_login_start_result = ClientLogin::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + ClientLoginStartParameters::default(), + )?; + + let mut state = client_login_start_result.state; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} + +#[test] +fn test_zeroize_server_login_start() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let mut server_rng = OsRng; + let server_kp = RistrettoSha5123dhNoSlowHash::generate_random_keypair(&mut server_rng); + let client_registration_start_result = + ClientRegistration::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + )?; + let server_registration_start_result = + ServerRegistration::::start( + &mut server_rng, + client_registration_start_result.message, + server_kp.public(), + )?; + let client_registration_finish_result = client_registration_start_result.state.finish( + &mut client_rng, + server_registration_start_result.message, + ClientRegistrationFinishParameters::default(), + )?; + let p_file = server_registration_start_result + .state + .finish(client_registration_finish_result.message)?; + let client_login_start_result = ClientLogin::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + ClientLoginStartParameters::default(), + )?; + let server_login_start_result = ServerLogin::::start( + &mut server_rng, + p_file, + &server_kp.private(), + client_login_start_result.message, + ServerLoginStartParameters::default(), + )?; + + let mut state = server_login_start_result.state; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} + +#[test] +fn test_zeroize_client_login_finish() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let mut server_rng = OsRng; + let server_kp = RistrettoSha5123dhNoSlowHash::generate_random_keypair(&mut server_rng); + let client_registration_start_result = + ClientRegistration::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + )?; + let server_registration_start_result = + ServerRegistration::::start( + &mut server_rng, + client_registration_start_result.message, + server_kp.public(), + )?; + let client_registration_finish_result = client_registration_start_result.state.finish( + &mut client_rng, + server_registration_start_result.message, + ClientRegistrationFinishParameters::default(), + )?; + let p_file = server_registration_start_result + .state + .finish(client_registration_finish_result.message)?; + let client_login_start_result = ClientLogin::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + ClientLoginStartParameters::default(), + )?; + let server_login_start_result = ServerLogin::::start( + &mut server_rng, + p_file, + &server_kp.private(), + client_login_start_result.message, + ServerLoginStartParameters::default(), + )?; + let client_login_finish_result = client_login_start_result.state.finish( + server_login_start_result.message, + ClientLoginFinishParameters::default(), + )?; + + let mut state = client_login_finish_result.state; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} + +#[test] +fn test_zeroize_server_login_finish() -> Result<(), ProtocolError> { + let mut client_rng = OsRng; + let mut server_rng = OsRng; + let server_kp = RistrettoSha5123dhNoSlowHash::generate_random_keypair(&mut server_rng); + let client_registration_start_result = + ClientRegistration::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + )?; + let server_registration_start_result = + ServerRegistration::::start( + &mut server_rng, + client_registration_start_result.message, + server_kp.public(), + )?; + let client_registration_finish_result = client_registration_start_result.state.finish( + &mut client_rng, + server_registration_start_result.message, + ClientRegistrationFinishParameters::default(), + )?; + let p_file = server_registration_start_result + .state + .finish(client_registration_finish_result.message)?; + let client_login_start_result = ClientLogin::::start( + &mut client_rng, + STR_PASSWORD.as_bytes(), + ClientLoginStartParameters::default(), + )?; + let server_login_start_result = ServerLogin::::start( + &mut server_rng, + p_file, + &server_kp.private(), + client_login_start_result.message, + ServerLoginStartParameters::default(), + )?; + let client_login_finish_result = client_login_start_result.state.finish( + server_login_start_result.message, + ClientLoginFinishParameters::default(), + )?; + let server_login_finish_result = server_login_start_result + .state + .finish(client_login_finish_result.message)?; + + let mut state = server_login_finish_result.state; + let ptrs = state.as_byte_ptrs(); + state.zeroize(); + + for (ptr, len) in ptrs { + let bytes = unsafe { from_raw_parts(ptr, len) }; + assert!(bytes.iter().all(|&x| x == 0)); + } + + Ok(()) +} diff --git a/src/tests/opaque_test_vectors.rs b/src/tests/opaque_test_vectors.rs index 2c0c3aa..0e9d0a3 100644 --- a/src/tests/opaque_test_vectors.rs +++ b/src/tests/opaque_test_vectors.rs @@ -576,7 +576,7 @@ fn test_server_login_finish() -> Result<(), ProtocolError> { assert_eq!( hex::encode(parameters.session_key), - hex::encode(server_login_result.session_key) + hex::encode(&server_login_result.session_key) ); } Ok(())