diff --git a/.env.example b/.env.example index 4101cb6..88bcfed 100644 --- a/.env.example +++ b/.env.example @@ -2,9 +2,15 @@ # Generate SECRET_KEY with: python -c "import secrets; print(secrets.token_hex(32))" SECRET_KEY= -# Set ADMIN_PASSWORD to create or rotate the admin account on startup. +# Required for first deployment. Used to create or rotate the admin account on startup. ADMIN_USERNAME=admin ADMIN_PASSWORD= # Default: 16 MiB MAX_UPLOAD_BYTES=16777216 + +# Default model path inside the Docker image. +MODEL_PATH=/app/my_survival_forest_model_quxi-10-0331.joblib + +# Keep public registration closed by default. +ALLOW_REGISTRATION=false diff --git a/Dockerfile b/Dockerfile index b716229..6b2e211 100644 --- a/Dockerfile +++ b/Dockerfile @@ -17,12 +17,15 @@ RUN conda create -n demo python=3.12 -y \ && conda run -n demo python -m pip install --no-cache-dir -r requirements.txt \ && conda clean -afy -COPY final_flask_app_0331_strict_fixed_F.py . -COPY my_survival_forest_model_quxi-10-0331.joblib . +COPY app ./app +COPY templates ./templates +COPY static ./static +COPY main.py . COPY example.xlsx . +COPY my_survival_forest_model_quxi-10-0331.joblib . RUN mkdir -p data static/images uploads EXPOSE 5005 -CMD ["conda", "run", "--no-capture-output", "-n", "demo", "python", "final_flask_app_0331_strict_fixed_F.py"] +CMD ["conda", "run", "--no-capture-output", "-n", "demo", "python", "main.py"] diff --git a/app/__init__.py b/app/__init__.py new file mode 100644 index 0000000..cb0289c --- /dev/null +++ b/app/__init__.py @@ -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) diff --git a/app/config.py b/app/config.py new file mode 100644 index 0000000..77734d5 --- /dev/null +++ b/app/config.py @@ -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) diff --git a/app/extensions.py b/app/extensions.py new file mode 100644 index 0000000..e0663ee --- /dev/null +++ b/app/extensions.py @@ -0,0 +1,5 @@ +from flask_login import LoginManager +from flask_sqlalchemy import SQLAlchemy + +db = SQLAlchemy() +login_manager = LoginManager() diff --git a/app/models.py b/app/models.py new file mode 100644 index 0000000..5e5eb98 --- /dev/null +++ b/app/models.py @@ -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)) diff --git a/app/prediction.py b/app/prediction.py new file mode 100644 index 0000000..af7e0a8 --- /dev/null +++ b/app/prediction.py @@ -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) diff --git a/app/routes.py b/app/routes.py new file mode 100644 index 0000000..9c643e0 --- /dev/null +++ b/app/routes.py @@ -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//") +@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, + } + ) diff --git a/app/security.py b/app/security.py new file mode 100644 index 0000000..2c7e6f0 --- /dev/null +++ b/app/security.py @@ -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)) diff --git a/docker-compose.yml b/docker-compose.yml index 8ad1989..e2b9ab3 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,13 +1,17 @@ services: pipeline-lifetime: + build: + context: . image: pipeline-lifetime:latest container_name: pipeline-lifetime restart: unless-stopped environment: - SECRET_KEY: ${SECRET_KEY:-} + SECRET_KEY: ${SECRET_KEY:?Set SECRET_KEY in .env} ADMIN_USERNAME: ${ADMIN_USERNAME:-admin} - ADMIN_PASSWORD: ${ADMIN_PASSWORD:-} + ADMIN_PASSWORD: ${ADMIN_PASSWORD:?Set ADMIN_PASSWORD in .env} MAX_UPLOAD_BYTES: ${MAX_UPLOAD_BYTES:-16777216} + MODEL_PATH: ${MODEL_PATH:-/app/my_survival_forest_model_quxi-10-0331.joblib} + ALLOW_REGISTRATION: ${ALLOW_REGISTRATION:-false} ports: - "5005:5005" volumes: diff --git a/main.py b/main.py index 4e632b4..47f8895 100644 --- a/main.py +++ b/main.py @@ -1,1389 +1,7 @@ -# -*- coding: utf-8 -*- -""" +from app import create_app - -运行要求: -1. 同目录放置 my_survival_forest_model_quxi-10-0331.joblib -2. pip install flask flask_sqlalchemy flask_login pandas joblib matplotlib openpyxl xlsxwriter xlrd -3. python final_flask_app_0331_strict.py -""" -from __future__ import annotations - -import logging -import os -import secrets -import sys -import uuid -from dataclasses import dataclass -from datetime import datetime -from io import BytesIO -from pathlib import Path -from typing import Any, Dict, List - -import joblib -import matplotlib -matplotlib.use("Agg") -import matplotlib.pyplot as plt -import numpy as np -import pandas as pd -from flask import ( - Flask, - abort, - flash, - jsonify, - redirect, - render_template_string, - request, - send_file, - session, - url_for, -) -from flask_login import ( - LoginManager, - UserMixin, - current_user, - login_required, - login_user, - logout_user, -) -from flask_sqlalchemy import SQLAlchemy -from matplotlib import font_manager, rcParams -from werkzeug.security import check_password_hash, generate_password_hash -from werkzeug.utils import secure_filename - -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", -] - - -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"] - - -configure_matplotlib_fonts() -rcParams["axes.unicode_minus"] = False - -BASE_DIR = Path(__file__).resolve().parent -DATA_DIR = BASE_DIR / "data" -STATIC_DIR = BASE_DIR / "static" -UPLOAD_DIR = BASE_DIR / "uploads" -IMAGE_DIR = STATIC_DIR / "images" - -FEATURES = [ - "管材", "管径", "流速", "压力", - "温度", "降雨量", "位置", - "结构缺陷", "功能缺陷", -] - -ID_COLUMN = "管道编号" -PIPE_AGE_COLUMN = "管龄" - -TEMPLATE_COLUMNS = [ - ID_COLUMN, - "管龄", "状态", "管材", "管径", "流速", - "压力", "温度", "降雨量", "位置", - "结构缺陷", "功能缺陷", -] - -DATA_DIR.mkdir(parents=True, exist_ok=True) - -logging.basicConfig( - filename=str(DATA_DIR / "app.log"), - level=logging.INFO, - format="%(asctime)s - %(levelname)s - %(message)s", -) - -app = Flask(__name__) - - -def env_int(name: str, default: int) -> int: - try: - return int(os.environ.get(name, str(default))) - except ValueError: - logging.warning("%s 配置无效,使用默认值 %s", name, default) - return default - - -secret_key = os.environ.get("SECRET_KEY") -if not secret_key: - secret_key = secrets.token_hex(32) - logging.warning("未设置 SECRET_KEY,已生成临时密钥;服务重启后登录会话将失效。") - -app.config["SECRET_KEY"] = secret_key -app.config["SQLALCHEMY_DATABASE_URI"] = f"sqlite:///{DATA_DIR / 'pipe_survival_0331.db'}" -app.config["SQLALCHEMY_TRACK_MODIFICATIONS"] = False -app.config["UPLOAD_FOLDER"] = str(UPLOAD_DIR) -app.config["MAX_CONTENT_LENGTH"] = env_int("MAX_UPLOAD_BYTES", 16 * 1024 * 1024) -app.config["SESSION_COOKIE_HTTPONLY"] = True -app.config["SESSION_COOKIE_SAMESITE"] = "Lax" - -db = SQLAlchemy(app) -login_manager = LoginManager(app) -login_manager.login_view = "login" -login_manager.login_message = "请先登录后再访问该页面。" - - -@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 - - -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)) - - -@login_manager.user_loader -def load_user(user_id: str): - return db.session.get(User, int(user_id)) - - -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)) - - -@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) - - -@dataclass -class PipeSummary: - pipe_id: str - pipe_age: str - grade_label: str - - -def resource_path(relative_path: str) -> str: - try: - base_path = Path(sys._MEIPASS) # type: ignore[attr-defined] - except Exception: - base_path = BASE_DIR - return str((base_path / relative_path).resolve()) - - -def safe_unlink(path: Path) -> None: - try: - path.unlink(missing_ok=True) - except OSError as exc: - logging.warning("删除文件失败 %s: %s", path, exc) - - -def ensure_dirs() -> None: - for path in [DATA_DIR, STATIC_DIR, IMAGE_DIR, UPLOAD_DIR]: - path.mkdir(parents=True, exist_ok=True) - - -def init_app() -> None: - ensure_dirs() - with app.app_context(): - db.create_all() - admin_username = os.environ.get("ADMIN_USERNAME", "admin").strip() or "admin" - admin_password = os.environ.get("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 更新。") - - try: - app.config["RSF_MODEL"] = load_model() - logging.info("模型加载成功") - except Exception as exc: - app.config["RSF_MODEL"] = None - logging.exception("模型加载失败: %s", exc) - - -def load_model(): - model_path = resource_path("my_survival_forest_model_quxi-10-0331.joblib") - if not os.path.exists(model_path): - raise FileNotFoundError(f"未找到模型文件: {model_path}") - return joblib.load(model_path) - - -@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} - - -BASE_TEMPLATE_HEAD = r""" - - - - - - - -""" - -LOGIN_TEMPLATE = r""" - - - - {{ '注册' if mode == 'register' else '登录' }} | - """ + BASE_TEMPLATE_HEAD + r""" - - -
- - -
-
-
-

系统门户

- -
- 登录 - 注册 -
- - {% with messages = get_flashed_messages(with_categories=true) %} - {% if messages %} -
- {% for category, message in messages %} -
{{ message }}
- {% endfor %} -
- {% endif %} - {% endwith %} - - {% if mode == 'login' %} -
- -
- -
- person - -
-
- -
-
- - 找回密码 -
-
- lock - -
-
- -
- -
-
- verified_user - -
-
{{ captcha }}
- - refresh - -
-
- - - - -
- {% else %} -
- -
- -
- badge - -
-
-
- -
- lock - -
-
- -
- {% endif %} -
-
- -
-
- 系统状态 - 服务条款 - API 文档 -
-
-
-
- - -""" - -DASHBOARD_TEMPLATE = r""" - - - - - """ + BASE_TEMPLATE_HEAD + r""" - - -
-
-
-
-
- water_drop - 供水管道健康评估系统 -
- -
-
- notifications - help -
-
-
{{ current_user.username[:1]|upper }}
-
- - -
-
-
-
-
-
- -
-
-

供水管道健康状态与剩余寿命评估技术导则

-

上传您的数据以生成预测结果供参考。

-
- - {% with messages = get_flashed_messages(with_categories=true) %} - {% if messages %} -
- {% for category, message in messages %} -
{{ message }}
- {% endfor %} -
- {% endif %} - {% endwith %} - - -
-
-
-
-

- upload_file - 文件上传 -

- - download - 下载模板 - -
- -
- - - -
- -
-
-
- - -
- - - - -
-
- -
- 预测结果仅供参考 -
- 文档 -
-
- - - - -""" - -RESULT_TEMPLATE = r""" - - - - 预测结果 - """ + BASE_TEMPLATE_HEAD + r""" - - -
-
-
-
-
- water_drop - 供水管道健康评估系统 -
- -
-
- notifications - help -
-
-
{{ current_user.username[:1]|upper }}
-
- - -
-
-
-
-
-
- -
- {% if result %} -
-
-
- description -
-
-
数据分析报告: {{ result.original_filename }}
-
基于上传数据生成的实时分析报告 • 生成于: {{ result.generated_at }}
-
-
-
-
check_circle数据源已验证
- 返回主页 -
-
- - - -
-
-
-

供水管道健康状态评估等级

- -
- -
-
I级
(0, 0.2]
管道安全风险十分严重,需立刻进行抢修或更新改造
-
II级
(0.2, 0.4]
管道安全风险较为严重,需尽快安排检修及加频巡检
-
III级
(0.4, 0.6]
管道安全风险较低,需安排定期巡检
-
IV级
(0.6, 0.8]
管道安全风险较小,维持常规巡视
-
V级
(0.8, 1]
管道安全,维持常规巡视
-
- -
- - - - - - - - - {% for item in result.summary_rows %} - - - - - {% endfor %} - -
管道编号健康等级
{{ item.pipe_id }} - {{ item.grade_label }} -
-
- 共 {{ result.sample_count }} 个样本,显示 {{ result.summary_rows|length }} 条 - 查看完整样本列表 -
-
-
- -
-
-
-

管道剩余寿命动态评估

-

生存曲线拟合

-
-
- -
- 生存概率阶梯图 -
- -
- info -

分析说明:{{ result.analysis_text }}

-
-
-
- - {% if result.importance_url %} -
-
-
-

模型输入因素重要性排序

-

各输入因素对预测结果的相对影响程度

-
-
-
- 模型输入因素重要性排序图 -
-
- insights -

说明:柱状图按重要性从高到低展示各输入因素对模型预测结果的相对贡献,数值为归一化占比,可用于辅助识别影响管道健康状态的关键因素。

-
-
- {% endif %} - {% else %} -
当前还没有预测结果,请先从主页上传文件并运行预测。
- {% endif %} -
- -
- 预测结果仅供参考 -
-
- - -""" - -ADMIN_TEMPLATE = r""" - - - - 管理台 - """ + BASE_TEMPLATE_HEAD + r""" - - -
-
-

管理员查看上传记录

- 返回主页 -
-
- - - - - - - - - - - {% for record in records %} - - - - - - - {% else %} - - {% endfor %} - -
用户原始文件上传时间下载
{{ record.user.username }}{{ record.original_filename }}{{ record.upload_time.strftime('%Y-%m-%d %H:%M:%S') }} - 原始文件 - 预测结果 -
暂无上传记录
-
-
- - -""" - -HISTORY_TEMPLATE = r""" - - - - 预测历史 - """ + BASE_TEMPLATE_HEAD + r""" - - -
-
-

预测历史

- 返回主页 -
-
- {% for record in records %} -
-
-
{{ record.original_filename }}
-
{{ record.upload_time.strftime('%Y-%m-%d %H:%M:%S') }}
-
- -
- {% else %} -
还没有历史记录。
- {% endfor %} -
-
- - -""" - - -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"阶梯状曲线表示模型对不同管道随时间推移维持在安全健康状态概率的动态预测。" - ) - - -def compute_feature_importance(model, x_test: "pd.DataFrame") -> "np.ndarray | None": - """返回各输入因素的相对重要性(与 FEATURES 顺序一致,归一化为占比)。 - - 优先使用模型自带的 feature_importances_;若不可用,则采用与标签无关的 - 置换重要性:打乱单个特征后观察模型风险评分的平均变化幅度。 - """ - 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) - n_repeats = 5 - vals = np.zeros(len(FEATURES), dtype=float) - for j, feat in enumerate(FEATURES): - diffs = [] - for _ in range(n_repeats): - 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 - - # 仅保留有实际贡献的因素(占比 >= 0.05%,即不会显示为 0.0%);若全部过低则回退展示全部 - 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) - - height = max(3.0, 0.62 * n + 1.6) - fig, ax = plt.subplots(figsize=(9, height)) - - 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, v in zip(bars, sorted_vals): - ax.text(bar.get_width() + max_val * 0.012, - bar.get_y() + bar.get_height() / 2, - f"{v:.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) - - -@app.route("/") -def index(): - if current_user.is_authenticated: - return redirect(url_for("home")) - return redirect(url_for("login")) - - -@app.route("/login", methods=["GET", "POST"]) -def login(): - if current_user.is_authenticated: - return redirect(url_for("home")) - - if request.method == "GET": - session["captcha"] = new_captcha() - return render_template_string(LOGIN_TEMPLATE, 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_string(LOGIN_TEMPLATE, 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_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"]), 400 - - login_user(user, remember=bool(request.form.get("remember"))) - return redirect(url_for("home")) - - -@app.route("/register", methods=["GET", "POST"]) -def register(): - if request.method == "GET": - return render_template_string(LOGIN_TEMPLATE, mode="register", captcha="") - - username = request.form.get("username", "").strip() - password = request.form.get("password", "") - - if not username: - flash("用户名不能为空", "error") - return render_template_string(LOGIN_TEMPLATE, mode="register", captcha=""), 400 - if len(password) < 6: - flash("密码至少需要 6 位", "error") - return render_template_string(LOGIN_TEMPLATE, mode="register", captcha=""), 400 - if User.query.filter_by(username=username).first(): - flash("用户名已存在", "error") - return render_template_string(LOGIN_TEMPLATE, 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_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"]) - - -@app.route("/logout", methods=["POST"]) -@login_required -def logout(): - logout_user() - return redirect(url_for("login")) - - -@app.route("/home") -@login_required -def home(): - return render_template_string(DASHBOARD_TEMPLATE) - - -@app.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_string(HISTORY_TEMPLATE, records=records) - - -@app.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_string(ADMIN_TEMPLATE, records=records) - - -@app.route("/download//") -@login_required -def download_file(record_id: int, file_type: str): - record = UploadRecord.query.get_or_404(record_id) - 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) - - -@app.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") - - -@app.route("/result") -@login_required -def result_page(): - result = session.get("last_result") - return render_template_string(RESULT_TEMPLATE, result=result) - - -@app.route("/predict", methods=["POST"]) -@login_required -def predict(): - model = app.config.get("RSF_MODEL") - if model is None: - return jsonify({"error": "模型未成功加载,请检查 my_survival_forest_model_quxi-10-0331.joblib 文件。"}), 500 - - uploaded = request.files.get("file") - if uploaded is None or uploaded.filename == "": - return jsonify({"error": "未选择文件"}), 400 - - filename = secure_filename(uploaded.filename) - if not filename: - return jsonify({"error": "文件名无效"}), 400 - if not filename.lower().endswith((".csv", ".xls", ".xlsx")): - return jsonify({"error": "不支持的格式,仅支持 CSV / XLS / XLSX"}), 400 - - user_dir = UPLOAD_DIR / f"user_{current_user.id}" - user_dir.mkdir(parents=True, exist_ok=True) - timestamp = datetime.now().strftime("%Y%m%d%H%M%S") - run_id = f"{timestamp}_{uuid.uuid4().hex[:8]}" - stem, ext = os.path.splitext(filename) - original_path = user_dir / f"{stem}_{run_id}{ext}" - uploaded.save(original_path) - - try: - if ext.lower() == ".csv": - df = pd.read_csv(original_path) - else: - df = pd.read_excel(original_path) - except Exception as exc: - logging.exception("文件解析失败: %s", exc) - safe_unlink(original_path) - return jsonify({"error": "文件解析失败,请检查编码或表格格式。"}), 400 - - required_columns = [ID_COLUMN, *FEATURES] - missing = [col for col in required_columns if col not in df.columns] - if missing: - safe_unlink(original_path) - return jsonify({"error": f"缺少必要字段: {', '.join(missing)}"}), 400 - - x_test = df[FEATURES].copy() - try: - curves = model.predict_survival_function(x_test) - except Exception as exc: - logging.exception("预测失败: %s", exc) - safe_unlink(original_path) - return jsonify({"error": "模型预测失败,请检查输入字段类型是否正确。"}), 500 - - 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() - - image_filename = f"plot_{current_user.id}_{run_id}.png" - image_path = IMAGE_DIR / image_filename - plt.savefig(image_path, dpi=160, bbox_inches="tight") - plt.close() - - importance_filename = None - try: - importance_values = compute_feature_importance(model, x_test) - if importance_values is not None: - importance_filename = f"importance_{current_user.id}_{run_id}.png" - render_importance_chart(importance_values, IMAGE_DIR / importance_filename) - except Exception as exc: - logging.exception("生成特征重要性图失败: %s", exc) - importance_filename = None - - excel_filename = f"{stem}_pre_{run_id}.xlsx" - excel_path = user_dir / excel_filename - 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) - - record = UploadRecord( - user_id=current_user.id, - original_filename=filename, - saved_path=str(original_path), - prediction_path=str(excel_path), - image_path=str(image_path), - ) - db.session.add(record) - db.session.commit() - - analysis_text = make_analysis_text(summary_rows) - last_result = { - "original_filename": filename, - "generated_at": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), - "image_url": url_for("static", filename=f"images/{image_filename}"), - "importance_url": url_for("static", filename=f"images/{importance_filename}") if importance_filename else None, - "excel_url": url_for("download_file", record_id=record.id, file_type="prediction"), - "result_url": url_for("result_page"), - "sample_count": int(len(summary_rows)), - "summary_rows": summary_rows[:3], - "analysis_text": 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": filename, - }) +app = create_app() if __name__ == "__main__": - init_app() app.run(host="0.0.0.0", port=5005, debug=False, use_reloader=False) diff --git a/static/js/dashboard.js b/static/js/dashboard.js new file mode 100644 index 0000000..fb8e010 --- /dev/null +++ b/static/js/dashboard.js @@ -0,0 +1,89 @@ +const form = document.getElementById('predictForm'); +const fileInput = document.getElementById('fileInput'); +const dropZone = document.getElementById('dropZone'); +const selectedFileName = document.getElementById('selectedFileName'); +const submitBtn = document.getElementById('submitBtn'); +const submitText = document.getElementById('submitText'); +let submitIcon = document.getElementById('submitIcon'); +const alertBox = document.getElementById('alertBox'); +const mainGrid = document.getElementById('mainGrid'); +const inlineResult = document.getElementById('inlineResult'); +const resultImage = document.getElementById('resultImage'); +const resultImportanceWrap = document.getElementById('resultImportanceWrap'); +const resultImportanceImage = document.getElementById('resultImportanceImage'); +const excelBtn = document.getElementById('excelBtn'); +const resultPageBtn = document.getElementById('resultPageBtn'); +const summaryFilename = document.getElementById('summaryFilename'); +const summaryCount = document.getElementById('summaryCount'); + +function showAlert(message, type='error') { + alertBox.classList.remove('hidden', 'bg-dangerSoft', 'text-dangerText', 'border-red-200', 'bg-blueSoft', 'text-primary', 'border-blue-200'); + if (type === 'error') { + alertBox.classList.add('bg-dangerSoft', 'text-dangerText', 'border-red-200'); + } else { + alertBox.classList.add('bg-blueSoft', 'text-primary', 'border-blue-200'); + } + alertBox.textContent = message; +} + +fileInput.addEventListener('change', () => { + const file = fileInput.files[0]; + if (!file) return; + selectedFileName.textContent = '已选择文件:' + file.name; + selectedFileName.classList.remove('hidden'); +}); +['dragenter', 'dragover'].forEach(evt => dropZone.addEventListener(evt, e => { + e.preventDefault(); + dropZone.classList.add('border-primary', 'bg-blue-50'); +})); +['dragleave', 'drop'].forEach(evt => dropZone.addEventListener(evt, e => { + e.preventDefault(); + dropZone.classList.remove('border-primary', 'bg-blue-50'); +})); + +form.addEventListener('submit', async (e) => { + e.preventDefault(); + alertBox.classList.add('hidden'); + if (!fileInput.files.length) { + showAlert('请先选择要上传的文件。'); + return; + } + + submitBtn.disabled = true; + submitText.textContent = '预测中...'; + submitIcon.replaceWith(Object.assign(document.createElement('span'), { id: 'submitSpinner', className: 'spinner' })); + + try { + const formData = new FormData(form); + const resp = await fetch(form.action, { method: 'POST', body: formData }); + const data = await resp.json(); + if (!resp.ok) { + showAlert(data.error || '预测失败,请稍后重试。'); + return; + } + resultImage.src = data.image_url; + if (data.importance_url) { + resultImportanceImage.src = data.importance_url; + resultImportanceWrap.classList.remove('hidden'); + } else { + resultImportanceWrap.classList.add('hidden'); + } + excelBtn.href = data.excel_url; + resultPageBtn.href = data.result_url; + summaryFilename.textContent = data.original_filename; + summaryCount.textContent = data.sample_count; + inlineResult.classList.remove('hidden'); + mainGrid.classList.remove('xl:grid-cols-[1.45fr_1fr]'); + mainGrid.classList.add('xl:grid-cols-[1.25fr_0.95fr_1.05fr]'); + showAlert('预测完成,已生成图表与 Excel 报告。', 'success'); + inlineResult.scrollIntoView({ behavior: 'smooth', block: 'nearest' }); + } catch (err) { + showAlert('请求失败,请检查后端服务是否正常。'); + } finally { + submitBtn.disabled = false; + submitText.textContent = '分析并预测'; + const restoredIcon = Object.assign(document.createElement('span'), { id: 'submitIcon', className: 'material-symbols-outlined', textContent: 'analytics' }); + document.getElementById('submitSpinner')?.replaceWith(restoredIcon); + submitIcon = restoredIcon; + } +}); diff --git a/templates/admin.html b/templates/admin.html new file mode 100644 index 0000000..df6b672 --- /dev/null +++ b/templates/admin.html @@ -0,0 +1,110 @@ + + + + + 管理台 + + + + + + + + + + + +
+
+

管理员查看上传记录

+ 返回主页 +
+
+ + + + + + + + + + + {% for record in records %} + + + + + + + {% else %} + + {% endfor %} + +
用户原始文件上传时间下载
{{ record.user.username }}{{ record.original_filename }}{{ record.upload_time.strftime('%Y-%m-%d %H:%M:%S') }} + 原始文件 + 预测结果 +
暂无上传记录
+
+
+ + diff --git a/templates/history.html b/templates/history.html new file mode 100644 index 0000000..e543c40 --- /dev/null +++ b/templates/history.html @@ -0,0 +1,99 @@ + + + + + 预测历史 + + + + + + + + + + + +
+
+

预测历史

+ 返回主页 +
+
+ {% for record in records %} +
+
+
{{ record.original_filename }}
+
{{ record.upload_time.strftime('%Y-%m-%d %H:%M:%S') }}
+
+ +
+ {% else %} +
还没有历史记录。
+ {% endfor %} +
+
+ + diff --git a/templates/home.html b/templates/home.html new file mode 100644 index 0000000..8d64819 --- /dev/null +++ b/templates/home.html @@ -0,0 +1,252 @@ + + + + + + + + + + + + + + + + +
+
+
+
+
+ water_drop + 供水管道健康评估系统 +
+ +
+
+ notifications + help +
+
+
{{ current_user.username[:1]|upper }}
+
+ + +
+
+
+
+
+
+ +
+
+

供水管道健康状态与剩余寿命评估技术导则

+

上传您的数据以生成预测结果供参考。

+
+ + {% with messages = get_flashed_messages(with_categories=true) %} + {% if messages %} +
+ {% for category, message in messages %} +
{{ message }}
+ {% endfor %} +
+ {% endif %} + {% endwith %} + + +
+
+
+
+

+ upload_file + 文件上传 +

+ + download + 下载模板 + +
+ +
+ + + +
+ +
+
+
+ + +
+ + + + +
+
+ +
+ 预测结果仅供参考 +
+ 文档 +
+
+ + + + + diff --git a/templates/login.html b/templates/login.html new file mode 100644 index 0000000..0bca6d1 --- /dev/null +++ b/templates/login.html @@ -0,0 +1,197 @@ + + + + + {{ '注册' if mode == 'register' else '登录' }} | + + + + + + + + + + + +
+ + +
+
+
+

系统门户

+ +
+ 登录 + {% if allow_registration %} + 注册 + {% endif %} +
+ + {% with messages = get_flashed_messages(with_categories=true) %} + {% if messages %} +
+ {% for category, message in messages %} +
{{ message }}
+ {% endfor %} +
+ {% endif %} + {% endwith %} + + {% if mode == 'login' %} +
+ +
+ +
+ person + +
+
+ +
+
+ + 找回密码 +
+
+ lock + +
+
+ +
+ +
+
+ verified_user + +
+
{{ captcha }}
+ + refresh + +
+
+ + + + +
+ {% else %} +
+ +
+ +
+ badge + +
+
+
+ +
+ lock + +
+
+ +
+ {% endif %} +
+
+ +
+
+ 系统状态 + 服务条款 + API 文档 +
+
+
+
+ + diff --git a/templates/result.html b/templates/result.html new file mode 100644 index 0000000..cf8dc2e --- /dev/null +++ b/templates/result.html @@ -0,0 +1,226 @@ + + + + + 预测结果 + + + + + + + + + + + +
+
+
+
+
+ water_drop + 供水管道健康评估系统 +
+ +
+
+ notifications + help +
+
+
{{ current_user.username[:1]|upper }}
+
+ + +
+
+
+
+
+
+ +
+ {% if result %} +
+
+
+ description +
+
+
数据分析报告: {{ result.original_filename }}
+
基于上传数据生成的实时分析报告 • 生成于: {{ result.generated_at }}
+
+
+
+
check_circle数据源已验证
+ 返回主页 +
+
+ + + +
+
+
+

供水管道健康状态评估等级

+ +
+ +
+
I级
(0, 0.2]
管道安全风险十分严重,需立刻进行抢修或更新改造
+
II级
(0.2, 0.4]
管道安全风险较为严重,需尽快安排检修及加频巡检
+
III级
(0.4, 0.6]
管道安全风险较低,需安排定期巡检
+
IV级
(0.6, 0.8]
管道安全风险较小,维持常规巡视
+
V级
(0.8, 1]
管道安全,维持常规巡视
+
+ +
+ + + + + + + + + {% for item in result.summary_rows %} + + + + + {% endfor %} + +
管道编号健康等级
{{ item.pipe_id }} + {{ item.grade_label }} +
+
+ 共 {{ result.sample_count }} 个样本,显示 {{ result.summary_rows|length }} 条 + 查看完整样本列表 +
+
+
+ +
+
+
+

管道剩余寿命动态评估

+

生存曲线拟合

+
+
+ +
+ 生存概率阶梯图 +
+ +
+ info +

分析说明:{{ result.analysis_text }}

+
+
+
+ + {% if result.importance_url %} +
+
+
+

模型输入因素重要性排序

+

各输入因素对预测结果的相对影响程度

+
+
+
+ 模型输入因素重要性排序图 +
+
+ insights +

说明:柱状图按重要性从高到低展示各输入因素对模型预测结果的相对贡献,数值为归一化占比,可用于辅助识别影响管道健康状态的关键因素。

+
+
+ {% endif %} + {% else %} +
当前还没有预测结果,请先从主页上传文件并运行预测。
+ {% endif %} +
+ + + + diff --git a/tests/test_prediction.py b/tests/test_prediction.py new file mode 100644 index 0000000..6bd7512 --- /dev/null +++ b/tests/test_prediction.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +import unittest + +import pandas as pd + +from app.prediction import ( + FEATURES, + ID_COLUMN, + PredictionError, + estimate_remaining_life, + grade_info, + interpolate_probability, + secure_upload_name, + validate_input_frame, +) + + +class PredictionHelpersTest(unittest.TestCase): + def test_secure_upload_name_accepts_chinese_filename(self) -> None: + filename, suffix = secure_upload_name("管道数据.xlsx", "run123") + + self.assertEqual(suffix, ".xlsx") + self.assertTrue(filename.endswith("_run123.xlsx")) + + def test_secure_upload_name_rejects_unsupported_extension(self) -> None: + with self.assertRaises(PredictionError): + secure_upload_name("管道数据.txt", "run123") + + def test_probability_helpers(self) -> None: + self.assertEqual(interpolate_probability([1, 5, 10], [0.9, 0.8, 0.6], 6), 0.6) + self.assertEqual(estimate_remaining_life([1, 5, 10], [0.9, 0.4, 0.2]), 5.0) + + def test_grade_boundaries(self) -> None: + self.assertEqual(grade_info(0.2)[0], "I级") + self.assertEqual(grade_info(0.8)[0], "IV级") + self.assertEqual(grade_info(0.81)[0], "V级") + + def test_validate_input_frame_reports_missing_columns(self) -> None: + df = pd.DataFrame({ID_COLUMN: [1], FEATURES[0]: [1]}) + + with self.assertRaises(PredictionError) as ctx: + validate_input_frame(df) + + self.assertIn("缺少必要字段", ctx.exception.message) + + +if __name__ == "__main__": + unittest.main()