parser.py 21 KB

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