Baseline: pr1 HYC下载站 v2.3 before security/functional fixes
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
# API模块初始化
|
||||
from .router import APIRouter
|
||||
from .v1 import APIv1
|
||||
from .v2 import APIv2
|
||||
|
||||
__all__ = ['APIRouter', 'APIv1', 'APIv2']
|
||||
+128
@@ -0,0 +1,128 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
管理员API处理器
|
||||
提供认证相关API(无 keys 管理)
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
from core.api_auth import APIAuthManager, require_auth
|
||||
|
||||
|
||||
class AdminAPI:
|
||||
"""管理员API处理器"""
|
||||
|
||||
def __init__(self, config: dict):
|
||||
self.config = config
|
||||
self.auth_manager = APIAuthManager(config)
|
||||
|
||||
def handle_request(self, handler, method, path, query_params):
|
||||
"""处理管理员API请求"""
|
||||
|
||||
# 解析路径
|
||||
parts = path.strip('/').split('/')
|
||||
|
||||
# 根路径 - 列出API
|
||||
if len(parts) == 0 or parts[0] == '':
|
||||
self._api_overview(handler)
|
||||
return
|
||||
|
||||
action = parts[0]
|
||||
|
||||
if action == 'auth':
|
||||
self._handle_auth(handler, method, query_params)
|
||||
elif action == 'sessions':
|
||||
self._handle_sessions(handler, method, parts[1:] if len(parts) > 1 else [])
|
||||
elif action == 'stats':
|
||||
self._handle_stats(handler)
|
||||
else:
|
||||
handler.send_json_response({
|
||||
"error": f"Unknown admin action: {action}",
|
||||
"available_actions": ["auth", "sessions", "stats"]
|
||||
}, 404)
|
||||
|
||||
def _api_overview(self, handler):
|
||||
"""API概览"""
|
||||
handler.send_json_response({
|
||||
"name": "HYC Admin API",
|
||||
"version": "1.0",
|
||||
"description": "管理员认证API",
|
||||
"endpoints": {
|
||||
"GET /api/v2/admin/sessions": "列出活跃会话",
|
||||
"DELETE /api/v2/admin/sessions/{session_id}": "销毁会话",
|
||||
"POST /api/v2/admin/auth/verify": "验证认证状态",
|
||||
"GET /api/v2/admin/stats": "获取认证统计"
|
||||
},
|
||||
"authentication": {
|
||||
"methods": [
|
||||
"Authorization: Bearer <token>",
|
||||
"X-API-Key: <token>",
|
||||
"Cookie: hyc_auth=<session>",
|
||||
"?key=<token>"
|
||||
]
|
||||
}
|
||||
})
|
||||
|
||||
def _handle_auth(self, handler, method, query_params):
|
||||
"""处理认证相关"""
|
||||
# POST /api/v2/admin/auth/verify - 验证当前认证状态
|
||||
if method == 'POST':
|
||||
auth_result = handler.auth_result if hasattr(handler, 'auth_result') else {}
|
||||
|
||||
if auth_result.get('authenticated'):
|
||||
handler.send_json_response({
|
||||
"authenticated": True,
|
||||
"level": auth_result.get('level'),
|
||||
"key_id": auth_result.get('key_id'),
|
||||
"name": auth_result.get('name'),
|
||||
"permissions": auth_result.get('permissions', [])
|
||||
})
|
||||
else:
|
||||
handler.send_json_response({
|
||||
"authenticated": False
|
||||
}, 401)
|
||||
|
||||
else:
|
||||
handler.send_json_response({"error": "Invalid method"}, 405)
|
||||
|
||||
def _handle_sessions(self, handler, method, parts):
|
||||
"""处理会话管理"""
|
||||
if method == 'GET':
|
||||
# 列出活跃会话
|
||||
sessions = []
|
||||
for session_id, session in self.auth_manager.sessions.items():
|
||||
if time.time() < session.expires_at:
|
||||
sessions.append({
|
||||
"session_id": session.session_id,
|
||||
"user_id": session.user_id,
|
||||
"level": session.level,
|
||||
"created_at": session.created_at,
|
||||
"expires_at": session.expires_at,
|
||||
"last_activity": session.last_activity
|
||||
})
|
||||
|
||||
handler.send_json_response({
|
||||
"sessions": sessions,
|
||||
"count": len(sessions)
|
||||
})
|
||||
|
||||
elif method == 'DELETE' and len(parts) >= 1 and parts[0]:
|
||||
session_id = parts[0]
|
||||
success = self.auth_manager.destroy_session(session_id)
|
||||
if success:
|
||||
handler.send_json_response({
|
||||
"success": True,
|
||||
"message": f"Session {session_id} destroyed"
|
||||
})
|
||||
else:
|
||||
handler.send_json_response({"error": "Session not found"}, 404)
|
||||
|
||||
else:
|
||||
handler.send_json_response({"error": "Invalid request"}, 400)
|
||||
|
||||
def _handle_stats(self, handler):
|
||||
"""获取认证统计"""
|
||||
stats = self.auth_manager.get_stats()
|
||||
handler.send_json_response(stats)
|
||||
@@ -0,0 +1,282 @@
|
||||
# 低端设备优化指南
|
||||
|
||||
## 支持的架构
|
||||
|
||||
| 架构 | 状态 | 说明 |
|
||||
|------|------|------|
|
||||
| X86_64/AMD64 | ✓ 完全支持 | 主流 PC/服务器 |
|
||||
| X86_32/i386/i686 | ✓ 支持 | 老旧 PC 设备 |
|
||||
| ARM64/AArch64 | ✓ 完全支持 | 树莓派4/5, Jetson |
|
||||
| ARMv7/armhf | ✓ 支持 | 树莓派3/2, Orange Pi |
|
||||
| ARMv6 | ⚠ 实验性 | 树莓派Zero |
|
||||
| MIPS | ⚠ 实验性 | 路由器等 |
|
||||
|
||||
## 设备预设
|
||||
|
||||
### ultra_low (极低端设备)
|
||||
- **内存**: < 256MB RAM
|
||||
- **示例**: 树莓派 Zero, 老旧路由器
|
||||
- **配置**:
|
||||
```bash
|
||||
python main.py --preset ultra_low
|
||||
```
|
||||
- **自动设置**:
|
||||
- Workers: 1
|
||||
- 最大缓存: 50MB
|
||||
- 块大小: 32KB
|
||||
- 禁用: WebSocket, SSE, 哈希计算
|
||||
|
||||
### low (低端设备)
|
||||
- **内存**: 256MB - 512MB RAM
|
||||
- **示例**: 树莓派 2, 老旧 VPS
|
||||
- **配置**:
|
||||
```bash
|
||||
python main.py --preset low
|
||||
```
|
||||
- **自动设置**:
|
||||
- Workers: 1
|
||||
- 最大缓存: 100MB
|
||||
- 块大小: 64KB
|
||||
- 启用所有功能
|
||||
|
||||
### medium (中等设备)
|
||||
- **内存**: 512MB - 1GB RAM
|
||||
- **示例**: 树莓派 4 (1GB), 低配 VPS
|
||||
- **配置**:
|
||||
```bash
|
||||
python main.py --preset medium
|
||||
```
|
||||
- **自动设置**:
|
||||
- Workers: 2
|
||||
- 最大缓存: 256MB
|
||||
- 块大小: 128KB
|
||||
|
||||
### high (高端设备)
|
||||
- **内存**: 1GB+ RAM
|
||||
- **示例**: 树莓派 4 (4GB/8GB), 家用服务器
|
||||
- **配置**:
|
||||
```bash
|
||||
python main.py --preset high
|
||||
```
|
||||
- **自动设置**:
|
||||
- Workers: 4
|
||||
- 最大缓存: 512MB
|
||||
- 块大小: 256KB
|
||||
|
||||
### auto (自动检测)
|
||||
- 根据系统资源自动选择预设
|
||||
- **配置**:
|
||||
```bash
|
||||
python main.py --preset auto # 默认
|
||||
```
|
||||
|
||||
## 手动配置
|
||||
|
||||
### 内存限制
|
||||
```bash
|
||||
# 设置 256MB 内存限制
|
||||
python main.py --memory-limit 256M
|
||||
|
||||
# 设置 512MB 内存限制
|
||||
python main.py --memory-limit 512M
|
||||
```
|
||||
|
||||
### 工作进程
|
||||
```bash
|
||||
# 单进程 (低端设备)
|
||||
python main.py --workers 1
|
||||
|
||||
# 双进程
|
||||
python main.py --workers 2
|
||||
```
|
||||
|
||||
### 传输优化
|
||||
```bash
|
||||
# 小块传输 (节省内存)
|
||||
python main.py --chunk-size 32K --buffer-size 64K
|
||||
```
|
||||
|
||||
### 禁用可选功能
|
||||
```bash
|
||||
# 禁用 WebSocket 和 SSE (节省内存)
|
||||
python main.py --disable-ws --disable-sse
|
||||
|
||||
# 禁用哈希计算 (节省 CPU)
|
||||
python main.py --disable-hash
|
||||
```
|
||||
|
||||
## 树莓派部署
|
||||
|
||||
### 方式一: 使用预编译镜像
|
||||
```bash
|
||||
# 拉取 ARM64 镜像
|
||||
docker pull hx100cv/hyc-download-station:v2.3-arm64
|
||||
|
||||
# 运行
|
||||
docker run -d \
|
||||
--name hyc-server \
|
||||
-p 8080:8080 \
|
||||
-v ./data:/data \
|
||||
-v ./downloads:/downloads \
|
||||
hyc-download-station:v2.3-arm64
|
||||
```
|
||||
|
||||
### 方式二: 使用 Docker Compose
|
||||
```bash
|
||||
# 树莓派专用配置
|
||||
docker-compose -f docker-compose.raspberry.yml up -d
|
||||
```
|
||||
|
||||
### 方式三: 轻量级配置
|
||||
```bash
|
||||
# 适用于 512MB RAM 的树莓派
|
||||
docker-compose -f docker-compose.lite.yml up -d
|
||||
```
|
||||
|
||||
## 系统兼容性检查
|
||||
|
||||
```bash
|
||||
# 检查系统兼容性
|
||||
python main.py --check-compat
|
||||
|
||||
# 或使用脚本
|
||||
./scripts/check-compat.sh
|
||||
```
|
||||
|
||||
输出示例:
|
||||
```
|
||||
========================================
|
||||
HYC下载站 v2.3 - 兼容性检查
|
||||
========================================
|
||||
|
||||
系统信息:
|
||||
- 架构: aarch64
|
||||
- 系统: Linux
|
||||
|
||||
✓ ARM64 (64位) - 完全支持
|
||||
✓ Python 3.11 - 支持
|
||||
|
||||
内存检查:
|
||||
- 总内存: 4096MB
|
||||
✓ 内存 1GB+ - 使用 high 预设
|
||||
|
||||
...
|
||||
|
||||
推荐启动命令:
|
||||
python main.py --preset high
|
||||
```
|
||||
|
||||
## Docker 多架构构建
|
||||
|
||||
### 环境准备
|
||||
```bash
|
||||
# 设置 QEMU 仿真 (x86_64 上构建 ARM)
|
||||
./scripts/setup-qemu.sh
|
||||
```
|
||||
|
||||
### 构建镜像
|
||||
```bash
|
||||
# 构建所有架构
|
||||
./scripts/build-multiarch.sh v2.3 hyc-download-station
|
||||
|
||||
# 或手动构建
|
||||
docker buildx build \
|
||||
--platform linux/amd64,linux/arm64,linux/arm/v7 \
|
||||
--tag hx100cv/hyc-download-station:v2.3 \
|
||||
--file docker/Dockerfile.multiarch \
|
||||
--push .
|
||||
```
|
||||
|
||||
### 手动构建特定架构
|
||||
```bash
|
||||
# ARMv7
|
||||
docker build \
|
||||
--platform linux/arm/v7 \
|
||||
--tag hx100cv/hyc-download-station:v2.3-armv7 \
|
||||
--file docker/Dockerfile.lite \
|
||||
--push .
|
||||
|
||||
# ARM64
|
||||
docker build \
|
||||
--platform linux/arm64 \
|
||||
--tag hx100cv/hyc-download-station:v2.3-arm64 \
|
||||
--file docker/Dockerfile.lite \
|
||||
--push .
|
||||
|
||||
# i386 (32位)
|
||||
docker build \
|
||||
--platform linux/386 \
|
||||
--tag hx100cv/hyc-download-station:v2.3-i386 \
|
||||
--file docker/Dockerfile.lite \
|
||||
--push .
|
||||
```
|
||||
|
||||
## 性能调优建议
|
||||
|
||||
### 树莓派 4 (4GB)
|
||||
```bash
|
||||
# 推荐配置
|
||||
python main.py \
|
||||
--preset high \
|
||||
--memory-limit 1G \
|
||||
--workers 2
|
||||
```
|
||||
|
||||
### 树莓派 3
|
||||
```bash
|
||||
# 推荐配置
|
||||
python main.py \
|
||||
--preset medium \
|
||||
--memory-limit 512M \
|
||||
--workers 2
|
||||
```
|
||||
|
||||
### 树莓派 2/Zero
|
||||
```bash
|
||||
# 推荐配置
|
||||
python main.py \
|
||||
--preset low \
|
||||
--memory-limit 256M \
|
||||
--workers 1 \
|
||||
--chunk-size 16K \
|
||||
--disable-ws \
|
||||
--disable-sse
|
||||
```
|
||||
|
||||
## 内存使用监控
|
||||
|
||||
启动后可以通过 API 查看内存使用:
|
||||
```bash
|
||||
# 查看内存状态
|
||||
curl http://localhost:8080/api/v1/monitor
|
||||
|
||||
# 或在 Web 界面查看
|
||||
# 访问 http://localhost:8080/api/ui/
|
||||
```
|
||||
|
||||
## 故障排除
|
||||
|
||||
### 内存不足
|
||||
```
|
||||
症状: OOM (Out of Memory) 错误
|
||||
解决:
|
||||
1. 使用 --preset ultra_low
|
||||
2. 减小 --memory-limit
|
||||
3. 增加 swap 空间
|
||||
```
|
||||
|
||||
### 构建失败
|
||||
```
|
||||
症状: 无法导入模块
|
||||
解决:
|
||||
1. 重新安装依赖: pip install -r requirements.txt
|
||||
2. 检查 Python 版本: python3 --version
|
||||
```
|
||||
|
||||
### ARM 镜像运行失败
|
||||
```
|
||||
症状: Illegal instruction
|
||||
解决:
|
||||
1. 确保使用正确的架构镜像
|
||||
2. 检查 QEMU 设置
|
||||
```
|
||||
+2055
File diff suppressed because it is too large
Load Diff
+1066
File diff suppressed because it is too large
Load Diff
+112
@@ -0,0 +1,112 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""API路由模块"""
|
||||
|
||||
import re
|
||||
from urllib.parse import urlparse, parse_qs
|
||||
|
||||
from .v1 import APIv1
|
||||
from .v2 import APIv2
|
||||
from .admin import AdminAPI
|
||||
from core.api_auth import APIAuthManager
|
||||
|
||||
|
||||
class APIRouter:
|
||||
"""API路由器 - 支持版本化API"""
|
||||
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.api_versions = {
|
||||
'v1': APIv1(config),
|
||||
'v2': APIv2(config)
|
||||
}
|
||||
self.default_version = config.get('api_version', 'v1')
|
||||
|
||||
# 初始化认证管理器(使用共享实例)
|
||||
if not config.get('_auth_manager'):
|
||||
config['_auth_manager'] = APIAuthManager(config)
|
||||
self.auth_manager = config['_auth_manager']
|
||||
|
||||
# 初始化Admin API
|
||||
self.admin_api = AdminAPI(config)
|
||||
|
||||
def handle_request(self, handler, method, path, query):
|
||||
"""处理API请求"""
|
||||
# 注入认证管理器到handler
|
||||
handler.auth_manager = self.auth_manager
|
||||
|
||||
# 解析路径,提取API版本
|
||||
# 格式: /api/v1/... 或 /api/v2/... 或 /api/...
|
||||
|
||||
# 移除 /api/ 前缀
|
||||
if path.startswith('api/'):
|
||||
api_path = path[4:]
|
||||
else:
|
||||
api_path = path
|
||||
|
||||
# 解析版本号
|
||||
parts = api_path.split('/')
|
||||
if parts[0] in ['v1', 'v2']:
|
||||
api_version = parts[0]
|
||||
api_action = '/'.join(parts[1:]) if len(parts) > 1 else ''
|
||||
else:
|
||||
# 检查是否是直接访问的 admin API (不带版本前缀)
|
||||
if api_path.startswith('admin/'):
|
||||
admin_action = api_path[6:] # 移除 'admin/'
|
||||
try:
|
||||
self.admin_api.handle_request(handler, method, admin_action, {})
|
||||
except Exception as e:
|
||||
handler.send_json_response({
|
||||
"error": f"Admin API处理错误: {str(e)}",
|
||||
"path": admin_action
|
||||
}, 500)
|
||||
return
|
||||
else:
|
||||
handler.send_error(400, "未指定API请求版本/指定版本错误")
|
||||
return
|
||||
|
||||
# 解析查询参数
|
||||
parsed_query = parse_qs(query)
|
||||
|
||||
# 检查是否是 admin API (带版本前缀,如 /api/v2/admin/stats)
|
||||
# 注意:auth/verify 应该交给 APIv2 处理,而不是 admin_api
|
||||
if api_version in ['v1', 'v2'] and api_action.startswith('admin/') and not api_action.startswith('admin/auth'):
|
||||
admin_action = api_action[6:] # 移除 'admin/'
|
||||
try:
|
||||
self.admin_api.handle_request(handler, method, admin_action, parsed_query)
|
||||
except Exception as e:
|
||||
handler.send_json_response({
|
||||
"error": f"Admin API处理错误: {str(e)}",
|
||||
"path": admin_action
|
||||
}, 500)
|
||||
return
|
||||
|
||||
# 获取对应的API处理器
|
||||
api_handler = self.api_versions.get(api_version)
|
||||
if not api_handler:
|
||||
handler.send_json_response({
|
||||
"error": f"不支持的API版本: {api_version}",
|
||||
"supported_versions": list(self.api_versions.keys())
|
||||
}, 400)
|
||||
return
|
||||
|
||||
# 调用对应的API处理器
|
||||
try:
|
||||
# 调试模式输出路由信息 (debug-api)
|
||||
if handler._is_debug_enabled('api'):
|
||||
msg = f"\n=== DEBUG API Router ===\n Version: {api_version}\n Action: {api_action}\n Method: {method}\n Query: {parsed_query}"
|
||||
handler._debug_log('api', msg, '\033[35m')
|
||||
|
||||
api_handler.handle_request(handler, method, api_action, parsed_query)
|
||||
except Exception as e:
|
||||
if handler._is_debug_enabled('error'):
|
||||
import traceback
|
||||
tb_str = traceback.format_exc()
|
||||
msg = f"\n=== DEBUG API ERROR ===\n{tb_str}"
|
||||
handler._debug_log('error', msg, '\033[31m')
|
||||
handler.send_json_response({
|
||||
"error": f"API处理错误: {str(e)}",
|
||||
"version": api_version,
|
||||
"path": api_action
|
||||
}, 500)
|
||||
@@ -0,0 +1,286 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
SSE (Server-Sent Events) 处理器模块
|
||||
提供单向事件推送功能
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from typing import Dict, Set, Optional
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass
|
||||
class SSEClient:
|
||||
"""SSE客户端"""
|
||||
client_id: str
|
||||
handler: any # HTTP请求处理器
|
||||
topics: Set[str] = field(default_factory=set)
|
||||
connected_at: float = field(default_factory=time.time)
|
||||
last_activity: float = field(default_factory=time.time)
|
||||
running: bool = False
|
||||
|
||||
|
||||
class SSEHandler:
|
||||
"""SSE事件处理器"""
|
||||
|
||||
# 预定义事件类型
|
||||
EVENT_STATS = 'stats'
|
||||
EVENT_MONITOR = 'monitor'
|
||||
EVENT_SYNC = 'sync'
|
||||
EVENT_DOWNLOAD = 'download'
|
||||
EVENT_SERVER = 'server'
|
||||
EVENT_ERROR = 'error'
|
||||
EVENT_PING = 'ping'
|
||||
|
||||
def __init__(self, config: dict = None):
|
||||
self.config = config or {}
|
||||
self.clients: Dict[str, SSEClient] = {}
|
||||
self.lock = threading.Lock()
|
||||
|
||||
# 心跳配置
|
||||
self.heartbeat_interval = 30 # 秒
|
||||
self.retry_interval = 3000 # 毫秒(客户端重连间隔)
|
||||
|
||||
# 消息缓冲区
|
||||
self.message_buffer_size = 100
|
||||
|
||||
# 统计
|
||||
self.stats = {
|
||||
'total_connections': 0,
|
||||
'total_events_sent': 0,
|
||||
'total_bytes_sent': 0
|
||||
}
|
||||
|
||||
def handle_connection(self, handler, topics: list = None) -> Optional[str]:
|
||||
"""
|
||||
处理新的SSE连接
|
||||
返回客户端ID
|
||||
"""
|
||||
client_id = self._generate_client_id()
|
||||
|
||||
# 设置SSE响应头
|
||||
handler.send_response(200)
|
||||
handler.send_header('Content-Type', 'text/event-stream')
|
||||
handler.send_header('Cache-Control', 'no-cache, no-store, must-revalidate')
|
||||
handler.send_header('Connection', 'keep-alive')
|
||||
handler.send_header('Access-Control-Allow-Origin', '*')
|
||||
handler.send_header('X-Accel-Buffering', 'no') # 禁用Nginx缓冲
|
||||
handler.end_headers()
|
||||
|
||||
# 创建客户端
|
||||
client = SSEClient(
|
||||
client_id=client_id,
|
||||
handler=handler,
|
||||
topics=set(topics) if topics else {self.EVENT_STATS, self.EVENT_SYNC}
|
||||
)
|
||||
|
||||
with self.lock:
|
||||
self.clients[client_id] = client
|
||||
self.stats['total_connections'] += 1
|
||||
|
||||
# 启动心跳
|
||||
client.running = True
|
||||
self._start_heartbeat(client_id)
|
||||
|
||||
# 发送初始连接事件
|
||||
self._send_event(client, self.EVENT_SERVER, {
|
||||
'type': 'connected',
|
||||
'client_id': client_id,
|
||||
'timestamp': time.time(),
|
||||
'topics': list(client.topics)
|
||||
})
|
||||
|
||||
return client_id
|
||||
|
||||
def close_connection(self, client_id: str):
|
||||
"""关闭连接"""
|
||||
with self.lock:
|
||||
client = self.clients.pop(client_id, None)
|
||||
if client:
|
||||
client.running = False
|
||||
|
||||
def subscribe(self, client_id: str, *topics: str):
|
||||
"""订阅主题"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if client:
|
||||
client.topics.update(topics)
|
||||
|
||||
def unsubscribe(self, client_id: str, *topics: str):
|
||||
"""取消订阅"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if client:
|
||||
for topic in topics:
|
||||
client.topics.discard(topic)
|
||||
|
||||
def is_subscribed(self, client_id: str, event_type: str) -> bool:
|
||||
"""检查是否订阅了事件类型"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if not client:
|
||||
return False
|
||||
return event_type in client.topics or '*' in client.topics
|
||||
|
||||
def send_event(self, client_id: str, event_type: str, data: dict):
|
||||
"""发送事件到指定客户端"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if not client or not client.running:
|
||||
return False
|
||||
|
||||
if event_type not in client.topics and '*' not in client.topics:
|
||||
return False
|
||||
|
||||
return self._send_event(client, event_type, data)
|
||||
|
||||
def broadcast(self, event_type: str, data: dict, topic: str = None):
|
||||
"""广播事件到所有客户端"""
|
||||
with self.lock:
|
||||
sent_count = 0
|
||||
dead_clients = []
|
||||
|
||||
for client_id, client in self.clients.items():
|
||||
if not client.running:
|
||||
dead_clients.append(client_id)
|
||||
continue
|
||||
|
||||
# 检查主题匹配
|
||||
if topic and event_type != topic:
|
||||
continue
|
||||
|
||||
if self._send_event(client, event_type, data):
|
||||
sent_count += 1
|
||||
|
||||
# 清理死掉的客户端
|
||||
for client_id in dead_clients:
|
||||
self.clients.pop(client_id, None)
|
||||
|
||||
return sent_count
|
||||
|
||||
def broadcast_to_topic(self, topic: str, event_type: str, data: dict):
|
||||
"""广播到订阅特定主题的客户端"""
|
||||
self.broadcast(event_type, data, topic)
|
||||
|
||||
def send_monitor_update(self, stats: dict):
|
||||
"""发送监控更新"""
|
||||
self.broadcast(self.EVENT_MONITOR, stats)
|
||||
|
||||
def send_sync_update(self, sync_data: dict):
|
||||
"""发送同步更新"""
|
||||
self.broadcast(self.EVENT_SYNC, sync_data)
|
||||
|
||||
def send_download_update(self, download_data: dict):
|
||||
"""发送下载更新"""
|
||||
self.broadcast(self.EVENT_DOWNLOAD, download_data)
|
||||
|
||||
def send_server_event(self, event_data: dict):
|
||||
"""发送服务器事件"""
|
||||
self.broadcast(self.EVENT_SERVER, event_data)
|
||||
|
||||
def get_client_count(self) -> int:
|
||||
"""获取客户端数量"""
|
||||
with self.lock:
|
||||
return len(self.clients)
|
||||
|
||||
def get_stats(self) -> dict:
|
||||
"""获取统计信息"""
|
||||
with self.lock:
|
||||
return {
|
||||
**self.stats,
|
||||
'connected_clients': len(self.clients),
|
||||
'topics': list(set(
|
||||
t for client in self.clients.values()
|
||||
for t in client.topics
|
||||
))
|
||||
}
|
||||
|
||||
def _send_event(self, client: SSEClient, event_type: str, data: dict) -> bool:
|
||||
"""发送单个事件"""
|
||||
try:
|
||||
event_data = {
|
||||
'event': event_type,
|
||||
'timestamp': time.time(),
|
||||
'data': data
|
||||
}
|
||||
|
||||
# SSE格式
|
||||
message = f"event: {event_type}\n"
|
||||
message += f"id: {uuid.uuid4().hex[:16]}\n"
|
||||
message += f"retry: {self.retry_interval}\n"
|
||||
message += "data: " + json.dumps(event_data, ensure_ascii=False) + "\n\n"
|
||||
|
||||
# 发送
|
||||
client.handler.wfile.write(message.encode('utf-8'))
|
||||
client.handler.wfile.flush()
|
||||
|
||||
# 统计
|
||||
self.stats['total_events_sent'] += len(message)
|
||||
client.last_activity = time.time()
|
||||
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"SSE发送失败 {client.client_id}: {e}")
|
||||
client.running = False
|
||||
return False
|
||||
|
||||
def _start_heartbeat(self, client_id: str):
|
||||
"""启动心跳"""
|
||||
def heartbeat():
|
||||
while True:
|
||||
time.sleep(self.heartbeat_interval)
|
||||
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if not client or not client.running:
|
||||
return
|
||||
|
||||
# 发送心跳
|
||||
self.send_event(client_id, self.EVENT_PING, {
|
||||
'timestamp': time.time()
|
||||
})
|
||||
|
||||
thread = threading.Thread(target=heartbeat, daemon=True)
|
||||
thread.start()
|
||||
|
||||
def _generate_client_id(self) -> str:
|
||||
"""生成客户端ID"""
|
||||
return f"sse_{uuid.uuid4().hex[:12]}"
|
||||
|
||||
def cleanup(self, max_idle_time: float = 300.0):
|
||||
"""清理空闲连接"""
|
||||
now = time.time()
|
||||
idle_clients = []
|
||||
|
||||
with self.lock:
|
||||
for client_id, client in self.clients.items():
|
||||
if now - client.last_activity > max_idle_time:
|
||||
idle_clients.append(client_id)
|
||||
|
||||
for client_id in idle_clients:
|
||||
self.send_event(client_id, self.EVENT_ERROR, {
|
||||
'type': 'timeout',
|
||||
'message': 'Connection timed out'
|
||||
})
|
||||
self.close_connection(client_id)
|
||||
|
||||
def get_client_info(self, client_id: str) -> Optional[dict]:
|
||||
"""获取客户端信息"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if not client:
|
||||
return None
|
||||
|
||||
return {
|
||||
'client_id': client.client_id,
|
||||
'topics': list(client.topics),
|
||||
'connected_at': client.connected_at,
|
||||
'last_activity': client.last_activity,
|
||||
'running': client.running
|
||||
}
|
||||
+5949
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,403 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
WebSocket处理器模块
|
||||
提供双向实时通信功能
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Dict, Set, Optional, Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass
|
||||
class WebSocketClient:
|
||||
"""WebSocket客户端"""
|
||||
client_id: str
|
||||
connection: any # WebSocket连接对象
|
||||
topics: Set[str] = field(default_factory=set)
|
||||
connected_at: float = field(default_factory=time.time)
|
||||
last_activity: float = field(default_factory=time.time)
|
||||
metadata: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
class WebSocketManager:
|
||||
"""WebSocket连接管理器"""
|
||||
|
||||
# 预定义主题
|
||||
TOPIC_MONITOR_CPU = 'monitor:cpu'
|
||||
TOPIC_MONITOR_MEMORY = 'monitor:memory'
|
||||
TOPIC_MONITOR_DISK = 'monitor:disk'
|
||||
TOPIC_MONITOR_NETWORK = 'monitor:network'
|
||||
TOPIC_MONITOR_ALL = 'monitor:*'
|
||||
TOPIC_SYNC_PROGRESS = 'sync:progress'
|
||||
TOPIC_SYNC_STATUS = 'sync:status'
|
||||
TOPIC_SYNC_ALL = 'sync:*'
|
||||
TOPIC_DOWNLOAD_PROGRESS = 'download:progress'
|
||||
TOPIC_DOWNLOAD_STATUS = 'download:status'
|
||||
TOPIC_DOWNLOAD_ALL = 'download:*'
|
||||
TOPIC_SERVER_STATUS = 'server:status'
|
||||
TOPIC_ALL = '*'
|
||||
|
||||
def __init__(self, config: dict = None):
|
||||
self.config = config or {}
|
||||
self.clients: Dict[str, WebSocketClient] = {}
|
||||
self.lock = threading.Lock()
|
||||
self._running = False
|
||||
|
||||
# 消息队列(用于批量发送)
|
||||
self.message_queues: Dict[str, list] = {}
|
||||
|
||||
# 心跳间隔(秒)
|
||||
self.heartbeat_interval = 30
|
||||
|
||||
# 统计
|
||||
self.stats = {
|
||||
'total_connections': 0,
|
||||
'total_messages_sent': 0,
|
||||
'total_messages_received': 0
|
||||
}
|
||||
|
||||
def register_client(self, client_id: str, connection, metadata: dict = None) -> WebSocketClient:
|
||||
"""注册新客户端"""
|
||||
with self.lock:
|
||||
client = WebSocketClient(
|
||||
client_id=client_id,
|
||||
connection=connection,
|
||||
topics={self.TOPIC_ALL}, # 默认订阅所有
|
||||
metadata=metadata or {}
|
||||
)
|
||||
self.clients[client_id] = client
|
||||
self.stats['total_connections'] += 1
|
||||
|
||||
# 启动心跳
|
||||
self._start_heartbeat(client_id)
|
||||
|
||||
return client
|
||||
|
||||
def unregister_client(self, client_id: str):
|
||||
"""注销客户端"""
|
||||
with self.lock:
|
||||
client = self.clients.pop(client_id, None)
|
||||
if client:
|
||||
self.message_queues.pop(client_id, None)
|
||||
|
||||
def subscribe(self, client_id: str, *topics: str):
|
||||
"""订阅主题"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if client:
|
||||
client.topics.update(topics)
|
||||
|
||||
def unsubscribe(self, client_id: str, *topics: str):
|
||||
"""取消订阅"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if client:
|
||||
for topic in topics:
|
||||
client.topics.discard(topic)
|
||||
|
||||
def is_subscribed(self, client_id: str, topic: str) -> bool:
|
||||
"""检查客户端是否订阅了主题"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if not client:
|
||||
return False
|
||||
|
||||
if self.TOPIC_ALL in client.topics:
|
||||
return True
|
||||
|
||||
# 检查精确匹配或通配符匹配
|
||||
if topic in client.topics:
|
||||
return True
|
||||
|
||||
# 检查通配符匹配
|
||||
topic_parts = topic.split(':')
|
||||
for subscribed_topic in client.topics:
|
||||
if subscribed_topic.endswith('*'):
|
||||
prefix = subscribed_topic.rstrip('*').rstrip(':')
|
||||
if topic.startswith(prefix):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def send_to_client(self, client_id: str, event_type: str, data: dict, callback: Callable = None):
|
||||
"""发送消息到指定客户端"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
if not client:
|
||||
return False
|
||||
|
||||
message = self._format_message(event_type, data)
|
||||
|
||||
try:
|
||||
if hasattr(client.connection, 'send'):
|
||||
client.connection.send(message)
|
||||
self.stats['total_messages_sent'] += 1
|
||||
client.last_activity = time.time()
|
||||
|
||||
if callback:
|
||||
callback(client_id, True)
|
||||
return True
|
||||
else:
|
||||
# 放入消息队列
|
||||
if client_id not in self.message_queues:
|
||||
self.message_queues[client_id] = []
|
||||
self.message_queues[client_id].append(message)
|
||||
|
||||
if callback:
|
||||
callback(client_id, True)
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"WebSocket发送失败 {client_id}: {e}")
|
||||
if callback:
|
||||
callback(client_id, False)
|
||||
return False
|
||||
|
||||
def broadcast(self, event_type: str, data: dict, topic: str = None):
|
||||
"""广播消息到所有客户端"""
|
||||
with self.lock:
|
||||
sent_count = 0
|
||||
failed_clients = []
|
||||
|
||||
for client_id, client in self.clients.items():
|
||||
# 检查是否匹配主题
|
||||
if topic and not self.is_subscribed(client_id, topic):
|
||||
continue
|
||||
|
||||
message = self._format_message(event_type, data)
|
||||
|
||||
try:
|
||||
if hasattr(client.connection, 'send'):
|
||||
client.connection.send(message)
|
||||
sent_count += 1
|
||||
else:
|
||||
if client_id not in self.message_queues:
|
||||
self.message_queues[client_id] = []
|
||||
self.message_queues[client_id].append(message)
|
||||
sent_count += 1
|
||||
|
||||
client.last_activity = time.time()
|
||||
|
||||
except Exception as e:
|
||||
print(f"广播到 {client_id} 失败: {e}")
|
||||
failed_clients.append(client_id)
|
||||
|
||||
self.stats['total_messages_sent'] += sent_count
|
||||
|
||||
# 清理失败的客户端
|
||||
for client_id in failed_clients:
|
||||
self.unregister_client(client_id)
|
||||
|
||||
return sent_count
|
||||
|
||||
def broadcast_to_topic(self, topic: str, event_type: str, data: dict):
|
||||
"""广播到订阅特定主题的客户端"""
|
||||
with self.lock:
|
||||
sent_count = 0
|
||||
|
||||
for client_id, client in self.clients.items():
|
||||
if self.is_subscribed(client_id, topic):
|
||||
message = self._format_message(event_type, data)
|
||||
|
||||
try:
|
||||
if hasattr(client.connection, 'send'):
|
||||
client.connection.send(message)
|
||||
sent_count += 1
|
||||
else:
|
||||
if client_id not in self.message_queues:
|
||||
self.message_queues[client_id] = []
|
||||
self.message_queues[client_id].append(message)
|
||||
sent_count += 1
|
||||
|
||||
except Exception as e:
|
||||
print(f"发送到 {client_id} 失败: {e}")
|
||||
|
||||
return sent_count
|
||||
|
||||
def broadcast_monitor_update(self, stats: dict):
|
||||
"""广播监控更新"""
|
||||
# 提取关键指标
|
||||
cpu_percent = stats.get('cpu', {}).get('percent', 0)
|
||||
memory_percent = stats.get('memory', {}).get('percent', 0)
|
||||
disk_percent = stats.get('disk', {}).get('percent', 0)
|
||||
|
||||
# 按阈值过滤
|
||||
if cpu_percent > 0 or memory_percent > 0 or disk_percent > 0:
|
||||
self.broadcast('monitor:stats', stats, 'monitor:*')
|
||||
|
||||
def broadcast_sync_progress(self, task_id: str, progress: dict):
|
||||
"""广播同步进度"""
|
||||
self.broadcast('sync:progress', progress, 'sync:*')
|
||||
|
||||
def broadcast_download_progress(self, download_id: str, progress: dict):
|
||||
"""广播下载进度"""
|
||||
self.broadcast('download:progress', progress, 'download:*')
|
||||
|
||||
def get_client_count(self) -> int:
|
||||
"""获取客户端数量"""
|
||||
with self.lock:
|
||||
return len(self.clients)
|
||||
|
||||
def get_client_topics(self, client_id: str) -> Set[str]:
|
||||
"""获取客户端订阅的主题"""
|
||||
with self.lock:
|
||||
client = self.clients.get(client_id)
|
||||
return client.topics.copy() if client else set()
|
||||
|
||||
def get_all_topics(self) -> Set[str]:
|
||||
"""获取所有被订阅的主题"""
|
||||
with self.lock:
|
||||
topics = set()
|
||||
for client in self.clients.values():
|
||||
topics.update(client.topics)
|
||||
return topics
|
||||
|
||||
def get_stats(self) -> dict:
|
||||
"""获取统计信息"""
|
||||
with self.lock:
|
||||
return {
|
||||
**self.stats,
|
||||
'connected_clients': len(self.clients),
|
||||
'active_topics': len(self.get_all_topics()),
|
||||
'queues_queued': sum(len(q) for q in self.message_queues.values())
|
||||
}
|
||||
|
||||
def _format_message(self, event_type: str, data: dict) -> str:
|
||||
"""格式化消息"""
|
||||
return json.dumps({
|
||||
'type': event_type,
|
||||
'timestamp': time.time(),
|
||||
'data': data
|
||||
}, ensure_ascii=False)
|
||||
|
||||
def _start_heartbeat(self, client_id: str):
|
||||
"""启动心跳"""
|
||||
def heartbeat():
|
||||
while client_id in self.clients:
|
||||
try:
|
||||
# 发送心跳
|
||||
self.send_to_client(
|
||||
client_id,
|
||||
'ping',
|
||||
{'timestamp': time.time()}
|
||||
)
|
||||
except Exception:
|
||||
break
|
||||
|
||||
time.sleep(self.heartbeat_interval)
|
||||
|
||||
thread = threading.Thread(target=heartbeat, daemon=True)
|
||||
thread.start()
|
||||
|
||||
def generate_client_id(self) -> str:
|
||||
"""生成客户端ID"""
|
||||
return f"ws_{uuid.uuid4().hex[:12]}"
|
||||
|
||||
def process_messages(self, client_id: str, messages: list):
|
||||
"""处理客户端消息"""
|
||||
for message in messages:
|
||||
self._handle_message(client_id, message)
|
||||
|
||||
def _handle_message(self, client_id: str, message: str):
|
||||
"""处理客户端消息"""
|
||||
self.stats['total_messages_received'] += 1
|
||||
|
||||
try:
|
||||
data = json.loads(message)
|
||||
|
||||
event_type = data.get('type')
|
||||
payload = data.get('data', {})
|
||||
|
||||
if event_type == 'subscribe':
|
||||
# 订阅主题
|
||||
topics = payload.get('topics', [])
|
||||
self.subscribe(client_id, *topics)
|
||||
|
||||
elif event_type == 'unsubscribe':
|
||||
# 取消订阅
|
||||
topics = payload.get('topics', [])
|
||||
self.unsubscribe(client_id, *topics)
|
||||
|
||||
elif event_type == 'ping':
|
||||
# 心跳响应
|
||||
self.send_to_client(client_id, 'pong', {'timestamp': time.time()})
|
||||
|
||||
elif event_type == 'status':
|
||||
# 请求状态 - 返回服务器和客户端状态
|
||||
response_data = {'status': 'ok'}
|
||||
|
||||
if payload.get('monitor'):
|
||||
# 获取实时监控数据
|
||||
try:
|
||||
if handler.monitor:
|
||||
response_data['monitor'] = handler.monitor.get_realtime_stats()
|
||||
else:
|
||||
# 回退到直接使用 psutil
|
||||
import psutil
|
||||
response_data['monitor'] = {
|
||||
'timestamp': datetime.now().isoformat(),
|
||||
'cpu': {
|
||||
'percent': psutil.cpu_percent(interval=0.1),
|
||||
'count': psutil.cpu_count()
|
||||
},
|
||||
'memory': psutil.virtual_memory()._asdict(),
|
||||
'disk': psutil.disk_usage(handler.config.get('base_dir', './downloads'))._asdict()
|
||||
}
|
||||
except Exception as e:
|
||||
response_data['monitor'] = {'error': str(e)}
|
||||
|
||||
if payload.get('ws'):
|
||||
# 获取 WebSocket 统计
|
||||
response_data['ws'] = self.get_stats()
|
||||
|
||||
if payload.get('client'):
|
||||
# 获取当前客户端信息
|
||||
client = self.clients.get(client_id)
|
||||
if client:
|
||||
response_data['client'] = {
|
||||
'client_id': client.client_id,
|
||||
'connected_at': client.connected_at,
|
||||
'last_activity': client.last_activity,
|
||||
'topics': list(client.topics),
|
||||
'metadata': client.metadata
|
||||
}
|
||||
|
||||
self.send_to_client(client_id, 'status:response', response_data)
|
||||
|
||||
elif event_type == 'sync':
|
||||
# 同步相关操作
|
||||
sync_action = payload.get('action', 'status')
|
||||
|
||||
if sync_action == 'status':
|
||||
# 获取同步状态
|
||||
response_data = {'action': 'status'}
|
||||
try:
|
||||
from core.sync_scheduler import SyncScheduler
|
||||
scheduler = SyncScheduler()
|
||||
response_data['sync_status'] = scheduler.get_status() if hasattr(scheduler, 'get_status') else {'message': 'sync scheduler running'}
|
||||
except Exception as e:
|
||||
response_data['sync_status'] = {'error': str(e)}
|
||||
|
||||
self.send_to_client(client_id, 'sync:response', response_data)
|
||||
|
||||
elif sync_action == 'list':
|
||||
# 获取同步任务列表
|
||||
response_data = {'action': 'list'}
|
||||
self.send_to_client(client_id, 'sync:response', response_data)
|
||||
|
||||
elif sync_action == 'trigger':
|
||||
# 触发手动同步
|
||||
response_data = {'action': 'trigger', 'status': 'pending'}
|
||||
self.send_to_client(client_id, 'sync:response', response_data)
|
||||
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
except Exception as e:
|
||||
print(f"处理WebSocket消息失败: {e}")
|
||||
Reference in New Issue
Block a user