// Package upgrade 实现 gitcat 的一键升级:检查更新、下载、校验并就地替换 // 当前运行的二进制文件。 // // 设计上有一条硬约束:**绝不在新二进制被验证可用之前动现有的那个**。 // 升级失败时,服务必须还能照常启动。因此流程是 // // 拉清单 → 选版本 → 下载到临时文件 → 校验 sha256 → 试运行 → 备份旧文件 // → 替换 → 写 pending 标记 → 重启 // // 中间任何一步失败都直接返回错误,旧二进制保持原样。 package upgrade import ( "context" "crypto/sha256" "encoding/hex" "encoding/json" "errors" "fmt" "io" "net/http" "net/url" "os" "os/exec" "path/filepath" "runtime" "strconv" "strings" "time" ) const ( // maxManifestSize 限制升级清单大小,避免被超大响应撑爆内存。 maxManifestSize = 1 << 20 // maxBinarySize 限制下载的二进制大小(200MB),同理。 maxBinarySize = 200 << 20 // fetchTimeout 是拉取清单 / 下载的默认上限。 fetchTimeout = 10 * time.Minute ) // Info 描述当前正在运行的二进制。 type Info struct { Version string // 语义化版本,如 1.2.0 Commit string // 构建时注入的提交号 BuildDate string // 构建时间 ExePath string // 可执行文件绝对路径 GOOS string GOARCH string } // Release 是发布清单中的一个版本。 // // 清单由管理员在后台配置的一个 URL 提供,格式示例: // // { // "latest": "1.2.0", // "min_supported": "1.0.0", // "releases": [ // {"version": "1.2.0", "url": "https://…/gitcat-linux-amd64", // "sha256": "…", "size": 8388608, "os": "linux", "arch": "amd64", // "notes": "修复推送越权"} // ] // } type Release struct { Version string `json:"version"` URL string `json:"url"` SHA256 string `json:"sha256"` Size int64 `json:"size"` OS string `json:"os"` Arch string `json:"arch"` Notes string `json:"notes"` } // Manifest 是升级源提供的版本清单。 type Manifest struct { Latest string `json:"latest"` // MinSupported 是允许直接跳到的最低版本:低于它的版本需要先中转升级, // 避免一次跨度过大的结构变更。 MinSupported string `json:"min_supported"` Releases []Release `json:"releases"` } // CheckResult 是一次检查更新的结论。 type CheckResult struct { Current string Latest string Upgradable bool Release *Release Manifest *Manifest Message string CheckedAt time.Time AssetExists bool } // InstallResult 描述一次替换的结果。 type InstallResult struct { FromVersion string ToVersion string ExePath string BackupPath string PendingFile string Restarted bool Note string } // platformKey 返回当前平台在清单里的标识。 func platformKey() string { return runtime.GOOS + "/" + runtime.GOARCH } // IsDev 判断当前是否是开发版(未注入正式版本号)。 func (i Info) IsDev() bool { v := strings.ToLower(strings.TrimSpace(i.Version)) return v == "" || v == "dev" || v == "devel" || strings.Contains(v, "-dirty") || strings.Contains(v, "-dev") } // ShortSHA 返回短提交号。 func (i Info) ShortSHA() string { if len(i.Commit) > 7 { return i.Commit[:7] } return i.Commit } // Compare 比较两个语义化版本:ab 返回 1,相等返回 0。 // // 允许带 "v" 前缀,比较时忽略;带预发布后缀(1.0.0-rc1)视为低于 1.0.0。 func Compare(a, b string) int { na, pa := splitVersion(a) nb, pb := splitVersion(b) for i := 0; i < 3; i++ { if na[i] != nb[i] { if na[i] < nb[i] { return -1 } return 1 } } switch { case pa == pb: return 0 case pa == "": return 1 // 正式版 > 预发布版 case pb == "": return -1 default: return strings.Compare(pa, pb) } } func splitVersion(v string) ([3]int, string) { var out [3]int v = strings.TrimSpace(v) v = strings.TrimPrefix(strings.TrimPrefix(v, "v"), "V") pre := "" if i := strings.IndexAny(v, "-+"); i >= 0 { pre = v[i+1:] v = v[:i] } for i, part := range strings.SplitN(v, ".", 3) { if i > 2 { break } n, err := strconv.Atoi(strings.TrimSpace(part)) if err != nil { continue } out[i] = n } return out, pre } // ValidVersion 判断字符串是否像一个可比较的版本号。 func ValidVersion(v string) bool { v = strings.TrimSpace(v) if v == "" { return false } nums, _ := splitVersion(v) return nums[0] > 0 || nums[1] > 0 || nums[2] > 0 } // httpClient 是升级流程专用的客户端:短超时、限制重定向。 func httpClient(timeout time.Duration) *http.Client { return &http.Client{ Timeout: timeout, CheckRedirect: func(req *http.Request, via []*http.Request) error { if len(via) >= 5 { return errors.New("重定向次数过多") } return nil }, } } // FetchManifest 拉取并解析版本清单。 func FetchManifest(ctx context.Context, manifestURL string) (*Manifest, error) { manifestURL = strings.TrimSpace(manifestURL) if manifestURL == "" { return nil, errors.New("未配置升级源") } u, err := url.Parse(manifestURL) if err != nil || !strings.HasPrefix(u.Scheme, "http") { return nil, errors.New("升级源必须是 http/https 地址") } ctx, cancel := context.WithTimeout(ctx, 60*time.Second) defer cancel() req, err := http.NewRequestWithContext(ctx, http.MethodGet, manifestURL, nil) if err != nil { return nil, err } req.Header.Set("Accept", "application/json") resp, err := httpClient(60 * time.Second).Do(req) if err != nil { return nil, err } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("升级源返回 HTTP %d", resp.StatusCode) } raw, err := io.ReadAll(io.LimitReader(resp.Body, maxManifestSize)) if err != nil { return nil, err } var m Manifest if err := json.Unmarshal(raw, &m); err != nil { return nil, fmt.Errorf("升级源返回的内容不是合法 JSON: %w", err) } if len(m.Releases) == 0 { return nil, errors.New("升级源里没有任何版本") } return &m, nil } // pickRelease 从清单中挑出适用于当前平台的目标版本。 func pickRelease(m *Manifest, version string) *Release { want := strings.TrimSpace(version) if want == "" { want = strings.TrimSpace(m.Latest) } for i := range m.Releases { r := m.Releases[i] if !strings.EqualFold(strings.TrimSpace(r.Version), want) { continue } if r.OS != "" && r.Arch != "" && r.OS+"/"+r.Arch != platformKey() { return nil // 清单里明确标注了其它平台 } rr := r return &rr } return nil } // Check 对比当前版本与升级源,给出升级结论。 func (i Info) Check(ctx context.Context, manifestURL string) (*CheckResult, error) { m, err := FetchManifest(ctx, manifestURL) if err != nil { return nil, err } res := &CheckResult{ Current: i.Version, Latest: m.Latest, Manifest: m, CheckedAt: time.Now(), } rel := pickRelease(m, m.Latest) if rel == nil { res.Message = fmt.Sprintf("升级源未提供适用于 %s 的版本 %s", platformKey(), m.Latest) return res, nil } res.AssetExists = true res.Release = rel switch { case i.IsDev(): res.Message = "当前是开发构建,无法自动判断新版本,请对照发布页确认。" case rel.Version == "" || !ValidVersion(rel.Version): res.Message = "升级源中的版本号不合法。" case Compare(i.Version, rel.Version) >= 0: res.Message = "已经是最新版本。" case m.MinSupported != "" && ValidVersion(m.MinSupported) && Compare(i.Version, m.MinSupported) < 0: res.Message = fmt.Sprintf("当前版本过旧,需先升级到 %s 或更高版本再继续。", m.MinSupported) default: res.Upgradable = true res.Message = fmt.Sprintf("发现新版本 %s", rel.Version) } return res, nil } // Download 下载指定版本到 dest,返回校验后的文件大小。 func Download(ctx context.Context, rel *Release, dest string) (int64, error) { if rel == nil || strings.TrimSpace(rel.URL) == "" { return 0, errors.New("该版本没有提供下载地址") } u, err := url.Parse(rel.URL) if err != nil || !strings.HasPrefix(u.Scheme, "http") { return 0, errors.New("下载地址必须是 http/https 地址") } ctx, cancel := context.WithTimeout(ctx, fetchTimeout) defer cancel() req, err := http.NewRequestWithContext(ctx, http.MethodGet, rel.URL, nil) if err != nil { return 0, err } req.Header.Set("User-Agent", "gitcat-upgrade") resp, err := httpClient(fetchTimeout).Do(req) if err != nil { return 0, err } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { return 0, fmt.Errorf("下载失败:HTTP %d", resp.StatusCode) } limit := int64(maxBinarySize) if rel.Size > 0 && rel.Size < limit { limit = rel.Size } f, err := os.OpenFile(dest, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o755) if err != nil { return 0, err } defer f.Close() h := sha256.New() n, err := io.Copy(io.MultiWriter(f, h), io.LimitReader(resp.Body, limit+1)) if err != nil { return 0, err } if n > limit { return 0, fmt.Errorf("下载内容超过预期体积(%d 字节),已中止", limit) } sum := hex.EncodeToString(h.Sum(nil)) if want := strings.ToLower(strings.TrimSpace(rel.SHA256)); want != "" && want != sum { return 0, fmt.Errorf("校验失败:期望 sha256 %s,实际 %s", want, sum) } if rel.Size > 0 && n != rel.Size { return 0, fmt.Errorf("校验失败:期望 %d 字节,实际 %d 字节", rel.Size, n) } if err := f.Sync(); err != nil { return 0, err } return n, nil } // Verify 对新下载的可执行文件做一次试运行,确认它真的能跑起来。 // // 这一步是"不损坏现网"的关键:架构不匹配、动态库缺失、文件被截断时, // 试运行会立刻失败,而此时旧二进制还没被碰过。 func Verify(binPath string) error { ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() cmd := exec.CommandContext(ctx, binPath, "-version") out, err := cmd.CombinedOutput() if err != nil { return fmt.Errorf("新版本无法执行(可能是平台不匹配或文件损坏):%v", err) } text := string(out) if !strings.Contains(text, "gitcat") { return fmt.Errorf("新版本试运行输出异常:%q", strings.TrimSpace(text)) } return nil } // pendingFile 是标记"已替换、待重启"的文件名。 const pendingFile = "upgrade-pending.json" // Pending 描述一次已落地但尚未生效的升级。 type Pending struct { FromVersion string `json:"from_version"` ToVersion string `json:"to_version"` BackupPath string `json:"backup_path"` AppliedAt time.Time `json:"applied_at"` } // PendingPath 返回 pending 标记文件路径。 func PendingPath(dataDir string) string { return filepath.Join(dataDir, pendingFile) } // ReadPending 读取待重启标记;不存在时返回 nil。 func ReadPending(dataDir string) *Pending { raw, err := os.ReadFile(PendingPath(dataDir)) if err != nil { return nil } var p Pending if err := json.Unmarshal(raw, &p); err != nil { return nil } return &p } // ClearPending 清除待重启标记。 func ClearPending(dataDir string) error { err := os.Remove(PendingPath(dataDir)) if errors.Is(err, os.ErrNotExist) { return nil } return err } // Install 执行替换:备份 → 替换 → 写标记。 // // exePath 为空(开发模式 `go run` 等)时只做备份并返回提示,不做替换。 func Install(dataDir, exePath, backupDir string, rel *Release, currentVersion string) (*InstallResult, error) { if rel == nil { return nil, errors.New("没有可安装的版本") } res := &InstallResult{ FromVersion: currentVersion, ToVersion: rel.Version, ExePath: exePath, } tmpDir := filepath.Join(dataDir, "tmp") if err := os.MkdirAll(tmpDir, 0o755); err != nil { return nil, err } tmp := filepath.Join(tmpDir, fmt.Sprintf("gitcat-%s-%s", rel.Version, platformKeyRepl())) if _, err := Download(context.Background(), rel, tmp); err != nil { os.Remove(tmp) return nil, err } if err := Verify(tmp); err != nil { os.Remove(tmp) return nil, err } if exePath == "" { res.Note = "当前不是以独立二进制方式运行(可能是 go run),已下载新版本到 " + tmp + ",请手工替换后重启。" return res, nil } if err := os.MkdirAll(backupDir, 0o755); err != nil { return nil, err } backup := filepath.Join(backupDir, "gitcat-"+time.Now().Format("20060102-150405")+".bak") if err := copyFile(exePath, backup); err != nil { os.Remove(tmp) return nil, fmt.Errorf("备份当前程序失败: %w", err) } res.BackupPath = backup if err := replaceExecutable(tmp, exePath); err != nil { os.Remove(tmp) return nil, err } res.PendingFile = PendingPath(dataDir) if err := os.WriteFile(res.PendingFile, mustJSON(Pending{ FromVersion: currentVersion, ToVersion: rel.Version, BackupPath: backup, AppliedAt: time.Now(), }), 0o644); err != nil { // 替换已经完成,标记写失败只影响"自动重启"这一步,不算失败。 res.Note = "程序已更新,但写入待重启标记失败,请手工重启服务。" return res, nil } return res, nil } func mustJSON(v any) []byte { b, err := json.MarshalIndent(v, "", " ") if err != nil { return []byte("{}") } return b } func platformKeyRepl() string { s := platformKey() return strings.ReplaceAll(s, "/", "-") } // replaceExecutable 用 newPath 覆盖 exePath。 // // Unix 上 rename 可以直接替换正在运行的可执行文件(内核持有的是 inode // 引用,替换只改目录项),这是最干净的做法。Windows 拒绝替换正在运行的 // exe,只能起一个后台辅助进程:等本进程退出后再完成搬运。 func replaceExecutable(newPath, exePath string) error { err := os.Rename(newPath, exePath) if err == nil { return nil } if runtime.GOOS != "windows" { return fmt.Errorf("替换程序文件失败: %w", err) } // Windows 兜底:让一个后台 PowerShell 等本进程退出后再完成替换。 // 之所以必须等退出,是因为映像文件在进程存活期间被系统独占锁定。 ps := fmt.Sprintf( `$ErrorActionPreference='Stop'; `+ `Wait-Process -Id %d -ErrorAction SilentlyContinue; `+ `Start-Sleep -Milliseconds 800; `+ `Copy-Item -LiteralPath '%s' -Destination '%s' -Force; `+ `Remove-Item -LiteralPath '%s' -Force -ErrorAction SilentlyContinue`, os.Getpid(), psQuote(newPath), psQuote(exePath), psQuote(newPath)) c := exec.Command("powershell", "-NoProfile", "-NonInteractive", "-WindowStyle", "Hidden", "-Command", ps) c.SysProcAttr = hideWindow() if err := c.Start(); err != nil { return fmt.Errorf("替换程序文件失败,且无法启动替换助手: %w", err) } // 不 Wait:助手必须在本次请求返回之后、本进程退出之前完成启动。 go func() { _ = c.Wait() }() return nil } // psQuote 转义单引号,供 PowerShell 单引号字符串使用。 func psQuote(s string) string { return "'" + strings.ReplaceAll(s, "'", "''") + "'" } func copyFile(src, dst string) error { in, err := os.Open(src) if err != nil { return err } defer in.Close() out, err := os.OpenFile(dst, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o755) if err != nil { return err } defer out.Close() if _, err := io.Copy(out, in); err != nil { return err } return out.Sync() }