Baseline: pr1 HYC下载站 v2.3 before security/functional fixes
This commit is contained in:
@@ -0,0 +1,587 @@
|
||||
#!/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
|
||||
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] = {}
|
||||
|
||||
# 确定基础目录(用于保存会话文件)
|
||||
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()
|
||||
|
||||
# 会话文件路径
|
||||
sessions_filename = config.get('auth_sessions_file', 'auth_sessions.json') if config else 'auth_sessions.json'
|
||||
self.sessions_file = os.path.join(base_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
|
||||
|
||||
# 加载已保存的会话
|
||||
self._load_sessions()
|
||||
|
||||
@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 是否过期
|
||||
if user.get('token_expires_at') 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"""
|
||||
if not cookie_value:
|
||||
return None
|
||||
|
||||
parts = cookie_value.split('.')
|
||||
if len(parts) != 3:
|
||||
return None
|
||||
|
||||
session_id, timestamp, signature = parts
|
||||
|
||||
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 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
|
||||
|
||||
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 = time.time()
|
||||
|
||||
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 ['*']
|
||||
)
|
||||
|
||||
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:
|
||||
"""销毁会话"""
|
||||
if session_id in self.sessions:
|
||||
del self.sessions[session_id]
|
||||
self._save_sessions()
|
||||
return True
|
||||
return False
|
||||
|
||||
# === 内部方法 ===
|
||||
|
||||
def _generate_signature(self, session_id: str, timestamp: str, user_id: str) -> str:
|
||||
"""生成签名"""
|
||||
secret = self.config.get('auth_secret', 'default_secret_change_me')
|
||||
data = f"{session_id}.{timestamp}.{user_id}.{secret}"
|
||||
return hashlib.sha256(data.encode()).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):
|
||||
"""保存会话"""
|
||||
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:
|
||||
with open(self.sessions_file, 'w', encoding='utf-8') as f:
|
||||
json.dump(data, f, ensure_ascii=False, indent=2)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def get_stats(self) -> dict:
|
||||
"""获取认证统计"""
|
||||
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端点定义 ===
|
||||
|
||||
ADMIN_API_ENDPOINTS = {
|
||||
# 同步管理
|
||||
'POST:/api/v2/sync/*': 'sync:manage',
|
||||
'POST:/api/v2/sync/*/start': 'sync:start',
|
||||
'POST:/api/v2/sync/*/stop': 'sync:stop',
|
||||
'DELETE:/api/v2/sync/*': 'sync:manage',
|
||||
|
||||
# 缓存管理
|
||||
'POST:/api/v2/cache/clean': 'cache:manage',
|
||||
'DELETE:/api/v2/cache/*': '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/*/trigger': 'webhook:trigger',
|
||||
|
||||
# 服务器配置
|
||||
'PUT:/api/v2/config': 'config:manage',
|
||||
'POST:/api/v2/server/reload': 'server:reload',
|
||||
|
||||
# 文件管理(高危操作)
|
||||
'DELETE:/api/v2/files/*': 'files:delete',
|
||||
'PUT:/api/v2/files/*/rename': 'files:rename',
|
||||
|
||||
# 用户管理
|
||||
'POST:/api/v2/users': 'users:create',
|
||||
'DELETE:/api/v2/users/*': 'users:delete',
|
||||
'PUT:/api/v2/users/*': 'users:update',
|
||||
|
||||
# === API v1 文件操作认证 ===
|
||||
'DELETE:/api/v1/file/*': 'files:delete',
|
||||
'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 check_endpoint_auth(method: str, path: str, auth_manager: APIAuthManager) -> dict:
|
||||
"""检查端点是否需要认证"""
|
||||
key = f"{method}:{path}"
|
||||
if key in ADMIN_API_ENDPOINTS:
|
||||
return {
|
||||
"required": True,
|
||||
"permission": ADMIN_API_ENDPOINTS[key]
|
||||
}
|
||||
|
||||
for pattern, permission in ADMIN_API_ENDPOINTS.items():
|
||||
if '*' in pattern:
|
||||
pat_method, pat_path = pattern.split(':', 1)
|
||||
if method == pat_method or pat_method == '*':
|
||||
if pat_path.endswith('*'):
|
||||
prefix = pat_path.rstrip('*').rstrip('/')
|
||||
if path.startswith(prefix):
|
||||
return {
|
||||
"required": True,
|
||||
"permission": permission
|
||||
}
|
||||
|
||||
return {
|
||||
"required": False,
|
||||
"permission": None
|
||||
}
|
||||
Reference in New Issue
Block a user