vm_export_rust.py 3.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102
  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, VariableRef};
  8. use bls12_381::Scalar;
  9. pub fn load_zkvm() -> ZKVirtualMachine {
  10. ZKVirtualMachine {
  11. constants: vec![""")
  12. constants = list(contract.constants.items())
  13. constants.sort(key=lambda obj: obj[1][0])
  14. constants = [(obj[0], obj[1][1]) for obj in constants]
  15. for symbol, value in constants:
  16. print("%s// %s" % (indent * 3, symbol))
  17. assert len(value) == 32*2
  18. chunk_str = lambda line, n: \
  19. [line[i:i + n] for i in range(0, len(line), n)]
  20. chunks = chunk_str(value, 2)
  21. # Reverse the endianness
  22. # We allow literal numbers but rust wants little endian
  23. chunks = chunks[::-1]
  24. print("%sScalar::from_bytes(&[" % (indent * 3))
  25. for i in range(0, 32, 4):
  26. print("%s0x%s, 0x%s, 0x%s, 0x%s," % (indent * 4,
  27. chunks[i], chunks[i + 1], chunks[i + 2], chunks[i + 3]))
  28. print("%s]).unwrap()," % (indent * 3))
  29. print("%s]," % (indent * 2))
  30. print("%salloc: vec![" % (indent * 2))
  31. for symbol, variable in contract.alloc.items():
  32. print("%s// %s" % (indent * 3, symbol))
  33. if variable.type.name == VariableType.PRIVATE.name:
  34. typestring = "Private"
  35. elif variable.type.name == VariableType.PUBLIC.name:
  36. typestring = "Public"
  37. else:
  38. assert False
  39. print("%s(AllocType::%s, %s)," % (indent * 3, typestring,
  40. variable.index))
  41. print("%s]," % (indent * 2))
  42. print("%sops: vec![" % (indent * 2))
  43. def var_ref_str(var_ref):
  44. if var_ref.type.name == VariableRefType.AUX.name:
  45. return "VariableRef::Aux(%s)" % var_ref.index
  46. elif var_ref.type.name == VariableRefType.LOCAL.name:
  47. return "VariableRef::Local(%s)" % var_ref.index
  48. else:
  49. assert False
  50. for op in contract.ops:
  51. print("%s// %s" % (indent * 3, op.line))
  52. args_part = ""
  53. if op.command == "load":
  54. assert len(op.args) == 2
  55. args_part = "(%s, %s)" % (var_ref_str(op.args[0]), op.args[1].index)
  56. elif op.args:
  57. args_part = ", ".join(var_ref_str(var_ref) for var_ref in op.args)
  58. args_part = "(%s)" % args_part
  59. print("%sCryptoOperation::%s%s," % (
  60. indent * 3,
  61. to_initial_caps(op.command),
  62. args_part
  63. ))
  64. print("%s]," % (indent * 2))
  65. print("%sconstraints: vec![" % (indent * 2))
  66. for constraint in contract.constraints:
  67. args_part = ""
  68. if constraint.args:
  69. print("%s// %s" % (indent *3, constraint.args_comment()))
  70. args = constraint.args[:]
  71. if (constraint.command == "lc0_add_coeff" or
  72. constraint.command == "lc1_add_coeff" or
  73. constraint.command == "lc2_add_coeff"):
  74. args[0] = args[0][0]
  75. args_part = ", ".join(str(index) for index in args)
  76. args_part = "(%s)" % args_part
  77. print("%sConstraintInstruction::%s%s," % (
  78. indent * 3,
  79. to_initial_caps(constraint.command),
  80. args_part
  81. ))
  82. print(r""" ],
  83. aux: vec![],
  84. params: None,
  85. verifying_key: None,
  86. }
  87. }""")