diff --git a/src/key_exchange/tripledh.rs b/src/key_exchange/tripledh.rs index d463f75..5d0da1f 100644 --- a/src/key_exchange/tripledh.rs +++ b/src/key_exchange/tripledh.rs @@ -209,29 +209,29 @@ impl KeyExchange for TripleDH { /// The client state produced after the first key exchange message #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] -pub struct Ke1State { +pub struct Ke1State { client_e_sk: PrivateKey, client_nonce: GenericArray, } impl_clone_for!( - struct Ke1State, + struct Ke1State, [client_e_sk, client_nonce], ); impl_debug_eq_hash_for!( - struct Ke1State, + struct Ke1State, [client_e_sk, client_nonce], ); // This can't be derived because of the use of a generic parameter -impl Zeroize for Ke1State { +impl Zeroize for Ke1State { fn zeroize(&mut self) { self.client_e_sk.zeroize(); self.client_nonce.zeroize(); } } -impl Drop for Ke1State { +impl Drop for Ke1State { fn drop(&mut self) { self.zeroize(); } @@ -240,7 +240,7 @@ impl Drop for Ke1State { /// The first key exchange message #[derive(PartialEq, Eq, Debug, Hash, Clone)] #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] -pub struct Ke1Message { +pub struct Ke1Message { pub(crate) client_nonce: GenericArray, pub(crate) client_e_pk: PublicKey, } @@ -344,7 +344,7 @@ impl> ToBytesWithPointers for Ke2State { #[derive(Clone, Debug, Eq, Hash, PartialEq)] #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] #[cfg_attr(feature = "serialize", serde(bound = ""))] -pub struct Ke2Message> { +pub struct Ke2Message> { server_nonce: GenericArray, server_e_pk: PublicKey, mac: GenericArray, @@ -408,7 +408,7 @@ impl> FromBytes for Ke2Message { #[allow(clippy::upper_case_acronyms)] // The triple of public and private components used in the 3DH computation -struct TripleDHComponents { +struct TripleDHComponents { pk1: PublicKey, sk1: PrivateKey, pk2: PublicKey, diff --git a/src/keypair.rs b/src/keypair.rs index b3c402c..62d72b7 100644 --- a/src/keypair.rs +++ b/src/keypair.rs @@ -19,7 +19,6 @@ use proptest::prelude::*; use rand::{rngs::StdRng, SeedableRng}; use rand::{CryptoRng, RngCore}; use std::fmt::Debug; -use std::marker::PhantomData; use std::ops::Deref; use zeroize::Zeroize; @@ -40,29 +39,29 @@ impl SizedBytesExt for T where T: SizedBytes {} derive(serde::Deserialize, serde::Serialize), serde(bound = "") )] -pub struct KeyPair { +pub struct KeyPair { pk: PublicKey, sk: PrivateKey, } impl_clone_for!( - struct KeyPair, + struct KeyPair, [pk, sk], ); impl_debug_eq_hash_for!( - struct KeyPair, + struct KeyPair, [pk, sk], ); // This can't be derived because of the use of a phantom parameter -impl Zeroize for KeyPair { +impl Zeroize for KeyPair { fn zeroize(&mut self) { self.pk.zeroize(); self.sk.zeroize(); } } -impl Drop for KeyPair { +impl Drop for KeyPair { fn drop(&mut self) { self.zeroize(); } @@ -85,8 +84,8 @@ impl KeyPair { let sk_bytes = G::scalar_as_bytes(sk); let pk = G::base_point().mult_by_slice(&sk_bytes); Self { - pk: PublicKey::new(Key(pk.to_arr().to_vec())), - sk: PrivateKey::new(Key(sk_bytes.to_vec())), + pk: PublicKey(Key(pk.to_arr())), + sk: PrivateKey(Key(sk_bytes)), } } @@ -94,10 +93,7 @@ impl KeyPair { /// &public_from_private(self.private()) == self.public() pub(crate) fn public_from_private(bytes: &PrivateKey) -> PublicKey { let bytes_data = GenericArray::::from_slice(&bytes.0[..]); - PublicKey::new(Key(G::base_point() - .mult_by_slice(bytes_data) - .to_arr() - .to_vec())) + PublicKey(Key(G::base_point().mult_by_slice(bytes_data).to_arr())) } /// Check whether a public key is valid. This is meant to be applied on @@ -121,9 +117,7 @@ impl KeyPair { /// Obtains a KeyPair from a slice representing the private key pub fn from_private_key_slice(input: &[u8]) -> Result { - let sk = PrivateKey::new(Key::from_arr::(GenericArray::from_slice( - input, - ))?); + let sk = PrivateKey(Key(GenericArray::clone_from_slice(input))); let pk = Self::public_from_private(&sk); Ok(Self { pk, sk }) } @@ -155,15 +149,55 @@ impl KeyPair { } /// A minimalist key type built around a \[u8; 32\] -#[derive(Debug, PartialEq, Eq, Clone, Hash, Zeroize)] -#[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] -// Ensure Key material is zeroed after use. -#[zeroize(drop)] +#[cfg_attr( + feature = "serialize", + derive(serde::Deserialize, serde::Serialize), + serde(bound = "") +)] #[repr(transparent)] -pub struct Key(Vec); +pub struct Key>(GenericArray); -impl Deref for Key { - type Target = Vec; +impl> Clone for Key { + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +impl> Debug for Key { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_tuple("Key").field(&self.0).finish() + } +} + +impl> Eq for Key {} + +impl> PartialEq for Key { + fn eq(&self, other: &Self) -> bool { + self.0.eq(&other.0) + } +} + +impl> std::hash::Hash for Key { + fn hash(&self, state: &mut H) { + self.0.hash(state); + } +} + +// This can't be derived because of the use of a generic parameter +impl> Zeroize for Key { + fn zeroize(&mut self) { + self.0.zeroize(); + } +} + +impl> Drop for Key { + fn drop(&mut self) { + self.zeroize(); + } +} + +impl> Deref for Key { + type Target = GenericArray; fn deref(&self) -> &Self::Target { &self.0 @@ -171,63 +205,49 @@ impl Deref for Key { } // Don't make it implement SizedBytes so that it's not constructible outside of this module. -impl Key { - fn to_arr>(&self) -> GenericArray { +impl> Key { + fn to_arr(&self) -> GenericArray { GenericArray::clone_from_slice(&self.0[..]) } #[allow(clippy::unnecessary_wraps)] - fn from_arr>( - key_bytes: &GenericArray, - ) -> Result { - Ok(Key(key_bytes.to_vec())) + fn from_arr(key_bytes: &GenericArray) -> Result { + Ok(Key(key_bytes.to_owned())) } } /// Wrapper around a Key to enforce that it's a private one. #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] #[repr(transparent)] -pub struct PrivateKey { - key: Key, - _g: PhantomData, -} +pub struct PrivateKey(Key); impl_clone_for!( - struct PrivateKey, - [key, _g], + tuple PrivateKey, + [0], ); impl_debug_eq_hash_for!( - struct PrivateKey, - [key, _g], + tuple PrivateKey, + [0], ); // This can't be derived because of the use of a phantom parameter -impl Zeroize for PrivateKey { +impl Zeroize for PrivateKey { fn zeroize(&mut self) { - self.key.zeroize(); + self.0.zeroize(); } } -impl Drop for PrivateKey { +impl Drop for PrivateKey { fn drop(&mut self) { self.zeroize(); } } -impl Deref for PrivateKey { - type Target = Key; +impl Deref for PrivateKey { + type Target = Key; fn deref(&self) -> &Self::Target { - &self.key - } -} - -impl PrivateKey { - fn new(key: Key) -> Self { - Self { - key, - _g: PhantomData, - } + &self.0 } } @@ -235,58 +255,46 @@ impl SizedBytes for PrivateKey { type Len = G::ScalarLen; fn to_arr(&self) -> GenericArray { - self.key.to_arr() + self.0.to_arr() } fn from_arr(key_bytes: &GenericArray) -> Result { - Ok(PrivateKey::new(Key::from_arr(key_bytes)?)) + Ok(PrivateKey(Key::from_arr(key_bytes)?)) } } /// Wrapper around a Key to enforce that it's a public one. #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] #[repr(transparent)] -pub struct PublicKey { - key: Key, - _g: PhantomData, -} +pub struct PublicKey(Key); impl_clone_for!( - struct PublicKey, - [key, _g], + tuple PublicKey, + [0], ); impl_debug_eq_hash_for!( - struct PublicKey, - [key, _g], + tuple PublicKey, + [0], ); -// This can't be derived because of the use of a phantom parameter -impl Zeroize for PublicKey { +// This can't be derived because of the use of a generic parameter +impl Zeroize for PublicKey { fn zeroize(&mut self) { - self.key.zeroize(); + self.0.zeroize(); } } -impl Drop for PublicKey { +impl Drop for PublicKey { fn drop(&mut self) { self.zeroize(); } } -impl Deref for PublicKey { - type Target = Key; +impl Deref for PublicKey { + type Target = Key; fn deref(&self) -> &Self::Target { - &self.key - } -} - -impl PublicKey { - fn new(key: Key) -> Self { - Self { - key, - _g: PhantomData, - } + &self.0 } } @@ -294,11 +302,11 @@ impl SizedBytes for PublicKey { type Len = G::ElemLen; fn to_arr(&self) -> GenericArray { - self.key.to_arr() + self.0.to_arr() } fn from_arr(key_bytes: &GenericArray) -> Result { - Ok(PublicKey::new(Key::from_arr(key_bytes)?)) + Ok(PublicKey(Key::from_arr(key_bytes)?)) } } @@ -314,7 +322,11 @@ mod tests { #[test] fn test_zeroize_key() -> Result<(), ProtocolError> { let key_len = ::ElemLen::to_usize(); - let mut key = Key(vec![1u8; key_len]); + let mut key = + Key::<::ElemLen>(GenericArray::clone_from_slice(&vec![ + 1u8; + key_len + ])); let ptr = key.as_ptr(); key.zeroize();