Replace Vec in Key with GenericArray

This commit is contained in:
daxpedda
2021-07-08 18:55:53 -07:00
committed by Kevin Lewi
parent 93d92133f9
commit 2955bcf5ac
2 changed files with 100 additions and 88 deletions
+8 -8
View File
@@ -209,29 +209,29 @@ impl<D: Hash, G: Group> KeyExchange<D, G> for TripleDH {
/// The client state produced after the first key exchange message
#[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))]
pub struct Ke1State<G> {
pub struct Ke1State<G: Group> {
client_e_sk: PrivateKey<G>,
client_nonce: GenericArray<u8, NonceLen>,
}
impl_clone_for!(
struct Ke1State<G>,
struct Ke1State<G: Group>,
[client_e_sk, client_nonce],
);
impl_debug_eq_hash_for!(
struct Ke1State<G>,
struct Ke1State<G: Group>,
[client_e_sk, client_nonce],
);
// This can't be derived because of the use of a generic parameter
impl<G> Zeroize for Ke1State<G> {
impl<G: Group> Zeroize for Ke1State<G> {
fn zeroize(&mut self) {
self.client_e_sk.zeroize();
self.client_nonce.zeroize();
}
}
impl<G> Drop for Ke1State<G> {
impl<G: Group> Drop for Ke1State<G> {
fn drop(&mut self) {
self.zeroize();
}
@@ -240,7 +240,7 @@ impl<G> Drop for Ke1State<G> {
/// The first key exchange message
#[derive(PartialEq, Eq, Debug, Hash, Clone)]
#[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))]
pub struct Ke1Message<G> {
pub struct Ke1Message<G: Group> {
pub(crate) client_nonce: GenericArray<u8, NonceLen>,
pub(crate) client_e_pk: PublicKey<G>,
}
@@ -344,7 +344,7 @@ impl<HashLen: ArrayLength<u8>> ToBytesWithPointers for Ke2State<HashLen> {
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
#[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))]
#[cfg_attr(feature = "serialize", serde(bound = ""))]
pub struct Ke2Message<G, HashLen: ArrayLength<u8>> {
pub struct Ke2Message<G: Group, HashLen: ArrayLength<u8>> {
server_nonce: GenericArray<u8, NonceLen>,
server_e_pk: PublicKey<G>,
mac: GenericArray<u8, HashLen>,
@@ -408,7 +408,7 @@ impl<G: Group, HashLen: ArrayLength<u8>> FromBytes for Ke2Message<G, HashLen> {
#[allow(clippy::upper_case_acronyms)]
// The triple of public and private components used in the 3DH computation
struct TripleDHComponents<G> {
struct TripleDHComponents<G: Group> {
pk1: PublicKey<G>,
sk1: PrivateKey<G>,
pk2: PublicKey<G>,
+92 -80
View File
@@ -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<T> SizedBytesExt for T where T: SizedBytes {}
derive(serde::Deserialize, serde::Serialize),
serde(bound = "")
)]
pub struct KeyPair<G> {
pub struct KeyPair<G: Group> {
pk: PublicKey<G>,
sk: PrivateKey<G>,
}
impl_clone_for!(
struct KeyPair<G>,
struct KeyPair<G: Group>,
[pk, sk],
);
impl_debug_eq_hash_for!(
struct KeyPair<G>,
struct KeyPair<G: Group>,
[pk, sk],
);
// This can't be derived because of the use of a phantom parameter
impl<G> Zeroize for KeyPair<G> {
impl<G: Group> Zeroize for KeyPair<G> {
fn zeroize(&mut self) {
self.pk.zeroize();
self.sk.zeroize();
}
}
impl<G> Drop for KeyPair<G> {
impl<G: Group> Drop for KeyPair<G> {
fn drop(&mut self) {
self.zeroize();
}
@@ -85,8 +84,8 @@ impl<G: Group> KeyPair<G> {
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<G: Group> KeyPair<G> {
/// &public_from_private(self.private()) == self.public()
pub(crate) fn public_from_private(bytes: &PrivateKey<G>) -> PublicKey<G> {
let bytes_data = GenericArray::<u8, G::ScalarLen>::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<G: Group> KeyPair<G> {
/// Obtains a KeyPair from a slice representing the private key
pub fn from_private_key_slice(input: &[u8]) -> Result<Self, InternalPakeError> {
let sk = PrivateKey::new(Key::from_arr::<G::ScalarLen>(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<G: Group + Debug> KeyPair<G> {
}
/// 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<u8>);
pub struct Key<L: ArrayLength<u8>>(GenericArray<u8, L>);
impl Deref for Key {
type Target = Vec<u8>;
impl<L: ArrayLength<u8>> Clone for Key<L> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<L: ArrayLength<u8>> Debug for Key<L> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_tuple("Key").field(&self.0).finish()
}
}
impl<L: ArrayLength<u8>> Eq for Key<L> {}
impl<L: ArrayLength<u8>> PartialEq for Key<L> {
fn eq(&self, other: &Self) -> bool {
self.0.eq(&other.0)
}
}
impl<L: ArrayLength<u8>> std::hash::Hash for Key<L> {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.0.hash(state);
}
}
// This can't be derived because of the use of a generic parameter
impl<L: ArrayLength<u8>> Zeroize for Key<L> {
fn zeroize(&mut self) {
self.0.zeroize();
}
}
impl<L: ArrayLength<u8>> Drop for Key<L> {
fn drop(&mut self) {
self.zeroize();
}
}
impl<L: ArrayLength<u8>> Deref for Key<L> {
type Target = GenericArray<u8, L>;
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<L: ArrayLength<u8>>(&self) -> GenericArray<u8, L> {
impl<L: ArrayLength<u8>> Key<L> {
fn to_arr(&self) -> GenericArray<u8, L> {
GenericArray::clone_from_slice(&self.0[..])
}
#[allow(clippy::unnecessary_wraps)]
fn from_arr<L: ArrayLength<u8>>(
key_bytes: &GenericArray<u8, L>,
) -> Result<Self, TryFromSizedBytesError> {
Ok(Key(key_bytes.to_vec()))
fn from_arr(key_bytes: &GenericArray<u8, L>) -> Result<Self, TryFromSizedBytesError> {
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<G> {
key: Key,
_g: PhantomData<G>,
}
pub struct PrivateKey<G: Group>(Key<G::ScalarLen>);
impl_clone_for!(
struct PrivateKey<G>,
[key, _g],
tuple PrivateKey<G: Group>,
[0],
);
impl_debug_eq_hash_for!(
struct PrivateKey<G>,
[key, _g],
tuple PrivateKey<G: Group>,
[0],
);
// This can't be derived because of the use of a phantom parameter
impl<G> Zeroize for PrivateKey<G> {
impl<G: Group> Zeroize for PrivateKey<G> {
fn zeroize(&mut self) {
self.key.zeroize();
self.0.zeroize();
}
}
impl<G> Drop for PrivateKey<G> {
impl<G: Group> Drop for PrivateKey<G> {
fn drop(&mut self) {
self.zeroize();
}
}
impl<G> Deref for PrivateKey<G> {
type Target = Key;
impl<G: Group> Deref for PrivateKey<G> {
type Target = Key<G::ScalarLen>;
fn deref(&self) -> &Self::Target {
&self.key
}
}
impl<G> PrivateKey<G> {
fn new(key: Key) -> Self {
Self {
key,
_g: PhantomData,
}
&self.0
}
}
@@ -235,58 +255,46 @@ impl<G: Group> SizedBytes for PrivateKey<G> {
type Len = G::ScalarLen;
fn to_arr(&self) -> GenericArray<u8, Self::Len> {
self.key.to_arr()
self.0.to_arr()
}
fn from_arr(key_bytes: &GenericArray<u8, Self::Len>) -> Result<Self, TryFromSizedBytesError> {
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<G> {
key: Key,
_g: PhantomData<G>,
}
pub struct PublicKey<G: Group>(Key<G::ElemLen>);
impl_clone_for!(
struct PublicKey<G>,
[key, _g],
tuple PublicKey<G: Group>,
[0],
);
impl_debug_eq_hash_for!(
struct PublicKey<G>,
[key, _g],
tuple PublicKey<G: Group>,
[0],
);
// This can't be derived because of the use of a phantom parameter
impl<G> Zeroize for PublicKey<G> {
// This can't be derived because of the use of a generic parameter
impl<G: Group> Zeroize for PublicKey<G> {
fn zeroize(&mut self) {
self.key.zeroize();
self.0.zeroize();
}
}
impl<G> Drop for PublicKey<G> {
impl<G: Group> Drop for PublicKey<G> {
fn drop(&mut self) {
self.zeroize();
}
}
impl<G> Deref for PublicKey<G> {
type Target = Key;
impl<G: Group> Deref for PublicKey<G> {
type Target = Key<G::ElemLen>;
fn deref(&self) -> &Self::Target {
&self.key
}
}
impl<G> PublicKey<G> {
fn new(key: Key) -> Self {
Self {
key,
_g: PhantomData,
}
&self.0
}
}
@@ -294,11 +302,11 @@ impl<G: Group> SizedBytes for PublicKey<G> {
type Len = G::ElemLen;
fn to_arr(&self) -> GenericArray<u8, Self::Len> {
self.key.to_arr()
self.0.to_arr()
}
fn from_arr(key_bytes: &GenericArray<u8, Self::Len>) -> Result<Self, TryFromSizedBytesError> {
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 = <RistrettoPoint as Group>::ElemLen::to_usize();
let mut key = Key(vec![1u8; key_len]);
let mut key =
Key::<<RistrettoPoint as Group>::ElemLen>(GenericArray::clone_from_slice(&vec![
1u8;
key_len
]));
let ptr = key.as_ptr();
key.zeroize();