"""Recursive-descent parser for MicroBasic.

Grammar and precedence follow the confirmed evidence in
docs/reverse_engineering.md, not just the manual's prose (which doesn't
state operator precedence explicitly). Precedence, loosest to tightest
binding -- confirmed by decoding real compiled output, not assumed:

    Or / XOr
    And
    Not                     (unary; confirmed to bind tighter than And)
    relational: = <> < > <= <=
    shift: << >>
    additive: + -
    multiplicative: * / Mod
    unary -                 (confirmed: unary + does not compile at all --
                              RoborunPlus itself rejects it)
    postfix ++ --
    primary (literals, var refs incl. array index, calls, parens,
              prefix ++/--)

Compound assignment (+=, -=, *=, /=, <<=, >>=) is desugared here into a
plain '=' with a synthesized BinaryOp RHS, since that's confirmed to be
exactly how the real compiler treats it -- codegen.py never sees a compound
assignment operator.
"""

from __future__ import annotations

from .ast_nodes import (
    AssignStmt,
    BinaryOp,
    BoolLiteral,
    Call,
    ContinueStmt,
    DimStmt,
    DoLoopUntilStmt,
    DoLoopWhileStmt,
    DoUntilStmt,
    ExitStmt,
    Expr,
    ExprStmt,
    ForCStyleStmt,
    ForTraditionalStmt,
    GosubStmt,
    GotoStmt,
    IfStmt,
    IncDecExpr,
    IncDecStmt,
    IntLiteral,
    Label,
    OptionExplicitStmt,
    PrintStmt,
    Program,
    ReturnStmt,
    Stmt,
    StringLiteral,
    TerminateStmt,
    UnaryOp,
    VarRef,
    WaitStmt,
    WhileStmt,
)
from .tokens import Token, TokenType

ASSIGN_OPS = {"=", "+=", "-=", "*=", "/=", "<<=", ">>="}
RELATIONAL_OPS = {"=", "<>", "<", ">", "<=", ">="}


class ParseError(Exception):
    def __init__(self, message: str, token: Token):
        super().__init__(f"{message} at {token!r}")
        self.token = token


def parse_int_literal(text: str) -> int:
    low = text.lower()
    if low.startswith("0x"):
        return int(text, 16)
    if low.startswith("0b"):
        return int(text, 2)
    return int(text, 10)


class Parser:
    def __init__(self, tokens: list[Token]):
        self.tokens = tokens
        self.pos = 0

    def parse(self) -> Program:
        statements = self.parse_block(terminators=set())
        return Program(statements=statements)

    # -- token stream helpers ----------------------------------------

    def _peek(self) -> Token:
        return self.tokens[self.pos]

    def _at_end(self) -> bool:
        return self._peek().type == TokenType.EOF

    def _advance(self) -> Token:
        tok = self.tokens[self.pos]
        if not self._at_end():
            self.pos += 1
        return tok

    def _check(self, type_: TokenType) -> bool:
        return self._peek().type == type_

    def _check_op(self, *values: str) -> bool:
        return self._check(TokenType.OP) and self._peek().value in values

    def _check_keyword(self, *words: str) -> bool:
        return self._check(TokenType.KEYWORD) and self._peek().value in words

    def _expect(self, type_: TokenType, message: str) -> Token:
        if not self._check(type_):
            raise ParseError(message, self._peek())
        return self._advance()

    def _expect_keyword(self, word: str) -> Token:
        if not self._check_keyword(word):
            raise ParseError(f"expected '{word}'", self._peek())
        return self._advance()

    def _expect_op(self, value: str) -> Token:
        if not self._check_op(value):
            raise ParseError(f"expected '{value}'", self._peek())
        return self._advance()

    def _skip_newlines(self) -> None:
        while self._check(TokenType.NEWLINE):
            self._advance()

    # -- statements ----------------------------------------------------

    def parse_block(self, terminators: set[str]) -> list[Stmt]:
        stmts: list[Stmt] = []
        while True:
            self._skip_newlines()
            if self._at_end():
                break
            if terminators and self._check_keyword(*terminators):
                break
            stmts.append(self.parse_statement())
        return stmts

    def parse_statement(self) -> Stmt:
        if self._check(TokenType.LABEL):
            tok = self._advance()
            return Label(tok.value)

        if self._check(TokenType.KEYWORD):
            kw = self._peek().value
            dispatch = {
                "dim": self.parse_dim,
                "option": self.parse_option_explicit,
                "if": self.parse_if,
                "while": self.parse_while,
                "do": self.parse_do,
                "for": self.parse_for,
                "goto": self.parse_goto,
                "gosub": self.parse_gosub,
                "return": self.parse_return,
                "terminate": self.parse_terminate,
                "print": self.parse_print,
                "exit": self.parse_exit,
                "continue": self.parse_continue,
            }
            if kw in dispatch:
                return dispatch[kw]()

        return self.parse_simple_statement()

    def parse_dim(self) -> DimStmt:
        self._advance()  # 'dim'
        name_tok = self._expect(TokenType.IDENT, "expected variable name after Dim")
        array_len = None
        if self._check(TokenType.LBRACKET):
            self._advance()
            len_tok = self._expect(TokenType.INTEGER, "expected array length")
            array_len = parse_int_literal(len_tok.value)
            self._expect(TokenType.RBRACKET, "expected ']'")
        self._expect_keyword("as")
        if self._check_keyword("integer"):
            self._advance()
            var_type = "Integer"
        elif self._check_keyword("boolean"):
            self._advance()
            var_type = "Boolean"
        else:
            raise ParseError("expected 'Integer' or 'Boolean'", self._peek())
        return DimStmt(name_tok.value, var_type, array_len)

    def parse_option_explicit(self) -> OptionExplicitStmt:
        self._advance()  # 'option'
        self._expect_keyword("explicit")
        return OptionExplicitStmt()

    def parse_if(self) -> IfStmt:
        self._advance()  # 'if'
        condition = self.parse_expression()
        if self._check_keyword("then"):
            self._advance()

        if self._check(TokenType.NEWLINE) or self._at_end():
            # block form
            self._skip_newlines()
            then_block = self.parse_block({"else", "elseif", "end"})
            elseif_clauses: list[tuple[Expr, list[Stmt]]] = []
            while self._check_keyword("elseif"):
                self._advance()
                cond2 = self.parse_expression()
                if self._check_keyword("then"):
                    self._advance()
                self._skip_newlines()
                body2 = self.parse_block({"else", "elseif", "end"})
                elseif_clauses.append((cond2, body2))
            else_block = None
            if self._check_keyword("else"):
                self._advance()
                self._skip_newlines()
                else_block = self.parse_block({"end"})
            self._expect_keyword("end")
            self._expect_keyword("if")
            return IfStmt(condition, then_block, elseif_clauses, else_block)

        # line form: If <cond> Then <stmt> [Else <stmt>]
        then_stmt = self.parse_statement()
        else_block = None
        if self._check_keyword("else"):
            self._advance()
            else_block = [self.parse_statement()]
        return IfStmt(condition, [then_stmt], [], else_block)

    def parse_while(self) -> WhileStmt:
        self._advance()  # 'while'
        condition = self.parse_expression()
        self._skip_newlines()
        body = self.parse_block({"end"})
        self._expect_keyword("end")
        self._expect_keyword("while")
        return WhileStmt(condition, body)

    def parse_do(self) -> Stmt:
        self._advance()  # 'do'
        if self._check_keyword("while"):
            self._advance()
            condition = self.parse_expression()
            self._skip_newlines()
            body = self.parse_block({"loop"})
            self._expect_keyword("loop")
            # confirmed byte-identical to While/End While
            return WhileStmt(condition, body)
        if self._check_keyword("until"):
            self._advance()
            condition = self.parse_expression()
            self._skip_newlines()
            body = self.parse_block({"loop"})
            self._expect_keyword("loop")
            return DoUntilStmt(condition, body)

        # post-test forms: Do / <body> / Loop {While|Until} <cond>
        self._skip_newlines()
        body = self.parse_block({"loop"})
        self._expect_keyword("loop")
        if self._check_keyword("while"):
            self._advance()
            condition = self.parse_expression()
            return DoLoopWhileStmt(body, condition)
        if self._check_keyword("until"):
            self._advance()
            condition = self.parse_expression()
            return DoLoopUntilStmt(body, condition)
        raise ParseError("expected 'While' or 'Until' after 'Loop'", self._peek())

    def parse_for(self) -> Stmt:
        self._advance()  # 'for'
        var_tok = self._expect(TokenType.IDENT, "expected loop variable name")
        self._expect_op("=")
        init = self.parse_expression()

        if self._check_keyword("to"):
            self._advance()
            limit = self.parse_expression()
            step = None
            if self._check_keyword("step"):
                self._advance()
                step = self.parse_expression()
            self._skip_newlines()
            body = self.parse_block({"next"})
            self._expect_keyword("next")
            if self._check(TokenType.IDENT):  # classic-BASIC "Next <var>" form
                self._advance()
            return ForTraditionalStmt(var_tok.value, init, limit, step, body)

        if self._check_keyword("andwhile"):
            self._advance()
            condition = self.parse_expression()
            evaluate = None
            if self._check_keyword("evaluate"):
                self._advance()
                evaluate = self.parse_simple_statement()
            self._skip_newlines()
            body = self.parse_block({"next"})
            self._expect_keyword("next")
            if self._check(TokenType.IDENT):  # classic-BASIC "Next <var>" form
                self._advance()
            return ForCStyleStmt(var_tok.value, init, condition, evaluate, body)

        raise ParseError("expected 'To' or 'AndWhile' in For statement", self._peek())

    def parse_goto(self) -> GotoStmt:
        self._advance()
        label_tok = self._expect(TokenType.IDENT, "expected label name after GoTo")
        return GotoStmt(label_tok.value)

    def parse_gosub(self) -> GosubStmt:
        self._advance()
        label_tok = self._expect(TokenType.IDENT, "expected label name after GoSub")
        return GosubStmt(label_tok.value)

    def parse_return(self) -> ReturnStmt:
        self._advance()
        return ReturnStmt()

    def parse_terminate(self) -> TerminateStmt:
        self._advance()
        return TerminateStmt()

    def parse_print(self) -> PrintStmt:
        self._advance()
        self._expect(TokenType.LPAREN, "expected '(' after Print")
        args = self._parse_arg_list()
        self._expect(TokenType.RPAREN, "expected ')'")
        return PrintStmt(args)

    def parse_exit(self) -> ExitStmt:
        self._advance()
        return ExitStmt(self._expect_loop_kind())

    def parse_continue(self) -> ContinueStmt:
        self._advance()
        return ContinueStmt(self._expect_loop_kind())

    def _expect_loop_kind(self) -> str:
        for kw, kind in (("for", "For"), ("while", "While"), ("do", "Do")):
            if self._check_keyword(kw):
                self._advance()
                return kind
        raise ParseError("expected 'For', 'While', or 'Do'", self._peek())

    def parse_simple_statement(self) -> Stmt:
        # prefix ++/-- as a bare statement
        if self._check_op("++", "--"):
            op = self._advance().value
            name_tok = self._expect(TokenType.IDENT, "expected variable name")
            return IncDecStmt(name_tok.value, op)

        if self._check(TokenType.IDENT):
            tok = self._advance()

            if self._check(TokenType.LPAREN):
                call = self._parse_call(tok.value)
                if tok.value.lower() == "wait":
                    if len(call.args) != 1:
                        raise ParseError("Wait expects exactly 1 argument", tok)
                    return WaitStmt(call.args[0])
                return ExprStmt(call)

            if self._check_op("++", "--"):
                op = self._advance().value
                return IncDecStmt(tok.value, op)

            index = None
            if self._check(TokenType.LBRACKET):
                self._advance()
                index = self.parse_expression()
                self._expect(TokenType.RBRACKET, "expected ']'")

            if self._check(TokenType.OP) and self._peek().value in ASSIGN_OPS:
                op = self._advance().value
                value = self.parse_expression()
                target = VarRef(tok.value, index)
                if op == "=":
                    return AssignStmt(target, value)
                binop = op[:-1]  # "+=" -> "+", "<<=" -> "<<"
                desugared = BinaryOp(binop, VarRef(tok.value, index), value)
                return AssignStmt(target, desugared)

            raise ParseError(f"unexpected token after '{tok.value}'", self._peek())

        raise ParseError("expected a statement", self._peek())

    def _parse_arg_list(self) -> list[Expr]:
        args: list[Expr] = []
        if not self._check(TokenType.RPAREN):
            args.append(self.parse_expression())
            while self._check(TokenType.COMMA):
                self._advance()
                args.append(self.parse_expression())
        return args

    def _parse_call(self, name: str) -> Call:
        self._expect(TokenType.LPAREN, "expected '('")
        args = self._parse_arg_list()
        self._expect(TokenType.RPAREN, "expected ')'")
        return Call(name, args)

    # -- expressions (precedence climbing) ------------------------------

    def parse_expression(self) -> Expr:
        return self.parse_or()

    def parse_or(self) -> Expr:
        left = self.parse_and()
        while self._check_keyword("or", "xor"):
            op = self._advance().value
            right = self.parse_and()
            left = BinaryOp(op, left, right)
        return left

    def parse_and(self) -> Expr:
        left = self.parse_not()
        while self._check_keyword("and"):
            self._advance()
            right = self.parse_not()
            left = BinaryOp("and", left, right)
        return left

    def parse_not(self) -> Expr:
        if self._check_keyword("not"):
            self._advance()
            return UnaryOp("not", self.parse_not())
        return self.parse_relational()

    def parse_relational(self) -> Expr:
        left = self.parse_shift()
        while self._check(TokenType.OP) and self._peek().value in RELATIONAL_OPS:
            op = self._advance().value
            right = self.parse_shift()
            left = BinaryOp(op, left, right)
        return left

    def parse_shift(self) -> Expr:
        left = self.parse_additive()
        while self._check_op("<<", ">>"):
            op = self._advance().value
            right = self.parse_additive()
            left = BinaryOp(op, left, right)
        return left

    def parse_additive(self) -> Expr:
        left = self.parse_multiplicative()
        while self._check_op("+", "-"):
            op = self._advance().value
            right = self.parse_multiplicative()
            left = BinaryOp(op, left, right)
        return left

    def parse_multiplicative(self) -> Expr:
        left = self.parse_unary()
        while self._check_op("*", "/") or self._check_keyword("mod"):
            op = self._advance().value
            right = self.parse_unary()
            left = BinaryOp(op, left, right)
        return left

    def parse_unary(self) -> Expr:
        if self._check_op("-"):
            self._advance()
            return UnaryOp("-", self.parse_unary())
        return self.parse_postfix()

    def parse_postfix(self) -> Expr:
        expr = self.parse_primary()
        if self._check_op("++", "--"):
            op = self._advance().value
            if not isinstance(expr, VarRef) or expr.index is not None:
                raise ParseError("'++'/'--' can only apply to a simple variable", self._peek())
            return IncDecExpr(expr.name, op, prefix=False)
        return expr

    def parse_primary(self) -> Expr:
        tok = self._peek()

        if self._check_op("++", "--"):
            op = self._advance().value
            name_tok = self._expect(TokenType.IDENT, "expected variable name")
            return IncDecExpr(name_tok.value, op, prefix=True)

        if self._check(TokenType.INTEGER):
            self._advance()
            return IntLiteral(parse_int_literal(tok.value))

        if self._check(TokenType.STRING):
            self._advance()
            return StringLiteral(tok.value)

        if self._check_keyword("true"):
            self._advance()
            return BoolLiteral(True)

        if self._check_keyword("false"):
            self._advance()
            return BoolLiteral(False)

        if self._check_keyword("tobool"):
            self._advance()
            self._expect(TokenType.LPAREN, "expected '(' after ToBool")
            inner = self.parse_expression()
            self._expect(TokenType.RPAREN, "expected ')'")
            return Call("tobool", [inner])

        if self._check(TokenType.LPAREN):
            self._advance()
            inner = self.parse_expression()
            self._expect(TokenType.RPAREN, "expected ')'")
            return inner

        if self._check(TokenType.IDENT):
            self._advance()
            if self._check(TokenType.LPAREN):
                return self._parse_call(tok.value)
            index = None
            if self._check(TokenType.LBRACKET):
                self._advance()
                index = self.parse_expression()
                self._expect(TokenType.RBRACKET, "expected ']'")
            return VarRef(tok.value, index)

        raise ParseError(f"unexpected token {tok!r} in expression", tok)


def parse(tokens: list[Token]) -> Program:
    return Parser(tokens).parse()
