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

zkvm: Implement cond_select opcode.

parazyd 3 лет назад
Родитель
Сommit
f4932072c6
7 измененных файлов с 53 добавлено и 3 удалено
  1. 1 0
      contrib/zk.lang
  2. 1 1
      contrib/zk.lua
  3. 2 2
      contrib/zk.vim
  4. 5 0
      proof/opcodes.zk
  5. 33 0
      src/zk/vm.rs
  6. 9 0
      src/zkas/opcode.rs
  7. 2 0
      tests/zkvm_opcodes.rs

+ 1 - 0
contrib/zk.lang

@@ -68,6 +68,7 @@
       <keyword>less_than_strict</keyword>
       <keyword>less_than_loose</keyword>
       <keyword>bool_check</keyword>
+      <keyword>cond_select</keyword>
       <keyword>witness_base</keyword>
       <keyword>constrain_equal_base</keyword>
       <keyword>constrain_equal_point</keyword>

+ 1 - 1
contrib/zk.lua

@@ -43,7 +43,7 @@ local instruction = token('instruction', word_match{
   'base_add', 'base_mul', 'base_sub',
   'poseidon_hash', 'merkle_root',
   'range_check', 'less_than_strict', 'less_than_loose', 'bool_check',
-  'witness_base',
+  'cond_select', 'witness_base',
   'constrain_equal_base', 'constrain_equal_point',
   'constrain_instance', 'debug',
 })

+ 2 - 2
contrib/zk.vim

@@ -22,8 +22,8 @@ syn keyword zkasInstruction
     \ ec_get_x ec_get_y
     \ base_add base_mul base_sub
     \ poseidon_hash merkle_root
-    \ range_check less_than_strict less_than_loose  bool_check
-    \ witness_base
+    \ range_check less_than_strict less_than_loose bool_check
+    \ cond_select witness_base
     \ constrain_equal_base constrain_equal_point
     \ constrain_instance debug
 

+ 5 - 0
proof/opcodes.zk

@@ -20,6 +20,8 @@ witness "Opcodes" {
 
 	Uint32 leaf_pos,
 	MerklePath path,
+
+	Base cond,
 }
 
 circuit "Opcodes" {
@@ -64,4 +66,7 @@ circuit "Opcodes" {
 	ephem_public = ec_mul_var_base(ephem_secret, pubkey);
 	constrain_instance(ec_get_x(ephem_public));
 	constrain_instance(ec_get_y(ephem_public));
+
+	out = cond_select(cond, a, b);
+	constrain_instance(out);
 }

+ 33 - 0
src/zk/vm.rs

@@ -53,6 +53,7 @@ use super::{
     assign_free_advice,
     gadget::{
         arithmetic::{ArithChip, ArithConfig, ArithInstruction},
+        cond_select::{ConditionalSelectChip, ConditionalSelectConfig},
         less_than::{LessThanChip, LessThanConfig},
         native_range_check::{NativeRangeCheckChip, NativeRangeCheckConfig},
         small_range_check::{SmallRangeCheckChip, SmallRangeCheckConfig},
@@ -79,6 +80,7 @@ pub struct VmConfig {
     native_253_range_check_config: NativeRangeCheckConfig<3, 253, 85>,
     lessthan_config: LessThanConfig<3, 253, 85>,
     boolcheck_config: SmallRangeCheckConfig,
+    condselect_config: ConditionalSelectConfig<pallas::Base>,
 }
 
 impl VmConfig {
@@ -105,6 +107,10 @@ impl VmConfig {
     fn arithmetic_chip(&self) -> ArithChip<pallas::Base> {
         ArithChip::construct(self.arith_config.clone())
     }
+
+    fn condselect_chip(&self) -> ConditionalSelectChip<pallas::Base> {
+        ConditionalSelectChip::construct(self.condselect_config.clone(), ())
+    }
 }
 
 pub struct ZkCircuit {
@@ -263,6 +269,10 @@ impl Circuit<pallas::Base> for ZkCircuit {
         // chip with a range of 2, which enforces one bit, i.e. 0 or 1.
         let boolcheck_config = SmallRangeCheckChip::configure(meta, advices[9], 2);
 
+        // Cnfiguration for the conditional selection chip
+        let condselect_config =
+            ConditionalSelectChip::configure(meta, advices[1..5].try_into().unwrap());
+
         VmConfig {
             primary,
             advices,
@@ -277,6 +287,7 @@ impl Circuit<pallas::Base> for ZkCircuit {
             native_253_range_check_config,
             lessthan_config,
             boolcheck_config,
+            condselect_config,
         }
     }
 
@@ -335,6 +346,9 @@ impl Circuit<pallas::Base> for ZkCircuit {
         // Construct the boolean check chip.
         let boolcheck_chip = SmallRangeCheckChip::construct(config.boolcheck_config.clone());
 
+        // Construct the conditional selectiono chip
+        let condselect_chip = config.condselect_chip();
+
         // ==========================
         // Constants setup
         // ==========================
@@ -805,6 +819,25 @@ impl Circuit<pallas::Base> for ZkCircuit {
                         .small_range_check(layouter.namespace(|| "copy boolean check"), w)?;
                 }
 
+                Opcode::CondSelect => {
+                    trace!(target: "zk::vm", "Executing `CondSelect{:?}` opcode", opcode.1);
+                    let args = &opcode.1;
+
+                    let cond: AssignedCell<Fp, Fp> = heap[args[0].1].clone().into();
+                    let lhs: AssignedCell<Fp, Fp> = heap[args[1].1].clone().into();
+                    let rhs: AssignedCell<Fp, Fp> = heap[args[2].1].clone().into();
+
+                    let out: AssignedCell<Fp, Fp> = condselect_chip.conditional_select(
+                        &mut layouter.namespace(|| "cond_select"),
+                        lhs,
+                        rhs,
+                        cond,
+                    )?;
+
+                    trace!(target: "zk::vm", "Pushing assignment to heap address {}", heap.len());
+                    heap.push(HeapVar::Base(out));
+                }
+
                 Opcode::ConstrainEqualBase => {
                     trace!(target: "zk::vm", "Executing `ConstrainEqualBase{:?}` opcode", opcode.1);
                     let args = &opcode.1;

+ 9 - 0
src/zkas/opcode.rs

@@ -78,6 +78,9 @@ pub enum Opcode {
     /// Check if a field element fits in a boolean (Either 0 or 1)
     BoolCheck = 0x53,
 
+    /// Conditionally select between two base field elements given a boolean
+    CondSelect = 0x60,
+
     /// Constrain equality of two Base field elements inside the circuit
     ConstrainEqualBase = 0xe0,
 
@@ -111,6 +114,7 @@ impl Opcode {
             "less_than_strict" => Some(Self::LessThanStrict),
             "less_than_loose" => Some(Self::LessThanLoose),
             "bool_check" => Some(Self::BoolCheck),
+            "cond_select" => Some(Self::CondSelect),
             "constrain_equal_base" => Some(Self::ConstrainEqualBase),
             "constrain_equal_point" => Some(Self::ConstrainEqualPoint),
             "constrain_instance" => Some(Self::ConstrainInstance),
@@ -138,6 +142,7 @@ impl Opcode {
             0x51 => Some(Self::LessThanStrict),
             0x52 => Some(Self::LessThanLoose),
             0x53 => Some(Self::BoolCheck),
+            0x60 => Some(Self::CondSelect),
             0xe0 => Some(Self::ConstrainEqualBase),
             0xe1 => Some(Self::ConstrainEqualPoint),
             0xf0 => Some(Self::ConstrainInstance),
@@ -194,6 +199,10 @@ impl Opcode {
 
             Opcode::BoolCheck => (vec![], vec![VarType::Base]),
 
+            Opcode::CondSelect => {
+                (vec![VarType::Base], vec![VarType::Base, VarType::Base, VarType::Base])
+            }
+
             Opcode::ConstrainEqualBase => (vec![], vec![VarType::Base, VarType::Base]),
 
             Opcode::ConstrainEqualPoint => (vec![], vec![VarType::EcPoint, VarType::EcPoint]),

+ 2 - 0
tests/zkvm_opcodes.rs

@@ -92,6 +92,7 @@ fn zkvm_opcodes() -> Result<()> {
         Witness::Base(Value::known(ephem_secret.inner())),
         Witness::Uint32(Value::known(leaf_pos.try_into().unwrap())),
         Witness::MerklePath(Value::known(merkle_path.try_into().unwrap())),
+        Witness::Base(Value::known(pallas::Base::ONE)),
     ];
 
     let value_commit = pedersen_commitment_u64(value, value_blind);
@@ -113,6 +114,7 @@ fn zkvm_opcodes() -> Result<()> {
         pub_y,
         ephem_x,
         ephem_y,
+        a,
     ];
 
     let circuit = ZkCircuit::new(prover_witnesses, zkbin.clone());