Files
mirror_server/core/server.py
T
HYC Fixer 82875b710a 复查修复(二): 同步模块与安全遗留
- 修复 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 删旧留新
2026-09-02 00:39:12 +08:00

241 lines
8.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""服务器核心模块 - 线程池版本"""
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)