#!/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: """从 version.json 读取本地版本号""" try: vf = _PROJECT_ROOT / "version.json" if vf.exists(): with open(vf) as f: data = json.load(f) return data.get("framework", {}).get("version", "0.0.0").lstrip("vV") except Exception: pass 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: # 同步版本号到 version.json vf = _PROJECT_ROOT / "version.json" if vf.exists(): with open(vf) as f: vdata = json.load(f) new_ver = vdata.get("framework", {}).get("version", "") if new_ver: 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()