Document all Errors (#55)

* Document all `Error`s

* Fix inaccuracies
This commit is contained in:
daxpedda
2022-01-24 20:55:02 -08:00
committed by GitHub
parent 16e072dcd4
commit b01b8ed409
9 changed files with 249 additions and 125 deletions
+21 -21
View File
@@ -10,32 +10,32 @@
use displaydoc::Display;
/// [`Result`](core::result::Result) shorthand that uses [`Error`].
pub type Result<T> = core::result::Result<T, Error>;
pub type Result<T, E = Error> = core::result::Result<T, E>;
/// Represents an error in the manipulation of internal cryptographic data
#[derive(Clone, Copy, Debug, Display, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub enum Error {
/// Could not parse byte sequence for key
InvalidByteSequence,
/// Could not deserialize element, or deserialized to the identity element
PointError,
/// Computing the hash-to-curve function failed
HashToCurveError,
/// Failure to serialize or deserialize bytes
SerializationError,
/// Use of incompatible modes (base vs. verifiable)
IncompatibleModeError,
/**
* Internal error thrown when different-lengthed slices are supplied
* to the compute_composites() function.
*/
MismatchedLengthsForCompositeInputs,
/// Size of input is empty or longer then [`u16::MAX`].
Input,
/// Size of metadata is longer then `u16::MAX - 21`.
Metadata,
/// Failure to deserialize bytes
Deserialization,
/// Batched items are more then [`u16::MAX`] or length don't match.
Batch,
/// In verifiable mode, occurs when the proof failed to verify
ProofVerificationError,
/// Encountered insufficient bytes when attempting to deserialize
SizeError,
/// Encountered an invalid scalar
ScalarError,
ProofVerification,
/// Size of seed is longer then [`u16::MAX`].
Seed,
}
/// Only used to implement [`Group`](crate::Group).
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub enum InternalError {
/// Size of input is empty or longer then [`u16::MAX`].
Input,
/// `input` is longer then [`u16::MAX`].
I2osp,
}
#[cfg(feature = "std")]
+14 -7
View File
@@ -21,7 +21,7 @@ use rand_core::{CryptoRng, RngCore};
use super::Group;
use crate::group::{STR_HASH_TO_GROUP, STR_HASH_TO_SCALAR};
use crate::voprf::{self, Mode};
use crate::{CipherSuite, Error, Result};
use crate::{CipherSuite, Error, InternalError, Result};
impl<C> Group for C
where
@@ -41,7 +41,10 @@ where
// Implements the `hash_to_curve()` function from
// https://datatracker.ietf.org/doc/html/draft-irtf-cfrg-hash-to-curve-11#section-3
fn hash_to_curve<CS: CipherSuite>(msg: &[&[u8]], mode: Mode) -> Result<Self::Elem>
fn hash_to_curve<CS: CipherSuite>(
input: &[&[u8]],
mode: Mode,
) -> Result<Self::Elem, InternalError>
where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
@@ -49,11 +52,15 @@ where
let dst =
GenericArray::from(STR_HASH_TO_GROUP).concat(voprf::get_context_string::<CS>(mode));
Self::hash_from_bytes::<ExpandMsgXmd<CS::Hash>>(msg, &dst).map_err(|_| Error::PointError)
Self::hash_from_bytes::<ExpandMsgXmd<CS::Hash>>(input, &dst)
.map_err(|_| InternalError::Input)
}
// Implements the `HashToScalar()` function
fn hash_to_scalar<CS: CipherSuite>(input: &[&[u8]], mode: Mode) -> Result<Self::Scalar>
fn hash_to_scalar<CS: CipherSuite>(
input: &[&[u8]],
mode: Mode,
) -> Result<Self::Scalar, InternalError>
where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
@@ -62,7 +69,7 @@ where
GenericArray::from(STR_HASH_TO_SCALAR).concat(voprf::get_context_string::<CS>(mode));
<Self as GroupDigest>::hash_to_scalar::<ExpandMsgXmd<CS::Hash>>(input, &dst)
.map_err(|_| Error::PointError)
.map_err(|_| InternalError::Input)
}
fn base_elem() -> Self::Elem {
@@ -85,7 +92,7 @@ where
fn deserialize_elem(element_bits: &GenericArray<u8, Self::ElemLen>) -> Result<Self::Elem> {
PublicKey::<Self>::from_sec1_bytes(element_bits)
.map(|public_key| public_key.to_projective())
.map_err(|_| Error::PointError)
.map_err(|_| Error::Deserialization)
}
fn random_scalar<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Scalar {
@@ -108,6 +115,6 @@ where
fn deserialize_scalar(scalar_bits: &GenericArray<u8, Self::ScalarLen>) -> Result<Self::Scalar> {
SecretKey::<Self>::from_be_bytes(scalar_bits)
.map(|secret_key| *secret_key.to_nonzero_scalar())
.map_err(|_| Error::ScalarError)
.map_err(|_| Error::Deserialization)
}
}
+26 -4
View File
@@ -24,7 +24,7 @@ use subtle::ConstantTimeEq;
use zeroize::Zeroize;
use crate::voprf::Mode;
use crate::{CipherSuite, Result};
use crate::{CipherSuite, InternalError, Result};
pub(crate) const STR_HASH_TO_SCALAR: [u8; 13] = *b"HashToScalar-";
pub(crate) const STR_HASH_TO_GROUP: [u8; 12] = *b"HashToGroup-";
@@ -52,14 +52,28 @@ pub trait Group {
/// The byte length necessary to represent scalars
type ScalarLen: ArrayLength<u8> + 'static;
/// transforms a password and domain separation tag (DST) into a curve point
fn hash_to_curve<CS: CipherSuite>(msg: &[&[u8]], mode: Mode) -> Result<Self::Elem>
/// Transforms a password and domain separation tag (DST) into a curve point
///
/// # Errors
/// [`Error::Input`](crate::Error::Input) if the `input` is empty or longer
/// then [`u16::MAX`].
fn hash_to_curve<CS: CipherSuite>(
input: &[&[u8]],
mode: Mode,
) -> Result<Self::Elem, InternalError>
where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>;
/// Hashes a slice of pseudo-random bytes to a scalar
fn hash_to_scalar<CS: CipherSuite>(input: &[&[u8]], mode: Mode) -> Result<Self::Scalar>
///
/// # Errors
/// [`Error::Input`](crate::Error::Input) if the `input` is empty or longer
/// then [`u16::MAX`].
fn hash_to_scalar<CS: CipherSuite>(
input: &[&[u8]],
mode: Mode,
) -> Result<Self::Scalar, InternalError>
where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>;
@@ -75,6 +89,10 @@ pub trait Group {
/// Return an element from its fixed-length bytes representation. If the
/// element is the identity element, return an error.
///
/// # Errors
/// [`Error::Deserialization`](crate::Error::Deserialization) if the element
/// is not a valid point on the group or the identity element.
fn deserialize_elem(element_bits: &GenericArray<u8, Self::ElemLen>) -> Result<Self::Elem>;
/// picks a scalar at random
@@ -92,6 +110,10 @@ pub trait Group {
/// Return a scalar from its fixed-length bytes representation. If the
/// scalar is zero or invalid, then return an error.
///
/// # Errors
/// [`Error::Deserialization`](crate::Error::Deserialization) if the scalar
/// is not a valid point on the group or zero.
fn deserialize_scalar(scalar_bits: &GenericArray<u8, Self::ScalarLen>) -> Result<Self::Scalar>;
}
+14 -8
View File
@@ -19,7 +19,7 @@ use rand_core::{CryptoRng, RngCore};
use super::{Group, STR_HASH_TO_GROUP, STR_HASH_TO_SCALAR};
use crate::voprf::{self, Mode};
use crate::{CipherSuite, Error, Result};
use crate::{CipherSuite, Error, InternalError, Result};
/// [`Group`] implementation for Ristretto255.
pub struct Ristretto255;
@@ -46,7 +46,10 @@ impl Group for Ristretto255 {
// Implements the `hash_to_ristretto255()` function from
// https://www.ietf.org/archive/id/draft-irtf-cfrg-hash-to-curve-10.txt
fn hash_to_curve<CS: CipherSuite>(msg: &[&[u8]], mode: Mode) -> Result<Self::Elem>
fn hash_to_curve<CS: CipherSuite>(
input: &[&[u8]],
mode: Mode,
) -> Result<Self::Elem, InternalError>
where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
@@ -55,8 +58,8 @@ impl Group for Ristretto255 {
GenericArray::from(STR_HASH_TO_GROUP).concat(voprf::get_context_string::<Self>(mode));
let mut uniform_bytes = GenericArray::<_, U64>::default();
ExpandMsgXmd::<CS::Hash>::expand_message(msg, &dst, 64)
.map_err(|_| Error::PointError)?
ExpandMsgXmd::<CS::Hash>::expand_message(input, &dst, 64)
.map_err(|_| InternalError::Input)?
.fill_bytes(&mut uniform_bytes);
Ok(RistrettoPoint::from_uniform_bytes(&uniform_bytes.into()))
@@ -64,7 +67,10 @@ impl Group for Ristretto255 {
// Implements the `HashToScalar()` function from
// https://www.ietf.org/archive/id/draft-irtf-cfrg-voprf-07.html#section-4.1
fn hash_to_scalar<'a, CS: CipherSuite>(input: &[&[u8]], mode: Mode) -> Result<Self::Scalar>
fn hash_to_scalar<'a, CS: CipherSuite>(
input: &[&[u8]],
mode: Mode,
) -> Result<Self::Scalar, InternalError>
where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
@@ -74,7 +80,7 @@ impl Group for Ristretto255 {
let mut uniform_bytes = GenericArray::<_, U64>::default();
ExpandMsgXmd::<CS::Hash>::expand_message(input, &dst, 64)
.map_err(|_| Error::PointError)?
.map_err(|_| InternalError::Input)?
.fill_bytes(&mut uniform_bytes);
Ok(Scalar::from_bytes_mod_order_wide(&uniform_bytes.into()))
@@ -97,7 +103,7 @@ impl Group for Ristretto255 {
CompressedRistretto::from_slice(element_bits)
.decompress()
.filter(|point| point != &RistrettoPoint::identity())
.ok_or(Error::PointError)
.ok_or(Error::Deserialization)
}
fn random_scalar<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Scalar {
@@ -130,6 +136,6 @@ impl Group for Ristretto255 {
fn deserialize_scalar(scalar_bits: &GenericArray<u8, Self::ScalarLen>) -> Result<Self::Scalar> {
Scalar::from_canonical_bytes((*scalar_bits).into())
.filter(|scalar| scalar != &Scalar::zero())
.ok_or(Error::ScalarError)
.ok_or(Error::Deserialization)
}
}
+2 -2
View File
@@ -34,7 +34,7 @@ fn test_group_properties() -> Result<()> {
fn test_identity_element_error<G: Group>() -> Result<()> {
let identity = G::identity_elem();
let result = G::deserialize_elem(&G::serialize_elem(identity));
assert!(matches!(result, Err(Error::PointError)));
assert!(matches!(result, Err(Error::Deserialization)));
Ok(())
}
@@ -43,7 +43,7 @@ fn test_identity_element_error<G: Group>() -> Result<()> {
fn test_zero_scalar_error<G: Group>() -> Result<()> {
let zero_scalar = G::zero_scalar();
let result = G::deserialize_scalar(&G::serialize_scalar(zero_scalar));
assert!(matches!(result, Err(Error::ScalarError)));
assert!(matches!(result, Err(Error::Deserialization)));
Ok(())
}
+11 -20
View File
@@ -59,8 +59,7 @@
//! use voprf::NonVerifiableServer;
//!
//! let mut server_rng = OsRng;
//! let server = NonVerifiableServer::<CipherSuite>::new(&mut server_rng)
//! .expect("Unable to construct server");
//! let server = NonVerifiableServer::<CipherSuite>::new(&mut server_rng);
//! ```
//!
//! ### Client Blinding
@@ -108,8 +107,7 @@
//! # ).expect("Unable to construct client");
//! # use voprf::NonVerifiableServer;
//! # let mut server_rng = OsRng;
//! # let server = NonVerifiableServer::<CipherSuite>::new(&mut server_rng)
//! # .expect("Unable to construct server");
//! # let server = NonVerifiableServer::<CipherSuite>::new(&mut server_rng);
//! let server_evaluate_result = server
//! .evaluate(&client_blind_result.message, None)
//! .expect("Unable to perform server evaluate");
@@ -136,8 +134,7 @@
//! # ).expect("Unable to construct client");
//! # use voprf::NonVerifiableServer;
//! # let mut server_rng = OsRng;
//! # let server = NonVerifiableServer::<CipherSuite>::new(&mut server_rng)
//! # .expect("Unable to construct server");
//! # let server = NonVerifiableServer::<CipherSuite>::new(&mut server_rng);
//! # let server_evaluate_result = server.evaluate(
//! # &client_blind_result.message,
//! # None,
@@ -178,8 +175,7 @@
//! use voprf::VerifiableServer;
//!
//! let mut server_rng = OsRng;
//! let server =
//! VerifiableServer::<CipherSuite>::new(&mut server_rng).expect("Unable to construct server");
//! let server = VerifiableServer::<CipherSuite>::new(&mut server_rng);
//!
//! // To be sent to the client
//! println!("Server public key: {:?}", server.get_public_key());
@@ -234,8 +230,7 @@
//! # ).expect("Unable to construct client");
//! # use voprf::VerifiableServer;
//! # let mut server_rng = OsRng;
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng)
//! # .expect("Unable to construct server");
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng);
//! let server_evaluate_result = server
//! .evaluate(&mut server_rng, &client_blind_result.message, None)
//! .expect("Unable to perform server evaluate");
@@ -263,8 +258,7 @@
//! # ).expect("Unable to construct client");
//! # use voprf::VerifiableServer;
//! # let mut server_rng = OsRng;
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng)
//! # .expect("Unable to construct server");
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng);
//! # let server_evaluate_result = server.evaluate(
//! # &mut server_rng,
//! # &client_blind_result.message,
@@ -346,8 +340,7 @@
//! # }
//! # use voprf::VerifiableServer;
//! let mut server_rng = OsRng;
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng)
//! # .expect("Unable to construct server");
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng);
//! let VerifiableServerBatchEvaluatePrepareResult {
//! prepared_evaluation_elements,
//! t,
@@ -385,8 +378,7 @@
//! # }
//! # use voprf::VerifiableServer;
//! let mut server_rng = OsRng;
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng)
//! # .expect("Unable to construct server");
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng);
//! let VerifiableServerBatchEvaluateResult { messages, proof } = server
//! .batch_evaluate(&mut server_rng, &client_messages, None)
//! .expect("Unable to perform server batch evaluate");
@@ -420,8 +412,7 @@
//! # }
//! # use voprf::VerifiableServer;
//! # let mut server_rng = OsRng;
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng)
//! # .expect("Unable to construct server");
//! # let server = VerifiableServer::<CipherSuite>::new(&mut server_rng);
//! # let VerifiableServerBatchEvaluateResult { messages, proof } = server
//! # .batch_evaluate(&mut server_rng, &client_messages, None)
//! # .expect("Unable to perform server batch evaluate");
@@ -488,7 +479,7 @@
#![deny(unsafe_code)]
#![no_std]
#![warn(clippy::cargo, missing_docs)]
#![warn(clippy::cargo, clippy::missing_errors_doc, missing_docs)]
#![allow(clippy::multiple_crate_versions)]
#[cfg(any(feature = "alloc", test))]
@@ -510,7 +501,7 @@ mod tests;
// Exports
pub use crate::ciphersuite::CipherSuite;
pub use crate::error::{Error, Result};
pub use crate::error::{Error, InternalError, Result};
pub use crate::group::Group;
#[cfg(feature = "ristretto255")]
pub use crate::group::Ristretto255;
+22 -1
View File
@@ -40,6 +40,9 @@ where
}
/// Deserialization from bytes
///
/// # Errors
/// [`Error::Deserialization`] if failed to deserialize `input`.
pub fn deserialize(input: &[u8]) -> Result<Self> {
let mut input = input.iter().copied();
@@ -71,6 +74,9 @@ where
}
/// Deserialization from bytes
///
/// # Errors
/// [`Error::Deserialization`] if failed to deserialize `input`.
pub fn deserialize(input: &[u8]) -> Result<Self> {
let mut input = input.iter().copied();
@@ -98,6 +104,9 @@ where
}
/// Deserialization from bytes
///
/// # Errors
/// [`Error::Deserialization`] if failed to deserialize `input`.
pub fn deserialize(input: &[u8]) -> Result<Self> {
let mut input = input.iter().copied();
@@ -128,6 +137,9 @@ where
}
/// Deserialization from bytes
///
/// # Errors
/// [`Error::Deserialization`] if failed to deserialize `input`.
pub fn deserialize(input: &[u8]) -> Result<Self> {
let mut input = input.iter().copied();
@@ -160,6 +172,9 @@ where
}
/// Deserialization from bytes
///
/// # Errors
/// [`Error::Deserialization`] if failed to deserialize `input`.
pub fn deserialize(input: &[u8]) -> Result<Self> {
let mut input = input.iter().copied();
@@ -184,6 +199,9 @@ where
}
/// Deserialization from bytes
///
/// # Errors
/// [`Error::Deserialization`] if failed to deserialize `input`.
pub fn deserialize(input: &[u8]) -> Result<Self> {
let mut input = input.iter().copied();
@@ -207,6 +225,9 @@ where
}
/// Deserialization from bytes
///
/// # Errors
/// [`Error::Deserialization`] if failed to deserialize `input`.
pub fn deserialize(input: &[u8]) -> Result<Self> {
let mut input = input.iter().copied();
@@ -220,5 +241,5 @@ fn deserialize<L: ArrayLength<u8>>(
input: &mut impl Iterator<Item = u8>,
) -> Result<GenericArray<u8, L>> {
let input = input.by_ref().take(L::USIZE);
GenericArray::from_exact_iter(input).ok_or(Error::SizeError)
GenericArray::from_exact_iter(input).ok_or(Error::Deserialization)
}
+4 -4
View File
@@ -12,16 +12,16 @@ use core::convert::TryFrom;
use generic_array::typenum::{IsLess, U2, U256};
use generic_array::{ArrayLength, GenericArray};
use crate::{Error, Result};
use crate::InternalError;
pub(crate) fn i2osp_2(input: usize) -> Result<GenericArray<u8, U2>> {
pub(crate) fn i2osp_2(input: usize) -> Result<GenericArray<u8, U2>, InternalError> {
u16::try_from(input)
.map(|input| input.to_be_bytes().into())
.map_err(|_| Error::SerializationError)
.map_err(|_| InternalError::I2osp)
}
pub(crate) fn i2osp_2_array<L: ArrayLength<u8> + IsLess<U256>>(
_: GenericArray<u8, L>,
_: &GenericArray<u8, L>,
) -> GenericArray<u8, U2> {
L::U16.to_be_bytes().into()
}
+135 -58
View File
@@ -220,6 +220,9 @@ where
{
/// Computes the first step for the multiplicative blinding version of
/// DH-OPRF.
///
/// # Errors
/// [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
pub fn blind<R: RngCore + CryptoRng>(
input: &[u8],
blinding_factor_rng: &mut R,
@@ -240,6 +243,9 @@ where
///
/// This should be used with caution, since it does not perform any checks
/// on the validity of the blinding factor!
///
/// # Errors
/// [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
pub fn deterministic_blind_unchecked(
input: &[u8],
blind: <CS::Group as Group>::Scalar,
@@ -253,6 +259,10 @@ where
/// Computes the third step for the multiplicative blinding version of
/// DH-OPRF, in which the client unblinds the server's message.
///
/// # Errors
/// - [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
/// - [`Error::Metadata`] if the `metadata` is longer then `u16::MAX - 21`.
pub fn finalize(
&self,
input: &[u8],
@@ -261,10 +271,10 @@ where
) -> Result<Output<CS::Hash>> {
let unblinded_element = evaluation_element.0 * &CS::Group::invert_scalar(self.blind);
let mut outputs = finalize_after_unblind::<CS, _, _>(
Some((input, unblinded_element)).into_iter(),
iter::once((input, unblinded_element)),
metadata.unwrap_or_default(),
Mode::Base,
)?;
);
outputs.next().unwrap()
}
@@ -288,6 +298,9 @@ where
{
/// Computes the first step for the multiplicative blinding version of
/// DH-OPRF.
///
/// # Errors
/// [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
pub fn blind<R: RngCore + CryptoRng>(
input: &[u8],
blinding_factor_rng: &mut R,
@@ -312,6 +325,9 @@ where
///
/// This should be used with caution, since it does not perform any checks
/// on the validity of the blinding factor!
///
/// # Errors
/// [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
pub fn deterministic_blind_unchecked(
input: &[u8],
blind: <CS::Group as Group>::Scalar,
@@ -328,6 +344,11 @@ where
/// Computes the third step for the multiplicative blinding version of
/// DH-OPRF, in which the client unblinds the server's message.
///
/// # Errors
/// - [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
/// - [`Error::Metadata`] if the `metadata` is longer then `u16::MAX - 21`.
/// - [`Error::ProofVerification`] if the `proof` failed to verify.
pub fn finalize(
&self,
input: &[u8],
@@ -336,9 +357,9 @@ where
pk: <CS::Group as Group>::Elem,
metadata: Option<&[u8]>,
) -> Result<Output<CS::Hash>> {
let inputs: &[&[u8]; 1] = core::array::from_ref(&input);
let clients: &[Self; 1] = core::array::from_ref(self);
let messages: &[EvaluationElement<CS>; 1] = core::array::from_ref(evaluation_element);
let inputs = core::array::from_ref(&input);
let clients = core::array::from_ref(self);
let messages = core::array::from_ref(evaluation_element);
let mut batch_result =
Self::batch_finalize(inputs, clients, messages, proof, pk, metadata)?;
@@ -347,6 +368,15 @@ where
/// Allows for batching of the finalization of multiple [VerifiableClient]
/// and [EvaluationElement] pairs
///
/// # Errors
/// - [`Error::Metadata`] if the `metadata` is longer then `u16::MAX - 21`.
/// - [`Error::Batch`] if the number of `clients` and `messages` don't match
/// or is longer then [`u16::MAX`].
/// - [`Error::ProofVerification`] if the `proof` failed to verify.
///
/// The resulting messages can each fail individually with [`Error::Input`]
/// if the `input` is empty or longer then [`u16::MAX`].
pub fn batch_finalize<'a, I: 'a, II, IC, IM>(
inputs: &'a II,
clients: &'a IC,
@@ -371,11 +401,11 @@ where
let inputs_and_unblinded_elements = inputs.into_iter().zip(unblinded_elements);
finalize_after_unblind::<CS, _, _>(
Ok(finalize_after_unblind::<CS, _, _>(
inputs_and_unblinded_elements,
metadata,
Mode::Verifiable,
)
))
}
#[cfg(test)]
@@ -403,14 +433,19 @@ where
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
{
/// Produces a new instance of a [NonVerifiableServer] using a supplied RNG
pub fn new<R: RngCore + CryptoRng>(rng: &mut R) -> Result<Self> {
pub fn new<R: RngCore + CryptoRng>(rng: &mut R) -> Self {
let mut seed = Output::<CS::Hash>::default();
rng.fill_bytes(&mut seed);
Self::new_from_seed(&seed)
// This can't fail as the hash output is type constrained.
Self::new_from_seed(&seed).unwrap()
}
/// Produces a new instance of a [NonVerifiableServer] using a supplied set
/// of bytes to represent the server's private key
///
/// # Errors
/// [`Error::Deserialization`] if the private key is not a valid point on
/// the group or zero.
pub fn new_with_key(private_key_bytes: &[u8]) -> Result<Self> {
let sk = CS::Group::deserialize_scalar(private_key_bytes.into())?;
Ok(Self { sk })
@@ -420,8 +455,11 @@ where
/// of bytes which are used as a seed to derive the server's private key.
///
/// Corresponds to DeriveKeyPair() function from the VOPRF specification.
///
/// # Errors
/// [`Error::Seed`] if the `seed` is empty or longer then [`u16::MAX`].
pub fn new_from_seed(seed: &[u8]) -> Result<Self> {
let sk = CS::Group::hash_to_scalar::<CS>(&[seed], Mode::Base)?;
let sk = CS::Group::hash_to_scalar::<CS>(&[seed], Mode::Base).map_err(|_| Error::Seed)?;
Ok(Self { sk })
}
@@ -434,6 +472,9 @@ where
/// Computes the second step for the multiplicative blinding version of
/// DH-OPRF. This message is sent from the server (who holds the OPRF key)
/// to the client.
///
/// # Errors
/// [`Error::Metadata`] if the `metadata` is longer then `u16::MAX - 21`.
pub fn evaluate(
&self,
blinded_element: &BlindedElement<CS>,
@@ -447,11 +488,12 @@ where
// context = "Context-" || contextString || I2OSP(len(info), 2) || info
let context = GenericArray::from(STR_CONTEXT)
.concat(context_string)
.concat(i2osp_2(metadata.len())?);
.concat(i2osp_2(metadata.len()).map_err(|_| Error::Metadata)?);
let context = [&context, metadata];
// m = GG.HashToScalar(context)
let m = CS::Group::hash_to_scalar::<CS>(&context, Mode::Base)?;
let m =
CS::Group::hash_to_scalar::<CS>(&context, Mode::Base).map_err(|_| Error::Metadata)?;
// t = skS + m
let t = self.sk + &m;
// Z = t^(-1) * R
@@ -469,14 +511,19 @@ where
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
{
/// Produces a new instance of a [VerifiableServer] using a supplied RNG
pub fn new<R: RngCore + CryptoRng>(rng: &mut R) -> Result<Self> {
pub fn new<R: RngCore + CryptoRng>(rng: &mut R) -> Self {
let mut seed = Output::<CS::Hash>::default();
rng.fill_bytes(&mut seed);
Self::new_from_seed(&seed)
// This can't fail as the hash output is type constrained.
Self::new_from_seed(&seed).unwrap()
}
/// Produces a new instance of a [VerifiableServer] using a supplied set of
/// bytes to represent the server's private key
///
/// # Errors
/// [`Error::Deserialization`] if the private key is not a valid point on
/// the group or zero.
pub fn new_with_key(key: &[u8]) -> Result<Self> {
let sk = CS::Group::deserialize_scalar(key.into())?;
let pk = CS::Group::base_elem() * &sk;
@@ -487,8 +534,12 @@ where
/// bytes which are used as a seed to derive the server's private key.
///
/// Corresponds to DeriveKeyPair() function from the VOPRF specification.
///
/// # Errors
/// [`Error::Seed`] if the `seed` is empty or longer then [`u16::MAX`].
pub fn new_from_seed(seed: &[u8]) -> Result<Self> {
let sk = CS::Group::hash_to_scalar::<CS>(&[seed], Mode::Verifiable)?;
let sk =
CS::Group::hash_to_scalar::<CS>(&[seed], Mode::Verifiable).map_err(|_| Error::Seed)?;
let pk = CS::Group::base_elem() * &sk;
Ok(Self { sk, pk })
}
@@ -502,6 +553,9 @@ where
/// Computes the second step for the multiplicative blinding version of
/// DH-OPRF. This message is sent from the server (who holds the OPRF key)
/// to the client.
///
/// # Errors
/// [`Error::Metadata`] if the `metadata` is longer then `u16::MAX - 21`.
pub fn evaluate<R: RngCore + CryptoRng>(
&self,
rng: &mut R,
@@ -511,19 +565,16 @@ where
let VerifiableServerBatchEvaluatePrepareResult {
prepared_evaluation_elements: mut evaluation_elements,
t,
} = self.batch_evaluate_prepare(Some(blinded_element).into_iter(), metadata)?;
} = self.batch_evaluate_prepare(iter::once(blinded_element), metadata)?;
let prepared_element = [evaluation_elements.next().unwrap()];
// This can't fail because we know the size of the inputs.
let VerifiableServerBatchEvaluateFinishResult {
mut messages,
proof,
} = Self::batch_evaluate_finish(
rng,
Some(blinded_element).into_iter(),
&prepared_element,
&t,
)?;
} = Self::batch_evaluate_finish(rng, iter::once(blinded_element), &prepared_element, &t)
.unwrap();
let message = messages.next().unwrap();
@@ -533,6 +584,9 @@ where
/// Allows for batching of the evaluation of multiple [BlindedElement]
/// messages from a [VerifiableClient]
///
/// # Errors
/// [`Error::Metadata`] if the `metadata` is longer then `u16::MAX - 21`.
#[cfg(feature = "alloc")]
pub fn batch_evaluate<'a, R: RngCore + CryptoRng, I>(
&self,
@@ -552,13 +606,15 @@ where
let prepared_elements = evaluation_elements.collect();
// This can't fail because we know the size of the inputs.
let VerifiableServerBatchEvaluateFinishResult { messages, proof } =
Self::batch_evaluate_finish::<_, _, Vec<_>>(
rng,
blinded_elements.into_iter(),
&prepared_elements,
&t,
)?;
)
.unwrap();
Ok(VerifiableServerBatchEvaluateResult {
messages: messages.collect(),
@@ -570,6 +626,9 @@ where
/// memory allocation. Returned [`PreparedEvaluationElement`] have to be
/// [`collect`](Iterator::collect)ed and passed into
/// [`batch_evaluate_finish`](Self::batch_evaluate_finish).
///
/// # Errors
/// [`Error::Metadata`] if the `metadata` is longer then `u16::MAX - 21`.
pub fn batch_evaluate_prepare<'a, I: Iterator<Item = &'a BlindedElement<CS>>>(
&self,
blinded_elements: I,
@@ -583,10 +642,11 @@ where
// context = "Context-" || contextString || I2OSP(len(info), 2) || info
let context = GenericArray::from(STR_CONTEXT)
.concat(context_string)
.concat(i2osp_2(metadata.len())?);
.concat(i2osp_2(metadata.len()).map_err(|_| Error::Metadata)?);
let context = [&context, metadata];
let m = CS::Group::hash_to_scalar::<CS>(&context, Mode::Verifiable)?;
let m = CS::Group::hash_to_scalar::<CS>(&context, Mode::Verifiable)
.map_err(|_| Error::Metadata)?;
let t = self.sk + &m;
let evaluation_elements = blinded_elements
// To make a return type possible, we have to convert to a `fn` pointer, which isn't
@@ -604,6 +664,10 @@ where
/// See [`batch_evaluate_prepare`](Self::batch_evaluate_prepare) for more
/// details.
///
/// # Errors
/// [`Error::Batch`] if the number of `blinded_elements` and
/// `evaluation_elements` don't match or is longer then [`u16::MAX`].
pub fn batch_evaluate_finish<'a, 'b, R: RngCore + CryptoRng, IB, IE>(
rng: &mut R,
blinded_elements: IB,
@@ -837,6 +901,8 @@ type BlindResult<C> = (
);
// Inner function for blind. Returns the blind scalar and the blinded element
//
// Can only fail with [`Error::Input`].
fn blind<CS: CipherSuite, R: RngCore + CryptoRng>(
input: &[u8],
blinding_factor_rng: &mut R,
@@ -855,6 +921,8 @@ where
// Inner function for blind that assumes that the blinding factor has already
// been chosen, and therefore takes it as input. Does not check if the blinding
// factor is non-zero.
//
// Can only fail with [`Error::Input`].
fn deterministic_blind_unchecked<CS: CipherSuite>(
input: &[u8],
blind: &<CS::Group as Group>::Scalar,
@@ -864,7 +932,7 @@ where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
{
let hashed_point = CS::Group::hash_to_curve::<CS>(&[input], mode)?;
let hashed_point = CS::Group::hash_to_curve::<CS>(&[input], mode).map_err(|_| Error::Input)?;
Ok(hashed_point * blind)
}
@@ -884,6 +952,8 @@ type VerifiableUnblindResult<'a, CS, IC, IM> = Map<
) -> <<CS as CipherSuite>::Group as Group>::Elem,
>;
// Can only fail with [`Error::Metadata`], [`Error::Batch] or
// [`Error::ProofVerification`].
fn verifiable_unblind<'a, CS: 'a + CipherSuite, IC, IM>(
clients: &'a IC,
messages: &'a IM,
@@ -906,10 +976,12 @@ where
// context = "Context-" || contextString || I2OSP(len(info), 2) || info
let context = GenericArray::from(STR_CONTEXT)
.concat(context_string)
.concat(i2osp_2(info.len())?);
.concat(i2osp_2(info.len()).map_err(|_| Error::Metadata)?);
let context = [&context, info];
let m = CS::Group::hash_to_scalar::<CS>(&context, Mode::Verifiable)?;
// The `input` used here is the metadata.
let m =
CS::Group::hash_to_scalar::<CS>(&context, Mode::Verifiable).map_err(|_| Error::Metadata)?;
let g = CS::Group::base_elem();
let t = g * &m;
@@ -931,6 +1003,7 @@ where
.map(|(blind, x)| x.0 * &CS::Group::invert_scalar(blind)))
}
// Can only fail with [`Error::Batch`].
#[allow(clippy::many_single_char_names)]
fn generate_proof<CS: CipherSuite, R: RngCore + CryptoRng>(
rng: &mut R,
@@ -968,7 +1041,7 @@ where
// challengeDST = "Challenge-" || contextString
let challenge_dst =
GenericArray::from(STR_CHALLENGE).concat(get_context_string::<CS>(Mode::Verifiable));
let challenge_dst_len = i2osp_2_array(challenge_dst);
let challenge_dst_len = i2osp_2_array(&challenge_dst);
// h2Input = I2OSP(len(Bm), 2) || Bm ||
// I2OSP(len(a0), 2) || a0 ||
// I2OSP(len(a1), 2) || a1 ||
@@ -990,12 +1063,14 @@ where
&challenge_dst,
];
let c_scalar = CS::Group::hash_to_scalar::<CS>(&h2_input, Mode::Verifiable)?;
// This can't fail, the size of the `input` is known.
let c_scalar = CS::Group::hash_to_scalar::<CS>(&h2_input, Mode::Verifiable).unwrap();
let s_scalar = r - &(c_scalar * &k);
Ok(Proof { c_scalar, s_scalar })
}
// Can only fail with [`Error::ProofVerification`] or [`Error::Batch`].
#[allow(clippy::many_single_char_names)]
fn verify_proof<CS: CipherSuite>(
a: <CS::Group as Group>::Elem,
@@ -1029,7 +1104,7 @@ where
// challengeDST = "Challenge-" || contextString
let challenge_dst =
GenericArray::from(STR_CHALLENGE).concat(get_context_string::<CS>(Mode::Verifiable));
let challenge_dst_len = i2osp_2_array(challenge_dst);
let challenge_dst_len = i2osp_2_array(&challenge_dst);
// h2Input = I2OSP(len(Bm), 2) || Bm ||
// I2OSP(len(a0), 2) || a0 ||
// I2OSP(len(a1), 2) || a1 ||
@@ -1051,11 +1126,12 @@ where
&challenge_dst,
];
let c = CS::Group::hash_to_scalar::<CS>(&h2_input, Mode::Verifiable)?;
// This can't fail, the size of the `input` is known.
let c = CS::Group::hash_to_scalar::<CS>(&h2_input, Mode::Verifiable).unwrap();
match c.ct_eq(&proof.c_scalar).into() {
true => Ok(()),
false => Err(Error::ProofVerificationError),
false => Err(Error::ProofVerification),
}
}
@@ -1069,6 +1145,7 @@ type FinalizeAfterUnblindResult<'a, C, I, IE> = Map<
) -> Result<Output<<C as CipherSuite>::Hash>>,
>;
// Returned values can only fail with [`Error::Input`] or [`Error::Metadata`].
fn finalize_after_unblind<
'a,
CS: CipherSuite,
@@ -1078,7 +1155,7 @@ fn finalize_after_unblind<
inputs_and_unblinded_elements: IE,
info: &'a [u8],
mode: Mode,
) -> Result<FinalizeAfterUnblindResult<CS, I, IE>>
) -> FinalizeAfterUnblindResult<CS, I, IE>
where
<CS::Hash as OutputSizeUser>::OutputSize:
IsLess<U256> + IsLessOrEqual<<CS::Hash as BlockSizeUser>::BlockSize>,
@@ -1089,12 +1166,12 @@ where
// finalizeDST = "Finalize-" || contextString
let finalize_dst = GenericArray::from(STR_FINALIZE).concat(get_context_string::<CS>(mode));
Ok(inputs_and_unblinded_elements
inputs_and_unblinded_elements
// To make a return type possible, we have to convert to a `fn` pointer,
// which isn't possible if we `move` from context.
.zip(iter::repeat((info, finalize_dst)))
.map(|((input, unblinded_element), (info, finalize_dst))| {
let finalize_dst_len = i2osp_2_array(finalize_dst);
let finalize_dst_len = i2osp_2_array(&finalize_dst);
let elem_len = <CS::Group as Group>::ElemLen::U16.to_be_bytes();
// hashInput = I2OSP(len(input), 2) || input ||
@@ -1103,16 +1180,16 @@ where
// I2OSP(len(finalizeDST), 2) || finalizeDST
// return Hash(hashInput)
Ok(CS::Hash::new()
.chain_update(i2osp_2(input.as_ref().len())?)
.chain_update(i2osp_2(input.as_ref().len()).map_err(|_| Error::Input)?)
.chain_update(input.as_ref())
.chain_update(i2osp_2(info.len())?)
.chain_update(i2osp_2(info.len()).map_err(|_| Error::Metadata)?)
.chain_update(info)
.chain_update(elem_len)
.chain_update(CS::Group::serialize_elem(unblinded_element))
.chain_update(finalize_dst_len)
.chain_update(finalize_dst)
.finalize())
}))
})
}
type ComputeCompositesResult<C> = (
@@ -1120,6 +1197,7 @@ type ComputeCompositesResult<C> = (
<<C as CipherSuite>::Group as Group>::Elem,
);
// Can only fail with [`Error::Batch`].
fn compute_composites<CS: CipherSuite>(
k_option: Option<<CS::Group as Group>::Scalar>,
b: <CS::Group as Group>::Elem,
@@ -1135,23 +1213,23 @@ where
let elem_len = <CS::Group as Group>::ElemLen::U16.to_be_bytes();
if c_slice.len() != d_slice.len() {
return Err(Error::MismatchedLengthsForCompositeInputs);
return Err(Error::Batch);
}
let len = u16::try_from(c_slice.len()).map_err(|_| Error::SerializationError)?;
let len = u16::try_from(c_slice.len()).map_err(|_| Error::Batch)?;
let seed_dst = GenericArray::from(STR_SEED).concat(get_context_string::<CS>(Mode::Verifiable));
let composite_dst =
GenericArray::from(STR_COMPOSITE).concat(get_context_string::<CS>(Mode::Verifiable));
let composite_dst_len = i2osp_2_array(composite_dst);
let composite_dst_len = i2osp_2_array(&composite_dst);
let seed = CS::Hash::new()
.chain_update(&elem_len)
.chain_update(CS::Group::serialize_elem(b))
.chain_update(i2osp_2_array(seed_dst))
.chain_update(i2osp_2_array(&seed_dst))
.chain_update(seed_dst)
.finalize();
let seed_len = i2osp_2(seed.len())?;
let seed_len = i2osp_2_array(&seed);
let mut m = CS::Group::identity_elem();
let mut z = CS::Group::identity_elem();
@@ -1166,8 +1244,8 @@ where
// I2OSP(len(Di), 2) || Di ||
// I2OSP(len(compositeDST), 2) || compositeDST
let h2_input = [
&seed_len,
seed.as_slice(),
seed_len.as_slice(),
&seed,
&i.to_be_bytes(),
&elem_len,
&ci,
@@ -1176,7 +1254,8 @@ where
&composite_dst_len,
&composite_dst,
];
let di = CS::Group::hash_to_scalar::<CS>(&h2_input, Mode::Verifiable)?;
// This can't fail, the size of the `input` is known.
let di = CS::Group::hash_to_scalar::<CS>(&h2_input, Mode::Verifiable).unwrap();
m = c.0 * &di + &m;
z = match k_option {
Some(_) => z,
@@ -1243,8 +1322,7 @@ mod tests {
let res = point * &CS::Group::invert_scalar(key + &m);
finalize_after_unblind::<CS, _, _>(Some((input, res)).into_iter(), info, mode)
.unwrap()
finalize_after_unblind::<CS, _, _>(iter::once((input, res)), info, mode)
.next()
.unwrap()
.unwrap()
@@ -1259,7 +1337,7 @@ mod tests {
let info = b"info";
let mut rng = OsRng;
let client_blind_result = NonVerifiableClient::<CS>::blind(input, &mut rng).unwrap();
let server = NonVerifiableServer::<CS>::new(&mut rng).unwrap();
let server = NonVerifiableServer::<CS>::new(&mut rng);
let server_result = server
.evaluate(&client_blind_result.message, Some(info))
.unwrap();
@@ -1280,7 +1358,7 @@ mod tests {
let info = b"info";
let mut rng = OsRng;
let client_blind_result = VerifiableClient::<CS>::blind(input, &mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng);
let server_result = server
.evaluate(&mut rng, &client_blind_result.message, Some(info))
.unwrap();
@@ -1307,7 +1385,7 @@ mod tests {
let info = b"info";
let mut rng = OsRng;
let client_blind_result = VerifiableClient::<CS>::blind(input, &mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng);
let server_result = server
.evaluate(&mut rng, &client_blind_result.message, Some(info))
.unwrap();
@@ -1344,7 +1422,7 @@ mod tests {
client_states.push(client_blind_result.state);
client_messages.push(client_blind_result.message);
}
let server = VerifiableServer::<CS>::new(&mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng);
let VerifiableServerBatchEvaluatePrepareResult {
prepared_evaluation_elements,
t,
@@ -1399,7 +1477,7 @@ mod tests {
client_states.push(client_blind_result.state);
client_messages.push(client_blind_result.message);
}
let server = VerifiableServer::<CS>::new(&mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng);
let VerifiableServerBatchEvaluatePrepareResult {
prepared_evaluation_elements,
t,
@@ -1452,11 +1530,10 @@ mod tests {
let point = CS::Group::hash_to_curve::<CS>(&[&input], Mode::Base).unwrap();
let res2 = finalize_after_unblind::<CS, _, _>(
Some((input.as_ref(), point)).into_iter(),
iter::once((input.as_ref(), point)),
info,
Mode::Base,
)
.unwrap()
.next()
.unwrap()
.unwrap();
@@ -1511,7 +1588,7 @@ mod tests {
let info = b"info";
let mut rng = OsRng;
let client_blind_result = NonVerifiableClient::<CS>::blind(input, &mut rng).unwrap();
let server = NonVerifiableServer::<CS>::new(&mut rng).unwrap();
let server = NonVerifiableServer::<CS>::new(&mut rng);
let server_result = server
.evaluate(&client_blind_result.message, Some(info))
.unwrap();
@@ -1538,7 +1615,7 @@ mod tests {
let info = b"info";
let mut rng = OsRng;
let client_blind_result = VerifiableClient::<CS>::blind(input, &mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng).unwrap();
let server = VerifiableServer::<CS>::new(&mut rng);
let server_result = server
.evaluate(&mut rng, &client_blind_result.message, Some(info))
.unwrap();