from dataclasses import dataclass, replace
import math


class EscapeToZero(Exception):
    pass


def escape_to_zero():
    raise EscapeToZero()


@dataclass
class ASTNode:
    pass


@dataclass
class OpBinary(ASTNode):
    lhs: any
    rhs: any


@dataclass
class OpAdd(OpBinary):
    rhs_flip_sign: bool = False


@dataclass
class OpMultiply(OpBinary):
    pass


@dataclass
class OpDivide(OpBinary):
    pass


@dataclass
class OpMod(OpBinary):
    pass


@dataclass
class OpAnd(OpBinary):
    pass


@dataclass
class OpOr(OpBinary):
    pass


@dataclass
class OpCompare(OpBinary):
    comparison: str


@dataclass
class OpVariableAccess(ASTNode):
    name: any


@dataclass
class OpArrayAccess(ASTNode):
    name: any
    index: any


@dataclass
class OpIn(OpBinary):
    pass


@dataclass
class VariableStub:
    pass


def eval(ast, env):
    def getenv(name):
        val = env.get(name, None)
        if callable(val):
            val = val()
        if val is None:
            raise Exception(f"Variable {ast.name} not found")
        return val

    if isinstance(ast, list):
        # OpIn rhs
        return [eval(a, env) for a in ast]
    if isinstance(ast, OpBinary):
        lhs_val = eval(ast.lhs, env)

        if isinstance(lhs_val, ASTNode):  # lhs contains variables
            try:
                rhs_val = eval(ast.rhs, env)
                return replace(ast, lhs=lhs_val, rhs=rhs_val)
            except (EscapeToZero, IndexError) as e:

                def raiser():
                    raise e

                return replace(ast, lhs=lhs_val, rhs=raiser)

        # short-circuit
        if isinstance(ast, OpAnd):
            if not lhs_val:
                return False
            return eval(ast.rhs, env)
        elif isinstance(ast, OpOr):
            if lhs_val:
                return True
            return eval(ast.rhs, env)

        rhs_val = eval(ast.rhs, env)
        if isinstance(rhs_val, ASTNode):  # rhs contains variables
            return replace(ast, lhs=lhs_val, rhs=rhs_val)

        if isinstance(ast, OpAdd):
            return lhs_val + (-rhs_val if ast.rhs_flip_sign else rhs_val)
        elif isinstance(ast, OpMultiply):
            return lhs_val * rhs_val
        elif isinstance(ast, OpDivide):
            return lhs_val / rhs_val
        elif isinstance(ast, OpMod):
            return lhs_val % rhs_val
        elif isinstance(ast, OpCompare):
            if ast.comparison == "==":
                if isinstance(lhs_val, float) or isinstance(rhs_val, float):
                    return math.isclose(float(lhs_val), float(rhs_val), rel_tol=1e-5)
                return lhs_val == rhs_val
            elif ast.comparison == "!=":
                return lhs_val != rhs_val
            elif ast.comparison == "<":
                return lhs_val < rhs_val
            elif ast.comparison == "<=":
                return lhs_val <= rhs_val
            elif ast.comparison == ">":
                return lhs_val > rhs_val
            elif ast.comparison == ">=":
                return lhs_val >= rhs_val
        elif isinstance(ast, OpIn):
            if isinstance(lhs_val, float):
                return any(
                    math.isclose(lhs_val, float(x), rel_tol=1e-5) for x in rhs_val
                )
            return lhs_val in rhs_val
    elif isinstance(ast, OpVariableAccess):
        val = getenv(ast.name)
        if isinstance(val, VariableStub):
            return ast
        if isinstance(val, list):
            raise Exception(f"Variable {ast.name} is a list, but was used as a scalar")
        return val
    elif isinstance(ast, OpArrayAccess):
        val = getenv(ast.name)
        if not isinstance(val, list):
            raise Exception(f"Variable {ast.name} is a scalar, but was used as a list")
        return val[ast.index]
    elif callable(ast):
        return ast()
    return ast


# pitch_bias_func = logitbias_simplify(self.bias, env, "pitch")
# take everything in env to be constant
# return a function that evaluates ast in terms of variable's value,
# or a constant if no expressions depend on variable
#
# pairs is a list of (ast, value)
def logitbias_simplify(pairs, env, variable):
    simplified_pairs = []
    bias = 0
    for ast, value in pairs:
        try:
            res = eval(ast, dict(env, **{variable: VariableStub()}))
        except EscapeToZero:
            continue
        if isinstance(res, ASTNode):
            simplified_pairs.append((res, value))
        elif res:
            bias += value

    simplified_env = {}

    if len(simplified_pairs) == 0:
        return bias

    def simplified_eval(val):
        simplified_env[variable] = val
        return bias + sum(
            bias_val if eval(ast, simplified_env) else 0
            for ast, bias_val in simplified_pairs
        )

    return simplified_eval
