fix(security): harden auth and uploads
This commit is contained in:
@@ -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' %}
|
||||
<form method="post" action="{{ url_for('login') }}" class="space-y-5">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token() }}">
|
||||
<div>
|
||||
<label class="block text-[11px] tracking-[0.18em] uppercase text-slate-500 mb-2">身份识别 / 用户名</label>
|
||||
<div class="relative">
|
||||
@@ -358,6 +438,7 @@ LOGIN_TEMPLATE = r"""
|
||||
</form>
|
||||
{% else %}
|
||||
<form method="post" action="{{ url_for('register') }}" class="space-y-5">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token() }}">
|
||||
<div>
|
||||
<label class="block text-[11px] tracking-[0.18em] uppercase text-slate-500 mb-2">用户名</label>
|
||||
<div class="relative">
|
||||
@@ -421,7 +502,10 @@ DASHBOARD_TEMPLATE = r"""
|
||||
<div class="h-8 w-px bg-slate-200"></div>
|
||||
<div class="flex items-center gap-3 text-sm font-semibold">
|
||||
<div class="w-8 h-8 rounded-full bg-primary flex items-center justify-center text-white text-xs">{{ current_user.username[:1]|upper }}</div>
|
||||
<a href="{{ url_for('logout') }}" class="text-textMain">退出登录</a>
|
||||
<form method="post" action="{{ url_for('logout') }}">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token() }}">
|
||||
<button class="text-textMain" type="submit">退出登录</button>
|
||||
</form>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -460,6 +544,7 @@ DASHBOARD_TEMPLATE = r"""
|
||||
</div>
|
||||
|
||||
<form id="predictForm" enctype="multipart/form-data">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token() }}">
|
||||
<label id="dropZone" class="block relative rounded-[14px] border-2 border-dashed border-outline p-12 bg-white hover:border-primary hover:bg-blue-50/40 transition-all cursor-pointer text-center">
|
||||
<input id="fileInput" name="file" type="file" accept=".csv,.xls,.xlsx" class="absolute inset-0 opacity-0 cursor-pointer" required>
|
||||
<div class="flex flex-col items-center gap-4 pointer-events-none">
|
||||
@@ -687,7 +772,10 @@ RESULT_TEMPLATE = r"""
|
||||
<div class="h-8 w-px bg-slate-200"></div>
|
||||
<div class="flex items-center gap-3 text-sm font-semibold">
|
||||
<div class="w-8 h-8 rounded-full bg-primary flex items-center justify-center text-white text-xs">{{ current_user.username[:1]|upper }}</div>
|
||||
<a href="{{ url_for('logout') }}" class="text-textMain">退出登录</a>
|
||||
<form method="post" action="{{ url_for('logout') }}">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token() }}">
|
||||
<button class="text-textMain" type="submit">退出登录</button>
|
||||
</form>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user