curve_tree_proofs.sage 9.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408
  1. ProofWitness = namedtuple("ProofWitness", [
  2. "v1", "v2", "v3", "v4", "r", "b", "σ"
  3. ])
  4. class ProofPublic:
  5. def __init__(self):
  6. self.C = None
  7. self.D = None
  8. self.X = None
  9. self.Y = None
  10. self.Z = None
  11. class ProofCommits:
  12. def __init__(self):
  13. self.v1 = None
  14. self.v2 = None
  15. self.v3 = None
  16. self.v4 = None
  17. self.r = None
  18. self.b = None
  19. self.σ = None
  20. self.σ_G1 = None
  21. self.σ_G2 = None
  22. self.v3_G1 = None
  23. self.v4_G2 = None
  24. self.blind_x = None
  25. self.blind_y = None
  26. self.blind_z = None
  27. # Used by inner product
  28. self.C0 = None
  29. self.C1 = None
  30. def transcript(self):
  31. points = [
  32. self.v1, self.v2, self.v3, self.v4, self.r, self.b, #self.σ,
  33. self.σ_G1, self.σ_G2, self.v3_G1, self.v4_G2,
  34. self.blind_x, self.blind_y, self.blind_z, self.C0, self.C1
  35. ]
  36. assert all(P is not None for P in points)
  37. points = [P.xy() for P in points]
  38. return list(zip(*points))
  39. class ProofResponses:
  40. def __init__(self):
  41. self.v1 = None
  42. self.v2 = None
  43. self.v3 = None
  44. self.v4 = None
  45. self.r = None
  46. self.b = None
  47. self.σ = None
  48. self.blind_x = None
  49. self.blind_y = None
  50. self.blind_z = None
  51. self.txy = None
  52. Proof = namedtuple("Proof", [
  53. "R", "s", "boolean_check"
  54. ])
  55. RingProof = namedtuple("RingProof", [
  56. "c0", "s0", "s1"
  57. ])
  58. def make_proof(Ei, witness):
  59. G1, G2, G3, G4, H = gens[Ei]
  60. S = Scalar[Ei]
  61. blind_x = int(S.random_element())
  62. blind_y = int(S.random_element())
  63. blind_xy = int(S.random_element())
  64. # We want that blind_xy + blind_z == witness.b
  65. blind_z = int(S(witness.b - blind_xy))
  66. k_v1 = int(S.random_element())
  67. k_v2 = int(S.random_element())
  68. k_v3 = int(S.random_element())
  69. k_v4 = int(S.random_element())
  70. k_r = int(S.random_element())
  71. k_b = int(S.random_element())
  72. k_σ = int(S.random_element())
  73. k_blind_x = int(S.random_element())
  74. k_blind_y = int(S.random_element())
  75. k_blind_xy = int(S.random_element())
  76. k_blind_z = int(S.random_element())
  77. # Used for inner product
  78. k_t0 = int(S.random_element())
  79. k_t1 = int(S.random_element())
  80. R = ProofCommits()
  81. R.v1 = k_v1 * G1
  82. R.v2 = k_v2 * G2
  83. R.v3 = k_v3 * G3
  84. R.v4 = k_v4 * G4
  85. R.r = k_r * H
  86. R.b = k_b * H
  87. # Used for 2nd proof
  88. R.σ_G1 = k_σ * G1
  89. R.σ_G2 = k_σ * G2
  90. R.v3_G1 = k_v3 * G1
  91. R.v4_G2 = k_v4 * G2
  92. R.blind_x = k_blind_x * H
  93. R.blind_y = k_blind_y * H
  94. R.blind_xy = k_blind_xy * H
  95. R.blind_z = k_blind_z * H
  96. # σ (v1 - v3)
  97. # σ (v2 - v4)
  98. # (k_σ + c σ)(k_v1 - k_v3 + c*(v1 - v3))
  99. #
  100. # sage: var("k_σ c σ k_v1 k_v3 v1 v3")
  101. # sage: ((k_σ + c*σ)*(k_v1 - k_v3 + c*(v1 - v3))).expand().collect(c)
  102. # (v1*σ - v3*σ)*c^2 + (k_σ*v1 - k_σ*v3 + k_v1*σ - k_v3*σ)*c + k_v1*k_σ - k_v3*k_σ
  103. R.C0 = (
  104. (k_v3*k_σ - k_v1*k_σ) * G1 +
  105. (k_v4*k_σ - k_v2*k_σ) * G2 +
  106. k_t0 * H
  107. )
  108. R.C1 = (
  109. (k_σ*witness.v3 - k_σ*witness.v1 + k_v3*witness.σ - k_v1*witness.σ) * G1 +
  110. (k_σ*witness.v4 - k_σ*witness.v2 + k_v4*witness.σ - k_v2*witness.σ) * G2 +
  111. k_t1 * H
  112. )
  113. c = hash_scalar(Ei, R.transcript())
  114. s = ProofResponses()
  115. s.v1 = int( k_v1 + c*witness.v1 )
  116. s.v2 = int( k_v2 + c*witness.v2 )
  117. s.v3 = int( k_v3 + c*witness.v3 )
  118. s.v4 = int( k_v4 + c*witness.v4 )
  119. s.r = int( k_r + c*witness.r )
  120. s.b = int( k_b + c*witness.b )
  121. s.σ = int( k_σ + c*witness.σ )
  122. s.blind_x = int(k_blind_x + c*blind_x)
  123. s.blind_y = int(k_blind_y + c*blind_y)
  124. s.blind_xy = int(k_blind_xy + c*blind_xy)
  125. s.blind_z = int(k_blind_z + c*blind_z)
  126. s.txy = c**2 * blind_xy + c * k_t1 + k_t0
  127. public = ProofPublic()
  128. public.X = ((witness.v3 - witness.v1) * G1 +
  129. (witness.v4 - witness.v2) * G2 +
  130. blind_x * H)
  131. public.Y = witness.σ * G1 + witness.σ * G2 + blind_y * H
  132. public.XY = (
  133. witness.σ * (witness.v3 - witness.v1) * G1 +
  134. witness.σ * (witness.v4 - witness.v2) * G2 +
  135. blind_xy * H
  136. )
  137. public.Z = witness.v1 * G1 + witness.v2 * G2 + blind_z * H
  138. assert witness.σ in (0, 1)
  139. if witness.σ == 0:
  140. assert public.XY == blind_xy * H
  141. assert (
  142. public.XY + public.Z
  143. ==
  144. witness.v1*G1 + witness.v2*G2 + (blind_xy + blind_z)*H
  145. )
  146. else:
  147. assert witness.σ == 1
  148. assert (
  149. public.XY
  150. ==
  151. (witness.v3 - witness.v1) * G1 +
  152. (witness.v4 - witness.v2) * G2 +
  153. blind_xy * H
  154. )
  155. assert (
  156. public.XY + public.Z
  157. ==
  158. witness.v3*G1 + witness.v4*G2 + (blind_xy + blind_z)*H
  159. )
  160. assert blind_xy + blind_z == witness.b
  161. P1 = public.Y
  162. P2 = public.Y - G1 - G2
  163. assert blind_y*H == [P1, P2][witness.σ]
  164. if witness.σ == 0:
  165. assert blind_y*H == P1
  166. assert blind_y*H - G1 - G2 == P2
  167. else:
  168. assert witness.σ == 1
  169. assert blind_y*H + G1 + G2 == P1
  170. assert blind_y*H == P2
  171. boolean_check = make_ring_sig(Ei, [P1, P2], blind_y, int(witness.σ))
  172. assert verify_ring_sig(Ei, boolean_check, [P1, P2])
  173. return Proof(R, s, boolean_check), public
  174. def make_ring_sig(Ei, public_keys, secret, j):
  175. H, S = gens[Ei][-1], Scalar[Ei]
  176. assert len(public_keys) == 2
  177. assert secret*H == public_keys[j]
  178. assert j in (0, 1)
  179. k0 = int(S.random_element())
  180. R0 = k0*H
  181. c1 = hash_scalar(Ei, R0.xy())
  182. s1 = int(S.random_element())
  183. R1 = s1*H - c1*public_keys[(j + 1) % 2]
  184. c0 = hash_scalar(Ei, R1.xy())
  185. s0 = k0 + c0*secret
  186. if j == 1:
  187. c0 = c1
  188. s0, s1 = s1, s0
  189. proof = RingProof(c0, s0, s1)
  190. return proof
  191. def verify_ring_sig(Ei, proof, public_keys):
  192. H = gens[Ei][-1]
  193. S = Scalar[Ei]
  194. assert len(public_keys) == 2
  195. R1 = proof.s0*H - proof.c0*public_keys[0]
  196. c1 = hash_scalar(Ei, R1.xy())
  197. R2 = proof.s1*H - c1*public_keys[1]
  198. c2 = hash_scalar(Ei, R2.xy())
  199. return c2 == proof.c0
  200. def verify_proof(Ei, proof, public):
  201. G1, G2, G3, G4, H = gens[Ei]
  202. S = Scalar[Ei]
  203. R, s = proof.R, proof.s
  204. c = hash_scalar(Ei, R.transcript())
  205. if (s.v1 * G1 +
  206. s.v2 * G2 +
  207. s.v3 * G3 +
  208. s.v4 * G4 +
  209. s.r * H
  210. !=
  211. R.v1 + R.v2 + R.v3 + R.v4 + R.r + c*public.C
  212. ):
  213. return False
  214. # Now we want to prove that
  215. # D = v1 G1 + v2 G2 + b H
  216. # or
  217. # D = v3 G1 + v4 G2 + b H
  218. # We do this by checking:
  219. # X = (v1 - v2)G + b_X H
  220. # Y = σ G + b_Y H
  221. # D = xy G + v2 G + b_D H
  222. # σ ∈ {0, 1}
  223. # X = (v1 - v2)G + b_X H
  224. if (s.v3 * G1 - s.v1 * G1 +
  225. s.v4 * G2 - s.v2 * G2 +
  226. s.blind_x * H
  227. !=
  228. R.v3_G1 - R.v1 + R.v4_G2 - R.v2 + R.blind_x + c*public.X
  229. ):
  230. return False
  231. # Y = σ G + b_Y H
  232. if (s.σ * G1 + s.σ * G2 + s.blind_y * H
  233. !=
  234. R.σ_G1 + R.σ_G2 + R.blind_y + c*public.Y
  235. ):
  236. return False
  237. # Z = v1 G1 + v2 G2
  238. if (s.v1 * G1 + s.v2 * G2 + s.blind_z * H
  239. !=
  240. R.v1 + R.v2 + R.blind_z + c*public.Z
  241. ):
  242. return False
  243. # Inner product verification. We select either P1 or P2
  244. # prove D1 = x1 y1 G1 + b1 H
  245. if (s.σ*(s.v3 - s.v1)*G1 + s.σ*(s.v4 - s.v2)*G2 + s.txy*H
  246. !=
  247. c**2*public.XY + c*R.C1 + R.C0
  248. ):
  249. return False
  250. # check D is correct
  251. if public.D != public.XY + public.Z:
  252. return False
  253. # boolean check proof for s
  254. P1 = public.Y
  255. P2 = public.Y - G1 - G2
  256. if not verify_ring_sig(Ei, proof.boolean_check, [P1, P2]):
  257. return False
  258. return True
  259. def hash_scalar(Ei, values):
  260. S = Scalar[Ei]
  261. hasher = hashlib.sha256()
  262. for value in values:
  263. hasher.update(str(value).encode())
  264. return S(int(hasher.hexdigest(), 16))
  265. # Test proving system
  266. def test_proof():
  267. G1, G2, G3, G4, H = gens[E1]
  268. S = Scalar[E1]
  269. P1, P2 = [E[E2].random_point() for _ in range(2)]
  270. (P1_x, P1_y), (P2_x, P2_y) = P1.xy(), P2.xy()
  271. r, b = [S.random_element() for _ in range(2)]
  272. C = hash_nodes(E1, P1, P2, r)
  273. # σ = 0 for P1, or σ = 1 for P2
  274. σ = S(1)
  275. D = hash_point(E1, P2, b)
  276. proof, public = make_proof(
  277. E1,
  278. ProofWitness(
  279. P1_x,
  280. P1_y,
  281. P2_x,
  282. P2_y,
  283. r,
  284. b,
  285. σ
  286. )
  287. )
  288. public.C = C
  289. public.D = D
  290. assert verify_proof(E1, proof, public)
  291. # Now try the other side too
  292. σ = S(0)
  293. D = hash_point(E1, P1, b)
  294. proof, public = make_proof(
  295. E1,
  296. ProofWitness(
  297. P1_x,
  298. P1_y,
  299. P2_x,
  300. P2_y,
  301. r,
  302. b,
  303. σ
  304. )
  305. )
  306. public.C = C
  307. public.D = D
  308. assert verify_proof(E1, proof, public)
  309. # Test the ring sigs too
  310. secret = int(S.random_element())
  311. P1 = secret*H
  312. P2 = E[E1].random_point()
  313. proof = make_ring_sig(E1, [P1, P2], secret, 0)
  314. assert verify_ring_sig(E1, proof, [P1, P2])
  315. # Also try in reverse
  316. P1, P2 = P2, P1
  317. proof = make_ring_sig(E1, [P1, P2], secret, 1)
  318. assert verify_ring_sig(E1, proof, [P1, P2])
  319. # Ring sigs is our boolean proof for σ
  320. σ = S(0)
  321. b = S.random_element()
  322. # Verifier only has P
  323. # We prove that σ ∈ {0, 1}
  324. P = σ*G1 + b*H
  325. # They can only make a ring signature on H
  326. # if σ is 0 or 1
  327. # P1 = P represents σ = 0
  328. P1 = P
  329. # P2 = P - G1 represents σ = 1
  330. P2 = P - G1
  331. proof = make_ring_sig(E1, [P1, P2], b, 0)
  332. assert verify_ring_sig(E1, proof, [P1, P2])
  333. # Also try σ = 1
  334. σ = S(1)
  335. P = σ*G1 + b*H
  336. P1 = P
  337. P2 = P - G1
  338. proof = make_ring_sig(E1, [P1, P2], b, 1)
  339. assert verify_ring_sig(E1, proof, [P1, P2])