Files
SenSu/services/update_service.py
T

264 lines
10 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
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}")
# 0. 备份当前版本
self.status["progress"] = 2
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.status["progress"] = 40
# 2. 解压到临时目录
tmp_dir = tempfile.mkdtemp()
with zipfile.ZipFile(tmp_zip.name, "r") as zf:
zf.extractall(tmp_dir)
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.status["progress"] = 80
# 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:
logger.warning(f"依赖安装警告: {r.stderr[-300:]}")
self.status["progress"] = 100
return {"ok": True, "msg": "更新完成, 框架将重启"}
except Exception as 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)
async def check_now(self):
await self._check_version()
async def stop(self):
self._running = False
if self._task:
self._task.cancel()