compile_export_rust.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122
  1. from vm import VariableType, VariableRefType
  2. def to_initial_caps(snake_str):
  3. components = snake_str.split("_")
  4. return "".join(x.title() for x in components)
  5. def display(contract):
  6. indent = " " * 4
  7. print(r"""use super::vm::{ZKVirtualMachine, CryptoOperation, AllocType, ConstraintInstruction, VariableIndex, VariableRef};
  8. use bls12_381::Scalar;
  9. pub fn load_params(params: Vec<Scalar>) -> Vec<(VariableIndex, Scalar)> {""")
  10. params = [(symbol, var) for symbol, var in contract.alloc.items() if var.is_param]
  11. print("%sassert_eq!(params.len(), %s);" % (indent, len(params)))
  12. print("%slet mut result = vec![(0, Scalar::zero()); %s];" % (
  13. indent, len(params)))
  14. for i, (symbol, variable) in enumerate(params):
  15. assert variable.is_param
  16. print("%s// %s" % (indent, symbol))
  17. print("%sresult[%s] = (%s, params[%s]);" % (
  18. indent, i, variable.index, i))
  19. print("%sresult" % indent)
  20. print("}\n")
  21. print(r"""pub fn load_zkvm() -> ZKVirtualMachine {
  22. ZKVirtualMachine {
  23. constants: vec![""")
  24. constants = list(contract.constants.items())
  25. constants.sort(key=lambda obj: obj[1][0])
  26. constants = [(obj[0], obj[1][1]) for obj in constants]
  27. for symbol, value in constants:
  28. print("%s// %s" % (indent * 3, symbol))
  29. assert len(value) == 32*2
  30. chunk_str = lambda line, n: \
  31. [line[i:i + n] for i in range(0, len(line), n)]
  32. chunks = chunk_str(value, 2)
  33. # Reverse the endianness
  34. # We allow literal numbers but rust wants little endian
  35. chunks = chunks[::-1]
  36. print("%sScalar::from_bytes(&[" % (indent * 3))
  37. for i in range(0, 32, 4):
  38. print("%s0x%s, 0x%s, 0x%s, 0x%s," % (indent * 4,
  39. chunks[i], chunks[i + 1], chunks[i + 2], chunks[i + 3]))
  40. print("%s]).unwrap()," % (indent * 3))
  41. print("%s]," % (indent * 2))
  42. print("%salloc: vec![" % (indent * 2))
  43. for symbol, variable in contract.alloc.items():
  44. print("%s// %s" % (indent * 3, symbol))
  45. if variable.type.name == VariableType.PRIVATE.name:
  46. typestring = "Private"
  47. elif variable.type.name == VariableType.PUBLIC.name:
  48. typestring = "Public"
  49. else:
  50. assert False
  51. print("%s(AllocType::%s, %s)," % (indent * 3, typestring,
  52. variable.index))
  53. print("%s]," % (indent * 2))
  54. print("%sops: vec![" % (indent * 2))
  55. def var_ref_str(var_ref):
  56. if var_ref.type.name == VariableRefType.AUX.name:
  57. return "VariableRef::Aux(%s)" % var_ref.index
  58. elif var_ref.type.name == VariableRefType.LOCAL.name:
  59. return "VariableRef::Local(%s)" % var_ref.index
  60. else:
  61. assert False
  62. for op in contract.ops:
  63. print("%s// %s" % (indent * 3, op.line))
  64. args_part = ""
  65. if op.command == "load":
  66. assert len(op.args) == 2
  67. args_part = "(%s, %s)" % (var_ref_str(op.args[0]), op.args[1].index)
  68. elif op.command == "debug":
  69. assert len(op.args) == 1
  70. args_part = '(String::from("%s"), %s)' % (
  71. op.line, var_ref_str(op.args[0]))
  72. elif op.args:
  73. args_part = ", ".join(var_ref_str(var_ref) for var_ref in op.args)
  74. args_part = "(%s)" % args_part
  75. print("%sCryptoOperation::%s%s," % (
  76. indent * 3,
  77. to_initial_caps(op.command),
  78. args_part
  79. ))
  80. print("%s]," % (indent * 2))
  81. print("%sconstraints: vec![" % (indent * 2))
  82. for constraint in contract.constraints:
  83. args_part = ""
  84. if constraint.args:
  85. print("%s// %s" % (indent *3, constraint.args_comment()))
  86. args = constraint.args[:]
  87. if (constraint.command == "lc0_add_coeff" or
  88. constraint.command == "lc1_add_coeff" or
  89. constraint.command == "lc2_add_coeff" or
  90. constraint.command == "lc0_add_one_coeff" or
  91. constraint.command == "lc1_add_one_coeff" or
  92. constraint.command == "lc2_add_one_coeff"):
  93. args[0] = args[0][0]
  94. args_part = ", ".join(str(index) for index in args)
  95. args_part = "(%s)" % args_part
  96. print("%sConstraintInstruction::%s%s," % (
  97. indent * 3,
  98. to_initial_caps(constraint.command),
  99. args_part
  100. ))
  101. print(r""" ],
  102. aux: vec![],
  103. params: None,
  104. verifying_key: None,
  105. }
  106. }""")