parser.py 21 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797
  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"""fn %s<CS>(
  199. mut cs: CS,
  200. %s
  201. ) -> Result<%s, SynthesisError>
  202. where
  203. CS: ConstraintSystem<bls12_381::Scalar>,
  204. {
  205. """ % (func_name, param_str, retstr)
  206. return header
  207. def as_expr(line, stack, consts, expr, code):
  208. var_from, type_to = expr.children
  209. if var_from not in stack:
  210. print("error: variable from not in stack frame:", var_from,
  211. file=sys.stderr)
  212. print("line:", line.text, "line:", line.lineno)
  213. return None
  214. type_from = stack[var_from]
  215. if type_from == "U64" and type_to == "Binary":
  216. code += "boolean::u64_into_boolean_vec_le(" + \
  217. "cs.namespace(|| \"" + line.text + "\"), " + var_from + \
  218. ")?;"
  219. elif type_from == "Scalar" and type_to == "Binary":
  220. code += "boolean::field_into_boolean_vec_le(" + \
  221. "cs.namespace(|| \"" + line.text + "\"), &" + var_from + \
  222. ")?;"
  223. else:
  224. print("error: unknown type conversion!", file=sys.stderr)
  225. print("line:", line.text, "line:", line.lineno)
  226. return None
  227. #print(var_from, type_from, type_to)
  228. return code, type_to
  229. def mul_expr(line, stack, consts, expr, code):
  230. var_a, var_b = expr.children
  231. #print("MUL", var_a, var_b)
  232. if var_b not in consts:
  233. print("error: unknown base!", file=sys.stderr)
  234. print("line:", line.text, "line:", line.lineno)
  235. return None
  236. base_type = consts[var_b]
  237. if base_type != "Point":
  238. print("error: unknown base type!", file=sys.stderr)
  239. print("line:", line.text, "line:", line.lineno)
  240. return None
  241. code += "ecc::fixed_base_multiplication(" + \
  242. "cs.namespace(|| \"" + line.text + "\"), &" + var_b + \
  243. ", &" + var_a + ")?;"
  244. return code, base_type
  245. def add_expr(line, stack, consts, expr, code):
  246. var_a, var_b = expr.children
  247. if var_a not in stack or var_b not in stack:
  248. print("error: missing stack item!", file=sys.stderr)
  249. print("line:", line.text, "line:", line.lineno)
  250. return None
  251. result_type = stack[var_a]
  252. if stack[var_b] != result_type:
  253. print("error: non matching items for addition!", file=sys.stderr)
  254. print("line:", line.text, "line:", line.lineno)
  255. return None
  256. code += var_a + ".add(cs.namespace(|| \"" + line.text \
  257. + "\"), &" + var_b + ")?;"
  258. return (code, result_type)
  259. def compile_let(line, stack, consts, statement):
  260. is_mutable = False
  261. if statement[0] == "mut":
  262. is_mutable = True
  263. statement = statement[1:]
  264. variable_name, variable_type = statement[0], statement[1]
  265. expr = statement[2]
  266. #print("LET", is_mutable, variable_name, variable_type)
  267. #print(" ", expr)
  268. code = "let " + ("mut " if is_mutable else "") + variable_name + " = "
  269. if expr.data == "as_expr":
  270. ceval = as_expr(line, stack, consts, expr, code)
  271. elif expr.data == "mul_expr":
  272. ceval = mul_expr(line, stack, consts, expr, code)
  273. elif expr.data == "add_expr":
  274. ceval = add_expr(line, stack, consts, expr, code)
  275. if ceval is None:
  276. return None
  277. code, type_to = ceval
  278. if variable_type != type_to:
  279. print("error: sub expr does not evaluate to correct type",
  280. file=sys.stderr)
  281. print("line:", line.text, "line:", line.lineno)
  282. return None
  283. stack[variable_name] = variable_type
  284. return code
  285. def interpret_func(func, consts):
  286. func_def = parse_func_def(func[0].text)
  287. header = compile_func_header(func_def)
  288. if header is None:
  289. return
  290. subroutine = header
  291. indent = " " * 4
  292. stack = dict(func_def[1])
  293. emitted_types = []
  294. for line in func[1:]:
  295. statement_type, statement = interpret_func_line(line.text, stack, consts)
  296. if statement_type == "let":
  297. code = compile_let(line, stack, consts, statement)
  298. if code is None:
  299. return
  300. subroutine += indent + code + "\n"
  301. elif statement_type == "return":
  302. for var in statement:
  303. if var not in stack:
  304. print("error: missing variable in stack!", file=sys.stderr)
  305. print("line:", line.text, "line:", line.lineno)
  306. return None
  307. if len(statement) == 1:
  308. code = "Ok(" + statement[0] + ")"
  309. else:
  310. code = "Ok(" + ",".join(statement) + ")"
  311. subroutine += indent + code + "\n"
  312. elif statement_type == "emit":
  313. assert len(statement) == 1
  314. variable = statement[0]
  315. if variable not in stack:
  316. print("error: missing variable in stack!", file=sys.stderr)
  317. print("line:", line.text, "line:", line.lineno)
  318. return None
  319. variable_type = stack[variable]
  320. if variable_type == "Point":
  321. code = variable + ".inputize(cs.namespace(|| \"" + \
  322. line.text + "\"))?;"
  323. else:
  324. print("error: unable to inputize type!", file=sys.stderr)
  325. print("line:", line.text, "line:", line.lineno)
  326. return None
  327. emitted_types.append(variable_type)
  328. subroutine += indent + code + "\n"
  329. subroutine += "}\n\n"
  330. return subroutine, emitted_types, func_def
  331. class CodeLineTransformer(lark.Transformer):
  332. def variable_name(self, name):
  333. return str(name[0])
  334. def let_statement(self, obj):
  335. return ("let", obj)
  336. def return_statement(self, obj):
  337. return ("return", obj)
  338. def emit_statement(self, obj):
  339. return ("emit", obj)
  340. def point(self, _):
  341. return "Point"
  342. def scalar(self, _):
  343. return "Scalar"
  344. def binary(self, _):
  345. return "Binary"
  346. def u64(self, _):
  347. return "U64"
  348. def type(self, typename):
  349. return str(typename[0])
  350. def mutable(self, _):
  351. return "mut"
  352. statement = list
  353. def interpret_func_line(text, stack, consts):
  354. parser = lark.Lark(r"""
  355. statement: let_statement
  356. | return_statement
  357. | emit_statement
  358. let_statement: "let" [mutable] variable_name ":" type "=" expr
  359. mutable: "mut"
  360. ?expr: as_expr
  361. | mul_expr
  362. | add_expr
  363. as_expr: variable_name "as" type
  364. mul_expr: variable_name "*" variable_name
  365. add_expr: variable_name "+" variable_name
  366. return_statement: "return" variable_name
  367. | "return" variable_tuple
  368. variable_tuple: "(" variable_name ("," variable_name)* ")"
  369. emit_statement: "emit" variable_name
  370. variable_name: NAME
  371. type: u64 | scalar | point | binary
  372. u64: "U64"
  373. scalar: "Scalar"
  374. point: "Point"
  375. binary: "Binary"
  376. %import common.CNAME -> NAME
  377. %import common.WS
  378. %ignore WS
  379. """, start="statement")
  380. tree = parser.parse(text)
  381. tokens = CodeLineTransformer().transform(tree)[0]
  382. return tokens
  383. class ContractDefTransformer(lark.Transformer):
  384. def contract_name(self, name):
  385. return str(name[0])
  386. def param(self, obj):
  387. return tuple(obj)
  388. def param_name(self, name):
  389. return str(name[0])
  390. def u64(self, _):
  391. return "U64"
  392. def scalar(self, _):
  393. return "Scalar"
  394. def point(self, _):
  395. return "Point"
  396. def binary(self, _):
  397. return "Binary"
  398. def type(self, obj):
  399. return obj[0]
  400. contract_def = list
  401. params = list
  402. type_list = list
  403. def parse_contract_def(text):
  404. parser = lark.Lark(r"""
  405. contract_def: "contract" contract_name "(" params+ ")" "->" type_list ":"
  406. contract_name: NAME
  407. params: param ("," param)*
  408. type_list: type
  409. | "(" type ("," type)* ")"
  410. param: param_name ":" type
  411. param_name: NAME
  412. type: u64 | scalar | point | binary
  413. u64: "U64"
  414. scalar: "Scalar"
  415. point: "Point"
  416. binary: "Binary"
  417. %import common.CNAME -> NAME
  418. %import common.WS
  419. %ignore WS
  420. """, start="contract_def")
  421. tree = parser.parse(text)
  422. tokens = ContractDefTransformer().transform(tree)
  423. assert len(tokens) == 3
  424. return tokens
  425. class ContractCodeLineTransformer(lark.Transformer):
  426. def variable_name(self, name):
  427. return str(name[0])
  428. def let_statement(self, obj):
  429. return ("let", obj)
  430. def return_statement(self, obj):
  431. return ("return", obj)
  432. def emit_statement(self, obj):
  433. return ("emit", obj)
  434. def point(self, _):
  435. return "Point"
  436. def scalar(self, _):
  437. return "Scalar"
  438. def binary(self, _):
  439. return "Binary"
  440. def u64(self, _):
  441. return "U64"
  442. def type(self, typename):
  443. return str(typename[0])
  444. def mutable(self, _):
  445. return "mut"
  446. def function_name(self, name):
  447. return str(name[0])
  448. statement = list
  449. variable_assign = list
  450. variable_decl = tuple
  451. def interpret_contract_line(text, stack, consts):
  452. parser = lark.Lark(r"""
  453. statement: let_statement
  454. | return_statement
  455. | emit_statement
  456. let_statement: "let" variable_assign "=" expr
  457. mutable: "mut"
  458. variable_assign: variable_decl
  459. | "(" variable_decl ("," variable_decl)* ")"
  460. variable_decl: [mutable] variable_name ":" type
  461. ?expr: as_expr
  462. | mul_expr
  463. | add_expr
  464. | funccall_expr
  465. as_expr: variable_name "as" type
  466. mul_expr: variable_name "*" variable_name
  467. add_expr: variable_name "+" variable_name
  468. funccall_expr: function_name "(" [variable_name ("," variable_name)*] ")"
  469. return_statement: "return" variable_name
  470. | "return" variable_tuple
  471. variable_tuple: "(" variable_name ("," variable_name)* ")"
  472. emit_statement: "emit" variable_name
  473. variable_name: NAME
  474. function_name: NAME
  475. type: u64 | scalar | point | binary
  476. u64: "U64"
  477. scalar: "Scalar"
  478. point: "Point"
  479. binary: "Binary"
  480. %import common.CNAME -> NAME
  481. %import common.WS
  482. %ignore WS
  483. """, start="statement")
  484. tree = parser.parse(text)
  485. tokens = ContractCodeLineTransformer().transform(tree)[0]
  486. return tokens
  487. def to_initial_caps(snake_str):
  488. components = snake_str.split("_")
  489. return "".join(x.title() for x in components)
  490. def create_contract_header(contract_def):
  491. contract_name, params, retvals = contract_def
  492. contract_name = to_initial_caps(contract_name)
  493. header = "pub struct %s {\n" % contract_name
  494. for param_name, param_type in params:
  495. header += " " * 4 + "pub %s: Option<%s>,\n" % (param_name, param_type)
  496. header += "}\n\n"
  497. header += r"""impl Circuit<bls12_381::Scalar> for %s {
  498. fn synthesize<CS: ConstraintSystem<bls12_381::Scalar>>(
  499. self,
  500. cs: &mut CS,
  501. ) -> Result<(), SynthesisError> {
  502. """ % contract_name
  503. return header
  504. # Worst code ever
  505. def compile_let2(line, stack, consts, funcs, statement):
  506. lhs = []
  507. for variable_decl in statement[0]:
  508. assert len(variable_decl) == 2 or \
  509. (len(variable_decl) == 3 and variable_decl[0] == "mut")
  510. if len(variable_decl) == 2:
  511. mutable = False
  512. elif len(variable_decl) == 3:
  513. assert variable_decl[0] == "mut"
  514. mutable = True
  515. variable_decl = variable_decl[1:]
  516. #else:
  517. # Error!
  518. lhs.append(list(variable_decl) + [mutable])
  519. variable_types = []
  520. code = "let "
  521. if len(lhs) == 1:
  522. name, type, is_mutable = lhs[0]
  523. variable_types.append(type)
  524. code += ("mut " if is_mutable else "") + name
  525. else:
  526. code += "("
  527. start = True
  528. for name, type, is_mutable in lhs:
  529. if not start:
  530. code += ", "
  531. start = False
  532. code += name
  533. variable_types.append(type)
  534. code += ")"
  535. code += " = "
  536. expr = statement[1]
  537. expr_type = expr.data
  538. expr = expr.children
  539. if expr_type == "funccall_expr":
  540. ceval = funccall_expr(line, stack, consts, funcs, expr, code)
  541. #code = "let " + ("mut " if is_mutable else "") + variable_name + " = "
  542. if ceval is None:
  543. return None
  544. code, types_to = ceval
  545. if variable_types != types_to:
  546. print("error: sub expr does not evaluate to correct type",
  547. file=sys.stderr)
  548. print("line:", line.text, "line:", line.lineno)
  549. return None
  550. for name, type, _ in lhs:
  551. stack[name] = type
  552. return code
  553. def funccall_expr(line, stack, consts, funcs, expr, code):
  554. func_name, arguments = expr[0], expr[1:]
  555. if func_name not in funcs:
  556. print("error: non-existant function call",
  557. file=sys.stderr)
  558. print("line:", line.text, "line:", line.lineno)
  559. return None
  560. code += "%s(cs.namespace(|| \"%s\"), %s)?;" % (
  561. func_name, line.text, ", ".join(arguments))
  562. return_type = funcs[func_name][-1][-1]
  563. return code, return_type
  564. def interpret_contract(contract, consts, funcs):
  565. contract_def = parse_contract_def(contract[0].text)
  566. contract_code = create_contract_header(contract_def)
  567. stack = dict(contract_def[1])
  568. for line in contract[1:2]:
  569. indent = " " * 4 * int(line.level + 1)
  570. statement_type, statement = interpret_contract_line(line.text, stack, consts)
  571. #pprint.pprint(statement_type)
  572. if statement_type == "let":
  573. code = compile_let2(line, stack, consts, funcs, statement)
  574. if code is None:
  575. return
  576. contract_code += indent + code + "\n"
  577. contract_code += " " * 8 + "Ok(())\n"
  578. contract_code += " " * 4 + "}\n"
  579. contract_code += "}\n\n"
  580. #print("-------------------------------")
  581. #print(contract_code)
  582. return contract_code, contract_def
  583. def main(argv):
  584. if len(argv) == 1:
  585. print("error: missing proof file", file=sys.stderr)
  586. return -1
  587. filename = sys.argv[1]
  588. text = open(filename, "r").read()
  589. if (linedescs := parse(text)) is None:
  590. return -1
  591. sections = section(linedescs)
  592. consts, funcs, contracts = classify(sections)
  593. consts = read_consts(consts)
  594. compiled_funcs = {}
  595. for func in funcs:
  596. if (compiled := interpret_func(func, consts)) is None:
  597. return -1
  598. _, _, func_def = compiled
  599. func_name, _, _ = func_def
  600. compiled_funcs[func_name] = compiled
  601. funcs = compiled_funcs
  602. compiled_contracts = {}
  603. for contract in contracts[1:]:
  604. if (compiled := interpret_contract(contract, consts, funcs)) is None:
  605. return -1
  606. #print(contract)
  607. _, contract_def = compiled
  608. contract_name, _, _ = contract_def
  609. compiled_contracts[contract_name] = compiled
  610. contracts = compiled_contracts
  611. # Concat
  612. output = ""
  613. for _, func in funcs.items():
  614. output += func[0]
  615. for _, contract in contracts.items():
  616. output += contract[0]
  617. print(output)
  618. if __name__ == "__main__":
  619. main(sys.argv)