use super::codegen::CodeGen;
use super::pedersen::{
convert_commitment, convert_randomness, random_scalars, recognize_linscalar,
recognize_pedersen_assignment, recognize_pubscalar, LinScalar, PedersenAssignment,
};
use super::sigma::combiners::*;
use super::sigma::types::{expr_type_tokens, VarDict};
use super::syntax::{collect_cind_points, taggedvardict_to_vardict};
use super::transform::paren_if_needed;
use super::TaggedVarDict;
use quote::{format_ident, quote};
use std::collections::HashMap;
use syn::spanned::Spanned;
use syn::{parse_quote, Error, Expr, Ident, Result};
fn subtract_expr(linscalar: LinScalar, subexpr: &Expr, subval: Option<i128>) -> LinScalar {
if subval != Some(0) {
let paren_sub = paren_if_needed(subexpr.clone());
if let Some(expr) = linscalar.pub_scalar_expr {
return LinScalar {
pub_scalar_expr: Some(parse_quote! {
#expr - #paren_sub
}),
..linscalar
};
} else {
return LinScalar {
pub_scalar_expr: Some(parse_quote! {
-#paren_sub
}),
..linscalar
};
}
}
linscalar
}
fn parse(vars: &TaggedVarDict, vardict: &VarDict, expr: &Expr) -> Option<LinScalar> {
let Expr::Binary(syn::ExprBinary {
left,
op: syn::BinOp::Ne(_),
right,
..
}) = expr
else {
return None;
};
let linscalar = recognize_linscalar(vars, vardict, left)?;
let (subexpr_is_vec, subval) = recognize_pubscalar(vars, vardict, right)?;
if linscalar.is_vec || subexpr_is_vec {
return None;
}
Some(subtract_expr(linscalar, right, subval))
}
#[allow(non_snake_case)] pub fn transform(
codegen: &mut CodeGen,
st: &mut StatementTree,
vars: &mut TaggedVarDict,
) -> Result<()> {
let mut vardict = taggedvardict_to_vardict(vars);
let mut randoms = random_scalars(vars, st);
let mut leaves = st.leaves_st_mut();
let cind_points = collect_cind_points(vars);
let pedersens: HashMap<Ident, PedersenAssignment> = leaves
.iter()
.filter_map(|leaf| {
if let StatementTree::Leaf(leafexpr) = leaf {
recognize_pedersen_assignment(vars, &randoms, &vardict, leafexpr)
.map(|ped_assign| (ped_assign.var(), ped_assign))
} else {
None
}
})
.collect();
let mut neq_stmt_index = 0usize;
let rng_var = codegen.gen_ident(&format_ident!("rng"));
for leaf in leaves.iter_mut() {
let StatementTree::Leaf(leafexpr) = leaf else {
continue;
};
let Some(neq_linscalar) = parse(vars, &vardict, leafexpr) else {
continue;
};
neq_stmt_index += 1;
if let Some(super::TaggedIdent::Scalar(super::TaggedScalar {
is_pub: false,
is_rand: true,
..
})) = vars.get(&neq_linscalar.id.to_string())
{
return Err(Error::new(
leafexpr.span(),
"target of not-equals expression cannot be rand",
));
}
let mut basic_statements: Vec<Expr> = Vec::new();
let neq_id = &neq_linscalar.id;
let ped_assign = if let Some(ped_assign) = pedersens.get(neq_id) {
ped_assign.clone()
} else {
if cind_points.len() < 2 {
return Err(Error::new(
proc_macro2::Span::call_site(),
"At least two cind Points must be declared to support not-equals statements",
));
}
let cind_A = &cind_points[0];
let cind_B = &cind_points[1];
let commitment_var = codegen.gen_point(
vars,
&format_ident!("neq{}_{}_genC", neq_stmt_index, neq_id),
false, true, );
let rand_var = codegen.gen_scalar(
vars,
&format_ident!("neq{}_{}_genr", neq_stmt_index, neq_id),
true, false, );
vardict = taggedvardict_to_vardict(vars);
randoms.insert(rand_var.to_string());
let ped_assign_expr: Expr = parse_quote! {
#commitment_var = #neq_id * #cind_A + #rand_var * #cind_B
};
let ped_assign =
recognize_pedersen_assignment(vars, &randoms, &vardict, &ped_assign_expr).unwrap();
codegen.prove_append(quote! {
let #rand_var = Scalar::random(#rng_var);
let #ped_assign_expr;
});
basic_statements.push(ped_assign_expr);
ped_assign
};
let commitment_var = codegen.gen_point(
vars,
&format_ident!("neq{}_{}_C", neq_stmt_index, neq_id),
false, false, );
let rand_var = codegen.gen_ident(&format_ident!("neq{}_{}_r", neq_stmt_index, neq_id));
vardict = taggedvardict_to_vardict(vars);
randoms.insert(rand_var.to_string());
codegen.prove_verify_append(convert_commitment(
&commitment_var,
&ped_assign,
&neq_linscalar,
&vardict,
)?);
codegen.prove_append(convert_randomness(
&rand_var,
&ped_assign,
&neq_linscalar,
&vardict,
)?);
let Lx_var = codegen.gen_ident(&format_ident!("neq{}_{}_var", neq_stmt_index, neq_id));
let Lx_code = expr_type_tokens(&vardict, &neq_linscalar.to_expr())?.1;
let j_var = codegen.gen_scalar(
vars,
&format_ident!("neq{}_{}_j", neq_stmt_index, neq_id),
false, false, );
let s_var = codegen.gen_scalar(
vars,
&format_ident!("neq{}_{}_s", neq_stmt_index, neq_id),
false, false, );
vardict = taggedvardict_to_vardict(vars);
let commit_generator = &ped_assign.pedersen.var_term.id;
let rand_generator = &ped_assign.pedersen.rand_term.id;
codegen.prove_append(quote! {
let #Lx_var = #Lx_code;
let #j_var = <Scalar as Field>::invert(&#Lx_var)
.into_option()
.ok_or(SigmaError::VerificationFailure)?;
let #s_var = -#rand_var * #j_var;
});
basic_statements.push(parse_quote! {
#commit_generator = #j_var * #commitment_var
+ #s_var * #rand_generator
});
let neq_st = StatementTree::And(
basic_statements
.into_iter()
.map(StatementTree::Leaf)
.collect(),
);
**leaf = neq_st;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::super::syntax::taggedvardict_from_strs;
use super::*;
fn parse_tester(vars: (&[&str], &[&str]), expr: Expr, expect: Option<LinScalar>) {
let taggedvardict = taggedvardict_from_strs(vars);
let vardict = taggedvardict_to_vardict(&taggedvardict);
let output = parse(&taggedvardict, &vardict, &expr);
assert_eq!(output, expect);
}
#[test]
fn parse_test() {
let vars = (
[
"x", "y", "z", "pub a", "pub b", "pub c", "rand r", "rand s", "rand t",
]
.as_slice(),
["C", "cind A", "cind B"].as_slice(),
);
parse_tester(
vars,
parse_quote! {
x != 0
},
Some(LinScalar {
coeff: 1,
pub_scalar_expr: None,
id: parse_quote! {x},
is_vec: false,
}),
);
parse_tester(
vars,
parse_quote! {
x != 5
},
Some(LinScalar {
coeff: 1,
pub_scalar_expr: Some(parse_quote! {-5}),
id: parse_quote! {x},
is_vec: false,
}),
);
parse_tester(
vars,
parse_quote! {
2*x != 5
},
Some(LinScalar {
coeff: 2,
pub_scalar_expr: Some(parse_quote! {-5}),
id: parse_quote! {x},
is_vec: false,
}),
);
parse_tester(
vars,
parse_quote! {
2*x + 12 != 5
},
Some(LinScalar {
coeff: 2,
pub_scalar_expr: Some(parse_quote! {12i128-5}),
id: parse_quote! {x},
is_vec: false,
}),
);
parse_tester(
vars,
parse_quote! {
2*x + a*a != 0
},
Some(LinScalar {
coeff: 2,
pub_scalar_expr: Some(parse_quote! {a*a}),
id: parse_quote! {x},
is_vec: false,
}),
);
parse_tester(
vars,
parse_quote! {
2*x + a*a != b*c + c
},
Some(LinScalar {
coeff: 2,
pub_scalar_expr: Some(parse_quote! {a*a-(b*c+c)}),
id: parse_quote! {x},
is_vec: false,
}),
);
}
}