修复 Agent 升级版本比对逻辑

This commit is contained in:
ryan
2026-06-22 11:26:55 +08:00
parent e8778a8641
commit 16b02fd3f1
3 changed files with 29 additions and 77 deletions
+2 -77
View File
@@ -12,12 +12,12 @@ import (
"io"
"net"
"net/http"
"strconv"
"strings"
"time"
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/utils"
)
const (
@@ -298,82 +298,7 @@ func isVersionNewer(current string, latest string) bool {
}
func compareVersions(local, remote string) int {
left := parseVersionInfo(local)
right := parseVersionInfo(remote)
if left.isDev {
if right.valid {
return -1
}
return 0
}
if !left.valid || !right.valid {
return 0
}
maxLen := len(left.numbers)
if len(right.numbers) > maxLen {
maxLen = len(right.numbers)
}
for index := 0; index < maxLen; index++ {
leftValue := 0
rightValue := 0
if index < len(left.numbers) {
leftValue = left.numbers[index]
}
if index < len(right.numbers) {
rightValue = right.numbers[index]
}
if leftValue < rightValue {
return -1
}
if leftValue > rightValue {
return 1
}
}
return 0
}
type versionInfo struct {
valid bool
isDev bool
numbers []int
}
func parseVersionInfo(version string) versionInfo {
normalized := strings.TrimSpace(strings.TrimPrefix(version, "v"))
if normalized == "" || normalized == "dev" {
return versionInfo{isDev: strings.EqualFold(normalized, "dev")}
}
base := normalized
if separator := strings.IndexRune(normalized, '-'); separator >= 0 {
base = normalized[:separator]
}
segments := strings.Split(base, ".")
parts := make([]int, 0, len(segments))
for _, segment := range segments {
segment = strings.TrimSpace(segment)
if segment == "" {
parts = append(parts, 0)
continue
}
numeric := strings.Builder{}
for _, r := range segment {
if r < '0' || r > '9' {
break
}
numeric.WriteRune(r)
}
if numeric.Len() == 0 {
parts = append(parts, 0)
continue
}
value, err := strconv.Atoi(numeric.String())
if err != nil {
return versionInfo{}
}
parts = append(parts, value)
}
return versionInfo{valid: len(parts) > 0, numbers: parts}
return utils.CompareVersions(local, remote)
}
func fetchLatestGitHubRelease(ctx context.Context, repo string, channel releaseChannel) (*githubReleaseResponse, error) {
@@ -331,3 +331,28 @@ type roundTripFunc func(req *http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestCompareVersions(t *testing.T) {
tests := []struct {
local string
remote string
expected int
}{
{"v3.0.0-beta", "v3.0.0-beta.1", -1},
{"v3.0.0-beta", "v3.0.0", -1},
{"v3.0.0-beta.1", "v3.0.0", -1},
{"dev", "v3.0.0", -1},
{"v3.0.0", "v3.0.0", 0},
{"v3.0.0", "v2.9.9", 1},
{"v3.0.0", "v3.0.1", -1},
{"v3.0.0-beta.1", "v3.0.0-beta.2", -1},
{"v3.0.0-beta.11", "v3.0.0-beta.2", 1},
}
for _, tt := range tests {
t.Run(tt.local+"_vs_"+tt.remote, func(t *testing.T) {
res := compareVersions(tt.local, tt.remote)
assert.Equal(t, tt.expected, res)
})
}
}