use super::Strategy;
use crate::mds_matrix::MDS_MATRIX;
use crate::WIDTH;
use dusk_bls12_381::BlsScalar;
use dusk_plonk::prelude::*;
pub struct GadgetStrategy<'a> {
cs: &'a mut Composer,
count: usize,
}
impl<'a> GadgetStrategy<'a> {
pub fn new(cs: &'a mut Composer) -> Self {
GadgetStrategy { cs, count: 0 }
}
pub fn gadget(composer: &'a mut Composer, x: &mut [Witness]) {
let mut strategy = GadgetStrategy::new(composer);
strategy.perm(x);
}
}
impl AsMut<Composer> for GadgetStrategy<'_> {
fn as_mut(&mut self) -> &mut Composer {
self.cs
}
}
impl<'a> Strategy<Witness> for GadgetStrategy<'a> {
fn add_round_key<'b, I>(&mut self, constants: &mut I, words: &mut [Witness])
where
I: Iterator<Item = &'b BlsScalar>,
{
if self.count == 0 {
words.iter_mut().for_each(|w| {
let constant = Self::next_c(constants);
let constraint = Constraint::new().left(1).a(*w).constant(constant);
*w = self.cs.gate_add(constraint);
});
}
}
fn quintic_s_box(&mut self, value: &mut Witness) {
let constraint = Constraint::new().mult(1).a(*value).b(*value);
let v2 = self.cs.gate_mul(constraint);
let constraint = Constraint::new().mult(1).a(v2).b(v2);
let v4 = self.cs.gate_mul(constraint);
let constraint = Constraint::new().mult(1).a(v4).b(*value);
*value = self.cs.gate_mul(constraint);
}
fn mul_matrix<'b, I>(&mut self, constants: &mut I, values: &mut [Witness])
where
I: Iterator<Item = &'b BlsScalar>,
{
let mut result = [Composer::ZERO; WIDTH];
self.count += 1;
for j in 0..WIDTH {
let c = if self.count < Self::rounds() {
Self::next_c(constants)
} else {
BlsScalar::zero()
};
let constraint = Constraint::new()
.left(MDS_MATRIX[j][0])
.a(values[0])
.right(MDS_MATRIX[j][1])
.b(values[1])
.fourth(MDS_MATRIX[j][2])
.d(values[2]);
result[j] = self.cs.gate_add(constraint);
let constraint = Constraint::new()
.left(MDS_MATRIX[j][3])
.a(values[3])
.right(MDS_MATRIX[j][4])
.b(values[4])
.fourth(1)
.d(result[j])
.constant(c);
result[j] = self.cs.gate_add(constraint);
}
values.copy_from_slice(&result);
}
}
#[cfg(test)]
mod tests {
use crate::{GadgetStrategy, ScalarStrategy, Strategy, WIDTH};
use core::result::Result;
use dusk_plonk::prelude::*;
use ff::Field;
use rand::rngs::StdRng;
use rand::SeedableRng;
#[derive(Default)]
struct TestCircuit {
i: [BlsScalar; WIDTH],
o: [BlsScalar; WIDTH],
}
impl Circuit for TestCircuit {
fn circuit(&self, composer: &mut Composer) -> Result<(), Error> {
let zero = Composer::ZERO;
let mut perm: [Witness; WIDTH] = [zero; WIDTH];
let mut i_var: [Witness; WIDTH] = [zero; WIDTH];
self.i.iter().zip(i_var.iter_mut()).for_each(|(i, v)| {
*v = composer.append_witness(*i);
});
let mut o_var: [Witness; WIDTH] = [zero; WIDTH];
self.o.iter().zip(o_var.iter_mut()).for_each(|(o, v)| {
*v = composer.append_witness(*o);
});
GadgetStrategy::gadget(composer, &mut i_var);
perm.copy_from_slice(&i_var);
i_var.iter().zip(o_var.iter()).for_each(|(p, o)| {
composer.assert_equal(*p, *o);
});
Ok(())
}
}
fn hades() -> ([BlsScalar; WIDTH], [BlsScalar; WIDTH]) {
let mut input = [BlsScalar::zero(); WIDTH];
input
.iter_mut()
.for_each(|s| *s = BlsScalar::random(&mut rand::thread_rng()));
let mut output = [BlsScalar::zero(); WIDTH];
output.copy_from_slice(&input);
ScalarStrategy::new().perm(&mut output);
(input, output)
}
fn setup() -> Result<(Prover, Verifier), Error> {
const CAPACITY: usize = 1 << 10;
let pp = PublicParameters::setup(CAPACITY, &mut rand::thread_rng())?;
let label = b"hades_gadget_tester";
Compiler::compile::<TestCircuit>(&pp, label)
}
#[test]
fn preimage() -> Result<(), Error> {
let (prover, verifier) = setup()?;
let (i, o) = hades();
let circuit = TestCircuit { i, o };
let mut rng = StdRng::seed_from_u64(0xbeef);
let (proof, public_inputs) = prover.prove(&mut rng, &circuit)?;
verifier.verify(&proof, &public_inputs)?;
Ok(())
}
#[test]
fn preimage_constant() -> Result<(), Error> {
let (prover, verifier) = setup()?;
let i = [BlsScalar::from(5000u64); WIDTH];
let mut o = [BlsScalar::from(5000u64); WIDTH];
ScalarStrategy::new().perm(&mut o);
let circuit = TestCircuit { i, o };
let mut rng = StdRng::seed_from_u64(0xbeef);
let (proof, public_inputs) = prover.prove(&mut rng, &circuit)?;
verifier.verify(&proof, &public_inputs)?;
Ok(())
}
#[test]
fn preimage_fails() -> Result<(), Error> {
let (prover, _) = setup()?;
let x_scalar = BlsScalar::from(31u64);
let mut i = [BlsScalar::zero(); WIDTH];
i[1] = x_scalar;
let mut o = [BlsScalar::from(31u64); WIDTH];
ScalarStrategy::new().perm(&mut o);
let circuit = TestCircuit { i, o };
let mut rng = StdRng::seed_from_u64(0xbeef);
assert!(
prover.prove(&mut rng, &circuit).is_err(),
"proving should fail since the circuit is invalid"
);
Ok(())
}
}