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))
