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

[research/buletproof-mpc] implement bulletproof over mpc]

ertosns 2 лет назад
Родитель
Сommit
27e9d70c85

+ 16 - 9
script/research/bulletproof-mpc/inner_product_proof.sage

@@ -2,15 +2,16 @@ load('../mpc/curve.sage')
 load('proof.sage')
 load('transcript.sage')
 
-n = 4
-Q = [CurvePoint.random()]
-H = [CurvePoint.random() for i in range(0,n)]
-G = [CurvePoint.random() for i in range(0,n)]
+n = 2
+Q = [CurvePoint.generator()]
+H = [CurvePoint.generator() for i in range(0,n)]
+G = [CurvePoint.generator() for i in range(0,n)]
 
-a = [K(random.randint(0,p)) for _ in range(0,n)]
-b = [K(random.randint(0,p)) for _ in range(0,n)]
+a = [1, 2]
+b = [2, 4]
 c = [sum([a*b for a, b in zip(a, b)])]
-y_inv = K(random.randint(0,p))
+print('c: {}'.format(c))
+y_inv = K(1)
 
 G_factors = [K(1)]*n
 H_factors = [y_inv**i for i in range(0,n)]
@@ -20,7 +21,13 @@ a_prime = a.copy()
 
 transcript = Transcript('bulletproof')
 proof = Proof(transcript, Q, G_factors, H_factors, G, H, a, b)
-
-P_res =  sum([CurvePoint.msm(G, a_prime), CurvePoint.msm(H, b_prime), CurvePoint.msm(Q, c)])
+g_a_prime = CurvePoint.msm(G, a_prime)
+h_b_prime = CurvePoint.msm(H, b_prime)
+q_c = CurvePoint.msm(Q, c)
+P_res =  sum([g_a_prime, h_b_prime, q_c])
+print('g_a_prime: {}'.format(g_a_prime))
+print('h_b_prime: {}'.format(h_b_prime))
+print('q_c: {}'.format(q_c))
+print('P: {}'.format(P_res))
 verifier = Transcript('bulletproof')
 pp, p, _ = proof.verify(n, verifier, G_factors, H_factors, P_res, Q, G, H)

+ 19 - 9
script/research/bulletproof-mpc/proof.sage

@@ -36,9 +36,10 @@ class Proof(object):
                 R_l += R
 
                 # choose true random challenges u, u^{-1}
-                transcript.append_message(b'L', bytes(''.join([l.__str__() for l in L]), encoding='utf-8'))
-                transcript.append_message(b'R', bytes(''.join([r.__str__() for r in R]), encoding='utf-8'))
-                u = K(transcript.challenge_bytes(b'u'))
+                #transcript.append_message(b'L', bytes(''.join([l.__str__() for l in L]), encoding='utf-8'))
+                #transcript.append_message(b'R', bytes(''.join([r.__str__() for r in R]), encoding='utf-8'))
+                #u = K(transcript.challenge_bytes(b'u'))
+                u = K(1)
                 u_inv = 1/u
 
                 for i in range(n):
@@ -74,10 +75,11 @@ class Proof(object):
                 R_l += R
 
                 # choose true random challenges u, u^{-1]}
-                transcript.append_message(b'L', bytes(''.join([l.__str__() for l in L]), encoding='utf-8'))
-                transcript.append_message(b'R', bytes(''.join([r.__str__() for r in R]), encoding='utf-8'))
+                #transcript.append_message(b'L', bytes(''.join([l.__str__() for l in L]), encoding='utf-8'))
+                #transcript.append_message(b'R', bytes(''.join([r.__str__() for r in R]), encoding='utf-8'))
 
-                u = K(transcript.challenge_bytes(b'u'))
+                #u = K(transcript.challenge_bytes(b'u'))
+                u = K(1)
                 u_inv = 1/u
                 for i in range(n):
                     # u * a_prime_l + u^{-1} * a_prime_r
@@ -103,9 +105,10 @@ class Proof(object):
           challenges_inv = []
           lg_n = len(self.lhs)
           for L, R in zip(self.lhs, self.rhs):
-              verifier.append_message(b'L', bytes(''.join([l.__str__() for l in [L]]), encoding='utf-8'))
-              verifier.append_message(b'R', bytes(''.join([r.__str__() for r in [R]]), encoding='utf-8'))
-              u = K(verifier.challenge_bytes(b'u'))
+              #verifier.append_message(b'L', bytes(''.join([l.__str__() for l in [L]]), encoding='utf-8'))
+              #verifier.append_message(b'R', bytes(''.join([r.__str__() for r in [R]]), encoding='utf-8'))
+              #u = K(verifier.challenge_bytes(b'u'))
+              u = K(1)
               u_inv = 1/u
               challenges += [u]
               challenges_inv += [1/u]
@@ -142,10 +145,17 @@ class Proof(object):
           ## h^{h_factor_b_s}
           res_p_3 = CurvePoint.msm(H, h_times_b_div_s)
           # L^(u^2)
+          print("L: {}".format(self.lhs))
           res_p_4 = CurvePoint.msm(self.lhs, neg_u_sq)
           # R^(u^-2)
+          print('R: {}'.format(self.rhs))
           res_p_5 = CurvePoint.msm(self.rhs, neg_u_inv_sq)
           # P prime  = L^{u^2} * P * R^{u^{-1}}
+          print('p_1: {}'.format(res_p_1))
+          print('p_2: {}'.format(res_p_2))
+          print('p_3: {}'.format(res_p_3))
+          print('p_4: {}'.format(res_p_4))
+          print('p_5: {}'.format(res_p_5))
           res_p = res_p_1 + res_p_2 + res_p_3 + res_p_4 + res_p_5;
           res = res_p == P
           # P prime == H(u^{-1} * a_prime_r, u * a_prime_l, u * b_prime_r, u ^ {-1} * b_prime_l, c_prime)

+ 23 - 0
script/research/bulletproof-mpc/utils.sage

@@ -1,3 +1,5 @@
+load('../mpc/share.sage')
+
 def countZeros(x):
     total_bits = 32
     res = 0
@@ -7,3 +9,24 @@ def countZeros(x):
         res += 1
         count += 1
     return res
+
+def sum_shares(shares, source, party_id):
+    zero_share = AuthenticatedShare(0, source, party_id)
+    for share in shares:
+        zero_share += share
+    return zero_share
+
+def sum_ec_shares(shares):
+    zero_share = ECAuthenticatedShare(0)
+    for share in shares:
+        zero_share += share
+    return zero_share
+
+def shares_mul(my_shares, peer_shares):
+    return sum([my_share * peer_share for my_share, peer_share in zip(my_shares, peer_shares)])
+
+def to_ec_shares(ec):
+    return ECAuthenticatedShare(ec)
+
+def to_ec_shares_list (ec_list):
+    return [to_ec_shares(ec) for ec in ec_list]