mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 15:46:37 +08:00
[优化] go 引用调整
This commit is contained in:
@@ -0,0 +1,51 @@
|
||||
//go:build !windows
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
if err := removeBackupBinary(backupPath); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(execPath, backupPath); err != nil {
|
||||
renameErr := err
|
||||
if err := os.Remove(tmpPath); err != nil && !os.IsNotExist(err) {
|
||||
slog.Error("remove tmp binary failed", "path", tmpPath, "error", err)
|
||||
return fmt.Errorf("backup current binary: %w; remove tmp binary: %v", renameErr, err)
|
||||
}
|
||||
return fmt.Errorf("backup current binary: %w", renameErr)
|
||||
}
|
||||
if err := os.Rename(tmpPath, execPath); err != nil {
|
||||
replaceErr := err
|
||||
if err := os.Rename(backupPath, execPath); err != nil {
|
||||
slog.Error("restore backup binary failed", "path", backupPath, "error", err)
|
||||
return fmt.Errorf("replace binary: %w; restore backup binary: %v", replaceErr, err)
|
||||
}
|
||||
return fmt.Errorf("replace binary: %w", replaceErr)
|
||||
}
|
||||
if err := removeBackupBinary(backupPath); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil {
|
||||
return fmt.Errorf("exec restart: %w", err)
|
||||
}
|
||||
return fmt.Errorf("unreachable after exec")
|
||||
}
|
||||
|
||||
func removeBackupBinary(path string) error {
|
||||
if err := os.Remove(path); err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
slog.Error("remove backup binary failed", "path", path, "error", err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
//go:build !windows
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRemoveBackupBinaryIgnoresMissingFile(t *testing.T) {
|
||||
backupPath := filepath.Join(t.TempDir(), "openflare-agent.bak")
|
||||
if err := removeBackupBinary(backupPath); err != nil {
|
||||
t.Fatalf("expected missing backup cleanup to be ignored: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
//go:build windows
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
scriptPath := execPath + ".update.cmd"
|
||||
script := fmt.Sprintf(`@echo off
|
||||
setlocal
|
||||
:waitloop
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 (
|
||||
ping 127.0.0.1 -n 2 >nul
|
||||
goto waitloop
|
||||
)
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 exit /b 1
|
||||
start "" %s
|
||||
del /Q "%s" >nul 2>nul
|
||||
del /Q "%%~f0" >nul 2>nul
|
||||
`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath)
|
||||
if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("write restart script: %w", err)
|
||||
}
|
||||
cmd := exec.Command("cmd", "/C", "start", "", scriptPath)
|
||||
if err := cmd.Start(); err != nil {
|
||||
os.Remove(scriptPath)
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("schedule restart: %w", err)
|
||||
}
|
||||
os.Exit(0)
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildWindowsCommandLine(execPath string, args []string) string {
|
||||
parts := []string{quoteWindowsArg(execPath)}
|
||||
for _, arg := range args {
|
||||
parts = append(parts, quoteWindowsArg(arg))
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func quoteWindowsArg(value string) string {
|
||||
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
|
||||
}
|
||||
@@ -0,0 +1,366 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-server/utils"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-agent/internal/agent"
|
||||
"github.com/rain-kl/openflare/openflare-agent/internal/config"
|
||||
)
|
||||
|
||||
const maxChecksumAssetSize = 64 * 1024
|
||||
|
||||
var replaceAndRestartFunc = replaceAndRestart
|
||||
|
||||
type Service struct {
|
||||
httpClient *http.Client
|
||||
lastCheckKey string
|
||||
}
|
||||
|
||||
func New() *Service {
|
||||
return &Service{
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
type githubRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
Assets []githubAsset `json:"assets"`
|
||||
}
|
||||
|
||||
type githubAsset struct {
|
||||
Name string `json:"name"`
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
if release == nil || release.TagName == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
remoteVersion := normalizeVersion(release.TagName)
|
||||
localVersion := normalizeVersion(config.Version)
|
||||
checkKey := buildReleaseCheckKey(options, remoteVersion)
|
||||
|
||||
if remoteVersion == localVersion {
|
||||
return nil
|
||||
}
|
||||
if !options.Force && checkKey != "" && checkKey == s.lastCheckKey {
|
||||
return nil
|
||||
}
|
||||
if !isNewer(localVersion, remoteVersion) {
|
||||
s.lastCheckKey = checkKey
|
||||
return nil
|
||||
}
|
||||
|
||||
slog.Info("agent update available", "from", localVersion, "to", remoteVersion)
|
||||
assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
|
||||
checksumAssetName := assetName + ".sha256"
|
||||
|
||||
var downloadURL string
|
||||
var checksumURL string
|
||||
for _, asset := range release.Assets {
|
||||
switch asset.Name {
|
||||
case assetName:
|
||||
downloadURL = asset.BrowserDownloadURL
|
||||
case checksumAssetName:
|
||||
checksumURL = asset.BrowserDownloadURL
|
||||
}
|
||||
}
|
||||
if downloadURL == "" {
|
||||
s.lastCheckKey = checkKey
|
||||
return fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName)
|
||||
}
|
||||
if checksumURL == "" {
|
||||
return fmt.Errorf("no matching checksum asset %q in release %s", checksumAssetName, release.TagName)
|
||||
}
|
||||
|
||||
expectedChecksum, err := s.downloadChecksum(ctx, checksumURL, assetName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("download checksum: %w", err)
|
||||
}
|
||||
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get executable path: %w", err)
|
||||
}
|
||||
if err = s.downloadAndRestart(ctx, downloadURL, expectedChecksum, execPath); err != nil {
|
||||
return fmt.Errorf("download and restart: %w", err)
|
||||
}
|
||||
s.lastCheckKey = checkKey
|
||||
return nil
|
||||
}
|
||||
|
||||
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)
|
||||
return s.fetchReleaseFromURL(ctx, url)
|
||||
}
|
||||
|
||||
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))
|
||||
return s.fetchReleaseFromURL(ctx, url)
|
||||
}
|
||||
|
||||
func (s *Service) fetchReleaseFromURL(ctx context.Context, url string) (*githubRelease, error) {
|
||||
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 func(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
slog.Error("failed to close response body", "error", err)
|
||||
}
|
||||
}(resp.Body)
|
||||
|
||||
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(reader).Decode(&release); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
func (s *Service) downloadChecksum(ctx context.Context, url string, assetName string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("checksum download returned %s", resp.Status)
|
||||
}
|
||||
|
||||
content, err := io.ReadAll(io.LimitReader(resp.Body, maxChecksumAssetSize+1))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if len(content) > maxChecksumAssetSize {
|
||||
return "", fmt.Errorf("checksum asset exceeds %d bytes", maxChecksumAssetSize)
|
||||
}
|
||||
checksum, err := parseSHA256Checksum(string(content), assetName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return checksum, nil
|
||||
}
|
||||
|
||||
func parseSHA256Checksum(content string, assetName string) (string, error) {
|
||||
assetName = strings.TrimSpace(assetName)
|
||||
for _, line := range strings.Split(content, "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "#") {
|
||||
continue
|
||||
}
|
||||
if checksum, ok := parseSHA256Line(line, assetName); ok {
|
||||
return checksum, nil
|
||||
}
|
||||
}
|
||||
if assetName == "" {
|
||||
return "", fmt.Errorf("checksum asset does not contain a valid sha256 digest")
|
||||
}
|
||||
return "", fmt.Errorf("checksum asset does not contain a sha256 digest for %q", assetName)
|
||||
}
|
||||
|
||||
func parseSHA256Line(line string, assetName string) (string, bool) {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) == 1 && isSHA256Hex(fields[0]) {
|
||||
return strings.ToLower(fields[0]), true
|
||||
}
|
||||
if len(fields) >= 2 && isSHA256Hex(fields[0]) {
|
||||
fileName := strings.TrimPrefix(strings.TrimSpace(fields[1]), "*")
|
||||
if assetName == "" || fileName == assetName {
|
||||
return strings.ToLower(fields[0]), true
|
||||
}
|
||||
}
|
||||
|
||||
prefix := "SHA256("
|
||||
if strings.HasPrefix(line, prefix) {
|
||||
closing := strings.Index(line, ")")
|
||||
if closing > len(prefix) && closing+1 < len(line) {
|
||||
fileName := strings.TrimSpace(line[len(prefix):closing])
|
||||
rest := strings.TrimSpace(line[closing+1:])
|
||||
rest = strings.TrimPrefix(rest, "=")
|
||||
rest = strings.TrimSpace(rest)
|
||||
if isSHA256Hex(rest) && (assetName == "" || fileName == assetName) {
|
||||
return strings.ToLower(rest), true
|
||||
}
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func isSHA256Hex(value string) bool {
|
||||
value = strings.TrimSpace(value)
|
||||
if len(value) != sha256.Size*2 {
|
||||
return false
|
||||
}
|
||||
_, err := hex.DecodeString(value)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedChecksum string, targetPath string) error {
|
||||
expectedChecksum = strings.ToLower(strings.TrimSpace(expectedChecksum))
|
||||
if !isSHA256Hex(expectedChecksum) {
|
||||
return fmt.Errorf("invalid expected sha256 checksum")
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("download returned %s", resp.Status)
|
||||
}
|
||||
|
||||
tmpPath := targetPath + ".update"
|
||||
if runtime.GOOS == "windows" && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") {
|
||||
tmpPath += ".exe"
|
||||
}
|
||||
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
hasher := sha256.New()
|
||||
if _, err = io.Copy(io.MultiWriter(tmpFile, hasher), resp.Body); err != nil {
|
||||
tmpFile.Close()
|
||||
os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
if err = tmpFile.Close(); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
actualChecksum := hex.EncodeToString(hasher.Sum(nil))
|
||||
if actualChecksum != expectedChecksum {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("sha256 checksum mismatch: expected %s, got %s", expectedChecksum, actualChecksum)
|
||||
}
|
||||
if err = os.Chmod(tmpPath, 0o755); err != nil && runtime.GOOS != "windows" {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("set executable permission: %w", err)
|
||||
}
|
||||
|
||||
slog.Info("agent binary updated, restarting")
|
||||
return replaceAndRestartFunc(targetPath, tmpPath)
|
||||
}
|
||||
|
||||
func assetNameForGOOSGOARCH(goos string, goarch string) string {
|
||||
name := fmt.Sprintf("openflare-agent-%s-%s", goos, goarch)
|
||||
if goos == "windows" {
|
||||
return name + ".exe"
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
func normalizeVersion(v string) string {
|
||||
v = strings.TrimSpace(v)
|
||||
v = strings.TrimPrefix(v, "v")
|
||||
return v
|
||||
}
|
||||
|
||||
func isNewer(local, remote string) bool {
|
||||
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
|
||||
}
|
||||
|
||||
func compareVersions(local string, remote string) int {
|
||||
return utils.CompareVersions(local, remote)
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/rain-kl/openflare/openflare-agent/internal/agent"
|
||||
"github.com/rain-kl/openflare/openflare-agent/internal/config"
|
||||
)
|
||||
|
||||
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/OpenFlare/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/OpenFlare", 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/OpenFlare/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/OpenFlare", 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 TestCheckAndUpdateRequiresChecksumAsset(t *testing.T) {
|
||||
originalVersion := config.Version
|
||||
config.Version = "v1.0.0"
|
||||
t.Cleanup(func() {
|
||||
config.Version = originalVersion
|
||||
})
|
||||
|
||||
assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
|
||||
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/OpenFlare/releases/latest" {
|
||||
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.1",
|
||||
"assets":[
|
||||
{"name":"` + assetName + `","browser_download_url":"https://example.test/agent"}
|
||||
]
|
||||
}`)),
|
||||
}, nil
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
err := service.CheckAndUpdate(context.Background(), "Rain-kl/OpenFlare", agent.UpdateOptions{})
|
||||
if err == nil || !strings.Contains(err.Error(), "no matching checksum asset") {
|
||||
t.Fatalf("expected missing checksum asset error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseSHA256Checksum(t *testing.T) {
|
||||
checksum := strings.Repeat("a", sha256.Size*2)
|
||||
testCases := []struct {
|
||||
name string
|
||||
content string
|
||||
asset string
|
||||
want string
|
||||
}{
|
||||
{name: "single digest", content: checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
{name: "sha256sum format", content: checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
{name: "bsd format", content: "SHA256(openflare-agent-linux-amd64)= " + checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
{name: "selects matching file", content: strings.Repeat("b", sha256.Size*2) + " other\n" + checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
got, err := parseSHA256Checksum(testCase.content, testCase.asset)
|
||||
if err != nil {
|
||||
t.Fatalf("expected checksum parse to succeed: %v", err)
|
||||
}
|
||||
if got != testCase.want {
|
||||
t.Fatalf("unexpected checksum: got %s want %s", got, testCase.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadAndRestartVerifiesChecksum(t *testing.T) {
|
||||
payload := []byte("new-agent-binary")
|
||||
sum := sha256.Sum256(payload)
|
||||
expectedChecksum := hex.EncodeToString(sum[:])
|
||||
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
|
||||
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
|
||||
t.Fatalf("write target: %v", err)
|
||||
}
|
||||
|
||||
var replacedTarget string
|
||||
var replacedTemp string
|
||||
originalReplace := replaceAndRestartFunc
|
||||
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
|
||||
replacedTarget = execPath
|
||||
replacedTemp = tmpPath
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
replaceAndRestartFunc = originalReplace
|
||||
})
|
||||
|
||||
service := &Service{
|
||||
httpClient: &http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(string(payload))),
|
||||
}, nil
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
if err := service.downloadAndRestart(context.Background(), "https://example.test/agent", expectedChecksum, targetPath); err != nil {
|
||||
t.Fatalf("expected verified download to succeed: %v", err)
|
||||
}
|
||||
if replacedTarget != targetPath {
|
||||
t.Fatalf("unexpected replace target: %s", replacedTarget)
|
||||
}
|
||||
if replacedTemp == "" {
|
||||
t.Fatal("expected replacement temp path to be recorded")
|
||||
}
|
||||
if _, err := os.Stat(replacedTemp); err != nil {
|
||||
t.Fatalf("expected verified temp binary to remain for replacement: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadAndRestartRejectsChecksumMismatch(t *testing.T) {
|
||||
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
|
||||
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
|
||||
t.Fatalf("write target: %v", err)
|
||||
}
|
||||
|
||||
originalReplace := replaceAndRestartFunc
|
||||
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
|
||||
t.Fatal("replace should not run on checksum mismatch")
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
replaceAndRestartFunc = originalReplace
|
||||
})
|
||||
|
||||
service := &Service{
|
||||
httpClient: &http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader("tampered")),
|
||||
}, nil
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
err := service.downloadAndRestart(context.Background(), "https://example.test/agent", strings.Repeat("0", sha256.Size*2), targetPath)
|
||||
if err == nil || !strings.Contains(err.Error(), "sha256 checksum mismatch") {
|
||||
t.Fatalf("expected checksum mismatch error, got %v", err)
|
||||
}
|
||||
if _, err = os.Stat(targetPath + ".update"); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected temp update file to be removed, stat err=%v", err)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user