refactor: split flask app structure
This commit is contained in:
+116
@@ -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)
|
||||
@@ -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)
|
||||
@@ -0,0 +1,5 @@
|
||||
from flask_login import LoginManager
|
||||
from flask_sqlalchemy import SQLAlchemy
|
||||
|
||||
db = SQLAlchemy()
|
||||
login_manager = LoginManager()
|
||||
@@ -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))
|
||||
@@ -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
@@ -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,
|
||||
}
|
||||
)
|
||||
@@ -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))
|
||||
Reference in New Issue
Block a user