Parser shaping up

Some tests passing - still crufty with debugging printouts
This commit is contained in:
Eli Bendersky 2015-01-28 06:24:43 -08:00
commit 8937d8d40a

View file

@ -76,7 +76,8 @@ class Lexer(object):
# AST hierarchy # AST hierarchy
class ASTNode(object): class ASTNode(object):
pass def dump(self, indent=0):
raise NotImplementedError
class ExprAST(ASTNode): class ExprAST(ASTNode):
@ -87,11 +88,19 @@ class NumberExprAST(ExprAST):
def __init__(self, val): def __init__(self, val):
self.val = val self.val = val
def dump(self, indent=0):
return '{0}{1}[{2}]'.format(
' ' * indent, self.__class__.__name__, self.val)
class VariableExprAST(ExprAST): class VariableExprAST(ExprAST):
def __init__(self, name): def __init__(self, name):
self.name = name self.name = name
def dump(self, indent=0):
return '{0}{1}[{2}]'.format(
' ' * indent, self.__class__.__name__, self.name)
class BinaryExprAST(ExprAST): class BinaryExprAST(ExprAST):
def __init__(self, op, lhs, rhs): def __init__(self, op, lhs, rhs):
@ -99,17 +108,35 @@ class BinaryExprAST(ExprAST):
self.lhs = lhs self.lhs = lhs
self.rhs = rhs self.rhs = rhs
def dump(self, indent=0):
s = '{0}{1}[{2}]\n'.format(
' ' * indent, self.__class__.__name__, self.op)
s += self.lhs.dump(indent + 2) + '\n'
s += self.rhs.dump(indent + 2)
return s
class CallExprAST(ExprAST): class CallExprAST(ExprAST):
def __init__(self, callee, args): def __init__(self, callee, args):
self.callee = callee self.callee = callee
self.args = args self.args = args
def dump(self, indent=0):
s = '{0}{1}[{2}]\n'.format(
' ' * indent, self.__class__.__name__, self.callee)
for arg in self.args:
s += arg.dump(indent + 2) + '\n'
return s[:-1] # snip out trailing '\n'
class PrototypeAST(ASTNode): class PrototypeAST(ASTNode):
def __init__(self, name, args): def __init__(self, name, argnames):
self.name = name self.name = name
self.args = args self.argnames = argnames
def dump(self, indent=0):
return '{0}{1}[{2}]'.format(
' ' * indent, self.__class__.__name__, ', '.join(self.argnames))
class FunctionAST(ASTNode): class FunctionAST(ASTNode):
@ -117,6 +144,12 @@ class FunctionAST(ASTNode):
self.proto = proto self.proto = proto
self.body = body self.body = body
def dump(self, indent=0):
s = '{0}{1}[{2}]\n'.format(
' ' * indent, self.__class__.__name__, self.proto.dump())
s += self.body.dump(indent + 2) + '\n'
return s
class ParseError(Exception): pass class ParseError(Exception): pass
@ -127,6 +160,18 @@ class Parser(object):
self.cur_tok = None self.cur_tok = None
self._get_next_token() self._get_next_token()
# toplevel ::= definition | external | expression | ';'
def parse_toplevel(self):
if self.cur_tok.kind == TokenKind.EXTERN:
return self._parse_external()
elif self.cur_tok.kind == TokenKind.DEF:
return self._parse_definition()
elif self._cur_tok_is_operator(';'):
self._get_next_token()
return None
else:
return self._parse_toplevel_expression()
def _get_next_token(self): def _get_next_token(self):
self.cur_tok = next(self.token_generator) self.cur_tok = next(self.token_generator)
@ -135,14 +180,14 @@ class Parser(object):
def _cur_tok_precedence(self): def _cur_tok_precedence(self):
"""Get the operator precedence of the current token.""" """Get the operator precedence of the current token."""
try: try:
return Parser._precedence_map[self.cur_tok] return Parser._precedence_map[self.cur_tok.value]
except KeyError: except KeyError:
return -1 return -1
def _cur_tok_is_operator(self, op): def _cur_tok_is_operator(self, op):
"""Query whether the current token is the operator 'op'""" """Query whether the current token is the operator op"""
return (self.cur_tok.kind == TokenKind.OPERATOR and return (self.cur_tok.kind == TokenKind.OPERATOR and
self.cur_tok.value == 'op') self.cur_tok.value == op)
# identifierexpr # identifierexpr
# ::= identifier # ::= identifier
@ -150,15 +195,16 @@ class Parser(object):
def _parse_identifier_expr(self): def _parse_identifier_expr(self):
id_name = self.cur_tok.value id_name = self.cur_tok.value
self._get_next_token() self._get_next_token()
print('* _parse_identifier_expr:', self.cur_tok)
# If followed by a '(' it's a call; otherwise, a simple variable ref. # If followed by a '(' it's a call; otherwise, a simple variable ref.
if self._cur_tok_is_operator('('): if not self._cur_tok_is_operator('('):
return VariableExprAST(id_name) return VariableExprAST(id_name)
self._get_next_token() self._get_next_token()
args = [] args = []
if not self._cur_tok_is_operator(')'): if not self._cur_tok_is_operator(')'):
while True: while True:
args.push_back(self._parse_expression()) args.append(self._parse_expression())
if self._cur_tok_is_operator(')'): if self._cur_tok_is_operator(')'):
break break
if not self._cur_tok_is_operator(','): if not self._cur_tok_is_operator(','):
@ -198,23 +244,25 @@ class Parser(object):
raise ParseError('Unknown token when expecting an expression') raise ParseError('Unknown token when expecting an expression')
# binoprhs ::= (<binop> primary)* # binoprhs ::= (<binop> primary)*
def _parse_binop_rhs(self, expr_prec, lhs_ast): def _parse_binop_rhs(self, expr_prec, lhs):
"""Parse the right-hand-side of a binary expression. """Parse the right-hand-side of a binary expression.
expr_prec: minimal precedence to keep going (precedence climbing). expr_prec: minimal precedence to keep going (precedence climbing).
lhs_ast: AST of the left-hand-side. lhs: AST of the left-hand-side.
""" """
while True: while True:
cur_prec = self._cur_tok_precedence() cur_prec = self._cur_tok_precedence()
print('* _parse_binop_rhs', self.cur_tok, ' --> prec=', cur_prec)
# If this is a binary operator with precedence lower than the # If this is a binary operator with precedence lower than the
# currently parsed sub-expression, bail out. If it binds at least # currently parsed sub-expression, bail out. If it binds at least
# as tightly, keep going. # as tightly, keep going.
# Note that the precedence of non-operators is defined to be -1, # Note that the precedence of non-operators is defined to be -1,
# so this condition handles cases when the expression ended. # so this condition handles cases when the expression ended.
if cur_prec < expr_prec: if cur_prec < expr_prec:
return lhs_ast return lhs
op = self.cur_tok.value op = self.cur_tok.value
self._get_next_token() # consume the operator self._get_next_token() # consume the operator
print(' after consume op:', self.cur_tok)
rhs = self._parse_primary() rhs = self._parse_primary()
next_prec = self._cur_tok_precedence() next_prec = self._cur_tok_precedence()
@ -232,11 +280,49 @@ class Parser(object):
# expression ::= primary binoprhs # expression ::= primary binoprhs
def _parse_expression(self): def _parse_expression(self):
print('* _parse_expression:', self.cur_tok)
lhs = self._parse_primary() lhs = self._parse_primary()
print(' after lhs:', self.cur_tok)
# Start with precedence 0 because we want to bind any operator to the # Start with precedence 0 because we want to bind any operator to the
# expression at this point. # expression at this point.
return self._parse_binop_rhs(0, lhs) return self._parse_binop_rhs(0, lhs)
# prototype ::= id '(' id* ')'
def _parse_prototype(self):
if self.cur_tok.kind != TokenKind.IDENTIFIER:
raise ParseError('Expected function name in prototype')
name = self.cur_tok.value
self._get_next_token() # consume the name
if not self._cur_tok_is_operator('('):
raise ParseError('Expected "(" in prototype')
self._get_next_token() # consume '('
argnames = []
while self.cur_tok.kind == TokenKind.IDENTIFIER:
argnames.append(self.cur_tok.value)
self._get_next_token()
if not self._cur_tok_is_operator(')'):
raise ParseError('Expected ")" in prototype')
self._get_next_token() # consume ')'
return PrototypeAST(name, argnames)
# external ::= 'extern' prototype
def _parse_external(self):
self._get_next_token() # consume 'extern'
return self._parse_prototype()
# definition ::= 'def' prototype expression
def _parse_definition(self):
self._get_next_token() # consume 'def'
proto = self._parse_prototype()
expr = self._parse_expression()
return FunctionAST(proto, expr)
# toplevel ::= expression
def _parse_toplevel_expression(self):
expr = self._parse_expression()
# Anonymous function
return FunctionAST(PrototypeAST('', []), expr)
#---- Some unit tests ----# #---- Some unit tests ----#
@ -283,7 +369,47 @@ class TestLexer(unittest.TestCase):
['DEF', 'IDENTIFIER', 'NUMBER', 'EOF']) ['DEF', 'IDENTIFIER', 'NUMBER', 'EOF'])
class TestParser(unittest.TestCase):
def _flatten(self, ast):
"""Test helper - flattens the AST into a sexpr-like nested list."""
if isinstance(ast, NumberExprAST):
return ['Number', ast.val]
elif isinstance(ast, VariableExprAST):
return ['Variable', ast.name]
elif isinstance(ast, BinaryExprAST):
return ['Binop', ast.op,
self._flatten(ast.lhs), self._flatten(ast.rhs)]
elif isinstance(ast, CallExprAST):
args = [self._flatten(arg) for arg in ast.args]
return ['Call', ast.callee, args]
elif isinstance(ast, PrototypeAST):
return ['Proto', ast.name, ', '.join(ast.argnames)]
elif isinstance(ast, FunctionAST):
return ['Function',
self._flatten(ast.proto), self._flatten(ast.body)]
else:
raise TypeError('unknown type in _flatten: {0}'.format(type(ast)))
def _assert_body(self, toplevel, expected):
"""Assert the flattened body of the given toplevel function"""
self.assertIsInstance(toplevel, FunctionAST)
self.assertEqual(self._flatten(toplevel.body), expected)
def test_basic(self):
p = Parser('2')
ast = p.parse_toplevel()
self.assertIsInstance(ast, FunctionAST)
self.assertIsInstance(ast.body, NumberExprAST)
self.assertEqual(ast.body.val, '2')
def test_basic_with_flattening(self):
ast = Parser('2').parse_toplevel()
self._assert_body(ast, ['Number', '2'])
ast = Parser('foobar').parse_toplevel()
self._assert_body(ast, ['Variable', 'foobar'])
if __name__ == '__main__': if __name__ == '__main__':
buf = '''2+3''' buf = '''2+3*fob(12, kwa)-.12'''
p = Parser(buf) #p = Parser(buf)
print(p._parse_expression()) #print(p.parse_toplevel().dump())