dce9e31bde
- 新增 JWT 认证模块,支持登录/注册/用户管理 - 新增文件上传接口,支持 .txt/.md/.html/.pdf/.docx 等格式解析入库 - 新增检索结果 AI 总结功能 - 新增文本去重缓存机制 - 新增全局认证夹具简化测试 - 新增配置项与环境变量支持 - 完善文档与测试覆盖
214 lines
7.9 KiB
Python
214 lines
7.9 KiB
Python
"""JWT 认证核心:密码哈希、token 签发/验签、用户存储、FastAPI 鉴权依赖
|
||
|
||
用户持久化在 Redis(key: auth:user:{username})。登录/注册需 Redis 可用;
|
||
鉴权(get_current_user)以 JWT 自包含信息为主,Redis 不可用时降级可用。
|
||
"""
|
||
|
||
from datetime import UTC, datetime, timedelta
|
||
|
||
import bcrypt
|
||
import jwt
|
||
import structlog
|
||
from fastapi import Depends
|
||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||
from redis import asyncio as redis_async
|
||
|
||
from app.api.response import ApiError
|
||
from app.config import settings
|
||
from app.models.auth import AuthUser, StoredUser
|
||
|
||
logger = structlog.get_logger()
|
||
|
||
# 错误码
|
||
ERR_UNAUTHORIZED = 1003 # 未认证(需要登录)
|
||
ERR_TOKEN_INVALID = 1005 # token 无效或已过期
|
||
ERR_FORBIDDEN = 1006 # 权限不足
|
||
ERR_USER_EXISTS = 1007 # 用户名已存在
|
||
ERR_BAD_CREDENTIALS = 1008 # 用户名或密码错误
|
||
ERR_REGISTER_DISABLED = 1009 # 注册已关闭
|
||
|
||
# Bearer token 提取器(auto_error=False,统一由 ApiError 处理)
|
||
_bearer = HTTPBearer(auto_error=False)
|
||
|
||
# 模块级 JWT secret(settings 为空时自动生成,仅开发用)
|
||
_auto_secret: str | None = None
|
||
|
||
# 模块级用户存储单例
|
||
_user_store: "UserStore | None" = None
|
||
|
||
|
||
def _get_secret() -> str:
|
||
"""获取 JWT 签名密钥,settings 为空时生成一次性随机密钥(仅开发)"""
|
||
global _auto_secret
|
||
if settings.jwt_secret_key:
|
||
return settings.jwt_secret_key
|
||
if _auto_secret is None:
|
||
import secrets
|
||
|
||
_auto_secret = secrets.token_urlsafe(48)
|
||
logger.warning(
|
||
"JWT_SECRET_KEY 未配置,已自动生成随机密钥(仅开发用,重启后旧 token 失效)"
|
||
)
|
||
return _auto_secret
|
||
|
||
|
||
def hash_password(password: str) -> str:
|
||
"""密码 bcrypt 哈希(返回 utf-8 字符串形式的哈希)"""
|
||
return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8")
|
||
|
||
|
||
def verify_password(password: str, hashed: str) -> bool:
|
||
"""校验明文密码与哈希是否匹配,异常或不匹配均返回 False"""
|
||
try:
|
||
return bcrypt.checkpw(password.encode("utf-8"), hashed.encode("utf-8"))
|
||
except Exception:
|
||
return False
|
||
|
||
|
||
def create_access_token(username: str, role: str) -> tuple[str, int]:
|
||
"""签发 JWT,返回 (token, expires_in_seconds)"""
|
||
now = datetime.now(UTC)
|
||
expire = now + timedelta(minutes=settings.jwt_expire_minutes)
|
||
payload = {
|
||
"sub": username,
|
||
"role": role,
|
||
"iat": int(now.timestamp()),
|
||
"exp": int(expire.timestamp()),
|
||
}
|
||
token = jwt.encode(payload, _get_secret(), algorithm=settings.jwt_algorithm)
|
||
return token, settings.jwt_expire_minutes * 60
|
||
|
||
|
||
def decode_token(token: str) -> dict:
|
||
"""验签并返回 payload,失败抛 ApiError(ERR_TOKEN_INVALID)"""
|
||
try:
|
||
return jwt.decode(token, _get_secret(), algorithms=[settings.jwt_algorithm])
|
||
except jwt.ExpiredSignatureError as exc:
|
||
raise ApiError(ERR_TOKEN_INVALID, "token 已过期") from exc
|
||
except jwt.InvalidTokenError as exc:
|
||
raise ApiError(ERR_TOKEN_INVALID, "token 无效") from exc
|
||
|
||
|
||
class UserStore:
|
||
"""Redis 用户存储"""
|
||
|
||
def __init__(self, redis_url: str | None = None) -> None:
|
||
self._redis_url = redis_url or settings.redis_url
|
||
self._client: redis_async.Redis | None = None
|
||
|
||
def _get_client(self) -> redis_async.Redis:
|
||
if self._client is None:
|
||
self._client = redis_async.from_url(self._redis_url, decode_responses=True)
|
||
return self._client
|
||
|
||
@staticmethod
|
||
def _key(username: str) -> str:
|
||
return f"auth:user:{username}"
|
||
|
||
async def get(self, username: str) -> StoredUser | None:
|
||
"""读取用户,Redis 异常或数据损坏均返回 None"""
|
||
try:
|
||
raw = await self._get_client().get(self._key(username))
|
||
except Exception:
|
||
logger.error("Redis 读取用户失败", username=username, exc_info=True)
|
||
return None
|
||
if raw is None:
|
||
return None
|
||
try:
|
||
return StoredUser.model_validate_json(raw)
|
||
except Exception:
|
||
logger.error("用户数据损坏", username=username, exc_info=True)
|
||
return None
|
||
|
||
async def exists(self, username: str) -> bool:
|
||
try:
|
||
return bool(await self._get_client().exists(self._key(username)))
|
||
except Exception:
|
||
logger.error("Redis 检查用户存在性失败", username=username, exc_info=True)
|
||
return False
|
||
|
||
async def create(
|
||
self, username: str, password: str, role: str = "user"
|
||
) -> StoredUser:
|
||
"""创建用户,已存在抛 ApiError(ERR_USER_EXISTS),Redis 写失败抛 ApiError(2000)"""
|
||
if await self.exists(username):
|
||
raise ApiError(ERR_USER_EXISTS, f"用户名已存在: {username}")
|
||
user = StoredUser(
|
||
username=username,
|
||
role=role,
|
||
created_at=datetime.now(UTC),
|
||
hashed_password=hash_password(password),
|
||
)
|
||
try:
|
||
await self._get_client().set(self._key(username), user.model_dump_json())
|
||
except Exception as exc:
|
||
logger.error("Redis 写入用户失败", username=username, exc_info=True)
|
||
raise ApiError(2000, f"用户创建失败: {exc}") from exc
|
||
logger.info("用户创建成功", username=username, role=role)
|
||
return user
|
||
|
||
async def authenticate(self, username: str, password: str) -> StoredUser:
|
||
"""验证用户名密码,失败抛 ApiError(ERR_BAD_CREDENTIALS)"""
|
||
user = await self.get(username)
|
||
if user is None or not verify_password(password, user.hashed_password):
|
||
raise ApiError(ERR_BAD_CREDENTIALS, "用户名或密码错误")
|
||
return user
|
||
|
||
|
||
def get_user_store() -> UserStore:
|
||
"""全局 UserStore 单例"""
|
||
global _user_store
|
||
if _user_store is None:
|
||
_user_store = UserStore()
|
||
return _user_store
|
||
|
||
|
||
async def ensure_default_admin() -> None:
|
||
"""启动时按配置创建默认管理员(密码为空则跳过)"""
|
||
pwd = settings.default_admin_password
|
||
if not pwd:
|
||
return
|
||
store = get_user_store()
|
||
try:
|
||
if await store.exists(settings.default_admin_username):
|
||
return
|
||
await store.create(settings.default_admin_username, pwd, role="admin")
|
||
logger.info("默认管理员已创建", username=settings.default_admin_username)
|
||
except ApiError:
|
||
raise
|
||
except Exception:
|
||
logger.warning("默认管理员创建失败,跳过", exc_info=True)
|
||
|
||
|
||
async def get_current_user(
|
||
credentials: HTTPAuthorizationCredentials | None = Depends(_bearer),
|
||
) -> AuthUser:
|
||
"""FastAPI 鉴权依赖:解析 Bearer token,返回 AuthUser
|
||
|
||
优先查 Redis 获取最新用户信息(支持删除/改角色生效);
|
||
Redis 不可用或用户未命中时降级为 JWT payload 构造 AuthUser,保证鉴权可用。
|
||
"""
|
||
if credentials is None or (credentials.scheme or "").lower() != "bearer":
|
||
raise ApiError(ERR_UNAUTHORIZED, "未提供认证凭证,请先登录")
|
||
payload = decode_token(credentials.credentials)
|
||
username = payload.get("sub")
|
||
role = payload.get("role", "user")
|
||
if not username:
|
||
raise ApiError(ERR_TOKEN_INVALID, "token 缺少用户标识")
|
||
|
||
stored = await get_user_store().get(username)
|
||
if stored is not None:
|
||
return AuthUser(
|
||
username=stored.username, role=stored.role, created_at=stored.created_at
|
||
)
|
||
# Redis 不可用或用户不存在(可能已删除):降级用 JWT payload
|
||
logger.warning("用户存储查询未命中,降级使用 JWT payload", username=username)
|
||
return AuthUser(username=username, role=role, created_at=datetime.now(UTC))
|
||
|
||
|
||
async def require_admin(user: AuthUser = Depends(get_current_user)) -> AuthUser:
|
||
"""要求管理员角色"""
|
||
if user.role != "admin":
|
||
raise ApiError(ERR_FORBIDDEN, "需要管理员权限")
|
||
return user
|