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 sqlalchemy.orm import joinedload from .config import BASE_DIR from .extensions import db from .models import AppSetting, UploadRecord, User from .prediction import PredictionError, run_prediction from .security import new_captcha bp = Blueprint("main", __name__) REFERENCE_PDF_NAME = "20260630标准文本——供水管道健康状态与剩余寿命评估技术导则.pdf" REGISTRATION_SETTING_KEY = "allow_registration" RECORDS_PER_PAGE = 10 def render_auth_template(mode: str, status_code: int = 200, captcha: str = ""): return render_template("login.html", mode=mode, captcha=captcha), status_code def render_login_error(message: str, status_code: int = 400): flash(message, "error") session["captcha"] = new_captcha() return render_auth_template("login", status_code, session["captcha"]) def require_admin() -> None: if not current_user.is_admin: abort(403) def registration_allowed() -> bool: return AppSetting.get_bool( REGISTRATION_SETTING_KEY, current_app.config["ALLOW_REGISTRATION"], ) def requested_page() -> int: try: return max(int(request.args.get("page", 1)), 1) except (TypeError, ValueError): return 1 def paginated_uploads(query, endpoint: str): page = requested_page() pagination = ( query.order_by(UploadRecord.upload_time.desc()) .paginate(page=page, per_page=RECORDS_PER_PAGE, error_out=False) ) if pagination.pages and page > pagination.pages: return pagination, redirect(url_for(endpoint, page=pagination.pages)) return pagination, None def prediction_result_payload(artifacts, record: UploadRecord) -> dict: 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 ) return { "original_filename": artifacts.original_filename, "generated_at": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), "image_url": image_url, "importance_url": importance_url, "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[:6], "analysis_text": artifacts.analysis_text, } @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_auth_template("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", ""): return render_login_error("验证码错误") user = User.query.filter_by(username=username).first() if not user or not user.check_password(password): return render_login_error("用户名或密码错误") login_user(user, remember=bool(request.form.get("remember"))) return redirect(url_for("main.home")) @bp.route("/register", methods=["GET", "POST"]) def register(): if request.method == "GET": return render_auth_template("register") if not registration_allowed(): flash("当前未开放自助注册,请联系管理员。", "error") return render_auth_template("register", 403) username = request.form.get("username", "").strip() password = request.form.get("password", "") if not username: flash("用户名不能为空", "error") return render_auth_template("register", 400) if len(password) < 6: flash("密码至少需要 6 位", "error") return render_auth_template("register", 400) if User.query.filter_by(username=username).first(): flash("用户名已存在", "error") return render_auth_template("register", 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_auth_template("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(): pagination, page_redirect = paginated_uploads( UploadRecord.query.filter_by(user_id=current_user.id), "main.history_page", ) if page_redirect: return page_redirect return render_template( "history.html", pagination=pagination, records=pagination.items, ) @bp.route("/admin") @login_required def admin_dashboard(): require_admin() pagination, page_redirect = paginated_uploads( UploadRecord.query.options(joinedload(UploadRecord.user)), "main.admin_dashboard", ) if page_redirect: return page_redirect return render_template( "admin.html", pagination=pagination, records=pagination.items, registration_allowed=registration_allowed(), ) @bp.route("/admin/registration", methods=["POST"]) @login_required def update_registration_setting(): require_admin() allow_registration = request.form.get("allow_registration") == "on" AppSetting.set_bool(REGISTRATION_SETTING_KEY, allow_registration) db.session.commit() message = "已开放用户自助注册" if allow_registration else "已关闭用户自助注册" if request.headers.get("X-Requested-With") == "XMLHttpRequest": return jsonify( { "message": message, "registration_allowed": allow_registration, "status_label": "已开放" if allow_registration else "已关闭", } ) flash(message, "info") return redirect(url_for("main.admin_dashboard")) @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) file_path = { "original": record.saved_path, "prediction": record.prediction_path, }.get(file_type) if file_path is None: 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("/reference_pdf") @login_required def reference_pdf(): pdf_path = BASE_DIR / REFERENCE_PDF_NAME if not pdf_path.exists(): abort(404) return send_file( pdf_path, as_attachment=False, download_name=REFERENCE_PDF_NAME, mimetype="application/pdf", ) @bp.route("/reference") @login_required def reference_page(): return render_template("reference.html") @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 = prediction_result_payload(artifacts, record) 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, } )