diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..4101cb6 --- /dev/null +++ b/.env.example @@ -0,0 +1,10 @@ +# Copy to .env for local deployment. +# 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. +ADMIN_USERNAME=admin +ADMIN_PASSWORD= + +# Default: 16 MiB +MAX_UPLOAD_BYTES=16777216 diff --git a/docker-compose.yml b/docker-compose.yml index 6780b89..8ad1989 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -3,6 +3,11 @@ services: image: pipeline-lifetime:latest container_name: pipeline-lifetime restart: unless-stopped + environment: + SECRET_KEY: ${SECRET_KEY:-} + ADMIN_USERNAME: ${ADMIN_USERNAME:-admin} + ADMIN_PASSWORD: ${ADMIN_PASSWORD:-} + MAX_UPLOAD_BYTES: ${MAX_UPLOAD_BYTES:-16777216} ports: - "5005:5005" volumes: diff --git a/final_flask_app_0331_strict_fixed_F.py b/final_flask_app_0331_strict_fixed_F.py index 80e0a0e..4e632b4 100644 --- a/final_flask_app_0331_strict_fixed_F.py +++ b/final_flask_app_0331_strict_fixed_F.py @@ -11,8 +11,9 @@ from __future__ import annotations import logging import os -import random +import secrets import sys +import uuid from dataclasses import dataclass from datetime import datetime from io import BytesIO @@ -84,6 +85,7 @@ FEATURES = [ ] ID_COLUMN = "管道编号" +PIPE_AGE_COLUMN = "管龄" TEMPLATE_COLUMNS = [ ID_COLUMN, @@ -101,10 +103,28 @@ logging.basicConfig( ) app = Flask(__name__) -app.config["SECRET_KEY"] = "pipe-survival-0331-strict-secret" + + +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) @@ -112,6 +132,14 @@ 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) @@ -145,6 +173,36 @@ 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 @@ -160,6 +218,13 @@ def resource_path(relative_path: str) -> str: 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) @@ -169,11 +234,25 @@ def init_app() -> None: ensure_dirs() with app.app_context(): db.create_all() - if not User.query.filter_by(username="admin").first(): - admin = User(username="admin", is_admin=True) - admin.set_password("admin123") - db.session.add(admin) + 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() @@ -194,7 +273,7 @@ def load_model(): def inject_helpers(): def now_year() -> int: return datetime.now().year - return {"feature_list": FEATURES, "now_year": now_year} + return {"feature_list": FEATURES, "now_year": now_year, "csrf_token": csrf_token} BASE_TEMPLATE_HEAD = r""" @@ -313,6 +392,7 @@ LOGIN_TEMPLATE = r""" {% if mode == 'login' %}
+
@@ -358,6 +438,7 @@ LOGIN_TEMPLATE = r""" {% else %}
+
@@ -421,7 +502,10 @@ DASHBOARD_TEMPLATE = r"""
{{ current_user.username[:1]|upper }}
- 退出登录 + + + +
@@ -460,6 +544,7 @@ DASHBOARD_TEMPLATE = r"""
+
@@ -1037,7 +1125,7 @@ def login(): return redirect(url_for("home")) if request.method == "GET": - session["captcha"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5)) + session["captcha"] = new_captcha() return render_template_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"]) username = request.form.get("username", "").strip() @@ -1046,13 +1134,13 @@ def login(): if captcha_input != session.get("captcha", ""): flash("验证码错误", "error") - session["captcha"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5)) + 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"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5)) + 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"))) @@ -1082,11 +1170,11 @@ def register(): db.session.add(user) db.session.commit() flash("注册成功,请登录", "info") - session["captcha"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5)) + session["captcha"] = new_captcha() return render_template_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"]) -@app.route("/logout") +@app.route("/logout", methods=["POST"]) @login_required def logout(): logout_user() @@ -1121,7 +1209,12 @@ 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) - file_path = record.saved_path if file_type == "original" else record.prediction_path + 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) @@ -1154,14 +1247,17 @@ def predict(): 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}_{timestamp}{ext}" + original_path = user_dir / f"{stem}_{run_id}{ext}" uploaded.save(original_path) try: @@ -1171,10 +1267,13 @@ def predict(): df = pd.read_excel(original_path) except Exception as exc: logging.exception("文件解析失败: %s", exc) + safe_unlink(original_path) return jsonify({"error": "文件解析失败,请检查编码或表格格式。"}), 400 - missing = [col for col in FEATURES if col not in df.columns] + 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() @@ -1182,6 +1281,7 @@ def predict(): 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)) @@ -1191,8 +1291,8 @@ def predict(): 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]["Status"]) if "Status" in df.columns and pd.notna(df.iloc[i]["Status"]) else f"Pipe_{i+1:03d}" - pipe_age = f"{df.iloc[i]['Pipeage']} 年" if "Pipeage" in df.columns and pd.notna(df.iloc[i]["Pipeage"]) else "-" + 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) @@ -1207,7 +1307,10 @@ def predict(): "grade_class": grade_class, }) summary_sheet_rows.append({ - "管道标识": i + 1, + "管道编号": pipe_id, + "管龄": pipe_age, + "健康概率": health_probability, + "预计剩余寿命": remaining_life, "健康等级": grade_label, }) plt.step(times, probs, where="post", linewidth=2, label=pipe_id) @@ -1220,7 +1323,7 @@ def predict(): plt.legend(loc="best", fontsize=8) plt.tight_layout() - image_filename = f"plot_{current_user.id}_{timestamp}.png" + 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() @@ -1229,13 +1332,13 @@ def predict(): try: importance_values = compute_feature_importance(model, x_test) if importance_values is not None: - importance_filename = f"importance_{current_user.id}_{timestamp}.png" + 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_{timestamp}.xlsx" + 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)