parser.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511
  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 compile_func_header(func_def):
  167. func_name, params, retvals = func_def
  168. #print("Function:", func_name)
  169. #print("Params:", params)
  170. #print("Return values:", retvals)
  171. #print()
  172. param_str = ""
  173. for param, type in params:
  174. if param_str:
  175. param_str += ", "
  176. param_str += param + ": "
  177. if type == "U64":
  178. param_str += "u64"
  179. elif type == "Scalar":
  180. param_str += "&jubjub::Fr"
  181. else:
  182. print("error: unsupported param type", file=sys.stderr)
  183. print("line:", line.text, "line:", line.lineno)
  184. return None
  185. converted_retvals = []
  186. for type in retvals:
  187. if type == "Binary":
  188. converted_retvals.append("boolean::Boolean")
  189. else:
  190. print("error: unsupported return type", file=sys.stderr)
  191. print("line:", line.text, "line:", line.lineno)
  192. return None
  193. retvals = converted_retvals
  194. if len(retvals) == 1:
  195. retstr = retvals[0]
  196. else:
  197. retstr = "(" + ", ".join(retvals) + ")"
  198. header = r"""
  199. fn %s<CS>(
  200. mut cs: CS,
  201. %s
  202. ) -> Result<%s, SynthesisError>
  203. where
  204. CS: ConstraintSystem<bls12_381::Scalar>,
  205. {
  206. """ % (func_name, param_str, retstr)
  207. return header
  208. def as_expr(line, stack, consts, expr, code):
  209. var_from, type_to = expr.children
  210. if var_from not in stack:
  211. print("error: variable from not in stack frame:", var_from,
  212. file=sys.stderr)
  213. print("line:", line.text, "line:", line.lineno)
  214. return None
  215. type_from = stack[var_from]
  216. if type_from == "U64" and type_to == "Binary":
  217. code += "boolean::u64_into_boolean_vec_le(" + \
  218. "cs.namespace(|| \"" + line.text + "\"), " + var_from + \
  219. ")?;"
  220. elif type_from == "Scalar" and type_to == "Binary":
  221. code += "boolean::field_into_boolean_vec_le(" + \
  222. "cs.namespace(|| \"" + line.text + "\"), &" + var_from + \
  223. ")?;"
  224. else:
  225. print("error: unknown type conversion!", file=sys.stderr)
  226. print("line:", line.text, "line:", line.lineno)
  227. return None
  228. #print(var_from, type_from, type_to)
  229. return code, type_to
  230. def mul_expr(line, stack, consts, expr, code):
  231. var_a, var_b = expr.children
  232. #print("MUL", var_a, var_b)
  233. if var_b not in consts:
  234. print("error: unknown base!", file=sys.stderr)
  235. print("line:", line.text, "line:", line.lineno)
  236. return None
  237. base_type = consts[var_b]
  238. if base_type != "Point":
  239. print("error: unknown base type!", file=sys.stderr)
  240. print("line:", line.text, "line:", line.lineno)
  241. return None
  242. code += "ecc::fixed_base_multiplication(" + \
  243. "cs.namespace(|| \"" + line.text + "\"), &" + var_b + \
  244. ", &" + var_a + ")?;"
  245. return code, base_type
  246. def add_expr(line, stack, consts, expr, code):
  247. var_a, var_b = expr.children
  248. if var_a not in stack or var_b not in stack:
  249. print("error: missing stack item!", file=sys.stderr)
  250. print("line:", line.text, "line:", line.lineno)
  251. return None
  252. result_type = stack[var_a]
  253. if stack[var_b] != result_type:
  254. print("error: non matching items for addition!", file=sys.stderr)
  255. print("line:", line.text, "line:", line.lineno)
  256. return None
  257. code += var_a + ".add(cs.namespace(|| \"" + line.text \
  258. + "\"), &" + var_b + ")?;"
  259. return (code, result_type)
  260. def compile_let(line, stack, consts, statement):
  261. is_mutable = False
  262. if statement[0] == "mut":
  263. is_mutable = True
  264. statement = statement[1:]
  265. variable_name, variable_type = statement[0], statement[1]
  266. expr = statement[2]
  267. #print("LET", is_mutable, variable_name, variable_type)
  268. #print(" ", expr)
  269. code = "let " + ("mut " if is_mutable else "") + variable_name + " = "
  270. if expr.data == "as_expr":
  271. ceval = as_expr(line, stack, consts, expr, code)
  272. elif expr.data == "mul_expr":
  273. ceval = mul_expr(line, stack, consts, expr, code)
  274. elif expr.data == "add_expr":
  275. ceval = add_expr(line, stack, consts, expr, code)
  276. if ceval is None:
  277. return None
  278. code, type_to = ceval
  279. if variable_type != type_to:
  280. print("error: sub expr does not evaluate to correct type",
  281. file=sys.stderr)
  282. print("line:", line.text, "line:", line.lineno)
  283. return None
  284. stack[variable_name] = variable_type
  285. return code
  286. def interpret_func(func, consts):
  287. func_def = parse_func_def(func[0].text)
  288. header = compile_func_header(func_def)
  289. if header is None:
  290. return
  291. subroutine = header
  292. indent = " " * 4
  293. stack = dict(func_def[1])
  294. emitted_types = []
  295. for line in func[1:]:
  296. statement_type, statement = interpret_func_line(line.text, stack, consts)
  297. if statement_type == "let":
  298. code = compile_let(line, stack, consts, statement)
  299. if code is None:
  300. return
  301. subroutine += indent + code + "\n"
  302. elif statement_type == "return":
  303. for var in statement:
  304. if var not in stack:
  305. print("error: missing variable in stack!", file=sys.stderr)
  306. print("line:", line.text, "line:", line.lineno)
  307. return None
  308. if len(statement) == 1:
  309. code = "Ok(" + statement[0] + ")"
  310. else:
  311. code = "Ok(" + ",".join(statement) + ")"
  312. subroutine += indent + code + "\n"
  313. elif statement_type == "emit":
  314. assert len(statement) == 1
  315. variable = statement[0]
  316. if variable not in stack:
  317. print("error: missing variable in stack!", file=sys.stderr)
  318. print("line:", line.text, "line:", line.lineno)
  319. return None
  320. variable_type = stack[variable]
  321. if variable_type == "Point":
  322. code = variable + ".inputize(cs.namespace(|| \"" + \
  323. line.text + "\"))?;"
  324. else:
  325. print("error: unable to inputize type!", file=sys.stderr)
  326. print("line:", line.text, "line:", line.lineno)
  327. return None
  328. emitted_types.append(variable_type)
  329. subroutine += indent + code + "\n"
  330. subroutine += "}"
  331. print(subroutine)
  332. class CodeLineTransformer(lark.Transformer):
  333. def variable_name(self, name):
  334. return str(name[0])
  335. def let_statement(self, obj):
  336. return ("let", obj)
  337. def return_statement(self, obj):
  338. return ("return", obj)
  339. def emit_statement(self, obj):
  340. return ("emit", obj)
  341. def point(self, _):
  342. return "Point"
  343. def scalar(self, _):
  344. return "Scalar"
  345. def binary(self, _):
  346. return "Binary"
  347. def u64(self, _):
  348. return "U64"
  349. def type(self, typename):
  350. return str(typename[0])
  351. def mutable(self, _):
  352. return "mut"
  353. statement = list
  354. def interpret_func_line(text, stack, consts):
  355. parser = lark.Lark(r"""
  356. statement: let_statement
  357. | return_statement
  358. | emit_statement
  359. let_statement: "let" [mutable] variable_name ":" type "=" expr
  360. mutable: "mut"
  361. ?expr: as_expr
  362. | mul_expr
  363. | add_expr
  364. as_expr: variable_name "as" type
  365. mul_expr: variable_name "*" variable_name
  366. add_expr: variable_name "+" variable_name
  367. return_statement: "return" variable_name
  368. | "return" variable_tuple
  369. variable_tuple: "(" variable_name ("," variable_name)* ")"
  370. emit_statement: "emit" variable_name
  371. variable_name: NAME
  372. type: u64 | scalar | point | binary
  373. u64: "U64"
  374. scalar: "Scalar"
  375. point: "Point"
  376. binary: "Binary"
  377. %import common.CNAME -> NAME
  378. %import common.WS
  379. %ignore WS
  380. """, start="statement")
  381. tree = parser.parse(text)
  382. tokens = CodeLineTransformer().transform(tree)[0]
  383. return tokens
  384. def main(argv):
  385. if len(argv) == 1:
  386. print("error: missing proof file", file=sys.stderr)
  387. return -1
  388. filename = sys.argv[1]
  389. text = open(filename, "r").read()
  390. if (linedescs := parse(text)) is None:
  391. return -1
  392. sections = section(linedescs)
  393. consts, funcs, contracts = classify(sections)
  394. consts = read_consts(consts)
  395. for func in funcs:
  396. interpret_func(func, consts)
  397. if __name__ == "__main__":
  398. main(sys.argv)