Adding trait implementations (#12)

This commit is contained in:
Kevin Lewi
2021-09-28 19:44:57 -07:00
committed by GitHub
parent 4206fdbf48
commit 85fb480b78
6 changed files with 330 additions and 175 deletions
+1 -2
View File
@@ -26,8 +26,7 @@ Rust **1.51** or higher.
Contributors
------------
The author of this code is Kevin Lewi
([@kevinlewi](https://github.com/kevinlewi)) .
The author of this code is Kevin Lewi ([@kevinlewi](https://github.com/kevinlewi)).
To learn more about contributing to this project, [see this document](./CONTRIBUTING.md).
License
+5
View File
@@ -119,6 +119,11 @@ pub trait Group:
/// Compares in constant time if the scalars are equal
fn ct_equal_scalar(s1: &Self::Scalar, s2: &Self::Scalar) -> bool;
/// Set the contents of self to the identity value
fn zeroize(&mut self) {
*self = <Self as Group>::identity();
}
}
#[cfg(test)]
+203
View File
@@ -0,0 +1,203 @@
// Copyright (c) Facebook, Inc. and its affiliates.
//
// This source code is licensed under both the MIT license found in the
// LICENSE-MIT file in the root directory of this source tree and the Apache
// License, Version 2.0 found in the LICENSE-APACHE file in the root directory
// of this source tree.
macro_rules! impl_debug_eq_hash_for {
(struct $name:ident$(<$($gen:ident$(: $bound:tt)?),+$(,)?>)?, [$field1:ident$(, $field2:ident)*$(,)?]$(, )?$([$($type:ty),+$(,)?]$(,)?)?) => {
impl$(<$($gen$(: $bound)?),+>)? core::fmt::Debug for $name$(<$($gen),+>)?
$(where $($type: core::fmt::Debug,)+)?
{
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("$name")
.field("$field1", &self.$field1)
$(.field("$field2", &self.$field2))*
.finish()
}
}
impl$(<$($gen$(: $bound)?),+>)? Eq for $name$(<$($gen),+>)?
$(where $($type: Eq,)+)?
{}
impl$(<$($gen$(: $bound)?),+>)? PartialEq for $name$(<$($gen),+>)?
$(where $($type: PartialEq,)+)?
{
fn eq(&self, other: &Self) -> bool {
PartialEq::eq(&self.$field1, &other.$field1)
$(&& PartialEq::eq(&self.$field2, &other.$field2))*
}
}
impl$(<$($gen$(: $bound)?),+>)? core::hash::Hash for $name$(<$($gen),+>)?
$(where $($type: core::hash::Hash,)+)?
{
fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
core::hash::Hash::hash(&self.$field1, state);
$(core::hash::Hash::hash(&self.$field2, state);)*
}
}
};
(tuple $name:ident$(<$($gen:ident$(: $bound:tt)?),+$(,)?>)?, [$field1:tt$(, $field2:tt)*$(,)?]$(, )?$([$($type:ty),+$(,)?]$(,)?)?) => {
impl$(<$($gen$(: $bound)?),+>)? core::fmt::Debug for $name$(<$($gen),+>)?
$(where $($type: core::fmt::Debug,)+)?
{
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_tuple("$name")
.field(&self.$field1)
$(.field(&self.$field2))*
.finish()
}
}
impl$(<$($gen$(: $bound)?),+>)? Eq for $name$(<$($gen),+>)?
$(where $($type: Eq,)+)?
{}
impl$(<$($gen$(: $bound)?),+>)? PartialEq for $name$(<$($gen),+>)?
$(where $($type: PartialEq,)+)?
{
fn eq(&self, other: &Self) -> bool {
PartialEq::eq(&self.$field1, &other.$field1)
$(&& PartialEq::eq(&self.$field2, &other.$field2))*
}
}
impl$(<$($gen$(: $bound)?),+>)? core::hash::Hash for $name$(<$($gen),+>)?
$(where $($type: core::hash::Hash,)+)?
{
fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
core::hash::Hash::hash(&self.$field1, state);
$(core::hash::Hash::hash(&self.$field2, state);)*
}
}
};
}
macro_rules! impl_clone_for {
(struct $name:ident$(<$($gen:ident$(: $bound:tt)?),+$(,)?>)?, [$field1:ident$(, $field2:ident)*$(,)?]$(, )?$([$($type:ty),+$(,)?]$(,)?)?) => {
impl$(<$($gen$(: $bound)?),+>)? Clone for $name$(<$($gen),+>)?
$(where $($type: Clone,)+)?
{
fn clone(&self) -> Self {
Self {
$field1: self.$field1.clone(),
$($field2: self.$field2.clone(),)*
}
}
}
};
(tuple $name:ident$(<$($gen:ident$(: $bound:tt)?),+$(,)?>)?, [$field1:tt$(, $field2:tt)*$(,)?]$(, )?$([$($type:ty),+$(,)?]$(,)?)?) => {
impl$(<$($gen$(: $bound)?),+>)? Clone for $name$(<$($gen),+>)?
$(where $($type: Clone,)+)?
{
fn clone(&self) -> Self {
Self(
self.$field1.clone(),
$(self.$field2.clone(),)*
)
}
}
};
}
macro_rules! impl_zeroize_on_drop_for {
(struct $name:ident$(<$($gen:ident$(: $bound:tt)?),+$(,)?>)?, [$field1:ident$(, $field2:ident)*$(,)?]$(, )?$([$($type:ty),+$(,)?]$(,)?)?) => {
impl$(<$($gen$(: $bound)?),+>)? zeroize::Zeroize for $name$(<$($gen),+>)?
{
fn zeroize(&mut self) {
self.$field1.zeroize();
$(self.$field2.zeroize();)*
}
}
impl$(<$($gen$(: $bound)?),+>)? Drop for $name$(<$($gen),+>)?
{
fn drop(&mut self) {
#[allow(unused_imports)]
use zeroize::Zeroize;
self.$field1.zeroize();
$(self.$field2.zeroize();)*
}
}
};
}
/// Inner macro used for deriving `serde`'s `Serialize` and `Deserialize` traits.
macro_rules! impl_serialize_and_deserialize_for {
($t:ident) => {
#[cfg(feature = "serialize")]
impl<CS: CipherSuite> serde::Serialize for $t<CS> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
if serializer.is_human_readable() {
serializer.serialize_str(&base64::encode(&self.serialize()))
} else {
serializer.serialize_bytes(&self.serialize())
}
}
}
#[cfg(feature = "serialize")]
impl<'de, CS: CipherSuite> serde::Deserialize<'de> for $t<CS> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
if deserializer.is_human_readable() {
let s = <&str>::deserialize(deserializer)?;
$t::<CS>::deserialize(&base64::decode(s).map_err(serde::de::Error::custom)?)
.map_err(serde::de::Error::custom)
} else {
struct ByteVisitor<CS: CipherSuite> {
marker: core::marker::PhantomData<CS>,
}
impl<'de, CS: CipherSuite> serde::de::Visitor<'de> for ByteVisitor<CS> {
type Value = $t<CS>;
fn expecting(
&self,
formatter: &mut core::fmt::Formatter,
) -> core::fmt::Result {
formatter.write_str(core::concat!(
"the byte representation of a ",
core::stringify!($t)
))
}
fn visit_bytes<E>(self, value: &[u8]) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
$t::<CS>::deserialize(value).map_err(|_| {
serde::de::Error::invalid_value(
serde::de::Unexpected::Bytes(value),
&core::concat!(
"invalid byte sequence for ",
core::stringify!($t)
),
)
})
}
}
deserializer.deserialize_bytes(ByteVisitor::<CS> {
marker: core::marker::PhantomData,
})
}
}
}
};
}
// Convenience macro for implementing all of the above traits
macro_rules! impl_traits_for {
(struct $name:ident$(<$($gen:ident$(: $bound:tt)?),+$(,)?>)?, [$field1:ident$(, $field2:ident)*$(,)?]$(, )?$([$($type:ty),+$(,)?]$(,)?)?) => {
impl_debug_eq_hash_for!(struct $name$(<$($gen$(: $bound)?),+>)?, [$field1$(, $field2)*], $([$($type),+])?);
impl_clone_for!(struct $name$(<$($gen$(: $bound)?),+>)?, [$field1$(, $field2)*], $([$($type),+])?);
impl_zeroize_on_drop_for!(struct $name$(<$($gen$(: $bound)?),+>)?, [$field1$(, $field2)*], $([$($type),+])?);
impl_serialize_and_deserialize_for!($name);
}
}
+4 -1
View File
@@ -459,6 +459,8 @@
//!
//! # Features
//!
//! - The `p256` feature enables using p256 as the underlying group for the [Ciphersuite] choice
//!
//! - The `serialize` feature, enabled by default, provides convenience functions for serializing and deserializing with
//! [serde](https://serde.rs/).
//!
@@ -472,9 +474,10 @@
extern crate alloc;
#[macro_use]
mod impls;
#[macro_use]
mod serialization;
mod ciphersuite;
pub mod errors;
pub mod group;
+8 -89
View File
@@ -20,80 +20,11 @@ use crate::{
use alloc::vec::Vec;
use generic_array::{typenum::Unsigned, GenericArray};
/// Inner macro used for deriving `serde`'s `Serialize` and `Deserialize` traits.
macro_rules! impl_serialize_and_deserialize_for {
($t:ident) => {
#[cfg(feature = "serialize")]
impl<CS: CipherSuite> serde::Serialize for $t<CS> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
if serializer.is_human_readable() {
serializer.serialize_str(&base64::encode(&self.serialize()))
} else {
serializer.serialize_bytes(&self.serialize())
}
}
}
#[cfg(feature = "serialize")]
impl<'de, CS: CipherSuite> serde::Deserialize<'de> for $t<CS> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
if deserializer.is_human_readable() {
let s = <&str>::deserialize(deserializer)?;
$t::<CS>::deserialize(&base64::decode(s).map_err(serde::de::Error::custom)?)
.map_err(serde::de::Error::custom)
} else {
struct ByteVisitor<CS: CipherSuite> {
marker: core::marker::PhantomData<CS>,
}
impl<'de, CS: CipherSuite> serde::de::Visitor<'de> for ByteVisitor<CS> {
type Value = $t<CS>;
fn expecting(
&self,
formatter: &mut core::fmt::Formatter,
) -> core::fmt::Result {
formatter.write_str(core::concat!(
"the byte representation of a ",
core::stringify!($t)
))
}
fn visit_bytes<E>(self, value: &[u8]) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
$t::<CS>::deserialize(value).map_err(|_| {
serde::de::Error::invalid_value(
serde::de::Unexpected::Bytes(value),
&core::concat!(
"invalid byte sequence for ",
core::stringify!($t)
),
)
})
}
}
deserializer.deserialize_bytes(ByteVisitor::<CS> {
marker: core::marker::PhantomData,
})
}
}
}
};
}
//////////////////////////////////////////////////////////
// Serialization and Deserialization for High-Level API //
// ==================================================== //
//////////////////////////////////////////////////////////
impl_serialize_and_deserialize_for!(NonVerifiableClient);
impl<CS: CipherSuite> NonVerifiableClient<CS> {
/// Serialization into bytes
pub fn serialize(&self) -> Vec<u8> {
@@ -118,8 +49,6 @@ impl<CS: CipherSuite> NonVerifiableClient<CS> {
}
}
impl_serialize_and_deserialize_for!(VerifiableClient);
impl<CS: CipherSuite> VerifiableClient<CS> {
/// Serialization into bytes
pub fn serialize(&self) -> Vec<u8> {
@@ -153,8 +82,6 @@ impl<CS: CipherSuite> VerifiableClient<CS> {
}
}
impl_serialize_and_deserialize_for!(NonVerifiableServer);
impl<CS: CipherSuite> NonVerifiableServer<CS> {
/// Serialization into bytes
pub fn serialize(&self) -> Vec<u8> {
@@ -174,8 +101,6 @@ impl<CS: CipherSuite> NonVerifiableServer<CS> {
}
}
impl_serialize_and_deserialize_for!(VerifiableServer);
impl<CS: CipherSuite> VerifiableServer<CS> {
/// Serialization into bytes
pub fn serialize(&self) -> Vec<u8> {
@@ -201,8 +126,6 @@ impl<CS: CipherSuite> VerifiableServer<CS> {
}
}
impl_serialize_and_deserialize_for!(Proof);
impl<CS: CipherSuite> Proof<CS> {
/// Serialization into bytes
pub fn serialize(&self) -> Vec<u8> {
@@ -226,35 +149,31 @@ impl<CS: CipherSuite> Proof<CS> {
}
}
impl_serialize_and_deserialize_for!(BlindedElement);
impl<CS: CipherSuite> BlindedElement<CS> {
/// Serialization into bytes
pub fn serialize(&self) -> Vec<u8> {
self.0.to_arr().to_vec()
self.value.to_arr().to_vec()
}
/// Deserialization from bytes
pub fn deserialize(input: &[u8]) -> Result<Self, InternalError> {
Ok(Self(CS::Group::from_element_slice(
GenericArray::from_slice(input),
)?))
Ok(Self {
value: CS::Group::from_element_slice(GenericArray::from_slice(input))?,
})
}
}
impl_serialize_and_deserialize_for!(EvaluationElement);
impl<CS: CipherSuite> EvaluationElement<CS> {
/// Serialization into bytes
pub fn serialize(&self) -> Vec<u8> {
self.0.to_arr().to_vec()
self.value.to_arr().to_vec()
}
/// Deserialization from bytes
pub fn deserialize(input: &[u8]) -> Result<Self, InternalError> {
Ok(Self(CS::Group::from_element_slice(
GenericArray::from_slice(input),
)?))
Ok(Self {
value: CS::Group::from_element_slice(GenericArray::from_slice(input))?,
})
}
}
+109 -83
View File
@@ -13,13 +13,12 @@ use crate::{
group::Group,
serialization::{i2osp, serialize},
};
use alloc::vec;
use alloc::vec::Vec;
use digest::Digest;
use generic_array::{typenum::Unsigned, GenericArray};
use rand::{CryptoRng, RngCore};
use alloc::vec;
use alloc::vec::Vec;
///////////////
// Constants //
// ========= //
@@ -53,6 +52,11 @@ pub struct NonVerifiableClient<CS: CipherSuite> {
pub(crate) blind: <CS::Group as Group>::Scalar,
pub(crate) data: Vec<u8>,
}
impl_traits_for!(
struct NonVerifiableClient<CS: CipherSuite>,
[blind, data],
[<CS::Group as Group>::Scalar],
);
/// A client which engages with a [VerifiableServer]
/// in verifiable mode, meaning that the OPRF outputs
@@ -62,6 +66,11 @@ pub struct VerifiableClient<CS: CipherSuite> {
pub(crate) blinded_element: CS::Group,
pub(crate) data: alloc::vec::Vec<u8>,
}
impl_traits_for!(
struct VerifiableClient<CS: CipherSuite>,
[blind, blinded_element, data],
[<CS::Group as Group>::Scalar, CS::Group],
);
/// A server which engages with a [NonVerifiableClient]
/// in base mode, meaning that the OPRF outputs are not
@@ -69,6 +78,12 @@ pub struct VerifiableClient<CS: CipherSuite> {
pub struct NonVerifiableServer<CS: CipherSuite> {
pub(crate) sk: <CS::Group as Group>::Scalar,
}
impl_traits_for!(
struct NonVerifiableServer<CS: CipherSuite>,
[sk],
[<CS::Group as Group>::Scalar],
);
/// A server which engages with a [VerifiableClient]
/// in verifiable mode, meaning that the OPRF outputs
/// can be checked against a server public key.
@@ -76,6 +91,11 @@ pub struct VerifiableServer<CS: CipherSuite> {
pub(crate) sk: <CS::Group as Group>::Scalar,
pub(crate) pk: CS::Group,
}
impl_traits_for!(
struct VerifiableServer<CS: CipherSuite>,
[sk, pk],
[<CS::Group as Group>::Scalar, CS::Group],
);
/// A proof produced by a [VerifiableServer] that
/// the OPRF output matches against a server public key.
@@ -83,15 +103,34 @@ pub struct Proof<CS: CipherSuite> {
pub(crate) c_scalar: <CS::Group as Group>::Scalar,
pub(crate) s_scalar: <CS::Group as Group>::Scalar,
}
impl_traits_for!(
struct Proof<CS: CipherSuite>,
[c_scalar, s_scalar],
[<CS::Group as Group>::Scalar],
);
/// The first client message sent from a client (either verifiable or not)
/// to a server (either verifiable or not).
pub struct BlindedElement<CS: CipherSuite>(pub(crate) CS::Group);
pub struct BlindedElement<CS: CipherSuite> {
pub(crate) value: CS::Group,
}
impl_traits_for!(
struct BlindedElement<CS: CipherSuite>,
[value],
[CS::Group],
);
/// The server's response to the [BlindedElement] message from
/// a client (either verifiable or not)
/// to a server (either verifiable or not).
pub struct EvaluationElement<CS: CipherSuite>(pub(crate) CS::Group);
pub struct EvaluationElement<CS: CipherSuite> {
pub(crate) value: CS::Group,
}
impl_traits_for!(
struct EvaluationElement<CS: CipherSuite>,
[value],
[CS::Group],
);
/////////////////////////
// API Implementations //
@@ -110,7 +149,9 @@ impl<CS: CipherSuite> NonVerifiableClient<CS> {
data: input.to_vec(),
blind,
},
message: BlindedElement(blinded_element),
message: BlindedElement {
value: blinded_element,
},
})
}
@@ -122,7 +163,7 @@ impl<CS: CipherSuite> NonVerifiableClient<CS> {
metadata: &Metadata,
) -> Result<NonVerifiableClientFinalizeResult<CS>, InternalError> {
let unblinded_element =
evaluation_element.0 * &<CS::Group as Group>::scalar_invert(&self.blind);
evaluation_element.value * &<CS::Group as Group>::scalar_invert(&self.blind);
let outputs = finalize_after_unblind::<CS>(
&[(self.data.clone(), unblinded_element)],
&metadata.0,
@@ -163,7 +204,9 @@ impl<CS: CipherSuite> VerifiableClient<CS> {
blind,
blinded_element,
},
message: BlindedElement(blinded_element),
message: BlindedElement {
value: blinded_element,
},
})
}
@@ -199,7 +242,9 @@ impl<CS: CipherSuite> VerifiableClient<CS> {
.map(|(client, evaluation_element)| BatchItems {
blind: client.blind,
evaluation_element: evaluation_element.clone(),
blinded_element: BlindedElement(client.blinded_element),
blinded_element: BlindedElement {
value: client.blinded_element,
},
})
.collect();
@@ -289,9 +334,11 @@ impl<CS: CipherSuite> NonVerifiableServer<CS> {
let dst = [STR_HASH_TO_SCALAR, &get_context_string::<CS>(Mode::Base)?].concat();
let m = CS::Group::hash_to_scalar::<CS::Hash>(&context, &dst)?;
let t = self.sk + &m;
let evaluation_element = blinded_element.0 * &CS::Group::scalar_invert(&t);
let evaluation_element = blinded_element.value * &CS::Group::scalar_invert(&t);
Ok(NonVerifiableServerEvaluateResult {
message: EvaluationElement(evaluation_element),
message: EvaluationElement {
value: evaluation_element,
},
})
}
}
@@ -370,7 +417,9 @@ impl<CS: CipherSuite> VerifiableServer<CS> {
let t = self.sk + &m;
let evaluation_elements: Vec<EvaluationElement<CS>> = blinded_elements
.iter()
.map(|x| EvaluationElement(x.0 * &CS::Group::scalar_invert(&t)))
.map(|x| EvaluationElement {
value: x.value * &CS::Group::scalar_invert(&t),
})
.collect();
let g = CS::Group::base_point();
@@ -499,28 +548,6 @@ struct BatchItems<CS: CipherSuite> {
blinded_element: BlindedElement<CS>,
}
impl<CS: CipherSuite> Clone for BlindedElement<CS> {
fn clone(&self) -> Self {
Self(self.0)
}
}
impl<CS: CipherSuite> Clone for EvaluationElement<CS> {
fn clone(&self) -> Self {
Self(self.0)
}
}
impl<CS: CipherSuite> Clone for VerifiableClient<CS> {
fn clone(&self) -> Self {
Self {
data: self.data.clone(),
blind: self.blind,
blinded_element: self.blinded_element,
}
}
}
// Inner function for blind. Returns the blind scalar and the blinded element
fn blind<CS: CipherSuite, R: RngCore + CryptoRng>(
input: &[u8],
@@ -574,7 +601,7 @@ fn verifiable_unblind<CS: CipherSuite>(
let unblinded_elements = blinds
.iter()
.zip(evaluation_elements.iter())
.map(|(&blind, x)| x.0 * &CS::Group::scalar_invert(&blind))
.map(|(&blind, x)| x.value * &CS::Group::scalar_invert(&blind))
.collect();
Ok(unblinded_elements)
}
@@ -705,8 +732,8 @@ fn compute_composites<CS: CipherSuite>(
let h2_input = [
serialize(&seed, 2)?,
i2osp(i, 2)?,
serialize(&c_slice[i].0.to_arr().to_vec(), 2)?,
serialize(&d_slice[i].0.to_arr().to_vec(), 2)?,
serialize(&c_slice[i].value.to_arr().to_vec(), 2)?,
serialize(&d_slice[i].value.to_arr().to_vec(), 2)?,
serialize(&composite_dst, 2)?,
]
.concat();
@@ -716,10 +743,10 @@ fn compute_composites<CS: CipherSuite>(
]
.concat();
let di = CS::Group::hash_to_scalar::<CS::Hash>(&h2_input, &dst)?;
m = c_slice[i].0 * &di + &m;
m = c_slice[i].value * &di + &m;
z = match k_option {
Some(_) => z,
None => d_slice[i].0 * &di + &z,
None => d_slice[i].value * &di + &z,
};
}
@@ -746,71 +773,55 @@ fn get_context_string<CS: CipherSuite>(mode: Mode) -> Result<alloc::vec::Vec<u8>
// Tests //
// ===== //
///////////
#[cfg(test)]
mod tests {
use super::*;
use crate::group::Group;
use curve25519_dalek::ristretto::RistrettoPoint;
use generic_array::{arr, GenericArray};
use rand::rngs::OsRng;
use sha2::Sha512;
struct Ristretto255Sha512;
impl CipherSuite for Ristretto255Sha512 {
type Group = RistrettoPoint;
type Hash = Sha512;
}
fn prf(
fn prf<CS: CipherSuite>(
input: &[u8],
oprf_key: &[u8],
info: &[u8],
) -> GenericArray<u8, <Sha512 as Digest>::OutputSize> {
) -> GenericArray<u8, <CS::Hash as Digest>::OutputSize> {
let dst = [
STR_HASH_TO_GROUP,
&get_context_string::<Ristretto255Sha512>(Mode::Base).unwrap(),
&get_context_string::<CS>(Mode::Base).unwrap(),
]
.concat();
let point = RistrettoPoint::hash_to_curve::<Sha512>(input, &dst).unwrap();
let scalar =
RistrettoPoint::from_scalar_slice(GenericArray::from_slice(&oprf_key[..])).unwrap();
let point = CS::Group::hash_to_curve::<CS::Hash>(input, &dst).unwrap();
let scalar = CS::Group::from_scalar_slice(GenericArray::from_slice(&oprf_key[..])).unwrap();
let context = [
STR_CONTEXT,
&get_context_string::<Ristretto255Sha512>(Mode::Base).unwrap(),
&get_context_string::<CS>(Mode::Base).unwrap(),
&serialize(info, 2).unwrap(),
]
.concat();
let dst = [
STR_HASH_TO_SCALAR,
&get_context_string::<Ristretto255Sha512>(Mode::Base).unwrap(),
&get_context_string::<CS>(Mode::Base).unwrap(),
]
.concat();
let m = <<Ristretto255Sha512 as CipherSuite>::Group as Group>::hash_to_scalar::<
<Ristretto255Sha512 as CipherSuite>::Hash,
>(&context, &dst)
.unwrap();
let m = <CS::Group as Group>::hash_to_scalar::<CS::Hash>(&context, &dst).unwrap();
let res = point
* &<<Ristretto255Sha512 as CipherSuite>::Group as Group>::scalar_invert(&(scalar + m));
let res = point * &<CS::Group as Group>::scalar_invert(&(scalar + &m));
finalize_after_unblind::<Ristretto255Sha512>(&[(input.to_vec(), res)], info, Mode::Base)
.unwrap()[0]
finalize_after_unblind::<CS>(&[(input.to_vec(), res)], info, Mode::Base).unwrap()[0].clone()
}
#[test]
fn oprf_retrieval() {
fn oprf_retrieval<CS: CipherSuite>() {
let input = b"hunter2";
let info = b"info";
let mut rng = OsRng;
let client_blind_result =
NonVerifiableClient::<Ristretto255Sha512>::blind(&input[..], &mut rng).unwrap();
let client_blind_result = NonVerifiableClient::<CS>::blind(&input[..], &mut rng).unwrap();
let oprf_key_bytes = arr![
u8; 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23,
24, 25, 26, 27, 28, 29, 30, 31, 32,
];
let server =
NonVerifiableServer::<Ristretto255Sha512>::new_with_key(&oprf_key_bytes).unwrap();
let server = NonVerifiableServer::<CS>::new_with_key(&oprf_key_bytes).unwrap();
let server_result = server
.evaluate(client_blind_result.message, &Metadata(info.to_vec()))
.unwrap();
@@ -818,39 +829,54 @@ mod tests {
.state
.finalize(server_result.message, &Metadata(info.to_vec()))
.unwrap();
let res2 = prf(&input[..], &oprf_key_bytes, info);
let res2 = prf::<CS>(&input[..], &oprf_key_bytes, info);
assert_eq!(client_finalize_result.output, res2);
}
#[test]
fn oprf_inversion_unsalted() {
fn oprf_inversion_unsalted<CS: CipherSuite>() {
let mut rng = OsRng;
let mut input = alloc::vec![0u8; 64];
rng.fill_bytes(&mut input);
let info = b"info";
let client_blind_result =
NonVerifiableClient::<Ristretto255Sha512>::blind(&input, &mut rng).unwrap();
let client_blind_result = NonVerifiableClient::<CS>::blind(&input, &mut rng).unwrap();
let client_finalize_result = client_blind_result
.state
.finalize(
EvaluationElement(client_blind_result.message.0),
EvaluationElement {
value: client_blind_result.message.value,
},
&Metadata(info.to_vec()),
)
.unwrap();
let dst = [
STR_HASH_TO_GROUP,
&get_context_string::<Ristretto255Sha512>(Mode::Base).unwrap(),
&get_context_string::<CS>(Mode::Base).unwrap(),
]
.concat();
let point = RistrettoPoint::hash_to_curve::<Sha512>(&input, &dst).unwrap();
let res2 = finalize_after_unblind::<Ristretto255Sha512>(
&[(input.to_vec(), point)],
info,
Mode::Base,
)
.unwrap()[0];
let point = CS::Group::hash_to_curve::<CS::Hash>(&input, &dst).unwrap();
let res2 = finalize_after_unblind::<CS>(&[(input.to_vec(), point)], info, Mode::Base)
.unwrap()[0]
.clone();
assert_eq!(client_finalize_result.output, res2);
}
#[test]
fn test_functionality() -> Result<(), InternalError> {
use crate::tests::Ristretto255Sha512;
oprf_retrieval::<Ristretto255Sha512>();
oprf_inversion_unsalted::<Ristretto255Sha512>();
#[cfg(feature = "p256")]
{
use crate::tests::P256Sha256;
oprf_retrieval::<P256Sha256>();
oprf_inversion_unsalted::<P256Sha256>();
}
Ok(())
}
}