[#12] Auto-update downloads and executes binary with no signature or checksum verification

This commit is contained in:
ryan
2026-05-29 11:28:51 +08:00
parent 806863f303
commit fa23cad9e9
7 changed files with 286 additions and 9 deletions
+126 -8
View File
@@ -2,6 +2,8 @@ package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
@@ -17,6 +19,10 @@ import (
"openflare-agent/internal/config"
)
const maxChecksumAssetSize = 64 * 1024
var replaceAndRestartFunc = replaceAndRestart
type Service struct {
httpClient *http.Client
lastCheckKey string
@@ -66,24 +72,36 @@ func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options agent
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 {
if asset.Name == assetName {
switch asset.Name {
case assetName:
downloadURL = asset.BrowserDownloadURL
break
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, execPath); err != nil {
if err = s.downloadAndRestart(ctx, downloadURL, expectedChecksum, execPath); err != nil {
return fmt.Errorf("download and restart: %w", err)
}
s.lastCheckKey = checkKey
@@ -189,7 +207,94 @@ func decodeRelease(reader io.Reader) (*githubRelease, error) {
return &release, nil
}
func (s *Service) downloadAndRestart(ctx context.Context, url string, targetPath string) error {
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
@@ -208,19 +313,32 @@ func (s *Service) downloadAndRestart(ctx context.Context, url string, targetPath
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, 0o755)
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600)
if err != nil {
return err
}
if _, err = io.Copy(tmpFile, resp.Body); err != nil {
hasher := sha256.New()
if _, err = io.Copy(io.MultiWriter(tmpFile, hasher), resp.Body); err != nil {
tmpFile.Close()
os.Remove(tmpPath)
return err
}
tmpFile.Close()
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 replaceAndRestart(targetPath, tmpPath)
return replaceAndRestartFunc(targetPath, tmpPath)
}
func assetNameForGOOSGOARCH(goos string, goarch string) string {
@@ -2,9 +2,15 @@ package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"openflare-agent/internal/agent"
"openflare-agent/internal/config"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
)
@@ -68,6 +74,150 @@ func TestGetReleaseByTag(t *testing.T) {
}
}
func TestCheckAndUpdateRequiresChecksumAsset(t *testing.T) {
originalVersion := config.AgentVersion
config.AgentVersion = "v1.0.0"
t.Cleanup(func() {
config.AgentVersion = 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