"""Safe arithmetic calculator for chat (no eval of arbitrary code)."""

from __future__ import annotations

import ast
import re
from decimal import Decimal, InvalidOperation, localcontext

_MATH_PREFIX_RE = re.compile(
    r"^(please\s+)?(can\s+you\s+)?(please\s+)?"
    r"(calculate|compute|solve|evaluate|simplify|what\s+is|whats|what\'s)\s+",
    re.IGNORECASE,
)
_PURE_MATH_RE = re.compile(r"^[\d\s+\-*/().%,^=xX÷×]+$")


def is_pure_math(text: str) -> bool:
    """True for calculator-style expressions with no accounting wording."""
    s = _normalize_expr(text)
    if len(s) < 3 or not re.search(r"\d", s):
        return False
    if not re.search(r"[+\-*/÷×xX^=]", s):
        return False
    if re.search(r"[a-wyzA-WYZ]", s):  # allow x/X as multiply
        return False
    return bool(_PURE_MATH_RE.fullmatch(s))


def _normalize_expr(text: str) -> str:
    s = text.strip().rstrip("?.!")
    s = _MATH_PREFIX_RE.sub("", s).strip()
    if s.count("=") == 1 and not s.strip().startswith("="):
        s = s.split("=", 1)[1].strip()
    s = s.replace("×", "*").replace("÷", "/").replace("^", "**")
    s = re.sub(r"(?<=\d)\s*[xX]\s*(?=\d)", "*", s)
    s = s.replace("%", "/100")
    return s.strip()


def _eval_node(node: ast.AST) -> Decimal:
    if isinstance(node, ast.Expression):
        return _eval_node(node.body)
    if isinstance(node, ast.Constant):
        # Exact integers only via Decimal(int) — never float (avoids e+55 rounding)
        if isinstance(node.value, int) and not isinstance(node.value, bool):
            return Decimal(node.value)
        if isinstance(node.value, float):
            raise ValueError("Float literals are not supported")
        raise ValueError("Only numbers are allowed")
    if isinstance(node, ast.UnaryOp):
        val = _eval_node(node.operand)
        if isinstance(node.op, ast.UAdd):
            return val
        if isinstance(node.op, ast.USub):
            return -val
        raise ValueError("Unsupported unary operator")
    if isinstance(node, ast.BinOp):
        left = _eval_node(node.left)
        right = _eval_node(node.right)
        if isinstance(node.op, ast.Add):
            return left + right
        if isinstance(node.op, ast.Sub):
            return left - right
        if isinstance(node.op, ast.Mult):
            return left * right
        if isinstance(node.op, ast.Div):
            if right == 0:
                raise ZeroDivisionError("Division by zero")
            return left / right
        if isinstance(node.op, ast.FloorDiv):
            if right == 0:
                raise ZeroDivisionError("Division by zero")
            return left // right
        if isinstance(node.op, ast.Mod):
            if right == 0:
                raise ZeroDivisionError("Division by zero")
            return left % right
        if isinstance(node.op, ast.Pow):
            if right != right.to_integral_value() or abs(int(right)) > 1000:
                raise ValueError("Exponent too large")
            return left ** int(right)
        raise ValueError("Unsupported operator")
    raise ValueError("Unsupported expression")


def _comma_int(n: int) -> str:
    sign = "-" if n < 0 else ""
    digits = str(abs(n))
    parts: list[str] = []
    while digits:
        parts.append(digits[-3:])
        digits = digits[:-3]
    return sign + ",".join(reversed(parts))


def _format_result(value: Decimal) -> str:
    """Always fixed-point / full integer — never scientific notation."""
    if value == value.to_integral_value():
        return _comma_int(int(value))

    text = format(value, "f")
    if "." in text:
        text = text.rstrip("0").rstrip(".")
    negative = text.startswith("-")
    if negative:
        text = text[1:]
    if "." in text:
        whole, frac = text.split(".", 1)
        out = f"{_comma_int(int(whole or '0'))}.{frac}"
    else:
        out = _comma_int(int(text or "0"))
    return f"-{out.lstrip('-')}" if negative else out


def calculate(text: str) -> str:
    """Evaluate a pure arithmetic expression and return a short reply."""
    expr = _normalize_expr(text)
    if not expr:
        return "I couldn't read that calculation. Try something like 12 + 5 * 3."
    if len(expr) > 4000:
        return "That expression is too long to calculate."

    digit_lens = [len(m.group(0)) for m in re.finditer(r"\d+", expr)]
    if digit_lens and max(digit_lens) > 200:
        return (
            "One of those numbers is too large for a reliable calculator answer. "
            "Try a shorter expression."
        )

    prec = max(50, (max(digit_lens) if digit_lens else 20) + 40)

    try:
        tree = ast.parse(expr, mode="eval")
        with localcontext() as ctx:
            ctx.prec = prec
            result = _eval_node(tree)
        return f"Result: {_format_result(result)}"
    except ZeroDivisionError:
        return "That expression divides by zero, so it can't be calculated."
    except (ValueError, SyntaxError, TypeError, InvalidOperation, OverflowError, MemoryError):
        return (
            "I couldn't calculate that. Use numbers with + − × ÷ and parentheses, "
            "for example (100 + 20) / 4."
        )
