Просмотр исходного кода

zk/vm: Proper separation for prover and verifier.

The witnesses the prover passes in should be correct, and the witnesses
that the verifier passes in should have the same types, and be None.

We'll see later how this will affect more complex circuits, but we
hope it won't bite us in the ass.
parazyd 4 лет назад
Родитель
Сommit
86d9aa81b2
6 измененных файлов с 192 добавлено и 100 удалено
  1. 143 53
      example/vm.rs
  2. 3 0
      src/zk/mod.rs
  3. 28 28
      src/zk/vm.rs
  4. 0 5
      src/zk/vm/mod.rs
  5. 17 13
      src/zk/vm_stack.rs
  6. 1 1
      src/zkas/decoder.rs

+ 143 - 53
example/vm.rs

@@ -1,3 +1,6 @@
+use std::time::Instant;
+
+#[allow(unused_imports)]
 use halo2::{
     arithmetic::{CurveAffine, Field},
     dev::MockProver,
@@ -17,7 +20,9 @@ use darkfi::{
         keypair::{PublicKey, SecretKey},
         merkle_node::MerkleNode,
         mint_proof::MintRevealedValues,
+        proof::{ProvingKey, VerifyingKey},
         spend_proof::SpendRevealedValues,
+        Proof,
     },
     zk::vm::{Witness, ZkCircuit},
     zkas::decoder::ZkBinary,
@@ -28,6 +33,9 @@ fn mint_proof() -> Result<()> {
     let bincode = include_bytes!("../proof/mint.zk.bin");
     let zkbin = ZkBinary::decode(bincode)?;
 
+    // ======
+    // Prover
+    // ======
     let value = 42;
     let token_id = pallas::Base::from(22);
     let value_blind = pallas::Scalar::random(&mut OsRng);
@@ -35,8 +43,20 @@ fn mint_proof() -> Result<()> {
     let serial = pallas::Base::random(&mut OsRng);
     let coin_blind = pallas::Base::random(&mut OsRng);
     let public_key = PublicKey::random(&mut OsRng);
+    let pk_coords = public_key.0.to_affine().coordinates().unwrap();
 
-    let revealed = MintRevealedValues::compute(
+    let witnesses_prover = vec![
+        Witness::Base(Some(*pk_coords.x())),
+        Witness::Base(Some(*pk_coords.y())),
+        Witness::Base(Some(pallas::Base::from(value))),
+        Witness::Base(Some(token_id)),
+        Witness::Base(Some(serial)),
+        Witness::Base(Some(coin_blind)),
+        Witness::Scalar(Some(value_blind)),
+        Witness::Scalar(Some(token_blind)),
+    ];
+
+    let public_inputs = MintRevealedValues::compute(
         value,
         token_id,
         value_blind,
@@ -44,31 +64,77 @@ fn mint_proof() -> Result<()> {
         serial,
         coin_blind,
         public_key,
-    );
-
-    let pk_coords = public_key.0.to_affine().coordinates().unwrap();
-    let witnesses = vec![
-        Witness::Base(*pk_coords.x()),
-        Witness::Base(*pk_coords.y()),
-        Witness::Base(pallas::Base::from(value)),
-        Witness::Base(token_id),
-        Witness::Base(serial),
-        Witness::Base(coin_blind),
-        Witness::Scalar(value_blind),
-        Witness::Scalar(token_blind),
+    )
+    .make_outputs()
+    .to_vec();
+
+    let circuit = ZkCircuit::new(witnesses_prover, zkbin.clone());
+
+    // let prover = MockProver::run(11, &circuit, vec![public_inputs.clone()]).unwrap();
+    // assert_eq!(prover.verify(), Ok(()));
+
+    let start = Instant::now();
+    let proving_key = ProvingKey::build(11, circuit.clone());
+    info!("Prover setup: [{:?}]", Instant::now() - start);
+
+    let start = Instant::now();
+    let proof = Proof::create(&proving_key, &[circuit], &public_inputs.clone())?;
+    info!("Prover prove: [{:?}]", Instant::now() - start);
+
+    // =======
+    // Verifier
+    // =======
+
+    let witnesses_verifier = vec![
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Scalar(None),
+        Witness::Scalar(None),
     ];
 
-    let circuit = ZkCircuit::new(witnesses, zkbin);
-    let prover = MockProver::run(11, &circuit, vec![revealed.make_outputs().to_vec()]).unwrap();
-    assert_eq!(prover.verify(), Ok(()));
+    let start = Instant::now();
+    let circuit = ZkCircuit::new(witnesses_verifier, zkbin);
+    let verifying_key = VerifyingKey::build(11, circuit);
+    info!("Verifier setup: [{:?}]", Instant::now() - start);
+
+    let start = Instant::now();
+    proof.verify(&verifying_key, &public_inputs)?;
+    info!("Verifier verify: [{:?}]", Instant::now() - start);
 
     Ok(())
 }
 
+fn fill_tree(coin2: pallas::Base) -> BridgeTree<MerkleNode, 32> {
+    let mut tree = BridgeTree::<MerkleNode, 32>::new(100);
+    let coin0 = pallas::Base::random(&mut OsRng);
+    let coin1 = pallas::Base::random(&mut OsRng);
+    let coin3 = pallas::Base::random(&mut OsRng);
+
+    tree.append(&MerkleNode(coin0));
+    tree.witness();
+
+    tree.append(&MerkleNode(coin1));
+
+    tree.append(&MerkleNode(coin2));
+    tree.witness();
+
+    tree.append(&MerkleNode(coin3));
+    tree.witness();
+
+    tree
+}
+
 fn burn_proof() -> Result<()> {
     let bincode = include_bytes!("../proof/burn.zk.bin");
     let zkbin = ZkBinary::decode(bincode)?;
 
+    // ======
+    // Prover
+    // ======
     let value = 42;
     let token_id = pallas::Base::from(22);
     let value_blind = pallas::Scalar::random(&mut OsRng);
@@ -78,14 +144,6 @@ fn burn_proof() -> Result<()> {
     let secret = SecretKey::random(&mut OsRng);
     let sig_secret = SecretKey::random(&mut OsRng);
 
-    let mut tree = BridgeTree::<MerkleNode, 32>::new(100);
-
-    let random_coin_1 = pallas::Base::random(&mut OsRng);
-    tree.append(&MerkleNode(random_coin_1));
-    tree.witness();
-    let random_coin_2 = pallas::Base::random(&mut OsRng);
-    tree.append(&MerkleNode(random_coin_2));
-
     let coin = {
         let coords = PublicKey::from_secret(secret).0.to_affine().coordinates().unwrap();
         let messages =
@@ -94,16 +152,27 @@ fn burn_proof() -> Result<()> {
         poseidon::Hash::init(P128Pow5T3, ConstantLength::<6>).hash(messages)
     };
 
-    tree.append(&MerkleNode(coin));
-    tree.witness();
+    let tree = fill_tree(coin);
+    let (leaf_position, merkle_path) = tree.authentication_path(&MerkleNode(coin)).unwrap();
 
-    let random_coin_3 = pallas::Base::random(&mut OsRng);
-    tree.append(&MerkleNode(random_coin_3));
-    tree.witness();
+    // Why are these types not matched in halo2 gadgets?
+    let leaf_pos: u64 = leaf_position.into();
+    let leaf_pos = leaf_pos as u32;
 
-    let (leaf_position, merkle_path) = tree.authentication_path(&MerkleNode(coin)).unwrap();
+    let witnesses_prover = vec![
+        Witness::Base(Some(secret.0)),
+        Witness::Base(Some(serial)),
+        Witness::Base(Some(pallas::Base::from(value))),
+        Witness::Base(Some(token_id)),
+        Witness::Base(Some(coin_blind)),
+        Witness::Scalar(Some(value_blind)),
+        Witness::Scalar(Some(token_blind)),
+        Witness::Uint32(Some(leaf_pos)),
+        Witness::MerklePath(Some(merkle_path.clone())),
+        Witness::Base(Some(sig_secret.0)),
+    ];
 
-    let revealed = SpendRevealedValues::compute(
+    let public_inputs = SpendRevealedValues::compute(
         value,
         token_id,
         value_blind,
@@ -112,37 +181,58 @@ fn burn_proof() -> Result<()> {
         coin_blind,
         secret,
         leaf_position,
-        merkle_path.clone(),
+        merkle_path,
         sig_secret,
-    );
-
-    // Why are these types not matched in halo2 gadgets?
-    let leaf_pos: u64 = leaf_position.into();
-    let leaf_pos = leaf_pos as u32;
-
-    let witnesses = vec![
-        Witness::Base(secret.0),
-        Witness::Base(serial),
-        Witness::Base(pallas::Base::from(value)),
-        Witness::Base(token_id),
-        Witness::Base(coin_blind),
-        Witness::Scalar(value_blind),
-        Witness::Scalar(token_blind),
-        Witness::Uint32(leaf_pos),
-        Witness::MerklePath(merkle_path),
-        Witness::Base(sig_secret.0),
+    )
+    .make_outputs()
+    .to_vec();
+
+    let circuit = ZkCircuit::new(witnesses_prover, zkbin.clone());
+
+    // let prover = MockProver::run(11, &circuit, vec![public_inputs.clone()])?;
+    // assert_eq!(prover.verify(), Ok(()));
+
+    let start = Instant::now();
+    let proving_key = ProvingKey::build(11, circuit.clone());
+    info!("Prover setup: [{:?}]", Instant::now() - start);
+
+    let start = Instant::now();
+    let proof = Proof::create(&proving_key, &[circuit], &public_inputs)?;
+    info!("Prover prove: [{:?}]", Instant::now() - start);
+
+    // ========
+    // Verifier
+    // ========
+
+    let witnesses_verifier = vec![
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Base(None),
+        Witness::Scalar(None),
+        Witness::Scalar(None),
+        Witness::Uint32(None),
+        Witness::MerklePath(None),
+        Witness::Base(None),
     ];
 
-    let circuit = ZkCircuit::new(witnesses, zkbin);
-    let prover = MockProver::run(11, &circuit, vec![revealed.make_outputs().to_vec()])?;
-    assert_eq!(prover.verify(), Ok(()));
+    let start = Instant::now();
+    let circuit = ZkCircuit::new(witnesses_verifier, zkbin);
+    let verifying_key = VerifyingKey::build(11, circuit);
+    info!("Verifier setup: [{:?}]", Instant::now() - start);
+
+    let start = Instant::now();
+    proof.verify(&verifying_key, &public_inputs)?;
+    info!("Verifier verify: [{:?}]", Instant::now() - start);
 
     Ok(())
 }
 
 fn main() -> Result<()> {
     TermLogger::init(
-        LevelFilter::Debug,
+        //LevelFilter::Debug,
+        LevelFilter::Info,
         simplelog::Config::default(),
         TerminalMode::Mixed,
         ColorChoice::Auto,

+ 3 - 0
src/zk/mod.rs

@@ -2,3 +2,6 @@ pub mod circuit;
 
 #[cfg(feature = "zkvm")]
 pub mod vm;
+
+#[cfg(feature = "zkvm")]
+mod vm_stack;

+ 28 - 28
src/zk/vm/vm.rs → src/zk/vm.rs

@@ -24,7 +24,7 @@ use halo2_gadgets::{
 use log::debug;
 use pasta_curves::pallas;
 
-use super::vm_stack::{Stack, Witness};
+pub use super::vm_stack::{Stack, Witness};
 use crate::{
     crypto::constants::{
         sinsemilla::{OrchardCommitDomains, OrchardHashDomains},
@@ -81,7 +81,7 @@ impl VmConfig {
     }
 }
 
-#[derive(Default)]
+#[derive(Clone, Default)]
 pub struct ZkCircuit {
     constants: Vec<String>,
     witnesses: Vec<Witness>,
@@ -257,17 +257,17 @@ impl Circuit<pallas::Base> for ZkCircuit {
                 "VALUE_COMMIT_VALUE" => {
                     let vcv = OrchardFixedBases::ValueCommitV;
                     let vcv = FixedPoint::from_inner(ecc_chip.clone(), vcv);
-                    stack.push(Stack::Var(Witness::EcFixedPoint(vcv)));
+                    stack.push(Stack::Var(Witness::EcFixedPoint(Some(vcv))));
                 }
                 "VALUE_COMMIT_RANDOM" => {
                     let vcr = OrchardFixedBases::ValueCommitR;
                     let vcr = FixedPoint::from_inner(ecc_chip.clone(), vcr);
-                    stack.push(Stack::Var(Witness::EcFixedPoint(vcr)));
+                    stack.push(Stack::Var(Witness::EcFixedPoint(Some(vcr))));
                 }
                 "NULLIFIER_K" => {
                     let nfk = OrchardFixedBases::NullifierK;
                     let nfk = FixedPoint::from_inner(ecc_chip.clone(), nfk);
-                    stack.push(Stack::Var(Witness::EcFixedPoint(nfk)));
+                    stack.push(Stack::Var(Witness::EcFixedPoint(Some(nfk))));
                 }
                 _ => unimplemented!(),
             }
@@ -278,46 +278,46 @@ impl Circuit<pallas::Base> for ZkCircuit {
         // table cell.
         for witness in &self.witnesses {
             match witness {
-                Witness::EcPoint(_) => {
+                Witness::EcPoint(w) => {
                     debug!("Pushing EcPoint to stack index {}", stack.len());
-                    stack.push(Stack::Var(witness.clone()));
+                    stack.push(Stack::Var(Witness::EcPoint(w.clone())));
                 }
 
-                Witness::EcFixedPoint(_) => {
+                Witness::EcFixedPoint(w) => {
                     debug!("Pushing EcFixedPoint to stack index {}", stack.len());
-                    stack.push(Stack::Var(witness.clone()));
+                    stack.push(Stack::Var(Witness::EcFixedPoint(w.clone())));
                 }
 
-                Witness::Base(v) => {
+                Witness::Base(w) => {
                     debug!("Loading Base element into cell");
-                    let w = self.load_private(
+                    let c = self.load_private(
                         layouter.namespace(|| "Load witness into cell"),
                         config.advices[0],
-                        Some(*v),
+                        *w,
                     )?;
 
-                    debug!("Pushing Base to stack index {}", stack.len());
-                    stack.push(Stack::Cell(w));
+                    debug!("Pushing Base/Cell to stack index {}", stack.len());
+                    stack.push(Stack::Cell(c));
                 }
 
-                Witness::Scalar(_) => {
+                Witness::Scalar(w) => {
                     debug!("Pushing Scalar to stack index {}", stack.len());
-                    stack.push(Stack::Var(witness.clone()));
+                    stack.push(Stack::Var(Witness::Scalar(*w)));
                 }
 
-                Witness::MerklePath(_) => {
+                Witness::MerklePath(w) => {
                     debug!("Pushing MerklePath to stack index {}", stack.len());
-                    stack.push(Stack::Var(witness.clone()));
+                    stack.push(Stack::Var(Witness::MerklePath(w.clone())));
                 }
 
-                Witness::Uint32(_) => {
+                Witness::Uint32(w) => {
                     debug!("Pushing Uint32 to stack index {}", stack.len());
-                    stack.push(Stack::Var(witness.clone()));
+                    stack.push(Stack::Var(Witness::Uint32(*w)));
                 }
 
-                Witness::Uint64(_) => {
+                Witness::Uint64(w) => {
                     debug!("Pushing Uint64 to stack index {}", stack.len());
-                    stack.push(Stack::Var(witness.clone()));
+                    stack.push(Stack::Var(Witness::Uint64(*w)));
                 }
             }
         }
@@ -338,7 +338,7 @@ impl Circuit<pallas::Base> for ZkCircuit {
                     let ret = lhs.add(layouter.namespace(|| "EcAdd()"), &rhs)?;
 
                     debug!("Pushing result to stack index {}", stack.len());
-                    stack.push(Stack::Var(Witness::EcPoint(ret)));
+                    stack.push(Stack::Var(Witness::EcPoint(Some(ret))));
                 }
 
                 Opcode::EcMul => {
@@ -348,12 +348,12 @@ impl Circuit<pallas::Base> for ZkCircuit {
                     let lhs: FixedPoint<pallas::Affine, EccChip<OrchardFixedBases>> =
                         stack[args[1]].clone().into();
 
-                    let rhs: pallas::Scalar = stack[args[0]].clone().into();
+                    let rhs: Option<pallas::Scalar> = stack[args[0]].clone().into();
 
-                    let (ret, _) = lhs.mul(layouter.namespace(|| "EcMul()"), Some(rhs))?;
+                    let (ret, _) = lhs.mul(layouter.namespace(|| "EcMul()"), rhs)?;
 
                     debug!("Pushing result to stack index {}", stack.len());
-                    stack.push(Stack::Var(Witness::EcPoint(ret)));
+                    stack.push(Stack::Var(Witness::EcPoint(Some(ret))));
                 }
 
                 Opcode::EcMulBase => {
@@ -368,7 +368,7 @@ impl Circuit<pallas::Base> for ZkCircuit {
                     let ret = lhs.mul_base_field(layouter.namespace(|| "EcMulBase()"), rhs)?;
 
                     debug!("Pushing result to stack index {}", stack.len());
-                    stack.push(Stack::Var(Witness::EcPoint(ret)));
+                    stack.push(Stack::Var(Witness::EcPoint(Some(ret))));
                 }
 
                 Opcode::EcMulShort => {
@@ -384,7 +384,7 @@ impl Circuit<pallas::Base> for ZkCircuit {
                         lhs.mul_short(layouter.namespace(|| "EcMulShort()"), (rhs, one))?;
 
                     debug!("Pushing result to stack index {}", stack.len());
-                    stack.push(Stack::Var(Witness::EcPoint(ret)));
+                    stack.push(Stack::Var(Witness::EcPoint(Some(ret))));
                 }
 
                 Opcode::EcGetX => {

+ 0 - 5
src/zk/vm/mod.rs

@@ -1,5 +0,0 @@
-pub mod vm;
-mod vm_stack;
-
-pub use vm::*;
-pub use vm_stack::Witness;

+ 17 - 13
src/zk/vm/vm_stack.rs → src/zk/vm_stack.rs

@@ -16,19 +16,19 @@ pub enum Stack {
 
 #[derive(Clone)]
 pub enum Witness {
-    EcPoint(Point<pallas::Affine, EccChip<OrchardFixedBases>>),
-    EcFixedPoint(FixedPoint<pallas::Affine, EccChip<OrchardFixedBases>>),
-    Base(pallas::Base),
-    Scalar(pallas::Scalar),
-    MerklePath(Vec<MerkleNode>),
-    Uint32(u32),
-    Uint64(u64),
+    EcPoint(Option<Point<pallas::Affine, EccChip<OrchardFixedBases>>>),
+    EcFixedPoint(Option<FixedPoint<pallas::Affine, EccChip<OrchardFixedBases>>>),
+    Base(Option<pallas::Base>),
+    Scalar(Option<pallas::Scalar>),
+    MerklePath(Option<Vec<MerkleNode>>),
+    Uint32(Option<u32>),
+    Uint64(Option<u64>),
 }
 
 impl From<Stack> for Point<pallas::Affine, EccChip<OrchardFixedBases>> {
     fn from(value: Stack) -> Self {
         match value {
-            Stack::Var(Witness::EcPoint(v)) => v,
+            Stack::Var(Witness::EcPoint(v)) => v.unwrap(),
             _ => unimplemented!(),
         }
     }
@@ -37,13 +37,13 @@ impl From<Stack> for Point<pallas::Affine, EccChip<OrchardFixedBases>> {
 impl From<Stack> for FixedPoint<pallas::Affine, EccChip<OrchardFixedBases>> {
     fn from(value: Stack) -> Self {
         match value {
-            Stack::Var(Witness::EcFixedPoint(v)) => v,
+            Stack::Var(Witness::EcFixedPoint(v)) => v.unwrap(),
             _ => unimplemented!(),
         }
     }
 }
 
-impl From<Stack> for pallas::Scalar {
+impl From<Stack> for std::option::Option<pallas::Scalar> {
     fn from(value: Stack) -> Self {
         match value {
             Stack::Var(Witness::Scalar(v)) => v,
@@ -64,7 +64,7 @@ impl From<Stack> for CellValue<pallas::Base> {
 impl From<Stack> for std::option::Option<u32> {
     fn from(value: Stack) -> Self {
         match value {
-            Stack::Var(Witness::Uint32(v)) => Some(v),
+            Stack::Var(Witness::Uint32(v)) => v,
             _ => unimplemented!(),
         }
     }
@@ -74,8 +74,12 @@ impl From<Stack> for std::option::Option<[pallas::Base; 32]> {
     fn from(value: Stack) -> Self {
         match value {
             Stack::Var(Witness::MerklePath(v)) => {
-                let ret: Vec<pallas::Base> = v.iter().map(|x| x.0).collect();
-                Some(ret.try_into().unwrap())
+                if let Some(path) = v {
+                    let ret: Vec<pallas::Base> = path.iter().map(|x| x.0).collect();
+                    return Some(ret.try_into().unwrap())
+                }
+
+                None
             }
             _ => unimplemented!(),
         }

+ 1 - 1
src/zkas/decoder.rs

@@ -5,7 +5,7 @@ use crate::{
     Result,
 };
 
-#[derive(Debug)]
+#[derive(Clone, Debug)]
 pub struct ZkBinary {
     pub constants: Vec<(Type, String)>,
     pub witnesses: Vec<Type>,