Files
Hulumath-Web/backend/toolbox/engine.py
Jacky 8e94f053b1
CI / test (pull_request) Canceled after 52s
PR合并自动部署 / release-check (pull_request) Successful in 13s
PR合并自动部署 / deploy (pull_request) Successful in 13s
feat: upgrade toolbox for v1.2.1
2026-08-10 02:05:21 +08:00

467 lines
19 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import ast
import math
from statistics import mean, median, pstdev, pvariance, quantiles
import sympy as sp
from mpmath.libmp.libhyper import NoConvergence
from rest_framework.exceptions import ValidationError
MAX_EXPRESSION_LENGTH = 500
MAX_AST_NODES = 120
MAX_MATRIX_CELLS = 36
MAX_SERIES_ORDER = 12
MAX_SYSTEM_EQUATIONS = 6
MAX_POLYNOMIAL_DEGREE = 12
SYMBOLS = {name: sp.Symbol(name, real=True) for name in ("x", "y", "z", "a", "b", "t", "n")}
CONSTANTS = {"pi": sp.pi, "e": sp.E, "E": sp.E, "i": sp.I, "I": sp.I}
FUNCTIONS = {
"sin": sp.sin,
"cos": sp.cos,
"tan": sp.tan,
"asin": sp.asin,
"acos": sp.acos,
"atan": sp.atan,
"sinh": sp.sinh,
"cosh": sp.cosh,
"tanh": sp.tanh,
"sqrt": sp.sqrt,
"exp": sp.exp,
"ln": sp.log,
"log": sp.log,
"abs": sp.Abs,
"factorial": sp.factorial,
"binomial": sp.binomial,
"gcd": sp.gcd,
"lcm": sp.lcm,
"floor": sp.floor,
"ceil": sp.ceiling,
}
FUNCTION_ARITY = {
"factorial": (1, 1),
"binomial": (2, 2),
"gcd": (2, 2),
"lcm": (2, 2),
}
UNIT_FACTORS = {
"mm": ("length", 0.001),
"cm": ("length", 0.01),
"m": ("length", 1.0),
"km": ("length", 1000.0),
"in": ("length", 0.0254),
"ft": ("length", 0.3048),
"g": ("mass", 0.001),
"kg": ("mass", 1.0),
"lb": ("mass", 0.45359237),
"s": ("time", 1.0),
"min": ("time", 60.0),
"h": ("time", 3600.0),
"rad": ("angle", 1.0),
"deg": ("angle", math.pi / 180),
}
class SafeExpressionParser:
def __init__(self, source):
source = str(source or "").strip()
if not source:
raise ValidationError({"expression": "请输入数学表达式"})
if len(source) > MAX_EXPRESSION_LENGTH:
raise ValidationError({"expression": "表达式不能超过 500 个字符"})
source = (
source.replace("π", "pi")
.replace("×", "*")
.replace("÷", "/")
.replace("", "-")
.replace("^", "**")
)
try:
self.tree = ast.parse(source, mode="eval")
except SyntaxError as exc:
raise ValidationError({"expression": "表达式语法无效"}) from exc
if sum(1 for _ in ast.walk(self.tree)) > MAX_AST_NODES:
raise ValidationError({"expression": "表达式过于复杂"})
def parse(self):
return self._convert(self.tree.body)
def _convert(self, node):
if isinstance(node, ast.Constant):
if isinstance(node.value, bool) or not isinstance(node.value, (int, float)):
raise ValidationError({"expression": "只允许数值常量"})
if isinstance(node.value, int) and len(str(abs(node.value))) > 50:
raise ValidationError({"expression": "整数位数过多"})
return sp.Integer(node.value) if isinstance(node.value, int) else sp.Float(node.value)
if isinstance(node, ast.Name):
if node.id in SYMBOLS:
return SYMBOLS[node.id]
if node.id in CONSTANTS:
return CONSTANTS[node.id]
raise ValidationError({"expression": f"不支持变量或常量 {node.id}"})
if isinstance(node, ast.UnaryOp) and isinstance(node.op, (ast.UAdd, ast.USub)):
value = self._convert(node.operand)
return value if isinstance(node.op, ast.UAdd) else -value
if isinstance(node, ast.BinOp):
left = self._convert(node.left)
right = self._convert(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):
return left / right
if isinstance(node.op, ast.Mod):
return sp.Mod(left, right)
if isinstance(node.op, ast.Pow):
if right.is_number and abs(float(right)) > 100:
raise ValidationError({"expression": "幂指数绝对值不能超过 100"})
return left**right
if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
function = FUNCTIONS.get(node.func.id)
if function is None:
raise ValidationError({"expression": f"不支持函数 {node.func.id}"})
minimum, maximum = FUNCTION_ARITY.get(node.func.id, (1, 2))
if node.keywords or not minimum <= len(node.args) <= maximum:
raise ValidationError({"expression": f"{node.func.id} 的参数数量无效"})
try:
return function(*(self._convert(item) for item in node.args))
except (TypeError, ValueError) as exc:
raise ValidationError({"expression": f"{node.func.id} 的参数无效"}) from exc
raise ValidationError({"expression": "表达式包含不允许的语法"})
def parse_expression(source):
return SafeExpressionParser(source).parse()
def parse_equation(source):
source = str(source or "")
if source.count("=") > 1:
raise ValidationError({"expression": "方程只能包含一个等号"})
if "=" not in source:
return parse_expression(source)
left, right = source.split("=", 1)
return sp.Eq(parse_expression(left), parse_expression(right))
def serialize_math(value):
if isinstance(value, dict):
return {
str(serialize_math(key)["exact"]): serialize_math(item)
for key, item in value.items()
}
if isinstance(value, (list, tuple)):
return [serialize_math(item) for item in value]
if isinstance(value, sp.MatrixBase):
return {
"exact": str(value.tolist()),
"decimal": str(value.evalf(12).tolist()),
"latex": sp.latex(value),
}
exact = str(value)
try:
decimal = str(sp.N(value, 12))
except Exception:
decimal = exact
return {"exact": exact, "decimal": decimal, "latex": sp.latex(value)}
def parse_matrix(source):
rows = [row.strip() for row in str(source or "").split(";") if row.strip()]
if not rows:
raise ValidationError({"expression": "矩阵格式示例:1,2;3,4"})
parsed = [[parse_expression(cell.strip()) for cell in row.split(",")] for row in rows]
width = len(parsed[0])
if width == 0 or any(len(row) != width for row in parsed):
raise ValidationError({"expression": "矩阵每行列数必须一致"})
if len(parsed) * width > MAX_MATRIX_CELLS:
raise ValidationError({"expression": "矩阵最多支持 36 个元素"})
return sp.Matrix(parsed)
def parse_number_list(source):
try:
values = [float(item.strip()) for item in str(source or "").split(",") if item.strip()]
except ValueError as exc:
raise ValidationError({"expression": "统计数据必须是逗号分隔的数字"}) from exc
if not 1 <= len(values) <= 500:
raise ValidationError({"expression": "统计数据数量必须在 1 到 500 之间"})
if not all(math.isfinite(item) for item in values):
raise ValidationError({"expression": "统计数据必须是有限数值"})
return values
def parse_variables(raw_variables, fallback="x"):
names = [
item.strip()
for item in str(raw_variables or fallback).split(",")
if item.strip()
]
if not names or len(names) > 4 or len(set(names)) != len(names):
raise ValidationError({"variables": "变量应为 1 到 4 个不重复名称"})
try:
return [SYMBOLS[name] for name in names]
except KeyError as exc:
raise ValidationError(
{"variables": "变量仅支持 x、y、z、a、b、t、n"}
) from exc
def parse_equation_system(source):
parts = [item.strip() for item in str(source or "").split(";") if item.strip()]
if not 1 <= len(parts) <= MAX_SYSTEM_EQUATIONS:
raise ValidationError(
{"expression": "方程组应使用分号分隔,最多支持 6 个方程"}
)
return [parse_equation(item) for item in parts]
def calculate(payload):
operation = str(payload.get("operation", "calculate"))
source = payload.get("expression", "")
variable_name = str(payload.get("variable", "x"))
variable = SYMBOLS.get(variable_name)
if variable is None:
raise ValidationError({"variable": "变量仅支持 x、y、z、a、b、t、n"})
if operation == "statistics":
values = parse_number_list(source)
quartile_values = (
quantiles(values, n=4, method="inclusive")
if len(values) > 1
else [values[0], values[0], values[0]]
)
result = {
"count": len(values),
"sum": sum(values),
"mean": mean(values),
"median": median(values),
"variance": pvariance(values),
"standard_deviation": pstdev(values),
"minimum": min(values),
"q1": quartile_values[0],
"q3": quartile_values[2],
"maximum": max(values),
"range": max(values) - min(values),
}
return {
"operation": operation,
"result": result,
"steps": ["读取数据", "计算集中趋势", "计算离散程度"],
}
if operation == "base":
try:
from_base = int(payload.get("from_base", 10))
to_base = int(payload.get("to_base", 2))
if not 2 <= from_base <= 36 or not 2 <= to_base <= 36:
raise ValueError
number = int(str(source).strip(), from_base)
except ValueError as exc:
raise ValidationError({"expression": "进制必须为 2 到 36,且输入应合法"}) from exc
digits = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ"
sign = "-" if number < 0 else ""
remaining = abs(number)
converted = "0"
if remaining:
pieces = []
while remaining:
remaining, index = divmod(remaining, to_base)
pieces.append(digits[index])
converted = "".join(reversed(pieces))
return {
"operation": operation,
"result": {"exact": f"{sign}{converted}", "decimal": str(number), "latex": sign + converted},
"steps": [f"按 {from_base} 进制读取", f"转换为 {to_base} 进制"],
}
if operation == "unit":
try:
value = float(str(source).strip())
from_unit = str(payload.get("from_unit", "m"))
to_unit = str(payload.get("to_unit", "cm"))
source_unit = UNIT_FACTORS[from_unit]
target_unit = UNIT_FACTORS[to_unit]
if source_unit[0] != target_unit[0] or not math.isfinite(value):
raise ValueError
except (KeyError, ValueError) as exc:
raise ValidationError({"expression": "单位不兼容或数值无效"}) from exc
converted = value * source_unit[1] / target_unit[1]
return {
"operation": operation,
"result": {
"exact": f"{converted:.12g} {to_unit}",
"decimal": f"{converted:.12g}",
"latex": f"{converted:.12g}\\,{to_unit}",
},
"steps": [f"将 {from_unit} 换算为标准单位", f"转换为 {to_unit}"],
}
if operation.startswith("matrix_"):
matrix = parse_matrix(source)
if operation == "matrix_det":
if not matrix.is_square:
raise ValidationError({"expression": "行列式要求方阵"})
result = matrix.det()
steps = ["读取矩阵", "按行列式规则计算"]
elif operation == "matrix_inverse":
if not matrix.is_square or matrix.det() == 0:
raise ValidationError({"expression": "矩阵不可逆"})
result = matrix.inv()
steps = ["读取矩阵", "验证行列式非零", "计算逆矩阵"]
elif operation == "matrix_rref":
result = matrix.rref()[0]
steps = ["读取矩阵", "执行初等行变换", "得到行最简形"]
elif operation == "matrix_transpose":
result = matrix.T
steps = ["读取矩阵", "交换行列"]
elif operation == "matrix_rank":
result = matrix.rank()
steps = ["读取矩阵", "执行行变换", "计算矩阵秩"]
elif operation == "matrix_nullspace":
result = matrix.nullspace()
steps = ["读取矩阵", "求解齐次线性方程组", "得到零空间基"]
elif operation == "matrix_eigenvalues":
if not matrix.is_square:
raise ValidationError({"expression": "特征值要求方阵"})
result = matrix.eigenvals()
steps = ["读取方阵", "构造特征多项式", "计算特征值及重数"]
else:
raise ValidationError({"operation": "不支持的矩阵操作"})
return {"operation": operation, "result": serialize_math(result), "steps": steps}
if operation == "solve_system":
variables = parse_variables(payload.get("variables"), variable_name)
equations = parse_equation_system(source)
result = sp.solve(equations, variables, dict=True)
if len(result) > 50:
raise ValidationError({"expression": "方程组解的数量过多"})
return {
"operation": operation,
"result": serialize_math(result),
"steps": [
f"读取 {len(equations)} 个方程",
f"以 {', '.join(str(item) for item in variables)} 为未知量",
"联立消元并求解",
],
}
equation_operations = {"solve", "polynomial_roots"}
expression = (
parse_equation(source)
if operation in equation_operations
else parse_expression(source)
)
steps = ["解析受限数学表达式"]
if operation == "calculate":
result = sp.simplify(expression)
steps.append("化简并保留精确值")
elif operation == "simplify":
result = sp.trigsimp(sp.cancel(expression))
steps.append("约分并进行代数/三角化简")
elif operation == "expand":
result = sp.expand(expression)
steps.append("展开乘积与幂")
elif operation == "factor":
result = sp.factor(expression)
steps.append("提取因式并分解")
elif operation == "solve":
result = sp.solve(expression, variable)
if len(result) > 50:
raise ValidationError({"expression": "解的数量过多"})
steps.extend([f"以 {variable_name} 为未知量", "求解方程"])
elif operation == "polynomial_roots":
polynomial_expression = (
expression.lhs - expression.rhs
if isinstance(expression, sp.Equality)
else expression
)
if polynomial_expression.free_symbols - {variable}:
raise ValidationError({"expression": "数值求根仅支持所选变量"})
try:
polynomial = sp.Poly(polynomial_expression, variable)
except sp.PolynomialError as exc:
raise ValidationError({"expression": "请输入单变量多项式"}) from exc
if not 1 <= polynomial.degree() <= MAX_POLYNOMIAL_DEGREE:
raise ValidationError(
{"expression": "数值求根支持 1 到 12 次单变量多项式"}
)
try:
result = list(sp.nroots(polynomial, n=12, maxsteps=100))
except NoConvergence as exc:
raise ValidationError(
{"expression": "数值求根未收敛,请简化多项式后重试"}
) from exc
steps.extend(
[
f"构造 {polynomial.degree()} 次多项式",
"使用高精度数值方法计算全部根",
]
)
elif operation == "derivative":
try:
order = int(payload.get("order", 1))
except (TypeError, ValueError) as exc:
raise ValidationError({"order": "导数阶数必须是整数"}) from exc
if not 1 <= order <= 5:
raise ValidationError({"order": "导数阶数必须在 1 到 5 之间"})
result = sp.diff(expression, variable, order)
steps.append(f"对 {variable_name}{order} 阶导数")
elif operation == "integral":
lower = str(payload.get("lower", "")).strip()
upper = str(payload.get("upper", "")).strip()
if lower or upper:
if not lower or not upper:
raise ValidationError({"bounds": "定积分必须同时填写上下限"})
result = sp.integrate(
expression,
(variable, parse_expression(lower), parse_expression(upper)),
)
steps.append(f"对 {variable_name} 计算定积分")
else:
result = sp.integrate(expression, variable)
steps.append(f"对 {variable_name} 计算不定积分")
elif operation == "limit":
point = parse_expression(payload.get("point", "0"))
direction = str(payload.get("direction", "+-"))
if direction not in {"+", "-", "+-"}:
raise ValidationError({"direction": "极限方向必须为 +、- 或 +-"})
result = sp.limit(expression, variable, point, dir=direction)
steps.append(f"令 {variable_name} 趋近 {point}")
elif operation == "series":
point = parse_expression(payload.get("point", "0"))
try:
order = int(payload.get("order", 6))
except (TypeError, ValueError) as exc:
raise ValidationError({"order": "级数展开阶数必须是整数"}) from exc
if not 1 <= order <= MAX_SERIES_ORDER:
raise ValidationError(
{"order": f"级数展开阶数必须在 1 到 {MAX_SERIES_ORDER} 之间"}
)
result = sp.series(expression, variable, point, order)
steps.append(
f"在 {variable_name} = {point} 附近展开到 {order - 1} 阶"
)
elif operation == "gradient":
variables = parse_variables(payload.get("variables"), variable_name)
result = sp.Matrix([sp.diff(expression, item) for item in variables])
steps.extend(
[
f"选取变量 {', '.join(str(item) for item in variables)}",
"分别计算一阶偏导并组成梯度",
]
)
elif operation == "hessian":
variables = parse_variables(payload.get("variables"), variable_name)
result = sp.hessian(expression, variables)
steps.extend(
[
f"选取变量 {', '.join(str(item) for item in variables)}",
"计算全部二阶偏导并组成 Hessian 矩阵",
]
)
else:
raise ValidationError({"operation": "不支持的计算类型"})
return {"operation": operation, "result": serialize_math(result), "steps": steps}