gitcat
1// Package upgrade 实现 gitcat 的一键升级:检查更新、下载、校验并就地替换
2// 当前运行的二进制文件。
3//
4// 设计上有一条硬约束:**绝不在新二进制被验证可用之前动现有的那个**。
5// 升级失败时,服务必须还能照常启动。因此流程是
6//
7// 拉清单 → 选版本 → 下载到临时文件 → 校验 sha256 → 试运行 → 备份旧文件
8// → 替换 → 写 pending 标记 → 重启
9//
10// 中间任何一步失败都直接返回错误,旧二进制保持原样。
11package upgrade
12
13import (
14 "context"
15 "crypto/sha256"
16 "encoding/hex"
17 "encoding/json"
18 "errors"
19 "fmt"
20 "io"
21 "net/http"
22 "net/url"
23 "os"
24 "os/exec"
25 "path/filepath"
26 "runtime"
27 "strconv"
28 "strings"
29 "time"
30)
31
32const (
33 // maxManifestSize 限制升级清单大小,避免被超大响应撑爆内存。
34 maxManifestSize = 1 << 20
35 // maxBinarySize 限制下载的二进制大小(200MB),同理。
36 maxBinarySize = 200 << 20
37 // fetchTimeout 是拉取清单 / 下载的默认上限。
38 fetchTimeout = 10 * time.Minute
39)
40
41// Info 描述当前正在运行的二进制。
42type Info struct {
43 Version string // 语义化版本,如 1.2.0
44 Commit string // 构建时注入的提交号
45 BuildDate string // 构建时间
46 ExePath string // 可执行文件绝对路径
47 GOOS string
48 GOARCH string
49}
50
51// Release 是发布清单中的一个版本。
52//
53// 清单由管理员在后台配置的一个 URL 提供,格式示例:
54//
55// {
56// "latest": "1.2.0",
57// "min_supported": "1.0.0",
58// "releases": [
59// {"version": "1.2.0", "url": "https://…/gitcat-linux-amd64",
60// "sha256": "…", "size": 8388608, "os": "linux", "arch": "amd64",
61// "notes": "修复推送越权"}
62// ]
63// }
64type Release struct {
65 Version string `json:"version"`
66 URL string `json:"url"`
67 SHA256 string `json:"sha256"`
68 Size int64 `json:"size"`
69 OS string `json:"os"`
70 Arch string `json:"arch"`
71 Notes string `json:"notes"`
72}
73
74// Manifest 是升级源提供的版本清单。
75type Manifest struct {
76 Latest string `json:"latest"`
77 // MinSupported 是允许直接跳到的最低版本:低于它的版本需要先中转升级,
78 // 避免一次跨度过大的结构变更。
79 MinSupported string `json:"min_supported"`
80 Releases []Release `json:"releases"`
81}
82
83// CheckResult 是一次检查更新的结论。
84type CheckResult struct {
85 Current string
86 Latest string
87 Upgradable bool
88 Release *Release
89 Manifest *Manifest
90 Message string
91 CheckedAt time.Time
92 AssetExists bool
93}
94
95// InstallResult 描述一次替换的结果。
96type InstallResult struct {
97 FromVersion string
98 ToVersion string
99 ExePath string
100 BackupPath string
101 PendingFile string
102 Restarted bool
103 Note string
104}
105
106// platformKey 返回当前平台在清单里的标识。
107func platformKey() string { return runtime.GOOS + "/" + runtime.GOARCH }
108
109// IsDev 判断当前是否是开发版(未注入正式版本号)。
110func (i Info) IsDev() bool {
111 v := strings.ToLower(strings.TrimSpace(i.Version))
112 return v == "" || v == "dev" || v == "devel" || strings.Contains(v, "-dirty") || strings.Contains(v, "-dev")
113}
114
115// ShortSHA 返回短提交号。
116func (i Info) ShortSHA() string {
117 if len(i.Commit) > 7 {
118 return i.Commit[:7]
119 }
120 return i.Commit
121}
122
123// Compare 比较两个语义化版本:a<b 返回 -1,a>b 返回 1,相等返回 0。
124//
125// 允许带 "v" 前缀,比较时忽略;带预发布后缀(1.0.0-rc1)视为低于 1.0.0。
126func Compare(a, b string) int {
127 na, pa := splitVersion(a)
128 nb, pb := splitVersion(b)
129 for i := 0; i < 3; i++ {
130 if na[i] != nb[i] {
131 if na[i] < nb[i] {
132 return -1
133 }
134 return 1
135 }
136 }
137 switch {
138 case pa == pb:
139 return 0
140 case pa == "":
141 return 1 // 正式版 > 预发布版
142 case pb == "":
143 return -1
144 default:
145 return strings.Compare(pa, pb)
146 }
147}
148
149func splitVersion(v string) ([3]int, string) {
150 var out [3]int
151 v = strings.TrimSpace(v)
152 v = strings.TrimPrefix(strings.TrimPrefix(v, "v"), "V")
153 pre := ""
154 if i := strings.IndexAny(v, "-+"); i >= 0 {
155 pre = v[i+1:]
156 v = v[:i]
157 }
158 for i, part := range strings.SplitN(v, ".", 3) {
159 if i > 2 {
160 break
161 }
162 n, err := strconv.Atoi(strings.TrimSpace(part))
163 if err != nil {
164 continue
165 }
166 out[i] = n
167 }
168 return out, pre
169}
170
171// ValidVersion 判断字符串是否像一个可比较的版本号。
172func ValidVersion(v string) bool {
173 v = strings.TrimSpace(v)
174 if v == "" {
175 return false
176 }
177 nums, _ := splitVersion(v)
178 return nums[0] > 0 || nums[1] > 0 || nums[2] > 0
179}
180
181// httpClient 是升级流程专用的客户端:短超时、限制重定向。
182func httpClient(timeout time.Duration) *http.Client {
183 return &http.Client{
184 Timeout: timeout,
185 CheckRedirect: func(req *http.Request, via []*http.Request) error {
186 if len(via) >= 5 {
187 return errors.New("重定向次数过多")
188 }
189 return nil
190 },
191 }
192}
193
194// FetchManifest 拉取并解析版本清单。
195func FetchManifest(ctx context.Context, manifestURL string) (*Manifest, error) {
196 manifestURL = strings.TrimSpace(manifestURL)
197 if manifestURL == "" {
198 return nil, errors.New("未配置升级源")
199 }
200 u, err := url.Parse(manifestURL)
201 if err != nil || !strings.HasPrefix(u.Scheme, "http") {
202 return nil, errors.New("升级源必须是 http/https 地址")
203 }
204 ctx, cancel := context.WithTimeout(ctx, 60*time.Second)
205 defer cancel()
206 req, err := http.NewRequestWithContext(ctx, http.MethodGet, manifestURL, nil)
207 if err != nil {
208 return nil, err
209 }
210 req.Header.Set("Accept", "application/json")
211 resp, err := httpClient(60 * time.Second).Do(req)
212 if err != nil {
213 return nil, err
214 }
215 defer resp.Body.Close()
216 if resp.StatusCode != http.StatusOK {
217 return nil, fmt.Errorf("升级源返回 HTTP %d", resp.StatusCode)
218 }
219 raw, err := io.ReadAll(io.LimitReader(resp.Body, maxManifestSize))
220 if err != nil {
221 return nil, err
222 }
223 var m Manifest
224 if err := json.Unmarshal(raw, &m); err != nil {
225 return nil, fmt.Errorf("升级源返回的内容不是合法 JSON: %w", err)
226 }
227 if len(m.Releases) == 0 {
228 return nil, errors.New("升级源里没有任何版本")
229 }
230 return &m, nil
231}
232
233// pickRelease 从清单中挑出适用于当前平台的目标版本。
234func pickRelease(m *Manifest, version string) *Release {
235 want := strings.TrimSpace(version)
236 if want == "" {
237 want = strings.TrimSpace(m.Latest)
238 }
239 for i := range m.Releases {
240 r := m.Releases[i]
241 if !strings.EqualFold(strings.TrimSpace(r.Version), want) {
242 continue
243 }
244 if r.OS != "" && r.Arch != "" && r.OS+"/"+r.Arch != platformKey() {
245 return nil // 清单里明确标注了其它平台
246 }
247 rr := r
248 return &rr
249 }
250 return nil
251}
252
253// Check 对比当前版本与升级源,给出升级结论。
254func (i Info) Check(ctx context.Context, manifestURL string) (*CheckResult, error) {
255 m, err := FetchManifest(ctx, manifestURL)
256 if err != nil {
257 return nil, err
258 }
259 res := &CheckResult{
260 Current: i.Version,
261 Latest: m.Latest,
262 Manifest: m,
263 CheckedAt: time.Now(),
264 }
265 rel := pickRelease(m, m.Latest)
266 if rel == nil {
267 res.Message = fmt.Sprintf("升级源未提供适用于 %s 的版本 %s", platformKey(), m.Latest)
268 return res, nil
269 }
270 res.AssetExists = true
271 res.Release = rel
272
273 switch {
274 case i.IsDev():
275 res.Message = "当前是开发构建,无法自动判断新版本,请对照发布页确认。"
276 case rel.Version == "" || !ValidVersion(rel.Version):
277 res.Message = "升级源中的版本号不合法。"
278 case Compare(i.Version, rel.Version) >= 0:
279 res.Message = "已经是最新版本。"
280 case m.MinSupported != "" && ValidVersion(m.MinSupported) && Compare(i.Version, m.MinSupported) < 0:
281 res.Message = fmt.Sprintf("当前版本过旧,需先升级到 %s 或更高版本再继续。", m.MinSupported)
282 default:
283 res.Upgradable = true
284 res.Message = fmt.Sprintf("发现新版本 %s", rel.Version)
285 }
286 return res, nil
287}
288
289// Download 下载指定版本到 dest,返回校验后的文件大小。
290func Download(ctx context.Context, rel *Release, dest string) (int64, error) {
291 if rel == nil || strings.TrimSpace(rel.URL) == "" {
292 return 0, errors.New("该版本没有提供下载地址")
293 }
294 u, err := url.Parse(rel.URL)
295 if err != nil || !strings.HasPrefix(u.Scheme, "http") {
296 return 0, errors.New("下载地址必须是 http/https 地址")
297 }
298 ctx, cancel := context.WithTimeout(ctx, fetchTimeout)
299 defer cancel()
300 req, err := http.NewRequestWithContext(ctx, http.MethodGet, rel.URL, nil)
301 if err != nil {
302 return 0, err
303 }
304 req.Header.Set("User-Agent", "gitcat-upgrade")
305 resp, err := httpClient(fetchTimeout).Do(req)
306 if err != nil {
307 return 0, err
308 }
309 defer resp.Body.Close()
310 if resp.StatusCode != http.StatusOK {
311 return 0, fmt.Errorf("下载失败:HTTP %d", resp.StatusCode)
312 }
313 limit := int64(maxBinarySize)
314 if rel.Size > 0 && rel.Size < limit {
315 limit = rel.Size
316 }
317
318 f, err := os.OpenFile(dest, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o755)
319 if err != nil {
320 return 0, err
321 }
322 defer f.Close()
323 h := sha256.New()
324 n, err := io.Copy(io.MultiWriter(f, h), io.LimitReader(resp.Body, limit+1))
325 if err != nil {
326 return 0, err
327 }
328 if n > limit {
329 return 0, fmt.Errorf("下载内容超过预期体积(%d 字节),已中止", limit)
330 }
331 sum := hex.EncodeToString(h.Sum(nil))
332 if want := strings.ToLower(strings.TrimSpace(rel.SHA256)); want != "" && want != sum {
333 return 0, fmt.Errorf("校验失败:期望 sha256 %s,实际 %s", want, sum)
334 }
335 if rel.Size > 0 && n != rel.Size {
336 return 0, fmt.Errorf("校验失败:期望 %d 字节,实际 %d 字节", rel.Size, n)
337 }
338 if err := f.Sync(); err != nil {
339 return 0, err
340 }
341 return n, nil
342}
343
344// Verify 对新下载的可执行文件做一次试运行,确认它真的能跑起来。
345//
346// 这一步是"不损坏现网"的关键:架构不匹配、动态库缺失、文件被截断时,
347// 试运行会立刻失败,而此时旧二进制还没被碰过。
348func Verify(binPath string) error {
349 ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
350 defer cancel()
351 cmd := exec.CommandContext(ctx, binPath, "-version")
352 out, err := cmd.CombinedOutput()
353 if err != nil {
354 return fmt.Errorf("新版本无法执行(可能是平台不匹配或文件损坏):%v", err)
355 }
356 text := string(out)
357 if !strings.Contains(text, "gitcat") {
358 return fmt.Errorf("新版本试运行输出异常:%q", strings.TrimSpace(text))
359 }
360 return nil
361}
362
363// pendingFile 是标记"已替换、待重启"的文件名。
364const pendingFile = "upgrade-pending.json"
365
366// Pending 描述一次已落地但尚未生效的升级。
367type Pending struct {
368 FromVersion string `json:"from_version"`
369 ToVersion string `json:"to_version"`
370 BackupPath string `json:"backup_path"`
371 AppliedAt time.Time `json:"applied_at"`
372}
373
374// PendingPath 返回 pending 标记文件路径。
375func PendingPath(dataDir string) string { return filepath.Join(dataDir, pendingFile) }
376
377// ReadPending 读取待重启标记;不存在时返回 nil。
378func ReadPending(dataDir string) *Pending {
379 raw, err := os.ReadFile(PendingPath(dataDir))
380 if err != nil {
381 return nil
382 }
383 var p Pending
384 if err := json.Unmarshal(raw, &p); err != nil {
385 return nil
386 }
387 return &p
388}
389
390// ClearPending 清除待重启标记。
391func ClearPending(dataDir string) error {
392 err := os.Remove(PendingPath(dataDir))
393 if errors.Is(err, os.ErrNotExist) {
394 return nil
395 }
396 return err
397}
398
399// Install 执行替换:备份 → 替换 → 写标记。
400//
401// exePath 为空(开发模式 `go run` 等)时只做备份并返回提示,不做替换。
402func Install(dataDir, exePath, backupDir string, rel *Release, currentVersion string) (*InstallResult, error) {
403 if rel == nil {
404 return nil, errors.New("没有可安装的版本")
405 }
406 res := &InstallResult{
407 FromVersion: currentVersion,
408 ToVersion: rel.Version,
409 ExePath: exePath,
410 }
411 tmpDir := filepath.Join(dataDir, "tmp")
412 if err := os.MkdirAll(tmpDir, 0o755); err != nil {
413 return nil, err
414 }
415 tmp := filepath.Join(tmpDir, fmt.Sprintf("gitcat-%s-%s", rel.Version, platformKeyRepl()))
416 if _, err := Download(context.Background(), rel, tmp); err != nil {
417 os.Remove(tmp)
418 return nil, err
419 }
420 if err := Verify(tmp); err != nil {
421 os.Remove(tmp)
422 return nil, err
423 }
424
425 if exePath == "" {
426 res.Note = "当前不是以独立二进制方式运行(可能是 go run),已下载新版本到 " + tmp + ",请手工替换后重启。"
427 return res, nil
428 }
429 if err := os.MkdirAll(backupDir, 0o755); err != nil {
430 return nil, err
431 }
432 backup := filepath.Join(backupDir, "gitcat-"+time.Now().Format("20060102-150405")+".bak")
433 if err := copyFile(exePath, backup); err != nil {
434 os.Remove(tmp)
435 return nil, fmt.Errorf("备份当前程序失败: %w", err)
436 }
437 res.BackupPath = backup
438
439 if err := replaceExecutable(tmp, exePath); err != nil {
440 os.Remove(tmp)
441 return nil, err
442 }
443 res.PendingFile = PendingPath(dataDir)
444 if err := os.WriteFile(res.PendingFile, mustJSON(Pending{
445 FromVersion: currentVersion,
446 ToVersion: rel.Version,
447 BackupPath: backup,
448 AppliedAt: time.Now(),
449 }), 0o644); err != nil {
450 // 替换已经完成,标记写失败只影响"自动重启"这一步,不算失败。
451 res.Note = "程序已更新,但写入待重启标记失败,请手工重启服务。"
452 return res, nil
453 }
454 return res, nil
455}
456
457func mustJSON(v any) []byte {
458 b, err := json.MarshalIndent(v, "", " ")
459 if err != nil {
460 return []byte("{}")
461 }
462 return b
463}
464
465func platformKeyRepl() string {
466 s := platformKey()
467 return strings.ReplaceAll(s, "/", "-")
468}
469
470// replaceExecutable 用 newPath 覆盖 exePath。
471//
472// Unix 上 rename 可以直接替换正在运行的可执行文件(内核持有的是 inode
473// 引用,替换只改目录项),这是最干净的做法。Windows 拒绝替换正在运行的
474// exe,只能起一个后台辅助进程:等本进程退出后再完成搬运。
475func replaceExecutable(newPath, exePath string) error {
476 err := os.Rename(newPath, exePath)
477 if err == nil {
478 return nil
479 }
480 if runtime.GOOS != "windows" {
481 return fmt.Errorf("替换程序文件失败: %w", err)
482 }
483 // Windows 兜底:让一个后台 PowerShell 等本进程退出后再完成替换。
484 // 之所以必须等退出,是因为映像文件在进程存活期间被系统独占锁定。
485 ps := fmt.Sprintf(
486 `$ErrorActionPreference='Stop'; `+
487 `Wait-Process -Id %d -ErrorAction SilentlyContinue; `+
488 `Start-Sleep -Milliseconds 800; `+
489 `Copy-Item -LiteralPath '%s' -Destination '%s' -Force; `+
490 `Remove-Item -LiteralPath '%s' -Force -ErrorAction SilentlyContinue`,
491 os.Getpid(), psQuote(newPath), psQuote(exePath), psQuote(newPath))
492 c := exec.Command("powershell", "-NoProfile", "-NonInteractive", "-WindowStyle", "Hidden", "-Command", ps)
493 c.SysProcAttr = hideWindow()
494 if err := c.Start(); err != nil {
495 return fmt.Errorf("替换程序文件失败,且无法启动替换助手: %w", err)
496 }
497 // 不 Wait:助手必须在本次请求返回之后、本进程退出之前完成启动。
498 go func() { _ = c.Wait() }()
499 return nil
500}
501
502// psQuote 转义单引号,供 PowerShell 单引号字符串使用。
503func psQuote(s string) string { return "'" + strings.ReplaceAll(s, "'", "''") + "'" }
504
505func copyFile(src, dst string) error {
506 in, err := os.Open(src)
507 if err != nil {
508 return err
509 }
510 defer in.Close()
511 out, err := os.OpenFile(dst, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o755)
512 if err != nil {
513 return err
514 }
515 defer out.Close()
516 if _, err := io.Copy(out, in); err != nil {
517 return err
518 }
519 return out.Sync()
520}