narodnik 5 سال پیش
والد
کامیت
f7b027c100
3فایلهای تغییر یافته به همراه353 افزوده شده و 35 حذف شده
  1. 31 29
      proofs/sapling3.prf
  2. 321 5
      scripts/parser.py
  3. 1 1
      scripts/sapvi.vim

+ 31 - 29
proofs/sapling3.prf

@@ -19,20 +19,21 @@ const:
     MERKLE: list<PedersenPersonalization>
     MERKLE: list<PedersenPersonalization>
     PRF_NF: Blake2sPersonalization
     PRF_NF: Blake2sPersonalization
 
 
-def value_commit(value: u64, randomness: Scalar) -> (Point, list<bool>):
-    let value_bits: list<bool> = value as list<bool>
+def value_commit(value: U64, randomness: Scalar) -> Binary:
+    let value_bits: Binary = value as Binary
     let value: Point = value * G_VCV
     let value: Point = value * G_VCV
 
 
-    let rcv: list<bool> = randomness as list<bool>
+    let rcv: Binary = randomness as Binary
     let rcv: Point = rcv * G_VCR
     let rcv: Point = rcv * G_VCR
 
 
     let cv: Point = value + rcv
     let cv: Point = value + rcv
-    return cv, value_bits
+    emit cv
+    return value_bits
 
 
 # The parameters to this function are the same as in:
 # The parameters to this function are the same as in:
 #   struct Spend
 #   struct Spend
 contract input_burn(
 contract input_burn(
-    value: u64,                 # ValueCommitment.value
+    value: U64,                 # ValueCommitment.value
     randomness: Scalar,         # ValueCommitment.randomness
     randomness: Scalar,         # ValueCommitment.randomness
 
 
     ak: Point,                  # from ProofGenerationKey
     ak: Point,                  # from ProofGenerationKey
@@ -44,34 +45,34 @@ contract input_burn(
 
 
     commitment_randomness: Scalar,
     commitment_randomness: Scalar,
 
 
-    auth_path: list<(Scalar, bool)>,
+    auth_path: list<(Scalar, Bool)>,
 
 
     anchor: Scalar
     anchor: Scalar
-) -> (Point, Point, Point, list<bool>):
+) -> (Point, Point, Point, Binary):
     let ak = witness(ak)
     let ak = witness(ak)
     ak.assert_not_small_order()
     ak.assert_not_small_order()
 
 
-    let ar: list<bool> = ar as list<bool>
+    let ar: Binary = ar as Binary
     let ar: Point = ar * G_SPEND
     let ar: Point = ar * G_SPEND
 
 
     let rk: Point = ak + ar
     let rk: Point = ak + ar
 
 
-    let nsk: list<bool> = nsk as list<bool>
+    let nsk: Binary = nsk as Binary
     let nk: Point = nsk * G_PROOF
     let nk: Point = nsk * G_PROOF
 
 
-    let mut ivk_preimage: list<bool> = []
-    # Must be list<bool> as well
+    let mut ivk_preimage: Binary = []
+    # Must be Binary as well
     ivk_preimage.extend(ak.repr())
     ivk_preimage.extend(ak.repr())
 
 
-    let mut nf_preimage: list<bool> = []
-    let nk_repr: list<bool> = nk.repr()
+    let mut nf_preimage: Binary = []
+    let nk_repr: Binary = nk.repr()
     ivk_preimage.extend(nk_repr)
     ivk_preimage.extend(nk_repr)
     nf_preimage.extend(nk_repr)
     nf_preimage.extend(nk_repr)
 
 
     assert len(ivk_preimage) == 512
     assert len(ivk_preimage) == 512
     assert len(nf_preimage) == 256
     assert len(nf_preimage) == 256
 
 
-    let mut ivk: list<bool> = blake2s(ivk_preimage, CRH_IVK)
+    let mut ivk: Binary = blake2s(ivk_preimage, CRH_IVK)
     ivk.truncate(Scalar::CAPACITY)
     ivk.truncate(Scalar::CAPACITY)
 
 
     let g_d: Point = witness g_d
     let g_d: Point = witness g_d
@@ -79,14 +80,14 @@ contract input_burn(
 
 
     let pk_d: Point = ivk * g_d
     let pk_d: Point = ivk * g_d
 
 
-    let mut note_contents: list<bool> = []
+    let mut note_contents: Binary = []
 
 
-    let (cv: Point, value_bits: list<bool>) = value_commit(value, randomness)
+    let (cv: Point, value_bits: Binary) = value_commit(value, randomness)
 
 
     let mut value_num: Num = Num.zero()
     let mut value_num: Num = Num.zero()
     let mut coeff: Scalar = Scalar.one()
     let mut coeff: Scalar = Scalar.one()
     for bit in value_bits:
     for bit in value_bits:
-        value_num = value_num.add_bool_with_coeff(bit, coeff)
+        value_num = value_num.add_Bool_with_coeff(bit, coeff)
         coeff = coeff.double()
         coeff = coeff.double()
     # Is this equivalent?
     # Is this equivalent?
     let value_num = value_bits as Num
     let value_num = value_bits as Num
@@ -98,21 +99,22 @@ contract input_burn(
     assert len(note_contents) == 64 + 256 + 256
     assert len(note_contents) == 64 + 256 + 256
 
 
     let mut cm: Point = pedersen_hash(NOTE_COMMIT, note_contents)
     let mut cm: Point = pedersen_hash(NOTE_COMMIT, note_contents)
-    let rcm: list<bool> = commitment_randomness as list<bool>
+    let rcm: Binary = commitment_randomness as Binary
     let rcm: Point = rcm * G_NOTE_COMMIT_R
     let rcm: Point = rcm * G_NOTE_COMMIT_R
     cm += rcm
     cm += rcm
 
 
-    let mut position_bits: list<bool> = []
+    let mut position_bits: Binary = []
     let mut cur: Scalar = cm.u
     let mut cur: Scalar = cm.u
 
 
+    # This should be fixed: for i in 0..N then use array indexing
     for i, (node, is_right) in enumerate(auth_path):
     for i, (node, is_right) in enumerate(auth_path):
         position_bits.push(is_right)
         position_bits.push(is_right)
 
 
         let node: EncryptedNum = EncryptedNum.from(node)
         let node: EncryptedNum = EncryptedNum.from(node)
         print(node)
         print(node)
-        let (left: list<bool>, right: list<bool>) = Num.swap_if(is_right, cur, node)
+        let (left: Binary, right: Binary) = Num.swap_if(is_right, cur, node)
 
 
-        let mut preimage: list<bool> = []
+        let mut preimage: Binary = []
         preimage.extend(left)
         preimage.extend(left)
         preimage.extend(right)
         preimage.extend(right)
 
 
@@ -127,12 +129,12 @@ contract input_burn(
 
 
     nf_preimage.extend(rho)
     nf_preimage.extend(rho)
     assert len(nf_preimage) == 512
     assert len(nf_preimage) == 512
-    let nf: list<bool> = blake2s(nf_preimage, PRF_NF)
+    let nf: Binary = blake2s(nf_preimage, PRF_NF)
 
 
     emit (rk, cv, rt, nf)
     emit (rk, cv, rt, nf)
 
 
 contract output_mint(
 contract output_mint(
-    value: u64,
+    value: U64,
     randomness: Scalar,
     randomness: Scalar,
 
 
     g_d: Point,
     g_d: Point,
@@ -142,20 +144,20 @@ contract output_mint(
 
 
     commitment_randomness: Scalar
     commitment_randomness: Scalar
 ) -> (Point, Point, Scalar):
 ) -> (Point, Point, Scalar):
-    let (cv: Point, value_bits: list<bool>) = value_commit(value, randomness)
+    let (cv: Point, value_bits: Binary) = value_commit(value, randomness)
 
 
-    let mut note_contents: list<bool> = []
+    let mut note_contents: Binary = []
     note_contents.extend(value_bits)
     note_contents.extend(value_bits)
 
 
     let g_d: Point = witness g_d
     let g_d: Point = witness g_d
     assert is_not_small_order(g_d)
     assert is_not_small_order(g_d)
 
 
-    let esk: list<bool> = esk as list<bool>
+    let esk: Binary = esk as Binary
     let epk: Point = esk * g_d
     let epk: Point = esk * g_d
 
 
-    let v_contents: list<bool> = pk_d.v as list<bool>
+    let v_contents: Binary = pk_d.v as Binary
 
 
-    let sign_bit: bool = pk_d.u.is_odd() as bool
+    let sign_bit: Bool = pk_d.u.is_odd() as Bool
 
 
     note_contents.extend(v_contents)
     note_contents.extend(v_contents)
     note_contents.push(sign_bit)
     note_contents.push(sign_bit)
@@ -164,7 +166,7 @@ contract output_mint(
 
 
     let mut cm: Point = pedersen_hash(NOTE_COMMIT, note_contents)
     let mut cm: Point = pedersen_hash(NOTE_COMMIT, note_contents)
 
 
-    let rcm: list<bool> = commitment_randomness as list<bool>
+    let rcm: Binary = commitment_randomness as Binary
     let rcm: Point = rcm * G_NOTE_COMMIT_R
     let rcm: Point = rcm * G_NOTE_COMMIT_R
 
 
     cm += rcm
     cm += rcm

+ 321 - 5
scripts/parser.py

@@ -110,9 +110,11 @@ def classify(sections):
 
 
 def tokenize_const(text):
 def tokenize_const(text):
     parser = lark.Lark(r"""
     parser = lark.Lark(r"""
-        value_map: NAME ":" type_def
+        value_map: name ":" type_def
 
 
-        type_def:   point
+        name: NAME
+
+        ?type_def:   point
                   | blake2s_personalization
                   | blake2s_personalization
                   | pedersen_personalization
                   | pedersen_personalization
                   | list
                   | list
@@ -131,13 +133,324 @@ def tokenize_const(text):
     """, start="value_map")
     """, start="value_map")
     return parser.parse(text)
     return parser.parse(text)
 
 
+class ConstTransformer(lark.Transformer):
+    def name(self, name):
+        return str(name[0])
+
+    def point(self, _):
+        return "Point"
+    def blake2s_personalization(self, _):
+        return "Blake2sPersonalization"
+    def pedersen_personalization(self, _):
+        return "PedersenPersonalization"
+    value_map = tuple
+    list = list
+
 def read_consts(consts):
 def read_consts(consts):
+    consts_map = {}
+
     for subsection in consts:
     for subsection in consts:
         assert subsection[0].text == "const:"
         assert subsection[0].text == "const:"
 
 
         for ldesc in subsection[1:]:
         for ldesc in subsection[1:]:
-            tokens = tokenize_const(ldesc.text)
-            print(tokens)
+            tree = tokenize_const(ldesc.text)
+            tokens = ConstTransformer().transform(tree)
+            #print(tokens)
+            name, typedesc = tokens
+            consts_map[name] = typedesc
+
+    #pprint.pprint(consts_map)
+    return consts_map
+
+class FuncDefTransformer(lark.Transformer):
+    def func_name(self, name):
+        return str(name[0])
+
+    def param(self, obj):
+        return tuple(obj)
+
+    def param_name(self, name):
+        return str(name[0])
+
+    def u64(self, _):
+        return "U64"
+    def scalar(self, _):
+        return "Scalar"
+    def point(self, _):
+        return "Point"
+    def binary(self, _):
+        return "Binary"
+
+    def type(self, obj):
+        return obj[0]
+
+    func_def = list
+    params = list
+    type_list = list
+
+def parse_func_def(text):
+    parser = lark.Lark(r"""
+        func_def: "def" func_name "(" params+ ")" "->" type_list ":"
+
+        func_name: NAME
+        params: param ("," param)*
+
+        type_list: type
+                 | "(" type ("," type)* ")"
+
+        param: param_name ":" type
+        param_name: NAME
+
+        type: u64 | scalar | point | binary
+
+        u64: "U64"
+        scalar: "Scalar"
+        point: "Point"
+        binary: "Binary"
+
+        %import common.CNAME -> NAME
+        %import common.WS
+        %ignore WS
+    """, start="func_def")
+    tree = parser.parse(text)
+    tokens = FuncDefTransformer().transform(tree)
+    assert len(tokens) == 3
+    return tokens
+
+def interpret_func(func, consts):
+    func_def = parse_func_def(func[0].text)
+
+    func_name, params, retvals = func_def
+    #print("Function:", func_name)
+    #print("Params:", params)
+    #print("Return values:", retvals)
+    #print()
+
+    param_str = ""
+    for param, type in params:
+        if param_str:
+            param_str += ", "
+        param_str += param + ": "
+        if type == "U64":
+            param_str += "u64"
+        elif type == "Scalar":
+            param_str += "&jubjub::Fr"
+        else:
+            print("error: unsupported param type", file=sys.stderr)
+            print("line:", line.text, "line:", line.lineno)
+            return None
+
+    converted_retvals = []
+    for type in retvals:
+        if type == "Binary":
+            converted_retvals.append("boolean::Boolean")
+        else:
+            print("error: unsupported return type", file=sys.stderr)
+            print("line:", line.text, "line:", line.lineno)
+            return None
+    retvals = converted_retvals
+
+    if len(retvals) == 1:
+        retstr = retvals[0]
+    else:
+        retstr = "(" + ", ".join(retvals) + ")"
+
+    subroutine = r"""
+fn %s<CS>(
+    mut cs: CS,
+    %s
+) -> Result<%s, SynthesisError>
+where
+    CS: ConstraintSystem<bls12_381::Scalar>,
+{
+""" % (func_name, param_str, retstr)
+
+    indent = " " * 4
+
+    stack = dict(params)
+    emitted_types = []
+    for line in func[1:]:
+        statement_type, statement = interpret_func_line(line.text, stack, consts)
+        if statement_type == "let":
+            is_mutable = False
+            if statement[0] == "mut":
+                is_mutable = True
+                statement = statement[1:]
+            variable_name, variable_type = statement[0], statement[1]
+            expr = statement[2]
+            #print("LET", is_mutable, variable_name, variable_type)
+            #print("  ", expr)
+            code = "let " + ("mut " if is_mutable else "") + variable_name + " = "
+
+            if expr.data == "as_expr":
+                var_from, type_to = expr.children
+                if var_from not in stack:
+                    print("error: variable from not in stack frame:", var_from,
+                          file=sys.stderr)
+                    print("line:", line.text, "line:", line.lineno)
+                    return None
+
+                type_from = stack[var_from]
+
+                if type_from == "U64" and type_to == "Binary":
+                    code += "boolean::u64_into_boolean_vec_le(" + \
+                        "cs.namespace(|| \"" + line.text + "\"), " + var_from + \
+                        ")?;"
+                elif type_from == "Scalar" and type_to == "Binary":
+                    code += "boolean::field_into_boolean_vec_le(" + \
+                        "cs.namespace(|| \"" + line.text + "\"), &" + var_from + \
+                        ")?;"
+                else:
+                    print("error: unknown type conversion!", file=sys.stderr)
+                    print("line:", line.text, "line:", line.lineno)
+                    return None
+
+                #print(var_from, type_from, type_to)
+                stack[variable_name] = type_to
+
+            elif expr.data == "mul_expr":
+                var_a, var_b = expr.children
+                #print("MUL", var_a, var_b)
+
+                if var_b not in consts:
+                    print("error: unknown base!", file=sys.stderr)
+                    print("line:", line.text, "line:", line.lineno)
+                    return None
+
+                base_type = consts[var_b]
+                if base_type != "Point":
+                    print("error: unknown base type!", file=sys.stderr)
+                    print("line:", line.text, "line:", line.lineno)
+                    return None
+
+                code += "ecc::fixed_base_multiplication(" + \
+                    "cs.namespace(|| \"" + line.text + "\"), &" + var_b + \
+                    ", &" + var_a + ")?;"
+                stack[variable_name] = "Point"
+
+            elif expr.data == "add_expr":
+                var_a, var_b = expr.children
+
+                if var_a not in stack or var_b not in stack:
+                    print("error: missing stack item!", file=sys.stderr)
+                    print("line:", line.text, "line:", line.lineno)
+                    return None
+
+                result_type = stack[var_a]
+                if stack[var_b] != result_type:
+                    print("error: non matching items for addition!", file=sys.stderr)
+                    print("line:", line.text, "line:", line.lineno)
+                    return None
+
+                code += var_a + ".add(cs.namespace(|| \"" + line.text \
+                    + "\"), &" + var_b + ")?;"
+                stack[variable_name] = result_type
+                    
+            subroutine += indent + code + "\n"
+
+        elif statement_type == "return":
+            for var in statement:
+                if var not in stack:
+                    print("error: missing variable in stack!", file=sys.stderr)
+                    print("line:", line.text, "line:", line.lineno)
+                    return None
+
+            if len(statement) == 1:
+                code = "Ok(" + statement[0] + ")"
+            else:
+                code = "Ok(" + ",".join(statement) + ")"
+            subroutine += indent + code + "\n"
+
+        elif statement_type == "emit":
+            assert len(statement) == 1
+            variable = statement[0]
+            if variable not in stack:
+                print("error: missing variable in stack!", file=sys.stderr)
+                print("line:", line.text, "line:", line.lineno)
+                return None
+
+            variable_type = stack[variable]
+
+            if variable_type == "Point":
+                code = variable + ".inputize(cs.namespace(|| \"" + \
+                    line.text + "\"))?;"
+            else:
+                print("error: unable to inputize type!", file=sys.stderr)
+                print("line:", line.text, "line:", line.lineno)
+                return None
+
+            emitted_types.append(variable_type)
+            subroutine += indent + code + "\n"
+
+    subroutine += "}"
+    print(subroutine)
+
+class CodeLineTransformer(lark.Transformer):
+    def variable_name(self, name):
+        return str(name[0])
+
+    def let_statement(self, obj):
+        return ("let", obj)
+    def return_statement(self, obj):
+        return ("return", obj)
+    def emit_statement(self, obj):
+        return ("emit", obj)
+
+    def point(self, _):
+        return "Point"
+    def scalar(self, _):
+        return "Scalar"
+    def binary(self, _):
+        return "Binary"
+    def u64(self, _):
+        return "U64"
+
+    def type(self, typename):
+        return str(typename[0])
+
+    def mutable(self, _):
+        return "mut"
+
+    statement = list
+
+def interpret_func_line(text, stack, consts):
+    parser = lark.Lark(r"""
+        statement: let_statement
+                 | return_statement
+                 | emit_statement
+
+        let_statement: "let" [mutable] variable_name ":" type "=" expr
+        mutable: "mut"
+
+        ?expr: as_expr
+            | mul_expr
+            | add_expr
+
+        as_expr: variable_name "as" type
+        mul_expr: variable_name "*" variable_name
+        add_expr: variable_name "+" variable_name
+
+        return_statement: "return" variable_name
+                        | "return" variable_tuple
+        variable_tuple: "(" variable_name ("," variable_name)* ")"
+
+        emit_statement: "emit" variable_name
+
+        variable_name: NAME
+        type: u64 | scalar | point | binary
+
+        u64: "U64"
+        scalar: "Scalar"
+        point: "Point"
+        binary: "Binary"
+
+        %import common.CNAME -> NAME
+        %import common.WS
+        %ignore WS
+    """, start="statement")
+    tree = parser.parse(text)
+    tokens = CodeLineTransformer().transform(tree)[0]
+    return tokens
 
 
 def main(argv):
 def main(argv):
     if len(argv) == 1:
     if len(argv) == 1:
@@ -153,7 +466,10 @@ def main(argv):
 
 
     consts, funcs, contracts = classify(sections)
     consts, funcs, contracts = classify(sections)
 
 
-    read_consts(consts)
+    consts = read_consts(consts)
+
+    for func in funcs:
+        interpret_func(func, consts)
 
 
 if __name__ == "__main__":
 if __name__ == "__main__":
     main(sys.argv)
     main(sys.argv)

+ 1 - 1
scripts/sapvi.vim

@@ -4,7 +4,7 @@ endif
 
 
 syn keyword sapviKeyword assert enforce for in def return const as let emit contract private proof
 syn keyword sapviKeyword assert enforce for in def return const as let emit contract private proof
 syn keyword sapviAttr mut
 syn keyword sapviAttr mut
-syn keyword sapviType Point Scalar EncryptedNum list bool u64 Num
+syn keyword sapviType Point Scalar EncryptedNum list Bool U64 Num Binary
 syn match sapviFunction "\zs[a-zA-Z0-9_]*\ze("
 syn match sapviFunction "\zs[a-zA-Z0-9_]*\ze("
 syn match sapviComment "#.*$"
 syn match sapviComment "#.*$"
 syn match sapviNumber '\d\+'
 syn match sapviNumber '\d\+'