Files
Hulumath-Web/backend/toolbox/engine.py
T

334 lines
13 KiB
Python
Raw 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
import sympy as sp
from rest_framework.exceptions import ValidationError
MAX_EXPRESSION_LENGTH = 500
MAX_AST_NODES = 120
MAX_MATRIX_CELLS = 36
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 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)
result = {
"count": len(values),
"mean": mean(values),
"median": median(values),
"variance": pvariance(values),
"standard_deviation": pstdev(values),
"minimum": min(values),
"maximum": max(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 = ["读取矩阵", "交换行列"]
else:
raise ValidationError({"operation": "不支持的矩阵操作"})
return {"operation": operation, "result": serialize_math(result), "steps": steps}
expression = parse_equation(source) if operation == "solve" 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 == "derivative":
order = int(payload.get("order", 1))
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}")
else:
raise ValidationError({"operation": "不支持的计算类型"})
return {"operation": operation, "result": serialize_math(result), "steps": steps}