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
+22 -217
View File
@@ -9,10 +9,11 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config"
"github.com/Rain-kl/Wavelet/internal/apps/agent/observability"
agentheartbeat "github.com/Rain-kl/Wavelet/internal/apps/agent/heartbeat"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
"github.com/Rain-kl/Wavelet/internal/apps/agent/wsclient"
edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat"
)
type HeartbeatService interface {
@@ -29,10 +30,6 @@ type SyncService interface {
ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error
}
type Updater interface {
CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error
}
type RuntimeManager interface {
CheckHealth(ctx context.Context) error
Restart(ctx context.Context) error
@@ -44,32 +41,23 @@ type WebSocketService interface {
URL() string
}
type UpdateOptions struct {
Channel string
TagName string
Force bool
}
type Runner struct {
Config *config.Config
StateStore *state.Store
ObservabilityBuffer *state.ObservabilityBufferStore
HeartbeatCycle *agentheartbeat.Cycle
HeartbeatService HeartbeatService
SyncService SyncService
Updater Updater
RuntimeManager RuntimeManager
WebSocketService WebSocketService
autoUpdate bool
updateNow bool
updateRepo string
updateChan string
updateTag string
restartOpenrestyNow bool
websocketUpgradeEnabled bool
}
func (r *Runner) Run(ctx context.Context) error {
if r.HeartbeatCycle != nil {
r.HeartbeatCycle.RecordSyncError = r.recordSyncError
}
nodeID, err := r.StateStore.EnsureNodeID()
if err != nil {
return err
@@ -149,36 +137,15 @@ func (r *Runner) Run(ctx context.Context) error {
func (r *Runner) performHeartbeatCycle(ctx context.Context, nodeID string, startup bool) (bool, error) {
r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(nodeID)
heartbeatResult, err := r.HeartbeatService.Heartbeat(ctx, payload)
if err != nil {
return false, err
}
r.ackObservabilityWindows(ackWindows)
if heartbeatResult == nil {
heartbeatResult = &protocol.HeartbeatResult{}
}
mode := "periodic"
if startup {
mode = "startup"
}
slog.Debug("agent heartbeat succeeded", "mode", mode, "node_id", nodeID)
changed := r.applySettings(heartbeatResult.AgentSettings)
r.applyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups)
if startup {
if err = r.SyncService.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil {
r.recordSyncError(err)
slog.Error("agent startup sync failed", "error", err)
} else {
slog.Debug("agent startup sync completed")
}
} else if err = r.SyncService.SyncOnce(ctx, heartbeatResult.ActiveConfig); err != nil {
r.recordSyncError(err)
slog.Error("agent sync failed", "error", err)
}
return r.HeartbeatCycle.Perform(ctx, nodeID, startup, r)
}
func (r *Runner) Apply(settings *protocol.AgentSettings) bool {
return r.applySettings(settings)
}
func (r *Runner) RestartOpenrestyIfNeeded(ctx context.Context) {
r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx)
return changed, nil
}
func (r *Runner) shouldUseWebSocket() bool {
@@ -254,7 +221,6 @@ func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol.
childCtx, cancel := context.WithCancel(ctx)
defer cancel()
// Start status ticker sender in background
go func() {
for {
select {
@@ -285,11 +251,11 @@ func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol.
func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(nodeID)
payload, ackWindows := r.HeartbeatCycle.PrepareHeartbeatPayload(nodeID)
if err := conn.SendStatus(payload); err != nil {
return err
}
r.ackObservabilityWindows(ackWindows)
r.HeartbeatCycle.AckObservabilityWindows(ackWindows)
return nil
}
@@ -303,7 +269,7 @@ func (r *Runner) handleWebSocketMessage(ctx context.Context, message protocol.WS
}
changed := r.applySettings(&settings)
r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx)
edgeheartbeat.TryAutoUpdate(ctx, r.HeartbeatCycle.Updater, agentheartbeat.AgentSettingsToAutoUpdate(&settings), "agent")
if !r.websocketUpgradeEnabled {
slog.Debug("agent ws disabled by server settings; falling back to http heartbeat")
return changed, errors.New("websocket upgrade disabled by server")
@@ -339,7 +305,7 @@ func (r *Runner) handleWebSocketMessage(ctx context.Context, message protocol.WS
slog.Debug("agent ws waf ip groups decode failed", "error", err)
return false, nil
}
r.applyWAFIPGroups(ctx, groups)
r.HeartbeatCycle.ApplyWAFIPGroups(ctx, groups)
return false, nil
case protocol.WSMessageTypePing:
slog.Debug("agent ws ping received")
@@ -409,11 +375,6 @@ func (r *Runner) applySettings(settings *protocol.AgentSettings) bool {
slog.Debug("agent websocket upgrade setting updated", "from", r.websocketUpgradeEnabled, "to", settings.WebsocketUpgradeEnabled)
}
r.websocketUpgradeEnabled = settings.WebsocketUpgradeEnabled
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)
r.restartOpenrestyNow = settings.RestartOpenrestyNow
return changed
}
@@ -436,37 +397,12 @@ func (r *Runner) tryRestartOpenresty(ctx context.Context) {
r.recordOpenrestyHealthy()
}
func (r *Runner) tryAutoUpdate(ctx context.Context) {
force := r.updateNow
shouldCheck := r.autoUpdate || force
r.updateNow = false
r.updateTag = strings.TrimSpace(r.updateTag)
if !shouldCheck || r.Updater == nil || r.updateRepo == "" {
return
}
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 {
slog.Error("agent update check failed", "error", err)
}
if force {
r.updateTag = ""
r.updateChan = ""
}
}
func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error {
if strings.TrimSpace(r.Config.DiscoveryToken) == "" {
return errors.New("agent_token 为空且未配置 discovery_token")
}
slog.Info("agent discovery registration started")
response, err := r.HeartbeatService.Register(ctx, r.nodePayload(*nodeID))
response, err := r.HeartbeatService.Register(ctx, r.HeartbeatCycle.NodePayload(*nodeID))
if err != nil {
return err
}
@@ -493,26 +429,9 @@ func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error {
*nodeID = response.NodeID
slog.Info("agent discovery registration succeeded", "node_id", response.NodeID)
r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(*nodeID)
heartbeatResult, heartbeatErr := r.HeartbeatService.Heartbeat(ctx, payload)
if heartbeatErr != nil {
slog.Error("agent post-register heartbeat failed", "error", heartbeatErr)
return nil
if _, err = r.HeartbeatCycle.Perform(ctx, *nodeID, true, r); err != nil {
slog.Error("agent post-register heartbeat failed", "error", err)
}
r.ackObservabilityWindows(ackWindows)
if heartbeatResult == nil {
heartbeatResult = &protocol.HeartbeatResult{}
}
r.applySettings(heartbeatResult.AgentSettings)
r.applyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups)
if err = r.SyncService.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil {
r.recordSyncError(err)
slog.Error("agent post-register startup sync failed", "error", err)
} else {
slog.Debug("agent post-register startup sync completed")
}
r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx)
return nil
}
@@ -582,118 +501,4 @@ func (r *Runner) recordOpenrestyUnhealthy(err error, fallbackOnly bool) {
if saveErr := r.StateStore.Save(snapshot); saveErr != nil {
slog.Error("save state after recording openresty error failed", "error", saveErr)
}
}
func (r *Runner) nodePayload(nodeID string) protocol.NodePayload {
snapshot, _ := r.StateStore.Load()
openrestyStatus := strings.TrimSpace(snapshot.OpenrestyStatus)
if openrestyStatus == "" {
openrestyStatus = protocol.OpenrestyStatusUnknown
}
profile := observability.BuildProfile(r.Config, r.StateStore)
managedOpenRestyMetrics := observability.CollectManagedOpenRestyMetrics(r.Config)
trafficReport, accessLogs, fallbackMetrics := observability.BuildTrafficObservability(r.Config, r.StateStore, managedOpenRestyMetrics)
if managedOpenRestyMetrics == nil {
managedOpenRestyMetrics = fallbackMetrics
}
metricSnapshot := observability.BuildSnapshot(r.Config, r.StateStore)
openrestyObservation := observability.BuildOpenrestyObservation(managedOpenRestyMetrics)
healthEvents := observability.BuildHealthEvents(snapshot)
payload := protocol.NodePayload{
NodeID: nodeID,
Name: r.Config.NodeName,
IP: r.Config.NodeIP,
Version: r.Config.Version,
ExtVersion: r.Config.ExtVersion,
CurrentVersion: snapshot.CurrentVersion,
LastError: snapshot.LastError,
OpenrestyStatus: openrestyStatus,
OpenrestyMessage: snapshot.OpenrestyMessage,
Profile: profile,
Snapshot: metricSnapshot,
OpenrestyObservation: openrestyObservation,
TrafficReport: trafficReport,
AccessLogs: accessLogs,
HealthEvents: healthEvents,
}
if r.SyncService != nil {
checksums, err := r.SyncService.WAFIPGroupChecksums()
if err != nil {
slog.Debug("load local waf ip group checksums failed", "error", err)
} else if len(checksums) > 0 {
payload.WAFIPGroupChecksums = checksums
}
}
return payload
}
func (r *Runner) applyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) {
if len(groups) == 0 || r.SyncService == nil {
return
}
if err := r.SyncService.ApplyWAFIPGroups(ctx, groups); err != nil {
r.recordSyncError(err)
slog.Error("agent apply waf ip groups failed", "error", err)
}
}
func (r *Runner) prepareHeartbeatPayload(nodeID string) (protocol.NodePayload, []int64) {
payload := r.nodePayload(nodeID)
if r.ObservabilityBuffer == nil || (payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0) {
return payload, nil
}
now := time.Now().UTC()
retainAfterUnix := now.Add(-time.Duration(r.Config.ObservabilityReplayMinutes) * time.Minute).Unix()
windowStartedAtUnix := state.ObservabilityWindowStartedAt(payload.Snapshot, payload.OpenrestyObservation, payload.TrafficReport)
if windowStartedAtUnix <= 0 {
return payload, nil
}
record := state.ObservabilityBufferRecord{
WindowStartedAtUnix: windowStartedAtUnix,
Snapshot: payload.Snapshot,
OpenrestyObservation: payload.OpenrestyObservation,
TrafficReport: payload.TrafficReport,
AccessLogs: payload.AccessLogs,
QueuedAtUnix: now.Unix(),
}
if err := r.ObservabilityBuffer.Upsert(record, retainAfterUnix); err != nil {
slog.Error("upsert observability buffer failed", "error", err)
return payload, nil
}
records, err := r.ObservabilityBuffer.Replayable(windowStartedAtUnix, retainAfterUnix)
if err != nil {
slog.Error("load replayable observability buffer failed", "error", err)
return payload, []int64{windowStartedAtUnix}
}
ackWindows := make([]int64, 0, len(records)+1)
buffered := make([]protocol.BufferedObservabilityRecord, 0, len(records))
for _, item := range records {
if item.WindowStartedAtUnix <= 0 {
continue
}
buffered = append(buffered, protocol.BufferedObservabilityRecord{
WindowStartedAtUnix: item.WindowStartedAtUnix,
Snapshot: item.Snapshot,
OpenrestyObservation: item.OpenrestyObservation,
TrafficReport: item.TrafficReport,
AccessLogs: item.AccessLogs,
})
ackWindows = append(ackWindows, item.WindowStartedAtUnix)
}
payload.BufferedObservability = buffered
ackWindows = append(ackWindows, windowStartedAtUnix)
return payload, ackWindows
}
func (r *Runner) ackObservabilityWindows(windowStartedAtUnix []int64) {
if r.ObservabilityBuffer == nil || len(windowStartedAtUnix) == 0 {
return
}
retainAfterUnix := time.Now().UTC().Add(-time.Duration(r.Config.ObservabilityReplayMinutes) * time.Minute).Unix()
if err := r.ObservabilityBuffer.Ack(windowStartedAtUnix, retainAfterUnix); err != nil {
slog.Error("ack observability buffer failed", "error", err)
}
}
}
+34 -21
View File
@@ -11,10 +11,24 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config"
agentheartbeat "github.com/Rain-kl/Wavelet/internal/apps/agent/heartbeat"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
"github.com/Rain-kl/Wavelet/internal/apps/agent/updater"
)
func withHeartbeatCycle(runner *Runner, observabilityBuffer *state.ObservabilityBufferStore) *Runner {
runner.HeartbeatCycle = &agentheartbeat.Cycle{
Config: runner.Config,
StateStore: runner.StateStore,
ObservabilityBuffer: observabilityBuffer,
Heartbeat: runner.HeartbeatService,
Sync: runner.SyncService,
Updater: updater.New(),
}
return runner
}
type fakeHeartbeatService struct {
mu sync.Mutex
registerCalls int
@@ -192,7 +206,7 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) {
syncService := &fakeSyncService{
startupErr: errors.New("当前没有激活版本,保持当前 OpenResty 配置"),
}
runner := &Runner{
runner := withHeartbeatCycle(&Runner{
Config: &config.Config{
AccessToken: "agent-token",
NodeName: "edge-01",
@@ -204,7 +218,7 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) {
StateStore: stateStore,
HeartbeatService: heartbeatService,
SyncService: syncService,
}
}, nil)
err := runner.Run(ctx)
if !errors.Is(err, context.Canceled) {
@@ -245,7 +259,7 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
}
},
}
runner := &Runner{
runner := withHeartbeatCycle(&Runner{
Config: &config.Config{
AccessToken: "agent-token",
NodeName: "edge-01",
@@ -257,7 +271,7 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
StateStore: stateStore,
HeartbeatService: heartbeatService,
SyncService: syncService,
}
}, nil)
err := runner.Run(ctx)
if !errors.Is(err, context.Canceled) {
@@ -303,7 +317,7 @@ func TestRunnerReportsOpenrestyHealthAndExecutesRestart(t *testing.T) {
healthErr: errors.New("docker openresty container is not running"),
clearHealthOnRestart: true,
}
runner := &Runner{
runner := withHeartbeatCycle(&Runner{
Config: &config.Config{
AccessToken: "agent-token",
NodeName: "edge-01",
@@ -316,7 +330,7 @@ func TestRunnerReportsOpenrestyHealthAndExecutesRestart(t *testing.T) {
HeartbeatService: heartbeatService,
SyncService: &fakeSyncService{},
RuntimeManager: runtimeManager,
}
}, nil)
err := runner.Run(ctx)
if !errors.Is(err, context.Canceled) {
@@ -357,7 +371,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
t.Fatalf("failed to seed state: %v", err)
}
runner := &Runner{
runner := withHeartbeatCycle(&Runner{
Config: &config.Config{
NodeName: "edge-observe-1",
NodeIP: "10.0.0.51",
@@ -369,7 +383,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
},
StateStore: stateStore,
}
}, nil)
if err := os.MkdirAll(filepath.Dir(runner.Config.AccessLogPath), 0o755); err != nil {
t.Fatalf("failed to prepare access log dir: %v", err)
}
@@ -381,7 +395,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
t.Fatalf("failed to prepare access log: %v", err)
}
firstPayload := runner.nodePayload("node-observe")
firstPayload := runner.HeartbeatCycle.NodePayload("node-observe")
if firstPayload.Profile == nil {
t.Fatal("expected first heartbeat payload to include system profile")
}
@@ -398,7 +412,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
t.Fatalf("expected health events for openresty and sync error, got %+v", firstPayload.HealthEvents)
}
secondPayload := runner.nodePayload("node-observe")
secondPayload := runner.HeartbeatCycle.NodePayload("node-observe")
if secondPayload.Profile != nil {
t.Fatal("expected unchanged profile to be omitted on subsequent heartbeat")
}
@@ -439,7 +453,7 @@ func TestRunnerReplaysBufferedObservabilityAfterHeartbeatRecovery(t *testing.T)
}
},
}
runner := &Runner{
runner := withHeartbeatCycle(&Runner{
Config: &config.Config{
AccessToken: "agent-token",
NodeName: "edge-buffer-01",
@@ -451,11 +465,10 @@ func TestRunnerReplaysBufferedObservabilityAfterHeartbeatRecovery(t *testing.T)
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
ObservabilityReplayMinutes: 15,
},
StateStore: stateStore,
ObservabilityBuffer: bufferStore,
HeartbeatService: heartbeatService,
SyncService: &fakeSyncService{},
}
StateStore: stateStore,
HeartbeatService: heartbeatService,
SyncService: &fakeSyncService{},
}, bufferStore)
if err := os.MkdirAll(filepath.Dir(runner.Config.RouteConfigPath), 0o755); err != nil {
t.Fatalf("failed to prepare route config dir: %v", err)
}
@@ -518,7 +531,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
if err != nil {
t.Fatalf("failed to load config: %v", err)
}
runner := &Runner{
runner := withHeartbeatCycle(&Runner{
Config: &config.Config{
ServerURL: cfg.ServerURL,
DiscoveryToken: cfg.DiscoveryToken,
@@ -531,7 +544,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
StateStore: stateStore,
HeartbeatService: heartbeatService,
SyncService: syncService,
}
}, nil)
runner.Config = cfg
runner.Config.Version = config.Version
runner.Config.ExtVersion = "1.27.1.2"
@@ -561,7 +574,7 @@ func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) {
syncService := &fakeSyncService{}
runner := &Runner{SyncService: syncService}
runner := withHeartbeatCycle(&Runner{SyncService: syncService}, nil)
payload, err := json.Marshal(protocol.ActiveConfigMeta{
Version: "20260529-001",
Checksum: "checksum-ws",
@@ -589,12 +602,12 @@ func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) {
}
func TestRunnerHandlesWebSocketSettingsDisabled(t *testing.T) {
runner := &Runner{
runner := withHeartbeatCycle(&Runner{
Config: &config.Config{
HeartbeatInterval: config.MillisecondDuration(10 * time.Second),
},
websocketUpgradeEnabled: true,
}
}, nil)
payload, err := json.Marshal(protocol.AgentSettings{
HeartbeatInterval: 15000,
WebsocketUpgradeEnabled: false,
+2 -70
View File
@@ -1,11 +1,9 @@
package config
import (
"context"
"encoding/json"
"errors"
"fmt"
"net"
"os"
pathpkg "path"
"path/filepath"
@@ -13,8 +11,7 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/geoip"
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
"github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip"
"github.com/Rain-kl/Wavelet/pkg/utils"
)
@@ -35,11 +32,6 @@ const (
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
)
var (
lookupOutboundIP = geoip.GetOutboundIP
lookupLocalIP = detectLocalNodeIP
)
type Config struct {
ServerURL string `json:"server_url"`
AccessToken string `json:"agent_token"`
@@ -166,7 +158,7 @@ func applyDefaults(cfg *Config, baseDir string) {
cfg.NodeName = detectHostname()
}
if cfg.NodeIP == "" {
cfg.NodeIP = detectNodeIP()
cfg.NodeIP = nodeip.Detect()
}
if cfg.MainConfigPath == "" {
cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultMainConfigRelativePath)
@@ -400,64 +392,4 @@ func detectHostname() string {
return strings.TrimSpace(host)
}
func detectNodeIP() string {
if ip := detectOutboundNodeIP(); ip != "" {
return ip
}
return lookupLocalIP()
}
func detectOutboundNodeIP() 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 detectLocalNodeIP() 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 := normalizeIPv4(ipNet.IP)
priority := nodeIPPriority(ipv4)
if priority > bestPriority {
bestIP = ipv4.String()
bestPriority = priority
}
if bestPriority == 2 {
return bestIP
}
}
}
return bestIP
}
func normalizeIPv4(ip net.IP) net.IP {
if ip == nil {
return nil
}
return ip.To4()
}
func nodeIPPriority(ip net.IP) int {
return iputil.Score(ip)
}
+12 -10
View File
@@ -10,7 +10,9 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/edge/nodeip"
"github.com/Rain-kl/Wavelet/pkg/geoip"
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
)
func TestLoadDefaultsToManagedBinaryPaths(t *testing.T) {
@@ -248,12 +250,12 @@ func TestLoadUsesEnvConfigWhenFileIsMissing(t *testing.T) {
}
func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) {
previousLookup := lookupOutboundIP
lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
previousLookup := nodeip.LookupOutboundIP
nodeip.LookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
return net.ParseIP("8.8.8.8"), nil
}
defer func() {
lookupOutboundIP = previousLookup
nodeip.LookupOutboundIP = previousLookup
}()
dir := t.TempDir()
@@ -281,17 +283,17 @@ func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) {
}
func TestLoadFallsBackToLocalIPWhenOutboundLookupFails(t *testing.T) {
previousOutboundLookup := lookupOutboundIP
previousLocalLookup := lookupLocalIP
lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
previousOutboundLookup := nodeip.LookupOutboundIP
previousLocalLookup := nodeip.LookupLocalIP
nodeip.LookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
return nil, errors.New("realip.cc unavailable")
}
lookupLocalIP = func() string {
nodeip.LookupLocalIP = func() string {
return "9.9.9.9"
}
defer func() {
lookupOutboundIP = previousOutboundLookup
lookupLocalIP = previousLocalLookup
nodeip.LookupOutboundIP = previousOutboundLookup
nodeip.LookupLocalIP = previousLocalLookup
}()
dir := t.TempDir()
@@ -512,7 +514,7 @@ func TestNodeIPPriority(t *testing.T) {
if tt.ip != "" {
parsed = net.ParseIP(tt.ip)
}
if got := nodeIPPriority(parsed); got != tt.expected {
if got := iputil.Score(parsed); got != tt.expected {
t.Fatalf("unexpected priority for %q: got %d want %d", tt.ip, got, tt.expected)
}
})
+2 -51
View File
@@ -1,54 +1,5 @@
package config
import (
"encoding/json"
"fmt"
"strconv"
"strings"
"time"
)
import edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config"
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())
}
type MillisecondDuration = edgeconfig.MillisecondDuration
+217
View File
@@ -0,0 +1,217 @@
package heartbeat
import (
"context"
"log/slog"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config"
"github.com/Rain-kl/Wavelet/internal/apps/agent/observability"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
"github.com/Rain-kl/Wavelet/internal/apps/agent/updater"
edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat"
)
type HeartbeatClient interface {
Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error)
}
type SyncService interface {
SyncOnStartup(ctx context.Context, target *protocol.ActiveConfigMeta) error
SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error
WAFIPGroupChecksums() (map[string]string, error)
ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error
}
type SettingsApplier interface {
Apply(settings *protocol.AgentSettings) (intervalChanged bool)
RestartOpenrestyIfNeeded(ctx context.Context)
}
type Cycle struct {
Config *config.Config
StateStore *state.Store
ObservabilityBuffer *state.ObservabilityBufferStore
Heartbeat HeartbeatClient
Sync SyncService
Updater *updater.Service
RecordSyncError func(err error)
}
func (c *Cycle) Perform(ctx context.Context, nodeID string, startup bool, settings SettingsApplier) (bool, error) {
payload, ackWindows := c.PrepareHeartbeatPayload(nodeID)
heartbeatResult, err := c.Heartbeat.Heartbeat(ctx, payload)
if err != nil {
return false, err
}
c.AckObservabilityWindows(ackWindows)
if heartbeatResult == nil {
heartbeatResult = &protocol.HeartbeatResult{}
}
mode := "periodic"
if startup {
mode = "startup"
}
slog.Debug("agent heartbeat succeeded", "mode", mode, "node_id", nodeID)
var changed bool
if settings != nil {
changed = settings.Apply(heartbeatResult.AgentSettings)
}
c.ApplyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups)
if startup {
if err = c.Sync.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil {
c.recordSyncError(err)
slog.Error("agent startup sync failed", "error", err)
} else {
slog.Debug("agent startup sync completed")
}
} else if err = c.Sync.SyncOnce(ctx, heartbeatResult.ActiveConfig); err != nil {
c.recordSyncError(err)
slog.Error("agent sync failed", "error", err)
}
if settings != nil {
settings.RestartOpenrestyIfNeeded(ctx)
}
edgeheartbeat.TryAutoUpdate(ctx, c.Updater, agentSettingsToAutoUpdate(heartbeatResult.AgentSettings), "agent")
return changed, nil
}
func (c *Cycle) NodePayload(nodeID string) protocol.NodePayload {
snapshot, _ := c.StateStore.Load()
openrestyStatus := strings.TrimSpace(snapshot.OpenrestyStatus)
if openrestyStatus == "" {
openrestyStatus = protocol.OpenrestyStatusUnknown
}
profile := observability.BuildProfile(c.Config, c.StateStore)
managedOpenRestyMetrics := observability.CollectManagedOpenRestyMetrics(c.Config)
trafficReport, accessLogs, fallbackMetrics := observability.BuildTrafficObservability(c.Config, c.StateStore, managedOpenRestyMetrics)
if managedOpenRestyMetrics == nil {
managedOpenRestyMetrics = fallbackMetrics
}
metricSnapshot := observability.BuildSnapshot(c.Config, c.StateStore)
openrestyObservation := observability.BuildOpenrestyObservation(managedOpenRestyMetrics)
healthEvents := observability.BuildHealthEvents(snapshot)
payload := protocol.NodePayload{
NodeID: nodeID,
Name: c.Config.NodeName,
IP: c.Config.NodeIP,
Version: c.Config.Version,
ExtVersion: c.Config.ExtVersion,
CurrentVersion: snapshot.CurrentVersion,
LastError: snapshot.LastError,
OpenrestyStatus: openrestyStatus,
OpenrestyMessage: snapshot.OpenrestyMessage,
Profile: profile,
Snapshot: metricSnapshot,
OpenrestyObservation: openrestyObservation,
TrafficReport: trafficReport,
AccessLogs: accessLogs,
HealthEvents: healthEvents,
}
if c.Sync != nil {
checksums, err := c.Sync.WAFIPGroupChecksums()
if err != nil {
slog.Debug("load local waf ip group checksums failed", "error", err)
} else if len(checksums) > 0 {
payload.WAFIPGroupChecksums = checksums
}
}
return payload
}
func (c *Cycle) PrepareHeartbeatPayload(nodeID string) (protocol.NodePayload, []int64) {
payload := c.NodePayload(nodeID)
if c.ObservabilityBuffer == nil || (payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0) {
return payload, nil
}
now := time.Now().UTC()
retainAfterUnix := now.Add(-time.Duration(c.Config.ObservabilityReplayMinutes) * time.Minute).Unix()
windowStartedAtUnix := state.ObservabilityWindowStartedAt(payload.Snapshot, payload.OpenrestyObservation, payload.TrafficReport)
if windowStartedAtUnix <= 0 {
return payload, nil
}
record := state.ObservabilityBufferRecord{
WindowStartedAtUnix: windowStartedAtUnix,
Snapshot: payload.Snapshot,
OpenrestyObservation: payload.OpenrestyObservation,
TrafficReport: payload.TrafficReport,
AccessLogs: payload.AccessLogs,
QueuedAtUnix: now.Unix(),
}
if err := c.ObservabilityBuffer.Upsert(record, retainAfterUnix); err != nil {
slog.Error("upsert observability buffer failed", "error", err)
return payload, nil
}
records, err := c.ObservabilityBuffer.Replayable(windowStartedAtUnix, retainAfterUnix)
if err != nil {
slog.Error("load replayable observability buffer failed", "error", err)
return payload, []int64{windowStartedAtUnix}
}
ackWindows := make([]int64, 0, len(records)+1)
buffered := make([]protocol.BufferedObservabilityRecord, 0, len(records))
for _, item := range records {
if item.WindowStartedAtUnix <= 0 {
continue
}
buffered = append(buffered, protocol.BufferedObservabilityRecord{
WindowStartedAtUnix: item.WindowStartedAtUnix,
Snapshot: item.Snapshot,
OpenrestyObservation: item.OpenrestyObservation,
TrafficReport: item.TrafficReport,
AccessLogs: item.AccessLogs,
})
ackWindows = append(ackWindows, item.WindowStartedAtUnix)
}
payload.BufferedObservability = buffered
ackWindows = append(ackWindows, windowStartedAtUnix)
return payload, ackWindows
}
func (c *Cycle) AckObservabilityWindows(windowStartedAtUnix []int64) {
if c.ObservabilityBuffer == nil || len(windowStartedAtUnix) == 0 {
return
}
retainAfterUnix := time.Now().UTC().Add(-time.Duration(c.Config.ObservabilityReplayMinutes) * time.Minute).Unix()
if err := c.ObservabilityBuffer.Ack(windowStartedAtUnix, retainAfterUnix); err != nil {
slog.Error("ack observability buffer failed", "error", err)
}
}
func (c *Cycle) ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) {
if len(groups) == 0 || c.Sync == nil {
return
}
if err := c.Sync.ApplyWAFIPGroups(ctx, groups); err != nil {
c.recordSyncError(err)
slog.Error("agent apply waf ip groups failed", "error", err)
}
}
func (c *Cycle) recordSyncError(err error) {
if c.RecordSyncError != nil {
c.RecordSyncError(err)
}
}
func AgentSettingsToAutoUpdate(settings *protocol.AgentSettings) *edgeheartbeat.AutoUpdateSettings {
if settings == nil {
return nil
}
return &edgeheartbeat.AutoUpdateSettings{
AutoUpdate: settings.AutoUpdate,
UpdateNow: settings.UpdateNow,
UpdateRepo: settings.UpdateRepo,
UpdateChannel: settings.UpdateChannel,
UpdateTag: settings.UpdateTag,
}
}
func agentSettingsToAutoUpdate(settings *protocol.AgentSettings) *edgeheartbeat.AutoUpdateSettings {
return AgentSettingsToAutoUpdate(settings)
}
+17 -116
View File
@@ -1,55 +1,44 @@
package httpclient
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"time"
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
)
type Client struct {
baseURL string
token string
httpClient *http.Client
base *edgehttp.Client
}
func New(baseURL string, token string, timeout time.Duration) *Client {
return &Client{
baseURL: strings.TrimRight(baseURL, "/"),
token: token,
httpClient: &http.Client{
Timeout: timeout,
},
base: edgehttp.New(baseURL, token, timeout, "X-Agent-Token"),
}
}
func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
slog.Debug("http register node request", "node_id", payload.NodeID, "current_version", payload.CurrentVersion)
resp := protocol.APIResponse[protocol.RegisterNodeResponse]{}
if err := c.postJSON(ctx, "/api/v1/agent/nodes/register", payload, &resp); err != nil {
if err := c.base.PostJSON(ctx, "/api/v1/agent/nodes/register", payload, &resp); err != nil {
return nil, err
}
if err := apiError(resp.ErrorMsg); err != nil {
if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err
}
slog.Debug("http register node response", "node_id", resp.Data.NodeID)
return &resp.Data, nil
}
func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) {
resp := protocol.APIResponse[protocol.HeartbeatData]{}
if err := c.postJSON(ctx, "/api/v1/agent/nodes/heartbeat", payload, &resp); err != nil {
if err := c.base.PostJSON(ctx, "/api/v1/agent/nodes/heartbeat", payload, &resp); err != nil {
return nil, err
}
if err := apiError(resp.ErrorMsg); err != nil {
if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err
}
return &protocol.HeartbeatResult{
@@ -61,134 +50,46 @@ func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*
func (c *Client) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) {
resp := protocol.APIResponse[protocol.ActiveConfigResponse]{}
if err := c.getJSON(ctx, "/api/v1/agent/config-versions/active", &resp); err != nil {
if err := c.base.GetJSON(ctx, "/api/v1/agent/config-versions/active", &resp); err != nil {
return nil, err
}
if err := apiError(resp.ErrorMsg); err != nil {
if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err
}
slog.Debug("http get active config response", "version", resp.Data.Version, "checksum", resp.Data.Checksum, "support_files", len(resp.Data.SupportFiles))
return &resp.Data, nil
}
func (c *Client) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
slog.Debug("http report apply log request", "node_id", payload.NodeID, "version", payload.Version, "result", payload.Result)
resp := protocol.APIResponse[json.RawMessage]{}
if err := c.postJSON(ctx, "/api/v1/agent/apply-logs", payload, &resp); err != nil {
if err := c.base.PostJSON(ctx, "/api/v1/agent/apply-logs", payload, &resp); err != nil {
return err
}
return apiError(resp.ErrorMsg)
return edgehttp.APIError(resp.ErrorMsg)
}
func (c *Client) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) {
resp := protocol.APIResponse[protocol.WAFIPGroupSyncResponse]{}
if err := c.postJSON(ctx, "/api/v1/agent/waf/ip-groups/sync", payload, &resp); err != nil {
if err := c.base.PostJSON(ctx, "/api/v1/agent/waf/ip-groups/sync", payload, &resp); err != nil {
return nil, err
}
if err := apiError(resp.ErrorMsg); err != nil {
if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
return nil, err
}
return &resp.Data, nil
}
func (c *Client) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID), nil)
if err != nil {
return nil, err
}
req.Header.Set("X-Agent-Token", c.token)
res, err := c.httpClient.Do(req)
res, err := c.base.DoRaw(ctx, http.MethodGet, fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID), nil)
if err != nil {
return nil, err
}
defer res.Body.Close()
if res.StatusCode != http.StatusOK {
return nil, readHTTPError(res)
return nil, edgehttp.ReadHTTPError(res)
}
return io.ReadAll(res.Body)
}
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
}
req.Header.Set("X-Agent-Token", c.token)
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")
req.Header.Set("X-Agent-Token", c.token)
return c.do(req, target)
}
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 readHTTPError(res *http.Response) error {
body, err := io.ReadAll(res.Body)
if err != nil {
return errors.New(res.Status)
}
return readBodyError(body, res.Status)
}
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)
}
c.base.SetToken(token)
}
+3 -25
View File
@@ -1,29 +1,7 @@
package logging
import (
"log/slog"
"os"
"strings"
)
import edgelogging "github.com/Rain-kl/Wavelet/internal/apps/edge/logging"
func Setup() {
opts := &slog.HandlerOptions{
AddSource: true,
Level: parseLevel(os.Getenv("LOG_LEVEL")),
}
handler := slog.NewTextHandler(os.Stdout, opts)
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
}
}
edgelogging.Setup(edgelogging.Options{AddSource: true})
}
+13 -267
View File
@@ -1,21 +1,18 @@
package observability
import (
"bufio"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"syscall"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
edgeobs "github.com/Rain-kl/Wavelet/internal/apps/edge/observability"
)
func BuildProfile(cfg *config.Config, stateStore *state.Store) *protocol.NodeSystemProfile {
@@ -47,22 +44,22 @@ func BuildSnapshot(cfg *config.Config, stateStore *state.Store) *protocol.NodeMe
CapturedAtUnix: now.Unix(),
}
memTotal, memUsed := readMemInfo()
memTotal, memUsed := edgeobs.ReadMemInfo()
metric.MemoryTotalBytes = memTotal
metric.MemoryUsedBytes = memUsed
storageTotal, storageUsed := statFilesystem(cfg.DataDir)
storageTotal, storageUsed := edgeobs.StatFilesystem(cfg.DataDir)
metric.StorageTotalBytes = storageTotal
metric.StorageUsedBytes = storageUsed
metric.NetworkRxBytes, metric.NetworkTxBytes = readLinuxNetworkTotals()
metric.DiskReadBytes, metric.DiskWriteBytes = readLinuxDiskTotals()
metric.NetworkRxBytes, metric.NetworkTxBytes = edgeobs.ReadLinuxNetworkTotals()
metric.DiskReadBytes, metric.DiskWriteBytes = edgeobs.ReadLinuxDiskTotals()
if stateStore == nil {
return metric
}
totalCPU, idleCPU := readLinuxCPUStat()
totalCPU, idleCPU := edgeobs.ReadLinuxCPUStat()
snapshot, err := stateStore.Load()
if err != nil {
return metric
@@ -121,12 +118,12 @@ func BuildHealthEvents(snapshot *state.Snapshot) []protocol.NodeHealthEvent {
func collectProfile(cfg *config.Config) *protocol.NodeSystemProfile {
hostname, _ := os.Hostname()
osName, osVersion := readLinuxOSRelease()
kernelVersion := readFirstLine("/proc/sys/kernel/osrelease")
cpuModel := readLinuxCPUModel()
totalMemory, _ := readMemInfo()
totalDisk, _ := statFilesystem(cfg.DataDir)
uptimeSeconds := readLinuxUptimeSeconds()
osName, osVersion := edgeobs.ReadLinuxOSRelease()
kernelVersion := edgeobs.ReadFirstLine("/proc/sys/kernel/osrelease")
cpuModel := edgeobs.ReadLinuxCPUModel()
totalMemory, _ := edgeobs.ReadMemInfo()
totalDisk, _ := edgeobs.StatFilesystem(cfg.DataDir)
uptimeSeconds := edgeobs.ReadLinuxUptimeSeconds()
return &protocol.NodeSystemProfile{
Hostname: strings.TrimSpace(hostname),
@@ -150,255 +147,4 @@ func fingerprintProfile(profile *protocol.NodeSystemProfile) string {
}
sum := sha256.Sum256(raw)
return hex.EncodeToString(sum[:])
}
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"])
}
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 ""
}
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
}
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)
}
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
if i == 4 {
// idle
}
}
idle, err := strconv.ParseUint(fields[4], 10, 64)
if err != nil {
return 0, 0
}
return total, idle
}
return 0, 0
}
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
}
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
}
}
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
}
func readFirstLine(path string) string {
content, err := os.ReadFile(path)
if err != nil {
return ""
}
return strings.TrimSpace(string(content))
}
}
-207
View File
@@ -1,207 +0,0 @@
package protocol
import "encoding/json"
type APIResponse[T any] struct {
ErrorMsg string `json:"error_msg"`
Data T `json:"data"`
}
type HeartbeatData struct {
AgentSettings *AgentSettings `json:"agent_settings"`
ActiveConfig *ActiveConfigMeta `json:"active_config"`
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
}
type HeartbeatResult struct {
AgentSettings *AgentSettings
ActiveConfig *ActiveConfigMeta
WAFIPGroups []WAFIPGroup
}
type AgentSettings struct {
HeartbeatInterval int `json:"heartbeat_interval"`
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
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"`
RestartOpenrestyNow bool `json:"restart_openresty_now"`
}
const (
WSMessageTypeStatus = "status"
WSMessageTypeSettings = "settings"
WSMessageTypeActiveConfig = "active_config"
WSMessageTypeForceSyncConfig = "force_sync_config"
WSMessageTypeWAFIPGroups = "waf_ip_groups"
WSMessageTypePing = "ping"
WSMessageTypePong = "pong"
)
type WSMessage struct {
Type string `json:"type"`
Payload json.RawMessage `json:"payload,omitempty"`
}
type WSOutboundMessage struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}
type WebSocketConnection interface {
URL() string
SendStatus(payload NodePayload) error
SendPong() error
Receive() (WSMessage, error)
Close() error
}
const (
OpenrestyStatusHealthy = "healthy"
OpenrestyStatusUnhealthy = "unhealthy"
OpenrestyStatusUnknown = "unknown"
)
type NodePayload struct {
NodeID string `json:"node_id"`
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
ExtVersion string `json:"ext_version"`
CurrentVersion string `json:"current_version"`
LastError string `json:"last_error"`
OpenrestyStatus string `json:"openresty_status"`
OpenrestyMessage string `json:"openresty_message"`
Profile *NodeSystemProfile `json:"profile,omitempty"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
BufferedObservability []BufferedObservabilityRecord `json:"buffered_observability,omitempty"`
HealthEvents []NodeHealthEvent `json:"health_events"`
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
}
type NodeSystemProfile struct {
Hostname string `json:"hostname"`
OSName string `json:"os_name"`
OSVersion string `json:"os_version"`
KernelVersion string `json:"kernel_version"`
Architecture string `json:"architecture"`
CPUModel string `json:"cpu_model"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAtUnix int64 `json:"reported_at_unix"`
}
type NodeMetricSnapshot struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
MemoryTotalBytes int64 `json:"memory_total_bytes"`
StorageUsedBytes int64 `json:"storage_used_bytes"`
StorageTotalBytes int64 `json:"storage_total_bytes"`
DiskReadBytes int64 `json:"disk_read_bytes"`
DiskWriteBytes int64 `json:"disk_write_bytes"`
NetworkRxBytes int64 `json:"network_rx_bytes"`
NetworkTxBytes int64 `json:"network_tx_bytes"`
}
type NodeOpenrestyObservation struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
OpenrestyConnections int64 `json:"openresty_connections"`
}
type NodeTrafficReport struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
RequestCount int64 `json:"request_count"`
ErrorCount int64 `json:"error_count"`
UniqueVisitorCount int64 `json:"unique_visitor_count"`
StatusCodes map[string]int64 `json:"status_codes"`
TopDomains map[string]int64 `json:"top_domains"`
SourceCountries map[string]int64 `json:"source_countries"`
}
type NodeAccessLog struct {
LoggedAtUnix int64 `json:"logged_at_unix"`
RemoteAddr string `json:"remote_addr"`
Host string `json:"host"`
Path string `json:"path"`
StatusCode int `json:"status_code"`
}
type BufferedObservabilityRecord struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
}
type NodeHealthEvent struct {
EventType string `json:"event_type"`
Severity string `json:"severity"`
Message string `json:"message"`
TriggeredAtUnix int64 `json:"triggered_at_unix"`
Metadata map[string]string `json:"metadata,omitempty"`
}
type RegisterNodeResponse struct {
NodeID string `json:"node_id"`
AccessToken string `json:"agent_token"`
Name string `json:"name"`
}
type ApplyLogPayload struct {
NodeID string `json:"node_id"`
Version string `json:"version"`
Result string `json:"result"`
Message string `json:"message"`
Checksum string `json:"checksum"`
MainConfigChecksum string `json:"main_config_checksum"`
RouteConfigChecksum string `json:"route_config_checksum"`
SupportFileCount int `json:"support_file_count"`
}
type ActiveConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
SourceConfigJSON string `json:"source_config_json"`
SupportFiles []SupportFile `json:"support_files"`
CreatedAt string `json:"created_at"`
}
type ActiveConfigMeta struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
Checksum string `json:"checksum"`
}
type WAFIPGroupSyncRequest struct {
IDs []uint `json:"ids,omitempty"`
Checksums map[string]string `json:"checksums,omitempty"`
}
type WAFIPGroupSyncResponse struct {
Groups []WAFIPGroup `json:"groups"`
}
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
+43
View File
@@ -0,0 +1,43 @@
package protocol
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
type APIResponse[T any] = pkgprotocol.APIResponse[T]
type HeartbeatData = pkgprotocol.HeartbeatData
type HeartbeatResult = pkgprotocol.HeartbeatResult
type AgentSettings = pkgprotocol.AgentSettings
type WSMessage = pkgprotocol.WSMessage
type WSOutboundMessage = pkgprotocol.WSOutboundMessage
type WebSocketConnection = pkgprotocol.WebSocketConnection
type NodePayload = pkgprotocol.NodePayload
type NodeSystemProfile = pkgprotocol.NodeSystemProfile
type NodeMetricSnapshot = pkgprotocol.NodeMetricSnapshot
type NodeOpenrestyObservation = pkgprotocol.NodeOpenrestyObservation
type NodeTrafficReport = pkgprotocol.NodeTrafficReport
type NodeAccessLog = pkgprotocol.NodeAccessLog
type BufferedObservabilityRecord = pkgprotocol.BufferedObservabilityRecord
type NodeHealthEvent = pkgprotocol.NodeHealthEvent
type RegisterNodeResponse = pkgprotocol.RegisterNodeResponse
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
type ActiveConfigResponse = pkgprotocol.ActiveConfigResponse
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
type WAFIPGroup = pkgprotocol.WAFIPGroup
type WAFIPGroupSyncRequest = pkgprotocol.WAFIPGroupSyncRequest
type WAFIPGroupSyncResponse = pkgprotocol.WAFIPGroupSyncResponse
type SupportFile = pkgprotocol.SupportFile
const (
WSMessageTypeStatus = pkgprotocol.WSMessageTypeStatus
WSMessageTypeSettings = pkgprotocol.WSMessageTypeSettings
WSMessageTypeActiveConfig = pkgprotocol.WSMessageTypeActiveConfig
WSMessageTypeForceSyncConfig = pkgprotocol.WSMessageTypeForceSyncConfig
WSMessageTypeWAFIPGroups = pkgprotocol.WSMessageTypeWAFIPGroups
WSMessageTypePing = pkgprotocol.WSMessageTypePing
WSMessageTypePong = pkgprotocol.WSMessageTypePong
)
const (
OpenrestyStatusHealthy = pkgprotocol.OpenrestyStatusHealthy
OpenrestyStatusUnhealthy = pkgprotocol.OpenrestyStatusUnhealthy
OpenrestyStatusUnknown = pkgprotocol.OpenrestyStatusUnknown
)
@@ -1,51 +0,0 @@
//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
}
@@ -1,15 +0,0 @@
//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)
}
}
@@ -1,53 +0,0 @@
//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, `"`, `""`) + `"`
}
+9 -358
View File
@@ -1,366 +1,17 @@
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"
"github.com/Rain-kl/Wavelet/internal/apps/agent/agent"
edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
"github.com/Rain-kl/Wavelet/internal/apps/agent/config"
)
const maxChecksumAssetSize = 64 * 1024
var replaceAndRestartFunc = replaceAndRestart
type Service struct {
httpClient *http.Client
lastCheckKey string
}
type Service = edgeupdater.Service
type UpdateOptions = edgeupdater.UpdateOptions
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)
}
return edgeupdater.New(edgeupdater.Config{
LocalVersion: config.Version,
AssetPrefix: "openflare-agent",
LogLabel: "agent",
})
}
-242
View File
@@ -1,242 +0,0 @@
package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/agent/agent"
"github.com/Rain-kl/Wavelet/internal/apps/agent/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)
}
})
}
}