Files
HYC Fixer abdbec85a4 复查修复(四): 独立审查发现的问题
- S1/S2: 会话锁改 RLock(持锁可重入调 _save_sessions);cookie 时间戳改整数+解析兼容(会话创建/验证往返已实测)
- M1: api_login 接入 verify_user(账号锁定/失败计数生效),DB 无用户时才回退 config 凭据
- M3+L5: handler _do_auth 统一入口加 IP 白名单检查;未知 auth_type 返回 401
- M4+L10: metrics 与无版本 /api/admin/ 加入受保护端点
- M6: debug 日志敏感头脱敏;main.py 不再打印 token 前缀
- M7: auth_token.txt / auth_sessions.json chmod 600
- M9: verify_user 统一错误消息防用户枚举
- M10: AdminAPI 复用共享 APIAuthManager(修复会话状态分裂)
- L7: check_auth 大小写不敏感匹配(防 /API/.. 大写绕过)
- L8: token_expires_at 显式 is not None 判断
- L11: verify_password 对非 bcrypt 哈希回退 PBKDF2(重写,修复 ValueError 分支不落回退的问题)
2026-09-02 00:45:01 +08:00

708 lines
25 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
API认证模块
使用数据库进行认证,支持:
- none: 无认证
- basic: Basic Auth(用户名密码)
- token: Token 认证(登录生成的token)
"""
import os
import sys
import json
import hashlib
import time
import secrets
import threading
from typing import Dict, List, Optional
from dataclasses import dataclass
from functools import wraps
@dataclass
class AuthSession:
"""认证会话"""
session_id: str
user_id: str
level: str
created_at: float
expires_at: float
last_activity: float
permissions: List[str]
class APIAuthManager:
"""API认证管理器"""
def __init__(self, config: dict = None):
self.config = config or {}
self.sessions: Dict[str, AuthSession] = {}
self._lock = threading.RLock() # 可重入: create/destroy 持锁时允许再调 _save_sessions
# 确定基础目录(用于保存会话文件)
base_dir = config.get('base_dir', '.') if config else '.'
# 检测是否是 PyInstaller 打包环境
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
base_dir = os.path.dirname(os.path.abspath(sys.executable))
elif base_dir == '.':
base_dir = os.getcwd()
# 数据目录:默认放到 base_dir 同级 data/ 下,
# 避免 auth_sessions.json 落入 web 静态根目录被公开下载
data_dir = config.get('data_dir') if config else None
if not data_dir:
data_dir = os.path.join(os.path.dirname(os.path.abspath(base_dir)), 'data')
try:
os.makedirs(data_dir, exist_ok=True)
except OSError:
data_dir = base_dir
self.data_dir = data_dir
# 会话文件路径
sessions_filename = config.get('auth_sessions_file', 'auth_sessions.json') if config else 'auth_sessions.json'
self.sessions_file = os.path.join(data_dir, sessions_filename)
# 会话超时时间
self.session_timeout = config.get('auth_session_timeout', 3600) if config else 3600
# Cookie名称
self.cookie_name = 'hyc_auth'
self.cookie_max_age = config.get('auth_cookie_max_age', 86400) if config else 86400
# IP 白名单
self.ip_whitelist = config.get('ip_whitelist', []) if config else []
self.ip_whitelist_enabled = config.get('ip_whitelist_enabled', False) if config else False
# Cookie 签名密钥:优先用配置;缺失则从文件读取或生成并持久化
self.auth_secret = self._load_or_create_secret(config)
# 加载已保存的会话
self._load_sessions()
def _load_or_create_secret(self, config) -> str:
"""获取或生成 auth_secret(持久化到数据目录,避免默认密钥公开可伪造)"""
secret = config.get('auth_secret') if config else None
if secret:
return secret
secret_file = os.path.join(self.data_dir, 'auth_secret.key')
try:
if os.path.exists(secret_file):
with open(secret_file, 'r', encoding='utf-8') as f:
secret = f.read().strip()
if secret:
return secret
secret = secrets.token_hex(32)
with open(secret_file, 'w', encoding='utf-8') as f:
f.write(secret)
os.chmod(secret_file, 0o600)
except Exception as e:
print(f"警告: 持久化 auth_secret 失败: {e}")
return secret
@property
def db(self):
"""动态获取数据库实例"""
return self.config.get('_db_instance')
def _get_client_ip(self, handler) -> str:
"""从 handler 获取客户端 IP"""
try:
forwarded = handler.headers.get('X-Forwarded-For')
if forwarded:
return forwarded.split(',')[0].strip()
return handler.client_address[0]
except Exception:
return None
def _check_ip_whitelist(self, ip: str) -> bool:
"""检查 IP 是否在白名单中"""
if not self.ip_whitelist_enabled or not self.ip_whitelist:
return True
if not ip:
return False
import ipaddress
for pattern in self.ip_whitelist:
try:
if '/' in pattern:
network = ipaddress.ip_network(pattern, strict=False)
if ipaddress.ip_address(ip) in network:
return True
elif pattern == ip:
return True
except ValueError:
continue
return False
# === Token 验证 ===
def validate_token(self, token: str, client_ip: str = None) -> Optional[dict]:
"""验证 token(从数据库)"""
if not token:
return None
# 优先从数据库验证
if self.db:
user = self.db.get_user_by_token(token)
if user:
# 检查 token 是否过期(显式 is not None,避免 0 被当作永不过期)
if user.get('token_expires_at') is not None and time.time() > user['token_expires_at']:
return {"valid": False, "reason": "Token已过期"}
# 检查用户是否启用
if not user.get('enabled', True):
return {"valid": False, "reason": "用户已被禁用"}
return {
"valid": True,
"key_id": f"user_{user['id']}",
"name": user['username'],
"level": user.get('role', 'admin'),
"permissions": ["*"],
"user_id": user['id'],
"username": user['username']
}
return {"valid": False, "reason": "无效的Token"}
# === Basic Auth 验证 ===
def validate_basic_auth(self, username: str, password: str, client_ip: str = None) -> dict:
"""验证 Basic Auth 用户名密码(从数据库)"""
# IP 白名单检查
if not self._check_ip_whitelist(client_ip):
if self.db:
self.db.add_login_log(username, client_ip, 'failed', 'IP不在白名单')
return {"valid": False, "reason": "IP不在白名单内"}
auth_type = self.config.get('auth_type', 'none')
# 如果认证类型为 none,任何用户都可以通过
if auth_type == 'none':
return {
"valid": True,
"user_id": 0,
"username": username or 'anonymous',
"level": "admin",
"key_id": "anonymous",
"name": f"Anonymous - {username or 'anonymous'}",
"permissions": ["*"]
}
# 从数据库验证
if self.db:
result = self.db.verify_user(username, password)
if result.get('valid'):
if self.db:
self.db.add_login_log(username, client_ip, 'success', '数据库验证')
return {
"valid": True,
"user_id": result.get('user_id'),
"username": username,
"level": result.get('role', 'admin'),
"key_id": f"user_{result.get('user_id')}",
"name": f"User - {username}",
"permissions": ["*"]
}
else:
if self.db:
self.db.add_login_log(username, client_ip, 'failed', result.get('reason', '验证失败'))
return {"valid": False, "reason": result.get('reason', '用户名或密码错误')}
return {"valid": False, "reason": "数据库不可用"}
# === Cookie 验证 ===
def validate_cookie(self, cookie_value: str, client_ip: str = None) -> Optional[dict]:
"""验证认证Cookie"""
import hmac
if not cookie_value:
return None
parts = cookie_value.split('.')
if len(parts) < 3:
return None
# session_id(64 hex, 无点) 是固定的第一段; 其余段为 timestamp + signature
# (旧版 timestamp 为浮点含点, 兼容拆分后的多段)
session_id = parts[0]
timestamp = '.'.join(parts[1:-1])
signature = parts[-1]
if not timestamp or not signature:
return None
with self._lock:
session = self.sessions.get(session_id)
if not session:
return None
if time.time() > session.expires_at:
del self.sessions[session_id]
return None
expected_sig = self._generate_signature(session_id, timestamp, session.user_id)
if not hmac.compare_digest(signature, expected_sig):
return None
session.last_activity = time.time()
return {
"valid": True,
"session_id": session_id,
"user_id": session.user_id,
"level": session.level,
"permissions": session.permissions
}
def validate_session_id(self, session_id: str) -> Optional[dict]:
"""验证会话ID"""
if not session_id:
return None
with self._lock:
session = self.sessions.get(session_id)
if not session:
return None
if time.time() > session.expires_at:
del self.sessions[session_id]
return None
session.last_activity = time.time()
return {
"valid": True,
"session_id": session_id,
"user_id": session.user_id,
"level": session.level,
"permissions": session.permissions
}
# === 请求验证 ===
def validate_request(self, handler, required_level: str = "admin") -> dict:
"""
验证请求的认证状态
支持的认证方式:
1. Authorization: Bearer <token>
2. Authorization: Basic <credentials>
3. X-API-Key: <token>
4. Cookie: hyc_auth=<session>
5. ?key=<token>
"""
import base64
auth_header = handler.headers.get('Authorization')
api_key = handler.headers.get('X-API-Key')
cookie = handler.headers.get('Cookie', '')
client_ip = handler.client_address[0] if hasattr(handler, 'client_address') else None
# 提取cookie值
cookie_value = None
for c in cookie.split(';'):
c = c.strip()
if c.startswith(f'{self.cookie_name}='):
cookie_value = c[len(self.cookie_name)+1:]
break
# 获取查询参数中的key
parsed_path = handler.path.split('?')
query_key = None
if len(parsed_path) > 1:
from urllib.parse import parse_qs
query = parse_qs(parsed_path[1])
query_key = query.get('key', [None])[0]
# 1. Bearer Token
if auth_header and auth_header.startswith('Bearer '):
token = auth_header[7:]
result = self.validate_token(token, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "bearer", **result}
# 2. Basic Auth
if auth_header and auth_header.startswith('Basic '):
try:
credentials = base64.b64decode(auth_header[6:]).decode('utf-8')
if ':' in credentials:
username, password = credentials.split(':', 1)
result = self.validate_basic_auth(username, password, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "basic", **result}
except Exception:
pass
# 3. API Key Header
if api_key:
result = self.validate_token(api_key, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "api_key", **result}
# 4. Cookie
if cookie_value:
result = self.validate_cookie(cookie_value, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "cookie", **result}
# 5. Query Parameter
if query_key:
result = self.validate_token(query_key, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "query", **result}
# 未认证
return {
"authenticated": False,
"error": "Authentication required",
"required_level": required_level
}
def check_permission(self, auth_result: dict, permission: str) -> bool:
"""检查是否有权限访问特定API"""
if not auth_result.get('authenticated'):
return False
permissions = auth_result.get('permissions', [])
if '*' in permissions:
return True
if permission in permissions:
return True
for p in permissions:
if p.endswith('*'):
prefix = p.rstrip('*')
if permission.startswith(prefix):
return True
return False
# === 会话管理 ===
def create_session(self, user_id: str, level: str,
permissions: List[str] = None) -> dict:
"""创建认证会话"""
session_id = secrets.token_hex(32)
timestamp = int(time.time()) # 整数时间戳: cookie 用 . 分隔时不会把浮点拆成多段
session = AuthSession(
session_id=session_id,
user_id=user_id,
level=level,
created_at=timestamp,
expires_at=timestamp + self.session_timeout,
last_activity=timestamp,
permissions=permissions or ['*']
)
with self._lock:
# 顺带清理过期会话(防止字典无限增长,配合懒清理)
now = time.time()
for sid in [sid for sid, sess in self.sessions.items() if sess.expires_at <= now]:
del self.sessions[sid]
self.sessions[session_id] = session
self._save_sessions()
signature = self._generate_signature(session_id, timestamp, user_id)
cookie_value = f"{session_id}.{timestamp}.{signature}"
return {
"session_id": session_id,
"cookie_name": self.cookie_name,
"cookie_value": cookie_value,
"cookie_max_age": self.cookie_max_age,
"expires": timestamp + self.session_timeout
}
def destroy_session(self, session_id: str) -> bool:
"""销毁会话"""
with self._lock:
if session_id in self.sessions:
del self.sessions[session_id]
self._save_sessions()
return True
return False
def list_sessions(self) -> List[dict]:
"""列出活跃会话(加锁遍历,供管理接口使用)"""
now = time.time()
with self._lock:
sessions = []
for session_id, session in self.sessions.items():
if now < session.expires_at:
sessions.append({
"session_id": session.session_id,
"user_id": session.user_id,
"level": session.level,
"created_at": session.created_at,
"expires_at": session.expires_at,
"last_activity": session.last_activity
})
return sessions
# === 内部方法 ===
def _generate_signature(self, session_id: str, timestamp: str, user_id: str) -> str:
"""生成签名(HMAC-SHA256,使用持久化的 auth_secret)"""
import hmac
data = f"{session_id}.{timestamp}.{user_id}".encode('utf-8')
return hmac.new(self.auth_secret.encode('utf-8'), data, hashlib.sha256).hexdigest()[:32]
def _load_sessions(self):
"""加载会话"""
if os.path.exists(self.sessions_file):
try:
with open(self.sessions_file, 'r', encoding='utf-8') as f:
data = json.load(f)
now = time.time()
for item in data:
if item.get('expires_at') and now > item['expires_at']:
continue
session = AuthSession(
session_id=item['session_id'],
user_id=item['user_id'],
level=item['level'],
created_at=item['created_at'],
expires_at=item['expires_at'],
last_activity=item['last_activity'],
permissions=item.get('permissions', ['*'])
)
self.sessions[session.session_id] = session
except Exception as e:
print(f"加载会话失败: {e}")
def _save_sessions(self):
"""保存会话(原子写: 临时文件 + os.replace)"""
with self._lock:
data = []
for session in self.sessions.values():
data.append({
"session_id": session.session_id,
"user_id": session.user_id,
"level": session.level,
"created_at": session.created_at,
"expires_at": session.expires_at,
"last_activity": session.last_activity,
"permissions": session.permissions
})
try:
tmp_file = self.sessions_file + '.tmp'
with open(tmp_file, 'w', encoding='utf-8') as f:
json.dump(data, f, ensure_ascii=False, indent=2)
os.replace(tmp_file, self.sessions_file)
try:
os.chmod(self.sessions_file, 0o600)
except OSError:
pass
except Exception as e:
print(f"警告: 保存会话失败: {e}")
def get_stats(self) -> dict:
"""获取认证统计"""
with self._lock:
active_sessions = sum(
1 for s in self.sessions.values()
if time.time() < s.expires_at
)
return {
"active_sessions": active_sessions,
"session_timeout": self.session_timeout
}
# === API认证装饰器 ===
def require_auth(required_level: str = "admin", permission: str = None):
"""
API认证装饰器
使用方式:
@require_auth()
def api_endpoint(self, handler):
...
@require_auth(permission="sync:start")
def api_sync_start(self, handler):
...
"""
def decorator(func):
@wraps(func)
def wrapper(self, handler, *args, **kwargs):
# 检查是否需要认证
if required_level == "none":
return func(self, handler, *args, **kwargs)
config = getattr(handler, 'config', {})
auth_type = config.get('auth_type', 'none')
# 如果auth_type为none,跳过认证
if auth_type == "none":
handler.auth_result = {
"authenticated": True,
"level": "admin",
"user_id": "anonymous",
"permissions": ["*"]
}
return func(self, handler, *args, **kwargs)
auth_manager = getattr(handler, 'auth_manager', None)
if not auth_manager:
handler.send_json_response({
"error": "认证系统未初始化",
"code": "AUTH_NOT_INITIALIZED"
}, 500)
return
auth_result = auth_manager.validate_request(handler, required_level)
if not auth_result.get('authenticated'):
handler.send_response(401)
handler.send_header('WWW-Authenticate', 'Bearer realm="HYC API"')
handler.send_header('Access-Control-Allow-Origin', '*')
handler.send_json_response({
"error": "未认证或认证已过期",
"code": "UNAUTHORIZED",
"required_level": required_level,
"auth_methods": [
"Authorization: Bearer <token>",
"X-API-Key: <token>",
f"Cookie: {auth_manager.cookie_name}=<session>",
"?key=<token>"
]
})
return
if permission:
if not auth_manager.check_permission(auth_result, permission):
handler.send_json_response({
"error": "权限不足",
"code": "FORBIDDEN",
"required_permission": permission
}, 403)
return
handler.auth_result = auth_result
return func(self, handler, *args, **kwargs)
return wrapper
return decorator
# === 需要认证的API端点定义 ===
# 规则格式: 'METHOD:/api/vN/path' 或 'METHOD:/api/vN/prefix/'(前缀规则,尾斜杠)
# 匹配时路径统一归一化为无前导斜杠形式,与 v1/v2 传入的 api_action 对齐
ADMIN_API_ENDPOINTS = {
# 同步管理
'POST:/api/v2/sync/': 'sync:manage',
'DELETE:/api/v2/sync/': 'sync:manage',
# 缓存管理
'POST:/api/v2/cache/clean': 'cache:manage',
'POST:/api/v2/cache/prewarm/': 'cache:manage',
'DELETE:/api/v2/cache/prewarm/': 'cache:manage',
# Webhook管理
'POST:/api/v2/webhooks': 'webhook:create',
'PUT:/api/v2/webhooks/': 'webhook:update',
'DELETE:/api/v2/webhooks/': 'webhook:delete',
'POST:/api/v2/webhooks/': 'webhook:trigger',
# 服务器配置
'PUT:/api/v2/config': 'config:manage',
'POST:/api/v2/server/reload': 'server:reload',
'POST:/api/v2/server/': 'server:manage',
# 文件管理(高危操作: 删除/重命名/元数据/版本)
'DELETE:/api/v1/file/': 'files:delete',
'DELETE:/api/v2/file/': 'files:delete',
'PUT:/api/v2/file/': 'files:update',
'POST:/api/v2/file/': 'files:update',
# 用户管理
'POST:/api/v2/users': 'users:create',
'DELETE:/api/v2/users/': 'users:delete',
'PUT:/api/v2/users/': 'users:update',
'POST:/api/v2/user/password': 'users:update',
# 镜像管理(写 settings.json,必须鉴权)
'POST:/api/v2/mirrors': 'mirrors:manage',
'PUT:/api/v2/mirrors/': 'mirrors:manage',
'DELETE:/api/v2/mirrors/': 'mirrors:manage',
# 告警配置与API文档生成
'PUT:/api/v2/alerts': 'config:manage',
'POST:/api/v2/alerts': 'config:manage',
'POST:/api/v2/api-docs/generate': 'config:manage',
# === API v1 文件操作认证 ===
'PUT:/api/v1/mkdir': 'files:create',
'POST:/api/v1/upload': 'files:upload',
'POST:/api/v1/batch': 'files:batch',
'POST:/api/v1/archive': 'files:archive',
}
def _normalize_endpoint_patterns():
"""把 ADMIN_API_ENDPOINTS 归一化为 (精确表, 前缀表)
规则格式: 'METHOD:/api/vN/path'
- 无通配符 → 精确匹配 (method, path)
- 尾斜杠或尾* → 前缀匹配, 匹配 norm_path.startswith(prefix + '/')
"""
exact = {}
prefix = []
for pattern, permission in ADMIN_API_ENDPOINTS.items():
pat_method, pat_path = pattern.split(':', 1)
pat_path = pat_path.lstrip('/')
if '*' in pattern or pat_path.endswith('/'):
if pat_path.endswith('*'):
pat_path = pat_path.rstrip('*')
pat_path = pat_path.rstrip('/')
prefix.append((pat_method, pat_path, permission))
else:
exact[(pat_method, pat_path)] = permission
return exact, prefix
_EXACT_ENDPOINTS, _PREFIX_ENDPOINTS = _normalize_endpoint_patterns()
def check_endpoint_auth(method: str, path: str, auth_manager: APIAuthManager) -> dict:
"""检查端点是否需要认证
路径统一归一化为无前导斜杠形式(如 'api/v2/sync/sources'),
与 v1/v2 传入的 api_action 保持一致。
"""
norm_path = path.lstrip('/')
# 精确匹配
key = (method, norm_path)
if key in _EXACT_ENDPOINTS:
return {
"required": True,
"permission": _EXACT_ENDPOINTS[key]
}
# 前缀匹配
for pat_method, pat_path, permission in _PREFIX_ENDPOINTS:
if method != pat_method and pat_method != '*':
continue
if norm_path.startswith(pat_path + '/'):
return {
"required": True,
"permission": permission
}
return {
"required": False,
"permission": None
}