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}