"""运费公式求值：ast 白名单 + Decimal。

只允许「数字常量、白名单变量名、+ - * /、一元正负号、括号」。
调用/属性/下标/比较/布尔/幂/取模/未知变量一律拒绝——公式来自租户后台输入，
是不可信输入，必须 fail-closed。金额一律 Decimal，禁止 float。
"""
import ast
from decimal import Decimal, DecimalException, localcontext
from typing import Mapping

MAX_FORMULA_CHARS = 500

#: 公式中允许出现的变量名（数值），由报价引擎在求值时全部提供。
ALLOWED_FORMULA_VARIABLES = frozenset({
    "cart_total",
    "cart_subtotal",
    "cart_quantity",
    "cart_weight_kg",
    "cart_volume_cm3",
    "matched_total",
    "matched_quantity",
    "matched_weight_kg",
    "shipping_fee",
})


class FormulaError(ValueError):
    """公式非法或求值失败。写库校验与运行期求值共用同一个异常类型。"""


def _to_decimal(value) -> Decimal:
    """转 Decimal：走 str() 保精度（Decimal(0.1) 会带上 float 误差）。

    变量取值也走这里，所以 None / "" / "abc" 这类脏值必须变成 FormulaError，
    不能漏成 decimal.InvalidOperation（那是 ArithmeticError，调用方 except ValueError 接不住）。
    """
    try:
        return Decimal(str(value))
    except DecimalException as exc:
        raise FormulaError(f"不是合法数值: {value!r}") from exc


def _finite(value: Decimal) -> Decimal:
    if not value.is_finite():
        raise FormulaError("公式结果不是有限数")
    return value


def _is_constant(node: ast.AST) -> bool:
    """子树里没有任何变量名。

    只有纯常量除数才能在校验期断定它真的是 0；**含变量的除数子树一律放行**——
    校验期没有取值，无法判断，交给运行期的除零检查。
    Admin 端文案不要把"公式已校验"说成"不会除零"。
    """
    return not any(isinstance(n, ast.Name) for n in ast.walk(node))


def _eval(node: ast.AST, variables: Mapping[str, Decimal] | None) -> Decimal:
    if isinstance(node, ast.Constant):
        # bool 是 int 的子类，必须先排除；字符串/None/复数一律拒绝
        if isinstance(node.value, bool) or not isinstance(node.value, (int, float)):
            raise FormulaError("公式只允许数字常量")
        return _finite(_to_decimal(node.value))

    if isinstance(node, ast.Name):
        if node.id not in ALLOWED_FORMULA_VARIABLES:
            raise FormulaError(f"未知变量: {node.id}")
        if variables is None:  # 纯语法校验模式，不需要真实取值
            # 占位取 1 而不是 0：0 会让 a/b 在校验期变成 0/0（DivisionUndefined）。
            # 除数究竟是不是 0 由 _is_constant 判定，见 Div 分支。
            return Decimal(1)
        if node.id not in variables:
            raise FormulaError(f"变量未提供取值: {node.id}")
        return _finite(_to_decimal(variables[node.id]))

    if isinstance(node, ast.UnaryOp) and isinstance(node.op, (ast.UAdd, ast.USub)):
        operand = _eval(node.operand, variables)
        return _finite(-operand if isinstance(node.op, ast.USub) else operand)

    if isinstance(node, ast.BinOp):
        op = node.op
        left = _eval(node.left, variables)
        right = _eval(node.right, variables)
        if isinstance(op, ast.Add):
            return _finite(left + right)
        if isinstance(op, ast.Sub):
            return _finite(left - right)
        if isinstance(op, ast.Mult):
            return _finite(left * right)
        if isinstance(op, ast.Div):
            # 校验模式下变量一律占位成 0，此时只有「纯常量除数」才能断定除零；
            # 否则 matched_total / matched_quantity（按件均摊，主流写法）会被误判为除零而存不下来。
            if right == 0 and (variables is not None or _is_constant(node.right)):
                raise FormulaError("公式除数为 0")
            try:
                return _finite(left / right)
            except DecimalException as exc:
                raise FormulaError(f"公式求值失败: {exc}") from exc
        raise FormulaError(f"不支持的运算符: {type(op).__name__}")

    raise FormulaError(f"不支持的公式语法: {type(node).__name__}")


def _parse(source: str) -> ast.AST:
    if not isinstance(source, str):
        raise FormulaError("公式必须是字符串")
    if len(source) > MAX_FORMULA_CHARS:
        raise FormulaError(f"formula 超过上限 {MAX_FORMULA_CHARS} 字符")
    try:
        tree = ast.parse(source.strip(), mode="eval")
    except (SyntaxError, ValueError, MemoryError, RecursionError) as exc:
        raise FormulaError(f"公式语法错误: {exc}") from exc
    return tree.body


def validate_formula(source: str) -> None:
    """写库前的纯语法校验：不需要变量取值，只确认整棵树都在白名单内。"""
    _eval(_parse(source), None)


def evaluate_formula(source: str, variables: Mapping[str, Decimal]) -> Decimal:
    """求值。variables 必须覆盖公式里用到的每个变量，缺一个就报错（不静默当 0）。

    变量取值请用 services.build_formula_variables() 生成，不要手写映射
    （公式变量名与事实字段名不同名，手写错了是"价格算错"而不是报错）。

    两套区间词汇不要混：**条件**用 {min, max}（min 闭、max 开），**阶梯**（Task 3）用 start/end。

    Task 3 注意：
    - 本函数**不给结果兜底**，`5 - cart_total * 2` 可以返回负数。运费为负没有意义，
      调用方必须自己把结果 floor 到 0（这里不做，是为了让 Task 3 能区分"公式写反了"和"就是免运费"）。
    - 求值失败（运行期除零、脏变量值）抛 FormulaError；Task 3 应把它转成
      REASON_KEYS 里的 'pricing_failed' 候选，而不是 500。
    """
    node = _parse(source)
    # ponytail: 固定 28 位精度足够运费场景；真要按币种精度取整交给 Task 3 的量化步骤
    with localcontext() as ctx:
        ctx.prec = 28
        return _eval(node, variables)
