use crate::error::SchnorrError;
use ark_ec::{AffineRepr, CurveGroup, VariableBaseMSM};
use ark_ff::PrimeField;
use ark_serialize::{CanonicalDeserialize, CanonicalSerialize};
use ark_std::{
cfg_iter,
collections::{BTreeMap, BTreeSet},
io::Write,
iter,
vec::Vec,
};
use core::ops::Add;
use digest::Digest;
use dock_crypto_utils::{
expect_equality, hashing_utils::field_elem_from_try_and_incr,
randomized_mult_checker::RandomizedMultChecker, serde_utils::ArkObjectBytes,
};
use serde::{Deserialize, Serialize};
use serde_with::serde_as;
use zeroize::{Zeroize, ZeroizeOnDrop};
#[cfg(feature = "parallel")]
use rayon::prelude::*;
pub trait SchnorrChallengeContributor {
fn challenge_contribution<W: Write>(&self, writer: W) -> Result<(), SchnorrError>;
}
#[serde_as]
#[derive(
Clone,
Debug,
PartialEq,
Eq,
Zeroize,
ZeroizeOnDrop,
CanonicalSerialize,
CanonicalDeserialize,
Serialize,
Deserialize,
)]
pub struct SchnorrCommitment<G: AffineRepr> {
#[serde_as(as = "Vec<ArkObjectBytes>")]
pub blindings: Vec<G::ScalarField>,
#[zeroize(skip)]
#[serde_as(as = "ArkObjectBytes")]
pub t: G,
}
impl<G: AffineRepr> SchnorrCommitment<G> {
pub fn new(bases: &[G], blindings: Vec<G::ScalarField>) -> Self {
let t = G::Group::msm_unchecked(bases, &blindings).into_affine();
Self { blindings, t }
}
pub fn response(
&self,
witnesses: &[G::ScalarField],
challenge: &G::ScalarField,
) -> Result<SchnorrResponse<G>, SchnorrError> {
expect_equality!(
self.blindings.len(),
witnesses.len(),
SchnorrError::ExpectedSameSizeSequences
);
let responses = cfg_iter!(self.blindings)
.zip(cfg_iter!(witnesses))
.map(|(b, w)| *b + (*w * *challenge))
.collect::<Vec<_>>();
Ok(SchnorrResponse(responses))
}
}
impl<G: AffineRepr> SchnorrChallengeContributor for SchnorrCommitment<G> {
fn challenge_contribution<W: Write>(&self, writer: W) -> Result<(), SchnorrError> {
self.t.serialize_compressed(writer).map_err(|e| e.into())
}
}
#[serde_as]
#[derive(
Clone, Debug, PartialEq, Eq, CanonicalSerialize, CanonicalDeserialize, Serialize, Deserialize,
)]
pub struct SchnorrResponse<G: AffineRepr>(
#[serde_as(as = "Vec<ArkObjectBytes>")] pub Vec<G::ScalarField>,
);
impl<G: AffineRepr> SchnorrResponse<G> {
pub fn is_valid(
&self,
bases: &[G],
y: &G,
t: &G,
challenge: &G::ScalarField,
) -> Result<(), SchnorrError> {
expect_equality!(
self.0.len(),
bases.len(),
SchnorrError::ExpectedSameSizeSequences
);
if (G::Group::msm_unchecked(bases, &self.0).add(y.mul_bigint((-*challenge).into_bigint())))
.into_affine()
== *t
{
Ok(())
} else {
Err(SchnorrError::InvalidResponse)
}
}
pub fn verify_using_randomized_mult_checker(
&self,
bases: Vec<G>,
y: G,
t: G,
challenge: &G::ScalarField,
rmc: &mut RandomizedMultChecker<G>,
) -> Result<(), SchnorrError> {
expect_equality!(
self.0.len(),
bases.len(),
SchnorrError::ExpectedSameSizeSequences
);
rmc.add_many(
bases.into_iter().chain(iter::once(y)),
self.0.iter().chain(iter::once(&-*challenge)),
t,
);
Ok(())
}
pub fn get_response(&self, idx: usize) -> Result<&G::ScalarField, SchnorrError> {
if idx >= self.0.len() {
Err(SchnorrError::IndexOutOfBounds(idx, self.0.len()))
} else {
Ok(&self.0[idx])
}
}
pub fn get_responses(
&self,
ids: &BTreeSet<usize>,
) -> Result<BTreeMap<usize, G::ScalarField>, SchnorrError> {
let mut resp = BTreeMap::new();
for i in ids {
match self.0.get(*i) {
Some(r) => {
resp.insert(*i, *r);
}
_ => return Err(SchnorrError::IndexOutOfBounds(*i, self.0.len())),
}
}
Ok(resp)
}
pub fn len(&self) -> usize {
self.0.len()
}
}
pub fn compute_random_oracle_challenge<F: PrimeField, D: Digest>(challenge_bytes: &[u8]) -> F {
field_elem_from_try_and_incr::<F, D>(challenge_bytes)
}
#[cfg(test)]
mod tests {
use super::*;
use ark_bls12_381::{Fr, G1Affine, G1Projective, G2Affine, G2Projective};
use ark_ec::VariableBaseMSM;
use ark_std::{
rand::{rngs::StdRng, SeedableRng},
UniformRand,
};
#[macro_export]
macro_rules! test_serialization {
($obj_type:ty, $obj: ident) => {
let mut serz = vec![];
ark_serialize::CanonicalSerialize::serialize_compressed(&$obj, &mut serz).unwrap();
let deserz: $obj_type =
ark_serialize::CanonicalDeserialize::deserialize_compressed(&serz[..]).unwrap();
assert_eq!(deserz, $obj);
let mut serz = vec![];
$obj.serialize_compressed(&mut serz).unwrap();
let deserz: $obj_type =
CanonicalDeserialize::deserialize_compressed(&serz[..]).unwrap();
assert_eq!(deserz, $obj);
let obj_ser = serde_json::to_string(&$obj).unwrap();
let obj_deser = serde_json::from_str::<$obj_type>(&obj_ser).unwrap();
assert_eq!($obj, obj_deser);
let ser = rmp_serde::to_vec_named(&$obj).unwrap();
let deser = rmp_serde::from_slice::<$obj_type>(&ser).unwrap();
assert_eq!($obj, deser);
};
}
macro_rules! test_schnorr_in_group {
( $group_element_proj:ident, $group_element_affine:ident ) => {
let mut rng = StdRng::seed_from_u64(0u64);
let count = 10;
let bases = (0..count)
.into_iter()
.map(|_| $group_element_proj::rand(&mut rng).into_affine())
.collect::<Vec<_>>();
let witnesses = (0..count)
.into_iter()
.map(|_| Fr::rand(&mut rng))
.collect::<Vec<_>>();
let y = $group_element_proj::msm_unchecked(&bases, &witnesses).into_affine();
let blindings = (0..count)
.into_iter()
.map(|_| Fr::rand(&mut rng))
.collect::<Vec<_>>();
let comm = SchnorrCommitment::new(&bases, blindings.clone());
test_serialization!(SchnorrCommitment<$group_element_affine>, comm);
let challenge = Fr::rand(&mut rng);
let resp = comm.response(&witnesses, &challenge).unwrap();
resp.is_valid(&bases, &y, &comm.t, &challenge).unwrap();
test_serialization!(SchnorrResponse<$group_element_affine>, resp);
let mut checker = RandomizedMultChecker::new_using_rng(&mut rng);
resp.verify_using_randomized_mult_checker(bases, y, comm.t, &challenge, &mut checker)
.unwrap();
assert!(checker.verify());
};
}
#[test]
fn schnorr_vector() {
test_schnorr_in_group!(G1Projective, G1Affine);
test_schnorr_in_group!(G2Projective, G2Affine);
}
}