文件
workbuddy-portal/workbuddy_portal/security.py
T
wangchuanli 23799b4ea5 feat(oss): 补齐开源声明体系(MIT + 第三方声明 + 贡献/安全/行为准则)
* LICENSE              —— MIT
  * THIRD-PARTY-NOTICES  —— 依赖清单、再分发合规说明(含随仓库分发的
                            Apache ECharts 5.6.0 / Apache-2.0)与自查清单
  * CONTRIBUTING.md      —— 开发环境、五层验证、必须遵守的不变量、提交规范
  * SECURITY.md          —— 漏洞私有报告渠道、已有措施、已知非目标
  * CODE_OF_CONDUCT.md   —— 改编自 Contributor Covenant 2.1
  * .github/             —— Bug 报告 / 功能建议表单 + PR 模板
  * .editorconfig        —— 与 .gitattributes 保持一致
  * 全部 Python / Shell 源文件加 SPDX-License-Identifier: MIT 头
  * README 增加「开源与许可」章节与许可标识
2026-09-14 16:15:26 +08:00

197 行
5.9 KiB
Python

# -*- coding: utf-8 -*-
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 Wang Chuanli
"""密码哈希、登录装饰器、CSRF。
局域网可访问 ⇒ 必须有鉴权。这里用 Werkzeug 自带的 PBKDF2,不引第三方依赖。
"""
import functools
import hmac
import secrets
import time
from flask import (current_app, flash, jsonify, redirect, render_template, request,
session, url_for)
from werkzeug.security import check_password_hash, generate_password_hash
from . import config, db
# 简易失败计数(内存即可:单进程部署,重启清零可接受)
_fails = {} # ip -> [count, first_ts]
_FAILS_MAX_IPS = 4096 # 上限,防止大量来源 IP 把字典撑爆
_FAILS_TTL = 3600 # 超过 1 小时无更新的条目会被清理
def _prune_fails(now=None):
"""清掉过期条目;条目数超上限时按时间淘汰最旧的。"""
now = now or time.time()
dead = [ip for ip, c in _fails.items() if now - c[1] > _FAILS_TTL]
for ip in dead:
_fails.pop(ip, None)
if len(_fails) > _FAILS_MAX_IPS:
for ip, _ in sorted(_fails.items(), key=lambda kv: kv[1][1])[:len(_fails) - _FAILS_MAX_IPS]:
_fails.pop(ip, None)
def hash_password(p):
return generate_password_hash(p, method="pbkdf2:sha256:200000")
def verify_password(hashed, p):
try:
return check_password_hash(hashed, p)
except (ValueError, TypeError):
return False
def login_ok(conn, username, password):
row = conn.execute("SELECT * FROM users WHERE username=?", (username,)).fetchone()
if row is None or not verify_password(row["password_hash"], password):
return None
conn.execute("UPDATE users SET last_login_at=?, login_count=login_count+1 WHERE id=?",
(db.now_str(), row["id"]))
return row
# ---------------- 跳转目标白名单(防开放重定向) ----------------
def safe_next(target, fallback="/"):
"""只允许站内相对路径。
`//evil.com`、`/\\evil.com`、`https://evil.com` 都必须拒绝:
`//` 开头是协议相对 URL,浏览器会把 `//evil.com` 当成外站跳转。
"""
if not target:
return fallback
t = str(target).strip()
if not t.startswith("/"):
return fallback
if t.startswith("//") or t.startswith("/\\") or "\\" in t:
return fallback
# 去重斜杠后仍以 // 开头的(如 "/\t/evil")一并拒绝
if t.lstrip("/").startswith("//"):
return fallback
if "\r" in t or "\n" in t:
return fallback
return t
# ---------------- 登录失败限速 ----------------
def note_fail(ip):
now = time.time()
_prune_fails(now)
c = _fails.get(ip)
if c is None or now - c[1] > config.LOGIN_LOCK_MINUTES * 60:
_fails[ip] = [1, now]
return 1
c[0] += 1
return c[0]
def is_locked(ip):
c = _fails.get(ip)
if not c or c[0] < config.MAX_LOGIN_FAILS:
return False
return time.time() - c[1] <= config.LOGIN_LOCK_MINUTES * 60
def clear_fail(ip):
_fails.pop(ip, None)
def lock_left(ip):
c = _fails.get(ip)
if not c:
return 0
return max(0, int(config.LOGIN_LOCK_MINUTES * 60 - (time.time() - c[1])))
# ---------------- 会话 ----------------
def current_user():
uid = session.get("uid")
if not uid:
return None
return {"id": uid, "username": session.get("uname"), "display_name": session.get("dname"),
"is_admin": bool(session.get("adm", 1))}
def is_admin():
u = current_user()
return bool(u and u.get("is_admin"))
def login_session(user):
session.clear()
session["uid"] = user["id"]
session["uname"] = user["username"]
session["dname"] = user["display_name"] or user["username"]
try:
session["adm"] = 1 if user["is_admin"] else 0
except (KeyError, IndexError, TypeError):
session["adm"] = 1
session.permanent = True
def logout_session():
session.clear()
def wants_json():
return (request.path.startswith("/api/")
or request.accept_mimetypes.best == "application/json")
def login_required(fn):
@functools.wraps(fn)
def wrapper(*a, **kw):
if current_user() is None:
if wants_json():
return jsonify({"ok": False, "error": "unauthorized",
"message": "登录已失效,请重新登录"}), 401
return redirect(url_for("views.login", next=request.full_path))
return fn(*a, **kw)
return wrapper
def admin_required(fn):
"""管理员专属操作(用户管理等)。非管理员返回 403。"""
@functools.wraps(fn)
@login_required
def wrapper(*a, **kw):
if not is_admin():
if wants_json():
return jsonify({"ok": False, "error": "forbidden",
"message": "只有管理员可以执行该操作"}), 403
return render_template("error.html", code=403, message="只有管理员可以执行该操作"), 403
return fn(*a, **kw)
return wrapper
# ---------------- CSRF ----------------
def csrf_token():
t = session.get("_csrf")
if not t:
t = session["_csrf"] = secrets.token_urlsafe(24)
return t
def check_csrf():
"""对所有 POST/PUT/DELETE 生效,失败直接 400。"""
if request.method in ("GET", "HEAD", "OPTIONS"):
return None
sent = request.form.get("_csrf") or request.headers.get("X-CSRF-Token") or ""
if not sent or not hmac.compare_digest(sent, session.get("_csrf", "")):
if wants_json():
return jsonify({"ok": False, "error": "csrf", "message": "CSRF 校验失败,请刷新页面"}), 400
return "CSRF 校验失败,请刷新页面后重试", 400
return None
def init_app(app):
app.jinja_env.globals["csrf_token"] = csrf_token
app.jinja_env.globals["current_user"] = current_user
@app.before_request
def _guard():
return check_csrf()