From 82875b710a76181af8ae6d7990f619c23ac0cdd9 Mon Sep 17 00:00:00 2001 From: HYC Fixer Date: Wed, 2 Sep 2026 00:39:12 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A4=8D=E6=9F=A5=E4=BF=AE=E5=A4=8D(=E4=BA=8C)?= =?UTF-8?q?:=20=E5=90=8C=E6=AD=A5=E6=A8=A1=E5=9D=97=E4=B8=8E=E5=AE=89?= =?UTF-8?q?=E5=85=A8=E9=81=97=E7=95=99?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 修复 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 删旧留新 --- api/v1.py | 9 +- api/v2.py | 82 +++-- core/api_auth.py | 4 + core/config_hotreload.py | 687 ++++++++++++++++++----------------- core/database.py | 20 +- core/mirror_sync.py | 169 ++++++--- core/scheduler.py | 767 ++++++++++++++++++++------------------- core/server.py | 2 +- core/sync_engine.py | 6 +- handlers/http_handler.py | 21 +- 10 files changed, 957 insertions(+), 810 deletions(-) diff --git a/api/v1.py b/api/v1.py index 7baf976..c2ae244 100644 --- a/api/v1.py +++ b/api/v1.py @@ -1297,6 +1297,10 @@ class APIv1: return 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) 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): handler.send_json_response({"error": "Unsafe archive contents"}, 403) return + members = zipf.namelist() zipf.extractall(extract_dir) # 同步提取的文件到数据库 if self.db: extracted_files = [] - for member in zipf.namelist(): + for member in members: member_path = os.path.join(extract_dir, member) if os.path.isfile(member_path): rel_member_path = os.path.relpath(member_path, self.config['base_dir']).replace("\\", "/") @@ -1327,7 +1332,7 @@ class APIv1: handler.send_json_response({ "operation": "extract", "extract_dir": target_dir, - "extracted_files": len(zipf.namelist()) + "extracted_files": len(members) }) except Exception as e: diff --git a/api/v2.py b/api/v2.py index 6c6b9cb..06e042d 100644 --- a/api/v2.py +++ b/api/v2.py @@ -24,6 +24,37 @@ class APIv2(APIv1): # 管理员API处理器(始终创建,auth_type检查在装饰器中处理) 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): """处理API v2请求""" import sys @@ -271,7 +302,7 @@ class APIv2(APIv1): elif path == 'health/stats': if method == 'GET': 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()) else: handler.send_error(405) @@ -2283,11 +2314,12 @@ class APIv2(APIv1): if source_name is None: return super().api_start_sync(handler) if hasattr(handler, 'sync_manager') and handler.sync_manager: - task_id = handler.sync_manager.start_sync(source_name) - if task_id: + # MirrorSyncManager.start_sync 返回 bool(True=已启动/已在运行) + started = handler.sync_manager.start_sync(source_name) + if started: handler.send_json_response({ "success": True, - "task_id": task_id, + "started": True, "source_name": source_name }) else: @@ -2573,7 +2605,7 @@ class APIv2(APIv1): from core.health_check import HealthChecker, HealthStatus mirrors = self.config.get('mirrors', {}) - checker = HealthChecker(self.config.get('health_check', {})) + checker = self._get_health_checker() results = [] for mirror_type, mirror_config in mirrors.items(): @@ -2614,7 +2646,7 @@ class APIv2(APIv1): try: from core.health_check import HealthChecker, HealthStatus - checker = HealthChecker(self.config.get('health_check', {})) + checker = self._get_health_checker() mirrors = self.config.get('mirrors', {}) # 查找源对应的镜像类型 @@ -2659,7 +2691,7 @@ class APIv2(APIv1): try: from core.health_check import MirrorFailoverManager - failover = MirrorFailoverManager(self.config) + failover = self._get_failover_manager() failover.initialize() handler.send_json_response({ @@ -2679,7 +2711,7 @@ class APIv2(APIv1): try: from core.health_check import MirrorFailoverManager - failover = MirrorFailoverManager(self.config) + failover = self._get_failover_manager() failover.initialize() success = failover.perform_failover(mirror_type) @@ -3285,7 +3317,7 @@ class APIv2(APIv1): # 数据库验证 if db: - user = db.get_user(username) + user = db.get_user_with_password(username) if user and db.verify_password(password, user['password_hash']): # 数据库验证成功,生成 token import secrets @@ -3388,7 +3420,7 @@ class APIv2(APIv1): 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 not db.verify_password(old_password, user['password_hash']): @@ -3985,7 +4017,7 @@ class APIv2(APIv1): """获取重启状态""" 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() handler.send_json_response({ @@ -3999,7 +4031,7 @@ class APIv2(APIv1): """获取待处理请求""" 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() handler.send_json_response({ @@ -4018,7 +4050,7 @@ class APIv2(APIv1): except ValueError: strategy = RestartStrategy.GRACEFUL - restart_manager = GracefulRestartManager(self.config.get('restart', {})) + restart_manager = self._get_restart_manager() # 准备重启 prepare_result = restart_manager.prepare_restart() @@ -4061,7 +4093,7 @@ class APIv2(APIv1): except json.JSONDecodeError: 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) @@ -4072,7 +4104,7 @@ class APIv2(APIv1): """立即重启服务器""" 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') @@ -4088,7 +4120,7 @@ class APIv2(APIv1): """获取重启历史""" 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() handler.send_json_response({ @@ -4182,7 +4214,7 @@ class APIv2(APIv1): """获取缓存预热状态""" from core.cache_prewarm import CachePrewarmer - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer = self._get_prewarmer() status = prewarmer.get_status() handler.send_json_response(status) @@ -4191,7 +4223,7 @@ class APIv2(APIv1): """获取缓存预热统计""" from core.cache_prewarm import CachePrewarmer - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer = self._get_prewarmer() stats = prewarmer.get_stats() handler.send_json_response(stats) @@ -4205,7 +4237,7 @@ class APIv2(APIv1): mirror_type = query_params.get('mirror_type', [None])[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) handler.send_json_response({ @@ -4217,7 +4249,7 @@ class APIv2(APIv1): """获取预热历史""" from core.cache_prewarm import CachePrewarmer - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer = self._get_prewarmer() history = prewarmer.get_history() handler.send_json_response({ @@ -4234,7 +4266,7 @@ class APIv2(APIv1): limit = int(query_params.get('limit', [50])[0]) priority = query_params.get('priority', ['medium'])[0] - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer = self._get_prewarmer() # 如果指定了镜像类型,只预热该类型 targets = None @@ -4277,7 +4309,7 @@ class APIv2(APIv1): 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) handler.send_json_response({ @@ -4313,7 +4345,7 @@ class APIv2(APIv1): }, 400) return - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer = self._get_prewarmer() popular = prewarmer.get_popular_items(mirror_type) if limit: @@ -4335,7 +4367,7 @@ class APIv2(APIv1): mirror_type = query_params.get('mirror_type', [None])[0] - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer = self._get_prewarmer() if mirror_type: items = prewarmer.get_popular_items(mirror_type) @@ -4355,7 +4387,7 @@ class APIv2(APIv1): """清空预热队列""" from core.cache_prewarm import CachePrewarmer - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer = self._get_prewarmer() prewarmer.clear_items() handler.send_json_response({ diff --git a/core/api_auth.py b/core/api_auth.py index cf964da..35abb6e 100644 --- a/core/api_auth.py +++ b/core/api_auth.py @@ -396,6 +396,10 @@ class APIAuthManager: ) with self._lock: + # 顺带清理过期会话(防止字典无限增长,配合懒清理) + now = time.time() + for sid in [sid for sid, sess in self.sessions.items() if sess.expires_at <= now]: + del self.sessions[sid] self.sessions[session_id] = session self._save_sessions() diff --git a/core/config_hotreload.py b/core/config_hotreload.py index e92b1c7..63ec349 100644 --- a/core/config_hotreload.py +++ b/core/config_hotreload.py @@ -1,340 +1,347 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - -""" -配置热更新模块 -支持不重启服务的情况下重新加载配置 -""" - -import os -import sys -import json -import time -import threading -import logging -from typing import Dict, Any, Optional, Callable -from datetime import datetime - -logger = logging.getLogger(__name__) - - -class ConfigHotReloader: - """配置热重载管理器""" - - def __init__(self, config_path: str, callback: Callable = None): - """ - 初始化热重载管理器 - - Args: - config_path: 配置文件路径 - callback: 配置变更时的回调函数,接收 (config, change_type) 参数 - """ - self.config_path = config_path - self.callback = callback - - self._config: Dict[str, Any] = {} - self._last_modified: float = 0 - self._last_load_time: Optional[datetime] = None - self._lock = threading.RLock() - - # 配置变更历史 - self._change_history: list = [] - - # 监听器 - self._listeners: Dict[str, list] = { - 'on_change': [], - 'on_error': [] - } - - # 加载初始配置 - self.reload() - - def reload(self, silent: bool = False) -> bool: - """ - 重新加载配置 - - Args: - silent: 静默模式,不触发变更通知 - - Returns: - 是否加载成功 - """ - try: - if not os.path.exists(self.config_path): - if not silent: - logger.warning(f"配置文件不存在: {self.config_path}") - return False - - # 获取文件修改时间 - current_mtime = os.path.getmtime(self.config_path) - - # 检查是否有变化 - if current_mtime == self._last_modified and not silent: - return True - - # 加载配置 - with open(self.config_path, 'r', encoding='utf-8') as f: - new_config = json.load(f) - - # 计算变更 - changes = self._compute_changes(self._config, new_config) - - with self._lock: - old_config = self._config.copy() - self._config = new_config - self._last_modified = current_mtime - self._last_load_time = datetime.now() - - # 记录变更 - if changes: - change_record = { - 'timestamp': self._last_load_time.isoformat(), - 'changes': changes, - 'old_config_keys': list(old_config.keys()), - 'new_config_keys': list(new_config.keys()) - } - self._change_history.append(change_record) - - # 保持历史记录在合理范围内 - if len(self._change_history) > 100: - self._change_history = self._change_history[-50:] - - if not silent and changes: - self._notify_change(changes) - - logger.info(f"配置已重新加载: {self.config_path}") - return True - - except json.JSONDecodeError as e: - error_msg = f"配置 JSON 格式错误: {e}" - logger.error(error_msg) - self._notify_error(error_msg) - return False - except Exception as e: - error_msg = f"加载配置失败: {e}" - logger.error(error_msg) - self._notify_error(error_msg) - return False - - def _compute_changes(self, old: Dict, new: Dict) -> Dict: - """计算配置变更""" - changes = { - 'added': [], - 'removed': [], - 'modified': [] - } - - old_keys = set(old.keys()) - new_keys = set(new.keys()) - - # 新增的键 - for key in new_keys - old_keys: - changes['added'].append(key) - - # 移除的键 - for key in old_keys - new_keys: - changes['removed'].append(key) - - # 修改的键 - for key in old_keys & new_keys: - if old[key] != new[key]: - # 检查是否是嵌套字典 - if isinstance(old[key], dict) and isinstance(new[key], dict): - nested = self._compute_nested_changes(old[key], new[key], f"{key}.") - if nested['added'] or nested['removed'] or nested['modified']: - changes['modified'].append({ - 'key': key, - 'type': 'nested', - 'changes': nested - }) - else: - changes['modified'].append({ - 'key': key, - 'type': 'value', - 'old_value': old[key], - 'new_value': new[key] - }) - - return changes - - def _compute_nested_changes(self, old: Dict, new: Dict, prefix: str = "") -> Dict: - """计算嵌套字典的变更""" - changes = { - 'added': [], - 'removed': [], - 'modified': [] - } - - old_keys = set(old.keys()) - new_keys = set(new.keys()) - - for key in new_keys - old_keys: - changes['added'].append(f"{prefix}{key}") - - for key in old_keys - new_keys: - changes['removed'].append(f"{prefix}{key}") - - for key in old_keys & new_keys: - if old[key] != new[key]: - changes['modified'].append(f"{prefix}{key}") - - return changes - - def _notify_change(self, changes: Dict): - """通知配置变更""" - for listener in self._listeners['on_change']: - try: - if callable(listener): - listener(self._config, changes) - except Exception as e: - logger.error(f"配置变更监听器执行失败: {e}") - - if self.callback: - try: - self.callback(self._config, changes) - except Exception as e: - logger.error(f"配置回调函数执行失败: {e}") - - def _notify_error(self, error: str): - """通知错误""" - for listener in self._listeners['on_error']: - try: - if callable(listener): - listener(error) - except Exception as e: - logger.error(f"错误监听器执行失败: {e}") - - def add_change_listener(self, callback: Callable): - """添加配置变更监听器""" - self._listeners['on_change'].append(callback) - - def add_error_listener(self, callback: Callable): - """添加错误监听器""" - self._listeners['on_error'].append(callback) - - def get(self, key: str, default: Any = None) -> Any: - """获取配置值""" - with self._lock: - return self._config.get(key, default) - - def get_all(self) -> Dict: - """获取完整配置""" - with self._lock: - return self._config.copy() - - def set(self, key: str, value: Any, save: bool = True) -> bool: - """设置配置值(仅内存中)""" - with self._lock: - self._config[key] = value - - if save: - return self.save() - - return True - - def save(self, path: str = None) -> bool: - """保存配置到文件""" - save_path = path or self.config_path - - try: - with open(save_path, 'w', encoding='utf-8') as f: - json.dump(self._config, f, ensure_ascii=False, indent=4) - self._last_modified = os.path.getmtime(save_path) - return True - except Exception as e: - logger.error(f"保存配置失败: {e}") - return False - - def get_change_history(self, limit: int = 10) -> list: - """获取配置变更历史""" - return self._change_history[-limit:] - - def watch(self, interval: float = 5.0): - """ - 启动后台监控线程 - - Args: - interval: 检查间隔(秒) - """ - def _watch_loop(): - while True: - try: - self.reload() - except Exception as e: - logger.error(f"配置监控错误: {e}") - time.sleep(interval) - - thread = threading.Thread(target=_watch_loop, daemon=True) - thread.start() - logger.info(f"配置热监控已启动,间隔: {interval}秒") - - -class ConfigManager: - """配置管理器 - 支持热更新""" - - def __init__(self, config: Dict[str, Any] = None): - self.config = config or {} - self._hot_reloader: Optional[ConfigHotReloader] = None - - def load_from_file(self, path: str, enable_watch: bool = False) -> bool: - """从文件加载配置""" - if not os.path.exists(path): - return False - - try: - with open(path, 'r', encoding='utf-8') as f: - self.config = json.load(f) - - if enable_watch: - self._hot_reloader = ConfigHotReloader(path) - self._hot_reloader.watch() - - return True - except Exception as e: - logger.error(f"加载配置失败: {e}") - return False - - def hot_reload(self, path: str = None) -> bool: - """触发热重载""" - if self._hot_reloader: - return self._hot_reloader.reload() - return False - - def get(self, key: str, default: Any = None) -> Any: - """获取配置值""" - return self.config.get(key, default) - - def set(self, key: str, value: Any, persist: bool = False, path: str = None) -> bool: - """设置配置值""" - keys = key.split('.') - current = self.config - - for k in keys[:-1]: - if k not in current: - current[k] = {} - current = current[k] - - current[keys[-1]] = value - - if persist: - if self._hot_reloader: - return self._hot_reloader.save(path) - elif path: - try: - with open(path, 'w', encoding='utf-8') as f: - json.dump(self.config, f, ensure_ascii=False, indent=4) - return True - except Exception as e: - logger.error(f"保存配置失败: {e}") - return False - - return True - - 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.config.copy() +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +配置热更新模块 +支持不重启服务的情况下重新加载配置 +""" + +import os +import sys +import json +import time +import threading +import logging +from typing import Dict, Any, Optional, Callable +from datetime import datetime + +logger = logging.getLogger(__name__) + + +class ConfigHotReloader: + """配置热重载管理器""" + + def __init__(self, config_path: str, callback: Callable = None): + """ + 初始化热重载管理器 + + Args: + config_path: 配置文件路径 + callback: 配置变更时的回调函数,接收 (config, change_type) 参数 + """ + self.config_path = config_path + self.callback = callback + + self._config: Dict[str, Any] = {} + self._last_modified: float = 0 + self._last_load_time: Optional[datetime] = None + self._lock = threading.RLock() + + # 配置变更历史 + self._change_history: list = [] + + # 监听器 + self._listeners: Dict[str, list] = { + 'on_change': [], + 'on_error': [] + } + + # 加载初始配置 + self.reload() + + def reload(self, silent: bool = False) -> bool: + """ + 重新加载配置 + + Args: + silent: 静默模式,不触发变更通知 + + Returns: + 是否加载成功 + """ + try: + if not os.path.exists(self.config_path): + if not silent: + logger.warning(f"配置文件不存在: {self.config_path}") + return False + + # 获取文件修改时间 + current_mtime = os.path.getmtime(self.config_path) + + # 检查是否有变化 + if current_mtime == self._last_modified and not silent: + return True + + # 加载配置 + with open(self.config_path, 'r', encoding='utf-8') as f: + new_config = json.load(f) + + # 计算变更 + changes = self._compute_changes(self._config, new_config) + + with self._lock: + old_config = self._config.copy() + self._config = new_config + self._last_modified = current_mtime + self._last_load_time = datetime.now() + + # 记录变更 + if changes: + change_record = { + 'timestamp': self._last_load_time.isoformat(), + 'changes': changes, + 'old_config_keys': list(old_config.keys()), + 'new_config_keys': list(new_config.keys()) + } + self._change_history.append(change_record) + + # 保持历史记录在合理范围内 + if len(self._change_history) > 100: + self._change_history = self._change_history[-50:] + + if not silent and changes: + self._notify_change(changes) + + logger.info(f"配置已重新加载: {self.config_path}") + return True + + except json.JSONDecodeError as e: + error_msg = f"配置 JSON 格式错误: {e}" + logger.error(error_msg) + self._notify_error(error_msg) + return False + except Exception as e: + error_msg = f"加载配置失败: {e}" + logger.error(error_msg) + self._notify_error(error_msg) + return False + + def _compute_changes(self, old: Dict, new: Dict) -> Dict: + """计算配置变更""" + changes = { + 'added': [], + 'removed': [], + 'modified': [] + } + + old_keys = set(old.keys()) + new_keys = set(new.keys()) + + # 新增的键 + for key in new_keys - old_keys: + changes['added'].append(key) + + # 移除的键 + for key in old_keys - new_keys: + changes['removed'].append(key) + + # 修改的键 + for key in old_keys & new_keys: + if old[key] != new[key]: + # 检查是否是嵌套字典 + if isinstance(old[key], dict) and isinstance(new[key], dict): + nested = self._compute_nested_changes(old[key], new[key], f"{key}.") + if nested['added'] or nested['removed'] or nested['modified']: + changes['modified'].append({ + 'key': key, + 'type': 'nested', + 'changes': nested + }) + else: + changes['modified'].append({ + 'key': key, + 'type': 'value', + 'old_value': old[key], + 'new_value': new[key] + }) + + return changes + + def _compute_nested_changes(self, old: Dict, new: Dict, prefix: str = "") -> Dict: + """计算嵌套字典的变更""" + changes = { + 'added': [], + 'removed': [], + 'modified': [] + } + + old_keys = set(old.keys()) + new_keys = set(new.keys()) + + for key in new_keys - old_keys: + changes['added'].append(f"{prefix}{key}") + + for key in old_keys - new_keys: + changes['removed'].append(f"{prefix}{key}") + + for key in old_keys & new_keys: + if old[key] != new[key]: + changes['modified'].append(f"{prefix}{key}") + + return changes + + def _notify_change(self, changes: Dict): + """通知配置变更""" + for listener in self._listeners['on_change']: + try: + if callable(listener): + listener(self._config, changes) + except Exception as e: + logger.error(f"配置变更监听器执行失败: {e}") + + if self.callback: + try: + self.callback(self._config, changes) + except Exception as e: + logger.error(f"配置回调函数执行失败: {e}") + + def _notify_error(self, error: str): + """通知错误""" + for listener in self._listeners['on_error']: + try: + if callable(listener): + listener(error) + except Exception as e: + logger.error(f"错误监听器执行失败: {e}") + + def add_change_listener(self, callback: Callable): + """添加配置变更监听器""" + self._listeners['on_change'].append(callback) + + def add_error_listener(self, callback: Callable): + """添加错误监听器""" + self._listeners['on_error'].append(callback) + + def get(self, key: str, default: Any = None) -> Any: + """获取配置值""" + with self._lock: + return self._config.get(key, default) + + def get_all(self) -> Dict: + """获取完整配置""" + with self._lock: + return self._config.copy() + + def set(self, key: str, value: Any, save: bool = True) -> bool: + """设置配置值(仅内存中)""" + with self._lock: + self._config[key] = value + + if save: + return self.save() + + return True + + def save(self, path: str = None) -> bool: + """保存配置到文件""" + save_path = path or self.config_path + + try: + with open(save_path, 'w', encoding='utf-8') as f: + json.dump(self._config, f, ensure_ascii=False, indent=4) + self._last_modified = os.path.getmtime(save_path) + return True + except Exception as e: + logger.error(f"保存配置失败: {e}") + return False + + def get_change_history(self, limit: int = 10) -> list: + """获取配置变更历史""" + return self._change_history[-limit:] + + def watch(self, interval: float = 5.0): + """ + 启动后台监控线程 + + Args: + interval: 检查间隔(秒) + """ + def _watch_loop(): + while True: + try: + self.reload() + except Exception as e: + logger.error(f"配置监控错误: {e}") + time.sleep(interval) + + thread = threading.Thread(target=_watch_loop, daemon=True) + thread.start() + logger.info(f"配置热监控已启动,间隔: {interval}秒") + + +class ConfigManager: + """配置管理器 - 支持热更新(单配置源: 热重载开启时以 reloader 的配置为准)""" + + def __init__(self, config: Dict[str, Any] = None): + self.config = config or {} + self._hot_reloader: Optional[ConfigHotReloader] = None + + def _effective_config(self) -> Dict[str, Any]: + """获取当前生效配置(避免与 reloader._config 双配置源分裂)""" + if self._hot_reloader is not None: + return self._hot_reloader._config + return self.config + + def load_from_file(self, path: str, enable_watch: bool = False) -> bool: + """从文件加载配置""" + if not os.path.exists(path): + return False + + try: + with open(path, 'r', encoding='utf-8') as f: + self.config = json.load(f) + + if enable_watch: + self._hot_reloader = ConfigHotReloader(path) + self._hot_reloader.watch() + + return True + except Exception as e: + logger.error(f"加载配置失败: {e}") + return False + + def hot_reload(self, path: str = None) -> bool: + """触发热重载""" + if self._hot_reloader: + return self._hot_reloader.reload() + return False + + def get(self, key: str, default: Any = None) -> Any: + """获取配置值(读当前生效配置)""" + return self._effective_config().get(key, default) + + def set(self, key: str, value: Any, persist: bool = False, path: str = None) -> bool: + """设置配置值(写入当前生效配置,持久化时与 reloader 一致)""" + keys = key.split('.') + current = self._effective_config() + + for k in keys[:-1]: + if k not in current: + current[k] = {} + current = current[k] + + current[keys[-1]] = value + + if persist: + if self._hot_reloader: + # set 已写入 reloader._config, save 落盘的就是新值 + return self._hot_reloader.save(path) + elif path: + try: + with open(path, 'w', encoding='utf-8') as f: + json.dump(self._effective_config(), f, ensure_ascii=False, indent=4) + return True + except Exception as e: + logger.error(f"保存配置失败: {e}") + return False + + return True + + 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() diff --git a/core/database.py b/core/database.py index e077d85..232030b 100644 --- a/core/database.py +++ b/core/database.py @@ -369,11 +369,10 @@ class UserRecord(Base): enabled = Column(Boolean, default=True) def to_dict(self) -> dict: + """安全序列化: 不返回密码哈希与 token(防止泄露给前端/日志)""" return { 'id': self.id, 'username': self.username, - 'password_hash': self.password_hash, - 'token': self.token, 'token_expires_at': self.token_expires_at, 'role': self.role, 'email': self.email, @@ -387,6 +386,13 @@ class UserRecord(Base): '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): """登录日志记录表""" @@ -1212,13 +1218,21 @@ class DatabaseManager: return {'success': True, 'user_id': user.id} def get_user(self, username: str) -> dict: - """获取用户信息""" + """获取用户信息(安全版,不含 password_hash/token)""" with self.session() as session: user = session.query(UserRecord).filter_by(username=username).first() if user: return user.to_dict() 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: """通过ID获取用户信息""" with self.session() as session: diff --git a/core/mirror_sync.py b/core/mirror_sync.py index aed78d3..b18c5b4 100644 --- a/core/mirror_sync.py +++ b/core/mirror_sync.py @@ -97,10 +97,14 @@ class MirrorSyncManager: self.sync_status = {} def save_sync_state(self): - """保存同步状态""" + """保存同步状态(快照 + 原子写,可在持锁状态下调用,不会死锁)""" try: - with open(self.sync_state_file, 'w', encoding='utf-8') as f: - json.dump(self.sync_status, f, ensure_ascii=False, indent=2) + # dict() 浅拷贝是单条 C 操作,并发修改安全 + 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: print(f"保存同步状态失败: {e}") @@ -249,18 +253,18 @@ class MirrorSyncManager: if name not in self.sync_sources: return False - if name in self.sync_threads and self.sync_threads[name].is_alive(): - return True # 已经在运行 - - # 设置运行标志,确保同步循环可以执行 - self.running = True - 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]['error'] = None - thread = Thread(target=self._sync_worker, args=(name,), daemon=True) - self.sync_threads[name] = thread + thread = Thread(target=self._sync_worker, args=(name,), daemon=True) + self.sync_threads[name] = thread thread.start() return True @@ -280,12 +284,14 @@ class MirrorSyncManager: if not packages or not isinstance(packages, list): return {"success": False, "error": "请提供有效的包名列表"} - # 生成临时任务ID + # 生成临时任务ID(随机后缀防同秒碰撞) 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: @@ -333,17 +339,27 @@ class MirrorSyncManager: self.sync_status[temp_name]['status'] = 'error' self.sync_status[temp_name]['error'] = str(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): - """停止同步指定源 - 立即停止""" - # 立即设置状态为停止 + """停止同步指定源 - 请求停止 + + 注意: 不删除 sync_threads 条目(删除后 start_sync 会再起一个 worker, + 与仍在运行的旧线程形成双 worker)。线程结束后 is_alive() 为 False, + start_sync 会自然复用该条目。 + """ + # 设置状态为停止(源可能已从 sync_status 移除,容错) 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 - # 立即返回,不等待线程结束 - if name in self.sync_threads: - del self.sync_threads[name] # 短暂等待后重置运行标志 import time time.sleep(0.5) @@ -364,7 +380,7 @@ class MirrorSyncManager: print(f"[SYNC] 开始同步: {name}") source_config = self.sync_sources[name] 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: if sync_type in ('http', 'https'): @@ -394,8 +410,12 @@ class MirrorSyncManager: raise ValueError(f"不支持的同步类型: {sync_type}") with self.sync_lock: - self.sync_status[name]['status'] = 'completed' - self.sync_status[name]['last_sync'] = datetime.now().isoformat() + if not self.running: + # 被 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: with self.sync_lock: @@ -403,6 +423,9 @@ class MirrorSyncManager: self.sync_status[name]['error'] = str(e) finally: + # 清理线程条目(start_sync 通过 is_alive 判断, 清理后允许重新 start) + with self.sync_lock: + self.sync_threads.pop(name, None) self.save_sync_state() def _sync_http(self, name, config, specific_packages=None): @@ -420,7 +443,7 @@ class MirrorSyncManager: # 获取文件列表 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) total_files = len(file_list) @@ -445,10 +468,13 @@ class MirrorSyncManager: with self.sync_lock: self.sync_status[name]['files_synced'] = synced_count - # 获取本地路径 + # 获取本地路径(先做路径穿越校验) filename = file_info.get('name', '') if not filename: continue + if self._safe_remote_name(filename) is None: + print(f"[HTTP Sync] 跳过不安全的文件名: {filename!r}") + continue # 处理子目录 local_path = os.path.join(target_dir, filename) @@ -476,6 +502,29 @@ class MirrorSyncManager: 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): """FTP同步实现""" try: @@ -519,15 +568,19 @@ class MirrorSyncManager: for filename, file_info in remote_files.items(): if not self.running: 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']: sub_config = config.copy() - sub_config['remote_path'] = os.path.join(remote_path, filename).replace('\\', '/') - sub_config['target'] = os.path.join(config.get('target', name), filename) - self._sync_ftp(name + '/' + filename, sub_config) + sub_config['remote_path'] = os.path.join(remote_path, safe_name).replace('\\', '/') + sub_config['target'] = os.path.join(config.get('target', name), safe_name) + # 递归沿用顶层 name,避免子目录用假名导致 sync_status KeyError + self._sync_ftp(name, sub_config) else: - if self._need_sync_ftp(ftp, filename, local_path, file_info): - if self._download_ftp_file(ftp, filename, local_path, file_info): + if self._need_sync_ftp(ftp, safe_name, local_path, file_info): + if self._download_ftp_file(ftp, safe_name, local_path, file_info): synced_count += 1 with self.sync_lock: @@ -603,13 +656,12 @@ class MirrorSyncManager: for file_attr in file_list: if not self.running: 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('\\', '/') local_filepath = os.path.join(local_path, remote_filename) - if remote_filename in ['.', '..']: - continue - if file_attr.st_mode & 0o40000: sub_synced = self._sync_sftp_directory(sftp, remote_filepath, local_filepath, sync_name) synced_count += sub_synced @@ -679,7 +731,13 @@ class MirrorSyncManager: import subprocess 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) print(f"开始Rsync同步 {name} -> {target_dir}") @@ -725,10 +783,15 @@ class MirrorSyncManager: if os.path.exists(os.path.join(target_dir, '.git')): # 已存在,执行git pull try: - subprocess.run(['git', 'fetch', '--all'], cwd=target_dir, check=True, capture_output=True) - subprocess.run(['git', 'reset', '--hard', f'origin/{branch}'], cwd=target_dir, check=True, capture_output=True) + git_timeout = config.get('git_timeout', 1800) + 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: raise Exception(f"Git pull失败: {e}") + except subprocess.TimeoutExpired: + raise Exception(f"Git pull超时({git_timeout}s)") else: # 克隆新仓库 cmd = ['git', 'clone'] @@ -739,9 +802,12 @@ class MirrorSyncManager: cmd.extend([repo_url, target_dir]) 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: raise Exception(f"Git clone失败: {e}") + except subprocess.TimeoutExpired: + raise Exception(f"Git clone超时({git_timeout}s)") def _sync_s3(self, name, config): """AWS S3兼容存储同步 (S3/OSS/COS/MinIO等)""" @@ -888,7 +954,7 @@ class MirrorSyncManager: if is_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: try: @@ -898,7 +964,7 @@ class MirrorSyncManager: encoded_auth = base64.b64encode(auth_string.encode()).decode() 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()) if 'files' in api_data: for file_info in api_data['files']: @@ -920,7 +986,7 @@ class MirrorSyncManager: 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文件列表 Args: @@ -986,12 +1052,10 @@ class MirrorSyncManager: if not self.running: break - # 更新进度 + # 更新进度(只更新当前源,避免污染其他源状态) with self.sync_lock: - if hasattr(self, 'sync_status') and self.sync_status: - for name in self.sync_status: - if 'total_files' in self.sync_status[name]: - self.sync_status[name]['files_synced'] = i + if sync_name and sync_name in self.sync_status: + self.sync_status[sync_name]['files_synced'] = i # 获取每个包的文件列表 package_url = f"{package_index_url}{package_name}/" @@ -1098,8 +1162,11 @@ class MirrorSyncManager: """检查HTTP文件是否需要同步""" if not os.path.exists(local_path): return True - local_size = os.path.getsize(local_path) 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 def _need_sync_ftp(self, ftp, remote_filename, local_path, file_info): @@ -1213,7 +1280,7 @@ class MirrorSyncManager: return False except Exception as e: - print(f"下载HTTP文件失败 {url}: {e}") + print(f"下载HTTP文件失败 {self._redact_url(url)}: {e}") if os.path.exists(temp_path): os.remove(temp_path) return False @@ -1230,6 +1297,10 @@ class MirrorSyncManager: else: 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: if start_pos > 0: ftp.voidcmd('TYPE I') diff --git a/core/scheduler.py b/core/scheduler.py index 69afd24..ef11c62 100644 --- a/core/scheduler.py +++ b/core/scheduler.py @@ -1,380 +1,387 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - -""" -定时任务调度器 -支持 cron 表达式和简单间隔的定时任务 -""" - -import os -import time -import logging -import threading -from datetime import datetime, timedelta -from typing import Dict, List, Optional, Callable -from enum import Enum - -logger = logging.getLogger(__name__) - - -class TaskStatus(Enum): - """任务状态""" - IDLE = "idle" - RUNNING = "running" - ERROR = "error" - DISABLED = "disabled" - - -class ScheduledTask: - """定时任务""" - - def __init__(self, name: str, task_type: str, config: dict, - callback: Callable, logger=None): - """ - 初始化定时任务 - - Args: - name: 任务名称 - task_type: 任务类型 ('cron' 或 'interval') - config: 任务配置 - callback: 回调函数 - logger: 日志器 - """ - self.name = name - self.task_type = task_type # 'cron' 或 'interval' - self.config = config or {} - self.callback = callback - self.logger = logger or logging.getLogger(__name__) - - # 状态 - self.status = TaskStatus.IDLE - self.last_run: Optional[datetime] = None - self.next_run: Optional[datetime] = None - self.last_error: Optional[str] = None - self.run_count = 0 - - # 配置解析 - self._parse_config() - - def _parse_config(self): - """解析任务配置""" - if self.task_type == 'cron': - # Cron 表达式: "minute hour day month weekday" - # 例如: "0 3 * * *" 每天凌晨3点 - cron = self.config.get('cron', '0 0 * * *') - parts = cron.split() - if len(parts) == 5: - self.cron_parts = { - 'minute': self._parse_cron_part(parts[0], 0, 59), - 'hour': self._parse_cron_part(parts[1], 0, 23), - 'day': self._parse_cron_part(parts[2], 1, 31), - 'month': self._parse_cron_part(parts[3], 1, 12), - 'weekday': self._parse_cron_part(parts[4], 0, 6) - } - else: - self.logger.warning(f"无效的 cron 表达式: {cron}") - self.cron_parts = None - - elif self.task_type == 'interval': - # 间隔: seconds, minutes, hours - interval = self.config.get('interval', {}) - self.interval_seconds = ( - interval.get('seconds', 0) + - interval.get('minutes', 0) * 60 + - interval.get('hours', 0) * 3600 + - interval.get('days', 0) * 86400 - ) - if self.interval_seconds <= 0: - self.interval_seconds = 3600 # 默认1小时 - - # 是否启用 - self.enabled = self.config.get('enabled', True) - - def _parse_cron_part(self, part: str, min_val: int, max_val: int) -> List[int]: - """解析 cron 表达式的一部分""" - result = [] - if part == '*': - return list(range(min_val, max_val + 1)) - - # 处理列表: "1,2,3" - if ',' in part: - result = [] - for sub in part.split(','): - result.extend(self._parse_cron_part(sub.strip(), min_val, max_val)) - return result - - # 处理范围: "1-5" - if '-' in part: - start, end = part.split('-') - return list(range(int(start), int(end) + 1)) - - # 处理步进: "*/5" - if '/' in part: - base, step = part.split('/') - base_list = self._parse_cron_part(base or '*', min_val, max_val) - step = int(step) - return base_list[::step] - - # 单个值 - try: - val = int(part) - if min_val <= val <= max_val: - return [val] - except ValueError: - pass - - return [] - - def should_run_now(self) -> bool: - """检查是否应该在当前时刻运行""" - if not self.enabled: - return False - - now = datetime.now() - - if self.task_type == 'cron' and self.cron_parts: - return self._matches_cron(now) - elif self.task_type == 'interval': - if self.last_run is None: - return True - elapsed = (now - self.last_run).total_seconds() - return elapsed >= self.interval_seconds - - return False - - def _matches_cron(self, dt: datetime) -> bool: - """检查时间是否匹配 cron 表达式""" - if not self.cron_parts: - return False - - # cron 约定: 0=周日...6=周六; datetime.weekday(): 0=周一...6=周日 - cron_weekday = (dt.weekday() + 1) % 7 - - return ( - dt.minute in self.cron_parts['minute'] and - dt.hour in self.cron_parts['hour'] and - dt.day in self.cron_parts['day'] and - dt.month in self.cron_parts['month'] and - cron_weekday in self.cron_parts['weekday'] - ) - - def get_next_run_time(self) -> Optional[datetime]: - """计算下次运行时间""" - if not self.enabled: - return None - - now = datetime.now() - - if self.task_type == 'cron' and self.cron_parts: - # 找到下一个匹配的时间点 - for i in range(365 * 24 * 60): # 最多查找1年 - candidate = now + timedelta(minutes=i) - if self._matches_cron(candidate): - return candidate - elif self.task_type == 'interval': - if self.last_run: - return self.last_run + timedelta(seconds=self.interval_seconds) - return now - - return None - - def run(self) -> bool: - """执行任务""" - if self.status == TaskStatus.RUNNING: - self.logger.warning(f"任务 {self.name} 已在运行中") - return False - - self.status = TaskStatus.RUNNING - self.last_run = datetime.now() - self.last_error = None - - try: - self.logger.info(f"开始执行定时任务: {self.name}") - result = self.callback(self.name, self.config) - self.run_count += 1 - self.logger.info(f"定时任务 {self.name} 执行完成") - return True - except Exception as e: - self.last_error = str(e) - self.status = TaskStatus.ERROR - self.logger.error(f"定时任务 {self.name} 执行失败: {e}") - return False - finally: - if self.status != TaskStatus.ERROR: - self.status = TaskStatus.IDLE - - def to_dict(self) -> dict: - """转换为字典""" - return { - 'name': self.name, - 'type': self.task_type, - 'enabled': self.enabled, - 'status': self.status.value, - 'config': self.config, - '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, - 'last_error': self.last_error - } - - -class Scheduler: - """定时任务调度器""" - - def __init__(self, config: dict = None): - self.config = config or {} - self.tasks: Dict[str, ScheduledTask] = {} - self._running = False - self._thread: Optional[threading.Thread] = None - self._lock = threading.Lock() - - # 默认检查间隔 - self.check_interval = self.config.get('check_interval', 10) - - # 事件回调 - self.on_task_start: Optional[Callable] = None - self.on_task_complete: Optional[Callable] = None - self.on_task_error: Optional[Callable] = None - - def add_task(self, name: str, task_type: str, config: dict, - callback: Callable) -> bool: - """ - 添加定时任务 - - Args: - name: 任务名称 - task_type: 任务类型 ('cron' 或 'interval') - config: 任务配置 - callback: 回调函数 - - Returns: - 是否成功 - """ - with self._lock: - if name in self.tasks: - logger.warning(f"任务 {name} 已存在,将被替换") - self.tasks[name] = ScheduledTask(name, task_type, config, callback, logger) - return True - - def remove_task(self, name: str) -> bool: - """移除任务""" - with self._lock: - if name in self.tasks: - del self.tasks[name] - return True - return False - - def get_task(self, name: str) -> Optional[ScheduledTask]: - """获取任务""" - return self.tasks.get(name) - - def get_all_tasks(self) -> List[dict]: - """获取所有任务状态""" - with self._lock: - for task in self.tasks.values(): - task.next_run = task.get_next_run_time() - return [task.to_dict() for task in self.tasks.values()] - - def start(self): - """启动调度器""" - if self._running: - logger.warning("调度器已在运行中") - return - - self._running = True - self._thread = threading.Thread(target=self._run_loop, daemon=True) - self._thread.start() - logger.info("定时任务调度器已启动") - - def stop(self): - """停止调度器""" - self._running = False - if self._thread: - self._thread.join(timeout=5) - logger.info("定时任务调度器已停止") - - def _run_loop(self): - """运行循环""" - while self._running: - try: - now = datetime.now() - - with self._lock: - for name, task in self.tasks.items(): - if task.should_run_now(): - # 使用线程池执行任务 - from concurrent.futures import ThreadPoolExecutor - with ThreadPoolExecutor(max_workers=1) as executor: - executor.submit(task.run) - - time.sleep(self.check_interval) - - except Exception as e: - logger.error(f"调度器循环错误: {e}") - time.sleep(5) - - def run_task_now(self, name: str) -> bool: - """立即运行指定任务""" - task = self.get_task(name) - if task: - return task.run() - return False - - def enable_task(self, name: str, enabled: bool = True) -> bool: - """启用/禁用任务""" - task = self.get_task(name) - if task: - task.enabled = enabled - 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) - task._parse_config() - return True - return False - - -# ==================== 同步任务工厂 ==================== - -# def create_sync_task_callback(sync_manager): -# """创建同步任务的回调函数""" -# def sync_task_callback(task_name: str, config: dict): -# """同步任务回调""" -# sync_manager.start_sync(task_name) -# return True -# return sync_task_callback - - -# ==================== 默认任务配置 ==================== -# DEFAULT_SCHEDULED_TASKS = { ... } - -DEFAULT_SCHEDULED_TASKS = { - # 数据库清理 - 每天凌晨2点 - 'cleanup_db': { - 'type': 'cron', - 'config': { - 'cron': '0 2 * * *', - 'enabled': True - } - }, - # 缓存清理 - 每6小时 - 'cleanup_cache': { - 'type': 'interval', - 'config': { - 'interval': {'hours': 6}, - 'enabled': True - } - }, - # 健康检查 - 每5分钟 - 'health_check': { - 'type': 'interval', - 'config': { - 'interval': {'minutes': 5}, - 'enabled': True - } - } -} +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +定时任务调度器 +支持 cron 表达式和简单间隔的定时任务 +""" + +import os +import time +import logging +import threading +from datetime import datetime, timedelta +from typing import Dict, List, Optional, Callable +from enum import Enum + +logger = logging.getLogger(__name__) + + +class TaskStatus(Enum): + """任务状态""" + IDLE = "idle" + RUNNING = "running" + ERROR = "error" + DISABLED = "disabled" + + +class ScheduledTask: + """定时任务""" + + def __init__(self, name: str, task_type: str, config: dict, + callback: Callable, logger=None): + """ + 初始化定时任务 + + Args: + name: 任务名称 + task_type: 任务类型 ('cron' 或 'interval') + config: 任务配置 + callback: 回调函数 + logger: 日志器 + """ + self.name = name + self.task_type = task_type # 'cron' 或 'interval' + self.config = config or {} + self.callback = callback + self.logger = logger or logging.getLogger(__name__) + + # 状态 + self.status = TaskStatus.IDLE + self.last_run: Optional[datetime] = None + self.next_run: Optional[datetime] = None + self.last_error: Optional[str] = None + self.run_count = 0 + + # 配置解析 + self._parse_config() + + def _parse_config(self): + """解析任务配置""" + if self.task_type == 'cron': + # Cron 表达式: "minute hour day month weekday" + # 例如: "0 3 * * *" 每天凌晨3点 + cron = self.config.get('cron', '0 0 * * *') + parts = cron.split() + if len(parts) == 5: + self.cron_parts = { + 'minute': self._parse_cron_part(parts[0], 0, 59), + 'hour': self._parse_cron_part(parts[1], 0, 23), + 'day': self._parse_cron_part(parts[2], 1, 31), + 'month': self._parse_cron_part(parts[3], 1, 12), + 'weekday': self._parse_cron_part(parts[4], 0, 6) + } + else: + self.logger.warning(f"无效的 cron 表达式: {cron}") + self.cron_parts = None + + elif self.task_type == 'interval': + # 间隔: seconds, minutes, hours + interval = self.config.get('interval', {}) + self.interval_seconds = ( + interval.get('seconds', 0) + + interval.get('minutes', 0) * 60 + + interval.get('hours', 0) * 3600 + + interval.get('days', 0) * 86400 + ) + if self.interval_seconds <= 0: + self.interval_seconds = 3600 # 默认1小时 + + # 是否启用 + self.enabled = self.config.get('enabled', True) + + def _parse_cron_part(self, part: str, min_val: int, max_val: int) -> List[int]: + """解析 cron 表达式的一部分""" + result = [] + if part == '*': + return list(range(min_val, max_val + 1)) + + # 处理列表: "1,2,3" + if ',' in part: + result = [] + for sub in part.split(','): + result.extend(self._parse_cron_part(sub.strip(), min_val, max_val)) + return result + + # 处理范围: "1-5" + if '-' in part: + start, end = part.split('-') + return list(range(int(start), int(end) + 1)) + + # 处理步进: "*/5" + if '/' in part: + base, step = part.split('/') + base_list = self._parse_cron_part(base or '*', min_val, max_val) + step = int(step) + return base_list[::step] + + # 单个值 + try: + val = int(part) + if min_val <= val <= max_val: + return [val] + except ValueError: + pass + + return [] + + def should_run_now(self) -> bool: + """检查是否应该在当前时刻运行""" + if not self.enabled: + return False + + now = datetime.now() + + if self.task_type == 'cron' and self.cron_parts: + if not self._matches_cron(now): + return False + # 同一分钟内只触发一次(轮询间隔 < 60s 时避免重复触发) + if self.last_run is not None: + last_minute = self.last_run.strftime('%Y-%m-%d %H:%M') + if last_minute == now.strftime('%Y-%m-%d %H:%M'): + return False + return True + elif self.task_type == 'interval': + if self.last_run is None: + return True + elapsed = (now - self.last_run).total_seconds() + return elapsed >= self.interval_seconds + + return False + + def _matches_cron(self, dt: datetime) -> bool: + """检查时间是否匹配 cron 表达式""" + if not self.cron_parts: + return False + + # cron 约定: 0=周日...6=周六; datetime.weekday(): 0=周一...6=周日 + cron_weekday = (dt.weekday() + 1) % 7 + + return ( + dt.minute in self.cron_parts['minute'] and + dt.hour in self.cron_parts['hour'] and + dt.day in self.cron_parts['day'] and + dt.month in self.cron_parts['month'] and + cron_weekday in self.cron_parts['weekday'] + ) + + def get_next_run_time(self) -> Optional[datetime]: + """计算下次运行时间""" + if not self.enabled: + return None + + now = datetime.now() + + if self.task_type == 'cron' and self.cron_parts: + # 找到下一个匹配的时间点 + for i in range(365 * 24 * 60): # 最多查找1年 + candidate = now + timedelta(minutes=i) + if self._matches_cron(candidate): + return candidate + elif self.task_type == 'interval': + if self.last_run: + return self.last_run + timedelta(seconds=self.interval_seconds) + return now + + return None + + def run(self) -> bool: + """执行任务""" + if self.status == TaskStatus.RUNNING: + self.logger.warning(f"任务 {self.name} 已在运行中") + return False + + self.status = TaskStatus.RUNNING + self.last_run = datetime.now() + self.last_error = None + + try: + self.logger.info(f"开始执行定时任务: {self.name}") + result = self.callback(self.name, self.config) + self.run_count += 1 + self.logger.info(f"定时任务 {self.name} 执行完成") + return True + except Exception as e: + self.last_error = str(e) + self.status = TaskStatus.ERROR + self.logger.error(f"定时任务 {self.name} 执行失败: {e}") + return False + finally: + if self.status != TaskStatus.ERROR: + self.status = TaskStatus.IDLE + + def to_dict(self) -> dict: + """转换为字典""" + return { + 'name': self.name, + 'type': self.task_type, + 'enabled': self.enabled, + 'status': self.status.value, + 'config': self.config, + '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, + 'last_error': self.last_error + } + + +class Scheduler: + """定时任务调度器""" + + def __init__(self, config: dict = None): + self.config = config or {} + self.tasks: Dict[str, ScheduledTask] = {} + self._running = False + self._thread: Optional[threading.Thread] = None + self._lock = threading.Lock() + + # 默认检查间隔 + self.check_interval = self.config.get('check_interval', 10) + + # 事件回调 + self.on_task_start: Optional[Callable] = None + self.on_task_complete: Optional[Callable] = None + self.on_task_error: Optional[Callable] = None + + def add_task(self, name: str, task_type: str, config: dict, + callback: Callable) -> bool: + """ + 添加定时任务 + + Args: + name: 任务名称 + task_type: 任务类型 ('cron' 或 'interval') + config: 任务配置 + callback: 回调函数 + + Returns: + 是否成功 + """ + with self._lock: + if name in self.tasks: + logger.warning(f"任务 {name} 已存在,将被替换") + self.tasks[name] = ScheduledTask(name, task_type, config, callback, logger) + return True + + def remove_task(self, name: str) -> bool: + """移除任务""" + with self._lock: + if name in self.tasks: + del self.tasks[name] + return True + return False + + def get_task(self, name: str) -> Optional[ScheduledTask]: + """获取任务""" + return self.tasks.get(name) + + def get_all_tasks(self) -> List[dict]: + """获取所有任务状态""" + with self._lock: + for task in self.tasks.values(): + task.next_run = task.get_next_run_time() + return [task.to_dict() for task in self.tasks.values()] + + def start(self): + """启动调度器""" + if self._running: + logger.warning("调度器已在运行中") + return + + self._running = True + self._thread = threading.Thread(target=self._run_loop, daemon=True) + self._thread.start() + logger.info("定时任务调度器已启动") + + def stop(self): + """停止调度器""" + self._running = False + if self._thread: + self._thread.join(timeout=5) + logger.info("定时任务调度器已停止") + + def _run_loop(self): + """运行循环""" + while self._running: + try: + now = datetime.now() + + with self._lock: + for name, task in self.tasks.items(): + if task.should_run_now(): + # 使用线程池执行任务 + from concurrent.futures import ThreadPoolExecutor + with ThreadPoolExecutor(max_workers=1) as executor: + executor.submit(task.run) + + time.sleep(self.check_interval) + + except Exception as e: + logger.error(f"调度器循环错误: {e}") + time.sleep(5) + + def run_task_now(self, name: str) -> bool: + """立即运行指定任务""" + task = self.get_task(name) + if task: + return task.run() + return False + + def enable_task(self, name: str, enabled: bool = True) -> bool: + """启用/禁用任务""" + task = self.get_task(name) + if task: + task.enabled = enabled + 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) + task._parse_config() + return True + return False + + +# ==================== 同步任务工厂 ==================== + +# def create_sync_task_callback(sync_manager): +# """创建同步任务的回调函数""" +# def sync_task_callback(task_name: str, config: dict): +# """同步任务回调""" +# sync_manager.start_sync(task_name) +# return True +# return sync_task_callback + + +# ==================== 默认任务配置 ==================== +# DEFAULT_SCHEDULED_TASKS = { ... } + +DEFAULT_SCHEDULED_TASKS = { + # 数据库清理 - 每天凌晨2点 + 'cleanup_db': { + 'type': 'cron', + 'config': { + 'cron': '0 2 * * *', + 'enabled': True + } + }, + # 缓存清理 - 每6小时 + 'cleanup_cache': { + 'type': 'interval', + 'config': { + 'interval': {'hours': 6}, + 'enabled': True + } + }, + # 健康检查 - 每5分钟 + 'health_check': { + 'type': 'interval', + 'config': { + 'interval': {'minutes': 5}, + 'enabled': True + } + } +} diff --git a/core/server.py b/core/server.py index d65043d..a069846 100644 --- a/core/server.py +++ b/core/server.py @@ -95,7 +95,7 @@ class MirrorServer: else: self.config_manager = config - self.config = self.config_manager.config + self.config = self.config_manager.get_all() self.server = None self.sync_manager = None self.is_running = False diff --git a/core/sync_engine.py b/core/sync_engine.py index 3c67ff7..1926161 100644 --- a/core/sync_engine.py +++ b/core/sync_engine.py @@ -935,8 +935,10 @@ class SyncEngine: if task.status in [SyncStatus.COMPLETED, SyncStatus.FAILED, SyncStatus.CANCELLED] ] - for task_id in completed_ids[-keep_count:]: - del self.active_tasks[task_id] + # 删除最旧的已完成任务,保留最近 keep_count 个(原逻辑删最新留最旧,逻辑倒置) + 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: diff --git a/handlers/http_handler.py b/handlers/http_handler.py index 8a3b1f0..92ee2b8 100644 --- a/handlers/http_handler.py +++ b/handlers/http_handler.py @@ -12,6 +12,7 @@ import mimetypes import base64 import hashlib import shutil +import threading from datetime import datetime from http.server import BaseHTTPRequestHandler from urllib.parse import unquote, urlparse, parse_qs @@ -40,6 +41,7 @@ class MirrorServerHandler(BaseHTTPRequestHandler): debug_log_file = None # 调试日志文件路径 _debug_categories = set() # 启用的调试类别 _mirror_handlers = {} # 镜像处理器实例缓存 + _stats_lock = threading.Lock() # 保护 JSON 统计/历史文件读改写 @classmethod def _setup_debug(cls, config): @@ -1421,10 +1423,11 @@ class MirrorServerHandler(BaseHTTPRequestHandler): except Exception as e: print(f"Error updating download count in database: {e}") - # 回退到 JSON 文件 - stats = self.load_stats() - stats[filepath] = stats.get(filepath, 0) + 1 - self.save_stats(stats) + # 回退到 JSON 文件(加锁防并发读改写丢计数) + with self._stats_lock: + stats = self.load_stats() + stats[filepath] = stats.get(filepath, 0) + 1 + self.save_stats(stats) # ==================== 下载历史记录 ==================== @@ -1501,8 +1504,9 @@ class MirrorServerHandler(BaseHTTPRequestHandler): except Exception as e: print(f"Error logging download to database: {e}") - # 回退到 JSON 文件 - history = self.load_download_history(1000) + # 回退到 JSON 文件(加锁防并发写坏/丢记录) + with self._stats_lock: + history = self.load_download_history(1000) entry = { 'timestamp': datetime.now().isoformat(), @@ -1513,5 +1517,6 @@ class MirrorServerHandler(BaseHTTPRequestHandler): 'method': self.command if hasattr(self, 'command') else 'GET' } - history.append(entry) - self.save_download_history(history) + with self._stats_lock: + history.append(entry) + self.save_download_history(history)