use super::codegen::CodeGen;
use super::sigma::combiners::*;
use super::sigma::types::{expr_type_tokens, AExprType};
use super::syntax::{collect_cind_points, taggedvardict_to_vardict};
use super::transform::prune_statement_tree;
use super::{TaggedIdent, TaggedScalar, TaggedVarDict};
use quote::quote;
use syn::{parse_quote, Error, Expr, Result};
#[allow(non_snake_case)] pub fn transform(
codegen: &mut CodeGen,
st: &mut StatementTree,
vars: &mut TaggedVarDict,
) -> Result<()> {
let vardict = taggedvardict_to_vardict(vars);
let cind_points = collect_cind_points(vars);
st.for_each_disjunction_branch(&mut |branch, path| {
let in_root_disjunction_branch = path.is_empty();
branch.for_each_disjunction_branch_leaf(&mut |leaf| {
if let StatementTree::Leaf(Expr::Assign(syn::ExprAssign { left, right, .. })) = leaf {
if let Expr::Path(syn::ExprPath { path, .. }) = left.as_ref() {
if let Some(id) = path.get_ident() {
let idstr = id.to_string();
if let Some(TaggedIdent::Scalar(TaggedScalar {
is_pub: true,
is_vec: l_is_vec,
..
})) = vars.get(&idstr)
{
if let (
AExprType::Scalar {
is_pub: true,
is_vec: r_is_vec,
..
},
right_tokens,
) = expr_type_tokens(&vardict, right)?
{
if *l_is_vec != r_is_vec {
return Err(Error::new(
proc_macro2::Span::call_site(),
"Only one side of the public equality statement is a vector",
));
}
if in_root_disjunction_branch {
codegen.prove_verify_append(quote! {
if #id != #right_tokens {
return Err(SigmaError::VerificationFailure);
}
});
*leaf = StatementTree::leaf_true();
} else {
if cind_points.is_empty() {
return Err(Error::new(
proc_macro2::Span::call_site(),
"At least one cind Point must be declared to support public Scalar equality statements inside disjunctions",
));
}
let cind_A = &cind_points[0];
*leaf = StatementTree::Leaf(parse_quote! {
#id * #cind_A = (#right) * #cind_A
});
}
}
}
}
}
}
Ok(())
})
})?;
prune_statement_tree(st);
Ok(())
}