feat: upgrade toolbox for v1.2.1
This commit is contained in:
+136
-3
@@ -1,13 +1,17 @@
|
||||
import ast
|
||||
import math
|
||||
from statistics import mean, median, pstdev, pvariance
|
||||
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 = {
|
||||
@@ -188,6 +192,31 @@ def parse_number_list(source):
|
||||
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", "")
|
||||
@@ -198,14 +227,23 @@ def calculate(payload):
|
||||
|
||||
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,
|
||||
@@ -278,11 +316,43 @@ def calculate(payload):
|
||||
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}
|
||||
|
||||
expression = parse_equation(source) if operation == "solve" else parse_expression(source)
|
||||
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)
|
||||
@@ -301,8 +371,39 @@ def calculate(payload):
|
||||
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":
|
||||
order = int(payload.get("order", 1))
|
||||
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)
|
||||
@@ -328,6 +429,38 @@ def calculate(payload):
|
||||
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}
|
||||
|
||||
Reference in New Issue
Block a user