proof.sage 6.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163
  1. '''
  2. bulletproof protocol 2 with multi-exponentiation.
  3. '''
  4. load('../mpc/curve.sage')
  5. load('../mpc/beaver.sage')
  6. load('utils.sage')
  7. class Proof(object):
  8. def __init__(self, transcript, Q, G_factors, H_factors, G, H, a, b):
  9. '''
  10. create inner product proof
  11. '''
  12. self.source = Source(p)
  13. n = len(G)
  14. assert (n == len(H) == len(H_factors) == len(a) == len(b))
  15. L_l = []
  16. R_l = []
  17. if n!=1:
  18. n /=2
  19. a_l, a_r = a[0:n], a[n:]
  20. b_l, b_r = b[0:n], b[n:]
  21. G_l, G_r = G[0:n], G[n:]
  22. H_l, H_r = H[0:n], H[n:]
  23. c_l = [sum([a*b for a,b in zip(a_l, b_r)])]
  24. c_r = [sum([a*b for a,b in zip(a_r, b_l)])]
  25. al_g = [al*g for al, g in zip(a_l, G_factors[n:2*n])]
  26. br_h = [br*h for br,h in zip(b_r, H_factors[0:n])]
  27. L_gr_al_g = CurvePoint.msm(G_r, al_g)
  28. L_hl_br_h = CurvePoint.msm(H_l, br_h)
  29. L_q_cl = CurvePoint.msm(Q, c_l)
  30. # L, R
  31. # note that P = L*R
  32. L = [sum([L_gr_al_g, L_hl_br_h , L_q_cl])]
  33. R = [sum([CurvePoint.msm(G_l, [ar*g for ar, g in zip(a_r, G_factors[0:n])]), CurvePoint.msm(H_r, [bl*h for bl,h in zip(b_l, H_factors[n:2*n])]), CurvePoint.msm(Q, c_r)])]
  34. L_l += L
  35. R_l += R
  36. # choose true random challenges u, u^{-1}
  37. #transcript.append_message(b'L', bytes(''.join([l.__str__() for l in L]), encoding='utf-8'))
  38. #transcript.append_message(b'R', bytes(''.join([r.__str__() for r in R]), encoding='utf-8'))
  39. #u = K(transcript.challenge_bytes(b'u'))
  40. u = K(1)
  41. u_inv = 1/u
  42. for i in range(n):
  43. # a_prime
  44. a_l[i] = a_l[i] * u + u_inv * a_r[i]
  45. # p_prime
  46. b_l[i] = b_l[i] * u_inv + u * b_r[i]
  47. # G_prime
  48. G_l[i] = CurvePoint.msm([G_l[i], G_r[i]], [u_inv * G_factors[i], u * G_factors[n+i]])
  49. # H_prime
  50. H_l[i] = CurvePoint.msm([H_l[i], H_r[i]], [u * H_factors[i], u_inv * H_factors[n+i]])
  51. a = a_l # a is a_prime
  52. b = b_l # b is b_prime
  53. G = G_l # G is G_prime
  54. H = H_l # H is H_prime
  55. while n!=1:
  56. n /=2
  57. a_l, a_r = a[0:n], a[n:] # a_prime_l, a_prime_r
  58. b_l, b_r = b[0:n], b[n:] # b_prime_l, b_prime_r
  59. G_l, G_r = G[0:n], G[n:] # G_prime_l, G_prime_r
  60. H_l, H_r = H[0:n], H[n:] # H_prime_l, H_prime_r
  61. c_l = [sum([a*b for (a,b) in zip(a_l, b_r)])] # c_prime_l
  62. c_r = [sum([a*b for (a,b) in zip(a_r, b_l)])] # c_prime_r
  63. # L_prime
  64. L = [sum([CurvePoint.msm(G_r, a_l), CurvePoint.msm(H_l, b_r), CurvePoint.msm(Q, c_l)])]
  65. # R_prime
  66. R = [sum([CurvePoint.msm(G_l, a_r), CurvePoint.msm(H_r, b_l), CurvePoint.msm(Q, c_r)])]
  67. L_l += L
  68. R_l += R
  69. # choose true random challenges u, u^{-1]}
  70. #transcript.append_message(b'L', bytes(''.join([l.__str__() for l in L]), encoding='utf-8'))
  71. #transcript.append_message(b'R', bytes(''.join([r.__str__() for r in R]), encoding='utf-8'))
  72. #u = K(transcript.challenge_bytes(b'u'))
  73. u = K(1)
  74. u_inv = 1/u
  75. for i in range(n):
  76. # u * a_prime_l + u^{-1} * a_prime_r
  77. a_l[i] = a_l[i] * u + u_inv * a_r[i]
  78. # u^{-1} * b_prime_l + u * b_prime_r
  79. b_l[i] = b_l[i] * u_inv + u * b_r[i]
  80. # G_l_prime
  81. G_l[i] = CurvePoint.msm([G_l[i], G_r[i]], [u_inv, u])
  82. # H_l_prime
  83. H_l[i] = CurvePoint.msm([H_l[i], H_r[i]], [u, u_inv])
  84. a = a_l
  85. b = b_l
  86. G = G_l
  87. H = H_l
  88. #
  89. self.lhs = L_l
  90. self.rhs = R_l
  91. self.a = a[0]
  92. self.b = b[0]
  93. def challenges(self, n, verifier):
  94. challenges = []
  95. challenges_inv = []
  96. lg_n = len(self.lhs)
  97. for L, R in zip(self.lhs, self.rhs):
  98. #verifier.append_message(b'L', bytes(''.join([l.__str__() for l in [L]]), encoding='utf-8'))
  99. #verifier.append_message(b'R', bytes(''.join([r.__str__() for r in [R]]), encoding='utf-8'))
  100. #u = K(verifier.challenge_bytes(b'u'))
  101. u = K(1)
  102. u_inv = 1/u
  103. challenges += [u]
  104. challenges_inv += [1/u]
  105. inv_prod = K(1)
  106. for u_inv in challenges_inv:
  107. inv_prod *=K(1)
  108. challenges_sq = [i*i for i in challenges]
  109. challenges_inv_sq = [i*i for i in challenges_inv]
  110. mul_inv = K(1)
  111. for i in challenges_inv:
  112. mul_inv *=i
  113. S = [mul_inv]
  114. for i in range(1,n):
  115. lg_i = 32 - 1 - countZeros(i)
  116. k = 1 << lg_i
  117. u_lg_i_sq = challenges_sq[(lg_n -1) - lg_i]
  118. S += [S[i-k] * u_lg_i_sq]
  119. return challenges_sq, challenges_inv_sq, S
  120. def verify(self, n, verifier, G_factors, H_factors, P, Q, G, H):
  121. u_sq, u_inv_sq, s = self.challenges(n, verifier)
  122. g_times_a_times_s = [self.a * s_i * g_i for g_i, s_i in zip(G_factors, s)][:n]
  123. # inverse of count is reverse
  124. inv_s = reversed(s)
  125. h_times_b_div_s = [self.b * s_i_inv * h_i for h_i, s_i_inv in zip(H_factors, inv_s)]
  126. neg_u_sq = [i*K(-1) for i in u_sq]
  127. neg_u_inv_sq = [i*K(-1) for i in u_inv_sq]
  128. # P
  129. ## u^c
  130. res_p_1 = CurvePoint.msm(Q, [self.a*self.b])
  131. ## g^{g_factor_a_s}
  132. res_p_2 = CurvePoint.msm(G, g_times_a_times_s)
  133. ## h^{h_factor_b_s}
  134. res_p_3 = CurvePoint.msm(H, h_times_b_div_s)
  135. # L^(u^2)
  136. print("L: {}".format(self.lhs))
  137. res_p_4 = CurvePoint.msm(self.lhs, neg_u_sq)
  138. # R^(u^-2)
  139. print('R: {}'.format(self.rhs))
  140. res_p_5 = CurvePoint.msm(self.rhs, neg_u_inv_sq)
  141. # P prime = L^{u^2} * P * R^{u^{-1}}
  142. print('p_1: {}'.format(res_p_1))
  143. print('p_2: {}'.format(res_p_2))
  144. print('p_3: {}'.format(res_p_3))
  145. print('p_4: {}'.format(res_p_4))
  146. print('p_5: {}'.format(res_p_5))
  147. res_p = res_p_1 + res_p_2 + res_p_3 + res_p_4 + res_p_5;
  148. res = res_p == P
  149. # P prime == H(u^{-1} * a_prime_r, u * a_prime_l, u * b_prime_r, u ^ {-1} * b_prime_l, c_prime)
  150. assert (res), 'P: {}, expected P: {}'.format(res_p, P)
  151. return res_p, P, res