diff --git a/Cargo.lock b/Cargo.lock index 8df1ebc..8a06222 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3774,7 +3774,7 @@ dependencies = [ [[package]] name = "pod2_onchain" version = "0.1.0" -source = "git+https://github.com/0xPARC/pod2-onchain.git?rev=0df0e99ddf37842eb8395ae06ab92a8d79b76cfb#0df0e99ddf37842eb8395ae06ab92a8d79b76cfb" +source = "git+https://github.com/0xPARC/pod2-onchain.git?rev=36c1b426b05e3a5e13f2ba251d1ea3e8eed5bb66#36c1b426b05e3a5e13f2ba251d1ea3e8eed5bb66" dependencies = [ "anyhow", "bindgen", diff --git a/ad-server/src/endpoints.rs b/ad-server/src/endpoints.rs index 63c463d..1828c56 100644 --- a/ad-server/src/endpoints.rs +++ b/ad-server/src/endpoints.rs @@ -268,13 +268,16 @@ mod tests { println!("Prebuilding circuits to calculate vd_set..."); let vd_set = &*DEFAULT_VD_SET; println!("vd_set calculation complete"); - let (state_predicates, rev_predicates) = app::build_predicates(¶ms); + let batches = app::build_predicates(¶ms); + let pred_update = batches[0] + .predicate_ref_by_name("update") + .expect("update defined"); let shrunk_main_pod_build = ShrunkMainPodSetup::new(¶ms).build()?; let pod_config = PodConfig { params, vd_set: vd_set.clone(), - state_predicates, - rev_predicates, + batches, + pred_update, }; let (queue_tx, queue_rx) = mpsc::channel::(8); diff --git a/ad-server/src/main.rs b/ad-server/src/main.rs index a91ea90..6359608 100644 --- a/ad-server/src/main.rs +++ b/ad-server/src/main.rs @@ -3,14 +3,14 @@ use std::{collections::HashMap, str::FromStr, sync::Arc}; use alloy::primitives::Address; use anyhow::{Context as _, Result}; -use app::{Predicates, RevPredicates, build_predicates}; +use app::build_predicates; use common::{ ProofType, shrink::{ShrunkMainPodBuild, ShrunkMainPodSetup}, }; use pod2::{ backends::plonky2::basetypes::DEFAULT_VD_SET, - middleware::{Params, VDSet}, + middleware::{CustomPredicateBatch, CustomPredicateRef, Params, VDSet}, }; use sqlx::{ migrate::MigrateDatabase, @@ -70,8 +70,8 @@ impl Config { pub struct PodConfig { params: Params, vd_set: VDSet, - state_predicates: Predicates, - rev_predicates: RevPredicates, + batches: Vec>, + pred_update: CustomPredicateRef, } pub struct Context { @@ -137,13 +137,16 @@ async fn main() -> Result<()> { info!("Prebuilding circuits to calculate vd_set..."); let vd_set = &*DEFAULT_VD_SET; info!("vd_set calculation complete"); - let (state_predicates, rev_predicates) = build_predicates(¶ms); + let batches = build_predicates(¶ms); + let pred_update = batches[0] + .predicate_ref_by_name("update") + .expect("update defined"); let shrunk_main_pod_build = ShrunkMainPodSetup::new(¶ms).build()?; let pod_config = PodConfig { params, vd_set: vd_set.clone(), - state_predicates, - rev_predicates, + batches, + pred_update, }; if cfg.proof_type == ProofType::Groth16 { diff --git a/ad-server/src/queue.rs b/ad-server/src/queue.rs index b595ba5..5ca2df2 100644 --- a/ad-server/src/queue.rs +++ b/ad-server/src/queue.rs @@ -160,7 +160,7 @@ async fn handle_create(ctx: Arc, req_id: Uuid) -> Result<()> { // send the payload to ethereum let payload_bytes = Payload::Create(PayloadCreate { id: Hash::from(RawValue::from(new_id)), // TODO hash - custom_predicate_ref: ctx.pod_config.state_predicates.update.clone(), + custom_predicate_ref: ctx.pod_config.pred_update.clone(), vds_root: ctx.pod_config.vd_set.root(), }) .to_bytes(); @@ -204,7 +204,7 @@ async fn handle_update(ctx: Arc, req_id: Uuid, id: i64, op: Op) -> Resu let start = std::time::Instant::now(); let mut builder = MainPodBuilder::new(&ctx.pod_config.params, &ctx.pod_config.vd_set); - let mut helper = Helper::new(&mut builder, &ctx.pod_config.state_predicates); + let mut helper = Helper::new(&mut builder, &ctx.pod_config.batches); let op = Dictionary::from(op); let op_raw = RawValue::from(op.commitment()); @@ -319,11 +319,7 @@ async fn handle_update_rev(ctx: Arc, req_id: Uuid, id: i64, num: i64) - Statement::None }; - let mut rev_helper = RevHelper::new( - &mut builder, - &ctx.pod_config.state_predicates, - &ctx.pod_config.rev_predicates, - ); + let mut rev_helper = RevHelper::new(&mut builder, &ctx.pod_config.batches); let (rev_state, rev_st_update) = rev_helper.st_rev_sync(rev_state, op, st_update, old_st_rev_sync); diff --git a/app/src/lib.rs b/app/src/lib.rs index 5fba430..265d65d 100644 --- a/app/src/lib.rs +++ b/app/src/lib.rs @@ -1,19 +1,17 @@ #![allow(clippy::uninlined_format_args)] -use std::{ - collections::{HashMap, HashSet}, - fmt, - str::FromStr, -}; +mod macros; + +use std::{collections::HashSet, fmt, str::FromStr, sync::Arc}; -use anyhow::{Context, Result}; +use anyhow::{Result, bail}; use common::set_from_value; use hex::ToHex; use pod2::{ - frontend::{MainPodBuilder, Operation}, + frontend::MainPodBuilder, lang::parse, middleware::{ - CustomPredicateRef, EMPTY_VALUE, Key, Params, Statement, TypedValue, Value, + CustomPredicateBatch, EMPTY_VALUE, Key, Params, Statement, TypedValue, Value, containers::{Dictionary, Set}, }, }; @@ -21,38 +19,6 @@ use serde::{Deserialize, Serialize}; pub const DEPTH: usize = 32; -#[macro_export] -macro_rules! dict { - ({ $($key:expr => $val:expr),* , }) => ( - $crate::dict!({ $($key => $val),* }).unwrap() - ); - ({ $($key:expr => $val:expr),* }) => ({ - pod2::dict!(DEPTH, { $($key => $val),* }).unwrap() - }); -} - -#[derive(Debug, Clone)] -pub struct Predicates { - pub init: CustomPredicateRef, - pub add: CustomPredicateRef, - pub del: CustomPredicateRef, - pub update: CustomPredicateRef, -} - -#[derive(Debug, Clone)] -pub struct RevPredicates { - pub add_fresh: CustomPredicateRef, - pub add_existing: CustomPredicateRef, - pub add: CustomPredicateRef, - pub del_singleton: CustomPredicateRef, - pub del_else: CustomPredicateRef, - pub del: CustomPredicateRef, - pub sync_init: CustomPredicateRef, - pub sync_add: CustomPredicateRef, - pub sync_del: CustomPredicateRef, - pub sync: CustomPredicateRef, -} - #[derive(PartialEq, Eq, Hash, Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub enum Op { @@ -117,7 +83,7 @@ impl From for TypedValue { /// "green" => Set(...), /// "blue" => Set(...), /// } -pub fn build_predicates(params: &Params) -> (Predicates, RevPredicates) { +pub fn build_predicates(params: &Params) -> Vec> { let empty = format!("Raw({:#})", EMPTY_VALUE); let empty_state = format!( r#"{{"{r}": {empty}, "{g}": {empty}, "{b}": {empty}}}"#, @@ -296,131 +262,94 @@ pub fn build_predicates(params: &Params) -> (Predicates, RevPredicates) { .custom_batch; // State batch predicates - - let state_preds = Predicates { - init: state_batch.predicate_ref_by_name("init").unwrap(), - add: state_batch.predicate_ref_by_name("add").unwrap(), - del: state_batch.predicate_ref_by_name("del").unwrap(), - update: state_batch.predicate_ref_by_name("update").unwrap(), - }; - - // Reverse index state predicates - - let rev_preds = RevPredicates { - add_fresh: rev_state_add_batch - .predicate_ref_by_name("rev_add_fresh") - .unwrap(), - add_existing: rev_state_add_batch - .predicate_ref_by_name("rev_add_existing") - .unwrap(), - add: rev_state_add_batch - .predicate_ref_by_name("rev_add") - .unwrap(), - del_singleton: rev_state_del_batch - .predicate_ref_by_name("rev_del_singleton") - .unwrap(), - del_else: rev_state_del_batch - .predicate_ref_by_name("rev_del_else") - .unwrap(), - del: rev_state_del_batch - .predicate_ref_by_name("rev_del") - .unwrap(), - sync_init: rev_state_batch - .predicate_ref_by_name("rev_sync_init") - .unwrap(), - sync_add: rev_state_batch - .predicate_ref_by_name("rev_sync_add") - .unwrap(), - sync_del: rev_state_batch - .predicate_ref_by_name("rev_sync_del") - .unwrap(), - sync: rev_state_batch.predicate_ref_by_name("rev_sync").unwrap(), - }; - - (state_preds, rev_preds) + vec![ + state_batch.clone(), + rev_state_add_batch.clone(), + rev_state_del_batch.clone(), + rev_state_batch.clone(), + ] } pub struct Helper<'a> { pub builder: &'a mut MainPodBuilder, - pub predicates: &'a Predicates, + pub batches: &'a [Arc], } impl<'a> Helper<'a> { - pub fn new(pod_builder: &'a mut MainPodBuilder, predicates: &'a Predicates) -> Self { + pub fn new( + pod_builder: &'a mut MainPodBuilder, + batches: &'a [Arc], + ) -> Self { Self { builder: pod_builder, - predicates, + batches, } } pub fn st_init(&mut self, old: Dictionary, op: Dictionary) -> Result<(Dictionary, Statement)> { let name = String::try_from(op.get(&Key::from("name")).unwrap().typed()).unwrap(); assert_eq!(name, "init"); - // DictContains(op, "name", "init") - let st0 = self - .builder - .priv_op(Operation::dict_contains(op.clone(), "name", "init")) - .unwrap(); - // Equal(old, EMPTY) - let st1 = self - .builder - .priv_op(Operation::eq(old.clone(), EMPTY_VALUE)) - .context("old state is not empty")?; - - let empty_group = Value::from(Set::new(DEPTH, HashSet::new()).unwrap()); + if Value::from(old.clone()) != Value::from(EMPTY_VALUE) { + bail!("old state is not empty") + } let init_state = dict!({ - "red" => empty_group.clone(), - "green" => empty_group.clone(), - "blue" => empty_group} + "red" => set!(), + "green" => set!(), + "blue" => set!()} ); - // Equal(new, {"red": EMPTY, "green": EMPTY, "blue": EMPTY}) - let st2 = self - .builder - .priv_op(Operation::eq(init_state.clone(), init_state.clone())) - .unwrap(); // init(new, old, op) - let st = self - .builder - .priv_op(Operation::custom( - self.predicates.init.clone(), - [st0, st1, st2], - )) - .unwrap(); + let st = st_custom!( + (self.builder, self.batches), + init( + DictContains(op, "name", "init"), + Equal(old, EMPTY_VALUE), + Equal(init_state, init_state), + ) + ); Ok((init_state, st)) } - pub fn st_add_del( - &mut self, - old: Dictionary, - op: Dictionary, - ) -> Result<(Dictionary, Statement)> { + pub fn st_add(&mut self, old: Dictionary, op: Dictionary) -> Result<(Dictionary, Statement)> { let name = String::try_from(op.get(&Key::from("name")).unwrap().typed()).unwrap(); - assert!(name == "add" || name == "del"); + assert!(name == "add"); - let st0 = if name == "add" { - // DictContains(op, "name", "add") - self.builder - .priv_op(Operation::dict_contains(op.clone(), "name", "add")) - .unwrap() + let group = Key::try_from(op.get(&Key::from("group")).unwrap().typed()).unwrap(); + let old_group = old.get(&group).unwrap(); + + let user = op.get(&Key::from("user")).unwrap(); + let mut new_group = if let TypedValue::Set(set) = old_group.typed() { + set.clone() } else { - // DictContains(op, "name", "del") - self.builder - .priv_op(Operation::dict_contains(op.clone(), "name", "del")) - .unwrap() + panic!("Value not a Set: {:?}", old_group) }; + if new_group.contains(user) { + bail!("old_group already contains user"); + } + new_group.insert(user).unwrap(); + + let mut new = old.clone(); + new.update(&group, &Value::from(new_group.clone())).unwrap(); + + // add(new, old, op, private: old_group, new_group) + let st = st_custom!( + (self.builder, self.batches), + add( + DictContains(op, "name", "add"), + DictContains(old, (&op, "group"), old_group), + SetInsert(new_group, old_group, (&op, "user")), + DictUpdate(new, old, (&op, "group"), new_group), + ) + ); + Ok((new, st)) + } + + pub fn st_del(&mut self, old: Dictionary, op: Dictionary) -> Result<(Dictionary, Statement)> { + let name = String::try_from(op.get(&Key::from("name")).unwrap().typed()).unwrap(); + assert!(name == "del"); let group = Key::try_from(op.get(&Key::from("group")).unwrap().typed()).unwrap(); let old_group = old.get(&group).unwrap(); - // DictContains(old, op.group, old_group) - let st1 = self - .builder - .priv_op(Operation::dict_contains( - old.clone(), - (&op, "group"), - old_group.clone(), - )) - .unwrap(); let user = op.get(&Key::from("user")).unwrap(); let mut new_group = if let TypedValue::Set(set) = old_group.typed() { @@ -428,58 +357,24 @@ impl<'a> Helper<'a> { } else { panic!("Value not a Set: {:?}", old_group) }; - let st2 = if name == "add" { - new_group.insert(user).unwrap(); - // SetInsert(new_group, old_group, op.user) - self.builder - .priv_op(Operation::set_insert( - new_group.clone(), - old_group.clone(), - (&op, "user"), - )) - .context("old_group already contains user")? - } else { - new_group.delete(user).unwrap(); - // SetDelete(new_group, old_group, op.user) - self.builder - .priv_op(Operation::set_delete( - new_group.clone(), - old_group.clone(), - (&op, "user"), - )) - .context("old_group doesn't contain user")? - }; + if !new_group.contains(user) { + bail!("old_group doesn't contain user"); + } + new_group.delete(user).unwrap(); let mut new = old.clone(); new.update(&group, &Value::from(new_group.clone())).unwrap(); - // DictUpdate(new, old, op.group, new_group) - let st3 = self - .builder - .priv_op(Operation::dict_update( - new.clone(), - old.clone(), - (&op, "group"), - new_group, - )) - .unwrap(); - let st = if name == "add" { - // add(new, old, op, private: old_group, new_group) - self.builder - .priv_op(Operation::custom( - self.predicates.add.clone(), - [st0, st1, st2, st3], - )) - .unwrap() - } else { - // del(new, old, op, private: old_group, new_group) - self.builder - .priv_op(Operation::custom( - self.predicates.del.clone(), - [st0, st1, st2, st3], - )) - .unwrap() - }; + // del(new, old, op, private: old_group, new_group) + let st = st_custom!( + (self.builder, self.batches), + del( + DictContains(op, "name", "del"), + DictContains(old, (&op, "group"), old_group), + SetDelete(new_group, old_group, (&op, "user")), + DictUpdate(new, old, (&op, "group"), new_group), + ) + ); Ok((new, st)) } @@ -489,51 +384,42 @@ impl<'a> Helper<'a> { op: Dictionary, ) -> Result<(Dictionary, Statement)> { let name = String::try_from(op.get(&Key::from("name")).unwrap().typed()).unwrap(); - let st_none = Statement::None; - let (new, sts) = match name.as_str() { + let (new, [st0, st1, st2]) = match name.as_str() { "init" => { // init(new, old, op) - let (new, st) = self.st_init(old, op)?; - (new, [st, st_none.clone(), st_none.clone()]) + let (new, st_init) = self.st_init(old, op)?; + (new, [st_init, Statement::None, Statement::None]) } "add" => { // add(new, old, op, private: old_group, new_group) - let (new, st) = self.st_add_del(old, op)?; - (new, [st_none.clone(), st, st_none.clone()]) + let (new, st_add) = self.st_add(old, op)?; + (new, [Statement::None, st_add, Statement::None]) } "del" => { // del(new, old, op, private: old_group, new_group) - let (new, st) = self.st_add_del(old, op)?; - (new, [st_none.clone(), st_none.clone(), st]) + let (new, st_del) = self.st_del(old, op)?; + (new, [Statement::None, Statement::None, st_del]) } _ => panic!("invalid op.name = {}", name), }; - - // update(new, old, op) - let st = self - .builder - .priv_op(Operation::custom(self.predicates.update.clone(), sts)) - .unwrap(); + let st = st_custom!((self.builder, self.batches), update(st0, st1, st2,)); Ok((new, st)) } } pub struct RevHelper<'a> { pub builder: &'a mut MainPodBuilder, - pub predicates: &'a Predicates, - pub rev_predicates: &'a RevPredicates, + pub batches: &'a [Arc], } impl<'a> RevHelper<'a> { pub fn new( pod_builder: &'a mut MainPodBuilder, - predicates: &'a Predicates, - rev_predicates: &'a RevPredicates, + batches: &'a [Arc], ) -> Self { Self { builder: pod_builder, - predicates, - rev_predicates, + batches, } } @@ -542,24 +428,16 @@ impl<'a> RevHelper<'a> { st_update: Statement, op: Dictionary, ) -> (Dictionary, Statement) { - let init_rev_state = Dictionary::new(DEPTH, HashMap::new()).unwrap(); - let st1 = self - .builder - .priv_op(Operation::dict_contains(op.clone(), "name", "init")) - .unwrap(); - let st2 = self - .builder - .priv_op(Operation::eq(init_rev_state.clone(), EMPTY_VALUE)) - .unwrap(); - ( - init_rev_state, - self.builder - .priv_op(Operation::custom( - self.rev_predicates.sync_init.clone(), - [st_update, st1, st2], - )) - .unwrap(), - ) + let init_rev_state = dict!({}); + let st = st_custom!( + (self.builder, self.batches), + rev_sync_init( + st_update, + DictContains(op, "name", "init"), + Equal(init_rev_state, EMPTY_VALUE), + ) + ); + (init_rev_state, st) } pub fn st_rev_add_fresh( @@ -576,32 +454,15 @@ impl<'a> RevHelper<'a> { new_rev .insert(user, &Value::from(user_groups.clone())) .unwrap(); - let st0 = self - .builder - .priv_op(Operation::set_insert( - user_groups.clone(), - empty_set, - (&op, "group"), - )) - .unwrap(); - let st1 = self - .builder - .priv_op(Operation::dict_insert( - new_rev.clone(), - old_rev, - (&op, "user"), - user_groups, - )) - .unwrap(); - ( - new_rev, - self.builder - .priv_op(Operation::custom( - self.rev_predicates.add_fresh.clone(), - [st0, st1], - )) - .unwrap(), - ) + + let st = st_custom!( + (self.builder, self.batches), + rev_add_fresh( + SetInsert(user_groups, empty_set, (&op, "group")), + DictInsert(new_rev, old_rev, (&op, "user"), user_groups), + ) + ); + (new_rev, st) } pub fn st_rev_add_existing( @@ -623,40 +484,15 @@ impl<'a> RevHelper<'a> { .update(user, &Value::from(user_groups.clone())) .unwrap(); - let st0 = self - .builder - .priv_op(Operation::dict_contains( - old_rev.clone(), - (&op, "user"), - old_user_groups.clone(), - )) - .unwrap(); - let st1 = self - .builder - .priv_op(Operation::set_insert( - user_groups.clone(), - old_user_groups.clone(), - (&op, "group"), - )) - .unwrap(); - let st2 = self - .builder - .priv_op(Operation::dict_update( - new_rev.clone(), - old_rev, - (&op, "user"), - user_groups, - )) - .unwrap(); - ( - new_rev, - self.builder - .priv_op(Operation::custom( - self.rev_predicates.add_existing.clone(), - [st0, st1, st2], - )) - .unwrap(), - ) + let st = st_custom!( + (self.builder, self.batches), + rev_add_existing( + DictContains(old_rev, (&op, "user"), old_user_groups), + SetInsert(user_groups, old_user_groups, (&op, "group")), + DictUpdate(new_rev, old_rev, (&op, "user"), user_groups), + ) + ); + (new_rev, st) } pub fn st_rev_add(&mut self, old_rev: Dictionary, op: Dictionary) -> (Dictionary, Statement) { @@ -664,23 +500,19 @@ impl<'a> RevHelper<'a> { Key::from(String::try_from(op.get(&Key::from("user")).unwrap().typed()).unwrap()); let group = Value::from(String::try_from(op.get(&Key::from("group")).unwrap().typed()).unwrap()); - let st_none = Statement::None; - let (new, sts) = match old_rev.get(&user) { + let (new, [st0, st1]) = match old_rev.get(&user) { Err(_) => { - let (new, st) = self.st_rev_add_fresh(old_rev, op, &user, &group); - (new, [st, st_none]) + let (new, st_rev_add_fresh) = self.st_rev_add_fresh(old_rev, op, &user, &group); + (new, [st_rev_add_fresh, Statement::None]) } Ok(_) => { - let (new, st) = self.st_rev_add_existing(old_rev, op, &user, &group); - (new, [st_none, st]) + let (new, st_rev_add_existing) = + self.st_rev_add_existing(old_rev, op, &user, &group); + (new, [Statement::None, st_rev_add_existing]) } }; - ( - new, - self.builder - .priv_op(Operation::custom(self.rev_predicates.add.clone(), sts)) - .unwrap(), - ) + let st = st_custom!((self.builder, self.batches), rev_add(st0, st1,)); + (new, st) } pub fn st_rev_del_singleton( @@ -690,43 +522,18 @@ impl<'a> RevHelper<'a> { user: &Key, ) -> (Dictionary, Statement) { let old_user_groups = old_rev.get(user).unwrap(); - let empty_set = Set::new(DEPTH, HashSet::new()).unwrap(); let mut new_rev = old_rev.clone(); new_rev.delete(user).unwrap(); - let st0 = self - .builder - .priv_op(Operation::dict_contains( - old_rev.clone(), - (&op, "user"), - old_user_groups.clone(), - )) - .unwrap(); - let st1 = self - .builder - .priv_op(Operation::set_delete( - empty_set, - old_user_groups.clone(), - (&op, "group"), - )) - .unwrap(); - let st2 = self - .builder - .priv_op(Operation::dict_delete( - new_rev.clone(), - old_rev, - (&op, "user"), - )) - .unwrap(); - ( - new_rev, - self.builder - .priv_op(Operation::custom( - self.rev_predicates.del_singleton.clone(), - [st0, st1, st2], - )) - .unwrap(), - ) + let st = st_custom!( + (self.builder, self.batches), + rev_del_singleton( + DictContains(old_rev, (&op, "user"), old_user_groups), + SetDelete(set!(), old_user_groups, (&op, "group")), + DictDelete(new_rev, old_rev, (&op, "user")), + ) + ); + (new_rev, st) } pub fn st_rev_del_else( @@ -748,40 +555,15 @@ impl<'a> RevHelper<'a> { .update(user, &Value::from(user_groups.clone())) .unwrap(); - let st0 = self - .builder - .priv_op(Operation::dict_contains( - old_rev.clone(), - (&op, "user"), - old_user_groups.clone(), - )) - .unwrap(); - let st1 = self - .builder - .priv_op(Operation::set_delete( - user_groups.clone(), - old_user_groups.clone(), - (&op, "group"), - )) - .unwrap(); - let st2 = self - .builder - .priv_op(Operation::dict_update( - new_rev.clone(), - old_rev, - (&op, "user"), - user_groups, - )) - .unwrap(); - ( - new_rev, - self.builder - .priv_op(Operation::custom( - self.rev_predicates.del_else.clone(), - [st0, st1, st2], - )) - .unwrap(), - ) + let st = st_custom!( + (self.builder, self.batches), + rev_del_else( + DictContains(old_rev, (&op, "user"), old_user_groups), + SetDelete(user_groups, old_user_groups, (&op, "group")), + DictUpdate(new_rev, old_rev, (&op, "user"), user_groups), + ) + ); + (new_rev, st) } pub fn st_rev_del(&mut self, old_rev: Dictionary, op: Dictionary) -> (Dictionary, Statement) { @@ -789,30 +571,23 @@ impl<'a> RevHelper<'a> { Key::from(String::try_from(op.get(&Key::from("user")).unwrap().typed()).unwrap()); let group = Value::from(String::try_from(op.get(&Key::from("group")).unwrap().typed()).unwrap()); - let st_none = Statement::None; let groups = set_from_value(old_rev.get(&user).unwrap()).unwrap(); - let (new, sts) = match groups.set().len() { + let (new, [st0, st1]) = match groups.set().len() { 1 => { - if groups.contains(&group) { - let (new, st) = self.st_rev_del_singleton(old_rev, op, &user); - (new, [st, st_none]) - } else { + if !groups.contains(&group) { panic!("User is not a member of the specified group.") } + let (new, st_rev_del_singleton) = self.st_rev_del_singleton(old_rev, op, &user); + (new, [st_rev_del_singleton, Statement::None]) } _ => { - let (new, st) = self.st_rev_del_else(old_rev, op, &user, &group); - (new, [st_none, st]) + let (new, st_rev_del_else) = self.st_rev_del_else(old_rev, op, &user, &group); + (new, [Statement::None, st_rev_del_else]) } }; - - ( - new, - self.builder - .priv_op(Operation::custom(self.rev_predicates.del.clone(), sts)) - .unwrap(), - ) + let st = st_custom!((self.builder, self.batches), rev_del(st0, st1,)); + (new, st) } pub fn st_rev_sync_add( @@ -822,20 +597,17 @@ impl<'a> RevHelper<'a> { old_st_rev_sync: Statement, op: Dictionary, ) -> (Dictionary, Statement) { - let st2 = self - .builder - .priv_op(Operation::dict_contains(op.clone(), "name", "add")) - .unwrap(); - let (new, st3) = self.st_rev_add(old_rev, op); - ( - new, - self.builder - .priv_op(Operation::custom( - self.rev_predicates.sync_add.clone(), - [old_st_rev_sync, st_update, st2, st3], - )) - .unwrap(), - ) + let (new, st_rev_add) = self.st_rev_add(old_rev, op.clone()); + let st = st_custom!( + (self.builder, self.batches), + rev_sync_add( + old_st_rev_sync, + st_update, + DictContains(op, "name", "add"), + st_rev_add, + ) + ); + (new, st) } pub fn st_rev_sync_del( @@ -845,20 +617,17 @@ impl<'a> RevHelper<'a> { old_st_rev_sync: Statement, op: Dictionary, ) -> (Dictionary, Statement) { - let st2 = self - .builder - .priv_op(Operation::dict_contains(op.clone(), "name", "del")) - .unwrap(); - let (new, st3) = self.st_rev_del(old_rev, op); - ( - new, - self.builder - .priv_op(Operation::custom( - self.rev_predicates.sync_del.clone(), - [old_st_rev_sync, st_update, st2, st3], - )) - .unwrap(), - ) + let (new, st_rev_del) = self.st_rev_del(old_rev, op.clone()); + let st = st_custom!( + (self.builder, self.batches), + rev_sync_del( + old_st_rev_sync, + st_update, + DictContains(op.clone(), "name", "del"), + st_rev_del, + ) + ); + (new, st) } pub fn st_rev_sync( @@ -869,33 +638,28 @@ impl<'a> RevHelper<'a> { old_st_rev_sync: Statement, ) -> (Dictionary, Statement) { let name = String::try_from(op.get(&Key::from("name")).unwrap().typed()).unwrap(); - let st_none = Statement::None; - let (new, sts) = match name.as_str() { + let (new, [st0, st1, st2]) = match name.as_str() { "init" => { // rev_sync_init(rev_state, state) - let (new, st) = self.st_rev_sync_init(st_update, op); - (new, [st, st_none.clone(), st_none.clone()]) + let (new, st_rev_sync_init) = self.st_rev_sync_init(st_update, op); + (new, [st_rev_sync_init, Statement::None, Statement::None]) } "add" => { // rev_sync_add(rev_state, state) - let (new, st) = self.st_rev_sync_add(old_rev, st_update, old_st_rev_sync, op); - (new, [st_none.clone(), st, st_none.clone()]) + let (new, st_rev_sync_add) = + self.st_rev_sync_add(old_rev, st_update, old_st_rev_sync, op); + (new, [Statement::None, st_rev_sync_add, Statement::None]) } "del" => { // rev_sync_del(rev_state, state) - let (new, st) = self.st_rev_sync_del(old_rev, st_update, old_st_rev_sync, op); - (new, [st_none.clone(), st_none.clone(), st]) + let (new, st_rev_sync_del) = + self.st_rev_sync_del(old_rev, st_update, old_st_rev_sync, op); + (new, [Statement::None, Statement::None, st_rev_sync_del]) } _ => panic!("invalid op.name = {}", name), }; - - ( - new, - // rev_sync(rev_state, state) - self.builder - .priv_op(Operation::custom(self.rev_predicates.sync.clone(), sts)) - .unwrap(), - ) + let st = st_custom!((self.builder, self.batches), rev_sync(st0, st1, st2,)); + (new, st) } } @@ -915,15 +679,14 @@ mod tests { params: &Params, vd_set: &VDSet, prover: &dyn MainPodProver, - predicates: &Predicates, - rev_predicates: &RevPredicates, + batches: &[Arc], state: Dictionary, rev_state: Dictionary, op: Op, old_rev_state_pod: Option, ) -> (Dictionary, Dictionary, Option) { let mut builder = MainPodBuilder::new(params, vd_set); - let mut helper = Helper::new(&mut builder, predicates); + let mut helper = Helper::new(&mut builder, batches); // State Pod let (state, st_update) = helper @@ -948,7 +711,7 @@ mod tests { } else { Statement::None }; - let mut rev_helper = RevHelper::new(&mut builder, predicates, rev_predicates); + let mut rev_helper = RevHelper::new(&mut builder, batches); let (rev_state, rev_st_update) = rev_helper.st_rev_sync(rev_state, Dictionary::from(op), st_update, old_st_rev_sync); builder.reveal(&rev_st_update); @@ -967,11 +730,14 @@ mod tests { #[test] fn test_app() { env_logger::init(); - // let (vd_set, prover) = (&VDSet::new(8, &[]).unwrap(), &MockProver {}); + // let (vd_set, prover) = ( + // &VDSet::new(8, &[]).unwrap(), + // &pod2::backends::plonky2::mock::mainpod::MockProver {}, + // ); let (vd_set, prover) = (&*DEFAULT_VD_SET, &Prover {}); let params = Params::default(); - let (state_predicates, rev_predicates) = build_predicates(¶ms); + let batches = build_predicates(¶ms); // Initial state let mut state = dict!({}); @@ -1008,8 +774,7 @@ mod tests { ¶ms, vd_set, prover, - &state_predicates, - &rev_predicates, + &batches, state, rev_state, op, diff --git a/app/src/macros.rs b/app/src/macros.rs new file mode 100644 index 0000000..08ce1ef --- /dev/null +++ b/app/src/macros.rs @@ -0,0 +1,115 @@ +use std::sync::Arc; + +use pod2::middleware::{CustomPredicateBatch, CustomPredicateRef}; + +#[macro_export] +macro_rules! set { + () => ({ + pod2::middleware::containers::Set::new(DEPTH, std::collections::HashSet::new()).unwrap() + }); + ($($val:expr),* ,) => ( + $crate::set!($($val),*).unwrap() + ); + ($($val:expr),*) => ({ + let mut set = std::collections::HashSet::new(); + $( set.insert($crate::middleware::Value::from($val)); )* + pod2::middleware::containers::Set::new(DEPTH, set).unwrap() + }); +} + +#[macro_export] +macro_rules! dict { + ({ }) => ( + pod2::middleware::containers::Dictionary::new(DEPTH, std::collections::HashMap::new()).unwrap() + ); + ({ $($key:expr => $val:expr),* , }) => ( + $crate::dict!({ $($key => $val),* }).unwrap() + ); + ({ $($key:expr => $val:expr),* }) => ({ + let mut map = std::collections::HashMap::new(); + $( map.insert(pod2::middleware::Key::from($key.clone()), pod2::middleware::Value::from($val.clone())); )* + pod2::middleware::containers::Dictionary::new(DEPTH, map).unwrap() + }); +} + +#[macro_export] +macro_rules! op { + (Equal($a:expr, $b:expr)) => { + pod2::frontend::Operation::eq($a.clone(), $b.clone()) + }; + (DictContains($dict:expr, $key:expr, $value:expr)) => { + pod2::frontend::Operation::dict_contains($dict.clone(), $key.clone(), $value.clone()) + }; + (DictUpdate($dict:expr, $old_dict:expr, $key:expr, $value:expr)) => { + pod2::frontend::Operation::dict_update( + $dict.clone(), + $old_dict.clone(), + $key.clone(), + $value.clone(), + ) + }; + (DictInsert($dict:expr, $old_dict:expr, $key:expr, $value:expr)) => { + pod2::frontend::Operation::dict_insert( + $dict.clone(), + $old_dict.clone(), + $key.clone(), + $value.clone(), + ) + }; + (DictDelete($dict:expr, $old_dict:expr, $key:expr)) => { + pod2::frontend::Operation::dict_delete($dict.clone(), $old_dict.clone(), $key.clone()) + }; + (SetInsert($set:expr, $old_set:expr, $value:expr)) => { + pod2::frontend::Operation::set_insert($set.clone(), $old_set.clone(), $value.clone()) + }; + (SetDelete($set:expr, $old_set:expr, $value:expr)) => { + pod2::frontend::Operation::set_delete($set.clone(), $old_set.clone(), $value.clone()) + }; +} + +pub fn find_custom_pred_by_name( + batches: &[Arc], + name: &str, +) -> Option { + for batch in batches { + for (index, predicate) in batch.predicates().iter().enumerate() { + if predicate.name == name { + return Some(CustomPredicateRef { + batch: batch.clone(), + index, + }); + } + } + } + None +} + +#[macro_export] +macro_rules! _st_custom_args { + ($builder:expr, $input_sts:expr,) => {{ + }}; + ($builder:expr, $input_sts:expr, $pred:ident($($args:expr),+), $($tail:tt)*) => {{ + $input_sts.push($builder.priv_op(op!($pred($($args),+))).unwrap()); + _st_custom_args!($builder, $input_sts, $($tail)*) + }}; + ($builder:expr, $input_sts:expr, $st:expr, $($tail:tt)*) => {{ + $input_sts.push($st); + _st_custom_args!($builder, $input_sts, $($tail)*) + }}; +} + +/// Argument types: +/// $builder: &mut MainPodBuilder +/// $batches: &[Arc {{ + let custom_pred = $crate::macros::find_custom_pred_by_name($batches, stringify!($pred)).unwrap(); + let mut input_sts = Vec::new(); + _st_custom_args!($builder, &mut input_sts, $($args)*); + $builder + .priv_op(pod2::frontend::Operation::custom(custom_pred, input_sts)) + .unwrap() + }}; +} diff --git a/common/src/groth.rs b/common/src/groth.rs index 63d1ef9..174b6b3 100644 --- a/common/src/groth.rs +++ b/common/src/groth.rs @@ -74,10 +74,10 @@ mod tests { fn compute_pod_proof() -> Result { let params = Params::default(); let vd_set = &*DEFAULT_VD_SET; - let (state_predicates, _) = app::build_predicates(¶ms); + let batches = app::build_predicates(¶ms); let mut builder = MainPodBuilder::new(¶ms, vd_set); - let mut helper = app::Helper::new(&mut builder, &state_predicates); + let mut helper = app::Helper::new(&mut builder, &batches); let initial_state = Dictionary::new( params.max_depth_mt_containers, diff --git a/common/src/payload.rs b/common/src/payload.rs index 2db7e7c..f1e48e0 100644 --- a/common/src/payload.rs +++ b/common/src/payload.rs @@ -256,14 +256,11 @@ mod tests { println!("ShrunkMainPod setup"); let shrunk_main_pod_build = ShrunkMainPodSetup::new(¶ms).build().unwrap(); let common_data = &shrunk_main_pod_build.circuit_data.common; - let (state_predicates, _rev_predicates) = app::build_predicates(¶ms); + let batches = app::build_predicates(¶ms); let id = Hash([F(1), F(2), F(3), F(4)]); let custom_predicate_ref = CustomPredicateRef { - batch: CustomPredicateBatch::new_opaque( - "unknown".to_string(), - state_predicates.update.batch.id(), - ), - index: state_predicates.update.index, + batch: CustomPredicateBatch::new_opaque("unknown".to_string(), batches[0].id()), + index: batches[0].predicate_ref_by_name("update").unwrap().index, }; let vd_set = &*DEFAULT_VD_SET; let vds_root = vd_set.root(); @@ -280,8 +277,8 @@ mod tests { assert_eq!(payload_create, payload_create_decoded); let mut builder = MainPodBuilder::new(¶ms, vd_set); - let (state_predicates, _rev_predicates) = app::build_predicates(¶ms); - let mut helper = app::Helper::new(&mut builder, &state_predicates); + let batches = app::build_predicates(¶ms); + let mut helper = app::Helper::new(&mut builder, &batches); let state = containers::Dictionary::new(params.max_depth_mt_containers, HashMap::new()).unwrap();