From 8457e8b9003bd8b02114e9dfe735c307569f81c1 Mon Sep 17 00:00:00 2001 From: Kevin Lewi Date: Thu, 14 Oct 2021 13:55:43 -0700 Subject: [PATCH] Adding "danger" feature to expose underlying elements (#27) * Adding internal feature to expose underlying elements * Applying @daxpedda's comments and suggestions --- .github/workflows/main.yml | 2 + Cargo.toml | 8 ++- src/lib.rs | 7 ++- src/serialization.rs | 51 ++++++++++++++- src/voprf.rs | 125 +++++++++++++------------------------ 5 files changed, 108 insertions(+), 85 deletions(-) diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml index 1e31d5a..bbf7011 100644 --- a/.github/workflows/main.yml +++ b/.github/workflows/main.yml @@ -18,6 +18,7 @@ jobs: - p256,ristretto255_u64 frontend_feature: - serde + - danger toolchain: - stable - 1.51.0 @@ -64,6 +65,7 @@ jobs: frontend_feature: - - --features serde + - --features danger steps: - uses: actions/checkout@v2 - uses: hecrj/setup-rust-action@v1 diff --git a/Cargo.toml b/Cargo.toml index bff254a..e34a1ac 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,6 +13,7 @@ resolver = "2" [features] default = ["ristretto255_u64", "serde"] +danger = [] ristretto255_u64 = ["curve25519-dalek/u64_backend"] ristretto255_u32 = ["curve25519-dalek/u32_backend"] ristretto255_fiat_u64 = ["curve25519-dalek/fiat_u64_backend"] @@ -42,12 +43,13 @@ zeroize = { version = "1", default-features = false } generic-array = { version = "0.14", features = ["more_lengths"] } hex = "0.4" json = "0.12" +proptest = "1" rand = "0.8" -sha2 = "0.9" regex = "1" -voprf = { path = "", default-features = false, features = ["std"] } +sha2 = "0.9" +voprf = { path = "", default-features = false, features = ["std", "danger"] } [package.metadata.docs.rs] -features = ["p256", "std"] +features = ["danger", "p256", "std"] targets = [] rustdoc-args = ["--cfg", "docsrs"] diff --git a/src/lib.rs b/src/lib.rs index e41d7f0..3840982 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -401,6 +401,11 @@ //! - The `serde` feature, enabled by default, provides convenience functions for serializing and deserializing with //! [serde](https://serde.rs/). //! +//! - The `danger` feature, disabled by default, exposes functions for setting and getting +//! internal values not available in the default API. These functions are intended for use in +//! by higher-level cryptographic protocols that need access to these raw values and are able to +//! perform the necessary validations on them (such as being valid group elements). +//! //! - The backend features are re-exported from //! [curve25519-dalek](https://doc.dalek.rs/curve25519_dalek/index.html#backends-and-features) and allow for selecting //! the corresponding backend for the curve arithmetic used. The `ristretto255_u64` feature is included as the default. @@ -408,7 +413,7 @@ //! //! - The `ristretto255_simd` feature is re-exported from //! [curve25519-dalek](https://doc.dalek.rs/curve25519_dalek/index.html#backends-and-features) and enables parallel formulas, -//! using either AVX2 or AVX512-IFMA. This will automatically enable the `ristretto255_u64` and requires Rust nightly. +//! using either AVX2 or AVX512-IFMA. This will automatically enable the `ristretto255_u64` feature and requires Rust nightly. #![deny(unsafe_code)] #![warn(clippy::cargo, missing_docs)] diff --git a/src/serialization.rs b/src/serialization.rs index cd98a1d..4caff32 100644 --- a/src/serialization.rs +++ b/src/serialization.rs @@ -142,7 +142,7 @@ impl Proof { /// Deserialization from bytes pub fn deserialize(input: &[u8]) -> Result { let scalar_len = ::ScalarLen::USIZE; - if input.len() < scalar_len + scalar_len { + if input.len() != scalar_len + scalar_len { return Err(InternalError::SizeError); } Ok(Proof { @@ -161,6 +161,10 @@ impl BlindedElement { /// Deserialization from bytes pub fn deserialize(input: &[u8]) -> Result { + let elem_len = ::ElemLen::USIZE; + if input.len() != elem_len { + return Err(InternalError::SizeError); + } Ok(Self { value: G::from_element_slice(input)?, hash: PhantomData, @@ -176,6 +180,10 @@ impl EvaluationElement { /// Deserialization from bytes pub fn deserialize(input: &[u8]) -> Result { + let elem_len = ::ElemLen::USIZE; + if input.len() != elem_len { + return Err(InternalError::SizeError); + } Ok(Self { value: G::from_element_slice(input)?, hash: PhantomData, @@ -218,7 +226,10 @@ pub(crate) fn serialize>(input: &[u8]) -> Result, Int #[cfg(test)] mod unit_tests { use super::*; + use curve25519_dalek::ristretto::RistrettoPoint; use generic_array::typenum::{U1, U2}; + use proptest::{collection::vec, prelude::*}; + use sha2::Sha512; // Test the error condition for I2OSP #[test] @@ -233,4 +244,42 @@ mod unit_tests { assert!(i2osp::(256 * 256).is_err()); assert!(i2osp::(256 * 256 + 1).is_err()); } + + proptest! { + #[test] + fn test_nocrash_nonverifiable_client(bytes in vec(any::(), 0..200)) { + NonVerifiableClient::::deserialize(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_verifiable_client(bytes in vec(any::(), 0..200)) { + VerifiableClient::::deserialize(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_nonverifiable_server(bytes in vec(any::(), 0..200)) { + NonVerifiableServer::::deserialize(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_verifiable_server(bytes in vec(any::(), 0..200)) { + VerifiableServer::::deserialize(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_blinded_element(bytes in vec(any::(), 0..200)) { + BlindedElement::::deserialize(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_evaluation_element(bytes in vec(any::(), 0..200)) { + EvaluationElement::::deserialize(&bytes[..]).map_or(true, |_| true); + } + + #[test] + fn test_nocrash_proof(bytes in vec(any::(), 0..200)) { + Proof::::deserialize(&bytes[..]).map_or(true, |_| true); + } + + } } diff --git a/src/voprf.rs b/src/voprf.rs index 848410c..8c37dae 100644 --- a/src/voprf.rs +++ b/src/voprf.rs @@ -192,20 +192,11 @@ impl NonVerifiableClient { } } - #[cfg(test)] - /// Only used for test functions + #[cfg(feature = "danger")] + /// Exposes the blind group element pub fn get_blind(&self) -> ::Scalar { self.blind } - - #[cfg(test)] - /// Only used for testing zeroize - pub fn as_ptrs(&self) -> Vec> { - vec![ - self.data.clone(), - ::scalar_as_bytes(self.blind).to_vec(), - ] - } } impl VerifiableClient { @@ -330,16 +321,6 @@ impl VerifiableClient { pub fn get_blind(&self) -> ::Scalar { self.blind } - - #[cfg(test)] - /// Only used for testing zeroize - pub fn as_ptrs(&self) -> Vec> { - vec![ - self.data.clone(), - ::scalar_as_bytes(self.blind).to_vec(), - self.blinded_element.to_arr().to_vec(), - ] - } } impl NonVerifiableServer { @@ -405,12 +386,6 @@ impl NonVerifiableServer { }, }) } - - #[cfg(test)] - /// Only used for testing zeroize - pub fn as_ptrs(&self) -> Vec> { - vec![::scalar_as_bytes(self.sk).to_vec()] - } } impl VerifiableServer { @@ -523,15 +498,6 @@ impl VerifiableServer { pub fn get_public_key(&self) -> G { self.pk } - - #[cfg(test)] - /// Only used for testing zeroize - pub fn as_ptrs(&self) -> Vec> { - vec![ - ::scalar_as_bytes(self.sk).to_vec(), - self.pk.to_arr().to_vec(), - ] - } } ///////////////////////// @@ -598,10 +564,24 @@ impl BlindedElement { } } - #[cfg(test)] - /// Only used for testing zeroize - pub fn as_ptrs(&self) -> Vec> { - vec![self.value.to_arr().to_vec()] + #[cfg(feature = "danger")] + /// Creates a [BlindedElement] from a raw group element. + /// + /// # Caution + /// + /// This should be used with caution, since + /// it does not perform any checks on the validity of the value itself! + pub fn from_value_unchecked(value: G) -> Self { + Self { + value, + hash: PhantomData, + } + } + + #[cfg(feature = "danger")] + /// Exposes the internal value + pub fn value(&self) -> G { + self.value } } @@ -614,21 +594,24 @@ impl EvaluationElement { } } - #[cfg(test)] - /// Only used for testing zeroize - pub fn as_ptrs(&self) -> Vec> { - vec![self.value.to_arr().to_vec()] + #[cfg(feature = "danger")] + /// Creates an [EvaluationElement] from a raw group element. + /// + /// # Caution + /// + /// This should be used with caution, since + /// it does not perform any checks on the validity of the value itself! + pub fn from_value_unchecked(value: G) -> Self { + Self { + value, + hash: PhantomData, + } } -} -impl Proof { - #[cfg(test)] - /// Only used for testing zeroize - pub fn as_ptrs(&self) -> Vec> { - vec![ - ::scalar_as_bytes(self.c_scalar).to_vec(), - ::scalar_as_bytes(self.s_scalar).to_vec(), - ] + #[cfg(feature = "danger")] + /// Exposes the internal value + pub fn value(&self) -> G { + self.value } } @@ -1053,15 +1036,11 @@ mod tests { let mut state = client_blind_result.state; Zeroize::zeroize(&mut state); - for bytes in state.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(state.serialize().iter().all(|&x| x == 0)); let mut message = client_blind_result.message; Zeroize::zeroize(&mut message); - for bytes in message.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(message.serialize().iter().all(|&x| x == 0)); } fn zeroize_verifiable_client() { @@ -1072,15 +1051,11 @@ mod tests { let mut state = client_blind_result.state; Zeroize::zeroize(&mut state); - for bytes in state.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(state.serialize().iter().all(|&x| x == 0)); let mut message = client_blind_result.message; Zeroize::zeroize(&mut message); - for bytes in message.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(message.serialize().iter().all(|&x| x == 0)); } fn zeroize_base_server() { @@ -1096,15 +1071,11 @@ mod tests { let mut state = server; Zeroize::zeroize(&mut state); - for bytes in state.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(state.serialize().iter().all(|&x| x == 0)); let mut message = server_result.message; Zeroize::zeroize(&mut message); - for bytes in message.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(message.serialize().iter().all(|&x| x == 0)); } fn zeroize_verifiable_server() { @@ -1120,21 +1091,15 @@ mod tests { let mut state = server; Zeroize::zeroize(&mut state); - for bytes in state.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(state.serialize().iter().all(|&x| x == 0)); let mut message = server_result.message; Zeroize::zeroize(&mut message); - for bytes in message.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(message.serialize().iter().all(|&x| x == 0)); let mut proof = server_result.proof; Zeroize::zeroize(&mut proof); - for bytes in proof.as_ptrs() { - assert!(bytes.iter().all(|&x| x == 0)); - } + assert!(proof.serialize().iter().all(|&x| x == 0)); } #[test]