use ark_ff::{BigInteger, PrimeField};
use serde::{Deserialize, Serialize};
#[cfg(feature = "debug")]
use ark_std::collections::HashMap;
use crate::{
anemoi::{AnemoiJive, N_ANEMOI_ROUNDS},
errors::ZkpError,
plonk::constraint_system::{ConstraintSystem, CsIndex, VarIndex},
utils::serialization::{ark_deserialize, ark_serialize},
};
pub const N_WIRES_PER_GATE: usize = 5;
pub const N_SELECTORS: usize = 8;
pub const N_WIRE_SELECTORS: usize = 3;
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct TurboCS<F: PrimeField> {
#[serde(serialize_with = "ark_serialize", deserialize_with = "ark_deserialize")]
pub selectors: Vec<Vec<F>>,
pub wiring: [Vec<VarIndex>; N_WIRES_PER_GATE],
#[serde(serialize_with = "ark_serialize", deserialize_with = "ark_deserialize")]
pub edwards_a: F,
#[serde(serialize_with = "ark_serialize", deserialize_with = "ark_deserialize")]
pub anemoi_preprocessed_round_keys_x: [[F; 2]; N_ANEMOI_ROUNDS],
#[serde(serialize_with = "ark_serialize", deserialize_with = "ark_deserialize")]
pub anemoi_preprocessed_round_keys_y: [[F; 2]; N_ANEMOI_ROUNDS],
#[serde(serialize_with = "ark_serialize", deserialize_with = "ark_deserialize")]
pub anemoi_generator: F,
#[serde(serialize_with = "ark_serialize", deserialize_with = "ark_deserialize")]
pub anemoi_generator_inv: F,
pub anemoi_constraints_indices: Vec<CsIndex>,
pub num_vars: usize,
pub size: usize,
pub public_vars_constraint_indices: Vec<CsIndex>,
pub public_vars_witness_indices: Vec<VarIndex>,
pub boolean_constraint_indices: Vec<CsIndex>,
pub verifier_only: bool,
#[serde(serialize_with = "ark_serialize", deserialize_with = "ark_deserialize")]
pub witness: Vec<F>,
#[cfg(feature = "debug")]
#[serde(skip)]
pub witness_backtrace: HashMap<VarIndex, std::backtrace::Backtrace>,
}
impl<F: PrimeField> ConstraintSystem<F> for TurboCS<F> {
fn size(&self) -> usize {
self.size
}
fn num_vars(&self) -> usize {
self.num_vars
}
fn wiring(&self) -> &[Vec<usize>] {
&self.wiring[..]
}
fn quot_eval_dom_size(&self) -> usize {
if self.size > 8 {
self.size * 6
} else {
self.size * 16
}
}
fn n_wires_per_gate() -> usize {
N_WIRES_PER_GATE
}
fn num_selectors() -> usize {
N_SELECTORS
}
fn num_wire_selectors() -> usize {
N_WIRE_SELECTORS
}
fn get_edwards_a(&self) -> F {
self.edwards_a
}
fn public_vars_constraint_indices(&self) -> &[CsIndex] {
&self.public_vars_constraint_indices
}
fn public_vars_witness_indices(&self) -> &[VarIndex] {
&self.public_vars_witness_indices
}
fn boolean_constraint_indices(&self) -> &[CsIndex] {
&self.boolean_constraint_indices
}
fn selector(&self, index: usize) -> Result<&[F], ZkpError> {
if index >= self.selectors.len() {
return Err(ZkpError::SelectorIndexOutOfBound);
}
Ok(&self.selectors[index])
}
fn compute_witness_selectors(&self) -> [Vec<F>; N_WIRE_SELECTORS] {
let empty_poly = vec![F::ZERO; self.size];
let polys = [empty_poly.clone(), empty_poly.clone(), empty_poly];
polys
}
fn eval_gate_func(wire_vals: &[&F], sel_vals: &[&F], pub_input: &F) -> Result<F, ZkpError> {
if wire_vals.len() != N_WIRES_PER_GATE || sel_vals.len() != N_SELECTORS {
return Err(ZkpError::SelectorIndexOutOfBound);
}
let add1 = sel_vals[0].mul(wire_vals[0]);
let add2 = sel_vals[1].mul(wire_vals[1]);
let add3 = sel_vals[2].mul(wire_vals[2]);
let add4 = sel_vals[3].mul(wire_vals[3]);
let mul1 = sel_vals[4].mul(wire_vals[0].mul(wire_vals[1]));
let mul2 = sel_vals[5].mul(wire_vals[2].mul(wire_vals[3]));
let constant = sel_vals[6].add(pub_input);
let out = sel_vals[7].mul(wire_vals[4]);
let mut r = add1;
r.add_assign(&add2);
r.add_assign(&add3);
r.add_assign(&add4);
r.add_assign(&mul1);
r.add_assign(&mul2);
r.add_assign(&constant);
r.sub_assign(&out);
Ok(r)
}
fn eval_selector_multipliers(wire_vals: &[&F]) -> Result<Vec<F>, ZkpError> {
if wire_vals.len() < N_WIRES_PER_GATE {
return Err(ZkpError::SelectorIndexOutOfBound);
}
let mut w0w1w2w3w4 = *wire_vals[0];
w0w1w2w3w4.mul_assign(wire_vals[1]);
w0w1w2w3w4.mul_assign(wire_vals[2]);
w0w1w2w3w4.mul_assign(wire_vals[3]);
w0w1w2w3w4.mul_assign(wire_vals[4]);
Ok(vec![
*wire_vals[0],
*wire_vals[1],
*wire_vals[2],
*wire_vals[3],
wire_vals[0].mul(wire_vals[1]),
wire_vals[2].mul(wire_vals[3]),
F::ONE,
wire_vals[4].neg(),
])
}
fn is_verifier_only(&self) -> bool {
self.verifier_only
}
fn shrink_to_verifier_only(&self) -> Self {
Self {
selectors: vec![],
wiring: [vec![], vec![], vec![], vec![], vec![]],
edwards_a: F::ZERO,
anemoi_preprocessed_round_keys_x: [[F::ZERO; 2]; N_ANEMOI_ROUNDS],
anemoi_preprocessed_round_keys_y: [[F::ZERO; 2]; N_ANEMOI_ROUNDS],
anemoi_generator: F::ZERO,
anemoi_generator_inv: F::ZERO,
anemoi_constraints_indices: vec![],
num_vars: self.num_vars,
size: self.size,
public_vars_constraint_indices: vec![],
public_vars_witness_indices: vec![],
boolean_constraint_indices: vec![],
verifier_only: true,
witness: vec![],
#[cfg(feature = "debug")]
witness_backtrace: HashMap::new(),
}
}
fn compute_anemoi_jive_selectors(&self) -> [Vec<F>; 4] {
let empty_poly = vec![F::ZERO; self.size];
let mut polys = [
empty_poly.clone(),
empty_poly.clone(),
empty_poly.clone(),
empty_poly,
];
for i in self.anemoi_constraints_indices.iter() {
for j in 0..N_ANEMOI_ROUNDS {
polys[0][*i + j] = self.anemoi_preprocessed_round_keys_x[j][0];
polys[1][*i + j] = self.anemoi_preprocessed_round_keys_x[j][1];
polys[2][*i + j] = self.anemoi_preprocessed_round_keys_y[j][0];
polys[3][*i + j] = self.anemoi_preprocessed_round_keys_y[j][1];
}
}
polys
}
fn get_anemoi_parameters(&self) -> (F, F) {
(self.anemoi_generator, self.anemoi_generator_inv)
}
fn get_hiding_degree(&self, idx: usize) -> usize {
if idx < 3 {
return 3;
} else {
return 2;
}
}
}
fn compute_binary_le<F: PrimeField>(bytes: &[u8]) -> Vec<F> {
let mut res = vec![];
for byte in bytes.iter() {
let mut tmp = *byte;
for _ in 0..8 {
if (tmp & 1) == 0 {
res.push(F::ZERO);
} else {
res.push(F::ONE);
}
tmp >>= 1;
}
}
res
}
impl<F: PrimeField> Default for TurboCS<F> {
fn default() -> Self {
Self::new()
}
}
impl<F: PrimeField> TurboCS<F> {
pub fn new() -> TurboCS<F> {
let selectors: Vec<Vec<F>> = core::iter::repeat(vec![]).take(N_SELECTORS).collect();
let mut cs = Self {
selectors,
wiring: [vec![], vec![], vec![], vec![], vec![]],
edwards_a: F::ZERO,
anemoi_preprocessed_round_keys_x: [[F::ZERO; 2]; N_ANEMOI_ROUNDS],
anemoi_preprocessed_round_keys_y: [[F::ZERO; 2]; N_ANEMOI_ROUNDS],
anemoi_generator: F::ZERO,
anemoi_generator_inv: F::ZERO,
anemoi_constraints_indices: vec![],
num_vars: 2,
size: 0,
public_vars_constraint_indices: vec![],
public_vars_witness_indices: vec![],
boolean_constraint_indices: vec![],
verifier_only: false,
witness: vec![F::ZERO, F::ONE],
#[cfg(feature = "debug")]
witness_backtrace: HashMap::new(),
};
cs.insert_constant_gate(cs.zero_var(), F::zero());
cs.insert_constant_gate(cs.one_var(), F::one());
cs
}
pub fn zero_var(&self) -> VarIndex {
0
}
pub fn one_var(&self) -> VarIndex {
1
}
pub fn insert_lc_gate(
&mut self,
wires_in: &[VarIndex; 4],
wire_out: VarIndex,
q1: F,
q2: F,
q3: F,
q4: F,
) {
assert!(
wires_in.iter().all(|&x| x < self.num_vars),
"input wire index out of bound"
);
assert!(wire_out < self.num_vars, "wire_out index out of bound");
let zero = F::ZERO;
self.push_add_selectors(q1, q2, q3, q4);
self.push_mul_selectors(zero, zero);
self.push_constant_selector(zero);
self.push_out_selector(F::ONE);
for (i, wire) in wires_in.iter().enumerate() {
self.wiring[i].push(*wire);
}
self.wiring[4].push(wire_out);
self.finish_new_gate();
}
pub fn insert_add_gate(&mut self, left_var: VarIndex, right_var: VarIndex, out_var: VarIndex) {
self.insert_lc_gate(
&[left_var, right_var, 0, 0],
out_var,
F::ONE,
F::ONE,
F::ZERO,
F::ZERO,
);
}
pub fn insert_sub_gate(&mut self, left_var: VarIndex, right_var: VarIndex, out_var: VarIndex) {
self.insert_lc_gate(
&[left_var, right_var, 0, 0],
out_var,
F::ONE,
F::ONE.neg(),
F::ZERO,
F::ZERO,
);
}
pub fn insert_mul_gate(&mut self, left_var: VarIndex, right_var: VarIndex, out_var: VarIndex) {
assert!(left_var < self.num_vars, "left_var index out of bound");
assert!(right_var < self.num_vars, "right_var index out of bound");
assert!(out_var < self.num_vars, "out_var index out of bound");
let zero = F::ZERO;
self.push_add_selectors(zero, zero, zero, zero);
self.push_mul_selectors(F::ONE, zero);
self.push_constant_selector(zero);
self.push_out_selector(F::ONE);
self.wiring[0].push(left_var);
self.wiring[1].push(right_var);
self.wiring[2].push(0);
self.wiring[3].push(0);
self.wiring[4].push(out_var);
self.finish_new_gate();
}
pub fn new_variable(&mut self, value: F) -> VarIndex {
self.num_vars += 1;
self.witness.push(value);
#[cfg(feature = "debug")]
{
self.witness_backtrace
.insert(self.num_vars - 1, std::backtrace::Backtrace::capture());
}
self.num_vars - 1
}
pub fn add_variables(&mut self, values: &[F]) {
self.num_vars += values.len();
for value in values.iter() {
self.witness.push((*value).clone());
}
#[cfg(feature = "debug")]
{
for var in self.num_vars - values.len()..self.num_vars {
self.witness_backtrace
.insert(var, std::backtrace::Backtrace::capture());
}
}
}
#[cfg(feature = "debug")]
pub fn finish_new_gate(&mut self) {
self.size += 1;
let wiring_0_var = self.wiring[0][self.size - 1];
let wiring_1_var = self.wiring[1][self.size - 1];
let wiring_2_var = self.wiring[2][self.size - 1];
let wiring_3_var = self.wiring[3][self.size - 1];
let wiring_4_var = self.wiring[4][self.size - 1];
let wiring_0 = self.witness[wiring_0_var];
let wiring_1 = self.witness[wiring_1_var];
let wiring_2 = self.witness[wiring_2_var];
let wiring_3 = self.witness[wiring_3_var];
let wiring_4 = self.witness[wiring_4_var];
let selector_0 = self.selectors[0][self.size - 1];
let selector_1 = self.selectors[1][self.size - 1];
let selector_2 = self.selectors[2][self.size - 1];
let selector_3 = self.selectors[3][self.size - 1];
let selector_4 = self.selectors[4][self.size - 1];
let selector_5 = self.selectors[5][self.size - 1];
let selector_6 = self.selectors[6][self.size - 1];
let selector_7 = self.selectors[7][self.size - 1];
let selector_8 = self.selectors[8][self.size - 1];
let add1 = selector_0.mul(wiring_0);
let add2 = selector_1.mul(wiring_1);
let add3 = selector_2.mul(wiring_2);
let add4 = selector_3.mul(wiring_3);
let mul1 = selector_4.mul(wiring_0.mul(wiring_1));
let mul2 = selector_5.mul(wiring_2.mul(wiring_3));
let constant = selector_6;
let out = selector_7.mul(wiring_4);
let mut r = add1;
r.add_assign(&add2);
r.add_assign(&add3);
r.add_assign(&add4);
r.add_assign(&mul1);
r.add_assign(&mul2);
r.add_assign(&constant);
r.sub_assign(&out);
if !r.is_zero() {
println!("{}", std::backtrace::Backtrace::capture());
println!("cs constraint not satisfied.");
}
if !(selector_0.is_zero() && selector_4.is_zero()) {
self.witness_backtrace.remove(&wiring_0_var);
}
if !(selector_1.is_zero() && selector_4.is_zero()) {
self.witness_backtrace.remove(&wiring_1_var);
}
if !(selector_2.is_zero() && selector_5.is_zero()) {
self.witness_backtrace.remove(&wiring_2_var);
}
if !(selector_3.is_zero() && selector_5.is_zero()) {
self.witness_backtrace.remove(&wiring_3_var);
}
if !selector_7.is_zero() {
self.witness_backtrace.remove(&wiring_4_var);
}
}
#[cfg(not(feature = "debug"))]
#[inline]
pub fn finish_new_gate(&mut self) {
self.size += 1;
}
pub fn linear_combine(
&mut self,
wires_in: &[VarIndex; 4],
q1: F,
q2: F,
q3: F,
q4: F,
) -> VarIndex {
assert!(
wires_in.iter().all(|&x| x < self.num_vars),
"input wire index out of bound"
);
let w0q1 = self.witness[wires_in[0]].mul(&q1);
let w1q2 = self.witness[wires_in[1]].mul(&q2);
let w2q3 = self.witness[wires_in[2]].mul(&q3);
let w3q4 = self.witness[wires_in[3]].mul(&q4);
let mut lc = w0q1;
lc.add_assign(&w1q2);
lc.add_assign(&w2q3);
lc.add_assign(&w3q4);
let wire_out = self.new_variable(lc);
self.insert_lc_gate(wires_in, wire_out, q1, q2, q3, q4);
wire_out
}
pub fn add(&mut self, left_var: VarIndex, right_var: VarIndex) -> VarIndex {
assert!(left_var < self.num_vars, "left_var index out of bound");
assert!(right_var < self.num_vars, "right_var index out of bound");
let out_var = self.new_variable(self.witness[left_var].add(&self.witness[right_var]));
self.insert_add_gate(left_var, right_var, out_var);
out_var
}
pub fn sub(&mut self, left_var: VarIndex, right_var: VarIndex) -> VarIndex {
assert!(left_var < self.num_vars, "left_var index out of bound");
assert!(right_var < self.num_vars, "right_var index out of bound");
let out_var = self.new_variable(self.witness[left_var].sub(&self.witness[right_var]));
self.insert_sub_gate(left_var, right_var, out_var);
out_var
}
pub fn equal(&mut self, left_var: VarIndex, right_var: VarIndex) {
let zero_var = self.zero_var();
self.insert_sub_gate(left_var, right_var, zero_var);
}
pub fn mul(&mut self, left_var: VarIndex, right_var: VarIndex) -> VarIndex {
assert!(left_var < self.num_vars, "left_var index out of bound");
assert!(right_var < self.num_vars, "right_var index out of bound");
let out_var = self.new_variable(self.witness[left_var].mul(&self.witness[right_var]));
self.insert_mul_gate(left_var, right_var, out_var);
out_var
}
pub fn insert_boolean_gate(&mut self, var: VarIndex) {
self.insert_mul_gate(var, var, var);
}
pub fn range_check(&mut self, var: VarIndex, n_bits: usize) -> Vec<VarIndex> {
assert!(var < self.num_vars, "var index out of bound");
assert!(n_bits >= 2, "the number of bits is less than two");
let witness_bytes = self.witness[var].into_bigint().to_bytes_le();
let mut binary_repr = compute_binary_le::<F>(&witness_bytes);
while binary_repr.len() < n_bits {
binary_repr.push(F::ZERO);
}
let b: Vec<VarIndex> = binary_repr
.into_iter()
.take(n_bits)
.map(|val| self.new_variable(val))
.collect();
let one = F::ONE;
let two = one.add(&one);
let four = two.add(&two);
let eight = four.add(&four);
let bin = vec![one, two, four, eight];
let mut acc = b[n_bits - 1];
self.insert_boolean_gate(b[n_bits - 1]);
let m = (n_bits - 2) / 3;
for i in 0..m {
acc = self.linear_combine(
&[
acc,
b[n_bits - 1 - i * 3 - 1],
b[n_bits - 1 - i * 3 - 2],
b[n_bits - 1 - i * 3 - 3],
],
bin[3],
bin[2],
bin[1],
bin[0],
);
self.attach_boolean_constraint_to_gate();
}
let zero = F::ZERO;
match (n_bits - 1) - 3 * m {
1 => self.insert_lc_gate(&[acc, b[0], 0, 0], var, bin[1], bin[0], zero, zero),
2 => self.insert_lc_gate(&[acc, b[1], b[0], 0], var, bin[2], bin[1], bin[0], zero),
_ => self.insert_lc_gate(
&[acc, b[2], b[1], b[0]],
var,
bin[3],
bin[2],
bin[1],
bin[0],
),
}
self.attach_boolean_constraint_to_gate();
b
}
pub fn select(&mut self, var0: VarIndex, var1: VarIndex, bit: VarIndex) -> VarIndex {
assert!(var0 < self.num_vars, "var0 index out of bound");
assert!(var1 < self.num_vars, "var1 index out of bound");
assert!(bit < self.num_vars, "bit var index out of bound");
let zero = F::ZERO;
let one = F::ONE;
self.push_add_selectors(zero, one, zero, zero);
self.push_mul_selectors(one.neg(), one);
self.push_constant_selector(zero);
self.push_out_selector(one);
let out = if self.witness[bit] == zero {
self.witness[var0].clone()
} else {
self.witness[var1].clone()
};
let out_var = self.new_variable(out);
self.wiring[0].push(bit);
self.wiring[1].push(var0);
self.wiring[2].push(bit);
self.wiring[3].push(var1);
self.wiring[4].push(out_var);
self.finish_new_gate();
out_var
}
pub fn is_equal(&mut self, left_var: VarIndex, right_var: VarIndex) -> VarIndex {
let (is_equal, _) = self.is_equal_or_not_equal(left_var, right_var);
is_equal
}
pub fn is_not_equal(&mut self, left_var: VarIndex, right_var: VarIndex) -> VarIndex {
let (_, is_not_equal) = self.is_equal_or_not_equal(left_var, right_var);
is_not_equal
}
pub fn is_equal_or_not_equal(
&mut self,
left_var: VarIndex,
right_var: VarIndex,
) -> (VarIndex, VarIndex) {
let diff = self.sub(left_var, right_var);
let inv_diff_scalar = self.witness[diff].inverse().unwrap_or(F::ZERO);
let inv_diff = self.new_variable(inv_diff_scalar);
let mul_var = self.mul(diff, inv_diff);
let one_var = self.one_var();
let diff_is_zero = self.sub(one_var, mul_var);
let zero_var = self.zero_var();
self.insert_mul_gate(diff, diff_is_zero, zero_var);
(diff_is_zero, mul_var)
}
pub fn insert_constant_gate(&mut self, var: VarIndex, constant: F) {
assert!(var < self.num_vars, "variable index out of bound");
let zero = F::ZERO;
self.push_add_selectors(zero, zero, zero, zero);
self.push_mul_selectors(zero, zero);
self.push_constant_selector(constant);
self.push_out_selector(F::ONE);
for i in 0..N_WIRES_PER_GATE {
self.wiring[i].push(var);
}
#[cfg(feature = "debug")]
let backtrace = { self.witness_backtrace.remove(&var) };
self.finish_new_gate();
#[cfg(feature = "debug")]
{
match backtrace {
Some(v) => self.witness_backtrace.insert(var, v),
None => None,
};
}
}
pub fn insert_constant_gate_for_input(&mut self, var: VarIndex, constant: F) {
assert!(var < self.num_vars, "variable index out of bound");
let zero = F::ZERO;
self.push_add_selectors(zero, zero, zero, zero);
self.push_mul_selectors(zero, zero);
self.push_constant_selector(constant);
self.push_out_selector(F::ONE);
for i in 0..N_WIRES_PER_GATE {
self.wiring[i].push(var);
}
self.size += 1;
}
pub fn prepare_pi_variable(&mut self, var: VarIndex) {
self.public_vars_witness_indices.push(var);
self.public_vars_constraint_indices.push(self.size);
self.insert_constant_gate_for_input(var, F::ZERO);
}
pub fn attach_boolean_constraint_to_gate(&mut self) {
self.boolean_constraint_indices.push(self.size - 1);
}
pub fn attach_anemoi_jive_constraints_to_gate(&mut self) {
debug_assert!(!self.anemoi_generator.is_zero());
self.anemoi_constraints_indices.push(self.size - 1);
}
pub fn load_anemoi_parameters<H: AnemoiJive<F, 2, N_ANEMOI_ROUNDS>>(&mut self) {
self.anemoi_preprocessed_round_keys_x = H::PREPROCESSED_ROUND_KEYS_X;
self.anemoi_preprocessed_round_keys_y = H::PREPROCESSED_ROUND_KEYS_Y;
self.anemoi_generator = H::GENERATOR;
self.anemoi_generator_inv = H::GENERATOR_INV;
}
pub fn pad(&mut self) {
let n = self.size.next_power_of_two();
let diff = n - self.size();
for selector in self.selectors.iter_mut() {
selector.extend(vec![F::ZERO; diff]);
}
for wire in self.wiring.iter_mut() {
wire.extend(vec![0; diff]);
}
self.size += diff;
#[cfg(feature = "debug")]
{
if !self.witness_backtrace.is_empty() {
let mut animoi_witness_var = Vec::new();
for cs_index in self.anemoi_constraints_indices.iter() {
for r in 0..N_ANEMOI_ROUNDS {
animoi_witness_var.push(self.get_witness_index(0, cs_index + r));
animoi_witness_var.push(self.get_witness_index(1, cs_index + r));
animoi_witness_var.push(self.get_witness_index(2, cs_index + r));
animoi_witness_var.push(self.get_witness_index(3, cs_index + r));
animoi_witness_var.push(self.get_witness_index(4, cs_index + r));
}
}
for (var, backtrace) in &self.witness_backtrace {
if animoi_witness_var.contains(var) {
continue;
}
panic!("dangling witness:\n{}", backtrace);
}
}
}
}
pub fn push_add_selectors(&mut self, q1: F, q2: F, q3: F, q4: F) {
self.selectors[0].push(q1);
self.selectors[1].push(q2);
self.selectors[2].push(q3);
self.selectors[3].push(q4);
}
pub fn push_mul_selectors(&mut self, q_mul12: F, q_mul34: F) {
self.selectors[4].push(q_mul12);
self.selectors[5].push(q_mul34);
}
pub fn push_constant_selector(&mut self, q_c: F) {
self.selectors[6].push(q_c);
}
pub fn push_out_selector(&mut self, q_out: F) {
self.selectors[7].push(q_out);
}
fn get_witness_index(&self, wire_index: usize, cs_index: CsIndex) -> VarIndex {
assert!(wire_index < N_WIRES_PER_GATE, "wire index out of bound");
assert!(cs_index < self.size, "constraint index out of bound");
self.wiring[wire_index][cs_index]
}
pub fn verify_witness(&self, witness: &[F], online_vars: &[F]) -> Result<(), ZkpError> {
if witness.len() != self.num_vars {
return Err(ZkpError::Message(format!(
"witness len = {}, num_vars = {}",
witness.len(),
self.num_vars
)));
}
if online_vars.len() != self.public_vars_witness_indices.len()
|| online_vars.len() != self.public_vars_constraint_indices.len()
{
return Err(ZkpError::Message(
"wrong number of online variables".to_owned(),
));
}
if !self.anemoi_constraints_indices.is_empty() {
assert!(!self.anemoi_generator.is_zero());
}
for cs_index in self.anemoi_constraints_indices.iter() {
for r in 0..N_ANEMOI_ROUNDS {
let a_i = witness[self.get_witness_index(0, cs_index + r)].clone();
let b_i = witness[self.get_witness_index(1, cs_index + r)].clone();
let c_i = witness[self.get_witness_index(2, cs_index + r)].clone();
let d_i = witness[self.get_witness_index(3, cs_index + r)].clone();
let o_i = witness[self.get_witness_index(4, cs_index + r)].clone();
let a_i_next = witness[self.get_witness_index(0, cs_index + 1 + r)].clone();
let b_i_next = witness[self.get_witness_index(1, cs_index + 1 + r)].clone();
let c_i_next = witness[self.get_witness_index(2, cs_index + 1 + r)].clone();
let d_i_next = witness[self.get_witness_index(3, cs_index + 1 + r)].clone();
if o_i != d_i_next {
return Err(ZkpError::Message(format!(
"cs index {} round {}: the output wire {:?} does not equal to the fourth wire {:?} in the next constraint",
cs_index,
r,
o_i,
d_i_next
)));
}
let prk_i_a = self.anemoi_preprocessed_round_keys_x[r][0].clone();
let prk_i_b = self.anemoi_preprocessed_round_keys_x[r][1].clone();
let prk_i_c = self.anemoi_preprocessed_round_keys_y[r][0].clone();
let prk_i_d = self.anemoi_preprocessed_round_keys_y[r][1].clone();
let g = self.anemoi_generator.clone();
let g2 = g.square().add(F::ONE);
let da_i = a_i + d_i;
let cb_i = b_i + c_i;
let d2a_i = da_i + a_i;
let c2b_i = cb_i + b_i;
let left = (da_i + g * cb_i + prk_i_c - &c_i_next).pow(&[5u64])
+ g * (da_i + g * cb_i + prk_i_c).square();
let right = d2a_i + g * c2b_i + prk_i_a;
if left != right {
return Err(ZkpError::Message(format!(
"cs index {} round {}: the first of anemoi equation does not equal: {:?} != {:?}",
cs_index, r, left, right
)));
}
let left = (g * da_i + g2 * cb_i + prk_i_d - &d_i_next).pow(&[5u64])
+ g * (g * da_i + g2 * cb_i + prk_i_d).square();
let right = g * d2a_i + g2 * c2b_i + prk_i_b;
if left != right {
return Err(ZkpError::Message(format!(
"cs index {} round {}: the second equation of anemoi does not equal: {:?} != {:?}",
cs_index, r, left, right
)));
}
let left = (da_i + g * cb_i + prk_i_c - &c_i_next).pow(&[5u64])
+ g * c_i_next.square()
+ &self.anemoi_generator_inv;
let right = a_i_next;
if left != right {
return Err(ZkpError::Message(format!(
"cs index {} round {}: the third equation of anemoi does not equal: {:?} != {:?}",
cs_index, r, left, right
)));
}
let left = (g * da_i + g2 * cb_i + prk_i_d - &d_i_next).pow(&[5u64])
+ g * d_i_next.square()
+ &self.anemoi_generator_inv;
let right = b_i_next;
if left != right {
return Err(ZkpError::Message(format!(
"cs index {} round {}: the fourth equation of anemoi does not equal: {:?} != {:?}",
cs_index, r, left, right
)));
}
}
}
for cs_index in 0..self.size() {
let mut public_online = F::ZERO;
for ((c_i, w_i), online_var) in self
.public_vars_constraint_indices
.iter()
.zip(self.public_vars_witness_indices.iter())
.zip(online_vars.iter())
{
if *c_i == cs_index {
public_online = (*online_var).clone();
if witness[*w_i] != public_online {
return Err(ZkpError::Message(format!(
"cs index {}: online var {:?} does not match witness {:?}",
cs_index,
public_online,
witness[*w_i].clone()
)));
}
}
}
let w1_value = &witness[self.get_witness_index(0, cs_index)];
let w2_value = &witness[self.get_witness_index(1, cs_index)];
let w3_value = &witness[self.get_witness_index(2, cs_index)];
let w4_value = &witness[self.get_witness_index(3, cs_index)];
let w_out_value = &witness[self.get_witness_index(4, cs_index)];
let wire_vals = vec![w1_value, w2_value, w3_value, w4_value, w_out_value];
let sel_vals: Vec<&F> = (0..Self::num_selectors())
.map(|i| &self.selectors[i][cs_index])
.collect();
let eval_gate = Self::eval_gate_func(&wire_vals, &sel_vals, &public_online)?;
if eval_gate != F::ZERO {
return Err(ZkpError::Message(format!(
"cs index {}: wire_vals = ({:?}), sel_vals = ({:?})",
cs_index, wire_vals, sel_vals
)));
}
if self.boolean_constraint_indices.contains(&cs_index) {
if !w2_value.is_zero() && !w2_value.is_one() {
return Err(ZkpError::Message(format!(
"cs index {}: the second wire {:?} is not one or zero",
cs_index, w2_value
)));
}
if !w3_value.is_zero() && !w3_value.is_one() {
return Err(ZkpError::Message(format!(
"cs index {}: the third wire {:?} is not one or zero",
cs_index, w3_value
)));
}
if !w4_value.is_zero() && !w4_value.is_one() {
return Err(ZkpError::Message(format!(
"cs index {}: the fourth wire {:?} is not one or zero",
cs_index, w4_value
)));
}
}
}
Ok(())
}
pub fn get_and_clear_witness(&mut self) -> Vec<F> {
let res = self.witness.clone();
self.witness.clear();
res
}
}