narodnik hace 5 años
padre
commit
f7b027c100
Se han modificado 3 ficheros con 353 adiciones y 35 borrados
  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>
     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 rcv: list<bool> = randomness as list<bool>
+    let rcv: Binary = randomness as Binary
     let rcv: Point = rcv * G_VCR
 
     let cv: Point = value + rcv
-    return cv, value_bits
+    emit cv
+    return value_bits
 
 # The parameters to this function are the same as in:
 #   struct Spend
 contract input_burn(
-    value: u64,                 # ValueCommitment.value
+    value: U64,                 # ValueCommitment.value
     randomness: Scalar,         # ValueCommitment.randomness
 
     ak: Point,                  # from ProofGenerationKey
@@ -44,34 +45,34 @@ contract input_burn(
 
     commitment_randomness: Scalar,
 
-    auth_path: list<(Scalar, bool)>,
+    auth_path: list<(Scalar, Bool)>,
 
     anchor: Scalar
-) -> (Point, Point, Point, list<bool>):
+) -> (Point, Point, Point, Binary):
     let ak = witness(ak)
     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 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 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())
 
-    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)
     nf_preimage.extend(nk_repr)
 
     assert len(ivk_preimage) == 512
     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)
 
     let g_d: Point = witness g_d
@@ -79,14 +80,14 @@ contract input_burn(
 
     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 coeff: Scalar = Scalar.one()
     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()
     # Is this equivalent?
     let value_num = value_bits as Num
@@ -98,21 +99,22 @@ contract input_burn(
     assert len(note_contents) == 64 + 256 + 256
 
     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
     cm += rcm
 
-    let mut position_bits: list<bool> = []
+    let mut position_bits: Binary = []
     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):
         position_bits.push(is_right)
 
         let node: EncryptedNum = EncryptedNum.from(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(right)
 
@@ -127,12 +129,12 @@ contract input_burn(
 
     nf_preimage.extend(rho)
     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)
 
 contract output_mint(
-    value: u64,
+    value: U64,
     randomness: Scalar,
 
     g_d: Point,
@@ -142,20 +144,20 @@ contract output_mint(
 
     commitment_randomness: 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)
 
     let g_d: Point = witness 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 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.push(sign_bit)
@@ -164,7 +166,7 @@ contract output_mint(
 
     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
 
     cm += rcm

+ 321 - 5
scripts/parser.py

@@ -110,9 +110,11 @@ def classify(sections):
 
 def tokenize_const(text):
     parser = lark.Lark(r"""
-        value_map: NAME ":" type_def
+        value_map: name ":" type_def
 
-        type_def:   point
+        name: NAME
+
+        ?type_def:   point
                   | blake2s_personalization
                   | pedersen_personalization
                   | list
@@ -131,13 +133,324 @@ def tokenize_const(text):
     """, start="value_map")
     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):
+    consts_map = {}
+
     for subsection in consts:
         assert subsection[0].text == "const:"
 
         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):
     if len(argv) == 1:
@@ -153,7 +466,10 @@ def main(argv):
 
     consts, funcs, contracts = classify(sections)
 
-    read_consts(consts)
+    consts = read_consts(consts)
+
+    for func in funcs:
+        interpret_func(func, consts)
 
 if __name__ == "__main__":
     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 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 sapviComment "#.*$"
 syn match sapviNumber '\d\+'