- 修复 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 删旧留新
241 lines
8.5 KiB
Python
241 lines
8.5 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
|
||
"""服务器核心模块 - 线程池版本"""
|
||
|
||
import os
|
||
import sys
|
||
import ssl
|
||
import signal
|
||
import socket
|
||
import mimetypes
|
||
import threading
|
||
import socketserver
|
||
from datetime import datetime
|
||
from concurrent.futures import ThreadPoolExecutor
|
||
|
||
from .config import ConfigManager
|
||
from .mirror_sync import MirrorSyncManager
|
||
|
||
|
||
class ThreadPoolHTTPServer(socketserver.ThreadingMixIn, socketserver.TCPServer):
|
||
"""使用线程池的 HTTP 服务器(有界队列 + 连接超时,防慢速 DoS)"""
|
||
allow_reuse_address = True
|
||
daemon_threads = True # 使用守护线程
|
||
|
||
def __init__(self, server_address, RequestHandlerClass, max_workers=50,
|
||
queue_size=50, socket_timeout=30):
|
||
self.max_workers = max_workers
|
||
self.queue_size = queue_size
|
||
self.socket_timeout = socket_timeout
|
||
self._executor = ThreadPoolExecutor(
|
||
max_workers=max_workers,
|
||
thread_name_prefix="http_handler"
|
||
)
|
||
# 有界并发槽位(活跃 + 排队),超过上限直接拒绝新连接,
|
||
# 防止连接洪水导致任务无限堆积在内存
|
||
self._slots = threading.BoundedSemaphore(max_workers + queue_size)
|
||
super().__init__(server_address, RequestHandlerClass)
|
||
|
||
def process_request(self, request, client_address):
|
||
"""使用线程池处理请求(有界队列)"""
|
||
if not self._slots.acquire(blocking=False):
|
||
# 队列已满:拒绝并立即关闭连接
|
||
try:
|
||
request.shutdown(socket.SHUT_RDWR)
|
||
except Exception:
|
||
pass
|
||
try:
|
||
request.close()
|
||
except Exception:
|
||
pass
|
||
return
|
||
try:
|
||
self._executor.submit(self._handle_request, request, client_address)
|
||
except Exception:
|
||
self._slots.release()
|
||
try:
|
||
request.close()
|
||
except Exception:
|
||
pass
|
||
|
||
def _handle_request(self, request, client_address):
|
||
"""实际处理请求"""
|
||
try:
|
||
# 连接级读写超时,防慢客户端无限占用工作线程
|
||
try:
|
||
request.settimeout(self.socket_timeout)
|
||
except Exception:
|
||
pass
|
||
self.finish_request(request, client_address)
|
||
except Exception:
|
||
self.handle_error(request, client_address)
|
||
finally:
|
||
self._slots.release()
|
||
self.shutdown_request(request)
|
||
|
||
def server_close(self):
|
||
"""关闭服务器和线程池"""
|
||
# Python 3.8 兼容处理
|
||
import sys
|
||
try:
|
||
self._executor.shutdown(wait=False, cancel_futures=True)
|
||
except TypeError:
|
||
# Python 3.8 不支持 cancel_futures 参数
|
||
self._executor.shutdown(wait=False)
|
||
super().server_close()
|
||
|
||
|
||
class MirrorServer:
|
||
"""镜像服务器主类"""
|
||
|
||
def __init__(self, config):
|
||
if isinstance(config, dict):
|
||
self.config_manager = ConfigManager(config)
|
||
else:
|
||
self.config_manager = config
|
||
|
||
self.config = self.config_manager.get_all()
|
||
self.server = None
|
||
self.sync_manager = None
|
||
self.is_running = False
|
||
|
||
def _validate_config(self, config):
|
||
"""验证和修复配置"""
|
||
return self.config_manager._validate_config(config)
|
||
|
||
def start(self):
|
||
"""启动服务器"""
|
||
try:
|
||
# 创建下载目录
|
||
base_dir = self.config['base_dir']
|
||
if not os.path.exists(base_dir):
|
||
os.makedirs(base_dir)
|
||
print(f"创建下载目录: {os.path.abspath(base_dir)}")
|
||
|
||
# 初始化MIME类型
|
||
mimetypes.init()
|
||
|
||
# 记录启动时间
|
||
self.config['start_time'] = __import__('time').time()
|
||
|
||
# 创建同步管理器(仅当启用时)
|
||
if self.config.get('enable_sync', True):
|
||
self.sync_manager = MirrorSyncManager(self.config)
|
||
self.sync_manager.start()
|
||
# 注入到配置,供 SyncScheduler 定时回调使用
|
||
self.config['_sync_manager'] = self.sync_manager
|
||
|
||
# 创建系统监控器(仅当启用时)
|
||
self.monitor = None
|
||
if self.config.get('enable_monitor', True):
|
||
try:
|
||
from .monitor import SystemMonitor
|
||
self.monitor = SystemMonitor(self.config)
|
||
print(f" ✓ 系统监控已启用 (间隔: {self.config.get('monitor_interval', 5)}秒)")
|
||
except ImportError as e:
|
||
print(f" ✗ 系统监控导入失败: {e}")
|
||
except Exception as e:
|
||
print(f" ✗ 系统监控初始化失败: {e}")
|
||
|
||
# 延迟导入 handler(避免循环导入)
|
||
from handlers.http_handler import MirrorServerHandler
|
||
|
||
# 获取线程数配置
|
||
max_workers = min(self.config.get('max_workers', 10), 10) # 限制最大线程数
|
||
|
||
# 创建服务器
|
||
server_address = (self.config['host'], self.config['port'])
|
||
self.server = ThreadPoolHTTPServer(
|
||
server_address,
|
||
MirrorServerHandler,
|
||
max_workers=max_workers,
|
||
queue_size=max_workers * 5,
|
||
socket_timeout=self.config.get('timeout', 30)
|
||
)
|
||
|
||
# 传递配置到处理器
|
||
MirrorServerHandler.config = self.config
|
||
MirrorServerHandler.sync_manager = self.sync_manager
|
||
MirrorServerHandler.monitor = self.monitor
|
||
# 设置调试模式
|
||
MirrorServerHandler._setup_debug(self.config)
|
||
|
||
# 设置 accept 轮询超时
|
||
self.server.timeout = self.config.get('timeout', 30)
|
||
|
||
# 启用HTTPS
|
||
if self.config.get('ssl_cert') and self.config.get('ssl_key'):
|
||
if not self._setup_ssl():
|
||
return False
|
||
|
||
self.is_running = True
|
||
return True
|
||
|
||
except Exception as e:
|
||
import traceback
|
||
print(f"服务器启动失败: {e}")
|
||
traceback.print_exc()
|
||
return False
|
||
|
||
def _setup_ssl(self):
|
||
"""设置SSL"""
|
||
try:
|
||
context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||
context.load_cert_chain(self.config['ssl_cert'], self.config['ssl_key'])
|
||
self.server.socket = context.wrap_socket(self.server.socket, server_side=True)
|
||
print(f"启用HTTPS,证书: {self.config['ssl_cert']}")
|
||
return True
|
||
except Exception as e:
|
||
print(f"启用HTTPS失败: {e}")
|
||
return False
|
||
|
||
def stop(self):
|
||
"""停止服务器"""
|
||
print("正在停止服务器...")
|
||
|
||
if self.sync_manager:
|
||
try:
|
||
self.sync_manager.stop()
|
||
except Exception as e:
|
||
print(f"停止同步管理器时出错: {e}")
|
||
|
||
if self.server:
|
||
try:
|
||
self.server.server_close()
|
||
except Exception as e:
|
||
print(f"关闭服务器连接时出错: {e}")
|
||
|
||
self.is_running = False
|
||
print("服务器已停止")
|
||
|
||
def serve_forever(self):
|
||
"""运行服务器"""
|
||
if not self.server:
|
||
print("服务器未启动")
|
||
return
|
||
|
||
# 打印服务器信息(已在 main.py 中显示,此处仅保留最简信息)
|
||
protocol = "https" if self.config.get('ssl_cert') else "http"
|
||
sync_count = len(self.sync_manager.sync_sources) if self.sync_manager else 0
|
||
print(f"\n▶ 服务器运行于: {protocol}://{self.config['host']}:{self.config['port']}")
|
||
print(f"▶ 同步源数: {sync_count} | 最大线程: {self.server.max_workers}")
|
||
print("▶ 按 Ctrl+C 停止服务器")
|
||
|
||
try:
|
||
self.server.serve_forever()
|
||
except KeyboardInterrupt:
|
||
print("\n正在关闭服务器...")
|
||
finally:
|
||
self.stop()
|
||
|
||
|
||
# 全局变量,用于信号处理器访问服务器实例(预留)
|
||
# _server_instance = None
|
||
|
||
# def signal_handler(signum, frame):
|
||
# """处理退出信号"""
|
||
# print(f"\n收到信号 {signum},正在关闭服务器...")
|
||
# import os
|
||
# os._exit(0)
|