feat: add support for release channels in version upgrade and node agent updates

- Introduced ReleaseChannel type to manage stable and preview releases.
- Updated DashboardTopbar to handle version upgrades based on selected release channel.
- Enhanced node detail page to allow manual checks for agent updates on stable and preview channels.
- Modified API endpoints to support fetching and upgrading based on release channels.
- Updated UI components to reflect changes in version checking and upgrade processes.
- Added tests for new functionality related to preview releases and agent updates.
This commit is contained in:
ryan
2026-03-12 11:25:13 +08:00
parent 93e43fb3b0
commit f50eb9adee
25 changed files with 1448 additions and 239 deletions
+27 -3
View File
@@ -24,7 +24,13 @@ type SyncService interface {
}
type Updater interface {
CheckAndUpdate(ctx context.Context, repo string) error
CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error
}
type UpdateOptions struct {
Channel string
TagName string
Force bool
}
type Runner struct {
@@ -37,6 +43,8 @@ type Runner struct {
autoUpdate bool
updateNow bool
updateRepo string
updateChan string
updateTag string
}
func (r *Runner) Run(ctx context.Context) error {
@@ -133,18 +141,34 @@ func (r *Runner) applySettings(settings *protocol.AgentSettings) bool {
r.autoUpdate = settings.AutoUpdate
r.updateNow = settings.UpdateNow
r.updateRepo = strings.TrimSpace(settings.UpdateRepo)
r.updateChan = strings.TrimSpace(settings.UpdateChannel)
r.updateTag = strings.TrimSpace(settings.UpdateTag)
return changed
}
func (r *Runner) tryAutoUpdate(ctx context.Context) {
shouldCheck := r.autoUpdate || r.updateNow
force := r.updateNow
shouldCheck := r.autoUpdate || force
r.updateNow = false
r.updateTag = strings.TrimSpace(r.updateTag)
if !shouldCheck || r.Updater == nil || r.updateRepo == "" {
return
}
if err := r.Updater.CheckAndUpdate(ctx, r.updateRepo); err != nil {
channel := "stable"
if force && r.updateChan != "" {
channel = r.updateChan
}
if err := r.Updater.CheckAndUpdate(ctx, r.updateRepo, UpdateOptions{
Channel: channel,
TagName: r.updateTag,
Force: force,
}); err != nil {
log.Printf("agent update check failed: %v", err)
}
if force {
r.updateTag = ""
r.updateChan = ""
}
}
func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error {
@@ -19,6 +19,8 @@ type AgentSettings struct {
AutoUpdate bool `json:"auto_update"`
UpdateRepo string `json:"update_repo"`
UpdateNow bool `json:"update_now"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
}
type NodePayload struct {
+246 -32
View File
@@ -9,15 +9,17 @@ import (
"net/http"
"os"
"runtime"
"strconv"
"strings"
"time"
"atsflare-agent/internal/agent"
"atsflare-agent/internal/config"
)
type Service struct {
httpClient *http.Client
lastCheckTag string
lastCheckKey string
}
func New() *Service {
@@ -27,8 +29,10 @@ func New() *Service {
}
type githubRelease struct {
TagName string `json:"tag_name"`
Assets []githubAsset `json:"assets"`
TagName string `json:"tag_name"`
Prerelease bool `json:"prerelease"`
Draft bool `json:"draft"`
Assets []githubAsset `json:"assets"`
}
type githubAsset struct {
@@ -36,8 +40,8 @@ type githubAsset struct {
BrowserDownloadURL string `json:"browser_download_url"`
}
func (s *Service) CheckAndUpdate(ctx context.Context, repo string) error {
release, err := s.getLatestRelease(ctx, repo)
func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options agent.UpdateOptions) error {
release, err := s.getRelease(ctx, repo, options)
if err != nil {
return fmt.Errorf("check latest release: %w", err)
}
@@ -47,12 +51,16 @@ func (s *Service) CheckAndUpdate(ctx context.Context, repo string) error {
remoteVersion := normalizeVersion(release.TagName)
localVersion := normalizeVersion(config.AgentVersion)
checkKey := buildReleaseCheckKey(options, remoteVersion)
if remoteVersion == localVersion || remoteVersion == s.lastCheckTag {
if remoteVersion == localVersion {
return nil
}
if !options.Force && checkKey != "" && checkKey == s.lastCheckKey {
return nil
}
if !isNewer(localVersion, remoteVersion) {
s.lastCheckTag = remoteVersion
s.lastCheckKey = checkKey
return nil
}
@@ -67,7 +75,7 @@ func (s *Service) CheckAndUpdate(ctx context.Context, repo string) error {
}
}
if downloadURL == "" {
s.lastCheckTag = remoteVersion
s.lastCheckKey = checkKey
return fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName)
}
@@ -78,10 +86,22 @@ func (s *Service) CheckAndUpdate(ctx context.Context, repo string) error {
if err = s.downloadAndRestart(ctx, downloadURL, execPath); err != nil {
return fmt.Errorf("download and restart: %w", err)
}
s.lastCheckKey = checkKey
return nil
}
func (s *Service) getLatestRelease(ctx context.Context, repo string) (*githubRelease, error) {
func (s *Service) getRelease(ctx context.Context, repo string, options agent.UpdateOptions) (*githubRelease, error) {
tagName := strings.TrimSpace(options.TagName)
if tagName != "" {
return s.getReleaseByTag(ctx, repo, tagName)
}
if strings.EqualFold(strings.TrimSpace(options.Channel), "preview") {
return s.getLatestPreviewRelease(ctx, repo)
}
return s.getLatestStableRelease(ctx, repo)
}
func (s *Service) getLatestStableRelease(ctx context.Context, repo string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/latest", repo)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
@@ -102,8 +122,68 @@ func (s *Service) getLatestRelease(ctx context.Context, repo string) (*githubRel
return nil, fmt.Errorf("github api returned %s", resp.Status)
}
return decodeRelease(resp.Body)
}
func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases?per_page=20", repo)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/vnd.github+json")
resp, err := s.httpClient.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("github api returned %s", resp.Status)
}
var releases []githubRelease
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
return nil, err
}
for _, release := range releases {
if release.Draft || !release.Prerelease {
continue
}
releaseCopy := release
return &releaseCopy, nil
}
return nil, nil
}
func (s *Service) getReleaseByTag(ctx context.Context, repo string, tag string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/tags/%s", repo, strings.TrimSpace(tag))
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/vnd.github+json")
resp, err := s.httpClient.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusNotFound {
return nil, nil
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("github api returned %s", resp.Status)
}
return decodeRelease(resp.Body)
}
func decodeRelease(reader io.Reader) (*githubRelease, error) {
var release githubRelease
if err = json.NewDecoder(resp.Body).Decode(&release); err != nil {
if err := json.NewDecoder(reader).Decode(&release); err != nil {
return nil, err
}
return &release, nil
@@ -158,26 +238,160 @@ func normalizeVersion(v string) string {
}
func isNewer(local, remote string) bool {
localParts := strings.Split(local, ".")
remoteParts := strings.Split(remote, ".")
maxLen := len(localParts)
if len(remoteParts) > maxLen {
maxLen = len(remoteParts)
}
for i := 0; i < maxLen; i++ {
lp, rp := "0", "0"
if i < len(localParts) {
lp = localParts[i]
}
if i < len(remoteParts) {
rp = remoteParts[i]
}
if rp > lp {
return true
}
if rp < lp {
return false
}
}
return false
return compareVersions(local, remote) < 0
}
func buildReleaseCheckKey(options agent.UpdateOptions, remoteVersion string) string {
channel := strings.TrimSpace(options.Channel)
if channel == "" {
channel = "stable"
}
if tagName := strings.TrimSpace(options.TagName); tagName != "" {
return channel + ":" + tagName
}
return channel + ":" + remoteVersion
}
type versionInfo struct {
valid bool
isDev bool
numbers []int
prerelease []string
}
func parseVersionInfo(version string) versionInfo {
normalized := normalizeVersion(version)
if normalized == "" || strings.EqualFold(normalized, "dev") {
return versionInfo{isDev: strings.EqualFold(normalized, "dev")}
}
base := normalized
prerelease := ""
if index := strings.IndexRune(normalized, '-'); index >= 0 {
base = normalized[:index]
prerelease = normalized[index+1:]
}
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 {
return versionInfo{}
}
value, err := strconv.Atoi(numeric.String())
if err != nil {
return versionInfo{}
}
parts = append(parts, value)
}
info := versionInfo{valid: len(parts) > 0, numbers: parts}
if prerelease != "" {
info.prerelease = splitPrereleaseIdentifiers(prerelease)
}
return info
}
func splitPrereleaseIdentifiers(value string) []string {
parts := strings.FieldsFunc(strings.TrimSpace(value), func(r rune) bool {
return r == '.' || r == '-'
})
filtered := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
filtered = append(filtered, part)
}
}
return filtered
}
func compareVersions(local string, 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
}
}
if len(left.prerelease) == 0 && len(right.prerelease) == 0 {
return 0
}
if len(left.prerelease) == 0 {
return 1
}
if len(right.prerelease) == 0 {
return -1
}
maxLen = len(left.prerelease)
if len(right.prerelease) > maxLen {
maxLen = len(right.prerelease)
}
for index := 0; index < maxLen; index++ {
if index >= len(left.prerelease) {
return -1
}
if index >= len(right.prerelease) {
return 1
}
leftPart := left.prerelease[index]
rightPart := right.prerelease[index]
leftNumber, leftErr := strconv.Atoi(leftPart)
rightNumber, rightErr := strconv.Atoi(rightPart)
switch {
case leftErr == nil && rightErr == nil:
if leftNumber < rightNumber {
return -1
}
if leftNumber > rightNumber {
return 1
}
case leftErr == nil && rightErr != nil:
return -1
case leftErr != nil && rightErr == nil:
return 1
default:
if leftPart < rightPart {
return -1
}
if leftPart > rightPart {
return 1
}
}
}
return 0
}
@@ -0,0 +1,91 @@
package updater
import (
"atsflare-agent/internal/agent"
"context"
"io"
"net/http"
"strings"
"testing"
)
type roundTripFunc func(req *http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestGetLatestPreviewRelease(t *testing.T) {
service := &Service{
httpClient: &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/ATSFlare/releases?per_page=20" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`[
{"tag_name":"v1.0.0","prerelease":false},
{"tag_name":"v1.1.0-rc.1","prerelease":true}
]`)),
}, nil
}),
},
}
release, err := service.getRelease(context.Background(), "Rain-kl/ATSFlare", agent.UpdateOptions{Channel: "preview"})
if err != nil {
t.Fatalf("expected preview release query to succeed: %v", err)
}
if release == nil || release.TagName != "v1.1.0-rc.1" {
t.Fatalf("unexpected preview release: %#v", release)
}
}
func TestGetReleaseByTag(t *testing.T) {
service := &Service{
httpClient: &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/ATSFlare/releases/tags/v1.1.0-rc.1" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)),
}, nil
}),
},
}
release, err := service.getRelease(context.Background(), "Rain-kl/ATSFlare", agent.UpdateOptions{Channel: "preview", TagName: "v1.1.0-rc.1", Force: true})
if err != nil {
t.Fatalf("expected tag release query to succeed: %v", err)
}
if release == nil || release.TagName != "v1.1.0-rc.1" {
t.Fatalf("unexpected tag release: %#v", release)
}
}
func TestIsNewerSupportsPrerelease(t *testing.T) {
testCases := []struct {
name string
local string
remote string
expected bool
}{
{name: "stable newer than prerelease", local: "1.2.3-rc.1", remote: "1.2.3", expected: true},
{name: "same stable not newer", local: "1.2.3", remote: "1.2.3-rc.1", expected: false},
{name: "higher prerelease sequence", local: "1.2.3-rc.1", remote: "1.2.3-rc.2", expected: true},
{name: "higher minor", local: "1.2.3", remote: "1.3.0-rc.1", expected: true},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
if actual := isNewer(testCase.local, testCase.remote); actual != testCase.expected {
t.Fatalf("unexpected compare result: local=%s remote=%s actual=%v expected=%v", testCase.local, testCase.remote, actual, testCase.expected)
}
})
}
}