refactor: split flask app structure

This commit is contained in:
2026-07-06 14:41:16 +08:00
parent c2198aad51
commit a75e857c71
18 changed files with 1871 additions and 1390 deletions
+116
View File
@@ -0,0 +1,116 @@
from __future__ import annotations
import logging
from datetime import datetime
from flask import Flask, abort, jsonify, request
from .config import Config, DATA_DIR, ensure_dirs
from .extensions import db, login_manager
from .models import User
from .prediction import FEATURES, load_model
from .security import csrf_token, validate_csrf_token
def create_app(config_object: type[Config] = Config, *, load_model_on_start: bool = True) -> Flask:
ensure_dirs()
configure_logging()
app = Flask(__name__, template_folder="../templates", static_folder="../static")
app.config.from_object(config_object)
if app.config.get("SECRET_KEY_GENERATED"):
logging.warning("未设置 SECRET_KEY,已生成临时密钥;服务重启后登录会话将失效。")
db.init_app(app)
login_manager.init_app(app)
login_manager.login_view = "main.login"
login_manager.login_message = "请先登录后再访问该页面。"
register_app_hooks(app)
from .routes import bp
app.register_blueprint(bp)
with app.app_context():
db.create_all()
init_admin_user(app)
if load_model_on_start:
try:
app.config["RSF_MODEL"] = load_model(app.config["MODEL_PATH"])
logging.info("模型加载成功")
except Exception as exc:
app.config["RSF_MODEL"] = None
logging.exception("模型加载失败: %s", exc)
return app
def configure_logging() -> None:
logging.basicConfig(
filename=str(DATA_DIR / "app.log"),
level=logging.INFO,
format="%(asctime)s - %(levelname)s - %(message)s",
)
def init_admin_user(app: Flask) -> None:
admin_username = app.config["ADMIN_USERNAME"]
admin_password = app.config["ADMIN_PASSWORD"]
if admin_password:
admin = User.query.filter_by(username=admin_username).first()
if admin is None:
admin = User(username=admin_username, is_admin=True)
admin.set_password(admin_password)
db.session.add(admin)
else:
admin.is_admin = True
if not admin.check_password(admin_password):
admin.set_password(admin_password)
db.session.commit()
elif not User.query.filter_by(is_admin=True).first():
logging.warning("未设置 ADMIN_PASSWORD,跳过自动创建管理员账号。")
default_admin = User.query.filter_by(username="admin", is_admin=True).first()
if default_admin and default_admin.check_password("admin123"):
logging.warning("检测到默认管理员密码 admin123,请立即通过 ADMIN_PASSWORD 更新。")
def register_app_hooks(app: Flask) -> None:
@login_manager.user_loader
def load_user(user_id: str):
try:
return db.session.get(User, int(user_id))
except (TypeError, ValueError):
return None
@app.context_processor
def inject_helpers():
def now_year() -> int:
return datetime.now().year
return {
"feature_list": FEATURES,
"now_year": now_year,
"csrf_token": csrf_token,
"allow_registration": app.config["ALLOW_REGISTRATION"],
}
@app.errorhandler(413)
def handle_file_too_large(_exc):
max_mb = app.config["MAX_CONTENT_LENGTH"] // (1024 * 1024)
if request.path == "/predict":
return jsonify({"error": f"文件过大,请上传 {max_mb}MB 以内的文件。"}), 413
return f"文件过大,请上传 {max_mb}MB 以内的文件。", 413
@app.before_request
def protect_csrf():
if request.method not in {"POST", "PUT", "PATCH", "DELETE"}:
return None
if validate_csrf_token():
return None
if request.path == "/predict":
return jsonify({"error": "CSRF 校验失败,请刷新页面后重试。"}), 400
abort(400)
+50
View File
@@ -0,0 +1,50 @@
from __future__ import annotations
import os
import secrets
from pathlib import Path
BASE_DIR = Path(__file__).resolve().parent.parent
DATA_DIR = BASE_DIR / "data"
STATIC_DIR = BASE_DIR / "static"
UPLOAD_DIR = BASE_DIR / "uploads"
IMAGE_DIR = STATIC_DIR / "images"
def env_int(name: str, default: int) -> int:
try:
return int(os.environ.get(name, str(default)))
except ValueError:
return default
def env_bool(name: str, default: bool = False) -> bool:
value = os.environ.get(name)
if value is None:
return default
return value.strip().lower() in {"1", "true", "yes", "on"}
class Config:
SECRET_KEY = os.environ.get("SECRET_KEY") or secrets.token_hex(32)
SECRET_KEY_GENERATED = not bool(os.environ.get("SECRET_KEY"))
SQLALCHEMY_DATABASE_URI = os.environ.get(
"DATABASE_URL",
f"sqlite:///{DATA_DIR / 'pipe_survival_0331.db'}",
)
SQLALCHEMY_TRACK_MODIFICATIONS = False
MAX_CONTENT_LENGTH = env_int("MAX_UPLOAD_BYTES", 16 * 1024 * 1024)
SESSION_COOKIE_HTTPONLY = True
SESSION_COOKIE_SAMESITE = "Lax"
MODEL_PATH = os.environ.get(
"MODEL_PATH",
str(BASE_DIR / "my_survival_forest_model_quxi-10-0331.joblib"),
)
ALLOW_REGISTRATION = env_bool("ALLOW_REGISTRATION", False)
ADMIN_USERNAME = os.environ.get("ADMIN_USERNAME", "admin").strip() or "admin"
ADMIN_PASSWORD = os.environ.get("ADMIN_PASSWORD")
def ensure_dirs() -> None:
for path in (DATA_DIR, STATIC_DIR, IMAGE_DIR, UPLOAD_DIR):
path.mkdir(parents=True, exist_ok=True)
+5
View File
@@ -0,0 +1,5 @@
from flask_login import LoginManager
from flask_sqlalchemy import SQLAlchemy
db = SQLAlchemy()
login_manager = LoginManager()
+38
View File
@@ -0,0 +1,38 @@
from __future__ import annotations
from datetime import datetime
from flask_login import UserMixin
from werkzeug.security import check_password_hash, generate_password_hash
from .extensions import db
class User(UserMixin, db.Model):
__tablename__ = "users"
id = db.Column(db.Integer, primary_key=True)
username = db.Column(db.String(100), unique=True, nullable=False)
password_hash = db.Column(db.String(255), nullable=False)
is_admin = db.Column(db.Boolean, default=False, nullable=False)
created_at = db.Column(db.DateTime, default=datetime.utcnow)
def set_password(self, password: str) -> None:
self.password_hash = generate_password_hash(password)
def check_password(self, password: str) -> bool:
return check_password_hash(self.password_hash, password)
class UploadRecord(db.Model):
__tablename__ = "upload_records"
id = db.Column(db.Integer, primary_key=True)
user_id = db.Column(db.Integer, db.ForeignKey("users.id"), nullable=False)
original_filename = db.Column(db.String(255), nullable=False)
saved_path = db.Column(db.String(500), nullable=False)
prediction_path = db.Column(db.String(500), nullable=False)
image_path = db.Column(db.String(500), nullable=False)
upload_time = db.Column(db.DateTime, default=datetime.utcnow)
user = db.relationship("User", backref=db.backref("uploads", lazy=True))
+380
View File
@@ -0,0 +1,380 @@
from __future__ import annotations
import logging
import os
import uuid
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any
import joblib
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from matplotlib import font_manager, rcParams
from werkzeug.datastructures import FileStorage
from werkzeug.utils import secure_filename
from .config import IMAGE_DIR, UPLOAD_DIR
CHINESE_FONT_CANDIDATES = [
"Noto Sans CJK SC",
"Noto Sans CJK JP",
"Noto Sans CJK TC",
"Source Han Sans SC",
"WenQuanYi Micro Hei",
"SimHei",
"Microsoft YaHei",
"Arial Unicode MS",
]
FEATURES = [
"管材",
"管径",
"流速",
"压力",
"温度",
"降雨量",
"位置",
"结构缺陷",
"功能缺陷",
]
ID_COLUMN = "管道编号"
PIPE_AGE_COLUMN = "管龄"
SUPPORTED_EXTENSIONS = {".csv", ".xls", ".xlsx"}
class PredictionError(Exception):
def __init__(self, message: str, status_code: int = 400) -> None:
super().__init__(message)
self.message = message
self.status_code = status_code
@dataclass
class PredictionArtifacts:
original_filename: str
saved_path: Path
excel_path: Path
image_path: Path
image_filename: str
importance_filename: str | None
sample_count: int
summary_rows: list[dict[str, Any]]
analysis_text: str
def configure_matplotlib_fonts() -> None:
available_fonts = {font.name for font in font_manager.fontManager.ttflist}
selected_fonts = [font for font in CHINESE_FONT_CANDIDATES if font in available_fonts]
rcParams["font.sans-serif"] = selected_fonts + ["DejaVu Sans"]
rcParams["axes.unicode_minus"] = False
configure_matplotlib_fonts()
def load_model(model_path: str):
if not os.path.exists(model_path):
raise FileNotFoundError(f"未找到模型文件: {model_path}")
model = joblib.load(model_path)
feature_names = getattr(model, "feature_names_in_", None)
if feature_names is not None and list(feature_names) != FEATURES:
raise ValueError(
"模型输入特征与系统配置不一致: "
f"model={list(feature_names)}, app={FEATURES}"
)
n_features = getattr(model, "n_features_in_", None)
if n_features is not None and int(n_features) != len(FEATURES):
raise ValueError(f"模型特征数量不一致: model={n_features}, app={len(FEATURES)}")
return model
def safe_unlink(path: Path) -> None:
try:
path.unlink(missing_ok=True)
except OSError as exc:
logging.warning("删除文件失败 %s: %s", path, exc)
def secure_upload_name(original_filename: str, run_id: str) -> tuple[str, str]:
suffix = Path(original_filename).suffix.lower()
if suffix not in SUPPORTED_EXTENSIONS:
raise PredictionError("不支持的格式,仅支持 CSV / XLS / XLSX")
safe_full_name = secure_filename(original_filename)
safe_stem = Path(safe_full_name).stem if safe_full_name else ""
if not safe_stem:
safe_stem = "upload"
return f"{safe_stem}_{run_id}{suffix}", suffix
def read_input_file(path: Path, suffix: str) -> pd.DataFrame:
try:
if suffix == ".csv":
return pd.read_csv(path)
return pd.read_excel(path)
except Exception as exc:
logging.exception("文件解析失败: %s", exc)
raise PredictionError("文件解析失败,请检查编码或表格格式。")
def validate_input_frame(df: pd.DataFrame) -> None:
if df.empty:
raise PredictionError("上传文件没有可预测的数据。")
required_columns = [ID_COLUMN, *FEATURES]
missing = [col for col in required_columns if col not in df.columns]
if missing:
raise PredictionError(f"缺少必要字段: {', '.join(missing)}")
def grade_info(probability: float) -> tuple[str, str, str]:
if probability <= 0.2:
return ("I级", "管道安全风险十分严重,需立刻进行抢修或更新改造", "bg-dangerSoft text-dangerText")
if probability <= 0.4:
return ("II级", "管道安全风险较为严重,需尽快安排检修及加频巡检", "bg-orange-50 text-orange-600")
if probability <= 0.6:
return ("III级", "管道安全风险较低,需安排定期巡检", "bg-amber-50 text-amber-600")
if probability <= 0.8:
return ("IV级", "管道安全风险较小,维持常规巡视", "bg-blue-50 text-blue-600")
return ("V级", "管道安全,维持常规巡视", "bg-blueSoft text-primary")
def interpolate_probability(times: list[float], probs: list[float], target: float) -> float:
if not times:
return 0.0
if target <= times[0]:
return float(probs[0])
for idx in range(1, len(times)):
if times[idx] >= target:
return float(probs[idx])
return float(probs[-1])
def estimate_remaining_life(times: list[float], probs: list[float]) -> float:
for t, p in zip(times, probs):
if p <= 0.5:
return float(t)
return float(times[-1]) if times else 0.0
def make_analysis_text(summary_rows: list[dict[str, Any]]) -> str:
if not summary_rows:
return "当前结果为空,暂无可供解释的样本。"
worst = min(summary_rows, key=lambda x: x["health_probability"])
best = max(summary_rows, key=lambda x: x["health_probability"])
return (
"阶梯状曲线表示模型对不同管道随时间推移维持在安全健康状态概率的动态预测。"
f"当前样本中风险最高管道为 {worst['pipe_id']}{worst['grade_label']}),"
f"健康概率最高管道为 {best['pipe_id']}{best['health_probability']:.1%})。"
)
def compute_feature_importance(model, x_test: pd.DataFrame) -> np.ndarray | None:
try:
importances = model.feature_importances_
except Exception:
importances = None
if importances is not None and len(importances) == len(FEATURES):
vals = np.asarray(importances, dtype=float)
else:
try:
baseline = np.asarray(model.predict(x_test), dtype=float)
except Exception as exc:
logging.exception("特征重要性基线预测失败: %s", exc)
return None
rng = np.random.default_rng(42)
vals = np.zeros(len(FEATURES), dtype=float)
for j, feat in enumerate(FEATURES):
diffs = []
for _ in range(5):
x_perm = x_test.copy()
x_perm[feat] = rng.permutation(x_perm[feat].to_numpy())
try:
perm_pred = np.asarray(model.predict(x_perm), dtype=float)
except Exception:
perm_pred = baseline
diffs.append(float(np.mean(np.abs(perm_pred - baseline))))
vals[j] = float(np.mean(diffs)) if diffs else 0.0
vals = np.clip(vals, a_min=0.0, a_max=None)
total = float(vals.sum())
if total > 0:
vals = vals / total
return vals
def render_importance_chart(values: np.ndarray, save_path: Path) -> None:
from matplotlib.colors import LinearSegmentedColormap
keep = values >= 5e-4
if not bool(np.any(keep)):
keep = np.ones_like(values, dtype=bool)
kept_feats = [FEATURES[i] for i in range(len(FEATURES)) if keep[i]]
kept_vals = values[keep]
order = np.argsort(kept_vals)
sorted_feats = [kept_feats[k] for k in order]
sorted_vals = kept_vals[order]
n = len(sorted_vals)
fig, ax = plt.subplots(figsize=(9, max(3.0, 0.62 * n + 1.6)))
cmap = LinearSegmentedColormap.from_list("brand", ["#7fb2e6", "#005EB8", "#0c4188"])
colors = cmap(np.linspace(0.15, 1.0, n)) if n else None
bars = ax.barh(sorted_feats, sorted_vals, color=colors, height=0.66, edgecolor="white", linewidth=0.8, zorder=3)
ax.set_xlabel("相对重要性", fontsize=11, color="#475569")
ax.set_title("模型输入因素重要性排序", fontsize=15, fontweight="bold", color="#0f172a", pad=14)
ax.grid(axis="x", color="#e2e8f0", linewidth=1, zorder=0)
ax.set_axisbelow(True)
for spine in ("top", "right", "left"):
ax.spines[spine].set_visible(False)
ax.spines["bottom"].set_color("#cbd5e1")
ax.tick_params(axis="y", length=0, labelsize=11)
ax.tick_params(axis="x", colors="#94a3b8", labelsize=9)
max_val = float(sorted_vals.max()) if n else 0.0
for bar, value in zip(bars, sorted_vals):
ax.text(
bar.get_width() + max_val * 0.012,
bar.get_y() + bar.get_height() / 2,
f"{value:.1%}",
va="center",
ha="left",
fontsize=10,
fontweight="bold",
color="#1e293b",
)
if max_val > 0:
ax.set_xlim(0, max_val * 1.18)
fig.tight_layout()
fig.savefig(save_path, dpi=160, bbox_inches="tight")
plt.close(fig)
def run_prediction(uploaded: FileStorage, user_id: int, model) -> PredictionArtifacts:
original_filename = uploaded.filename or ""
if not original_filename:
raise PredictionError("未选择文件")
timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
run_id = f"{timestamp}_{uuid.uuid4().hex[:8]}"
saved_filename, suffix = secure_upload_name(original_filename, run_id)
user_dir = UPLOAD_DIR / f"user_{user_id}"
user_dir.mkdir(parents=True, exist_ok=True)
original_path = user_dir / saved_filename
uploaded.save(original_path)
try:
df = read_input_file(original_path, suffix)
validate_input_frame(df)
x_test = df[FEATURES].copy()
try:
curves = model.predict_survival_function(x_test)
except Exception as exc:
logging.exception("预测失败: %s", exc)
raise PredictionError("模型预测失败,请检查输入字段类型是否正确。", 500)
image_filename = f"plot_{user_id}_{run_id}.png"
image_path = IMAGE_DIR / image_filename
summary_rows, summary_sheet_rows = render_survival_chart(df, curves, image_path)
importance_filename = None
try:
importance_values = compute_feature_importance(model, x_test)
if importance_values is not None:
importance_filename = f"importance_{user_id}_{run_id}.png"
render_importance_chart(importance_values, IMAGE_DIR / importance_filename)
except Exception as exc:
logging.exception("生成特征重要性图失败: %s", exc)
safe_stem = Path(saved_filename).stem
excel_path = user_dir / f"{safe_stem}_pre.xlsx"
write_prediction_workbook(excel_path, curves, summary_rows, summary_sheet_rows)
return PredictionArtifacts(
original_filename=original_filename,
saved_path=original_path,
excel_path=excel_path,
image_path=image_path,
image_filename=image_filename,
importance_filename=importance_filename,
sample_count=len(summary_rows),
summary_rows=summary_rows,
analysis_text=make_analysis_text(summary_rows),
)
except PredictionError:
safe_unlink(original_path)
raise
def render_survival_chart(df: pd.DataFrame, curves, image_path: Path) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
plt.figure(figsize=(10, 5.6))
summary_rows: list[dict[str, Any]] = []
summary_sheet_rows: list[dict[str, Any]] = []
for i, curve in enumerate(curves):
times = [float(x) for x in list(curve.x)]
probs = [float(y) for y in list(curve.y)]
pipe_id = str(df.iloc[i][ID_COLUMN]) if pd.notna(df.iloc[i][ID_COLUMN]) else f"Pipe_{i+1:03d}"
pipe_age = f"{df.iloc[i][PIPE_AGE_COLUMN]}" if PIPE_AGE_COLUMN in df.columns and pd.notna(df.iloc[i][PIPE_AGE_COLUMN]) else "-"
health_probability = interpolate_probability(times, probs, 10)
remaining_life = estimate_remaining_life(times, probs)
grade_label, grade_desc, grade_class = grade_info(health_probability)
summary_rows.append(
{
"pipe_id": pipe_id,
"pipe_age": pipe_age,
"health_probability": health_probability,
"remaining_life": remaining_life,
"grade_label": grade_label,
"grade_desc": grade_desc,
"grade_class": grade_class,
}
)
summary_sheet_rows.append(
{
"管道编号": pipe_id,
"管龄": pipe_age,
"健康概率": health_probability,
"预计剩余寿命": remaining_life,
"健康等级": grade_label,
}
)
plt.step(times, probs, where="post", linewidth=2, label=pipe_id)
plt.xlabel("预测时间轴(年)")
plt.ylabel("生存概率")
plt.title("预测分析图")
plt.grid(alpha=0.18)
if len(summary_rows) <= 12:
plt.legend(loc="best", fontsize=8)
plt.tight_layout()
plt.savefig(image_path, dpi=160, bbox_inches="tight")
plt.close()
return summary_rows, summary_sheet_rows
def write_prediction_workbook(
excel_path: Path,
curves,
summary_rows: list[dict[str, Any]],
summary_sheet_rows: list[dict[str, Any]],
) -> None:
with pd.ExcelWriter(excel_path, engine="xlsxwriter") as writer:
pd.DataFrame(summary_sheet_rows).to_excel(writer, sheet_name="结果摘要", index=False)
for i, curve in enumerate(curves):
times = [float(x) for x in list(curve.x)]
probs = [float(y) for y in list(curve.y)]
pipe_id = summary_rows[i]["pipe_id"]
out_df = pd.DataFrame({"时间(年)": times, f"{pipe_id}生存概率": probs})
out_df.to_excel(writer, sheet_name=f"样本{i+1}", index=False)
+215
View File
@@ -0,0 +1,215 @@
from __future__ import annotations
import os
from datetime import datetime
from flask import (
Blueprint,
abort,
current_app,
flash,
jsonify,
redirect,
render_template,
request,
send_file,
session,
url_for,
)
from flask_login import current_user, login_required, login_user, logout_user
from .config import BASE_DIR
from .extensions import db
from .models import UploadRecord, User
from .prediction import PredictionError, run_prediction
from .security import new_captcha
bp = Blueprint("main", __name__)
@bp.route("/")
def index():
if current_user.is_authenticated:
return redirect(url_for("main.home"))
return redirect(url_for("main.login"))
@bp.route("/login", methods=["GET", "POST"])
def login():
if current_user.is_authenticated:
return redirect(url_for("main.home"))
if request.method == "GET":
session["captcha"] = new_captcha()
return render_template("login.html", mode="login", captcha=session["captcha"])
username = request.form.get("username", "").strip()
password = request.form.get("password", "")
captcha_input = request.form.get("captcha", "").strip().upper()
if captcha_input != session.get("captcha", ""):
flash("验证码错误", "error")
session["captcha"] = new_captcha()
return render_template("login.html", mode="login", captcha=session["captcha"]), 400
user = User.query.filter_by(username=username).first()
if not user or not user.check_password(password):
flash("用户名或密码错误", "error")
session["captcha"] = new_captcha()
return render_template("login.html", mode="login", captcha=session["captcha"]), 400
login_user(user, remember=bool(request.form.get("remember")))
return redirect(url_for("main.home"))
@bp.route("/register", methods=["GET", "POST"])
def register():
if not current_app.config["ALLOW_REGISTRATION"]:
abort(404)
if request.method == "GET":
return render_template("login.html", mode="register", captcha="")
username = request.form.get("username", "").strip()
password = request.form.get("password", "")
if not username:
flash("用户名不能为空", "error")
return render_template("login.html", mode="register", captcha=""), 400
if len(password) < 6:
flash("密码至少需要 6 位", "error")
return render_template("login.html", mode="register", captcha=""), 400
if User.query.filter_by(username=username).first():
flash("用户名已存在", "error")
return render_template("login.html", mode="register", captcha=""), 400
user = User(username=username, is_admin=False)
user.set_password(password)
db.session.add(user)
db.session.commit()
flash("注册成功,请登录", "info")
session["captcha"] = new_captcha()
return render_template("login.html", mode="login", captcha=session["captcha"])
@bp.route("/logout", methods=["POST"])
@login_required
def logout():
logout_user()
return redirect(url_for("main.login"))
@bp.route("/home")
@login_required
def home():
return render_template("home.html")
@bp.route("/history")
@login_required
def history_page():
records = (
UploadRecord.query.filter_by(user_id=current_user.id)
.order_by(UploadRecord.upload_time.desc())
.all()
)
return render_template("history.html", records=records)
@bp.route("/admin")
@login_required
def admin_dashboard():
if not current_user.is_admin:
abort(403)
records = UploadRecord.query.order_by(UploadRecord.upload_time.desc()).all()
return render_template("admin.html", records=records)
@bp.route("/download/<int:record_id>/<file_type>")
@login_required
def download_file(record_id: int, file_type: str):
record = db.session.get(UploadRecord, record_id)
if record is None:
abort(404)
if not (current_user.is_admin or current_user.id == record.user_id):
abort(403)
if file_type == "original":
file_path = record.saved_path
elif file_type == "prediction":
file_path = record.prediction_path
else:
abort(404)
if not os.path.exists(file_path):
abort(404)
return send_file(file_path, as_attachment=True)
@bp.route("/download_template")
def download_template():
template_path = BASE_DIR / "example.xlsx"
if not template_path.exists():
abort(404)
return send_file(template_path, as_attachment=True, download_name="example.xlsx")
@bp.route("/result")
@login_required
def result_page():
result = session.get("last_result")
return render_template("result.html", result=result)
@bp.route("/predict", methods=["POST"])
@login_required
def predict():
model = current_app.config.get("RSF_MODEL")
if model is None:
return jsonify({"error": "模型未成功加载,请检查模型文件。"}), 500
uploaded = request.files.get("file")
if uploaded is None or uploaded.filename == "":
return jsonify({"error": "未选择文件"}), 400
try:
artifacts = run_prediction(uploaded, int(current_user.id), model)
except PredictionError as exc:
return jsonify({"error": exc.message}), exc.status_code
record = UploadRecord(
user_id=current_user.id,
original_filename=artifacts.original_filename,
saved_path=str(artifacts.saved_path),
prediction_path=str(artifacts.excel_path),
image_path=str(artifacts.image_path),
)
db.session.add(record)
db.session.commit()
last_result = {
"original_filename": artifacts.original_filename,
"generated_at": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"image_url": url_for("static", filename=f"images/{artifacts.image_filename}"),
"importance_url": (
url_for("static", filename=f"images/{artifacts.importance_filename}")
if artifacts.importance_filename
else None
),
"excel_url": url_for("main.download_file", record_id=record.id, file_type="prediction"),
"result_url": url_for("main.result_page"),
"sample_count": int(artifacts.sample_count),
"summary_rows": artifacts.summary_rows[:3],
"analysis_text": artifacts.analysis_text,
}
session["last_result"] = last_result
return jsonify(
{
"message": "预测成功",
"image_url": last_result["image_url"],
"importance_url": last_result["importance_url"],
"excel_url": last_result["excel_url"],
"result_url": last_result["result_url"],
"sample_count": last_result["sample_count"],
"original_filename": artifacts.original_filename,
}
)
+24
View File
@@ -0,0 +1,24 @@
from __future__ import annotations
import secrets
from flask import request, session
def csrf_token() -> str:
token = session.get("_csrf_token")
if not token:
token = secrets.token_urlsafe(32)
session["_csrf_token"] = token
return token
def validate_csrf_token() -> bool:
expected = session.get("_csrf_token", "")
received = request.form.get("csrf_token") or request.headers.get("X-CSRFToken", "")
return bool(expected and received and secrets.compare_digest(expected, received))
def new_captcha() -> str:
alphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
return "".join(secrets.choice(alphabet) for _ in range(5))