307 lines
10 KiB
Python
307 lines
10 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""统一 Web 鉴权模块:OAuth2 Password + SQLite 存储。
|
||
|
||
- 用户/访问令牌存 SQLite:hexi/data/web_auth.db
|
||
- 提供 OAuth2 密码流(tokenUrl 指向 /hub/api/auth/token)
|
||
- 供统一 Web /hub 与各插件 API 共用(/api/<plugin>/...)
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import hashlib
|
||
import json
|
||
import os
|
||
import secrets
|
||
import sqlite3
|
||
import time
|
||
from pathlib import Path
|
||
from typing import Optional
|
||
|
||
from fastapi import Depends, HTTPException
|
||
from fastapi.security import OAuth2PasswordBearer
|
||
|
||
basic_path = Path(__file__).resolve().parents[1] # hexi/
|
||
DATA_DIR = basic_path / "data"
|
||
DB_PATH = DATA_DIR / "web_auth.db"
|
||
|
||
TOKEN_TTL_SECONDS = 60 * 60 * 24 # 24 小时
|
||
PBKDF2_ITERATIONS = 100_000
|
||
|
||
oauth2_scheme = OAuth2PasswordBearer(
|
||
tokenUrl="/hub/api/auth/token",
|
||
auto_error=False,
|
||
)
|
||
|
||
|
||
def _conn() -> sqlite3.Connection:
|
||
conn = sqlite3.connect(DB_PATH)
|
||
conn.row_factory = sqlite3.Row
|
||
return conn
|
||
|
||
|
||
def _now() -> str:
|
||
return time.strftime("%Y-%m-%d %H:%M:%S")
|
||
|
||
|
||
def _read_env(key: str, default: str) -> str:
|
||
"""取值优先级:os.environ > .env 文件 > default(NoneBot 不一定把 .env 全塞进 env)。"""
|
||
val = os.getenv(key)
|
||
if val:
|
||
return val
|
||
env_path = basic_path.parent / ".env"
|
||
if env_path.exists():
|
||
for line in env_path.read_text(encoding="utf-8").splitlines():
|
||
line = line.strip()
|
||
if line.startswith(key + "=") or line.startswith(key + " ="):
|
||
value = line.split("=", 1)[1].strip()
|
||
return value.strip('"').strip("'")
|
||
return default
|
||
|
||
|
||
def _hash_password(password: str) -> str:
|
||
salt = secrets.token_hex(16)
|
||
digest = hashlib.pbkdf2_hmac(
|
||
"sha256", password.encode("utf-8"), salt.encode("utf-8"), PBKDF2_ITERATIONS
|
||
).hex()
|
||
return f"{salt}${digest}"
|
||
|
||
|
||
def _verify_password(password: str, stored: str) -> bool:
|
||
try:
|
||
salt, digest = stored.split("$", 1)
|
||
except ValueError:
|
||
return False
|
||
calc = hashlib.pbkdf2_hmac(
|
||
"sha256", password.encode("utf-8"), salt.encode("utf-8"), PBKDF2_ITERATIONS
|
||
).hex()
|
||
return secrets.compare_digest(digest, calc)
|
||
|
||
|
||
def init_auth_db() -> None:
|
||
"""只建表(轻量,可每次调用);用户播种/同步由 sync_admin() 在启动时做一次。"""
|
||
DATA_DIR.mkdir(parents=True, exist_ok=True)
|
||
with _conn() as conn:
|
||
conn.execute(
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS web_users (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
username TEXT UNIQUE NOT NULL,
|
||
password_hash TEXT NOT NULL,
|
||
is_active INTEGER NOT NULL DEFAULT 1,
|
||
created_at TEXT
|
||
)
|
||
"""
|
||
)
|
||
conn.execute(
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS web_tokens (
|
||
token TEXT PRIMARY KEY,
|
||
user_id INTEGER NOT NULL,
|
||
expires_at REAL NOT NULL,
|
||
created_at TEXT
|
||
)
|
||
"""
|
||
)
|
||
conn.commit()
|
||
|
||
|
||
def _env_has_password_key() -> bool:
|
||
""".env 中是否显式配置了 hexi_web_password(区别于代码默认值 admin)。"""
|
||
env_path = basic_path.parent / ".env"
|
||
if not env_path.exists():
|
||
return False
|
||
for line in env_path.read_text(encoding="utf-8").splitlines():
|
||
line = line.strip()
|
||
if line.startswith("hexi_web_password=") or line.startswith(
|
||
"hexi_web_password ="
|
||
):
|
||
return True
|
||
return False
|
||
|
||
|
||
def _write_env_password(password: str) -> None:
|
||
"""把 Web 端修改后的密码回写 .env,保持 .env 永远是权威源(重启后不还原)。"""
|
||
env_path = basic_path.parent / ".env"
|
||
try:
|
||
lines = (
|
||
env_path.read_text(encoding="utf-8").splitlines()
|
||
if env_path.exists()
|
||
else []
|
||
)
|
||
replaced = False
|
||
out: list[str] = []
|
||
for line in lines:
|
||
stripped = line.strip()
|
||
if stripped.startswith("hexi_web_password=") or stripped.startswith(
|
||
"hexi_web_password ="
|
||
):
|
||
if not replaced:
|
||
# JSON 字符串写法可安全转义 # 等特殊字符,dotenv 解析后与裸值一致
|
||
out.append("hexi_web_password=" + json.dumps(password))
|
||
replaced = True
|
||
else:
|
||
out.append(line)
|
||
if not replaced:
|
||
out.append("hexi_web_password=" + json.dumps(password))
|
||
env_path.write_text(
|
||
"\n".join(out).rstrip("\n") + "\n", encoding="utf-8"
|
||
)
|
||
except OSError:
|
||
pass
|
||
|
||
|
||
def _seed_admin_if_missing(username: str, password: str) -> None:
|
||
with _conn() as conn:
|
||
row = conn.execute(
|
||
"SELECT id FROM web_users WHERE username=?", (username,)
|
||
).fetchone()
|
||
if row is None:
|
||
conn.execute(
|
||
"INSERT INTO web_users(username,password_hash,is_active,created_at) VALUES(?,?,1,?)",
|
||
(username, _hash_password(password), _now()),
|
||
)
|
||
conn.commit()
|
||
|
||
|
||
def sync_admin() -> None:
|
||
"""启动时调用:以 .env 为准创建/同步超级管理员。
|
||
|
||
仅当 .env 显式配置了 hexi_web_password 时才会在每次启动同步密码
|
||
(此时 .env 是权威源,重启始终以 .env 覆盖);若 .env 未显式配置,
|
||
则只在用户不存在时播种默认账号,保证 Web 端「修改密码」能持久生效、
|
||
不被每次启动静默还原。
|
||
"""
|
||
init_auth_db()
|
||
username = _read_env("hexi_web_username", "admin")
|
||
password = _read_env("hexi_web_password", "admin")
|
||
if not _env_has_password_key():
|
||
_seed_admin_if_missing(username, password)
|
||
return
|
||
with _conn() as conn:
|
||
row = conn.execute(
|
||
"SELECT id FROM web_users WHERE username=?", (username,)
|
||
).fetchone()
|
||
if row is None:
|
||
conn.execute(
|
||
"INSERT INTO web_users(username,password_hash,is_active,created_at) VALUES(?,?,1,?)",
|
||
(username, _hash_password(password), _now()),
|
||
)
|
||
else:
|
||
# 以 .env 为准同步超级管理员密码:解决改 .env 后库里还是旧密码的问题
|
||
conn.execute(
|
||
"UPDATE web_users SET password_hash=?, is_active=1 WHERE id=?",
|
||
(_hash_password(password), row["id"]),
|
||
)
|
||
conn.commit()
|
||
|
||
|
||
def create_user(username: str, password: str) -> int:
|
||
init_auth_db()
|
||
with _conn() as conn:
|
||
cursor = conn.execute(
|
||
"INSERT INTO web_users(username,password_hash,is_active,created_at) VALUES(?,?,1,?)",
|
||
(username, _hash_password(password), _now()),
|
||
)
|
||
conn.commit()
|
||
return int(cursor.lastrowid)
|
||
|
||
|
||
def revoke_user_tokens(user_id: int, except_token: Optional[str] = None) -> None:
|
||
"""吊销某用户访问令牌(可保留当前令牌)。改密码后调用,旧登录态全部失效。"""
|
||
init_auth_db()
|
||
with _conn() as conn:
|
||
if except_token:
|
||
conn.execute(
|
||
"DELETE FROM web_tokens WHERE user_id=? AND token<>?",
|
||
(user_id, except_token),
|
||
)
|
||
else:
|
||
conn.execute("DELETE FROM web_tokens WHERE user_id=?", (user_id,))
|
||
conn.commit()
|
||
|
||
|
||
def change_password(
|
||
user_id: int, old_password: str, new_password: str, keep_token: Optional[str] = None
|
||
) -> bool:
|
||
"""校验旧密码后改新密码,并吊销旧令牌(保留当前令牌可选)。成功返回 True。
|
||
|
||
成功后会回写 .env(hexi_web_password),保证重启后仍是新密码:
|
||
否则 sync_admin 会按 .env 默认值把密码还原。
|
||
"""
|
||
if not new_password:
|
||
return False
|
||
init_auth_db()
|
||
with _conn() as conn:
|
||
row = conn.execute(
|
||
"SELECT * FROM web_users WHERE id=? AND is_active=1", (user_id,)
|
||
).fetchone()
|
||
if not row or not _verify_password(old_password, row["password_hash"]):
|
||
return False
|
||
conn.execute(
|
||
"UPDATE web_users SET password_hash=? WHERE id=?",
|
||
(_hash_password(new_password), user_id),
|
||
)
|
||
conn.commit()
|
||
_write_env_password(new_password)
|
||
revoke_user_tokens(user_id, except_token=keep_token)
|
||
return True
|
||
|
||
|
||
def authenticate(username: str, password: str) -> Optional[int]:
|
||
"""校验成功返回 user_id,否则 None。"""
|
||
init_auth_db()
|
||
with _conn() as conn:
|
||
row = conn.execute(
|
||
"SELECT * FROM web_users WHERE username=? AND is_active=1", (username,)
|
||
).fetchone()
|
||
if row and _verify_password(password, row["password_hash"]):
|
||
return int(row["id"])
|
||
return None
|
||
|
||
|
||
def issue_token(user_id: int) -> tuple[str, int]:
|
||
"""签发不透明访问令牌(存库可吊销),返回 (token, expires_in)。"""
|
||
init_auth_db()
|
||
token = secrets.token_urlsafe(32)
|
||
expires_at = time.time() + TOKEN_TTL_SECONDS
|
||
with _conn() as conn:
|
||
conn.execute(
|
||
"INSERT INTO web_tokens(token,user_id,expires_at,created_at) VALUES(?,?,?,?)",
|
||
(token, user_id, expires_at, _now()),
|
||
)
|
||
conn.commit()
|
||
return token, TOKEN_TTL_SECONDS
|
||
|
||
|
||
def revoke_token(token: str) -> None:
|
||
init_auth_db()
|
||
with _conn() as conn:
|
||
conn.execute("DELETE FROM web_tokens WHERE token=?", (token,))
|
||
conn.commit()
|
||
|
||
|
||
def get_user_by_token(token: str) -> Optional[dict]:
|
||
init_auth_db()
|
||
with _conn() as conn:
|
||
row = conn.execute(
|
||
"""
|
||
SELECT u.id, u.username, u.created_at
|
||
FROM web_tokens t JOIN web_users u ON u.id = t.user_id
|
||
WHERE t.token=? AND t.expires_at > ?
|
||
""",
|
||
(token, time.time()),
|
||
).fetchone()
|
||
return dict(row) if row else None
|
||
|
||
|
||
async def get_current_user(token: Optional[str] = Depends(oauth2_scheme)) -> dict:
|
||
if not token:
|
||
raise HTTPException(status_code=401, detail="未登录")
|
||
user = get_user_by_token(token)
|
||
if not user:
|
||
raise HTTPException(status_code=401, detail="登录态无效或已过期")
|
||
return user
|
||
|
||
|
||
# 各插件 API 用的统一鉴权依赖
|
||
require_admin = Depends(get_current_user) |