use super::codegen::CodeGen;
use super::pedersen::{
convert_commitment, convert_randomness, recognize_linscalar, recognize_pedersen_assignment,
recognize_pubscalar, unique_random_scalars, 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::{parse_quote, Error, Expr, Ident, Result};
#[derive(Clone, Debug, PartialEq, Eq)]
struct RangeStatement {
upper: Expr,
linscalar: LinScalar,
}
fn subtract_expr(
expr: Option<&Expr>,
exprval: Option<i128>,
lower: &Expr,
lowerval: Option<i128>,
) -> (Expr, Option<i128>) {
if let (Some(ev), Some(lv)) = (exprval, lowerval) {
if let Some(diffv) = ev.checked_sub(lv) {
return (parse_quote! { #diffv }, Some(diffv));
}
}
let paren_lower = paren_if_needed(lower.clone());
(
if let Some(e) = expr {
parse_quote! { #e - #paren_lower }
} else {
parse_quote! { -#paren_lower }
},
None,
)
}
fn parse(vars: &TaggedVarDict, vardict: &VarDict, expr: &Expr) -> Option<RangeStatement> {
if let Expr::MethodCall(syn::ExprMethodCall {
receiver,
method,
turbofish: None,
args,
..
}) = expr
{
if &method.to_string() != "contains" {
return None;
}
let mut range_expr = receiver.as_ref();
if let Expr::Paren(syn::ExprParen {
expr: parened_expr, ..
}) = range_expr
{
range_expr = parened_expr;
}
if let Expr::Range(syn::ExprRange {
start, limits, end, ..
}) = range_expr
{
let lower = start.as_ref()?.as_ref().clone();
let mut upper = end.as_ref()?.as_ref().clone();
let Some((false, lowerval)) = recognize_pubscalar(vars, vardict, &lower) else {
return None;
};
let Some((false, mut upperval)) = recognize_pubscalar(vars, vardict, &upper) else {
return None;
};
let inclusive_upper = matches!(limits, syn::RangeLimits::Closed(_));
if args.len() != 1 {
return None;
}
let priv_expr = args.first().unwrap();
let mut linscalar = recognize_linscalar(vars, vardict, priv_expr)?;
let linscalar_pubscalar_val = if let Some(ref pse) = linscalar.pub_scalar_expr {
let Some((false, pubscalar_val)) = recognize_pubscalar(vars, vardict, pse) else {
return None;
};
pubscalar_val
} else {
Some(0)
};
if inclusive_upper {
let mut added_numerically = false;
if let Some(uv) = upperval {
if let Some(new_uv) = uv.checked_add(1) {
upper = parse_quote! { #new_uv };
upperval = Some(new_uv);
added_numerically = true;
}
}
if !added_numerically {
upper = parse_quote! { #upper + 1 };
upperval = None;
}
}
if lowerval != Some(0) {
(upper, _) = subtract_expr(Some(&upper), upperval, &lower, lowerval);
let pubscalar_expr;
(pubscalar_expr, _) = subtract_expr(
linscalar.pub_scalar_expr.as_ref(),
linscalar_pubscalar_val,
&lower,
lowerval,
);
linscalar.pub_scalar_expr = Some(pubscalar_expr);
}
return Some(RangeStatement { upper, linscalar });
}
}
None
}
#[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 = unique_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 range_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(range_stmt) = parse(vars, &vardict, leafexpr) else {
continue;
};
range_stmt_index += 1;
let mut basic_statements: Vec<Expr> = Vec::new();
let range_id = &range_stmt.linscalar.id;
let ped_assign = if let Some(ped_assign) = pedersens.get(range_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 range statements",
));
}
let cind_A = &cind_points[0];
let cind_B = &cind_points[1];
let commitment_var = codegen.gen_point(
vars,
&format_ident!("range{}_{}_genC", range_stmt_index, range_id),
false, true, );
let rand_var = codegen.gen_scalar(
vars,
&format_ident!("range{}_{}_genr", range_stmt_index, range_id),
true, false, );
vardict = taggedvardict_to_vardict(vars);
randoms.insert(rand_var.to_string());
let ped_assign_expr: Expr = parse_quote! {
#commitment_var = #range_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_ident(&format_ident!("range{}_{}_C", range_stmt_index, range_id));
let rand_var =
codegen.gen_ident(&format_ident!("range{}_{}_r", range_stmt_index, range_id));
vardict = taggedvardict_to_vardict(vars);
randoms.insert(rand_var.to_string());
codegen.verify_append(convert_commitment(
&commitment_var,
&ped_assign,
&range_stmt.linscalar,
&vardict,
)?);
codegen.prove_append(convert_randomness(
&rand_var,
&ped_assign,
&range_stmt.linscalar,
&vardict,
)?);
let upper_var = codegen.gen_ident(&format_ident!(
"range{}_{}_upper",
range_stmt_index,
range_id
));
let upper_code = expr_type_tokens(&vardict, &range_stmt.upper)?.1;
let bitrep_scalars_var = codegen.gen_ident(&format_ident!(
"range{}_{}_bitrep_scalars",
range_stmt_index,
range_id
));
let nbits_var = codegen.gen_ident(&format_ident!(
"range{}_{}_nbits",
range_stmt_index,
range_id
));
codegen.prove_verify_pre_instance_append(quote! {
let #upper_var = #upper_code;
let #bitrep_scalars_var =
sigma_compiler::rangeutils::bitrep_scalars_vartime(#upper_var)?;
if #bitrep_scalars_var.is_empty() {
return Err(SigmaError::VerificationFailure);
}
let #nbits_var = #bitrep_scalars_var.len();
});
let x_var = codegen.gen_ident(&format_ident!("range{}_{}_var", range_stmt_index, range_id));
let bitrep_var = codegen.gen_ident(&format_ident!(
"range{}_{}_bitrep",
range_stmt_index,
range_id
));
let x_code = expr_type_tokens(&vardict, &range_stmt.linscalar.to_expr())?.1;
codegen.prove_append(quote! {
let #x_var = #x_code;
let #bitrep_var =
sigma_compiler::rangeutils::compute_bitrep(#x_var, &#bitrep_scalars_var);
});
let bitcomm_var = codegen.gen_point(
vars,
&format_ident!("range{}_{}_bitC", range_stmt_index, range_id),
true, true, );
let bits_var = codegen.gen_scalar(
vars,
&format_ident!("range{}_{}_bit", range_stmt_index, range_id),
false, true, );
let bitrand_var = codegen.gen_scalar(
vars,
&format_ident!("range{}_{}_bitrand", range_stmt_index, range_id),
false, true, );
let bitrandsq_var = codegen.gen_scalar(
vars,
&format_ident!("range{}_{}_bitrandsq", range_stmt_index, range_id),
false, true, );
let firstbitcomm_var = codegen.gen_point(
vars,
&format_ident!("range{}_{}_firstbitC", range_stmt_index, range_id),
false, false, );
let firstbit_var = codegen.gen_scalar(
vars,
&format_ident!("range{}_{}_firstbit", range_stmt_index, range_id),
false, false, );
let firstbitrand_var = codegen.gen_scalar(
vars,
&format_ident!("range{}_{}_firstbitrand", range_stmt_index, range_id),
false, false, );
let firstbitrandsq_var = codegen.gen_scalar(
vars,
&format_ident!("range{}_{}_firstbitrandsq", range_stmt_index, range_id),
false, false, );
vardict = taggedvardict_to_vardict(vars);
randoms.insert(bitrand_var.to_string());
randoms.insert(firstbitrand_var.to_string());
let commit_generator = &ped_assign.pedersen.var_term.id;
let rand_generator = &ped_assign.pedersen.rand_term.id;
codegen.verify_pre_instance_append(quote! {
let mut #bitcomm_var = Vec::<Point>::new();
#bitcomm_var.resize(#nbits_var - 1, Point::default());
});
codegen.prove_append(quote! {
let #bits_var: Vec<Scalar> =
#bitrep_var
.iter()
.skip(1)
.map(|b| Scalar::conditional_select(
&Scalar::ZERO,
&Scalar::ONE,
*b,
))
.collect();
let #bitrand_var: Vec<Scalar> =
(0..(#nbits_var-1))
.map(|_| Scalar::random(#rng_var))
.collect();
let #bitrandsq_var: Vec<Scalar> =
(0..(#nbits_var-1))
.map(|i| Scalar::conditional_select(
&#bitrand_var[i],
&Scalar::ZERO,
#bitrep_var[i+1],
))
.collect();
let #bitcomm_var: Vec<Point> =
(0..(#nbits_var-1))
.map(|i| #bits_var[i] * #commit_generator +
#bitrand_var[i] * #rand_generator)
.collect();
let #firstbit_var =
Scalar::conditional_select(
&Scalar::ZERO,
&Scalar::ONE,
#bitrep_var[0],
);
let mut #firstbitrand_var = #rand_var;
for i in 0..(#nbits_var-1) {
#firstbitrand_var -=
#bitrand_var[i] * #bitrep_scalars_var[i+1];
}
let #firstbitrandsq_var = Scalar::conditional_select(
&#firstbitrand_var,
&Scalar::ZERO,
#bitrep_var[0],
);
let #firstbitcomm_var =
#firstbit_var * #commit_generator +
#firstbitrand_var * #rand_generator;
});
codegen.verify_append(quote! {
let mut #firstbitcomm_var = #commitment_var;
for i in 0..(#nbits_var-1) {
#firstbitcomm_var -=
#bitcomm_var[i] * #bitrep_scalars_var[i+1];
}
});
basic_statements.push(parse_quote! {
#bitcomm_var = #bits_var * #commit_generator
+ #bitrand_var * #rand_generator
});
basic_statements.push(parse_quote! {
#bitcomm_var = #bits_var * #bitcomm_var
+ #bitrandsq_var * #rand_generator
});
basic_statements.push(parse_quote! {
#firstbitcomm_var = #firstbit_var * #commit_generator
+ #firstbitrand_var * #rand_generator
});
basic_statements.push(parse_quote! {
#firstbitcomm_var = #firstbit_var * #firstbitcomm_var
+ #firstbitrandsq_var * #rand_generator
});
let range_st = StatementTree::And(
basic_statements
.into_iter()
.map(StatementTree::Leaf)
.collect(),
);
**leaf = range_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<RangeStatement>) {
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! {
(0..100).contains(x)
},
Some(RangeStatement {
upper: parse_quote! { 100 },
linscalar: LinScalar {
coeff: 1,
pub_scalar_expr: None,
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(0..=100).contains(x)
},
Some(RangeStatement {
upper: parse_quote! { 101i128 },
linscalar: LinScalar {
coeff: 1,
pub_scalar_expr: None,
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(-12..100).contains(x)
},
Some(RangeStatement {
upper: parse_quote! { 112i128 },
linscalar: LinScalar {
coeff: 1,
pub_scalar_expr: Some(parse_quote! { 12i128 }),
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(-12..(1<<20)).contains(x)
},
Some(RangeStatement {
upper: parse_quote! { 1048588i128 },
linscalar: LinScalar {
coeff: 1,
pub_scalar_expr: Some(parse_quote! { 12i128 }),
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(12..(1<<20)).contains(x+7)
},
Some(RangeStatement {
upper: parse_quote! { 1048564i128 },
linscalar: LinScalar {
coeff: 1,
pub_scalar_expr: Some(parse_quote! { -5i128 }),
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(12..(1<<20)).contains(2*x+7)
},
Some(RangeStatement {
upper: parse_quote! { 1048564i128 },
linscalar: LinScalar {
coeff: 2,
pub_scalar_expr: Some(parse_quote! { -5i128 }),
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(-1..(((1<<126)-1)*2)).contains(x)
},
Some(RangeStatement {
upper: parse_quote! { 170141183460469231731687303715884105727i128 },
linscalar: LinScalar {
coeff: 1,
pub_scalar_expr: Some(parse_quote! { 1i128 }),
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(-2..(((1<<126)-1)*2)).contains(x)
},
Some(RangeStatement {
upper: parse_quote! { (((1<<126)-1)*2)-(-2) },
linscalar: LinScalar {
coeff: 1,
pub_scalar_expr: Some(parse_quote! { 2i128 }),
id: parse_quote! {x},
is_vec: false,
},
}),
);
parse_tester(
vars,
parse_quote! {
(a*b..b+c*c+7).contains(3*x+c*(a+b+2))
},
Some(RangeStatement {
upper: parse_quote! { b+c*c+7-(a*b) },
linscalar: LinScalar {
coeff: 3,
pub_scalar_expr: Some(parse_quote! { c*(a+b+2i128)-(a*b) }),
id: parse_quote! {x},
is_vec: false,
},
}),
);
}
}