复查修复(二): 同步模块与安全遗留
- 修复 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:
@@ -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:
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
+347
-340
@@ -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()
|
||||
|
||||
+17
-3
@@ -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:
|
||||
|
||||
+120
-49
@@ -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')
|
||||
|
||||
+387
-380
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
@@ -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
|
||||
|
||||
+4
-2
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user