复查修复(二): 同步模块与安全遗留

- 修复 remove_sync_source 引入的 stop_sync KeyError 回归(容错 get)
- stop_sync 不再删除存活线程条目(防双 worker);FTP 子目录递归传真名(不再 KeyError)
- 远程文件名统一 _safe_remote_name 校验(FTP/SFTP/HTTP 防路径穿越)
- temp 任务名加随机后缀防碰撞;temp 状态/线程完成后清理(防 sync_state.json 膨胀)
- save_sync_state 快照+原子写(临时文件+os.replace),持锁调用不死锁
- URL 打印脱敏(user:pass@ -> ***@)
- _need_sync_http size=0 不再全量重下(仅按存在性)
- cron 同一分钟去重;PyPI 进度只更新当前源
- v2 start_sync 不再把 bool 当 task_id
- UserRecord.to_dict 脱敏(不返回 password_hash/token),新增 to_dict_private/get_user_with_password
- config_hotreload 单配置源(set+persist 与热重载一致);server.py 用 get_all()
- 会话创建时顺带清理过期项;JSON stats/历史读改写加锁
- 解压目标目录先校验;FTP RETR 命令注入防护;rsync --delete 目标保护
- git/urlopen 补超时;cleanup_completed_tasks 删旧留新
This commit is contained in:
HYC Fixer
2026-09-02 00:39:12 +08:00
parent fc72759a92
commit 82875b710a
10 changed files with 957 additions and 810 deletions
+7 -2
View File
@@ -1297,6 +1297,10 @@ class APIv1:
return return
extract_dir = os.path.join(self.config['base_dir'], target_dir) extract_dir = os.path.join(self.config['base_dir'], target_dir)
# 先校验目标目录本身,防止 target_dir 含 .. 把目录建到 base_dir 外
if not is_safe_path(self.config['base_dir'], extract_dir):
handler.send_json_response({"error": "Invalid target directory"}, 403)
return
os.makedirs(extract_dir, exist_ok=True) os.makedirs(extract_dir, exist_ok=True)
with zipfile.ZipFile(archive_path, 'r') as zipf: with zipfile.ZipFile(archive_path, 'r') as zipf:
@@ -1305,12 +1309,13 @@ class APIv1:
if not is_safe_path(self.config['base_dir'], member_path): if not is_safe_path(self.config['base_dir'], member_path):
handler.send_json_response({"error": "Unsafe archive contents"}, 403) handler.send_json_response({"error": "Unsafe archive contents"}, 403)
return return
members = zipf.namelist()
zipf.extractall(extract_dir) zipf.extractall(extract_dir)
# 同步提取的文件到数据库 # 同步提取的文件到数据库
if self.db: if self.db:
extracted_files = [] extracted_files = []
for member in zipf.namelist(): for member in members:
member_path = os.path.join(extract_dir, member) member_path = os.path.join(extract_dir, member)
if os.path.isfile(member_path): if os.path.isfile(member_path):
rel_member_path = os.path.relpath(member_path, self.config['base_dir']).replace("\\", "/") rel_member_path = os.path.relpath(member_path, self.config['base_dir']).replace("\\", "/")
@@ -1327,7 +1332,7 @@ class APIv1:
handler.send_json_response({ handler.send_json_response({
"operation": "extract", "operation": "extract",
"extract_dir": target_dir, "extract_dir": target_dir,
"extracted_files": len(zipf.namelist()) "extracted_files": len(members)
}) })
except Exception as e: except Exception as e:
+57 -25
View File
@@ -24,6 +24,37 @@ class APIv2(APIv1):
# 管理员API处理器(始终创建,auth_type检查在装饰器中处理) # 管理员API处理器(始终创建,auth_type检查在装饰器中处理)
self.admin_api = AdminAPI(config) self.admin_api = AdminAPI(config)
# ==================== 状态型管理器(惰性单例,跨请求保留状态) ====================
def _get_health_checker(self):
"""健康检查器单例(原每请求新建导致状态丢失)"""
if not hasattr(self, '_health_checker'):
from core.health_check import HealthChecker
self._health_checker = self._get_health_checker()
return self._health_checker
def _get_failover_manager(self):
"""故障转移管理器单例"""
if not hasattr(self, '_failover_manager'):
from core.health_check import MirrorFailoverManager
self._failover_manager = MirrorFailoverManager(self.config)
self._failover_manager.initialize()
return self._failover_manager
def _get_restart_manager(self):
"""优雅重启管理器单例"""
if not hasattr(self, '_restart_manager'):
from core.graceful_restart import GracefulRestartManager
self._restart_manager = self._get_restart_manager()
return self._restart_manager
def _get_prewarmer(self):
"""缓存预热器单例(原每请求新建导致队列/状态丢失)"""
if not hasattr(self, '_prewarmer'):
from core.cache_prewarm import CachePrewarmer
self._prewarmer = self._get_prewarmer()
return self._prewarmer
def handle_request(self, handler, method, path, query_params): def handle_request(self, handler, method, path, query_params):
"""处理API v2请求""" """处理API v2请求"""
import sys import sys
@@ -271,7 +302,7 @@ class APIv2(APIv1):
elif path == 'health/stats': elif path == 'health/stats':
if method == 'GET': if method == 'GET':
from core.health_check import HealthChecker from core.health_check import HealthChecker
checker = HealthChecker(self.config.get('health_check', {})) checker = self._get_health_checker()
handler.send_json_response(checker.get_stats()) handler.send_json_response(checker.get_stats())
else: else:
handler.send_error(405) handler.send_error(405)
@@ -2283,11 +2314,12 @@ class APIv2(APIv1):
if source_name is None: if source_name is None:
return super().api_start_sync(handler) return super().api_start_sync(handler)
if hasattr(handler, 'sync_manager') and handler.sync_manager: if hasattr(handler, 'sync_manager') and handler.sync_manager:
task_id = handler.sync_manager.start_sync(source_name) # MirrorSyncManager.start_sync 返回 bool(True=已启动/已在运行)
if task_id: started = handler.sync_manager.start_sync(source_name)
if started:
handler.send_json_response({ handler.send_json_response({
"success": True, "success": True,
"task_id": task_id, "started": True,
"source_name": source_name "source_name": source_name
}) })
else: else:
@@ -2573,7 +2605,7 @@ class APIv2(APIv1):
from core.health_check import HealthChecker, HealthStatus from core.health_check import HealthChecker, HealthStatus
mirrors = self.config.get('mirrors', {}) mirrors = self.config.get('mirrors', {})
checker = HealthChecker(self.config.get('health_check', {})) checker = self._get_health_checker()
results = [] results = []
for mirror_type, mirror_config in mirrors.items(): for mirror_type, mirror_config in mirrors.items():
@@ -2614,7 +2646,7 @@ class APIv2(APIv1):
try: try:
from core.health_check import HealthChecker, HealthStatus from core.health_check import HealthChecker, HealthStatus
checker = HealthChecker(self.config.get('health_check', {})) checker = self._get_health_checker()
mirrors = self.config.get('mirrors', {}) mirrors = self.config.get('mirrors', {})
# 查找源对应的镜像类型 # 查找源对应的镜像类型
@@ -2659,7 +2691,7 @@ class APIv2(APIv1):
try: try:
from core.health_check import MirrorFailoverManager from core.health_check import MirrorFailoverManager
failover = MirrorFailoverManager(self.config) failover = self._get_failover_manager()
failover.initialize() failover.initialize()
handler.send_json_response({ handler.send_json_response({
@@ -2679,7 +2711,7 @@ class APIv2(APIv1):
try: try:
from core.health_check import MirrorFailoverManager from core.health_check import MirrorFailoverManager
failover = MirrorFailoverManager(self.config) failover = self._get_failover_manager()
failover.initialize() failover.initialize()
success = failover.perform_failover(mirror_type) success = failover.perform_failover(mirror_type)
@@ -3285,7 +3317,7 @@ class APIv2(APIv1):
# 数据库验证 # 数据库验证
if db: if db:
user = db.get_user(username) user = db.get_user_with_password(username)
if user and db.verify_password(password, user['password_hash']): if user and db.verify_password(password, user['password_hash']):
# 数据库验证成功,生成 token # 数据库验证成功,生成 token
import secrets import secrets
@@ -3388,7 +3420,7 @@ class APIv2(APIv1):
config_pass = config.get('auth_pass', '') config_pass = config.get('auth_pass', '')
# 验证旧密码(优先验证数据库,没有则验证配置文件)—— 强制要求,防止无旧密码改密 # 验证旧密码(优先验证数据库,没有则验证配置文件)—— 强制要求,防止无旧密码改密
user = db.get_user(username) if db else None user = db.get_user_with_password(username) if db else None
if user: if user:
# 验证数据库密码 # 验证数据库密码
if not db.verify_password(old_password, user['password_hash']): if not db.verify_password(old_password, user['password_hash']):
@@ -3985,7 +4017,7 @@ class APIv2(APIv1):
"""获取重启状态""" """获取重启状态"""
from core.graceful_restart import GracefulRestartManager, ServerState from core.graceful_restart import GracefulRestartManager, ServerState
restart_manager = GracefulRestartManager(self.config.get('restart', {})) restart_manager = self._get_restart_manager()
stats = restart_manager.get_stats() stats = restart_manager.get_stats()
handler.send_json_response({ handler.send_json_response({
@@ -3999,7 +4031,7 @@ class APIv2(APIv1):
"""获取待处理请求""" """获取待处理请求"""
from core.graceful_restart import GracefulRestartManager from core.graceful_restart import GracefulRestartManager
restart_manager = GracefulRestartManager(self.config.get('restart', {})) restart_manager = self._get_restart_manager()
pending = restart_manager.get_pending_requests() pending = restart_manager.get_pending_requests()
handler.send_json_response({ handler.send_json_response({
@@ -4018,7 +4050,7 @@ class APIv2(APIv1):
except ValueError: except ValueError:
strategy = RestartStrategy.GRACEFUL strategy = RestartStrategy.GRACEFUL
restart_manager = GracefulRestartManager(self.config.get('restart', {})) restart_manager = self._get_restart_manager()
# 准备重启 # 准备重启
prepare_result = restart_manager.prepare_restart() prepare_result = restart_manager.prepare_restart()
@@ -4061,7 +4093,7 @@ class APIv2(APIv1):
except json.JSONDecodeError: except json.JSONDecodeError:
restart_strategy = RestartStrategy.GRACEFUL restart_strategy = RestartStrategy.GRACEFUL
restart_manager = GracefulRestartManager(self.config.get('restart', {})) restart_manager = self._get_restart_manager()
# 执行重启 # 执行重启
result = restart_manager.perform_restart(strategy=restart_strategy) result = restart_manager.perform_restart(strategy=restart_strategy)
@@ -4072,7 +4104,7 @@ class APIv2(APIv1):
"""立即重启服务器""" """立即重启服务器"""
from core.graceful_restart import GracefulRestartManager from core.graceful_restart import GracefulRestartManager
restart_manager = GracefulRestartManager(self.config.get('restart', {})) restart_manager = self._get_restart_manager()
# 获取脚本路径 # 获取脚本路径
script_path = self.config.get('main_script', 'main.py') script_path = self.config.get('main_script', 'main.py')
@@ -4088,7 +4120,7 @@ class APIv2(APIv1):
"""获取重启历史""" """获取重启历史"""
from core.graceful_restart import GracefulRestartManager from core.graceful_restart import GracefulRestartManager
restart_manager = GracefulRestartManager(self.config.get('restart', {})) restart_manager = self._get_restart_manager()
history = restart_manager.get_restart_history() history = restart_manager.get_restart_history()
handler.send_json_response({ handler.send_json_response({
@@ -4182,7 +4214,7 @@ class APIv2(APIv1):
"""获取缓存预热状态""" """获取缓存预热状态"""
from core.cache_prewarm import CachePrewarmer from core.cache_prewarm import CachePrewarmer
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
status = prewarmer.get_status() status = prewarmer.get_status()
handler.send_json_response(status) handler.send_json_response(status)
@@ -4191,7 +4223,7 @@ class APIv2(APIv1):
"""获取缓存预热统计""" """获取缓存预热统计"""
from core.cache_prewarm import CachePrewarmer from core.cache_prewarm import CachePrewarmer
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
stats = prewarmer.get_stats() stats = prewarmer.get_stats()
handler.send_json_response(stats) handler.send_json_response(stats)
@@ -4205,7 +4237,7 @@ class APIv2(APIv1):
mirror_type = query_params.get('mirror_type', [None])[0] mirror_type = query_params.get('mirror_type', [None])[0]
limit = int(query_params.get('limit', [50])[0]) limit = int(query_params.get('limit', [50])[0])
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
items = prewarmer.get_items(status=status, mirror_type=mirror_type, limit=limit) items = prewarmer.get_items(status=status, mirror_type=mirror_type, limit=limit)
handler.send_json_response({ handler.send_json_response({
@@ -4217,7 +4249,7 @@ class APIv2(APIv1):
"""获取预热历史""" """获取预热历史"""
from core.cache_prewarm import CachePrewarmer from core.cache_prewarm import CachePrewarmer
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
history = prewarmer.get_history() history = prewarmer.get_history()
handler.send_json_response({ handler.send_json_response({
@@ -4234,7 +4266,7 @@ class APIv2(APIv1):
limit = int(query_params.get('limit', [50])[0]) limit = int(query_params.get('limit', [50])[0])
priority = query_params.get('priority', ['medium'])[0] priority = query_params.get('priority', ['medium'])[0]
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
# 如果指定了镜像类型,只预热该类型 # 如果指定了镜像类型,只预热该类型
targets = None targets = None
@@ -4277,7 +4309,7 @@ class APIv2(APIv1):
from core.cache_prewarm import CachePrewarmer from core.cache_prewarm import CachePrewarmer
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
prewarmer.add_items_batch(mirror_type, items, priority) prewarmer.add_items_batch(mirror_type, items, priority)
handler.send_json_response({ handler.send_json_response({
@@ -4313,7 +4345,7 @@ class APIv2(APIv1):
}, 400) }, 400)
return return
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
popular = prewarmer.get_popular_items(mirror_type) popular = prewarmer.get_popular_items(mirror_type)
if limit: if limit:
@@ -4335,7 +4367,7 @@ class APIv2(APIv1):
mirror_type = query_params.get('mirror_type', [None])[0] mirror_type = query_params.get('mirror_type', [None])[0]
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
if mirror_type: if mirror_type:
items = prewarmer.get_popular_items(mirror_type) items = prewarmer.get_popular_items(mirror_type)
@@ -4355,7 +4387,7 @@ class APIv2(APIv1):
"""清空预热队列""" """清空预热队列"""
from core.cache_prewarm import CachePrewarmer from core.cache_prewarm import CachePrewarmer
prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) prewarmer = self._get_prewarmer()
prewarmer.clear_items() prewarmer.clear_items()
handler.send_json_response({ handler.send_json_response({
+4
View File
@@ -396,6 +396,10 @@ class APIAuthManager:
) )
with self._lock: 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.sessions[session_id] = session
self._save_sessions() self._save_sessions()
+347 -340
View File
@@ -1,340 +1,347 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
""" """
配置热更新模块 配置热更新模块
支持不重启服务的情况下重新加载配置 支持不重启服务的情况下重新加载配置
""" """
import os import os
import sys import sys
import json import json
import time import time
import threading import threading
import logging import logging
from typing import Dict, Any, Optional, Callable from typing import Dict, Any, Optional, Callable
from datetime import datetime from datetime import datetime
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class ConfigHotReloader: class ConfigHotReloader:
"""配置热重载管理器""" """配置热重载管理器"""
def __init__(self, config_path: str, callback: Callable = None): def __init__(self, config_path: str, callback: Callable = None):
""" """
初始化热重载管理器 初始化热重载管理器
Args: Args:
config_path: 配置文件路径 config_path: 配置文件路径
callback: 配置变更时的回调函数,接收 (config, change_type) 参数 callback: 配置变更时的回调函数,接收 (config, change_type) 参数
""" """
self.config_path = config_path self.config_path = config_path
self.callback = callback self.callback = callback
self._config: Dict[str, Any] = {} self._config: Dict[str, Any] = {}
self._last_modified: float = 0 self._last_modified: float = 0
self._last_load_time: Optional[datetime] = None self._last_load_time: Optional[datetime] = None
self._lock = threading.RLock() self._lock = threading.RLock()
# 配置变更历史 # 配置变更历史
self._change_history: list = [] self._change_history: list = []
# 监听器 # 监听器
self._listeners: Dict[str, list] = { self._listeners: Dict[str, list] = {
'on_change': [], 'on_change': [],
'on_error': [] 'on_error': []
} }
# 加载初始配置 # 加载初始配置
self.reload() self.reload()
def reload(self, silent: bool = False) -> bool: def reload(self, silent: bool = False) -> bool:
""" """
重新加载配置 重新加载配置
Args: Args:
silent: 静默模式,不触发变更通知 silent: 静默模式,不触发变更通知
Returns: Returns:
是否加载成功 是否加载成功
""" """
try: try:
if not os.path.exists(self.config_path): if not os.path.exists(self.config_path):
if not silent: if not silent:
logger.warning(f"配置文件不存在: {self.config_path}") logger.warning(f"配置文件不存在: {self.config_path}")
return False return False
# 获取文件修改时间 # 获取文件修改时间
current_mtime = os.path.getmtime(self.config_path) current_mtime = os.path.getmtime(self.config_path)
# 检查是否有变化 # 检查是否有变化
if current_mtime == self._last_modified and not silent: if current_mtime == self._last_modified and not silent:
return True return True
# 加载配置 # 加载配置
with open(self.config_path, 'r', encoding='utf-8') as f: with open(self.config_path, 'r', encoding='utf-8') as f:
new_config = json.load(f) new_config = json.load(f)
# 计算变更 # 计算变更
changes = self._compute_changes(self._config, new_config) changes = self._compute_changes(self._config, new_config)
with self._lock: with self._lock:
old_config = self._config.copy() old_config = self._config.copy()
self._config = new_config self._config = new_config
self._last_modified = current_mtime self._last_modified = current_mtime
self._last_load_time = datetime.now() self._last_load_time = datetime.now()
# 记录变更 # 记录变更
if changes: if changes:
change_record = { change_record = {
'timestamp': self._last_load_time.isoformat(), 'timestamp': self._last_load_time.isoformat(),
'changes': changes, 'changes': changes,
'old_config_keys': list(old_config.keys()), 'old_config_keys': list(old_config.keys()),
'new_config_keys': list(new_config.keys()) 'new_config_keys': list(new_config.keys())
} }
self._change_history.append(change_record) self._change_history.append(change_record)
# 保持历史记录在合理范围内 # 保持历史记录在合理范围内
if len(self._change_history) > 100: if len(self._change_history) > 100:
self._change_history = self._change_history[-50:] self._change_history = self._change_history[-50:]
if not silent and changes: if not silent and changes:
self._notify_change(changes) self._notify_change(changes)
logger.info(f"配置已重新加载: {self.config_path}") logger.info(f"配置已重新加载: {self.config_path}")
return True return True
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
error_msg = f"配置 JSON 格式错误: {e}" error_msg = f"配置 JSON 格式错误: {e}"
logger.error(error_msg) logger.error(error_msg)
self._notify_error(error_msg) self._notify_error(error_msg)
return False return False
except Exception as e: except Exception as e:
error_msg = f"加载配置失败: {e}" error_msg = f"加载配置失败: {e}"
logger.error(error_msg) logger.error(error_msg)
self._notify_error(error_msg) self._notify_error(error_msg)
return False return False
def _compute_changes(self, old: Dict, new: Dict) -> Dict: def _compute_changes(self, old: Dict, new: Dict) -> Dict:
"""计算配置变更""" """计算配置变更"""
changes = { changes = {
'added': [], 'added': [],
'removed': [], 'removed': [],
'modified': [] 'modified': []
} }
old_keys = set(old.keys()) old_keys = set(old.keys())
new_keys = set(new.keys()) new_keys = set(new.keys())
# 新增的键 # 新增的键
for key in new_keys - old_keys: for key in new_keys - old_keys:
changes['added'].append(key) changes['added'].append(key)
# 移除的键 # 移除的键
for key in old_keys - new_keys: for key in old_keys - new_keys:
changes['removed'].append(key) changes['removed'].append(key)
# 修改的键 # 修改的键
for key in old_keys & new_keys: for key in old_keys & new_keys:
if old[key] != new[key]: if old[key] != new[key]:
# 检查是否是嵌套字典 # 检查是否是嵌套字典
if isinstance(old[key], dict) and isinstance(new[key], dict): if isinstance(old[key], dict) and isinstance(new[key], dict):
nested = self._compute_nested_changes(old[key], new[key], f"{key}.") nested = self._compute_nested_changes(old[key], new[key], f"{key}.")
if nested['added'] or nested['removed'] or nested['modified']: if nested['added'] or nested['removed'] or nested['modified']:
changes['modified'].append({ changes['modified'].append({
'key': key, 'key': key,
'type': 'nested', 'type': 'nested',
'changes': nested 'changes': nested
}) })
else: else:
changes['modified'].append({ changes['modified'].append({
'key': key, 'key': key,
'type': 'value', 'type': 'value',
'old_value': old[key], 'old_value': old[key],
'new_value': new[key] 'new_value': new[key]
}) })
return changes return changes
def _compute_nested_changes(self, old: Dict, new: Dict, prefix: str = "") -> Dict: def _compute_nested_changes(self, old: Dict, new: Dict, prefix: str = "") -> Dict:
"""计算嵌套字典的变更""" """计算嵌套字典的变更"""
changes = { changes = {
'added': [], 'added': [],
'removed': [], 'removed': [],
'modified': [] 'modified': []
} }
old_keys = set(old.keys()) old_keys = set(old.keys())
new_keys = set(new.keys()) new_keys = set(new.keys())
for key in new_keys - old_keys: for key in new_keys - old_keys:
changes['added'].append(f"{prefix}{key}") changes['added'].append(f"{prefix}{key}")
for key in old_keys - new_keys: for key in old_keys - new_keys:
changes['removed'].append(f"{prefix}{key}") changes['removed'].append(f"{prefix}{key}")
for key in old_keys & new_keys: for key in old_keys & new_keys:
if old[key] != new[key]: if old[key] != new[key]:
changes['modified'].append(f"{prefix}{key}") changes['modified'].append(f"{prefix}{key}")
return changes return changes
def _notify_change(self, changes: Dict): def _notify_change(self, changes: Dict):
"""通知配置变更""" """通知配置变更"""
for listener in self._listeners['on_change']: for listener in self._listeners['on_change']:
try: try:
if callable(listener): if callable(listener):
listener(self._config, changes) listener(self._config, changes)
except Exception as e: except Exception as e:
logger.error(f"配置变更监听器执行失败: {e}") logger.error(f"配置变更监听器执行失败: {e}")
if self.callback: if self.callback:
try: try:
self.callback(self._config, changes) self.callback(self._config, changes)
except Exception as e: except Exception as e:
logger.error(f"配置回调函数执行失败: {e}") logger.error(f"配置回调函数执行失败: {e}")
def _notify_error(self, error: str): def _notify_error(self, error: str):
"""通知错误""" """通知错误"""
for listener in self._listeners['on_error']: for listener in self._listeners['on_error']:
try: try:
if callable(listener): if callable(listener):
listener(error) listener(error)
except Exception as e: except Exception as e:
logger.error(f"错误监听器执行失败: {e}") logger.error(f"错误监听器执行失败: {e}")
def add_change_listener(self, callback: Callable): def add_change_listener(self, callback: Callable):
"""添加配置变更监听器""" """添加配置变更监听器"""
self._listeners['on_change'].append(callback) self._listeners['on_change'].append(callback)
def add_error_listener(self, callback: Callable): def add_error_listener(self, callback: Callable):
"""添加错误监听器""" """添加错误监听器"""
self._listeners['on_error'].append(callback) self._listeners['on_error'].append(callback)
def get(self, key: str, default: Any = None) -> Any: def get(self, key: str, default: Any = None) -> Any:
"""获取配置值""" """获取配置值"""
with self._lock: with self._lock:
return self._config.get(key, default) return self._config.get(key, default)
def get_all(self) -> Dict: def get_all(self) -> Dict:
"""获取完整配置""" """获取完整配置"""
with self._lock: with self._lock:
return self._config.copy() return self._config.copy()
def set(self, key: str, value: Any, save: bool = True) -> bool: def set(self, key: str, value: Any, save: bool = True) -> bool:
"""设置配置值(仅内存中)""" """设置配置值(仅内存中)"""
with self._lock: with self._lock:
self._config[key] = value self._config[key] = value
if save: if save:
return self.save() return self.save()
return True return True
def save(self, path: str = None) -> bool: def save(self, path: str = None) -> bool:
"""保存配置到文件""" """保存配置到文件"""
save_path = path or self.config_path save_path = path or self.config_path
try: try:
with open(save_path, 'w', encoding='utf-8') as f: with open(save_path, 'w', encoding='utf-8') as f:
json.dump(self._config, f, ensure_ascii=False, indent=4) json.dump(self._config, f, ensure_ascii=False, indent=4)
self._last_modified = os.path.getmtime(save_path) self._last_modified = os.path.getmtime(save_path)
return True return True
except Exception as e: except Exception as e:
logger.error(f"保存配置失败: {e}") logger.error(f"保存配置失败: {e}")
return False return False
def get_change_history(self, limit: int = 10) -> list: def get_change_history(self, limit: int = 10) -> list:
"""获取配置变更历史""" """获取配置变更历史"""
return self._change_history[-limit:] return self._change_history[-limit:]
def watch(self, interval: float = 5.0): def watch(self, interval: float = 5.0):
""" """
启动后台监控线程 启动后台监控线程
Args: Args:
interval: 检查间隔(秒) interval: 检查间隔(秒)
""" """
def _watch_loop(): def _watch_loop():
while True: while True:
try: try:
self.reload() self.reload()
except Exception as e: except Exception as e:
logger.error(f"配置监控错误: {e}") logger.error(f"配置监控错误: {e}")
time.sleep(interval) time.sleep(interval)
thread = threading.Thread(target=_watch_loop, daemon=True) thread = threading.Thread(target=_watch_loop, daemon=True)
thread.start() thread.start()
logger.info(f"配置热监控已启动,间隔: {interval}秒") logger.info(f"配置热监控已启动,间隔: {interval}秒")
class ConfigManager: class ConfigManager:
"""配置管理器 - 支持热更新""" """配置管理器 - 支持热更新(单配置源: 热重载开启时以 reloader 的配置为准)"""
def __init__(self, config: Dict[str, Any] = None): def __init__(self, config: Dict[str, Any] = None):
self.config = config or {} self.config = config or {}
self._hot_reloader: Optional[ConfigHotReloader] = None self._hot_reloader: Optional[ConfigHotReloader] = None
def load_from_file(self, path: str, enable_watch: bool = False) -> bool: def _effective_config(self) -> Dict[str, Any]:
"""从文件加载配置""" """获取当前生效配置(避免与 reloader._config 双配置源分裂)"""
if not os.path.exists(path): if self._hot_reloader is not None:
return False return self._hot_reloader._config
return self.config
try:
with open(path, 'r', encoding='utf-8') as f: def load_from_file(self, path: str, enable_watch: bool = False) -> bool:
self.config = json.load(f) """从文件加载配置"""
if not os.path.exists(path):
if enable_watch: return False
self._hot_reloader = ConfigHotReloader(path)
self._hot_reloader.watch() try:
with open(path, 'r', encoding='utf-8') as f:
return True self.config = json.load(f)
except Exception as e:
logger.error(f"加载配置失败: {e}") if enable_watch:
return False self._hot_reloader = ConfigHotReloader(path)
self._hot_reloader.watch()
def hot_reload(self, path: str = None) -> bool:
"""触发热重载""" return True
if self._hot_reloader: except Exception as e:
return self._hot_reloader.reload() logger.error(f"加载配置失败: {e}")
return False return False
def get(self, key: str, default: Any = None) -> Any: def hot_reload(self, path: str = None) -> bool:
"""获取配置值""" """触发热重载"""
return self.config.get(key, default) if self._hot_reloader:
return self._hot_reloader.reload()
def set(self, key: str, value: Any, persist: bool = False, path: str = None) -> bool: return False
"""设置配置值"""
keys = key.split('.') def get(self, key: str, default: Any = None) -> Any:
current = self.config """获取配置值(读当前生效配置)"""
return self._effective_config().get(key, default)
for k in keys[:-1]:
if k not in current: def set(self, key: str, value: Any, persist: bool = False, path: str = None) -> bool:
current[k] = {} """设置配置值(写入当前生效配置,持久化时与 reloader 一致)"""
current = current[k] keys = key.split('.')
current = self._effective_config()
current[keys[-1]] = value
for k in keys[:-1]:
if persist: if k not in current:
if self._hot_reloader: current[k] = {}
return self._hot_reloader.save(path) current = current[k]
elif path:
try: current[keys[-1]] = value
with open(path, 'w', encoding='utf-8') as f:
json.dump(self.config, f, ensure_ascii=False, indent=4) if persist:
return True if self._hot_reloader:
except Exception as e: # set 已写入 reloader._config, save 落盘的就是新值
logger.error(f"保存配置失败: {e}") return self._hot_reloader.save(path)
return False elif path:
try:
return True with open(path, 'w', encoding='utf-8') as f:
json.dump(self._effective_config(), f, ensure_ascii=False, indent=4)
def add_change_listener(self, callback: Callable): return True
"""添加变更监听器""" except Exception as e:
if self._hot_reloader: logger.error(f"保存配置失败: {e}")
self._hot_reloader.add_change_listener(callback) return False
def get_all(self) -> Dict: return True
"""获取完整配置"""
return self.config.copy() def add_change_listener(self, callback: Callable):
"""添加变更监听器"""
if self._hot_reloader:
self._hot_reloader.add_change_listener(callback)
def get_all(self) -> Dict:
"""获取完整配置(当前生效版本)"""
return self._effective_config().copy()
+17 -3
View File
@@ -369,11 +369,10 @@ class UserRecord(Base):
enabled = Column(Boolean, default=True) enabled = Column(Boolean, default=True)
def to_dict(self) -> dict: def to_dict(self) -> dict:
"""安全序列化: 不返回密码哈希与 token(防止泄露给前端/日志)"""
return { return {
'id': self.id, 'id': self.id,
'username': self.username, 'username': self.username,
'password_hash': self.password_hash,
'token': self.token,
'token_expires_at': self.token_expires_at, 'token_expires_at': self.token_expires_at,
'role': self.role, 'role': self.role,
'email': self.email, 'email': self.email,
@@ -387,6 +386,13 @@ class UserRecord(Base):
'enabled': self.enabled 'enabled': self.enabled
} }
def to_dict_private(self) -> dict:
"""内部序列化: 含密码哈希与 token,仅供认证内部逻辑使用"""
data = self.to_dict()
data['password_hash'] = self.password_hash
data['token'] = self.token
return data
class LoginLogRecord(Base): class LoginLogRecord(Base):
"""登录日志记录表""" """登录日志记录表"""
@@ -1212,13 +1218,21 @@ class DatabaseManager:
return {'success': True, 'user_id': user.id} return {'success': True, 'user_id': user.id}
def get_user(self, username: str) -> dict: def get_user(self, username: str) -> dict:
"""获取用户信息""" """获取用户信息(安全版,不含 password_hash/token)"""
with self.session() as session: with self.session() as session:
user = session.query(UserRecord).filter_by(username=username).first() user = session.query(UserRecord).filter_by(username=username).first()
if user: if user:
return user.to_dict() return user.to_dict()
return None return None
def get_user_with_password(self, username: str) -> dict:
"""获取用户信息(含密码哈希,仅供密码验证内部使用)"""
with self.session() as session:
user = session.query(UserRecord).filter_by(username=username).first()
if user:
return user.to_dict_private()
return None
def get_user_by_id(self, user_id: int) -> dict: def get_user_by_id(self, user_id: int) -> dict:
"""通过ID获取用户信息""" """通过ID获取用户信息"""
with self.session() as session: with self.session() as session:
+120 -49
View File
@@ -97,10 +97,14 @@ class MirrorSyncManager:
self.sync_status = {} self.sync_status = {}
def save_sync_state(self): def save_sync_state(self):
"""保存同步状态""" """保存同步状态(快照 + 原子写,可在持锁状态下调用,不会死锁)"""
try: try:
with open(self.sync_state_file, 'w', encoding='utf-8') as f: # dict() 浅拷贝是单条 C 操作,并发修改安全
json.dump(self.sync_status, f, ensure_ascii=False, indent=2) snapshot = dict(self.sync_status)
tmp_file = self.sync_state_file + '.tmp'
with open(tmp_file, 'w', encoding='utf-8') as f:
json.dump(snapshot, f, ensure_ascii=False, indent=2)
os.replace(tmp_file, self.sync_state_file)
except Exception as e: except Exception as e:
print(f"保存同步状态失败: {e}") print(f"保存同步状态失败: {e}")
@@ -249,18 +253,18 @@ class MirrorSyncManager:
if name not in self.sync_sources: if name not in self.sync_sources:
return False return False
if name in self.sync_threads and self.sync_threads[name].is_alive():
return True # 已经在运行
# 设置运行标志,确保同步循环可以执行
self.running = True
with self.sync_lock: with self.sync_lock:
existing = self.sync_threads.get(name)
if existing and existing.is_alive():
return True # 已经在运行(检查+记录在同一把锁内,防双 worker 竞态)
# 设置运行标志,确保同步循环可以执行
self.running = True
self.sync_status[name]['status'] = 'syncing' self.sync_status[name]['status'] = 'syncing'
self.sync_status[name]['error'] = None self.sync_status[name]['error'] = None
thread = Thread(target=self._sync_worker, args=(name,), daemon=True) thread = Thread(target=self._sync_worker, args=(name,), daemon=True)
self.sync_threads[name] = thread self.sync_threads[name] = thread
thread.start() thread.start()
return True return True
@@ -280,12 +284,14 @@ class MirrorSyncManager:
if not packages or not isinstance(packages, list): if not packages or not isinstance(packages, list):
return {"success": False, "error": "请提供有效的包名列表"} return {"success": False, "error": "请提供有效的包名列表"}
# 生成临时任务ID # 生成临时任务ID(随机后缀防同秒碰撞)
import time import time
task_id = f"temp_sync_{int(time.time())}" import uuid
_ts = int(time.time())
task_id = f"temp_sync_{_ts}_{uuid.uuid4().hex[:6]}"
# 使用临时任务名进行同步 # 使用临时任务名进行同步
temp_name = f"{source_name}_temp_{int(time.time())}" temp_name = f"{source_name}_temp_{_ts}_{uuid.uuid4().hex[:6]}"
# 设置状态 # 设置状态
with self.sync_lock: with self.sync_lock:
@@ -333,17 +339,27 @@ class MirrorSyncManager:
self.sync_status[temp_name]['status'] = 'error' self.sync_status[temp_name]['status'] = 'error'
self.sync_status[temp_name]['error'] = str(e) self.sync_status[temp_name]['error'] = str(e)
print(f"[Temp Sync] 临时同步失败: {e}") print(f"[Temp Sync] 临时同步失败: {e}")
finally:
# 清理临时条目,防止 sync_state.json 无限膨胀
with self.sync_lock:
self.sync_status.pop(temp_name, None)
self.sync_threads.pop(temp_name, None)
self.save_sync_state()
def stop_sync(self, name): def stop_sync(self, name):
"""停止同步指定源 - 立即停止""" """停止同步指定源 - 请求停止
# 立即设置状态为停止
注意: 不删除 sync_threads 条目(删除后 start_sync 会再起一个 worker,
与仍在运行的旧线程形成双 worker)。线程结束后 is_alive() 为 False,
start_sync 会自然复用该条目。
"""
# 设置状态为停止(源可能已从 sync_status 移除,容错)
with self.sync_lock: with self.sync_lock:
self.sync_status[name]['status'] = 'stopped' status = self.sync_status.get(name)
if status is not None:
status['status'] = 'stopped'
# 设置停止标志,让线程提前退出 # 设置停止标志,让线程提前退出
self.running = False self.running = False
# 立即返回,不等待线程结束
if name in self.sync_threads:
del self.sync_threads[name]
# 短暂等待后重置运行标志 # 短暂等待后重置运行标志
import time import time
time.sleep(0.5) time.sleep(0.5)
@@ -364,7 +380,7 @@ class MirrorSyncManager:
print(f"[SYNC] 开始同步: {name}") print(f"[SYNC] 开始同步: {name}")
source_config = self.sync_sources[name] source_config = self.sync_sources[name]
sync_type = source_config.get('type', 'http') sync_type = source_config.get('type', 'http')
print(f"[SYNC] 类型: {sync_type}, URL: {source_config.get('url', 'N/A')}") print(f"[SYNC] 类型: {sync_type}, URL: {self._redact_url(source_config.get('url', 'N/A'))}")
try: try:
if sync_type in ('http', 'https'): if sync_type in ('http', 'https'):
@@ -394,8 +410,12 @@ class MirrorSyncManager:
raise ValueError(f"不支持的同步类型: {sync_type}") raise ValueError(f"不支持的同步类型: {sync_type}")
with self.sync_lock: with self.sync_lock:
self.sync_status[name]['status'] = 'completed' if not self.running:
self.sync_status[name]['last_sync'] = datetime.now().isoformat() # 被 stop_sync 中断:标记为停止而非完成
self.sync_status[name]['status'] = 'stopped'
else:
self.sync_status[name]['status'] = 'completed'
self.sync_status[name]['last_sync'] = datetime.now().isoformat()
except Exception as e: except Exception as e:
with self.sync_lock: with self.sync_lock:
@@ -403,6 +423,9 @@ class MirrorSyncManager:
self.sync_status[name]['error'] = str(e) self.sync_status[name]['error'] = str(e)
finally: finally:
# 清理线程条目(start_sync 通过 is_alive 判断, 清理后允许重新 start)
with self.sync_lock:
self.sync_threads.pop(name, None)
self.save_sync_state() self.save_sync_state()
def _sync_http(self, name, config, specific_packages=None): def _sync_http(self, name, config, specific_packages=None):
@@ -420,7 +443,7 @@ class MirrorSyncManager:
# 获取文件列表 # 获取文件列表
source_url = config.get('url', '') source_url = config.get('url', '')
print(f"[HTTP Sync] 源URL: {source_url}") print(f"[HTTP Sync] 源URL: {self._redact_url(source_url)}")
file_list = self._get_http_file_list(source_url, config, specific_packages) file_list = self._get_http_file_list(source_url, config, specific_packages)
total_files = len(file_list) total_files = len(file_list)
@@ -445,10 +468,13 @@ class MirrorSyncManager:
with self.sync_lock: with self.sync_lock:
self.sync_status[name]['files_synced'] = synced_count self.sync_status[name]['files_synced'] = synced_count
# 获取本地路径 # 获取本地路径(先做路径穿越校验)
filename = file_info.get('name', '') filename = file_info.get('name', '')
if not filename: if not filename:
continue continue
if self._safe_remote_name(filename) is None:
print(f"[HTTP Sync] 跳过不安全的文件名: {filename!r}")
continue
# 处理子目录 # 处理子目录
local_path = os.path.join(target_dir, filename) local_path = os.path.join(target_dir, filename)
@@ -476,6 +502,29 @@ class MirrorSyncManager:
print(f"[HTTP Sync] 完成: 成功 {synced_count}, 失败 {failed_count}") print(f"[HTTP Sync] 完成: 成功 {synced_count}, 失败 {failed_count}")
@staticmethod
def _safe_remote_name(filename):
"""校验远程文件名/相对路径,拒绝路径穿越(.. 段、绝对路径、反斜杠)"
允许 '/' 分隔的多级子路径(HTTP 列表常见),但拒绝任何 '..' 段。
"""
if not filename or filename in ('.', '..'):
return None
if filename.startswith('/') or '\\' in filename:
return None
for seg in filename.split('/'):
if seg in ('', '.', '..'):
return None
return filename
@staticmethod
def _redact_url(url):
"""脱敏 URL 中的 user:pass@ 段(防止凭证进日志/状态文件)"""
if not url:
return url
import re as _re
return _re.sub(r'(://)([^/@\s]+)@', r'\1***@', str(url))
def _sync_ftp(self, name, config): def _sync_ftp(self, name, config):
"""FTP同步实现""" """FTP同步实现"""
try: try:
@@ -519,15 +568,19 @@ class MirrorSyncManager:
for filename, file_info in remote_files.items(): for filename, file_info in remote_files.items():
if not self.running: if not self.running:
break break
local_path = os.path.join(target_dir, filename) safe_name = self._safe_remote_name(filename)
if safe_name is None:
continue
local_path = os.path.join(target_dir, safe_name)
if file_info['is_dir']: if file_info['is_dir']:
sub_config = config.copy() sub_config = config.copy()
sub_config['remote_path'] = os.path.join(remote_path, filename).replace('\\', '/') sub_config['remote_path'] = os.path.join(remote_path, safe_name).replace('\\', '/')
sub_config['target'] = os.path.join(config.get('target', name), filename) sub_config['target'] = os.path.join(config.get('target', name), safe_name)
self._sync_ftp(name + '/' + filename, sub_config) # 递归沿用顶层 name,避免子目录用假名导致 sync_status KeyError
self._sync_ftp(name, sub_config)
else: else:
if self._need_sync_ftp(ftp, filename, local_path, file_info): if self._need_sync_ftp(ftp, safe_name, local_path, file_info):
if self._download_ftp_file(ftp, filename, local_path, file_info): if self._download_ftp_file(ftp, safe_name, local_path, file_info):
synced_count += 1 synced_count += 1
with self.sync_lock: with self.sync_lock:
@@ -603,13 +656,12 @@ class MirrorSyncManager:
for file_attr in file_list: for file_attr in file_list:
if not self.running: if not self.running:
break break
remote_filename = file_attr.filename remote_filename = self._safe_remote_name(file_attr.filename)
if remote_filename is None:
continue
remote_filepath = os.path.join(remote_path, remote_filename).replace('\\', '/') remote_filepath = os.path.join(remote_path, remote_filename).replace('\\', '/')
local_filepath = os.path.join(local_path, remote_filename) local_filepath = os.path.join(local_path, remote_filename)
if remote_filename in ['.', '..']:
continue
if file_attr.st_mode & 0o40000: if file_attr.st_mode & 0o40000:
sub_synced = self._sync_sftp_directory(sftp, remote_filepath, local_filepath, sync_name) sub_synced = self._sync_sftp_directory(sftp, remote_filepath, local_filepath, sync_name)
synced_count += sub_synced synced_count += sub_synced
@@ -679,7 +731,13 @@ class MirrorSyncManager:
import subprocess import subprocess
source = config.get('source', config.get('url', '')) source = config.get('source', config.get('url', ''))
target_dir = os.path.join(self.config['base_dir'], config.get('target', name)) target_rel = config.get('target', name) or name
if target_rel in ('.', './', '') or os.path.isabs(target_rel):
raise Exception(f"Rsync 同步目标不安全: {target_rel!r} (禁止同步到 '.' 或绝对路径, 防止 --delete 清空目录)")
target_dir = os.path.join(self.config['base_dir'], target_rel)
# 目标目录不得等于 base_dir 本身
if os.path.realpath(target_dir) == os.path.realpath(self.config['base_dir']):
raise Exception("Rsync 同步目标不允许是下载根目录本身 (--delete 会清空全部文件)")
os.makedirs(target_dir, exist_ok=True) os.makedirs(target_dir, exist_ok=True)
print(f"开始Rsync同步 {name} -> {target_dir}") print(f"开始Rsync同步 {name} -> {target_dir}")
@@ -725,10 +783,15 @@ class MirrorSyncManager:
if os.path.exists(os.path.join(target_dir, '.git')): if os.path.exists(os.path.join(target_dir, '.git')):
# 已存在,执行git pull # 已存在,执行git pull
try: try:
subprocess.run(['git', 'fetch', '--all'], cwd=target_dir, check=True, capture_output=True) git_timeout = config.get('git_timeout', 1800)
subprocess.run(['git', 'reset', '--hard', f'origin/{branch}'], cwd=target_dir, check=True, capture_output=True) subprocess.run(['git', 'fetch', '--all'], cwd=target_dir, check=True,
capture_output=True, timeout=git_timeout)
subprocess.run(['git', 'reset', '--hard', f'origin/{branch}'], cwd=target_dir,
check=True, capture_output=True, timeout=git_timeout)
except subprocess.CalledProcessError as e: except subprocess.CalledProcessError as e:
raise Exception(f"Git pull失败: {e}") raise Exception(f"Git pull失败: {e}")
except subprocess.TimeoutExpired:
raise Exception(f"Git pull超时({git_timeout}s)")
else: else:
# 克隆新仓库 # 克隆新仓库
cmd = ['git', 'clone'] cmd = ['git', 'clone']
@@ -739,9 +802,12 @@ class MirrorSyncManager:
cmd.extend([repo_url, target_dir]) cmd.extend([repo_url, target_dir])
try: try:
subprocess.run(cmd, check=True, capture_output=True) git_timeout = config.get('git_timeout', 1800)
subprocess.run(cmd, check=True, capture_output=True, timeout=git_timeout)
except subprocess.CalledProcessError as e: except subprocess.CalledProcessError as e:
raise Exception(f"Git clone失败: {e}") raise Exception(f"Git clone失败: {e}")
except subprocess.TimeoutExpired:
raise Exception(f"Git clone超时({git_timeout}s)")
def _sync_s3(self, name, config): def _sync_s3(self, name, config):
"""AWS S3兼容存储同步 (S3/OSS/COS/MinIO等)""" """AWS S3兼容存储同步 (S3/OSS/COS/MinIO等)"""
@@ -888,7 +954,7 @@ class MirrorSyncManager:
if is_pypi: if is_pypi:
print(f"[HTTP] 检测到PyPI索引,使用PyPI专用解析") print(f"[HTTP] 检测到PyPI索引,使用PyPI专用解析")
return self._get_pypi_file_list(base_url, config, specific_packages) return self._get_pypi_file_list(base_url, config, specific_packages, sync_name=name)
if 'api_url' in config: if 'api_url' in config:
try: try:
@@ -898,7 +964,7 @@ class MirrorSyncManager:
encoded_auth = base64.b64encode(auth_string.encode()).decode() encoded_auth = base64.b64encode(auth_string.encode()).decode()
req.add_header('Authorization', f'Basic {encoded_auth}') req.add_header('Authorization', f'Basic {encoded_auth}')
with urllib.request.urlopen(req) as response: with urllib.request.urlopen(req, timeout=60) as response:
api_data = json.loads(response.read().decode()) api_data = json.loads(response.read().decode())
if 'files' in api_data: if 'files' in api_data:
for file_info in api_data['files']: for file_info in api_data['files']:
@@ -920,7 +986,7 @@ class MirrorSyncManager:
return file_list return file_list
def _get_pypi_file_list(self, base_url, config, specific_packages=None): def _get_pypi_file_list(self, base_url, config, specific_packages=None, sync_name=None):
"""获取PyPI文件列表 """获取PyPI文件列表
Args: Args:
@@ -986,12 +1052,10 @@ class MirrorSyncManager:
if not self.running: if not self.running:
break break
# 更新进度 # 更新进度(只更新当前源,避免污染其他源状态)
with self.sync_lock: with self.sync_lock:
if hasattr(self, 'sync_status') and self.sync_status: if sync_name and sync_name in self.sync_status:
for name in self.sync_status: self.sync_status[sync_name]['files_synced'] = i
if 'total_files' in self.sync_status[name]:
self.sync_status[name]['files_synced'] = i
# 获取每个包的文件列表 # 获取每个包的文件列表
package_url = f"{package_index_url}{package_name}/" package_url = f"{package_index_url}{package_name}/"
@@ -1098,8 +1162,11 @@ class MirrorSyncManager:
"""检查HTTP文件是否需要同步""" """检查HTTP文件是否需要同步"""
if not os.path.exists(local_path): if not os.path.exists(local_path):
return True return True
local_size = os.path.getsize(local_path)
remote_size = file_info.get('size', 0) remote_size = file_info.get('size', 0)
if remote_size <= 0:
# 远端大小未知(PyPI/HTML 列表恒为 0):仅按存在性判断,避免每次全量重下
return False
local_size = os.path.getsize(local_path)
return local_size != remote_size return local_size != remote_size
def _need_sync_ftp(self, ftp, remote_filename, local_path, file_info): def _need_sync_ftp(self, ftp, remote_filename, local_path, file_info):
@@ -1213,7 +1280,7 @@ class MirrorSyncManager:
return False return False
except Exception as e: except Exception as e:
print(f"下载HTTP文件失败 {url}: {e}") print(f"下载HTTP文件失败 {self._redact_url(url)}: {e}")
if os.path.exists(temp_path): if os.path.exists(temp_path):
os.remove(temp_path) os.remove(temp_path)
return False return False
@@ -1230,6 +1297,10 @@ class MirrorSyncManager:
else: else:
return True return True
# 防 FTP 命令注入: 文件名含 CR/LF 时拒绝
if any(c in remote_filename for c in ('\r', '\n')):
raise Exception(f"文件名含非法字符: {remote_filename!r}")
with open(local_path, mode) as f: with open(local_path, mode) as f:
if start_pos > 0: if start_pos > 0:
ftp.voidcmd('TYPE I') ftp.voidcmd('TYPE I')
+387 -380
View File
@@ -1,380 +1,387 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
""" """
定时任务调度器 定时任务调度器
支持 cron 表达式和简单间隔的定时任务 支持 cron 表达式和简单间隔的定时任务
""" """
import os import os
import time import time
import logging import logging
import threading import threading
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Dict, List, Optional, Callable from typing import Dict, List, Optional, Callable
from enum import Enum from enum import Enum
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class TaskStatus(Enum): class TaskStatus(Enum):
"""任务状态""" """任务状态"""
IDLE = "idle" IDLE = "idle"
RUNNING = "running" RUNNING = "running"
ERROR = "error" ERROR = "error"
DISABLED = "disabled" DISABLED = "disabled"
class ScheduledTask: class ScheduledTask:
"""定时任务""" """定时任务"""
def __init__(self, name: str, task_type: str, config: dict, def __init__(self, name: str, task_type: str, config: dict,
callback: Callable, logger=None): callback: Callable, logger=None):
""" """
初始化定时任务 初始化定时任务
Args: Args:
name: 任务名称 name: 任务名称
task_type: 任务类型 ('cron' 或 'interval') task_type: 任务类型 ('cron' 或 'interval')
config: 任务配置 config: 任务配置
callback: 回调函数 callback: 回调函数
logger: 日志器 logger: 日志器
""" """
self.name = name self.name = name
self.task_type = task_type # 'cron' 或 'interval' self.task_type = task_type # 'cron' 或 'interval'
self.config = config or {} self.config = config or {}
self.callback = callback self.callback = callback
self.logger = logger or logging.getLogger(__name__) self.logger = logger or logging.getLogger(__name__)
# 状态 # 状态
self.status = TaskStatus.IDLE self.status = TaskStatus.IDLE
self.last_run: Optional[datetime] = None self.last_run: Optional[datetime] = None
self.next_run: Optional[datetime] = None self.next_run: Optional[datetime] = None
self.last_error: Optional[str] = None self.last_error: Optional[str] = None
self.run_count = 0 self.run_count = 0
# 配置解析 # 配置解析
self._parse_config() self._parse_config()
def _parse_config(self): def _parse_config(self):
"""解析任务配置""" """解析任务配置"""
if self.task_type == 'cron': if self.task_type == 'cron':
# Cron 表达式: "minute hour day month weekday" # Cron 表达式: "minute hour day month weekday"
# 例如: "0 3 * * *" 每天凌晨3点 # 例如: "0 3 * * *" 每天凌晨3点
cron = self.config.get('cron', '0 0 * * *') cron = self.config.get('cron', '0 0 * * *')
parts = cron.split() parts = cron.split()
if len(parts) == 5: if len(parts) == 5:
self.cron_parts = { self.cron_parts = {
'minute': self._parse_cron_part(parts[0], 0, 59), 'minute': self._parse_cron_part(parts[0], 0, 59),
'hour': self._parse_cron_part(parts[1], 0, 23), 'hour': self._parse_cron_part(parts[1], 0, 23),
'day': self._parse_cron_part(parts[2], 1, 31), 'day': self._parse_cron_part(parts[2], 1, 31),
'month': self._parse_cron_part(parts[3], 1, 12), 'month': self._parse_cron_part(parts[3], 1, 12),
'weekday': self._parse_cron_part(parts[4], 0, 6) 'weekday': self._parse_cron_part(parts[4], 0, 6)
} }
else: else:
self.logger.warning(f"无效的 cron 表达式: {cron}") self.logger.warning(f"无效的 cron 表达式: {cron}")
self.cron_parts = None self.cron_parts = None
elif self.task_type == 'interval': elif self.task_type == 'interval':
# 间隔: seconds, minutes, hours # 间隔: seconds, minutes, hours
interval = self.config.get('interval', {}) interval = self.config.get('interval', {})
self.interval_seconds = ( self.interval_seconds = (
interval.get('seconds', 0) + interval.get('seconds', 0) +
interval.get('minutes', 0) * 60 + interval.get('minutes', 0) * 60 +
interval.get('hours', 0) * 3600 + interval.get('hours', 0) * 3600 +
interval.get('days', 0) * 86400 interval.get('days', 0) * 86400
) )
if self.interval_seconds <= 0: if self.interval_seconds <= 0:
self.interval_seconds = 3600 # 默认1小时 self.interval_seconds = 3600 # 默认1小时
# 是否启用 # 是否启用
self.enabled = self.config.get('enabled', True) self.enabled = self.config.get('enabled', True)
def _parse_cron_part(self, part: str, min_val: int, max_val: int) -> List[int]: def _parse_cron_part(self, part: str, min_val: int, max_val: int) -> List[int]:
"""解析 cron 表达式的一部分""" """解析 cron 表达式的一部分"""
result = [] result = []
if part == '*': if part == '*':
return list(range(min_val, max_val + 1)) return list(range(min_val, max_val + 1))
# 处理列表: "1,2,3" # 处理列表: "1,2,3"
if ',' in part: if ',' in part:
result = [] result = []
for sub in part.split(','): for sub in part.split(','):
result.extend(self._parse_cron_part(sub.strip(), min_val, max_val)) result.extend(self._parse_cron_part(sub.strip(), min_val, max_val))
return result return result
# 处理范围: "1-5" # 处理范围: "1-5"
if '-' in part: if '-' in part:
start, end = part.split('-') start, end = part.split('-')
return list(range(int(start), int(end) + 1)) return list(range(int(start), int(end) + 1))
# 处理步进: "*/5" # 处理步进: "*/5"
if '/' in part: if '/' in part:
base, step = part.split('/') base, step = part.split('/')
base_list = self._parse_cron_part(base or '*', min_val, max_val) base_list = self._parse_cron_part(base or '*', min_val, max_val)
step = int(step) step = int(step)
return base_list[::step] return base_list[::step]
# 单个值 # 单个值
try: try:
val = int(part) val = int(part)
if min_val <= val <= max_val: if min_val <= val <= max_val:
return [val] return [val]
except ValueError: except ValueError:
pass pass
return [] return []
def should_run_now(self) -> bool: def should_run_now(self) -> bool:
"""检查是否应该在当前时刻运行""" """检查是否应该在当前时刻运行"""
if not self.enabled: if not self.enabled:
return False return False
now = datetime.now() now = datetime.now()
if self.task_type == 'cron' and self.cron_parts: if self.task_type == 'cron' and self.cron_parts:
return self._matches_cron(now) if not self._matches_cron(now):
elif self.task_type == 'interval': return False
if self.last_run is None: # 同一分钟内只触发一次(轮询间隔 < 60s 时避免重复触发)
return True if self.last_run is not None:
elapsed = (now - self.last_run).total_seconds() last_minute = self.last_run.strftime('%Y-%m-%d %H:%M')
return elapsed >= self.interval_seconds if last_minute == now.strftime('%Y-%m-%d %H:%M'):
return False
return False return True
elif self.task_type == 'interval':
def _matches_cron(self, dt: datetime) -> bool: if self.last_run is None:
"""检查时间是否匹配 cron 表达式""" return True
if not self.cron_parts: elapsed = (now - self.last_run).total_seconds()
return False return elapsed >= self.interval_seconds
# cron 约定: 0=周日...6=周六; datetime.weekday(): 0=周一...6=周日 return False
cron_weekday = (dt.weekday() + 1) % 7
def _matches_cron(self, dt: datetime) -> bool:
return ( """检查时间是否匹配 cron 表达式"""
dt.minute in self.cron_parts['minute'] and if not self.cron_parts:
dt.hour in self.cron_parts['hour'] and return False
dt.day in self.cron_parts['day'] and
dt.month in self.cron_parts['month'] and # cron 约定: 0=周日...6=周六; datetime.weekday(): 0=周一...6=周日
cron_weekday in self.cron_parts['weekday'] cron_weekday = (dt.weekday() + 1) % 7
)
return (
def get_next_run_time(self) -> Optional[datetime]: dt.minute in self.cron_parts['minute'] and
"""计算下次运行时间""" dt.hour in self.cron_parts['hour'] and
if not self.enabled: dt.day in self.cron_parts['day'] and
return None dt.month in self.cron_parts['month'] and
cron_weekday in self.cron_parts['weekday']
now = datetime.now() )
if self.task_type == 'cron' and self.cron_parts: def get_next_run_time(self) -> Optional[datetime]:
# 找到下一个匹配的时间点 """计算下次运行时间"""
for i in range(365 * 24 * 60): # 最多查找1年 if not self.enabled:
candidate = now + timedelta(minutes=i) return None
if self._matches_cron(candidate):
return candidate now = datetime.now()
elif self.task_type == 'interval':
if self.last_run: if self.task_type == 'cron' and self.cron_parts:
return self.last_run + timedelta(seconds=self.interval_seconds) # 找到下一个匹配的时间点
return now for i in range(365 * 24 * 60): # 最多查找1年
candidate = now + timedelta(minutes=i)
return None if self._matches_cron(candidate):
return candidate
def run(self) -> bool: elif self.task_type == 'interval':
"""执行任务""" if self.last_run:
if self.status == TaskStatus.RUNNING: return self.last_run + timedelta(seconds=self.interval_seconds)
self.logger.warning(f"任务 {self.name} 已在运行中") return now
return False
return None
self.status = TaskStatus.RUNNING
self.last_run = datetime.now() def run(self) -> bool:
self.last_error = None """执行任务"""
if self.status == TaskStatus.RUNNING:
try: self.logger.warning(f"任务 {self.name} 已在运行中")
self.logger.info(f"开始执行定时任务: {self.name}") return False
result = self.callback(self.name, self.config)
self.run_count += 1 self.status = TaskStatus.RUNNING
self.logger.info(f"定时任务 {self.name} 执行完成") self.last_run = datetime.now()
return True self.last_error = None
except Exception as e:
self.last_error = str(e) try:
self.status = TaskStatus.ERROR self.logger.info(f"开始执行定时任务: {self.name}")
self.logger.error(f"定时任务 {self.name} 执行失败: {e}") result = self.callback(self.name, self.config)
return False self.run_count += 1
finally: self.logger.info(f"定时任务 {self.name} 执行完成")
if self.status != TaskStatus.ERROR: return True
self.status = TaskStatus.IDLE except Exception as e:
self.last_error = str(e)
def to_dict(self) -> dict: self.status = TaskStatus.ERROR
"""转换为字典""" self.logger.error(f"定时任务 {self.name} 执行失败: {e}")
return { return False
'name': self.name, finally:
'type': self.task_type, if self.status != TaskStatus.ERROR:
'enabled': self.enabled, self.status = TaskStatus.IDLE
'status': self.status.value,
'config': self.config, def to_dict(self) -> dict:
'last_run': self.last_run.isoformat() if self.last_run else None, """转换为字典"""
'next_run': self.next_run.isoformat() if self.next_run else None, return {
'run_count': self.run_count, 'name': self.name,
'last_error': self.last_error 'type': self.task_type,
} 'enabled': self.enabled,
'status': self.status.value,
'config': self.config,
class Scheduler: 'last_run': self.last_run.isoformat() if self.last_run else None,
"""定时任务调度器""" 'next_run': self.next_run.isoformat() if self.next_run else None,
'run_count': self.run_count,
def __init__(self, config: dict = None): 'last_error': self.last_error
self.config = config or {} }
self.tasks: Dict[str, ScheduledTask] = {}
self._running = False
self._thread: Optional[threading.Thread] = None class Scheduler:
self._lock = threading.Lock() """定时任务调度器"""
# 默认检查间隔 def __init__(self, config: dict = None):
self.check_interval = self.config.get('check_interval', 10) self.config = config or {}
self.tasks: Dict[str, ScheduledTask] = {}
# 事件回调 self._running = False
self.on_task_start: Optional[Callable] = None self._thread: Optional[threading.Thread] = None
self.on_task_complete: Optional[Callable] = None self._lock = threading.Lock()
self.on_task_error: Optional[Callable] = None
# 默认检查间隔
def add_task(self, name: str, task_type: str, config: dict, self.check_interval = self.config.get('check_interval', 10)
callback: Callable) -> bool:
""" # 事件回调
添加定时任务 self.on_task_start: Optional[Callable] = None
self.on_task_complete: Optional[Callable] = None
Args: self.on_task_error: Optional[Callable] = None
name: 任务名称
task_type: 任务类型 ('cron' 或 'interval') def add_task(self, name: str, task_type: str, config: dict,
config: 任务配置 callback: Callable) -> bool:
callback: 回调函数 """
添加定时任务
Returns:
是否成功 Args:
""" name: 任务名称
with self._lock: task_type: 任务类型 ('cron' 或 'interval')
if name in self.tasks: config: 任务配置
logger.warning(f"任务 {name} 已存在,将被替换") callback: 回调函数
self.tasks[name] = ScheduledTask(name, task_type, config, callback, logger)
return True Returns:
是否成功
def remove_task(self, name: str) -> bool: """
"""移除任务""" with self._lock:
with self._lock: if name in self.tasks:
if name in self.tasks: logger.warning(f"任务 {name} 已存在,将被替换")
del self.tasks[name] self.tasks[name] = ScheduledTask(name, task_type, config, callback, logger)
return True return True
return False
def remove_task(self, name: str) -> bool:
def get_task(self, name: str) -> Optional[ScheduledTask]: """移除任务"""
"""获取任务""" with self._lock:
return self.tasks.get(name) if name in self.tasks:
del self.tasks[name]
def get_all_tasks(self) -> List[dict]: return True
"""获取所有任务状态""" return False
with self._lock:
for task in self.tasks.values(): def get_task(self, name: str) -> Optional[ScheduledTask]:
task.next_run = task.get_next_run_time() """获取任务"""
return [task.to_dict() for task in self.tasks.values()] return self.tasks.get(name)
def start(self): def get_all_tasks(self) -> List[dict]:
"""启动调度器""" """获取所有任务状态"""
if self._running: with self._lock:
logger.warning("调度器已在运行中") for task in self.tasks.values():
return task.next_run = task.get_next_run_time()
return [task.to_dict() for task in self.tasks.values()]
self._running = True
self._thread = threading.Thread(target=self._run_loop, daemon=True) def start(self):
self._thread.start() """启动调度器"""
logger.info("定时任务调度器已启动") if self._running:
logger.warning("调度器已在运行中")
def stop(self): return
"""停止调度器"""
self._running = False self._running = True
if self._thread: self._thread = threading.Thread(target=self._run_loop, daemon=True)
self._thread.join(timeout=5) self._thread.start()
logger.info("定时任务调度器已停止") logger.info("定时任务调度器已启动")
def _run_loop(self): def stop(self):
"""运行循环""" """停止调度器"""
while self._running: self._running = False
try: if self._thread:
now = datetime.now() self._thread.join(timeout=5)
logger.info("定时任务调度器已停止")
with self._lock:
for name, task in self.tasks.items(): def _run_loop(self):
if task.should_run_now(): """运行循环"""
# 使用线程池执行任务 while self._running:
from concurrent.futures import ThreadPoolExecutor try:
with ThreadPoolExecutor(max_workers=1) as executor: now = datetime.now()
executor.submit(task.run)
with self._lock:
time.sleep(self.check_interval) for name, task in self.tasks.items():
if task.should_run_now():
except Exception as e: # 使用线程池执行任务
logger.error(f"调度器循环错误: {e}") from concurrent.futures import ThreadPoolExecutor
time.sleep(5) with ThreadPoolExecutor(max_workers=1) as executor:
executor.submit(task.run)
def run_task_now(self, name: str) -> bool:
"""立即运行指定任务""" time.sleep(self.check_interval)
task = self.get_task(name)
if task: except Exception as e:
return task.run() logger.error(f"调度器循环错误: {e}")
return False time.sleep(5)
def enable_task(self, name: str, enabled: bool = True) -> bool: def run_task_now(self, name: str) -> bool:
"""启用/禁用任务""" """立即运行指定任务"""
task = self.get_task(name) task = self.get_task(name)
if task: if task:
task.enabled = enabled return task.run()
return True return False
return False
def enable_task(self, name: str, enabled: bool = True) -> bool:
def update_task_config(self, name: str, config: dict) -> bool: """启用/禁用任务"""
"""更新任务配置""" task = self.get_task(name)
task = self.get_task(name) if task:
if task: task.enabled = enabled
task.config.update(config) return True
task._parse_config() return False
return True
return False def update_task_config(self, name: str, config: dict) -> bool:
"""更新任务配置"""
task = self.get_task(name)
# ==================== 同步任务工厂 ==================== if task:
task.config.update(config)
# def create_sync_task_callback(sync_manager): task._parse_config()
# """创建同步任务的回调函数""" return True
# def sync_task_callback(task_name: str, config: dict): return False
# """同步任务回调"""
# sync_manager.start_sync(task_name)
# return True # ==================== 同步任务工厂 ====================
# return sync_task_callback
# def create_sync_task_callback(sync_manager):
# """创建同步任务的回调函数"""
# ==================== 默认任务配置 ==================== # def sync_task_callback(task_name: str, config: dict):
# DEFAULT_SCHEDULED_TASKS = { ... } # """同步任务回调"""
# sync_manager.start_sync(task_name)
DEFAULT_SCHEDULED_TASKS = { # return True
# 数据库清理 - 每天凌晨2点 # return sync_task_callback
'cleanup_db': {
'type': 'cron',
'config': { # ==================== 默认任务配置 ====================
'cron': '0 2 * * *', # DEFAULT_SCHEDULED_TASKS = { ... }
'enabled': True
} DEFAULT_SCHEDULED_TASKS = {
}, # 数据库清理 - 每天凌晨2点
# 缓存清理 - 每6小时 'cleanup_db': {
'cleanup_cache': { 'type': 'cron',
'type': 'interval', 'config': {
'config': { 'cron': '0 2 * * *',
'interval': {'hours': 6}, 'enabled': True
'enabled': True }
} },
}, # 缓存清理 - 每6小时
# 健康检查 - 每5分钟 'cleanup_cache': {
'health_check': { 'type': 'interval',
'type': 'interval', 'config': {
'config': { 'interval': {'hours': 6},
'interval': {'minutes': 5}, 'enabled': True
'enabled': True }
} },
} # 健康检查 - 每5分钟
} 'health_check': {
'type': 'interval',
'config': {
'interval': {'minutes': 5},
'enabled': True
}
}
}
+1 -1
View File
@@ -95,7 +95,7 @@ class MirrorServer:
else: else:
self.config_manager = config self.config_manager = config
self.config = self.config_manager.config self.config = self.config_manager.get_all()
self.server = None self.server = None
self.sync_manager = None self.sync_manager = None
self.is_running = False self.is_running = False
+4 -2
View File
@@ -935,8 +935,10 @@ class SyncEngine:
if task.status in [SyncStatus.COMPLETED, SyncStatus.FAILED, SyncStatus.CANCELLED] if task.status in [SyncStatus.COMPLETED, SyncStatus.FAILED, SyncStatus.CANCELLED]
] ]
for task_id in completed_ids[-keep_count:]: # 删除最旧的已完成任务,保留最近 keep_count 个(原逻辑删最新留最旧,逻辑倒置)
del self.active_tasks[task_id] if len(completed_ids) > keep_count:
for task_id in completed_ids[:len(completed_ids) - keep_count]:
self.active_tasks.pop(task_id, None)
class SyncSource: class SyncSource:
+13 -8
View File
@@ -12,6 +12,7 @@ import mimetypes
import base64 import base64
import hashlib import hashlib
import shutil import shutil
import threading
from datetime import datetime from datetime import datetime
from http.server import BaseHTTPRequestHandler from http.server import BaseHTTPRequestHandler
from urllib.parse import unquote, urlparse, parse_qs from urllib.parse import unquote, urlparse, parse_qs
@@ -40,6 +41,7 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
debug_log_file = None # 调试日志文件路径 debug_log_file = None # 调试日志文件路径
_debug_categories = set() # 启用的调试类别 _debug_categories = set() # 启用的调试类别
_mirror_handlers = {} # 镜像处理器实例缓存 _mirror_handlers = {} # 镜像处理器实例缓存
_stats_lock = threading.Lock() # 保护 JSON 统计/历史文件读改写
@classmethod @classmethod
def _setup_debug(cls, config): def _setup_debug(cls, config):
@@ -1421,10 +1423,11 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
except Exception as e: except Exception as e:
print(f"Error updating download count in database: {e}") print(f"Error updating download count in database: {e}")
# 回退到 JSON 文件 # 回退到 JSON 文件(加锁防并发读改写丢计数)
stats = self.load_stats() with self._stats_lock:
stats[filepath] = stats.get(filepath, 0) + 1 stats = self.load_stats()
self.save_stats(stats) stats[filepath] = stats.get(filepath, 0) + 1
self.save_stats(stats)
# ==================== 下载历史记录 ==================== # ==================== 下载历史记录 ====================
@@ -1501,8 +1504,9 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
except Exception as e: except Exception as e:
print(f"Error logging download to database: {e}") print(f"Error logging download to database: {e}")
# 回退到 JSON 文件 # 回退到 JSON 文件(加锁防并发写坏/丢记录)
history = self.load_download_history(1000) with self._stats_lock:
history = self.load_download_history(1000)
entry = { entry = {
'timestamp': datetime.now().isoformat(), 'timestamp': datetime.now().isoformat(),
@@ -1513,5 +1517,6 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
'method': self.command if hasattr(self, 'command') else 'GET' 'method': self.command if hasattr(self, 'command') else 'GET'
} }
history.append(entry) with self._stats_lock:
self.save_download_history(history) history.append(entry)
self.save_download_history(history)