Baseline: pr1 HYC下载站 v2.3 before security/functional fixes

This commit is contained in:
HYC Fixer
2026-08-30 12:12:58 +08:00
commit a8e773839b
77 changed files with 38568 additions and 0 deletions
+15
View File
@@ -0,0 +1,15 @@
# 核心模块初始化
from .config import ConfigManager
from .mirror_sync import MirrorSyncManager
from .server import MirrorServer
from .utils import format_file_size, get_file_hash, parse_size, sanitize_filename
__all__ = [
'ConfigManager',
'MirrorSyncManager',
'MirrorServer',
'format_file_size',
'get_file_hash',
'parse_size',
'sanitize_filename'
]
+591
View File
@@ -0,0 +1,591 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
告警模块
支持邮件告警、Webhook 告警、告警规则引擎
"""
import os
import sys
import json
import time
import threading
import logging
import smtplib
import requests
from datetime import datetime
from typing import Dict, List, Optional, Callable
from email.mime.text import MIMEText
from email.mime.multipart import MIMEMultipart
from enum import Enum
logger = logging.getLogger(__name__)
class AlertSeverity(Enum):
"""告警级别"""
INFO = "info"
WARNING = "warning"
ERROR = "error"
CRITICAL = "critical"
class AlertType(Enum):
"""告警类型"""
DISK_HIGH = "disk_high"
DISK_CRITICAL = "disk_critical"
SYNC_FAILED = "sync_failed"
SOURCE_UNHEALTHY = "source_unhealthy"
CACHE_FULL = "cache_full"
SERVICE_DOWN = "service_down"
CUSTOM = "custom"
class Alert:
"""告警对象"""
def __init__(
self,
alert_type: str,
severity: AlertSeverity,
title: str,
message: str,
details: Dict = None,
source: str = None
):
self.id = f"{int(time.time())}_{threading.get_ident()}"
self.type = alert_type
self.severity = severity
self.title = title
self.message = message
self.details = details or {}
self.source = source
self.timestamp = datetime.now()
self.sent = False
self.acknowledged = False
def to_dict(self) -> Dict:
"""转换为字典"""
return {
'id': self.id,
'type': self.type,
'severity': self.severity.value,
'title': self.title,
'message': self.message,
'details': self.details,
'source': self.source,
'timestamp': self.timestamp.isoformat(),
'sent': self.sent,
'acknowledged': self.acknowledged
}
class EmailAlerter:
"""邮件告警器"""
def __init__(self, config: Dict = None):
"""
初始化邮件告警器
Args:
config: 邮件配置
"""
self.config = config or {}
self.enabled = self.config.get('enabled', False)
self.smtp_host = self.config.get('smtp_host', 'localhost')
self.smtp_port = self.config.get('smtp_port', 587)
self.smtp_user = self.config.get('smtp_user', '')
self.smtp_password = self.config.get('smtp_password', '')
self.from_address = self.config.get('from_address', 'hyc-mirror@localhost')
self.to_addresses = self.config.get('to_addresses', [])
self.use_tls = self.config.get('use_tls', True)
# 连接池
self._connection: Optional[smtplib.SMTP] = None
self._last_connect_time: Optional[datetime] = None
self._connection_timeout = 30
def _get_connection(self) -> smtplib.SMTP:
"""获取 SMTP 连接"""
if self._connection:
# 检查连接是否仍然有效
try:
self._connection.noop()
return self._connection
except Exception:
try:
self._connection.quit()
except Exception:
pass
self._connection = None
# 创建新连接
try:
self._connection = smtplib.SMTP(self.smtp_host, self.smtp_port, timeout=self._connection_timeout)
if self.use_tls:
self._connection.starttls()
if self.smtp_user and self.smtp_password:
self._connection.login(self.smtp_user, self.smtp_password)
self._last_connect_time = datetime.now()
return self._connection
except Exception as e:
logger.error(f"Failed to connect to SMTP server: {e}")
raise
def send(self, alert: Alert) -> bool:
"""
发送告警邮件
Args:
alert: 告警对象
Returns:
是否发送成功
"""
if not self.enabled:
logger.debug("Email alerts disabled")
return False
if not self.to_addresses:
logger.warning("No recipients configured for email alerts")
return False
try:
msg = MIMEMultipart('alternative')
msg['Subject'] = f"[{alert.severity.value.upper()}] {alert.title}"
msg['From'] = self.from_address
msg['To'] = ', '.join(self.to_addresses)
# HTML 格式
html_content = self._format_html(alert)
msg.attach(MIMEText(html_content, 'html', 'utf-8'))
# 纯文本格式
text_content = self._format_text(alert)
msg.attach(MIMEText(text_content, 'plain', 'utf-8'))
# 发送邮件
server = self._get_connection()
server.send_message(msg)
logger.info(f"Alert email sent: {alert.title}")
return True
except Exception as e:
logger.error(f"Failed to send alert email: {e}")
return False
def _format_html(self, alert: Alert) -> str:
"""格式化 HTML 内容"""
severity_colors = {
'info': '#2196F3',
'warning': '#FF9800',
'error': '#F44336',
'critical': '#9C27B0'
}
color = severity_colors.get(alert.severity.value, '#666666')
details_html = ''
if alert.details:
details_html = '<h3>Details</h3><table>'
for key, value in alert.details.items():
details_html += f'<tr><td><b>{key}:</b></td><td>{value}</td></tr>'
details_html += '</table>'
return f"""
<html>
<head>
<style>
body {{ font-family: Arial, sans-serif; margin: 20px; }}
.header {{ background-color: {color}; color: white; padding: 10px; }}
.title {{ font-size: 24px; margin: 20px 0; }}
.message {{ background-color: #f5f5f5; padding: 15px; border-radius: 5px; margin: 20px 0; }}
.details {{ margin-top: 20px; }}
.footer {{ margin-top: 30px; font-size: 12px; color: #666; }}
</style>
</head>
<body>
<div class="header">
<h1>HYC Mirror Alert</h1>
</div>
<div class="title">[{alert.severity.value.upper()}] {alert.title}</div>
<div class="message">
<p><b>Message:</b> {alert.message}</p>
<p><b>Time:</b> {alert.timestamp.strftime('%Y-%m-%d %H:%M:%S')}</p>
<p><b>Type:</b> {alert.type}</p>
{details_html}
</div>
<div class="footer">
<p>This is an automated alert from HYC Mirror Server</p>
</div>
</body>
</html>
"""
def _format_text(self, alert: Alert) -> str:
"""格式化纯文本内容"""
details_text = ''
if alert.details:
details_text = '\nDetails:\n'
for key, value in alert.details.items():
details_text += f" {key}: {value}\n"
return f"""
HYC Mirror Alert
================
Severity: {alert.severity.value.upper()}
Title: {alert.title}
Message: {alert.message}
Time: {alert.timestamp.strftime('%Y-%m-%d %H:%M:%S')}
Type: {alert.type}
{details_text}
---
This is an automated alert from HYC Mirror Server
"""
def test_connection(self) -> Dict:
"""测试 SMTP 连接"""
try:
server = self._get_connection()
return {
'success': True,
'message': 'SMTP connection successful'
}
except Exception as e:
return {
'success': False,
'message': f'SMTP connection failed: {str(e)}'
}
def close(self):
"""关闭连接"""
if self._connection:
try:
self._connection.quit()
except Exception:
pass
self._connection = None
class WebhookAlerter:
"""Webhook 告警器"""
def __init__(self, config: Dict = None):
"""
初始化 Webhook 告警器
Args:
config: Webhook 配置
"""
self.config = config or {}
self.enabled = self.config.get('enabled', False)
self.webhook_url = self.config.get('webhook_url', '')
def send(self, alert: Alert) -> bool:
"""
发送告警到 Webhook
Args:
alert: 告警对象
Returns:
是否发送成功
"""
if not self.enabled:
logger.debug("Webhook alerts disabled")
return False
if not self.webhook_url:
logger.warning("No webhook URL configured")
return False
try:
payload = {
'event': 'alert',
'alert': alert.to_dict(),
'timestamp': datetime.now().isoformat()
}
headers = {
'Content-Type': 'application/json',
'User-Agent': 'HYC-Mirror-Alerts/1.0'
}
response = requests.post(
self.webhook_url,
json=payload,
headers=headers,
timeout=30
)
if response.status_code < 400:
logger.info(f"Alert webhook sent: {alert.title}")
return True
else:
logger.error(f"Webhook returned error: {response.status_code}")
return False
except Exception as e:
logger.error(f"Failed to send webhook alert: {e}")
return False
class AlertManager:
"""告警管理器"""
def __init__(self, config: Dict = None):
"""
初始化告警管理器
Args:
config: 告警配置
"""
self.config = config or {}
self.enabled = self.config.get('enabled', False)
# 初始化告警器
self.email_alerter = EmailAlerter(self.config.get('email', {}))
self.webhook_alerter = WebhookAlerter(self.config.get('webhook', {}))
# 告警规则
self.rules = self.config.get('rules', {})
# 告警历史
self._alerts: List[Alert] = []
self._alerts_lock = threading.Lock()
self._max_history = 100
# 回调函数
self._on_alert: Optional[Callable] = None
self._on_ack: Optional[Callable] = None
# 告警冷却(防止重复告警)
self._alert_cooldowns: Dict[str, float] = {}
self._default_cooldown = 300 # 5 分钟
def set_alert_callback(self, callback: Callable):
"""设置告警回调"""
self._on_alert = callback
def set_ack_callback(self, callback: Callable):
"""设置确认回调"""
self._on_ack = callback
def check_rule(self, rule_name: str, data: Dict) -> Optional[Alert]:
"""
检查规则并生成告警
Args:
rule_name: 规则名称
data: 检查数据
Returns:
告警对象或 None
"""
if not self.enabled:
return None
rule = self.rules.get(rule_name, {})
if not rule.get('enabled', False):
return None
severity = AlertSeverity(rule.get('severity', 'warning'))
threshold = rule.get('threshold')
# 磁盘空间检查
if rule_name == 'disk_high' and threshold:
disk_percent = data.get('disk_percent', 0)
if disk_percent >= threshold:
return Alert(
alert_type=AlertType.DISK_HIGH.value,
severity=severity,
title=f"Disk usage is high: {disk_percent}%",
message=f"Disk usage has reached {disk_percent}%, which is above the {threshold}% threshold.",
details={'disk_percent': disk_percent, 'threshold': threshold},
source='monitor'
)
if rule_name == 'disk_critical' and threshold:
disk_percent = data.get('disk_percent', 0)
if disk_percent >= threshold:
return Alert(
alert_type=AlertType.DISK_CRITICAL.value,
severity=severity,
title=f"Disk usage is critical: {disk_percent}%",
message=f"Disk usage has reached {disk_percent}%, which is above the {critical_threshold}% threshold. Immediate action required!",
details={'disk_percent': disk_percent, 'threshold': threshold},
source='monitor'
)
# 同步失败检查
if rule_name == 'sync_failed':
sync_result = data.get('sync_result')
if sync_result and not sync_result.get('success', True):
return Alert(
alert_type=AlertType.SYNC_FAILED.value,
severity=severity,
title=f"Sync failed: {sync_result.get('source', 'unknown')}",
message=sync_result.get('error', 'Unknown sync error'),
details=sync_result,
source='sync'
)
# 源不健康检查
if rule_name == 'source_unhealthy':
unhealthy_sources = data.get('unhealthy_sources', [])
if unhealthy_sources:
return Alert(
alert_type=AlertType.SOURCE_UNHEALTHY.value,
severity=severity,
title=f"Unhealthy mirror sources detected: {len(unhealthy_sources)}",
message=f"The following mirror sources are unhealthy: {', '.join(unhealthy_sources)}",
details={'unhealthy_sources': unhealthy_sources},
source='health_check'
)
return None
def trigger_alert(self, alert: Alert) -> bool:
"""
触发告警
Args:
alert: 告警对象
Returns:
是否发送成功
"""
if not self.enabled:
return False
# 检查冷却时间
cooldown_key = f"{alert.type}:{alert.source or 'unknown'}"
last_alert = self._alert_cooldowns.get(cooldown_key, 0)
if time.time() - last_alert < self._default_cooldown:
logger.debug(f"Alert {alert.type} in cooldown, skipping")
return False
# 发送告警
email_sent = self.email_alerter.send(alert)
webhook_sent = self.webhook_alerter.send(alert)
alert.sent = email_sent or webhook_sent
# 记录告警
with self._alerts_lock:
self._alerts.append(alert)
if len(self._alerts) > self._max_history:
self._alerts = self._alerts[-self._max_history:]
# 更新冷却时间
self._alert_cooldowns[cooldown_key] = time.time()
# 触发回调
if self._on_alert and alert.sent:
try:
self._on_alert(alert)
except Exception as e:
logger.error(f"Alert callback failed: {e}")
return alert.sent
def acknowledge_alert(self, alert_id: str) -> bool:
"""
确认告警
Args:
alert_id: 告警 ID
Returns:
是否成功
"""
with self._alerts_lock:
for alert in self._alerts:
if alert.id == alert_id:
alert.acknowledged = True
if self._on_ack:
try:
self._on_ack(alert)
except Exception as e:
logger.error(f"Ack callback failed: {e}")
return True
return False
def get_alerts(
self,
acknowledged: bool = None,
severity: str = None,
limit: int = 50
) -> List[Dict]:
"""
获取告警列表
Args:
acknowledged: 过滤已确认状态
severity: 过滤级别
limit: 返回数量限制
Returns:
告警列表
"""
with self._alerts_lock:
alerts = [a.to_dict() for a in self._alerts]
# 过滤
if acknowledged is not None:
alerts = [a for a in alerts if a['acknowledged'] == acknowledged]
if severity:
alerts = [a for a in alerts if a['severity'] == severity]
# 返回最近的告警
return alerts[-limit:]
def get_stats(self) -> Dict:
"""获取告警统计"""
with self._alerts_lock:
total = len(self._alerts)
unack = sum(1 for a in self._alerts if not a['acknowledged'])
by_severity = {}
for a in self._alerts:
by_severity[a['severity']] = by_severity.get(a['severity'], 0) + 1
return {
'total_alerts': total,
'unacknowledged': unack,
'by_severity': by_severity,
'email_enabled': self.email_alerter.enabled,
'webhook_enabled': self.webhook_alerter.enabled,
'rules_enabled': sum(1 for r in self.rules.values() if r.get('enabled', False))
}
def clear_history(self) -> bool:
"""清除告警历史"""
with self._alerts_lock:
self._alerts = []
return True
def test_email(self, to_address: str) -> Dict:
"""测试邮件发送"""
test_alert = Alert(
alert_type=AlertType.CUSTOM.value,
severity=AlertSeverity.INFO,
title="Test Alert",
message="This is a test alert from HYC Mirror Server",
details={'test': True}
)
# 临时添加收件人
original_recipients = self.email_alerter.to_addresses
self.email_alerter.to_addresses = [to_address]
success = self.email_alerter.send(test_alert)
# 恢复收件人
self.email_alerter.to_addresses = original_recipients
return {
'success': success,
'message': 'Test email sent successfully' if success else 'Failed to send test email'
}
+587
View File
@@ -0,0 +1,587 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
API认证模块
使用数据库进行认证,支持:
- none: 无认证
- basic: Basic Auth(用户名密码)
- token: Token 认证(登录生成的token)
"""
import os
import sys
import json
import hashlib
import time
import secrets
from typing import Dict, List, Optional
from dataclasses import dataclass
from functools import wraps
@dataclass
class AuthSession:
"""认证会话"""
session_id: str
user_id: str
level: str
created_at: float
expires_at: float
last_activity: float
permissions: List[str]
class APIAuthManager:
"""API认证管理器"""
def __init__(self, config: dict = None):
self.config = config or {}
self.sessions: Dict[str, AuthSession] = {}
# 确定基础目录(用于保存会话文件)
base_dir = config.get('base_dir', '.') if config else '.'
# 检测是否是 PyInstaller 打包环境
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
base_dir = os.path.dirname(os.path.abspath(sys.executable))
elif base_dir == '.':
base_dir = os.getcwd()
# 会话文件路径
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.session_timeout = config.get('auth_session_timeout', 3600) if config else 3600
# Cookie名称
self.cookie_name = 'hyc_auth'
self.cookie_max_age = config.get('auth_cookie_max_age', 86400) if config else 86400
# IP 白名单
self.ip_whitelist = config.get('ip_whitelist', []) if config else []
self.ip_whitelist_enabled = config.get('ip_whitelist_enabled', False) if config else False
# 加载已保存的会话
self._load_sessions()
@property
def db(self):
"""动态获取数据库实例"""
return self.config.get('_db_instance')
def _get_client_ip(self, handler) -> str:
"""从 handler 获取客户端 IP"""
try:
forwarded = handler.headers.get('X-Forwarded-For')
if forwarded:
return forwarded.split(',')[0].strip()
return handler.client_address[0]
except Exception:
return None
def _check_ip_whitelist(self, ip: str) -> bool:
"""检查 IP 是否在白名单中"""
if not self.ip_whitelist_enabled or not self.ip_whitelist:
return True
if not ip:
return False
import ipaddress
for pattern in self.ip_whitelist:
try:
if '/' in pattern:
network = ipaddress.ip_network(pattern, strict=False)
if ipaddress.ip_address(ip) in network:
return True
elif pattern == ip:
return True
except ValueError:
continue
return False
# === Token 验证 ===
def validate_token(self, token: str, client_ip: str = None) -> Optional[dict]:
"""验证 token(从数据库)"""
if not token:
return None
# 优先从数据库验证
if self.db:
user = self.db.get_user_by_token(token)
if user:
# 检查 token 是否过期
if user.get('token_expires_at') and time.time() > user['token_expires_at']:
return {"valid": False, "reason": "Token已过期"}
# 检查用户是否启用
if not user.get('enabled', True):
return {"valid": False, "reason": "用户已被禁用"}
return {
"valid": True,
"key_id": f"user_{user['id']}",
"name": user['username'],
"level": user.get('role', 'admin'),
"permissions": ["*"],
"user_id": user['id'],
"username": user['username']
}
return {"valid": False, "reason": "无效的Token"}
# === Basic Auth 验证 ===
def validate_basic_auth(self, username: str, password: str, client_ip: str = None) -> dict:
"""验证 Basic Auth 用户名密码(从数据库)"""
# IP 白名单检查
if not self._check_ip_whitelist(client_ip):
if self.db:
self.db.add_login_log(username, client_ip, 'failed', 'IP不在白名单')
return {"valid": False, "reason": "IP不在白名单内"}
auth_type = self.config.get('auth_type', 'none')
# 如果认证类型为 none,任何用户都可以通过
if auth_type == 'none':
return {
"valid": True,
"user_id": 0,
"username": username or 'anonymous',
"level": "admin",
"key_id": "anonymous",
"name": f"Anonymous - {username or 'anonymous'}",
"permissions": ["*"]
}
# 从数据库验证
if self.db:
result = self.db.verify_user(username, password)
if result.get('valid'):
if self.db:
self.db.add_login_log(username, client_ip, 'success', '数据库验证')
return {
"valid": True,
"user_id": result.get('user_id'),
"username": username,
"level": result.get('role', 'admin'),
"key_id": f"user_{result.get('user_id')}",
"name": f"User - {username}",
"permissions": ["*"]
}
else:
if self.db:
self.db.add_login_log(username, client_ip, 'failed', result.get('reason', '验证失败'))
return {"valid": False, "reason": result.get('reason', '用户名或密码错误')}
return {"valid": False, "reason": "数据库不可用"}
# === Cookie 验证 ===
def validate_cookie(self, cookie_value: str, client_ip: str = None) -> Optional[dict]:
"""验证认证Cookie"""
if not cookie_value:
return None
parts = cookie_value.split('.')
if len(parts) != 3:
return None
session_id, timestamp, signature = parts
session = self.sessions.get(session_id)
if not session:
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
session.last_activity = time.time()
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
if time.time() > session.expires_at:
del self.sessions[session_id]
return None
session.last_activity = time.time()
return {
"valid": True,
"session_id": session_id,
"user_id": session.user_id,
"level": session.level,
"permissions": session.permissions
}
# === 请求验证 ===
def validate_request(self, handler, required_level: str = "admin") -> dict:
"""
验证请求的认证状态
支持的认证方式:
1. Authorization: Bearer <token>
2. Authorization: Basic <credentials>
3. X-API-Key: <token>
4. Cookie: hyc_auth=<session>
5. ?key=<token>
"""
import base64
auth_header = handler.headers.get('Authorization')
api_key = handler.headers.get('X-API-Key')
cookie = handler.headers.get('Cookie', '')
client_ip = handler.client_address[0] if hasattr(handler, 'client_address') else None
# 提取cookie值
cookie_value = None
for c in cookie.split(';'):
c = c.strip()
if c.startswith(f'{self.cookie_name}='):
cookie_value = c[len(self.cookie_name)+1:]
break
# 获取查询参数中的key
parsed_path = handler.path.split('?')
query_key = None
if len(parsed_path) > 1:
from urllib.parse import parse_qs
query = parse_qs(parsed_path[1])
query_key = query.get('key', [None])[0]
# 1. Bearer Token
if auth_header and auth_header.startswith('Bearer '):
token = auth_header[7:]
result = self.validate_token(token, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "bearer", **result}
# 2. Basic Auth
if auth_header and auth_header.startswith('Basic '):
try:
credentials = base64.b64decode(auth_header[6:]).decode('utf-8')
if ':' in credentials:
username, password = credentials.split(':', 1)
result = self.validate_basic_auth(username, password, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "basic", **result}
except Exception:
pass
# 3. API Key Header
if api_key:
result = self.validate_token(api_key, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "api_key", **result}
# 4. Cookie
if cookie_value:
result = self.validate_cookie(cookie_value, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "cookie", **result}
# 5. Query Parameter
if query_key:
result = self.validate_token(query_key, client_ip)
if result and result.get('valid'):
return {"authenticated": True, "method": "query", **result}
# 未认证
return {
"authenticated": False,
"error": "Authentication required",
"required_level": required_level
}
def check_permission(self, auth_result: dict, permission: str) -> bool:
"""检查是否有权限访问特定API"""
if not auth_result.get('authenticated'):
return False
permissions = auth_result.get('permissions', [])
if '*' in permissions:
return True
if permission in permissions:
return True
for p in permissions:
if p.endswith('*'):
prefix = p.rstrip('*')
if permission.startswith(prefix):
return True
return False
# === 会话管理 ===
def create_session(self, user_id: str, level: str,
permissions: List[str] = None) -> dict:
"""创建认证会话"""
session_id = secrets.token_hex(32)
timestamp = time.time()
session = AuthSession(
session_id=session_id,
user_id=user_id,
level=level,
created_at=timestamp,
expires_at=timestamp + self.session_timeout,
last_activity=timestamp,
permissions=permissions or ['*']
)
self.sessions[session_id] = session
self._save_sessions()
signature = self._generate_signature(session_id, timestamp, user_id)
cookie_value = f"{session_id}.{timestamp}.{signature}"
return {
"session_id": session_id,
"cookie_name": self.cookie_name,
"cookie_value": cookie_value,
"cookie_max_age": self.cookie_max_age,
"expires": timestamp + self.session_timeout
}
def destroy_session(self, session_id: str) -> bool:
"""销毁会话"""
if session_id in self.sessions:
del self.sessions[session_id]
self._save_sessions()
return True
return False
# === 内部方法 ===
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]
def _load_sessions(self):
"""加载会话"""
if os.path.exists(self.sessions_file):
try:
with open(self.sessions_file, 'r', encoding='utf-8') as f:
data = json.load(f)
now = time.time()
for item in data:
if item.get('expires_at') and now > item['expires_at']:
continue
session = AuthSession(
session_id=item['session_id'],
user_id=item['user_id'],
level=item['level'],
created_at=item['created_at'],
expires_at=item['expires_at'],
last_activity=item['last_activity'],
permissions=item.get('permissions', ['*'])
)
self.sessions[session.session_id] = session
except Exception as e:
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
})
try:
with open(self.sessions_file, 'w', encoding='utf-8') as f:
json.dump(data, f, ensure_ascii=False, indent=2)
except Exception:
pass
def get_stats(self) -> dict:
"""获取认证统计"""
active_sessions = sum(
1 for s in self.sessions.values()
if time.time() < s.expires_at
)
return {
"active_sessions": active_sessions,
"session_timeout": self.session_timeout
}
# === API认证装饰器 ===
def require_auth(required_level: str = "admin", permission: str = None):
"""
API认证装饰器
使用方式:
@require_auth()
def api_endpoint(self, handler):
...
@require_auth(permission="sync:start")
def api_sync_start(self, handler):
...
"""
def decorator(func):
@wraps(func)
def wrapper(self, handler, *args, **kwargs):
# 检查是否需要认证
if required_level == "none":
return func(self, handler, *args, **kwargs)
config = getattr(handler, 'config', {})
auth_type = config.get('auth_type', 'none')
# 如果auth_type为none,跳过认证
if auth_type == "none":
handler.auth_result = {
"authenticated": True,
"level": "admin",
"user_id": "anonymous",
"permissions": ["*"]
}
return func(self, handler, *args, **kwargs)
auth_manager = getattr(handler, 'auth_manager', None)
if not auth_manager:
handler.send_json_response({
"error": "认证系统未初始化",
"code": "AUTH_NOT_INITIALIZED"
}, 500)
return
auth_result = auth_manager.validate_request(handler, required_level)
if not auth_result.get('authenticated'):
handler.send_response(401)
handler.send_header('WWW-Authenticate', 'Bearer realm="HYC API"')
handler.send_header('Access-Control-Allow-Origin', '*')
handler.send_json_response({
"error": "未认证或认证已过期",
"code": "UNAUTHORIZED",
"required_level": required_level,
"auth_methods": [
"Authorization: Bearer <token>",
"X-API-Key: <token>",
f"Cookie: {auth_manager.cookie_name}=<session>",
"?key=<token>"
]
})
return
if permission:
if not auth_manager.check_permission(auth_result, permission):
handler.send_json_response({
"error": "权限不足",
"code": "FORBIDDEN",
"required_permission": permission
}, 403)
return
handler.auth_result = auth_result
return func(self, handler, *args, **kwargs)
return wrapper
return decorator
# === 需要认证的API端点定义 ===
ADMIN_API_ENDPOINTS = {
# 同步管理
'POST:/api/v2/sync/*': 'sync:manage',
'POST:/api/v2/sync/*/start': 'sync:start',
'POST:/api/v2/sync/*/stop': 'sync:stop',
'DELETE:/api/v2/sync/*': 'sync:manage',
# 缓存管理
'POST:/api/v2/cache/clean': 'cache:manage',
'DELETE:/api/v2/cache/*': 'cache:manage',
# Webhook管理
'POST:/api/v2/webhooks': 'webhook:create',
'PUT:/api/v2/webhooks/*': 'webhook:update',
'DELETE:/api/v2/webhooks/*': 'webhook:delete',
'POST:/api/v2/webhooks/*/trigger': 'webhook:trigger',
# 服务器配置
'PUT:/api/v2/config': 'config:manage',
'POST:/api/v2/server/reload': 'server:reload',
# 文件管理(高危操作)
'DELETE:/api/v2/files/*': 'files:delete',
'PUT:/api/v2/files/*/rename': 'files:rename',
# 用户管理
'POST:/api/v2/users': 'users:create',
'DELETE:/api/v2/users/*': 'users:delete',
'PUT:/api/v2/users/*': 'users:update',
# === API v1 文件操作认证 ===
'DELETE:/api/v1/file/*': 'files:delete',
'PUT:/api/v1/mkdir': 'files:create',
'POST:/api/v1/upload': 'files:upload',
'POST:/api/v1/batch': 'files:batch',
'POST:/api/v1/archive': 'files:archive',
}
def check_endpoint_auth(method: str, path: str, auth_manager: APIAuthManager) -> dict:
"""检查端点是否需要认证"""
key = f"{method}:{path}"
if key in ADMIN_API_ENDPOINTS:
return {
"required": True,
"permission": ADMIN_API_ENDPOINTS[key]
}
for pattern, permission in ADMIN_API_ENDPOINTS.items():
if '*' in pattern:
pat_method, pat_path = pattern.split(':', 1)
if method == pat_method or pat_method == '*':
if pat_path.endswith('*'):
prefix = pat_path.rstrip('*').rstrip('/')
if path.startswith(prefix):
return {
"required": True,
"permission": permission
}
return {
"required": False,
"permission": None
}
+816
View File
@@ -0,0 +1,816 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
API 文档生成器
自动生成 OpenAPI/Swagger 格式的 API 文档
"""
import os
import sys
import json
import inspect
from datetime import datetime
from typing import Dict, List, Any, Optional
from dataclasses import dataclass, field
@dataclass
class APIDoc:
"""API 文档信息"""
title: str = "HYC下载站 API"
version: str = "2.2.0"
description: str = "HYC镜像下载站 REST API 文档"
servers: List[Dict] = field(default_factory=list)
tags: List[Dict] = field(default_factory=list)
paths: Dict = field(default_factory=dict)
components: Dict = field(default_factory=dict)
class APIDocGenerator:
"""API 文档生成器"""
def __init__(self, config: Dict = None):
self.config = config or {}
self.api_doc = APIDoc()
# 初始化组件
self._init_components()
def _init_components(self):
"""初始化文档组件"""
self.api_doc.components = {
'securitySchemes': {
'BearerAuth': {
'type': 'http',
'scheme': 'bearer',
'bearerFormat': 'JWT',
'description': 'JWT token 认证'
},
'ApiKeyAuth': {
'type': 'apiKey',
'in': 'header',
'name': 'X-API-Key',
'description': 'API Key 认证'
}
},
'schemas': {
'Error': {
'type': 'object',
'properties': {
'error': {'type': 'string', 'description': '错误信息'},
'code': {'type': 'string', 'description': '错误代码'},
'message': {'type': 'string', 'description': '详细描述'}
}
},
'Success': {
'type': 'object',
'properties': {
'success': {'type': 'boolean'},
'message': {'type': 'string'},
'data': {'type': 'object'}
}
},
'HealthStatus': {
'type': 'object',
'properties': {
'status': {'type': 'string', 'enum': ['healthy', 'degraded', 'unhealthy']},
'components': {'type': 'object'},
'timestamp': {'type': 'string', 'format': 'date-time'}
}
}
}
}
def generate(self) -> Dict:
"""生成完整的 API 文档"""
doc = {
'openapi': '3.0.3',
'info': {
'title': self.api_doc.title,
'version': self.api_doc.version,
'description': self.api_doc.description,
'contact': {
'name': 'HYC Mirror Support',
'email': '[email protected]'
},
'license': {
'name': 'MIT',
'url': 'https://opensource.org/licenses/MIT'
}
},
'servers': self.api_doc.servers,
'tags': self.api_doc.tags,
'paths': self.api_doc.paths,
'components': self.api_doc.components
}
return doc
def add_server(self, url: str, description: str = ''):
"""添加服务器"""
self.api_doc.servers.append({
'url': url,
'description': description
})
def add_tag(self, name: str, description: str = ''):
"""添加标签"""
self.api_doc.tags.append({
'name': name,
'description': description
})
def add_endpoint(
self,
method: str,
path: str,
summary: str,
description: str = '',
tags: List[str] = None,
parameters: List[Dict] = None,
requestBody: Dict = None,
responses: Dict = None,
security: List[Dict] = None,
deprecated: bool = False
):
"""
添加 API 端点
Args:
method: HTTP 方法 (GET, POST, PUT, DELETE, PATCH)
path: API 路径
summary: 简要描述
description: 详细描述
tags: 标签列表
parameters: 参数列表
requestBody: 请求体
responses: 响应定义
security: 安全要求
deprecated: 是否废弃
"""
if parameters is None:
parameters = []
if responses is None:
responses = self._default_responses()
if security is None:
security = []
# 转换路径参数
path_params = self._extract_path_params(path)
for param in path_params:
parameters.append({
'name': param,
'in': 'path',
'required': True,
'schema': {'type': 'string'},
'description': f'Path parameter: {param}'
})
# 转换查询参数
query_params = self._extract_query_params(path)
for param in query_params:
parameters.append({
'name': param,
'in': 'query',
'required': False,
'schema': {'type': 'string'},
'description': f'Query parameter: {param}'
})
# 构建路径
clean_path = path.format(**{p: f'{{{p}}}' for p in path_params})
if clean_path not in self.api_doc.paths:
self.api_doc.paths[clean_path] = {}
endpoint = {
'summary': summary,
'description': description,
'tags': tags or [],
'parameters': parameters,
'responses': responses,
'deprecated': deprecated
}
if security:
endpoint['security'] = security
if requestBody:
endpoint['requestBody'] = requestBody
self.api_doc.paths[clean_path][method.lower()] = endpoint
def _default_responses(self) -> Dict:
"""获取默认响应"""
return {
'200': {
'description': 'Successful response',
'content': {
'application/json': {
'schema': {'type': 'object'}
}
}
},
'400': {
'description': 'Bad request',
'content': {
'application/json': {
'schema': {'$ref': '#/components/schemas/Error'}
}
}
},
'401': {
'description': 'Unauthorized',
'content': {
'application/json': {
'schema': {'$ref': '#/components/schemas/Error'}
}
}
},
'403': {
'description': 'Forbidden',
'content': {
'application/json': {
'schema': {'$ref': '#/components/schemas/Error'}
}
}
},
'404': {
'description': 'Not found',
'content': {
'application/json': {
'schema': {'$ref': '#/components/schemas/Error'}
}
}
},
'500': {
'description': 'Internal server error',
'content': {
'application/json': {
'schema': {'$ref': '#/components/schemas/Error'}
}
}
}
}
def _extract_path_params(self, path: str) -> List[str]:
"""提取路径参数"""
import re
return re.findall(r'\{(\w+)\}', path)
def _extract_query_params(self, path: str) -> List[str]:
"""提取查询参数"""
import re
return re.findall(r':(\w+)', path)
def save(self, filepath: str, format: str = 'json'):
"""
保存 API 文档
Args:
filepath: 保存路径
format: 格式 (json, yaml)
"""
doc = self.generate()
if format == 'yaml':
try:
import yaml
with open(filepath, 'w', encoding='utf-8') as f:
yaml.dump(doc, f, default_flow_style=False, allow_unicode=True)
except ImportError:
# 回退为 JSON
format = 'json'
if format == 'json':
with open(filepath, 'w', encoding='utf-8') as f:
json.dump(doc, f, ensure_ascii=False, indent=2)
return filepath
class APIEndpointRegistry:
"""API 端点注册表"""
def __init__(self):
self.endpoints: List[Dict] = []
def register(self, method: str, path: str, handler_name: str, description: str = ''):
"""注册端点"""
self.endpoints.append({
'method': method.upper(),
'path': path,
'handler': handler_name,
'description': description
})
def get_all(self) -> List[Dict]:
"""获取所有端点"""
return self.endpoints
def generate_docs(self) -> Dict:
"""生成文档"""
generator = APIDocGenerator()
for ep in self.endpoints:
generator.add_endpoint(
method=ep['method'],
path=ep['path'],
summary=ep['description'],
description=ep['description']
)
return generator.generate()
def generate_api_docs(config: Dict = None) -> Dict:
"""
生成完整的 API 文档
Args:
config: 服务器配置
Returns:
OpenAPI 格式的文档
"""
generator = APIDocGenerator(config)
# 设置服务器信息
host = config.get('host', 'localhost')
port = config.get('port', 8080)
protocol = 'https' if config.get('ssl_cert') else 'http'
generator.add_server(f'{protocol}://{host}:{port}', 'Production server')
# 添加标签
generator.add_tag('Server', '服务器信息')
generator.add_tag('Monitoring', '监控与指标')
generator.add_tag('Mirrors', '镜像源管理')
generator.add_tag('Sync', '同步管理')
generator.add_tag('Cache', '缓存管理')
generator.add_tag('Health', '健康检查')
generator.add_tag('Alerts', '告警管理')
generator.add_tag('Webhooks', 'Webhook管理')
generator.add_tag('Configuration', '配置管理')
generator.add_tag('Authentication', '认证管理')
# ========== 服务器信息端点 ==========
generator.add_endpoint(
'GET', '/api/v2/server/info',
summary='获取服务器信息',
description='返回服务器的详细信息,包括版本、运行时间、配置等',
tags=['Server']
)
# ========== 监控端点 ==========
generator.add_endpoint(
'GET', '/api/v2/monitor/realtime',
summary='获取实时监控数据',
description='返回服务器的实时监控数据,包括 CPU、内存、磁盘使用情况',
tags=['Monitoring']
)
generator.add_endpoint(
'GET', '/api/v2/monitor/history',
summary='获取历史监控数据',
description='返回指定时间段内的历史监控数据',
tags=['Monitoring'],
parameters=[{
'name': 'period',
'in': 'query',
'schema': {'type': 'string'},
'description': '时间周期: 1h, 6h, 24h, 7d, 30d'
}]
)
generator.add_endpoint(
'GET', '/api/v2/metrics',
summary='获取 Prometheus 指标',
description='返回 Prometheus 格式的监控指标',
tags=['Monitoring']
)
# ========== 镜像源端点 ==========
generator.add_endpoint(
'GET', '/api/v2/mirrors',
summary='列出所有镜像加速源',
description='返回所有可用的镜像加速源列表',
tags=['Mirrors']
)
generator.add_endpoint(
'GET', '/api/v2/mirrors/:name',
summary='获取镜像源详情',
description='返回指定镜像源的详细信息',
tags=['Mirrors']
)
generator.add_endpoint(
'POST', '/api/v2/mirrors',
summary='添加自定义镜像源',
description='添加新的自定义镜像加速源',
tags=['Mirrors'],
requestBody={
'required': True,
'content': {
'application/json': {
'schema': {
'type': 'object',
'properties': {
'name': {'type': 'string'},
'type': {'type': 'string'},
'url': {'type': 'string'},
'enabled': {'type': 'boolean'}
},
'required': ['name', 'url']
}
}
}
}
)
generator.add_endpoint(
'DELETE', '/api/v2/mirrors/:name',
summary='删除自定义镜像源',
description='删除指定的自定义镜像加速源',
tags=['Mirrors']
)
generator.add_endpoint(
'POST', '/api/v2/mirrors/:name/refresh',
summary='刷新镜像源缓存',
description='刷新指定镜像源的缓存数据',
tags=['Mirrors']
)
# ========== 同步管理端点 ==========
generator.add_endpoint(
'GET', '/api/v2/sync/sources',
summary='获取同步源列表',
description='返回所有配置的同步源',
tags=['Sync']
)
generator.add_endpoint(
'POST', '/api/v2/sync/sources',
summary='添加同步源',
description='添加新的同步源配置',
tags=['Sync'],
requestBody={
'required': True,
'content': {
'application/json': {
'schema': {
'type': 'object',
'properties': {
'name': {'type': 'string'},
'type': {'type': 'string'},
'url': {'type': 'string'},
'schedule': {'type': 'string'}
}
}
}
}
}
)
generator.add_endpoint(
'POST', '/api/v2/sync/:source_name/start',
summary='启动同步任务',
description='启动指定源的同步任务',
tags=['Sync']
)
generator.add_endpoint(
'POST', '/api/v2/sync/:source_name/stop',
summary='停止同步任务',
description='停止指定源的同步任务',
tags=['Sync']
)
generator.add_endpoint(
'GET', '/api/v2/sync/:source_name/status',
summary='获取同步状态',
description='返回指定同步源的当前状态',
tags=['Sync']
)
# ========== 缓存管理端点 ==========
generator.add_endpoint(
'GET', '/api/v2/cache/stats',
summary='获取缓存统计',
description='返回缓存的使用统计信息',
tags=['Cache']
)
generator.add_endpoint(
'GET', '/api/v2/cache/usage',
summary='获取缓存使用详情',
description='返回缓存的详细使用情况',
tags=['Cache']
)
generator.add_endpoint(
'POST', '/api/v2/cache/clean',
summary='清理缓存',
description='清理指定或全部缓存',
tags=['Cache'],
requestBody={
'content': {
'application/json': {
'schema': {
'type': 'object',
'properties': {
'source': {'type': 'string'}
}
}
}
}
}
)
# ========== 缓存预热端点 ==========
generator.add_endpoint(
'GET', '/api/v2/cache/prewarm',
summary='获取预热状态',
description='返回缓存预热的当前状态',
tags=['Cache']
)
generator.add_endpoint(
'POST', '/api/v2/cache/prewarm',
summary='执行缓存预热',
description='手动执行缓存预热任务',
tags=['Cache']
)
generator.add_endpoint(
'GET', '/api/v2/cache/prewarm/items',
summary='获取预热项目列表',
description='返回待预热的项目列表',
tags=['Cache']
)
generator.add_endpoint(
'POST', '/api/v2/cache/prewarm/clear',
summary='清空预热队列',
description='清空待预热的项目队列',
tags=['Cache']
)
# ========== 健康检查端点 ==========
generator.add_endpoint(
'GET', '/api/v2/health',
summary='获取健康状态',
description='返回服务器的整体健康状态',
tags=['Health']
)
generator.add_endpoint(
'GET', '/api/v2/health/sources',
summary='获取镜像源健康状态',
description='返回所有镜像源的健康检查结果',
tags=['Health']
)
generator.add_endpoint(
'GET', '/api/v2/health/check/:source_name',
summary='检查指定源健康',
description='手动触发指定镜像源的健康检查',
tags=['Health']
)
generator.add_endpoint(
'GET', '/api/v2/health/failover',
summary='获取故障切换状态',
description='返回故障切换系统的当前状态',
tags=['Health']
)
generator.add_endpoint(
'POST', '/api/v2/health/failover/:mirror_type',
summary='触发故障切换',
description='手动触发指定镜像类型的故障切换',
tags=['Health']
)
# ========== 告警管理端点 ==========
generator.add_endpoint(
'GET', '/api/v2/alerts',
summary='获取告警列表',
description='返回当前告警列表',
tags=['Alerts'],
parameters=[{
'name': 'limit',
'in': 'query',
'schema': {'type': 'integer'},
'description': '返回数量限制'
}]
)
generator.add_endpoint(
'POST', '/api/v2/alerts/:alert_id/acknowledge',
summary='确认告警',
description='确认指定告警',
tags=['Alerts']
)
generator.add_endpoint(
'POST', '/api/v2/alerts/clear',
summary='清除告警历史',
description='清除所有告警历史记录',
tags=['Alerts']
)
generator.add_endpoint(
'POST', '/api/v2/alerts/test',
summary='测试告警发送',
description='发送测试告警以验证配置',
tags=['Alerts']
)
generator.add_endpoint(
'GET', '/api/v2/alerts/config',
summary='获取告警配置',
description='返回当前的告警配置',
tags=['Alerts']
)
generator.add_endpoint(
'PUT', '/api/v2/alerts/config',
summary='更新告警配置',
description='更新告警配置(邮件、Webhook 等)',
tags=['Alerts']
)
# ========== Webhook 端点 ==========
generator.add_endpoint(
'GET', '/api/v2/webhooks',
summary='列出所有 Webhook',
description='返回所有配置的 Webhook',
tags=['Webhooks']
)
generator.add_endpoint(
'POST', '/api/v2/webhooks',
summary='创建 Webhook',
description='创建新的 Webhook 配置',
tags=['Webhooks']
)
generator.add_endpoint(
'GET', '/api/v2/webhooks/:webhook_id',
summary='获取 Webhook 详情',
description='返回指定 Webhook 的详细信息',
tags=['Webhooks']
)
generator.add_endpoint(
'PUT', '/api/v2/webhooks/:webhook_id',
summary='更新 Webhook',
description='更新指定 Webhook 的配置',
tags=['Webhooks']
)
generator.add_endpoint(
'DELETE', '/api/v2/webhooks/:webhook_id',
summary='删除 Webhook',
description='删除指定的 Webhook',
tags=['Webhooks']
)
generator.add_endpoint(
'POST', '/api/v2/webhooks/:webhook_id/test',
summary='测试 Webhook',
description='发送测试请求到指定的 Webhook',
tags=['Webhooks']
)
generator.add_endpoint(
'GET', '/api/v2/webhooks/:webhook_id/deliveries',
summary='获取 Webhook 交付历史',
description='返回指定 Webhook 的交付历史记录',
tags=['Webhooks']
)
generator.add_endpoint(
'GET', '/api/v2/webhooks/:webhook_id/stats',
summary='获取 Webhook 统计',
description='返回指定 Webhook 的交付统计信息',
tags=['Webhooks']
)
# ========== 配置管理端点 ==========
generator.add_endpoint(
'GET', '/api/v2/config',
summary='获取配置',
description='返回当前的服务器配置',
tags=['Configuration']
)
generator.add_endpoint(
'PUT', '/api/v2/config',
summary='保存配置',
description='保存配置到 settings.json',
tags=['Configuration']
)
generator.add_endpoint(
'POST', '/api/v2/config/reload',
summary='重新加载配置',
description='重新加载配置文件(热更新)',
tags=['Configuration']
)
generator.add_endpoint(
'GET', '/api/v2/config/changes',
summary='获取配置变更历史',
description='返回配置变更的历史记录',
tags=['Configuration']
)
# ========== 重启管理端点 ==========
generator.add_endpoint(
'GET', '/api/v2/server/restart',
summary='获取重启状态',
description='返回服务器重启管理的当前状态',
tags=['Server']
)
generator.add_endpoint(
'POST', '/api/v2/server/restart',
summary='准备重启',
description='准备执行服务器重启',
tags=['Server']
)
generator.add_endpoint(
'POST', '/api/v2/server/restart/confirm',
summary='确认执行重启',
description='确认并执行服务器重启',
tags=['Server']
)
generator.add_endpoint(
'POST', '/api/v2/server/restart/immediate',
summary='立即重启',
description='立即重启服务器(不等待请求完成)',
tags=['Server']
)
generator.add_endpoint(
'GET', '/api/v2/server/restart/pending',
summary='获取待处理请求',
description='返回当前待处理的请求列表',
tags=['Server']
)
generator.add_endpoint(
'GET', '/api/v2/server/restart/history',
summary='获取重启历史',
description='返回服务器重启的历史记录',
tags=['Server']
)
# ========== 认证端点 ==========
generator.add_endpoint(
'POST', '/api/v2/admin/auth/verify',
summary='验证认证状态',
description='验证当前请求的认证状态',
tags=['Authentication']
)
return generator.generate()
def save_api_docs(config: Dict, filepath: str = None, format: str = 'json'):
"""
生成并保存 API 文档
Args:
config: 服务器配置
filepath: 保存路径
format: 输出格式 (json, yaml)
"""
doc = generate_api_docs(config)
if filepath is None:
# 默认保存到项目根目录
script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
filepath = os.path.join(script_dir, 'docs', 'api-docs.json')
# 确保目录存在
os.makedirs(os.path.dirname(filepath), exist_ok=True)
if format == 'yaml':
try:
import yaml
with open(filepath, 'w', encoding='utf-8') as f:
yaml.dump(doc, f, default_flow_style=False, allow_unicode=True)
except ImportError:
format = 'json'
if format == 'json':
with open(filepath, 'w', encoding='utf-8') as f:
json.dump(doc, f, ensure_ascii=False, indent=2)
return filepath
+677
View File
@@ -0,0 +1,677 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
缓存预热模块
预热常用镜像缓存提高命中率
"""
import os
import sys
import json
import time
import threading
import logging
import requests
from datetime import datetime
from typing import Dict, List, Optional, Callable
from dataclasses import dataclass, field
from enum import Enum
from concurrent.futures import ThreadPoolExecutor, as_completed
logger = logging.getLogger(__name__)
class PrewarmPriority(Enum):
"""预热优先级"""
CRITICAL = "critical" # 关键,必须预热
HIGH = "high" # 高优先级
MEDIUM = "medium" # 中优先级
LOW = "low" # 低优先级
class ItemStatus(Enum):
"""项目状态"""
PENDING = "pending"
IN_PROGRESS = "in_progress"
SUCCESS = "success"
FAILED = "failed"
SKIPPED = "skipped"
@dataclass
class PrewarmItem:
"""预热项目"""
id: str
mirror_type: str
item_name: str
url: str
priority: str
status: str = field(default=ItemStatus.PENDING.value)
attempts: int = 0
max_attempts: int = 3
response_time_ms: float = 0
error_message: str = None
prewarmed_at: float = None
size_bytes: int = 0
def to_dict(self) -> Dict:
return {
'id': self.id,
'mirror_type': self.mirror_type,
'item_name': self.item_name,
'url': self.url,
'priority': self.priority,
'status': self.status,
'attempts': self.attempts,
'response_time_ms': self.response_time_ms,
'error_message': self.error_message,
'prewarmed_at': self.prewarmed_at,
'size_bytes': self.size_bytes
}
@dataclass
class PrewarmTarget:
"""预热目标配置"""
mirror_type: str
priority: str
limit: int
items: List[str] = field(default_factory=list) # 指定的预热项目列表
tags: List[str] = field(default_factory=list) # 按标签筛选
class CachePrewarmer:
"""缓存预热管理器"""
def __init__(self, config: Dict = None):
"""
初始化缓存预热管理器
Args:
config: 预热配置
"""
self.config = config or {}
self.enabled = self.config.get('enabled', False)
# 预热目标
self.targets = self._parse_targets(self.config.get('targets', []))
# 调度配置
self.schedule = self.config.get('schedule', '0 3 * * *') # 默认每天凌晨3点
self.batch_size = self.config.get('batch_size', 10)
# 状态管理
self._items: Dict[str, PrewarmItem] = {}
self._items_lock = threading.Lock()
self._is_running = False
self._run_lock = threading.Lock()
# 历史记录
self._history: List[Dict] = []
self._history_lock = threading.Lock()
# 回调函数
self._on_start: Optional[Callable] = None
self._on_complete: Optional[Callable] = None
self._on_item_complete: Optional[Callable] = None
self._on_error: Optional[Callable] = None
# HTTP 会话
self._session = requests.Session()
self._session.headers.update({
'User-Agent': 'HYC-Mirror-Prewarmer/1.0'
})
# 超时设置
self._request_timeout = self.config.get('request_timeout', 30)
# 常用镜像包列表
self._popular_items = self._load_popular_items()
def _parse_targets(self, targets_config: List[Dict]) -> List[PrewarmTarget]:
"""解析预热目标配置"""
targets = []
for t in targets_config:
targets.append(PrewarmTarget(
mirror_type=t.get('mirror_type', 'docker'),
priority=t.get('priority', 'medium'),
limit=t.get('limit', 100),
items=t.get('items', []),
tags=t.get('tags', [])
))
return targets
def _load_popular_items(self) -> Dict[str, List[str]]:
"""加载常用镜像包列表"""
return {
'docker': [
'library/alpine:latest',
'library/ubuntu:latest',
'library/debian:latest',
'library/centos:latest',
'library/nginx:latest',
'library/python:3.9',
'library/python:3.10',
'library/node:18',
'library/node:20',
'library/go:1.20',
'library/redis:alpine',
'library/mysql:8',
'library/postgres:15',
],
'pip': [
'requests',
'numpy',
'pandas',
'flask',
'django',
'scipy',
'scikit-learn',
'torch',
'tensorflow',
'celery',
'pytest',
'black',
],
'npm': [
'react',
'vue',
'angular',
'lodash',
'express',
'axios',
'typescript',
'webpack',
'vite',
'eslint',
],
'apt': [
'ubuntu-desktop',
'ubuntu-standard',
'nginx',
'python3-pip',
'nodejs',
'docker.io',
],
'yum': [
'epel-release',
'nginx',
'docker-ce',
'python3-pip',
'nodejs',
],
'go': [
'golang.org/x/tools',
'github.com/gin-gonic/gin',
'github.com/beego/beego',
'github.com/gorilla/mux',
],
}
def _get_base_url(self, mirror_type: str) -> str:
"""获取镜像源基础 URL"""
mirrors = self.config.get('mirrors', {})
if isinstance(mirrors, dict):
mirror_config = mirrors.get(mirror_type, {})
if isinstance(mirror_config, dict):
sources = mirror_config.get('sources', [])
if sources:
source_config = mirror_config.get('sources_config', {}).get(sources[0], {})
return source_config.get('url', '')
return ''
def _generate_url(self, mirror_type: str, item_name: str) -> str:
"""生成预热 URL"""
base_url = self._get_base_url(mirror_type)
if not base_url:
return ''
if mirror_type == 'docker':
return f"{base_url}/v2/{item_name}/manifests/latest"
elif mirror_type == 'pip':
return f"{base_url}/simple/{item_name}/"
elif mirror_type == 'npm':
return f"{base_url}/{item_name}"
elif mirror_type == 'apt':
return f"{base_url}/dists/{item_name}/InRelease"
elif mirror_type == 'yum':
return f"{base_url}/repodata/repomd.xml"
elif mirror_type == 'go':
return f"{base_url}/{item_name}?go-get=1"
else:
return f"{base_url}/{item_name}"
def _create_item(
self,
mirror_type: str,
item_name: str,
priority: str = 'medium'
) -> PrewarmItem:
"""创建预热项目"""
url = self._generate_url(mirror_type, item_name)
return PrewarmItem(
id=f"{mirror_type}_{item_name}_{int(time.time())}",
mirror_type=mirror_type,
item_name=item_name,
url=url,
priority=priority
)
def add_item(self, item: PrewarmItem):
"""添加预热项目"""
with self._items_lock:
self._items[item.id] = item
def add_items_batch(self, mirror_type: str, items: List[str], priority: str = 'medium'):
"""批量添加预热项目"""
for item_name in items:
item = self._create_item(mirror_type, item_name, priority)
self.add_item(item)
def _prewarm_item(self, item: PrewarmItem) -> PrewarmItem:
"""
预热单个项目
Args:
item: 预热项目
Returns:
更新后的项目
"""
item.attempts += 1
item.status = ItemStatus.IN_PROGRESS.value
try:
start_time = time.time()
response = self._session.get(
item.url,
timeout=self._request_timeout,
allow_redirects=True
)
response.raise_for_status()
item.response_time_ms = round((time.time() - start_time) * 1000, 2)
item.size_bytes = len(response.content)
item.status = ItemStatus.SUCCESS.value
item.prewarmed_at = time.time()
logger.debug(f"Prewarmed {item.mirror_type}/{item.item_name}: {item.response_time_ms}ms")
except requests.exceptions.Timeout:
item.error_message = f"Timeout after {self._request_timeout}s"
if item.attempts < item.max_attempts:
item.status = ItemStatus.PENDING.value
else:
item.status = ItemStatus.FAILED.value
except requests.exceptions.HTTPError as e:
item.error_message = f"HTTP Error: {e.response.status_code}"
if item.attempts < item.max_attempts:
item.status = ItemStatus.PENDING.value
else:
item.status = ItemStatus.FAILED.value
except Exception as e:
item.error_message = str(e)
item.status = ItemStatus.FAILED.value
return item
def run(self, targets: List[PrewarmTarget] = None) -> Dict:
"""
执行预热
Args:
targets: 预热目标列表,None 表示使用配置的默认目标
Returns:
预热结果
"""
with self._run_lock:
if self._is_running:
return {
'success': False,
'error': 'Prewarm already running'
}
self._is_running = True
start_time = time.time()
result = {
'success': True,
'total_items': 0,
'success_count': 0,
'failed_count': 0,
'skipped_count': 0,
'elapsed_seconds': 0,
'errors': []
}
try:
# 通知开始
if self._on_start:
try:
self._on_start()
except Exception as e:
logger.error(f"Prewarm start callback failed: {e}")
# 确定要预热的目标
if targets is None:
targets = self.targets
# 如果没有指定目标,使用流行项目
if not targets:
for mirror_type, items in self._popular_items.items():
targets.append(PrewarmTarget(
mirror_type=mirror_type,
priority='medium',
limit=len(items),
items=items
))
# 添加项目到队列
total_added = 0
for target in targets:
if target.items:
# 使用指定的预热项目
for item_name in target.items[:target.limit]:
item = self._create_item(target.mirror_type, item_name, target.priority)
self.add_item(item)
total_added += 1
else:
# 使用流行项目列表
popular = self._popular_items.get(target.mirror_type, [])
for item_name in popular[:target.limit]:
item = self._create_item(target.mirror_type, item_name, target.priority)
self.add_item(item)
total_added += 1
result['total_items'] = total_added
# 按优先级排序
priority_order = {
'critical': 0,
'high': 1,
'medium': 2,
'low': 3
}
with self._items_lock:
sorted_items = sorted(
self._items.values(),
key=lambda x: (priority_order.get(x.priority, 99), x.id)
)
# 分批执行
completed = 0
with ThreadPoolExecutor(max_workers=self.batch_size) as executor:
futures = {
executor.submit(self._prewarm_item, item): item
for item in sorted_items
}
for future in as_completed(futures):
item = futures[future]
try:
updated_item = future.result()
# 更新项目状态
with self._items_lock:
self._items[item.id] = updated_item
# 统计
if updated_item.status == ItemStatus.SUCCESS.value:
result['success_count'] += 1
elif updated_item.status == ItemStatus.FAILED.value:
result['failed_count'] += 1
result['errors'].append({
'item': updated_item.item_name,
'error': updated_item.error_message
})
else:
result['skipped_count'] += 1
completed += 1
# 回调
if self._on_item_complete:
try:
self._on_item_complete(updated_item)
except Exception as e:
logger.error(f"Item complete callback failed: {e}")
except Exception as e:
logger.error(f"Prewarm execution error: {e}")
result['failed_count'] += 1
result['errors'].append({
'item': item.item_name,
'error': str(e)
})
result['elapsed_seconds'] = round(time.time() - start_time, 2)
# 记录历史
self._add_to_history(result)
# 通知完成
if self._on_complete:
try:
self._on_complete(result)
except Exception as e:
logger.error(f"Prewarm complete callback failed: {e}")
except Exception as e:
logger.error(f"Prewarm failed: {e}")
result['success'] = False
result['errors'].append({'error': str(e)})
finally:
self._is_running = False
return result
def _add_to_history(self, result: Dict):
"""添加历史记录"""
record = {
'timestamp': datetime.now().isoformat(),
'success': result.get('success', False),
'total_items': result.get('total_items', 0),
'success_count': result.get('success_count', 0),
'failed_count': result.get('failed_count', 0),
'elapsed_seconds': result.get('elapsed_seconds', 0)
}
with self._history_lock:
self._history.append(record)
# 只保留最近 50 条记录
self._history = self._history[-50:]
def get_status(self) -> Dict:
"""获取预热状态"""
with self._items_lock:
items = list(self._items.values())
total = len(items)
success = sum(1 for i in items if i.status == ItemStatus.SUCCESS.value)
failed = sum(1 for i in items if i.status == ItemStatus.FAILED.value)
in_progress = sum(1 for i in items if i.status == ItemStatus.IN_PROGRESS.value)
pending = sum(1 for i in items if i.status == ItemStatus.PENDING.value)
return {
'enabled': self.enabled,
'is_running': self._is_running,
'total_items': total,
'success_count': success,
'failed_count': failed,
'in_progress_count': in_progress,
'pending_count': pending,
'success_rate': (success / total * 100) if total > 0 else 0,
'targets_count': len(self.targets)
}
def get_items(
self,
status: str = None,
mirror_type: str = None,
limit: int = 50
) -> List[Dict]:
"""获取预热项目列表"""
with self._items_lock:
items = [i.to_dict() for i in self._items.values()]
# 过滤
if status:
items = [i for i in items if i['status'] == status]
if mirror_type:
items = [i for i in items if i['mirror_type'] == mirror_type]
return items[-limit:]
def get_history(self, limit: int = 20) -> List[Dict]:
"""获取预热历史"""
with self._history_lock:
return list(self._history[-limit:])
def get_stats(self) -> Dict:
"""获取统计信息"""
status = self.get_status()
history = self.get_history(10)
avg_duration = 0
if history:
durations = [h['elapsed_seconds'] for h in history if 'elapsed_seconds' in h]
if durations:
avg_duration = sum(durations) / len(durations)
return {
'enabled': self.enabled,
'is_running': self._is_running,
'total_prewarmed': status['success_count'],
'total_failed': status['failed_count'],
'success_rate': status['success_rate'],
'avg_duration_seconds': round(avg_duration, 2),
'recent_runs': len(history),
'targets': [
{
'mirror_type': t.mirror_type,
'priority': t.priority,
'limit': t.limit
}
for t in self.targets
]
}
def clear_items(self, status: str = None):
"""清除预热项目"""
with self._items_lock:
if status:
self._items = {
k: v for k, v in self._items.items()
if v.status != status
}
else:
self._items = {}
def set_start_callback(self, callback: Callable):
"""设置开始回调"""
self._on_start = callback
def set_complete_callback(self, callback: Callable):
"""设置完成回调"""
self._on_complete = callback
def set_item_complete_callback(self, callback: Callable):
"""设置项目完成回调"""
self._on_item_complete = callback
def set_error_callback(self, callback: Callable):
"""设置错误回调"""
self._on_error = callback
def get_popular_items(self, mirror_type: str) -> List[str]:
"""获取指定镜像类型的流行项目列表"""
return self._popular_items.get(mirror_type, [])
def add_popular_items_to_queue(
self,
mirror_type: str,
limit: int = None,
priority: str = 'medium'
):
"""
添加流行项目到预热队列
Args:
mirror_type: 镜像类型
limit: 数量限制
priority: 优先级
"""
popular = self._popular_items.get(mirror_type, [])
if limit:
popular = popular[:limit]
self.add_items_batch(mirror_type, popular, priority)
class PrewarmScheduler:
"""预热调度器"""
def __init__(self, config: Dict = None):
self.config = config or {}
self._scheduler = None
self._running = False
def start(self, prewarmer: CachePrewarmer):
"""启动调度器"""
if self._running:
return
schedule = self.config.get('schedule', '0 3 * * *')
logger.info(f"Starting prewarm scheduler with schedule: {schedule}")
# 使用简单的时间间隔检查
self._running = True
self._scheduler_thread = threading.Thread(
target=self._run_scheduler,
args=(prewarmer,),
daemon=True
)
self._scheduler_thread.start()
def _run_scheduler(self, prewarmer: CachePrewarmer):
"""运行调度器"""
import croniter
schedule = self.config.get('schedule', '0 3 * * *')
try:
cron = croniter.croniter(schedule, datetime.now())
next_run = cron.get_next(datetime)
except Exception as e:
logger.error(f"Invalid cron schedule: {e}")
return
while self._running:
now = datetime.now()
if now >= next_run:
logger.info("Running scheduled prewarm")
try:
prewarmer.run()
except Exception as e:
logger.error(f"Scheduled prewarm failed: {e}")
cron = croniter.croniter(schedule, datetime.now())
next_run = cron.get_next(datetime)
# 检查间隔
time.sleep(60)
def stop(self):
"""停止调度器"""
self._running = False
+350
View File
@@ -0,0 +1,350 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""配置管理模块"""
import os
import sys
import json
import hashlib
import time
from typing import Dict, Any, Optional
from .utils import parse_size
def get_resource_path(relative_path: str) -> str:
"""获取打包后的资源路径"""
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)
# 打包模式
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
# 外部目录:与 exe 同级
external_path = os.path.join(os.path.dirname(sys.executable), relative_path)
if os.path.exists(external_path):
return external_path
# 打包后的资源路径(_MEIPASS)
bundled_path = os.path.join(sys._MEIPASS, relative_path)
if os.path.exists(bundled_path):
return bundled_path
return external_path
# 开发模式
return os.path.join(project_root, relative_path)
def deep_merge(base: Dict[str, Any], override: Dict[str, Any]) -> Dict[str, Any]:
"""
深度合并配置
- 如果 override 中有的键,会完全替换 base 中的值(除非都是 dict)
- 如果都是 dict,则递归合并
- 不会修改原参数
Args:
base: 默认配置(基础配置)
override: 要合并的配置(优先级更高)
Returns:
合并后的配置
"""
result = base.copy()
for key, value in override.items():
if key in result and isinstance(result[key], dict) and isinstance(value, dict):
# 两者都是字典,递归合并
result[key] = deep_merge(result[key], value)
else:
# 直接覆盖
result[key] = value
return result
def load_json_config(file_path: str) -> Optional[Dict[str, Any]]:
"""加载 JSON 配置文件"""
if not os.path.exists(file_path):
return None
try:
with open(file_path, 'r', encoding='utf-8') as f:
return json.load(f)
except json.JSONDecodeError as e:
print(f"警告: 配置文件 {file_path} JSON 格式错误: {e}")
return None
except Exception as e:
print(f"警告: 无法读取配置文件 {file_path}: {e}")
return None
class ConfigManager:
"""配置管理器"""
def __init__(self, config: Dict[str, Any] = None, settings_path: str = None):
"""
初始化配置管理器
Args:
config: 传入的配置(会覆盖默认配置)
settings_path: 默认配置文件路径(可以是位置参数或关键字参数)
"""
# 支持 settings_path 作为位置参数
if isinstance(config, str):
settings_path = config
config = None
elif settings_path is None:
# 使用打包后的资源路径
settings_path = get_resource_path('settings.json')
elif settings_path:
settings_path = settings_path
self._settings_path = settings_path
# 加载默认配置
self.default_config = self._load_default_config()
# 合并传入的配置
if config:
self.config = self._validate_config(deep_merge(self.default_config, config))
else:
self.config = self._validate_config(self.default_config.copy())
def _load_default_config(self) -> Dict[str, Any]:
"""加载默认配置文件"""
default_config = load_json_config(self._settings_path)
if default_config is None:
# 如果找不到默认配置,使用内联最小配置
default_config = {
'server_name': 'HYC下载站',
'host': '0.0.0.0',
'port': 8080,
'base_dir': './downloads',
'api_version': 'v2',
'directory_listing': True,
'enable_stats': True,
'auth_type': 'none',
'max_upload_size': 1024 * 1024 * 1024,
'timeout': 30,
'verbose': 0,
'enable_range': True,
'ignore_hidden': True,
'show_hash': False,
'calculate_hash': False,
'max_search_results': 100,
'enable_ws': True,
'enable_sse': True,
'enable_monitor': True,
'monitor_interval': 5,
'enable_sync': True,
'enable_mirrors': True,
'database': {
'enabled': True,
'type': 'sqlite',
'sqlite': {'path': './data/hyc.db'}
},
'mirrors': {
'docker': {'enabled': True},
'apt': {'enabled': True},
'yum': {'enabled': True},
'pypi': {'enabled': True},
'npm': {'enabled': True},
'go': {'enabled': True}
},
'sync_sources': {},
'webhooks': {'enabled': False, 'storage': 'webhooks.json'},
'auth_sessions_file': 'auth_sessions.json',
'auth_session_timeout': 3600,
'auth_cookie_max_age': 86400
}
print(f"警告: 未找到默认配置文件 ({self._settings_path}),使用内联默认配置")
return default_config
@classmethod
def from_settings(cls, custom_config: Dict[str, Any] = None, settings_path: str = None) -> 'ConfigManager':
"""
从默认配置创建配置管理器
Args:
custom_config: 自定义配置,会覆盖默认配置
settings_path: 默认配置文件路径
Returns:
ConfigManager 实例
"""
return cls(config=custom_config, settings_path=settings_path)
def _validate_config(self, config: Dict[str, Any]) -> Dict[str, Any]:
"""验证和修复配置"""
# 确保必要配置存在
required = ['base_dir', 'host', 'port']
for key in required:
if key not in config:
raise ValueError(f"缺少必要配置: {key}")
# 修复路径配置
config['base_dir'] = os.path.abspath(config['base_dir'])
# 设置默认值(不在 _validate_config 中处理,由默认配置提供)
# 验证认证配置
auth_type = config.get('auth_type', 'none')
if auth_type == 'basic':
if 'auth_user' not in config:
config['auth_user'] = 'admin'
if 'auth_pass' not in config:
config['auth_pass'] = 'admin123'
elif auth_type == 'token':
# 每次运行都重新生成标准的 token
import secrets
config['auth_token'] = secrets.token_hex(32)
# 验证上传大小配置
if 'max_upload_size' in config:
try:
if isinstance(config['max_upload_size'], str):
config['max_upload_size'] = parse_size(config['max_upload_size'])
except ValueError as e:
print(f"警告: 无效的上传大小配置: {e}")
config['max_upload_size'] = 1024 * 1024 * 1024
# 验证端口范围
if 'port' in config:
port = config['port']
if not (1 <= port <= 65535):
raise ValueError(f"无效的端口号: {port}")
# 验证并创建必要目录
base_dir = config['base_dir']
try:
if not os.path.exists(base_dir):
os.makedirs(base_dir, exist_ok=True)
# 测试写入权限
test_file = os.path.join(base_dir, '.write_test')
with open(test_file, 'w') as f:
f.write('test')
os.remove(test_file)
except Exception as e:
raise ValueError(f"基础目录无法访问: {e}")
# 获取项目根目录(脚本所在目录)
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
# 在项目根目录创建必要的数据目录
necessary_dirs = [
os.path.join(project_root, 'data'),
os.path.join(project_root, 'logs'),
]
for dir_path in necessary_dirs:
if not os.path.exists(dir_path):
try:
os.makedirs(dir_path, exist_ok=True)
except Exception as e:
print(f"警告: 无法创建目录 {dir_path}: {e}")
# 更新配置指向项目根目录
config['data_dir'] = os.path.join(project_root, 'data')
config['logs_dir'] = os.path.join(project_root, 'logs')
return config
def get(self, key: str, default: Any = None) -> Any:
"""获取配置项"""
return self.config.get(key, default)
def update(self, updates: Dict[str, Any]):
"""更新配置"""
self.config.update(updates)
self.config = self._validate_config(self.config)
def get_full_config(self) -> Dict[str, Any]:
"""获取完整配置字典"""
return self.config.copy()
def to_dict(self) -> Dict[str, Any]:
"""返回配置字典(不包含敏感信息)"""
safe_config = {
"server_name": self.config.get("server_name", "Mirror Server"),
"version": "2.2",
"base_dir": self.config['base_dir'],
"directory_listing": self.config.get('directory_listing', True),
"max_upload_size": self.config.get('max_upload_size'),
"enable_stats": self.config.get('enable_stats', True),
"auth_type": self.config.get('auth_type', 'none'),
"sort_by": self.config.get('sort_by', 'name'),
"sort_reverse": self.config.get('sort_reverse', False),
"ignore_hidden": self.config.get('ignore_hidden', True),
"enable_range": self.config.get('enable_range', True),
"show_hash": self.config.get('show_hash', False),
"calculate_hash": self.config.get('calculate_hash', False),
"max_search_results": self.config.get('max_search_results', 100),
"api_version": self.config.get('api_version', 'v1'),
"verbose": self.config.get('verbose', 0)
}
return safe_config
def load_config_file(config_path: str) -> Dict[str, Any]:
"""加载配置文件"""
if not os.path.exists(config_path):
return {}
try:
with open(config_path, 'r', encoding='utf-8') as f:
return json.load(f)
except Exception as e:
print(f"错误: 无法加载配置文件 {config_path}: {e}")
return {}
def load_settings_with_override(settings_path: str, override_path: str = None) -> Dict[str, Any]:
"""
加载默认配置并合并覆盖配置
Args:
settings_path: 默认配置文件路径
override_path: 覆盖配置文件路径(可选)
Returns:
合并后的完整配置
"""
# 加载默认配置
default_config = load_json_config(settings_path) or {}
# 加载覆盖配置
override_config = {}
if override_path:
override_config = load_json_config(override_path) or {}
# 深度合并
return deep_merge(default_config, override_config)
def save_config_file(config_path: str, config: Dict[str, Any]) -> bool:
"""
保存配置文件
Args:
config_path: 保存路径
config: 配置字典
Returns:
是否保存成功
"""
try:
# 创建目录
os.makedirs(os.path.dirname(config_path), exist_ok=True)
with open(config_path, 'w', encoding='utf-8') as f:
json.dump(config, f, ensure_ascii=False, indent=4)
return True
except Exception as e:
print(f"错误: 无法保存配置文件 {config_path}: {e}")
return False
+340
View File
@@ -0,0 +1,340 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
配置热更新模块
支持不重启服务的情况下重新加载配置
"""
import os
import sys
import json
import time
import threading
import logging
from typing import Dict, Any, Optional, Callable
from datetime import datetime
logger = logging.getLogger(__name__)
class ConfigHotReloader:
"""配置热重载管理器"""
def __init__(self, config_path: str, callback: Callable = None):
"""
初始化热重载管理器
Args:
config_path: 配置文件路径
callback: 配置变更时的回调函数,接收 (config, change_type) 参数
"""
self.config_path = config_path
self.callback = callback
self._config: Dict[str, Any] = {}
self._last_modified: float = 0
self._last_load_time: Optional[datetime] = None
self._lock = threading.RLock()
# 配置变更历史
self._change_history: list = []
# 监听器
self._listeners: Dict[str, list] = {
'on_change': [],
'on_error': []
}
# 加载初始配置
self.reload()
def reload(self, silent: bool = False) -> bool:
"""
重新加载配置
Args:
silent: 静默模式,不触发变更通知
Returns:
是否加载成功
"""
try:
if not os.path.exists(self.config_path):
if not silent:
logger.warning(f"配置文件不存在: {self.config_path}")
return False
# 获取文件修改时间
current_mtime = os.path.getmtime(self.config_path)
# 检查是否有变化
if current_mtime == self._last_modified and not silent:
return True
# 加载配置
with open(self.config_path, 'r', encoding='utf-8') as f:
new_config = json.load(f)
# 计算变更
changes = self._compute_changes(self._config, new_config)
with self._lock:
old_config = self._config.copy()
self._config = new_config
self._last_modified = current_mtime
self._last_load_time = datetime.now()
# 记录变更
if changes:
change_record = {
'timestamp': self._last_load_time.isoformat(),
'changes': changes,
'old_config_keys': list(old_config.keys()),
'new_config_keys': list(new_config.keys())
}
self._change_history.append(change_record)
# 保持历史记录在合理范围内
if len(self._change_history) > 100:
self._change_history = self._change_history[-50:]
if not silent and changes:
self._notify_change(changes)
logger.info(f"配置已重新加载: {self.config_path}")
return True
except json.JSONDecodeError as e:
error_msg = f"配置 JSON 格式错误: {e}"
logger.error(error_msg)
self._notify_error(error_msg)
return False
except Exception as e:
error_msg = f"加载配置失败: {e}"
logger.error(error_msg)
self._notify_error(error_msg)
return False
def _compute_changes(self, old: Dict, new: Dict) -> Dict:
"""计算配置变更"""
changes = {
'added': [],
'removed': [],
'modified': []
}
old_keys = set(old.keys())
new_keys = set(new.keys())
# 新增的键
for key in new_keys - old_keys:
changes['added'].append(key)
# 移除的键
for key in old_keys - new_keys:
changes['removed'].append(key)
# 修改的键
for key in old_keys & new_keys:
if old[key] != new[key]:
# 检查是否是嵌套字典
if isinstance(old[key], dict) and isinstance(new[key], dict):
nested = self._compute_nested_changes(old[key], new[key], f"{key}.")
if nested['added'] or nested['removed'] or nested['modified']:
changes['modified'].append({
'key': key,
'type': 'nested',
'changes': nested
})
else:
changes['modified'].append({
'key': key,
'type': 'value',
'old_value': old[key],
'new_value': new[key]
})
return changes
def _compute_nested_changes(self, old: Dict, new: Dict, prefix: str = "") -> Dict:
"""计算嵌套字典的变更"""
changes = {
'added': [],
'removed': [],
'modified': []
}
old_keys = set(old.keys())
new_keys = set(new.keys())
for key in new_keys - old_keys:
changes['added'].append(f"{prefix}{key}")
for key in old_keys - new_keys:
changes['removed'].append(f"{prefix}{key}")
for key in old_keys & new_keys:
if old[key] != new[key]:
changes['modified'].append(f"{prefix}{key}")
return changes
def _notify_change(self, changes: Dict):
"""通知配置变更"""
for listener in self._listeners['on_change']:
try:
if callable(listener):
listener(self._config, changes)
except Exception as e:
logger.error(f"配置变更监听器执行失败: {e}")
if self.callback:
try:
self.callback(self._config, changes)
except Exception as e:
logger.error(f"配置回调函数执行失败: {e}")
def _notify_error(self, error: str):
"""通知错误"""
for listener in self._listeners['on_error']:
try:
if callable(listener):
listener(error)
except Exception as e:
logger.error(f"错误监听器执行失败: {e}")
def add_change_listener(self, callback: Callable):
"""添加配置变更监听器"""
self._listeners['on_change'].append(callback)
def add_error_listener(self, callback: Callable):
"""添加错误监听器"""
self._listeners['on_error'].append(callback)
def get(self, key: str, default: Any = None) -> Any:
"""获取配置值"""
with self._lock:
return self._config.get(key, default)
def get_all(self) -> Dict:
"""获取完整配置"""
with self._lock:
return self._config.copy()
def set(self, key: str, value: Any, save: bool = True) -> bool:
"""设置配置值(仅内存中)"""
with self._lock:
self._config[key] = value
if save:
return self.save()
return True
def save(self, path: str = None) -> bool:
"""保存配置到文件"""
save_path = path or self.config_path
try:
with open(save_path, 'w', encoding='utf-8') as f:
json.dump(self._config, f, ensure_ascii=False, indent=4)
self._last_modified = os.path.getmtime(save_path)
return True
except Exception as e:
logger.error(f"保存配置失败: {e}")
return False
def get_change_history(self, limit: int = 10) -> list:
"""获取配置变更历史"""
return self._change_history[-limit:]
def watch(self, interval: float = 5.0):
"""
启动后台监控线程
Args:
interval: 检查间隔(秒)
"""
def _watch_loop():
while True:
try:
self.reload()
except Exception as e:
logger.error(f"配置监控错误: {e}")
time.sleep(interval)
thread = threading.Thread(target=_watch_loop, daemon=True)
thread.start()
logger.info(f"配置热监控已启动,间隔: {interval}秒")
class ConfigManager:
"""配置管理器 - 支持热更新"""
def __init__(self, config: Dict[str, Any] = None):
self.config = config or {}
self._hot_reloader: Optional[ConfigHotReloader] = None
def load_from_file(self, path: str, enable_watch: bool = False) -> bool:
"""从文件加载配置"""
if not os.path.exists(path):
return False
try:
with open(path, 'r', encoding='utf-8') as f:
self.config = json.load(f)
if enable_watch:
self._hot_reloader = ConfigHotReloader(path)
self._hot_reloader.watch()
return True
except Exception as e:
logger.error(f"加载配置失败: {e}")
return False
def hot_reload(self, path: str = None) -> bool:
"""触发热重载"""
if self._hot_reloader:
return self._hot_reloader.reload()
return False
def get(self, key: str, default: Any = None) -> Any:
"""获取配置值"""
return self.config.get(key, default)
def set(self, key: str, value: Any, persist: bool = False, path: str = None) -> bool:
"""设置配置值"""
keys = key.split('.')
current = self.config
for k in keys[:-1]:
if k not in current:
current[k] = {}
current = current[k]
current[keys[-1]] = value
if persist:
if self._hot_reloader:
return self._hot_reloader.save(path)
elif path:
try:
with open(path, 'w', encoding='utf-8') as f:
json.dump(self.config, f, ensure_ascii=False, indent=4)
return True
except Exception as e:
logger.error(f"保存配置失败: {e}")
return False
return True
def add_change_listener(self, callback: Callable):
"""添加变更监听器"""
if self._hot_reloader:
self._hot_reloader.add_change_listener(callback)
def get_all(self) -> Dict:
"""获取完整配置"""
return self.config.copy()
+1545
View File
File diff suppressed because it is too large Load Diff
+603
View File
@@ -0,0 +1,603 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
平滑重启模块
支持优雅停止、零停机重启、滚动更新
"""
import os
import sys
import signal
import time
import threading
import logging
import subprocess
from typing import Dict, List, Optional, Callable
from dataclasses import dataclass, field
from datetime import datetime
from enum import Enum
logger = logging.getLogger(__name__)
class RestartStrategy(Enum):
"""重启策略"""
GRACEFUL = "graceful" # 优雅停止,等待请求完成
ROLLING = "rolling" # 滚动更新,零停机
IMMEDIATE = "immediate" # 立即重启
class ServerState(Enum):
"""服务器状态"""
RUNNING = "running"
STOPPING = "stopping"
STOPPED = "stopped"
STARTING = "starting"
RESTARTING = "restarting"
@dataclass
class PendingRequest:
"""待处理的请求"""
request_id: str
start_time: float
endpoint: str
method: str
client_address: tuple
class GracefulRestartManager:
"""平滑重启管理器"""
def __init__(self, config: Dict = None):
"""
初始化平滑重启管理器
Args:
config: 重启配置
"""
self.config = config or {}
self.state = ServerState.STOPPED
# 等待超时时间
self.graceful_timeout = self.config.get('graceful_timeout', 30) # 秒
self.shutdown_timeout = self.config.get('shutdown_timeout', 10) # 秒
# 待处理请求追踪
self._pending_requests: Dict[str, PendingRequest] = {}
self._requests_lock = threading.Lock()
# 请求计数器
self._request_counter = 0
self._counter_lock = threading.Lock()
# 回调函数
self._on_prepare_restart: Optional[Callable] = None
self._on_start_restart: Optional[Callable] = None
self._on_complete_restart: Optional[Callable] = None
self._on_restart_failed: Optional[Callable] = None
# 状态锁
self._state_lock = threading.Lock()
# 重启历史
self._restart_history: List[Dict] = []
self._history_lock = threading.Lock()
# 信号处理器
self._setup_signal_handlers()
def _setup_signal_handlers(self):
"""设置信号处理器"""
signal.signal(signal.SIGTERM, self._handle_signal)
signal.signal(signal.SIGINT, self._handle_signal)
signal.signal(signal.SIGHUP, self._handle_signal)
def _handle_signal(self, signum, frame):
"""处理退出信号"""
signal_names = {
signal.SIGTERM: 'SIGTERM',
signal.SIGINT: 'SIGINT',
signal.SIGHUP: 'SIGHUP'
}
sig_name = signal_names.get(signum, str(signum))
logger.info(f"收到信号 {sig_name},准备优雅关闭")
# 设置状态为停止中
self.state = ServerState.STOPPING
def register_request(self, endpoint: str, method: str, client: tuple) -> str:
"""
注册一个待处理的请求
Args:
endpoint: 请求端点
method: 请求方法
client: 客户端地址
Returns:
请求 ID
"""
with self._counter_lock:
self._request_counter += 1
request_id = f"req_{self._request_counter}_{int(time.time())}"
request = PendingRequest(
request_id=request_id,
start_time=time.time(),
endpoint=endpoint,
method=method,
client_address=client
)
with self._requests_lock:
self._pending_requests[request_id] = request
return request_id
def complete_request(self, request_id: str):
"""标记请求完成"""
with self._requests_lock:
self._pending_requests.pop(request_id, None)
def get_pending_count(self) -> int:
"""获取待处理请求数量"""
with self._requests_lock:
return len(self._pending_requests)
def get_pending_requests(self) -> List[Dict]:
"""获取待处理请求详情"""
with self._requests_lock:
return [
{
'request_id': r.request_id,
'endpoint': r.endpoint,
'method': r.method,
'duration': round(time.time() - r.start_time, 2),
'client': str(r.client_address[0]) if r.client_address else 'unknown'
}
for r in self._pending_requests.values()
]
def is_safe_to_restart(self) -> bool:
"""检查是否安全重启(没有待处理请求)"""
return self.get_pending_count() == 0
def wait_for_requests(
self,
timeout: int = None,
progress_callback: Callable[[int, int], None] = None
) -> bool:
"""
等待所有请求完成
Args:
timeout: 超时时间(秒)
progress_callback: 进度回调函数(remaining, total)
Returns:
是否在超时前完成所有请求
"""
timeout = timeout or self.graceful_timeout
start_time = time.time()
while time.time() - start_time < timeout:
pending = self.get_pending_count()
if progress_callback:
progress_callback(pending, 0)
if pending == 0:
return True
time.sleep(0.5)
return False
def set_prepare_restart_callback(self, callback: Callable):
"""设置准备重启回调"""
self._on_prepare_restart = callback
def set_start_restart_callback(self, callback: Callable):
"""设置开始重启回调"""
self._on_start_restart = callback
def set_complete_restart_callback(self, callback: Callable):
"""设置完成重启回调"""
self._on_complete_restart = callback
def set_restart_failed_callback(self, callback: Callable):
"""设置重启失败回调"""
self._on_restart_failed = callback
def prepare_restart(self) -> Dict:
"""
准备重启(通知各模块准备)
Returns:
准备结果
"""
with self._state_lock:
self.state = ServerState.STOPPING
result = {
'success': True,
'pending_requests': self.get_pending_count(),
'message': ''
}
# 通知回调
if self._on_prepare_restart:
try:
self._on_prepare_restart()
except Exception as e:
logger.error(f"Prepare restart callback failed: {e}")
result['success'] = False
result['message'] = str(e)
return result
def perform_restart(
self,
strategy: RestartStrategy = RestartStrategy.GRACEFUL,
new_config: Dict = None,
script_path: str = None
) -> Dict:
"""
执行重启
Args:
strategy: 重启策略
new_config: 新配置(用于配置热更新)
script_path: 服务器脚本路径
Returns:
重启结果
"""
with self._state_lock:
if self.state == ServerState.RESTARTING:
return {
'success': False,
'error': 'Restart already in progress'
}
self.state = ServerState.RESTARTING
start_time = time.time()
result = {
'success': False,
'strategy': strategy.value,
'elapsed_seconds': 0,
'message': ''
}
try:
# 准备阶段
prepare_result = self.prepare_restart()
if not prepare_result['success']:
result['message'] = f"Prepare failed: {prepare_result['message']}"
self.state = ServerState.RUNNING
return result
# 根据策略执行重启
if strategy == RestartStrategy.GRACEFUL:
result = self._graceful_restart(prepare_result)
elif strategy == RestartStrategy.ROLLING:
result = self._rolling_restart(script_path, new_config)
elif strategy == RestartStrategy.IMMEDIATE:
result = self._immediate_restart(script_path)
result['elapsed_seconds'] = round(time.time() - start_time, 2)
# 记录重启历史
self._add_to_history(result)
except Exception as e:
logger.error(f"Restart failed: {e}")
result['success'] = False
result['message'] = str(e)
self.state = ServerState.RUNNING
if self._on_restart_failed:
try:
self._on_restart_failed(str(e))
except Exception:
pass
return result
def _graceful_restart(self, prepare_result: Dict) -> Dict:
"""优雅重启(等待请求完成)"""
result = {
'success': True,
'strategy': 'graceful',
'pending_requests': prepare_result['pending_requests'],
'message': ''
}
# 等待请求完成
pending = prepare_result['pending_requests']
if pending > 0:
logger.info(f"等待 {pending} 个请求完成,超时 {self.graceful_timeout} 秒")
def progress_callback(remaining, total):
if remaining % 5 == 0:
logger.info(f"还有 {remaining} 个请求待处理")
success = self.wait_for_requests(
timeout=self.graceful_timeout,
progress_callback=progress_callback
)
if not success:
remaining = self.get_pending_count()
result['success'] = False
result['message'] = f"{remaining} 个请求未在 {self.graceful_timeout} 秒内完成"
logger.warning(result['message'])
self.state = ServerState.RUNNING
return result
result['message'] = '所有请求已完成'
logger.info('所有请求已完成')
# 通知开始重启
if self._on_start_restart:
try:
self._on_start_restart()
except Exception as e:
logger.error(f"Start restart callback failed: {e}")
result['success'] = True
result['message'] = 'Graceful restart prepared'
self.state = ServerState.STOPPED
return result
def _rolling_restart(self, script_path: str, new_config: Dict = None) -> Dict:
"""滚动重启(零停机)"""
result = {
'success': True,
'strategy': 'rolling',
'message': ''
}
if not script_path:
script_path = os.path.join(os.path.dirname(__file__), '..', 'main.py')
script_path = os.path.abspath(script_path)
# 检查是否有新配置需要应用
if new_config:
result['config_updated'] = True
logger.info("配置将在重启后应用")
else:
result['config_updated'] = False
# 通知开始重启
if self._on_start_restart:
try:
self._on_start_restart()
except Exception as e:
logger.error(f"Start restart callback failed: {e}")
# 启动新进程(在同一端口,但使用不同进程ID)
# 注意:实际实现需要在负载均衡器层面处理
logger.info("滚动重启:建议在负载均衡器层面处理零停机")
self.state = ServerState.STOPPED
result['message'] = 'Rolling restart prepared - use load balancer for zero-downtime'
return result
def _immediate_restart(self, script_path: str = None) -> Dict:
"""立即重启"""
result = {
'success': True,
'strategy': 'immediate',
'message': ''
}
if not script_path:
script_path = os.path.join(os.path.dirname(__file__), '..', 'main.py')
script_path = os.path.abspath(script_path)
# 通知开始重启
if self._on_start_restart:
try:
self._on_start_restart()
except Exception as e:
logger.error(f"Start restart callback failed: {e}")
# 发送重启信号给主进程
logger.info("立即重启服务器")
# 在子进程中重启
try:
# 启动新进程
cmd = [sys.executable, script_path]
env = os.environ.copy()
env['HYC_RESTARTED'] = '1'
subprocess.Popen(cmd, env=env)
self.state = ServerState.STOPPED
result['message'] = 'Immediate restart initiated'
except Exception as e:
result['success'] = False
result['message'] = f"Failed to restart: {str(e)}"
self.state = ServerState.RUNNING
return result
def _add_to_history(self, result: Dict):
"""添加重启历史记录"""
record = {
'timestamp': datetime.now().isoformat(),
'strategy': result.get('strategy', 'unknown'),
'success': result.get('success', False),
'elapsed_seconds': result.get('elapsed_seconds', 0),
'message': result.get('message', '')
}
with self._history_lock:
self._restart_history.append(record)
# 只保留最近 20 条记录
self._restart_history = self._restart_history[-20:]
def get_restart_history(self) -> List[Dict]:
"""获取重启历史"""
with self._history_lock:
return list(self._restart_history)
def get_stats(self) -> Dict:
"""获取统计信息"""
return {
'state': self.state.value,
'pending_requests': self.get_pending_count(),
'graceful_timeout': self.graceful_timeout,
'shutdown_timeout': self.shutdown_timeout,
'restart_count': len(self._restart_history),
'recent_restarts': self.get_restart_history()[-5:]
}
def update_config(self, new_config: Dict):
"""更新配置(热更新)"""
if 'graceful_timeout' in new_config:
self.graceful_timeout = new_config['graceful_timeout']
if 'shutdown_timeout' in new_config:
self.shutdown_timeout = new_config['shutdown_timeout']
logger.info(f"重启配置已更新: timeout={self.graceful_timeout}s")
class ServerHealthChecker:
"""服务器健康检查器"""
def __init__(self, host: str = 'localhost', port: int = 8080):
self.host = host
self.port = port
self._last_check = None
self._is_healthy = False
def check(self) -> Dict:
"""
检查服务器健康状态
Returns:
健康检查结果
"""
import socket
result = {
'healthy': False,
'latency_ms': 0,
'error': None,
'timestamp': datetime.now().isoformat()
}
try:
start_time = time.time()
# 尝试连接
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.settimeout(5)
sock.connect((self.host, self.port))
sock.close()
result['latency_ms'] = round((time.time() - start_time) * 1000, 2)
result['healthy'] = True
except Exception as e:
result['error'] = str(e)
self._last_check = result
return result
def get_latency(self) -> float:
"""获取延迟(毫秒)"""
if self._last_check:
return self._last_check.get('latency_ms', 0)
return 0
class RollingRestartManager:
"""滚动重启管理器(支持多实例)"""
def __init__(self, config: Dict = None):
self.config = config or {}
self.instances: Dict[str, Dict] = {} # instance_id -> info
self.current_instance_id = None
self._lock = threading.Lock()
def register_instance(self, instance_id: str, info: Dict):
"""注册实例"""
with self._lock:
self.instances[instance_id] = {
**info,
'registered_at': datetime.now().isoformat(),
'status': 'active'
}
def unregister_instance(self, instance_id: str):
"""注销实例"""
with self._lock:
if instance_id in self.instances:
self.instances[instance_id]['status'] = 'draining'
def get_active_instances(self) -> List[str]:
"""获取活跃实例列表"""
with self._lock:
return [
i for i, info in self.instances.items()
if info['status'] == 'active'
]
def perform_rolling_restart(
self,
instance_id: str,
restart_func: Callable
) -> Dict:
"""
对单个实例执行滚动重启
Args:
instance_id: 实例 ID
restart_func: 重启函数
Returns:
重启结果
"""
with self._lock:
if instance_id not in self.instances:
return {
'success': False,
'error': f'Instance {instance_id} not found'
}
# 标记为排水中
self.instances[instance_id]['status'] = 'draining'
# 等待连接耗尽
time.sleep(5)
# 执行重启
try:
restart_func(instance_id)
self.instances[instance_id]['status'] = 'active'
self.instances[instance_id]['restarted_at'] = datetime.now().isoformat()
return {
'success': True,
'instance_id': instance_id
}
except Exception as e:
self.instances[instance_id]['status'] = 'error'
return {
'success': False,
'instance_id': instance_id,
'error': str(e)
}
+426
View File
@@ -0,0 +1,426 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
镜像源健康检查模块
自动检测上游镜像源的可用性,支持故障切换
"""
import os
import sys
import json
import time
import threading
import logging
import urllib.request
import urllib.error
from datetime import datetime, timedelta
from typing import Dict, List, Optional, Callable
from enum import Enum
from dataclasses import dataclass, field
logger = logging.getLogger(__name__)
class HealthStatus(Enum):
"""健康状态"""
UNKNOWN = "unknown"
HEALTHY = "healthy"
DEGRADED = "degraded"
UNHEALTHY = "unhealthy"
@dataclass
class HealthCheckResult:
"""健康检查结果"""
source_name: str
status: HealthStatus
response_time: float # 毫秒
http_status: Optional[int] = None
error_message: Optional[str] = None
last_check: Optional[datetime] = None
consecutive_failures: int = 0
total_checks: int = 0
success_rate: float = 100.0
details: Dict = field(default_factory=dict)
class HealthChecker:
"""健康检查器"""
def __init__(self, config: dict = None):
"""
初始化健康检查器
Args:
config: 健康检查配置
"""
self.config = config or {}
self.default_timeout = self.config.get('timeout', 10) # 秒
self.default_interval = self.config.get('interval', 60) # 秒
self.max_retries = self.config.get('max_retries', 3)
self.failure_threshold = self.config.get('failure_threshold', 3) # 连续失败次数阈值
# 检查结果存储
self._results: Dict[str, HealthCheckResult] = {}
self._lock = threading.Lock()
# 状态变化回调
self._on_status_change: Optional[Callable] = None
def set_status_change_callback(self, callback: Callable):
"""设置状态变化回调"""
self._on_status_change = callback
def check_source(self, source_name: str, source_config: dict) -> HealthCheckResult:
"""
检查单个镜像源的健康状态
Args:
source_name: 镜像源名称
source_config: 镜像源配置
Returns:
HealthCheckResult: 检查结果
"""
url = source_config.get('url', '')
if not url:
return HealthCheckResult(
source_name=source_name,
status=HealthStatus.UNHEALTHY,
response_time=0,
error_message="No URL configured",
last_check=datetime.now()
)
start_time = time.time()
http_status = None
error_message = None
details = {}
try:
# 发送 HEAD 请求(更轻量)
req = urllib.request.Request(
url.rstrip('/') + '/',
method='HEAD',
headers={
'User-Agent': 'HYC-Mirror-HealthCheck/1.0',
'Accept': '*/*'
}
)
with urllib.request.urlopen(req, timeout=self.default_timeout) as response:
http_status = response.status
details['headers'] = dict(response.headers)
except urllib.error.HTTPError as e:
http_status = e.code
error_message = f"HTTP {e.code}"
except urllib.error.URLError as e:
error_message = f"Connection error: {str(e.reason)}"
except Exception as e:
error_message = str(e)
response_time = (time.time() - start_time) * 1000 # 转换为毫秒
# 判断健康状态
if error_message:
if http_status and 400 <= http_status < 500:
status = HealthStatus.DEGRADED # 客户端错误,可能暂时
else:
status = HealthStatus.UNHEALTHY
elif http_status and 200 <= http_status < 400:
status = HealthStatus.HEALTHY
elif http_status:
status = HealthStatus.DEGRADED
else:
status = HealthStatus.UNHEALTHY
# 计算统计数据
with self._lock:
if source_name not in self._results:
self._results[source_name] = HealthCheckResult(
source_name=source_name,
status=status,
response_time=response_time,
last_check=datetime.now(),
consecutive_failures=0,
total_checks=1
)
else:
old_result = self._results[source_name]
consecutive_failures = old_result.consecutive_failures + (1 if status == HealthStatus.UNHEALTHY else 0)
total_checks = old_result.total_checks + 1
success_rate = ((total_checks - consecutive_failures) / total_checks) * 100
self._results[source_name] = HealthCheckResult(
source_name=source_name,
status=status,
response_time=response_time,
http_status=http_status,
error_message=error_message,
last_check=datetime.now(),
consecutive_failures=consecutive_failures,
total_checks=total_checks,
success_rate=success_rate,
details=details
)
# 检查状态变化,触发回调
if self._on_status_change and source_name in self._results:
old_status = self._results[source_name].status
if old_status != status:
try:
self._on_status_change(source_name, old_status, status)
except Exception as e:
logger.error(f"状态变化回调执行失败: {e}")
return self._results[source_name]
def get_all_results(self) -> List[HealthCheckResult]:
"""获取所有检查结果"""
with self._lock:
return list(self._results.values())
def get_result(self, source_name: str) -> Optional[HealthCheckResult]:
"""获取指定源的结果"""
with self._lock:
return self._results.get(source_name)
def is_healthy(self, source_name: str) -> bool:
"""检查源是否健康"""
result = self.get_result(source_name)
if result is None:
return True # 未检查过的默认健康
return result.status == HealthStatus.HEALTHY
def get_unhealthy_sources(self) -> List[str]:
"""获取不健康的源列表"""
with self._lock:
return [name for name, result in self._results.items()
if result.status == HealthStatus.UNHEALTHY]
def get_stats(self) -> dict:
"""获取健康检查统计"""
with self._lock:
total = len(self._results)
healthy = sum(1 for r in self._results.values() if r.status == HealthStatus.HEALTHY)
degraded = sum(1 for r in self._results.values() if r.status == HealthStatus.DEGRADED)
unhealthy = sum(1 for r in self._results.values() if r.status == HealthStatus.UNHEALTHY)
avg_response_time = 0
if total > 0:
avg_response_time = sum(r.response_time for r in self._results.values()) / total
return {
'total_sources': total,
'healthy': healthy,
'degraded': degraded,
'unhealthy': unhealthy,
'avg_response_time_ms': round(avg_response_time, 2),
'timestamp': datetime.now().isoformat()
}
class MirrorFailoverManager:
"""镜像源故障切换管理器"""
def __init__(self, config: dict = None):
"""
初始化故障切换管理器
Args:
config: 配置,包含镜像源列表
"""
self.config = config or {}
self.mirrors: Dict[str, Dict] = self.config.get('mirrors', {})
# 启用故障切换
self.failover_enabled = self.config.get('failover_enabled', True)
self.failover_threshold = self.config.get('failover_threshold', 3) # 连续失败次数
# 健康检查器
self.health_checker = HealthChecker(self.config.get('health_check', {}))
# 当前活跃源
self._active_source: Dict[str, str] = {} # mirror_type -> source_name
self._source_priority: Dict[str, List[str]] = {} # mirror_type -> [优先列表]
# 故障切换历史
self._failover_history: List[dict] = []
# 回调
self._on_failover: Optional[Callable] = None
def set_failover_callback(self, callback: Callable):
"""设置故障切换回调"""
self._on_failover = callback
def initialize(self):
"""初始化,确定各镜像类型的首选源"""
for mirror_type, mirror_config in self.mirrors.items():
if not isinstance(mirror_config, dict):
continue
sources = mirror_config.get('sources', [])
if sources:
# 使用配置的优先列表
self._source_priority[mirror_type] = sources
else:
# 使用内置的默认优先列表
self._source_priority[mirror_type] = self._get_default_priority(mirror_type)
# 选择首选源
if self._source_priority[mirror_type]:
self._active_source[mirror_type] = self._source_priority[mirror_type][0]
def _get_default_priority(self, mirror_type: str) -> List[str]:
"""获取默认的镜像源优先列表"""
priorities = {
'docker': ['docker.io', 'docker.mirrors.aliyun.com', 'dockerhub.azk8s.cn'],
'apt': ['archive.ubuntu.com', 'mirrors.aliyun.com', 'security.ubuntu.com'],
'yum': ['mirror.centos.org', 'mirrors.aliyun.com'],
'pypi': ['pypi.org', 'pypi.mirrors.aliyun.com'],
'npm': ['registry.npmjs.org', 'registry.npmmirror.com'],
'go': ['proxy.golang.org', 'goproxy.cn']
}
return priorities.get(mirror_type, [])
def check_all(self) -> Dict[str, HealthCheckResult]:
"""检查所有镜像源"""
results = {}
for mirror_type, mirror_config in self.mirrors.items():
if not isinstance(mirror_config, dict):
continue
sources = mirror_config.get('sources', [])
for source_name in sources:
if source_name not in results:
result = self.health_checker.check_source(source_name, {'url': self._get_source_url(mirror_type, source_name)})
results[source_name] = result
return results
def _get_source_url(self, mirror_type: str, source_name: str) -> str:
"""获取源 URL"""
# 从配置中获取
sources_config = self.mirrors.get(mirror_type, {}).get('sources_config', {})
if source_name in sources_config:
return sources_config[source_name].get('url', '')
# 从 URL 模板生成
url_template = self.mirrors.get(mirror_type, {}).get('url_template', '')
if url_template and '{mirror}' in url_template:
return url_template.replace('{mirror}', source_name)
return ''
def get_active_source(self, mirror_type: str) -> Optional[str]:
"""获取当前活跃的镜像源"""
return self._active_source.get(mirror_type)
def get_source_for_request(self, mirror_type: str, original_url: str) -> str:
"""
根据故障切换策略获取请求的源 URL
Args:
mirror_type: 镜像类型
original_url: 原始 URL
Returns:
str: 实际请求的 URL
"""
if not self.failover_enabled:
return original_url
active_source = self._active_source.get(mirror_type)
if not active_source:
return original_url
source_url = self._get_source_url(mirror_type, active_source)
if not source_url:
return original_url
# 替换 URL 中的主机部分
try:
from urllib.parse import urlparse
parsed = urlparse(original_url)
# 构建新 URL
new_url = f"{parsed.scheme}://{source_url}{parsed.path}"
if parsed.query:
new_url += f"?{parsed.query}"
return new_url
except Exception:
return original_url
def perform_failover(self, mirror_type: str) -> bool:
"""
对指定镜像类型执行故障切换
Args:
mirror_type: 镜像类型
Returns:
bool: 是否成功切换
"""
priority_list = self._source_priority.get(mirror_type, [])
if not priority_list:
return False
current_index = 0
if mirror_type in self._active_source:
try:
current_index = priority_list.index(self._active_source[mirror_type])
except ValueError:
pass
# 查找下一个健康的源
old_source = self._active_source.get(mirror_type)
for i in range(current_index + 1, len(priority_list)):
source_name = priority_list[i]
result = self.health_checker.get_result(source_name)
if result and result.status == HealthStatus.HEALTHY:
self._active_source[mirror_type] = source_name
# 记录故障切换
failover_record = {
'timestamp': datetime.now().isoformat(),
'mirror_type': mirror_type,
'old_source': old_source,
'new_source': source_name,
'reason': 'Health check failed'
}
self._failover_history.append(failover_record)
# 保持历史记录在合理范围内
if len(self._failover_history) > 100:
self._failover_history = self._failover_history[-50:]
# 触发回调
if self._on_failover:
try:
self._on_failover(mirror_type, old_source, source_name)
except Exception as e:
logger.error(f"故障切换回调执行失败: {e}")
return True
return False
def get_failover_history(self, limit: int = 10) -> List[dict]:
"""获取故障切换历史"""
return self._failover_history[-limit:]
def get_health_summary(self) -> dict:
"""获取健康状态摘要"""
return {
'failover_enabled': self.failover_enabled,
'health': self.health_checker.get_stats(),
'active_sources': self._active_source.copy(),
'failover_history_count': len(self._failover_history)
}
+1258
View File
File diff suppressed because it is too large Load Diff
+376
View File
@@ -0,0 +1,376 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
系统监控模块 - 提供实时系统监控功能
支持CPU、内存、磁盘、网络等指标的实时采集和历史记录
"""
import os
import json
import time
import threading
from datetime import datetime
from typing import Dict, List, Optional, Any
class SystemMonitor:
"""系统监控器 - 采集和提供系统运行指标"""
def __init__(self, config: dict):
self.config = config
self.base_dir = config.get('base_dir', './downloads')
# 历史数据配置
self.history_file = config.get('monitor_history_file', 'monitor_history.json')
self.history_hours = config.get('monitor_history_hours', 168) # 默认7天
self.history = []
self.history_lock = threading.Lock()
# SSE客户端管理
self.sse_clients = {}
self.sse_lock = threading.Lock()
# 监控配置
self.collection_interval = config.get('monitor_interval', 5) # 采集间隔(秒)
self.enabled = True
# 负载历史(用于计算趋势)
self.load_history = []
# 加载历史数据
self._load_history()
def get_realtime_stats(self) -> dict:
"""获取实时系统状态"""
import psutil
import traceback
stats = {
"timestamp": datetime.now().isoformat(),
"errors": []
}
try:
# CPU信息
cpu_percent = psutil.cpu_percent(interval=0.1)
cpu_count = psutil.cpu_count()
cpu_freq = psutil.cpu_freq()
stats["cpu"] = {
"percent": cpu_percent,
"count": cpu_count,
"freq_current": round(cpu_freq.current, 0) if cpu_freq else None,
"freq_max": round(cpu_freq.max, 0) if cpu_freq else None,
"freq_min": round(cpu_freq.min, 0) if cpu_freq else None,
"per_core": psutil.cpu_percent(interval=None, percpu=True)
}
except Exception as e:
stats["cpu"] = {"error": str(e)}
stats["errors"].append(f"CPU: {str(e)}")
try:
# 内存信息
memory = psutil.virtual_memory()
swap = psutil.swap_memory()
stats["memory"] = {
"total": memory.total,
"available": memory.available,
"used": memory.used,
"percent": memory.percent,
"swap_total": swap.total,
"swap_used": swap.used,
"swap_percent": swap.percent
}
except Exception as e:
stats["memory"] = {"error": str(e)}
stats["errors"].append(f"内存: {str(e)}")
try:
# 负载平均值(可能在 Termux 中不可用)
load_avg = os.getloadavg() if hasattr(os, 'getloadavg') else [0, 0, 0]
if "cpu" in stats:
stats["cpu"]["load_avg_1m"] = round(load_avg[0], 2)
stats["cpu"]["load_avg_5m"] = round(load_avg[1], 2)
stats["cpu"]["load_avg_15m"] = round(load_avg[2], 2)
except Exception:
pass
try:
# 磁盘信息
disk_usage = psutil.disk_usage(self.base_dir)
disk_io = psutil.disk_io_counters()
stats["disk"] = {
"total": disk_usage.total,
"used": disk_usage.used,
"free": disk_usage.free,
"percent": disk_usage.percent,
"read_bytes": disk_io.read_bytes if disk_io else 0,
"write_bytes": disk_io.write_bytes if disk_io else 0,
"read_count": disk_io.read_count if disk_io else 0,
"write_count": disk_io.write_count if disk_io else 0
}
except Exception as e:
stats["disk"] = {"error": str(e)}
stats["errors"].append(f"磁盘: {str(e)}")
try:
# 网络信息(可能在受限环境中失败)
net_io = psutil.net_io_counters()
connections = psutil.net_connections()
stats["network"] = {
"bytes_sent": net_io.bytes_sent,
"bytes_recv": net_io.bytes_recv,
"packets_sent": net_io.packets_sent,
"packets_recv": net_io.packets_recv,
"connections_count": len(connections),
"connections_established": len([c for c in connections if c.status == 'ESTABLISHED'])
}
except PermissionError as e:
stats["network"] = {
"note": "权限不足,无法访问网络信息",
"error": str(e)
}
except Exception as e:
stats["network"] = {"error": str(e)}
stats["errors"].append(f"网络: {str(e)}")
try:
# 进程信息
process = psutil.Process()
proc_mem = process.memory_info()
proc_cpu = process.cpu_percent(interval=0.1)
stats["process"] = {
"memory_rss": proc_mem.rss,
"memory_vms": proc_mem.vms,
"cpu_percent": proc_cpu,
"thread_count": process.num_threads(),
"open_files": process.num_fds() if hasattr(process, 'num_fds') else 0
}
except Exception as e:
stats["process"] = {"error": str(e)}
stats["errors"].append(f"进程: {str(e)}")
# 计算运行时间
try:
uptime = time.time() - self.config.get('start_time', time.time())
stats["uptime"] = round(uptime, 2)
except Exception:
stats["uptime"] = None
return stats
def get_monitor_history(self, hours: int = 24) -> dict:
"""获取历史监控数据"""
cutoff_time = time.time() - (hours * 3600)
with self.history_lock:
filtered_history = [
point for point in self.history
if point.get('timestamp_unix', 0) >= cutoff_time
]
return {
"hours": hours,
"total_points": len(filtered_history),
"data": filtered_history
}
def get_stats_summary(self) -> dict:
"""获取统计摘要"""
history = self.get_monitor_history(24).get('data', [])
if not history:
return {
"status": "no_data",
"message": "暂无监控数据"
}
# 计算各项指标的平均值和最大值
cpu_values = [p.get('cpu', {}).get('percent', 0) for p in history]
memory_values = [p.get('memory', {}).get('percent', 0) for p in history]
disk_values = [p.get('disk', {}).get('percent', 0) for p in history]
return {
"status": "ok",
"period_hours": 24,
"cpu": {
"avg": round(sum(cpu_values) / len(cpu_values), 1) if cpu_values else 0,
"max": max(cpu_values) if cpu_values else 0,
"min": min(cpu_values) if cpu_values else 0
},
"memory": {
"avg": round(sum(memory_values) / len(memory_values), 1) if memory_values else 0,
"max": max(memory_values) if memory_values else 0,
"min": min(memory_values) if memory_values else 0
},
"disk": {
"avg": round(sum(disk_values) / len(disk_values), 1) if disk_values else 0,
"max": max(disk_values) if disk_values else 0,
"min": min(disk_values) if disk_values else 0
},
"total_downloads": history[-1].get('downloads', {}).get('total', 0) if history else 0,
"total_connections": sum(p.get('network', {}).get('connections_count', 0) for p in history)
}
def start_monitoring(self, callback=None):
"""启动监控循环(后台线程)"""
def monitor_loop():
while self.enabled:
try:
stats = self.get_realtime_stats()
# 添加时间戳
stats['timestamp_unix'] = time.time()
# 保存历史
self._add_history_point(stats)
# SSE广播
if callback:
callback(stats)
self._broadcast_sse('stats', stats)
except Exception as e:
print(f"监控采集错误: {e}")
time.sleep(self.collection_interval)
thread = threading.Thread(target=monitor_loop, daemon=True)
thread.start()
return thread
def stop_monitoring(self):
"""停止监控"""
self.enabled = False
def register_sse_client(self, client_id: str, topics: List[str] = None) -> None:
"""注册SSE客户端"""
with self.sse_lock:
self.sse_clients[client_id] = {
'topics': set(topics) if topics else {'*'},
'last_ping': time.time()
}
def unregister_sse_client(self, client_id: str) -> None:
"""注销SSE客户端"""
with self.sse_lock:
self.sse_clients.pop(client_id, None)
def get_sse_clients_count(self) -> int:
"""获取SSE客户端数量"""
with self.sse_lock:
return len(self.sse_clients)
def broadcast_event(self, event_type: str, data: dict) -> None:
"""广播事件到所有SSE客户端"""
self._broadcast_sse(event_type, data)
def _broadcast_sse(self, event_type: str, data: dict) -> None:
"""SSE广播(内部方法)"""
message = f"event: {event_type}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
with self.sse_lock:
disconnected = []
for client_id, client in self.sse_clients.items():
try:
if client['topics'] == {'*'} or event_type in client['topics']:
# 实际发送需要在request handler中处理
# 这里只记录消息
pass
except Exception:
disconnected.append(client_id)
# 清理断开的客户端
for client_id in disconnected:
self.sse_clients.pop(client_id, None)
def _add_history_point(self, stats: dict) -> None:
"""添加历史数据点"""
# 精简数据以减少存储
point = {
"timestamp": stats.get('timestamp'),
"timestamp_unix": stats.get('timestamp_unix'),
"cpu_percent": stats.get('cpu', {}).get('percent', 0),
"memory_percent": stats.get('memory', {}).get('percent', 0),
"disk_percent": stats.get('disk', {}).get('percent', 0),
"network_bytes_sent": stats.get('network', {}).get('bytes_sent', 0),
"network_bytes_recv": stats.get('network', {}).get('bytes_recv', 0),
"connections_count": stats.get('network', {}).get('connections_count', 0),
"load_avg_1m": stats.get('cpu', {}).get('load_avg_1m', 0)
}
with self.history_lock:
self.history.append(point)
# 清理过期数据
cutoff_time = time.time() - (self.history_hours * 3600)
self.history = [
p for p in self.history
if p.get('timestamp_unix', 0) >= cutoff_time
]
# 保存到文件
self._save_history()
def _save_history(self) -> None:
"""保存历史数据到文件"""
try:
with open(self.history_file, 'w', encoding='utf-8') as f:
json.dump(self.history, f, ensure_ascii=False, indent=2)
except Exception as e:
print(f"保存监控历史失败: {e}")
def _load_history(self) -> None:
"""从文件加载历史数据"""
try:
if os.path.exists(self.history_file):
with open(self.history_file, 'r', encoding='utf-8') as f:
data = json.load(f)
if isinstance(data, list):
self.history = data[-10000:] # 限制历史数量
except Exception as e:
print(f"加载监控历史失败: {e}")
self.history = []
def get_health_status(self) -> dict:
"""获取健康状态"""
stats = self.get_realtime_stats()
if 'error' in stats:
return {
"status": "unhealthy",
"error": stats['error']
}
# 检查各项指标
warnings = []
cpu_percent = stats.get('cpu', {}).get('percent', 0)
if cpu_percent > 90:
warnings.append(f"CPU使用率过高: {cpu_percent}%")
elif cpu_percent > 70:
warnings.append(f"CPU使用率较高: {cpu_percent}%")
memory_percent = stats.get('memory', {}).get('percent', 0)
if memory_percent > 90:
warnings.append(f"内存使用率过高: {memory_percent}%")
elif memory_percent > 80:
warnings.append(f"内存使用率较高: {memory_percent}%")
disk_percent = stats.get('disk', {}).get('percent', 0)
if disk_percent > 90:
warnings.append(f"磁盘使用率过高: {disk_percent}%")
elif disk_percent > 80:
warnings.append(f"磁盘使用率较高: {disk_percent}%")
if warnings:
return {
"status": "degraded",
"warnings": warnings,
"timestamp": stats.get('timestamp')
}
return {
"status": "healthy",
"timestamp": stats.get('timestamp')
}
+722
View File
@@ -0,0 +1,722 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
内存优化模块
专为低端设备优化 (2CPU/1G 内存等)
"""
import os
import sys
import gc
import time
import threading
from typing import Dict, List, Optional, Callable
from dataclasses import dataclass
from contextlib import contextmanager
from functools import wraps
# 跨平台兼容处理:resource 模块仅在 Unix/Linux 上可用
try:
import resource
_HAS_RESOURCE = True
except ImportError:
_HAS_RESOURCE = False
resource = None
# 内存限制配置
DEFAULT_MEMORY_LIMIT = 512 * 1024 * 1024 # 512MB
LOW_MEMORY_LIMIT = 256 * 1024 * 1024 # 256MB
VERY_LOW_MEMORY_LIMIT = 128 * 1024 * 1024 # 128MB
class MemoryManager:
"""内存管理器 - 支持定时垃圾回收"""
def __init__(self, config: dict = None):
self.config = config or {}
self.enabled = self.config.get('enabled', True)
self.memory_limit = self.config.get('memory_limit', DEFAULT_MEMORY_LIMIT)
self.soft_limit = self.memory_limit * 0.8 # 80% 时触发警告
self.check_interval = self.config.get('check_interval', 10) # 秒
# 定时垃圾回收配置
self.gc_interval = self.config.get('gc_interval', 300) # 默认 5 分钟
self.enable_scheduled_gc = self.config.get('enable_scheduled_gc', True)
# 缓存清理回调
self.cache_cleaners: List[Callable] = []
# 回调函数
self.on_memory_warning: Optional[Callable] = None
self.on_memory_critical: Optional[Callable] = None
self._running = False
self._monitor_thread = None
self._gc_thread = None
def start(self):
"""启动内存监控"""
if not self.enabled:
return
self._running = True
# 启动内存监控线程
self._monitor_thread = threading.Thread(target=self._monitor_loop, daemon=True)
self._monitor_thread.start()
# 启动定时垃圾回收线程
if self.enable_scheduled_gc:
self._gc_thread = threading.Thread(target=self._gc_loop, daemon=True)
self._gc_thread.start()
print(f"[内存管理] 定时GC: {self.gc_interval}秒")
# 设置内存限制
self.set_memory_limit(self.memory_limit)
print(f"[内存管理] 已启动, 限制: {self.memory_limit // 1024 // 1024}MB")
def stop(self):
"""停止内存监控"""
self._running = False
if self._monitor_thread:
self._monitor_thread.join(timeout=2)
if self._gc_thread:
self._gc_thread.join(timeout=2)
def register_cache_cleaner(self, cleaner: Callable):
"""注册缓存清理回调函数"""
self.cache_cleaners.append(cleaner)
def set_memory_limit(self, limit: int):
"""设置内存限制 (Linux/Unix)"""
if not _HAS_RESOURCE or resource is None:
# Windows 平台不支持内存限制,跳过
print(f"[内存管理] 跳过内存限制设置 (Windows平台不支持)")
return
try:
# 软限制
resource.setrlimit(resource.RLIMIT_AS, (limit, limit))
print(f"[内存管理] 已设置内存限制: {limit // 1024 // 1024}MB")
except Exception as e:
print(f"[内存管理] 设置内存限制失败: {e}")
def get_memory_usage(self) -> dict:
"""获取内存使用情况"""
try:
# 进程内存
import psutil
process = psutil.Process(os.getpid())
mem_info = process.memory_info()
# 系统内存
sys_mem = psutil.virtual_memory()
return {
'process_rss': mem_info.rss,
'process_vms': mem_info.vms,
'process_percent': process.memory_percent(),
'system_total': sys_mem.total,
'system_available': sys_mem.available,
'system_percent': sys_mem.percent,
'process_rss_mb': mem_info.rss / 1024 / 1024,
'system_available_mb': sys_mem.available / 1024 / 1024
}
except Exception as e:
# 备用方法:使用 resource (Unix) 或返回估计值 (Windows)
if _HAS_RESOURCE and resource is not None:
try:
return {
'process_rss': resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * 1024,
'process_rss_mb': resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
}
except Exception:
pass
# Windows 无 psutil 的极端情况
return {
'process_rss': 0,
'process_rss_mb': 0,
'error': str(e)
}
def _monitor_loop(self):
"""监控循环"""
while self._running:
try:
usage = self.get_memory_usage()
# 检查是否达到软限制
if usage['process_rss'] >= self.soft_limit:
if self.on_memory_warning:
self.on_memory_warning(usage)
self._aggressive_cleanup()
# 检查是否达到硬限制
if usage['process_rss'] >= self.memory_limit:
if self.on_memory_critical:
self.on_memory_critical(usage)
self._emergency_cleanup()
time.sleep(self.check_interval)
except Exception:
pass
def _gc_loop(self):
"""定时垃圾回收循环"""
while self._running:
try:
# 执行垃圾回收
self._scheduled_gc()
# 清理注册过的缓存
for cleaner in self.cache_cleaners:
try:
cleaner()
except Exception:
pass
except Exception:
pass
time.sleep(self.gc_interval)
def _scheduled_gc(self):
"""定时 GC 执行"""
# 标准 GC
collected = gc.collect()
# 清理 Python 内部缓存
if hasattr(sys, 'exc_clear'):
sys.exc_clear()
def _aggressive_cleanup(self):
"""激进清理"""
# 强制垃圾回收
gc.collect()
# 清理 Python 缓存
if hasattr(gc, 'set_threshold'):
gc.set_threshold(500, 10, 5)
# 尝试释放内存
try:
import psutil
process = psutil.Process(os.getpid())
process.memory_info().rss # 刷新
except Exception:
pass
def _emergency_cleanup(self):
"""紧急清理"""
print("[内存管理] ⚠️ 达到内存限制,尝试紧急清理...")
# 完全垃圾回收
gc.collect()
gc.collect()
gc.collect()
# 清理所有缓存
if hasattr(gc, 'garbage'):
del gc.garbage[:]
# 触发警告
print("[内存管理] ⚠️ 内存仍过高,考虑重启服务")
def get_status(self) -> dict:
"""获取状态"""
usage = self.get_memory_usage()
return {
'enabled': self.enabled,
'memory_limit_mb': self.memory_limit // 1024 // 1024,
'current_mb': usage['process_rss_mb'],
'available_mb': usage.get('system_available_mb', 0),
'percent': (usage['process_rss'] / self.memory_limit * 100) if self.memory_limit else 0
}
# ==================== 低内存配置 ====================
class LowMemoryConfig:
"""低端设备配置"""
# 可禁用的功能列表
FEATURES = {
'ws': 'enable_ws', # WebSocket
'sse': 'enable_sse', # Server-Sent Events
'hash_calc': 'calculate_hash', # 文件哈希计算
'stats': 'enable_stats', # 统计功能
'monitor': 'enable_monitor', # 系统监控
'sync': 'enable_sync', # 同步功能
'mirrors': 'enable_mirrors', # 加速源
}
# 预设配置
PRESETS = {
'ultra_low': {
'description': '极低端设备 (<256MB RAM)',
'workers': 1,
'max_cache_size': 32 * 1024 * 1024, # 32MB
'chunk_size': 16 * 1024, # 16KB
'buffer_size': 32 * 1024, # 32KB
'db_pool_size': 1,
'max_connections': 3,
'monitor_interval': 60,
'gc_interval': 180, # 3分钟
'timeout': 15,
'disable_optional_features': ['ws', 'sse', 'hash_calc', 'stats', 'monitor']
},
'low': {
'description': '低端设备 (256-512MB RAM)',
'workers': 1,
'max_cache_size': 64 * 1024 * 1024, # 64MB
'chunk_size': 32 * 1024, # 32KB
'buffer_size': 64 * 1024, # 64KB
'db_pool_size': 1,
'max_connections': 5,
'monitor_interval': 30,
'gc_interval': 300, # 5分钟
'timeout': 20,
'disable_optional_features': ['ws', 'sse', 'hash_calc', 'stats']
},
'medium': {
'description': '中等设备 (512MB-1GB RAM)',
'workers': 2,
'max_cache_size': 128 * 1024 * 1024, # 128MB
'chunk_size': 64 * 1024, # 64KB
'buffer_size': 128 * 1024, # 128KB
'db_pool_size': 2,
'max_connections': 15,
'monitor_interval': 15,
'gc_interval': 600, # 10分钟
'timeout': 30,
'disable_optional_features': []
},
'high': {
'description': '高端设备 (1GB+ RAM)',
'workers': 4,
'max_cache_size': 256 * 1024 * 1024, # 256MB
'chunk_size': 128 * 1024, # 128KB
'buffer_size': 256 * 1024, # 256KB
'db_pool_size': 4,
'max_connections': 50,
'monitor_interval': 5,
'gc_interval': 900, # 15分钟
'timeout': 30,
'disable_optional_features': []
},
'performance': {
'description': '高性能设备 (4GB+ RAM)',
'workers': 8,
'max_cache_size': 1024 * 1024 * 1024, # 1GB
'chunk_size': 256 * 1024, # 256KB
'buffer_size': 512 * 1024, # 512KB
'db_pool_size': 8,
'max_connections': 200,
'monitor_interval': 3,
'gc_interval': 1800, # 30分钟
'timeout': 60,
'disable_optional_features': []
}
}
def __init__(self, preset: str = 'auto', custom_config: dict = None):
"""
初始化配置
Args:
preset: 预设 ('ultra_low', 'low', 'medium', 'high', 'auto')
custom_config: 自定义配置
"""
if preset == 'auto':
preset = self._detect_preset()
self.preset = preset
self.config = self.PRESETS.get(preset, self.PRESETS['low']).copy()
if custom_config:
self.config.update(custom_config)
def _detect_preset(self) -> str:
"""自动检测设备配置
检测逻辑:
1. 获取总内存和可用内存
2. 计算可用内存占比
3. 结合总内存和可用内存占比综合判断
"""
try:
import psutil
mem = psutil.virtual_memory()
total_ram = mem.total
available_ram = mem.available
percent_used = mem.percent # 已使用百分比
# 转换为MB
total_mb = total_ram / (1024 * 1024)
# 根据总内存和可用内存占比综合判断
if total_mb < 200:
# 低于 200MB 总内存
return 'ultra_low'
elif total_mb < 400:
# 200MB - 400MB
return 'ultra_low'
elif total_mb < 700:
# 400MB - 700MB
if percent_used > 80:
return 'ultra_low'
return 'low'
elif total_mb < 1200:
# 700MB - 1.2GB
if percent_used > 70:
return 'ultra_low'
elif percent_used > 50:
return 'low'
return 'medium'
elif total_mb < 2500:
# 1.2GB - 2.5GB
if percent_used > 70:
return 'low'
elif percent_used > 40:
return 'medium'
return 'high'
elif total_mb < 5000:
# 2.5GB - 5GB
if percent_used > 60:
return 'medium'
return 'high'
else:
# 5GB+
return 'performance'
except ImportError:
# 如果没有 psutil,使用保守的 low 配置
return 'low'
except Exception:
return 'low'
def get_device_info(self) -> dict:
"""获取设备详细信息用于显示"""
try:
import psutil
mem = psutil.virtual_memory()
cpu_count = psutil.cpu_count(logical=True) or 1
return {
'total_ram_mb': mem.total / (1024 * 1024),
'available_ram_mb': mem.available / (1024 * 1024),
'percent_used': mem.percent,
'cpu_count': cpu_count,
'preset': self.preset
}
except Exception:
return {
'total_ram_mb': 0,
'available_ram_mb': 0,
'percent_used': 0,
'cpu_count': 1,
'preset': self.preset
}
def apply_to_config(self, base_config: dict) -> dict:
"""应用配置到基础配置"""
config = base_config.copy()
# 应用通用设置
config['workers'] = self.config.get('workers', 1)
config['max_cache_size'] = self.config.get('max_cache_size', 64 * 1024 * 1024)
config['chunk_size'] = self.config.get('chunk_size', 64 * 1024)
config['buffer_size'] = self.config.get('buffer_size', 128 * 1024)
config['timeout'] = self.config.get('timeout', 30)
# 数据库池
if 'database' not in config:
config['database'] = {}
config['database']['db_pool_size'] = self.config.get('db_pool_size', 2)
# 定时 GC 配置
config['gc_interval'] = self.config.get('gc_interval', 300)
# 禁用可选功能
for feature in self.config.get('disable_optional_features', []):
feature_key = self.FEATURES.get(feature, feature)
if feature_key.startswith('enable_') or feature_key == 'calculate_hash':
config[feature_key] = False
return config
def get_config(self) -> dict:
"""获取配置"""
return self.config.copy()
def get_status(self) -> dict:
"""获取状态"""
return {
'preset': self.preset,
'description': self.config['description'],
'settings': self.config
}
# ==================== 流式处理优化 ====================
class StreamingOptimizer:
"""流式处理优化器"""
def __init__(self, config: dict = None):
self.config = config or {}
self.chunk_size = self.config.get('chunk_size', 128 * 1024)
self.buffer_size = self.config.get('buffer_size', 256 * 1024)
# 内存池
self._chunk_pool = None
self._use_memory_pool = self.config.get('use_memory_pool', True)
def get_optimized_chunk_size(self, file_size: int) -> int:
"""根据文件大小获取优化的块大小"""
if file_size < 1024 * 1024: # < 1MB
return 16 * 1024 # 16KB
elif file_size < 10 * 1024 * 1024: # < 10MB
return 32 * 1024 # 32KB
elif file_size < 100 * 1024 * 1024: # < 100MB
return 64 * 1024 # 64KB
else:
return self.chunk_size
@contextmanager
def memory_efficient_file_read(self, file_path: str, chunk_size: int = None):
"""
内存高效的文件读取
Usage:
with optimizer.memory_efficient_file_read('/path/to/file') as f:
for chunk in f:
process(chunk)
"""
chunk_size = chunk_size or self.chunk_size
file_size = os.path.getsize(file_path)
chunk_size = self.get_optimized_chunk_size(file_size)
file = open(file_path, 'rb')
try:
yield file
finally:
file.close()
@contextmanager
def memory_efficient_file_write(self, file_path: str, chunk_size: int = None):
"""内存高效的文件写入"""
chunk_size = chunk_size or self.chunk_size
file = open(file_path, 'wb')
try:
yield file
finally:
file.close()
# ==================== 架构检测 ====================
class ArchitectureDetector:
"""架构检测器"""
@staticmethod
def get_architecture() -> dict:
"""
获取架构信息
Returns:
dict: 包含架构信息的字典
"""
info = {
'platform': sys.platform,
'architecture': 'unknown',
'machine': 'unknown',
'processor': 'unknown',
'python_version': sys.version,
'byte_order': sys.byteorder
}
# 机器类型
info['machine'] = os.uname().machine if hasattr(os, 'uname') else 'unknown'
# 检测 32位/64位
if info['machine'] in ['x86_64', 'amd64', 'aarch64', 'arm64']:
info['architecture'] = '64bit'
elif info['machine'] in ['i386', 'i686', 'armv7l', 'armv6l']:
info['architecture'] = '32bit'
elif info['machine'] in ['armv8l', 'aarch32']:
info['architecture'] = '32bit' # 32位ARM
# ARM 变体
if info['machine'].startswith('arm'):
if info['machine'] in ['armv7l', 'armv7hl']:
info['arm_variant'] = 'armv7'
elif info['machine'].startswith('armv8'):
info['arm_variant'] = 'armv8'
elif info['machine'].startswith('armv6'):
info['arm_variant'] = 'armv6'
else:
info['arm_variant'] = 'unknown'
# x86 变体
if info['machine'] in ['i386', 'i686']:
info['x86_variant'] = 'i386'
elif info['machine'] == 'x86_64':
info['x86_variant'] = 'x86_64'
return info
@staticmethod
def is_low_end_device() -> bool:
"""检测是否为低端设备"""
try:
import psutil
mem = psutil.virtual_memory()
return mem.total < 1024 * 1024 * 1024 # < 1GB
except Exception:
return False
@staticmethod
def get_recommended_config() -> dict:
"""获取推荐的配置"""
arch = ArchitectureDetector.get_architecture()
if arch['architecture'] == '32bit':
return {
'max_workers': 2,
'max_cache_size': 100 * 1024 * 1024, # 100MB
'enable_threading': True,
'use_processes': False, # 32位进程数有限制
'max_file_handles': 256
}
elif ArchitectureDetector.is_low_end_device():
return {
'max_workers': 1,
'max_cache_size': 50 * 1024 * 1024, # 50MB
'enable_threading': True,
'use_processes': False,
'max_file_handles': 128
}
else:
return {
'max_workers': 4,
'max_cache_size': 500 * 1024 * 1024, # 500MB
'enable_threading': True,
'use_processes': True,
'max_file_handles': 1024
}
# ==================== 兼容性检查 ====================
def check_compatibility() -> dict:
"""
检查系统兼容性
Returns:
dict: 兼容性检查结果
"""
results = {
'compatible': True,
'warnings': [],
'errors': [],
'info': {}
}
# Python 版本检查
if sys.version_info < (3, 8):
results['compatible'] = False
results['errors'].append(f"Python 3.8+ 所需, 当前版本: {sys.version}")
# 架构信息
arch_info = ArchitectureDetector.get_architecture()
results['info']['architecture'] = arch_info
# 检查必需模块
required_modules = [
('os', '标准库'),
('json', '标准库'),
('http', '标准库'),
('sqlite3', '标准库')
]
optional_modules = [
('psutil', '系统监控 (推荐)'),
('sqlalchemy', '数据库 (推荐)'),
(' cryptography', '加密 (推荐)'),
('aiohttp', '异步HTTP (可选)'),
('paramiko', 'SSH/SFTP (可选)')
]
for module, desc in required_modules:
try:
__import__(module)
except ImportError:
results['compatible'] = False
results['errors'].append(f"必需模块缺失: {module} ({desc})")
for module, desc in optional_modules:
try:
__import__(module)
except ImportError:
results['warnings'].append(f"可选模块缺失: {module} ({desc})")
# 内存检查
try:
import psutil
mem = psutil.virtual_memory()
if mem.total < 256 * 1024 * 1024:
results['warnings'].append("内存低于 256MB,可能无法正常运行")
except Exception:
results['warnings'].append("无法检测内存,可能内存不足")
# 磁盘空间检查
try:
disk = psutil.disk_usage('.')
if disk.free < 100 * 1024 * 1024: # 100MB
results['warnings'].append("可用磁盘空间不足 100MB")
except Exception:
pass
return results
# ==================== 便捷函数 ====================
def get_memory_manager(config: dict = None) -> MemoryManager:
"""获取内存管理器"""
return MemoryManager(config)
def get_low_memory_config(preset: str = 'auto') -> LowMemoryConfig:
"""获取低端设备配置"""
return LowMemoryConfig(preset)
def detect_and_configure() -> dict:
"""
自动检测并配置
Returns:
dict: 配置信息
"""
# 检查兼容性
compat = check_compatibility()
if not compat['compatible']:
print("⚠️ 系统兼容性警告:")
for error in compat['errors']:
print(f" - {error}")
# 获取推荐配置
arch_config = ArchitectureDetector.get_recommended_config()
# 获取低端设备配置
low_mem_config = get_low_memory_config('auto')
arch_info = ArchitectureDetector.get_architecture()
return {
'compatible': compat['compatible'],
'architecture': arch_info,
'recommended': arch_config,
'low_memory': low_mem_config.get_status(),
'warnings': compat['warnings']
}
+386
View File
@@ -0,0 +1,386 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Prometheus 指标导出模块
提供 /metrics 端点暴露监控数据
"""
import os
import sys
import time
from datetime import datetime
from typing import Dict, Optional
# 添加项目根目录到路径
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
class PrometheusMetrics:
"""Prometheus 指标收集器"""
def __init__(self, config: dict = None):
self.config = config or {}
self._metrics: Dict[str, dict] = {}
# 初始化指标
self._init_metrics()
def _init_metrics(self):
"""初始化指标定义"""
self._metrics = {
# 服务器指标
'hyc_server_uptime_seconds': {
'type': 'gauge',
'description': 'Server uptime in seconds',
'value': 0
},
'hyc_server_start_time': {
'type': 'gauge',
'description': 'Server start timestamp',
'value': time.time()
},
# 文件指标
'hyc_files_total': {
'type': 'gauge',
'description': 'Total number of files in the system',
'value': 0
},
'hyc_files_size_bytes': {
'type': 'gauge',
'description': 'Total size of all files in bytes',
'value': 0
},
'hyc_downloads_total': {
'type': 'counter',
'description': 'Total number of downloads',
'value': 0
},
'hyc_downloads_today': {
'type': 'counter',
'description': 'Number of downloads today',
'value': 0
},
# 缓存指标
'hyc_cache_size_bytes': {
'type': 'gauge',
'description': 'Current cache size in bytes',
'value': 0
},
'hyc_cache_entries': {
'type': 'gauge',
'description': 'Number of cache entries',
'value': 0
},
'hyc_cache_hits_total': {
'type': 'counter',
'description': 'Total number of cache hits',
'value': 0
},
'hyc_cache_misses_total': {
'type': 'counter',
'description': 'Total number of cache misses',
'value': 0
},
'hyc_cache_hit_ratio': {
'type': 'gauge',
'description': 'Cache hit ratio (0-1)',
'value': 0
},
# 同步指标
'hyc_sync_sources_total': {
'type': 'gauge',
'description': 'Total number of sync sources',
'value': 0
},
'hyc_sync_running': {
'type': 'gauge',
'description': 'Number of currently running sync operations',
'value': 0
},
'hyc_sync_files_total': {
'type': 'counter',
'description': 'Total number of synced files',
'value': 0
},
'hyc_sync_last_timestamp': {
'type': 'gauge',
'description': 'Timestamp of last successful sync',
'value': 0
},
# 数据库指标
'hyc_db_files_total': {
'type': 'gauge',
'description': 'Number of files in database',
'value': 0
},
'hyc_db_sync_records': {
'type': 'gauge',
'description': 'Number of sync records in database',
'value': 0
},
'hyc_db_cache_records': {
'type': 'gauge',
'description': 'Number of cache records in database',
'value': 0
},
# 系统资源指标 (从监控模块获取)
'hyc_cpu_percent': {
'type': 'gauge',
'description': 'CPU usage percentage',
'value': 0
},
'hyc_memory_percent': {
'type': 'gauge',
'description': 'Memory usage percentage',
'value': 0
},
'hyc_disk_percent': {
'type': 'gauge',
'description': 'Disk usage percentage',
'value': 0
},
'hyc_disk_free_bytes': {
'type': 'gauge',
'description': 'Free disk space in bytes',
'value': 0
},
'hyc_disk_total_bytes': {
'type': 'gauge',
'description': 'Total disk space in bytes',
'value': 0
},
'hyc_network_rx_bytes': {
'type': 'counter',
'description': 'Network receive bytes',
'value': 0
},
'hyc_network_tx_bytes': {
'type': 'counter',
'description': 'Network transmit bytes',
'value': 0
},
'hyc_active_connections': {
'type': 'gauge',
'description': 'Number of active connections',
'value': 0
},
# 镜像源指标
'hyc_mirror_enabled': {
'type': 'gauge',
'description': 'Whether a mirror is enabled (1=enabled, 0=disabled)',
'labels': ['mirror_type'],
'value': {}
},
'hyc_mirror_last_sync': {
'type': 'gauge',
'description': 'Timestamp of last mirror sync',
'labels': ['mirror_type'],
'value': {}
},
}
def set_uptime(self, seconds: float):
"""设置运行时间"""
self._metrics['hyc_server_uptime_seconds']['value'] = seconds
def set_files(self, count: int, size_bytes: int):
"""设置文件统计"""
self._metrics['hyc_files_total']['value'] = count
self._metrics['hyc_files_size_bytes']['value'] = size_bytes
def set_downloads(self, total: int, today: int):
"""设置下载统计"""
self._metrics['hyc_downloads_total']['value'] = total
self._metrics['hyc_downloads_today']['value'] = today
def set_cache(self, size_bytes: int, entries: int, hits: int, misses: int):
"""设置缓存统计"""
self._metrics['hyc_cache_size_bytes']['value'] = size_bytes
self._metrics['hyc_cache_entries']['value'] = entries
self._metrics['hyc_cache_hits_total']['value'] = hits
self._metrics['hyc_cache_misses_total']['value'] = misses
# 计算命中率
total = hits + misses
if total > 0:
self._metrics['hyc_cache_hit_ratio']['value'] = hits / total
else:
self._metrics['hyc_cache_hit_ratio']['value'] = 0
def set_sync(self, running: int, files_total: int, last_timestamp: float):
"""设置同步统计"""
self._metrics['hyc_sync_running']['value'] = running
self._metrics['hyc_sync_files_total']['value'] = files_total
self._metrics['hyc_sync_last_timestamp']['value'] = last_timestamp
def set_db_stats(self, files: int, sync_records: int, cache_records: int):
"""设置数据库统计"""
self._metrics['hyc_db_files_total']['value'] = files
self._metrics['hyc_db_sync_records']['value'] = sync_records
self._metrics['hyc_db_cache_records']['value'] = cache_records
def set_system(self, cpu: float, memory: float, disk: float,
disk_free: int, disk_total: int,
network_rx: int, network_tx: int):
"""设置系统资源统计"""
self._metrics['hyc_cpu_percent']['value'] = cpu
self._metrics['hyc_memory_percent']['value'] = memory
self._metrics['hyc_disk_percent']['value'] = disk
self._metrics['hyc_disk_free_bytes']['value'] = disk_free
self._metrics['hyc_disk_total_bytes']['value'] = disk_total
self._metrics['hyc_network_rx_bytes']['value'] = network_rx
self._metrics['hyc_network_tx_bytes']['value'] = network_tx
def set_connections(self, count: int):
"""设置连接数"""
self._metrics['hyc_active_connections']['value'] = count
def set_mirror_status(self, mirror_type: str, enabled: bool, last_sync: float):
"""设置镜像源状态"""
key = 'hyc_mirror_enabled'
if 'labels' not in self._metrics[key]:
self._metrics[key]['labels'] = ['mirror_type']
if 'value' not in self._metrics[key]:
self._metrics[key]['value'] = {}
self._metrics[key]['value'][mirror_type] = 1 if enabled else 0
key = 'hyc_mirror_last_sync'
if 'labels' not in self._metrics[key]:
self._metrics[key]['labels'] = ['mirror_type']
if 'value' not in self._metrics[key]:
self._metrics[key]['value'] = {}
self._metrics[key]['value'][mirror_type] = last_sync
def increment_downloads(self, count: int = 1):
"""增加下载计数"""
self._metrics['hyc_downloads_total']['value'] += count
def increment_cache_hits(self, count: int = 1):
"""增加缓存命中计数"""
self._metrics['hyc_cache_hits_total']['value'] += count
def increment_cache_misses(self, count: int = 1):
"""增加缓存未命中计数"""
self._metrics['hyc_cache_misses_total']['value'] += count
def increment_sync_files(self, count: int = 1):
"""增加同步文件计数"""
self._metrics['hyc_sync_files_total']['value'] += count
def generate_metrics(self) -> str:
"""生成 Prometheus 格式的指标输出"""
output = []
output.append("# Prometheus metrics for HYC Mirror Server")
output.append(f"# Generated at: {datetime.now().isoformat()}")
output.append("")
for name, metric in self._metrics.items():
desc = metric.get('description', '')
mtype = metric.get('type', 'gauge')
output.append(f"# TYPE {name} {mtype}")
output.append(f"# HELP {name} {desc}")
value = metric.get('value')
# 处理带标签的指标
if 'labels' in metric and isinstance(value, dict):
labels = metric['labels']
for label_values, v in value.items():
if isinstance(label_values, str):
label_str = f'{",".join([f"{labels[0]}={label_values}"])}'
else:
label_str = ','.join([f"{l}={v}" for l, v in zip(labels, label_values)])
output.append(f"{name}{{{label_str}}} {v}")
# 处理普通指标
elif isinstance(value, dict):
# 旧格式,可能是直接存储
for k, v in value.items():
output.append(f"{name}{{type=\"{k}\"}} {v}")
else:
output.append(f"{name} {value}")
output.append("")
return '\n'.join(output)
# ==================== 指标中间件 ====================
class MetricsMiddleware:
"""HTTP 请求指标中间件"""
def __init__(self, metrics: PrometheusMetrics = None):
self.metrics = metrics or PrometheusMetrics()
self._request_count = 0
self._request_duration_total = 0
self._errors_4xx = 0
self._errors_5xx = 0
def record_request(self, duration: float, status_code: int):
"""记录请求"""
self._request_count += 1
self._request_duration_total += duration
if 400 <= status_code < 500:
self._errors_4xx += 1
elif status_code >= 500:
self._errors_5xx += 1
def get_request_count(self) -> int:
return self._request_count
def get_request_duration_total(self) -> float:
return self._request_duration_total
def get_errors_4xx(self) -> int:
return self._errors_4xx
def get_errors_5xx(self) -> int:
return self._errors_5xx
def update_metrics(self):
"""更新 Prometheus 指标"""
# 添加请求相关指标
if 'hyc_http_requests_total' not in self.metrics._metrics:
self.metrics._metrics['hyc_http_requests_total'] = {
'type': 'counter',
'description': 'Total HTTP requests',
'value': self._request_count
}
else:
self.metrics._metrics['hyc_http_requests_total']['value'] = self._request_count
if 'hyc_http_request_duration_seconds_total' not in self.metrics._metrics:
self.metrics._metrics['hyc_http_request_duration_seconds_total'] = {
'type': 'counter',
'description': 'Total HTTP request duration in seconds',
'value': self._request_duration_total
}
else:
self.metrics._metrics['hyc_http_request_duration_seconds_total']['value'] = self._request_duration_total
if 'hyc_http_requests_4xx_total' not in self.metrics._metrics:
self.metrics._metrics['hyc_http_requests_4xx_total'] = {
'type': 'counter',
'description': 'Total HTTP 4xx errors',
'value': self._errors_4xx
}
else:
self.metrics._metrics['hyc_http_requests_4xx_total']['value'] = self._errors_4xx
if 'hyc_http_requests_5xx_total' not in self.metrics._metrics:
self.metrics._metrics['hyc_http_requests_5xx_total'] = {
'type': 'counter',
'description': 'Total HTTP 5xx errors',
'value': self._errors_5xx
}
else:
self.metrics._metrics['hyc_http_requests_5xx_total']['value'] = self._errors_5xx
+374
View File
@@ -0,0 +1,374 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
定时任务调度器
支持 cron 表达式和简单间隔的定时任务
"""
import os
import time
import logging
import threading
from datetime import datetime, timedelta
from typing import Dict, List, Optional, Callable
from enum import Enum
logger = logging.getLogger(__name__)
class TaskStatus(Enum):
"""任务状态"""
IDLE = "idle"
RUNNING = "running"
ERROR = "error"
DISABLED = "disabled"
class ScheduledTask:
"""定时任务"""
def __init__(self, name: str, task_type: str, config: dict,
callback: Callable, logger=None):
"""
初始化定时任务
Args:
name: 任务名称
task_type: 任务类型 ('cron' 或 'interval')
config: 任务配置
callback: 回调函数
logger: 日志器
"""
self.name = name
self.task_type = task_type # 'cron' 或 'interval'
self.config = config or {}
self.callback = callback
self.logger = logger or logging.getLogger(__name__)
# 状态
self.status = TaskStatus.IDLE
self.last_run: Optional[datetime] = None
self.next_run: Optional[datetime] = None
self.last_error: Optional[str] = None
self.run_count = 0
# 配置解析
self._parse_config()
def _parse_config(self):
"""解析任务配置"""
if self.task_type == 'cron':
# Cron 表达式: "minute hour day month weekday"
# 例如: "0 3 * * *" 每天凌晨3点
cron = self.config.get('cron', '0 0 * * *')
parts = cron.split()
if len(parts) == 5:
self.cron_parts = {
'minute': self._parse_cron_part(parts[0], 0, 59),
'hour': self._parse_cron_part(parts[1], 0, 23),
'day': self._parse_cron_part(parts[2], 1, 31),
'month': self._parse_cron_part(parts[3], 1, 12),
'weekday': self._parse_cron_part(parts[4], 0, 6)
}
else:
self.logger.warning(f"无效的 cron 表达式: {cron}")
self.cron_parts = None
elif self.task_type == 'interval':
# 间隔: seconds, minutes, hours
interval = self.config.get('interval', {})
self.interval_seconds = (
interval.get('seconds', 0) +
interval.get('minutes', 0) * 60 +
interval.get('hours', 0) * 3600 +
interval.get('days', 0) * 86400
)
if self.interval_seconds <= 0:
self.interval_seconds = 3600 # 默认1小时
# 是否启用
self.enabled = self.config.get('enabled', True)
def _parse_cron_part(self, part: str, min_val: int, max_val: int) -> List[int]:
"""解析 cron 表达式的一部分"""
result = []
if part == '*':
return list(range(min_val, max_val + 1))
# 处理列表: "1,2,3"
if ',' in part:
return self._parse_cron_part(part.replace(',', ' '), min_val, max_val)
# 处理范围: "1-5"
if '-' in part:
start, end = part.split('-')
return list(range(int(start), int(end) + 1))
# 处理步进: "*/5"
if '/' in part:
base, step = part.split('/')
base_list = self._parse_cron_part(base or '*', min_val, max_val)
step = int(step)
return base_list[::step]
# 单个值
try:
val = int(part)
if min_val <= val <= max_val:
return [val]
except ValueError:
pass
return []
def should_run_now(self) -> bool:
"""检查是否应该在当前时刻运行"""
if not self.enabled:
return False
now = datetime.now()
if self.task_type == 'cron' and self.cron_parts:
return self._matches_cron(now)
elif self.task_type == 'interval':
if self.last_run is None:
return True
elapsed = (now - self.last_run).total_seconds()
return elapsed >= self.interval_seconds
return False
def _matches_cron(self, dt: datetime) -> bool:
"""检查时间是否匹配 cron 表达式"""
if not self.cron_parts:
return False
return (
dt.minute in self.cron_parts['minute'] and
dt.hour in self.cron_parts['hour'] and
dt.day in self.cron_parts['day'] and
dt.month in self.cron_parts['month'] and
dt.weekday() in self.cron_parts['weekday']
)
def get_next_run_time(self) -> Optional[datetime]:
"""计算下次运行时间"""
if not self.enabled:
return None
now = datetime.now()
if self.task_type == 'cron' and self.cron_parts:
# 找到下一个匹配的时间点
for i in range(365 * 24 * 60): # 最多查找1年
candidate = now + timedelta(minutes=i)
if self._matches_cron(candidate):
return candidate
elif self.task_type == 'interval':
if self.last_run:
return self.last_run + timedelta(seconds=self.interval_seconds)
return now
return None
def run(self) -> bool:
"""执行任务"""
if self.status == TaskStatus.RUNNING:
self.logger.warning(f"任务 {self.name} 已在运行中")
return False
self.status = TaskStatus.RUNNING
self.last_run = datetime.now()
self.last_error = None
try:
self.logger.info(f"开始执行定时任务: {self.name}")
result = self.callback(self.name, self.config)
self.run_count += 1
self.logger.info(f"定时任务 {self.name} 执行完成")
return True
except Exception as e:
self.last_error = str(e)
self.status = TaskStatus.ERROR
self.logger.error(f"定时任务 {self.name} 执行失败: {e}")
return False
finally:
if self.status != TaskStatus.ERROR:
self.status = TaskStatus.IDLE
def to_dict(self) -> dict:
"""转换为字典"""
return {
'name': self.name,
'type': self.task_type,
'enabled': self.enabled,
'status': self.status.value,
'config': self.config,
'last_run': self.last_run.isoformat() if self.last_run else None,
'next_run': self.next_run.isoformat() if self.next_run else None,
'run_count': self.run_count,
'last_error': self.last_error
}
class Scheduler:
"""定时任务调度器"""
def __init__(self, config: dict = None):
self.config = config or {}
self.tasks: Dict[str, ScheduledTask] = {}
self._running = False
self._thread: Optional[threading.Thread] = None
self._lock = threading.Lock()
# 默认检查间隔
self.check_interval = self.config.get('check_interval', 10)
# 事件回调
self.on_task_start: Optional[Callable] = None
self.on_task_complete: Optional[Callable] = None
self.on_task_error: Optional[Callable] = None
def add_task(self, name: str, task_type: str, config: dict,
callback: Callable) -> bool:
"""
添加定时任务
Args:
name: 任务名称
task_type: 任务类型 ('cron' 或 'interval')
config: 任务配置
callback: 回调函数
Returns:
是否成功
"""
with self._lock:
if name in self.tasks:
logger.warning(f"任务 {name} 已存在,将被替换")
self.tasks[name] = ScheduledTask(name, task_type, config, callback, logger)
return True
def remove_task(self, name: str) -> bool:
"""移除任务"""
with self._lock:
if name in self.tasks:
del self.tasks[name]
return True
return False
def get_task(self, name: str) -> Optional[ScheduledTask]:
"""获取任务"""
return self.tasks.get(name)
def get_all_tasks(self) -> List[dict]:
"""获取所有任务状态"""
with self._lock:
for task in self.tasks.values():
task.next_run = task.get_next_run_time()
return [task.to_dict() for task in self.tasks.values()]
def start(self):
"""启动调度器"""
if self._running:
logger.warning("调度器已在运行中")
return
self._running = True
self._thread = threading.Thread(target=self._run_loop, daemon=True)
self._thread.start()
logger.info("定时任务调度器已启动")
def stop(self):
"""停止调度器"""
self._running = False
if self._thread:
self._thread.join(timeout=5)
logger.info("定时任务调度器已停止")
def _run_loop(self):
"""运行循环"""
while self._running:
try:
now = datetime.now()
with self._lock:
for name, task in self.tasks.items():
if task.should_run_now():
# 使用线程池执行任务
from concurrent.futures import ThreadPoolExecutor
with ThreadPoolExecutor(max_workers=1) as executor:
executor.submit(task.run)
time.sleep(self.check_interval)
except Exception as e:
logger.error(f"调度器循环错误: {e}")
time.sleep(5)
def run_task_now(self, name: str) -> bool:
"""立即运行指定任务"""
task = self.get_task(name)
if task:
return task.run()
return False
def enable_task(self, name: str, enabled: bool = True) -> bool:
"""启用/禁用任务"""
task = self.get_task(name)
if task:
task.enabled = enabled
return True
return False
def update_task_config(self, name: str, config: dict) -> bool:
"""更新任务配置"""
task = self.get_task(name)
if task:
task.config.update(config)
task._parse_config()
return True
return False
# ==================== 同步任务工厂 ====================
# def create_sync_task_callback(sync_manager):
# """创建同步任务的回调函数"""
# def sync_task_callback(task_name: str, config: dict):
# """同步任务回调"""
# sync_manager.start_sync(task_name)
# return True
# return sync_task_callback
# ==================== 默认任务配置 ====================
# DEFAULT_SCHEDULED_TASKS = { ... }
DEFAULT_SCHEDULED_TASKS = {
# 数据库清理 - 每天凌晨2点
'cleanup_db': {
'type': 'cron',
'config': {
'cron': '0 2 * * *',
'enabled': True
}
},
# 缓存清理 - 每6小时
'cleanup_cache': {
'type': 'interval',
'config': {
'interval': {'hours': 6},
'enabled': True
}
},
# 健康检查 - 每5分钟
'health_check': {
'type': 'interval',
'config': {
'interval': {'minutes': 5},
'enabled': True
}
}
}
+480
View File
@@ -0,0 +1,480 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
安全增强模块
提供 IP 白名单/黑名单、请求速率限制、HTTPS 支持
"""
import os
import sys
import time
import json
import hashlib
import threading
import logging
import ssl
import ipaddress
from datetime import datetime, timedelta
from typing import Dict, List, Optional, Tuple
from dataclasses import dataclass, field
from functools import wraps
from collections import defaultdict
from queue import Queue
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
logger = logging.getLogger(__name__)
# ==================== IP 管理器 ====================
class IPManager:
"""IP 管理器 - 白名单/黑名单"""
def __init__(self, config: dict = None):
self.config = config or {}
self.whitelist: List[str] = []
self.blacklist: List[str] = []
self._load_lists()
def _load_lists(self):
"""加载 IP 列表"""
# 加载白名单
whitelist_file = self.config.get('whitelist_file', 'whitelist.txt')
if os.path.exists(whitelist_file):
with open(whitelist_file, 'r') as f:
for line in f:
line = line.strip()
if line and not line.startswith('#'):
self.whitelist.append(line)
# 加载黑名单
blacklist_file = self.config.get('blacklist_file', 'blacklist.txt')
if os.path.exists(blacklist_file):
with open(blacklist_file, 'r') as f:
for line in f:
line = line.strip()
if line and not line.startswith('#'):
self.blacklist.append(line)
def is_allowed(self, ip: str) -> Tuple[bool, str]:
"""
检查 IP 是否允许访问
返回: (是否允许, 原因)
"""
# 检查白名单
if self.whitelist:
for pattern in self.whitelist:
if self._match_ip(ip, pattern):
return True, "白名单"
# 检查黑名单
for pattern in self.blacklist:
if self._match_ip(ip, pattern):
return False, "黑名单"
return True, "允许"
def _match_ip(self, ip: str, pattern: str) -> bool:
"""匹配 IP"""
try:
# 单个 IP
if pattern == ip:
return True
# CIDR 范围
if '/' in pattern:
network = ipaddress.ip_network(pattern, strict=False)
return ipaddress.ip_address(ip) in network
# 通配符 (例如: 192.168.1.*)
if pattern.endswith('*'):
prefix = pattern.rstrip('*').rstrip('.')
return ip.startswith(prefix)
except Exception:
pass
return False
def add_to_whitelist(self, ip: str):
"""添加到白名单"""
if ip not in self.whitelist:
self.whitelist.append(ip)
self._save_list('whitelist')
def add_to_blacklist(self, ip: str):
"""添加到黑名单"""
if ip not in self.blacklist:
self.blacklist.append(ip)
self._save_list('blacklist')
def remove_from_whitelist(self, ip: str):
"""从白名单移除"""
if ip in self.whitelist:
self.whitelist.remove(ip)
self._save_list('whitelist')
def remove_from_blacklist(self, ip: str):
"""从黑名单移除"""
if ip in self.blacklist:
self.blacklist.remove(ip)
self._save_list('blacklist')
def _save_list(self, list_type: str):
"""保存列表到文件"""
filename = f'{list_type}.txt'
data = '\n'.join(self.whitelist if list_type == 'whitelist' else self.blacklist)
with open(filename, 'w') as f:
f.write(data)
def get_status(self) -> dict:
"""获取状态"""
return {
'whitelist_count': len(self.whitelist),
'blacklist_count': len(self.blacklist),
'whitelist': self.whitelist[:10],
'blacklist': self.blacklist[:10]
}
# ==================== 速率限制器 ====================
class RateLimiter:
"""请求速率限制器"""
def __init__(self, config: dict = None):
self.config = config or {}
self.requests: Dict[str, List[float]] = defaultdict(list)
self.lock = threading.Lock()
# 配置
self.requests_per_minute = self.config.get('requests_per_minute', 100)
self.burst_limit = self.config.get('burst_limit', 20)
self.window_seconds = 60
def is_allowed(self, identifier: str) -> Tuple[bool, int]:
"""
检查请求是否允许
返回: (是否允许, 剩余配额)
"""
now = time.time()
window_start = now - self.window_seconds
with self.lock:
# 清理过期记录
self.requests[identifier] = [
t for t in self.requests[identifier]
if t > window_start
]
# 检查限制
if len(self.requests[identifier]) >= self.requests_per_minute:
return False, 0
# 记录请求
self.requests[identifier].append(now)
remaining = self.requests_per_minute - len(self.requests[identifier])
return True, remaining
def get_usage(self, identifier: str) -> dict:
"""获取使用情况"""
now = time.time()
window_start = now - self.window_seconds
with self.lock:
requests = [
t for t in self.requests[identifier]
if t > window_start
]
return {
'requests': len(requests),
'limit': self.requests_per_minute,
'remaining': self.requests_per_minute - len(requests),
'reset_in': int(self.window_seconds - (now - min(requests) if requests else now))
}
def get_status(self) -> dict:
"""获取全局状态"""
total_requests = sum(len(v) for v in self.requests.values())
return {
'active_ips': len(self.requests),
'total_requests': total_requests,
'limit_per_minute': self.requests_per_minute,
'burst_limit': self.burst_limit
}
def reset(self):
"""重置所有记录"""
with self.lock:
self.requests.clear()
# ==================== HTTPS 管理器 ====================
class HTTPSManager:
"""HTTPS 证书管理器"""
def __init__(self, config: dict = None):
self.config = config or {}
self.cert_file = self.config.get('ssl_cert')
self.key_file = self.config.get('ssl_key')
self.context = None
def is_enabled(self) -> bool:
"""是否启用 HTTPS"""
return bool(self.cert_file and self.key_file)
def create_context(self) -> Optional[ssl.SSLContext]:
"""创建 SSL 上下文"""
if not self.is_enabled():
return None
try:
context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
# 加载证书
context.load_cert_chain(self.cert_file, self.key_file)
# 安全配置
context.minimum_version = ssl.TLSVersion.TLSv1_2
context.set_ciphers('ECDHE+AESGCM:DHE+AESGCM:ECDHE+CHACHA20')
self.context = context
return context
except Exception as e:
logger.error(f"创建 SSL 上下文失败: {e}")
return None
@staticmethod
def generate_self_signed_cert(cert_path: str, key_path: str, common_name: str = 'localhost') -> bool:
"""生成自签名证书(用于测试)"""
try:
from cryptography import x509
from cryptography.x509.oid import NameOID
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.backends import default_backend
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.hazmat.primitives import serialization
import datetime as dt
# 生成私钥
private_key = rsa.generate_private_key(
public_exponent=65537,
key_size=2048,
backend=default_backend()
)
# 生成证书
subject = issuer = x509.Name([
x509.NameAttribute(NameOID.COUNTRY_NAME, "CN"),
x509.NameAttribute(NameOID.STATE_OR_PROVINCE_NAME, "Shanghai"),
x509.NameAttribute(NameOID.LOCALITY_NAME, "Shanghai"),
x509.NameAttribute(NameOID.ORGANIZATION_NAME, "HYC Download Station"),
x509.NameAttribute(NameOID.COMMON_NAME, common_name),
])
cert = (
x509.CertificateBuilder()
.subject_name(subject)
.issuer_name(issuer)
.public_key(private_key.public_key())
.serial_number(x509.random_serial_number())
.not_valid_before(dt.datetime.utcnow())
.not_valid_after(dt.datetime.utcnow() + dt.timedelta(days=365))
.add_extension(
x509.SubjectAlternativeName([
x509.DNSName(common_name),
x509.DNSName("localhost"),
]),
critical=False,
)
.sign(private_key, hashes.SHA256(), default_backend())
)
# 保存证书
with open(cert_path, 'wb') as f:
f.write(cert.public_bytes(serialization.Encoding.PEM))
# 保存私钥
with open(key_path, 'wb') as f:
f.write(private_key.private_bytes(
encoding=serialization.Encoding.PEM,
format=serialization.PrivateFormat.TraditionalOpenSSL,
encryption_algorithm=serialization.NoEncryption()
))
return True
except Exception as e:
logger.error(f"生成自签名证书失败: {e}")
return False
def get_status(self) -> dict:
"""获取状态"""
return {
'enabled': self.is_enabled(),
'cert_file': self.cert_file,
'key_file': self.key_file
}
# ==================== 安全中间件 ====================
class SecurityMiddleware:
"""安全中间件"""
def __init__(self, config: dict = None):
self.config = config or {}
# 初始化各组件
self.ip_manager = IPManager(self.config.get('ip', {}))
self.rate_limiter = RateLimiter(self.config.get('rate_limit', {}))
self.https_manager = HTTPSManager(self.config.get('ssl', {}))
def check_request(self, handler) -> Tuple[bool, str]:
"""
检查请求是否安全
返回: (是否通过, 原因)
"""
client_ip = self._get_client_ip(handler)
# IP 检查
allowed, reason = self.ip_manager.is_allowed(client_ip)
if not allowed:
return False, f"IP被阻止: {reason}"
# 速率限制
allowed, _ = self.rate_limiter.is_allowed(client_ip)
if not allowed:
return False, "请求过于频繁"
return True, "通过"
def _get_client_ip(self, handler) -> str:
"""获取客户端 IP"""
# 检查代理头
forwarded = handler.headers.get('X-Forwarded-For')
if forwarded:
return forwarded.split(',')[0].strip()
real_ip = handler.headers.get('X-Real-IP')
if real_ip:
return real_ip
return handler.client_address[0] if hasattr(handler, 'client_address') else 'unknown'
def get_security_headers(self) -> dict:
"""获取安全响应头"""
return {
'X-Content-Type-Options': 'nosniff',
'X-Frame-Options': 'SAMEORIGIN',
'X-XSS-Protection': '1; mode=block',
'Strict-Transport-Security': 'max-age=31536000; includeSubDomains',
'Content-Security-Policy': "default-src 'self'; script-src 'self' 'unsafe-inline'; style-src 'self' 'unsafe-inline';"
}
def get_status(self) -> dict:
"""获取安全状态"""
return {
'ip': self.ip_manager.get_status(),
'rate_limit': self.rate_limiter.get_status(),
'ssl': self.https_manager.get_status()
}
# ==================== 审计日志 ====================
class AuditLogger:
"""审计日志记录器"""
def __init__(self, config: dict = None):
self.config = config or {}
self.logs: Queue = Queue(maxsize=10000)
self.log_file = self.config.get('audit_log', 'audit.log')
self.enabled = self.config.get('enabled', True)
# 启动日志写入线程
if self.enabled:
self._start_writer()
def log(self, event_type: str, data: dict):
"""记录事件"""
if not self.enabled:
return
entry = {
'timestamp': time.time(),
'type': event_type,
**data
}
if self.logs.full():
self.logs.get() # 移除最旧的
self.logs.put(entry)
def _start_writer(self):
"""启动日志写入线程"""
def writer():
while True:
try:
entry = self.logs.get(timeout=1)
self._write_entry(entry)
except Exception:
continue
thread = threading.Thread(target=writer, daemon=True)
thread.start()
def _write_entry(self, entry: dict):
"""写入日志条目"""
try:
with open(self.log_file, 'a', encoding='utf-8') as f:
line = json.dumps(entry, ensure_ascii=False)
f.write(line + '\n')
except Exception as e:
logger.error(f"写入审计日志失败: {e}")
def get_recent_logs(self, event_type: str = None, limit: int = 100) -> List[dict]:
"""获取最近的日志"""
result = []
with self.logs.mutex:
for entry in list(self.logs.queue):
if event_type and entry.get('type') != event_type:
continue
result.append(entry)
if len(result) >= limit:
break
return result[-limit:]
def get_status(self) -> dict:
"""获取状态"""
with self.logs.mutex:
return {
'enabled': self.enabled,
'log_file': self.log_file,
'pending_logs': self.logs.qsize(),
'max_size': self.logs.maxsize
}
# ==================== 便捷函数 ====================
def get_security_middleware(config: dict = None) -> SecurityMiddleware:
"""获取安全中间件"""
return SecurityMiddleware(config)
def get_ip_manager(config: dict = None) -> IPManager:
"""获取 IP 管理器"""
return IPManager(config)
def get_rate_limiter(config: dict = None) -> RateLimiter:
"""获取速率限制器"""
return RateLimiter(config)
+205
View File
@@ -0,0 +1,205 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""服务器核心模块 - 线程池版本"""
import os
import sys
import ssl
import signal
import mimetypes
import threading
import socketserver
from datetime import datetime
from concurrent.futures import ThreadPoolExecutor
from .config import ConfigManager
from .mirror_sync import MirrorSyncManager
class ThreadPoolHTTPServer(socketserver.ThreadingMixIn, socketserver.TCPServer):
"""使用线程池的 HTTP 服务器"""
allow_reuse_address = True
daemon_threads = True # 使用守护线程
def __init__(self, server_address, RequestHandlerClass, max_workers=50):
self.max_workers = max_workers
self._executor = ThreadPoolExecutor(
max_workers=max_workers,
thread_name_prefix="http_handler"
)
super().__init__(server_address, RequestHandlerClass)
def process_request(self, request, client_address):
"""使用线程池处理请求"""
self._executor.submit(self._handle_request, request, client_address)
def _handle_request(self, request, client_address):
"""实际处理请求"""
try:
self.finish_request(request, client_address)
except Exception:
self.handle_error(request, client_address)
finally:
self.shutdown_request(request)
def server_close(self):
"""关闭服务器和线程池"""
# Python 3.8 兼容处理
import sys
try:
self._executor.shutdown(wait=False, cancel_futures=True)
except TypeError:
# Python 3.8 不支持 cancel_futures 参数
self._executor.shutdown(wait=False)
super().server_close()
class MirrorServer:
"""镜像服务器主类"""
def __init__(self, config):
if isinstance(config, dict):
self.config_manager = ConfigManager(config)
else:
self.config_manager = config
self.config = self.config_manager.config
self.server = None
self.sync_manager = None
self.is_running = False
def _validate_config(self, config):
"""验证和修复配置"""
return self.config_manager._validate_config(config)
def start(self):
"""启动服务器"""
try:
# 创建下载目录
base_dir = self.config['base_dir']
if not os.path.exists(base_dir):
os.makedirs(base_dir)
print(f"创建下载目录: {os.path.abspath(base_dir)}")
# 初始化MIME类型
mimetypes.init()
# 记录启动时间
self.config['start_time'] = __import__('time').time()
# 创建同步管理器(仅当启用时)
if self.config.get('enable_sync', True):
self.sync_manager = MirrorSyncManager(self.config)
self.sync_manager.start()
# 创建系统监控器(仅当启用时)
self.monitor = None
if self.config.get('enable_monitor', True):
try:
from .monitor import SystemMonitor
self.monitor = SystemMonitor(self.config)
print(f" ✓ 系统监控已启用 (间隔: {self.config.get('monitor_interval', 5)}秒)")
except ImportError as e:
print(f" ✗ 系统监控导入失败: {e}")
except Exception as e:
print(f" ✗ 系统监控初始化失败: {e}")
# 延迟导入 handler(避免循环导入)
from handlers.http_handler import MirrorServerHandler
# 获取线程数配置
max_workers = min(self.config.get('max_workers', 10), 10) # 限制最大线程数
# 创建服务器
server_address = (self.config['host'], self.config['port'])
self.server = ThreadPoolHTTPServer(
server_address,
MirrorServerHandler,
max_workers=max_workers
)
# 传递配置到处理器
MirrorServerHandler.config = self.config
MirrorServerHandler.sync_manager = self.sync_manager
MirrorServerHandler.monitor = self.monitor
# 设置调试模式
MirrorServerHandler._setup_debug(self.config)
# 设置超时
self.server.timeout = self.config.get('timeout', 30)
# 启用HTTPS
if self.config.get('ssl_cert') and self.config.get('ssl_key'):
if not self._setup_ssl():
return False
self.is_running = True
return True
except Exception as e:
import traceback
print(f"服务器启动失败: {e}")
traceback.print_exc()
return False
def _setup_ssl(self):
"""设置SSL"""
try:
context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
context.load_cert_chain(self.config['ssl_cert'], self.config['ssl_key'])
self.server.socket = context.wrap_socket(self.server.socket, server_side=True)
print(f"启用HTTPS,证书: {self.config['ssl_cert']}")
return True
except Exception as e:
print(f"启用HTTPS失败: {e}")
return False
def stop(self):
"""停止服务器"""
print("正在停止服务器...")
if self.sync_manager:
try:
self.sync_manager.stop()
except Exception as e:
print(f"停止同步管理器时出错: {e}")
if self.server:
try:
self.server.server_close()
except Exception as e:
print(f"关闭服务器连接时出错: {e}")
self.is_running = False
print("服务器已停止")
def serve_forever(self):
"""运行服务器"""
if not self.server:
print("服务器未启动")
return
# 打印服务器信息(已在 main.py 中显示,此处仅保留最简信息)
protocol = "https" if self.config.get('ssl_cert') else "http"
sync_count = len(self.sync_manager.sync_sources) if self.sync_manager else 0
print(f"\n▶ 服务器运行于: {protocol}://{self.config['host']}:{self.config['port']}")
print(f"▶ 同步源数: {sync_count} | 最大线程: {self.server.max_workers}")
print("▶ 按 Ctrl+C 停止服务器")
try:
self.server.serve_forever()
except KeyboardInterrupt:
print("\n正在关闭服务器...")
finally:
self.stop()
# 全局变量,用于信号处理器访问服务器实例(预留)
# _server_instance = None
# def signal_handler(signum, frame):
# """处理退出信号"""
# print(f"\n收到信号 {signum},正在关闭服务器...")
# import os
# os._exit(0)
+1031
View File
File diff suppressed because it is too large Load Diff
+475
View File
@@ -0,0 +1,475 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
定时同步调度器
负责本地数据和数据库之间的定时同步
"""
import os
import sys
import time
import json
import hashlib
import threading
import logging
from datetime import datetime
from typing import Dict, List, Optional, Callable
from concurrent.futures import ThreadPoolExecutor
# 添加项目根目录到路径
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from core.database import DatabaseManager, get_db
from core.scheduler import Scheduler, ScheduledTask
logger = logging.getLogger(__name__)
class SyncScheduler:
"""同步调度器"""
def __init__(self, config: dict, db: DatabaseManager = None):
self.config = config
self.db = db
self._running = False
self._executor = ThreadPoolExecutor(max_workers=4)
# 同步配置
self.sync_interval = config.get('database', {}).get('sync_interval', 60)
self.auto_scan = config.get('auto_scan', True)
self.scan_interval = config.get('scan_interval', 300) # 5分钟扫描一次
# 同步状态
self.last_sync_time = 0
self.last_scan_time = 0
self.sync_in_progress = False
self.scan_in_progress = False
# 回调函数
self.on_file_added: Optional[Callable] = None
self.on_file_deleted: Optional[Callable] = None
self.on_file_updated: Optional[Callable] = None
self.on_sync_complete: Optional[Callable] = None
# 待同步队列
self._pending_add = [] # 待添加的文件
self._pending_update = [] # 待更新的文件
self._pending_delete = [] # 待删除的文件
# 定时任务调度器
self.task_scheduler = None
self.scheduled_syncs: Dict[str, dict] = {}
def _init_scheduled_syncs(self):
"""初始化定时同步任务"""
if not self.config.get('enable_sync', True):
return
# 从配置加载定时同步设置
sync_sources = self.config.get('sync_sources', {})
scheduled_sources = {}
for name, source_config in sync_sources.items():
schedule = source_config.get('schedule', {})
if schedule.get('enabled', False):
scheduled_sources[name] = {
'type': schedule.get('type', 'interval'), # 'cron' 或 'interval'
'config': {
'cron': schedule.get('cron'),
'interval': schedule.get('interval', {}),
'enabled': True
}
}
if scheduled_sources:
self.task_scheduler = Scheduler()
for name, sched_config in scheduled_sources.items():
self.task_scheduler.add_task(
name=f"sync_{name}",
task_type=sched_config['type'],
config=sched_config['config'],
callback=self._create_sync_callback(name)
)
self.scheduled_syncs = scheduled_sources
def _create_sync_callback(self, source_name: str):
"""创建同步回调函数"""
def sync_callback(task_name: str, config: dict):
logger.info(f"定时同步任务触发: {source_name}")
self.start_sync(source_name)
return True
return sync_callback
def start(self):
"""启动同步调度器"""
if self._running:
logger.warning("SyncScheduler 已经运行中")
return
self._running = True
self._executor.submit(self._sync_loop)
self._executor.submit(self._scan_loop)
# 初始化并启动定时同步
self._init_scheduled_syncs()
if self.task_scheduler:
self.task_scheduler.start()
logger.info(f"同步调度器已启动,间隔: {self.sync_interval}秒")
def stop(self):
"""停止同步调度器"""
self._running = False
# 停止定时任务调度器
if self.task_scheduler:
self.task_scheduler.stop()
self.task_scheduler = None
self._executor.shutdown(wait=True)
logger.info("同步调度器已停止")
def _sync_loop(self):
"""同步循环"""
while self._running:
try:
if time.time() - self.last_sync_time >= self.sync_interval:
self.perform_sync()
time.sleep(1)
except Exception as e:
logger.error(f"同步循环错误: {e}")
time.sleep(5)
def _scan_loop(self):
"""扫描循环 - 检测本地文件变化"""
while self._running:
try:
if self.auto_scan and time.time() - self.last_scan_time >= self.scan_interval:
self.scan_local_files()
time.sleep(5)
except Exception as e:
logger.error(f"扫描循环错误: {e}")
time.sleep(10)
def queue_add(self, file_info: dict):
"""队列添加文件"""
self._pending_add.append(file_info)
def queue_update(self, file_info: dict):
"""队列更新文件"""
self._pending_update.append(file_info)
def queue_delete(self, file_id: str):
"""队列删除文件"""
self._pending_delete.append(file_id)
def perform_sync(self):
"""执行同步"""
if self.sync_in_progress:
logger.warning("同步已在进行中,跳过")
return
self.sync_in_progress = True
start_time = time.time()
try:
logger.info("开始执行数据库同步...")
# 同步待添加的文件
added = 0
for file_info in self._pending_add[:]:
try:
self._sync_add_file(file_info)
self._pending_add.remove(file_info)
added += 1
except Exception as e:
logger.error(f"同步添加文件失败: {e}")
# 同步待更新的文件
updated = 0
for file_info in self._pending_update[:]:
try:
self._sync_update_file(file_info)
self._pending_update.remove(file_info)
updated += 1
except Exception as e:
logger.error(f"同步更新文件失败: {e}")
# 同步待删除的文件
deleted = 0
for file_id in self._pending_delete[:]:
try:
self._sync_delete_file(file_id)
self._pending_delete.remove(file_id)
deleted += 1
except Exception as e:
logger.error(f"同步删除文件失败: {e}")
# 同步统计
self.db.reset_pending_count()
self.last_sync_time = time.time()
duration = time.time() - start_time
logger.info(f"同步完成: 添加{added}, 更新{updated}, 删除{deleted}, 耗时{duration:.2f}秒")
# 回调
if self.on_sync_complete:
self.on_sync_complete({
'added': added,
'updated': updated,
'deleted': deleted,
'duration': duration
})
except Exception as e:
logger.error(f"同步过程错误: {e}")
finally:
self.sync_in_progress = False
def scan_local_files(self):
"""扫描本地文件"""
if self.scan_in_progress:
return
self.scan_in_progress = True
try:
base_dir = self.config.get('base_dir', './downloads')
if not os.path.exists(base_dir):
self.last_scan_time = time.time()
return
# 扫描文件
scanned_files = []
for root, dirs, files in os.walk(base_dir):
for filename in files:
full_path = os.path.join(root, filename)
rel_path = os.path.relpath(full_path, base_dir).replace("\\", "/")
stat = os.stat(full_path)
file_info = {
'path': rel_path,
'name': filename,
'size': stat.st_size,
'mtime': stat.st_mtime,
'ctime': stat.st_ctime
}
scanned_files.append(file_info)
# 与数据库对比
db_files = self.db.list_files(limit=100000)
db_paths = {f.path for f in db_files if not f.is_dir}
# 检测新增
local_paths = {f['path'] for f in scanned_files}
new_paths = local_paths - db_paths
for path in new_paths:
file_info = next((f for f in scanned_files if f['path'] == path), None)
if file_info:
file_id = hashlib.md5(path.encode()).hexdigest()
self._sync_add_file({
'file_id': file_id,
'path': path,
'name': file_info['name'],
'size': file_info['size'],
'updated_at': file_info['mtime']
})
# 检测删除
deleted_paths = db_paths - local_paths
for path in deleted_paths:
record = self.db.get_file_by_path(path)
if record:
self.db.delete_file(record.file_id)
self.last_scan_time = time.time()
except Exception as e:
logger.error(f"扫描本地文件错误: {e}")
finally:
self.scan_in_progress = False
def _sync_add_file(self, file_info: dict):
"""同步添加文件"""
existing = self.db.get_file_by_path(file_info['path'])
if existing:
# 已存在,更新
self.db.update_file(
existing.file_id,
size=file_info.get('size', 0),
updated_at=file_info.get('updated_at', time.time()),
hash=file_info.get('hash'),
sync_status='synced'
)
else:
# 新增
file_id = file_info.get('file_id') or hashlib.md5(
file_info['path'].encode()
).hexdigest()
self.db.add_file(
file_id=file_id,
path=file_info['path'],
name=file_info['name'],
size=file_info.get('size', 0),
hash=file_info.get('hash'),
is_dir=False,
created_at=file_info.get('created_at'),
updated_at=file_info.get('updated_at', time.time())
)
if self.on_file_added:
self.on_file_added(file_info)
def _sync_update_file(self, file_info: dict):
"""同步更新文件"""
file_id = file_info.get('file_id')
if file_id:
self.db.update_file(
file_id,
size=file_info.get('size'),
updated_at=file_info.get('updated_at', time.time()),
hash=file_info.get('hash'),
sync_status='synced'
)
if self.on_file_updated:
self.on_file_updated(file_info)
def _sync_delete_file(self, file_id: str):
"""同步删除文件"""
self.db.delete_file(file_id)
if self.on_file_deleted:
self.on_file_deleted({'file_id': file_id})
def get_status(self) -> dict:
"""获取同步状态"""
return {
'running': self._running,
'last_sync_time': self.last_sync_time,
'last_scan_time': self.last_scan_time,
'sync_in_progress': self.sync_in_progress,
'scan_in_progress': self.scan_in_progress,
'pending_add': len(self._pending_add),
'pending_update': len(self._pending_update),
'pending_delete': len(self._pending_delete),
'pending_operations': self.db.get_pending_operations() if self.db else 0
}
def force_sync(self):
"""强制立即同步"""
self.last_sync_time = 0
self.perform_sync()
# ==================== 文件操作包装器 ====================
class DatabaseBackedFileOperations:
"""数据库支持的文件操作"""
def __init__(self, config: dict, db: DatabaseManager, scheduler: SyncScheduler = None):
self.config = config
self.db = db
self.scheduler = scheduler
self.base_dir = config.get('base_dir', './downloads')
def add_file_record(self, path: str, name: str, size: int = 0,
hash: str = None, is_dir: bool = False) -> dict:
"""添加文件记录到数据库"""
import hashlib
file_id = hashlib.md5(path.encode()).hexdigest()
file_info = {
'file_id': file_id,
'path': path,
'name': name,
'size': size,
'hash': hash,
'is_dir': is_dir,
'created_at': time.time(),
'updated_at': time.time()
}
if self.scheduler:
self.scheduler.queue_add(file_info)
else:
self.db.add_file(
file_id=file_id,
path=path,
name=name,
size=size,
hash=hash,
is_dir=is_dir,
created_at=time.time(),
updated_at=time.time()
)
return file_info
def update_file_record(self, file_id: str, **kwargs):
"""更新文件记录"""
if self.scheduler:
self.scheduler.queue_update({'file_id': file_id, **kwargs})
else:
self.db.update_file(file_id, **kwargs)
def delete_file_record(self, file_id: str, hard: bool = False):
"""删除文件记录"""
if self.scheduler:
self.scheduler.queue_delete(file_id)
else:
self.db.delete_file(file_id, hard=hard)
def record_download(self, file_path: str, file_size: int = 0,
client_ip: str = None, duration: float = 0,
success: bool = True, error_message: str = None):
"""记录下载"""
self.db.add_download_record(
file_path=file_path,
file_size=file_size,
client_ip=client_ip,
duration=duration,
success=success,
error_message=error_message
)
# 更新下载计数
record = self.db.get_file_by_path(file_path)
if record:
self.db.increment_download_count(record.file_id)
def record_cache_hit(self, cache_key: str, cache_type: str):
"""记录缓存命中"""
record = self.db.get_cache_record(cache_key)
if record:
self.db.increment_cache_hits(cache_key)
else:
self.db.add_cache_record(
cache_key=cache_key,
cache_type=cache_type,
hits=1,
last_hit=time.time()
)
# ==================== 便捷函数 ====================
def get_sync_scheduler(config: dict) -> SyncScheduler:
"""获取同步调度器"""
db = get_db(config)
return SyncScheduler(config, db)
def init_database_sync(config: dict, db=None) -> tuple:
"""初始化数据库和同步"""
if db is None:
db = get_db(config)
scheduler = SyncScheduler(config, db)
file_ops = DatabaseBackedFileOperations(config, db, scheduler)
return db, scheduler, file_ops
+94
View File
@@ -0,0 +1,94 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""工具函数模块"""
import os
import re
import hashlib
def format_file_size(size_bytes: int) -> str:
"""格式化文件大小为可读格式"""
if size_bytes == 0:
return "0 B"
size_names = ["B", "KB", "MB", "GB", "TB", "PB"]
i = 0
while size_bytes >= 1024 and i < len(size_names) - 1:
size_bytes /= 1024.0
i += 1
return f"{size_bytes:.2f} {size_names[i]}"
def get_file_hash(filepath: str, algorithm: str = 'sha256') -> str:
"""计算文件哈希值"""
if not filepath or not isinstance(filepath, str):
return "Error: Invalid filepath"
if not os.path.exists(filepath):
return "Error: File not found"
try:
hash_func = hashlib.new(algorithm)
with open(filepath, 'rb') as f:
for chunk in iter(lambda: f.read(4096), b""):
hash_func.update(chunk)
return hash_func.hexdigest()
except Exception as e:
return f"Error: {str(e)}"
def parse_size(size_str: str) -> int:
"""解析文件大小字符串为字节数"""
size_str = size_str.upper().strip()
match = re.match(r'^(\d+(?:\.\d+)?)\s*([BKMGT]?)B?$', size_str)
if not match:
raise ValueError(f"无效的文件大小格式: {size_str}")
number = float(match.group(1))
unit = match.group(2) or 'B'
units = {
'B': 1,
'K': 1024,
'M': 1024 ** 2,
'G': 1024 ** 3,
'T': 1024 ** 4
}
if unit not in units:
raise ValueError(f"无效的单位: {unit}")
return int(number * units[unit])
def sanitize_filename(filename: str) -> str:
"""清理文件名,防止路径遍历和安全问题"""
filename = os.path.basename(filename) # 移除路径分隔符
# 移除危险字符
filename = re.sub(r'[<>:"|?*\\\x00-\x1f]', '_', filename)
# 限制长度
if len(filename) > 255:
name, ext = os.path.splitext(filename)
filename = name[:255 - len(ext)] + ext
# 防止空文件名
if not filename or filename in ('.', '..'):
filename = 'unnamed_file'
return filename
def is_safe_path(base_dir: str, path: str) -> bool:
"""检查路径是否安全(防止目录遍历)"""
try:
abs_path = os.path.abspath(path)
abs_base = os.path.abspath(base_dir)
common_path = os.path.commonpath([abs_path, abs_base])
return common_path == abs_base
except ValueError:
return False