310 lines
13 KiB
Python
310 lines
13 KiB
Python
#!/usr/bin/env python3
|
|
"""版本更新服务 — 下载 release ZIP + 解压 + 保护配置 + 热重启"""
|
|
import logging
|
|
import asyncio
|
|
import os
|
|
import shutil
|
|
import subprocess
|
|
import tempfile
|
|
import urllib.request
|
|
import zipfile
|
|
import json
|
|
import time
|
|
from pathlib import Path
|
|
from typing import Optional
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
CHECK_INTERVAL = 3600 * 6 # 6 小时
|
|
_PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
|
_RESTART_FLAG = _PROJECT_ROOT / ".restart_flag"
|
|
|
|
# 更新时跳过的路径 (用户配置/数据)
|
|
_PRESERVE_PATHS = {
|
|
"config/framework/base_config.yaml",
|
|
"config/permissions/",
|
|
"data/",
|
|
"plugins/*/config.yaml",
|
|
"logs/",
|
|
".restart_flag",
|
|
}
|
|
|
|
|
|
def _parse_ver(v: str) -> tuple:
|
|
v = v.strip().lstrip("vV")
|
|
try:
|
|
return tuple(int(p) for p in v.split(".")[:3])
|
|
except Exception:
|
|
return (0, 0, 0)
|
|
|
|
|
|
_BACKUP_SKIP = {"data/", "logs/", "__pycache__/", ".git/", "node_modules/", ".venv/", "venv/",
|
|
"*.pyc", "*.pyo", ".restart_flag"}
|
|
|
|
|
|
def _backup_current(dest_path: Path):
|
|
"""将当前项目打包为 ZIP (跳过数据/日志/缓存)"""
|
|
from fnmatch import fnmatch
|
|
with zipfile.ZipFile(str(dest_path), "w", zipfile.ZIP_DEFLATED) as zf:
|
|
for root, dirs, files in os.walk(str(_PROJECT_ROOT)):
|
|
# 跳过不需要备份的目录
|
|
dirs[:] = [d for d in dirs if not any(
|
|
fnmatch(d + "/", p.rstrip("/") + "/") or fnmatch(d, p.rstrip("/"))
|
|
for p in _BACKUP_SKIP if "/" in p or p.endswith("/")
|
|
) and d not in {p.rstrip("/") for p in _BACKUP_SKIP if not any(c in p for c in "*?[")}]
|
|
for f in files:
|
|
rel = os.path.relpath(os.path.join(root, f), str(_PROJECT_ROOT))
|
|
if any(fnmatch(f, p) or fnmatch(rel, p) for p in _BACKUP_SKIP):
|
|
continue
|
|
zf.write(os.path.join(root, f), rel)
|
|
|
|
|
|
def _should_preserve(rel_path: str) -> bool:
|
|
"""检查路径是否应跳过 (保护用户数据)"""
|
|
from fnmatch import fnmatch
|
|
for pattern in _PRESERVE_PATHS:
|
|
if fnmatch(rel_path, pattern) or rel_path.startswith(pattern.rstrip("/") + "/"):
|
|
return True
|
|
# 匹配目录前缀
|
|
if rel_path == pattern.rstrip("/"):
|
|
return True
|
|
return False
|
|
|
|
|
|
class UpdateService:
|
|
def __init__(self, service_manager):
|
|
self.sm = service_manager
|
|
self._task: Optional[asyncio.Task] = None
|
|
self._running = False
|
|
# 日志目录
|
|
import logging as _logging
|
|
self.update_logger = _logging.getLogger("senSu.update")
|
|
log_dir = _PROJECT_ROOT / "logs" / "updates_logs"
|
|
log_dir.mkdir(parents=True, exist_ok=True)
|
|
fh = _logging.FileHandler(
|
|
log_dir / f"update_{time.strftime('%Y%m%d_%H%M%S')}.log",
|
|
encoding='utf-8'
|
|
)
|
|
fh.setFormatter(_logging.Formatter('%(asctime)s [%(levelname)s] %(message)s'))
|
|
self.update_logger.addHandler(fh)
|
|
self.update_logger.setLevel(_logging.INFO)
|
|
self.status: dict = {"state": "idle", "progress": 0, "error": "",
|
|
"last_check": 0, "checked": False, "update_available": False,
|
|
"current_version": "", "latest_version": ""}
|
|
# 从配置读取分支和仓库地址
|
|
try:
|
|
init = self.sm.get_service("init")
|
|
cfg = init.get_config("base").get("framework", {})
|
|
except Exception:
|
|
cfg = {}
|
|
self.branch = cfg.get("branch", "main")
|
|
self.repo_url = cfg.get("repo_url", "https://git.yeij.top/AskaEth/SenSu").rstrip("/")
|
|
self.repo_name = self.repo_url.rstrip("/").split("/")[-1]
|
|
# ZIP 下载地址: Gitea 格式
|
|
self.download_url = f"{self.repo_url}/archive/{self.branch}.zip"
|
|
# 版本检查 API
|
|
self.version_url = f"{self.repo_url}/raw/branch/{self.branch}/version.json"
|
|
|
|
async def start(self):
|
|
self._running = True
|
|
self.status["current_version"] = self._current_version()
|
|
self._task = asyncio.create_task(self._loop())
|
|
logger.info(f"🔄 更新服务已启动 — 分支: {self.branch}, 仓库: {self.repo_url}")
|
|
|
|
async def _loop(self):
|
|
await asyncio.sleep(30) # 等框架就绪
|
|
while self._running:
|
|
await self._check_version()
|
|
await asyncio.sleep(CHECK_INTERVAL)
|
|
|
|
async def _check_version(self):
|
|
"""检查是否有新版本 (非阻塞)"""
|
|
try:
|
|
loop = asyncio.get_event_loop()
|
|
data = await loop.run_in_executor(None, self._sync_fetch_json, self.version_url)
|
|
self.status["last_check"] = time.time()
|
|
self.status["checked"] = True
|
|
current = self._current_version()
|
|
self.status["current_version"] = current
|
|
if not data:
|
|
return
|
|
remote_fw = data.get("framework", {})
|
|
remote_ver = remote_fw.get("version", "")
|
|
if not remote_ver:
|
|
return
|
|
if _parse_ver(remote_ver) > _parse_ver(current):
|
|
self.status["update_available"] = True
|
|
self.status["latest_version"] = remote_ver
|
|
self.status["download_url"] = self.download_url
|
|
logger.info(f"🔄 新版本可用: {remote_ver} (当前: {current})")
|
|
except Exception as e:
|
|
logger.debug(f"版本检查失败: {e}")
|
|
|
|
def _sync_fetch_json(self, url: str) -> Optional[dict]:
|
|
try:
|
|
req = urllib.request.Request(url, headers={"User-Agent": "SenSu-Update/1.0"})
|
|
with urllib.request.urlopen(req, timeout=15) as resp:
|
|
return json.loads(resp.read().decode())
|
|
except Exception:
|
|
return None
|
|
|
|
def _current_version(self) -> str:
|
|
try:
|
|
init = self.sm.get_service("init")
|
|
return init.get_config("base").get("framework", {}).get("version", "0.0.0").lstrip("vV")
|
|
except Exception:
|
|
return "0.0.0"
|
|
|
|
# ── 执行更新 ──
|
|
|
|
async def apply_update(self) -> dict:
|
|
"""下载并应用更新 (在线程池中执行磁盘操作)"""
|
|
if self.status["state"] == "downloading":
|
|
return {"ok": False, "error": "更新已在进行中"}
|
|
|
|
self.status = {"state": "downloading", "progress": 0, "error": "",
|
|
"last_check": self.status.get("last_check", 0),
|
|
"update_available": self.status.get("update_available", False)}
|
|
try:
|
|
loop = asyncio.get_event_loop()
|
|
result = await loop.run_in_executor(None, self._sync_apply)
|
|
if result.get("ok"):
|
|
self.status = {"state": "done", "progress": 100, "error": "",
|
|
"last_check": time.time(), "update_available": False}
|
|
# 写重启标记
|
|
_RESTART_FLAG.touch()
|
|
logger.info("✅ 更新完成, 框架将在下次主循环检测后重启")
|
|
else:
|
|
self.status["state"] = "error"
|
|
self.status["error"] = result.get("error", "未知错误")
|
|
return result
|
|
except Exception as e:
|
|
self.status["state"] = "error"
|
|
self.status["error"] = str(e)
|
|
return {"ok": False, "error": str(e)}
|
|
|
|
def _sync_apply(self) -> dict:
|
|
"""同步执行更新 (在线程池中)"""
|
|
tmp_zip = None
|
|
tmp_dir = None
|
|
try:
|
|
url = self.download_url
|
|
logger.info(f"📥 下载更新包: {url}")
|
|
|
|
self._log("开始更新...")
|
|
# 0. 备份当前版本
|
|
self.status["progress"] = 2
|
|
self._log("创建备份...")
|
|
backup_name = f"SenSu_backup_{self._current_version()}_{time.strftime('%Y%m%d_%H%M%S')}.zip"
|
|
backup_dir = _PROJECT_ROOT / "data" / "backups"
|
|
backup_dir.mkdir(parents=True, exist_ok=True)
|
|
backup_path = backup_dir / backup_name
|
|
_backup_current(backup_path)
|
|
logger.info(f"📦 已备份到: {backup_path}")
|
|
# 清理旧备份 (保留最新 5 个)
|
|
backups = sorted(backup_dir.glob("SenSu_backup_*.zip"), key=os.path.getmtime, reverse=True)
|
|
for old in backups[5:]:
|
|
old.unlink()
|
|
|
|
# 1. 下载 ZIP
|
|
self.status["progress"] = 10
|
|
tmp_zip = tempfile.NamedTemporaryFile(suffix=".zip", delete=False)
|
|
req = urllib.request.Request(url, headers={"User-Agent": "SenSu-Update/1.0"})
|
|
with urllib.request.urlopen(req, timeout=300) as resp:
|
|
shutil.copyfileobj(resp, tmp_zip)
|
|
tmp_zip.close()
|
|
self._log("下载完成")
|
|
self.status["progress"] = 40
|
|
|
|
# 2. 解压到临时目录
|
|
tmp_dir = tempfile.mkdtemp()
|
|
with zipfile.ZipFile(tmp_zip.name, "r") as zf:
|
|
zf.extractall(tmp_dir)
|
|
self._log("解压更新包...")
|
|
self.status["progress"] = 60
|
|
|
|
# 3. Gitea zip 内层目录: 通常是 repo_name-branch/
|
|
extracted_root = tmp_dir
|
|
contents = os.listdir(tmp_dir)
|
|
if len(contents) == 1:
|
|
inner = os.path.join(tmp_dir, contents[0])
|
|
if os.path.isdir(inner):
|
|
extracted_root = inner
|
|
|
|
# 4. 覆盖文件 (跳过用户配置)
|
|
for root, dirs, files in os.walk(extracted_root):
|
|
rel = os.path.relpath(root, extracted_root)
|
|
if rel == ".":
|
|
rel = ""
|
|
for f in files:
|
|
fp = os.path.join(rel, f) if rel else f
|
|
if _should_preserve(fp):
|
|
continue
|
|
src = os.path.join(root, f)
|
|
dst = os.path.join(str(_PROJECT_ROOT), fp)
|
|
os.makedirs(os.path.dirname(dst), exist_ok=True)
|
|
shutil.copy2(src, dst)
|
|
|
|
self._log("覆盖文件 (跳过用户配置)...")
|
|
self.status["progress"] = 80
|
|
|
|
# 4.5 同步版本号到旧配置文件
|
|
try:
|
|
old_cfg_path = str(_PROJECT_ROOT / "config" / "framework" / "base_config.yaml")
|
|
new_tpl_path = str(_PROJECT_ROOT / "config" / "framework" / "base_config.yaml.example")
|
|
if os.path.exists(old_cfg_path) and os.path.exists(new_tpl_path):
|
|
import yaml as _yaml
|
|
with open(new_tpl_path) as f:
|
|
new_tpl = _yaml.safe_load(f)
|
|
new_ver = new_tpl.get("framework", {}).get("version", "")
|
|
if new_ver:
|
|
with open(old_cfg_path) as f:
|
|
old_cfg = _yaml.safe_load(f) or {}
|
|
old_cfg.setdefault("framework", {})["version"] = new_ver
|
|
with open(old_cfg_path, "w") as f:
|
|
_yaml.dump(old_cfg, f, default_flow_style=False, allow_unicode=True)
|
|
self._log(f"版本号已同步: {new_ver}")
|
|
logger.info(f"📝 版本号已同步: {new_ver}")
|
|
except Exception as e:
|
|
logger.warning(f"版本号同步失败: {e}")
|
|
|
|
self._log("安装依赖...")
|
|
# 5. 安装依赖 (子进程)
|
|
req_file = str(_PROJECT_ROOT / "requirements.txt")
|
|
if os.path.exists(req_file):
|
|
r = subprocess.run(
|
|
["pip", "install", "-r", req_file, "--break-system-packages"],
|
|
capture_output=True, text=True, timeout=120,
|
|
cwd=str(_PROJECT_ROOT),
|
|
)
|
|
if r.returncode != 0:
|
|
self._log(f"依赖安装警告: {r.stderr[-200:]}")
|
|
logger.warning(f"依赖安装警告: {r.stderr[-300:]}")
|
|
|
|
self.status["progress"] = 100
|
|
self._log("更新完成, 框架将重启")
|
|
return {"ok": True, "msg": "更新完成, 框架将重启"}
|
|
|
|
except Exception as e:
|
|
self._log(f"更新失败: {e}")
|
|
logger.error(f"更新失败: {e}")
|
|
return {"ok": False, "error": str(e)}
|
|
finally:
|
|
if tmp_zip and os.path.exists(tmp_zip.name):
|
|
os.unlink(tmp_zip.name)
|
|
if tmp_dir and os.path.exists(tmp_dir):
|
|
shutil.rmtree(tmp_dir, ignore_errors=True)
|
|
|
|
def _log(self, msg: str):
|
|
"""写入更新日志 + 更新SSE状态"""
|
|
self.update_logger.info(msg)
|
|
self.status["log"] = self.status.get("log", "") + msg + "\n"
|
|
|
|
async def check_now(self):
|
|
await self._check_version()
|
|
|
|
async def stop(self):
|
|
self._running = False
|
|
if self._task:
|
|
self._task.cancel()
|