Files
mirror_server/core/api_auth.py
T

695 lines
24 KiB
Python
Raw 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.Lock()
# 确定基础目录(用于保存会话文件)
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 是否过期
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"""
import hmac
if not cookie_value:
return None
parts = cookie_value.split('.')
if len(parts) != 3:
return None
session_id, timestamp, signature = parts
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 = 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 ['*']
)
with self._lock:
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)
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
}