use bacteria::Transcript;
use mohan::{
ser,
dalek::{
ristretto::RistrettoPoint,
scalar::Scalar,
constants::RISTRETTO_BASEPOINT_COMPRESSED,
traits::Identity
},
mohan_rand,
hash::hash_to_ristretto
};
use crate::ZeiError;
use serde::{ Serialize, Deserialize };
#[derive(Clone, Debug, Serialize, Deserialize)]
#[allow(non_snake_case)]
pub struct BatchEqualityProof {
pub(crate) c: Scalar,
pub(crate) z: Scalar,
}
impl Default for BatchEqualityProof {
fn default() -> BatchEqualityProof {
BatchEqualityProof {
c: Scalar::zero(),
z: Scalar::zero()
}
}
}
impl BatchEqualityProof {
pub fn proove(
cin_asset: &Vec<RistrettoPoint>,
cin_blind: &Vec<Scalar>,
cout_asset: &Vec<RistrettoPoint>,
cout_blind: &Vec<Scalar>
) -> Result<BatchEqualityProof, ZeiError> {
let mut transcript = Transcript::new(b"zei_equality_proof");
for i in cin_asset.iter() {
transcript.commit_point(b"commitment", &i.compress());
}
for i in cout_asset.iter() {
transcript.commit_point(b"commitment", &i.compress());
}
let mut all_comms = vec![];
all_comms.extend(cin_asset.iter().cloned());
all_comms.extend(cout_asset.iter().cloned());
let mut all_blinds = vec![];
all_blinds.extend(cin_blind.iter().cloned());
all_blinds.extend(cout_blind.iter().cloned());
let mut big_d = RistrettoPoint::identity();
let mut a_blind : Scalar = Scalar::zero();
assert!(all_comms.len() == all_blinds.len());
for i in 0..all_comms.len() {
transcript.append_u64(b"beta_index", i as u64);
let beta_i = transcript.challenge_scalar(b"c");
big_d = beta_i * (all_comms[i] - all_comms[0]);
a_blind = beta_i * (all_blinds[i] - all_blinds[0]);
}
let base_h = hash_to_ristretto(RISTRETTO_BASEPOINT_COMPRESSED.as_bytes());
let mut rng = transcript
.build_rng()
.finalize(&mut mohan_rand());
let r = Scalar::random(&mut rng);
let rh = &r * &base_h;
let c = {
transcript.commit_point(b"A", &big_d.compress());
transcript.commit_point(b"H", &base_h.compress());
transcript.commit_point(b"factor", &rh.compress());
transcript.challenge_scalar(b"c")
};
let z = (&a_blind * &c) + r;
Ok(
BatchEqualityProof {
c: c,
z: z
}
)
}
pub fn verify(&self, cin_asset: &Vec<RistrettoPoint>, cout_asset: &Vec<RistrettoPoint>) -> Result<(), ZeiError> {
let mut transcript = Transcript::new(b"zei_equality_proof");
for i in cin_asset.iter() {
transcript.commit_point(b"commitment", &i.compress())
}
for i in cout_asset.iter() {
transcript.commit_point(b"commitment", &i.compress())
}
let mut all_comms = vec![];
all_comms.extend(cin_asset.iter().cloned());
all_comms.extend(cout_asset.iter().cloned());
let mut big_d = RistrettoPoint::identity();
for i in 0..all_comms.len() {
transcript.append_u64(b"beta_index", i as u64);
let beta_i = transcript.challenge_scalar(b"c");
big_d = beta_i * (all_comms[i] - all_comms[0]);
}
let base_h = hash_to_ristretto(RISTRETTO_BASEPOINT_COMPRESSED.as_bytes());
let factor = (&self.z * base_h) - (self.c * big_d);
let c = {
transcript.commit_point(b"A", &big_d.compress());
transcript.commit_point(b"H", &base_h.compress());
transcript.commit_point(b"factor", &factor.compress());
transcript.challenge_scalar(b"c")
};
if c == self.c {
Ok(())
} else {
Err(ZeiError::VerificationError)
}
}
}
impl ser::Writeable for BatchEqualityProof {
fn write<W: ser::Writer>(&self, writer: &mut W) -> Result<(), ser::Error> {
self.c.write(writer)?;
self.z.write(writer)?;
Ok(())
}
}
impl ser::Readable for BatchEqualityProof {
fn read(reader: &mut dyn ser::Reader) -> Result<BatchEqualityProof, ser::Error> {
Ok(BatchEqualityProof {
c: Scalar::read(reader)?,
z: Scalar::read(reader)?
})
}
}
#[cfg(test)]
mod test {
use super::*;
use crate::PedersenGens;
use crate::core::AssetId;
use mohan::hash::H256;
#[test]
fn test_batch_simple() {
let mut transcript = Transcript::new(b"zei_equality_proof_test");
let mut rng = transcript
.build_rng()
.finalize(&mut mohan_rand());
let value = Scalar::from(11119989891u64);
let mut input_blinds = vec![];
let mut output_blinds = vec![];
for i in 0..10 {
input_blinds.push(Scalar::random(&mut rng));
output_blinds.push(Scalar::random(&mut rng));
}
let pc_gens = PedersenGens::default();
let mut inputs_comm = vec![];
let mut outputs_comm = vec![];
for i in 0..10 {
inputs_comm.push(
pc_gens.commit(&value, &input_blinds[i])
);
outputs_comm.push(
pc_gens.commit(&value, &output_blinds[i])
);
}
let proof = BatchEqualityProof::proove(&inputs_comm, &input_blinds, &outputs_comm, &output_blinds).unwrap();
assert!(proof.verify(&inputs_comm, &outputs_comm).is_ok());
}
#[test]
fn test_batch_simple_asset_mismatch() {
let mut transcript = Transcript::new(b"zei_equality_proof_test");
let mut rng = transcript
.build_rng()
.finalize(&mut mohan_rand());
let flav_id = AssetId::from_inner(H256::from_vec(b"satoshi"));
let flav_id2 = AssetId::from_inner(H256::from_vec(b"turing"));
let mut input_blinds = vec![];
let mut output_blinds = vec![];
for i in 0..10 {
input_blinds.push(Scalar::random(&mut rng));
output_blinds.push(Scalar::random(&mut rng));
}
let pc_gens = PedersenGens::default();
let mut inputs_comm = vec![];
let mut outputs_comm = vec![];
for i in 0..10 {
inputs_comm.push(
pc_gens.commit(&flav_id.into_scalar(), &input_blinds[i])
);
outputs_comm.push(
pc_gens.commit(&flav_id2.into_scalar(), &output_blinds[i])
);
}
let proof = BatchEqualityProof::proove(&inputs_comm, &input_blinds, &outputs_comm, &output_blinds).unwrap();
assert!(proof.verify(&inputs_comm, &outputs_comm).is_err());
}
}