Baseline: pr1 HYC下载站 v2.3 before security/functional fixes
This commit is contained in:
@@ -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
@@ -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'
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
+376
@@ -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')
|
||||
}
|
||||
@@ -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']
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user