from logitbias import ( OpBinary, OpAdd, OpMultiply, OpDivide, OpMod, OpAnd, OpOr, OpCompare, OpIn, OpVariableAccess, OpArrayAccess ) reserved = { 'in': 'IN', 'true': 'BOOL', 'false': 'BOOL', } tokens = ( 'NAME', 'NUMBER', 'PLUS','MINUS','TIMES','DIVIDE','MOD', 'LPAREN','RPAREN', 'LBRACK', 'RBRACK', 'AND', 'OR', 'GT', 'LT', 'GE', 'LE', 'EQ', 'NE', 'COMMA', 'NOT', ) + tuple(set(reserved.values())) # Tokens t_PLUS = r'\+' t_MINUS = r'-' t_TIMES = r'\*' t_DIVIDE = r'/' t_MOD = r'%' t_LPAREN = r'\(' t_RPAREN = r'\)' t_LBRACK = r'\[' t_RBRACK = r'\]' t_AND = r'&&' t_OR = r'\|\|' t_GT = r'>' t_LT = r'<' t_GE = r'>=' t_LE = r'<=' t_EQ = r'==?=?' t_NE = r'!==?' t_NOT = r'!' t_COMMA = r',' def t_NUMBER(t): r'[0-9]*\.?[0-9]+' t.value = float(t.value) return t def t_NAME(t): r'[a-zA-Z_][a-zA-Z0-9_]*' t.type = reserved.get(t.value,'NAME') return t # Ignored characters t_ignore = " \t\n" def t_error(t): raise Exception(f"Illegal character {t.value[0]}") # Build the lexer import ply.lex as lex lexer = lex.lex() # Parsing rules precedence = ( ('left','OR'), ('left','AND'), ('left','EQ','NE'), ('left','GT','LT','GE','LE','IN'), ('left','PLUS','MINUS'), ('left','TIMES','DIVIDE','MOD'), ('right','UMINUS','NOT'), ) def p_statement_expr(t): 'statement : expression' t[0] = t[1] def p_expression_binop(t): '''expression : expression PLUS expression | expression MINUS expression | expression TIMES expression | expression DIVIDE expression | expression MOD expression | expression AND expression | expression OR expression | expression GT expression | expression LT expression | expression GE expression | expression LE expression | expression EQ expression | expression NE expression''' if t[2] == '+' : t[0] = OpAdd(t[1], t[3]) elif t[2] == '-': t[0] = OpAdd(t[1], t[3], rhs_flip_sign=True) elif t[2] == '*': t[0] = OpMultiply(t[1], t[3]) elif t[2] == '/': t[0] = OpDivide(t[1], t[3]) elif t[2] == '%': t[0] = OpMod(t[1], t[3]) elif t[2] == '&&': t[0] = OpAnd(t[1], t[3]) elif t[2] == '||': t[0] = OpOr(t[1], t[3]) elif t[2] == '>': t[0] = OpCompare(t[1], t[3], '>') elif t[2] == '<': t[0] = OpCompare(t[1], t[3], '<') elif t[2] == '>=': t[0] = OpCompare(t[1], t[3], '>=') elif t[2] == '<=': t[0] = OpCompare(t[1], t[3], '<=') elif t[2][0] == '=': t[0] = OpCompare(t[1], t[3], '==') elif t[2][:2] == '!=': t[0] = OpCompare(t[1], t[3], '!=') def p_expression_in(t): 'expression : expression IN list' t[0] = OpIn(t[1], t[3]) def p_list(t): 'list : LBRACK list_inner RBRACK' t[0] = t[2] def p_list_inner(t): '''list_inner : list_inner COMMA expression | expression''' if len(t) == 4: t[0] = t[1] + [t[3]] else: t[0] = [t[1]] def p_expression_uminus_not(t): '''expression : MINUS expression %prec UMINUS | NOT expression %prec UMINUS''' if t[1] == '-': t[0] = OpMultiply(-1, t[2]) else: t[0] = OpCompare(False, t[2], '==') def p_expression_group(t): 'expression : LPAREN expression RPAREN' t[0] = t[2] def p_arrayindex(t): '''arrayindex : NUMBER | MINUS NUMBER''' if len(t) == 3: idx = -t[2] else: idx = t[1] assert idx == int(idx) t[0] = int(idx) def p_expression_arrayaccess(t): 'expression : NAME LBRACK arrayindex RBRACK' t[0] = OpArrayAccess(t[1], t[3]) def p_expression_name(t): 'expression : NAME' t[0] = OpVariableAccess(t[1]) def p_expression_number(t): 'expression : NUMBER' t[0] = t[1] def p_expression_bool(t): 'expression : BOOL' t[0] = t[1].lower() == 'true' def p_error(t): raise Exception(f"Syntax error at '{t.value}'") import ply.yacc as yacc parser = yacc.yacc() import dataclasses def fixup(ast): range_comparisons = ('<', '<=', '>', '>=') if not isinstance(ast, OpBinary): return ast fixed_lhs = fixup(ast.lhs) fixed_rhs = fixup(ast.rhs) if isinstance(ast, OpCompare) and ast.comparison in range_comparisons and isinstance(fixed_lhs, OpCompare) and fixed_lhs.comparison in range_comparisons: return OpAnd(fixed_lhs, OpCompare(fixed_lhs.rhs, fixed_rhs, ast.comparison)) return ast def parse(s): return fixup(parser.parse(s))