main.py 9.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296
  1. import hashlib
  2. import sys
  3. from collections import namedtuple
  4. from classnamespace import ClassNamespace
  5. from crypto import pallas_curve
  6. class TransactionBuilder:
  7. def __init__(self, ec):
  8. self.clear_inputs = []
  9. self.inputs = []
  10. self.outputs = []
  11. self.ec = ec
  12. def add_clear_input(self, value, token_id, signature_secret):
  13. clear_input = ClassNamespace()
  14. clear_input.value = value
  15. clear_input.token_id = token_id
  16. clear_input.signature_secret = signature_secret
  17. self.clear_inputs.append(clear_input)
  18. def add_input(self, input):
  19. self.inputs.append(input)
  20. def add_output(self, value, token_id, public):
  21. output = ClassNamespace()
  22. output.value = value
  23. output.token_id = token_id
  24. output.public = public
  25. self.outputs.append(output)
  26. def compute_remainder_blind(self, clear_inputs, input_blinds,
  27. output_blinds):
  28. total = 0
  29. total += sum(input.value_blind for input in clear_inputs)
  30. total += sum(input_blinds)
  31. total -= sum(output_blinds)
  32. return total % self.ec.order
  33. def build(self):
  34. tx = Transaction(self.ec)
  35. token_blind = self.ec.random_scalar()
  36. for input in self.clear_inputs:
  37. tx_clear_input = ClassNamespace()
  38. tx_clear_input.__name__ = "TransactionClearInput"
  39. tx_clear_input.value = input.value
  40. tx_clear_input.token_id = input.token_id
  41. tx_clear_input.value_blind = self.ec.random_scalar()
  42. tx_clear_input.token_blind = token_blind
  43. tx_clear_input.signature_public = self.ec.multiply(
  44. input.signature_secret, self.ec.G)
  45. tx.clear_inputs.append(tx_clear_input)
  46. input_blinds = []
  47. for input in self.inputs:
  48. tx_input = ClassNamespace()
  49. tx.inputs.append(tx_input)
  50. assert self.outputs
  51. output_blinds = []
  52. for i, output in enumerate(self.outputs):
  53. if i == len(self.outputs) - 1:
  54. value_blind = self.compute_remainder_blind(
  55. tx.clear_inputs, input_blinds, output_blinds)
  56. else:
  57. value_blind = self.ec.random_scalar()
  58. output_blinds.append(value_blind)
  59. note = ClassNamespace()
  60. note.serial = self.ec.random_base()
  61. note.value = output.value
  62. note.token_id = output.token_id
  63. note.coin_blind = self.ec.random_base()
  64. note.value_blind = value_blind
  65. note.token_blind = token_blind
  66. tx_output = ClassNamespace()
  67. tx_output.__name__ = "TransactionOutput"
  68. tx_output.mint_proof = MintProof(
  69. note.value, note.token_id, note.value_blind,
  70. note.token_blind, note.serial, note.coin_blind,
  71. output.public, self.ec)
  72. tx_output.revealed = tx_output.mint_proof.get_revealed()
  73. assert tx_output.mint_proof.verify(tx_output.revealed)
  74. # Is normally encrypted
  75. tx_output.enc_note = note
  76. tx_output.enc_note.__name__ = "TransactionOutputEncryptedNote"
  77. tx.outputs.append(tx_output)
  78. unsigned_tx_data = tx.partial_encode()
  79. for (input, info) in zip(tx.clear_inputs, self.clear_inputs):
  80. secret = info.signature_secret
  81. signature = sign(unsigned_tx_data, secret, self.ec)
  82. input.signature = signature
  83. for (input, info) in zip(tx.inputs, self.inputs):
  84. secret = info.signature_secret
  85. signature = sign(unsigned_tx_data, secret, self.ec)
  86. input.signature = signature
  87. return tx
  88. class MintProof:
  89. def __init__(self, value, token_id, value_blind, token_blind, serial,
  90. coin_blind, public, ec):
  91. self.value = value
  92. self.token_id = token_id
  93. self.value_blind = value_blind
  94. self.token_blind = token_blind
  95. self.serial = serial
  96. self.coin_blind = coin_blind
  97. self.public = public
  98. self.ec = ec
  99. def get_revealed(self):
  100. revealed = ClassNamespace()
  101. revealed.coin = ff_hash(
  102. self.ec.p,
  103. self.public[0],
  104. self.public[1],
  105. self.value,
  106. self.token_id,
  107. self.serial,
  108. self.coin_blind
  109. )
  110. revealed.value_commit = pedersen_encrypt(
  111. self.value, self.value_blind, self.ec
  112. )
  113. revealed.token_commit = pedersen_encrypt(
  114. self.token_id, self.token_blind, self.ec
  115. )
  116. return revealed
  117. def verify(self, public):
  118. revealed = self.get_revealed()
  119. return all([
  120. revealed.coin == public.coin,
  121. revealed.value_commit == public.value_commit,
  122. revealed.token_commit == public.token_commit
  123. ])
  124. def pedersen_encrypt(x, y, ec):
  125. vcv = ec.multiply(x, ec.G)
  126. vcr = ec.multiply(y, ec.H)
  127. return ec.add(vcv, vcr)
  128. def ff_hash(p, *args):
  129. hasher = hashlib.sha256()
  130. for arg in args:
  131. match arg:
  132. case int() as arg:
  133. hasher.update(arg.to_bytes(32, byteorder="little"))
  134. case bytes() as arg:
  135. hasher.update(arg)
  136. case _:
  137. raise Exception(f"unknown hash arg '{arg}' type: {type(arg)}")
  138. value = int.from_bytes(hasher.digest(), byteorder="little")
  139. return value % p
  140. def hash_point(point, message=None):
  141. hasher = hashlib.sha256()
  142. for x_i in point:
  143. hasher.update(x_i.to_bytes(32, byteorder="little"))
  144. # Optional message
  145. if message is not None:
  146. hasher.update(message)
  147. value = int.from_bytes(hasher.digest(), byteorder="little")
  148. return value
  149. def sign(message, secret, ec):
  150. ephem_secret = ec.random_scalar()
  151. ephem_public = ec.multiply(ephem_secret, ec.G)
  152. challenge = hash_point(ephem_public, message) % ec.order
  153. response = (ephem_secret + challenge * secret) % ec.order
  154. return ephem_public, response
  155. def verify(message, signature, public, ec):
  156. ephem_public, response = signature
  157. challenge = hash_point(ephem_public, message) % ec.order
  158. # sG
  159. lhs = ec.multiply(response, ec.G)
  160. # R + cP
  161. rhs_cP = ec.multiply(challenge, public)
  162. rhs = ec.add(ephem_public, rhs_cP)
  163. return lhs == rhs
  164. class Transaction:
  165. def __init__(self, ec):
  166. self.clear_inputs = []
  167. self.inputs = []
  168. self.outputs = []
  169. self.ec = ec
  170. def partial_encode(self):
  171. return b"hello"
  172. def verify(self):
  173. if not self._check_value_commits():
  174. return False, "value commits do not match"
  175. if not self._check_proofs():
  176. return False, "proofs failed to verify"
  177. if not self._verify_token_commitments():
  178. return False, "token ID mismatch"
  179. unsigned_tx_data = self.partial_encode()
  180. for input in self.clear_inputs:
  181. public = input.signature_public
  182. if not verify(unsigned_tx_data, input.signature, public, self.ec):
  183. return False
  184. for input in self.inputs:
  185. public = input.revealed.signature_public
  186. if not verify(unsigned_tx_data, input.signature, public, self.ec):
  187. return False
  188. return True, None
  189. def _check_value_commits(self):
  190. valcom_total = (0, 1, 0)
  191. for input in self.clear_inputs:
  192. value_commit = pedersen_encrypt(input.value, input.value_blind,
  193. self.ec)
  194. valcom_total = self.ec.add(valcom_total, value_commit)
  195. for input in self.inputs:
  196. value_commit = input.revealed.value_commit
  197. valcom_total = self.ec.add(valcom_total, value_commit)
  198. for output in self.outputs:
  199. v = output.revealed.value_commit
  200. value_commit = (v[0], -v[1], v[2])
  201. valcom_total = self.ec.add(valcom_total, value_commit)
  202. return valcom_total == (0, 1, 0)
  203. def _check_proofs(self):
  204. for input in self.inputs:
  205. if not input.burn_proof.verify(input.revealed):
  206. return False
  207. for output in self.outputs:
  208. if not output.mint_proof.verify(output.revealed):
  209. return False
  210. return True
  211. def _verify_token_commitments(self):
  212. assert len(self.outputs) > 0
  213. token_commit_value = self.outputs[0].revealed.token_commit
  214. for input in self.clear_inputs:
  215. token_commit = pedersen_encrypt(input.token_id, input.token_blind,
  216. self.ec)
  217. if token_commit != token_commit_value:
  218. return False
  219. for input in self.inputs:
  220. if input.revealed.token_commit != token_commit_value:
  221. return False
  222. for output in self.outputs:
  223. if output.revealed.token_commit != token_commit_value:
  224. return False
  225. return True
  226. def main(argv):
  227. ec = pallas_curve()
  228. secret = ec.random_scalar()
  229. public = ec.multiply(secret, ec.G)
  230. initial_supply = 21000
  231. token_id = 110
  232. signature_secret = ec.random_scalar()
  233. builder = TransactionBuilder(ec)
  234. builder.add_clear_input(initial_supply, token_id, signature_secret)
  235. builder.add_output(initial_supply, token_id, public)
  236. tx = builder.build()
  237. is_verify, reason = tx.verify()
  238. if not is_verify:
  239. print(f"tx verify failed: {reason}", file=sys.stderr)
  240. return -1
  241. return 0
  242. if __name__ == "__main__":
  243. sys.exit(main(sys.argv))