refactor(edge): 抽取边缘运行时共享包并完成 Phase 3 重构

- 新增 internal/apps/edge/,三组件改为薄包装,删除 3000+ 行重复代码
- Agent 心跳周期下沉至 heartbeat/cycle.go
- 协议类型迁入 pkg/protocol/agent.go
- 补充设计文档与 changelog
This commit is contained in:
ryan
2026-06-19 14:56:00 +08:00
parent cc5e53c51e
commit db9a9f98fd
53 changed files with 2091 additions and 2947 deletions
+54
View File
@@ -0,0 +1,54 @@
package config
import (
"encoding/json"
"fmt"
"strconv"
"strings"
"time"
)
type MillisecondDuration time.Duration
func (d MillisecondDuration) Duration() time.Duration {
return time.Duration(d)
}
func (d MillisecondDuration) String() string {
return time.Duration(d).String()
}
func (d *MillisecondDuration) UnmarshalJSON(data []byte) error {
raw := strings.TrimSpace(string(data))
if raw == "" || raw == "null" {
*d = 0
return nil
}
if strings.HasPrefix(raw, "\"") {
var text string
if err := json.Unmarshal(data, &text); err != nil {
return err
}
text = strings.TrimSpace(text)
if text == "" {
*d = 0
return nil
}
parsed, err := time.ParseDuration(text)
if err != nil {
return fmt.Errorf("invalid duration string %q: %w", text, err)
}
*d = MillisecondDuration(parsed)
return nil
}
ms, err := strconv.ParseInt(raw, 10, 64)
if err != nil {
return fmt.Errorf("invalid duration milliseconds %q: %w", raw, err)
}
*d = MillisecondDuration(time.Duration(ms) * time.Millisecond)
return nil
}
func (d MillisecondDuration) MarshalJSON() ([]byte, error) {
return json.Marshal(time.Duration(d).Milliseconds())
}
@@ -0,0 +1,53 @@
package config
import (
"encoding/json"
"testing"
"time"
)
func TestMillisecondDurationUnmarshalJSON(t *testing.T) {
tests := []struct {
name string
input string
want time.Duration
wantErr bool
}{
{name: "null", input: "null", want: 0},
{name: "empty", input: `""`, want: 0},
{name: "integer milliseconds", input: "30000", want: 30 * time.Second},
{name: "duration string", input: `"5s"`, want: 5 * time.Second},
{name: "invalid number", input: "not-a-number", wantErr: true},
{name: "invalid duration string", input: `"not-a-duration"`, wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var d MillisecondDuration
err := json.Unmarshal([]byte(tt.input), &d)
if tt.wantErr {
if err == nil {
t.Fatal("expected error")
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if d.Duration() != tt.want {
t.Fatalf("got %s, want %s", d, tt.want)
}
})
}
}
func TestMillisecondDurationMarshalJSON(t *testing.T) {
d := MillisecondDuration(7 * time.Second)
data, err := json.Marshal(d)
if err != nil {
t.Fatalf("Marshal failed: %v", err)
}
if string(data) != "7000" {
t.Fatalf("unexpected marshaled value: %s", string(data))
}
}
@@ -0,0 +1,41 @@
package heartbeat
import (
"context"
"log/slog"
"strings"
edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
)
type AutoUpdateSettings struct {
AutoUpdate bool
UpdateNow bool
UpdateRepo string
UpdateChannel string
UpdateTag string
}
func TryAutoUpdate(ctx context.Context, updater *edgeupdater.Service, settings *AutoUpdateSettings, logLabel string) {
if settings == nil || updater == nil {
return
}
force := settings.UpdateNow
shouldCheck := settings.AutoUpdate || force
if !shouldCheck || strings.TrimSpace(settings.UpdateRepo) == "" {
return
}
channel := "stable"
if force && strings.TrimSpace(settings.UpdateChannel) != "" {
channel = settings.UpdateChannel
}
slog.Info("checking for "+logLabel+" updates", "repo", settings.UpdateRepo, "channel", channel, "force", force)
err := updater.CheckAndUpdate(ctx, settings.UpdateRepo, edgeupdater.UpdateOptions{
Channel: channel,
TagName: settings.UpdateTag,
Force: force,
})
if err != nil {
slog.Error(logLabel+" update check failed", "error", err)
}
}
+23
View File
@@ -0,0 +1,23 @@
package heartbeat
import (
"context"
"time"
)
// RunLoop invokes fn immediately, then on each interval tick until ctx is cancelled.
func RunLoop(ctx context.Context, interval time.Duration, fn func(context.Context)) {
fn(ctx)
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
fn(ctx)
}
}
}
+41
View File
@@ -0,0 +1,41 @@
package heartbeat
import (
"context"
"sync/atomic"
"testing"
"time"
)
func TestRunLoopImmediateAndTicker(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var calls atomic.Int32
interval := 20 * time.Millisecond
done := make(chan struct{})
go func() {
RunLoop(ctx, interval, func(context.Context) {
calls.Add(1)
})
close(done)
}()
time.Sleep(5 * time.Millisecond)
if got := calls.Load(); got != 1 {
t.Fatalf("expected 1 immediate call, got %d", got)
}
time.Sleep(35 * time.Millisecond)
if got := calls.Load(); got < 2 {
t.Fatalf("expected at least 2 calls after ticker, got %d", got)
}
cancel()
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("RunLoop did not exit after context cancellation")
}
}
+132
View File
@@ -0,0 +1,132 @@
package httpclient
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"strings"
"time"
)
type Client struct {
baseURL string
token string
authHeader string
httpClient *http.Client
}
func New(baseURL, token string, timeout time.Duration, authHeader string) *Client {
return &Client{
baseURL: strings.TrimRight(baseURL, "/"),
token: token,
authHeader: authHeader,
httpClient: &http.Client{Timeout: timeout},
}
}
func (c *Client) SetToken(token string) {
c.token = strings.TrimSpace(token)
slog.Debug("http client token updated")
}
func (c *Client) GetJSON(ctx context.Context, path string, target any) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil)
if err != nil {
return err
}
c.setAuthHeader(req)
return c.do(req, target)
}
func (c *Client) PostJSON(ctx context.Context, path string, body any, target any) error {
data, err := json.Marshal(body)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+path, bytes.NewReader(data))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
c.setAuthHeader(req)
return c.do(req, target)
}
func (c *Client) DoRaw(ctx context.Context, method, path string, headers map[string]string) (*http.Response, error) {
req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, nil)
if err != nil {
return nil, err
}
c.setAuthHeader(req)
for key, value := range headers {
req.Header.Set(key, value)
}
return c.httpClient.Do(req)
}
func (c *Client) setAuthHeader(req *http.Request) {
if c.authHeader != "" {
req.Header.Set(c.authHeader, c.token)
}
}
func (c *Client) do(req *http.Request, target any) error {
res, err := c.httpClient.Do(req)
if err != nil {
slog.Error("http request failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
defer func(Body io.ReadCloser) {
err := Body.Close()
if err != nil {
slog.Error("failed to close response body", "error", err)
}
}(res.Body)
body, err := io.ReadAll(res.Body)
if err != nil {
slog.Error("http response read failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
if res.StatusCode != http.StatusOK {
slog.Warn("http request returned non-200", "method", req.Method, "path", req.URL.Path, "status", res.Status)
return ReadBodyError(body, res.Status)
}
if target == nil {
return nil
}
if err = json.Unmarshal(body, target); err != nil {
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
return nil
}
func APIError(msg string) error {
if strings.TrimSpace(msg) == "" {
return nil
}
return errors.New(msg)
}
func ReadBodyError(body []byte, fallback string) error {
var errBody struct {
ErrorMsg string `json:"error_msg"`
}
if err := json.Unmarshal(body, &errBody); err == nil && strings.TrimSpace(errBody.ErrorMsg) != "" {
return errors.New(errBody.ErrorMsg)
}
return errors.New(fallback)
}
func ReadHTTPError(res *http.Response) error {
body, err := io.ReadAll(res.Body)
if err != nil {
return errors.New(res.Status)
}
return ReadBodyError(body, res.Status)
}
+33
View File
@@ -0,0 +1,33 @@
package logging
import (
"log/slog"
"os"
"strings"
)
type Options struct {
AddSource bool
}
func Setup(opts Options) {
handlerOpts := &slog.HandlerOptions{
AddSource: opts.AddSource,
Level: ParseLevel(os.Getenv("LOG_LEVEL")),
}
handler := slog.NewTextHandler(os.Stdout, handlerOpts)
slog.SetDefault(slog.New(handler))
}
func ParseLevel(value string) slog.Level {
switch strings.ToLower(strings.TrimSpace(value)) {
case "debug":
return slog.LevelDebug
case "warn", "warning":
return slog.LevelWarn
case "error":
return slog.LevelError
default:
return slog.LevelInfo
}
}
+69
View File
@@ -0,0 +1,69 @@
package nodeip
import (
"context"
"net"
"time"
"github.com/Rain-kl/Wavelet/pkg/geoip"
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
)
var (
LookupOutboundIP = geoip.GetOutboundIP
LookupLocalIP = DetectLocal
)
func Detect() string {
if ip := detectOutbound(); ip != "" {
return ip
}
return LookupLocalIP()
}
func detectOutbound() string {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
ip, err := LookupOutboundIP(ctx)
if err != nil || ip == nil {
return ""
}
return ip.String()
}
func DetectLocal() string {
interfaces, err := net.Interfaces()
if err != nil {
return ""
}
bestIP := ""
bestPriority := -1
for _, iface := range interfaces {
if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 {
continue
}
addrs, err := iface.Addrs()
if err != nil {
continue
}
for _, addr := range addrs {
ipNet, ok := addr.(*net.IPNet)
if !ok || ipNet.IP == nil || ipNet.IP.IsLoopback() {
continue
}
ipv4 := ipNet.IP.To4()
if ipv4 == nil {
continue
}
priority := iputil.Score(ipv4)
if priority > bestPriority {
bestIP = ipv4.String()
bestPriority = priority
}
if bestPriority == 2 {
return bestIP
}
}
}
return bestIP
}
+268
View File
@@ -0,0 +1,268 @@
package observability
import (
"bufio"
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"syscall"
)
// ReadLinuxOSRelease returns the OS name and version from /etc/os-release.
func ReadLinuxOSRelease() (string, string) {
file, err := os.Open("/etc/os-release")
if err != nil {
return runtime.GOOS, ""
}
defer file.Close()
values := make(map[string]string)
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
key, value, ok := strings.Cut(line, "=")
if !ok {
continue
}
values[key] = strings.Trim(value, `"`)
}
if pretty := strings.TrimSpace(values["PRETTY_NAME"]); pretty != "" {
return pretty, strings.TrimSpace(values["VERSION_ID"])
}
name := strings.TrimSpace(values["NAME"])
if name == "" {
name = runtime.GOOS
}
return name, strings.TrimSpace(values["VERSION_ID"])
}
// ReadLinuxCPUModel returns the CPU model name from /proc/cpuinfo.
func ReadLinuxCPUModel() string {
file, err := os.Open("/proc/cpuinfo")
if err != nil {
return ""
}
defer file.Close()
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(strings.ToLower(line), "model name") {
_, value, ok := strings.Cut(line, ":")
if ok {
return strings.TrimSpace(value)
}
}
}
return ""
}
// ReadMemInfo returns total and used memory bytes from /proc/meminfo.
func ReadMemInfo() (int64, int64) {
file, err := os.Open("/proc/meminfo")
if err != nil {
return 0, 0
}
defer file.Close()
var memTotalKB int64
var memAvailableKB int64
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := scanner.Text()
switch {
case strings.HasPrefix(line, "MemTotal:"):
memTotalKB = parseMemInfoValue(line)
case strings.HasPrefix(line, "MemAvailable:"):
memAvailableKB = parseMemInfoValue(line)
}
}
total := memTotalKB * 1024
if total == 0 {
return 0, 0
}
used := total - (memAvailableKB * 1024)
if used < 0 {
used = 0
}
return total, used
}
func parseMemInfoValue(line string) int64 {
fields := strings.Fields(line)
if len(fields) < 2 {
return 0
}
value, err := strconv.ParseInt(fields[1], 10, 64)
if err != nil {
return 0
}
return value
}
// ReadLinuxUptimeSeconds returns system uptime in seconds from /proc/uptime.
func ReadLinuxUptimeSeconds() int64 {
content, err := os.ReadFile("/proc/uptime")
if err != nil {
return 0
}
fields := strings.Fields(string(content))
if len(fields) == 0 {
return 0
}
value, err := strconv.ParseFloat(fields[0], 64)
if err != nil {
return 0
}
return int64(value)
}
// ReadLinuxCPUStat returns aggregate CPU jiffies and idle jiffies from /proc/stat.
func ReadLinuxCPUStat() (uint64, uint64) {
content, err := os.ReadFile("/proc/stat")
if err != nil {
return 0, 0
}
lines := strings.Split(string(content), "\n")
for _, line := range lines {
if !strings.HasPrefix(line, "cpu ") {
continue
}
fields := strings.Fields(line)
if len(fields) < 5 {
return 0, 0
}
var total uint64
for i := 1; i < len(fields); i++ {
value, err := strconv.ParseUint(fields[i], 10, 64)
if err != nil {
return 0, 0
}
total += value
}
idle, err := strconv.ParseUint(fields[4], 10, 64)
if err != nil {
return 0, 0
}
return total, idle
}
return 0, 0
}
// ReadLinuxNetworkTotals returns aggregate RX and TX byte totals from /proc/net/dev.
func ReadLinuxNetworkTotals() (int64, int64) {
file, err := os.Open("/proc/net/dev")
if err != nil {
return 0, 0
}
defer file.Close()
var rx int64
var tx int64
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if !strings.Contains(line, ":") {
continue
}
name, data, ok := strings.Cut(line, ":")
if !ok {
continue
}
if strings.TrimSpace(name) == "lo" {
continue
}
fields := strings.Fields(data)
if len(fields) < 16 {
continue
}
rxValue, err := strconv.ParseInt(fields[0], 10, 64)
if err == nil {
rx += rxValue
}
txValue, err := strconv.ParseInt(fields[8], 10, 64)
if err == nil {
tx += txValue
}
}
return rx, tx
}
// ReadLinuxDiskTotals returns aggregate disk read and write byte totals from /proc/diskstats.
func ReadLinuxDiskTotals() (int64, int64) {
file, err := os.Open("/proc/diskstats")
if err != nil {
return 0, 0
}
defer file.Close()
var readBytes int64
var writeBytes int64
scanner := bufio.NewScanner(file)
for scanner.Scan() {
fields := strings.Fields(scanner.Text())
if len(fields) < 14 {
continue
}
device := fields[2]
if shouldSkipDiskDevice(device) {
continue
}
readSectors, err := strconv.ParseInt(fields[5], 10, 64)
if err == nil {
readBytes += readSectors * 512
}
writeSectors, err := strconv.ParseInt(fields[9], 10, 64)
if err == nil {
writeBytes += writeSectors * 512
}
}
return readBytes, writeBytes
}
func shouldSkipDiskDevice(device string) bool {
switch {
case device == "":
return true
case strings.HasPrefix(device, "loop"),
strings.HasPrefix(device, "ram"),
strings.HasPrefix(device, "dm-"):
return true
default:
return false
}
}
// StatFilesystem returns total and used bytes for the filesystem containing path.
func StatFilesystem(path string) (int64, int64) {
if strings.TrimSpace(path) == "" {
path = string(os.PathSeparator)
}
absPath := filepath.Clean(path)
var stat syscall.Statfs_t
if err := syscall.Statfs(absPath, &stat); err != nil {
return 0, 0
}
total := int64(stat.Blocks) * int64(stat.Bsize)
free := int64(stat.Bavail) * int64(stat.Bsize)
used := total - free
if used < 0 {
used = 0
}
return total, used
}
// ReadFirstLine reads and returns the trimmed first line of a file.
func ReadFirstLine(path string) string {
content, err := os.ReadFile(path)
if err != nil {
return ""
}
return strings.TrimSpace(string(content))
}
@@ -0,0 +1,78 @@
package observability
import (
"os"
"path/filepath"
"testing"
)
func TestParseMemInfoValue(t *testing.T) {
t.Parallel()
tests := []struct {
line string
want int64
}{
{line: "MemTotal: 16384000 kB", want: 16384000},
{line: "MemAvailable: 8192000 kB", want: 8192000},
{line: "invalid", want: 0},
{line: "MemTotal: not-a-number kB", want: 0},
}
for _, tt := range tests {
if got := parseMemInfoValue(tt.line); got != tt.want {
t.Fatalf("parseMemInfoValue(%q) = %d, want %d", tt.line, got, tt.want)
}
}
}
func TestShouldSkipDiskDevice(t *testing.T) {
t.Parallel()
tests := []struct {
device string
want bool
}{
{device: "", want: true},
{device: "loop0", want: true},
{device: "ram0", want: true},
{device: "dm-0", want: true},
{device: "sda", want: false},
{device: "nvme0n1", want: false},
}
for _, tt := range tests {
if got := shouldSkipDiskDevice(tt.device); got != tt.want {
t.Fatalf("shouldSkipDiskDevice(%q) = %v, want %v", tt.device, got, tt.want)
}
}
}
func TestReadFirstLine(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, "sample.txt")
if err := os.WriteFile(path, []byte(" first line\nsecond line\n"), 0o644); err != nil {
t.Fatalf("write file: %v", err)
}
if got := ReadFirstLine(path); got != "first line\nsecond line" {
t.Fatalf("ReadFirstLine() = %q, want trimmed first line content", got)
}
if got := ReadFirstLine(filepath.Join(dir, "missing.txt")); got != "" {
t.Fatalf("ReadFirstLine(missing) = %q, want empty string", got)
}
}
func TestStatFilesystem(t *testing.T) {
t.Parallel()
total, used := StatFilesystem(t.TempDir())
if total <= 0 {
t.Fatalf("StatFilesystem() total = %d, want > 0", total)
}
if used < 0 || used > total {
t.Fatalf("StatFilesystem() used = %d, total = %d", used, total)
}
}
+64
View File
@@ -0,0 +1,64 @@
package runner
import (
"context"
"log/slog"
"time"
)
type WSConnection interface {
Close() error
}
type WSReconnectConfig struct {
ComponentName string
ConnectBackoff time.Duration
ReconnectDelay time.Duration
OnShutdown func()
}
func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig,
connect func(context.Context) (WSConnection, error),
handle func(context.Context, WSConnection),
) error {
if cfg.ConnectBackoff <= 0 {
cfg.ConnectBackoff = 5 * time.Second
}
if cfg.ReconnectDelay <= 0 {
cfg.ReconnectDelay = 2 * time.Second
}
label := cfg.ComponentName
if label == "" {
label = "edge"
}
for {
select {
case <-ctx.Done():
if cfg.OnShutdown != nil {
cfg.OnShutdown()
}
return ctx.Err()
default:
}
conn, err := connect(ctx)
if err != nil {
slog.Error(label+" ws connect failed, will retry", "error", err)
SleepContext(ctx, cfg.ConnectBackoff)
continue
}
handle(ctx, conn)
_ = conn.Close()
slog.Info(label + " ws connection closed, reconnecting...")
SleepContext(ctx, cfg.ReconnectDelay)
}
}
func SleepContext(ctx context.Context, d time.Duration) {
select {
case <-ctx.Done():
case <-time.After(d):
}
}
@@ -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, `"`, `""`) + `"`
}
+381
View File
@@ -0,0 +1,381 @@
package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"os"
"runtime"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/utils"
)
const maxChecksumAssetSize = 64 * 1024
var replaceAndRestartFunc = replaceAndRestart
type Config struct {
LocalVersion string
AssetPrefix string
LogLabel string
}
type Service struct {
httpClient *http.Client
lastCheckKey string
localVersion string
assetPrefix string
logLabel string
}
func New(cfg Config) *Service {
return &Service{
httpClient: &http.Client{Timeout: 30 * time.Second},
localVersion: cfg.LocalVersion,
assetPrefix: cfg.AssetPrefix,
logLabel: cfg.LogLabel,
}
}
type UpdateOptions struct {
Channel string
TagName string
Force bool
}
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 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(s.localVersion)
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(s.logLabel+" update available", "from", localVersion, "to", remoteVersion)
assetName := s.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 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(s.logLabel + " binary updated, restarting")
return replaceAndRestartFunc(targetPath, tmpPath)
}
func (s *Service) assetNameForGOOSGOARCH(goos string, goarch string) string {
name := fmt.Sprintf("%s-%s-%s", s.assetPrefix, 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 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)
}
+232
View File
@@ -0,0 +1,232 @@
package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"os"
"path/filepath"
"runtime"
"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 testService(httpClient *http.Client) *Service {
return &Service{
httpClient: httpClient,
localVersion: "v1.0.0",
assetPrefix: "openflare-agent",
logLabel: "agent",
}
}
func TestGetLatestPreviewRelease(t *testing.T) {
service := testService(&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", 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 := testService(&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", 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) {
assetName := testService(nil).assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
service := testService(&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", 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 := testService(&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 := testService(&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)
}
})
}
}