From 2a899c411e35e6b1b81573fd619df8ec106667c4 Mon Sep 17 00:00:00 2001 From: HYC Fixer Date: Sun, 30 Aug 2026 12:20:17 +0800 Subject: [PATCH] =?UTF-8?q?P0-C:=20=E5=AE=89=E5=85=A8=E4=BF=AE=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 会话文件移出 web 根(data/ 目录),auth_secret 启动自动生成并持久化(0600) - cookie 签名改 HMAC-SHA256 + compare_digest;会话表加锁 + 原子写 - serve_path 增加敏感文件黑名单(纵深防御) - PyPI: 移除 ?url= 任意 URL 回退(SSRF),限制 scheme,缓存键防路径穿越 - v2 文件元数据/版本/缩略图端点全部加 is_safe_path;缩略图尺寸与像素上限 - api-docs/generate 限写 docs 目录;user/password 强制旧密码;GET /api/v2/config 脱敏 - 目录列表/错误页 HTML 转义(防存储型 XSS);Content-Disposition 文件名清洗 - 信号处理改优雅退出(移除 os._exit);启动时默认凭据安全警告 --- api/admin.py | 14 +- api/v2.py | 8886 +++++++++++++++++++------------------- core/api_auth.py | 184 +- handlers/http_handler.py | 2998 ++++++------- main.py | 26 +- mirrors/pypi.py | 1596 +++---- 6 files changed, 6938 insertions(+), 6766 deletions(-) diff --git a/api/admin.py b/api/admin.py index 9f9b350..e9767f4 100644 --- a/api/admin.py +++ b/api/admin.py @@ -90,18 +90,8 @@ class AdminAPI: 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 - }) + # 列出活跃会话(auth_manager 内部加锁遍历) + sessions = self.auth_manager.list_sessions() handler.send_json_response({ "sessions": sessions, diff --git a/api/v2.py b/api/v2.py index 128f766..8d868f6 100644 --- a/api/v2.py +++ b/api/v2.py @@ -1,4414 +1,4472 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - -"""API v2 版本处理模块 - 增强版本""" - -import os -import json -import re -import time -from datetime import datetime - -from .v1 import APIv1 -from .admin import AdminAPI -from core.utils import format_file_size -from core.api_auth import require_auth, check_endpoint_auth - - -class APIv2(APIv1): - """API v2 - 增强版本,继承v1并添加新功能""" - - def __init__(self, config): - super().__init__(config) - self.api_version = "v2" - # 管理员API处理器(始终创建,auth_type检查在装饰器中处理) - self.admin_api = AdminAPI(config) - - def handle_request(self, handler, method, path, query_params): - """处理API v2请求""" - import sys - - # 调试输出 (debug-v2) - if handler._is_debug_enabled('v2'): - msg = f"\n=== DEBUG APIv2.handle_request ===\n path: '{path}'\n method: '{method}'" - handler._debug_log('v2', msg, '\033[36m') - - # v2 端点列表(这些是 v2 专有端点,不应该交给 APIv1 处理) - v2_endpoints = [ - 'admin/', - 'search/enhanced', - 'search/by-tag', - 'search/by-date', - 'stats/detailed', - 'stats/trending', - 'stats/download-trend', - 'stats/download-by-period', - 'stats/rank', - 'cache/popular', - 'file/', - 'metadata/', - 'monitor/', - 'health/', - 'health/', - 'alerts', - 'alerts/', - 'webhooks', - 'sync/', - 'cache/', - 'mirrors', - 'pypi', - 'config', - 'server/', - 'health', - 'downloads/', - 'metrics', - 'activity', - 'user/', - 'jdk/', - ] - - # 检查是否是 v2 专有端点 - is_v2_endpoint = False - for ep in v2_endpoints: - if path == ep or path.startswith(ep): - is_v2_endpoint = True - break - - # 如果是 v2 端点,调用 auth_manager 时要跳过 v1 路径 - auth_manager = getattr(handler, 'auth_manager', None) - - # ========== 首先处理 v2 专有端点 ========== - if is_v2_endpoint: - # 公开端点列表(无需认证) - public_endpoints = [ - 'admin/auth/verify', - 'user/login', - 'search/enhanced', - 'search/by-tag', - 'search/by-date', - 'jdk/retrieve/', - 'jdk/list', - ] - - # 检查是否公开端点 - is_public = path in public_endpoints or any(path.startswith(ep + '/') for ep in public_endpoints) - - # 如果auth_type为none,跳过所有认证检查 - auth_type = self.config.get('auth_type', 'none') - skip_auth = auth_type == 'none' - - # 如果需要认证(不是公开端点且auth_type不是none) - if auth_manager and not is_public and not skip_auth: - # 传入完整路径(api/v2/ 前缀),与 ADMIN_API_ENDPOINTS 规则匹配 - auth_check = check_endpoint_auth(method, f"api/v2/{path}", auth_manager) - if auth_check['required']: - auth_result = auth_manager.validate_request(handler, 'admin') - if not auth_result.get('authenticated'): - handler.send_response(401) - handler.send_header('WWW-Authenticate', 'Bearer') - handler.send_header('Access-Control-Allow-Origin', '*') - handler.send_json_response({ - "error": "认证Required", - "code": "UNAUTHORIZED", - "required_permission": auth_check['permission'] - }) - return - - if auth_check['permission']: - if not auth_manager.check_permission(auth_result, auth_check['permission']): - handler.send_header('Access-Control-Allow-Origin', '*') - handler.send_json_response({ - "error": "权限不足", - "code": "FORBIDDEN", - "required_permission": auth_check['permission'] - }, 403) - return - - handler.auth_result = auth_result - - # ========== v2 端点处理 ========== - # 认证验证 API - if path == 'admin/auth/verify': - if method == 'POST': - self.api_verify_auth(handler) - else: - handler.send_error(405) - return - - # 管理员统计 API - if path == 'admin/stats': - self.admin_api.handle_request(handler, method, 'stats', query_params) - return - - # v2新增的增强功能 - # 增强搜索API(无需认证) - if path == 'search/enhanced': - if method == 'GET': - self.api_search_files_enhanced(handler, query_params) - else: - handler.send_error(405) - return - elif path == 'search/by-tag': - if method == 'GET': - self.api_search_by_tag(handler, query_params) - else: - handler.send_error(405) - return - elif path == 'search/by-date': - if method == 'GET': - self.api_search_by_date(handler, query_params) - else: - handler.send_error(405) - return - - # JDK API (调用 v1 中的实现) - elif path.startswith('jdk/'): - jdk_path = path[4:] # 移除 'jdk/' 前缀 - self.handle_jdk_api(handler, method, jdk_path) - return - - # 增强统计API - elif path == 'stats/detailed': - if method == 'GET': - self.api_get_stats_detailed(handler) - else: - handler.send_error(405) - return - elif path == 'stats/trending': - if method == 'GET': - self.api_get_trending_files(handler, query_params) - else: - handler.send_error(405) - return - - # 增强文件操作API - elif path.startswith('file/') and path.endswith('/metadata'): - filename = path[5:-9] # 移除 'file/' 和 '/metadata' - if method == 'GET': - self.api_get_file_metadata(handler, filename) - elif method == 'PUT': - self.api_update_file_metadata(handler, filename) - else: - handler.send_error(405) - return - - # 批量元数据操作 - elif path == 'metadata/batch': - if method == 'GET': - self.api_get_batch_metadata(handler, query_params) - elif method == 'PUT': - self.api_update_batch_metadata(handler) - else: - handler.send_error(405) - return - - # 文件版本控制 - elif path.startswith('file/') and '/versions' in path: - parts = path.split('/') - if len(parts) >= 3 and parts[-1] == 'versions': - filename = '/'.join(parts[1:-1]) - if method == 'GET': - self.api_get_file_versions(handler, filename) - elif method == 'POST': - self.api_create_file_version(handler, filename) - else: - handler.send_error(405) - return - - # 缩略图API - elif path.startswith('file/') and path.endswith('/thumbnail'): - filename = path[5:-10] # 移除 'file/' 和 '/thumbnail' - if method == 'GET': - self.api_get_file_thumbnail(handler, filename, query_params) - else: - handler.send_error(405) - return - - # 服务器监控(实时数据) - elif path == 'monitor/realtime': - if method == 'GET': - self.api_get_realtime_stats(handler) - else: - handler.send_error(405) - return - elif path == 'monitor/history': - if method == 'GET': - self.api_get_monitor_history(handler, query_params) - else: - handler.send_error(405) - return - elif path == 'monitor/detailed': - if method == 'GET': - self.api_get_monitor_detailed(handler) - else: - handler.send_error(405) - return - - # ========== 镜像源健康检查 ========== - elif path == 'health/sources': - if method == 'GET': - self.api_get_source_health(handler, query_params) - else: - handler.send_error(405) - return - elif path.startswith('health/check/'): - source_name = path[13:] # 移除 'health/check/' - if method == 'GET': - self.api_check_source(handler, source_name) - else: - handler.send_error(405) - return - elif path == 'health/failover': - if method == 'GET': - self.api_get_failover_status(handler) - else: - handler.send_error(405) - return - elif path.startswith('health/failover/') and method == 'POST': - mirror_type = path[18:] # 移除 'health/failover/' - self.api_trigger_failover(handler, mirror_type) - return - elif path == 'health/stats': - if method == 'GET': - from core.health_check import HealthChecker - checker = HealthChecker(self.config.get('health_check', {})) - handler.send_json_response(checker.get_stats()) - else: - handler.send_error(405) - return - - # Webhook支持 - elif path == 'webhooks': - if method == 'GET': - self.api_list_webhooks(handler) - elif method == 'POST': - self.api_create_webhook(handler) - else: - handler.send_error(405) - return - elif path.startswith('webhooks/'): - webhook_id = path[9:] - - # 交付历史: webhooks/{id}/deliveries - if webhook_id.endswith('/deliveries'): - actual_id = webhook_id[:-11] - if method == 'GET': - self.api_get_webhook_deliveries(handler, actual_id) - else: - handler.send_error(405) - return - - # 统计: webhooks/{id}/stats - if webhook_id.endswith('/stats'): - actual_id = webhook_id[:-6] - if method == 'GET': - self.api_get_webhook_stats(handler, actual_id) - else: - handler.send_error(405) - return - - # 单个 webhook 操作 - if method == 'GET': - self.api_get_webhook(handler, webhook_id) - elif method == 'DELETE': - self.api_delete_webhook(handler, webhook_id) - elif method == 'POST': - # 检查是否是测试请求 - if '/test' in webhook_id: - actual_id = webhook_id.split('/')[0] - self.api_test_webhook(handler, actual_id) - else: - self.api_test_webhook(handler, webhook_id) - elif method == 'PUT': - self.api_update_webhook(handler, webhook_id) - else: - handler.send_error(405) - return - - # ========== 以下端点需要认证 ========== - - # 同步管理API - elif path == 'sync/sources': - if method == 'GET': - self.api_get_sync_sources(handler) - elif method == 'POST': - self.api_add_sync_source(handler) - else: - handler.send_error(405) - return - elif path.startswith('sync/') and path.endswith('/start'): - source_name = path[5:-6] - if method == 'POST': - self.api_start_sync(handler, source_name) - else: - handler.send_error(405) - return - elif path.startswith('sync/') and path.endswith('/stop'): - source_name = path[5:-5] - if method == 'POST': - self.api_stop_sync(handler, source_name) - else: - handler.send_error(405) - return - elif path.startswith('sync/') and path.endswith('/status'): - source_name = path[5:-7] - if method == 'GET': - self.api_get_sync_status(handler, source_name) - else: - handler.send_error(405) - return - elif path == 'sync/history': - if method == 'GET': - self.api_get_sync_history(handler, query_params) - else: - handler.send_error(405) - return - elif path == 'sync/packages': - if method == 'POST': - self.api_sync_packages(handler) - else: - handler.send_error(405) - return - elif path.startswith('sync/packages/') and path.endswith('/status'): - source_name = path[14:-8] - if method == 'GET': - self.api_get_temp_sync_status(handler, source_name) - else: - handler.send_error(405) - return - - # 缓存管理API - elif path == 'cache/stats': - if method == 'GET': - self.api_get_cache_stats(handler) - else: - handler.send_error(405) - return - elif path == 'cache/clean' and method == 'POST': - self.api_clean_cache(handler) - return - elif path == 'cache/usage': - if method == 'GET': - self.api_get_cache_usage(handler) - else: - handler.send_error(405) - return - - # 镜像加速源API - elif path == 'mirrors': - if method == 'GET': - self.api_list_mirrors(handler) - else: - handler.send_error(405) - return - - # 镜像管理 API - 必须放在特殊镜像处理之前 - # mirrors/xxx/enable, mirrors/xxx/refresh, mirrors/xxx (PUT/DELETE) - elif path.endswith('/enable'): - # 格式: mirrors/xxx/enable - parts = path.split('/') - if len(parts) >= 3: - mirror_name = parts[1] - if method == 'PUT': - self.api_enable_mirror(handler, mirror_name, query_params) - else: - handler.send_error(405) - return - elif path.endswith('/refresh'): - # 格式: mirrors/xxx/refresh - parts = path.split('/') - if len(parts) >= 3: - mirror_name = parts[1] - if method == 'POST': - self.api_refresh_mirror(handler, mirror_name) - else: - handler.send_error(405) - return - - # 加速源访问 API - mirrors/pypi/*, mirrors/npm/*, mirrors/go/* 等 - # 注意:排除 mirrors/pypi/enable, mirrors/pypi/refresh 等管理路径 - elif (path.startswith('mirrors/pypi/') and not '/enable' in path and not '/refresh' in path) or path == 'mirrors/pypi': - # PyPI 加速源 - 优先本地,没有再从上游拉取 - from mirrors import get_mirror_handler - mirror_path = path.replace('mirrors/pypi/', '', 1) - if path == 'mirrors/pypi': - mirror_path = 'pypi' - handler_class = get_mirror_handler('pypi') - if handler_class: - # 创建实例 - mirror_config = self.config.get('mirrors', {}).get('pypi', {}) - base_dir = self.config.get('base_dir', './downloads') - storage_subdir = mirror_config.get('storage_dir', 'pypi') - # 配置: base_dir=downloads目录, storage_dir=pypi子目录 - mirror_handler = handler_class({ - 'base_dir': base_dir, - 'storage_dir': storage_subdir, - 'upstream_url': mirror_config.get('url', 'https://pypi.org') - }) - mirror_handler.handle_request(handler, mirror_path) - else: - handler.send_error(404, "PyPI mirror not configured") - return - elif path.startswith('mirrors/npm/') or path == 'mirrors/npm': - # NPM 加速源 - 优先本地,没有再从上游拉取 - from mirrors import get_mirror_handler - mirror_path = path.replace('mirrors/npm/', '', 1) - if path == 'mirrors/npm': - mirror_path = 'npm' - handler_class = get_mirror_handler('npm') - if handler_class: - mirror_config = self.config.get('mirrors', {}).get('npm', {}) - base_dir = self.config.get('base_dir', './downloads') - storage_subdir = mirror_config.get('storage_dir', 'npm') - mirror_handler = handler_class({ - 'base_dir': base_dir, - 'storage_dir': storage_subdir, - 'upstream_url': mirror_config.get('url', 'https://registry.npmjs.org') - }) - mirror_handler.handle_request(handler, mirror_path) - else: - handler.send_error(404, "NPM mirror not configured") - return - elif path.startswith('mirrors/go/') or path == 'mirrors/go': - # Go 加速源 - 优先本地,没有再从上游拉取 - from mirrors import get_mirror_handler - mirror_path = path.replace('mirrors/go/', '', 1) - if path == 'mirrors/go': - mirror_path = 'go' - handler_class = get_mirror_handler('go') - if handler_class: - mirror_config = self.config.get('mirrors', {}).get('go', {}) - base_dir = self.config.get('base_dir', './downloads') - storage_subdir = mirror_config.get('storage_dir', 'go') - mirror_handler = handler_class({ - 'base_dir': base_dir, - 'storage_dir': storage_subdir, - 'upstream_url': mirror_config.get('url', 'https://goproxy.cn') - }) - mirror_handler.handle_request(handler, mirror_path) - else: - handler.send_error(404, "Go mirror not configured") - return - elif path.startswith('mirrors/docker/') or path == 'mirrors/docker': - # Docker 加速源 - 优先本地,没有再从上游拉取 - from mirrors import get_mirror_handler - mirror_path = path.replace('mirrors/docker/', '', 1) - if path == 'mirrors/docker': - mirror_path = 'docker' - handler_class = get_mirror_handler('docker') - if handler_class: - mirror_config = self.config.get('mirrors', {}).get('docker', {}) - base_dir = self.config.get('base_dir', './downloads') - storage_subdir = mirror_config.get('storage_dir', 'docker') - mirror_handler = handler_class({ - 'base_dir': base_dir, - 'storage_dir': storage_subdir, - 'upstream_url': mirror_config.get('url', 'https://registry.hub.docker.com') - }) - mirror_handler.handle_request(handler, mirror_path) - else: - handler.send_error(404, "Docker mirror not configured") - return - - # PyPI包路径处理 - 处理 /pypi/packages/... 和 /pypi/web/... 路径 - # 这些路径来自镜像返回的HTML中的绝对链接 - elif path.startswith('pypi/packages/') or path.startswith('pypi/web/') or path.startswith('pypi/simple/'): - import re - import sys - from mirrors import PyPIMirror - # 从Referer中提取镜像名称 - referer = handler.headers.get('Referer', '') - mirror_name = None - if 'mirrors/' in referer: - match = re.search(r'mirrors/([^/]+)', referer) - if match: - mirror_name = match.group(1) - - # 优先使用Referer中指定的镜像,否则查找任意可用的pypi类型镜像 - mirrors_config = self.config.get('mirrors', {}) - if mirror_name and mirror_name in mirrors_config: - pypi_config = mirrors_config[mirror_name] - else: - # 尝试查找任意pypi类型的镜像(按优先级:pypi-cn, pypi) - pypi_config = None - for pref_name in ['pypi-cn', 'pypi']: - if pref_name in mirrors_config and mirrors_config[pref_name].get('type') == 'pypi': - pypi_config = mirrors_config[pref_name] - break - if not pypi_config: - # 尝试查找任意pypi类型的镜像 - for name, cfg in mirrors_config.items(): - if cfg.get('type') == 'pypi': - pypi_config = cfg - break - base_dir = self.config.get('base_dir', './downloads') - storage_subdir = pypi_config.get('storage_dir', 'pypi') if pypi_config else 'pypi' - pypi_handler = PyPIMirror({ - 'base_dir': base_dir, - 'storage_dir': storage_subdir, - 'upstream_url': pypi_config.get('url', 'https://pypi.org') if pypi_config else 'https://pypi.org' - }) - pypi_handler.handle_request(handler, path) - return - - # 通用镜像处理 - 支持自定义镜像名称如 pypi-cn, npm-cn 等 - elif path.startswith('mirrors/') and method == 'PUT': - # 更新自定义加速源 - mirror_name = path[8:] - self.api_update_mirror(handler, mirror_name) - return - elif path.startswith('mirrors/') and method == 'DELETE': - # 删除自定义加速源 - mirror_name = path[8:] - self.api_delete_mirror(handler, mirror_name) - return - elif path == 'mirrors' and method == 'POST': - # 添加自定义加速源 - self.api_add_mirror(handler) - return - - # 通用镜像处理 - 支持自定义镜像名称如 pypi-cn, npm-cn 等 - elif path.startswith('mirrors/'): - - from mirrors import HttpMirror, get_mirror_handler, get_default_upstream - # 解析镜像名称和路径 - # 格式: mirrors/{mirror_name}/... 或 mirrors/{mirror_name} - # 例如: mirrors/pypi-cn/simple/setuptools -> mirror_name=pypi-cn, mirror_path=simple/setuptools - parts = path[8:].split('/', 1) - mirror_name = parts[0] - # 去掉 mirror_name 前缀,只保留后面的路径 - if len(parts) > 1: - full_path = parts[1] - # 去掉路径开头的 mirror_name(如 pypi-cn/) - if full_path.startswith(mirror_name + '/'): - mirror_path = full_path[len(mirror_name)+1:] - else: - mirror_path = full_path - else: - mirror_path = mirror_name - - # 获取镜像配置 - mirror_config = self.config.get('mirrors', {}).get(mirror_name, {}) - - if not mirror_config: - handler.send_error(404, f"Mirror '{mirror_name}' not configured") - return - - # 获取镜像类型 - mirror_type = mirror_config.get('type', 'http') - - # 调试 - debug_file = '/tmp/pypi_debug.log' - with open(debug_file, 'a') as f: - f.write(f"[V2] mirror_type={mirror_type}, handler_class={get_mirror_handler(mirror_type)}\n") - - # 获取处理器类 - handler_class = get_mirror_handler(mirror_type) - if not handler_class: - handler_class = HttpMirror - - # 创建处理器实例 - base_dir = self.config.get('base_dir', './downloads') - storage_subdir = mirror_config.get('storage_dir', mirror_name) - # 拼接完整存储路径 - storage_dir = os.path.join(base_dir, storage_subdir) - - # 根据镜像类型确定上游URL - default_upstream = 'https://pypi.org/simple' - if mirror_type in ('pypi', 'pip', 'pipenv', 'poetry'): - default_upstream = 'https://pypi.org/simple' - elif mirror_type == 'npm': - default_upstream = 'https://registry.npmjs.org' - elif mirror_type == 'go': - default_upstream = 'https://goproxy.cn' - elif mirror_type == 'docker': - default_upstream = 'https://registry.hub.docker.com' - else: - default_upstream = get_default_upstream(mirror_type) - - mirror_handler = handler_class({ - 'base_dir': base_dir, - 'storage_dir': storage_dir, - 'upstream_url': mirror_config.get('url', default_upstream), - 'type': mirror_type, - 'cache_enabled': mirror_config.get('cache_enabled', True), - 'cache_ttl': mirror_config.get('cache_ttl', 3600) - }) - mirror_handler.handle_request(handler, mirror_path) - return - - # 用户管理 - elif path == 'user/login': - if method == 'POST': - self.api_login(handler) - else: - handler.send_error(405) - return - elif path == 'user/password': - if method == 'POST': - self.api_change_password(handler) - else: - handler.send_error(405) - return - elif path == 'user/login-logs': - if method == 'GET': - self.api_get_login_logs(handler, query_params) - else: - handler.send_error(405) - return - - # 配置文件管理 - elif path == 'config': - if method == 'GET': - self.api_get_config(handler) - elif method == 'PUT': - self.api_save_config(handler) - else: - handler.send_error(405) - return - elif path == 'config/reload': - if method == 'POST': - self.api_reload_config(handler) - else: - handler.send_error(405) - return - elif path == 'config/changes': - if method == 'GET': - self.api_get_config_changes(handler) - else: - handler.send_error(405) - return - - # 告警管理 - elif path == 'alerts': - if method == 'GET': - self.api_get_alerts(handler, query_params) - else: - handler.send_error(405) - return - elif path.startswith('alerts/') and '/acknowledge' in path: - alert_id = path[7:].split('/')[0] - if method == 'POST': - self.api_acknowledge_alert(handler, alert_id) - else: - handler.send_error(405) - return - elif path == 'alerts/clear': - if method == 'POST': - self.api_clear_alerts(handler) - else: - handler.send_error(405) - return - elif path == 'alerts/test': - if method == 'POST': - self.api_test_alert(handler) - else: - handler.send_error(405) - return - elif path == 'alerts/config': - if method == 'GET': - self.api_get_alert_config(handler) - elif method == 'PUT': - self.api_save_alert_config(handler) - else: - handler.send_error(405) - return - - # Prometheus 指标 - elif path == 'metrics': - if method == 'GET': - self.api_get_metrics(handler) - else: - handler.send_error(405) - return - - # 下载趋势 - elif path == 'stats/download-trend': - if method == 'GET': - self.api_get_download_trend(handler, query_params) - else: - handler.send_error(405) - return - - # 按周期下载统计 - elif path == 'stats/download-by-period': - if method == 'GET': - self.api_get_download_by_period(handler, query_params) - else: - handler.send_error(405) - return - - # 下载排行 - elif path == 'stats/rank': - if method == 'GET': - self.api_get_download_rank(handler, query_params) - else: - handler.send_error(405) - return - - # 热门缓存 - elif path == 'cache/popular': - if method == 'GET': - self.api_get_hot_cache(handler, query_params) - else: - handler.send_error(405) - return - - # 最近活动 - elif path == 'activity': - if method == 'GET': - self.api_get_recent_activity(handler, query_params) - else: - handler.send_error(405) - return - - # 服务器重启管理 - elif path == 'server/restart': - if method == 'GET': - self.api_get_restart_status(handler) - elif method == 'POST': - self.api_graceful_restart(handler, query_params) - else: - handler.send_error(405) - return - elif path == 'server/restart/confirm': - if method == 'POST': - self.api_confirm_restart(handler) - else: - handler.send_error(405) - return - elif path == 'server/restart/immediate': - if method == 'POST': - self.api_immediate_restart(handler) - else: - handler.send_error(405) - return - elif path == 'server/restart/pending': - if method == 'GET': - self.api_get_pending_requests(handler) - else: - handler.send_error(405) - return - elif path == 'server/restart/history': - if method == 'GET': - self.api_get_restart_history(handler) - else: - handler.send_error(405) - return - elif path == 'server/restart/config': - if method == 'GET': - self.api_get_restart_config(handler) - elif method == 'PUT': - self.api_update_restart_config(handler) - else: - handler.send_error(405) - return - - # 缓存预热管理 - elif path == 'cache/prewarm': - if method == 'GET': - self.api_get_prewarm_status(handler) - elif method == 'POST': - self.api_run_prewarm(handler, query_params) - else: - handler.send_error(405) - return - elif path == 'cache/prewarm/stats': - if method == 'GET': - self.api_get_prewarm_stats(handler) - else: - handler.send_error(405) - return - elif path == 'cache/prewarm/items': - if method == 'GET': - self.api_get_prewarm_items(handler, query_params) - elif method == 'POST': - self.api_add_prewarm_items(handler) - else: - handler.send_error(405) - return - elif path == 'cache/prewarm/history': - if method == 'GET': - self.api_get_prewarm_history(handler) - else: - handler.send_error(405) - return - elif path == 'cache/prewarm/clear': - if method == 'POST': - self.api_clear_prewarm_queue(handler) - else: - handler.send_error(405) - return - elif path == 'cache/prewarm/popular': - if method == 'GET': - self.api_get_popular_items(handler, query_params) - elif method == 'POST': - self.api_add_popular_items(handler, query_params) - else: - handler.send_error(405) - return - elif path == 'cache/prewarm/config': - if method == 'GET': - self.api_get_prewarm_config(handler) - elif method == 'PUT': - self.api_save_prewarm_config(handler) - else: - handler.send_error(405) - return - - # API 文档 - elif path == 'api-docs.json': - if method == 'GET': - self.api_get_api_docs(handler) - else: - handler.send_error(405) - return - elif path == 'api-docs.yaml': - if method == 'GET': - self.api_get_api_docs(handler, format='yaml') - else: - handler.send_error(405) - return - elif path == 'api-docs/generate': - if method == 'POST': - self.api_generate_api_docs(handler) - else: - handler.send_error(405) - return - - # 服务器信息 - elif path == 'server/info': - if method == 'GET': - self.api_get_server_info(handler) - else: - handler.send_error(405) - return - - # v2 端点都没匹配到,尝试调用 APIv1 - try: - return super().handle_request(handler, method, path, query_params) - except: - pass - - # 404 - 未找到端点 - handler.send_error(404) - - # ==================== 认证 API ==================== - - def api_verify_auth(self, handler): - """验证认证状态""" - # 调试模式输出 (debug-v2) - if handler._is_debug_enabled('v2'): - auth_header = handler.headers.get('Authorization', '') - api_key = handler.headers.get('X-API-Key', '') - auth_header_display = auth_header[:30] + '...' if len(auth_header) > 30 else auth_header - api_key_display = api_key[:20] + '...' if len(api_key) > 20 else api_key - msg = f"\n=== DEBUG api_verify_auth ===\n auth_header: '{auth_header_display}'\n api_key: '{api_key_display}'\n auth_type: {self.config.get('auth_type', 'none')}" - handler._debug_log('v2', msg, '\033[32m') - - auth_manager = getattr(handler, 'auth_manager', None) - - # 如果 auth_type 为 none,则允许访问 - if self.config.get('auth_type', 'none') == 'none': - handler.send_json_response({ - "valid": True, - "level": "admin", - "user_id": "anonymous", - "permissions": ["admin:*", "files:*", "sync:*", "keys:*"], - "expires_at": None, - "message": "Auth disabled - full access" - }) - return - - # 使用 auth_manager 验证(支持 Bearer Token、API Key、Cookie、Query Parameter) - if auth_manager: - result = auth_manager.validate_request(handler) - if result.get('authenticated'): - handler.send_json_response({ - "valid": True, - "level": result.get('level', 'admin'), - "user_id": result.get('user_id', result.get('key_id', 'unknown')), - "permissions": result.get('permissions', ["*"]), - "method": result.get('method', 'unknown') - }) - return - - # 无效认证 - handler.send_json_response({ - "valid": False, - "error": "Invalid or expired credentials", - "auth_type": self.config.get('auth_type', 'token') - }, 401) - - # ==================== 增强搜索API ==================== - - def api_search_files_enhanced(self, handler, query_params): - """增强的文件搜索 - 支持更多选项""" - search_term = query_params.get('q', [''])[0].lower() - search_type = query_params.get('type', ['all'])[0] - search_mode = query_params.get('mode', ['fuzzy'])[0] # fuzzy, exact, regex - max_results = int(query_params.get('limit', ['100'])[0]) - offset = int(query_params.get('offset', ['0'])[0]) - include_content = query_params.get('include_content', ['false'])[0].lower() == 'true' - - if not search_term: - handler.send_json_response({"error": "No search term provided"}, 400) - return - - results = [] - search_time = 0 - import time as time_module - start_time = time_module.time() - - for root, dirs, files in os.walk(self.config['base_dir']): - if search_type in ['all', 'dir']: - for dir_name in dirs: - match = False - if search_mode == 'exact': - match = search_term == dir_name.lower() - elif search_mode == 'regex': - try: - match = re.search(search_term, dir_name.lower()) is not None - except re.error: - match = False - else: # fuzzy - match = search_term in dir_name.lower() - - if match: - full_path = os.path.join(root, dir_name) - rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") - try: - mtime = os.path.getmtime(full_path) - result = { - "name": dir_name, - "path": rel_path + "/", - "type": "directory", - "size": 0, - "modified": datetime.fromtimestamp(mtime).isoformat(), - "match_score": self._calculate_match_score(dir_name, search_term, search_mode) - } - if include_content: - result['item_count'] = len(os.listdir(full_path)) - results.append(result) - except OSError: - continue - - if search_type in ['all', 'file']: - for file_name in files: - match = False - if search_mode == 'exact': - match = search_term == file_name.lower() - elif search_mode == 'regex': - try: - match = re.search(search_term, file_name.lower()) is not None - except re.error: - match = False - else: # fuzzy - match = search_term in file_name.lower() - - if match: - full_path = os.path.join(root, file_name) - rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") - try: - size = os.path.getsize(full_path) - mtime = os.path.getmtime(full_path) - mime_type, _ = os.path.splitext(file_name) - mime_type = mime_type[1:].lower() if mime_type else 'unknown' - - result = { - "name": file_name, - "path": rel_path, - "type": mime_type, - "size": size, - "size_formatted": self.format_file_size(size), - "modified": datetime.fromtimestamp(mtime).isoformat(), - "match_score": self._calculate_match_score(file_name, search_term, search_mode) - } - - if include_content: - if mime_type in ['txt', 'log', 'md', 'json', 'xml', 'html']: - try: - with open(full_path, 'r', encoding='utf-8', errors='ignore') as f: - content = f.read(5000) - result['content_preview'] = content - except: - pass - - results.append(result) - except OSError: - continue - - if len(results) >= max_results + offset: - break - - search_time = time_module.time() - start_time - - # 按匹配分数排序 - results.sort(key=lambda x: x.get('match_score', 0), reverse=True) - - paginated_results = results[offset:offset + max_results] - - handler.send_json_response({ - "query": search_term, - "search_type": search_type, - "search_mode": search_mode, - "total_count": len(results), - "returned_count": len(paginated_results), - "offset": offset, - "limit": max_results, - "search_time": round(search_time, 3), - "results": paginated_results - }) - - def api_search_by_tag(self, handler, query_params): - """按标签搜索文件""" - tag = query_params.get('tag', [''])[0] - if not tag: - handler.send_json_response({"error": "No tag specified"}, 400) - return - - # 这里假设有一个标签存储系统 - # 实际实现需要维护文件标签数据库 - results = [] - handler.send_json_response({ - "tag": tag, - "total_count": 0, - "results": results - }) - - def api_search_by_date(self, handler, query_params): - """按日期范围搜索文件""" - start_date = query_params.get('start', [''])[0] - end_date = query_params.get('end', [''])[0] - search_type = query_params.get('type', ['all'])[0] - - if not start_date: - handler.send_json_response({"error": "Start date required"}, 400) - return - - try: - start_dt = datetime.fromisoformat(start_date) - end_dt = datetime.fromisoformat(end_date) if end_date else datetime.now() - except ValueError: - handler.send_json_response({"error": "Invalid date format"}, 400) - return - - results = [] - - for root, dirs, files in os.walk(self.config['base_dir']): - if search_type in ['all', 'dir']: - for dir_name in dirs: - full_path = os.path.join(root, dir_name) - try: - mtime = datetime.fromtimestamp(os.path.getmtime(full_path)) - if start_dt <= mtime <= end_dt: - rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") - results.append({ - "name": dir_name, - "path": rel_path + "/", - "type": "directory", - "modified": mtime.isoformat() - }) - except OSError: - continue - - if search_type in ['all', 'file']: - for file_name in files: - full_path = os.path.join(root, file_name) - try: - mtime = datetime.fromtimestamp(os.path.getmtime(full_path)) - if start_dt <= mtime <= end_dt: - rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") - results.append({ - "name": file_name, - "path": rel_path, - "type": "file", - "modified": mtime.isoformat() - }) - except OSError: - continue - - if len(results) > 1000: - break - - handler.send_json_response({ - "start_date": start_date, - "end_date": end_date.isoformat(), - "total_count": len(results), - "results": results - }) - - def _calculate_match_score(self, text, search_term, mode): - """计算匹配分数""" - if mode == 'exact': - return 100 if text.lower() == search_term else 0 - elif mode == 'regex': - try: - return 80 if re.search(search_term, text.lower()) else 0 - except: - return 0 - else: # fuzzy - text_lower = text.lower() - if search_term == text_lower: - return 100 - elif text_lower.startswith(search_term): - return 90 - elif text_lower.endswith(search_term): - return 80 - elif search_term in text_lower: - return 70 - return 0 - - # ==================== 增强统计API ==================== - - def api_get_stats_detailed(self, handler): - """获取详细统计信息""" - import mimetypes - - total_files = 0 - total_dirs = 0 - total_size = 0 - file_types = {} - size_distribution = { - "small": 0, # < 1MB - "medium": 0, # 1MB - 100MB - "large": 0, # 100MB - 1GB - "xlarge": 0 # > 1GB - } - oldest_file = None - newest_file = None - largest_file = None - - for root, dirs, files in os.walk(self.config['base_dir']): - total_dirs += len(dirs) - total_files += len(files) - for filename in files: - try: - file_path = os.path.join(root, filename) - size = os.path.getsize(file_path) - mtime = os.path.getmtime(file_path) - total_size += size - - mime_type, _ = mimetypes.guess_type(file_path) - if mime_type is None: - mime_type = "application/octet-stream" - file_types[mime_type] = file_types.get(mime_type, 0) + 1 - - # 大小分布 - if size < 1024 * 1024: - size_distribution["small"] += 1 - elif size < 100 * 1024 * 1024: - size_distribution["medium"] += 1 - elif size < 1024 * 1024 * 1024: - size_distribution["large"] += 1 - else: - size_distribution["xlarge"] += 1 - - # 最旧/最新文件 - if oldest_file is None or mtime < oldest_file[1]: - oldest_file = (filename, mtime) - if newest_file is None or mtime > newest_file[1]: - newest_file = (filename, mtime) - - # 最大文件 - if largest_file is None or size > largest_file[1]: - largest_file = (filename, size) - - except OSError: - continue - - handler.send_json_response({ - "summary": { - "total_files": total_files, - "total_dirs": total_dirs, - "total_size": total_size, - "total_size_formatted": self.format_file_size(total_size) - }, - "file_types": dict(sorted(file_types.items(), key=lambda x: x[1], reverse=True)), - "size_distribution": size_distribution, - "extremes": { - "oldest_file": { - "name": oldest_file[0] if oldest_file else None, - "modified": datetime.fromtimestamp(oldest_file[1]).isoformat() if oldest_file else None - }, - "newest_file": { - "name": newest_file[0] if newest_file else None, - "modified": datetime.fromtimestamp(newest_file[1]).isoformat() if newest_file else None - }, - "largest_file": { - "name": largest_file[0] if largest_file else None, - "size": largest_file[1] if largest_file else None, - "size_formatted": self.format_file_size(largest_file[1]) if largest_file else None - } - }, - "updated": datetime.now().isoformat() - }) - - def api_get_trending_files(self, handler, query_params): - """获取热门文件(按下载次数)""" - limit = int(query_params.get('limit', ['20'])[0]) - time_range = query_params.get('range', ['24h'])[0] # 24h, 7d, 30d - - # 获取下载统计 - stats = handler.load_stats() - - # 转换为列表并排序 - trending = [] - for filepath, count in stats.items(): - full_path = os.path.join(self.config['base_dir'], filepath) - if os.path.exists(full_path) and os.path.isfile(full_path): - try: - info = os.stat(full_path) - trending.append({ - "path": filepath, - "name": os.path.basename(filepath), - "size": info.st_size, - "size_formatted": self.format_file_size(info.st_size), - "modified": datetime.fromtimestamp(info.st_mtime).isoformat(), - "downloads": count - }) - except OSError: - continue - - # 按下载次数排序 - trending.sort(key=lambda x: x['downloads'], reverse=True) - trending = trending[:limit] - - handler.send_json_response({ - "time_range": time_range, - "limit": limit, - "trending": trending - }) - - def api_get_download_trend(self, handler, query_params): - """获取下载趋势(按天统计)""" - days = int(query_params.get('days', [7])[0]) - days = min(max(days, 1), 90) # 限制 1-90 天 - - trend = [] - - # 从数据库获取下载记录进行统计 - db = self.get_db() - if db: - try: - from datetime import timedelta - from core.database import DownloadRecord - - now = datetime.now() - start_time = (now - timedelta(days=days)).timestamp() - - with db.session() as session: - records = session.query(DownloadRecord).filter( - DownloadRecord.download_time >= start_time, - DownloadRecord.success == True - ).all() - - # 按日期聚合 - daily_stats = {} - for r in records: - if r.download_time: - date = datetime.fromtimestamp(r.download_time).strftime('%Y-%m-%d') - daily_stats[date] = daily_stats.get(date, 0) + 1 - - # 填充所有日期 - for i in range(days): - date = (now - timedelta(days=days - 1 - i)).strftime('%Y-%m-%d') - trend.append({ - "date": date, - "downloads": daily_stats.get(date, 0) - }) - except Exception as e: - print(f"Error getting download trend from database: {e}") - - # 如果没有数据库数据,返回模拟数据 - if not trend: - now = datetime.now() - for i in range(days): - date = (now - timedelta(days=days - 1 - i)).strftime('%m-%d') - trend.append({ - "date": date, - "downloads": 0 - }) - - handler.send_json_response({ - "days": days, - "trend": trend - }) - - def api_get_download_by_period(self, handler, query_params): - """按周期获取下载统计(年/月/日)""" - from datetime import datetime, timedelta - from sqlalchemy import func - - period = query_params.get('period', ['day'])[0] - year = int(query_params.get('year', [datetime.now().year])[0]) - month = int(query_params.get('month', [datetime.now().month])[0]) - - # 限制 period 值 - if period not in ['year', 'month', 'day']: - period = 'day' - - result = { - "period": period, - "year": year, - "month": month, - "data": [] - } - - db = self.get_db() - if db: - try: - from core.database import DownloadRecord - now = datetime.now() - - with db.session() as session: - if period == 'year': - # 按月统计全年数据 - start_time = datetime(year, 1, 1).timestamp() - end_time = datetime(year + 1, 1, 1).timestamp() - - records = session.query(DownloadRecord).filter( - DownloadRecord.download_time >= start_time, - DownloadRecord.download_time < end_time, - DownloadRecord.success == True - ).all() - - monthly_stats = {m: 0 for m in range(1, 13)} - total = 0 - for r in records: - if r.download_time: - dt = datetime.fromtimestamp(r.download_time) - monthly_stats[dt.month] += 1 - total += 1 - - result["data"] = [ - {"label": f"{m}月", "value": monthly_stats[m], "month": m} - for m in range(1, 13) - ] - result["total"] = total - - elif period == 'month': - # 按日统计当月数据 - start_time = datetime(year, month, 1).timestamp() - if month == 12: - end_time = datetime(year + 1, 1, 1).timestamp() - else: - end_time = datetime(year, month + 1, 1).timestamp() - - records = session.query(DownloadRecord).filter( - DownloadRecord.download_time >= start_time, - DownloadRecord.download_time < end_time, - DownloadRecord.success == True - ).all() - - days_in_month = (datetime(year, month + 1, 1) - timedelta(days=1)).day if month < 12 else 31 - daily_stats = {d: 0 for d in range(1, days_in_month + 1)} - total = 0 - for r in records: - if r.download_time: - dt = datetime.fromtimestamp(r.download_time) - daily_stats[dt.day] += 1 - total += 1 - - result["data"] = [ - {"label": f"{d}日", "value": daily_stats[d], "day": d} - for d in range(1, days_in_month + 1) - ] - result["total"] = total - - else: # day - 按日统计当月数据 - # 获取当月第一天和最后一天 - if month == 12: - start_date = datetime(year, month, 1) - end_date = datetime(year + 1, 1, 1) - else: - start_date = datetime(year, month, 1) - end_date = datetime(year, month + 1, 1) - - start_time = start_date.timestamp() - end_time = end_date.timestamp() - - records = session.query(DownloadRecord).filter( - DownloadRecord.download_time >= start_time, - DownloadRecord.download_time < end_time, - DownloadRecord.success == True - ).all() - - # 按天统计 - days_in_month = (end_date - timedelta(days=1)).day - daily_stats = {d: 0 for d in range(1, days_in_month + 1)} - total = 0 - for r in records: - if r.download_time: - dt = datetime.fromtimestamp(r.download_time) - daily_stats[dt.day] += 1 - total += 1 - - result["data"] = [ - {"label": f"{d}日", "value": daily_stats[d], "day": d} - for d in range(1, days_in_month + 1) - ] - result["total"] = total - - # 获取历史年份列表 - if period == 'year': - with db.session() as session: - # SQLite 不支持 from_unixtime,使用 Python 处理 - records = session.query(DownloadRecord.download_time).filter( - DownloadRecord.success == True - ).distinct().all() - years_set = set() - for (ts,) in records: - if ts: - dt = datetime.fromtimestamp(ts) - years_set.add(dt.year) - result["available_years"] = sorted(list(years_set), reverse=True)[:10] - - except Exception as e: - print(f"Error getting download by period: {e}") - result["error"] = str(e) - - # 如果没有数据,返回空数据结构 - if not result.get("data"): - if period == 'year': - result["data"] = [{"label": f"{m}月", "value": 0, "month": m} for m in range(1, 13)] - elif period == 'month': - result["data"] = [{"label": f"{d}日", "value": 0, "day": d} for d in range(1, 32)] - else: - result["data"] = [{"label": f"{h}:00", "value": 0, "hour": h} for h in range(24)] - - handler.send_json_response(result) - - def api_get_download_rank(self, handler, query_params): - """获取下载排行 TOP 20""" - import traceback - import os - limit = int(query_params.get('limit', [20])[0]) - limit = min(max(limit, 1), 100) - - rank_data = [] - errors = [] - - # 从数据库获取下载记录进行统计 - db = self.get_db() - if db: - try: - from core.database import DownloadRecord - - with db.session() as session: - # 按文件路径分组统计下载次数 - from sqlalchemy import func - results = session.query( - DownloadRecord.file_path, - func.count(DownloadRecord.id).label('download_count') - ).filter( - DownloadRecord.success == True - ).group_by( - DownloadRecord.file_path - ).order_by( - func.count(DownloadRecord.id).desc() - ).limit(limit).all() - - total_downloads = sum(r[1] for r in results) if results else 0 - - # 获取 base_dir - base_dir = self.config.get('base_dir', './downloads') - base_dir = os.path.abspath(base_dir) if base_dir else './downloads' - - for idx, (file_path, count) in enumerate(results, 1): - # 获取文件大小 - file_size = 0 - full_path = os.path.join(base_dir, file_path) if file_path else '' - if full_path and os.path.exists(full_path): - try: - file_size = os.path.getsize(full_path) - except OSError: - pass - - # 计算占比 - percentage = (count / total_downloads * 100) if total_downloads > 0 else 0 - - rank_data.append({ - "rank": idx, - "file": file_path or 'Unknown', - "downloads": count, - "size": file_size, - "percentage": round(percentage, 1) - }) - except Exception as e: - errors.append(str(e)) - traceback.print_exc() - - handler.send_json_response({ - "rank": rank_data, - "count": len(rank_data), - "total_records": len(rank_data), - "_debug": {"errors": errors} if errors else {} - }) - - def api_get_hot_cache(self, handler, query_params): - """获取热门缓存文件""" - import traceback - import os - limit = int(query_params.get('limit', [20])[0]) - limit = min(max(limit, 1), 100) - - cache_data = [] - errors = [] - - # 从数据库获取缓存访问记录 - db = self.get_db() - if db: - try: - from core.database import DownloadRecord - - with db.session() as session: - from sqlalchemy import func - # 获取被访问过的缓存文件(按访问次数排序) - results = session.query( - DownloadRecord.file_path, - func.count(DownloadRecord.id).label('access_count'), - func.max(DownloadRecord.download_time).label('last_access') - ).filter( - DownloadRecord.success == True - ).group_by( - DownloadRecord.file_path - ).order_by( - func.count(DownloadRecord.id).desc() - ).limit(limit).all() - - # 获取 base_dir - base_dir = self.config.get('base_dir', './downloads') - base_dir = os.path.abspath(base_dir) if base_dir else './downloads' - - for idx, (file_path, access_count, last_access) in enumerate(results, 1): - # 获取文件信息 - full_path = os.path.join(base_dir, file_path) if file_path else '' - file_size = 0 - if full_path and os.path.exists(full_path): - try: - file_size = os.path.getsize(full_path) - except OSError: - pass - - cache_data.append({ - "rank": idx, - "file": file_path or 'Unknown', - "access_count": access_count, - "size": file_size, - "last_access": datetime.fromtimestamp(last_access).isoformat() if last_access else '-' - }) - except Exception as e: - errors.append(str(e)) - traceback.print_exc() - - handler.send_json_response({ - "cache": cache_data, - "count": len(cache_data), - "_debug": {"errors": errors} if errors else {} - }) - - def get_db(self): - """获取数据库实例""" - from core.database import get_db - try: - return get_db() - except Exception: - return None - - # ==================== 文件元数据API ==================== - - def api_get_file_metadata(self, handler, filename): - """获取文件元数据""" - full_path = os.path.join(self.config['base_dir'], filename) - if not os.path.exists(full_path): - handler.send_json_response({"error": "File not found"}, 404) - return - - # 这里可以从单独的元数据文件中读取 - metadata_file = full_path + '.meta' - metadata = {} - - if os.path.exists(metadata_file): - try: - with open(metadata_file, 'r', encoding='utf-8') as f: - metadata = json.load(f) - except: - pass - - # 添加基本文件信息 - import os - stat = os.stat(full_path) - metadata['_file_info'] = { - "size": stat.st_size, - "modified": datetime.fromtimestamp(stat.st_mtime).isoformat(), - "created": datetime.fromtimestamp(stat.st_ctime).isoformat() - } - - handler.send_json_response(metadata) - - def api_update_file_metadata(self, handler, filename): - """更新文件元数据""" - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "No metadata provided"}, 400) - return - - try: - metadata = json.loads(handler.rfile.read(content_length)) - full_path = os.path.join(self.config['base_dir'], filename) - metadata_file = full_path + '.meta' - - with open(metadata_file, 'w', encoding='utf-8') as f: - json.dump(metadata, f, ensure_ascii=False, indent=2) - - handler.send_json_response({"success": True}) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - def api_get_batch_metadata(self, handler, query_params): - """批量获取文件元数据""" - paths = query_params.get('paths', []) - if not paths: - handler.send_json_response({"error": "No file paths provided"}, 400) - return - - results = {} - for path in paths: - full_path = os.path.join(self.config['base_dir'], path) - if os.path.exists(full_path): - metadata_file = full_path + '.meta' - metadata = {} - if os.path.exists(metadata_file): - try: - with open(metadata_file, 'r', encoding='utf-8') as f: - metadata = json.load(f) - except: - pass - results[path] = metadata - - handler.send_json_response(results) - - def api_update_batch_metadata(self, handler): - """批量更新文件元数据""" - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "No data provided"}, 400) - return - - try: - data = json.loads(handler.rfile.read(content_length)) - results = {} - - for path, metadata in data.items(): - full_path = os.path.join(self.config['base_dir'], path) - metadata_file = full_path + '.meta' - try: - with open(metadata_file, 'w', encoding='utf-8') as f: - json.dump(metadata, f, ensure_ascii=False, indent=2) - results[path] = {"success": True} - except Exception as e: - results[path] = {"success": False, "error": str(e)} - - handler.send_json_response(results) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - # ==================== 文件版本控制API ==================== - - def api_get_file_versions(self, handler, filename): - """获取文件版本列表""" - # 这里可以实现版本控制系统 - # 简化实现:检查备份文件 - full_path = os.path.join(self.config['base_dir'], filename) - versions = [] - - # 查找备份文件 - dir_path = os.path.dirname(full_path) - base_name = os.path.basename(full_path) - - if os.path.exists(dir_path): - for item in os.listdir(dir_path): - if item.startswith(base_name) and item != base_name: - backup_path = os.path.join(dir_path, item) - try: - stat = os.stat(backup_path) - versions.append({ - "name": item, - "size": stat.st_size, - "modified": datetime.fromtimestamp(stat.st_mtime).isoformat() - }) - except OSError: - continue - - handler.send_json_response({ - "filename": filename, - "versions": versions - }) - - def api_create_file_version(self, handler, filename): - """创建文件版本(备份)""" - import time - full_path = os.path.join(self.config['base_dir'], filename) - if not os.path.exists(full_path): - handler.send_json_response({"error": "File not found"}, 404) - return - - # 创建备份 - timestamp = int(time.time()) - backup_path = f"{full_path}.v{timestamp}" - - try: - import shutil - shutil.copy2(full_path, backup_path) - handler.send_json_response({ - "success": True, - "version": backup_path - }) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - # ==================== 缩略图API ==================== - - def api_get_file_thumbnail(self, handler, filename, query_params): - """获取文件缩略图""" - import mimetypes - - full_path = os.path.join(self.config['base_dir'], filename) - if not os.path.exists(full_path): - handler.send_json_response({"error": "File not found"}, 404) - return - - mime_type, _ = mimetypes.guess_type(full_path) - if not mime_type or not mime_type.startswith('image/'): - handler.send_json_response({"error": "Not an image file"}, 400) - return - - width = int(query_params.get('width', ['200'])[0]) - height = int(query_params.get('height', ['200'])[0]) - - try: - from PIL import Image - - with Image.open(full_path) as img: - img.thumbnail((width, height), Image.LANCZOS) - - import io - buffer = io.BytesIO() - img.save(buffer, format='JPEG', quality=85) - thumbnail_data = buffer.getvalue() - - handler.send_response(200) - handler.send_header("Content-Type", "image/jpeg") - handler.send_header("Content-Length", str(len(thumbnail_data))) - handler.send_header("Cache-Control", "public, max-age=86400") - handler.end_headers() - handler.wfile.write(thumbnail_data) - - except ImportError: - handler.send_json_response({"error": "PIL/Pillow not installed"}, 500) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - # ==================== 服务器监控API ==================== - - def api_get_realtime_stats(self, handler): - """获取实时服务器统计""" - try: - import psutil - result = { - "timestamp": datetime.now().isoformat(), - "cpu": { - "percent": psutil.cpu_percent(interval=0.1), - "count": psutil.cpu_count() - }, - "memory": { - "total": psutil.virtual_memory().total, - "available": psutil.virtual_memory().available, - "percent": psutil.virtual_memory().percent, - "used": psutil.virtual_memory().used, - "free": psutil.virtual_memory().free - }, - "disk": { - "total": psutil.disk_usage(self.config['base_dir']).total, - "used": psutil.disk_usage(self.config['base_dir']).used, - "free": psutil.disk_usage(self.config['base_dir']).free, - "percent": psutil.disk_usage(self.config['base_dir']).percent - } - } - # 尝试获取网络数据,失败时忽略 - try: - result["network"] = { - "connections": len(psutil.net_connections()), - "io": psutil.net_io_counters()._asdict() if psutil.net_io_counters() else None - } - except (PermissionError, OSError): - result["network"] = {"connections": 0, "io": None, "note": "Permission denied"} - - # 尝试获取 CPU 频率,失败时忽略 - try: - result["cpu"]["freq"] = psutil.cpu_freq()._asdict() if psutil.cpu_freq() else None - except (PermissionError, OSError): - result["cpu"]["freq"] = None - - handler.send_json_response(result) - except ImportError: - handler.send_json_response({"error": "psutil not installed"}, 500) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - def api_get_monitor_history(self, handler, query_params): - """获取历史监控数据""" - hours = int(query_params.get('hours', ['24'])[0]) - - # 尝试从数据库获取 - if self.db_enabled and self.db: - try: - records = self.db.get_monitor_history(hours) - data = [r.to_dict() for r in records] - handler.send_json_response({ - "hours": hours, - "data": data, - "source": "database" - }) - return - except Exception as e: - print(f"Error getting monitor history from database: {e}") - - # 回退到空数据 - handler.send_json_response({ - "hours": hours, - "data": [], - "source": "none" - }) - - # ==================== Webhook API ==================== - - def api_list_webhooks(self, handler): - """列出所有webhook""" - # 从数据库获取 - if self.db_enabled and self.db: - try: - webhooks = self.db.get_webhooks() - handler.send_json_response({ - "webhooks": [w.to_dict() for w in webhooks], - "count": len(webhooks), - "source": "database" - }) - return - except Exception as e: - print(f"Error getting webhooks from database: {e}") - - # 回退到空列表 - handler.send_json_response({ - "webhooks": [], - "count": 0, - "source": "none" - }) - - def api_create_webhook(self, handler): - """创建webhook""" - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "No webhook data provided"}, 400) - return - - try: - webhook_data = json.loads(handler.rfile.read(content_length)) - - # 验证必要字段 - name = webhook_data.get('name', '').strip() - url = webhook_data.get('url', '').strip() - - if not name: - handler.send_json_response({"error": "Webhook name is required"}, 400) - return - if not url: - handler.send_json_response({"error": "Webhook URL is required"}, 400) - return - - # 保存到数据库 - if self.db_enabled and self.db: - record = self.db.add_webhook( - name=name, - url=url, - events=webhook_data.get('events', []), - secret=webhook_data.get('secret'), - enabled=webhook_data.get('enabled', True) - ) - handler.send_json_response({ - "success": True, - "message": "Webhook created successfully", - "webhook": record.to_dict() - }) - else: - handler.send_json_response({ - "success": False, - "error": "Database not available" - }, 503) - - except json.JSONDecodeError: - handler.send_json_response({"error": "Invalid JSON format"}, 400) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - def api_get_webhook(self, handler, webhook_id): - """获取webhook详情""" - webhook_id = int(webhook_id) - - if self.db_enabled and self.db: - try: - webhook = self.db.get_webhook(webhook_id) - if webhook: - handler.send_json_response({ - "webhook": webhook.to_dict() - }) - else: - handler.send_json_response({ - "error": "Webhook not found" - }, 404) - return - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - handler.send_json_response({ - "error": "Database not available" - }, 503) - - def api_delete_webhook(self, handler, webhook_id): - """删除webhook""" - webhook_id = int(webhook_id) - - if self.db_enabled and self.db: - try: - success = self.db.delete_webhook(webhook_id) - if success: - handler.send_json_response({ - "success": True, - "message": "Webhook deleted successfully" - }) - else: - handler.send_json_response({ - "error": "Webhook not found" - }, 404) - return - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - handler.send_json_response({ - "success": False, - "error": "Database not available" - }, 503) - - def api_test_webhook(self, handler, webhook_id): - """测试webhook""" - webhook_id = int(webhook_id) - - if self.db_enabled and self.db: - try: - webhook = self.db.get_webhook(webhook_id) - if not webhook: - handler.send_json_response({ - "error": "Webhook not found" - }, 404) - return - - # 发送测试请求 - import urllib.request - import urllib.parse - - test_payload = { - "event": "test", - "timestamp": datetime.now().isoformat(), - "data": { - "message": "This is a test webhook from HYC下载站" - } - } - - try: - data = json.dumps(test_payload).encode('utf-8') - req = urllib.request.Request( - webhook.url, - data=data, - headers={ - 'Content-Type': 'application/json', - 'X-Webhook-Secret': webhook.secret or '' - }, - method='POST' - ) - with urllib.request.urlopen(req, timeout=10) as response: - handler.send_json_response({ - "success": True, - "webhook_id": webhook_id, - "test_result": "Webhook test successful", - "status_code": response.status - }) - except urllib.error.HTTPError as e: - handler.send_json_response({ - "success": False, - "webhook_id": webhook_id, - "error": f"HTTP Error: {e.code} {e.reason}" - }, 400) - except Exception as e: - handler.send_json_response({ - "success": False, - "webhook_id": webhook_id, - "error": str(e) - }, 400) - - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - else: - handler.send_json_response({ - "success": False, - "error": "Database not available" - }, 503) - - def api_get_webhook_deliveries(self, handler, webhook_id): - """获取 webhook 交付历史""" - webhook_id = int(webhook_id) - - # 验证 webhook 存在 - if self.db_enabled and self.db: - try: - webhook = self.db.get_webhook(webhook_id) - if not webhook: - handler.send_json_response({ - "error": "Webhook not found" - }, 404) - return - - # 获取交付历史 - deliveries = self.db.get_webhook_deliveries(webhook_id=webhook_id, limit=50) - - handler.send_json_response({ - "webhook_id": webhook_id, - "webhook_name": webhook.name, - "deliveries": [d.to_dict() for d in deliveries], - "count": len(deliveries) - }) - - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - else: - handler.send_json_response({ - "error": "Database not available" - }, 503) - - def api_get_webhook_stats(self, handler, webhook_id): - """获取 webhook 交付统计""" - webhook_id = int(webhook_id) - - if self.db_enabled and self.db: - try: - stats = self.db.get_webhook_stats(webhook_id) - - handler.send_json_response({ - "webhook_id": webhook_id, - "stats": stats - }) - - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - else: - handler.send_json_response({ - "error": "Database not available" - }, 503) - - def api_update_webhook(self, handler, webhook_id): - """更新 webhook 配置""" - webhook_id = int(webhook_id) - - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "No webhook data provided"}, 400) - return - - try: - webhook_data = json.loads(handler.rfile.read(content_length)) - - if self.db_enabled and self.db: - webhook = self.db.get_webhook(webhook_id) - if not webhook: - handler.send_json_response({ - "error": "Webhook not found" - }, 404) - return - - # 更新 webhook - updated = self.db.update_webhook( - webhook_id, - name=webhook_data.get('name', webhook.name), - url=webhook_data.get('url', webhook.url), - events=webhook_data.get('events'), - secret=webhook_data.get('secret'), - enabled=webhook_data.get('enabled', webhook.enabled) - ) - - if updated: - handler.send_json_response({ - "success": True, - "message": "Webhook updated successfully", - "webhook": updated.to_dict() - }) - else: - handler.send_json_response({ - "success": False, - "error": "Failed to update webhook" - }, 500) - else: - handler.send_json_response({ - "success": False, - "error": "Database not available" - }, 503) - - except json.JSONDecodeError: - handler.send_json_response({"error": "Invalid JSON format"}, 400) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - # ==================== 同步管理 API (真实数据) ==================== - - def api_get_sync_sources(self, handler): - """获取所有同步源(真实数据)""" - # 从 sync_manager 获取真实数据 - if hasattr(handler, 'sync_manager') and handler.sync_manager: - sources = getattr(handler.sync_manager, 'sync_sources', {}) - # 转换为数组格式 - sources_list = [] - for name, config in sources.items(): - item = {"name": name} - item.update(config) - sources_list.append(item) - handler.send_json_response({ - "sources": sources_list, - "count": len(sources_list) - }) - else: - handler.send_json_response({ - "sources": [], - "count": 0, - "message": "Sync manager not available" - }) - - def api_add_sync_source(self, handler): - """添加同步源""" - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "No data provided"}, 400) - return - - try: - data = json.loads(handler.rfile.read(content_length)) - - if hasattr(handler, 'sync_manager') and handler.sync_manager: - name = data.get('name') - config = data.get('config', {}) - - if not name: - handler.send_json_response({"error": "Source name required"}, 400) - return - - success = handler.sync_manager.add_source(name, config) - handler.send_json_response({ - "success": success, - "name": name - }) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - def api_start_sync(self, handler, source_name): - """启动同步""" - if hasattr(handler, 'sync_manager') and handler.sync_manager: - task_id = handler.sync_manager.start_sync(source_name) - if task_id: - handler.send_json_response({ - "success": True, - "task_id": task_id, - "source_name": source_name - }) - else: - handler.send_json_response({ - "success": False, - "error": "Failed to start sync", - "source_name": source_name - }, 500) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - def api_stop_sync(self, handler, source_name): - """停止同步""" - if hasattr(handler, 'sync_manager') and handler.sync_manager: - handler.sync_manager.stop_all_tasks_for_source(source_name) - handler.send_json_response({ - "success": True, - "source_name": source_name - }) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - def api_get_sync_status(self, handler, source_name): - """获取同步状态(真实数据)""" - # source_name 格式: source_name/status,需要提取 - if source_name.endswith('/status'): - source_name = source_name[:-7] - - if hasattr(handler, 'sync_manager') and handler.sync_manager: - status = handler.sync_manager.get_source_status(source_name) - handler.send_json_response(status) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - def api_get_sync_history(self, handler, query_params): - """获取同步历史""" - limit = int(query_params.get('limit', ['100'])[0]) - - if hasattr(handler, 'sync_manager') and handler.sync_manager: - history = handler.sync_manager.get_sync_history(limit=limit) - handler.send_json_response({ - "history": history, - "count": len(history) - }) - else: - handler.send_json_response({ - "history": [], - "count": 0 - }) - - # ==================== 定时任务调度 API ==================== - - def api_get_scheduled_tasks(self, handler): - """获取所有定时任务状态""" - if hasattr(handler, 'sync_manager') and handler.sync_manager: - scheduler = getattr(handler.sync_manager, 'task_scheduler', None) - if scheduler and hasattr(scheduler, 'get_all_tasks'): - tasks = scheduler.get_all_tasks() - handler.send_json_response({ - "tasks": tasks, - "count": len(tasks) - }) - return - - handler.send_json_response({ - "tasks": [], - "count": 0, - "message": "Scheduler not available" - }) - - def api_update_sync_schedule(self, handler, source_name): - """更新同步源的定时配置""" - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "No schedule data provided"}, 400) - return - - try: - data = json.loads(handler.rfile.read(content_length)) - - if hasattr(handler, 'sync_manager') and handler.sync_manager: - # 更新同步源的定时配置 - sync_manager = handler.sync_manager - - if hasattr(sync_manager, 'scheduled_syncs'): - schedule_config = data.get('schedule', {}) - if source_name not in sync_manager.scheduled_syncs: - sync_manager.scheduled_syncs[source_name] = {} - - sync_manager.scheduled_syncs[source_name] = { - 'type': schedule_config.get('type', 'interval'), - 'config': { - 'cron': schedule_config.get('cron'), - 'interval': schedule_config.get('interval', {}), - 'enabled': schedule_config.get('enabled', True) - } - } - - # 如果调度器已运行,更新任务 - if sync_manager.task_scheduler: - task_name = f"sync_{source_name}" - existing_task = sync_manager.task_scheduler.get_task(task_name) - if existing_task: - sync_manager.task_scheduler.update_task_config( - task_name, - sync_manager.scheduled_syncs[source_name]['config'] - ) - if schedule_config.get('enabled'): - sync_manager.task_scheduler.enable_task(task_name, True) - else: - sync_manager.task_scheduler.enable_task(task_name, False) - - handler.send_json_response({ - "success": True, - "source_name": source_name, - "schedule": sync_manager.scheduled_syncs.get(source_name, {}) - }) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - except json.JSONDecodeError: - handler.send_json_response({"error": "Invalid JSON format"}, 400) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - def api_run_sync_now(self, handler, source_name): - """立即触发同步(覆盖定时)""" - if hasattr(handler, 'sync_manager') and handler.sync_manager: - success = handler.sync_manager.start_sync(source_name) - handler.send_json_response({ - "success": success, - "source_name": source_name, - "message": "Sync started" if success else "Failed to start sync" - }) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - def api_sync_packages(self, handler): - """临时单次同步指定源的特定包""" - try: - # 读取请求体 - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "请求体不能为空"}, 400) - return - - body = handler.rfile.read(content_length) - data = json.loads(body.decode('utf-8')) - - # 获取参数 - source = data.get('source') - packages = data.get('packages', []) - - if not source: - handler.send_json_response({"error": "缺少 'source' 参数"}, 400) - return - - if not packages or not isinstance(packages, list): - handler.send_json_response({"error": "请提供有效的 'packages' 列表"}, 400) - return - - # 调用 sync_manager - if hasattr(handler, 'sync_manager') and handler.sync_manager: - result = handler.sync_manager.sync_packages(source, packages) - handler.send_json_response(result) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - except json.JSONDecodeError: - handler.send_json_response({"error": "Invalid JSON format"}, 400) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - def api_get_temp_sync_status(self, handler, source_name): - """获取临时同步状态""" - if hasattr(handler, 'sync_manager') and handler.sync_manager: - sync_manager = handler.sync_manager - # 查找匹配的临时同步任务 - temp_status = None - if hasattr(sync_manager, 'sync_status'): - for task_name, status in sync_manager.sync_status.items(): - if status.get('source_name') == source_name and status.get('is_temp_sync'): - temp_status = status - break - - if temp_status: - handler.send_json_response({ - "success": True, - "source": source_name, - "status": temp_status.get('status'), - "packages": temp_status.get('packages', []), - "files_synced": temp_status.get('files_synced', 0), - "total_files": temp_status.get('total_files', 0), - "last_sync": temp_status.get('last_sync'), - "error": temp_status.get('error') - }) - else: - handler.send_json_response({ - "success": False, - "error": f"没有找到 {source_name} 的临时同步任务" - }) - else: - handler.send_json_response({"error": "Sync manager not available"}, 500) - - # ==================== 系统监控 API (真实数据) ==================== - - def api_get_monitor_detailed(self, handler): - """获取详细监控数据""" - if hasattr(handler, 'monitor') and handler.monitor: - stats = handler.monitor.get_realtime_stats() - handler.send_json_response(stats) - else: - handler.send_json_response({ - "error": "Monitor not available", - "timestamp": datetime.now().isoformat() - }, 500) - - def api_get_monitor_history_detailed(self, handler, query_params): - """获取历史监控数据(详细版)""" - hours = int(query_params.get('hours', ['24'])[0]) - - # 尝试从数据库获取详细统计 - if self.db_enabled and self.db: - try: - stats = self.db.get_monitor_stats(hours) - history = self.db.get_monitor_history(hours) - handler.send_json_response({ - "hours": hours, - "stats": stats, - "history": [r.to_dict() for r in history], - "source": "database" - }) - return - except Exception as e: - print(f"Error getting monitor stats from database: {e}") - - # 回退到原有实现 - if hasattr(handler, 'monitor') and handler.monitor: - history = handler.monitor.get_monitor_history(hours) - handler.send_json_response({ - "hours": hours, - "data": history, - "source": "monitor" - }) - else: - handler.send_json_response({ - "hours": hours, - "data": [], - "source": "none" - }) - - def api_get_monitor_summary(self, handler): - """获取监控摘要""" - if hasattr(handler, 'monitor') and handler.monitor: - summary = handler.monitor.get_stats_summary() - handler.send_json_response(summary) - else: - handler.send_json_response({ - "status": "unavailable", - "message": "Monitor not available" - }, 500) - - def api_get_health_status(self, handler): - """获取健康状态""" - if hasattr(handler, 'monitor') and handler.monitor: - health = handler.monitor.get_health_status() - handler.send_json_response(health) - else: - handler.send_json_response({ - "status": "unknown", - "message": "Monitor not available" - }) - - # ==================== 镜像源健康检查 API ==================== - - def api_get_source_health(self, handler, query_params): - """获取镜像源健康状态""" - try: - from core.health_check import HealthChecker, HealthStatus - - mirrors = self.config.get('mirrors', {}) - checker = HealthChecker(self.config.get('health_check', {})) - - results = [] - for mirror_type, mirror_config in mirrors.items(): - if not isinstance(mirror_config, dict): - continue - - # 获取该镜像类型的所有源 - sources = mirror_config.get('sources', []) - for source_name in sources: - result = checker.check_source(source_name, { - 'url': self._get_source_url(mirror_type, source_name) - }) - results.append({ - 'mirror_type': mirror_type, - 'source_name': source_name, - 'status': result.status.value, - 'response_time_ms': round(result.response_time, 2), - 'http_status': result.http_status, - 'error': result.error_message, - 'success_rate': round(result.success_rate, 2), - 'last_check': result.last_check.isoformat() if result.last_check else None, - 'consecutive_failures': result.consecutive_failures - }) - - handler.send_json_response({ - 'sources': results, - 'count': len(results), - 'summary': checker.get_stats() - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def api_check_source(self, handler, source_name): - """手动触发单个源的健康检查""" - try: - from core.health_check import HealthChecker, HealthStatus - - checker = HealthChecker(self.config.get('health_check', {})) - mirrors = self.config.get('mirrors', {}) - - # 查找源对应的镜像类型 - mirror_type = None - source_url = '' - for mtype, mconfig in mirrors.items(): - if not isinstance(mconfig, dict): - continue - sources = mconfig.get('sources', []) - if source_name in sources: - mirror_type = mtype - source_config = mconfig.get('sources_config', {}).get(source_name, {}) - source_url = source_config.get('url', '') - break - - if not mirror_type: - handler.send_json_response({ - 'error': f"Source '{source_name}' not found" - }, 404) - return - - result = checker.check_source(source_name, {'url': source_url}) - - handler.send_json_response({ - 'source_name': source_name, - 'mirror_type': mirror_type, - 'status': result.status.value, - 'response_time_ms': round(result.response_time, 2), - 'http_status': result.http_status, - 'error': result.error_message, - 'success_rate': round(result.success_rate, 2), - 'last_check': result.last_check.isoformat() if result.last_check else None - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def api_get_failover_status(self, handler): - """获取故障切换状态""" - try: - from core.health_check import MirrorFailoverManager - - failover = MirrorFailoverManager(self.config) - failover.initialize() - - handler.send_json_response({ - 'failover_enabled': failover.failover_enabled, - 'active_sources': failover._active_source, - 'health_summary': failover.get_health_summary(), - 'failover_history': failover.get_failover_history() - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def api_trigger_failover(self, handler, mirror_type): - """手动触发故障切换""" - try: - from core.health_check import MirrorFailoverManager - - failover = MirrorFailoverManager(self.config) - failover.initialize() - - success = failover.perform_failover(mirror_type) - - handler.send_json_response({ - 'mirror_type': mirror_type, - 'success': success, - 'active_source': failover.get_active_source(mirror_type), - 'failover_history': failover.get_failover_history() - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def _get_source_url(self, mirror_type: str, source_name: str) -> str: - """获取源的 URL""" - mirrors = self.config.get('mirrors', {}) - mirror_config = mirrors.get(mirror_type, {}) - - # 先检查 sources_config - sources_config = mirror_config.get('sources_config', {}) - if source_name in sources_config: - return sources_config[source_name].get('url', '') - - # 使用 URL 模板 - url_template = mirror_config.get('url_template', '') - if url_template and '{mirror}' in url_template: - return url_template.replace('{mirror}', source_name) - - return '' - - # ==================== 缓存管理 API ==================== - - def api_get_cache_stats(self, handler): - """获取缓存统计""" - if hasattr(handler, 'cache_manager') and handler.cache_manager: - stats = handler.cache_manager.get_stats() - # 转换字段名以匹配前端期望 - handler.send_json_response({ - "size": stats.get('total_size', 0), - "count": stats.get('file_count', 0), - "hit_rate": stats.get('hit_rate', 0), - "last_clean": stats.get('last_clean', None), - "strategy": stats.get('strategy', 'unknown') - }) - else: - handler.send_json_response({ - "size": 0, - "count": 0, - "hit_rate": 0, - "last_clean": None, - "strategy": "unknown" - }) - - def api_clean_cache(self, handler): - """清理缓存""" - source = None # 可以从请求中获取 - - if hasattr(handler, 'cache_manager') and handler.cache_manager: - count = handler.cache_manager.clear(source) - handler.send_json_response({ - "success": True, - "deleted_count": count - }) - else: - handler.send_json_response({"error": "Cache manager not available"}, 500) - - def api_get_cache_usage(self, handler): - """获取缓存使用详情""" - if hasattr(handler, 'cache_manager') and handler.cache_manager: - usage = handler.cache_manager.get_cache_usage() - handler.send_json_response({ - "items": usage, - "count": len(usage) - }) - else: - handler.send_json_response({ - "items": [], - "count": 0 - }) - - def api_get_recent_activity(self, handler, query_params): - """获取最近活动""" - import traceback - limit = int(query_params.get('limit', [20])[0]) - limit = min(max(limit, 1), 100) # 限制在 1-100 之间 - offset = int(query_params.get('offset', [0])[0]) - offset = max(offset, 0) # 确保 offset 不为负数 - - activities = [] - errors = [] - all_activities = [] # 收集所有活动用于统一排序 - - # 从同步记录获取 - if hasattr(handler, 'sync_manager') and handler.sync_manager: - try: - if hasattr(handler.sync_manager, 'db_enabled') and handler.sync_manager.db_enabled: - from core.database import get_db - db = get_db() - if db: - with db.session() as session: - from core.database import SyncRecord - # 获取足够多的记录用于分页 - fetch_limit = limit + offset - records = session.query(SyncRecord).order_by( - SyncRecord.start_time.desc() - ).limit(fetch_limit).all() - for r in records: - all_activities.append({ - "time": datetime.fromtimestamp(r.start_time).isoformat() if r.start_time else '', - "timestamp": r.start_time or 0, - "type": "同步", - "content": f"同步任务: {r.source_name or r.sync_id}", - "status": "成功" if r.status == "completed" else ("进行中" if r.status == "running" else "失败"), - "status_type": "success" if r.status == "completed" else ("running" if r.status == "running" else "error") - }) - else: - errors.append("同步记录: db 为空") - except Exception as e: - errors.append(f"同步记录: {str(e)}") - traceback.print_exc() - - # 从下载记录获取 - try: - # 优先使用 handler.db - db = getattr(handler, 'db', None) - if db is None: - from core.database import get_db - db = get_db() - - if db: - from core.database import DownloadRecord - with db.session() as session: - # 获取足够多的记录用于分页 - fetch_limit = limit + offset - records = session.query(DownloadRecord).order_by( - DownloadRecord.download_time.desc() - ).limit(fetch_limit).all() - for r in records: - all_activities.append({ - "time": datetime.fromtimestamp(r.download_time).isoformat() if r.download_time else '', - "timestamp": r.download_time or 0, - "type": "下载", - "content": f"下载: {r.file_path or 'Unknown'}", - "status": "成功" if r.success else "失败", - "status_type": "success" if r.success else "error" - }) - else: - errors.append("下载记录: db 为空") - except Exception as e: - errors.append(f"下载记录: {str(e)}") - traceback.print_exc() - - # 从告警记录获取 - try: - from core.alerts import AlertManager - alert_manager = AlertManager(self.config.get('alerts', {})) - fetch_limit = limit + offset - alerts = alert_manager.get_alerts(limit=fetch_limit) - for a in alerts: - all_activities.append({ - "time": a.get('timestamp', ''), - "timestamp": 0, - "type": "告警", - "content": a.get('message', ''), - "status": a.get('severity', 'info').upper(), - "status_type": "error" if a.get('severity') == 'error' else ("warning" if a.get('severity') == 'warning' else "info") - }) - except Exception as e: - errors.append(f"告警记录: {str(e)}") - traceback.print_exc() - - # 按时间戳排序并应用分页 - all_activities.sort(key=lambda x: x.get('timestamp', 0), reverse=True) - - # 应用 offset 和 limit - paged_activities = all_activities[offset:offset + limit] - - handler.send_json_response({ - "activities": paged_activities, - "count": len(paged_activities), - "total": len(all_activities), - "_debug": {"errors": errors} if errors else {} - }) - - # ==================== 镜像加速源 API ==================== - - def api_list_mirrors(self, handler): - """列出所有镜像加速源(仅从配置文件读取,无内置预设)""" - # 仅从配置文件读取镜像配置 - mirrors_config = {} - try: - import json - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - # 优先使用 settings.json - settings_path = os.path.join(project_root, 'settings.json') - config_path = settings_path - if not os.path.exists(config_path): - # 回退到 config.json - config_path = os.path.join(project_root, 'config.json') - - if os.path.exists(config_path): - with open(config_path, 'r', encoding='utf-8') as f: - config_data = json.load(f) - mirrors_config = config_data.get('mirrors', {}) - except Exception: - pass - - # 如果文件没有配置,回退到内存配置 - if not mirrors_config: - mirrors_config = self.config.get('mirrors', {}) - - # 收集所有镜像(从 mirrors 配置读取) - available_mirrors = {} - - # 从 mirrors 配置读取所有镜像 - for mirror_name, mirror_config in mirrors_config.items(): - available_mirrors[mirror_name] = { - "name": mirror_config.get('name', mirror_name), - "type": mirror_config.get('type', 'http'), - "description": mirror_config.get('description', f"{mirror_name} 镜像"), - "url": mirror_config.get('url', ''), - "target": mirror_config.get('target', ''), - "enabled": mirror_config.get('enabled', True), - "custom": mirror_config.get('custom', True), - "auto_sync": mirror_config.get('auto_sync', False), - "schedule": mirror_config.get('schedule', {}), - "last_sync": mirror_config.get('last_sync'), - "storage_dir": mirror_config.get('storage_dir', mirror_name) - } - - handler.send_json_response({ - "mirrors": available_mirrors, - "count": len(available_mirrors) - }) - - def api_get_mirror_info(self, handler, mirror_name): - """获取镜像加速源信息""" - from mirrors import get_mirror_handler - - handler_class = get_mirror_handler(mirror_name) - if not handler_class: - handler.send_json_response({ - "error": f"Unknown mirror type: {mirror_name}" - }, 400) - return - - # 获取镜像配置 - mirror_config = self.config.get('mirrors', {}).get(mirror_name, {}) - - handler.send_json_response({ - "name": mirror_name, - "enabled": mirror_config.get('enabled', False), - "config": mirror_config - }) - - def api_refresh_mirror(self, handler, mirror_name): - """刷新镜像元数据""" - handler.send_json_response({ - "success": True, - "message": f"Mirror {mirror_name} refresh initiated" - }) - - def api_enable_mirror(self, handler, mirror_name, query_params): - """启用/禁用镜像源""" - # 获取 enabled 参数 - enabled = query_params.get('enabled', ['true'])[0].lower() == 'true' - - # 更新内存中的配置 - if 'mirrors' not in self.config: - self.config['mirrors'] = {} - if mirror_name not in self.config['mirrors']: - self.config['mirrors'][mirror_name] = {} - self.config['mirrors'][mirror_name]['enabled'] = enabled - - # 保存到配置文件 - 使用项目根目录的绝对路径 - try: - import json - import os - # 获取项目根目录 (vs1 目录) - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - # 优先使用 settings.json - settings_path = os.path.join(project_root, 'settings.json') - config_path = settings_path - if not os.path.exists(config_path): - config_path = os.path.join(project_root, 'config.json') - - config_data = {} - - # 如果配置文件存在,读取它 - if os.path.exists(config_path): - with open(config_path, 'r', encoding='utf-8') as f: - config_data = json.load(f) - - # 更新配置 - if 'mirrors' not in config_data: - config_data['mirrors'] = {} - if mirror_name not in config_data['mirrors']: - config_data['mirrors'][mirror_name] = {} - config_data['mirrors'][mirror_name]['enabled'] = enabled - - # 保存配置 - with open(config_path, 'w', encoding='utf-8') as f: - json.dump(config_data, f, ensure_ascii=False, indent=2) - - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Failed to save config: {str(e)}" - }, 500) - return - - handler.send_json_response({ - "success": True, - "mirror": mirror_name, - "enabled": enabled - }) - - def api_add_mirror(self, handler): - """添加自定义加速源""" - # 读取请求体 - try: - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length > 0: - body = handler.rfile.read(content_length) - import json - data = json.loads(body.decode('utf-8')) - else: - data = {} - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Invalid request body: {str(e)}" - }, 400) - return - - # 验证必要参数 - mirror_name = data.get('name', '').strip() - mirror_type = data.get('type', 'custom').strip() - mirror_url = data.get('url', '').strip() - - if not mirror_name: - handler.send_json_response({ - "success": False, - "error": "Mirror name is required" - }, 400) - return - - # 验证名称格式(只允许字母、数字、下划线、连字符) - if not mirror_name.replace('_', '').replace('-', '').isalnum(): - handler.send_json_response({ - "success": False, - "error": "Mirror name can only contain letters, numbers, underscores and hyphens" - }, 400) - return - - # 保存到配置文件 - try: - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - # 优先使用 settings.json - settings_path = os.path.join(project_root, 'settings.json') - config_path = settings_path - if not os.path.exists(config_path): - config_path = os.path.join(project_root, 'config.json') - - config_data = {} - if os.path.exists(config_path): - with open(config_path, 'r', encoding='utf-8') as f: - config_data = json.load(f) - - # 初始化 mirrors 节 - if 'mirrors' not in config_data: - config_data['mirrors'] = {} - - # 添加新镜像 - config_data['mirrors'][mirror_name] = { - "type": mirror_type, - "url": mirror_url, - "enabled": data.get('enabled', True), - "description": data.get('description', f"Custom mirror: {mirror_name}"), - "storage_dir": data.get('storage_dir', mirror_name), - "custom": True, - "created_at": datetime.now().isoformat() - } - - # 保存配置 - with open(config_path, 'w', encoding='utf-8') as f: - json.dump(config_data, f, ensure_ascii=False, indent=2) - - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Failed to save config: {str(e)}" - }, 500) - return - - handler.send_json_response({ - "success": True, - "message": f"Mirror '{mirror_name}' added successfully", - "mirror": { - "name": mirror_name, - "type": mirror_type, - "url": mirror_url, - "enabled": data.get('enabled', True) - } - }) - - def api_update_mirror(self, handler, mirror_name): - """更新自定义加速源""" - # 读取请求体 - try: - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length > 0: - body = handler.rfile.read(content_length) - import json - data = json.loads(body.decode('utf-8')) - else: - data = {} - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Invalid request body: {str(e)}" - }, 400) - return - - # 读取并更新配置文件 - try: - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - settings_path = os.path.join(project_root, 'settings.json') - config_path = settings_path - if not os.path.exists(config_path): - config_path = os.path.join(project_root, 'config.json') - - if not os.path.exists(config_path): - handler.send_json_response({ - "success": False, - "error": "Config file not found" - }, 404) - return - - with open(config_path, 'r', encoding='utf-8') as f: - config_data = json.load(f) - - # 检查mirrors节是否存在 - if 'mirrors' not in config_data: - config_data['mirrors'] = {} - - # 检查镜像是否存在 - if mirror_name not in config_data['mirrors']: - handler.send_json_response({ - "success": False, - "error": f"Mirror '{mirror_name}' not found" - }, 404) - return - - # 更新镜像信息(只更新提供的字段) - mirror_data = config_data['mirrors'][mirror_name] - if 'type' in data: - mirror_data['type'] = data['type'] - if 'url' in data: - mirror_data['url'] = data['url'] - if 'target' in data: - mirror_data['target'] = data['target'] - if 'description' in data: - mirror_data['description'] = data['description'] - if 'enabled' in data: - mirror_data['enabled'] = data['enabled'] - if 'storage_dir' in data: - mirror_data['storage_dir'] = data['storage_dir'] - mirror_data['updated_at'] = datetime.now().isoformat() - - config_data['mirrors'][mirror_name] = mirror_data - - # 保存配置 - with open(config_path, 'w', encoding='utf-8') as f: - json.dump(config_data, f, ensure_ascii=False, indent=2) - - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Failed to save config: {str(e)}" - }, 500) - return - - handler.send_json_response({ - "success": True, - "message": f"Mirror '{mirror_name}' updated successfully", - "mirror": { - "name": mirror_name, - "type": mirror_data.get('type'), - "url": mirror_data.get('url'), - "enabled": mirror_data.get('enabled') - } - }) - - def api_delete_mirror(self, handler, mirror_name): - """删除加速源""" - # 从配置文件删除 - try: - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - # 优先使用 settings.json - settings_path = os.path.join(project_root, 'settings.json') - config_path = settings_path - if not os.path.exists(config_path): - config_path = os.path.join(project_root, 'config.json') - - if not os.path.exists(config_path): - handler.send_json_response({ - "success": False, - "error": "Config file not found" - }, 404) - return - - with open(config_path, 'r', encoding='utf-8') as f: - config_data = json.load(f) - - if 'mirrors' not in config_data or mirror_name not in config_data['mirrors']: - handler.send_json_response({ - "success": False, - "error": f"Mirror '{mirror_name}' not found" - }, 404) - return - - # 删除镜像 - del config_data['mirrors'][mirror_name] - - # 保存配置 - with open(config_path, 'w', encoding='utf-8') as f: - json.dump(config_data, f, ensure_ascii=False, indent=2) - - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Failed to delete mirror: {str(e)}" - }, 500) - return - - handler.send_json_response({ - "success": True, - "message": f"Mirror '{mirror_name}' deleted successfully" - }) - - # ==================== 用户管理 API ==================== - - def api_login(self, handler): - """用户登录""" - try: - # 获取配置和数据库 - config = handler.config if hasattr(handler, 'config') else {} - db = getattr(handler, 'db', None) or (hasattr(handler, 'config') and handler.config.get('_db_instance')) - - # 读取请求体 - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"success": False, "error": "请求体不能为空"}, 400) - return - - body = handler.rfile.read(content_length).decode('utf-8') - import json - data = json.loads(body) - username = data.get('username', '') - password = data.get('password', '') - - if not username or not password: - handler.send_json_response({"success": False, "error": "用户名和密码不能为空"}, 400) - return - - # 获取客户端IP - client_ip = None - try: - client_ip = handler.client_address[0] - except Exception: - pass - - # 获取 auth_type - auth_type = config.get('auth_type', 'none') - - # 如果 auth_type 为 none,任何用户都可以通过 - if auth_type == 'none': - # 生成 token - import secrets - token = secrets.token_hex(32) - token_expires_at = time.time() + 86400 # 24小时过期 - - handler.send_json_response({ - "success": True, - "token": token, - "token_expires_at": token_expires_at, - "username": username or 'anonymous', - "level": "admin" - }) - return - - # 验证凭据 - config_user = config.get('auth_user', '') - config_pass = config.get('auth_pass', '') - - # 数据库验证 - if db: - user = db.get_user(username) - if user and db.verify_password(password, user['password_hash']): - # 数据库验证成功,生成 token - import secrets - token = secrets.token_hex(32) - token_expires_at = time.time() + 86400 # 24小时过期 - - # 保存 token 到数据库 - db.update_user_token(username, token, token_expires_at) - - db.add_login_log(username, client_ip, 'success', '登录成功(数据库)') - - handler.send_json_response({ - "success": True, - "token": token, - "token_expires_at": token_expires_at, - "username": username, - "level": user.get('role', 'admin') - }) - return - - # 配置文件验证(仅当数据库中没有该用户时) - if username == config_user and password == config_pass: - # 验证成功,生成 token - import secrets - token = secrets.token_hex(32) - token_expires_at = time.time() + 86400 # 24小时过期 - - # 保存 token 到数据库 - if db: - user = db.get_user(username) - if user: - # 数据库中有用户,更新 token - db.update_user_token(username, token, token_expires_at) - else: - # 数据库中没有该用户,创建新用户并保存 token - from core.database import UserRecord - password_hash = db.hash_password(password) - new_user = UserRecord( - username=username, - password_hash=password_hash, - token=token, - token_expires_at=token_expires_at, - role='admin', - enabled=True - ) - with db.session() as session: - session.add(new_user) - db.add_login_log(username, client_ip, 'success', '登录成功(配置验证)') - - handler.send_json_response({ - "success": True, - "token": token, - "token_expires_at": token_expires_at, - "username": username, - "level": "admin" - }) - else: - # 验证失败 - if db: - db.add_login_log(username, client_ip, 'failed', '用户名或密码错误') - handler.send_json_response({"success": False, "error": "用户名或密码错误"}, 401) - - except json.JSONDecodeError: - handler.send_json_response({"success": False, "error": "无效的 JSON"}, 400) - except Exception as e: - import traceback - traceback.print_exc() - handler.send_json_response({"success": False, "error": str(e)}, 500) - - def api_change_password(self, handler): - """修改用户密码""" - try: - # 获取数据库实例 - db = getattr(handler, 'db', None) or (hasattr(handler, 'config') and handler.config.get('_db_instance')) - config = handler.config if hasattr(handler, 'config') else {} - - if not db: - handler.send_json_response({"success": False, "error": "数据库不可用"}, 500) - return - - # 读取请求体 - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({"error": "请求体不能为空"}, 400) - return - - body = handler.rfile.read(content_length) - data = json.loads(body.decode('utf-8')) - - username = data.get('username') - old_password = data.get('old_password') - new_password = data.get('new_password') - - if not username or not new_password: - handler.send_json_response({"error": "缺少必要参数"}, 400) - return - - # 获取配置中的账号密码 - config_user = config.get('auth_user', '') - config_pass = config.get('auth_pass', '') - - # 验证旧密码(优先验证数据库,没有则验证配置文件) - if old_password: - user = db.get_user(username) - if user: - # 验证数据库密码 - if not db.verify_password(old_password, user['password_hash']): - handler.send_json_response({"success": False, "error": "原密码错误"}, 400) - return - elif username == config_user and old_password != config_pass: - # 数据库没有用户,验证配置文件 - handler.send_json_response({"success": False, "error": "原密码错误"}, 400) - return - - # 使用 bcrypt 加密新密码 - new_hash = db.hash_password(new_password) - - # 更新数据库 - existing_user = db.get_user(username) - if existing_user: - success = db.update_password(username, new_hash) - else: - result = db.create_user(username, new_hash, 'admin') - success = result.get('success', False) - - if success: - # 修改密码后清除 token,强制重新登录 - if db and username: - db.clear_user_token(username) - - handler.send_json_response({ - "success": True, - "message": "密码修改成功(已存储到数据库),请重新登录" - }) - else: - handler.send_json_response({"success": False, "error": "用户不存在"}, 404) - - except json.JSONDecodeError: - handler.send_json_response({"error": "Invalid JSON format"}, 400) - except Exception as e: - handler.send_json_response({"error": str(e)}, 500) - - def api_get_login_logs(self, handler, query_params): - """获取登录日志""" - try: - # 尝试从多个位置获取数据库实例 - db = None - if hasattr(handler, 'db'): - db = handler.db - elif hasattr(handler, 'config') and handler.config: - db = handler.config.get('_db_instance') - - if not db: - handler.send_json_response({"success": False, "error": "数据库不可用,请确保数据库已启用"}, 500) - return - - # 安全获取 limit 参数 - try: - limit = int(query_params.get('limit', ['50'])[0]) - except (ValueError, TypeError): - limit = 50 - - logs = db.get_login_logs(limit=limit) - - handler.send_json_response({ - "success": True, - "logs": logs, - "count": len(logs) - }) - except Exception as e: - import traceback - trace = traceback.format_exc() - print(f"[ERROR] api_get_login_logs: {e}") - print(trace) - handler.send_json_response({"error": str(e)}, 500) - - def api_get_config(self, handler): - """获取配置文件内容 (settings.json)""" - try: - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - settings_path = os.path.join(project_root, 'settings.json') - - # 优先使用 settings.json,如果不存在则尝试 config.json - config_path = settings_path - if not os.path.exists(config_path): - config_path = os.path.join(project_root, 'config.json') - - if not os.path.exists(config_path): - handler.send_json_response({ - "success": False, - "error": "Config file not found (settings.json or config.json)" - }, 404) - return - - with open(config_path, 'r', encoding='utf-8') as f: - config_content = f.read() - - handler.send_json_response({ - "success": True, - "config": config_content, - "path": config_path, - "filename": os.path.basename(config_path) - }) - - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Failed to read config: {str(e)}" - }, 500) - - def api_save_config(self, handler): - """保存配置文件 (settings.json) - 只更新部分配置""" - try: - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length > 0: - config_content = handler.rfile.read(content_length).decode('utf-8') - else: - handler.send_json_response({ - "success": False, - "error": "No config content provided" - }, 400) - return - - # 验证 JSON 格式 - try: - updates = json.loads(config_content) - except json.JSONDecodeError as e: - handler.send_json_response({ - "success": False, - "error": f"Invalid JSON format: {str(e)}" - }, 400) - return - - if not isinstance(updates, dict): - handler.send_json_response({ - "success": False, - "error": "Config must be a JSON object" - }, 400) - return - - # 读取现有配置 - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - settings_path = os.path.join(project_root, 'settings.json') - - if not os.path.exists(settings_path): - handler.send_json_response({ - "success": False, - "error": "settings.json not found" - }, 404) - return - - # 加载现有配置 - from core.config import load_json_config - current_config = load_json_config(settings_path) or {} - - # 使用 deep_merge 合并更新(只更新传入的字段) - from core.config import deep_merge - new_config = deep_merge(current_config, updates) - - # 备份现有配置 - backup_path = settings_path + '.bak' - try: - with open(settings_path, 'r', encoding='utf-8') as f: - backup_content = f.read() - with open(backup_path, 'w', encoding='utf-8') as f: - f.write(backup_content) - except Exception: - pass - - # 保存合并后的配置 - with open(settings_path, 'w', encoding='utf-8') as f: - json.dump(new_config, f, indent=4, ensure_ascii=False) - - handler.send_json_response({ - "success": True, - "message": "Config updated successfully (partial update)", - "path": settings_path, - "updated_keys": list(updates.keys()) - }) - - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Failed to save config: {str(e)}" - }, 500) - - def api_reload_config(self, handler): - """重新加载配置文件(热更新)""" - try: - import os - from core.config import load_json_config - - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - settings_path = os.path.join(project_root, 'settings.json') - - if not os.path.exists(settings_path): - handler.send_json_response({ - "success": False, - "error": "Config file not found" - }, 404) - return - - # 加载并验证配置 - new_config = load_json_config(settings_path) - if new_config is None: - handler.send_json_response({ - "success": False, - "error": "Failed to load config" - }, 500) - return - - # 计算变更 - old_config = self.config.copy() - changes = self._compute_config_changes(old_config, new_config) - - # 更新内存配置 - self.config.update(new_config) - - # 通知各模块配置变更 - change_notifications = [] - if 'enable_monitor' in changes.get('modified', []): - change_notifications.append("Monitor settings changed") - if 'enable_sync' in changes.get('modified', []): - change_notifications.append("Sync settings changed") - if 'mirrors' in changes.get('modified', []): - change_notifications.append("Mirrors configuration changed") - - handler.send_json_response({ - "success": True, - "message": "Configuration reloaded successfully", - "changes": changes, - "change_notifications": change_notifications, - "timestamp": datetime.now().isoformat() - }) - - except Exception as e: - handler.send_json_response({ - "success": False, - "error": f"Failed to reload config: {str(e)}" - }, 500) - - def _compute_config_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]: - changes['modified'].append(key) - - return changes - - def api_get_config_changes(self, handler): - """获取配置变更历史""" - handler.send_json_response({ - "changes": [], - "message": "Config change history tracking requires hot reload enabled" - }) - - # ==================== WebSocket/SSE 状态 API ==================== - - def api_get_ws_clients(self, handler): - """获取WebSocket客户端状态""" - # 需要访问全局ws_manager - handler.send_json_response({ - "ws_clients": 0, - "sse_clients": 0 - }) - - # ==================== Prometheus 指标 API ==================== - - def api_get_metrics(self, handler): - """获取 Prometheus 格式的指标""" - from core.prometheus import PrometheusMetrics - - metrics = PrometheusMetrics(self.config) - - # 设置运行时间 - uptime = time.time() - self.config.get('start_time', time.time()) - metrics.set_uptime(uptime) - - # 从数据库获取统计 - if self.db_enabled and self.db: - try: - db_stats = self.db.get_stats() - metrics.set_files( - db_stats.get('total_files', 0), - db_stats.get('total_size', 0) - ) - metrics.set_db_stats( - db_stats.get('total_files', 0), - db_stats.get('total_sync_records', 0), - db_stats.get('total_cache_records', 0) - ) - except Exception as e: - pass - - # 从缓存获取统计 - if hasattr(handler, 'cache_manager') and handler.cache_manager: - try: - cache_stats = handler.cache_manager.get_stats() - metrics.set_cache( - cache_stats.get('size', 0), - cache_stats.get('count', 0), - cache_stats.get('hits', 0), - cache_stats.get('misses', 0) - ) - except Exception as e: - pass - - # 从监控模块获取系统指标 - if hasattr(handler, 'monitor') and handler.monitor: - try: - monitor_stats = handler.monitor.get_realtime_stats() - metrics.set_system( - cpu=monitor_stats.get('cpu', {}).get('percent', 0), - memory=monitor_stats.get('memory', {}).get('percent', 0), - disk=monitor_stats.get('disk', {}).get('percent', 0), - disk_free=monitor_stats.get('disk', {}).get('free', 0), - disk_total=monitor_stats.get('disk', {}).get('total', 0), - network_rx=monitor_stats.get('network', {}).get('rx', 0), - network_tx=monitor_stats.get('network', {}).get('tx', 0) - ) - except Exception as e: - pass - - # 从同步管理器获取镜像状态 - if hasattr(handler, 'sync_manager') and handler.sync_manager: - mirrors = self.config.get('mirrors', {}) - for mirror_type, mirror_config in mirrors.items(): - if isinstance(mirror_config, dict): - enabled = mirror_config.get('enabled', True) - last_sync = 0 - metrics.set_mirror_status(mirror_type, enabled, last_sync) - - # 生成 Prometheus 格式输出 - output = metrics.generate_metrics() - - handler.send_response(200) - handler.send_header('Content-Type', 'text/plain; charset=utf-8') - handler.send_header('Content-Length', str(len(output))) - handler.end_headers() - handler.wfile.write(output.encode('utf-8')) - - # ==================== 服务器信息 API ==================== - - def api_get_server_info(self, handler): - """获取服务器完整信息""" - import psutil - - uptime_seconds = time.time() - self.config.get('start_time', time.time()) - uptime_str = self._format_uptime(uptime_seconds) - - handler.send_json_response({ - "name": self.config.get('server_name', 'HYC下载站'), - "version": "2.2.0", - "uptime_seconds": round(uptime_seconds, 2), - "uptime_formatted": uptime_str, - "api_version": "v2", - "config": { - "host": self.config.get('host'), - "port": self.config.get('port'), - "base_dir": self.config.get('base_dir'), - "auth_type": self.config.get('auth_type'), - "directory_listing": self.config.get('directory_listing'), - "max_upload_size": self.config.get('max_upload_size') - }, - "features": { - "websocket": True, - "sse": True, - "sync": True, - "mirrors": True, - "cache": True, - "monitor": True - } - }) - - # ==================== 告警管理 API ==================== - - def api_get_alerts(self, handler, query_params): - """获取告警列表""" - try: - from core.alerts import AlertManager - - alert_manager = AlertManager(self.config.get('alerts', {})) - - # 解析查询参数 - limit = int(query_params.get('limit', [50])[0]) - acknowledged = query_params.get('acknowledged', [None])[0] - severity = query_params.get('severity', [None])[0] - - if acknowledged is not None: - acknowledged = acknowledged.lower() == 'true' - - alerts = alert_manager.get_alerts( - acknowledged=acknowledged, - severity=severity, - limit=limit - ) - - stats = alert_manager.get_stats() - - handler.send_json_response({ - 'alerts': alerts, - 'count': len(alerts), - 'stats': stats - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def api_acknowledge_alert(self, handler, alert_id): - """确认告警""" - try: - from core.alerts import AlertManager - - alert_manager = AlertManager(self.config.get('alerts', {})) - success = alert_manager.acknowledge_alert(alert_id) - - handler.send_json_response({ - 'success': success, - 'message': f"Alert {alert_id} acknowledged" if success else "Alert not found" - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def api_clear_alerts(self, handler): - """清除告警历史""" - try: - from core.alerts import AlertManager - - alert_manager = AlertManager(self.config.get('alerts', {})) - success = alert_manager.clear_history() - - handler.send_json_response({ - 'success': success, - 'message': 'Alert history cleared' - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def api_test_alert(self, handler): - """测试告警发送""" - try: - from core.alerts import AlertManager, Alert, AlertSeverity - - alert_manager = AlertManager(self.config.get('alerts', {})) - - # 创建测试告警 - test_alert = Alert( - alert_type='test', - severity=AlertSeverity.INFO, - title='Test Alert', - message='This is a test alert from HYC Mirror Server', - details={'test': True, 'timestamp': datetime.now().isoformat()} - ) - - success = alert_manager.trigger_alert(test_alert) - - handler.send_json_response({ - 'success': success, - 'message': 'Test alert sent successfully' if success else 'Failed to send test alert (check configuration)' - }) - - except Exception as e: - handler.send_json_response({ - 'error': str(e) - }, 500) - - def api_get_alert_config(self, handler): - """获取告警配置""" - alerts_config = self.config.get('alerts', {}) - - # 隐藏敏感信息 - config = { - 'enabled': alerts_config.get('enabled', False), - 'email': { - 'enabled': alerts_config.get('email', {}).get('enabled', False), - 'smtp_host': alerts_config.get('email', {}).get('smtp_host', ''), - 'smtp_port': alerts_config.get('email', {}).get('smtp_port', 587), - 'from_address': alerts_config.get('email', {}).get('from_address', ''), - 'to_addresses': alerts_config.get('email', {}).get('to_addresses', []), - 'use_tls': alerts_config.get('email', {}).get('use_tls', True) - }, - 'webhook': { - 'enabled': alerts_config.get('webhook', {}).get('enabled', False), - 'url': '***' if alerts_config.get('webhook', {}).get('url') else '' - }, - 'rules': alerts_config.get('rules', {}) - } - - handler.send_json_response(config) - - def api_save_alert_config(self, handler): - """保存告警配置""" - try: - import os - from core.alerts import AlertManager - - # 获取请求体 - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length > 0: - post_data = handler.rfile.read(content_length) - try: - new_config = json.loads(post_data.decode('utf-8')) - except json.JSONDecodeError as e: - handler.send_json_response({ - 'success': False, - 'error': f"Invalid JSON format: {str(e)}" - }, 400) - return - else: - handler.send_json_response({ - 'success': False, - 'error': "No configuration data provided" - }, 400) - return - - # 更新配置 - if 'alerts' not in self.config: - self.config['alerts'] = {} - - # 更新告警配置 - if 'enabled' in new_config: - self.config['alerts']['enabled'] = new_config['enabled'] - if 'email' in new_config: - self.config['alerts']['email'] = new_config['email'] - if 'webhook' in new_config: - self.config['alerts']['webhook'] = new_config['webhook'] - if 'rules' in new_config: - self.config['alerts']['rules'] = new_config['rules'] - - # 保存到文件 - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - settings_path = os.path.join(project_root, 'settings.json') - - with open(settings_path, 'r', encoding='utf-8') as f: - settings_data = json.load(f) - - settings_data['alerts'] = self.config['alerts'] - - with open(settings_path, 'w', encoding='utf-8') as f: - json.dump(settings_data, f, ensure_ascii=False, indent=4) - - handler.send_json_response({ - 'success': True, - 'message': 'Alert configuration saved successfully', - 'config': self.config.get('alerts', {}) - }) - - except Exception as e: - handler.send_json_response({ - 'success': False, - 'error': str(e) - }, 500) - - # ==================== 平滑重启 API ==================== - - def api_get_restart_status(self, handler): - """获取重启状态""" - from core.graceful_restart import GracefulRestartManager, ServerState - - restart_manager = GracefulRestartManager(self.config.get('restart', {})) - stats = restart_manager.get_stats() - - handler.send_json_response({ - 'state': stats['state'], - 'pending_requests': stats['pending_requests'], - 'graceful_timeout': stats['graceful_timeout'], - 'recent_restarts': stats['recent_restarts'] - }) - - def api_get_pending_requests(self, handler): - """获取待处理请求""" - from core.graceful_restart import GracefulRestartManager - - restart_manager = GracefulRestartManager(self.config.get('restart', {})) - pending = restart_manager.get_pending_requests() - - handler.send_json_response({ - 'count': len(pending), - 'requests': pending - }) - - def api_graceful_restart(self, handler, query_params): - """执行优雅重启""" - from core.graceful_restart import GracefulRestartManager, RestartStrategy - - # 解析策略参数 - strategy_param = query_params.get('strategy', ['graceful'])[0] - try: - strategy = RestartStrategy(strategy_param) - except ValueError: - strategy = RestartStrategy.GRACEFUL - - restart_manager = GracefulRestartManager(self.config.get('restart', {})) - - # 准备重启 - prepare_result = restart_manager.prepare_restart() - if not prepare_result['success']: - handler.send_json_response({ - 'success': False, - 'error': prepare_result['message'] - }, 500) - return - - # 返回待处理请求信息,让客户端决定是否继续 - pending_count = prepare_result['pending_requests'] - - handler.send_json_response({ - 'success': True, - 'pending_requests': pending_count, - 'message': f'Ready to restart with {pending_count} pending requests', - 'strategy': strategy.value, - 'graceful_timeout': restart_manager.graceful_timeout, - 'continue_url': '/api/v2/server/restart/confirm' - }) - - def api_confirm_restart(self, handler): - """确认执行重启""" - from core.graceful_restart import GracefulRestartManager, RestartStrategy - - try: - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length > 0: - post_data = json.loads(handler.rfile.read(content_length)) - strategy = post_data.get('strategy', 'graceful') - else: - strategy = 'graceful' - - try: - restart_strategy = RestartStrategy(strategy) - except ValueError: - restart_strategy = RestartStrategy.GRACEFUL - - except json.JSONDecodeError: - restart_strategy = RestartStrategy.GRACEFUL - - restart_manager = GracefulRestartManager(self.config.get('restart', {})) - - # 执行重启 - result = restart_manager.perform_restart(strategy=restart_strategy) - - handler.send_json_response(result) - - def api_immediate_restart(self, handler): - """立即重启服务器""" - from core.graceful_restart import GracefulRestartManager - - restart_manager = GracefulRestartManager(self.config.get('restart', {})) - - # 获取脚本路径 - script_path = self.config.get('main_script', 'main.py') - - result = restart_manager.perform_restart( - strategy='immediate', - script_path=script_path - ) - - handler.send_json_response(result) - - def api_get_restart_history(self, handler): - """获取重启历史""" - from core.graceful_restart import GracefulRestartManager - - restart_manager = GracefulRestartManager(self.config.get('restart', {})) - history = restart_manager.get_restart_history() - - handler.send_json_response({ - 'count': len(history), - 'history': history - }) - - def api_get_restart_config(self, handler): - """获取重启配置""" - restart_config = self.config.get('restart', {}) - - handler.send_json_response({ - 'graceful_timeout': restart_config.get('graceful_timeout', 30), - 'shutdown_timeout': restart_config.get('shutdown_timeout', 10), - 'enabled': True - }) - - def api_update_restart_config(self, handler): - """更新重启配置""" - try: - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({ - 'success': False, - 'error': 'No configuration data provided' - }, 400) - return - - new_config = json.loads(handler.rfile.read(content_length)) - - # 更新内存配置 - if 'restart' not in self.config: - self.config['restart'] = {} - - if 'graceful_timeout' in new_config: - self.config['restart']['graceful_timeout'] = new_config['graceful_timeout'] - if 'shutdown_timeout' in new_config: - self.config['restart']['shutdown_timeout'] = new_config['shutdown_timeout'] - - # 保存到文件 - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - settings_path = os.path.join(project_root, 'settings.json') - - with open(settings_path, 'r', encoding='utf-8') as f: - settings_data = json.load(f) - - if 'restart' not in settings_data: - settings_data['restart'] = {} - settings_data['restart'] = {**settings_data.get('restart', {}), **new_config} - - with open(settings_path, 'w', encoding='utf-8') as f: - json.dump(settings_data, f, ensure_ascii=False, indent=4) - - handler.send_json_response({ - 'success': True, - 'message': 'Restart configuration updated', - 'config': self.config.get('restart', {}) - }) - - except json.JSONDecodeError as e: - handler.send_json_response({ - 'success': False, - 'error': f'Invalid JSON format: {str(e)}' - }, 400) - except Exception as e: - handler.send_json_response({ - 'success': False, - 'error': str(e) - }, 500) - - def _format_uptime(self, seconds: float) -> str: - """格式化运行时间""" - if seconds < 60: - return f"{int(seconds)}秒" - elif seconds < 3600: - minutes = int(seconds // 60) - return f"{minutes}分钟" - elif seconds < 86400: - hours = int(seconds // 3600) - minutes = int((seconds % 3600) // 60) - return f"{hours}小时{minutes}分钟" - else: - days = int(seconds // 86400) - hours = int((seconds % 86400) // 3600) - return f"{days}天{hours}小时" - - # ==================== 缓存预热 API ==================== - - def api_get_prewarm_status(self, handler): - """获取缓存预热状态""" - from core.cache_prewarm import CachePrewarmer - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - status = prewarmer.get_status() - - handler.send_json_response(status) - - def api_get_prewarm_stats(self, handler): - """获取缓存预热统计""" - from core.cache_prewarm import CachePrewarmer - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - stats = prewarmer.get_stats() - - handler.send_json_response(stats) - - def api_get_prewarm_items(self, handler, query_params): - """获取预热项目列表""" - from core.cache_prewarm import CachePrewarmer - - # 解析查询参数 - status = query_params.get('status', [None])[0] - mirror_type = query_params.get('mirror_type', [None])[0] - limit = int(query_params.get('limit', [50])[0]) - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - items = prewarmer.get_items(status=status, mirror_type=mirror_type, limit=limit) - - handler.send_json_response({ - 'count': len(items), - 'items': items - }) - - def api_get_prewarm_history(self, handler): - """获取预热历史""" - from core.cache_prewarm import CachePrewarmer - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - history = prewarmer.get_history() - - handler.send_json_response({ - 'count': len(history), - 'history': history - }) - - def api_run_prewarm(self, handler, query_params): - """执行缓存预热""" - from core.cache_prewarm import CachePrewarmer - - # 解析参数 - mirror_type = query_params.get('mirror_type', [None])[0] - limit = int(query_params.get('limit', [50])[0]) - priority = query_params.get('priority', ['medium'])[0] - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - - # 如果指定了镜像类型,只预热该类型 - targets = None - if mirror_type: - targets = [] - from core.cache_prewarm import PrewarmTarget - targets.append(PrewarmTarget( - mirror_type=mirror_type, - priority=priority, - limit=limit - )) - - result = prewarmer.run(targets=targets) - - handler.send_json_response(result) - - def api_add_prewarm_items(self, handler): - """添加预热项目""" - try: - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({ - 'success': False, - 'error': 'No data provided' - }, 400) - return - - data = json.loads(handler.rfile.read(content_length)) - - mirror_type = data.get('mirror_type') - items = data.get('items', []) - priority = data.get('priority', 'medium') - - if not mirror_type or not items: - handler.send_json_response({ - 'success': False, - 'error': 'mirror_type and items are required' - }, 400) - return - - from core.cache_prewarm import CachePrewarmer - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - prewarmer.add_items_batch(mirror_type, items, priority) - - handler.send_json_response({ - 'success': True, - 'message': f'Added {len(items)} items to prewarm queue', - 'mirror_type': mirror_type, - 'count': len(items) - }) - - except json.JSONDecodeError as e: - handler.send_json_response({ - 'success': False, - 'error': f'Invalid JSON format: {str(e)}' - }, 400) - except Exception as e: - handler.send_json_response({ - 'success': False, - 'error': str(e) - }, 500) - - def api_add_popular_items(self, handler, query_params): - """添加流行项目到预热队列""" - from core.cache_prewarm import CachePrewarmer - - mirror_type = query_params.get('mirror_type', [None])[0] - limit = int(query_params.get('limit', [20])[0]) - priority = query_params.get('priority', ['medium'])[0] - - if not mirror_type: - handler.send_json_response({ - 'success': False, - 'error': 'mirror_type is required' - }, 400) - return - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - popular = prewarmer.get_popular_items(mirror_type) - - if limit: - popular = popular[:limit] - - prewarmer.add_popular_items_to_queue(mirror_type, limit, priority) - - handler.send_json_response({ - 'success': True, - 'message': f'Added {len(popular)} popular items', - 'mirror_type': mirror_type, - 'count': len(popular), - 'items': popular - }) - - def api_get_popular_items(self, handler, query_params): - """获取流行项目列表""" - from core.cache_prewarm import CachePrewarmer - - mirror_type = query_params.get('mirror_type', [None])[0] - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - - if mirror_type: - items = prewarmer.get_popular_items(mirror_type) - handler.send_json_response({ - 'mirror_type': mirror_type, - 'count': len(items), - 'items': items - }) - else: - all_popular = prewarmer._popular_items - handler.send_json_response({ - 'mirror_types': list(all_popular.keys()), - 'total_types': len(all_popular) - }) - - def api_clear_prewarm_queue(self, handler): - """清空预热队列""" - from core.cache_prewarm import CachePrewarmer - - prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) - prewarmer.clear_items() - - handler.send_json_response({ - 'success': True, - 'message': 'Prewarm queue cleared' - }) - - def api_get_prewarm_config(self, handler): - """获取缓存预热配置""" - prewarm_config = self.config.get('cache_prewarm', {}) - - handler.send_json_response({ - 'enabled': prewarm_config.get('enabled', False), - 'schedule': prewarm_config.get('schedule', '0 3 * * *'), - 'batch_size': prewarm_config.get('batch_size', 10), - 'targets': prewarm_config.get('targets', []) - }) - - def api_save_prewarm_config(self, handler): - """保存缓存预热配置""" - try: - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length == 0: - handler.send_json_response({ - 'success': False, - 'error': 'No configuration data provided' - }, 400) - return - - new_config = json.loads(handler.rfile.read(content_length)) - - # 更新内存配置 - if 'cache_prewarm' not in self.config: - self.config['cache_prewarm'] = {} - - if 'enabled' in new_config: - self.config['cache_prewarm']['enabled'] = new_config['enabled'] - if 'schedule' in new_config: - self.config['cache_prewarm']['schedule'] = new_config['schedule'] - if 'batch_size' in new_config: - self.config['cache_prewarm']['batch_size'] = new_config['batch_size'] - if 'targets' in new_config: - self.config['cache_prewarm']['targets'] = new_config['targets'] - - # 保存到文件 - import os - project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - settings_path = os.path.join(project_root, 'settings.json') - - with open(settings_path, 'r', encoding='utf-8') as f: - settings_data = json.load(f) - - if 'cache_prewarm' not in settings_data: - settings_data['cache_prewarm'] = {} - settings_data['cache_prewarm'] = {**settings_data.get('cache_prewarm', {}), **new_config} - - with open(settings_path, 'w', encoding='utf-8') as f: - json.dump(settings_data, f, ensure_ascii=False, indent=4) - - handler.send_json_response({ - 'success': True, - 'message': 'Cache prewarm configuration saved', - 'config': self.config.get('cache_prewarm', {}) - }) - - except json.JSONDecodeError as e: - handler.send_json_response({ - 'success': False, - 'error': f'Invalid JSON format: {str(e)}' - }, 400) - except Exception as e: - handler.send_json_response({ - 'success': False, - 'error': str(e) - }, 500) - - # ==================== API 文档 ==================== - - def api_get_api_docs(self, handler, format: str = 'json'): - """获取 API 文档""" - from core.api_docs import generate_api_docs - - docs = generate_api_docs(self.config) - - if format == 'yaml': - try: - import yaml - content = yaml.dump(docs, default_flow_style=False, allow_unicode=True) - handler.send_response(200) - handler.send_header('Content-Type', 'text/yaml') - handler.send_header('Content-Length', len(content.encode('utf-8'))) - handler.end_headers() - handler.wfile.write(content.encode('utf-8')) - return - except ImportError: - format = 'json' - - content = json.dumps(docs, ensure_ascii=False, indent=2) - handler.send_response(200) - handler.send_header('Content-Type', 'application/json') - handler.send_header('Content-Length', len(content.encode('utf-8'))) - handler.send_header('Access-Control-Allow-Origin', '*') - handler.end_headers() - handler.wfile.write(content.encode('utf-8')) - - def api_generate_api_docs(self, handler): - """生成并保存 API 文档""" - try: - from core.api_docs import save_api_docs - - # 获取保存路径 - content_length = int(handler.headers.get('Content-Length', 0)) - if content_length > 0: - data = json.loads(handler.rfile.read(content_length)) - filepath = data.get('filepath', 'docs/api-docs.json') - format = data.get('format', 'json') - else: - filepath = 'docs/api-docs.json' - format = 'json' - - # 生成并保存文档 - saved_path = save_api_docs(self.config, filepath, format) - - handler.send_json_response({ - 'success': True, - 'message': f'API documentation generated', - 'path': saved_path, - 'format': format - }) - - except Exception as e: - handler.send_json_response({ - 'success': False, - 'error': str(e) - }, 500) +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +"""API v2 版本处理模块 - 增强版本""" + +import os +import json +import re +import time +from datetime import datetime + +from .v1 import APIv1 +from .admin import AdminAPI +from core.utils import format_file_size, is_safe_path +from core.api_auth import require_auth, check_endpoint_auth + + +class APIv2(APIv1): + """API v2 - 增强版本,继承v1并添加新功能""" + + def __init__(self, config): + super().__init__(config) + self.api_version = "v2" + # 管理员API处理器(始终创建,auth_type检查在装饰器中处理) + self.admin_api = AdminAPI(config) + + def handle_request(self, handler, method, path, query_params): + """处理API v2请求""" + import sys + + # 调试输出 (debug-v2) + if handler._is_debug_enabled('v2'): + msg = f"\n=== DEBUG APIv2.handle_request ===\n path: '{path}'\n method: '{method}'" + handler._debug_log('v2', msg, '\033[36m') + + # v2 端点列表(这些是 v2 专有端点,不应该交给 APIv1 处理) + v2_endpoints = [ + 'admin/', + 'search/enhanced', + 'search/by-tag', + 'search/by-date', + 'stats/detailed', + 'stats/trending', + 'stats/download-trend', + 'stats/download-by-period', + 'stats/rank', + 'cache/popular', + 'file/', + 'metadata/', + 'monitor/', + 'health/', + 'health/', + 'alerts', + 'alerts/', + 'webhooks', + 'sync/', + 'cache/', + 'mirrors', + 'pypi', + 'config', + 'server/', + 'health', + 'downloads/', + 'metrics', + 'activity', + 'user/', + 'jdk/', + ] + + # 检查是否是 v2 专有端点 + is_v2_endpoint = False + for ep in v2_endpoints: + if path == ep or path.startswith(ep): + is_v2_endpoint = True + break + + # 如果是 v2 端点,调用 auth_manager 时要跳过 v1 路径 + auth_manager = getattr(handler, 'auth_manager', None) + + # ========== 首先处理 v2 专有端点 ========== + if is_v2_endpoint: + # 公开端点列表(无需认证) + public_endpoints = [ + 'admin/auth/verify', + 'user/login', + 'search/enhanced', + 'search/by-tag', + 'search/by-date', + 'jdk/retrieve/', + 'jdk/list', + ] + + # 检查是否公开端点 + is_public = path in public_endpoints or any(path.startswith(ep + '/') for ep in public_endpoints) + + # 如果auth_type为none,跳过所有认证检查 + auth_type = self.config.get('auth_type', 'none') + skip_auth = auth_type == 'none' + + # 如果需要认证(不是公开端点且auth_type不是none) + if auth_manager and not is_public and not skip_auth: + # 传入完整路径(api/v2/ 前缀),与 ADMIN_API_ENDPOINTS 规则匹配 + auth_check = check_endpoint_auth(method, f"api/v2/{path}", auth_manager) + if auth_check['required']: + auth_result = auth_manager.validate_request(handler, 'admin') + if not auth_result.get('authenticated'): + handler.send_response(401) + handler.send_header('WWW-Authenticate', 'Bearer') + handler.send_header('Access-Control-Allow-Origin', '*') + handler.send_json_response({ + "error": "认证Required", + "code": "UNAUTHORIZED", + "required_permission": auth_check['permission'] + }) + return + + if auth_check['permission']: + if not auth_manager.check_permission(auth_result, auth_check['permission']): + handler.send_header('Access-Control-Allow-Origin', '*') + handler.send_json_response({ + "error": "权限不足", + "code": "FORBIDDEN", + "required_permission": auth_check['permission'] + }, 403) + return + + handler.auth_result = auth_result + + # ========== v2 端点处理 ========== + # 认证验证 API + if path == 'admin/auth/verify': + if method == 'POST': + self.api_verify_auth(handler) + else: + handler.send_error(405) + return + + # 管理员统计 API + if path == 'admin/stats': + self.admin_api.handle_request(handler, method, 'stats', query_params) + return + + # v2新增的增强功能 + # 增强搜索API(无需认证) + if path == 'search/enhanced': + if method == 'GET': + self.api_search_files_enhanced(handler, query_params) + else: + handler.send_error(405) + return + elif path == 'search/by-tag': + if method == 'GET': + self.api_search_by_tag(handler, query_params) + else: + handler.send_error(405) + return + elif path == 'search/by-date': + if method == 'GET': + self.api_search_by_date(handler, query_params) + else: + handler.send_error(405) + return + + # JDK API (调用 v1 中的实现) + elif path.startswith('jdk/'): + jdk_path = path[4:] # 移除 'jdk/' 前缀 + self.handle_jdk_api(handler, method, jdk_path) + return + + # 增强统计API + elif path == 'stats/detailed': + if method == 'GET': + self.api_get_stats_detailed(handler) + else: + handler.send_error(405) + return + elif path == 'stats/trending': + if method == 'GET': + self.api_get_trending_files(handler, query_params) + else: + handler.send_error(405) + return + + # 增强文件操作API + elif path.startswith('file/') and path.endswith('/metadata'): + filename = path[5:-9] # 移除 'file/' 和 '/metadata' + if method == 'GET': + self.api_get_file_metadata(handler, filename) + elif method == 'PUT': + self.api_update_file_metadata(handler, filename) + else: + handler.send_error(405) + return + + # 批量元数据操作 + elif path == 'metadata/batch': + if method == 'GET': + self.api_get_batch_metadata(handler, query_params) + elif method == 'PUT': + self.api_update_batch_metadata(handler) + else: + handler.send_error(405) + return + + # 文件版本控制 + elif path.startswith('file/') and '/versions' in path: + parts = path.split('/') + if len(parts) >= 3 and parts[-1] == 'versions': + filename = '/'.join(parts[1:-1]) + if method == 'GET': + self.api_get_file_versions(handler, filename) + elif method == 'POST': + self.api_create_file_version(handler, filename) + else: + handler.send_error(405) + return + + # 缩略图API + elif path.startswith('file/') and path.endswith('/thumbnail'): + filename = path[5:-10] # 移除 'file/' 和 '/thumbnail' + if method == 'GET': + self.api_get_file_thumbnail(handler, filename, query_params) + else: + handler.send_error(405) + return + + # 服务器监控(实时数据) + elif path == 'monitor/realtime': + if method == 'GET': + self.api_get_realtime_stats(handler) + else: + handler.send_error(405) + return + elif path == 'monitor/history': + if method == 'GET': + self.api_get_monitor_history(handler, query_params) + else: + handler.send_error(405) + return + elif path == 'monitor/detailed': + if method == 'GET': + self.api_get_monitor_detailed(handler) + else: + handler.send_error(405) + return + + # ========== 镜像源健康检查 ========== + elif path == 'health/sources': + if method == 'GET': + self.api_get_source_health(handler, query_params) + else: + handler.send_error(405) + return + elif path.startswith('health/check/'): + source_name = path[13:] # 移除 'health/check/' + if method == 'GET': + self.api_check_source(handler, source_name) + else: + handler.send_error(405) + return + elif path == 'health/failover': + if method == 'GET': + self.api_get_failover_status(handler) + else: + handler.send_error(405) + return + elif path.startswith('health/failover/') and method == 'POST': + mirror_type = path[18:] # 移除 'health/failover/' + self.api_trigger_failover(handler, mirror_type) + return + elif path == 'health/stats': + if method == 'GET': + from core.health_check import HealthChecker + checker = HealthChecker(self.config.get('health_check', {})) + handler.send_json_response(checker.get_stats()) + else: + handler.send_error(405) + return + + # Webhook支持 + elif path == 'webhooks': + if method == 'GET': + self.api_list_webhooks(handler) + elif method == 'POST': + self.api_create_webhook(handler) + else: + handler.send_error(405) + return + elif path.startswith('webhooks/'): + webhook_id = path[9:] + + # 交付历史: webhooks/{id}/deliveries + if webhook_id.endswith('/deliveries'): + actual_id = webhook_id[:-11] + if method == 'GET': + self.api_get_webhook_deliveries(handler, actual_id) + else: + handler.send_error(405) + return + + # 统计: webhooks/{id}/stats + if webhook_id.endswith('/stats'): + actual_id = webhook_id[:-6] + if method == 'GET': + self.api_get_webhook_stats(handler, actual_id) + else: + handler.send_error(405) + return + + # 单个 webhook 操作 + if method == 'GET': + self.api_get_webhook(handler, webhook_id) + elif method == 'DELETE': + self.api_delete_webhook(handler, webhook_id) + elif method == 'POST': + # 检查是否是测试请求 + if '/test' in webhook_id: + actual_id = webhook_id.split('/')[0] + self.api_test_webhook(handler, actual_id) + else: + self.api_test_webhook(handler, webhook_id) + elif method == 'PUT': + self.api_update_webhook(handler, webhook_id) + else: + handler.send_error(405) + return + + # ========== 以下端点需要认证 ========== + + # 同步管理API + elif path == 'sync/sources': + if method == 'GET': + self.api_get_sync_sources(handler) + elif method == 'POST': + self.api_add_sync_source(handler) + else: + handler.send_error(405) + return + elif path.startswith('sync/') and path.endswith('/start'): + source_name = path[5:-6] + if method == 'POST': + self.api_start_sync(handler, source_name) + else: + handler.send_error(405) + return + elif path.startswith('sync/') and path.endswith('/stop'): + source_name = path[5:-5] + if method == 'POST': + self.api_stop_sync(handler, source_name) + else: + handler.send_error(405) + return + elif path.startswith('sync/') and path.endswith('/status'): + source_name = path[5:-7] + if method == 'GET': + self.api_get_sync_status(handler, source_name) + else: + handler.send_error(405) + return + elif path == 'sync/history': + if method == 'GET': + self.api_get_sync_history(handler, query_params) + else: + handler.send_error(405) + return + elif path == 'sync/packages': + if method == 'POST': + self.api_sync_packages(handler) + else: + handler.send_error(405) + return + elif path.startswith('sync/packages/') and path.endswith('/status'): + source_name = path[14:-8] + if method == 'GET': + self.api_get_temp_sync_status(handler, source_name) + else: + handler.send_error(405) + return + + # 缓存管理API + elif path == 'cache/stats': + if method == 'GET': + self.api_get_cache_stats(handler) + else: + handler.send_error(405) + return + elif path == 'cache/clean' and method == 'POST': + self.api_clean_cache(handler) + return + elif path == 'cache/usage': + if method == 'GET': + self.api_get_cache_usage(handler) + else: + handler.send_error(405) + return + + # 镜像加速源API + elif path == 'mirrors': + if method == 'GET': + self.api_list_mirrors(handler) + else: + handler.send_error(405) + return + + # 镜像管理 API - 必须放在特殊镜像处理之前 + # mirrors/xxx/enable, mirrors/xxx/refresh, mirrors/xxx (PUT/DELETE) + elif path.endswith('/enable'): + # 格式: mirrors/xxx/enable + parts = path.split('/') + if len(parts) >= 3: + mirror_name = parts[1] + if method == 'PUT': + self.api_enable_mirror(handler, mirror_name, query_params) + else: + handler.send_error(405) + return + elif path.endswith('/refresh'): + # 格式: mirrors/xxx/refresh + parts = path.split('/') + if len(parts) >= 3: + mirror_name = parts[1] + if method == 'POST': + self.api_refresh_mirror(handler, mirror_name) + else: + handler.send_error(405) + return + + # 加速源访问 API - mirrors/pypi/*, mirrors/npm/*, mirrors/go/* 等 + # 注意:排除 mirrors/pypi/enable, mirrors/pypi/refresh 等管理路径 + elif (path.startswith('mirrors/pypi/') and not '/enable' in path and not '/refresh' in path) or path == 'mirrors/pypi': + # PyPI 加速源 - 优先本地,没有再从上游拉取 + from mirrors import get_mirror_handler + mirror_path = path.replace('mirrors/pypi/', '', 1) + if path == 'mirrors/pypi': + mirror_path = 'pypi' + handler_class = get_mirror_handler('pypi') + if handler_class: + # 创建实例 + mirror_config = self.config.get('mirrors', {}).get('pypi', {}) + base_dir = self.config.get('base_dir', './downloads') + storage_subdir = mirror_config.get('storage_dir', 'pypi') + # 配置: base_dir=downloads目录, storage_dir=pypi子目录 + mirror_handler = handler_class({ + 'base_dir': base_dir, + 'storage_dir': storage_subdir, + 'upstream_url': mirror_config.get('url', 'https://pypi.org') + }) + mirror_handler.handle_request(handler, mirror_path) + else: + handler.send_error(404, "PyPI mirror not configured") + return + elif path.startswith('mirrors/npm/') or path == 'mirrors/npm': + # NPM 加速源 - 优先本地,没有再从上游拉取 + from mirrors import get_mirror_handler + mirror_path = path.replace('mirrors/npm/', '', 1) + if path == 'mirrors/npm': + mirror_path = 'npm' + handler_class = get_mirror_handler('npm') + if handler_class: + mirror_config = self.config.get('mirrors', {}).get('npm', {}) + base_dir = self.config.get('base_dir', './downloads') + storage_subdir = mirror_config.get('storage_dir', 'npm') + mirror_handler = handler_class({ + 'base_dir': base_dir, + 'storage_dir': storage_subdir, + 'upstream_url': mirror_config.get('url', 'https://registry.npmjs.org') + }) + mirror_handler.handle_request(handler, mirror_path) + else: + handler.send_error(404, "NPM mirror not configured") + return + elif path.startswith('mirrors/go/') or path == 'mirrors/go': + # Go 加速源 - 优先本地,没有再从上游拉取 + from mirrors import get_mirror_handler + mirror_path = path.replace('mirrors/go/', '', 1) + if path == 'mirrors/go': + mirror_path = 'go' + handler_class = get_mirror_handler('go') + if handler_class: + mirror_config = self.config.get('mirrors', {}).get('go', {}) + base_dir = self.config.get('base_dir', './downloads') + storage_subdir = mirror_config.get('storage_dir', 'go') + mirror_handler = handler_class({ + 'base_dir': base_dir, + 'storage_dir': storage_subdir, + 'upstream_url': mirror_config.get('url', 'https://goproxy.cn') + }) + mirror_handler.handle_request(handler, mirror_path) + else: + handler.send_error(404, "Go mirror not configured") + return + elif path.startswith('mirrors/docker/') or path == 'mirrors/docker': + # Docker 加速源 - 优先本地,没有再从上游拉取 + from mirrors import get_mirror_handler + mirror_path = path.replace('mirrors/docker/', '', 1) + if path == 'mirrors/docker': + mirror_path = 'docker' + handler_class = get_mirror_handler('docker') + if handler_class: + mirror_config = self.config.get('mirrors', {}).get('docker', {}) + base_dir = self.config.get('base_dir', './downloads') + storage_subdir = mirror_config.get('storage_dir', 'docker') + mirror_handler = handler_class({ + 'base_dir': base_dir, + 'storage_dir': storage_subdir, + 'upstream_url': mirror_config.get('url', 'https://registry.hub.docker.com') + }) + mirror_handler.handle_request(handler, mirror_path) + else: + handler.send_error(404, "Docker mirror not configured") + return + + # PyPI包路径处理 - 处理 /pypi/packages/... 和 /pypi/web/... 路径 + # 这些路径来自镜像返回的HTML中的绝对链接 + elif path.startswith('pypi/packages/') or path.startswith('pypi/web/') or path.startswith('pypi/simple/'): + import re + import sys + from mirrors import PyPIMirror + # 从Referer中提取镜像名称 + referer = handler.headers.get('Referer', '') + mirror_name = None + if 'mirrors/' in referer: + match = re.search(r'mirrors/([^/]+)', referer) + if match: + mirror_name = match.group(1) + + # 优先使用Referer中指定的镜像,否则查找任意可用的pypi类型镜像 + mirrors_config = self.config.get('mirrors', {}) + if mirror_name and mirror_name in mirrors_config: + pypi_config = mirrors_config[mirror_name] + else: + # 尝试查找任意pypi类型的镜像(按优先级:pypi-cn, pypi) + pypi_config = None + for pref_name in ['pypi-cn', 'pypi']: + if pref_name in mirrors_config and mirrors_config[pref_name].get('type') == 'pypi': + pypi_config = mirrors_config[pref_name] + break + if not pypi_config: + # 尝试查找任意pypi类型的镜像 + for name, cfg in mirrors_config.items(): + if cfg.get('type') == 'pypi': + pypi_config = cfg + break + base_dir = self.config.get('base_dir', './downloads') + storage_subdir = pypi_config.get('storage_dir', 'pypi') if pypi_config else 'pypi' + pypi_handler = PyPIMirror({ + 'base_dir': base_dir, + 'storage_dir': storage_subdir, + 'upstream_url': pypi_config.get('url', 'https://pypi.org') if pypi_config else 'https://pypi.org' + }) + pypi_handler.handle_request(handler, path) + return + + # 通用镜像处理 - 支持自定义镜像名称如 pypi-cn, npm-cn 等 + elif path.startswith('mirrors/') and method == 'PUT': + # 更新自定义加速源 + mirror_name = path[8:] + self.api_update_mirror(handler, mirror_name) + return + elif path.startswith('mirrors/') and method == 'DELETE': + # 删除自定义加速源 + mirror_name = path[8:] + self.api_delete_mirror(handler, mirror_name) + return + elif path == 'mirrors' and method == 'POST': + # 添加自定义加速源 + self.api_add_mirror(handler) + return + + # 通用镜像处理 - 支持自定义镜像名称如 pypi-cn, npm-cn 等 + elif path.startswith('mirrors/'): + + from mirrors import HttpMirror, get_mirror_handler, get_default_upstream + # 解析镜像名称和路径 + # 格式: mirrors/{mirror_name}/... 或 mirrors/{mirror_name} + # 例如: mirrors/pypi-cn/simple/setuptools -> mirror_name=pypi-cn, mirror_path=simple/setuptools + parts = path[8:].split('/', 1) + mirror_name = parts[0] + # 去掉 mirror_name 前缀,只保留后面的路径 + if len(parts) > 1: + full_path = parts[1] + # 去掉路径开头的 mirror_name(如 pypi-cn/) + if full_path.startswith(mirror_name + '/'): + mirror_path = full_path[len(mirror_name)+1:] + else: + mirror_path = full_path + else: + mirror_path = mirror_name + + # 获取镜像配置 + mirror_config = self.config.get('mirrors', {}).get(mirror_name, {}) + + if not mirror_config: + handler.send_error(404, f"Mirror '{mirror_name}' not configured") + return + + # 获取镜像类型 + mirror_type = mirror_config.get('type', 'http') + + # 调试 + debug_file = '/tmp/pypi_debug.log' + with open(debug_file, 'a') as f: + f.write(f"[V2] mirror_type={mirror_type}, handler_class={get_mirror_handler(mirror_type)}\n") + + # 获取处理器类 + handler_class = get_mirror_handler(mirror_type) + if not handler_class: + handler_class = HttpMirror + + # 创建处理器实例 + base_dir = self.config.get('base_dir', './downloads') + storage_subdir = mirror_config.get('storage_dir', mirror_name) + # 拼接完整存储路径 + storage_dir = os.path.join(base_dir, storage_subdir) + + # 根据镜像类型确定上游URL + default_upstream = 'https://pypi.org/simple' + if mirror_type in ('pypi', 'pip', 'pipenv', 'poetry'): + default_upstream = 'https://pypi.org/simple' + elif mirror_type == 'npm': + default_upstream = 'https://registry.npmjs.org' + elif mirror_type == 'go': + default_upstream = 'https://goproxy.cn' + elif mirror_type == 'docker': + default_upstream = 'https://registry.hub.docker.com' + else: + default_upstream = get_default_upstream(mirror_type) + + mirror_handler = handler_class({ + 'base_dir': base_dir, + 'storage_dir': storage_dir, + 'upstream_url': mirror_config.get('url', default_upstream), + 'type': mirror_type, + 'cache_enabled': mirror_config.get('cache_enabled', True), + 'cache_ttl': mirror_config.get('cache_ttl', 3600) + }) + mirror_handler.handle_request(handler, mirror_path) + return + + # 用户管理 + elif path == 'user/login': + if method == 'POST': + self.api_login(handler) + else: + handler.send_error(405) + return + elif path == 'user/password': + if method == 'POST': + self.api_change_password(handler) + else: + handler.send_error(405) + return + elif path == 'user/login-logs': + if method == 'GET': + self.api_get_login_logs(handler, query_params) + else: + handler.send_error(405) + return + + # 配置文件管理 + elif path == 'config': + if method == 'GET': + self.api_get_config(handler) + elif method == 'PUT': + self.api_save_config(handler) + else: + handler.send_error(405) + return + elif path == 'config/reload': + if method == 'POST': + self.api_reload_config(handler) + else: + handler.send_error(405) + return + elif path == 'config/changes': + if method == 'GET': + self.api_get_config_changes(handler) + else: + handler.send_error(405) + return + + # 告警管理 + elif path == 'alerts': + if method == 'GET': + self.api_get_alerts(handler, query_params) + else: + handler.send_error(405) + return + elif path.startswith('alerts/') and '/acknowledge' in path: + alert_id = path[7:].split('/')[0] + if method == 'POST': + self.api_acknowledge_alert(handler, alert_id) + else: + handler.send_error(405) + return + elif path == 'alerts/clear': + if method == 'POST': + self.api_clear_alerts(handler) + else: + handler.send_error(405) + return + elif path == 'alerts/test': + if method == 'POST': + self.api_test_alert(handler) + else: + handler.send_error(405) + return + elif path == 'alerts/config': + if method == 'GET': + self.api_get_alert_config(handler) + elif method == 'PUT': + self.api_save_alert_config(handler) + else: + handler.send_error(405) + return + + # Prometheus 指标 + elif path == 'metrics': + if method == 'GET': + self.api_get_metrics(handler) + else: + handler.send_error(405) + return + + # 下载趋势 + elif path == 'stats/download-trend': + if method == 'GET': + self.api_get_download_trend(handler, query_params) + else: + handler.send_error(405) + return + + # 按周期下载统计 + elif path == 'stats/download-by-period': + if method == 'GET': + self.api_get_download_by_period(handler, query_params) + else: + handler.send_error(405) + return + + # 下载排行 + elif path == 'stats/rank': + if method == 'GET': + self.api_get_download_rank(handler, query_params) + else: + handler.send_error(405) + return + + # 热门缓存 + elif path == 'cache/popular': + if method == 'GET': + self.api_get_hot_cache(handler, query_params) + else: + handler.send_error(405) + return + + # 最近活动 + elif path == 'activity': + if method == 'GET': + self.api_get_recent_activity(handler, query_params) + else: + handler.send_error(405) + return + + # 服务器重启管理 + elif path == 'server/restart': + if method == 'GET': + self.api_get_restart_status(handler) + elif method == 'POST': + self.api_graceful_restart(handler, query_params) + else: + handler.send_error(405) + return + elif path == 'server/restart/confirm': + if method == 'POST': + self.api_confirm_restart(handler) + else: + handler.send_error(405) + return + elif path == 'server/restart/immediate': + if method == 'POST': + self.api_immediate_restart(handler) + else: + handler.send_error(405) + return + elif path == 'server/restart/pending': + if method == 'GET': + self.api_get_pending_requests(handler) + else: + handler.send_error(405) + return + elif path == 'server/restart/history': + if method == 'GET': + self.api_get_restart_history(handler) + else: + handler.send_error(405) + return + elif path == 'server/restart/config': + if method == 'GET': + self.api_get_restart_config(handler) + elif method == 'PUT': + self.api_update_restart_config(handler) + else: + handler.send_error(405) + return + + # 缓存预热管理 + elif path == 'cache/prewarm': + if method == 'GET': + self.api_get_prewarm_status(handler) + elif method == 'POST': + self.api_run_prewarm(handler, query_params) + else: + handler.send_error(405) + return + elif path == 'cache/prewarm/stats': + if method == 'GET': + self.api_get_prewarm_stats(handler) + else: + handler.send_error(405) + return + elif path == 'cache/prewarm/items': + if method == 'GET': + self.api_get_prewarm_items(handler, query_params) + elif method == 'POST': + self.api_add_prewarm_items(handler) + else: + handler.send_error(405) + return + elif path == 'cache/prewarm/history': + if method == 'GET': + self.api_get_prewarm_history(handler) + else: + handler.send_error(405) + return + elif path == 'cache/prewarm/clear': + if method == 'POST': + self.api_clear_prewarm_queue(handler) + else: + handler.send_error(405) + return + elif path == 'cache/prewarm/popular': + if method == 'GET': + self.api_get_popular_items(handler, query_params) + elif method == 'POST': + self.api_add_popular_items(handler, query_params) + else: + handler.send_error(405) + return + elif path == 'cache/prewarm/config': + if method == 'GET': + self.api_get_prewarm_config(handler) + elif method == 'PUT': + self.api_save_prewarm_config(handler) + else: + handler.send_error(405) + return + + # API 文档 + elif path == 'api-docs.json': + if method == 'GET': + self.api_get_api_docs(handler) + else: + handler.send_error(405) + return + elif path == 'api-docs.yaml': + if method == 'GET': + self.api_get_api_docs(handler, format='yaml') + else: + handler.send_error(405) + return + elif path == 'api-docs/generate': + if method == 'POST': + self.api_generate_api_docs(handler) + else: + handler.send_error(405) + return + + # 服务器信息 + elif path == 'server/info': + if method == 'GET': + self.api_get_server_info(handler) + else: + handler.send_error(405) + return + + # v2 端点都没匹配到,尝试调用 APIv1 + try: + return super().handle_request(handler, method, path, query_params) + except: + pass + + # 404 - 未找到端点 + handler.send_error(404) + + # ==================== 认证 API ==================== + + def api_verify_auth(self, handler): + """验证认证状态""" + # 调试模式输出 (debug-v2) + if handler._is_debug_enabled('v2'): + auth_header = handler.headers.get('Authorization', '') + api_key = handler.headers.get('X-API-Key', '') + auth_header_display = auth_header[:30] + '...' if len(auth_header) > 30 else auth_header + api_key_display = api_key[:20] + '...' if len(api_key) > 20 else api_key + msg = f"\n=== DEBUG api_verify_auth ===\n auth_header: '{auth_header_display}'\n api_key: '{api_key_display}'\n auth_type: {self.config.get('auth_type', 'none')}" + handler._debug_log('v2', msg, '\033[32m') + + auth_manager = getattr(handler, 'auth_manager', None) + + # 如果 auth_type 为 none,则允许访问 + if self.config.get('auth_type', 'none') == 'none': + handler.send_json_response({ + "valid": True, + "level": "admin", + "user_id": "anonymous", + "permissions": ["admin:*", "files:*", "sync:*", "keys:*"], + "expires_at": None, + "message": "Auth disabled - full access" + }) + return + + # 使用 auth_manager 验证(支持 Bearer Token、API Key、Cookie、Query Parameter) + if auth_manager: + result = auth_manager.validate_request(handler) + if result.get('authenticated'): + handler.send_json_response({ + "valid": True, + "level": result.get('level', 'admin'), + "user_id": result.get('user_id', result.get('key_id', 'unknown')), + "permissions": result.get('permissions', ["*"]), + "method": result.get('method', 'unknown') + }) + return + + # 无效认证 + handler.send_json_response({ + "valid": False, + "error": "Invalid or expired credentials", + "auth_type": self.config.get('auth_type', 'token') + }, 401) + + # ==================== 增强搜索API ==================== + + def api_search_files_enhanced(self, handler, query_params): + """增强的文件搜索 - 支持更多选项""" + search_term = query_params.get('q', [''])[0].lower() + search_type = query_params.get('type', ['all'])[0] + search_mode = query_params.get('mode', ['fuzzy'])[0] # fuzzy, exact, regex + max_results = int(query_params.get('limit', ['100'])[0]) + offset = int(query_params.get('offset', ['0'])[0]) + include_content = query_params.get('include_content', ['false'])[0].lower() == 'true' + + if not search_term: + handler.send_json_response({"error": "No search term provided"}, 400) + return + + results = [] + search_time = 0 + import time as time_module + start_time = time_module.time() + + for root, dirs, files in os.walk(self.config['base_dir']): + if search_type in ['all', 'dir']: + for dir_name in dirs: + match = False + if search_mode == 'exact': + match = search_term == dir_name.lower() + elif search_mode == 'regex': + try: + match = re.search(search_term, dir_name.lower()) is not None + except re.error: + match = False + else: # fuzzy + match = search_term in dir_name.lower() + + if match: + full_path = os.path.join(root, dir_name) + rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") + try: + mtime = os.path.getmtime(full_path) + result = { + "name": dir_name, + "path": rel_path + "/", + "type": "directory", + "size": 0, + "modified": datetime.fromtimestamp(mtime).isoformat(), + "match_score": self._calculate_match_score(dir_name, search_term, search_mode) + } + if include_content: + result['item_count'] = len(os.listdir(full_path)) + results.append(result) + except OSError: + continue + + if search_type in ['all', 'file']: + for file_name in files: + match = False + if search_mode == 'exact': + match = search_term == file_name.lower() + elif search_mode == 'regex': + try: + match = re.search(search_term, file_name.lower()) is not None + except re.error: + match = False + else: # fuzzy + match = search_term in file_name.lower() + + if match: + full_path = os.path.join(root, file_name) + rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") + try: + size = os.path.getsize(full_path) + mtime = os.path.getmtime(full_path) + mime_type, _ = os.path.splitext(file_name) + mime_type = mime_type[1:].lower() if mime_type else 'unknown' + + result = { + "name": file_name, + "path": rel_path, + "type": mime_type, + "size": size, + "size_formatted": self.format_file_size(size), + "modified": datetime.fromtimestamp(mtime).isoformat(), + "match_score": self._calculate_match_score(file_name, search_term, search_mode) + } + + if include_content: + if mime_type in ['txt', 'log', 'md', 'json', 'xml', 'html']: + try: + with open(full_path, 'r', encoding='utf-8', errors='ignore') as f: + content = f.read(5000) + result['content_preview'] = content + except: + pass + + results.append(result) + except OSError: + continue + + if len(results) >= max_results + offset: + break + + search_time = time_module.time() - start_time + + # 按匹配分数排序 + results.sort(key=lambda x: x.get('match_score', 0), reverse=True) + + paginated_results = results[offset:offset + max_results] + + handler.send_json_response({ + "query": search_term, + "search_type": search_type, + "search_mode": search_mode, + "total_count": len(results), + "returned_count": len(paginated_results), + "offset": offset, + "limit": max_results, + "search_time": round(search_time, 3), + "results": paginated_results + }) + + def api_search_by_tag(self, handler, query_params): + """按标签搜索文件""" + tag = query_params.get('tag', [''])[0] + if not tag: + handler.send_json_response({"error": "No tag specified"}, 400) + return + + # 这里假设有一个标签存储系统 + # 实际实现需要维护文件标签数据库 + results = [] + handler.send_json_response({ + "tag": tag, + "total_count": 0, + "results": results + }) + + def api_search_by_date(self, handler, query_params): + """按日期范围搜索文件""" + start_date = query_params.get('start', [''])[0] + end_date = query_params.get('end', [''])[0] + search_type = query_params.get('type', ['all'])[0] + + if not start_date: + handler.send_json_response({"error": "Start date required"}, 400) + return + + try: + start_dt = datetime.fromisoformat(start_date) + end_dt = datetime.fromisoformat(end_date) if end_date else datetime.now() + except ValueError: + handler.send_json_response({"error": "Invalid date format"}, 400) + return + + results = [] + + for root, dirs, files in os.walk(self.config['base_dir']): + if search_type in ['all', 'dir']: + for dir_name in dirs: + full_path = os.path.join(root, dir_name) + try: + mtime = datetime.fromtimestamp(os.path.getmtime(full_path)) + if start_dt <= mtime <= end_dt: + rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") + results.append({ + "name": dir_name, + "path": rel_path + "/", + "type": "directory", + "modified": mtime.isoformat() + }) + except OSError: + continue + + if search_type in ['all', 'file']: + for file_name in files: + full_path = os.path.join(root, file_name) + try: + mtime = datetime.fromtimestamp(os.path.getmtime(full_path)) + if start_dt <= mtime <= end_dt: + rel_path = os.path.relpath(full_path, self.config['base_dir']).replace("\\", "/") + results.append({ + "name": file_name, + "path": rel_path, + "type": "file", + "modified": mtime.isoformat() + }) + except OSError: + continue + + if len(results) > 1000: + break + + handler.send_json_response({ + "start_date": start_date, + "end_date": end_date.isoformat(), + "total_count": len(results), + "results": results + }) + + def _calculate_match_score(self, text, search_term, mode): + """计算匹配分数""" + if mode == 'exact': + return 100 if text.lower() == search_term else 0 + elif mode == 'regex': + try: + return 80 if re.search(search_term, text.lower()) else 0 + except: + return 0 + else: # fuzzy + text_lower = text.lower() + if search_term == text_lower: + return 100 + elif text_lower.startswith(search_term): + return 90 + elif text_lower.endswith(search_term): + return 80 + elif search_term in text_lower: + return 70 + return 0 + + # ==================== 增强统计API ==================== + + def api_get_stats_detailed(self, handler): + """获取详细统计信息""" + import mimetypes + + total_files = 0 + total_dirs = 0 + total_size = 0 + file_types = {} + size_distribution = { + "small": 0, # < 1MB + "medium": 0, # 1MB - 100MB + "large": 0, # 100MB - 1GB + "xlarge": 0 # > 1GB + } + oldest_file = None + newest_file = None + largest_file = None + + for root, dirs, files in os.walk(self.config['base_dir']): + total_dirs += len(dirs) + total_files += len(files) + for filename in files: + try: + file_path = os.path.join(root, filename) + size = os.path.getsize(file_path) + mtime = os.path.getmtime(file_path) + total_size += size + + mime_type, _ = mimetypes.guess_type(file_path) + if mime_type is None: + mime_type = "application/octet-stream" + file_types[mime_type] = file_types.get(mime_type, 0) + 1 + + # 大小分布 + if size < 1024 * 1024: + size_distribution["small"] += 1 + elif size < 100 * 1024 * 1024: + size_distribution["medium"] += 1 + elif size < 1024 * 1024 * 1024: + size_distribution["large"] += 1 + else: + size_distribution["xlarge"] += 1 + + # 最旧/最新文件 + if oldest_file is None or mtime < oldest_file[1]: + oldest_file = (filename, mtime) + if newest_file is None or mtime > newest_file[1]: + newest_file = (filename, mtime) + + # 最大文件 + if largest_file is None or size > largest_file[1]: + largest_file = (filename, size) + + except OSError: + continue + + handler.send_json_response({ + "summary": { + "total_files": total_files, + "total_dirs": total_dirs, + "total_size": total_size, + "total_size_formatted": self.format_file_size(total_size) + }, + "file_types": dict(sorted(file_types.items(), key=lambda x: x[1], reverse=True)), + "size_distribution": size_distribution, + "extremes": { + "oldest_file": { + "name": oldest_file[0] if oldest_file else None, + "modified": datetime.fromtimestamp(oldest_file[1]).isoformat() if oldest_file else None + }, + "newest_file": { + "name": newest_file[0] if newest_file else None, + "modified": datetime.fromtimestamp(newest_file[1]).isoformat() if newest_file else None + }, + "largest_file": { + "name": largest_file[0] if largest_file else None, + "size": largest_file[1] if largest_file else None, + "size_formatted": self.format_file_size(largest_file[1]) if largest_file else None + } + }, + "updated": datetime.now().isoformat() + }) + + def api_get_trending_files(self, handler, query_params): + """获取热门文件(按下载次数)""" + limit = int(query_params.get('limit', ['20'])[0]) + time_range = query_params.get('range', ['24h'])[0] # 24h, 7d, 30d + + # 获取下载统计 + stats = handler.load_stats() + + # 转换为列表并排序 + trending = [] + for filepath, count in stats.items(): + full_path = os.path.join(self.config['base_dir'], filepath) + if os.path.exists(full_path) and os.path.isfile(full_path): + try: + info = os.stat(full_path) + trending.append({ + "path": filepath, + "name": os.path.basename(filepath), + "size": info.st_size, + "size_formatted": self.format_file_size(info.st_size), + "modified": datetime.fromtimestamp(info.st_mtime).isoformat(), + "downloads": count + }) + except OSError: + continue + + # 按下载次数排序 + trending.sort(key=lambda x: x['downloads'], reverse=True) + trending = trending[:limit] + + handler.send_json_response({ + "time_range": time_range, + "limit": limit, + "trending": trending + }) + + def api_get_download_trend(self, handler, query_params): + """获取下载趋势(按天统计)""" + days = int(query_params.get('days', [7])[0]) + days = min(max(days, 1), 90) # 限制 1-90 天 + + trend = [] + + # 从数据库获取下载记录进行统计 + db = self.get_db() + if db: + try: + from datetime import timedelta + from core.database import DownloadRecord + + now = datetime.now() + start_time = (now - timedelta(days=days)).timestamp() + + with db.session() as session: + records = session.query(DownloadRecord).filter( + DownloadRecord.download_time >= start_time, + DownloadRecord.success == True + ).all() + + # 按日期聚合 + daily_stats = {} + for r in records: + if r.download_time: + date = datetime.fromtimestamp(r.download_time).strftime('%Y-%m-%d') + daily_stats[date] = daily_stats.get(date, 0) + 1 + + # 填充所有日期 + for i in range(days): + date = (now - timedelta(days=days - 1 - i)).strftime('%Y-%m-%d') + trend.append({ + "date": date, + "downloads": daily_stats.get(date, 0) + }) + except Exception as e: + print(f"Error getting download trend from database: {e}") + + # 如果没有数据库数据,返回模拟数据 + if not trend: + now = datetime.now() + for i in range(days): + date = (now - timedelta(days=days - 1 - i)).strftime('%m-%d') + trend.append({ + "date": date, + "downloads": 0 + }) + + handler.send_json_response({ + "days": days, + "trend": trend + }) + + def api_get_download_by_period(self, handler, query_params): + """按周期获取下载统计(年/月/日)""" + from datetime import datetime, timedelta + from sqlalchemy import func + + period = query_params.get('period', ['day'])[0] + year = int(query_params.get('year', [datetime.now().year])[0]) + month = int(query_params.get('month', [datetime.now().month])[0]) + + # 限制 period 值 + if period not in ['year', 'month', 'day']: + period = 'day' + + result = { + "period": period, + "year": year, + "month": month, + "data": [] + } + + db = self.get_db() + if db: + try: + from core.database import DownloadRecord + now = datetime.now() + + with db.session() as session: + if period == 'year': + # 按月统计全年数据 + start_time = datetime(year, 1, 1).timestamp() + end_time = datetime(year + 1, 1, 1).timestamp() + + records = session.query(DownloadRecord).filter( + DownloadRecord.download_time >= start_time, + DownloadRecord.download_time < end_time, + DownloadRecord.success == True + ).all() + + monthly_stats = {m: 0 for m in range(1, 13)} + total = 0 + for r in records: + if r.download_time: + dt = datetime.fromtimestamp(r.download_time) + monthly_stats[dt.month] += 1 + total += 1 + + result["data"] = [ + {"label": f"{m}月", "value": monthly_stats[m], "month": m} + for m in range(1, 13) + ] + result["total"] = total + + elif period == 'month': + # 按日统计当月数据 + start_time = datetime(year, month, 1).timestamp() + if month == 12: + end_time = datetime(year + 1, 1, 1).timestamp() + else: + end_time = datetime(year, month + 1, 1).timestamp() + + records = session.query(DownloadRecord).filter( + DownloadRecord.download_time >= start_time, + DownloadRecord.download_time < end_time, + DownloadRecord.success == True + ).all() + + days_in_month = (datetime(year, month + 1, 1) - timedelta(days=1)).day if month < 12 else 31 + daily_stats = {d: 0 for d in range(1, days_in_month + 1)} + total = 0 + for r in records: + if r.download_time: + dt = datetime.fromtimestamp(r.download_time) + daily_stats[dt.day] += 1 + total += 1 + + result["data"] = [ + {"label": f"{d}日", "value": daily_stats[d], "day": d} + for d in range(1, days_in_month + 1) + ] + result["total"] = total + + else: # day - 按日统计当月数据 + # 获取当月第一天和最后一天 + if month == 12: + start_date = datetime(year, month, 1) + end_date = datetime(year + 1, 1, 1) + else: + start_date = datetime(year, month, 1) + end_date = datetime(year, month + 1, 1) + + start_time = start_date.timestamp() + end_time = end_date.timestamp() + + records = session.query(DownloadRecord).filter( + DownloadRecord.download_time >= start_time, + DownloadRecord.download_time < end_time, + DownloadRecord.success == True + ).all() + + # 按天统计 + days_in_month = (end_date - timedelta(days=1)).day + daily_stats = {d: 0 for d in range(1, days_in_month + 1)} + total = 0 + for r in records: + if r.download_time: + dt = datetime.fromtimestamp(r.download_time) + daily_stats[dt.day] += 1 + total += 1 + + result["data"] = [ + {"label": f"{d}日", "value": daily_stats[d], "day": d} + for d in range(1, days_in_month + 1) + ] + result["total"] = total + + # 获取历史年份列表 + if period == 'year': + with db.session() as session: + # SQLite 不支持 from_unixtime,使用 Python 处理 + records = session.query(DownloadRecord.download_time).filter( + DownloadRecord.success == True + ).distinct().all() + years_set = set() + for (ts,) in records: + if ts: + dt = datetime.fromtimestamp(ts) + years_set.add(dt.year) + result["available_years"] = sorted(list(years_set), reverse=True)[:10] + + except Exception as e: + print(f"Error getting download by period: {e}") + result["error"] = str(e) + + # 如果没有数据,返回空数据结构 + if not result.get("data"): + if period == 'year': + result["data"] = [{"label": f"{m}月", "value": 0, "month": m} for m in range(1, 13)] + elif period == 'month': + result["data"] = [{"label": f"{d}日", "value": 0, "day": d} for d in range(1, 32)] + else: + result["data"] = [{"label": f"{h}:00", "value": 0, "hour": h} for h in range(24)] + + handler.send_json_response(result) + + def api_get_download_rank(self, handler, query_params): + """获取下载排行 TOP 20""" + import traceback + import os + limit = int(query_params.get('limit', [20])[0]) + limit = min(max(limit, 1), 100) + + rank_data = [] + errors = [] + + # 从数据库获取下载记录进行统计 + db = self.get_db() + if db: + try: + from core.database import DownloadRecord + + with db.session() as session: + # 按文件路径分组统计下载次数 + from sqlalchemy import func + results = session.query( + DownloadRecord.file_path, + func.count(DownloadRecord.id).label('download_count') + ).filter( + DownloadRecord.success == True + ).group_by( + DownloadRecord.file_path + ).order_by( + func.count(DownloadRecord.id).desc() + ).limit(limit).all() + + total_downloads = sum(r[1] for r in results) if results else 0 + + # 获取 base_dir + base_dir = self.config.get('base_dir', './downloads') + base_dir = os.path.abspath(base_dir) if base_dir else './downloads' + + for idx, (file_path, count) in enumerate(results, 1): + # 获取文件大小 + file_size = 0 + full_path = os.path.join(base_dir, file_path) if file_path else '' + if full_path and os.path.exists(full_path): + try: + file_size = os.path.getsize(full_path) + except OSError: + pass + + # 计算占比 + percentage = (count / total_downloads * 100) if total_downloads > 0 else 0 + + rank_data.append({ + "rank": idx, + "file": file_path or 'Unknown', + "downloads": count, + "size": file_size, + "percentage": round(percentage, 1) + }) + except Exception as e: + errors.append(str(e)) + traceback.print_exc() + + handler.send_json_response({ + "rank": rank_data, + "count": len(rank_data), + "total_records": len(rank_data), + "_debug": {"errors": errors} if errors else {} + }) + + def api_get_hot_cache(self, handler, query_params): + """获取热门缓存文件""" + import traceback + import os + limit = int(query_params.get('limit', [20])[0]) + limit = min(max(limit, 1), 100) + + cache_data = [] + errors = [] + + # 从数据库获取缓存访问记录 + db = self.get_db() + if db: + try: + from core.database import DownloadRecord + + with db.session() as session: + from sqlalchemy import func + # 获取被访问过的缓存文件(按访问次数排序) + results = session.query( + DownloadRecord.file_path, + func.count(DownloadRecord.id).label('access_count'), + func.max(DownloadRecord.download_time).label('last_access') + ).filter( + DownloadRecord.success == True + ).group_by( + DownloadRecord.file_path + ).order_by( + func.count(DownloadRecord.id).desc() + ).limit(limit).all() + + # 获取 base_dir + base_dir = self.config.get('base_dir', './downloads') + base_dir = os.path.abspath(base_dir) if base_dir else './downloads' + + for idx, (file_path, access_count, last_access) in enumerate(results, 1): + # 获取文件信息 + full_path = os.path.join(base_dir, file_path) if file_path else '' + file_size = 0 + if full_path and os.path.exists(full_path): + try: + file_size = os.path.getsize(full_path) + except OSError: + pass + + cache_data.append({ + "rank": idx, + "file": file_path or 'Unknown', + "access_count": access_count, + "size": file_size, + "last_access": datetime.fromtimestamp(last_access).isoformat() if last_access else '-' + }) + except Exception as e: + errors.append(str(e)) + traceback.print_exc() + + handler.send_json_response({ + "cache": cache_data, + "count": len(cache_data), + "_debug": {"errors": errors} if errors else {} + }) + + def get_db(self): + """获取数据库实例""" + from core.database import get_db + try: + return get_db() + except Exception: + return None + + # ==================== 文件元数据API ==================== + + def api_get_file_metadata(self, handler, filename): + """获取文件元数据""" + full_path = os.path.join(self.config['base_dir'], filename) + if not is_safe_path(self.config['base_dir'], full_path): + handler.send_json_response({"error": "Access denied"}, 403) + return + if not os.path.exists(full_path): + handler.send_json_response({"error": "File not found"}, 404) + return + + # 这里可以从单独的元数据文件中读取 + metadata_file = full_path + '.meta' + metadata = {} + + if os.path.exists(metadata_file): + try: + with open(metadata_file, 'r', encoding='utf-8') as f: + metadata = json.load(f) + except: + pass + + # 添加基本文件信息 + import os + stat = os.stat(full_path) + metadata['_file_info'] = { + "size": stat.st_size, + "modified": datetime.fromtimestamp(stat.st_mtime).isoformat(), + "created": datetime.fromtimestamp(stat.st_ctime).isoformat() + } + + handler.send_json_response(metadata) + + def api_update_file_metadata(self, handler, filename): + """更新文件元数据""" + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "No metadata provided"}, 400) + return + + try: + metadata = json.loads(handler.rfile.read(content_length)) + full_path = os.path.join(self.config['base_dir'], filename) + if not is_safe_path(self.config['base_dir'], full_path): + handler.send_json_response({"error": "Access denied"}, 403) + return + metadata_file = full_path + '.meta' + + with open(metadata_file, 'w', encoding='utf-8') as f: + json.dump(metadata, f, ensure_ascii=False, indent=2) + + handler.send_json_response({"success": True}) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + def api_get_batch_metadata(self, handler, query_params): + """批量获取文件元数据""" + paths = query_params.get('paths', []) + if not paths: + handler.send_json_response({"error": "No file paths provided"}, 400) + return + + results = {} + for path in paths: + full_path = os.path.join(self.config['base_dir'], path) + if not is_safe_path(self.config['base_dir'], full_path): + continue + if os.path.exists(full_path): + metadata_file = full_path + '.meta' + metadata = {} + if os.path.exists(metadata_file): + try: + with open(metadata_file, 'r', encoding='utf-8') as f: + metadata = json.load(f) + except: + pass + results[path] = metadata + + handler.send_json_response(results) + + def api_update_batch_metadata(self, handler): + """批量更新文件元数据""" + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "No data provided"}, 400) + return + + try: + data = json.loads(handler.rfile.read(content_length)) + results = {} + + for path, metadata in data.items(): + full_path = os.path.join(self.config['base_dir'], path) + if not is_safe_path(self.config['base_dir'], full_path): + results[path] = {"success": False, "error": "Access denied"} + continue + metadata_file = full_path + '.meta' + try: + with open(metadata_file, 'w', encoding='utf-8') as f: + json.dump(metadata, f, ensure_ascii=False, indent=2) + results[path] = {"success": True} + except Exception as e: + results[path] = {"success": False, "error": str(e)} + + handler.send_json_response(results) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + # ==================== 文件版本控制API ==================== + + def api_get_file_versions(self, handler, filename): + """获取文件版本列表""" + # 这里可以实现版本控制系统 + # 简化实现:检查备份文件 + full_path = os.path.join(self.config['base_dir'], filename) + if not is_safe_path(self.config['base_dir'], full_path): + handler.send_json_response({"error": "Access denied"}, 403) + return + versions = [] + + # 查找备份文件 + dir_path = os.path.dirname(full_path) + base_name = os.path.basename(full_path) + + if os.path.exists(dir_path): + for item in os.listdir(dir_path): + if item.startswith(base_name) and item != base_name: + backup_path = os.path.join(dir_path, item) + try: + stat = os.stat(backup_path) + versions.append({ + "name": item, + "size": stat.st_size, + "modified": datetime.fromtimestamp(stat.st_mtime).isoformat() + }) + except OSError: + continue + + handler.send_json_response({ + "filename": filename, + "versions": versions + }) + + def api_create_file_version(self, handler, filename): + """创建文件版本(备份)""" + import time + full_path = os.path.join(self.config['base_dir'], filename) + if not is_safe_path(self.config['base_dir'], full_path): + handler.send_json_response({"error": "Access denied"}, 403) + return + if not os.path.exists(full_path): + handler.send_json_response({"error": "File not found"}, 404) + return + + # 创建备份 + timestamp = int(time.time()) + backup_path = f"{full_path}.v{timestamp}" + + try: + import shutil + shutil.copy2(full_path, backup_path) + handler.send_json_response({ + "success": True, + "version": backup_path + }) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + # ==================== 缩略图API ==================== + + def api_get_file_thumbnail(self, handler, filename, query_params): + """获取文件缩略图""" + import mimetypes + + full_path = os.path.join(self.config['base_dir'], filename) + if not is_safe_path(self.config['base_dir'], full_path): + handler.send_json_response({"error": "Access denied"}, 403) + return + if not os.path.exists(full_path): + handler.send_json_response({"error": "File not found"}, 404) + return + + mime_type, _ = mimetypes.guess_type(full_path) + if not mime_type or not mime_type.startswith('image/'): + handler.send_json_response({"error": "Not an image file"}, 400) + return + + try: + width = min(max(int(query_params.get('width', ['200'])[0]), 16), 2048) + height = min(max(int(query_params.get('height', ['200'])[0]), 16), 2048) + except ValueError: + width, height = 200, 200 + # 限制解码像素上限,防解压炸弹 + try: + from PIL import Image as _PILImage + _PILImage.MAX_IMAGE_PIXELS = 50_000_000 + except Exception: + pass + + try: + from PIL import Image + + with Image.open(full_path) as img: + img.thumbnail((width, height), Image.LANCZOS) + + import io + buffer = io.BytesIO() + img.save(buffer, format='JPEG', quality=85) + thumbnail_data = buffer.getvalue() + + handler.send_response(200) + handler.send_header("Content-Type", "image/jpeg") + handler.send_header("Content-Length", str(len(thumbnail_data))) + handler.send_header("Cache-Control", "public, max-age=86400") + handler.end_headers() + handler.wfile.write(thumbnail_data) + + except ImportError: + handler.send_json_response({"error": "PIL/Pillow not installed"}, 500) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + # ==================== 服务器监控API ==================== + + def api_get_realtime_stats(self, handler): + """获取实时服务器统计""" + try: + import psutil + result = { + "timestamp": datetime.now().isoformat(), + "cpu": { + "percent": psutil.cpu_percent(interval=0.1), + "count": psutil.cpu_count() + }, + "memory": { + "total": psutil.virtual_memory().total, + "available": psutil.virtual_memory().available, + "percent": psutil.virtual_memory().percent, + "used": psutil.virtual_memory().used, + "free": psutil.virtual_memory().free + }, + "disk": { + "total": psutil.disk_usage(self.config['base_dir']).total, + "used": psutil.disk_usage(self.config['base_dir']).used, + "free": psutil.disk_usage(self.config['base_dir']).free, + "percent": psutil.disk_usage(self.config['base_dir']).percent + } + } + # 尝试获取网络数据,失败时忽略 + try: + result["network"] = { + "connections": len(psutil.net_connections()), + "io": psutil.net_io_counters()._asdict() if psutil.net_io_counters() else None + } + except (PermissionError, OSError): + result["network"] = {"connections": 0, "io": None, "note": "Permission denied"} + + # 尝试获取 CPU 频率,失败时忽略 + try: + result["cpu"]["freq"] = psutil.cpu_freq()._asdict() if psutil.cpu_freq() else None + except (PermissionError, OSError): + result["cpu"]["freq"] = None + + handler.send_json_response(result) + except ImportError: + handler.send_json_response({"error": "psutil not installed"}, 500) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + def api_get_monitor_history(self, handler, query_params): + """获取历史监控数据""" + hours = int(query_params.get('hours', ['24'])[0]) + + # 尝试从数据库获取 + if self.db_enabled and self.db: + try: + records = self.db.get_monitor_history(hours) + data = [r.to_dict() for r in records] + handler.send_json_response({ + "hours": hours, + "data": data, + "source": "database" + }) + return + except Exception as e: + print(f"Error getting monitor history from database: {e}") + + # 回退到空数据 + handler.send_json_response({ + "hours": hours, + "data": [], + "source": "none" + }) + + # ==================== Webhook API ==================== + + def api_list_webhooks(self, handler): + """列出所有webhook""" + # 从数据库获取 + if self.db_enabled and self.db: + try: + webhooks = self.db.get_webhooks() + handler.send_json_response({ + "webhooks": [w.to_dict() for w in webhooks], + "count": len(webhooks), + "source": "database" + }) + return + except Exception as e: + print(f"Error getting webhooks from database: {e}") + + # 回退到空列表 + handler.send_json_response({ + "webhooks": [], + "count": 0, + "source": "none" + }) + + def api_create_webhook(self, handler): + """创建webhook""" + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "No webhook data provided"}, 400) + return + + try: + webhook_data = json.loads(handler.rfile.read(content_length)) + + # 验证必要字段 + name = webhook_data.get('name', '').strip() + url = webhook_data.get('url', '').strip() + + if not name: + handler.send_json_response({"error": "Webhook name is required"}, 400) + return + if not url: + handler.send_json_response({"error": "Webhook URL is required"}, 400) + return + + # 保存到数据库 + if self.db_enabled and self.db: + record = self.db.add_webhook( + name=name, + url=url, + events=webhook_data.get('events', []), + secret=webhook_data.get('secret'), + enabled=webhook_data.get('enabled', True) + ) + handler.send_json_response({ + "success": True, + "message": "Webhook created successfully", + "webhook": record.to_dict() + }) + else: + handler.send_json_response({ + "success": False, + "error": "Database not available" + }, 503) + + except json.JSONDecodeError: + handler.send_json_response({"error": "Invalid JSON format"}, 400) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + def api_get_webhook(self, handler, webhook_id): + """获取webhook详情""" + webhook_id = int(webhook_id) + + if self.db_enabled and self.db: + try: + webhook = self.db.get_webhook(webhook_id) + if webhook: + handler.send_json_response({ + "webhook": webhook.to_dict() + }) + else: + handler.send_json_response({ + "error": "Webhook not found" + }, 404) + return + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + handler.send_json_response({ + "error": "Database not available" + }, 503) + + def api_delete_webhook(self, handler, webhook_id): + """删除webhook""" + webhook_id = int(webhook_id) + + if self.db_enabled and self.db: + try: + success = self.db.delete_webhook(webhook_id) + if success: + handler.send_json_response({ + "success": True, + "message": "Webhook deleted successfully" + }) + else: + handler.send_json_response({ + "error": "Webhook not found" + }, 404) + return + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + handler.send_json_response({ + "success": False, + "error": "Database not available" + }, 503) + + def api_test_webhook(self, handler, webhook_id): + """测试webhook""" + webhook_id = int(webhook_id) + + if self.db_enabled and self.db: + try: + webhook = self.db.get_webhook(webhook_id) + if not webhook: + handler.send_json_response({ + "error": "Webhook not found" + }, 404) + return + + # 发送测试请求 + import urllib.request + import urllib.parse + + test_payload = { + "event": "test", + "timestamp": datetime.now().isoformat(), + "data": { + "message": "This is a test webhook from HYC下载站" + } + } + + try: + data = json.dumps(test_payload).encode('utf-8') + req = urllib.request.Request( + webhook.url, + data=data, + headers={ + 'Content-Type': 'application/json', + 'X-Webhook-Secret': webhook.secret or '' + }, + method='POST' + ) + with urllib.request.urlopen(req, timeout=10) as response: + handler.send_json_response({ + "success": True, + "webhook_id": webhook_id, + "test_result": "Webhook test successful", + "status_code": response.status + }) + except urllib.error.HTTPError as e: + handler.send_json_response({ + "success": False, + "webhook_id": webhook_id, + "error": f"HTTP Error: {e.code} {e.reason}" + }, 400) + except Exception as e: + handler.send_json_response({ + "success": False, + "webhook_id": webhook_id, + "error": str(e) + }, 400) + + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + else: + handler.send_json_response({ + "success": False, + "error": "Database not available" + }, 503) + + def api_get_webhook_deliveries(self, handler, webhook_id): + """获取 webhook 交付历史""" + webhook_id = int(webhook_id) + + # 验证 webhook 存在 + if self.db_enabled and self.db: + try: + webhook = self.db.get_webhook(webhook_id) + if not webhook: + handler.send_json_response({ + "error": "Webhook not found" + }, 404) + return + + # 获取交付历史 + deliveries = self.db.get_webhook_deliveries(webhook_id=webhook_id, limit=50) + + handler.send_json_response({ + "webhook_id": webhook_id, + "webhook_name": webhook.name, + "deliveries": [d.to_dict() for d in deliveries], + "count": len(deliveries) + }) + + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + else: + handler.send_json_response({ + "error": "Database not available" + }, 503) + + def api_get_webhook_stats(self, handler, webhook_id): + """获取 webhook 交付统计""" + webhook_id = int(webhook_id) + + if self.db_enabled and self.db: + try: + stats = self.db.get_webhook_stats(webhook_id) + + handler.send_json_response({ + "webhook_id": webhook_id, + "stats": stats + }) + + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + else: + handler.send_json_response({ + "error": "Database not available" + }, 503) + + def api_update_webhook(self, handler, webhook_id): + """更新 webhook 配置""" + webhook_id = int(webhook_id) + + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "No webhook data provided"}, 400) + return + + try: + webhook_data = json.loads(handler.rfile.read(content_length)) + + if self.db_enabled and self.db: + webhook = self.db.get_webhook(webhook_id) + if not webhook: + handler.send_json_response({ + "error": "Webhook not found" + }, 404) + return + + # 更新 webhook + updated = self.db.update_webhook( + webhook_id, + name=webhook_data.get('name', webhook.name), + url=webhook_data.get('url', webhook.url), + events=webhook_data.get('events'), + secret=webhook_data.get('secret'), + enabled=webhook_data.get('enabled', webhook.enabled) + ) + + if updated: + handler.send_json_response({ + "success": True, + "message": "Webhook updated successfully", + "webhook": updated.to_dict() + }) + else: + handler.send_json_response({ + "success": False, + "error": "Failed to update webhook" + }, 500) + else: + handler.send_json_response({ + "success": False, + "error": "Database not available" + }, 503) + + except json.JSONDecodeError: + handler.send_json_response({"error": "Invalid JSON format"}, 400) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + # ==================== 同步管理 API (真实数据) ==================== + + def api_get_sync_sources(self, handler): + """获取所有同步源(真实数据)""" + # 从 sync_manager 获取真实数据 + if hasattr(handler, 'sync_manager') and handler.sync_manager: + sources = getattr(handler.sync_manager, 'sync_sources', {}) + # 转换为数组格式 + sources_list = [] + for name, config in sources.items(): + item = {"name": name} + item.update(config) + sources_list.append(item) + handler.send_json_response({ + "sources": sources_list, + "count": len(sources_list) + }) + else: + handler.send_json_response({ + "sources": [], + "count": 0, + "message": "Sync manager not available" + }) + + def api_add_sync_source(self, handler): + """添加同步源""" + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "No data provided"}, 400) + return + + try: + data = json.loads(handler.rfile.read(content_length)) + + if hasattr(handler, 'sync_manager') and handler.sync_manager: + name = data.get('name') + config = data.get('config', {}) + + if not name: + handler.send_json_response({"error": "Source name required"}, 400) + return + + success = handler.sync_manager.add_source(name, config) + handler.send_json_response({ + "success": success, + "name": name + }) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + def api_start_sync(self, handler, source_name): + """启动同步""" + if hasattr(handler, 'sync_manager') and handler.sync_manager: + task_id = handler.sync_manager.start_sync(source_name) + if task_id: + handler.send_json_response({ + "success": True, + "task_id": task_id, + "source_name": source_name + }) + else: + handler.send_json_response({ + "success": False, + "error": "Failed to start sync", + "source_name": source_name + }, 500) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + def api_stop_sync(self, handler, source_name): + """停止同步""" + if hasattr(handler, 'sync_manager') and handler.sync_manager: + handler.sync_manager.stop_all_tasks_for_source(source_name) + handler.send_json_response({ + "success": True, + "source_name": source_name + }) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + def api_get_sync_status(self, handler, source_name): + """获取同步状态(真实数据)""" + # source_name 格式: source_name/status,需要提取 + if source_name.endswith('/status'): + source_name = source_name[:-7] + + if hasattr(handler, 'sync_manager') and handler.sync_manager: + status = handler.sync_manager.get_source_status(source_name) + handler.send_json_response(status) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + def api_get_sync_history(self, handler, query_params): + """获取同步历史""" + limit = int(query_params.get('limit', ['100'])[0]) + + if hasattr(handler, 'sync_manager') and handler.sync_manager: + history = handler.sync_manager.get_sync_history(limit=limit) + handler.send_json_response({ + "history": history, + "count": len(history) + }) + else: + handler.send_json_response({ + "history": [], + "count": 0 + }) + + # ==================== 定时任务调度 API ==================== + + def api_get_scheduled_tasks(self, handler): + """获取所有定时任务状态""" + if hasattr(handler, 'sync_manager') and handler.sync_manager: + scheduler = getattr(handler.sync_manager, 'task_scheduler', None) + if scheduler and hasattr(scheduler, 'get_all_tasks'): + tasks = scheduler.get_all_tasks() + handler.send_json_response({ + "tasks": tasks, + "count": len(tasks) + }) + return + + handler.send_json_response({ + "tasks": [], + "count": 0, + "message": "Scheduler not available" + }) + + def api_update_sync_schedule(self, handler, source_name): + """更新同步源的定时配置""" + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "No schedule data provided"}, 400) + return + + try: + data = json.loads(handler.rfile.read(content_length)) + + if hasattr(handler, 'sync_manager') and handler.sync_manager: + # 更新同步源的定时配置 + sync_manager = handler.sync_manager + + if hasattr(sync_manager, 'scheduled_syncs'): + schedule_config = data.get('schedule', {}) + if source_name not in sync_manager.scheduled_syncs: + sync_manager.scheduled_syncs[source_name] = {} + + sync_manager.scheduled_syncs[source_name] = { + 'type': schedule_config.get('type', 'interval'), + 'config': { + 'cron': schedule_config.get('cron'), + 'interval': schedule_config.get('interval', {}), + 'enabled': schedule_config.get('enabled', True) + } + } + + # 如果调度器已运行,更新任务 + if sync_manager.task_scheduler: + task_name = f"sync_{source_name}" + existing_task = sync_manager.task_scheduler.get_task(task_name) + if existing_task: + sync_manager.task_scheduler.update_task_config( + task_name, + sync_manager.scheduled_syncs[source_name]['config'] + ) + if schedule_config.get('enabled'): + sync_manager.task_scheduler.enable_task(task_name, True) + else: + sync_manager.task_scheduler.enable_task(task_name, False) + + handler.send_json_response({ + "success": True, + "source_name": source_name, + "schedule": sync_manager.scheduled_syncs.get(source_name, {}) + }) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + except json.JSONDecodeError: + handler.send_json_response({"error": "Invalid JSON format"}, 400) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + def api_run_sync_now(self, handler, source_name): + """立即触发同步(覆盖定时)""" + if hasattr(handler, 'sync_manager') and handler.sync_manager: + success = handler.sync_manager.start_sync(source_name) + handler.send_json_response({ + "success": success, + "source_name": source_name, + "message": "Sync started" if success else "Failed to start sync" + }) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + def api_sync_packages(self, handler): + """临时单次同步指定源的特定包""" + try: + # 读取请求体 + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "请求体不能为空"}, 400) + return + + body = handler.rfile.read(content_length) + data = json.loads(body.decode('utf-8')) + + # 获取参数 + source = data.get('source') + packages = data.get('packages', []) + + if not source: + handler.send_json_response({"error": "缺少 'source' 参数"}, 400) + return + + if not packages or not isinstance(packages, list): + handler.send_json_response({"error": "请提供有效的 'packages' 列表"}, 400) + return + + # 调用 sync_manager + if hasattr(handler, 'sync_manager') and handler.sync_manager: + result = handler.sync_manager.sync_packages(source, packages) + handler.send_json_response(result) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + except json.JSONDecodeError: + handler.send_json_response({"error": "Invalid JSON format"}, 400) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + def api_get_temp_sync_status(self, handler, source_name): + """获取临时同步状态""" + if hasattr(handler, 'sync_manager') and handler.sync_manager: + sync_manager = handler.sync_manager + # 查找匹配的临时同步任务 + temp_status = None + if hasattr(sync_manager, 'sync_status'): + for task_name, status in sync_manager.sync_status.items(): + if status.get('source_name') == source_name and status.get('is_temp_sync'): + temp_status = status + break + + if temp_status: + handler.send_json_response({ + "success": True, + "source": source_name, + "status": temp_status.get('status'), + "packages": temp_status.get('packages', []), + "files_synced": temp_status.get('files_synced', 0), + "total_files": temp_status.get('total_files', 0), + "last_sync": temp_status.get('last_sync'), + "error": temp_status.get('error') + }) + else: + handler.send_json_response({ + "success": False, + "error": f"没有找到 {source_name} 的临时同步任务" + }) + else: + handler.send_json_response({"error": "Sync manager not available"}, 500) + + # ==================== 系统监控 API (真实数据) ==================== + + def api_get_monitor_detailed(self, handler): + """获取详细监控数据""" + if hasattr(handler, 'monitor') and handler.monitor: + stats = handler.monitor.get_realtime_stats() + handler.send_json_response(stats) + else: + handler.send_json_response({ + "error": "Monitor not available", + "timestamp": datetime.now().isoformat() + }, 500) + + def api_get_monitor_history_detailed(self, handler, query_params): + """获取历史监控数据(详细版)""" + hours = int(query_params.get('hours', ['24'])[0]) + + # 尝试从数据库获取详细统计 + if self.db_enabled and self.db: + try: + stats = self.db.get_monitor_stats(hours) + history = self.db.get_monitor_history(hours) + handler.send_json_response({ + "hours": hours, + "stats": stats, + "history": [r.to_dict() for r in history], + "source": "database" + }) + return + except Exception as e: + print(f"Error getting monitor stats from database: {e}") + + # 回退到原有实现 + if hasattr(handler, 'monitor') and handler.monitor: + history = handler.monitor.get_monitor_history(hours) + handler.send_json_response({ + "hours": hours, + "data": history, + "source": "monitor" + }) + else: + handler.send_json_response({ + "hours": hours, + "data": [], + "source": "none" + }) + + def api_get_monitor_summary(self, handler): + """获取监控摘要""" + if hasattr(handler, 'monitor') and handler.monitor: + summary = handler.monitor.get_stats_summary() + handler.send_json_response(summary) + else: + handler.send_json_response({ + "status": "unavailable", + "message": "Monitor not available" + }, 500) + + def api_get_health_status(self, handler): + """获取健康状态""" + if hasattr(handler, 'monitor') and handler.monitor: + health = handler.monitor.get_health_status() + handler.send_json_response(health) + else: + handler.send_json_response({ + "status": "unknown", + "message": "Monitor not available" + }) + + # ==================== 镜像源健康检查 API ==================== + + def api_get_source_health(self, handler, query_params): + """获取镜像源健康状态""" + try: + from core.health_check import HealthChecker, HealthStatus + + mirrors = self.config.get('mirrors', {}) + checker = HealthChecker(self.config.get('health_check', {})) + + results = [] + for mirror_type, mirror_config in mirrors.items(): + if not isinstance(mirror_config, dict): + continue + + # 获取该镜像类型的所有源 + sources = mirror_config.get('sources', []) + for source_name in sources: + result = checker.check_source(source_name, { + 'url': self._get_source_url(mirror_type, source_name) + }) + results.append({ + 'mirror_type': mirror_type, + 'source_name': source_name, + 'status': result.status.value, + 'response_time_ms': round(result.response_time, 2), + 'http_status': result.http_status, + 'error': result.error_message, + 'success_rate': round(result.success_rate, 2), + 'last_check': result.last_check.isoformat() if result.last_check else None, + 'consecutive_failures': result.consecutive_failures + }) + + handler.send_json_response({ + 'sources': results, + 'count': len(results), + 'summary': checker.get_stats() + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def api_check_source(self, handler, source_name): + """手动触发单个源的健康检查""" + try: + from core.health_check import HealthChecker, HealthStatus + + checker = HealthChecker(self.config.get('health_check', {})) + mirrors = self.config.get('mirrors', {}) + + # 查找源对应的镜像类型 + mirror_type = None + source_url = '' + for mtype, mconfig in mirrors.items(): + if not isinstance(mconfig, dict): + continue + sources = mconfig.get('sources', []) + if source_name in sources: + mirror_type = mtype + source_config = mconfig.get('sources_config', {}).get(source_name, {}) + source_url = source_config.get('url', '') + break + + if not mirror_type: + handler.send_json_response({ + 'error': f"Source '{source_name}' not found" + }, 404) + return + + result = checker.check_source(source_name, {'url': source_url}) + + handler.send_json_response({ + 'source_name': source_name, + 'mirror_type': mirror_type, + 'status': result.status.value, + 'response_time_ms': round(result.response_time, 2), + 'http_status': result.http_status, + 'error': result.error_message, + 'success_rate': round(result.success_rate, 2), + 'last_check': result.last_check.isoformat() if result.last_check else None + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def api_get_failover_status(self, handler): + """获取故障切换状态""" + try: + from core.health_check import MirrorFailoverManager + + failover = MirrorFailoverManager(self.config) + failover.initialize() + + handler.send_json_response({ + 'failover_enabled': failover.failover_enabled, + 'active_sources': failover._active_source, + 'health_summary': failover.get_health_summary(), + 'failover_history': failover.get_failover_history() + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def api_trigger_failover(self, handler, mirror_type): + """手动触发故障切换""" + try: + from core.health_check import MirrorFailoverManager + + failover = MirrorFailoverManager(self.config) + failover.initialize() + + success = failover.perform_failover(mirror_type) + + handler.send_json_response({ + 'mirror_type': mirror_type, + 'success': success, + 'active_source': failover.get_active_source(mirror_type), + 'failover_history': failover.get_failover_history() + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def _get_source_url(self, mirror_type: str, source_name: str) -> str: + """获取源的 URL""" + mirrors = self.config.get('mirrors', {}) + mirror_config = mirrors.get(mirror_type, {}) + + # 先检查 sources_config + sources_config = mirror_config.get('sources_config', {}) + if source_name in sources_config: + return sources_config[source_name].get('url', '') + + # 使用 URL 模板 + url_template = mirror_config.get('url_template', '') + if url_template and '{mirror}' in url_template: + return url_template.replace('{mirror}', source_name) + + return '' + + # ==================== 缓存管理 API ==================== + + def api_get_cache_stats(self, handler): + """获取缓存统计""" + if hasattr(handler, 'cache_manager') and handler.cache_manager: + stats = handler.cache_manager.get_stats() + # 转换字段名以匹配前端期望 + handler.send_json_response({ + "size": stats.get('total_size', 0), + "count": stats.get('file_count', 0), + "hit_rate": stats.get('hit_rate', 0), + "last_clean": stats.get('last_clean', None), + "strategy": stats.get('strategy', 'unknown') + }) + else: + handler.send_json_response({ + "size": 0, + "count": 0, + "hit_rate": 0, + "last_clean": None, + "strategy": "unknown" + }) + + def api_clean_cache(self, handler): + """清理缓存""" + source = None # 可以从请求中获取 + + if hasattr(handler, 'cache_manager') and handler.cache_manager: + count = handler.cache_manager.clear(source) + handler.send_json_response({ + "success": True, + "deleted_count": count + }) + else: + handler.send_json_response({"error": "Cache manager not available"}, 500) + + def api_get_cache_usage(self, handler): + """获取缓存使用详情""" + if hasattr(handler, 'cache_manager') and handler.cache_manager: + usage = handler.cache_manager.get_cache_usage() + handler.send_json_response({ + "items": usage, + "count": len(usage) + }) + else: + handler.send_json_response({ + "items": [], + "count": 0 + }) + + def api_get_recent_activity(self, handler, query_params): + """获取最近活动""" + import traceback + limit = int(query_params.get('limit', [20])[0]) + limit = min(max(limit, 1), 100) # 限制在 1-100 之间 + offset = int(query_params.get('offset', [0])[0]) + offset = max(offset, 0) # 确保 offset 不为负数 + + activities = [] + errors = [] + all_activities = [] # 收集所有活动用于统一排序 + + # 从同步记录获取 + if hasattr(handler, 'sync_manager') and handler.sync_manager: + try: + if hasattr(handler.sync_manager, 'db_enabled') and handler.sync_manager.db_enabled: + from core.database import get_db + db = get_db() + if db: + with db.session() as session: + from core.database import SyncRecord + # 获取足够多的记录用于分页 + fetch_limit = limit + offset + records = session.query(SyncRecord).order_by( + SyncRecord.start_time.desc() + ).limit(fetch_limit).all() + for r in records: + all_activities.append({ + "time": datetime.fromtimestamp(r.start_time).isoformat() if r.start_time else '', + "timestamp": r.start_time or 0, + "type": "同步", + "content": f"同步任务: {r.source_name or r.sync_id}", + "status": "成功" if r.status == "completed" else ("进行中" if r.status == "running" else "失败"), + "status_type": "success" if r.status == "completed" else ("running" if r.status == "running" else "error") + }) + else: + errors.append("同步记录: db 为空") + except Exception as e: + errors.append(f"同步记录: {str(e)}") + traceback.print_exc() + + # 从下载记录获取 + try: + # 优先使用 handler.db + db = getattr(handler, 'db', None) + if db is None: + from core.database import get_db + db = get_db() + + if db: + from core.database import DownloadRecord + with db.session() as session: + # 获取足够多的记录用于分页 + fetch_limit = limit + offset + records = session.query(DownloadRecord).order_by( + DownloadRecord.download_time.desc() + ).limit(fetch_limit).all() + for r in records: + all_activities.append({ + "time": datetime.fromtimestamp(r.download_time).isoformat() if r.download_time else '', + "timestamp": r.download_time or 0, + "type": "下载", + "content": f"下载: {r.file_path or 'Unknown'}", + "status": "成功" if r.success else "失败", + "status_type": "success" if r.success else "error" + }) + else: + errors.append("下载记录: db 为空") + except Exception as e: + errors.append(f"下载记录: {str(e)}") + traceback.print_exc() + + # 从告警记录获取 + try: + from core.alerts import AlertManager + alert_manager = AlertManager(self.config.get('alerts', {})) + fetch_limit = limit + offset + alerts = alert_manager.get_alerts(limit=fetch_limit) + for a in alerts: + all_activities.append({ + "time": a.get('timestamp', ''), + "timestamp": 0, + "type": "告警", + "content": a.get('message', ''), + "status": a.get('severity', 'info').upper(), + "status_type": "error" if a.get('severity') == 'error' else ("warning" if a.get('severity') == 'warning' else "info") + }) + except Exception as e: + errors.append(f"告警记录: {str(e)}") + traceback.print_exc() + + # 按时间戳排序并应用分页 + all_activities.sort(key=lambda x: x.get('timestamp', 0), reverse=True) + + # 应用 offset 和 limit + paged_activities = all_activities[offset:offset + limit] + + handler.send_json_response({ + "activities": paged_activities, + "count": len(paged_activities), + "total": len(all_activities), + "_debug": {"errors": errors} if errors else {} + }) + + # ==================== 镜像加速源 API ==================== + + def api_list_mirrors(self, handler): + """列出所有镜像加速源(仅从配置文件读取,无内置预设)""" + # 仅从配置文件读取镜像配置 + mirrors_config = {} + try: + import json + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + # 优先使用 settings.json + settings_path = os.path.join(project_root, 'settings.json') + config_path = settings_path + if not os.path.exists(config_path): + # 回退到 config.json + config_path = os.path.join(project_root, 'config.json') + + if os.path.exists(config_path): + with open(config_path, 'r', encoding='utf-8') as f: + config_data = json.load(f) + mirrors_config = config_data.get('mirrors', {}) + except Exception: + pass + + # 如果文件没有配置,回退到内存配置 + if not mirrors_config: + mirrors_config = self.config.get('mirrors', {}) + + # 收集所有镜像(从 mirrors 配置读取) + available_mirrors = {} + + # 从 mirrors 配置读取所有镜像 + for mirror_name, mirror_config in mirrors_config.items(): + available_mirrors[mirror_name] = { + "name": mirror_config.get('name', mirror_name), + "type": mirror_config.get('type', 'http'), + "description": mirror_config.get('description', f"{mirror_name} 镜像"), + "url": mirror_config.get('url', ''), + "target": mirror_config.get('target', ''), + "enabled": mirror_config.get('enabled', True), + "custom": mirror_config.get('custom', True), + "auto_sync": mirror_config.get('auto_sync', False), + "schedule": mirror_config.get('schedule', {}), + "last_sync": mirror_config.get('last_sync'), + "storage_dir": mirror_config.get('storage_dir', mirror_name) + } + + handler.send_json_response({ + "mirrors": available_mirrors, + "count": len(available_mirrors) + }) + + def api_get_mirror_info(self, handler, mirror_name): + """获取镜像加速源信息""" + from mirrors import get_mirror_handler + + handler_class = get_mirror_handler(mirror_name) + if not handler_class: + handler.send_json_response({ + "error": f"Unknown mirror type: {mirror_name}" + }, 400) + return + + # 获取镜像配置 + mirror_config = self.config.get('mirrors', {}).get(mirror_name, {}) + + handler.send_json_response({ + "name": mirror_name, + "enabled": mirror_config.get('enabled', False), + "config": mirror_config + }) + + def api_refresh_mirror(self, handler, mirror_name): + """刷新镜像元数据""" + handler.send_json_response({ + "success": True, + "message": f"Mirror {mirror_name} refresh initiated" + }) + + def api_enable_mirror(self, handler, mirror_name, query_params): + """启用/禁用镜像源""" + # 获取 enabled 参数 + enabled = query_params.get('enabled', ['true'])[0].lower() == 'true' + + # 更新内存中的配置 + if 'mirrors' not in self.config: + self.config['mirrors'] = {} + if mirror_name not in self.config['mirrors']: + self.config['mirrors'][mirror_name] = {} + self.config['mirrors'][mirror_name]['enabled'] = enabled + + # 保存到配置文件 - 使用项目根目录的绝对路径 + try: + import json + import os + # 获取项目根目录 (vs1 目录) + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + # 优先使用 settings.json + settings_path = os.path.join(project_root, 'settings.json') + config_path = settings_path + if not os.path.exists(config_path): + config_path = os.path.join(project_root, 'config.json') + + config_data = {} + + # 如果配置文件存在,读取它 + if os.path.exists(config_path): + with open(config_path, 'r', encoding='utf-8') as f: + config_data = json.load(f) + + # 更新配置 + if 'mirrors' not in config_data: + config_data['mirrors'] = {} + if mirror_name not in config_data['mirrors']: + config_data['mirrors'][mirror_name] = {} + config_data['mirrors'][mirror_name]['enabled'] = enabled + + # 保存配置 + with open(config_path, 'w', encoding='utf-8') as f: + json.dump(config_data, f, ensure_ascii=False, indent=2) + + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Failed to save config: {str(e)}" + }, 500) + return + + handler.send_json_response({ + "success": True, + "mirror": mirror_name, + "enabled": enabled + }) + + def api_add_mirror(self, handler): + """添加自定义加速源""" + # 读取请求体 + try: + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length > 0: + body = handler.rfile.read(content_length) + import json + data = json.loads(body.decode('utf-8')) + else: + data = {} + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Invalid request body: {str(e)}" + }, 400) + return + + # 验证必要参数 + mirror_name = data.get('name', '').strip() + mirror_type = data.get('type', 'custom').strip() + mirror_url = data.get('url', '').strip() + + if not mirror_name: + handler.send_json_response({ + "success": False, + "error": "Mirror name is required" + }, 400) + return + + # 验证名称格式(只允许字母、数字、下划线、连字符) + if not mirror_name.replace('_', '').replace('-', '').isalnum(): + handler.send_json_response({ + "success": False, + "error": "Mirror name can only contain letters, numbers, underscores and hyphens" + }, 400) + return + + # 保存到配置文件 + try: + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + # 优先使用 settings.json + settings_path = os.path.join(project_root, 'settings.json') + config_path = settings_path + if not os.path.exists(config_path): + config_path = os.path.join(project_root, 'config.json') + + config_data = {} + if os.path.exists(config_path): + with open(config_path, 'r', encoding='utf-8') as f: + config_data = json.load(f) + + # 初始化 mirrors 节 + if 'mirrors' not in config_data: + config_data['mirrors'] = {} + + # 添加新镜像 + config_data['mirrors'][mirror_name] = { + "type": mirror_type, + "url": mirror_url, + "enabled": data.get('enabled', True), + "description": data.get('description', f"Custom mirror: {mirror_name}"), + "storage_dir": data.get('storage_dir', mirror_name), + "custom": True, + "created_at": datetime.now().isoformat() + } + + # 保存配置 + with open(config_path, 'w', encoding='utf-8') as f: + json.dump(config_data, f, ensure_ascii=False, indent=2) + + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Failed to save config: {str(e)}" + }, 500) + return + + handler.send_json_response({ + "success": True, + "message": f"Mirror '{mirror_name}' added successfully", + "mirror": { + "name": mirror_name, + "type": mirror_type, + "url": mirror_url, + "enabled": data.get('enabled', True) + } + }) + + def api_update_mirror(self, handler, mirror_name): + """更新自定义加速源""" + # 读取请求体 + try: + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length > 0: + body = handler.rfile.read(content_length) + import json + data = json.loads(body.decode('utf-8')) + else: + data = {} + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Invalid request body: {str(e)}" + }, 400) + return + + # 读取并更新配置文件 + try: + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + settings_path = os.path.join(project_root, 'settings.json') + config_path = settings_path + if not os.path.exists(config_path): + config_path = os.path.join(project_root, 'config.json') + + if not os.path.exists(config_path): + handler.send_json_response({ + "success": False, + "error": "Config file not found" + }, 404) + return + + with open(config_path, 'r', encoding='utf-8') as f: + config_data = json.load(f) + + # 检查mirrors节是否存在 + if 'mirrors' not in config_data: + config_data['mirrors'] = {} + + # 检查镜像是否存在 + if mirror_name not in config_data['mirrors']: + handler.send_json_response({ + "success": False, + "error": f"Mirror '{mirror_name}' not found" + }, 404) + return + + # 更新镜像信息(只更新提供的字段) + mirror_data = config_data['mirrors'][mirror_name] + if 'type' in data: + mirror_data['type'] = data['type'] + if 'url' in data: + mirror_data['url'] = data['url'] + if 'target' in data: + mirror_data['target'] = data['target'] + if 'description' in data: + mirror_data['description'] = data['description'] + if 'enabled' in data: + mirror_data['enabled'] = data['enabled'] + if 'storage_dir' in data: + mirror_data['storage_dir'] = data['storage_dir'] + mirror_data['updated_at'] = datetime.now().isoformat() + + config_data['mirrors'][mirror_name] = mirror_data + + # 保存配置 + with open(config_path, 'w', encoding='utf-8') as f: + json.dump(config_data, f, ensure_ascii=False, indent=2) + + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Failed to save config: {str(e)}" + }, 500) + return + + handler.send_json_response({ + "success": True, + "message": f"Mirror '{mirror_name}' updated successfully", + "mirror": { + "name": mirror_name, + "type": mirror_data.get('type'), + "url": mirror_data.get('url'), + "enabled": mirror_data.get('enabled') + } + }) + + def api_delete_mirror(self, handler, mirror_name): + """删除加速源""" + # 从配置文件删除 + try: + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + # 优先使用 settings.json + settings_path = os.path.join(project_root, 'settings.json') + config_path = settings_path + if not os.path.exists(config_path): + config_path = os.path.join(project_root, 'config.json') + + if not os.path.exists(config_path): + handler.send_json_response({ + "success": False, + "error": "Config file not found" + }, 404) + return + + with open(config_path, 'r', encoding='utf-8') as f: + config_data = json.load(f) + + if 'mirrors' not in config_data or mirror_name not in config_data['mirrors']: + handler.send_json_response({ + "success": False, + "error": f"Mirror '{mirror_name}' not found" + }, 404) + return + + # 删除镜像 + del config_data['mirrors'][mirror_name] + + # 保存配置 + with open(config_path, 'w', encoding='utf-8') as f: + json.dump(config_data, f, ensure_ascii=False, indent=2) + + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Failed to delete mirror: {str(e)}" + }, 500) + return + + handler.send_json_response({ + "success": True, + "message": f"Mirror '{mirror_name}' deleted successfully" + }) + + # ==================== 用户管理 API ==================== + + def api_login(self, handler): + """用户登录""" + try: + # 获取配置和数据库 + config = handler.config if hasattr(handler, 'config') else {} + db = getattr(handler, 'db', None) or (hasattr(handler, 'config') and handler.config.get('_db_instance')) + + # 读取请求体 + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"success": False, "error": "请求体不能为空"}, 400) + return + + body = handler.rfile.read(content_length).decode('utf-8') + import json + data = json.loads(body) + username = data.get('username', '') + password = data.get('password', '') + + if not username or not password: + handler.send_json_response({"success": False, "error": "用户名和密码不能为空"}, 400) + return + + # 获取客户端IP + client_ip = None + try: + client_ip = handler.client_address[0] + except Exception: + pass + + # 获取 auth_type + auth_type = config.get('auth_type', 'none') + + # 如果 auth_type 为 none,任何用户都可以通过 + if auth_type == 'none': + # 生成 token + import secrets + token = secrets.token_hex(32) + token_expires_at = time.time() + 86400 # 24小时过期 + + handler.send_json_response({ + "success": True, + "token": token, + "token_expires_at": token_expires_at, + "username": username or 'anonymous', + "level": "admin" + }) + return + + # 验证凭据 + config_user = config.get('auth_user', '') + config_pass = config.get('auth_pass', '') + + # 数据库验证 + if db: + user = db.get_user(username) + if user and db.verify_password(password, user['password_hash']): + # 数据库验证成功,生成 token + import secrets + token = secrets.token_hex(32) + token_expires_at = time.time() + 86400 # 24小时过期 + + # 保存 token 到数据库 + db.update_user_token(username, token, token_expires_at) + + db.add_login_log(username, client_ip, 'success', '登录成功(数据库)') + + handler.send_json_response({ + "success": True, + "token": token, + "token_expires_at": token_expires_at, + "username": username, + "level": user.get('role', 'admin') + }) + return + + # 配置文件验证(仅当数据库中没有该用户时) + if username == config_user and password == config_pass: + # 验证成功,生成 token + import secrets + token = secrets.token_hex(32) + token_expires_at = time.time() + 86400 # 24小时过期 + + # 保存 token 到数据库 + if db: + user = db.get_user(username) + if user: + # 数据库中有用户,更新 token + db.update_user_token(username, token, token_expires_at) + else: + # 数据库中没有该用户,创建新用户并保存 token + from core.database import UserRecord + password_hash = db.hash_password(password) + new_user = UserRecord( + username=username, + password_hash=password_hash, + token=token, + token_expires_at=token_expires_at, + role='admin', + enabled=True + ) + with db.session() as session: + session.add(new_user) + db.add_login_log(username, client_ip, 'success', '登录成功(配置验证)') + + handler.send_json_response({ + "success": True, + "token": token, + "token_expires_at": token_expires_at, + "username": username, + "level": "admin" + }) + else: + # 验证失败 + if db: + db.add_login_log(username, client_ip, 'failed', '用户名或密码错误') + handler.send_json_response({"success": False, "error": "用户名或密码错误"}, 401) + + except json.JSONDecodeError: + handler.send_json_response({"success": False, "error": "无效的 JSON"}, 400) + except Exception as e: + import traceback + traceback.print_exc() + handler.send_json_response({"success": False, "error": str(e)}, 500) + + def api_change_password(self, handler): + """修改用户密码""" + try: + # 获取数据库实例 + db = getattr(handler, 'db', None) or (hasattr(handler, 'config') and handler.config.get('_db_instance')) + config = handler.config if hasattr(handler, 'config') else {} + + if not db: + handler.send_json_response({"success": False, "error": "数据库不可用"}, 500) + return + + # 读取请求体 + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({"error": "请求体不能为空"}, 400) + return + + body = handler.rfile.read(content_length) + data = json.loads(body.decode('utf-8')) + + username = data.get('username') + old_password = data.get('old_password') + new_password = data.get('new_password') + + if not username or not new_password or not old_password: + handler.send_json_response({"error": "缺少必要参数(username / old_password / new_password)"}, 400) + return + + # 获取配置中的账号密码 + config_user = config.get('auth_user', '') + config_pass = config.get('auth_pass', '') + + # 验证旧密码(优先验证数据库,没有则验证配置文件)—— 强制要求,防止无旧密码改密 + user = db.get_user(username) if db else None + if user: + # 验证数据库密码 + if not db.verify_password(old_password, user['password_hash']): + handler.send_json_response({"success": False, "error": "原密码错误"}, 400) + return + elif username == config_user and old_password != config_pass: + # 数据库没有用户,验证配置文件 + handler.send_json_response({"success": False, "error": "原密码错误"}, 400) + return + + # 使用 bcrypt 加密新密码 + new_hash = db.hash_password(new_password) + + # 更新数据库 + existing_user = db.get_user(username) + if existing_user: + success = db.update_password(username, new_hash) + else: + result = db.create_user(username, new_hash, 'admin') + success = result.get('success', False) + + if success: + # 修改密码后清除 token,强制重新登录 + if db and username: + db.clear_user_token(username) + + handler.send_json_response({ + "success": True, + "message": "密码修改成功(已存储到数据库),请重新登录" + }) + else: + handler.send_json_response({"success": False, "error": "用户不存在"}, 404) + + except json.JSONDecodeError: + handler.send_json_response({"error": "Invalid JSON format"}, 400) + except Exception as e: + handler.send_json_response({"error": str(e)}, 500) + + def api_get_login_logs(self, handler, query_params): + """获取登录日志""" + try: + # 尝试从多个位置获取数据库实例 + db = None + if hasattr(handler, 'db'): + db = handler.db + elif hasattr(handler, 'config') and handler.config: + db = handler.config.get('_db_instance') + + if not db: + handler.send_json_response({"success": False, "error": "数据库不可用,请确保数据库已启用"}, 500) + return + + # 安全获取 limit 参数 + try: + limit = int(query_params.get('limit', ['50'])[0]) + except (ValueError, TypeError): + limit = 50 + + logs = db.get_login_logs(limit=limit) + + handler.send_json_response({ + "success": True, + "logs": logs, + "count": len(logs) + }) + except Exception as e: + import traceback + trace = traceback.format_exc() + print(f"[ERROR] api_get_login_logs: {e}") + print(trace) + handler.send_json_response({"error": str(e)}, 500) + + def api_get_config(self, handler): + """获取配置文件内容 (settings.json)""" + try: + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + settings_path = os.path.join(project_root, 'settings.json') + + # 优先使用 settings.json,如果不存在则尝试 config.json + config_path = settings_path + if not os.path.exists(config_path): + config_path = os.path.join(project_root, 'config.json') + + if not os.path.exists(config_path): + handler.send_json_response({ + "success": False, + "error": "Config file not found (settings.json or config.json)" + }, 404) + return + + # 脱敏: 不返回 settings.json 原文(含 auth_pass/auth_token 等密钥), + # 只返回精选非敏感字段 + try: + with open(config_path, 'r', encoding='utf-8') as f: + raw = json.load(f) + except Exception: + raw = {} + + safe_keys = [ + 'server_name', 'host', 'port', 'base_dir', 'api_version', + 'directory_listing', 'enable_stats', 'show_hash', 'ignore_hidden', + 'enable_range', 'max_workers', 'timeout', 'max_upload_size', + 'enable_ws', 'enable_sse', 'enable_monitor', 'monitor_interval', + 'enable_sync', 'enable_mirrors', 'auth_type', 'log_level', + 'cache_size', 'cache_ttl', 'sort_by', 'sort_reverse', + 'max_search_results', 'session_timeout', 'sync_interval', + ] + safe_config = {k: raw.get(k) for k in safe_keys if k in raw} + + handler.send_json_response({ + "success": True, + "config": safe_config, + "path": config_path, + "filename": os.path.basename(config_path) + }) + + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Failed to read config: {str(e)}" + }, 500) + + def api_save_config(self, handler): + """保存配置文件 (settings.json) - 只更新部分配置""" + try: + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length > 0: + config_content = handler.rfile.read(content_length).decode('utf-8') + else: + handler.send_json_response({ + "success": False, + "error": "No config content provided" + }, 400) + return + + # 验证 JSON 格式 + try: + updates = json.loads(config_content) + except json.JSONDecodeError as e: + handler.send_json_response({ + "success": False, + "error": f"Invalid JSON format: {str(e)}" + }, 400) + return + + if not isinstance(updates, dict): + handler.send_json_response({ + "success": False, + "error": "Config must be a JSON object" + }, 400) + return + + # 读取现有配置 + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + settings_path = os.path.join(project_root, 'settings.json') + + if not os.path.exists(settings_path): + handler.send_json_response({ + "success": False, + "error": "settings.json not found" + }, 404) + return + + # 加载现有配置 + from core.config import load_json_config + current_config = load_json_config(settings_path) or {} + + # 使用 deep_merge 合并更新(只更新传入的字段) + from core.config import deep_merge + new_config = deep_merge(current_config, updates) + + # 备份现有配置 + backup_path = settings_path + '.bak' + try: + with open(settings_path, 'r', encoding='utf-8') as f: + backup_content = f.read() + with open(backup_path, 'w', encoding='utf-8') as f: + f.write(backup_content) + except Exception: + pass + + # 保存合并后的配置 + with open(settings_path, 'w', encoding='utf-8') as f: + json.dump(new_config, f, indent=4, ensure_ascii=False) + + handler.send_json_response({ + "success": True, + "message": "Config updated successfully (partial update)", + "path": settings_path, + "updated_keys": list(updates.keys()) + }) + + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Failed to save config: {str(e)}" + }, 500) + + def api_reload_config(self, handler): + """重新加载配置文件(热更新)""" + try: + import os + from core.config import load_json_config + + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + settings_path = os.path.join(project_root, 'settings.json') + + if not os.path.exists(settings_path): + handler.send_json_response({ + "success": False, + "error": "Config file not found" + }, 404) + return + + # 加载并验证配置 + new_config = load_json_config(settings_path) + if new_config is None: + handler.send_json_response({ + "success": False, + "error": "Failed to load config" + }, 500) + return + + # 计算变更 + old_config = self.config.copy() + changes = self._compute_config_changes(old_config, new_config) + + # 更新内存配置 + self.config.update(new_config) + + # 通知各模块配置变更 + change_notifications = [] + if 'enable_monitor' in changes.get('modified', []): + change_notifications.append("Monitor settings changed") + if 'enable_sync' in changes.get('modified', []): + change_notifications.append("Sync settings changed") + if 'mirrors' in changes.get('modified', []): + change_notifications.append("Mirrors configuration changed") + + handler.send_json_response({ + "success": True, + "message": "Configuration reloaded successfully", + "changes": changes, + "change_notifications": change_notifications, + "timestamp": datetime.now().isoformat() + }) + + except Exception as e: + handler.send_json_response({ + "success": False, + "error": f"Failed to reload config: {str(e)}" + }, 500) + + def _compute_config_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]: + changes['modified'].append(key) + + return changes + + def api_get_config_changes(self, handler): + """获取配置变更历史""" + handler.send_json_response({ + "changes": [], + "message": "Config change history tracking requires hot reload enabled" + }) + + # ==================== WebSocket/SSE 状态 API ==================== + + def api_get_ws_clients(self, handler): + """获取WebSocket客户端状态""" + # 需要访问全局ws_manager + handler.send_json_response({ + "ws_clients": 0, + "sse_clients": 0 + }) + + # ==================== Prometheus 指标 API ==================== + + def api_get_metrics(self, handler): + """获取 Prometheus 格式的指标""" + from core.prometheus import PrometheusMetrics + + metrics = PrometheusMetrics(self.config) + + # 设置运行时间 + uptime = time.time() - self.config.get('start_time', time.time()) + metrics.set_uptime(uptime) + + # 从数据库获取统计 + if self.db_enabled and self.db: + try: + db_stats = self.db.get_stats() + metrics.set_files( + db_stats.get('total_files', 0), + db_stats.get('total_size', 0) + ) + metrics.set_db_stats( + db_stats.get('total_files', 0), + db_stats.get('total_sync_records', 0), + db_stats.get('total_cache_records', 0) + ) + except Exception as e: + pass + + # 从缓存获取统计 + if hasattr(handler, 'cache_manager') and handler.cache_manager: + try: + cache_stats = handler.cache_manager.get_stats() + metrics.set_cache( + cache_stats.get('size', 0), + cache_stats.get('count', 0), + cache_stats.get('hits', 0), + cache_stats.get('misses', 0) + ) + except Exception as e: + pass + + # 从监控模块获取系统指标 + if hasattr(handler, 'monitor') and handler.monitor: + try: + monitor_stats = handler.monitor.get_realtime_stats() + metrics.set_system( + cpu=monitor_stats.get('cpu', {}).get('percent', 0), + memory=monitor_stats.get('memory', {}).get('percent', 0), + disk=monitor_stats.get('disk', {}).get('percent', 0), + disk_free=monitor_stats.get('disk', {}).get('free', 0), + disk_total=monitor_stats.get('disk', {}).get('total', 0), + network_rx=monitor_stats.get('network', {}).get('rx', 0), + network_tx=monitor_stats.get('network', {}).get('tx', 0) + ) + except Exception as e: + pass + + # 从同步管理器获取镜像状态 + if hasattr(handler, 'sync_manager') and handler.sync_manager: + mirrors = self.config.get('mirrors', {}) + for mirror_type, mirror_config in mirrors.items(): + if isinstance(mirror_config, dict): + enabled = mirror_config.get('enabled', True) + last_sync = 0 + metrics.set_mirror_status(mirror_type, enabled, last_sync) + + # 生成 Prometheus 格式输出 + output = metrics.generate_metrics() + + handler.send_response(200) + handler.send_header('Content-Type', 'text/plain; charset=utf-8') + handler.send_header('Content-Length', str(len(output))) + handler.end_headers() + handler.wfile.write(output.encode('utf-8')) + + # ==================== 服务器信息 API ==================== + + def api_get_server_info(self, handler): + """获取服务器完整信息""" + import psutil + + uptime_seconds = time.time() - self.config.get('start_time', time.time()) + uptime_str = self._format_uptime(uptime_seconds) + + handler.send_json_response({ + "name": self.config.get('server_name', 'HYC下载站'), + "version": "2.2.0", + "uptime_seconds": round(uptime_seconds, 2), + "uptime_formatted": uptime_str, + "api_version": "v2", + "config": { + "host": self.config.get('host'), + "port": self.config.get('port'), + "base_dir": self.config.get('base_dir'), + "auth_type": self.config.get('auth_type'), + "directory_listing": self.config.get('directory_listing'), + "max_upload_size": self.config.get('max_upload_size') + }, + "features": { + "websocket": True, + "sse": True, + "sync": True, + "mirrors": True, + "cache": True, + "monitor": True + } + }) + + # ==================== 告警管理 API ==================== + + def api_get_alerts(self, handler, query_params): + """获取告警列表""" + try: + from core.alerts import AlertManager + + alert_manager = AlertManager(self.config.get('alerts', {})) + + # 解析查询参数 + limit = int(query_params.get('limit', [50])[0]) + acknowledged = query_params.get('acknowledged', [None])[0] + severity = query_params.get('severity', [None])[0] + + if acknowledged is not None: + acknowledged = acknowledged.lower() == 'true' + + alerts = alert_manager.get_alerts( + acknowledged=acknowledged, + severity=severity, + limit=limit + ) + + stats = alert_manager.get_stats() + + handler.send_json_response({ + 'alerts': alerts, + 'count': len(alerts), + 'stats': stats + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def api_acknowledge_alert(self, handler, alert_id): + """确认告警""" + try: + from core.alerts import AlertManager + + alert_manager = AlertManager(self.config.get('alerts', {})) + success = alert_manager.acknowledge_alert(alert_id) + + handler.send_json_response({ + 'success': success, + 'message': f"Alert {alert_id} acknowledged" if success else "Alert not found" + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def api_clear_alerts(self, handler): + """清除告警历史""" + try: + from core.alerts import AlertManager + + alert_manager = AlertManager(self.config.get('alerts', {})) + success = alert_manager.clear_history() + + handler.send_json_response({ + 'success': success, + 'message': 'Alert history cleared' + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def api_test_alert(self, handler): + """测试告警发送""" + try: + from core.alerts import AlertManager, Alert, AlertSeverity + + alert_manager = AlertManager(self.config.get('alerts', {})) + + # 创建测试告警 + test_alert = Alert( + alert_type='test', + severity=AlertSeverity.INFO, + title='Test Alert', + message='This is a test alert from HYC Mirror Server', + details={'test': True, 'timestamp': datetime.now().isoformat()} + ) + + success = alert_manager.trigger_alert(test_alert) + + handler.send_json_response({ + 'success': success, + 'message': 'Test alert sent successfully' if success else 'Failed to send test alert (check configuration)' + }) + + except Exception as e: + handler.send_json_response({ + 'error': str(e) + }, 500) + + def api_get_alert_config(self, handler): + """获取告警配置""" + alerts_config = self.config.get('alerts', {}) + + # 隐藏敏感信息 + config = { + 'enabled': alerts_config.get('enabled', False), + 'email': { + 'enabled': alerts_config.get('email', {}).get('enabled', False), + 'smtp_host': alerts_config.get('email', {}).get('smtp_host', ''), + 'smtp_port': alerts_config.get('email', {}).get('smtp_port', 587), + 'from_address': alerts_config.get('email', {}).get('from_address', ''), + 'to_addresses': alerts_config.get('email', {}).get('to_addresses', []), + 'use_tls': alerts_config.get('email', {}).get('use_tls', True) + }, + 'webhook': { + 'enabled': alerts_config.get('webhook', {}).get('enabled', False), + 'url': '***' if alerts_config.get('webhook', {}).get('url') else '' + }, + 'rules': alerts_config.get('rules', {}) + } + + handler.send_json_response(config) + + def api_save_alert_config(self, handler): + """保存告警配置""" + try: + import os + from core.alerts import AlertManager + + # 获取请求体 + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length > 0: + post_data = handler.rfile.read(content_length) + try: + new_config = json.loads(post_data.decode('utf-8')) + except json.JSONDecodeError as e: + handler.send_json_response({ + 'success': False, + 'error': f"Invalid JSON format: {str(e)}" + }, 400) + return + else: + handler.send_json_response({ + 'success': False, + 'error': "No configuration data provided" + }, 400) + return + + # 更新配置 + if 'alerts' not in self.config: + self.config['alerts'] = {} + + # 更新告警配置 + if 'enabled' in new_config: + self.config['alerts']['enabled'] = new_config['enabled'] + if 'email' in new_config: + self.config['alerts']['email'] = new_config['email'] + if 'webhook' in new_config: + self.config['alerts']['webhook'] = new_config['webhook'] + if 'rules' in new_config: + self.config['alerts']['rules'] = new_config['rules'] + + # 保存到文件 + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + settings_path = os.path.join(project_root, 'settings.json') + + with open(settings_path, 'r', encoding='utf-8') as f: + settings_data = json.load(f) + + settings_data['alerts'] = self.config['alerts'] + + with open(settings_path, 'w', encoding='utf-8') as f: + json.dump(settings_data, f, ensure_ascii=False, indent=4) + + handler.send_json_response({ + 'success': True, + 'message': 'Alert configuration saved successfully', + 'config': self.config.get('alerts', {}) + }) + + except Exception as e: + handler.send_json_response({ + 'success': False, + 'error': str(e) + }, 500) + + # ==================== 平滑重启 API ==================== + + def api_get_restart_status(self, handler): + """获取重启状态""" + from core.graceful_restart import GracefulRestartManager, ServerState + + restart_manager = GracefulRestartManager(self.config.get('restart', {})) + stats = restart_manager.get_stats() + + handler.send_json_response({ + 'state': stats['state'], + 'pending_requests': stats['pending_requests'], + 'graceful_timeout': stats['graceful_timeout'], + 'recent_restarts': stats['recent_restarts'] + }) + + def api_get_pending_requests(self, handler): + """获取待处理请求""" + from core.graceful_restart import GracefulRestartManager + + restart_manager = GracefulRestartManager(self.config.get('restart', {})) + pending = restart_manager.get_pending_requests() + + handler.send_json_response({ + 'count': len(pending), + 'requests': pending + }) + + def api_graceful_restart(self, handler, query_params): + """执行优雅重启""" + from core.graceful_restart import GracefulRestartManager, RestartStrategy + + # 解析策略参数 + strategy_param = query_params.get('strategy', ['graceful'])[0] + try: + strategy = RestartStrategy(strategy_param) + except ValueError: + strategy = RestartStrategy.GRACEFUL + + restart_manager = GracefulRestartManager(self.config.get('restart', {})) + + # 准备重启 + prepare_result = restart_manager.prepare_restart() + if not prepare_result['success']: + handler.send_json_response({ + 'success': False, + 'error': prepare_result['message'] + }, 500) + return + + # 返回待处理请求信息,让客户端决定是否继续 + pending_count = prepare_result['pending_requests'] + + handler.send_json_response({ + 'success': True, + 'pending_requests': pending_count, + 'message': f'Ready to restart with {pending_count} pending requests', + 'strategy': strategy.value, + 'graceful_timeout': restart_manager.graceful_timeout, + 'continue_url': '/api/v2/server/restart/confirm' + }) + + def api_confirm_restart(self, handler): + """确认执行重启""" + from core.graceful_restart import GracefulRestartManager, RestartStrategy + + try: + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length > 0: + post_data = json.loads(handler.rfile.read(content_length)) + strategy = post_data.get('strategy', 'graceful') + else: + strategy = 'graceful' + + try: + restart_strategy = RestartStrategy(strategy) + except ValueError: + restart_strategy = RestartStrategy.GRACEFUL + + except json.JSONDecodeError: + restart_strategy = RestartStrategy.GRACEFUL + + restart_manager = GracefulRestartManager(self.config.get('restart', {})) + + # 执行重启 + result = restart_manager.perform_restart(strategy=restart_strategy) + + handler.send_json_response(result) + + def api_immediate_restart(self, handler): + """立即重启服务器""" + from core.graceful_restart import GracefulRestartManager + + restart_manager = GracefulRestartManager(self.config.get('restart', {})) + + # 获取脚本路径 + script_path = self.config.get('main_script', 'main.py') + + result = restart_manager.perform_restart( + strategy='immediate', + script_path=script_path + ) + + handler.send_json_response(result) + + def api_get_restart_history(self, handler): + """获取重启历史""" + from core.graceful_restart import GracefulRestartManager + + restart_manager = GracefulRestartManager(self.config.get('restart', {})) + history = restart_manager.get_restart_history() + + handler.send_json_response({ + 'count': len(history), + 'history': history + }) + + def api_get_restart_config(self, handler): + """获取重启配置""" + restart_config = self.config.get('restart', {}) + + handler.send_json_response({ + 'graceful_timeout': restart_config.get('graceful_timeout', 30), + 'shutdown_timeout': restart_config.get('shutdown_timeout', 10), + 'enabled': True + }) + + def api_update_restart_config(self, handler): + """更新重启配置""" + try: + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({ + 'success': False, + 'error': 'No configuration data provided' + }, 400) + return + + new_config = json.loads(handler.rfile.read(content_length)) + + # 更新内存配置 + if 'restart' not in self.config: + self.config['restart'] = {} + + if 'graceful_timeout' in new_config: + self.config['restart']['graceful_timeout'] = new_config['graceful_timeout'] + if 'shutdown_timeout' in new_config: + self.config['restart']['shutdown_timeout'] = new_config['shutdown_timeout'] + + # 保存到文件 + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + settings_path = os.path.join(project_root, 'settings.json') + + with open(settings_path, 'r', encoding='utf-8') as f: + settings_data = json.load(f) + + if 'restart' not in settings_data: + settings_data['restart'] = {} + settings_data['restart'] = {**settings_data.get('restart', {}), **new_config} + + with open(settings_path, 'w', encoding='utf-8') as f: + json.dump(settings_data, f, ensure_ascii=False, indent=4) + + handler.send_json_response({ + 'success': True, + 'message': 'Restart configuration updated', + 'config': self.config.get('restart', {}) + }) + + except json.JSONDecodeError as e: + handler.send_json_response({ + 'success': False, + 'error': f'Invalid JSON format: {str(e)}' + }, 400) + except Exception as e: + handler.send_json_response({ + 'success': False, + 'error': str(e) + }, 500) + + def _format_uptime(self, seconds: float) -> str: + """格式化运行时间""" + if seconds < 60: + return f"{int(seconds)}秒" + elif seconds < 3600: + minutes = int(seconds // 60) + return f"{minutes}分钟" + elif seconds < 86400: + hours = int(seconds // 3600) + minutes = int((seconds % 3600) // 60) + return f"{hours}小时{minutes}分钟" + else: + days = int(seconds // 86400) + hours = int((seconds % 86400) // 3600) + return f"{days}天{hours}小时" + + # ==================== 缓存预热 API ==================== + + def api_get_prewarm_status(self, handler): + """获取缓存预热状态""" + from core.cache_prewarm import CachePrewarmer + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + status = prewarmer.get_status() + + handler.send_json_response(status) + + def api_get_prewarm_stats(self, handler): + """获取缓存预热统计""" + from core.cache_prewarm import CachePrewarmer + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + stats = prewarmer.get_stats() + + handler.send_json_response(stats) + + def api_get_prewarm_items(self, handler, query_params): + """获取预热项目列表""" + from core.cache_prewarm import CachePrewarmer + + # 解析查询参数 + status = query_params.get('status', [None])[0] + mirror_type = query_params.get('mirror_type', [None])[0] + limit = int(query_params.get('limit', [50])[0]) + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + items = prewarmer.get_items(status=status, mirror_type=mirror_type, limit=limit) + + handler.send_json_response({ + 'count': len(items), + 'items': items + }) + + def api_get_prewarm_history(self, handler): + """获取预热历史""" + from core.cache_prewarm import CachePrewarmer + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + history = prewarmer.get_history() + + handler.send_json_response({ + 'count': len(history), + 'history': history + }) + + def api_run_prewarm(self, handler, query_params): + """执行缓存预热""" + from core.cache_prewarm import CachePrewarmer + + # 解析参数 + mirror_type = query_params.get('mirror_type', [None])[0] + limit = int(query_params.get('limit', [50])[0]) + priority = query_params.get('priority', ['medium'])[0] + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + + # 如果指定了镜像类型,只预热该类型 + targets = None + if mirror_type: + targets = [] + from core.cache_prewarm import PrewarmTarget + targets.append(PrewarmTarget( + mirror_type=mirror_type, + priority=priority, + limit=limit + )) + + result = prewarmer.run(targets=targets) + + handler.send_json_response(result) + + def api_add_prewarm_items(self, handler): + """添加预热项目""" + try: + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({ + 'success': False, + 'error': 'No data provided' + }, 400) + return + + data = json.loads(handler.rfile.read(content_length)) + + mirror_type = data.get('mirror_type') + items = data.get('items', []) + priority = data.get('priority', 'medium') + + if not mirror_type or not items: + handler.send_json_response({ + 'success': False, + 'error': 'mirror_type and items are required' + }, 400) + return + + from core.cache_prewarm import CachePrewarmer + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer.add_items_batch(mirror_type, items, priority) + + handler.send_json_response({ + 'success': True, + 'message': f'Added {len(items)} items to prewarm queue', + 'mirror_type': mirror_type, + 'count': len(items) + }) + + except json.JSONDecodeError as e: + handler.send_json_response({ + 'success': False, + 'error': f'Invalid JSON format: {str(e)}' + }, 400) + except Exception as e: + handler.send_json_response({ + 'success': False, + 'error': str(e) + }, 500) + + def api_add_popular_items(self, handler, query_params): + """添加流行项目到预热队列""" + from core.cache_prewarm import CachePrewarmer + + mirror_type = query_params.get('mirror_type', [None])[0] + limit = int(query_params.get('limit', [20])[0]) + priority = query_params.get('priority', ['medium'])[0] + + if not mirror_type: + handler.send_json_response({ + 'success': False, + 'error': 'mirror_type is required' + }, 400) + return + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + popular = prewarmer.get_popular_items(mirror_type) + + if limit: + popular = popular[:limit] + + prewarmer.add_popular_items_to_queue(mirror_type, limit, priority) + + handler.send_json_response({ + 'success': True, + 'message': f'Added {len(popular)} popular items', + 'mirror_type': mirror_type, + 'count': len(popular), + 'items': popular + }) + + def api_get_popular_items(self, handler, query_params): + """获取流行项目列表""" + from core.cache_prewarm import CachePrewarmer + + mirror_type = query_params.get('mirror_type', [None])[0] + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + + if mirror_type: + items = prewarmer.get_popular_items(mirror_type) + handler.send_json_response({ + 'mirror_type': mirror_type, + 'count': len(items), + 'items': items + }) + else: + all_popular = prewarmer._popular_items + handler.send_json_response({ + 'mirror_types': list(all_popular.keys()), + 'total_types': len(all_popular) + }) + + def api_clear_prewarm_queue(self, handler): + """清空预热队列""" + from core.cache_prewarm import CachePrewarmer + + prewarmer = CachePrewarmer(self.config.get('cache_prewarm', {})) + prewarmer.clear_items() + + handler.send_json_response({ + 'success': True, + 'message': 'Prewarm queue cleared' + }) + + def api_get_prewarm_config(self, handler): + """获取缓存预热配置""" + prewarm_config = self.config.get('cache_prewarm', {}) + + handler.send_json_response({ + 'enabled': prewarm_config.get('enabled', False), + 'schedule': prewarm_config.get('schedule', '0 3 * * *'), + 'batch_size': prewarm_config.get('batch_size', 10), + 'targets': prewarm_config.get('targets', []) + }) + + def api_save_prewarm_config(self, handler): + """保存缓存预热配置""" + try: + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length == 0: + handler.send_json_response({ + 'success': False, + 'error': 'No configuration data provided' + }, 400) + return + + new_config = json.loads(handler.rfile.read(content_length)) + + # 更新内存配置 + if 'cache_prewarm' not in self.config: + self.config['cache_prewarm'] = {} + + if 'enabled' in new_config: + self.config['cache_prewarm']['enabled'] = new_config['enabled'] + if 'schedule' in new_config: + self.config['cache_prewarm']['schedule'] = new_config['schedule'] + if 'batch_size' in new_config: + self.config['cache_prewarm']['batch_size'] = new_config['batch_size'] + if 'targets' in new_config: + self.config['cache_prewarm']['targets'] = new_config['targets'] + + # 保存到文件 + import os + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + settings_path = os.path.join(project_root, 'settings.json') + + with open(settings_path, 'r', encoding='utf-8') as f: + settings_data = json.load(f) + + if 'cache_prewarm' not in settings_data: + settings_data['cache_prewarm'] = {} + settings_data['cache_prewarm'] = {**settings_data.get('cache_prewarm', {}), **new_config} + + with open(settings_path, 'w', encoding='utf-8') as f: + json.dump(settings_data, f, ensure_ascii=False, indent=4) + + handler.send_json_response({ + 'success': True, + 'message': 'Cache prewarm configuration saved', + 'config': self.config.get('cache_prewarm', {}) + }) + + except json.JSONDecodeError as e: + handler.send_json_response({ + 'success': False, + 'error': f'Invalid JSON format: {str(e)}' + }, 400) + except Exception as e: + handler.send_json_response({ + 'success': False, + 'error': str(e) + }, 500) + + # ==================== API 文档 ==================== + + def api_get_api_docs(self, handler, format: str = 'json'): + """获取 API 文档""" + from core.api_docs import generate_api_docs + + docs = generate_api_docs(self.config) + + if format == 'yaml': + try: + import yaml + content = yaml.dump(docs, default_flow_style=False, allow_unicode=True) + handler.send_response(200) + handler.send_header('Content-Type', 'text/yaml') + handler.send_header('Content-Length', len(content.encode('utf-8'))) + handler.end_headers() + handler.wfile.write(content.encode('utf-8')) + return + except ImportError: + format = 'json' + + content = json.dumps(docs, ensure_ascii=False, indent=2) + handler.send_response(200) + handler.send_header('Content-Type', 'application/json') + handler.send_header('Content-Length', len(content.encode('utf-8'))) + handler.send_header('Access-Control-Allow-Origin', '*') + handler.end_headers() + handler.wfile.write(content.encode('utf-8')) + + def api_generate_api_docs(self, handler): + """生成并保存 API 文档""" + try: + from core.api_docs import save_api_docs + + # 获取保存路径 + content_length = int(handler.headers.get('Content-Length', 0)) + if content_length > 0: + data = json.loads(handler.rfile.read(content_length)) + filepath = data.get('filepath', 'docs/api-docs.json') + format = data.get('format', 'json') + else: + filepath = 'docs/api-docs.json' + format = 'json' + + # 路径安全校验: 只允许写入项目 docs/ 或 api/docs/ 目录 + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + allowed_dirs = [ + os.path.realpath(os.path.join(project_root, 'docs')), + os.path.realpath(os.path.join(project_root, 'api', 'docs')), + ] + resolved = os.path.realpath(os.path.join(project_root, filepath)) + if not any(resolved == d or resolved.startswith(d + os.sep) for d in allowed_dirs): + handler.send_json_response({ + 'success': False, + 'error': 'Invalid filepath: must be under docs/' + }, 403) + return + + # 生成并保存文档 + saved_path = save_api_docs(self.config, filepath, format) + + handler.send_json_response({ + 'success': True, + 'message': f'API documentation generated', + 'path': saved_path, + 'format': format + }) + + except Exception as e: + handler.send_json_response({ + 'success': False, + 'error': str(e) + }, 500) diff --git a/core/api_auth.py b/core/api_auth.py index 1820eb7..730e659 100644 --- a/core/api_auth.py +++ b/core/api_auth.py @@ -38,6 +38,7 @@ class APIAuthManager: def __init__(self, config: dict = None): self.config = config or {} self.sessions: Dict[str, AuthSession] = {} + self._lock = threading.Lock() # 确定基础目录(用于保存会话文件) base_dir = config.get('base_dir', '.') if config else '.' @@ -48,9 +49,20 @@ class APIAuthManager: elif base_dir == '.': base_dir = os.getcwd() + # 数据目录:默认放到 base_dir 同级 data/ 下, + # 避免 auth_sessions.json 落入 web 静态根目录被公开下载 + data_dir = config.get('data_dir') if config else None + if not data_dir: + data_dir = os.path.join(os.path.dirname(os.path.abspath(base_dir)), 'data') + try: + os.makedirs(data_dir, exist_ok=True) + except OSError: + data_dir = base_dir + self.data_dir = data_dir + # 会话文件路径 sessions_filename = config.get('auth_sessions_file', 'auth_sessions.json') if config else 'auth_sessions.json' - self.sessions_file = os.path.join(base_dir, sessions_filename) + self.sessions_file = os.path.join(data_dir, sessions_filename) # 会话超时时间 self.session_timeout = config.get('auth_session_timeout', 3600) if config else 3600 @@ -63,9 +75,33 @@ class APIAuthManager: self.ip_whitelist = config.get('ip_whitelist', []) if config else [] self.ip_whitelist_enabled = config.get('ip_whitelist_enabled', False) if config else False + # Cookie 签名密钥:优先用配置;缺失则从文件读取或生成并持久化 + self.auth_secret = self._load_or_create_secret(config) + # 加载已保存的会话 self._load_sessions() + def _load_or_create_secret(self, config) -> str: + """获取或生成 auth_secret(持久化到数据目录,避免默认密钥公开可伪造)""" + secret = config.get('auth_secret') if config else None + if secret: + return secret + + secret_file = os.path.join(self.data_dir, 'auth_secret.key') + try: + if os.path.exists(secret_file): + with open(secret_file, 'r', encoding='utf-8') as f: + secret = f.read().strip() + if secret: + return secret + secret = secrets.token_hex(32) + with open(secret_file, 'w', encoding='utf-8') as f: + f.write(secret) + os.chmod(secret_file, 0o600) + except Exception as e: + print(f"警告: 持久化 auth_secret 失败: {e}") + return secret + @property def db(self): """动态获取数据库实例""" @@ -182,6 +218,7 @@ class APIAuthManager: def validate_cookie(self, cookie_value: str, client_ip: str = None) -> Optional[dict]: """验证认证Cookie""" + import hmac if not cookie_value: return None @@ -191,50 +228,52 @@ class APIAuthManager: session_id, timestamp, signature = parts - session = self.sessions.get(session_id) - if not session: - return None + with self._lock: + session = self.sessions.get(session_id) + if not session: + return None - if time.time() > session.expires_at: - del self.sessions[session_id] - return None + if time.time() > session.expires_at: + del self.sessions[session_id] + return None - expected_sig = self._generate_signature(session_id, timestamp, session.user_id) - if signature != expected_sig: - return None + expected_sig = self._generate_signature(session_id, timestamp, session.user_id) + if not hmac.compare_digest(signature, expected_sig): + return None - session.last_activity = time.time() + session.last_activity = time.time() - return { - "valid": True, - "session_id": session_id, - "user_id": session.user_id, - "level": session.level, - "permissions": session.permissions - } + return { + "valid": True, + "session_id": session_id, + "user_id": session.user_id, + "level": session.level, + "permissions": session.permissions + } def validate_session_id(self, session_id: str) -> Optional[dict]: """验证会话ID""" if not session_id: return None - session = self.sessions.get(session_id) - if not session: - return None + with self._lock: + session = self.sessions.get(session_id) + if not session: + return None - if time.time() > session.expires_at: - del self.sessions[session_id] - return None + if time.time() > session.expires_at: + del self.sessions[session_id] + return None - session.last_activity = time.time() + session.last_activity = time.time() - return { - "valid": True, - "session_id": session_id, - "user_id": session.user_id, - "level": session.level, - "permissions": session.permissions - } + return { + "valid": True, + "session_id": session_id, + "user_id": session.user_id, + "level": session.level, + "permissions": session.permissions + } # === 请求验证 === @@ -355,8 +394,9 @@ class APIAuthManager: permissions=permissions or ['*'] ) - self.sessions[session_id] = session - self._save_sessions() + with self._lock: + self.sessions[session_id] = session + self._save_sessions() signature = self._generate_signature(session_id, timestamp, user_id) cookie_value = f"{session_id}.{timestamp}.{signature}" @@ -371,19 +411,37 @@ class APIAuthManager: def destroy_session(self, session_id: str) -> bool: """销毁会话""" - if session_id in self.sessions: - del self.sessions[session_id] - self._save_sessions() - return True + with self._lock: + if session_id in self.sessions: + del self.sessions[session_id] + self._save_sessions() + return True return False + def list_sessions(self) -> List[dict]: + """列出活跃会话(加锁遍历,供管理接口使用)""" + now = time.time() + with self._lock: + sessions = [] + for session_id, session in self.sessions.items(): + if now < 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 + }) + return sessions + # === 内部方法 === def _generate_signature(self, session_id: str, timestamp: str, user_id: str) -> str: - """生成签名""" - secret = self.config.get('auth_secret', 'default_secret_change_me') - data = f"{session_id}.{timestamp}.{user_id}.{secret}" - return hashlib.sha256(data.encode()).hexdigest()[:32] + """生成签名(HMAC-SHA256,使用持久化的 auth_secret)""" + import hmac + data = f"{session_id}.{timestamp}.{user_id}".encode('utf-8') + return hmac.new(self.auth_secret.encode('utf-8'), data, hashlib.sha256).hexdigest()[:32] def _load_sessions(self): """加载会话""" @@ -409,31 +467,35 @@ class APIAuthManager: print(f"加载会话失败: {e}") def _save_sessions(self): - """保存会话""" - data = [] - for session in self.sessions.values(): - data.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, - "permissions": session.permissions - }) + """保存会话(原子写: 临时文件 + os.replace)""" + with self._lock: + data = [] + for session in self.sessions.values(): + data.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, + "permissions": session.permissions + }) try: - with open(self.sessions_file, 'w', encoding='utf-8') as f: + tmp_file = self.sessions_file + '.tmp' + with open(tmp_file, 'w', encoding='utf-8') as f: json.dump(data, f, ensure_ascii=False, indent=2) - except Exception: - pass + os.replace(tmp_file, self.sessions_file) + except Exception as e: + print(f"警告: 保存会话失败: {e}") def get_stats(self) -> dict: """获取认证统计""" - active_sessions = sum( - 1 for s in self.sessions.values() - if time.time() < s.expires_at - ) + with self._lock: + active_sessions = sum( + 1 for s in self.sessions.values() + if time.time() < s.expires_at + ) return { "active_sessions": active_sessions, diff --git a/handlers/http_handler.py b/handlers/http_handler.py index febf84d..1da2952 100644 --- a/handlers/http_handler.py +++ b/handlers/http_handler.py @@ -1,1487 +1,1511 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - -"""HTTP请求处理模块""" - -import os -import sys -import json -import re -import time -import mimetypes -import base64 -import hashlib -import shutil -from datetime import datetime -from http.server import BaseHTTPRequestHandler -from urllib.parse import unquote, urlparse, parse_qs - -from core.utils import format_file_size, get_file_hash, sanitize_filename, is_safe_path -from api.router import APIRouter -from mirrors import get_mirror_handler - - -# PyInstaller 资源路径处理 -def get_resource_path(relative_path): - """获取打包后的资源路径""" - if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'): - return os.path.join(sys._MEIPASS, relative_path) - return os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), relative_path) - - -class MirrorServerHandler(BaseHTTPRequestHandler): - """镜像服务器请求处理器""" - - config = None - sync_manager = None - monitor = None # 系统监控器实例 - protocol_version = 'HTTP/1.1' - api_router = None - debug_log_file = None # 调试日志文件路径 - _debug_categories = set() # 启用的调试类别 - _mirror_handlers = {} # 镜像处理器实例缓存 - - @classmethod - def _setup_debug(cls, config): - """根据配置设置调试模式""" - if config is None: - cls._debug_categories = set() - cls.debug_log_file = None - return - - # debug 可以是: - # - true/false: 全局开启/关闭 - # - 列表: 只开启指定的类别 - debug_setting = config.get('debug', False) - if debug_setting is True: - # 全局开启所有 - cls._debug_categories = {'http', 'api', 'auth', 'v2', 'error', 'download'} - elif isinstance(debug_setting, list): - cls._debug_categories = set(debug_setting) - else: - cls._debug_categories = set() - - # 设置 debug 日志文件 - cls.debug_log_file = config.get('debug_log_file') - - @classmethod - def _write_debug_log(cls, msg): - """写入调试日志到文件""" - if cls.debug_log_file: - try: - with open(cls.debug_log_file, 'a', encoding='utf-8') as f: - f.write(msg + '\n') - except Exception: - pass - - def __init__(self, *args, **kwargs): - if 'config' in kwargs: - self.config = kwargs.pop('config') - else: - self.config = MirrorServerHandler.config - - if 'sync_manager' in kwargs: - self.sync_manager = kwargs.pop('sync_manager') - else: - self.sync_manager = MirrorServerHandler.sync_manager - - if 'monitor' in kwargs: - self.monitor = kwargs.pop('monitor') - else: - self.monitor = MirrorServerHandler.monitor - - super().__init__(*args, **kwargs) - - # 初始化API路由(使用共享的 auth_manager) - if self.config is not None: - # 从 config 中获取共享的 auth_manager - self.auth_manager = self.config.get('_auth_manager') - if not self.auth_manager: - # 如果没有,创建新的并保存到 config - MirrorServerHandler.api_router = APIRouter(self.config) - self.auth_manager = self.config.get('_auth_manager') - if self.api_router is None: - MirrorServerHandler.api_router = APIRouter(self.config) - - def _is_debug_enabled(self, category): - """检查特定调试类别是否启用""" - # 如果 debug_categories 为空,检查单个配置 - if not MirrorServerHandler._debug_categories: - return self.config and self.config.get(f'debug_{category}', False) - return category in MirrorServerHandler._debug_categories - - def _debug_log(self, category, msg, color='\033[36m'): - """Output debug log - - If debug_log_file is not set, output to terminal by default - - If debug_log_file is set, output to file only (except errors) - """ - if not self._is_debug_enabled(category): - return - - # 格式化消息 - timestamp = datetime.now().strftime('%Y-%m-%d %H:%M:%S.%f')[:-3] - formatted_msg = f"[DEBUG {timestamp}] [{category.upper()}] {msg}" - - # 写入日志文件 - self._write_debug_log(formatted_msg) - - # 是否在终端输出 - # 如果设置了 debug_log_file,不在终端输出(除非是错误) - debug_log_file = MirrorServerHandler.debug_log_file - if not debug_log_file: - # 没有设置日志文件,默认在终端输出 - print(f"{color}{formatted_msg}\033[0m") - - def log_message(self, format_str, *args): - """自定义日志输出""" - is_verbose = self.config and self.config.get('verbose', 0) > 0 - - # 调试模式输出详细日志 (debug-http) - if self._is_debug_enabled('http'): - timestamp = datetime.now().strftime('%Y-%m-%d %H:%M:%S.%f')[:-3] - msg = f"{self.address_string()} - {format_str % args}" - self._debug_log('http', msg, '\033[36m') - - # 详细模式输出(仅当没有设置 debug_log_file 时在终端输出) - if is_verbose and not self._is_debug_enabled('http'): - msg = f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] {self.address_string()} - {format_str % args}" - print(msg) - - # 访问日志 - if self.config and self.config.get('access_log'): - log_entry = f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] {self.address_string()} - {self.command} {self.path} {self.protocol_version} {self.headers.get('User-Agent', 'Unknown')}\n" - try: - with open(self.config['access_log'], 'a', encoding='utf-8') as f: - f.write(log_entry) - except Exception as e: - print(f"写入访问日志失败: {e}") - - def _handle_mirror_request(self, path: str): - """处理镜像加速源请求""" - # 调试模式输出 (debug-http) - if self._is_debug_enabled('http'): - msg = f"\n=== DEBUG Mirror Request ===\n Path: {path}" - self._debug_log('http', msg, '\033[33m') - - # 确定镜像类型 - if path.startswith("pypi/") or path.startswith("simple/"): - mirror_type = "pypi" - mirror_path = path # 保留完整路径,让 pypi.py 来处理 - elif path.startswith("npm/"): - mirror_type = "npm" - mirror_path = path.replace("npm/", "") - elif path.startswith("go/"): - mirror_type = "go" - mirror_path = path.replace("go/", "") - else: - self.send_error(404, "Unknown mirror type") - return - - # 尝试从Referer中提取镜像名称 - mirror_name = None - referer = self.headers.get('Referer', '') - if 'mirrors/' in referer: - # 例如: http://localhost:8080/api/v2/mirrors/pypi-cn/simple/ - import re - match = re.search(r'mirrors/([^/]+)', referer) - if match: - mirror_name = match.group(1) - - # 获取镜像处理器 - import sys - handler = self._get_mirror_handler(mirror_type, mirror_name) - if not handler: - self.send_error(404, f"Mirror type not available: {mirror_type}") - return - - # 处理请求 - try: - handler.handle_request(self, mirror_path) - except Exception as e: - if self._is_debug_enabled('error'): - import traceback - tb_str = traceback.format_exc() - msg = f"\n=== DEBUG Mirror Handler ERROR ===\n{tb_str}" - self._debug_log('error', msg, '\033[31m') - self.send_error(500, f"Mirror handler error: {str(e)}") - - def _get_mirror_handler(self, mirror_type: str, mirror_name: str = None): - """获取或创建镜像处理器实例""" - import sys - - # 如果指定了镜像名称,优先使用该镜像的配置 - cache_key = f"{mirror_type}:{mirror_name}" if mirror_name else mirror_type - - # 检查缓存 - if cache_key in MirrorServerHandler._mirror_handlers: - return MirrorServerHandler._mirror_handlers[cache_key] - - # 检查配置中是否启用了该镜像 - mirrors_config = self.config.get('mirrors', {}) if self.config else {} - mirror_config = None - - # 优先使用指定的镜像名称 - if mirror_name and mirror_name in mirrors_config: - mirror_config = mirrors_config[mirror_name] - else: - # 否则查找匹配类型的镜像 - for name, config in mirrors_config.items(): - if config.get('type') == mirror_type and config.get('enabled'): - mirror_config = config - break - - if not mirror_config: - return None - - # 创建处理器实例 - handler_class = get_mirror_handler(mirror_type) - if not handler_class: - return None - - # 配置处理器 - 使用 base_dir 作为存储目录基础 - base_dir = self.config.get('base_dir', './downloads') if self.config else './downloads' - # 获取镜像配置的存储目录(相对路径),拼接到 base_dir 下 - storage_subdir = mirror_config.get('storage_dir', mirror_type) - storage_dir = os.path.join(base_dir, storage_subdir) - handler = handler_class({ - 'upstream_url': mirror_config.get('url', ''), - 'storage_dir': storage_dir, - 'base_dir': base_dir - }) - - # 缓存处理器 - MirrorServerHandler._mirror_handlers[cache_key] = handler - return handler - - def check_auth(self, path=None): - """检查认证(方法感知: 读操作与写操作区分对待)""" - # 获取检查路径 - if path is not None: - check_path = path - elif hasattr(self, 'path'): - check_path = unquote(self.path).lstrip('/') - else: - check_path = '' - - # 去掉查询串(do_POST/do_HEAD 传入的路径可能带 ?query) - if '?' in check_path: - check_path = check_path.split('?', 1)[0] - - # 根路径直接放行 - if not check_path or check_path == '/': - return True - - if not self.config: - return True - - auth_type = self.config.get('auth_type', 'none') - if auth_type == 'none': - return True - - method = (getattr(self, 'command', '') or 'GET').upper() - - # 任何方法都公开的端点(登录、认证状态查询等) - public_any = [ - 'api/v2/user/login', - 'api/v2/admin/auth/verify', - ] - - # 只读公开端点(仅 GET/HEAD 放行) - public_get = [ - # 文件只读:列表/搜索/下载/mirror/mc - 'api/v1/files', - 'api/v1/file/', - 'api/v1/search', - 'api/v1/mirror/', - 'api/v1/mc/', - 'api/v1/stats', - 'api/v1/health', - 'api/v1/cache/stats', - # v2只读 - 'api/v2/search/', - 'api/v2/health', - 'api/v2/stats/', - 'api/v2/cache/stats', - 'api/v2/cache/popular', - 'api/v2/api-docs.yaml', - # 镜像加速源代理下载(只读) - 'api/v2/mirrors/pypi', - 'api/v2/mirrors/npm', - 'api/v2/mirrors/go', - 'api/v2/mirrors/docker', - ] - - # 受保护前缀(GET/HEAD 命中也需要认证;写操作另有兜底) - protected_endpoints = [ - # 用户操作 - 'api/v2/user/password', # 改密码 - 'api/v2/user/login-logs', # 登录日志(含IP) - 'api/v2/users', - # 文件修改操作 - 'api/v1/upload', - 'api/v1/mkdir', - 'api/v1/batch', - 'api/v1/archive', - # 同步操作 - 'api/v1/sync/start', - 'api/v1/sync/stop', - 'api/v1/sync/sources', - # v2管理 - 'api/v2/admin/', - 'api/v2/config', - 'api/v2/server/', - 'api/v2/cache/clean', - 'api/v2/cache/prewarm', - 'api/v2/webhooks', - 'api/v2/sync/', - 'api/v2/file/', # 文件删除/重命名/元数据/版本/缩略图 - 'api/v2/mirrors', # 镜像管理(写 settings.json) - 'api/v2/alerts', - 'api/v2/activity', - 'api/v2/monitor', - ] - - # 任何方法都公开 - for endpoint in public_any: - if check_path == endpoint or check_path.startswith(endpoint + '/'): - return True - - if method in ('GET', 'HEAD'): - # 只读公开端点放行 - for endpoint in public_get: - if check_path == endpoint or check_path.startswith(endpoint + '/'): - return True - # 其余 GET: 命中受保护前缀才需要认证 - for endpoint in protected_endpoints: - if check_path.startswith(endpoint): - return self._do_auth(auth_type) - return True - - # 写方法(POST/PUT/DELETE): API 路径一律要求认证,防止新增端点漏配 - if check_path.startswith('api/'): - return self._do_auth(auth_type) - - return True - - def _do_auth(self, auth_type): - """执行指定类型的认证检查""" - if auth_type == 'basic': - return self._check_basic_auth() - elif auth_type == 'token': - return self._check_token_auth() - return False - - def _check_basic_auth(self): - """检查基本认证""" - import hmac - auth_header = self.headers.get('Authorization') - if not auth_header or not auth_header.startswith('Basic '): - self.send_auth_required() - return False - - try: - auth_decoded = base64.b64decode(auth_header[6:]).decode('utf-8') - username, password = auth_decoded.split(':', 1) - expected_user = self.config.get('auth_user', 'admin') if self.config else 'admin' - expected_pass = self.config.get('auth_pass', 'admin123') if self.config else 'admin123' - - # 恒定时间比较,防时序攻击 - if hmac.compare_digest(username, expected_user) and hmac.compare_digest(password, expected_pass): - return True - else: - self.send_auth_required() - return False - except Exception: - self.send_auth_required() - return False - - def _check_token_auth(self): - """检查令牌认证""" - token = None - # 从多个来源获取 token - headers_dict = dict(self.headers) - if 'token' in headers_dict: - token = headers_dict['token'] - elif 'X-API-Key' in headers_dict: - token = headers_dict['X-API-Key'] - elif 'Authorization' in headers_dict and headers_dict['Authorization'].startswith('Bearer '): - token = headers_dict['Authorization'][7:] - elif '?' in self.path: - parsed = urlparse(self.path) - query_params = parse_qs(parsed.query) - token = query_params.get('token', [None])[0] - - if not token: - self.send_json_response({"error": "Invalid or missing token", "code": "UNAUTHORIZED"}, 401) - return False - - # 首先检查是否是会话 token - if hasattr(self, 'auth_manager') and self.auth_manager: - session = self.auth_manager.validate_session_id(token) - if session and session.get('valid'): - return True - - # 检查是否是静态 token - expected_token = self.config.get('auth_token') if self.config else None - import hmac - if token and expected_token and hmac.compare_digest(token, expected_token): - return True - - # 验证失败 - self.send_json_response({"error": "Invalid or missing token", "code": "UNAUTHORIZED"}, 401) - return False - - def send_auth_required(self): - """发送认证要求""" - self.send_response(401) - self.send_header('WWW-Authenticate', 'Basic realm="Mirror Server"') - self.send_header('Content-Type', 'application/json') - self.send_header('Access-Control-Allow-Origin', '*') - self.end_headers() - self.wfile.write(b'{"error": "Authentication Required", "code": "UNAUTHORIZED"}') - - def do_GET(self): - """处理GET请求""" - import sys - sys.stderr.flush() - # 调试模式输出请求详情 (debug-http) - if self._is_debug_enabled('http'): - msg = f"\n=== DEBUG GET Request ===\n Path: {self.path}\n Headers: {dict(self.headers)}" - self._debug_log('http', msg, '\033[33m') - - try: - # 确保配置已加载 - if self.config is None: - self.config = MirrorServerHandler.config - if self.sync_manager is None: - self.sync_manager = MirrorServerHandler.sync_manager - if self.api_router is None and self.config is not None: - MirrorServerHandler.api_router = APIRouter(self.config) - - parsed_path = urlparse(self.path) - path = unquote(parsed_path.path).lstrip('/') - query = parsed_path.query - - # 检查认证(公开端点不需要认证) - if not self.check_auth(path): - return - - # 处理 /api/docs 和 /api/ui 路径(返回静态页面) - if path.startswith("api/docs"): - # 提供 api/docs 目录下的静态文件 - rel_path = path[9:] # 去掉 "api/docs" - if rel_path and not rel_path.startswith('/'): - rel_path = '/' + rel_path - self.serve_docs(rel_path) - return - elif path.startswith("api/ui"): - # 提供 api/ui 目录下的静态文件 - rel_path = path[7:] # 去掉 "api/ui" - if rel_path and not rel_path.startswith('/'): - rel_path = '/' + rel_path - self.serve_ui(rel_path) - return - elif path.startswith("ui/") or path == "ui": - # /ui/ 路径已废弃,返回 404 - self.send_error(404, "UI moved to /api/ui/") - return - elif path.startswith("docs/") or path == "docs": - # /docs/ 路径已废弃,返回 404 - self.send_error(404, "Docs moved to /api/docs/") - return - - # 处理 PyPI 包文件路径 - 转发到 API 路由 - # 这些路径来自 pip 下载请求,如 /pypi/packages/hash/file.tar.gz - if path.startswith("pypi/packages/") or path.startswith("pypi/web/") or path.startswith("pypi/simple/"): - # 转发到 API v2 路由 - api_path = "api/v2/" + path - self.api_router.handle_request(self, 'GET', api_path, query) - return - - # 处理 API 路径 - if path.startswith("api/"): - # 使用API路由处理 - self.api_router.handle_request(self, 'GET', path, query) - # 文件夹/文件访问(/pypi/ 也是本地文件夹) - elif path == "": - # 根路径显示文件列表 - self.serve_path("") - elif path.startswith("file/"): - # 文件下载路由 /file/path/to/file -> serve_path(path/to/file) - file_rel_path = path[5:] # 去掉 "file/" 前缀 - self.serve_path(file_rel_path) - else: - self.serve_path(path) - except Exception as e: - if self._is_debug_enabled('error'): - import traceback - tb_str = traceback.format_exc() - msg = f"\n=== DEBUG GET ERROR ===\n{tb_str}" - self._debug_log('error', msg, '\033[31m') - self.handle_error(500, f"服务器内部错误: {str(e)}") - - def do_POST(self): - """处理POST请求""" - import time - request_id = int(time.time() * 1000000) - - # 调试模式输出请求详情 (debug-http) - if self._is_debug_enabled('http'): - content_length = self.headers.get('Content-Length', 0) - msg = f"\n=== DEBUG POST Request #{request_id} ===\n Path: {self.path}\n Content-Length: {content_length}" - self._debug_log('http', msg, '\033[33m') - - try: - # 确保配置已加载 - if self.config is None: - self.config = MirrorServerHandler.config - if self.api_router is None and self.config is not None: - MirrorServerHandler.api_router = APIRouter(self.config) - - # 解析路径(去掉查询串,避免把 ?query 拼进 API 路径) - parsed_path = urlparse(self.path) - path = unquote(parsed_path.path).lstrip('/') - query = parsed_path.query - - # 调试模式输出完整路径 (debug-http) - if self._is_debug_enabled('http'): - msg = f"\n=== DEBUG POST Path Check ===\n path: '{path}'\n starts with api/: {path.startswith('api/')}" - self._debug_log('http', msg, '\033[33m') - - # 检查认证(公开端点不需要认证) - if not self.check_auth(path): - return - - if path.startswith("api/"): - self.api_router.handle_request(self, 'POST', path, query) - else: - self.send_error(405) - except Exception as e: - if self._is_debug_enabled('error'): - import traceback - tb_str = traceback.format_exc() - msg = f"\n=== DEBUG POST ERROR ===\n{tb_str}" - self._debug_log('error', msg, '\033[31m') - self.handle_error(500, f"服务器内部错误: {str(e)}") - - def do_OPTIONS(self): - """处理OPTIONS请求(CORS预检)""" - # 调试模式输出 (debug-http) - if self._is_debug_enabled('http'): - msg = f"\n=== DEBUG OPTIONS Request ===\n Path: {self.path}" - self._debug_log('http', msg, '\033[34m') - - self.send_response(200) - self.send_header('Access-Control-Allow-Origin', '*') - self.send_header('Access-Control-Allow-Methods', - 'GET, POST, PUT, DELETE, OPTIONS') - self.send_header('Access-Control-Allow-Headers', - 'Content-Type, Authorization, Token') - self.send_header('Access-Control-Max-Age', '86400') - self.end_headers() - - def do_DELETE(self): - """处理DELETE请求""" - try: - # 确保配置已加载 - if self.config is None: - self.config = MirrorServerHandler.config - if self.api_router is None and self.config is not None: - MirrorServerHandler.api_router = APIRouter(self.config) - - if not self.check_auth(): - return - parsed_path = urlparse(self.path) - path = unquote(parsed_path.path).lstrip('/') - query = parsed_path.query - if path.startswith("api/"): - self.api_router.handle_request(self, 'DELETE', path, query) - else: - self.send_error(405) - except Exception as e: - if self._is_debug_enabled('error'): - import traceback - tb_str = traceback.format_exc() - msg = f"\n=== DEBUG DELETE ERROR ===\n{tb_str}" - self._debug_log('error', msg, '\033[31m') - self.handle_error(500, f"服务器内部错误: {str(e)}") - - def do_PUT(self): - """处理PUT请求""" - try: - # 确保配置已加载 - if self.config is None: - self.config = MirrorServerHandler.config - if self.api_router is None and self.config is not None: - MirrorServerHandler.api_router = APIRouter(self.config) - - if not self.check_auth(): - return - parsed_path = urlparse(self.path) - path = unquote(parsed_path.path).lstrip('/') - query = parsed_path.query - if path.startswith("api/"): - self.api_router.handle_request(self, 'PUT', path, query) - else: - self.send_error(405) - except Exception as e: - if self._is_debug_enabled('error'): - import traceback - tb_str = traceback.format_exc() - msg = f"\n=== DEBUG PUT ERROR ===\n{tb_str}" - self._debug_log('error', msg, '\033[31m') - self.handle_error(500, f"服务器内部错误: {str(e)}") - - def do_HEAD(self): - """处理HEAD请求""" - if not self.check_auth(): - return - # 解析路径(去掉查询串),并复用 serve_path 的安全检查防路径穿越 - parsed_path = urlparse(self.path) - rel_path = unquote(parsed_path.path).lstrip('/') - if self.config is None: - self.send_error(500) - return - file_path = os.path.join(self.config['base_dir'], rel_path) - if not is_safe_path(self.config['base_dir'], file_path): - self.send_error(403, "Access denied") - return - if os.path.isfile(file_path): - self.send_file_headers(file_path) - else: - self.send_error(404) - - # ==================== 静态文件服务 ==================== - - def serve_docs(self, rel_path): - """提供 api/docs 目录下的静态文件""" - docs_dir = get_resource_path('api/docs') - - if not os.path.isdir(docs_dir): - self.send_error(404, "Docs directory not found") - return - - # 默认提供 index.html - if not rel_path or rel_path == '/': - rel_path = 'index.html' - - file_path = os.path.join(docs_dir, rel_path) - - # 防止目录遍历 - if not os.path.realpath(file_path).startswith(os.path.realpath(docs_dir)): - self.send_error(403, "Access denied") - return - - if os.path.isfile(file_path): - self.serve_file(file_path, f'api/docs/{rel_path}') - else: - # 提供 docs 目录索引 - self.serve_docs_index(docs_dir) - - def serve_docs_index(self, docs_dir): - """提供 docs 目录索引页面""" - try: - items = [] - for name in sorted(os.listdir(docs_dir)): - full_path = os.path.join(docs_dir, name) - rel_path = f'docs/{name}' - is_dir = os.path.isdir(full_path) - items.append({ - "name": name, - "path": rel_path + ("/" if is_dir else ""), - "is_dir": is_dir - }) - except OSError: - self.send_error(403) - return - - # 生成 HTML - title = "HYC下载站 - 文档" - html = f''' - - - - - {title} - - - -

{title}

-''' - for item in items: - if item['is_dir']: - html += f''' -
- {item['name']}/ -
目录
-
''' - else: - # 根据文件类型添加描述 - desc = "" - if item['name'].endswith('.md'): - desc = "Markdown 文档" - elif item['name'].endswith('.yaml') or item['name'].endswith('.yml'): - desc = "OpenAPI 配置" - elif item['name'].endswith('.json'): - desc = "JSON 配置" - else: - desc = "文件" - html += f''' -
- {item['name']} -
{desc}
-
''' - - html += ''' - -''' - - self.send_response(200) - self.send_header("Content-Type", "text/html; charset=utf-8") - self.send_header("Content-Length", str(len(html.encode('utf-8')))) - self.end_headers() - self.wfile.write(html.encode('utf-8')) - - def serve_ui(self, rel_path): - """提供 api/ui 目录下的静态文件""" - ui_dir = get_resource_path('api/ui') - - if not os.path.isdir(ui_dir): - self.send_error(404, "UI directory not found") - return - - # 默认提供 index.html - if not rel_path or rel_path == '/': - rel_path = 'index.html' - - file_path = os.path.join(ui_dir, rel_path) - - # 防止目录遍历 - if not os.path.realpath(file_path).startswith(os.path.realpath(ui_dir)): - self.send_error(403, "Access denied") - return - - if os.path.isfile(file_path): - self.serve_file(file_path, f'api/ui/{rel_path}') - else: - self.send_error(404, f"File not found: {rel_path}") - - def serve_path(self, rel_path): - """处理路径请求(文件或目录)""" - if self.config is None: - self.send_error(500) - return - - file_path = os.path.join(self.config['base_dir'], rel_path) - - if not is_safe_path(self.config['base_dir'], file_path): - self.send_error(403, "Access denied") - return - - if os.path.isdir(file_path): - if self.config.get('directory_listing', True): - self.serve_directory(file_path, rel_path) - else: - self.send_error(403, "Directory listing is disabled") - elif os.path.isfile(file_path): - self.serve_file(file_path, rel_path) - else: - self.send_error(404) - - def serve_directory(self, dir_path, rel_dir): - """提供目录浏览(镜像站风格)""" - try: - # 检查是否有索引文件 - index_files = ['index.html', 'index.htm'] - for index_file in index_files: - index_path = os.path.join(dir_path, index_file) - if os.path.isfile(index_path): - self.serve_file(index_path, os.path.join(rel_dir, index_file)) - return - - # 获取目录内容 - items = [] - for name in os.listdir(dir_path): - full_path = os.path.join(dir_path, name) - rel_item_path = os.path.join(rel_dir, name).replace("\\", "/") - is_dir = os.path.isdir(full_path) - - if self.config.get('ignore_hidden', True) and name.startswith('.'): - continue - - try: - size = "-" if is_dir else format_file_size(os.path.getsize(full_path)) - mtime = datetime.fromtimestamp(os.path.getmtime(full_path)).strftime("%Y-%m-%d %H:%M") - - sha256 = "" - if self.config.get('show_hash') and not is_dir: - sha256 = get_file_hash(full_path)[:16] + "..." - items.append({ - "name": name, - "path": rel_item_path + ("/" if is_dir else ""), - "size": size, - "modified": mtime, - "is_dir": is_dir, - "sha256": sha256 - }) - except OSError: - continue - - # 排序 - Windows 文件管理器风格:文件夹在前,按名称递增排序 - sort_by = self.config.get('sort_by', 'name') - reverse = self.config.get('sort_reverse', False) # 默认为 False(递增) - if sort_by == 'name': - # 文件夹优先,然后按名称递增排序(不区分大小写) - items.sort(key=lambda x: (not x["is_dir"], x["name"].lower()), reverse=False) - elif sort_by == 'size': - # 文件夹优先,然后按大小递增排序 - items.sort(key=lambda x: (not x["is_dir"], os.path.getsize(os.path.join(dir_path, x["name"])) if not x["is_dir"] else 0), reverse=False) - elif sort_by == 'modified': - # 文件夹优先,然后按修改时间递增排序 - items.sort(key=lambda x: (not x["is_dir"], os.path.getmtime(os.path.join(dir_path, x["name"]))), reverse=False) - - except OSError: - self.send_error(403) - return - - # 构建面包屑 - breadcrumbs = [] - parts = [p for p in rel_dir.split("/") if p] - current = "" - breadcrumbs.append({"name": "HOME", "path": "/"}) - for part in parts: - current = os.path.join(current, part).replace("\\", "/") - breadcrumbs.append({"name": part, "path": "/" + current + "/"}) - - # 生成HTML - title = "HYC下载站" - html = self._generate_directory_html(title, breadcrumbs, items, rel_dir) - - self.send_response(200) - self.send_header("Content-Type", "text/html; charset=utf-8") - self.send_header("Content-Length", str(len(html.encode('utf-8')))) - self.send_header("Cache-Control", "no-cache") - self.end_headers() - self.wfile.write(html.encode('utf-8')) - - def _generate_directory_html(self, title, breadcrumbs, items, rel_dir): - """生成目录浏览HTML""" - if rel_dir: # 如果不是根目录 - parts = [p for p in rel_dir.split('/') if p] # 过滤空部分 - if len(parts) > 1: - parent_path = '/' + '/'.join(parts[:-1]) + '/' - elif len(parts) == 1: - parent_path = '/' - else: - parent_path = '/' - else: - parent_path = '/' # 根目录没有上一级 - - # 动态计算列数 - colspan = 4 if self.config.get('show_hash') else 3 - - html = f""" - - - - {title} - - - - -
-

{title}

- - - - - - - - {('' if self.config.get('show_hash') else '')} - - - """ - - # 修复:只在非根目录显示上一级目录链接 - if rel_dir: # 如果不是根目录 - html += f'\n' - - for item in items: - html += f'' - html += f'' - html += f'' - html += f'' - if self.config.get('show_hash'): - html += f'' - html += '\n' - - html += f""" - -
NameLast ModifiedSizeSHA256
../
{item["name"]}{" /" if item["is_dir"] else ""}{item["modified"]}{item["size"]}{item["sha256"]}
-
-

Files: {len(items)} | {self.config.get("server_name", "Mirror Server")}

-
-
- - """ - return html - - def send_file_headers(self, file_path): - """发送文件头信息(用于HEAD请求)""" - if not os.path.exists(file_path) or not os.path.isfile(file_path): - self.send_error(404) - return - file_size = os.path.getsize(file_path) - mime_type, _ = mimetypes.guess_type(file_path) - if mime_type is None: - mime_type = "application/octet-stream" - - self.send_response(200) - self.send_header("Content-Type", mime_type) - self.send_header("Content-Length", str(file_size)) - self.send_header("Content-Disposition", - f'attachment; filename="{os.path.basename(file_path)}"') - self.send_header("Accept-Ranges", "bytes") - self.send_header("Cache-Control", "public, max-age=3600") - self.send_header( - "Last-Modified", self.date_time_string(os.path.getmtime(file_path))) - self.end_headers() - - def serve_file(self, file_path, rel_path): - """提供文件下载,支持断点续传和流式传输""" - if not os.path.exists(file_path) or not os.path.isfile(file_path): - self.send_error(404) - return - - # 获取文件信息 - file_size = os.path.getsize(file_path) - mime_type, _ = mimetypes.guess_type(file_path) - if mime_type is None: - mime_type = "application/octet-stream" - - # 获取客户端IP - client_ip = self.client_address[0] if hasattr(self, 'client_address') else 'unknown' - - # 检查Range头部(支持断点续传) - range_header = self.headers.get('Range') - range_start = 0 - range_end = file_size - 1 - - if range_header and self.config.get('enable_range', True): - match = re.match(r'bytes=(\d+)-(\d*)', range_header) - if match: - range_start = int(match.group(1)) - range_end_str = match.group(2) - if range_end_str: - range_end = int(range_end_str) - - if range_start >= file_size or range_end >= file_size or range_start > range_end: - self.send_error(416, "Requested Range Not Satisfiable") - return - - # 计算传输内容 - content_length = range_end - range_start + 1 - - # 发送响应头 - if range_start == 0 and range_end == file_size - 1: - # 完整文件下载 - self.send_response(200) - else: - # 部分内容(206) - self.send_response(206) - self.send_header("Content-Range", f"bytes {range_start}-{range_end}/{file_size}") - - self.send_header("Content-Type", mime_type) - self.send_header("Content-Length", str(content_length)) - # HTML 文件直接在浏览器中显示,不强制下载 - if mime_type == 'text/html': - self.send_header("Content-Disposition", f'inline; filename="{os.path.basename(file_path)}"') - else: - self.send_header("Content-Disposition", - f'attachment; filename="{os.path.basename(file_path)}"') - self.send_header("Accept-Ranges", "bytes") - self.send_header("Cache-Control", "public, max-age=3600") - self.send_header("Last-Modified", self.date_time_string(os.path.getmtime(file_path))) - self.send_header("X-Download-IP", client_ip) - self.send_header("X-File-Size", str(file_size)) - self.end_headers() - - # 流式传输文件 - chunk_size = 64 * 1024 # 64KB chunks for better performance - bytes_sent = 0 - - try: - with open(file_path, 'rb') as f: - f.seek(range_start) - remaining = content_length - - while remaining > 0: - chunk = f.read(min(chunk_size, remaining)) - if not chunk: - break - - self.wfile.write(chunk) - bytes_sent += len(chunk) - remaining -= len(chunk) - - except Exception as e: - print(f"文件传输错误: {e}") - - # 只对真正的下载(非 HTML 页面)更新统计和记录 - if self.config.get('enable_stats', True) and mime_type != 'text/html': - self.record_download(rel_path, file_size, client_ip) - - def record_download(self, filepath, file_size=0, client_ip='unknown'): - """记录下载(同时更新计数和创建下载记录)""" - if not self.config.get('enable_stats', True): - if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): - self._debug_log('download', f"Stats disabled, skipping download record for: {filepath}") - return - - db = self._get_db() - if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): - self._debug_log('download', f"record_download called for: {filepath}, db: {db}") - user_agent = self.headers.get('User-Agent', 'Unknown') if hasattr(self, 'headers') else 'Unknown' - - if db: - try: - # 尝试更新 FileRecord 的下载计数(通过路径查找) - try: - record = db.get_file_by_path(filepath) - if record: - db.increment_download_count(record.file_id) - except Exception as e: - pass # 忽略更新计数错误 - - # 创建下载记录 - try: - new_record = db.add_download_record( - file_path=filepath, - file_size=file_size, - client_ip=client_ip, - user_agent=user_agent, - success=True - ) - if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): - self._debug_log('download', f"Download record created successfully: {filepath}") - except Exception as e: - if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): - self._debug_log('download', f"Error creating download record: {e}") - except Exception as e: - if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): - self._debug_log('download', f"Error recording download: {e}") - - def _serve_file_chunked(self, file_path, rel_path, chunk_size=64*1024): - """流式分块传输文件(用于大文件)- 备用功能""" - if not os.path.exists(file_path) or not os.path.isfile(file_path): - self.send_error(404) - return - - file_size = os.path.getsize(file_path) - mime_type, _ = mimetypes.guess_type(file_path) - self.send_response(200) - self.send_header("Content-Type", mime_type) - self.send_header("Content-Length", str(file_size)) - self.send_header("Content-Disposition", - f'attachment; filename="{os.path.basename(file_path)}"') - self.send_header("Accept-Ranges", "bytes") - self.send_header("Transfer-Encoding", "chunked") - self.end_headers() - - try: - with open(file_path, 'rb') as f: - while True: - chunk = f.read(chunk_size) - if not chunk: - break - self.wfile.write(chunk) - except Exception as e: - print(f"流式传输错误: {e}") - - # 更新统计(serve_file_chunked 只用于真正的下载) - if self.config.get('enable_stats', True): - self.update_download_count(rel_path) - - def send_json_response(self, data, status_code=200): - """发送JSON响应""" - json_data = json.dumps(data, ensure_ascii=False, indent=2) - self.send_response(status_code) - self.send_header("Content-Type", "application/json; charset=utf-8") - self.send_header("Content-Length", str(len(json_data.encode('utf-8')))) - self.send_header("Cache-Control", "no-cache") - self.end_headers() - self.wfile.write(json_data.encode('utf-8')) - - def date_time_string(self, timestamp=None): - """重写日期时间字符串格式化""" - if timestamp is None: - timestamp = time.time() - return datetime.fromtimestamp(timestamp).strftime('%a, %d %b %Y %H:%M:%S GMT') - - def handle_error(self, code, message=None): - """自定义错误处理""" - error_messages = { - 400: "错误的请求", - 401: "未经授权", - 403: "禁止访问", - 404: "文件未找到", - 405: "方法不允许", - 413: "文件太大", - 416: "请求范围不符合要求", - 500: "内部服务器错误" - } - - if message is None: - message = error_messages.get(code, "未知错误") - - error_page = f""" - - - - - - {code} {message} - - - -
-

{code}

-

{message}

-

请求的页面遇到问题,请稍后重试.

- 返回首页 -
- -""" - - self.send_response(code) - self.send_header("Content-Type", "text/html; charset=utf-8") - self.send_header("Content-Length", - str(len(error_page.encode('utf-8')))) - self.end_headers() - self.wfile.write(error_page.encode('utf-8')) - - def send_error(self, code, message=None): - """发送错误响应""" - self.handle_error(code, message) - - # 统计相关方法 - def _get_db(self): - """获取数据库实例""" - if hasattr(self, 'config') and self.config: - return self.config.get('_db_instance') - return None - - def load_stats(self): - """加载下载统计信息(优先使用数据库,回退到JSON)""" - db = self._get_db() - if db: - try: - # 使用专门的方法获取下载统计,避免会话问题 - return db.get_download_stats(limit=10000) - except Exception as e: - print(f"Error loading stats from database: {e}") - - # 回退到 JSON 文件 - stats_file = self.config.get('stats_file', 'stats.json') - try: - if os.path.exists(stats_file): - with open(stats_file, 'r', encoding='utf-8') as f: - return json.load(f) - except Exception as e: - print(f"Error loading stats: {e}") - return {} - - def save_stats(self, stats): - """保存下载统计信息(优先使用数据库,回退到JSON)""" - db = self._get_db() - if db: - # 数据库模式下,stats 由数据库直接管理,不需要手动保存 - return - - # 回退到 JSON 文件 - stats_file = self.config.get('stats_file', 'stats.json') - try: - with open(stats_file, 'w', encoding='utf-8') as f: - json.dump(stats, f, ensure_ascii=False, indent=2) - except Exception as e: - print(f"Error saving stats: {e}") - - def get_download_count(self, filepath): - """获取特定文件的下载次数""" - db = self._get_db() - if db: - try: - file_record = db.get_file_by_path(filepath) - if file_record: - return file_record.download_count if hasattr(file_record, 'download_count') else 0 - except Exception as e: - print(f"Error getting download count from database: {e}") - - stats = self.load_stats() - return stats.get(filepath, 0) - - def get_total_downloads(self): - """获取总下载次数""" - db = self._get_db() - if db: - try: - stats = db.get_stats() - return stats.get('total_downloads', 0) - except Exception as e: - print(f"Error getting total downloads from database: {e}") - - stats = self.load_stats() - return sum(stats.values()) - - def update_download_count(self, filepath): - """更新文件的下载计数(优先使用数据库,通过路径查找)""" - if not self.config.get('enable_stats', True): - return - - db = self._get_db() - if db: - try: - # 通过路径查找记录,获取 file_id - record = db.get_file_by_path(filepath) - if record: - db.increment_download_count(record.file_id) - return - 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) - - # ==================== 下载历史记录 ==================== - - def load_download_history(self, limit=100): - """加载下载历史记录(优先使用数据库,回退到JSON)""" - db = self._get_db() - if db: - try: - records = db.get_download_records(limit=limit) - history = [] - for r in records: - history.append({ - 'timestamp': r.created_at.isoformat() if hasattr(r.created_at, 'isoformat') else str(r.created_at), - 'filepath': r.file_path if hasattr(r, 'file_path') else getattr(r, 'filepath', str(r)), - 'file_size': r.file_size if hasattr(r, 'file_size') else 0, - 'client_ip': r.client_ip if hasattr(r, 'client_ip') else 'unknown', - 'user_agent': r.user_agent if hasattr(r, 'user_agent') else 'Unknown', - 'method': 'GET' - }) - return history - except Exception as e: - print(f"Error loading download history from database: {e}") - - # 回退到 JSON 文件 - history_file = self.config.get('download_history_file', 'download_history.json') - try: - if os.path.exists(history_file): - with open(history_file, 'r', encoding='utf-8') as f: - data = json.load(f) - return data[-limit:] - except Exception as e: - print(f"Error loading download history: {e}") - return [] - - def save_download_history(self, history): - """保存下载历史记录(数据库模式下不需要)""" - db = self._get_db() - if db: - # 数据库模式下,history 由数据库直接管理 - return - - # 回退到 JSON 文件 - history_file = self.config.get('download_history_file', 'download_history.json') - max_history = self.config.get('max_history_count', 1000) - - try: - # 保留最近的记录 - history = history[-max_history:] - with open(history_file, 'w', encoding='utf-8') as f: - json.dump(history, f, ensure_ascii=False, indent=2) - except Exception as e: - print(f"Error saving download history: {e}") - - def _log_download(self, filepath, file_size=0): - """记录下载历史(优先使用数据库)- 备用功能""" - if not self.config.get('enable_stats', True): - return - - client_ip = self.client_address[0] if hasattr(self, 'client_address') else 'unknown' - user_agent = self.headers.get('User-Agent', 'Unknown') - - db = self._get_db() - if db: - try: - db.add_download_record( - file_path=filepath, - file_size=file_size, - client_ip=client_ip, - user_agent=user_agent, - success=True, - duration=0 - ) - return - except Exception as e: - print(f"Error logging download to database: {e}") - - # 回退到 JSON 文件 - history = self.load_download_history(1000) - - entry = { - 'timestamp': datetime.now().isoformat(), - 'filepath': filepath, - 'file_size': file_size, - 'client_ip': client_ip, - 'user_agent': user_agent, - 'method': self.command if hasattr(self, 'command') else 'GET' - } - - history.append(entry) - self.save_download_history(history) +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +"""HTTP请求处理模块""" + +import os +import sys +import json +import re +import time +import mimetypes +import base64 +import hashlib +import shutil +from datetime import datetime +from http.server import BaseHTTPRequestHandler +from urllib.parse import unquote, urlparse, parse_qs + +from core.utils import format_file_size, get_file_hash, sanitize_filename, is_safe_path +from api.router import APIRouter +from mirrors import get_mirror_handler + + +# PyInstaller 资源路径处理 +def get_resource_path(relative_path): + """获取打包后的资源路径""" + if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'): + return os.path.join(sys._MEIPASS, relative_path) + return os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), relative_path) + + +class MirrorServerHandler(BaseHTTPRequestHandler): + """镜像服务器请求处理器""" + + config = None + sync_manager = None + monitor = None # 系统监控器实例 + protocol_version = 'HTTP/1.1' + api_router = None + debug_log_file = None # 调试日志文件路径 + _debug_categories = set() # 启用的调试类别 + _mirror_handlers = {} # 镜像处理器实例缓存 + + @classmethod + def _setup_debug(cls, config): + """根据配置设置调试模式""" + if config is None: + cls._debug_categories = set() + cls.debug_log_file = None + return + + # debug 可以是: + # - true/false: 全局开启/关闭 + # - 列表: 只开启指定的类别 + debug_setting = config.get('debug', False) + if debug_setting is True: + # 全局开启所有 + cls._debug_categories = {'http', 'api', 'auth', 'v2', 'error', 'download'} + elif isinstance(debug_setting, list): + cls._debug_categories = set(debug_setting) + else: + cls._debug_categories = set() + + # 设置 debug 日志文件 + cls.debug_log_file = config.get('debug_log_file') + + @classmethod + def _write_debug_log(cls, msg): + """写入调试日志到文件""" + if cls.debug_log_file: + try: + with open(cls.debug_log_file, 'a', encoding='utf-8') as f: + f.write(msg + '\n') + except Exception: + pass + + def __init__(self, *args, **kwargs): + if 'config' in kwargs: + self.config = kwargs.pop('config') + else: + self.config = MirrorServerHandler.config + + if 'sync_manager' in kwargs: + self.sync_manager = kwargs.pop('sync_manager') + else: + self.sync_manager = MirrorServerHandler.sync_manager + + if 'monitor' in kwargs: + self.monitor = kwargs.pop('monitor') + else: + self.monitor = MirrorServerHandler.monitor + + super().__init__(*args, **kwargs) + + # 初始化API路由(使用共享的 auth_manager) + if self.config is not None: + # 从 config 中获取共享的 auth_manager + self.auth_manager = self.config.get('_auth_manager') + if not self.auth_manager: + # 如果没有,创建新的并保存到 config + MirrorServerHandler.api_router = APIRouter(self.config) + self.auth_manager = self.config.get('_auth_manager') + if self.api_router is None: + MirrorServerHandler.api_router = APIRouter(self.config) + + def _is_debug_enabled(self, category): + """检查特定调试类别是否启用""" + # 如果 debug_categories 为空,检查单个配置 + if not MirrorServerHandler._debug_categories: + return self.config and self.config.get(f'debug_{category}', False) + return category in MirrorServerHandler._debug_categories + + def _debug_log(self, category, msg, color='\033[36m'): + """Output debug log + - If debug_log_file is not set, output to terminal by default + - If debug_log_file is set, output to file only (except errors) + """ + if not self._is_debug_enabled(category): + return + + # 格式化消息 + timestamp = datetime.now().strftime('%Y-%m-%d %H:%M:%S.%f')[:-3] + formatted_msg = f"[DEBUG {timestamp}] [{category.upper()}] {msg}" + + # 写入日志文件 + self._write_debug_log(formatted_msg) + + # 是否在终端输出 + # 如果设置了 debug_log_file,不在终端输出(除非是错误) + debug_log_file = MirrorServerHandler.debug_log_file + if not debug_log_file: + # 没有设置日志文件,默认在终端输出 + print(f"{color}{formatted_msg}\033[0m") + + def log_message(self, format_str, *args): + """自定义日志输出""" + is_verbose = self.config and self.config.get('verbose', 0) > 0 + + # 调试模式输出详细日志 (debug-http) + if self._is_debug_enabled('http'): + timestamp = datetime.now().strftime('%Y-%m-%d %H:%M:%S.%f')[:-3] + msg = f"{self.address_string()} - {format_str % args}" + self._debug_log('http', msg, '\033[36m') + + # 详细模式输出(仅当没有设置 debug_log_file 时在终端输出) + if is_verbose and not self._is_debug_enabled('http'): + msg = f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] {self.address_string()} - {format_str % args}" + print(msg) + + # 访问日志 + if self.config and self.config.get('access_log'): + log_entry = f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] {self.address_string()} - {self.command} {self.path} {self.protocol_version} {self.headers.get('User-Agent', 'Unknown')}\n" + try: + with open(self.config['access_log'], 'a', encoding='utf-8') as f: + f.write(log_entry) + except Exception as e: + print(f"写入访问日志失败: {e}") + + def _handle_mirror_request(self, path: str): + """处理镜像加速源请求""" + # 调试模式输出 (debug-http) + if self._is_debug_enabled('http'): + msg = f"\n=== DEBUG Mirror Request ===\n Path: {path}" + self._debug_log('http', msg, '\033[33m') + + # 确定镜像类型 + if path.startswith("pypi/") or path.startswith("simple/"): + mirror_type = "pypi" + mirror_path = path # 保留完整路径,让 pypi.py 来处理 + elif path.startswith("npm/"): + mirror_type = "npm" + mirror_path = path.replace("npm/", "") + elif path.startswith("go/"): + mirror_type = "go" + mirror_path = path.replace("go/", "") + else: + self.send_error(404, "Unknown mirror type") + return + + # 尝试从Referer中提取镜像名称 + mirror_name = None + referer = self.headers.get('Referer', '') + if 'mirrors/' in referer: + # 例如: http://localhost:8080/api/v2/mirrors/pypi-cn/simple/ + import re + match = re.search(r'mirrors/([^/]+)', referer) + if match: + mirror_name = match.group(1) + + # 获取镜像处理器 + import sys + handler = self._get_mirror_handler(mirror_type, mirror_name) + if not handler: + self.send_error(404, f"Mirror type not available: {mirror_type}") + return + + # 处理请求 + try: + handler.handle_request(self, mirror_path) + except Exception as e: + if self._is_debug_enabled('error'): + import traceback + tb_str = traceback.format_exc() + msg = f"\n=== DEBUG Mirror Handler ERROR ===\n{tb_str}" + self._debug_log('error', msg, '\033[31m') + self.send_error(500, f"Mirror handler error: {str(e)}") + + def _get_mirror_handler(self, mirror_type: str, mirror_name: str = None): + """获取或创建镜像处理器实例""" + import sys + + # 如果指定了镜像名称,优先使用该镜像的配置 + cache_key = f"{mirror_type}:{mirror_name}" if mirror_name else mirror_type + + # 检查缓存 + if cache_key in MirrorServerHandler._mirror_handlers: + return MirrorServerHandler._mirror_handlers[cache_key] + + # 检查配置中是否启用了该镜像 + mirrors_config = self.config.get('mirrors', {}) if self.config else {} + mirror_config = None + + # 优先使用指定的镜像名称 + if mirror_name and mirror_name in mirrors_config: + mirror_config = mirrors_config[mirror_name] + else: + # 否则查找匹配类型的镜像 + for name, config in mirrors_config.items(): + if config.get('type') == mirror_type and config.get('enabled'): + mirror_config = config + break + + if not mirror_config: + return None + + # 创建处理器实例 + handler_class = get_mirror_handler(mirror_type) + if not handler_class: + return None + + # 配置处理器 - 使用 base_dir 作为存储目录基础 + base_dir = self.config.get('base_dir', './downloads') if self.config else './downloads' + # 获取镜像配置的存储目录(相对路径),拼接到 base_dir 下 + storage_subdir = mirror_config.get('storage_dir', mirror_type) + storage_dir = os.path.join(base_dir, storage_subdir) + handler = handler_class({ + 'upstream_url': mirror_config.get('url', ''), + 'storage_dir': storage_dir, + 'base_dir': base_dir + }) + + # 缓存处理器 + MirrorServerHandler._mirror_handlers[cache_key] = handler + return handler + + def check_auth(self, path=None): + """检查认证(方法感知: 读操作与写操作区分对待)""" + # 获取检查路径 + if path is not None: + check_path = path + elif hasattr(self, 'path'): + check_path = unquote(self.path).lstrip('/') + else: + check_path = '' + + # 去掉查询串(do_POST/do_HEAD 传入的路径可能带 ?query) + if '?' in check_path: + check_path = check_path.split('?', 1)[0] + + # 根路径直接放行 + if not check_path or check_path == '/': + return True + + if not self.config: + return True + + auth_type = self.config.get('auth_type', 'none') + if auth_type == 'none': + return True + + method = (getattr(self, 'command', '') or 'GET').upper() + + # 任何方法都公开的端点(登录、认证状态查询等) + public_any = [ + 'api/v2/user/login', + 'api/v2/admin/auth/verify', + ] + + # 只读公开端点(仅 GET/HEAD 放行) + public_get = [ + # 文件只读:列表/搜索/下载/mirror/mc + 'api/v1/files', + 'api/v1/file/', + 'api/v1/search', + 'api/v1/mirror/', + 'api/v1/mc/', + 'api/v1/stats', + 'api/v1/health', + 'api/v1/cache/stats', + # v2只读 + 'api/v2/search/', + 'api/v2/health', + 'api/v2/stats/', + 'api/v2/cache/stats', + 'api/v2/cache/popular', + 'api/v2/api-docs.yaml', + # 镜像加速源代理下载(只读) + 'api/v2/mirrors/pypi', + 'api/v2/mirrors/npm', + 'api/v2/mirrors/go', + 'api/v2/mirrors/docker', + ] + + # 受保护前缀(GET/HEAD 命中也需要认证;写操作另有兜底) + protected_endpoints = [ + # 用户操作 + 'api/v2/user/password', # 改密码 + 'api/v2/user/login-logs', # 登录日志(含IP) + 'api/v2/users', + # 文件修改操作 + 'api/v1/upload', + 'api/v1/mkdir', + 'api/v1/batch', + 'api/v1/archive', + # 同步操作 + 'api/v1/sync/start', + 'api/v1/sync/stop', + 'api/v1/sync/sources', + # v2管理 + 'api/v2/admin/', + 'api/v2/config', + 'api/v2/server/', + 'api/v2/cache/clean', + 'api/v2/cache/prewarm', + 'api/v2/webhooks', + 'api/v2/sync/', + 'api/v2/file/', # 文件删除/重命名/元数据/版本/缩略图 + 'api/v2/mirrors', # 镜像管理(写 settings.json) + 'api/v2/alerts', + 'api/v2/activity', + 'api/v2/monitor', + ] + + # 任何方法都公开 + for endpoint in public_any: + if check_path == endpoint or check_path.startswith(endpoint + '/'): + return True + + if method in ('GET', 'HEAD'): + # 只读公开端点放行 + for endpoint in public_get: + if check_path == endpoint or check_path.startswith(endpoint + '/'): + return True + # 其余 GET: 命中受保护前缀才需要认证 + for endpoint in protected_endpoints: + if check_path.startswith(endpoint): + return self._do_auth(auth_type) + return True + + # 写方法(POST/PUT/DELETE): API 路径一律要求认证,防止新增端点漏配 + if check_path.startswith('api/'): + return self._do_auth(auth_type) + + return True + + def _do_auth(self, auth_type): + """执行指定类型的认证检查""" + if auth_type == 'basic': + return self._check_basic_auth() + elif auth_type == 'token': + return self._check_token_auth() + return False + + def _check_basic_auth(self): + """检查基本认证""" + import hmac + auth_header = self.headers.get('Authorization') + if not auth_header or not auth_header.startswith('Basic '): + self.send_auth_required() + return False + + try: + auth_decoded = base64.b64decode(auth_header[6:]).decode('utf-8') + username, password = auth_decoded.split(':', 1) + expected_user = self.config.get('auth_user', 'admin') if self.config else 'admin' + expected_pass = self.config.get('auth_pass', 'admin123') if self.config else 'admin123' + + # 恒定时间比较,防时序攻击 + if hmac.compare_digest(username, expected_user) and hmac.compare_digest(password, expected_pass): + return True + else: + self.send_auth_required() + return False + except Exception: + self.send_auth_required() + return False + + def _check_token_auth(self): + """检查令牌认证""" + token = None + # 从多个来源获取 token + headers_dict = dict(self.headers) + if 'token' in headers_dict: + token = headers_dict['token'] + elif 'X-API-Key' in headers_dict: + token = headers_dict['X-API-Key'] + elif 'Authorization' in headers_dict and headers_dict['Authorization'].startswith('Bearer '): + token = headers_dict['Authorization'][7:] + elif '?' in self.path: + parsed = urlparse(self.path) + query_params = parse_qs(parsed.query) + token = query_params.get('token', [None])[0] + + if not token: + self.send_json_response({"error": "Invalid or missing token", "code": "UNAUTHORIZED"}, 401) + return False + + # 首先检查是否是会话 token + if hasattr(self, 'auth_manager') and self.auth_manager: + session = self.auth_manager.validate_session_id(token) + if session and session.get('valid'): + return True + + # 检查是否是静态 token + expected_token = self.config.get('auth_token') if self.config else None + import hmac + if token and expected_token and hmac.compare_digest(token, expected_token): + return True + + # 验证失败 + self.send_json_response({"error": "Invalid or missing token", "code": "UNAUTHORIZED"}, 401) + return False + + def send_auth_required(self): + """发送认证要求""" + self.send_response(401) + self.send_header('WWW-Authenticate', 'Basic realm="Mirror Server"') + self.send_header('Content-Type', 'application/json') + self.send_header('Access-Control-Allow-Origin', '*') + self.end_headers() + self.wfile.write(b'{"error": "Authentication Required", "code": "UNAUTHORIZED"}') + + def do_GET(self): + """处理GET请求""" + import sys + sys.stderr.flush() + # 调试模式输出请求详情 (debug-http) + if self._is_debug_enabled('http'): + msg = f"\n=== DEBUG GET Request ===\n Path: {self.path}\n Headers: {dict(self.headers)}" + self._debug_log('http', msg, '\033[33m') + + try: + # 确保配置已加载 + if self.config is None: + self.config = MirrorServerHandler.config + if self.sync_manager is None: + self.sync_manager = MirrorServerHandler.sync_manager + if self.api_router is None and self.config is not None: + MirrorServerHandler.api_router = APIRouter(self.config) + + parsed_path = urlparse(self.path) + path = unquote(parsed_path.path).lstrip('/') + query = parsed_path.query + + # 检查认证(公开端点不需要认证) + if not self.check_auth(path): + return + + # 处理 /api/docs 和 /api/ui 路径(返回静态页面) + if path.startswith("api/docs"): + # 提供 api/docs 目录下的静态文件 + rel_path = path[9:] # 去掉 "api/docs" + if rel_path and not rel_path.startswith('/'): + rel_path = '/' + rel_path + self.serve_docs(rel_path) + return + elif path.startswith("api/ui"): + # 提供 api/ui 目录下的静态文件 + rel_path = path[7:] # 去掉 "api/ui" + if rel_path and not rel_path.startswith('/'): + rel_path = '/' + rel_path + self.serve_ui(rel_path) + return + elif path.startswith("ui/") or path == "ui": + # /ui/ 路径已废弃,返回 404 + self.send_error(404, "UI moved to /api/ui/") + return + elif path.startswith("docs/") or path == "docs": + # /docs/ 路径已废弃,返回 404 + self.send_error(404, "Docs moved to /api/docs/") + return + + # 处理 PyPI 包文件路径 - 转发到 API 路由 + # 这些路径来自 pip 下载请求,如 /pypi/packages/hash/file.tar.gz + if path.startswith("pypi/packages/") or path.startswith("pypi/web/") or path.startswith("pypi/simple/"): + # 转发到 API v2 路由 + api_path = "api/v2/" + path + self.api_router.handle_request(self, 'GET', api_path, query) + return + + # 处理 API 路径 + if path.startswith("api/"): + # 使用API路由处理 + self.api_router.handle_request(self, 'GET', path, query) + # 文件夹/文件访问(/pypi/ 也是本地文件夹) + elif path == "": + # 根路径显示文件列表 + self.serve_path("") + elif path.startswith("file/"): + # 文件下载路由 /file/path/to/file -> serve_path(path/to/file) + file_rel_path = path[5:] # 去掉 "file/" 前缀 + self.serve_path(file_rel_path) + else: + self.serve_path(path) + except Exception as e: + if self._is_debug_enabled('error'): + import traceback + tb_str = traceback.format_exc() + msg = f"\n=== DEBUG GET ERROR ===\n{tb_str}" + self._debug_log('error', msg, '\033[31m') + self.handle_error(500, f"服务器内部错误: {str(e)}") + + def do_POST(self): + """处理POST请求""" + import time + request_id = int(time.time() * 1000000) + + # 调试模式输出请求详情 (debug-http) + if self._is_debug_enabled('http'): + content_length = self.headers.get('Content-Length', 0) + msg = f"\n=== DEBUG POST Request #{request_id} ===\n Path: {self.path}\n Content-Length: {content_length}" + self._debug_log('http', msg, '\033[33m') + + try: + # 确保配置已加载 + if self.config is None: + self.config = MirrorServerHandler.config + if self.api_router is None and self.config is not None: + MirrorServerHandler.api_router = APIRouter(self.config) + + # 解析路径(去掉查询串,避免把 ?query 拼进 API 路径) + parsed_path = urlparse(self.path) + path = unquote(parsed_path.path).lstrip('/') + query = parsed_path.query + + # 调试模式输出完整路径 (debug-http) + if self._is_debug_enabled('http'): + msg = f"\n=== DEBUG POST Path Check ===\n path: '{path}'\n starts with api/: {path.startswith('api/')}" + self._debug_log('http', msg, '\033[33m') + + # 检查认证(公开端点不需要认证) + if not self.check_auth(path): + return + + if path.startswith("api/"): + self.api_router.handle_request(self, 'POST', path, query) + else: + self.send_error(405) + except Exception as e: + if self._is_debug_enabled('error'): + import traceback + tb_str = traceback.format_exc() + msg = f"\n=== DEBUG POST ERROR ===\n{tb_str}" + self._debug_log('error', msg, '\033[31m') + self.handle_error(500, f"服务器内部错误: {str(e)}") + + def do_OPTIONS(self): + """处理OPTIONS请求(CORS预检)""" + # 调试模式输出 (debug-http) + if self._is_debug_enabled('http'): + msg = f"\n=== DEBUG OPTIONS Request ===\n Path: {self.path}" + self._debug_log('http', msg, '\033[34m') + + self.send_response(200) + self.send_header('Access-Control-Allow-Origin', '*') + self.send_header('Access-Control-Allow-Methods', + 'GET, POST, PUT, DELETE, OPTIONS') + self.send_header('Access-Control-Allow-Headers', + 'Content-Type, Authorization, Token') + self.send_header('Access-Control-Max-Age', '86400') + self.end_headers() + + def do_DELETE(self): + """处理DELETE请求""" + try: + # 确保配置已加载 + if self.config is None: + self.config = MirrorServerHandler.config + if self.api_router is None and self.config is not None: + MirrorServerHandler.api_router = APIRouter(self.config) + + if not self.check_auth(): + return + parsed_path = urlparse(self.path) + path = unquote(parsed_path.path).lstrip('/') + query = parsed_path.query + if path.startswith("api/"): + self.api_router.handle_request(self, 'DELETE', path, query) + else: + self.send_error(405) + except Exception as e: + if self._is_debug_enabled('error'): + import traceback + tb_str = traceback.format_exc() + msg = f"\n=== DEBUG DELETE ERROR ===\n{tb_str}" + self._debug_log('error', msg, '\033[31m') + self.handle_error(500, f"服务器内部错误: {str(e)}") + + def do_PUT(self): + """处理PUT请求""" + try: + # 确保配置已加载 + if self.config is None: + self.config = MirrorServerHandler.config + if self.api_router is None and self.config is not None: + MirrorServerHandler.api_router = APIRouter(self.config) + + if not self.check_auth(): + return + parsed_path = urlparse(self.path) + path = unquote(parsed_path.path).lstrip('/') + query = parsed_path.query + if path.startswith("api/"): + self.api_router.handle_request(self, 'PUT', path, query) + else: + self.send_error(405) + except Exception as e: + if self._is_debug_enabled('error'): + import traceback + tb_str = traceback.format_exc() + msg = f"\n=== DEBUG PUT ERROR ===\n{tb_str}" + self._debug_log('error', msg, '\033[31m') + self.handle_error(500, f"服务器内部错误: {str(e)}") + + def do_HEAD(self): + """处理HEAD请求""" + if not self.check_auth(): + return + # 解析路径(去掉查询串),并复用 serve_path 的安全检查防路径穿越 + parsed_path = urlparse(self.path) + rel_path = unquote(parsed_path.path).lstrip('/') + if self.config is None: + self.send_error(500) + return + file_path = os.path.join(self.config['base_dir'], rel_path) + if not is_safe_path(self.config['base_dir'], file_path): + self.send_error(403, "Access denied") + return + if os.path.isfile(file_path): + self.send_file_headers(file_path) + else: + self.send_error(404) + + # ==================== 静态文件服务 ==================== + + def serve_docs(self, rel_path): + """提供 api/docs 目录下的静态文件""" + docs_dir = get_resource_path('api/docs') + + if not os.path.isdir(docs_dir): + self.send_error(404, "Docs directory not found") + return + + # 默认提供 index.html + if not rel_path or rel_path == '/': + rel_path = 'index.html' + + file_path = os.path.join(docs_dir, rel_path) + + # 防止目录遍历 + if not os.path.realpath(file_path).startswith(os.path.realpath(docs_dir)): + self.send_error(403, "Access denied") + return + + if os.path.isfile(file_path): + self.serve_file(file_path, f'api/docs/{rel_path}') + else: + # 提供 docs 目录索引 + self.serve_docs_index(docs_dir) + + def serve_docs_index(self, docs_dir): + """提供 docs 目录索引页面""" + try: + items = [] + for name in sorted(os.listdir(docs_dir)): + full_path = os.path.join(docs_dir, name) + rel_path = f'docs/{name}' + is_dir = os.path.isdir(full_path) + items.append({ + "name": name, + "path": rel_path + ("/" if is_dir else ""), + "is_dir": is_dir + }) + except OSError: + self.send_error(403) + return + + # 生成 HTML + title = "HYC下载站 - 文档" + html = f''' + + + + + {title} + + + +

{title}

+''' + for item in items: + if item['is_dir']: + html += f''' +
+ {item['name']}/ +
目录
+
''' + else: + # 根据文件类型添加描述 + desc = "" + if item['name'].endswith('.md'): + desc = "Markdown 文档" + elif item['name'].endswith('.yaml') or item['name'].endswith('.yml'): + desc = "OpenAPI 配置" + elif item['name'].endswith('.json'): + desc = "JSON 配置" + else: + desc = "文件" + html += f''' +
+ {item['name']} +
{desc}
+
''' + + html += ''' + +''' + + self.send_response(200) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.send_header("Content-Length", str(len(html.encode('utf-8')))) + self.end_headers() + self.wfile.write(html.encode('utf-8')) + + def serve_ui(self, rel_path): + """提供 api/ui 目录下的静态文件""" + ui_dir = get_resource_path('api/ui') + + if not os.path.isdir(ui_dir): + self.send_error(404, "UI directory not found") + return + + # 默认提供 index.html + if not rel_path or rel_path == '/': + rel_path = 'index.html' + + file_path = os.path.join(ui_dir, rel_path) + + # 防止目录遍历 + if not os.path.realpath(file_path).startswith(os.path.realpath(ui_dir)): + self.send_error(403, "Access denied") + return + + if os.path.isfile(file_path): + self.serve_file(file_path, f'api/ui/{rel_path}') + else: + self.send_error(404, f"File not found: {rel_path}") + + def serve_path(self, rel_path): + """处理路径请求(文件或目录)""" + if self.config is None: + self.send_error(500) + return + + # 敏感文件黑名单(纵深防御:禁止下载会话/密钥等文件) + SENSITIVE_FILES = { + 'auth_sessions.json', 'auth_secret.key', 'auth_token.txt', + 'settings.json', 'sync_state.json', 'webhooks.json', + } + base_name = os.path.basename(rel_path).lower() + if base_name in SENSITIVE_FILES: + self.send_error(403, "Access denied") + return + + file_path = os.path.join(self.config['base_dir'], rel_path) + + if not is_safe_path(self.config['base_dir'], file_path): + self.send_error(403, "Access denied") + return + + if os.path.isdir(file_path): + if self.config.get('directory_listing', True): + self.serve_directory(file_path, rel_path) + else: + self.send_error(403, "Directory listing is disabled") + elif os.path.isfile(file_path): + self.serve_file(file_path, rel_path) + else: + self.send_error(404) + + def serve_directory(self, dir_path, rel_dir): + """提供目录浏览(镜像站风格)""" + try: + # 检查是否有索引文件 + index_files = ['index.html', 'index.htm'] + for index_file in index_files: + index_path = os.path.join(dir_path, index_file) + if os.path.isfile(index_path): + self.serve_file(index_path, os.path.join(rel_dir, index_file)) + return + + # 获取目录内容 + items = [] + for name in os.listdir(dir_path): + full_path = os.path.join(dir_path, name) + rel_item_path = os.path.join(rel_dir, name).replace("\\", "/") + is_dir = os.path.isdir(full_path) + + if self.config.get('ignore_hidden', True) and name.startswith('.'): + continue + + try: + size = "-" if is_dir else format_file_size(os.path.getsize(full_path)) + mtime = datetime.fromtimestamp(os.path.getmtime(full_path)).strftime("%Y-%m-%d %H:%M") + + sha256 = "" + if self.config.get('show_hash') and not is_dir: + sha256 = get_file_hash(full_path)[:16] + "..." + items.append({ + "name": name, + "path": rel_item_path + ("/" if is_dir else ""), + "size": size, + "modified": mtime, + "is_dir": is_dir, + "sha256": sha256 + }) + except OSError: + continue + + # 排序 - Windows 文件管理器风格:文件夹在前,按名称递增排序 + sort_by = self.config.get('sort_by', 'name') + reverse = self.config.get('sort_reverse', False) # 默认为 False(递增) + if sort_by == 'name': + # 文件夹优先,然后按名称递增排序(不区分大小写) + items.sort(key=lambda x: (not x["is_dir"], x["name"].lower()), reverse=False) + elif sort_by == 'size': + # 文件夹优先,然后按大小递增排序 + items.sort(key=lambda x: (not x["is_dir"], os.path.getsize(os.path.join(dir_path, x["name"])) if not x["is_dir"] else 0), reverse=False) + elif sort_by == 'modified': + # 文件夹优先,然后按修改时间递增排序 + items.sort(key=lambda x: (not x["is_dir"], os.path.getmtime(os.path.join(dir_path, x["name"]))), reverse=False) + + except OSError: + self.send_error(403) + return + + # 构建面包屑 + breadcrumbs = [] + parts = [p for p in rel_dir.split("/") if p] + current = "" + breadcrumbs.append({"name": "HOME", "path": "/"}) + for part in parts: + current = os.path.join(current, part).replace("\\", "/") + breadcrumbs.append({"name": part, "path": "/" + current + "/"}) + + # 生成HTML + title = "HYC下载站" + html = self._generate_directory_html(title, breadcrumbs, items, rel_dir) + + self.send_response(200) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.send_header("Content-Length", str(len(html.encode('utf-8')))) + self.send_header("Cache-Control", "no-cache") + self.end_headers() + self.wfile.write(html.encode('utf-8')) + + def _generate_directory_html(self, title, breadcrumbs, items, rel_dir): + """生成目录浏览HTML""" + if rel_dir: # 如果不是根目录 + parts = [p for p in rel_dir.split('/') if p] # 过滤空部分 + if len(parts) > 1: + parent_path = '/' + '/'.join(parts[:-1]) + '/' + elif len(parts) == 1: + parent_path = '/' + else: + parent_path = '/' + else: + parent_path = '/' # 根目录没有上一级 + + # 动态计算列数 + colspan = 4 if self.config.get('show_hash') else 3 + + import html as _html + title_safe = _html.escape(str(title)) + html = f""" + + + + {title_safe} + + + + +
+

{title}

+ + + + + + + + {('' if self.config.get('show_hash') else '')} + + + """ + + # 修复:只在非根目录显示上一级目录链接 + if rel_dir: # 如果不是根目录 + html += f'\n' + + for item in items: + item_name_safe = _html.escape(str(item["name"])) + item_path_safe = _html.escape(item["path"], quote=True) + html += f'' + html += f'' + html += f'' + html += f'' + if self.config.get('show_hash'): + html += f'' + html += '\n' + + html += f""" + +
NameLast ModifiedSizeSHA256
../
{item_name_safe}{" /" if item["is_dir"] else ""}{_html.escape(str(item["modified"]))}{_html.escape(str(item["size"]))}{_html.escape(str(item["sha256"]))}
+
+

Files: {len(items)} | {_html.escape(str(self.config.get("server_name", "Mirror Server")))}

+
+
+ + """ + return html + + @staticmethod + def _safe_disposition_filename(file_path): + """清洗 Content-Disposition 文件名,防响应头注入(引号/CRLF/控制字符)""" + name = os.path.basename(file_path) + # 去掉引号与换行等危险字符 + name = re.sub(r'["\r\n\x00-\x1f]', '_', name) + return name + + def send_file_headers(self, file_path): + """发送文件头信息(用于HEAD请求)""" + if not os.path.exists(file_path) or not os.path.isfile(file_path): + self.send_error(404) + return + file_size = os.path.getsize(file_path) + mime_type, _ = mimetypes.guess_type(file_path) + if mime_type is None: + mime_type = "application/octet-stream" + + self.send_response(200) + self.send_header("Content-Type", mime_type) + self.send_header("Content-Length", str(file_size)) + self.send_header("Content-Disposition", + f'attachment; filename="{self._safe_disposition_filename(file_path)}"') + self.send_header("Accept-Ranges", "bytes") + self.send_header("Cache-Control", "public, max-age=3600") + self.send_header( + "Last-Modified", self.date_time_string(os.path.getmtime(file_path))) + self.end_headers() + + def serve_file(self, file_path, rel_path): + """提供文件下载,支持断点续传和流式传输""" + if not os.path.exists(file_path) or not os.path.isfile(file_path): + self.send_error(404) + return + + # 获取文件信息 + file_size = os.path.getsize(file_path) + mime_type, _ = mimetypes.guess_type(file_path) + if mime_type is None: + mime_type = "application/octet-stream" + + # 获取客户端IP + client_ip = self.client_address[0] if hasattr(self, 'client_address') else 'unknown' + + # 检查Range头部(支持断点续传) + range_header = self.headers.get('Range') + range_start = 0 + range_end = file_size - 1 + + if range_header and self.config.get('enable_range', True): + match = re.match(r'bytes=(\d+)-(\d*)', range_header) + if match: + range_start = int(match.group(1)) + range_end_str = match.group(2) + if range_end_str: + range_end = int(range_end_str) + + if range_start >= file_size or range_end >= file_size or range_start > range_end: + self.send_error(416, "Requested Range Not Satisfiable") + return + + # 计算传输内容 + content_length = range_end - range_start + 1 + + # 发送响应头 + if range_start == 0 and range_end == file_size - 1: + # 完整文件下载 + self.send_response(200) + else: + # 部分内容(206) + self.send_response(206) + self.send_header("Content-Range", f"bytes {range_start}-{range_end}/{file_size}") + + self.send_header("Content-Type", mime_type) + self.send_header("Content-Length", str(content_length)) + # HTML 文件直接在浏览器中显示,不强制下载 + if mime_type == 'text/html': + self.send_header("Content-Disposition", f'inline; filename="{self._safe_disposition_filename(file_path)}"') + else: + self.send_header("Content-Disposition", + f'attachment; filename="{self._safe_disposition_filename(file_path)}"') + self.send_header("Accept-Ranges", "bytes") + self.send_header("Cache-Control", "public, max-age=3600") + self.send_header("Last-Modified", self.date_time_string(os.path.getmtime(file_path))) + self.send_header("X-Download-IP", client_ip) + self.send_header("X-File-Size", str(file_size)) + self.end_headers() + + # 流式传输文件 + chunk_size = 64 * 1024 # 64KB chunks for better performance + bytes_sent = 0 + + try: + with open(file_path, 'rb') as f: + f.seek(range_start) + remaining = content_length + + while remaining > 0: + chunk = f.read(min(chunk_size, remaining)) + if not chunk: + break + + self.wfile.write(chunk) + bytes_sent += len(chunk) + remaining -= len(chunk) + + except Exception as e: + print(f"文件传输错误: {e}") + + # 只对真正的下载(非 HTML 页面)更新统计和记录 + if self.config.get('enable_stats', True) and mime_type != 'text/html': + self.record_download(rel_path, file_size, client_ip) + + def record_download(self, filepath, file_size=0, client_ip='unknown'): + """记录下载(同时更新计数和创建下载记录)""" + if not self.config.get('enable_stats', True): + if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): + self._debug_log('download', f"Stats disabled, skipping download record for: {filepath}") + return + + db = self._get_db() + if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): + self._debug_log('download', f"record_download called for: {filepath}, db: {db}") + user_agent = self.headers.get('User-Agent', 'Unknown') if hasattr(self, 'headers') else 'Unknown' + + if db: + try: + # 尝试更新 FileRecord 的下载计数(通过路径查找) + try: + record = db.get_file_by_path(filepath) + if record: + db.increment_download_count(record.file_id) + except Exception as e: + pass # 忽略更新计数错误 + + # 创建下载记录 + try: + new_record = db.add_download_record( + file_path=filepath, + file_size=file_size, + client_ip=client_ip, + user_agent=user_agent, + success=True + ) + if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): + self._debug_log('download', f"Download record created successfully: {filepath}") + except Exception as e: + if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): + self._debug_log('download', f"Error creating download record: {e}") + except Exception as e: + if hasattr(self, '_debug_log') and self._is_debug_enabled('download'): + self._debug_log('download', f"Error recording download: {e}") + + def _serve_file_chunked(self, file_path, rel_path, chunk_size=64*1024): + """流式分块传输文件(用于大文件)- 备用功能""" + if not os.path.exists(file_path) or not os.path.isfile(file_path): + self.send_error(404) + return + + file_size = os.path.getsize(file_path) + mime_type, _ = mimetypes.guess_type(file_path) + self.send_response(200) + self.send_header("Content-Type", mime_type) + self.send_header("Content-Length", str(file_size)) + self.send_header("Content-Disposition", + f'attachment; filename="{self._safe_disposition_filename(file_path)}"') + self.send_header("Accept-Ranges", "bytes") + self.send_header("Transfer-Encoding", "chunked") + self.end_headers() + + try: + with open(file_path, 'rb') as f: + while True: + chunk = f.read(chunk_size) + if not chunk: + break + self.wfile.write(chunk) + except Exception as e: + print(f"流式传输错误: {e}") + + # 更新统计(serve_file_chunked 只用于真正的下载) + if self.config.get('enable_stats', True): + self.update_download_count(rel_path) + + def send_json_response(self, data, status_code=200): + """发送JSON响应""" + json_data = json.dumps(data, ensure_ascii=False, indent=2) + self.send_response(status_code) + self.send_header("Content-Type", "application/json; charset=utf-8") + self.send_header("Content-Length", str(len(json_data.encode('utf-8')))) + self.send_header("Cache-Control", "no-cache") + self.end_headers() + self.wfile.write(json_data.encode('utf-8')) + + def date_time_string(self, timestamp=None): + """重写日期时间字符串格式化""" + if timestamp is None: + timestamp = time.time() + return datetime.fromtimestamp(timestamp).strftime('%a, %d %b %Y %H:%M:%S GMT') + + def handle_error(self, code, message=None): + """自定义错误处理""" + error_messages = { + 400: "错误的请求", + 401: "未经授权", + 403: "禁止访问", + 404: "文件未找到", + 405: "方法不允许", + 413: "文件太大", + 416: "请求范围不符合要求", + 500: "内部服务器错误" + } + + if message is None: + message = error_messages.get(code, "未知错误") + + import html as _html + message_safe = _html.escape(str(message)) + error_page = f""" + + + + + + {code} {message_safe} + + + +
+

{code}

+

{message_safe}

+

请求的页面遇到问题,请稍后重试.

+ 返回首页 +
+ +""" + + self.send_response(code) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.send_header("Content-Length", + str(len(error_page.encode('utf-8')))) + self.end_headers() + self.wfile.write(error_page.encode('utf-8')) + + def send_error(self, code, message=None): + """发送错误响应""" + self.handle_error(code, message) + + # 统计相关方法 + def _get_db(self): + """获取数据库实例""" + if hasattr(self, 'config') and self.config: + return self.config.get('_db_instance') + return None + + def load_stats(self): + """加载下载统计信息(优先使用数据库,回退到JSON)""" + db = self._get_db() + if db: + try: + # 使用专门的方法获取下载统计,避免会话问题 + return db.get_download_stats(limit=10000) + except Exception as e: + print(f"Error loading stats from database: {e}") + + # 回退到 JSON 文件 + stats_file = self.config.get('stats_file', 'stats.json') + try: + if os.path.exists(stats_file): + with open(stats_file, 'r', encoding='utf-8') as f: + return json.load(f) + except Exception as e: + print(f"Error loading stats: {e}") + return {} + + def save_stats(self, stats): + """保存下载统计信息(优先使用数据库,回退到JSON)""" + db = self._get_db() + if db: + # 数据库模式下,stats 由数据库直接管理,不需要手动保存 + return + + # 回退到 JSON 文件 + stats_file = self.config.get('stats_file', 'stats.json') + try: + with open(stats_file, 'w', encoding='utf-8') as f: + json.dump(stats, f, ensure_ascii=False, indent=2) + except Exception as e: + print(f"Error saving stats: {e}") + + def get_download_count(self, filepath): + """获取特定文件的下载次数""" + db = self._get_db() + if db: + try: + file_record = db.get_file_by_path(filepath) + if file_record: + return file_record.download_count if hasattr(file_record, 'download_count') else 0 + except Exception as e: + print(f"Error getting download count from database: {e}") + + stats = self.load_stats() + return stats.get(filepath, 0) + + def get_total_downloads(self): + """获取总下载次数""" + db = self._get_db() + if db: + try: + stats = db.get_stats() + return stats.get('total_downloads', 0) + except Exception as e: + print(f"Error getting total downloads from database: {e}") + + stats = self.load_stats() + return sum(stats.values()) + + def update_download_count(self, filepath): + """更新文件的下载计数(优先使用数据库,通过路径查找)""" + if not self.config.get('enable_stats', True): + return + + db = self._get_db() + if db: + try: + # 通过路径查找记录,获取 file_id + record = db.get_file_by_path(filepath) + if record: + db.increment_download_count(record.file_id) + return + 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) + + # ==================== 下载历史记录 ==================== + + def load_download_history(self, limit=100): + """加载下载历史记录(优先使用数据库,回退到JSON)""" + db = self._get_db() + if db: + try: + records = db.get_download_records(limit=limit) + history = [] + for r in records: + history.append({ + 'timestamp': r.created_at.isoformat() if hasattr(r.created_at, 'isoformat') else str(r.created_at), + 'filepath': r.file_path if hasattr(r, 'file_path') else getattr(r, 'filepath', str(r)), + 'file_size': r.file_size if hasattr(r, 'file_size') else 0, + 'client_ip': r.client_ip if hasattr(r, 'client_ip') else 'unknown', + 'user_agent': r.user_agent if hasattr(r, 'user_agent') else 'Unknown', + 'method': 'GET' + }) + return history + except Exception as e: + print(f"Error loading download history from database: {e}") + + # 回退到 JSON 文件 + history_file = self.config.get('download_history_file', 'download_history.json') + try: + if os.path.exists(history_file): + with open(history_file, 'r', encoding='utf-8') as f: + data = json.load(f) + return data[-limit:] + except Exception as e: + print(f"Error loading download history: {e}") + return [] + + def save_download_history(self, history): + """保存下载历史记录(数据库模式下不需要)""" + db = self._get_db() + if db: + # 数据库模式下,history 由数据库直接管理 + return + + # 回退到 JSON 文件 + history_file = self.config.get('download_history_file', 'download_history.json') + max_history = self.config.get('max_history_count', 1000) + + try: + # 保留最近的记录 + history = history[-max_history:] + with open(history_file, 'w', encoding='utf-8') as f: + json.dump(history, f, ensure_ascii=False, indent=2) + except Exception as e: + print(f"Error saving download history: {e}") + + def _log_download(self, filepath, file_size=0): + """记录下载历史(优先使用数据库)- 备用功能""" + if not self.config.get('enable_stats', True): + return + + client_ip = self.client_address[0] if hasattr(self, 'client_address') else 'unknown' + user_agent = self.headers.get('User-Agent', 'Unknown') + + db = self._get_db() + if db: + try: + db.add_download_record( + file_path=filepath, + file_size=file_size, + client_ip=client_ip, + user_agent=user_agent, + success=True, + duration=0 + ) + return + except Exception as e: + print(f"Error logging download to database: {e}") + + # 回退到 JSON 文件 + history = self.load_download_history(1000) + + entry = { + 'timestamp': datetime.now().isoformat(), + 'filepath': filepath, + 'file_size': file_size, + 'client_ip': client_ip, + 'user_agent': user_agent, + 'method': self.command if hasattr(self, 'command') else 'GET' + } + + history.append(entry) + self.save_download_history(history) diff --git a/main.py b/main.py index 4d17f1c..d634c23 100644 --- a/main.py +++ b/main.py @@ -36,10 +36,9 @@ from core.optimization import ( def signal_handler(signum, _frame): - """处理退出信号""" - print(f"\n收到信号 {signum},正在关闭服务器...") - import os - os._exit(0) + """处理退出信号(优雅退出:触发 KeyboardInterrupt 走正常关闭流程)""" + print(f"\n收到信号 {signum},正在优雅关闭服务器...") + raise KeyboardInterrupt def parse_arguments(): @@ -587,6 +586,21 @@ def main(): # 设置服务器启动时间(用于计算运行时间) config['start_time'] = time.time() + # 默认凭据警告 + try: + if config.get('auth_type') == 'basic' and config.get('auth_pass') in (None, '', 'admin123'): + print("\n" + "!" * 60) + print("! 安全警告: 正在使用默认/空认证密码 (admin/admin123)") + print("! 请立即修改 settings.json 中的 auth_pass 或使用 --auth-pass 指定") + print("!" * 60) + if config.get('auth_type') == 'token' and config.get('auth_token') in (None, '', 'admin_token_123456'): + print("\n" + "!" * 60) + print("! 安全警告: 正在使用默认/空认证 token") + print("! 请立即修改 settings.json 中的 auth_token 或使用 --auth-token 指定") + print("!" * 60) + except Exception: + pass + # 创建并启动服务器 try: server = MirrorServer(config) @@ -595,6 +609,10 @@ def main(): else: print("服务器启动失败") sys.exit(1) + except KeyboardInterrupt: + # 信号处理已触发优雅关闭(server.stop + atexit cleanup 已执行) + print("\n服务器已正常退出") + sys.exit(0) except Exception as e: print(f"错误: {e}") import traceback diff --git a/mirrors/pypi.py b/mirrors/pypi.py index a2d69b3..1d73db3 100644 --- a/mirrors/pypi.py +++ b/mirrors/pypi.py @@ -1,788 +1,808 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - -""" -PyPI镜像代理处理器 -支持Python包索引 -""" - -import os -import json -import re -import time -import urllib.request -import urllib.parse -import urllib.error -from typing import Dict, List, Optional -from datetime import datetime - - -class PyPIMirror: - """PyPI镜像代理""" - - def __init__(self, config: dict): - self.config = config - - # 配置 - storage_dir 基于 base_dir - self.upstream_url = config.get('upstream_url', 'https://pypi.org') - self.base_dir = config.get('base_dir', './downloads') - storage_subdir = config.get('storage_dir', 'pypi') - self.storage_dir = os.path.join(self.base_dir, storage_subdir) - self.simple_dir = os.path.join(self.storage_dir, 'simple') - self.web_dir = os.path.join(self.storage_dir, 'web') - - # 确保存储目录存在 - os.makedirs(self.simple_dir, exist_ok=True) - os.makedirs(self.web_dir, exist_ok=True) - - def handle_request(self, handler, path: str) -> bool: - """ - 处理PyPI请求 - 路径格式: /simple/requests/ 或 /packages/xxx.tar.gz 或 /pypi/web/package 或 /pypi/packages/hash/file - """ - try: - import sys - # 如果路径以 pypi/ 开头,也需要去掉 - path = path.lstrip('/') - if path.startswith('pypi/'): - path = path[5:] - parts = path.strip('/').split('/') - - # 过滤空字符串 - parts = [p for p in parts if p] - - if not parts: - return self._handle_index(handler) - - if parts[0] == 'simple': - # Simple API - import sys - if len(parts) == 1: - # /simple/ - 返回根索引 - return self._handle_index(handler) - elif len(parts) == 2: - # /simple/package/ - return self._handle_simple_index(handler, parts[1]) - elif len(parts) >= 3: - # /simple/package/version/ 或 /simple/package/version#egg=... - return self._handle_package_file(handler, parts[1], '/'.join(parts[2:])) - else: - handler.send_error(400, "Invalid simple API path") - return False - - elif parts[0] == 'web': - # /web/package/ 或 /web/package/json - # pip sends /pypi/web//json - package = parts[1] if len(parts) >= 2 else '' - return self._handle_web_api(handler, package) - - elif parts[0] == 'packages': - # 包下载 - filename = '/'.join(parts[1:]) - return self._handle_package_download(handler, filename) - - elif parts[0] == 'legacy': - # 旧版PyPI兼容 - return self._handle_legacy(handler, '/'.join(parts[1:])) - - else: - handler.send_error(404, "Unknown API") - return False - - except Exception as e: - handler.send_error(500, str(e)) - return False - - def _handle_index(self, handler) -> bool: - """处理索引请求 - 返回所有可用包的列表""" - # 从上游获取包列表 - url = self.upstream_url.rstrip('/') - if url.endswith('/simple'): - url = url # 保持 /simple - else: - url = url + '/simple' - - try: - req = urllib.request.Request(url) - req.add_header('Accept', 'text/html') - req.add_header('User-Agent', 'PyPI-Mirror/1.0') - - with urllib.request.urlopen(req, timeout=30) as response: - data = response.read().decode('utf-8') - - # 转换相对链接 - # 清华源返回的可能是完整的HTML,需要转换链接 - data = self._convert_simple_index_html(data) - data_bytes = data.encode('utf-8') - - handler.send_response(200) - handler.send_header('Content-Type', 'text/html; charset=utf-8') - handler.send_header('Content-Length', str(len(data_bytes))) - handler.end_headers() - handler.wfile.write(data_bytes) - return True - - except Exception as e: - handler.send_error(502, f"Failed to fetch package index: {str(e)}") - return False - - def _convert_simple_index_html(self, html: str) -> str: - """转换根索引页面的HTML""" - import re - # 替换上游链接 - def convert_link(match): - href = match.group(1) - text = match.group(2) - if href.startswith('/simple/'): - return match.group(0) # 已经是相对路径 - elif href.startswith('https://pypi.tuna.tsinghua.edu.cn/simple/'): - simple_part = href.split("/simple/")[-1] - return f'{text}' - elif href.startswith('https://'): - # 其他上游链接,提取包名 - pkg_name = href.rstrip('/').split('/')[-1] - return f'{text}' - return match.group(0) - - # 匹配 text - return re.sub(r']+href="([^"]+)"[^>]*>([^<]*)', convert_link, html) - - def _handle_simple_index(self, handler, package: str) -> bool: - """处理Simple API索引请求""" - import json - import time - - package = package.lower() - - # 调试 - 确保函数被调用 - debug_file = '/tmp/pypi_debug.log' - with open(debug_file, 'a') as f: - f.write(f"[HANDLE_SIMPLE] START package={package}\n") - - # 检查客户端Accept header - accept = handler.headers.get('Accept', '') - wants_json = 'application/vnd.pypi.simple.v1+json' in accept - - # 调试 - debug_file = '/tmp/pypi_debug.log' - with open(debug_file, 'a') as f: - f.write(f"[SIMPLE_INDEX] package={package}, wants_json={wants_json}, accept={accept[:50]}\n") - - # 根据请求格式选择正确的缓存key,统一使用 simple/ 前缀 - cache_key = f"simple/{package}" - - cached = self._get_cache(cache_key) - - # 调试缓存 - debug_file = '/tmp/pypi_debug.log' - with open(debug_file, 'a') as f: - f.write(f"[CACHE_CHECK] cache_key={cache_key}, cached={'YES' if cached else 'NO'}\n") - - if cached: - # 返回缓存,使用正确的Content-Type - if wants_json: - handler.send_response(200) - handler.send_header('Content-Type', 'application/vnd.pypi.simple.v1+json; charset=utf-8') - else: - handler.send_response(200) - handler.send_header('Content-Type', 'text/html; charset=utf-8') - handler.send_header('Content-Length', str(len(cached))) - handler.end_headers() - handler.wfile.write(cached) - return True - - # 从上游获取HTML - # upstream_url 已经是完整路径(如 https://pypi.tuna.tsinghua.edu.cn/simple) - # 所以只需要添加 /package/ - url = f"{self.upstream_url}/{package}/" - - try: - req = urllib.request.Request(url) - req.add_header('Accept', 'text/html') - req.add_header('User-Agent', 'PyPI-Mirror/1.0') - - with urllib.request.urlopen(req, timeout=30) as response: - data = response.read().decode('utf-8') - - # 根据客户端请求返回不同格式,统一使用 simple/ 路径 - if wants_json: - # 转换为JSON格式 - json_data = self._convert_to_json(package, data) - proxy_data = json.dumps(json_data) - content_type = 'application/vnd.pypi.simple.v1+json; charset=utf-8' - cache_key = f"simple/{package}" - else: - # 转换为HTML格式 - debug_file = '/tmp/pypi_debug.log' - with open(debug_file, 'a') as f: - f.write(f"[CONVERT_CALL] Before conversion\n") - try: - proxy_data = self._convert_simple_html(package, data) - with open(debug_file, 'a') as f: - f.write(f"[CONVERT_CALL] After conversion\n") - except Exception as e: - with open(debug_file, 'a') as f: - f.write(f"[CONVERT_ERROR] {e}\n") - proxy_data = data # fallback to raw data - content_type = 'application/vnd.pypi.simple.v1+html; charset=utf-8' - - if True: - self._set_cache(cache_key, proxy_data.encode('utf-8')) - - proxy_bytes = proxy_data.encode('utf-8') - handler.send_response(200) - handler.send_header('Content-Type', content_type) - handler.send_header('Content-Length', str(len(proxy_bytes))) - handler.end_headers() - handler.wfile.write(proxy_bytes) - return True - - except urllib.error.HTTPError as e: - if e.code == 404: - handler.send_error(404, f"Package not found: {package}") - else: - handler.send_error(502, f"Failed to fetch from upstream: {str(e)}") - return False - - def _convert_to_json(self, package: str, html: str) -> dict: - """将HTML转换为JSON格式""" - import re - - # 解析HTML中的链接 - links = [] - - # 匹配 text - pattern = r']+href="([^"]+)"[^>]*>([^<]*)' - matches = re.findall(pattern, html) - - for href, text in matches: - # 提取文件名(从链接文本) - filename = text.strip() if text.strip() else '' - - # 如果没有链接文本,从URL中提取 - if not filename: - if '#' in href: - filename = href.split('#')[0].split('/')[-1] - else: - filename = href.split('/')[-1] - - # JSON格式中URL不应该包含fragment - # pip从filename字段提取版本号(如 Flask-1.0.0.tar.gz -> 1.0.0) - - # 解析URL并转换为代理路径 - url = href - if href.startswith('../'): - # 相对路径 - 需要转换为代理路径 - # 格式: ../../packages/hash1/hash2/fullhash/filename#sha256=... - # 或: ../../packages/hash1/hash2/filename - # parts = ['..', '..', 'packages', 'hash1', 'hash2', 'fullhash', 'filename', ...] - parts = href.split('/') - try: - pkg_idx = parts.index('packages') - # 提取从 packages 后面到文件名之前的所有部分作为 hash 路径 - # 文件名是最后一个非空部分(可能包含 #fragment) - # 找到文件名的位置(最后一个部分) - filename_idx = len(parts) - 1 - while filename_idx > pkg_idx and not parts[filename_idx]: - filename_idx -= 1 - # hash_path 是 packages 后面到文件名之前的所有部分 - if filename_idx > pkg_idx + 1: - hash_path = '/'.join(parts[pkg_idx+1:filename_idx]) - else: - hash_path = parts[pkg_idx+1] if pkg_idx + 1 < len(parts) else '' - # 文件名: 检查 pkg_idx+3 是否存在且不是哈希 - if pkg_idx + 3 < len(parts): - fname_full = parts[pkg_idx+3] - # 如果 fname_full 看起来像哈希(包含 sha256= 或长度>=32的十六进制),则使用原始 filename - # 清华源格式: .../hash/filename#sha256=... - # 其中 hash 是 28-30 位十六进制 - is_hash_like = ('sha256=' in fname_full or 'sha512=' in fname_full or - (len(fname_full) >= 28 and all(c in '0123456789abcdef' for c in fname_full[:28].lower()))) - if is_hash_like: - # 这是哈希,不是文件名 - fname = filename if filename else '' - else: - # 这是文件名 - fname = fname_full.split('#')[0] - if not filename: - filename = fname - else: - fname = filename if filename else '' - # URL不包含fragment - url = f"/pypi/packages/{hash_path}/{fname}" - except ValueError: - pass - elif href.startswith('/pypi/'): - # 绝对路径(如 /pypi/packages/hash/filename#egg=package-version) - # 去掉fragment - href_clean = href.split('#')[0] - parts = href_clean.split('/') - # parts = ['', 'pypi', 'packages', 'hash', 'filename'] - if len(parts) >= 5: - hash_path = parts[3] - fname = parts[4] - if not filename: - filename = fname - url = f"/pypi/packages/{hash_path}/{fname}" - elif href.startswith('http'): - # 绝对URL - 转换为代理路径 - if 'files.pythonhosted.org' in href or 'files.pypi.org' in href: - # 格式: https://files.pythonhosted.org/packages/hash1/hash2/完整哈希/filename - # 例如: https://files.pythonhosted.org/packages/ec/f9/7f9263c5695f4bd0023734af91bedb2ff8209e8de6ead162f35d8dc762fd/flask-3.1.2-py3-none-any.whl - parts = href.split('/packages/') - if len(parts) >= 2: - path_after_packages = parts[1] - # 完整路径: hash1/hash2/完整哈希/filename - url = f"/pypi/packages/{path_after_packages}" - fname = path_after_packages.split('/')[-1] - if not filename: - filename = fname - elif 'pypi.tuna.tsinghua.edu.cn' in href or 'mirrors.tuna.tsinghua.edu.cn' in href: - parts = href.rsplit('/', 1) - if len(parts) == 2: - path_part = parts[0] - fname = parts[1] - hash_path = path_part.split('/')[-1] - if not filename: - filename = fname - url = f"/pypi/packages/{hash_path}/{fname}" - - link_entry = { - "filename": filename, - "url": url - } - - links.append(link_entry) - - return { - "meta": { - "api-version": "1.0", - "repository-version": "1.0" - }, - "name": package, - "files": links - } - - def _handle_package_file(self, handler, package: str, filename: str) -> bool: - """处理包文件请求""" - import urllib.parse - - # 解析文件名 - # 新格式: /pypi/packages/hash/filename#pip=package-version - # 或旧格式: /pypi/packages/filename?url=... - parsed = urllib.parse.urlparse(f"/{filename}") - actual_filename = parsed.path.lstrip('/') - query_params = urllib.parse.parse_qs(parsed.query) - fragment = urllib.parse.parse_qs(parsed.fragment) if parsed.fragment else {} - - # 从fragment中提取包信息(用于缓存键) - actual_package = package - if 'pip' in fragment: - # 格式: #pip=flask-2.0.0 - pip_info = fragment['pip'][0] - if '-' in pip_info: - # 提取版本号 - parts = pip_info.split('-', 1) - if len(parts) == 2: - actual_package = parts[0] - - cache_key = f"packages/{actual_filename}" - - cached = self._get_cache(cache_key) - if cached: - handler.send_response(200) - handler.send_header('Content-Type', 'application/octet-stream') - handler.send_header('Content-Length', str(len(cached))) - handler.end_headers() - handler.wfile.write(cached) - return True - - # 尝试从上游获取 - # 格式: /pypi/packages/hash/filename -> 构造上游URL - possible_urls = [] - - # 获取基础URL(去掉 /simple 后缀) - base_url = self.upstream_url.rstrip('/') - if base_url.endswith('/simple'): - base_url = base_url[:-7] - - # 新格式: hash/filename -> 尝试清华源 - # 使用 actual_filename(已解析的纯文件路径) - possible_urls.append(f"{base_url}/packages/{actual_filename}") - - # 尝试官方源 - possible_urls.append(f"https://files.pythonhosted.org/packages/{actual_filename}") - - # 如果有查询参数中的URL,也尝试 - if 'url' in query_params: - possible_urls.insert(0, urllib.parse.unquote(query_params['url'][0])) - - data = None - last_error = None - - for url in possible_urls: - try: - import sys - req = urllib.request.Request(url) - req.add_header('User-Agent', 'PyPI-Mirror/1.0') - - with urllib.request.urlopen(req, timeout=60) as response: - data = response.read() - break # 成功获取,退出循环 - except Exception as e: - import sys - last_error = e - continue - - if data is None: - handler.send_error(502, f"Failed to fetch package: {last_error}") - return False - - if True: - self._set_cache(cache_key, data) - - handler.send_response(200) - handler.send_header('Content-Type', 'application/octet-stream') - handler.send_header('Content-Length', str(len(data))) - handler.end_headers() - handler.wfile.write(data) - return True - - def _handle_web_api(self, handler, package: str) -> bool: - """处理Web API请求""" - import sys - package = package.lower() - cache_key = f"web/{package}" - - cached = self._get_cache(cache_key) - if cached: - handler.send_response(200) - handler.send_header('Content-Type', 'application/vnd.pypi.simple.v1+json; charset=utf-8') - handler.send_header('Content-Length', str(len(cached))) - handler.end_headers() - handler.wfile.write(cached) - return True - - # 从上游获取 - 需要去掉 /simple 后缀 - base_url = self.upstream_url.rstrip('/') - if base_url.endswith('/simple'): - base_url = base_url[:-7] - url = f"{base_url}/pypi/{package}/json" - - try: - req = urllib.request.Request(url) - req.add_header('Accept', 'application/json') - req.add_header('User-Agent', 'PyPI-Mirror/1.0') - - with urllib.request.urlopen(req, timeout=30) as response: - data = response.read().decode('utf-8') - data_json = json.loads(data) - - # 转换URL - data_json = self._convert_package_json(package, data_json) - - proxy_data = json.dumps(data_json) - - if True: - self._set_cache(cache_key, proxy_data.encode('utf-8')) - - proxy_bytes = proxy_data.encode('utf-8') - handler.send_response(200) - handler.send_header('Content-Type', 'application/vnd.pypi.simple.v1+json; charset=utf-8') - handler.send_header('Content-Length', str(len(proxy_bytes))) - handler.end_headers() - handler.wfile.write(proxy_bytes) - return True - - except urllib.error.HTTPError as e: - handler.send_error(502, f"Failed to fetch package info: {str(e)}") - return False - - def _handle_package_download(self, handler, filename: str) -> bool: - """处理包下载请求""" - import sys - # 使用 packages/ 前缀,保持标准 PyPI 目录结构 - cache_key = f"packages/{filename}" - - cached = self._get_cache(cache_key) - if cached: - handler.send_response(200) - handler.send_header('Content-Type', 'application/octet-stream') - handler.send_header('Content-Length', str(len(cached))) - handler.end_headers() - handler.wfile.write(cached) - return True - - # 从上游获取 - 注意去掉 /simple 后缀 - base_url = self.upstream_url.rstrip('/') - if base_url.endswith('/simple'): - base_url = base_url[:-7] # 去掉 /simple - url = f"{base_url}/packages/{filename}" - - try: - req = urllib.request.Request(url) - req.add_header('User-Agent', 'PyPI-Mirror/1.0') - - with urllib.request.urlopen(req, timeout=120) as response: - data = response.read() - - if True: - self._set_cache(cache_key, data) - - handler.send_response(200) - handler.send_header('Content-Type', 'application/octet-stream') - handler.send_header('Content-Length', str(len(data))) - handler.end_headers() - handler.wfile.write(data) - return True - - except Exception as e: - handler.send_error(502, f"Failed to download: {str(e)}") - return False - - def _handle_legacy(self, handler, path: str) -> bool: - """处理旧版PyPI兼容""" - handler.send_error(410, "Legacy PyPI API is deprecated") - return False - - def _convert_simple_html(self, package: str, html: str) -> str: - """转换Simple API HTML,替换URL为代理地址""" - import urllib.parse - import sys - import os - # 写入调试文件 - debug_file = '/tmp/pypi_debug.log' - with open(debug_file, 'a') as f: - f.write(f"[PyPI] Converting HTML for package: {package}\n") - # 打印前几个链接用于调试 - import re - test_matches = re.findall(r'href="([^"]+)"', html)[:3] - with open(debug_file, 'a') as f: - f.write(f"[PyPI] Sample links: {test_matches}\n") - - def convert_absolute_url(match): - """转换绝对URL为代理链接""" - original_url = match.group(1) if match.lastindex else match.group(0) - # 提取文件名和完整的hash路径 - # 格式: https://pypi.tuna.tsinghua.edu.cn/packages/hash1/hash2/fullhash/filename - # 我们将其转换为: /pypi/packages/hash1/hash2/fullhash/filename#pip= - parts = original_url.rsplit('/', 1) - if len(parts) == 2: - path_part = parts[0] - filename = parts[1] - # 提取从 packages/ 后面的完整路径(包含完整hash) - try: - pkg_idx = path_part.index('/packages/') - hash_path = path_part[pkg_idx + 10:] # 去掉 /packages/ - except ValueError: - hash_path = filename - return f'href="/pypi/packages/{hash_path}/{filename}#pip={package}-{filename.split("-")[1] if "-" in filename else ""}"' - return match.group(0) - - # 替换绝对URL - pypi.tuna.tsinghua.edu.cn (清华源) - html = re.sub( - r'(https://pypi\.tuna\.tsinghua\.edu\.cn/packages/[^"\']+)', - convert_absolute_url, - html - ) - # 替换绝对URL - mirrors.tuna.tsinghua.edu.cn - html = re.sub( - r'(https://mirrors\.tuna\.tsinghua\.edu\.cn/pypi/packages/[^"\']+)', - convert_absolute_url, - html - ) - # 替换绝对URL - files.pypi.org - html = re.sub( - r'(https://files\.pypi\.org/packages/[^"\']+)', - convert_absolute_url, - html - ) - # 替换绝对URL - files.pythonhosted.org - html = re.sub( - r'(https://files\.pythonhosted\.org/packages/[^"\']+)', - convert_absolute_url, - html - ) - - # 替换相对路径链接 - ../../packages/hash1/hash2/fullhash/filename -> /pypi/packages/hash1/hash2/fullhash/filename#pip=... - def convert_relative_match(match): - """转换相对路径链接""" - href = match.group(1) - debug_file = '/tmp/pypi_debug.log' - with open(debug_file, 'a') as f: - f.write(f"[CONVERT] Input href: {href[:80]}...\n") - # 提取文件名 - filename = href.split('/')[-1].split('#')[0] - # 提取完整的hash路径(从 packages/ 后面的所有部分除了文件名) - parts = href.split('/') - with open(debug_file, 'a') as f: - f.write(f"[CONVERT] Parts: {parts}\n") - try: - pkg_idx = parts.index('packages') - # packages 后面到倒数第二个是 hash 路径,最后一个是文件名 - hash_parts = parts[pkg_idx+1:-1] # 除了最后一个(文件名) - with open(debug_file, 'a') as f: - f.write(f"[CONVERT] hash_parts: {hash_parts}\n") - hash_path = '/'.join(hash_parts) if hash_parts else filename - with open(debug_file, 'a') as f: - f.write(f"[CONVERT] hash_path: {hash_path}\n") - except ValueError: - hash_path = filename - with open(debug_file, 'a') as f: - f.write(f"[CONVERT] ValueError, hash_path: {hash_path}\n") - - # 提取版本号 - 从文件名中提取,如 Flask-0.1.tar.gz -> 0.1 - base_name = filename - # 去掉扩展名 - for ext in ['.tar.gz', '.whl', '.tar.bz2', '.tar.xz']: - if base_name.endswith(ext): - base_name = base_name[:-len(ext)] - break - - # 尝试多种大小写组合来去掉包名前缀 - version = base_name - for pkg_name in [package, package.lower(), package.upper(), package.capitalize()]: - if base_name.lower().startswith(pkg_name.lower() + '-'): - version = base_name[len(pkg_name)+1:] - break - - return f'href="/pypi/packages/{hash_path}/{filename}#egg={package}-{version}"' - - # 匹配相对路径的链接 - html = re.sub( - r'href="(\.\./\.\./packages/[^"]+)"', - convert_relative_match, - html - ) - html = re.sub( - r"href='(\.\./\.\./packages/[^']+)'", - convert_relative_match, - html - ) - - return html - - def _convert_package_json(self, package: str, data: dict) -> dict: - """转换Package JSON,替换URL为代理地址""" - # 转换URL函数 - def convert_url(url): - # 处理 files.pythonhosted.org 和 files.pypi.org - # 格式: https://files.pythonhosted.org/packages///<完整哈希>/ - # 例如: https://files.pythonhosted.org/packages/ec/f9/7f9263c5695f4bd0023734af91bedb2ff8209e8de6ead162f35d8dc762fd/flask-3.1.2-py3-none-any.whl - if 'files.pythonhosted.org' in url or 'files.pypi.org' in url: - path_parts = url.split('/packages/') - if len(path_parts) >= 2: - # 直接使用 /packages/ 后的完整路径 - path_after_packages = path_parts[1] - return f'/pypi/packages/{path_after_packages}' - return url - - # 转换urls - if 'urls' in data: - for item in data['urls']: - if 'url' in item: - item['url'] = convert_url(item['url']) - - return data - - def _fetch(self, url: str) -> Optional[bytes]: - """从URL获取数据""" - try: - req = urllib.request.Request(url) - req.add_header('User-Agent', 'PyPI-Mirror/1.0') - - with urllib.request.urlopen(req, timeout=30) as response: - return response.read() - - except Exception: - return None - - def _get_cache(self, cache_key: str) -> Optional[bytes]: - """获取缓存""" - if not True: - return None - - cache_path = self._get_cache_path(cache_key) - meta_path = cache_path + '.meta' - - if not os.path.exists(cache_path): - return None - - if os.path.exists(meta_path): - try: - with open(meta_path, 'r') as f: - meta = json.load(f) - if time.time() > meta.get('expires', 0): - return None - except Exception: - pass - - try: - with open(cache_path, 'rb') as f: - return f.read() - except Exception: - return None - - def _set_cache(self, cache_key: str, data: bytes): - """设置缓存""" - cache_path = self._get_cache_path(cache_key) - meta_path = cache_path + '.meta' - - os.makedirs(os.path.dirname(cache_path), exist_ok=True) - - try: - with open(cache_path, 'wb') as f: - f.write(data) - - meta = { - 'cached_at': time.time(), - 'expires': time.time() + 86400, - 'size': len(data) - } - - with open(meta_path, 'w') as f: - json.dump(meta, f) - - except Exception as e: - print(f"PyPI缓存写入失败: {e}") - - def _get_cache_path(self, cache_key: str) -> str: - """获取缓存路径""" - # cache_key 格式: packages/fe/df/88ccbee.../filename - # 存储路径: downloads/pypi-cn/packages/fe/df/88ccbee.../filename - # 确保使用正斜杠 - safe_key = cache_key.replace('\\', '/') - return os.path.join(self.storage_dir, safe_key) - - def get_cache_stats(self) -> dict: - """获取缓存统计""" - if not os.path.exists(self.storage_dir): - return {'files': 0, 'size': 0} - - total_size = 0 - file_count = 0 - - for root, dirs, files in os.walk(self.storage_dir): - for f in files: - if not f.endswith('.meta'): - file_count += 1 - total_size += os.path.getsize(os.path.join(root, f)) - - return { - 'files': file_count, - 'size': total_size, - 'size_formatted': self._format_size(total_size) - } - - def _format_size(self, size_bytes: int) -> str: - """格式化文件大小""" - if size_bytes == 0: - return "0 B" - - units = ["B", "KB", "MB", "GB"] - i = 0 - while size_bytes >= 1024 and i < len(units) - 1: - size_bytes /= 1024.0 - i += 1 - - return f"{size_bytes:.2f} {units[i]}" +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +PyPI镜像代理处理器 +支持Python包索引 +""" + +import os +import json +import re +import time +import urllib.request +import urllib.parse +import urllib.error +from typing import Dict, List, Optional +from datetime import datetime + + +class PyPIMirror: + """PyPI镜像代理""" + + def __init__(self, config: dict): + self.config = config + + # 配置 - storage_dir 基于 base_dir + self.upstream_url = config.get('upstream_url', 'https://pypi.org') + self.base_dir = config.get('base_dir', './downloads') + storage_subdir = config.get('storage_dir', 'pypi') + self.storage_dir = os.path.join(self.base_dir, storage_subdir) + self.simple_dir = os.path.join(self.storage_dir, 'simple') + self.web_dir = os.path.join(self.storage_dir, 'web') + + # 确保存储目录存在 + os.makedirs(self.simple_dir, exist_ok=True) + os.makedirs(self.web_dir, exist_ok=True) + + def handle_request(self, handler, path: str) -> bool: + """ + 处理PyPI请求 + 路径格式: /simple/requests/ 或 /packages/xxx.tar.gz 或 /pypi/web/package 或 /pypi/packages/hash/file + """ + try: + import sys + # 如果路径以 pypi/ 开头,也需要去掉 + path = path.lstrip('/') + if path.startswith('pypi/'): + path = path[5:] + parts = path.strip('/').split('/') + + # 过滤空字符串 + parts = [p for p in parts if p] + + if not parts: + return self._handle_index(handler) + + if parts[0] == 'simple': + # Simple API + import sys + if len(parts) == 1: + # /simple/ - 返回根索引 + return self._handle_index(handler) + elif len(parts) == 2: + # /simple/package/ + return self._handle_simple_index(handler, parts[1]) + elif len(parts) >= 3: + # /simple/package/version/ 或 /simple/package/version#egg=... + return self._handle_package_file(handler, parts[1], '/'.join(parts[2:])) + else: + handler.send_error(400, "Invalid simple API path") + return False + + elif parts[0] == 'web': + # /web/package/ 或 /web/package/json + # pip sends /pypi/web//json + package = parts[1] if len(parts) >= 2 else '' + return self._handle_web_api(handler, package) + + elif parts[0] == 'packages': + # 包下载 + filename = '/'.join(parts[1:]) + return self._handle_package_download(handler, filename) + + elif parts[0] == 'legacy': + # 旧版PyPI兼容 + return self._handle_legacy(handler, '/'.join(parts[1:])) + + else: + handler.send_error(404, "Unknown API") + return False + + except Exception as e: + handler.send_error(500, str(e)) + return False + + def _handle_index(self, handler) -> bool: + """处理索引请求 - 返回所有可用包的列表""" + # 从上游获取包列表 + url = self.upstream_url.rstrip('/') + if url.endswith('/simple'): + url = url # 保持 /simple + else: + url = url + '/simple' + + try: + req = urllib.request.Request(url) + req.add_header('Accept', 'text/html') + req.add_header('User-Agent', 'PyPI-Mirror/1.0') + + with urllib.request.urlopen(req, timeout=30) as response: + data = response.read().decode('utf-8') + + # 转换相对链接 + # 清华源返回的可能是完整的HTML,需要转换链接 + data = self._convert_simple_index_html(data) + data_bytes = data.encode('utf-8') + + handler.send_response(200) + handler.send_header('Content-Type', 'text/html; charset=utf-8') + handler.send_header('Content-Length', str(len(data_bytes))) + handler.end_headers() + handler.wfile.write(data_bytes) + return True + + except Exception as e: + handler.send_error(502, f"Failed to fetch package index: {str(e)}") + return False + + def _convert_simple_index_html(self, html: str) -> str: + """转换根索引页面的HTML""" + import re + # 替换上游链接 + def convert_link(match): + href = match.group(1) + text = match.group(2) + if href.startswith('/simple/'): + return match.group(0) # 已经是相对路径 + elif href.startswith('https://pypi.tuna.tsinghua.edu.cn/simple/'): + simple_part = href.split("/simple/")[-1] + return f'{text}' + elif href.startswith('https://'): + # 其他上游链接,提取包名 + pkg_name = href.rstrip('/').split('/')[-1] + return f'{text}' + return match.group(0) + + # 匹配 text + return re.sub(r']+href="([^"]+)"[^>]*>([^<]*)', convert_link, html) + + def _handle_simple_index(self, handler, package: str) -> bool: + """处理Simple API索引请求""" + import json + import time + + package = package.lower() + + # 调试 - 确保函数被调用 + debug_file = '/tmp/pypi_debug.log' + with open(debug_file, 'a') as f: + f.write(f"[HANDLE_SIMPLE] START package={package}\n") + + # 检查客户端Accept header + accept = handler.headers.get('Accept', '') + wants_json = 'application/vnd.pypi.simple.v1+json' in accept + + # 调试 + debug_file = '/tmp/pypi_debug.log' + with open(debug_file, 'a') as f: + f.write(f"[SIMPLE_INDEX] package={package}, wants_json={wants_json}, accept={accept[:50]}\n") + + # 根据请求格式选择正确的缓存key,统一使用 simple/ 前缀 + cache_key = f"simple/{package}" + + cached = self._get_cache(cache_key) + + # 调试缓存 + debug_file = '/tmp/pypi_debug.log' + with open(debug_file, 'a') as f: + f.write(f"[CACHE_CHECK] cache_key={cache_key}, cached={'YES' if cached else 'NO'}\n") + + if cached: + # 返回缓存,使用正确的Content-Type + if wants_json: + handler.send_response(200) + handler.send_header('Content-Type', 'application/vnd.pypi.simple.v1+json; charset=utf-8') + else: + handler.send_response(200) + handler.send_header('Content-Type', 'text/html; charset=utf-8') + handler.send_header('Content-Length', str(len(cached))) + handler.end_headers() + handler.wfile.write(cached) + return True + + # 从上游获取HTML + # upstream_url 已经是完整路径(如 https://pypi.tuna.tsinghua.edu.cn/simple) + # 所以只需要添加 /package/ + url = f"{self.upstream_url}/{package}/" + + try: + req = urllib.request.Request(url) + req.add_header('Accept', 'text/html') + req.add_header('User-Agent', 'PyPI-Mirror/1.0') + + with urllib.request.urlopen(req, timeout=30) as response: + data = response.read().decode('utf-8') + + # 根据客户端请求返回不同格式,统一使用 simple/ 路径 + if wants_json: + # 转换为JSON格式 + json_data = self._convert_to_json(package, data) + proxy_data = json.dumps(json_data) + content_type = 'application/vnd.pypi.simple.v1+json; charset=utf-8' + cache_key = f"simple/{package}" + else: + # 转换为HTML格式 + debug_file = '/tmp/pypi_debug.log' + with open(debug_file, 'a') as f: + f.write(f"[CONVERT_CALL] Before conversion\n") + try: + proxy_data = self._convert_simple_html(package, data) + with open(debug_file, 'a') as f: + f.write(f"[CONVERT_CALL] After conversion\n") + except Exception as e: + with open(debug_file, 'a') as f: + f.write(f"[CONVERT_ERROR] {e}\n") + proxy_data = data # fallback to raw data + content_type = 'application/vnd.pypi.simple.v1+html; charset=utf-8' + + if True: + self._set_cache(cache_key, proxy_data.encode('utf-8')) + + proxy_bytes = proxy_data.encode('utf-8') + handler.send_response(200) + handler.send_header('Content-Type', content_type) + handler.send_header('Content-Length', str(len(proxy_bytes))) + handler.end_headers() + handler.wfile.write(proxy_bytes) + return True + + except urllib.error.HTTPError as e: + if e.code == 404: + handler.send_error(404, f"Package not found: {package}") + else: + handler.send_error(502, f"Failed to fetch from upstream: {str(e)}") + return False + + def _convert_to_json(self, package: str, html: str) -> dict: + """将HTML转换为JSON格式""" + import re + + # 解析HTML中的链接 + links = [] + + # 匹配 text + pattern = r']+href="([^"]+)"[^>]*>([^<]*)' + matches = re.findall(pattern, html) + + for href, text in matches: + # 提取文件名(从链接文本) + filename = text.strip() if text.strip() else '' + + # 如果没有链接文本,从URL中提取 + if not filename: + if '#' in href: + filename = href.split('#')[0].split('/')[-1] + else: + filename = href.split('/')[-1] + + # JSON格式中URL不应该包含fragment + # pip从filename字段提取版本号(如 Flask-1.0.0.tar.gz -> 1.0.0) + + # 解析URL并转换为代理路径 + url = href + if href.startswith('../'): + # 相对路径 - 需要转换为代理路径 + # 格式: ../../packages/hash1/hash2/fullhash/filename#sha256=... + # 或: ../../packages/hash1/hash2/filename + # parts = ['..', '..', 'packages', 'hash1', 'hash2', 'fullhash', 'filename', ...] + parts = href.split('/') + try: + pkg_idx = parts.index('packages') + # 提取从 packages 后面到文件名之前的所有部分作为 hash 路径 + # 文件名是最后一个非空部分(可能包含 #fragment) + # 找到文件名的位置(最后一个部分) + filename_idx = len(parts) - 1 + while filename_idx > pkg_idx and not parts[filename_idx]: + filename_idx -= 1 + # hash_path 是 packages 后面到文件名之前的所有部分 + if filename_idx > pkg_idx + 1: + hash_path = '/'.join(parts[pkg_idx+1:filename_idx]) + else: + hash_path = parts[pkg_idx+1] if pkg_idx + 1 < len(parts) else '' + # 文件名: 检查 pkg_idx+3 是否存在且不是哈希 + if pkg_idx + 3 < len(parts): + fname_full = parts[pkg_idx+3] + # 如果 fname_full 看起来像哈希(包含 sha256= 或长度>=32的十六进制),则使用原始 filename + # 清华源格式: .../hash/filename#sha256=... + # 其中 hash 是 28-30 位十六进制 + is_hash_like = ('sha256=' in fname_full or 'sha512=' in fname_full or + (len(fname_full) >= 28 and all(c in '0123456789abcdef' for c in fname_full[:28].lower()))) + if is_hash_like: + # 这是哈希,不是文件名 + fname = filename if filename else '' + else: + # 这是文件名 + fname = fname_full.split('#')[0] + if not filename: + filename = fname + else: + fname = filename if filename else '' + # URL不包含fragment + url = f"/pypi/packages/{hash_path}/{fname}" + except ValueError: + pass + elif href.startswith('/pypi/'): + # 绝对路径(如 /pypi/packages/hash/filename#egg=package-version) + # 去掉fragment + href_clean = href.split('#')[0] + parts = href_clean.split('/') + # parts = ['', 'pypi', 'packages', 'hash', 'filename'] + if len(parts) >= 5: + hash_path = parts[3] + fname = parts[4] + if not filename: + filename = fname + url = f"/pypi/packages/{hash_path}/{fname}" + elif href.startswith('http'): + # 绝对URL - 转换为代理路径 + if 'files.pythonhosted.org' in href or 'files.pypi.org' in href: + # 格式: https://files.pythonhosted.org/packages/hash1/hash2/完整哈希/filename + # 例如: https://files.pythonhosted.org/packages/ec/f9/7f9263c5695f4bd0023734af91bedb2ff8209e8de6ead162f35d8dc762fd/flask-3.1.2-py3-none-any.whl + parts = href.split('/packages/') + if len(parts) >= 2: + path_after_packages = parts[1] + # 完整路径: hash1/hash2/完整哈希/filename + url = f"/pypi/packages/{path_after_packages}" + fname = path_after_packages.split('/')[-1] + if not filename: + filename = fname + elif 'pypi.tuna.tsinghua.edu.cn' in href or 'mirrors.tuna.tsinghua.edu.cn' in href: + parts = href.rsplit('/', 1) + if len(parts) == 2: + path_part = parts[0] + fname = parts[1] + hash_path = path_part.split('/')[-1] + if not filename: + filename = fname + url = f"/pypi/packages/{hash_path}/{fname}" + + link_entry = { + "filename": filename, + "url": url + } + + links.append(link_entry) + + return { + "meta": { + "api-version": "1.0", + "repository-version": "1.0" + }, + "name": package, + "files": links + } + + def _handle_package_file(self, handler, package: str, filename: str) -> bool: + """处理包文件请求""" + import urllib.parse + + # 解析文件名 + # 新格式: /pypi/packages/hash/filename#pip=package-version + # 或旧格式: /pypi/packages/filename?url=... + parsed = urllib.parse.urlparse(f"/{filename}") + actual_filename = parsed.path.lstrip('/') + query_params = urllib.parse.parse_qs(parsed.query) + fragment = urllib.parse.parse_qs(parsed.fragment) if parsed.fragment else {} + + # 从fragment中提取包信息(用于缓存键) + actual_package = package + if 'pip' in fragment: + # 格式: #pip=flask-2.0.0 + pip_info = fragment['pip'][0] + if '-' in pip_info: + # 提取版本号 + parts = pip_info.split('-', 1) + if len(parts) == 2: + actual_package = parts[0] + + cache_key = f"packages/{actual_filename}" + + cached = self._get_cache(cache_key) + if cached: + handler.send_response(200) + handler.send_header('Content-Type', 'application/octet-stream') + handler.send_header('Content-Length', str(len(cached))) + handler.end_headers() + handler.wfile.write(cached) + return True + + # 尝试从上游获取 + # 格式: /pypi/packages/hash/filename -> 构造上游URL + possible_urls = [] + + # 获取基础URL(去掉 /simple 后缀) + base_url = self.upstream_url.rstrip('/') + if base_url.endswith('/simple'): + base_url = base_url[:-7] + + # 新格式: hash/filename -> 尝试清华源 + # 使用 actual_filename(已解析的纯文件路径) + possible_urls.append(f"{base_url}/packages/{actual_filename}") + + # 尝试官方源 + possible_urls.append(f"https://files.pythonhosted.org/packages/{actual_filename}") + + # 注意: 不再支持 ?url= 参数指定任意上游 URL(SSRF 风险), + # 只允许从配置的上游与官方源获取 + + data = None + last_error = None + + for url in possible_urls: + if not url.startswith(('http://', 'https://')): + continue + try: + req = urllib.request.Request(url) + req.add_header('User-Agent', 'PyPI-Mirror/1.0') + + with urllib.request.urlopen(req, timeout=60) as response: + data = response.read() + break # 成功获取,退出循环 + except Exception as e: + import sys + last_error = e + continue + + if data is None: + handler.send_error(502, f"Failed to fetch package: {last_error}") + return False + + if True: + self._set_cache(cache_key, data) + + handler.send_response(200) + handler.send_header('Content-Type', 'application/octet-stream') + handler.send_header('Content-Length', str(len(data))) + handler.end_headers() + handler.wfile.write(data) + return True + + def _handle_web_api(self, handler, package: str) -> bool: + """处理Web API请求""" + import sys + package = package.lower() + cache_key = f"web/{package}" + + cached = self._get_cache(cache_key) + if cached: + handler.send_response(200) + handler.send_header('Content-Type', 'application/vnd.pypi.simple.v1+json; charset=utf-8') + handler.send_header('Content-Length', str(len(cached))) + handler.end_headers() + handler.wfile.write(cached) + return True + + # 从上游获取 - 需要去掉 /simple 后缀 + base_url = self.upstream_url.rstrip('/') + if base_url.endswith('/simple'): + base_url = base_url[:-7] + url = f"{base_url}/pypi/{package}/json" + + try: + req = urllib.request.Request(url) + req.add_header('Accept', 'application/json') + req.add_header('User-Agent', 'PyPI-Mirror/1.0') + + with urllib.request.urlopen(req, timeout=30) as response: + data = response.read().decode('utf-8') + data_json = json.loads(data) + + # 转换URL + data_json = self._convert_package_json(package, data_json) + + proxy_data = json.dumps(data_json) + + if True: + self._set_cache(cache_key, proxy_data.encode('utf-8')) + + proxy_bytes = proxy_data.encode('utf-8') + handler.send_response(200) + handler.send_header('Content-Type', 'application/vnd.pypi.simple.v1+json; charset=utf-8') + handler.send_header('Content-Length', str(len(proxy_bytes))) + handler.end_headers() + handler.wfile.write(proxy_bytes) + return True + + except urllib.error.HTTPError as e: + handler.send_error(502, f"Failed to fetch package info: {str(e)}") + return False + + def _handle_package_download(self, handler, filename: str) -> bool: + """处理包下载请求""" + import sys + # 使用 packages/ 前缀,保持标准 PyPI 目录结构 + cache_key = f"packages/{filename}" + + cached = self._get_cache(cache_key) + if cached: + handler.send_response(200) + handler.send_header('Content-Type', 'application/octet-stream') + handler.send_header('Content-Length', str(len(cached))) + handler.end_headers() + handler.wfile.write(cached) + return True + + # 从上游获取 - 注意去掉 /simple 后缀 + base_url = self.upstream_url.rstrip('/') + if base_url.endswith('/simple'): + base_url = base_url[:-7] # 去掉 /simple + url = f"{base_url}/packages/{filename}" + + try: + req = urllib.request.Request(url) + req.add_header('User-Agent', 'PyPI-Mirror/1.0') + + with urllib.request.urlopen(req, timeout=120) as response: + data = response.read() + + if True: + self._set_cache(cache_key, data) + + handler.send_response(200) + handler.send_header('Content-Type', 'application/octet-stream') + handler.send_header('Content-Length', str(len(data))) + handler.end_headers() + handler.wfile.write(data) + return True + + except Exception as e: + handler.send_error(502, f"Failed to download: {str(e)}") + return False + + def _handle_legacy(self, handler, path: str) -> bool: + """处理旧版PyPI兼容""" + handler.send_error(410, "Legacy PyPI API is deprecated") + return False + + def _convert_simple_html(self, package: str, html: str) -> str: + """转换Simple API HTML,替换URL为代理地址""" + import urllib.parse + import sys + import os + # 写入调试文件 + debug_file = '/tmp/pypi_debug.log' + with open(debug_file, 'a') as f: + f.write(f"[PyPI] Converting HTML for package: {package}\n") + # 打印前几个链接用于调试 + import re + test_matches = re.findall(r'href="([^"]+)"', html)[:3] + with open(debug_file, 'a') as f: + f.write(f"[PyPI] Sample links: {test_matches}\n") + + def convert_absolute_url(match): + """转换绝对URL为代理链接""" + original_url = match.group(1) if match.lastindex else match.group(0) + # 提取文件名和完整的hash路径 + # 格式: https://pypi.tuna.tsinghua.edu.cn/packages/hash1/hash2/fullhash/filename + # 我们将其转换为: /pypi/packages/hash1/hash2/fullhash/filename#pip= + parts = original_url.rsplit('/', 1) + if len(parts) == 2: + path_part = parts[0] + filename = parts[1] + # 提取从 packages/ 后面的完整路径(包含完整hash) + try: + pkg_idx = path_part.index('/packages/') + hash_path = path_part[pkg_idx + 10:] # 去掉 /packages/ + except ValueError: + hash_path = filename + return f'href="/pypi/packages/{hash_path}/{filename}#pip={package}-{filename.split("-")[1] if "-" in filename else ""}"' + return match.group(0) + + # 替换绝对URL - pypi.tuna.tsinghua.edu.cn (清华源) + html = re.sub( + r'(https://pypi\.tuna\.tsinghua\.edu\.cn/packages/[^"\']+)', + convert_absolute_url, + html + ) + # 替换绝对URL - mirrors.tuna.tsinghua.edu.cn + html = re.sub( + r'(https://mirrors\.tuna\.tsinghua\.edu\.cn/pypi/packages/[^"\']+)', + convert_absolute_url, + html + ) + # 替换绝对URL - files.pypi.org + html = re.sub( + r'(https://files\.pypi\.org/packages/[^"\']+)', + convert_absolute_url, + html + ) + # 替换绝对URL - files.pythonhosted.org + html = re.sub( + r'(https://files\.pythonhosted\.org/packages/[^"\']+)', + convert_absolute_url, + html + ) + + # 替换相对路径链接 - ../../packages/hash1/hash2/fullhash/filename -> /pypi/packages/hash1/hash2/fullhash/filename#pip=... + def convert_relative_match(match): + """转换相对路径链接""" + href = match.group(1) + debug_file = '/tmp/pypi_debug.log' + with open(debug_file, 'a') as f: + f.write(f"[CONVERT] Input href: {href[:80]}...\n") + # 提取文件名 + filename = href.split('/')[-1].split('#')[0] + # 提取完整的hash路径(从 packages/ 后面的所有部分除了文件名) + parts = href.split('/') + with open(debug_file, 'a') as f: + f.write(f"[CONVERT] Parts: {parts}\n") + try: + pkg_idx = parts.index('packages') + # packages 后面到倒数第二个是 hash 路径,最后一个是文件名 + hash_parts = parts[pkg_idx+1:-1] # 除了最后一个(文件名) + with open(debug_file, 'a') as f: + f.write(f"[CONVERT] hash_parts: {hash_parts}\n") + hash_path = '/'.join(hash_parts) if hash_parts else filename + with open(debug_file, 'a') as f: + f.write(f"[CONVERT] hash_path: {hash_path}\n") + except ValueError: + hash_path = filename + with open(debug_file, 'a') as f: + f.write(f"[CONVERT] ValueError, hash_path: {hash_path}\n") + + # 提取版本号 - 从文件名中提取,如 Flask-0.1.tar.gz -> 0.1 + base_name = filename + # 去掉扩展名 + for ext in ['.tar.gz', '.whl', '.tar.bz2', '.tar.xz']: + if base_name.endswith(ext): + base_name = base_name[:-len(ext)] + break + + # 尝试多种大小写组合来去掉包名前缀 + version = base_name + for pkg_name in [package, package.lower(), package.upper(), package.capitalize()]: + if base_name.lower().startswith(pkg_name.lower() + '-'): + version = base_name[len(pkg_name)+1:] + break + + return f'href="/pypi/packages/{hash_path}/{filename}#egg={package}-{version}"' + + # 匹配相对路径的链接 + html = re.sub( + r'href="(\.\./\.\./packages/[^"]+)"', + convert_relative_match, + html + ) + html = re.sub( + r"href='(\.\./\.\./packages/[^']+)'", + convert_relative_match, + html + ) + + return html + + def _convert_package_json(self, package: str, data: dict) -> dict: + """转换Package JSON,替换URL为代理地址""" + # 转换URL函数 + def convert_url(url): + # 处理 files.pythonhosted.org 和 files.pypi.org + # 格式: https://files.pythonhosted.org/packages///<完整哈希>/ + # 例如: https://files.pythonhosted.org/packages/ec/f9/7f9263c5695f4bd0023734af91bedb2ff8209e8de6ead162f35d8dc762fd/flask-3.1.2-py3-none-any.whl + if 'files.pythonhosted.org' in url or 'files.pypi.org' in url: + path_parts = url.split('/packages/') + if len(path_parts) >= 2: + # 直接使用 /packages/ 后的完整路径 + path_after_packages = path_parts[1] + return f'/pypi/packages/{path_after_packages}' + return url + + # 转换urls + if 'urls' in data: + for item in data['urls']: + if 'url' in item: + item['url'] = convert_url(item['url']) + + return data + + def _fetch(self, url: str) -> Optional[bytes]: + """从URL获取数据""" + try: + req = urllib.request.Request(url) + req.add_header('User-Agent', 'PyPI-Mirror/1.0') + + with urllib.request.urlopen(req, timeout=30) as response: + return response.read() + + except Exception: + return None + + def _get_cache(self, cache_key: str) -> Optional[bytes]: + """获取缓存""" + if not True: + return None + + cache_path = self._get_cache_path(cache_key) + if cache_path is None: + return None + meta_path = cache_path + '.meta' + + if not os.path.exists(cache_path): + return None + + if os.path.exists(meta_path): + try: + with open(meta_path, 'r') as f: + meta = json.load(f) + if time.time() > meta.get('expires', 0): + return None + except Exception: + pass + + try: + with open(cache_path, 'rb') as f: + return f.read() + except Exception: + return None + + def _set_cache(self, cache_key: str, data: bytes): + """设置缓存""" + cache_path = self._get_cache_path(cache_key) + if cache_path is None: + return + meta_path = cache_path + '.meta' + + os.makedirs(os.path.dirname(cache_path), exist_ok=True) + + try: + with open(cache_path, 'wb') as f: + f.write(data) + + meta = { + 'cached_at': time.time(), + 'expires': time.time() + 86400, + 'size': len(data) + } + + with open(meta_path, 'w') as f: + json.dump(meta, f) + + except Exception as e: + print(f"PyPI缓存写入失败: {e}") + + def _sanitize_cache_key(self, cache_key: str): + """清洗缓存键,拒绝路径穿越(返回 None 表示不安全)""" + safe_key = cache_key.replace('\\', '/') + parts = [] + for seg in safe_key.split('/'): + if seg in ('', '.'): + continue + if seg == '..': + return None # 路径穿越 + parts.append(seg) + return '/'.join(parts) + + def _get_cache_path(self, cache_key: str): + """获取缓存路径(防路径穿越,不安全返回 None)""" + # cache_key 格式: packages/fe/df/88ccbee.../filename + # 存储路径: downloads/pypi-cn/packages/fe/df/88ccbee.../filename + safe_key = self._sanitize_cache_key(cache_key) + if safe_key is None: + return None + path = os.path.join(self.storage_dir, safe_key) + if not os.path.realpath(path).startswith(os.path.realpath(self.storage_dir) + os.sep): + return None + return path + + def get_cache_stats(self) -> dict: + """获取缓存统计""" + if not os.path.exists(self.storage_dir): + return {'files': 0, 'size': 0} + + total_size = 0 + file_count = 0 + + for root, dirs, files in os.walk(self.storage_dir): + for f in files: + if not f.endswith('.meta'): + file_count += 1 + total_size += os.path.getsize(os.path.join(root, f)) + + return { + 'files': file_count, + 'size': total_size, + 'size_formatted': self._format_size(total_size) + } + + def _format_size(self, size_bytes: int) -> str: + """格式化文件大小""" + if size_bytes == 0: + return "0 B" + + units = ["B", "KB", "MB", "GB"] + i = 0 + while size_bytes >= 1024 and i < len(units) - 1: + size_bytes /= 1024.0 + i += 1 + + return f"{size_bytes:.2f} {units[i]}"