parser.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476
  1. import lark
  2. import pprint
  3. import re
  4. import sys
  5. class LineDesc:
  6. def __init__(self, level, text, lineno):
  7. self.level = level
  8. self.text = text
  9. self.lineno = lineno
  10. assert self.text[0] != ' '
  11. def __repr__(self):
  12. return "<%s:'%s'>" % (self.level, self.text)
  13. def clean_line(line):
  14. lead_spaces = len(line) - len(line.lstrip(" "))
  15. level = lead_spaces / 4
  16. # Remove leading spaces
  17. if line.strip(" ") == "":
  18. return None
  19. line = line.lstrip(" ")
  20. # Remove all comments
  21. line = re.sub('#.*$', '', line).strip()
  22. if not line:
  23. return None
  24. return level, line
  25. def parse(text):
  26. lines = text.split("\n")
  27. linedescs = []
  28. # These are to join open parenthesis
  29. current_line = ""
  30. paren_level = 0
  31. for lineno, line in enumerate(lines):
  32. if (lineinfo := clean_line(line)) is None:
  33. continue
  34. level, line = lineinfo
  35. for c in line:
  36. if c == "(":
  37. paren_level += 1
  38. elif c == ")":
  39. paren_level -= 1
  40. #print(level, paren_level, current_line)
  41. if paren_level < 0:
  42. print("error: too many closing paren )", file=sys.stderr)
  43. print("line:", lineno)
  44. return
  45. if current_line:
  46. current_line += " " + line
  47. else:
  48. current_line = line
  49. if paren_level > 0:
  50. continue
  51. #print(level, current_line)
  52. ldesc = LineDesc(level, current_line, lineno)
  53. linedescs.append(ldesc)
  54. current_line = ""
  55. if paren_level > 0:
  56. print("error: missing closing paren )", file=sys.stderr)
  57. return None
  58. return linedescs
  59. def section(linedescs):
  60. sections = []
  61. current_section = None
  62. for desc in linedescs:
  63. if desc.level == 0:
  64. if current_section:
  65. sections.append(current_section)
  66. current_section = [desc]
  67. continue
  68. current_section.append(desc)
  69. sections.append(current_section)
  70. return sections
  71. def classify(sections):
  72. consts = []
  73. funcs = []
  74. contracts = []
  75. for section in sections:
  76. assert len(section)
  77. if section[0].text == "const:":
  78. consts.append(section)
  79. elif section[0].text.startswith("def"):
  80. funcs.append(section)
  81. elif section[0].text.startswith("contract"):
  82. contracts.append(section)
  83. return consts, funcs, contracts
  84. def tokenize_const(text):
  85. parser = lark.Lark(r"""
  86. value_map: name ":" type_def
  87. name: NAME
  88. ?type_def: point
  89. | blake2s_personalization
  90. | pedersen_personalization
  91. | list
  92. point: "Point"
  93. blake2s_personalization: "Blake2sPersonalization"
  94. pedersen_personalization: "PedersenPersonalization"
  95. list: "list<" type_def ">"
  96. %import common.CNAME -> NAME
  97. %import common.WS
  98. %ignore WS
  99. """, start="value_map")
  100. return parser.parse(text)
  101. class ConstTransformer(lark.Transformer):
  102. def name(self, name):
  103. return str(name[0])
  104. def point(self, _):
  105. return "Point"
  106. def blake2s_personalization(self, _):
  107. return "Blake2sPersonalization"
  108. def pedersen_personalization(self, _):
  109. return "PedersenPersonalization"
  110. value_map = tuple
  111. list = list
  112. def read_consts(consts):
  113. consts_map = {}
  114. for subsection in consts:
  115. assert subsection[0].text == "const:"
  116. for ldesc in subsection[1:]:
  117. tree = tokenize_const(ldesc.text)
  118. tokens = ConstTransformer().transform(tree)
  119. #print(tokens)
  120. name, typedesc = tokens
  121. consts_map[name] = typedesc
  122. #pprint.pprint(consts_map)
  123. return consts_map
  124. class FuncDefTransformer(lark.Transformer):
  125. def func_name(self, name):
  126. return str(name[0])
  127. def param(self, obj):
  128. return tuple(obj)
  129. def param_name(self, name):
  130. return str(name[0])
  131. def u64(self, _):
  132. return "U64"
  133. def scalar(self, _):
  134. return "Scalar"
  135. def point(self, _):
  136. return "Point"
  137. def binary(self, _):
  138. return "Binary"
  139. def type(self, obj):
  140. return obj[0]
  141. func_def = list
  142. params = list
  143. type_list = list
  144. def parse_func_def(text):
  145. parser = lark.Lark(r"""
  146. func_def: "def" func_name "(" params+ ")" "->" type_list ":"
  147. func_name: NAME
  148. params: param ("," param)*
  149. type_list: type
  150. | "(" type ("," type)* ")"
  151. param: param_name ":" type
  152. param_name: NAME
  153. type: u64 | scalar | point | binary
  154. u64: "U64"
  155. scalar: "Scalar"
  156. point: "Point"
  157. binary: "Binary"
  158. %import common.CNAME -> NAME
  159. %import common.WS
  160. %ignore WS
  161. """, start="func_def")
  162. tree = parser.parse(text)
  163. tokens = FuncDefTransformer().transform(tree)
  164. assert len(tokens) == 3
  165. return tokens
  166. def interpret_func(func, consts):
  167. func_def = parse_func_def(func[0].text)
  168. func_name, params, retvals = func_def
  169. #print("Function:", func_name)
  170. #print("Params:", params)
  171. #print("Return values:", retvals)
  172. #print()
  173. param_str = ""
  174. for param, type in params:
  175. if param_str:
  176. param_str += ", "
  177. param_str += param + ": "
  178. if type == "U64":
  179. param_str += "u64"
  180. elif type == "Scalar":
  181. param_str += "&jubjub::Fr"
  182. else:
  183. print("error: unsupported param type", file=sys.stderr)
  184. print("line:", line.text, "line:", line.lineno)
  185. return None
  186. converted_retvals = []
  187. for type in retvals:
  188. if type == "Binary":
  189. converted_retvals.append("boolean::Boolean")
  190. else:
  191. print("error: unsupported return type", file=sys.stderr)
  192. print("line:", line.text, "line:", line.lineno)
  193. return None
  194. retvals = converted_retvals
  195. if len(retvals) == 1:
  196. retstr = retvals[0]
  197. else:
  198. retstr = "(" + ", ".join(retvals) + ")"
  199. subroutine = r"""
  200. fn %s<CS>(
  201. mut cs: CS,
  202. %s
  203. ) -> Result<%s, SynthesisError>
  204. where
  205. CS: ConstraintSystem<bls12_381::Scalar>,
  206. {
  207. """ % (func_name, param_str, retstr)
  208. indent = " " * 4
  209. stack = dict(params)
  210. emitted_types = []
  211. for line in func[1:]:
  212. statement_type, statement = interpret_func_line(line.text, stack, consts)
  213. if statement_type == "let":
  214. is_mutable = False
  215. if statement[0] == "mut":
  216. is_mutable = True
  217. statement = statement[1:]
  218. variable_name, variable_type = statement[0], statement[1]
  219. expr = statement[2]
  220. #print("LET", is_mutable, variable_name, variable_type)
  221. #print(" ", expr)
  222. code = "let " + ("mut " if is_mutable else "") + variable_name + " = "
  223. if expr.data == "as_expr":
  224. var_from, type_to = expr.children
  225. if var_from not in stack:
  226. print("error: variable from not in stack frame:", var_from,
  227. file=sys.stderr)
  228. print("line:", line.text, "line:", line.lineno)
  229. return None
  230. type_from = stack[var_from]
  231. if type_from == "U64" and type_to == "Binary":
  232. code += "boolean::u64_into_boolean_vec_le(" + \
  233. "cs.namespace(|| \"" + line.text + "\"), " + var_from + \
  234. ")?;"
  235. elif type_from == "Scalar" and type_to == "Binary":
  236. code += "boolean::field_into_boolean_vec_le(" + \
  237. "cs.namespace(|| \"" + line.text + "\"), &" + var_from + \
  238. ")?;"
  239. else:
  240. print("error: unknown type conversion!", file=sys.stderr)
  241. print("line:", line.text, "line:", line.lineno)
  242. return None
  243. #print(var_from, type_from, type_to)
  244. stack[variable_name] = type_to
  245. elif expr.data == "mul_expr":
  246. var_a, var_b = expr.children
  247. #print("MUL", var_a, var_b)
  248. if var_b not in consts:
  249. print("error: unknown base!", file=sys.stderr)
  250. print("line:", line.text, "line:", line.lineno)
  251. return None
  252. base_type = consts[var_b]
  253. if base_type != "Point":
  254. print("error: unknown base type!", file=sys.stderr)
  255. print("line:", line.text, "line:", line.lineno)
  256. return None
  257. code += "ecc::fixed_base_multiplication(" + \
  258. "cs.namespace(|| \"" + line.text + "\"), &" + var_b + \
  259. ", &" + var_a + ")?;"
  260. stack[variable_name] = "Point"
  261. elif expr.data == "add_expr":
  262. var_a, var_b = expr.children
  263. if var_a not in stack or var_b not in stack:
  264. print("error: missing stack item!", file=sys.stderr)
  265. print("line:", line.text, "line:", line.lineno)
  266. return None
  267. result_type = stack[var_a]
  268. if stack[var_b] != result_type:
  269. print("error: non matching items for addition!", file=sys.stderr)
  270. print("line:", line.text, "line:", line.lineno)
  271. return None
  272. code += var_a + ".add(cs.namespace(|| \"" + line.text \
  273. + "\"), &" + var_b + ")?;"
  274. stack[variable_name] = result_type
  275. subroutine += indent + code + "\n"
  276. elif statement_type == "return":
  277. for var in statement:
  278. if var not in stack:
  279. print("error: missing variable in stack!", file=sys.stderr)
  280. print("line:", line.text, "line:", line.lineno)
  281. return None
  282. if len(statement) == 1:
  283. code = "Ok(" + statement[0] + ")"
  284. else:
  285. code = "Ok(" + ",".join(statement) + ")"
  286. subroutine += indent + code + "\n"
  287. elif statement_type == "emit":
  288. assert len(statement) == 1
  289. variable = statement[0]
  290. if variable not in stack:
  291. print("error: missing variable in stack!", file=sys.stderr)
  292. print("line:", line.text, "line:", line.lineno)
  293. return None
  294. variable_type = stack[variable]
  295. if variable_type == "Point":
  296. code = variable + ".inputize(cs.namespace(|| \"" + \
  297. line.text + "\"))?;"
  298. else:
  299. print("error: unable to inputize type!", file=sys.stderr)
  300. print("line:", line.text, "line:", line.lineno)
  301. return None
  302. emitted_types.append(variable_type)
  303. subroutine += indent + code + "\n"
  304. subroutine += "}"
  305. print(subroutine)
  306. class CodeLineTransformer(lark.Transformer):
  307. def variable_name(self, name):
  308. return str(name[0])
  309. def let_statement(self, obj):
  310. return ("let", obj)
  311. def return_statement(self, obj):
  312. return ("return", obj)
  313. def emit_statement(self, obj):
  314. return ("emit", obj)
  315. def point(self, _):
  316. return "Point"
  317. def scalar(self, _):
  318. return "Scalar"
  319. def binary(self, _):
  320. return "Binary"
  321. def u64(self, _):
  322. return "U64"
  323. def type(self, typename):
  324. return str(typename[0])
  325. def mutable(self, _):
  326. return "mut"
  327. statement = list
  328. def interpret_func_line(text, stack, consts):
  329. parser = lark.Lark(r"""
  330. statement: let_statement
  331. | return_statement
  332. | emit_statement
  333. let_statement: "let" [mutable] variable_name ":" type "=" expr
  334. mutable: "mut"
  335. ?expr: as_expr
  336. | mul_expr
  337. | add_expr
  338. as_expr: variable_name "as" type
  339. mul_expr: variable_name "*" variable_name
  340. add_expr: variable_name "+" variable_name
  341. return_statement: "return" variable_name
  342. | "return" variable_tuple
  343. variable_tuple: "(" variable_name ("," variable_name)* ")"
  344. emit_statement: "emit" variable_name
  345. variable_name: NAME
  346. type: u64 | scalar | point | binary
  347. u64: "U64"
  348. scalar: "Scalar"
  349. point: "Point"
  350. binary: "Binary"
  351. %import common.CNAME -> NAME
  352. %import common.WS
  353. %ignore WS
  354. """, start="statement")
  355. tree = parser.parse(text)
  356. tokens = CodeLineTransformer().transform(tree)[0]
  357. return tokens
  358. def main(argv):
  359. if len(argv) == 1:
  360. print("error: missing proof file", file=sys.stderr)
  361. return -1
  362. filename = sys.argv[1]
  363. text = open(filename, "r").read()
  364. if (linedescs := parse(text)) is None:
  365. return -1
  366. sections = section(linedescs)
  367. consts, funcs, contracts = classify(sections)
  368. consts = read_consts(consts)
  369. for func in funcs:
  370. interpret_func(func, consts)
  371. if __name__ == "__main__":
  372. main(sys.argv)