use super::codegen::CodeGen;
use super::sigma::combiners::*;
use super::sigma::types::expr_type_tokens;
use super::syntax::taggedvardict_to_vardict;
use super::transform::{paren_if_needed, prune_statement_tree};
use super::{TaggedIdent, TaggedScalar, TaggedVarDict};
use quote::quote;
use std::collections::{HashSet, VecDeque};
use syn::spanned::Spanned;
use syn::visit::Visit;
use syn::visit_mut::{self, VisitMut};
use syn::{Error, Expr, Ident, Result};
fn priv_scalar_set(e: &Expr, taggedvardict: &TaggedVarDict) -> HashSet<String> {
let mut set: HashSet<String> = HashSet::new();
let vardict = taggedvardict_to_vardict(taggedvardict);
let mut priv_map = PrivScalarMap {
vars: &vardict,
closure: &mut |ident| {
set.insert(ident.to_string());
Ok(())
},
result: Ok(()),
};
priv_map.visit_expr(e);
set
}
fn do_substitution<'a>(expr: &mut Expr, idstr: &'a str, replacement: &'a Expr) {
struct Subs<'a> {
idstr: &'a str,
replacement: &'a Expr,
}
impl<'a> VisitMut for Subs<'a> {
fn visit_expr_mut(&mut self, node: &mut Expr) {
if let Expr::Path(expath) = node {
if let Some(id) = expath.path.get_ident() {
if id.to_string().as_str() == self.idstr {
*node = self.replacement.clone();
return;
}
}
}
visit_mut::visit_expr_mut(self, node);
}
}
let mut subs = Subs { idstr, replacement };
subs.visit_expr_mut(expr);
}
pub fn transform(
codegen: &mut CodeGen,
st: &mut StatementTree,
vars: &mut TaggedVarDict,
) -> Result<()> {
let vardict = taggedvardict_to_vardict(vars);
let mut subs: VecDeque<(Ident, Expr, HashSet<String>)> = VecDeque::new();
let mut subs_vars: HashSet<String> = HashSet::new();
st.for_each_disjunction_branch(&mut |branch, path| {
let in_root_disjunction_branch = path.is_empty();
branch.for_each_disjunction_branch_leaf(&mut |leaf| {
let mut is_subs = None;
if let StatementTree::Leaf(Expr::Assign(syn::ExprAssign { left, .. })) = 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: false, .. })) =
vars.get(&idstr)
{
is_subs = Some(id.clone());
}
}
}
}
if let Some(id) = is_subs {
let old_leaf = std::mem::replace(leaf, StatementTree::leaf_true());
if let StatementTree::Leaf(Expr::Assign(syn::ExprAssign { right, .. })) = old_leaf {
if let Ok((_, right_tokens)) = expr_type_tokens(&vardict, &right) {
let used_priv_scalars = priv_scalar_set(&right, vars);
if !subs_vars.insert(id.to_string()) {
return Err(Error::new(
id.span(),
"variable substituted multiple times",
));
}
if in_root_disjunction_branch {
codegen.prove_append(quote! {
if #id != #right_tokens {
return Err(SigmaError::VerificationFailure);
}
});
}
let right = paren_if_needed(*right);
subs.push_back((id, right, used_priv_scalars));
} else {
return Err(Error::new(
right.span(),
format!(
"Unrecognized arithmetic expression in substitution: {} = {}",
id,
quote! {#right}
),
));
}
}
}
Ok(())
})
})?;
while !subs.is_empty() {
let (id, expr, priv_vars) = subs.pop_front().unwrap();
let idstr = id.to_string();
if priv_vars.contains(&idstr) {
return Err(Error::new(
id.span(),
"variable appears in its own substitution",
));
}
for (_sid, sexpr, spriv_vars) in subs.iter_mut() {
if spriv_vars.contains(&idstr) {
do_substitution(sexpr, &idstr, &expr);
spriv_vars.remove(&idstr);
spriv_vars.extend(priv_vars.clone().into_iter());
}
}
for leafexpr in st.leaves_mut().iter_mut() {
do_substitution(leafexpr, &idstr, &expr);
}
vars.remove(&idstr);
}
prune_statement_tree(st);
Ok(())
}
#[cfg(test)]
mod tests {
use super::super::syntax::taggedvardict_from_strs;
use super::*;
use syn::parse_quote;
fn substitution_tester(
vars: (&[&str], &[&str]),
e: Expr,
subbed_vars: (&[&str], &[&str]),
subbed_e: Expr,
) -> Result<()> {
let mut taggedvardict = taggedvardict_from_strs(vars);
let mut st = StatementTree::parse(&e).unwrap();
let mut codegen = CodeGen::new_empty();
transform(&mut codegen, &mut st, &mut taggedvardict)?;
let subbed_taggedvardict = taggedvardict_from_strs(subbed_vars);
let subbed_st = StatementTree::parse(&subbed_e).unwrap();
assert_eq!(st, subbed_st);
assert_eq!(taggedvardict, subbed_taggedvardict);
Ok(())
}
#[test]
fn apply_substitutions_test() {
let vars_a = (["a", "b", "pub c"].as_slice(), ["A", "B", "C"].as_slice());
substitution_tester(
vars_a,
parse_quote! {
A = b*B + c*C
},
vars_a,
parse_quote! {
A = b*B + c*C
},
)
.unwrap();
substitution_tester(
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
B = a*A + c*C,
)
},
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
B = a*A + c*C,
)
},
)
.unwrap();
substitution_tester(
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
c = a,
)
},
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
c = a,
)
},
)
.unwrap();
substitution_tester(
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
a = c,
a = b,
)
},
vars_a,
parse_quote! { true },
)
.unwrap_err();
substitution_tester(
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
a = c,
a = b,
)
},
vars_a,
parse_quote! { true },
)
.unwrap_err();
substitution_tester(
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
a = 2*a + 1,
)
},
vars_a,
parse_quote! { true },
)
.unwrap_err();
substitution_tester(
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
a = 2*b + 1,
b = a + 4,
)
},
vars_a,
parse_quote! { true },
)
.unwrap_err();
let vars_nob = (["a", "pub c"].as_slice(), ["A", "B", "C"].as_slice());
substitution_tester(
vars_a,
parse_quote! {
AND (
A = b*B + c*C,
b = c,
)
},
vars_nob,
parse_quote! { A = c*B + c*C },
)
.unwrap();
let vars_cd = (
[
"c", "d", "r", "s", "c0", "d0", "r0", "s0", "c1", "d1", "r1", "s1",
]
.as_slice(),
["A", "B", "C", "D"].as_slice(),
);
let vars_cd_noc01 = (
["c", "d", "r", "s", "d0", "r0", "s0", "d1", "r1", "s1"].as_slice(),
["A", "B", "C", "D"].as_slice(),
);
substitution_tester(
vars_cd,
parse_quote! {
AND (
C = c*B + r*A,
D = d*B + s*A,
OR (
AND (
C = c0*B + r0*A,
D = d0*B + s0*A,
c0 = d0,
),
AND (
C = c1*B + r1*A,
D = d1*B + s1*A,
c1 = d1 + 1,
),
)
)
},
vars_cd_noc01,
parse_quote! {
AND (
C = c*B + r*A,
D = d*B + s*A,
OR (
AND (
C = d0*B + r0*A,
D = d0*B + s0*A,
),
AND (
C = (d1+1)*B + r1*A,
D = d1*B + s1*A,
),
)
)
},
)
.unwrap();
}
}