groth16.sage 5.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177
  1. q = 0x40000000000000000000000000000000224698fc0994a8dd8c46eb2100000001
  2. K = GF(q)
  3. P.<X> = K[]
  4. # def foo(s, x, y):
  5. # if s:
  6. # return x * y
  7. # else:
  8. # return x + y
  9. # z = foo(s, x, y)
  10. # Arithmetization for:
  11. # sxy + (s - 1)(x + y) - z = 0
  12. # s(s - 1) = 0
  13. var_1 = K(1)
  14. var_x = K(4)
  15. var_y = K(6)
  16. var_s = K(1)
  17. var_xy = var_x*var_y
  18. var_sxy = var_s*var_xy
  19. # w1 = (s - 1)(x + y)
  20. var_w1 = (var_s - 1)*(var_x + var_y)
  21. var_z = var_sxy + var_w1
  22. # There are n = 8 variables
  23. # i = 0 1 2 3 4 5 6 7
  24. S = vector([var_1, var_x, var_y, var_s, var_xy, var_sxy, var_w1, var_z])
  25. # Row 1
  26. # var_x * var_y == var_xy
  27. L_1 = vector([0, 1, 0, 0, 0, 0, 0, 0])
  28. R_1 = vector([0, 0, 1, 0, 0, 0, 0, 0])
  29. O_1 = vector([0, 0, 0, 0, 1, 0, 0, 0])
  30. assert (L_1*S) * (R_1*S) == (O_1*S)
  31. # Row 2
  32. # var_s * var_xy == var_sxy
  33. L_2 = vector([0, 0, 0, 1, 0, 0, 0, 0])
  34. R_2 = vector([0, 0, 0, 0, 1, 0, 0, 0])
  35. O_2 = vector([0, 0, 0, 0, 0, 1, 0, 0])
  36. assert (L_2*S) * (R_2*S) == (O_2*S)
  37. # Row 3
  38. # (var_s - 1) * (var_x + var_y) == var_w1
  39. L_3 = vector([-1, 0, 0, 1, 0, 0, 0, 0])
  40. R_3 = vector([ 0, 1, 1, 0, 0, 0, 0, 0])
  41. O_3 = vector([ 0, 0, 0, 0, 0, 0, 1, 0])
  42. assert (L_3*S) * (R_3*S) == (O_3*S)
  43. # Row 4
  44. # Here we want to check z = sxy + w1, but we need every row to have
  45. # at least one multiplication, so we just use the constant 1 for the RHS.
  46. # (var_sxy + var_w1) * var_1 == var_z
  47. L_4 = vector([0, 0, 0, 0, 0, 1, 1, 0])
  48. R_4 = vector([1, 0, 0, 0, 0, 0, 0, 0])
  49. O_4 = vector([0, 0, 0, 0, 0, 0, 0, 1])
  50. assert (L_4*S) * (R_4*S) == (O_4*S)
  51. # Row 5
  52. # Boolean check for s.
  53. # var_s * (var_s - var_1) == 0
  54. L_5 = vector([ 0, 0, 0, 1, 0, 0, 0, 0])
  55. R_5 = vector([-1, 0, 0, 1, 0, 0, 0, 0])
  56. O_5 = vector([ 0, 0, 0, 0, 0, 0, 0, 0])
  57. assert (L_5*S) * (R_5*S) == (O_5*S)
  58. L = matrix([L_1, L_2, L_3, L_4, L_5])
  59. R = matrix([R_1, R_2, R_3, R_4, R_5])
  60. O = matrix([O_1, O_2, O_3, O_4, O_5])
  61. def hadamard_prod(A, B):
  62. result = []
  63. for a_i, b_i in zip(A, B):
  64. result.append(a_i * b_i)
  65. return vector(result)
  66. assert hadamard_prod(L*S, R*S) == O*S
  67. # Now extract columns from matrices
  68. L_i_1, L_i_2, L_i_3, L_i_4, L_i_5, L_i_6, L_i_7, L_i_8 = (
  69. L[:,0], L[:,1], L[:,2], L[:,3], L[:,4], L[:,5], L[:,6], L[:,7]
  70. )
  71. R_i_1, R_i_2, R_i_3, R_i_4, R_i_5, R_i_6, R_i_7, R_i_8 = (
  72. R[:,0], R[:,1], R[:,2], R[:,3], R[:,4], R[:,5], R[:,6], R[:,7]
  73. )
  74. O_i_1, O_i_2, O_i_3, O_i_4, O_i_5, O_i_6, O_i_7, O_i_8 = (
  75. O[:,0], O[:,1], O[:,2], O[:,3], O[:,4], O[:,5], O[:,6], O[:,7]
  76. )
  77. l_1_X = P.lagrange_polynomial((i, l_i_1) for i, (l_i_1,) in enumerate(L_i_1))
  78. l_2_X = P.lagrange_polynomial((i, l_i_2) for i, (l_i_2,) in enumerate(L_i_2))
  79. l_3_X = P.lagrange_polynomial((i, l_i_3) for i, (l_i_3,) in enumerate(L_i_3))
  80. l_4_X = P.lagrange_polynomial((i, l_i_4) for i, (l_i_4,) in enumerate(L_i_4))
  81. l_5_X = P.lagrange_polynomial((i, l_i_5) for i, (l_i_5,) in enumerate(L_i_5))
  82. l_6_X = P.lagrange_polynomial((i, l_i_6) for i, (l_i_6,) in enumerate(L_i_6))
  83. l_7_X = P.lagrange_polynomial((i, l_i_7) for i, (l_i_7,) in enumerate(L_i_7))
  84. l_8_X = P.lagrange_polynomial((i, l_i_8) for i, (l_i_8,) in enumerate(L_i_8))
  85. for i, row in enumerate(L):
  86. assert l_1_X(i) == row[0]
  87. assert l_2_X(i) == row[1]
  88. assert l_3_X(i) == row[2]
  89. assert l_4_X(i) == row[3]
  90. assert l_5_X(i) == row[4]
  91. assert l_6_X(i) == row[5]
  92. assert l_7_X(i) == row[6]
  93. assert l_8_X(i) == row[7]
  94. # l₁(X) represents var_1 which is i = 0
  95. # X=0 is row 1
  96. assert l_1_X(0) == L_1[0]
  97. # X=4 is row 5
  98. assert l_1_X(4) == L_5[0]
  99. # l₄(X) represents var_s which is i = 3
  100. # X=2 is row 3
  101. assert l_4_X(2) == L_3[3]
  102. r_1_X = P.lagrange_polynomial((i, r_i_1) for i, (r_i_1,) in enumerate(R_i_1))
  103. r_2_X = P.lagrange_polynomial((i, r_i_2) for i, (r_i_2,) in enumerate(R_i_2))
  104. r_3_X = P.lagrange_polynomial((i, r_i_3) for i, (r_i_3,) in enumerate(R_i_3))
  105. r_4_X = P.lagrange_polynomial((i, r_i_4) for i, (r_i_4,) in enumerate(R_i_4))
  106. r_5_X = P.lagrange_polynomial((i, r_i_5) for i, (r_i_5,) in enumerate(R_i_5))
  107. r_6_X = P.lagrange_polynomial((i, r_i_6) for i, (r_i_6,) in enumerate(R_i_6))
  108. r_7_X = P.lagrange_polynomial((i, r_i_7) for i, (r_i_7,) in enumerate(R_i_7))
  109. r_8_X = P.lagrange_polynomial((i, r_i_8) for i, (r_i_8,) in enumerate(R_i_8))
  110. for i, row in enumerate(R):
  111. assert r_1_X(i) == row[0]
  112. assert r_2_X(i) == row[1]
  113. assert r_3_X(i) == row[2]
  114. assert r_4_X(i) == row[3]
  115. assert r_5_X(i) == row[4]
  116. assert r_6_X(i) == row[5]
  117. assert r_7_X(i) == row[6]
  118. assert r_8_X(i) == row[7]
  119. # r₁(X) represents var_1 which is i = 0
  120. # X=4 is row 5
  121. assert r_1_X(4) == R_5[0]
  122. o_1_X = P.lagrange_polynomial((i, o_i_1) for i, (o_i_1,) in enumerate(O_i_1))
  123. o_2_X = P.lagrange_polynomial((i, o_i_2) for i, (o_i_2,) in enumerate(O_i_2))
  124. o_3_X = P.lagrange_polynomial((i, o_i_3) for i, (o_i_3,) in enumerate(O_i_3))
  125. o_4_X = P.lagrange_polynomial((i, o_i_4) for i, (o_i_4,) in enumerate(O_i_4))
  126. o_5_X = P.lagrange_polynomial((i, o_i_5) for i, (o_i_5,) in enumerate(O_i_5))
  127. o_6_X = P.lagrange_polynomial((i, o_i_6) for i, (o_i_6,) in enumerate(O_i_6))
  128. o_7_X = P.lagrange_polynomial((i, o_i_7) for i, (o_i_7,) in enumerate(O_i_7))
  129. o_8_X = P.lagrange_polynomial((i, o_i_8) for i, (o_i_8,) in enumerate(O_i_8))
  130. for i, row in enumerate(O):
  131. assert o_1_X(i) == row[0]
  132. assert o_2_X(i) == row[1]
  133. assert o_3_X(i) == row[2]
  134. assert o_4_X(i) == row[3]
  135. assert o_5_X(i) == row[4]
  136. assert o_6_X(i) == row[5]
  137. assert o_7_X(i) == row[6]
  138. assert o_8_X(i) == row[7]
  139. l_X = vector([l_1_X, l_2_X, l_3_X, l_4_X, l_5_X, l_6_X, l_7_X, l_8_X])
  140. r_X = vector([r_1_X, r_2_X, r_3_X, r_4_X, r_5_X, r_6_X, r_7_X, r_8_X])
  141. o_X = vector([o_1_X, o_2_X, o_3_X, o_4_X, o_5_X, o_6_X, o_7_X, o_8_X])
  142. # Evaluate each row
  143. for q in range(5):
  144. lhs = sum(S[i]*l_X[i](q) for i in range(8))
  145. rhs = sum(S[i]*r_X[i](q) for i in range(8))
  146. out = sum(S[i]*o_X[i](q) for i in range(8))
  147. assert lhs*rhs == out
  148. # So this and the matrix form are both equivalent
  149. t = (S*l_X) * (S*r_X) - S*o_X
  150. for i in range(5):
  151. assert t(i) == 0
  152. z = (X - 0)*(X - 1)*(X - 2)*(X - 3)*(X - 4)
  153. h, rem = t.quo_rem(z)
  154. assert rem == 0