mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 15:46:37 +08:00
[优化] 改名
This commit is contained in:
@@ -0,0 +1,404 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"openflare-agent/internal/config"
|
||||
"openflare-agent/internal/observability"
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
)
|
||||
|
||||
type HeartbeatService interface {
|
||||
Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error)
|
||||
Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error)
|
||||
SetToken(token string)
|
||||
}
|
||||
|
||||
type SyncService interface {
|
||||
SyncOnStartup(ctx context.Context, target *protocol.ActiveConfigMeta) error
|
||||
SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) 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
|
||||
}
|
||||
|
||||
type UpdateOptions struct {
|
||||
Channel string
|
||||
TagName string
|
||||
Force bool
|
||||
}
|
||||
|
||||
type Runner struct {
|
||||
Config *config.Config
|
||||
StateStore *state.Store
|
||||
ObservabilityBuffer *state.ObservabilityBufferStore
|
||||
HeartbeatService HeartbeatService
|
||||
SyncService SyncService
|
||||
Updater Updater
|
||||
RuntimeManager RuntimeManager
|
||||
|
||||
autoUpdate bool
|
||||
updateNow bool
|
||||
updateRepo string
|
||||
updateChan string
|
||||
updateTag string
|
||||
restartOpenrestyNow bool
|
||||
}
|
||||
|
||||
func (r *Runner) Run(ctx context.Context) error {
|
||||
nodeID, err := r.StateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
slog.Info("agent runner started", "node_id", nodeID, "node", r.Config.NodeName, "ip", r.Config.NodeIP)
|
||||
if r.hasAgentToken() {
|
||||
r.refreshOpenrestyHealth(ctx)
|
||||
payload, ackWindows := r.prepareHeartbeatPayload(nodeID)
|
||||
heartbeatResult, hbErr := r.HeartbeatService.Heartbeat(ctx, payload)
|
||||
if hbErr != nil {
|
||||
slog.Error("agent startup heartbeat failed", "error", hbErr)
|
||||
} else {
|
||||
r.ackObservabilityWindows(ackWindows)
|
||||
if heartbeatResult == nil {
|
||||
heartbeatResult = &protocol.HeartbeatResult{}
|
||||
}
|
||||
slog.Debug("agent startup heartbeat succeeded", "node_id", nodeID)
|
||||
r.applySettings(heartbeatResult.AgentSettings)
|
||||
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")
|
||||
}
|
||||
r.tryRestartOpenresty(ctx)
|
||||
r.tryAutoUpdate(ctx)
|
||||
}
|
||||
} else if err = r.tryRegister(ctx, &nodeID); err != nil {
|
||||
slog.Error("agent initial discovery register failed", "error", err)
|
||||
}
|
||||
|
||||
heartbeatTicker := time.NewTicker(r.Config.HeartbeatInterval.Duration())
|
||||
defer heartbeatTicker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
slog.Info("agent runner shutting down", "error", ctx.Err())
|
||||
return ctx.Err()
|
||||
case <-heartbeatTicker.C:
|
||||
if !r.hasAgentToken() {
|
||||
if err = r.tryRegister(ctx, &nodeID); err != nil {
|
||||
slog.Error("agent discovery register failed", "error", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
r.refreshOpenrestyHealth(ctx)
|
||||
payload, ackWindows := r.prepareHeartbeatPayload(nodeID)
|
||||
heartbeatResult, hbErr := r.HeartbeatService.Heartbeat(ctx, payload)
|
||||
if hbErr != nil {
|
||||
slog.Error("agent heartbeat failed", "error", hbErr)
|
||||
} else {
|
||||
r.ackObservabilityWindows(ackWindows)
|
||||
if heartbeatResult == nil {
|
||||
heartbeatResult = &protocol.HeartbeatResult{}
|
||||
}
|
||||
if changed := r.applySettings(heartbeatResult.AgentSettings); changed {
|
||||
heartbeatTicker.Reset(r.Config.HeartbeatInterval.Duration())
|
||||
}
|
||||
if err = r.SyncService.SyncOnce(ctx, heartbeatResult.ActiveConfig); err != nil {
|
||||
r.recordSyncError(err)
|
||||
slog.Error("agent sync failed", "error", err)
|
||||
}
|
||||
r.tryRestartOpenresty(ctx)
|
||||
r.tryAutoUpdate(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) hasAgentToken() bool {
|
||||
return strings.TrimSpace(r.Config.AgentToken) != ""
|
||||
}
|
||||
|
||||
func (r *Runner) applySettings(settings *protocol.AgentSettings) bool {
|
||||
if settings == nil {
|
||||
return false
|
||||
}
|
||||
changed := false
|
||||
if settings.HeartbeatInterval > 0 {
|
||||
newInterval := config.MillisecondDuration(time.Duration(settings.HeartbeatInterval) * time.Millisecond)
|
||||
if newInterval != r.Config.HeartbeatInterval {
|
||||
slog.Info("agent heartbeat interval updated", "from", r.Config.HeartbeatInterval, "to", newInterval)
|
||||
r.Config.HeartbeatInterval = newInterval
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
func (r *Runner) tryRestartOpenresty(ctx context.Context) {
|
||||
if !r.restartOpenrestyNow {
|
||||
return
|
||||
}
|
||||
r.restartOpenrestyNow = false
|
||||
if r.RuntimeManager == nil {
|
||||
return
|
||||
}
|
||||
slog.Info("agent openresty restart requested by server")
|
||||
if err := r.RuntimeManager.Restart(ctx); err != nil {
|
||||
slog.Error("agent openresty restart failed", "error", err)
|
||||
r.recordOpenrestyUnhealthy(err, false)
|
||||
return
|
||||
}
|
||||
slog.Info("agent openresty restart succeeded")
|
||||
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))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if response == nil || strings.TrimSpace(response.AgentToken) == "" || strings.TrimSpace(response.NodeID) == "" {
|
||||
return errors.New("discovery register response 缺少 node_id 或 agent_token")
|
||||
}
|
||||
snapshot, err := r.StateStore.Load()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
snapshot.NodeID = response.NodeID
|
||||
if err = r.StateStore.Save(snapshot); err != nil {
|
||||
return err
|
||||
}
|
||||
r.Config.AgentToken = response.AgentToken
|
||||
r.Config.DiscoveryToken = ""
|
||||
if err = r.Config.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
r.HeartbeatService.SetToken(response.AgentToken)
|
||||
*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
|
||||
}
|
||||
r.ackObservabilityWindows(ackWindows)
|
||||
if heartbeatResult == nil {
|
||||
heartbeatResult = &protocol.HeartbeatResult{}
|
||||
}
|
||||
r.applySettings(heartbeatResult.AgentSettings)
|
||||
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
|
||||
}
|
||||
|
||||
func (r *Runner) recordSyncError(err error) {
|
||||
if err == nil || r.StateStore == nil {
|
||||
return
|
||||
}
|
||||
snapshot, loadErr := r.StateStore.Load()
|
||||
if loadErr != nil {
|
||||
slog.Error("load state before recording sync error failed", "error", loadErr)
|
||||
return
|
||||
}
|
||||
snapshot.LastError = err.Error()
|
||||
slog.Warn("recording sync error into state", "error", snapshot.LastError)
|
||||
if saveErr := r.StateStore.Save(snapshot); saveErr != nil {
|
||||
slog.Error("save state after sync error failed", "error", saveErr)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) refreshOpenrestyHealth(ctx context.Context) {
|
||||
if r.RuntimeManager == nil || r.StateStore == nil {
|
||||
return
|
||||
}
|
||||
if err := r.RuntimeManager.CheckHealth(ctx); err != nil {
|
||||
r.recordOpenrestyUnhealthy(err, true)
|
||||
return
|
||||
}
|
||||
r.recordOpenrestyHealthy()
|
||||
}
|
||||
|
||||
func (r *Runner) recordOpenrestyHealthy() {
|
||||
if r.StateStore == nil {
|
||||
return
|
||||
}
|
||||
snapshot, err := r.StateStore.Load()
|
||||
if err != nil {
|
||||
slog.Error("load state before recording openresty health failed", "error", err)
|
||||
return
|
||||
}
|
||||
if snapshot.OpenrestyStatus == protocol.OpenrestyStatusHealthy && strings.TrimSpace(snapshot.OpenrestyMessage) == "" {
|
||||
return
|
||||
}
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
|
||||
snapshot.OpenrestyMessage = ""
|
||||
if err = r.StateStore.Save(snapshot); err != nil {
|
||||
slog.Error("save state after recording openresty health failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) recordOpenrestyUnhealthy(err error, fallbackOnly bool) {
|
||||
if err == nil || r.StateStore == nil {
|
||||
return
|
||||
}
|
||||
snapshot, loadErr := r.StateStore.Load()
|
||||
if loadErr != nil {
|
||||
slog.Error("load state before recording openresty error failed", "error", loadErr)
|
||||
return
|
||||
}
|
||||
message := strings.TrimSpace(err.Error())
|
||||
if !fallbackOnly || strings.TrimSpace(snapshot.OpenrestyMessage) == "" {
|
||||
snapshot.OpenrestyMessage = message
|
||||
}
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
|
||||
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, managedOpenRestyMetrics)
|
||||
healthEvents := observability.BuildHealthEvents(snapshot)
|
||||
return protocol.NodePayload{
|
||||
NodeID: nodeID,
|
||||
Name: r.Config.NodeName,
|
||||
IP: r.Config.NodeIP,
|
||||
AgentVersion: r.Config.AgentVersion,
|
||||
NginxVersion: r.Config.NginxVersion,
|
||||
CurrentVersion: snapshot.CurrentVersion,
|
||||
LastError: snapshot.LastError,
|
||||
OpenrestyStatus: openrestyStatus,
|
||||
OpenrestyMessage: snapshot.OpenrestyMessage,
|
||||
Profile: profile,
|
||||
Snapshot: metricSnapshot,
|
||||
TrafficReport: trafficReport,
|
||||
AccessLogs: accessLogs,
|
||||
HealthEvents: healthEvents,
|
||||
}
|
||||
}
|
||||
|
||||
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.TrafficReport)
|
||||
if windowStartedAtUnix <= 0 {
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
record := state.ObservabilityBufferRecord{
|
||||
WindowStartedAtUnix: windowStartedAtUnix,
|
||||
Snapshot: payload.Snapshot,
|
||||
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,
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,496 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"openflare-agent/internal/config"
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
)
|
||||
|
||||
type fakeHeartbeatService struct {
|
||||
mu sync.Mutex
|
||||
registerCalls int
|
||||
heartbeatCalls int
|
||||
registerErr error
|
||||
registerResp *protocol.RegisterNodeResponse
|
||||
heartbeatErrs []error
|
||||
heartbeatResults []*protocol.HeartbeatResult
|
||||
heartbeatPayloads []protocol.NodePayload
|
||||
onHeartbeat func(int)
|
||||
lastToken string
|
||||
}
|
||||
|
||||
func (f *fakeHeartbeatService) Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.registerCalls++
|
||||
return f.registerResp, f.registerErr
|
||||
}
|
||||
|
||||
func (f *fakeHeartbeatService) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) {
|
||||
f.mu.Lock()
|
||||
f.heartbeatCalls++
|
||||
callIndex := f.heartbeatCalls
|
||||
f.heartbeatPayloads = append(f.heartbeatPayloads, payload)
|
||||
var err error
|
||||
if len(f.heartbeatErrs) >= callIndex {
|
||||
err = f.heartbeatErrs[callIndex-1]
|
||||
}
|
||||
var result *protocol.HeartbeatResult
|
||||
if len(f.heartbeatResults) >= callIndex {
|
||||
result = f.heartbeatResults[callIndex-1]
|
||||
}
|
||||
onHeartbeat := f.onHeartbeat
|
||||
f.mu.Unlock()
|
||||
if onHeartbeat != nil {
|
||||
onHeartbeat(callIndex)
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
|
||||
func (f *fakeHeartbeatService) SetToken(token string) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.lastToken = token
|
||||
}
|
||||
|
||||
type fakeSyncService struct {
|
||||
mu sync.Mutex
|
||||
startupErr error
|
||||
syncOnceErr error
|
||||
startupCalls int
|
||||
syncOnceCalls int
|
||||
onSyncOnceCall func(int)
|
||||
}
|
||||
|
||||
type fakeRuntimeManager struct {
|
||||
mu sync.Mutex
|
||||
healthErr error
|
||||
restartErr error
|
||||
restartCalls int
|
||||
clearHealthOnRestart bool
|
||||
}
|
||||
|
||||
func (f *fakeRuntimeManager) CheckHealth(ctx context.Context) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return f.healthErr
|
||||
}
|
||||
|
||||
func (f *fakeRuntimeManager) Restart(ctx context.Context) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.restartCalls++
|
||||
if f.clearHealthOnRestart && f.restartErr == nil {
|
||||
f.healthErr = nil
|
||||
}
|
||||
return f.restartErr
|
||||
}
|
||||
|
||||
func (f *fakeSyncService) SyncOnStartup(ctx context.Context, target *protocol.ActiveConfigMeta) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.startupCalls++
|
||||
return f.startupErr
|
||||
}
|
||||
|
||||
func (f *fakeSyncService) SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error {
|
||||
f.mu.Lock()
|
||||
f.syncOnceCalls++
|
||||
callIndex := f.syncOnceCalls
|
||||
callback := f.onSyncOnceCall
|
||||
f.mu.Unlock()
|
||||
if callback != nil {
|
||||
callback(callIndex)
|
||||
}
|
||||
return f.syncOnceErr
|
||||
}
|
||||
|
||||
func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
heartbeatService := &fakeHeartbeatService{
|
||||
heartbeatResults: []*protocol.HeartbeatResult{{}},
|
||||
onHeartbeat: func(callCount int) {
|
||||
if callCount >= 2 {
|
||||
cancel()
|
||||
}
|
||||
},
|
||||
}
|
||||
syncService := &fakeSyncService{
|
||||
startupErr: errors.New("当前没有激活版本,保持当前 OpenResty 配置"),
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
AgentToken: "agent-token",
|
||||
NodeName: "edge-01",
|
||||
NodeIP: "10.0.0.8",
|
||||
AgentVersion: config.AgentVersion,
|
||||
NginxVersion: "1.27.1.2",
|
||||
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
|
||||
},
|
||||
StateStore: stateStore,
|
||||
HeartbeatService: heartbeatService,
|
||||
SyncService: syncService,
|
||||
}
|
||||
|
||||
err := runner.Run(ctx)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
if heartbeatService.registerCalls != 0 {
|
||||
t.Fatalf("expected no discovery register call, got %d", heartbeatService.registerCalls)
|
||||
}
|
||||
if heartbeatService.heartbeatCalls < 2 {
|
||||
t.Fatalf("expected heartbeat loop to continue, got %d heartbeat calls", heartbeatService.heartbeatCalls)
|
||||
}
|
||||
snapshot, loadErr := stateStore.Load()
|
||||
if loadErr != nil {
|
||||
t.Fatalf("failed to load state: %v", loadErr)
|
||||
}
|
||||
if snapshot.LastError != "当前没有激活版本,保持当前 OpenResty 配置" {
|
||||
t.Fatalf("expected startup sync error to be recorded, got %q", snapshot.LastError)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
heartbeatService := &fakeHeartbeatService{
|
||||
registerErr: errors.New("register timeout"),
|
||||
heartbeatErrs: []error{errors.New("heartbeat timeout")},
|
||||
heartbeatResults: []*protocol.HeartbeatResult{
|
||||
{},
|
||||
},
|
||||
}
|
||||
syncService := &fakeSyncService{
|
||||
syncOnceErr: errors.New("openresty reload failed"),
|
||||
onSyncOnceCall: func(callCount int) {
|
||||
if callCount >= 1 {
|
||||
cancel()
|
||||
}
|
||||
},
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
AgentToken: "agent-token",
|
||||
NodeName: "edge-01",
|
||||
NodeIP: "10.0.0.8",
|
||||
AgentVersion: config.AgentVersion,
|
||||
NginxVersion: "1.27.1.2",
|
||||
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
|
||||
},
|
||||
StateStore: stateStore,
|
||||
HeartbeatService: heartbeatService,
|
||||
SyncService: syncService,
|
||||
}
|
||||
|
||||
err := runner.Run(ctx)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
if heartbeatService.registerCalls != 0 {
|
||||
t.Fatalf("expected no register attempt, got %d", heartbeatService.registerCalls)
|
||||
}
|
||||
if syncService.syncOnceCalls == 0 {
|
||||
t.Fatal("expected sync loop to continue after heartbeat/register errors")
|
||||
}
|
||||
snapshot, loadErr := stateStore.Load()
|
||||
if loadErr != nil {
|
||||
t.Fatalf("failed to load state: %v", loadErr)
|
||||
}
|
||||
if snapshot.LastError != "openresty reload failed" {
|
||||
t.Fatalf("expected sync error to be recorded, got %q", snapshot.LastError)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerReportsOpenrestyHealthAndExecutesRestart(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
if err := stateStore.Save(&state.Snapshot{
|
||||
OpenrestyStatus: protocol.OpenrestyStatusUnhealthy,
|
||||
OpenrestyMessage: "docker run openresty failed: bind 80 already allocated",
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to seed state: %v", err)
|
||||
}
|
||||
heartbeatService := &fakeHeartbeatService{
|
||||
heartbeatResults: []*protocol.HeartbeatResult{{
|
||||
AgentSettings: &protocol.AgentSettings{RestartOpenrestyNow: true},
|
||||
}},
|
||||
onHeartbeat: func(callCount int) {
|
||||
if callCount >= 1 {
|
||||
cancel()
|
||||
}
|
||||
},
|
||||
}
|
||||
runtimeManager := &fakeRuntimeManager{
|
||||
healthErr: errors.New("docker openresty container is not running"),
|
||||
clearHealthOnRestart: true,
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
AgentToken: "agent-token",
|
||||
NodeName: "edge-01",
|
||||
NodeIP: "10.0.0.8",
|
||||
AgentVersion: config.AgentVersion,
|
||||
NginxVersion: "1.27.1.2",
|
||||
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
|
||||
},
|
||||
StateStore: stateStore,
|
||||
HeartbeatService: heartbeatService,
|
||||
SyncService: &fakeSyncService{},
|
||||
RuntimeManager: runtimeManager,
|
||||
}
|
||||
|
||||
err := runner.Run(ctx)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
if len(heartbeatService.heartbeatPayloads) == 0 {
|
||||
t.Fatal("expected at least one heartbeat payload")
|
||||
}
|
||||
payload := heartbeatService.heartbeatPayloads[0]
|
||||
if payload.OpenrestyStatus != protocol.OpenrestyStatusUnhealthy {
|
||||
t.Fatalf("expected unhealthy openresty status in heartbeat payload, got %q", payload.OpenrestyStatus)
|
||||
}
|
||||
if payload.OpenrestyMessage != "docker run openresty failed: bind 80 already allocated" {
|
||||
t.Fatalf("unexpected openresty message: %q", payload.OpenrestyMessage)
|
||||
}
|
||||
if runtimeManager.restartCalls != 1 {
|
||||
t.Fatalf("expected one openresty restart attempt, got %d", runtimeManager.restartCalls)
|
||||
}
|
||||
snapshot, loadErr := stateStore.Load()
|
||||
if loadErr != nil {
|
||||
t.Fatalf("failed to load state: %v", loadErr)
|
||||
}
|
||||
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusHealthy || snapshot.OpenrestyMessage != "" {
|
||||
t.Fatal("expected restart success to mark openresty healthy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
if err := stateStore.Save(&state.Snapshot{
|
||||
NodeID: "node-observe",
|
||||
CurrentVersion: "20260314-001",
|
||||
LastError: "sync failed",
|
||||
OpenrestyStatus: protocol.OpenrestyStatusUnhealthy,
|
||||
OpenrestyMessage: "reload failed",
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to seed state: %v", err)
|
||||
}
|
||||
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
NodeName: "edge-observe-1",
|
||||
NodeIP: "10.0.0.51",
|
||||
AgentVersion: config.AgentVersion,
|
||||
NginxVersion: "1.27.1.2",
|
||||
DataDir: tempDir,
|
||||
RouteConfigPath: filepath.Join(tempDir, "conf.d", "openflare_routes.conf"),
|
||||
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
|
||||
},
|
||||
StateStore: stateStore,
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(runner.Config.RouteConfigPath), 0o755); err != nil {
|
||||
t.Fatalf("failed to prepare route config dir: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(
|
||||
filepath.Join(filepath.Dir(runner.Config.RouteConfigPath), "openflare_access.log"),
|
||||
[]byte("{\"ts\":\""+time.Now().UTC().Format(time.RFC3339)+"\",\"host\":\"edge.example.com\",\"path\":\"/\",\"remote_addr\":\"10.0.0.8\",\"status\":200}\n"),
|
||||
0o644,
|
||||
); err != nil {
|
||||
t.Fatalf("failed to prepare access log: %v", err)
|
||||
}
|
||||
|
||||
firstPayload := runner.nodePayload("node-observe")
|
||||
if firstPayload.Profile == nil {
|
||||
t.Fatal("expected first heartbeat payload to include system profile")
|
||||
}
|
||||
if firstPayload.Snapshot == nil {
|
||||
t.Fatal("expected first heartbeat payload to include metric snapshot")
|
||||
}
|
||||
if firstPayload.TrafficReport == nil || firstPayload.TrafficReport.RequestCount != 1 {
|
||||
t.Fatalf("expected first heartbeat payload to include traffic report, got %+v", firstPayload.TrafficReport)
|
||||
}
|
||||
if len(firstPayload.AccessLogs) != 1 || firstPayload.AccessLogs[0].Path != "/" {
|
||||
t.Fatalf("expected first heartbeat payload to include access logs, got %+v", firstPayload.AccessLogs)
|
||||
}
|
||||
if len(firstPayload.HealthEvents) != 2 {
|
||||
t.Fatalf("expected health events for openresty and sync error, got %+v", firstPayload.HealthEvents)
|
||||
}
|
||||
|
||||
secondPayload := runner.nodePayload("node-observe")
|
||||
if secondPayload.Profile != nil {
|
||||
t.Fatal("expected unchanged profile to be omitted on subsequent heartbeat")
|
||||
}
|
||||
if secondPayload.Snapshot == nil {
|
||||
t.Fatal("expected metric snapshot to continue reporting on subsequent heartbeat")
|
||||
}
|
||||
if secondPayload.TrafficReport != nil {
|
||||
t.Fatalf("expected unchanged traffic window to be omitted on subsequent heartbeat, got %+v", secondPayload.TrafficReport)
|
||||
}
|
||||
if len(secondPayload.AccessLogs) != 0 {
|
||||
t.Fatalf("expected unchanged access log delta to be omitted on subsequent heartbeat, got %+v", secondPayload.AccessLogs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerReplaysBufferedObservabilityAfterHeartbeatRecovery(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
bufferStore := state.NewObservabilityBufferStore(filepath.Join(tempDir, "observability-buffer.json"))
|
||||
nowUnix := time.Now().UTC().Unix()
|
||||
bufferWindow := nowUnix - (nowUnix % 60) - 60
|
||||
if err := bufferStore.Upsert(state.ObservabilityBufferRecord{
|
||||
WindowStartedAtUnix: bufferWindow,
|
||||
Snapshot: &protocol.NodeMetricSnapshot{CapturedAtUnix: bufferWindow + 5, CPUUsagePercent: 30},
|
||||
TrafficReport: &protocol.NodeTrafficReport{WindowStartedAtUnix: bufferWindow, WindowEndedAtUnix: bufferWindow + 60, RequestCount: 8},
|
||||
QueuedAtUnix: bufferWindow + 60,
|
||||
}, 0); err != nil {
|
||||
t.Fatalf("failed to seed observability buffer: %v", err)
|
||||
}
|
||||
heartbeatService := &fakeHeartbeatService{
|
||||
heartbeatErrs: []error{errors.New("server offline"), nil},
|
||||
heartbeatResults: []*protocol.HeartbeatResult{{}, {}},
|
||||
onHeartbeat: func(callCount int) {
|
||||
if callCount >= 2 {
|
||||
cancel()
|
||||
}
|
||||
},
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
AgentToken: "agent-token",
|
||||
NodeName: "edge-buffer-01",
|
||||
NodeIP: "10.0.0.52",
|
||||
AgentVersion: config.AgentVersion,
|
||||
NginxVersion: "1.27.1.2",
|
||||
DataDir: tempDir,
|
||||
RouteConfigPath: filepath.Join(tempDir, "conf.d", "openflare_routes.conf"),
|
||||
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
|
||||
ObservabilityReplayMinutes: 15,
|
||||
},
|
||||
StateStore: stateStore,
|
||||
ObservabilityBuffer: bufferStore,
|
||||
HeartbeatService: heartbeatService,
|
||||
SyncService: &fakeSyncService{},
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(runner.Config.RouteConfigPath), 0o755); err != nil {
|
||||
t.Fatalf("failed to prepare route config dir: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(
|
||||
filepath.Join(filepath.Dir(runner.Config.RouteConfigPath), "openflare_access.log"),
|
||||
[]byte("{\"ts\":\""+time.Now().UTC().Format(time.RFC3339)+"\",\"host\":\"edge.example.com\",\"path\":\"/\",\"remote_addr\":\"10.0.0.8\",\"status\":200}\n"),
|
||||
0o644,
|
||||
); err != nil {
|
||||
t.Fatalf("failed to prepare access log: %v", err)
|
||||
}
|
||||
|
||||
runErr := runner.Run(ctx)
|
||||
if runErr != context.Canceled {
|
||||
t.Fatalf("expected run to stop by context cancellation, got %v", runErr)
|
||||
}
|
||||
if len(heartbeatService.heartbeatPayloads) != 2 {
|
||||
t.Fatalf("expected two heartbeat payloads, got %d", len(heartbeatService.heartbeatPayloads))
|
||||
}
|
||||
secondPayload := heartbeatService.heartbeatPayloads[1]
|
||||
if len(secondPayload.BufferedObservability) != 1 {
|
||||
t.Fatalf("expected second heartbeat to replay one buffered observation, got %+v", secondPayload.BufferedObservability)
|
||||
}
|
||||
if len(secondPayload.BufferedObservability[0].AccessLogs) != 0 {
|
||||
t.Fatalf("expected seeded buffered observation to keep empty access logs, got %+v", secondPayload.BufferedObservability[0].AccessLogs)
|
||||
}
|
||||
|
||||
replayable, err := bufferStore.Replayable(0, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Replayable after recovery failed: %v", err)
|
||||
}
|
||||
if len(replayable) != 0 {
|
||||
t.Fatalf("expected buffer to be acked after successful heartbeat, got %+v", replayable)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
heartbeatService := &fakeHeartbeatService{
|
||||
registerResp: &protocol.RegisterNodeResponse{
|
||||
NodeID: "node-server-assigned",
|
||||
AgentToken: "agent-token-issued",
|
||||
Name: "edge-01",
|
||||
},
|
||||
heartbeatResults: []*protocol.HeartbeatResult{{}},
|
||||
onHeartbeat: func(callCount int) {
|
||||
if callCount >= 1 {
|
||||
cancel()
|
||||
}
|
||||
},
|
||||
}
|
||||
syncService := &fakeSyncService{}
|
||||
configPath := filepath.Join(t.TempDir(), "agent.json")
|
||||
if err := os.WriteFile(configPath, []byte(`{"server_url":"http://127.0.0.1:3000","discovery_token":"discovery-token","node_name":"edge-01","node_ip":"10.0.0.8"}`), 0o644); err != nil {
|
||||
t.Fatalf("failed to seed config file: %v", err)
|
||||
}
|
||||
cfg, err := config.Load(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load config: %v", err)
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
ServerURL: cfg.ServerURL,
|
||||
DiscoveryToken: cfg.DiscoveryToken,
|
||||
NodeName: cfg.NodeName,
|
||||
NodeIP: cfg.NodeIP,
|
||||
AgentVersion: config.AgentVersion,
|
||||
NginxVersion: "1.27.1.2",
|
||||
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
|
||||
},
|
||||
StateStore: stateStore,
|
||||
HeartbeatService: heartbeatService,
|
||||
SyncService: syncService,
|
||||
}
|
||||
runner.Config = cfg
|
||||
runner.Config.AgentVersion = config.AgentVersion
|
||||
runner.Config.NginxVersion = "1.27.1.2"
|
||||
runner.Config.HeartbeatInterval = config.MillisecondDuration(10 * time.Millisecond)
|
||||
|
||||
err = runner.Run(ctx)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
if heartbeatService.registerCalls == 0 {
|
||||
t.Fatal("expected discovery register to be attempted")
|
||||
}
|
||||
if heartbeatService.lastToken != "agent-token-issued" {
|
||||
t.Fatalf("expected client token to be updated, got %q", heartbeatService.lastToken)
|
||||
}
|
||||
snapshot, loadErr := stateStore.Load()
|
||||
if loadErr != nil {
|
||||
t.Fatalf("failed to load state: %v", loadErr)
|
||||
}
|
||||
if snapshot.NodeID != "node-server-assigned" {
|
||||
t.Fatalf("expected node id to be replaced, got %q", snapshot.NodeID)
|
||||
}
|
||||
if runner.Config.AgentToken != "agent-token-issued" || runner.Config.DiscoveryToken != "" {
|
||||
t.Fatal("expected config token rotation to complete")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net"
|
||||
"os"
|
||||
pathpkg "path"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultDockerMainConfigRelativePath = "etc/nginx/nginx.conf"
|
||||
defaultDockerRouteConfigRelativePath = "etc/nginx/conf.d/openflare_routes.conf"
|
||||
defaultCertDirRelativePath = "etc/nginx/certs"
|
||||
defaultLuaDirRelativePath = "etc/nginx/lua"
|
||||
defaultDockerStateRelativePath = "var/lib/openflare/agent-state.json"
|
||||
defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json"
|
||||
defaultDockerOpenRestyCertDir = "/etc/nginx/openflare-certs"
|
||||
defaultDockerOpenRestyLuaDir = "/etc/nginx/openflare-lua"
|
||||
defaultOpenRestyObservabilityPort = 18081
|
||||
defaultObservabilityReplayMinutes = 15
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
ServerURL string `json:"server_url"`
|
||||
AgentToken string `json:"agent_token"`
|
||||
DiscoveryToken string `json:"discovery_token"`
|
||||
NodeName string `json:"node_name"`
|
||||
NodeIP string `json:"node_ip"`
|
||||
AgentVersion string `json:"-"`
|
||||
NginxVersion string `json:"-"`
|
||||
OpenrestyPath string `json:"openresty_path"`
|
||||
OpenrestyContainerName string `json:"openresty_container_name"`
|
||||
OpenrestyDockerImage string `json:"openresty_docker_image"`
|
||||
DockerBinary string `json:"docker_binary"`
|
||||
DataDir string `json:"data_dir"`
|
||||
MainConfigPath string `json:"main_config_path"`
|
||||
RouteConfigPath string `json:"route_config_path"`
|
||||
CertDir string `json:"cert_dir"`
|
||||
OpenrestyCertDir string `json:"openresty_cert_dir"`
|
||||
LuaDir string `json:"lua_dir"`
|
||||
OpenrestyLuaDir string `json:"openresty_lua_dir"`
|
||||
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
|
||||
ObservabilityBufferPath string `json:"observability_buffer_path"`
|
||||
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
|
||||
StatePath string `json:"state_path"`
|
||||
HeartbeatInterval MillisecondDuration `json:"heartbeat_interval"`
|
||||
RequestTimeout MillisecondDuration `json:"request_timeout"`
|
||||
configPath string
|
||||
}
|
||||
|
||||
type configFile struct {
|
||||
ServerURL string `json:"server_url"`
|
||||
AgentToken string `json:"agent_token"`
|
||||
DiscoveryToken string `json:"discovery_token"`
|
||||
NodeName string `json:"node_name"`
|
||||
NodeIP string `json:"node_ip"`
|
||||
OpenrestyPath string `json:"openresty_path"`
|
||||
OpenrestyContainerName string `json:"openresty_container_name"`
|
||||
OpenrestyDockerImage string `json:"openresty_docker_image"`
|
||||
DockerBinary string `json:"docker_binary"`
|
||||
DataDir string `json:"data_dir"`
|
||||
MainConfigPath string `json:"main_config_path"`
|
||||
RouteConfigPath string `json:"route_config_path"`
|
||||
CertDir string `json:"cert_dir"`
|
||||
OpenrestyCertDir string `json:"openresty_cert_dir"`
|
||||
LuaDir string `json:"lua_dir"`
|
||||
OpenrestyLuaDir string `json:"openresty_lua_dir"`
|
||||
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
|
||||
ObservabilityBufferPath string `json:"observability_buffer_path"`
|
||||
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
|
||||
StatePath string `json:"state_path"`
|
||||
HeartbeatInterval MillisecondDuration `json:"heartbeat_interval"`
|
||||
RequestTimeout MillisecondDuration `json:"request_timeout"`
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
file := &configFile{}
|
||||
if err = json.Unmarshal(data, file); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cfg := &Config{
|
||||
ServerURL: file.ServerURL,
|
||||
AgentToken: file.AgentToken,
|
||||
DiscoveryToken: file.DiscoveryToken,
|
||||
NodeName: file.NodeName,
|
||||
NodeIP: file.NodeIP,
|
||||
OpenrestyPath: file.OpenrestyPath,
|
||||
OpenrestyContainerName: file.OpenrestyContainerName,
|
||||
OpenrestyDockerImage: file.OpenrestyDockerImage,
|
||||
DockerBinary: file.DockerBinary,
|
||||
DataDir: file.DataDir,
|
||||
MainConfigPath: file.MainConfigPath,
|
||||
RouteConfigPath: file.RouteConfigPath,
|
||||
CertDir: file.CertDir,
|
||||
OpenrestyCertDir: file.OpenrestyCertDir,
|
||||
LuaDir: file.LuaDir,
|
||||
OpenrestyLuaDir: file.OpenrestyLuaDir,
|
||||
OpenrestyObservabilityPort: file.OpenrestyObservabilityPort,
|
||||
ObservabilityBufferPath: file.ObservabilityBufferPath,
|
||||
ObservabilityReplayMinutes: file.ObservabilityReplayMinutes,
|
||||
StatePath: file.StatePath,
|
||||
HeartbeatInterval: file.HeartbeatInterval,
|
||||
RequestTimeout: file.RequestTimeout,
|
||||
}
|
||||
cfg.configPath = path
|
||||
applyDefaults(cfg, filepath.Dir(path))
|
||||
if err = validate(cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func applyDefaults(cfg *Config, baseDir string) {
|
||||
baseDir = filepath.Clean(baseDir)
|
||||
cfg.AgentVersion = AgentVersion
|
||||
if cfg.OpenrestyContainerName == "" {
|
||||
cfg.OpenrestyContainerName = "openflare-openresty"
|
||||
}
|
||||
if cfg.OpenrestyDockerImage == "" {
|
||||
cfg.OpenrestyDockerImage = "openresty/openresty:alpine"
|
||||
}
|
||||
if cfg.DockerBinary == "" {
|
||||
cfg.DockerBinary = "docker"
|
||||
}
|
||||
if cfg.DataDir == "" {
|
||||
cfg.DataDir = filepath.Join(baseDir, "data")
|
||||
}
|
||||
if cfg.NodeName == "" {
|
||||
cfg.NodeName = detectHostname()
|
||||
}
|
||||
if cfg.NodeIP == "" {
|
||||
cfg.NodeIP = detectNodeIP()
|
||||
}
|
||||
if cfg.OpenrestyPath == "" {
|
||||
cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultDockerMainConfigRelativePath)
|
||||
cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultDockerRouteConfigRelativePath)
|
||||
cfg.StatePath = joinManagedPath(cfg.DataDir, defaultDockerStateRelativePath)
|
||||
} else {
|
||||
if cfg.MainConfigPath == "" {
|
||||
cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultDockerMainConfigRelativePath)
|
||||
}
|
||||
if cfg.RouteConfigPath == "" {
|
||||
cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultDockerRouteConfigRelativePath)
|
||||
}
|
||||
if cfg.StatePath == "" {
|
||||
cfg.StatePath = joinManagedPath(cfg.DataDir, defaultDockerStateRelativePath)
|
||||
}
|
||||
}
|
||||
if cfg.CertDir == "" {
|
||||
cfg.CertDir = joinManagedPath(cfg.DataDir, defaultCertDirRelativePath)
|
||||
}
|
||||
if cfg.OpenrestyCertDir == "" {
|
||||
if cfg.OpenrestyPath != "" {
|
||||
cfg.OpenrestyCertDir = cfg.CertDir
|
||||
} else {
|
||||
cfg.OpenrestyCertDir = defaultDockerOpenRestyCertDir
|
||||
}
|
||||
}
|
||||
if cfg.LuaDir == "" {
|
||||
cfg.LuaDir = joinManagedPath(cfg.DataDir, defaultLuaDirRelativePath)
|
||||
}
|
||||
if cfg.OpenrestyLuaDir == "" {
|
||||
if cfg.OpenrestyPath != "" {
|
||||
cfg.OpenrestyLuaDir = cfg.LuaDir
|
||||
} else {
|
||||
cfg.OpenrestyLuaDir = defaultDockerOpenRestyLuaDir
|
||||
}
|
||||
}
|
||||
if cfg.OpenrestyObservabilityPort <= 0 {
|
||||
cfg.OpenrestyObservabilityPort = defaultOpenRestyObservabilityPort
|
||||
}
|
||||
if cfg.ObservabilityBufferPath == "" {
|
||||
cfg.ObservabilityBufferPath = joinManagedPath(cfg.DataDir, defaultObservabilityBufferRelativePath)
|
||||
}
|
||||
if cfg.ObservabilityReplayMinutes <= 0 {
|
||||
cfg.ObservabilityReplayMinutes = defaultObservabilityReplayMinutes
|
||||
}
|
||||
if cfg.HeartbeatInterval <= 0 {
|
||||
cfg.HeartbeatInterval = MillisecondDuration(10 * time.Second)
|
||||
}
|
||||
if cfg.RequestTimeout <= 0 {
|
||||
cfg.RequestTimeout = MillisecondDuration(10 * time.Second)
|
||||
}
|
||||
normalizeManagedPaths(cfg)
|
||||
}
|
||||
|
||||
func normalizeManagedPaths(cfg *Config) {
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
if usesSlashPath(cfg.DataDir) {
|
||||
cfg.DataDir = filepath.ToSlash(cfg.DataDir)
|
||||
}
|
||||
if usesSlashPath(cfg.MainConfigPath) {
|
||||
cfg.MainConfigPath = filepath.ToSlash(cfg.MainConfigPath)
|
||||
}
|
||||
if usesSlashPath(cfg.RouteConfigPath) {
|
||||
cfg.RouteConfigPath = filepath.ToSlash(cfg.RouteConfigPath)
|
||||
}
|
||||
if usesSlashPath(cfg.CertDir) {
|
||||
cfg.CertDir = filepath.ToSlash(cfg.CertDir)
|
||||
}
|
||||
if usesSlashPath(cfg.OpenrestyCertDir) {
|
||||
cfg.OpenrestyCertDir = filepath.ToSlash(cfg.OpenrestyCertDir)
|
||||
}
|
||||
if usesSlashPath(cfg.LuaDir) {
|
||||
cfg.LuaDir = filepath.ToSlash(cfg.LuaDir)
|
||||
}
|
||||
if usesSlashPath(cfg.OpenrestyLuaDir) {
|
||||
cfg.OpenrestyLuaDir = filepath.ToSlash(cfg.OpenrestyLuaDir)
|
||||
}
|
||||
if usesSlashPath(cfg.StatePath) {
|
||||
cfg.StatePath = filepath.ToSlash(cfg.StatePath)
|
||||
}
|
||||
if usesSlashPath(cfg.ObservabilityBufferPath) {
|
||||
cfg.ObservabilityBufferPath = filepath.ToSlash(cfg.ObservabilityBufferPath)
|
||||
}
|
||||
}
|
||||
|
||||
func usesSlashPath(path string) bool {
|
||||
return strings.HasPrefix(path, "/")
|
||||
}
|
||||
|
||||
func joinManagedPath(base string, relative string) string {
|
||||
if usesSlashPath(base) {
|
||||
return pathpkg.Join(filepath.ToSlash(base), relative)
|
||||
}
|
||||
return filepath.Join(base, relative)
|
||||
}
|
||||
|
||||
func validate(cfg *Config) error {
|
||||
if cfg.ServerURL == "" {
|
||||
return errors.New("server_url 不能为空")
|
||||
}
|
||||
if strings.TrimSpace(cfg.AgentToken) == "" && strings.TrimSpace(cfg.DiscoveryToken) == "" {
|
||||
return errors.New("agent_token 和 discovery_token 不能同时为空")
|
||||
}
|
||||
if cfg.NodeName == "" {
|
||||
return errors.New("node_name 不能为空")
|
||||
}
|
||||
if cfg.NodeIP == "" {
|
||||
return errors.New("node_ip 不能为空")
|
||||
}
|
||||
if cfg.OpenrestyObservabilityPort <= 0 || cfg.OpenrestyObservabilityPort > 65535 {
|
||||
return errors.New("openresty_observability_port 必须在 1-65535 之间")
|
||||
}
|
||||
if cfg.ObservabilityReplayMinutes <= 0 {
|
||||
return errors.New("observability_replay_minutes 必须大于 0")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cfg *Config) InitialAuthToken() string {
|
||||
if cfg == nil {
|
||||
return ""
|
||||
}
|
||||
if token := strings.TrimSpace(cfg.AgentToken); token != "" {
|
||||
return token
|
||||
}
|
||||
return strings.TrimSpace(cfg.DiscoveryToken)
|
||||
}
|
||||
|
||||
func (cfg *Config) Save() error {
|
||||
if cfg == nil {
|
||||
return errors.New("config 不能为空")
|
||||
}
|
||||
if cfg.configPath == "" {
|
||||
return errors.New("config path 未初始化")
|
||||
}
|
||||
data, err := json.MarshalIndent(cfg, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(cfg.configPath, data, 0o644)
|
||||
}
|
||||
|
||||
func detectHostname() string {
|
||||
host, err := os.Hostname()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(host)
|
||||
}
|
||||
|
||||
func firstNonEmpty(values ...string) string {
|
||||
for _, value := range values {
|
||||
if strings.TrimSpace(value) != "" {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func detectNodeIP() string {
|
||||
interfaces, err := net.Interfaces()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
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 {
|
||||
return ipv4.String()
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,289 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestLoadDockerModeUsesManagedPaths(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, "agent.json")
|
||||
payload := map[string]any{
|
||||
"server_url": "http://127.0.0.1:3000",
|
||||
"agent_token": "token",
|
||||
"node_name": "edge-01",
|
||||
"node_ip": "10.0.0.8",
|
||||
}
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal config: %v", err)
|
||||
}
|
||||
if err = os.WriteFile(configPath, data, 0o644); err != nil {
|
||||
t.Fatalf("failed to write config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := Load(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Load failed: %v", err)
|
||||
}
|
||||
|
||||
if cfg.DataDir != filepath.Join(dir, "data") {
|
||||
t.Fatalf("unexpected data dir: %s", cfg.DataDir)
|
||||
}
|
||||
if cfg.MainConfigPath != filepath.Join(dir, "data", defaultDockerMainConfigRelativePath) {
|
||||
t.Fatalf("unexpected main config path: %s", cfg.MainConfigPath)
|
||||
}
|
||||
if cfg.RouteConfigPath != filepath.Join(dir, "data", defaultDockerRouteConfigRelativePath) {
|
||||
t.Fatalf("unexpected route config path: %s", cfg.RouteConfigPath)
|
||||
}
|
||||
if cfg.CertDir != filepath.Join(dir, "data", defaultCertDirRelativePath) {
|
||||
t.Fatalf("unexpected cert dir: %s", cfg.CertDir)
|
||||
}
|
||||
if cfg.LuaDir != filepath.Join(dir, "data", defaultLuaDirRelativePath) {
|
||||
t.Fatalf("unexpected lua dir: %s", cfg.LuaDir)
|
||||
}
|
||||
if cfg.OpenrestyContainerName != "openflare-openresty" {
|
||||
t.Fatalf("unexpected openresty container name: %s", cfg.OpenrestyContainerName)
|
||||
}
|
||||
if cfg.OpenrestyDockerImage != "openresty/openresty:alpine" {
|
||||
t.Fatalf("unexpected openresty image: %s", cfg.OpenrestyDockerImage)
|
||||
}
|
||||
if cfg.OpenrestyCertDir != defaultDockerOpenRestyCertDir {
|
||||
t.Fatalf("unexpected openresty cert dir: %s", cfg.OpenrestyCertDir)
|
||||
}
|
||||
if cfg.OpenrestyLuaDir != defaultDockerOpenRestyLuaDir {
|
||||
t.Fatalf("unexpected openresty lua dir: %s", cfg.OpenrestyLuaDir)
|
||||
}
|
||||
if cfg.StatePath != filepath.Join(dir, "data", defaultDockerStateRelativePath) {
|
||||
t.Fatalf("unexpected state path: %s", cfg.StatePath)
|
||||
}
|
||||
if cfg.ObservabilityBufferPath != filepath.Join(dir, "data", defaultObservabilityBufferRelativePath) {
|
||||
t.Fatalf("unexpected observability buffer path: %s", cfg.ObservabilityBufferPath)
|
||||
}
|
||||
if cfg.OpenrestyObservabilityPort != defaultOpenRestyObservabilityPort {
|
||||
t.Fatalf("unexpected openresty observability port: %d", cfg.OpenrestyObservabilityPort)
|
||||
}
|
||||
if cfg.ObservabilityReplayMinutes != defaultObservabilityReplayMinutes {
|
||||
t.Fatalf("unexpected observability replay minutes: %d", cfg.ObservabilityReplayMinutes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadPathModeKeepsExplicitPaths(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, "agent.json")
|
||||
payload := map[string]any{
|
||||
"server_url": "http://127.0.0.1:3000",
|
||||
"agent_token": "token",
|
||||
"node_name": "edge-01",
|
||||
"node_ip": "10.0.0.8",
|
||||
"openresty_path": "/usr/local/openresty/nginx/sbin/openresty",
|
||||
"main_config_path": "/tmp/nginx.conf",
|
||||
"route_config_path": "/tmp/routes.conf",
|
||||
"state_path": "/tmp/agent-state.json",
|
||||
}
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal config: %v", err)
|
||||
}
|
||||
if err = os.WriteFile(configPath, data, 0o644); err != nil {
|
||||
t.Fatalf("failed to write config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := Load(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Load failed: %v", err)
|
||||
}
|
||||
|
||||
if cfg.MainConfigPath != "/tmp/nginx.conf" {
|
||||
t.Fatalf("unexpected main config path: %s", cfg.MainConfigPath)
|
||||
}
|
||||
if cfg.RouteConfigPath != "/tmp/routes.conf" {
|
||||
t.Fatalf("unexpected route config path: %s", cfg.RouteConfigPath)
|
||||
}
|
||||
if cfg.StatePath != "/tmp/agent-state.json" {
|
||||
t.Fatalf("unexpected state path: %s", cfg.StatePath)
|
||||
}
|
||||
if cfg.ObservabilityBufferPath != filepath.Join(dir, "data", defaultObservabilityBufferRelativePath) {
|
||||
t.Fatalf("unexpected observability buffer path: %s", cfg.ObservabilityBufferPath)
|
||||
}
|
||||
if cfg.OpenrestyCertDir != cfg.CertDir {
|
||||
t.Fatalf("expected path mode openresty cert dir to equal cert dir, got %s / %s", cfg.OpenrestyCertDir, cfg.CertDir)
|
||||
}
|
||||
if cfg.OpenrestyLuaDir != cfg.LuaDir {
|
||||
t.Fatalf("expected path mode openresty lua dir to equal lua dir, got %s / %s", cfg.OpenrestyLuaDir, cfg.LuaDir)
|
||||
}
|
||||
if cfg.OpenrestyObservabilityPort != defaultOpenRestyObservabilityPort {
|
||||
t.Fatalf("unexpected path mode openresty observability port: %d", cfg.OpenrestyObservabilityPort)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadUsesCustomDataDirForGeneratedFiles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, "agent.json")
|
||||
payload := map[string]any{
|
||||
"server_url": "http://127.0.0.1:3000",
|
||||
"agent_token": "token",
|
||||
"node_name": "edge-01",
|
||||
"node_ip": "10.0.0.8",
|
||||
"data_dir": "/srv/openflare",
|
||||
}
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal config: %v", err)
|
||||
}
|
||||
if err = os.WriteFile(configPath, data, 0o644); err != nil {
|
||||
t.Fatalf("failed to write config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := Load(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Load failed: %v", err)
|
||||
}
|
||||
|
||||
if cfg.RouteConfigPath != "/srv/openflare/"+defaultDockerRouteConfigRelativePath {
|
||||
t.Fatalf("unexpected route config path: %s", cfg.RouteConfigPath)
|
||||
}
|
||||
if cfg.MainConfigPath != "/srv/openflare/"+defaultDockerMainConfigRelativePath {
|
||||
t.Fatalf("unexpected main config path: %s", cfg.MainConfigPath)
|
||||
}
|
||||
if cfg.StatePath != "/srv/openflare/"+defaultDockerStateRelativePath {
|
||||
t.Fatalf("unexpected state path: %s", cfg.StatePath)
|
||||
}
|
||||
if cfg.ObservabilityBufferPath != "/srv/openflare/"+defaultObservabilityBufferRelativePath {
|
||||
t.Fatalf("unexpected observability buffer path: %s", cfg.ObservabilityBufferPath)
|
||||
}
|
||||
if cfg.CertDir != "/srv/openflare/"+defaultCertDirRelativePath {
|
||||
t.Fatalf("unexpected cert dir: %s", cfg.CertDir)
|
||||
}
|
||||
if cfg.LuaDir != "/srv/openflare/"+defaultLuaDirRelativePath {
|
||||
t.Fatalf("unexpected lua dir: %s", cfg.LuaDir)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadUsesMillisecondsForIntervals(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, "agent.json")
|
||||
payload := map[string]any{
|
||||
"server_url": "http://127.0.0.1:3000",
|
||||
"agent_token": "token",
|
||||
"node_name": "edge-01",
|
||||
"node_ip": "10.0.0.8",
|
||||
"heartbeat_interval": 30000,
|
||||
"request_timeout": 1500,
|
||||
}
|
||||
data, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal config: %v", err)
|
||||
}
|
||||
if err = os.WriteFile(configPath, data, 0o644); err != nil {
|
||||
t.Fatalf("failed to write config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := Load(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Load failed: %v", err)
|
||||
}
|
||||
|
||||
if cfg.HeartbeatInterval.Duration() != 30*time.Second {
|
||||
t.Fatalf("unexpected heartbeat interval: %s", cfg.HeartbeatInterval)
|
||||
}
|
||||
if cfg.RequestTimeout.Duration() != 1500*time.Millisecond {
|
||||
t.Fatalf("unexpected request timeout: %s", cfg.RequestTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSavePersistsMillisecondsAndOmitsRuntimeVersions(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
configPath := filepath.Join(dir, "agent.json")
|
||||
if err := os.WriteFile(configPath, []byte(`{"server_url":"http://127.0.0.1:3000","agent_token":"token","node_name":"edge-01","node_ip":"10.0.0.8"}`), 0o644); err != nil {
|
||||
t.Fatalf("failed to write config: %v", err)
|
||||
}
|
||||
|
||||
cfg, err := Load(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Load failed: %v", err)
|
||||
}
|
||||
cfg.NginxVersion = "1.27.1.2"
|
||||
cfg.HeartbeatInterval = MillisecondDuration(5 * time.Second)
|
||||
cfg.RequestTimeout = MillisecondDuration(7 * time.Second)
|
||||
|
||||
if err = cfg.Save(); err != nil {
|
||||
t.Fatalf("Save failed: %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read saved config: %v", err)
|
||||
}
|
||||
var decoded map[string]any
|
||||
if err = json.Unmarshal(data, &decoded); err != nil {
|
||||
t.Fatalf("failed to decode saved config: %v", err)
|
||||
}
|
||||
if _, ok := decoded["agent_version"]; ok {
|
||||
t.Fatal("agent_version should not be persisted")
|
||||
}
|
||||
if _, ok := decoded["nginx_version"]; ok {
|
||||
t.Fatal("nginx_version should not be persisted")
|
||||
}
|
||||
if decoded["heartbeat_interval"] != float64(5000) {
|
||||
t.Fatalf("unexpected heartbeat interval: %#v", decoded["heartbeat_interval"])
|
||||
}
|
||||
if decoded["request_timeout"] != float64(7000) {
|
||||
t.Fatalf("unexpected request timeout: %#v", decoded["request_timeout"])
|
||||
}
|
||||
if decoded["openresty_observability_port"] != float64(defaultOpenRestyObservabilityPort) {
|
||||
t.Fatalf("unexpected observability port: %#v", decoded["openresty_observability_port"])
|
||||
}
|
||||
if decoded["observability_replay_minutes"] != float64(defaultObservabilityReplayMinutes) {
|
||||
t.Fatalf("unexpected observability replay minutes: %#v", decoded["observability_replay_minutes"])
|
||||
}
|
||||
if _, ok := decoded["nginx_path"]; ok {
|
||||
t.Fatal("legacy nginx_path should not be persisted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitialAuthToken(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
agentToken string
|
||||
discoveryToken string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "prefer agent token",
|
||||
agentToken: "agent-token",
|
||||
discoveryToken: "discovery-token",
|
||||
expected: "agent-token",
|
||||
},
|
||||
{
|
||||
name: "fallback to discovery token",
|
||||
agentToken: " ",
|
||||
discoveryToken: "discovery-token",
|
||||
expected: "discovery-token",
|
||||
},
|
||||
{
|
||||
name: "nil config returns empty string",
|
||||
agentToken: "",
|
||||
discoveryToken: "",
|
||||
expected: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var cfg *Config
|
||||
if tt.name != "nil config returns empty string" {
|
||||
cfg = &Config{
|
||||
AgentToken: tt.agentToken,
|
||||
DiscoveryToken: tt.discoveryToken,
|
||||
}
|
||||
}
|
||||
if token := cfg.InitialAuthToken(); token != tt.expected {
|
||||
t.Fatalf("unexpected initial auth token: %q", token)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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,3 @@
|
||||
package config
|
||||
|
||||
var AgentVersion = "dev"
|
||||
@@ -0,0 +1,33 @@
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"openflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
type Client interface {
|
||||
RegisterNode(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error)
|
||||
Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error)
|
||||
SetToken(token string)
|
||||
}
|
||||
|
||||
type Service struct {
|
||||
client Client
|
||||
}
|
||||
|
||||
func New(client Client) *Service {
|
||||
return &Service{client: client}
|
||||
}
|
||||
|
||||
func (s *Service) Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
|
||||
return s.client.RegisterNode(ctx, payload)
|
||||
}
|
||||
|
||||
func (s *Service) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) {
|
||||
return s.client.Heartbeat(ctx, payload)
|
||||
}
|
||||
|
||||
func (s *Service) SetToken(token string) {
|
||||
s.client.SetToken(token)
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"openflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
baseURL string
|
||||
token string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
token: token,
|
||||
httpClient: &http.Client{
|
||||
Timeout: timeout,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
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/agent/nodes/register", payload, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.Success {
|
||||
return nil, errors.New(resp.Message)
|
||||
}
|
||||
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.HeartbeatAPIResponse{}
|
||||
if err := c.postJSON(ctx, "/api/agent/nodes/heartbeat", payload, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.Success {
|
||||
return nil, errors.New(resp.Message)
|
||||
}
|
||||
return &protocol.HeartbeatResult{
|
||||
AgentSettings: resp.AgentSettings,
|
||||
ActiveConfig: resp.ActiveConfig,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *Client) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) {
|
||||
resp := protocol.APIResponse[protocol.ActiveConfigResponse]{}
|
||||
if err := c.getJSON(ctx, "/api/agent/config-versions/active", &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.Success {
|
||||
return nil, errors.New(resp.Message)
|
||||
}
|
||||
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)
|
||||
return c.postJSON(ctx, "/api/agent/apply-logs", payload, nil)
|
||||
}
|
||||
|
||||
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 res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
slog.Warn("http request returned non-200", "method", req.Method, "path", req.URL.Path, "status", res.Status)
|
||||
return errors.New(res.Status)
|
||||
}
|
||||
if target == nil {
|
||||
var wrapper protocol.APIResponse[json.RawMessage]
|
||||
if err = json.NewDecoder(res.Body).Decode(&wrapper); err != nil {
|
||||
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
|
||||
return err
|
||||
}
|
||||
if !wrapper.Success {
|
||||
slog.Warn("http api response failed", "method", req.Method, "path", req.URL.Path, "message", wrapper.Message)
|
||||
return errors.New(wrapper.Message)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err = json.NewDecoder(res.Body).Decode(target); err != nil {
|
||||
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package logging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type customTextHandler struct {
|
||||
writer io.Writer
|
||||
level slog.Level
|
||||
attrs []slog.Attr
|
||||
groups []string
|
||||
}
|
||||
|
||||
func Setup() {
|
||||
handler := &customTextHandler{
|
||||
writer: os.Stdout,
|
||||
level: parseLevel(os.Getenv("LOG_LEVEL")),
|
||||
}
|
||||
slog.SetDefault(slog.New(handler))
|
||||
}
|
||||
|
||||
func (h *customTextHandler) Enabled(_ context.Context, level slog.Level) bool {
|
||||
return level >= h.level
|
||||
}
|
||||
|
||||
func (h *customTextHandler) Handle(_ context.Context, record slog.Record) error {
|
||||
var builder strings.Builder
|
||||
builder.WriteString(record.Time.Format("2006-01-02 15:04:05.000"))
|
||||
builder.WriteString(" | ")
|
||||
builder.WriteString(fmt.Sprintf("%-8s", levelLabel(record.Level)))
|
||||
builder.WriteString(" | ")
|
||||
builder.WriteString(sourceLocation(record.PC))
|
||||
builder.WriteString(" - ")
|
||||
builder.WriteString(record.Message)
|
||||
|
||||
attrs := make([]slog.Attr, 0, len(h.attrs)+record.NumAttrs())
|
||||
attrs = append(attrs, h.attrs...)
|
||||
record.Attrs(func(attr slog.Attr) bool {
|
||||
attrs = append(attrs, attr)
|
||||
return true
|
||||
})
|
||||
if len(attrs) > 0 {
|
||||
builder.WriteString(" | ")
|
||||
builder.WriteString(formatAttrs(h.groups, attrs))
|
||||
}
|
||||
builder.WriteByte('\n')
|
||||
_, err := io.WriteString(h.writer, builder.String())
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *customTextHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
|
||||
cloned := *h
|
||||
cloned.attrs = append(slices.Clone(h.attrs), attrs...)
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func (h *customTextHandler) WithGroup(name string) slog.Handler {
|
||||
if strings.TrimSpace(name) == "" {
|
||||
return h
|
||||
}
|
||||
cloned := *h
|
||||
cloned.groups = append(slices.Clone(h.groups), name)
|
||||
return &cloned
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
func levelLabel(level slog.Level) string {
|
||||
switch {
|
||||
case level <= slog.LevelDebug:
|
||||
return "DEBUG"
|
||||
case level < slog.LevelWarn:
|
||||
return "INFO"
|
||||
case level < slog.LevelError:
|
||||
return "WARNING"
|
||||
default:
|
||||
return "ERROR"
|
||||
}
|
||||
}
|
||||
|
||||
func sourceLocation(pc uintptr) string {
|
||||
if pc == 0 {
|
||||
return "unknown:unknown:0"
|
||||
}
|
||||
frame, _ := runtime.CallersFrames([]uintptr{pc}).Next()
|
||||
fileName := strings.TrimSuffix(filepath.Base(frame.File), filepath.Ext(frame.File))
|
||||
if fileName == "" {
|
||||
fileName = "unknown"
|
||||
}
|
||||
functionName := "unknown"
|
||||
if frame.Function != "" {
|
||||
parts := strings.Split(frame.Function, "/")
|
||||
functionName = parts[len(parts)-1]
|
||||
if dot := strings.LastIndex(functionName, "."); dot >= 0 && dot < len(functionName)-1 {
|
||||
functionName = functionName[dot+1:]
|
||||
}
|
||||
}
|
||||
return fmt.Sprintf("%s:%s:%d", fileName, functionName, frame.Line)
|
||||
}
|
||||
|
||||
func formatAttrs(groups []string, attrs []slog.Attr) string {
|
||||
parts := make([]string, 0, len(attrs))
|
||||
for _, attr := range attrs {
|
||||
key := attr.Key
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
if len(groups) > 0 {
|
||||
key = strings.Join(append(slices.Clone(groups), key), ".")
|
||||
}
|
||||
parts = append(parts, fmt.Sprintf("%s=%v", key, attr.Value.Any()))
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
@@ -0,0 +1,803 @@
|
||||
package nginx
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"openflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
const CertDirPlaceholder = "__OPENFLARE_CERT_DIR__"
|
||||
const RouteConfigPlaceholder = "__OPENFLARE_ROUTE_CONFIG__"
|
||||
const AccessLogPlaceholder = "__OPENFLARE_ACCESS_LOG__"
|
||||
const LuaDirPlaceholder = "__OPENFLARE_LUA_DIR__"
|
||||
const ObservabilityListenPlaceholder = "__OPENFLARE_OBSERVABILITY_LISTEN__"
|
||||
const ObservabilityPortPlaceholder = "__OPENFLARE_OBSERVABILITY_PORT__"
|
||||
const DockerMainConfigPath = "/usr/local/openresty/nginx/conf/nginx.conf"
|
||||
const DockerRouteConfigPath = "/etc/nginx/conf.d/openflare_routes.conf"
|
||||
const DockerAccessLogPath = "/etc/nginx/conf.d/openflare_access.log"
|
||||
|
||||
const dockerRuntimeCommand = "openresty"
|
||||
|
||||
type Executor interface {
|
||||
Test(ctx context.Context) error
|
||||
Reload(ctx context.Context) error
|
||||
EnsureRuntime(ctx context.Context, recreate bool) error
|
||||
CheckHealth(ctx context.Context) error
|
||||
Restart(ctx context.Context) error
|
||||
}
|
||||
|
||||
type CommandRunner interface {
|
||||
Run(ctx context.Context, name string, args ...string) ([]byte, error)
|
||||
}
|
||||
|
||||
type OSCommandRunner struct{}
|
||||
|
||||
func (r *OSCommandRunner) Run(ctx context.Context, name string, args ...string) ([]byte, error) {
|
||||
cmd := exec.CommandContext(ctx, name, args...)
|
||||
output, err := cmd.CombinedOutput()
|
||||
return output, err
|
||||
}
|
||||
|
||||
type PathExecutor struct {
|
||||
Path string
|
||||
Runner CommandRunner
|
||||
}
|
||||
|
||||
func (e *PathExecutor) Test(ctx context.Context) error {
|
||||
slog.Debug("running openresty test with binary", "path", e.Path)
|
||||
output, err := e.Runner.Run(ctx, e.Path, "-t")
|
||||
if err != nil {
|
||||
return fmt.Errorf("openresty -t failed: %w: %s", err, string(output))
|
||||
}
|
||||
slog.Debug("openresty test succeeded with binary", "path", e.Path)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *PathExecutor) Reload(ctx context.Context) error {
|
||||
slog.Debug("running openresty reload with binary", "path", e.Path)
|
||||
output, err := e.Runner.Run(ctx, e.Path, "-s", "reload")
|
||||
if err != nil {
|
||||
return fmt.Errorf("openresty reload failed: %w: %s", err, string(output))
|
||||
}
|
||||
slog.Debug("openresty reload succeeded with binary", "path", e.Path)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *PathExecutor) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *PathExecutor) CheckHealth(ctx context.Context) error {
|
||||
return e.Test(ctx)
|
||||
}
|
||||
|
||||
func (e *PathExecutor) Restart(ctx context.Context) error {
|
||||
slog.Info("restarting openresty with binary", "path", e.Path)
|
||||
output, err := e.Runner.Run(ctx, e.Path, "-s", "quit")
|
||||
if err != nil {
|
||||
text := string(output)
|
||||
if !isIgnorableOpenrestyStopError(text) {
|
||||
return fmt.Errorf("openresty stop failed: %w: %s", err, text)
|
||||
}
|
||||
}
|
||||
output, err = e.Runner.Run(ctx, e.Path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("openresty start failed: %w: %s", err, string(output))
|
||||
}
|
||||
slog.Info("openresty restart succeeded with binary", "path", e.Path)
|
||||
return nil
|
||||
}
|
||||
|
||||
type DockerExecutor struct {
|
||||
DockerBinary string
|
||||
ContainerName string
|
||||
Image string
|
||||
MainConfigPath string
|
||||
RouteConfigDir string
|
||||
CertDir string
|
||||
NginxCertDir string
|
||||
LuaDir string
|
||||
NginxLuaDir string
|
||||
OpenrestyObservabilityPort int
|
||||
Runner CommandRunner
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) Test(ctx context.Context) error {
|
||||
slog.Debug("running docker openresty test", "container", e.ContainerName, "image", e.Image)
|
||||
output, err := e.runEphemeralRuntimeCommand(ctx, "-t")
|
||||
if err != nil {
|
||||
return fmt.Errorf("docker %s -t failed: %w: %s", dockerRuntimeCommand, err, string(output))
|
||||
}
|
||||
slog.Debug("docker openresty test succeeded", "container", e.ContainerName, "runtime", dockerRuntimeCommand)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) Reload(ctx context.Context) error {
|
||||
return e.EnsureRuntime(ctx, true)
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
slog.Info("ensuring docker openresty runtime", "container", e.ContainerName, "recreate", recreate)
|
||||
output, err := e.Runner.Run(ctx, e.DockerBinary, "inspect", "-f", "{{.State.Running}}", e.ContainerName)
|
||||
if err == nil {
|
||||
if recreate {
|
||||
if err := e.removeContainer(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return e.runContainer(ctx)
|
||||
}
|
||||
if strings.TrimSpace(string(output)) == "true" {
|
||||
slog.Debug("docker openresty runtime already healthy", "container", e.ContainerName)
|
||||
return nil
|
||||
}
|
||||
if err := e.removeContainer(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return e.runContainer(ctx)
|
||||
}
|
||||
return e.runContainer(ctx)
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) CheckHealth(ctx context.Context) error {
|
||||
slog.Debug("checking docker openresty runtime health", "container", e.ContainerName)
|
||||
output, err := e.Runner.Run(ctx, e.DockerBinary, "inspect", "-f", "{{.State.Running}}", e.ContainerName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("docker inspect openresty failed: %w: %s", err, string(output))
|
||||
}
|
||||
if strings.TrimSpace(string(output)) != "true" {
|
||||
return errors.New("docker openresty container is not running")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) Restart(ctx context.Context) error {
|
||||
return e.EnsureRuntime(ctx, true)
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) removeContainer(ctx context.Context) error {
|
||||
slog.Info("removing docker openresty container", "container", e.ContainerName)
|
||||
output, err := e.Runner.Run(ctx, e.DockerBinary, "rm", "-f", e.ContainerName)
|
||||
if err != nil {
|
||||
text := string(output)
|
||||
if strings.Contains(text, "No such container") {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("docker rm openresty failed: %w: %s", err, text)
|
||||
}
|
||||
slog.Info("docker openresty container removed", "container", e.ContainerName)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) runContainer(ctx context.Context) error {
|
||||
slog.Info("starting docker openresty container", "container", e.ContainerName, "image", e.Image)
|
||||
runArgs := []string{
|
||||
"run", "-d",
|
||||
"--name", e.ContainerName,
|
||||
"-p", "80:80",
|
||||
"-p", "443:443",
|
||||
"-p", fmt.Sprintf("127.0.0.1:%d:%d", e.OpenrestyObservabilityPort, e.OpenrestyObservabilityPort),
|
||||
"-v", fmt.Sprintf("%s:%s", e.MainConfigPath, DockerMainConfigPath),
|
||||
"-v", fmt.Sprintf("%s:/etc/nginx/conf.d", e.RouteConfigDir),
|
||||
"-v", fmt.Sprintf("%s:%s", e.CertDir, e.NginxCertDir),
|
||||
"-v", fmt.Sprintf("%s:%s", e.LuaDir, e.NginxLuaDir),
|
||||
e.Image,
|
||||
}
|
||||
runOutput, runErr := e.Runner.Run(ctx, e.DockerBinary, runArgs...)
|
||||
if runErr != nil {
|
||||
return fmt.Errorf("docker run openresty failed: %w: %s", runErr, string(runOutput))
|
||||
}
|
||||
slog.Info("docker openresty container started", "container", e.ContainerName)
|
||||
return nil
|
||||
}
|
||||
|
||||
type Manager struct {
|
||||
MainConfigPath string
|
||||
RouteConfigPath string
|
||||
RuntimeRouteConfigPath string
|
||||
CertDir string
|
||||
NginxCertDir string
|
||||
LuaDir string
|
||||
NginxLuaDir string
|
||||
OpenrestyObservabilityListen string
|
||||
OpenrestyObservabilityPort int
|
||||
Executor Executor
|
||||
}
|
||||
|
||||
func (m *Manager) Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) error {
|
||||
slog.Info("openresty apply started", "main_config", m.MainConfigPath, "route_config", m.RouteConfigPath, "cert_files", len(supportFiles))
|
||||
backup, err := m.backup()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err = m.EnsureLuaAssets(); err != nil {
|
||||
slog.Error("writing lua assets failed, restoring backup", "error", err)
|
||||
_ = m.restore(backup)
|
||||
return err
|
||||
}
|
||||
if err = m.writeCertFiles(supportFiles); err != nil {
|
||||
slog.Error("writing cert files failed, restoring backup", "error", err)
|
||||
_ = m.restore(backup)
|
||||
return err
|
||||
}
|
||||
renderedMainConfig := m.renderMainConfig(mainConfig)
|
||||
if err = os.WriteFile(m.MainConfigPath, []byte(renderedMainConfig), 0o644); err != nil {
|
||||
slog.Error("writing openresty main config failed, restoring backup", "error", err)
|
||||
_ = m.restore(backup)
|
||||
return err
|
||||
}
|
||||
renderedRouteConfig := m.renderRouteConfig(routeConfig)
|
||||
if err = os.WriteFile(m.RouteConfigPath, []byte(renderedRouteConfig), 0o644); err != nil {
|
||||
slog.Error("writing openresty route config failed, restoring backup", "error", err)
|
||||
_ = m.restore(backup)
|
||||
return err
|
||||
}
|
||||
if err = m.Executor.Test(ctx); err != nil {
|
||||
slog.Error("openresty test failed after config write, restoring backup", "error", err)
|
||||
_ = m.restore(backup)
|
||||
return err
|
||||
}
|
||||
if err = m.Executor.Reload(ctx); err != nil {
|
||||
slog.Error("openresty reload failed after config write, restoring backup", "error", err)
|
||||
_ = m.restore(backup)
|
||||
return err
|
||||
}
|
||||
slog.Info("openresty apply completed successfully", "main_config", m.MainConfigPath, "route_config", m.RouteConfigPath)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) EnsureLuaAssets() error {
|
||||
if strings.TrimSpace(m.LuaDir) == "" {
|
||||
return nil
|
||||
}
|
||||
if err := os.RemoveAll(m.LuaDir); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(m.LuaDir, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, file := range ManagedObservabilityLuaFiles() {
|
||||
targetPath, err := luaFileTargetPath(m.LuaDir, file.Path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(targetPath, []byte(file.Content), 0o644); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
if m.Executor == nil {
|
||||
return errors.New("executor 未配置")
|
||||
}
|
||||
slog.Info("openresty ensure runtime requested", "recreate", recreate)
|
||||
return m.Executor.EnsureRuntime(ctx, recreate)
|
||||
}
|
||||
|
||||
func (m *Manager) CheckHealth(ctx context.Context) error {
|
||||
if m.Executor == nil {
|
||||
return errors.New("executor 未配置")
|
||||
}
|
||||
return m.Executor.CheckHealth(ctx)
|
||||
}
|
||||
|
||||
func (m *Manager) Restart(ctx context.Context) error {
|
||||
if m.Executor == nil {
|
||||
return errors.New("executor 未配置")
|
||||
}
|
||||
slog.Info("openresty restart requested")
|
||||
return m.Executor.Restart(ctx)
|
||||
}
|
||||
|
||||
func (m *Manager) CurrentChecksum() (string, error) {
|
||||
if m.RouteConfigPath == "" {
|
||||
return "", errors.New("route config path 不能为空")
|
||||
}
|
||||
if m.MainConfigPath == "" {
|
||||
return "", errors.New("main config path 不能为空")
|
||||
}
|
||||
mainData, err := os.ReadFile(m.MainConfigPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", nil
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
data, err := os.ReadFile(m.RouteConfigPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", nil
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
normalizedMain := string(mainData)
|
||||
if includePath := m.routeConfigIncludePath(); includePath != "" {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, includePath, RouteConfigPlaceholder)
|
||||
}
|
||||
if accessLogPath := m.accessLogRuntimePath(); accessLogPath != "" {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, accessLogPath, AccessLogPlaceholder)
|
||||
}
|
||||
if luaDir := m.luaRuntimePath(); luaDir != "" {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, luaDir, LuaDirPlaceholder)
|
||||
}
|
||||
if listen := strings.TrimSpace(m.OpenrestyObservabilityListen); listen != "" {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, listen, ObservabilityListenPlaceholder)
|
||||
}
|
||||
if m.OpenrestyObservabilityPort > 0 {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, fmt.Sprintf("%d", m.OpenrestyObservabilityPort), ObservabilityPortPlaceholder)
|
||||
}
|
||||
normalizedRoute := string(data)
|
||||
if m.NginxCertDir != "" {
|
||||
normalizedRoute = strings.ReplaceAll(normalizedRoute, m.NginxCertDir, CertDirPlaceholder)
|
||||
}
|
||||
files, err := m.readCertFiles()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
result := bundleChecksum(normalizedMain, normalizedRoute, files)
|
||||
slog.Debug("openresty current checksum calculated", "main_config", m.MainConfigPath, "route_config", m.RouteConfigPath, "checksum", result, "cert_files", len(files))
|
||||
return result, nil
|
||||
}
|
||||
|
||||
type ExecutorOptions struct {
|
||||
NginxPath string
|
||||
DockerBinary string
|
||||
ContainerName string
|
||||
Image string
|
||||
MainConfigPath string
|
||||
RouteConfigPath string
|
||||
CertDir string
|
||||
NginxCertDir string
|
||||
LuaDir string
|
||||
NginxLuaDir string
|
||||
OpenrestyObservabilityPort int
|
||||
}
|
||||
|
||||
func NewExecutor(options ExecutorOptions) Executor {
|
||||
runner := &OSCommandRunner{}
|
||||
if options.NginxPath != "" {
|
||||
return &PathExecutor{
|
||||
Path: options.NginxPath,
|
||||
Runner: runner,
|
||||
}
|
||||
}
|
||||
mainConfigPath := options.MainConfigPath
|
||||
if mainConfigPath != "" {
|
||||
if absPath, err := filepath.Abs(mainConfigPath); err == nil {
|
||||
mainConfigPath = absPath
|
||||
}
|
||||
}
|
||||
routeConfigDir := filepath.Dir(options.RouteConfigPath)
|
||||
if options.RouteConfigPath != "" {
|
||||
if absDir, err := filepath.Abs(routeConfigDir); err == nil {
|
||||
routeConfigDir = absDir
|
||||
}
|
||||
}
|
||||
certDir := options.CertDir
|
||||
if certDir != "" {
|
||||
if absDir, err := filepath.Abs(certDir); err == nil {
|
||||
certDir = absDir
|
||||
}
|
||||
}
|
||||
luaDir := options.LuaDir
|
||||
if luaDir != "" {
|
||||
if absDir, err := filepath.Abs(luaDir); err == nil {
|
||||
luaDir = absDir
|
||||
}
|
||||
}
|
||||
return &DockerExecutor{
|
||||
DockerBinary: options.DockerBinary,
|
||||
ContainerName: options.ContainerName,
|
||||
Image: options.Image,
|
||||
MainConfigPath: mainConfigPath,
|
||||
RouteConfigDir: routeConfigDir,
|
||||
CertDir: certDir,
|
||||
NginxCertDir: options.NginxCertDir,
|
||||
LuaDir: luaDir,
|
||||
NginxLuaDir: options.NginxLuaDir,
|
||||
OpenrestyObservabilityPort: options.OpenrestyObservabilityPort,
|
||||
Runner: runner,
|
||||
}
|
||||
}
|
||||
|
||||
func DetectVersion(ctx context.Context, options ExecutorOptions) string {
|
||||
version, err := detectVersion(ctx, options, &OSCommandRunner{})
|
||||
if err != nil {
|
||||
slog.Error("detect openresty version failed", "error", err)
|
||||
return ""
|
||||
}
|
||||
slog.Info("detected openresty version", "version", version)
|
||||
return version
|
||||
}
|
||||
|
||||
func detectVersion(ctx context.Context, options ExecutorOptions, runner CommandRunner) (string, error) {
|
||||
if runner == nil {
|
||||
runner = &OSCommandRunner{}
|
||||
}
|
||||
if options.NginxPath != "" {
|
||||
output, err := runner.Run(ctx, options.NginxPath, "-v")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("run runtime -v failed: %w: %s", err, string(output))
|
||||
}
|
||||
version := parseNginxVersion(string(output))
|
||||
if version == "" {
|
||||
return "", errors.New("cannot parse runtime version from binary output")
|
||||
}
|
||||
return version, nil
|
||||
}
|
||||
output, err := runDockerVersionProbe(ctx, runner, options.DockerBinary, options.Image)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("run docker %s -v failed: %w: %s", dockerRuntimeCommand, err, string(output))
|
||||
}
|
||||
version := parseNginxVersion(string(output))
|
||||
if version == "" {
|
||||
return "", errors.New("cannot parse runtime version from docker output")
|
||||
}
|
||||
return version, nil
|
||||
}
|
||||
|
||||
func parseNginxVersion(output string) string {
|
||||
matches := nginxVersionPattern.FindStringSubmatch(output)
|
||||
if len(matches) != 2 {
|
||||
return ""
|
||||
}
|
||||
return matches[1]
|
||||
}
|
||||
|
||||
var nginxVersionPattern = regexp.MustCompile(`(?im)(?:nginx|openresty) version:\s*(?:nginx|openresty)/([^\s]+)`)
|
||||
|
||||
func isIgnorableOpenrestyStopError(output string) bool {
|
||||
text := strings.ToLower(strings.TrimSpace(output))
|
||||
if text == "" {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(text, "invalid pid") || strings.Contains(text, "no such process")
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) runEphemeralRuntimeCommand(ctx context.Context, args ...string) ([]byte, error) {
|
||||
return e.runEphemeralRuntimeCommandWithBinary(ctx, dockerRuntimeCommand, args...)
|
||||
}
|
||||
|
||||
func (e *DockerExecutor) runEphemeralRuntimeCommandWithBinary(ctx context.Context, runtimeBinary string, args ...string) ([]byte, error) {
|
||||
runtimeArgs := []string{
|
||||
"run",
|
||||
"--rm",
|
||||
"-v",
|
||||
fmt.Sprintf("%s:%s", e.MainConfigPath, DockerMainConfigPath),
|
||||
"-v",
|
||||
fmt.Sprintf("%s:/etc/nginx/conf.d", e.RouteConfigDir),
|
||||
"-v",
|
||||
fmt.Sprintf("%s:%s", e.CertDir, e.NginxCertDir),
|
||||
"-v",
|
||||
fmt.Sprintf("%s:%s", e.LuaDir, e.NginxLuaDir),
|
||||
e.Image,
|
||||
runtimeBinary,
|
||||
}
|
||||
runtimeArgs = append(runtimeArgs, args...)
|
||||
return e.Runner.Run(ctx, e.DockerBinary, runtimeArgs...)
|
||||
}
|
||||
|
||||
func runDockerVersionProbe(ctx context.Context, runner CommandRunner, dockerBinary string, image string) ([]byte, error) {
|
||||
return runner.Run(ctx, dockerBinary, "run", "--rm", image, dockerRuntimeCommand, "-v")
|
||||
}
|
||||
|
||||
type backupState struct {
|
||||
MainExisted bool
|
||||
MainData []byte
|
||||
RouteExisted bool
|
||||
RouteData []byte
|
||||
Files []protocol.SupportFile
|
||||
}
|
||||
|
||||
func (m *Manager) backup() (*backupState, error) {
|
||||
if m.MainConfigPath == "" {
|
||||
return nil, errors.New("main config path 不能为空")
|
||||
}
|
||||
if m.RouteConfigPath == "" {
|
||||
return nil, errors.New("route config path 不能为空")
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(m.MainConfigPath), 0o755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if m.CertDir != "" {
|
||||
if err := os.MkdirAll(m.CertDir, 0o755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
state := &backupState{}
|
||||
mainData, err := os.ReadFile(m.MainConfigPath)
|
||||
if err == nil {
|
||||
state.MainExisted = true
|
||||
state.MainData = mainData
|
||||
} else if !os.IsNotExist(err) {
|
||||
return nil, err
|
||||
}
|
||||
data, err := os.ReadFile(m.RouteConfigPath)
|
||||
if err == nil {
|
||||
state.RouteExisted = true
|
||||
state.RouteData = data
|
||||
} else if !os.IsNotExist(err) {
|
||||
return nil, err
|
||||
}
|
||||
files, err := m.readCertFiles()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
state.Files = files
|
||||
slog.Debug("backup captured", "main_exists", state.MainExisted, "route_exists", state.RouteExisted, "cert_files", len(state.Files))
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func (m *Manager) restore(state *backupState) error {
|
||||
if state == nil {
|
||||
return nil
|
||||
}
|
||||
slog.Warn("restoring nginx backup", "main_existed", state.MainExisted, "route_existed", state.RouteExisted, "cert_files", len(state.Files))
|
||||
if state.MainExisted {
|
||||
if err := os.WriteFile(m.MainConfigPath, state.MainData, 0o644); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err := os.Remove(m.MainConfigPath); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if state.RouteExisted {
|
||||
if err := os.WriteFile(m.RouteConfigPath, state.RouteData, 0o644); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err := os.Remove(m.RouteConfigPath); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if m.CertDir == "" {
|
||||
return nil
|
||||
}
|
||||
if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(m.CertDir, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, file := range state.Files {
|
||||
targetPath, err := m.certFileTargetPath(file.Path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(targetPath, []byte(file.Content), certFileMode(file.Path)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) writeCertFiles(certFiles []protocol.SupportFile) error {
|
||||
if m.CertDir == "" {
|
||||
return nil
|
||||
}
|
||||
if err := os.RemoveAll(m.CertDir); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(m.CertDir, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, file := range certFiles {
|
||||
targetPath, err := m.certFileTargetPath(file.Path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(targetPath, []byte(file.Content), certFileMode(file.Path)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) readCertFiles() ([]protocol.SupportFile, error) {
|
||||
if m.CertDir == "" {
|
||||
return nil, nil
|
||||
}
|
||||
if _, err := os.Stat(m.CertDir); err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
files := make([]protocol.SupportFile, 0)
|
||||
err := filepath.Walk(m.CertDir, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if info.IsDir() {
|
||||
return nil
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
relativePath, err := filepath.Rel(m.CertDir, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
files = append(files, protocol.SupportFile{
|
||||
Path: filepath.ToSlash(relativePath),
|
||||
Content: string(data),
|
||||
})
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sort.Slice(files, func(i int, j int) bool {
|
||||
return files[i].Path < files[j].Path
|
||||
})
|
||||
return files, nil
|
||||
}
|
||||
|
||||
func (m *Manager) certFileTargetPath(relativePath string) (string, error) {
|
||||
if strings.TrimSpace(m.CertDir) == "" {
|
||||
return "", errors.New("cert dir 不能为空")
|
||||
}
|
||||
candidate := strings.TrimSpace(relativePath)
|
||||
if strings.Contains(candidate, `\`) {
|
||||
candidate = strings.ReplaceAll(candidate, `\`, "/")
|
||||
}
|
||||
normalizedPath := filepath.Clean(filepath.FromSlash(candidate))
|
||||
if normalizedPath == "." || normalizedPath == "" {
|
||||
return "", errors.New("cert file path 不能为空")
|
||||
}
|
||||
if filepath.IsAbs(normalizedPath) || filepath.VolumeName(normalizedPath) != "" {
|
||||
return "", fmt.Errorf("cert file path %q must be relative", relativePath)
|
||||
}
|
||||
targetPath := filepath.Join(m.CertDir, normalizedPath)
|
||||
relativeToBase, err := filepath.Rel(m.CertDir, targetPath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if relativeToBase == ".." || strings.HasPrefix(relativeToBase, ".."+string(os.PathSeparator)) {
|
||||
return "", fmt.Errorf("cert file path %q escapes cert dir", relativePath)
|
||||
}
|
||||
return targetPath, nil
|
||||
}
|
||||
|
||||
func certFileMode(relativePath string) fs.FileMode {
|
||||
switch strings.ToLower(filepath.Ext(strings.TrimSpace(relativePath))) {
|
||||
case ".crt", ".pem":
|
||||
return 0o644
|
||||
case ".key":
|
||||
return 0o600
|
||||
default:
|
||||
return 0o644
|
||||
}
|
||||
}
|
||||
|
||||
func luaFileTargetPath(baseDir string, relativePath string) (string, error) {
|
||||
if strings.TrimSpace(baseDir) == "" {
|
||||
return "", errors.New("lua dir 不能为空")
|
||||
}
|
||||
candidate := strings.TrimSpace(relativePath)
|
||||
if strings.Contains(candidate, `\`) {
|
||||
candidate = strings.ReplaceAll(candidate, `\`, "/")
|
||||
}
|
||||
normalizedPath := filepath.Clean(filepath.FromSlash(candidate))
|
||||
if normalizedPath == "." || normalizedPath == "" {
|
||||
return "", errors.New("lua file path 不能为空")
|
||||
}
|
||||
if filepath.IsAbs(normalizedPath) || filepath.VolumeName(normalizedPath) != "" {
|
||||
return "", fmt.Errorf("lua file path %q must be relative", relativePath)
|
||||
}
|
||||
targetPath := filepath.Join(baseDir, normalizedPath)
|
||||
relativeToBase, err := filepath.Rel(baseDir, targetPath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if relativeToBase == ".." || strings.HasPrefix(relativeToBase, ".."+string(os.PathSeparator)) {
|
||||
return "", fmt.Errorf("lua file path %q escapes lua dir", relativePath)
|
||||
}
|
||||
return targetPath, nil
|
||||
}
|
||||
|
||||
func (m *Manager) renderRouteConfig(content string) string {
|
||||
if m.NginxCertDir == "" {
|
||||
return content
|
||||
}
|
||||
return strings.ReplaceAll(content, CertDirPlaceholder, m.NginxCertDir)
|
||||
}
|
||||
|
||||
func (m *Manager) renderMainConfig(content string) string {
|
||||
rendered := content
|
||||
if includePath := m.routeConfigIncludePath(); includePath != "" {
|
||||
rendered = strings.ReplaceAll(rendered, RouteConfigPlaceholder, includePath)
|
||||
}
|
||||
if accessLogPath := m.accessLogRuntimePath(); accessLogPath != "" {
|
||||
rendered = strings.ReplaceAll(rendered, AccessLogPlaceholder, accessLogPath)
|
||||
}
|
||||
if luaDir := m.luaRuntimePath(); luaDir != "" {
|
||||
rendered = strings.ReplaceAll(rendered, LuaDirPlaceholder, luaDir)
|
||||
}
|
||||
if listen := strings.TrimSpace(m.OpenrestyObservabilityListen); listen != "" {
|
||||
rendered = strings.ReplaceAll(rendered, ObservabilityListenPlaceholder, listen)
|
||||
}
|
||||
if m.OpenrestyObservabilityPort > 0 {
|
||||
rendered = strings.ReplaceAll(rendered, ObservabilityPortPlaceholder, fmt.Sprintf("%d", m.OpenrestyObservabilityPort))
|
||||
}
|
||||
return rendered
|
||||
}
|
||||
|
||||
func ObservabilityListenAddress(openrestyPath string, port int) string {
|
||||
if port <= 0 {
|
||||
return ""
|
||||
}
|
||||
if strings.TrimSpace(openrestyPath) != "" {
|
||||
return fmt.Sprintf("127.0.0.1:%d", port)
|
||||
}
|
||||
return fmt.Sprintf("%d", port)
|
||||
}
|
||||
|
||||
func (m *Manager) routeConfigIncludePath() string {
|
||||
if strings.TrimSpace(m.RuntimeRouteConfigPath) != "" {
|
||||
return strings.TrimSpace(m.RuntimeRouteConfigPath)
|
||||
}
|
||||
return strings.TrimSpace(m.RouteConfigPath)
|
||||
}
|
||||
|
||||
func (m *Manager) accessLogRuntimePath() string {
|
||||
includePath := m.routeConfigIncludePath()
|
||||
if strings.TrimSpace(includePath) == "" {
|
||||
return ""
|
||||
}
|
||||
return filepath.ToSlash(filepath.Join(filepath.Dir(includePath), "openflare_access.log"))
|
||||
}
|
||||
|
||||
func (m *Manager) luaRuntimePath() string {
|
||||
if strings.TrimSpace(m.NginxLuaDir) == "" {
|
||||
return ""
|
||||
}
|
||||
return filepath.ToSlash(m.NginxLuaDir)
|
||||
}
|
||||
|
||||
func checksum(content string) string {
|
||||
sum := sha256.Sum256([]byte(content))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func bundleChecksum(mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) string {
|
||||
files := append([]protocol.SupportFile(nil), supportFiles...)
|
||||
sort.Slice(files, func(i int, j int) bool {
|
||||
return files[i].Path < files[j].Path
|
||||
})
|
||||
var builder strings.Builder
|
||||
builder.WriteString(mainConfig)
|
||||
builder.WriteString("\n--route-config--\n")
|
||||
builder.WriteString(routeConfig)
|
||||
builder.WriteString("\n--support-files--\n")
|
||||
for _, file := range files {
|
||||
builder.WriteString(file.Path)
|
||||
builder.WriteString("\n")
|
||||
builder.WriteString(file.Content)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
return checksum(builder.String())
|
||||
}
|
||||
@@ -0,0 +1,697 @@
|
||||
package nginx
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"openflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
type runCall struct {
|
||||
name string
|
||||
args []string
|
||||
}
|
||||
|
||||
type fakeRunner struct {
|
||||
calls []runCall
|
||||
runFn func(name string, args ...string) ([]byte, error)
|
||||
}
|
||||
|
||||
type fakeExecutor struct {
|
||||
testErr error
|
||||
reloadErr error
|
||||
}
|
||||
|
||||
func (r *fakeRunner) Run(ctx context.Context, name string, args ...string) ([]byte, error) {
|
||||
r.calls = append(r.calls, runCall{name: name, args: append([]string{}, args...)})
|
||||
if r.runFn != nil {
|
||||
return r.runFn(name, args...)
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (e *fakeExecutor) Test(ctx context.Context) error {
|
||||
return e.testErr
|
||||
}
|
||||
|
||||
func (e *fakeExecutor) Reload(ctx context.Context) error {
|
||||
return e.reloadErr
|
||||
}
|
||||
|
||||
func (e *fakeExecutor) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *fakeExecutor) CheckHealth(ctx context.Context) error {
|
||||
return e.testErr
|
||||
}
|
||||
|
||||
func (e *fakeExecutor) Restart(ctx context.Context) error {
|
||||
return e.reloadErr
|
||||
}
|
||||
|
||||
func TestPathExecutorCommands(t *testing.T) {
|
||||
runner := &fakeRunner{}
|
||||
executor := &PathExecutor{
|
||||
Path: "/usr/local/openresty/nginx/sbin/openresty",
|
||||
Runner: runner,
|
||||
}
|
||||
|
||||
if err := executor.Test(context.Background()); err != nil {
|
||||
t.Fatalf("Test failed: %v", err)
|
||||
}
|
||||
if err := executor.Reload(context.Background()); err != nil {
|
||||
t.Fatalf("Reload failed: %v", err)
|
||||
}
|
||||
|
||||
expected := []runCall{
|
||||
{name: "/usr/local/openresty/nginx/sbin/openresty", args: []string{"-t"}},
|
||||
{name: "/usr/local/openresty/nginx/sbin/openresty", args: []string{"-s", "reload"}},
|
||||
}
|
||||
if !reflect.DeepEqual(runner.calls, expected) {
|
||||
t.Fatalf("unexpected calls: %#v", runner.calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathExecutorEnsureRuntimeNoop(t *testing.T) {
|
||||
executor := &PathExecutor{
|
||||
Path: "/usr/local/openresty/nginx/sbin/openresty",
|
||||
Runner: &fakeRunner{},
|
||||
}
|
||||
if err := executor.EnsureRuntime(context.Background(), true); err != nil {
|
||||
t.Fatalf("EnsureRuntime failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathExecutorRestartIgnoresMissingPID(t *testing.T) {
|
||||
runner := &fakeRunner{
|
||||
runFn: func(name string, args ...string) ([]byte, error) {
|
||||
if len(args) == 2 && args[0] == "-s" && args[1] == "quit" {
|
||||
return []byte("openresty: [error] invalid PID number \"\" in \"/usr/local/openresty/nginx/logs/nginx.pid\""), errors.New("exit status 1")
|
||||
}
|
||||
return []byte(""), nil
|
||||
},
|
||||
}
|
||||
executor := &PathExecutor{
|
||||
Path: "/usr/local/openresty/nginx/sbin/openresty",
|
||||
Runner: runner,
|
||||
}
|
||||
if err := executor.Restart(context.Background()); err != nil {
|
||||
t.Fatalf("Restart failed: %v", err)
|
||||
}
|
||||
if len(runner.calls) != 2 {
|
||||
t.Fatalf("expected 2 restart calls, got %d", len(runner.calls))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDockerExecutorCheckHealthFailsWhenContainerStopped(t *testing.T) {
|
||||
runner := &fakeRunner{
|
||||
runFn: func(name string, args ...string) ([]byte, error) {
|
||||
return []byte("false"), nil
|
||||
},
|
||||
}
|
||||
executor := &DockerExecutor{
|
||||
DockerBinary: "docker",
|
||||
ContainerName: "openflare-openresty",
|
||||
Image: "openresty/openresty:alpine",
|
||||
MainConfigPath: filepath.Clean("/tmp/nginx.conf"),
|
||||
RouteConfigDir: filepath.Clean("/tmp/routes"),
|
||||
CertDir: filepath.Clean("/tmp/certs"),
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Clean("/tmp/lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
Runner: runner,
|
||||
}
|
||||
if err := executor.CheckHealth(context.Background()); err == nil {
|
||||
t.Fatal("expected CheckHealth to fail when container is not running")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDockerExecutorStartsContainerWhenMissing(t *testing.T) {
|
||||
runner := &fakeRunner{
|
||||
runFn: func(name string, args ...string) ([]byte, error) {
|
||||
if len(args) >= 1 && args[0] == "inspect" {
|
||||
return []byte(""), errors.New("not found")
|
||||
}
|
||||
return []byte("ok"), nil
|
||||
},
|
||||
}
|
||||
executor := &DockerExecutor{
|
||||
DockerBinary: "docker",
|
||||
ContainerName: "openflare-openresty",
|
||||
Image: "openresty/openresty:alpine",
|
||||
MainConfigPath: filepath.Clean("/tmp/nginx.conf"),
|
||||
RouteConfigDir: filepath.Clean("/tmp/routes"),
|
||||
CertDir: filepath.Clean("/tmp/certs"),
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Clean("/tmp/lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
Runner: runner,
|
||||
}
|
||||
|
||||
if err := executor.Test(context.Background()); err != nil {
|
||||
t.Fatalf("Test failed: %v", err)
|
||||
}
|
||||
|
||||
if len(runner.calls) != 1 {
|
||||
t.Fatalf("expected 1 call, got %d", len(runner.calls))
|
||||
}
|
||||
if runner.calls[0].args[0] != "run" || runner.calls[0].args[1] != "--rm" {
|
||||
t.Fatalf("expected docker run --rm for test, got %#v", runner.calls[0])
|
||||
}
|
||||
if runner.calls[0].args[len(runner.calls[0].args)-2] != "openresty" {
|
||||
t.Fatalf("expected docker test command to invoke openresty, got %#v", runner.calls[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestDockerExecutorStartsStoppedContainer(t *testing.T) {
|
||||
runner := &fakeRunner{
|
||||
runFn: func(name string, args ...string) ([]byte, error) {
|
||||
if len(args) >= 2 && args[0] == "inspect" {
|
||||
return []byte("false"), nil
|
||||
}
|
||||
return []byte("ok"), nil
|
||||
},
|
||||
}
|
||||
executor := &DockerExecutor{
|
||||
DockerBinary: "docker",
|
||||
ContainerName: "openflare-openresty",
|
||||
Image: "openresty/openresty:alpine",
|
||||
MainConfigPath: filepath.Clean("/tmp/nginx.conf"),
|
||||
RouteConfigDir: filepath.Clean("/tmp/routes"),
|
||||
CertDir: filepath.Clean("/tmp/certs"),
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Clean("/tmp/lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
Runner: runner,
|
||||
}
|
||||
|
||||
if err := executor.Reload(context.Background()); err != nil {
|
||||
t.Fatalf("Reload failed: %v", err)
|
||||
}
|
||||
|
||||
if len(runner.calls) != 3 {
|
||||
t.Fatalf("expected 3 calls, got %d", len(runner.calls))
|
||||
}
|
||||
if runner.calls[0].args[0] != "inspect" {
|
||||
t.Fatalf("expected docker inspect on first call, got %#v", runner.calls[0])
|
||||
}
|
||||
if runner.calls[1].args[0] != "rm" {
|
||||
t.Fatalf("expected docker rm on second call, got %#v", runner.calls[1])
|
||||
}
|
||||
if runner.calls[2].args[0] != "run" {
|
||||
t.Fatalf("expected docker run on third call, got %#v", runner.calls[2])
|
||||
}
|
||||
}
|
||||
|
||||
func TestDockerExecutorRunContainerMountsManagedFiles(t *testing.T) {
|
||||
mainConfigPath := filepath.Clean("/tmp/managed/nginx.conf")
|
||||
routeConfigDir := filepath.Clean("/tmp/managed/conf.d")
|
||||
certDir := filepath.Clean("/tmp/managed/certs")
|
||||
luaDir := filepath.Clean("/tmp/managed/lua")
|
||||
runner := &fakeRunner{}
|
||||
executor := &DockerExecutor{
|
||||
DockerBinary: "docker",
|
||||
ContainerName: "openflare-openresty",
|
||||
Image: "openresty/openresty:alpine",
|
||||
MainConfigPath: mainConfigPath,
|
||||
RouteConfigDir: routeConfigDir,
|
||||
CertDir: certDir,
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: luaDir,
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
OpenrestyObservabilityPort: 18081,
|
||||
Runner: runner,
|
||||
}
|
||||
|
||||
if err := executor.runContainer(context.Background()); err != nil {
|
||||
t.Fatalf("runContainer failed: %v", err)
|
||||
}
|
||||
|
||||
if len(runner.calls) != 1 {
|
||||
t.Fatalf("expected one docker run call, got %d", len(runner.calls))
|
||||
}
|
||||
|
||||
expectedArgs := []string{
|
||||
"run", "-d",
|
||||
"--name", "openflare-openresty",
|
||||
"-p", "80:80",
|
||||
"-p", "443:443",
|
||||
"-p", "127.0.0.1:18081:18081",
|
||||
"-v", mainConfigPath + ":" + DockerMainConfigPath,
|
||||
"-v", routeConfigDir + ":/etc/nginx/conf.d",
|
||||
"-v", certDir + ":/etc/nginx/openflare-certs",
|
||||
"-v", luaDir + ":/etc/nginx/openflare-lua",
|
||||
"openresty/openresty:alpine",
|
||||
}
|
||||
if !reflect.DeepEqual(runner.calls[0].args, expectedArgs) {
|
||||
t.Fatalf("unexpected docker run args: %#v", runner.calls[0].args)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDockerExecutorRecreatesContainerOnStartup(t *testing.T) {
|
||||
runner := &fakeRunner{
|
||||
runFn: func(name string, args ...string) ([]byte, error) {
|
||||
if len(args) >= 1 && args[0] == "inspect" {
|
||||
return []byte("true"), nil
|
||||
}
|
||||
return []byte("ok"), nil
|
||||
},
|
||||
}
|
||||
executor := &DockerExecutor{
|
||||
DockerBinary: "docker",
|
||||
ContainerName: "openflare-openresty",
|
||||
Image: "openresty/openresty:alpine",
|
||||
MainConfigPath: filepath.Clean("/tmp/nginx.conf"),
|
||||
RouteConfigDir: filepath.Clean("/tmp/routes"),
|
||||
CertDir: filepath.Clean("/tmp/certs"),
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Clean("/tmp/lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
OpenrestyObservabilityPort: 18081,
|
||||
Runner: runner,
|
||||
}
|
||||
|
||||
if err := executor.EnsureRuntime(context.Background(), true); err != nil {
|
||||
t.Fatalf("EnsureRuntime failed: %v", err)
|
||||
}
|
||||
if len(runner.calls) != 3 {
|
||||
t.Fatalf("expected 3 calls, got %d", len(runner.calls))
|
||||
}
|
||||
if runner.calls[1].args[0] != "rm" {
|
||||
t.Fatalf("expected docker rm on second call, got %#v", runner.calls[1])
|
||||
}
|
||||
if runner.calls[2].args[0] != "run" {
|
||||
t.Fatalf("expected docker run on third call, got %#v", runner.calls[2])
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewExecutorUsesAbsoluteDockerMountPath(t *testing.T) {
|
||||
executor := NewExecutor(ExecutorOptions{
|
||||
DockerBinary: "docker",
|
||||
ContainerName: "openflare-openresty",
|
||||
Image: "openresty/openresty:alpine",
|
||||
MainConfigPath: "./data/etc/nginx/nginx.conf",
|
||||
RouteConfigPath: "./data/etc/nginx/conf.d/openflare_routes.conf",
|
||||
CertDir: "./data/etc/nginx/certs",
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: "./data/etc/nginx/lua",
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
OpenrestyObservabilityPort: 18081,
|
||||
})
|
||||
|
||||
dockerExecutor, ok := executor.(*DockerExecutor)
|
||||
if !ok {
|
||||
t.Fatal("expected docker executor")
|
||||
}
|
||||
if !filepath.IsAbs(dockerExecutor.RouteConfigDir) {
|
||||
t.Fatalf("expected absolute route config dir, got %s", dockerExecutor.RouteConfigDir)
|
||||
}
|
||||
if !filepath.IsAbs(dockerExecutor.MainConfigPath) {
|
||||
t.Fatalf("expected absolute main config path, got %s", dockerExecutor.MainConfigPath)
|
||||
}
|
||||
if !strings.HasSuffix(dockerExecutor.RouteConfigDir, filepath.Clean("data/etc/nginx/conf.d")) {
|
||||
t.Fatalf("unexpected route config dir: %s", dockerExecutor.RouteConfigDir)
|
||||
}
|
||||
if !strings.HasSuffix(dockerExecutor.MainConfigPath, filepath.Clean("data/etc/nginx/nginx.conf")) {
|
||||
t.Fatalf("unexpected main config path: %s", dockerExecutor.MainConfigPath)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectVersionFromBinary(t *testing.T) {
|
||||
version, err := detectVersion(context.Background(), ExecutorOptions{
|
||||
NginxPath: "/usr/local/openresty/nginx/sbin/openresty",
|
||||
}, &fakeRunner{
|
||||
runFn: func(name string, args ...string) ([]byte, error) {
|
||||
return []byte("nginx version: openresty/1.27.1.2\n"), nil
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("detectVersion failed: %v", err)
|
||||
}
|
||||
if version != "1.27.1.2" {
|
||||
t.Fatalf("unexpected version: %s", version)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerApplyAndChecksumIncludeMainConfig(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
mainPath := filepath.Join(tempDir, "nginx.conf")
|
||||
routePath := filepath.Join(tempDir, "conf.d", "openflare_routes.conf")
|
||||
certDir := filepath.Join(tempDir, "certs")
|
||||
manager := &Manager{
|
||||
MainConfigPath: mainPath,
|
||||
RouteConfigPath: routePath,
|
||||
CertDir: certDir,
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Join(tempDir, "lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
Executor: &fakeExecutor{},
|
||||
}
|
||||
|
||||
err := manager.Apply(
|
||||
context.Background(),
|
||||
"include __OPENFLARE_ROUTE_CONFIG__;\naccess_log __OPENFLARE_ACCESS_LOG__ openflare_json;\n",
|
||||
"ssl_certificate __OPENFLARE_CERT_DIR__/1.crt;\n",
|
||||
[]protocol.SupportFile{{Path: "1.crt", Content: "cert"}},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("Apply failed: %v", err)
|
||||
}
|
||||
|
||||
mainData, err := os.ReadFile(mainPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read main config: %v", err)
|
||||
}
|
||||
expectedMain := "include " + routePath + ";\naccess_log " + filepath.Join(filepath.Dir(routePath), "openflare_access.log") + " openflare_json;\n"
|
||||
if string(mainData) != expectedMain {
|
||||
t.Fatalf("unexpected main config: %s", string(mainData))
|
||||
}
|
||||
|
||||
routeData, err := os.ReadFile(routePath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read route config: %v", err)
|
||||
}
|
||||
if string(routeData) != "ssl_certificate /etc/nginx/openflare-certs/1.crt;\n" {
|
||||
t.Fatalf("unexpected route config: %s", string(routeData))
|
||||
}
|
||||
|
||||
value, err := manager.CurrentChecksum()
|
||||
if err != nil {
|
||||
t.Fatalf("CurrentChecksum failed: %v", err)
|
||||
}
|
||||
expected := bundleChecksum(
|
||||
"include __OPENFLARE_ROUTE_CONFIG__;\naccess_log __OPENFLARE_ACCESS_LOG__ openflare_json;\n",
|
||||
"ssl_certificate __OPENFLARE_CERT_DIR__/1.crt;\n",
|
||||
[]protocol.SupportFile{{Path: "1.crt", Content: "cert"}},
|
||||
)
|
||||
if value != expected {
|
||||
t.Fatalf("unexpected checksum: got %s want %s", value, expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerApplyUsesRuntimeRouteConfigPath(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
mainPath := filepath.Join(tempDir, "nginx.conf")
|
||||
routePath := filepath.Join(tempDir, "conf.d", "openflare_routes.conf")
|
||||
manager := &Manager{
|
||||
MainConfigPath: mainPath,
|
||||
RouteConfigPath: routePath,
|
||||
RuntimeRouteConfigPath: DockerRouteConfigPath,
|
||||
CertDir: filepath.Join(tempDir, "certs"),
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Join(tempDir, "lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
Executor: &fakeExecutor{},
|
||||
}
|
||||
|
||||
if err := manager.Apply(context.Background(), "include __OPENFLARE_ROUTE_CONFIG__;\naccess_log __OPENFLARE_ACCESS_LOG__ openflare_json;\n", "server { listen 80; }\n", nil); err != nil {
|
||||
t.Fatalf("Apply failed: %v", err)
|
||||
}
|
||||
|
||||
mainData, err := os.ReadFile(mainPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read main config: %v", err)
|
||||
}
|
||||
expectedMain := "include " + DockerRouteConfigPath + ";\naccess_log " + DockerAccessLogPath + " openflare_json;\n"
|
||||
if string(mainData) != expectedMain {
|
||||
t.Fatalf("unexpected main config include path: %s", string(mainData))
|
||||
}
|
||||
|
||||
value, err := manager.CurrentChecksum()
|
||||
if err != nil {
|
||||
t.Fatalf("CurrentChecksum failed: %v", err)
|
||||
}
|
||||
expected := bundleChecksum(
|
||||
"include __OPENFLARE_ROUTE_CONFIG__;\naccess_log __OPENFLARE_ACCESS_LOG__ openflare_json;\n",
|
||||
"server { listen 80; }\n",
|
||||
nil,
|
||||
)
|
||||
if value != expected {
|
||||
t.Fatalf("unexpected checksum: got %s want %s", value, expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectVersionFromDockerImage(t *testing.T) {
|
||||
runner := &fakeRunner{
|
||||
runFn: func(name string, args ...string) ([]byte, error) {
|
||||
return []byte("nginx version: openresty/1.27.1.2\n"), nil
|
||||
},
|
||||
}
|
||||
version, err := detectVersion(context.Background(), ExecutorOptions{
|
||||
DockerBinary: "docker",
|
||||
Image: "openresty/openresty:alpine",
|
||||
}, runner)
|
||||
if err != nil {
|
||||
t.Fatalf("detectVersion failed: %v", err)
|
||||
}
|
||||
if version != "1.27.1.2" {
|
||||
t.Fatalf("unexpected version: %s", version)
|
||||
}
|
||||
if len(runner.calls) != 1 {
|
||||
t.Fatalf("expected one command call, got %d", len(runner.calls))
|
||||
}
|
||||
expectedArgs := []string{"run", "--rm", "openresty/openresty:alpine", "openresty", "-v"}
|
||||
if !reflect.DeepEqual(runner.calls[0].args, expectedArgs) {
|
||||
t.Fatalf("unexpected docker args: %#v", runner.calls[0].args)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseNginxVersionIgnoresDockerEntrypointPaths(t *testing.T) {
|
||||
output := strings.Join([]string{
|
||||
"/docker-entrypoint.sh: /docker-entrypoint.d/10-listen-on-ipv6-by-default.sh: info: can not modify /etc/nginx/conf.d/default.conf (read-only file system?)",
|
||||
"nginx version: openresty/1.27.1.2",
|
||||
}, "\n")
|
||||
|
||||
version := parseNginxVersion(output)
|
||||
if version != "1.27.1.2" {
|
||||
t.Fatalf("unexpected version: %s", version)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerApplyWritesSupportFilesAndReplacesPlaceholder(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
manager := &Manager{
|
||||
MainConfigPath: filepath.Join(tempDir, "nginx.conf"),
|
||||
RouteConfigPath: filepath.Join(tempDir, "routes.conf"),
|
||||
CertDir: filepath.Join(tempDir, "certs"),
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Join(tempDir, "lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
OpenrestyObservabilityListen: "18081",
|
||||
Executor: &fakeExecutor{},
|
||||
}
|
||||
|
||||
err := manager.Apply(context.Background(), "include __OPENFLARE_ROUTE_CONFIG__;\nserver { listen __OPENFLARE_OBSERVABILITY_LISTEN__; }", "ssl_certificate __OPENFLARE_CERT_DIR__/1.crt;", []protocol.SupportFile{
|
||||
{Path: "1.crt", Content: "cert-data"},
|
||||
{Path: "1.key", Content: "key-data"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Apply failed: %v", err)
|
||||
}
|
||||
|
||||
routeData, err := os.ReadFile(manager.RouteConfigPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read route config: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(routeData), "/etc/nginx/openflare-certs/1.crt") {
|
||||
t.Fatalf("expected placeholder replacement in route config, got %s", string(routeData))
|
||||
}
|
||||
mainData, err := os.ReadFile(manager.MainConfigPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read main config: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(mainData), "listen 18081;") {
|
||||
t.Fatalf("expected observability listen placeholder replacement in main config, got %s", string(mainData))
|
||||
}
|
||||
certData, err := os.ReadFile(filepath.Join(manager.CertDir, "1.crt"))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read cert file: %v", err)
|
||||
}
|
||||
if string(certData) != "cert-data" {
|
||||
t.Fatalf("unexpected cert file content: %s", string(certData))
|
||||
}
|
||||
luaInfo, err := os.Stat(filepath.Join(manager.LuaDir, "log.lua"))
|
||||
if err != nil {
|
||||
t.Fatalf("expected managed lua file to exist, stat err = %v", err)
|
||||
}
|
||||
if luaInfo.Mode().Perm() != 0o644 {
|
||||
t.Fatalf("unexpected lua mode: %o", luaInfo.Mode().Perm())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCertFileMode(t *testing.T) {
|
||||
testCases := []struct {
|
||||
path string
|
||||
want os.FileMode
|
||||
}{
|
||||
{path: "1.crt", want: 0o644},
|
||||
{path: "1.pem", want: 0o644},
|
||||
{path: "1.key", want: 0o600},
|
||||
{path: "misc.txt", want: 0o644},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
if got := certFileMode(testCase.path); got != testCase.want {
|
||||
t.Fatalf("unexpected mode for %s: got %o want %o", testCase.path, got, testCase.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerEnsureLuaAssetsWritesReadableFiles(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
manager := &Manager{
|
||||
LuaDir: filepath.Join(tempDir, "lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
}
|
||||
|
||||
err := manager.EnsureLuaAssets()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureLuaAssets failed: %v", err)
|
||||
}
|
||||
|
||||
luaInfo, err := os.Stat(filepath.Join(manager.LuaDir, "log.lua"))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to stat lua file: %v", err)
|
||||
}
|
||||
if luaInfo.Mode().Perm() != 0o644 {
|
||||
t.Fatalf("unexpected lua mode: %o", luaInfo.Mode().Perm())
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerRollbackRestoresCertFiles(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
routePath := filepath.Join(tempDir, "routes.conf")
|
||||
mainPath := filepath.Join(tempDir, "nginx.conf")
|
||||
certDir := filepath.Join(tempDir, "certs")
|
||||
if err := os.MkdirAll(certDir, 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll failed: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(mainPath, []byte("old-main"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile failed: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(routePath, []byte("old-route"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile failed: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(certDir, "1.crt"), []byte("old-cert"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile failed: %v", err)
|
||||
}
|
||||
manager := &Manager{
|
||||
MainConfigPath: mainPath,
|
||||
RouteConfigPath: routePath,
|
||||
CertDir: certDir,
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Join(tempDir, "lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
Executor: &fakeExecutor{
|
||||
testErr: errors.New("openresty test failed"),
|
||||
},
|
||||
}
|
||||
|
||||
err := manager.Apply(context.Background(), "new-main", "new-route", []protocol.SupportFile{
|
||||
{Path: "1.crt", Content: "new-cert"},
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected Apply to fail")
|
||||
}
|
||||
|
||||
mainData, err := os.ReadFile(mainPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read main config: %v", err)
|
||||
}
|
||||
if string(mainData) != "old-main" {
|
||||
t.Fatalf("expected main rollback, got %s", string(mainData))
|
||||
}
|
||||
routeData, err := os.ReadFile(routePath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read route config: %v", err)
|
||||
}
|
||||
if string(routeData) != "old-route" {
|
||||
t.Fatalf("expected route rollback, got %s", string(routeData))
|
||||
}
|
||||
certData, err := os.ReadFile(filepath.Join(certDir, "1.crt"))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read cert file: %v", err)
|
||||
}
|
||||
if string(certData) != "old-cert" {
|
||||
t.Fatalf("expected cert rollback, got %s", string(certData))
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCertFileTargetPathRejectsEscapes(t *testing.T) {
|
||||
manager := &Manager{CertDir: filepath.Join(t.TempDir(), "certs")}
|
||||
if err := os.MkdirAll(manager.CertDir, 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll failed: %v", err)
|
||||
}
|
||||
|
||||
absolutePath := "/tmp/evil.crt"
|
||||
if runtime.GOOS == "windows" {
|
||||
absolutePath = `C:/tmp/evil.crt`
|
||||
}
|
||||
|
||||
testCases := []struct {
|
||||
path string
|
||||
shouldErr bool
|
||||
}{
|
||||
{path: "nested/1.crt", shouldErr: false},
|
||||
{path: "../escape.crt", shouldErr: true},
|
||||
{path: "..\\escape.crt", shouldErr: true},
|
||||
{path: absolutePath, shouldErr: true},
|
||||
{path: "", shouldErr: true},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
targetPath, err := manager.certFileTargetPath(testCase.path)
|
||||
if testCase.shouldErr {
|
||||
if err == nil {
|
||||
t.Fatalf("expected path %q to be rejected, got target %q", testCase.path, targetPath)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("expected path %q to be accepted: %v", testCase.path, err)
|
||||
}
|
||||
if !strings.HasPrefix(targetPath, manager.CertDir) {
|
||||
t.Fatalf("expected target path %q to stay under %q", targetPath, manager.CertDir)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerApplyRejectsCertFilePathTraversal(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
manager := &Manager{
|
||||
MainConfigPath: filepath.Join(tempDir, "nginx.conf"),
|
||||
RouteConfigPath: filepath.Join(tempDir, "routes.conf"),
|
||||
CertDir: filepath.Join(tempDir, "certs"),
|
||||
NginxCertDir: "/etc/nginx/openflare-certs",
|
||||
LuaDir: filepath.Join(tempDir, "lua"),
|
||||
NginxLuaDir: "/etc/nginx/openflare-lua",
|
||||
Executor: &fakeExecutor{},
|
||||
}
|
||||
|
||||
err := manager.Apply(context.Background(), "main", "route", []protocol.SupportFile{
|
||||
{Path: "../escape.crt", Content: "bad"},
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected Apply to reject traversal path")
|
||||
}
|
||||
|
||||
if _, statErr := os.Stat(filepath.Join(tempDir, "escape.crt")); !os.IsNotExist(statErr) {
|
||||
t.Fatalf("expected escaped file to not exist, stat err = %v", statErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestObservabilityListenAddress(t *testing.T) {
|
||||
if got := ObservabilityListenAddress("", 18081); got != "18081" {
|
||||
t.Fatalf("unexpected docker observability listen address: %s", got)
|
||||
}
|
||||
if got := ObservabilityListenAddress("/usr/local/openresty/nginx/sbin/openresty", 18081); got != "127.0.0.1:18081" {
|
||||
t.Fatalf("unexpected path observability listen address: %s", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
package nginx
|
||||
|
||||
import "openflare-agent/internal/protocol"
|
||||
|
||||
const (
|
||||
openRestyObservabilityWindowTTL = 7200
|
||||
openRestyObservabilityWindowSize = 60
|
||||
)
|
||||
|
||||
const openRestyObservabilityInitLua = `local dict = ngx.shared.openflare_observability
|
||||
if not dict then
|
||||
return
|
||||
end
|
||||
|
||||
return
|
||||
`
|
||||
|
||||
const openRestyObservabilityLogLua = `local dict = ngx.shared.openflare_observability
|
||||
if not dict then
|
||||
return
|
||||
end
|
||||
|
||||
local request_uri = tostring(ngx.var.uri or "")
|
||||
if request_uri == "/openflare/observability" or request_uri == "/openflare/stub_status" then
|
||||
return
|
||||
end
|
||||
|
||||
local ttl = ` + "7200" + `
|
||||
local now = ngx.time()
|
||||
local window_size = ` + "60" + `
|
||||
local window_start = now - (now % window_size)
|
||||
|
||||
local function ensure_counter(key)
|
||||
dict:add(key, 0, ttl)
|
||||
end
|
||||
|
||||
local function incr(key, delta)
|
||||
ensure_counter(key)
|
||||
local value, err = dict:incr(key, delta)
|
||||
if not value and err == "not found" then
|
||||
dict:set(key, delta, ttl)
|
||||
end
|
||||
end
|
||||
|
||||
local function remember_value(list_key, marker_key, value)
|
||||
if value == "" then
|
||||
return
|
||||
end
|
||||
if not dict:add(marker_key, 1, ttl) then
|
||||
return
|
||||
end
|
||||
local existing = dict:get(list_key)
|
||||
if not existing or existing == "" then
|
||||
dict:set(list_key, value, ttl)
|
||||
return
|
||||
end
|
||||
dict:set(list_key, existing .. "\n" .. value, ttl)
|
||||
end
|
||||
|
||||
local window_prefix = tostring(window_start)
|
||||
incr("request_count:" .. window_prefix, 1)
|
||||
|
||||
local status = tostring(ngx.status or 0)
|
||||
if status ~= "0" then
|
||||
incr("status:" .. window_prefix .. ":" .. status, 1)
|
||||
remember_value(
|
||||
"status_keys:" .. window_prefix,
|
||||
"status_marker:" .. window_prefix .. ":" .. status,
|
||||
status
|
||||
)
|
||||
if tonumber(status) and tonumber(status) >= 500 then
|
||||
incr("error_count:" .. window_prefix, 1)
|
||||
end
|
||||
end
|
||||
|
||||
local host = tostring(ngx.var.host or "")
|
||||
if host ~= "" then
|
||||
incr("domain:" .. window_prefix .. ":" .. host, 1)
|
||||
remember_value(
|
||||
"domain_keys:" .. window_prefix,
|
||||
"domain_marker:" .. window_prefix .. ":" .. host,
|
||||
host
|
||||
)
|
||||
end
|
||||
|
||||
local remote_addr = tostring(ngx.var.binary_remote_addr or ngx.var.remote_addr or "")
|
||||
if remote_addr ~= "" and dict:add("visitor:" .. window_prefix .. ":" .. remote_addr, 1, ttl) then
|
||||
incr("unique_visitor_count:" .. window_prefix, 1)
|
||||
end
|
||||
|
||||
local request_length = tonumber(ngx.var.request_length) or 0
|
||||
if request_length > 0 then
|
||||
incr("openresty_rx_bytes:" .. window_prefix, request_length)
|
||||
end
|
||||
|
||||
local bytes_sent = tonumber(ngx.var.bytes_sent) or tonumber(ngx.var.body_bytes_sent) or 0
|
||||
if bytes_sent > 0 then
|
||||
incr("openresty_tx_bytes:" .. window_prefix, bytes_sent)
|
||||
end
|
||||
`
|
||||
|
||||
const openRestyObservabilityReadLua = `local cjson = require "cjson.safe"
|
||||
|
||||
local dict = ngx.shared.openflare_observability
|
||||
if not dict then
|
||||
ngx.status = ngx.HTTP_SERVICE_UNAVAILABLE
|
||||
ngx.say(cjson.encode({ message = "shared dict unavailable" }))
|
||||
return
|
||||
end
|
||||
|
||||
local now = ngx.time()
|
||||
local window_size = ` + "60" + `
|
||||
local window_start = now - (now % window_size)
|
||||
local current_window = tostring(window_start)
|
||||
|
||||
local function read_counter(key)
|
||||
return tonumber(dict:get(key) or 0) or 0
|
||||
end
|
||||
|
||||
local function read_map(window_id, prefix, list_key)
|
||||
local result = {}
|
||||
local raw = dict:get(list_key .. ":" .. window_id)
|
||||
if not raw or raw == "" then
|
||||
return result
|
||||
end
|
||||
for value in string.gmatch(raw, "[^\n]+") do
|
||||
result[value] = read_counter(prefix .. ":" .. window_id .. ":" .. value)
|
||||
end
|
||||
return result
|
||||
end
|
||||
|
||||
local payload = {
|
||||
window_started_at_unix = window_start,
|
||||
window_ended_at_unix = now,
|
||||
request_count = read_counter("request_count:" .. current_window),
|
||||
error_count = read_counter("error_count:" .. current_window),
|
||||
unique_visitor_count = read_counter("unique_visitor_count:" .. current_window),
|
||||
status_codes = read_map(current_window, "status", "status_keys"),
|
||||
top_domains = read_map(current_window, "domain", "domain_keys"),
|
||||
source_countries = {},
|
||||
openresty_rx_bytes = read_counter("openresty_rx_bytes:" .. current_window),
|
||||
openresty_tx_bytes = read_counter("openresty_tx_bytes:" .. current_window)
|
||||
}
|
||||
|
||||
ngx.header.content_type = "application/json"
|
||||
ngx.say(cjson.encode(payload))
|
||||
`
|
||||
|
||||
func ManagedObservabilityLuaFiles() []protocol.SupportFile {
|
||||
return []protocol.SupportFile{
|
||||
{Path: "init.lua", Content: openRestyObservabilityInitLua},
|
||||
{Path: "log.lua", Content: openRestyObservabilityLogLua},
|
||||
{Path: "read.lua", Content: openRestyObservabilityReadLua},
|
||||
{Path: "observability/init.lua", Content: openRestyObservabilityInitLua},
|
||||
{Path: "observability/log.lua", Content: openRestyObservabilityLogLua},
|
||||
{Path: "observability/read.lua", Content: openRestyObservabilityReadLua},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,396 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"openflare-agent/internal/config"
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
func BuildProfile(cfg *config.Config, stateStore *state.Store) *protocol.NodeSystemProfile {
|
||||
profile := collectProfile(cfg)
|
||||
if profile == nil {
|
||||
return nil
|
||||
}
|
||||
fingerprint := fingerprintProfile(profile)
|
||||
if stateStore == nil {
|
||||
return profile
|
||||
}
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
return profile
|
||||
}
|
||||
if snapshot.LastProfileFingerprint == fingerprint {
|
||||
return nil
|
||||
}
|
||||
snapshot.LastProfileFingerprint = fingerprint
|
||||
if err = stateStore.Save(snapshot); err != nil {
|
||||
return profile
|
||||
}
|
||||
return profile
|
||||
}
|
||||
|
||||
func BuildSnapshot(cfg *config.Config, stateStore *state.Store, managed *managedOpenRestyMetrics) *protocol.NodeMetricSnapshot {
|
||||
now := time.Now().UTC()
|
||||
metric := &protocol.NodeMetricSnapshot{
|
||||
CapturedAtUnix: now.Unix(),
|
||||
}
|
||||
|
||||
memTotal, memUsed := readMemInfo()
|
||||
metric.MemoryTotalBytes = memTotal
|
||||
metric.MemoryUsedBytes = memUsed
|
||||
|
||||
storageTotal, storageUsed := statFilesystem(cfg.DataDir)
|
||||
metric.StorageTotalBytes = storageTotal
|
||||
metric.StorageUsedBytes = storageUsed
|
||||
|
||||
metric.NetworkRxBytes, metric.NetworkTxBytes = readLinuxNetworkTotals()
|
||||
metric.DiskReadBytes, metric.DiskWriteBytes = readLinuxDiskTotals()
|
||||
if managed != nil {
|
||||
metric.OpenrestyRxBytes = managed.OpenrestyRxBytes
|
||||
metric.OpenrestyTxBytes = managed.OpenrestyTxBytes
|
||||
metric.OpenrestyConnections = managed.OpenrestyConnections
|
||||
}
|
||||
|
||||
if stateStore == nil {
|
||||
return metric
|
||||
}
|
||||
|
||||
totalCPU, idleCPU := readLinuxCPUStat()
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
return metric
|
||||
}
|
||||
if snapshot.LastCPUStatTotal > 0 && totalCPU > snapshot.LastCPUStatTotal && idleCPU >= snapshot.LastCPUStatIdle {
|
||||
deltaTotal := totalCPU - snapshot.LastCPUStatTotal
|
||||
deltaIdle := idleCPU - snapshot.LastCPUStatIdle
|
||||
if deltaTotal > 0 && deltaIdle <= deltaTotal {
|
||||
metric.CPUUsagePercent = (float64(deltaTotal-deltaIdle) / float64(deltaTotal)) * 100
|
||||
}
|
||||
}
|
||||
snapshot.LastCPUStatTotal = totalCPU
|
||||
snapshot.LastCPUStatIdle = idleCPU
|
||||
snapshot.LastMetricAtUnix = now.Unix()
|
||||
_ = stateStore.Save(snapshot)
|
||||
|
||||
return metric
|
||||
}
|
||||
|
||||
func BuildHealthEvents(snapshot *state.Snapshot) []protocol.NodeHealthEvent {
|
||||
if snapshot == nil {
|
||||
return []protocol.NodeHealthEvent{}
|
||||
}
|
||||
events := make([]protocol.NodeHealthEvent, 0, 2)
|
||||
nowUnix := time.Now().UTC().Unix()
|
||||
if strings.TrimSpace(snapshot.OpenrestyStatus) == protocol.OpenrestyStatusUnhealthy {
|
||||
events = append(events, protocol.NodeHealthEvent{
|
||||
EventType: "openresty_unhealthy",
|
||||
Severity: "critical",
|
||||
Message: strings.TrimSpace(snapshot.OpenrestyMessage),
|
||||
TriggeredAtUnix: nowUnix,
|
||||
})
|
||||
}
|
||||
if strings.TrimSpace(snapshot.LastError) != "" {
|
||||
events = append(events, protocol.NodeHealthEvent{
|
||||
EventType: "sync_error",
|
||||
Severity: "warning",
|
||||
Message: strings.TrimSpace(snapshot.LastError),
|
||||
TriggeredAtUnix: nowUnix,
|
||||
})
|
||||
}
|
||||
return events
|
||||
}
|
||||
|
||||
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()
|
||||
|
||||
return &protocol.NodeSystemProfile{
|
||||
Hostname: strings.TrimSpace(hostname),
|
||||
OSName: osName,
|
||||
OSVersion: osVersion,
|
||||
KernelVersion: kernelVersion,
|
||||
Architecture: runtime.GOARCH,
|
||||
CPUModel: cpuModel,
|
||||
CPUCores: runtime.NumCPU(),
|
||||
TotalMemoryBytes: totalMemory,
|
||||
TotalDiskBytes: totalDisk,
|
||||
UptimeSeconds: uptimeSeconds,
|
||||
ReportedAtUnix: time.Now().UTC().Unix(),
|
||||
}
|
||||
}
|
||||
|
||||
func fingerprintProfile(profile *protocol.NodeSystemProfile) string {
|
||||
raw, err := json.Marshal(profile)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
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))
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"openflare-agent/internal/config"
|
||||
"openflare-agent/internal/protocol"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const openRestyObservabilityPath = "/openflare/observability"
|
||||
const openRestyStubStatusPath = "/openflare/stub_status"
|
||||
|
||||
var stubStatusActivePattern = regexp.MustCompile(`Active connections:\s+(\d+)`)
|
||||
|
||||
type managedOpenRestyMetrics struct {
|
||||
TrafficReport *protocol.NodeTrafficReport
|
||||
OpenrestyRxBytes int64
|
||||
OpenrestyTxBytes int64
|
||||
OpenrestyConnections int64
|
||||
}
|
||||
|
||||
type openRestyObservabilityResponse 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"`
|
||||
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
|
||||
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
|
||||
}
|
||||
|
||||
func CollectManagedOpenRestyMetrics(cfg *config.Config) *managedOpenRestyMetrics {
|
||||
if cfg == nil || cfg.OpenrestyObservabilityPort <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
baseURL := fmt.Sprintf("http://127.0.0.1:%d", cfg.OpenrestyObservabilityPort)
|
||||
client := &http.Client{Timeout: 1500 * time.Millisecond}
|
||||
|
||||
observabilityResp := openRestyObservabilityResponse{}
|
||||
if err := fetchLocalJSON(client, baseURL+openRestyObservabilityPath, &observabilityResp); err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
result := &managedOpenRestyMetrics{
|
||||
TrafficReport: &protocol.NodeTrafficReport{
|
||||
WindowStartedAtUnix: observabilityResp.WindowStartedAtUnix,
|
||||
WindowEndedAtUnix: observabilityResp.WindowEndedAtUnix,
|
||||
RequestCount: observabilityResp.RequestCount,
|
||||
ErrorCount: observabilityResp.ErrorCount,
|
||||
UniqueVisitorCount: observabilityResp.UniqueVisitorCount,
|
||||
StatusCodes: normalizeCountMap(observabilityResp.StatusCodes),
|
||||
TopDomains: normalizeCountMap(observabilityResp.TopDomains),
|
||||
SourceCountries: normalizeCountMap(observabilityResp.SourceCountries),
|
||||
},
|
||||
OpenrestyRxBytes: observabilityResp.OpenrestyRxBytes,
|
||||
OpenrestyTxBytes: observabilityResp.OpenrestyTxBytes,
|
||||
}
|
||||
|
||||
if text, err := fetchLocalText(client, baseURL+openRestyStubStatusPath); err == nil {
|
||||
result.OpenrestyConnections = parseStubStatusActiveConnections(text)
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func fetchLocalJSON(client *http.Client, url string, target any) error {
|
||||
resp, err := client.Get(url)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("unexpected local observability status: %s", resp.Status)
|
||||
}
|
||||
return json.NewDecoder(resp.Body).Decode(target)
|
||||
}
|
||||
|
||||
func fetchLocalText(client *http.Client, url string) (string, error) {
|
||||
resp, err := client.Get(url)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("unexpected local stub status: %s", resp.Status)
|
||||
}
|
||||
data, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func parseStubStatusActiveConnections(raw string) int64 {
|
||||
matches := stubStatusActivePattern.FindStringSubmatch(raw)
|
||||
if len(matches) != 2 {
|
||||
return 0
|
||||
}
|
||||
value, err := strconv.ParseInt(matches[1], 10, 64)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func normalizeCountMap(values map[string]int64) map[string]int64 {
|
||||
if len(values) == 0 {
|
||||
return map[string]int64{}
|
||||
}
|
||||
result := make(map[string]int64, len(values))
|
||||
for key, value := range values {
|
||||
key = strings.TrimSpace(key)
|
||||
if key == "" || value <= 0 {
|
||||
continue
|
||||
}
|
||||
result[key] = value
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"openflare-agent/internal/config"
|
||||
)
|
||||
|
||||
func TestCollectManagedOpenRestyMetrics(t *testing.T) {
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("Listen failed: %v", err)
|
||||
}
|
||||
port := listener.Addr().(*net.TCPAddr).Port
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc(openRestyObservabilityPath, func(writer http.ResponseWriter, request *http.Request) {
|
||||
writer.Header().Set("Content-Type", "application/json")
|
||||
_, _ = writer.Write([]byte(`{"window_started_at_unix":1710403200,"window_ended_at_unix":1710403210,"request_count":12,"error_count":2,"unique_visitor_count":5,"status_codes":{"200":10,"502":2},"top_domains":{"app.example.com":9,"api.example.com":3},"source_countries":{},"openresty_rx_bytes":4096,"openresty_tx_bytes":8192}`))
|
||||
})
|
||||
mux.HandleFunc(openRestyStubStatusPath, func(writer http.ResponseWriter, request *http.Request) {
|
||||
_, _ = writer.Write([]byte("Active connections: 7 \nserver accepts handled requests\n 10 10 12 \nReading: 1 Writing: 2 Waiting: 4 \n"))
|
||||
})
|
||||
|
||||
server := httptest.NewUnstartedServer(mux)
|
||||
server.Listener = listener
|
||||
server.Start()
|
||||
defer server.Close()
|
||||
|
||||
metrics := CollectManagedOpenRestyMetrics(&config.Config{
|
||||
OpenrestyObservabilityPort: port,
|
||||
})
|
||||
if metrics == nil || metrics.TrafficReport == nil {
|
||||
t.Fatalf("expected managed openresty metrics, got %+v", metrics)
|
||||
}
|
||||
if metrics.TrafficReport.RequestCount != 12 || metrics.TrafficReport.ErrorCount != 2 {
|
||||
t.Fatalf("unexpected traffic report: %+v", metrics.TrafficReport)
|
||||
}
|
||||
if metrics.OpenrestyRxBytes != 4096 || metrics.OpenrestyTxBytes != 8192 {
|
||||
t.Fatalf("unexpected openresty byte counters: %+v", metrics)
|
||||
}
|
||||
if metrics.OpenrestyConnections != 7 {
|
||||
t.Fatalf("unexpected openresty connections: %+v", metrics)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseStubStatusActiveConnections(t *testing.T) {
|
||||
if value := parseStubStatusActiveConnections("Active connections: 19\n"); value != 19 {
|
||||
t.Fatalf("unexpected active connections: %d", value)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeCountMapDropsEmptyKeys(t *testing.T) {
|
||||
normalized := normalizeCountMap(map[string]int64{
|
||||
"": 4,
|
||||
" 200 ": 3,
|
||||
"app.example.com": 0,
|
||||
})
|
||||
if len(normalized) != 1 || normalized["200"] != 3 {
|
||||
t.Fatalf("unexpected normalized map: %+v", normalized)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectManagedOpenRestyMetricsHandlesUnavailableEndpoint(t *testing.T) {
|
||||
cfg := &config.Config{OpenrestyObservabilityPort: 1}
|
||||
if metrics := CollectManagedOpenRestyMetrics(cfg); metrics != nil {
|
||||
t.Fatalf("expected nil metrics for unavailable endpoint, got %+v", metrics)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenRestyObservabilityPathsAreStable(t *testing.T) {
|
||||
if !strings.HasPrefix(openRestyObservabilityPath, "/openflare/") {
|
||||
t.Fatalf("unexpected observability path: %s", openRestyObservabilityPath)
|
||||
}
|
||||
if !strings.HasPrefix(openRestyStubStatusPath, "/openflare/") {
|
||||
t.Fatalf("unexpected stub status path: %s", openRestyStubStatusPath)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,346 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"openflare-agent/internal/config"
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type accessLogRecord struct {
|
||||
Timestamp string `json:"ts"`
|
||||
Host string `json:"host"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Path string `json:"path"`
|
||||
Status int `json:"status"`
|
||||
BytesSent int64 `json:"bytes_sent"`
|
||||
RequestLength int64 `json:"request_length"`
|
||||
}
|
||||
|
||||
var combinedAccessLogPattern = regexp.MustCompile(`^(\S+)\s+\S+\s+\S+\s+\[([^\]]+)\]\s+"(?:\S+)\s+(\S+)(?:\s+[^"]*)?"\s+(\d{3})\s+\S+`)
|
||||
|
||||
type trafficAggregate struct {
|
||||
windowStartedAt time.Time
|
||||
windowEndedAt time.Time
|
||||
requestCount int64
|
||||
errorCount int64
|
||||
openrestyRxBytes int64
|
||||
openrestyTxBytes int64
|
||||
statusCodes map[string]int64
|
||||
topDomains map[string]int64
|
||||
visitors map[string]struct{}
|
||||
logs []protocol.NodeAccessLog
|
||||
}
|
||||
|
||||
func BuildTrafficReport(cfg *config.Config, stateStore *state.Store, managed *managedOpenRestyMetrics) *protocol.NodeTrafficReport {
|
||||
report, _, _ := BuildTrafficObservability(cfg, stateStore, managed)
|
||||
return report
|
||||
}
|
||||
|
||||
func BuildTrafficObservability(cfg *config.Config, stateStore *state.Store, managed *managedOpenRestyMetrics) (*protocol.NodeTrafficReport, []protocol.NodeAccessLog, *managedOpenRestyMetrics) {
|
||||
if cfg == nil || stateStore == nil {
|
||||
if managed != nil && managed.TrafficReport != nil {
|
||||
return managed.TrafficReport, nil, managed
|
||||
}
|
||||
return nil, nil, managed
|
||||
}
|
||||
|
||||
aggregate := readAccessLogDelta(cfg, stateStore)
|
||||
accessLogs := []protocol.NodeAccessLog{}
|
||||
if aggregate != nil {
|
||||
accessLogs = aggregate.accessLogs()
|
||||
}
|
||||
if managed != nil && managed.TrafficReport != nil {
|
||||
return managed.TrafficReport, accessLogs, managed
|
||||
}
|
||||
if aggregate == nil {
|
||||
return nil, accessLogs, managed
|
||||
}
|
||||
fallbackManaged := aggregate.managedMetrics()
|
||||
return aggregate.report(), accessLogs, fallbackManaged
|
||||
}
|
||||
|
||||
func readAccessLogDelta(cfg *config.Config, stateStore *state.Store) *trafficAggregate {
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
logPath := managedAccessLogPath(cfg)
|
||||
file, err := os.Open(logPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
if snapshot.AccessLogOffset != 0 {
|
||||
snapshot.AccessLogOffset = 0
|
||||
_ = stateStore.Save(snapshot)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
offset := snapshot.AccessLogOffset
|
||||
if offset < 0 || offset > info.Size() {
|
||||
offset = 0
|
||||
}
|
||||
if _, err = file.Seek(offset, io.SeekStart); err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
reader := bufio.NewReader(file)
|
||||
currentOffset := offset
|
||||
aggregate := newTrafficAggregate()
|
||||
|
||||
for {
|
||||
line, readErr := reader.ReadBytes('\n')
|
||||
if len(line) > 0 {
|
||||
currentOffset += int64(len(line))
|
||||
aggregate.consume(line)
|
||||
}
|
||||
if errors.Is(readErr, io.EOF) {
|
||||
break
|
||||
}
|
||||
if readErr != nil {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
snapshot.AccessLogOffset = currentOffset
|
||||
_ = stateStore.Save(snapshot)
|
||||
|
||||
return aggregate
|
||||
}
|
||||
|
||||
func managedAccessLogPath(cfg *config.Config) string {
|
||||
if cfg == nil || strings.TrimSpace(cfg.RouteConfigPath) == "" {
|
||||
return ""
|
||||
}
|
||||
return filepath.Join(filepath.Dir(cfg.RouteConfigPath), "openflare_access.log")
|
||||
}
|
||||
|
||||
func newTrafficAggregate() *trafficAggregate {
|
||||
return &trafficAggregate{
|
||||
statusCodes: make(map[string]int64),
|
||||
topDomains: make(map[string]int64),
|
||||
visitors: make(map[string]struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (aggregate *trafficAggregate) consume(line []byte) {
|
||||
trimmed := strings.TrimSpace(string(line))
|
||||
if trimmed == "" {
|
||||
return
|
||||
}
|
||||
|
||||
record, ok := parseAccessLogRecord(trimmed)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
if aggregate.windowStartedAt.IsZero() || record.Timestamp.Before(aggregate.windowStartedAt) {
|
||||
aggregate.windowStartedAt = record.Timestamp
|
||||
}
|
||||
if aggregate.windowEndedAt.IsZero() || record.Timestamp.After(aggregate.windowEndedAt) {
|
||||
aggregate.windowEndedAt = record.Timestamp
|
||||
}
|
||||
|
||||
aggregate.requestCount++
|
||||
if record.Status >= 500 {
|
||||
aggregate.errorCount++
|
||||
}
|
||||
if record.Status > 0 {
|
||||
aggregate.statusCodes[strconv.Itoa(record.Status)]++
|
||||
}
|
||||
if record.RequestLength > 0 {
|
||||
aggregate.openrestyRxBytes += record.RequestLength
|
||||
}
|
||||
if record.BytesSent > 0 {
|
||||
aggregate.openrestyTxBytes += record.BytesSent
|
||||
}
|
||||
if host := strings.TrimSpace(record.Host); host != "" {
|
||||
aggregate.topDomains[host]++
|
||||
}
|
||||
if remoteAddr := strings.TrimSpace(record.RemoteAddr); remoteAddr != "" {
|
||||
aggregate.visitors[remoteAddr] = struct{}{}
|
||||
}
|
||||
aggregate.logs = append(aggregate.logs, protocol.NodeAccessLog{
|
||||
LoggedAtUnix: record.Timestamp.Unix(),
|
||||
RemoteAddr: strings.TrimSpace(record.RemoteAddr),
|
||||
Host: strings.TrimSpace(record.Host),
|
||||
Path: normalizeAccessLogPath(record.Path),
|
||||
StatusCode: record.Status,
|
||||
})
|
||||
}
|
||||
|
||||
type parsedAccessLogRecord struct {
|
||||
Timestamp time.Time
|
||||
Host string
|
||||
RemoteAddr string
|
||||
Path string
|
||||
Status int
|
||||
BytesSent int64
|
||||
RequestLength int64
|
||||
}
|
||||
|
||||
func parseAccessLogRecord(raw string) (parsedAccessLogRecord, bool) {
|
||||
record, ok := parseJSONAccessLogRecord(raw)
|
||||
if ok {
|
||||
return record, true
|
||||
}
|
||||
return parseCombinedAccessLogRecord(raw)
|
||||
}
|
||||
|
||||
func parseJSONAccessLogRecord(raw string) (parsedAccessLogRecord, bool) {
|
||||
var record accessLogRecord
|
||||
if err := json.Unmarshal([]byte(raw), &record); err != nil {
|
||||
return parsedAccessLogRecord{}, false
|
||||
}
|
||||
timestamp, err := parseAccessLogTime(record.Timestamp)
|
||||
if err != nil {
|
||||
return parsedAccessLogRecord{}, false
|
||||
}
|
||||
return parsedAccessLogRecord{
|
||||
Timestamp: timestamp,
|
||||
Host: strings.TrimSpace(record.Host),
|
||||
RemoteAddr: strings.TrimSpace(record.RemoteAddr),
|
||||
Path: normalizeAccessLogPath(record.Path),
|
||||
Status: record.Status,
|
||||
BytesSent: record.BytesSent,
|
||||
RequestLength: record.RequestLength,
|
||||
}, true
|
||||
}
|
||||
|
||||
func parseCombinedAccessLogRecord(raw string) (parsedAccessLogRecord, bool) {
|
||||
matches := combinedAccessLogPattern.FindStringSubmatch(raw)
|
||||
if len(matches) != 5 {
|
||||
return parsedAccessLogRecord{}, false
|
||||
}
|
||||
timestamp, err := parseAccessLogTime(matches[2])
|
||||
if err != nil {
|
||||
return parsedAccessLogRecord{}, false
|
||||
}
|
||||
status, err := strconv.Atoi(matches[4])
|
||||
if err != nil {
|
||||
return parsedAccessLogRecord{}, false
|
||||
}
|
||||
return parsedAccessLogRecord{
|
||||
Timestamp: timestamp,
|
||||
RemoteAddr: strings.TrimSpace(matches[1]),
|
||||
Path: normalizeAccessLogPath(matches[3]),
|
||||
Status: status,
|
||||
}, true
|
||||
}
|
||||
|
||||
func (aggregate *trafficAggregate) report() *protocol.NodeTrafficReport {
|
||||
if aggregate.requestCount == 0 || aggregate.windowStartedAt.IsZero() || aggregate.windowEndedAt.IsZero() {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &protocol.NodeTrafficReport{
|
||||
WindowStartedAtUnix: aggregate.windowStartedAt.Unix(),
|
||||
WindowEndedAtUnix: aggregate.windowEndedAt.Unix(),
|
||||
RequestCount: aggregate.requestCount,
|
||||
ErrorCount: aggregate.errorCount,
|
||||
UniqueVisitorCount: int64(len(aggregate.visitors)),
|
||||
StatusCodes: cloneTrafficCounts(aggregate.statusCodes, 0),
|
||||
TopDomains: topCounts(aggregate.topDomains, 8),
|
||||
SourceCountries: map[string]int64{},
|
||||
}
|
||||
}
|
||||
|
||||
func (aggregate *trafficAggregate) accessLogs() []protocol.NodeAccessLog {
|
||||
if aggregate == nil || len(aggregate.logs) == 0 {
|
||||
return []protocol.NodeAccessLog{}
|
||||
}
|
||||
return append([]protocol.NodeAccessLog(nil), aggregate.logs...)
|
||||
}
|
||||
|
||||
func (aggregate *trafficAggregate) managedMetrics() *managedOpenRestyMetrics {
|
||||
if aggregate == nil {
|
||||
return nil
|
||||
}
|
||||
report := aggregate.report()
|
||||
if report == nil && aggregate.openrestyRxBytes <= 0 && aggregate.openrestyTxBytes <= 0 {
|
||||
return nil
|
||||
}
|
||||
return &managedOpenRestyMetrics{
|
||||
TrafficReport: report,
|
||||
OpenrestyRxBytes: aggregate.openrestyRxBytes,
|
||||
OpenrestyTxBytes: aggregate.openrestyTxBytes,
|
||||
}
|
||||
}
|
||||
|
||||
func parseAccessLogTime(value string) (time.Time, error) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return time.Time{}, errors.New("empty access log time")
|
||||
}
|
||||
timestamp, err := time.Parse(time.RFC3339, trimmed)
|
||||
if err == nil {
|
||||
return timestamp, nil
|
||||
}
|
||||
return time.Parse("02/Jan/2006:15:04:05 -0700", trimmed)
|
||||
}
|
||||
|
||||
func cloneTrafficCounts(values map[string]int64, limit int) map[string]int64 {
|
||||
if len(values) == 0 {
|
||||
return map[string]int64{}
|
||||
}
|
||||
items := make([]trafficCountItem, 0, len(values))
|
||||
for key, value := range values {
|
||||
items = append(items, trafficCountItem{key: key, value: value})
|
||||
}
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
if items[i].value == items[j].value {
|
||||
return items[i].key < items[j].key
|
||||
}
|
||||
return items[i].value > items[j].value
|
||||
})
|
||||
if limit > 0 && len(items) > limit {
|
||||
items = items[:limit]
|
||||
}
|
||||
result := make(map[string]int64, len(items))
|
||||
for _, item := range items {
|
||||
result[item.key] = item.value
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
type trafficCountItem struct {
|
||||
key string
|
||||
value int64
|
||||
}
|
||||
|
||||
func normalizeAccessLogPath(value string) string {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
if strings.HasPrefix(trimmed, "http://") || strings.HasPrefix(trimmed, "https://") {
|
||||
return trimmed
|
||||
}
|
||||
if strings.HasPrefix(trimmed, "/") {
|
||||
return trimmed
|
||||
}
|
||||
return "/" + trimmed
|
||||
}
|
||||
|
||||
func topCounts(values map[string]int64, limit int) map[string]int64 {
|
||||
return cloneTrafficCounts(values, limit)
|
||||
}
|
||||
@@ -0,0 +1,162 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"openflare-agent/internal/config"
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
)
|
||||
|
||||
func TestBuildTrafficReportAggregatesManagedAccessLog(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
routeConfigPath := filepath.Join(tempDir, "conf.d", "openflare_routes.conf")
|
||||
if err := os.MkdirAll(filepath.Dir(routeConfigPath), 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll failed: %v", err)
|
||||
}
|
||||
logPath := filepath.Join(filepath.Dir(routeConfigPath), "openflare_access.log")
|
||||
content := []byte(
|
||||
"{\"ts\":\"2026-03-14T08:00:00Z\",\"host\":\"app.example.com\",\"path\":\"/\",\"remote_addr\":\"10.0.0.1\",\"status\":200}\n" +
|
||||
"{\"ts\":\"2026-03-14T08:00:05Z\",\"host\":\"app.example.com\",\"path\":\"/healthz\",\"remote_addr\":\"10.0.0.2\",\"status\":503}\n" +
|
||||
"{\"ts\":\"2026-03-14T08:00:08Z\",\"host\":\"api.example.com\",\"path\":\"/api\",\"remote_addr\":\"10.0.0.1\",\"status\":200}\n",
|
||||
)
|
||||
if err := os.WriteFile(logPath, content, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile failed: %v", err)
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
report := BuildTrafficReport(&config.Config{RouteConfigPath: routeConfigPath}, stateStore, nil)
|
||||
if report == nil {
|
||||
t.Fatal("expected traffic report")
|
||||
}
|
||||
if report.RequestCount != 3 || report.ErrorCount != 1 || report.UniqueVisitorCount != 2 {
|
||||
t.Fatalf("unexpected traffic report counters: %+v", report)
|
||||
}
|
||||
if report.StatusCodes["200"] != 2 || report.StatusCodes["503"] != 1 {
|
||||
t.Fatalf("unexpected status codes: %+v", report.StatusCodes)
|
||||
}
|
||||
if report.TopDomains["app.example.com"] != 2 || report.TopDomains["api.example.com"] != 1 {
|
||||
t.Fatalf("unexpected top domains: %+v", report.TopDomains)
|
||||
}
|
||||
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
t.Fatalf("Load failed: %v", err)
|
||||
}
|
||||
if snapshot.AccessLogOffset != int64(len(content)) {
|
||||
t.Fatalf("unexpected access log offset: %d", snapshot.AccessLogOffset)
|
||||
}
|
||||
|
||||
secondReport := BuildTrafficReport(&config.Config{RouteConfigPath: routeConfigPath}, stateStore, nil)
|
||||
if secondReport != nil {
|
||||
t.Fatalf("expected no report without appended lines, got %+v", secondReport)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTrafficReportResetsOffsetAfterTruncate(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
routeConfigPath := filepath.Join(tempDir, "conf.d", "openflare_routes.conf")
|
||||
if err := os.MkdirAll(filepath.Dir(routeConfigPath), 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll failed: %v", err)
|
||||
}
|
||||
logPath := filepath.Join(filepath.Dir(routeConfigPath), "openflare_access.log")
|
||||
if err := os.WriteFile(logPath, []byte("{\"ts\":\"2026-03-14T09:00:00Z\",\"host\":\"app.example.com\",\"path\":\"/\",\"remote_addr\":\"10.0.0.3\",\"status\":200}\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile failed: %v", err)
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
if err := stateStore.Save(&state.Snapshot{AccessLogOffset: 4096}); err != nil {
|
||||
t.Fatalf("Save failed: %v", err)
|
||||
}
|
||||
|
||||
report := BuildTrafficReport(&config.Config{RouteConfigPath: routeConfigPath}, stateStore, nil)
|
||||
if report == nil || report.RequestCount != 1 {
|
||||
t.Fatalf("expected one request after truncate reset, got %+v", report)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTrafficObservabilityReturnsAccessLogs(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
routeConfigPath := filepath.Join(tempDir, "conf.d", "openflare_routes.conf")
|
||||
if err := os.MkdirAll(filepath.Dir(routeConfigPath), 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll failed: %v", err)
|
||||
}
|
||||
logPath := filepath.Join(filepath.Dir(routeConfigPath), "openflare_access.log")
|
||||
content := []byte(
|
||||
"{\"ts\":\"2026-03-14T08:00:00Z\",\"host\":\"app.example.com\",\"path\":\"/login\",\"remote_addr\":\"10.0.0.1\",\"status\":200,\"request_length\":128,\"bytes_sent\":512}\n" +
|
||||
"{\"ts\":\"2026-03-14T08:00:05Z\",\"host\":\"api.example.com\",\"path\":\"/v1/ping\",\"remote_addr\":\"10.0.0.2\",\"status\":502,\"request_length\":64,\"bytes_sent\":256}\n",
|
||||
)
|
||||
if err := os.WriteFile(logPath, content, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile failed: %v", err)
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
report, accessLogs, fallbackMetrics := BuildTrafficObservability(&config.Config{RouteConfigPath: routeConfigPath}, stateStore, nil)
|
||||
if report == nil || report.RequestCount != 2 {
|
||||
t.Fatalf("expected traffic report, got %+v", report)
|
||||
}
|
||||
if len(accessLogs) != 2 {
|
||||
t.Fatalf("expected access logs, got %+v", accessLogs)
|
||||
}
|
||||
if fallbackMetrics == nil || fallbackMetrics.OpenrestyRxBytes != 192 || fallbackMetrics.OpenrestyTxBytes != 768 {
|
||||
t.Fatalf("expected fallback throughput metrics, got %+v", fallbackMetrics)
|
||||
}
|
||||
if accessLogs[0].Path != "/login" || accessLogs[1].Path != "/v1/ping" {
|
||||
t.Fatalf("unexpected access log paths: %+v", accessLogs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTrafficReportParsesCombinedAccessLog(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
routeConfigPath := filepath.Join(tempDir, "conf.d", "openflare_routes.conf")
|
||||
if err := os.MkdirAll(filepath.Dir(routeConfigPath), 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll failed: %v", err)
|
||||
}
|
||||
logPath := filepath.Join(filepath.Dir(routeConfigPath), "openflare_access.log")
|
||||
content := []byte(
|
||||
"10.0.0.1 - - [14/Mar/2026:08:00:00 +0000] \"GET / HTTP/1.1\" 200 123 \"-\" \"curl/8.0\"\n" +
|
||||
"10.0.0.2 - - [14/Mar/2026:08:00:05 +0000] \"GET /healthz HTTP/1.1\" 502 64 \"-\" \"curl/8.0\"\n" +
|
||||
"10.0.0.1 - - [14/Mar/2026:08:00:10 +0000] \"GET /api HTTP/1.1\" 200 256 \"-\" \"curl/8.0\"\n",
|
||||
)
|
||||
if err := os.WriteFile(logPath, content, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile failed: %v", err)
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
report := BuildTrafficReport(&config.Config{RouteConfigPath: routeConfigPath}, stateStore, nil)
|
||||
if report == nil {
|
||||
t.Fatal("expected traffic report from combined access log")
|
||||
}
|
||||
if report.RequestCount != 3 || report.ErrorCount != 1 || report.UniqueVisitorCount != 2 {
|
||||
t.Fatalf("unexpected combined log counters: %+v", report)
|
||||
}
|
||||
if report.StatusCodes["200"] != 2 || report.StatusCodes["502"] != 1 {
|
||||
t.Fatalf("unexpected combined log status codes: %+v", report.StatusCodes)
|
||||
}
|
||||
if len(report.TopDomains) != 0 {
|
||||
t.Fatalf("expected combined access log to omit top domains when host is unavailable, got %+v", report.TopDomains)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTrafficReportReturnsManagedWindowEvenWhenRequestCountZero(t *testing.T) {
|
||||
report := BuildTrafficReport(nil, nil, &managedOpenRestyMetrics{
|
||||
TrafficReport: &protocol.NodeTrafficReport{
|
||||
WindowStartedAtUnix: 1710403200,
|
||||
WindowEndedAtUnix: 1710403260,
|
||||
RequestCount: 0,
|
||||
ErrorCount: 0,
|
||||
UniqueVisitorCount: 0,
|
||||
StatusCodes: map[string]int64{},
|
||||
TopDomains: map[string]int64{},
|
||||
SourceCountries: map[string]int64{},
|
||||
},
|
||||
})
|
||||
if report == nil {
|
||||
t.Fatal("expected managed traffic report to be returned even when request count is zero")
|
||||
}
|
||||
if report.RequestCount != 0 || report.WindowStartedAtUnix != 1710403200 || report.WindowEndedAtUnix != 1710403260 {
|
||||
t.Fatalf("unexpected managed traffic report: %+v", report)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
package protocol
|
||||
|
||||
type APIResponse[T any] struct {
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
Data T `json:"data"`
|
||||
}
|
||||
|
||||
type HeartbeatAPIResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
Data any `json:"data"`
|
||||
AgentSettings *AgentSettings `json:"agent_settings,omitempty"`
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config,omitempty"`
|
||||
}
|
||||
|
||||
type HeartbeatResult struct {
|
||||
AgentSettings *AgentSettings
|
||||
ActiveConfig *ActiveConfigMeta
|
||||
}
|
||||
|
||||
type AgentSettings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
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 (
|
||||
OpenrestyStatusHealthy = "healthy"
|
||||
OpenrestyStatusUnhealthy = "unhealthy"
|
||||
OpenrestyStatusUnknown = "unknown"
|
||||
)
|
||||
|
||||
type NodePayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
AgentVersion string `json:"agent_version"`
|
||||
NginxVersion string `json:"nginx_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"`
|
||||
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
|
||||
BufferedObservability []BufferedObservabilityRecord `json:"buffered_observability,omitempty"`
|
||||
HealthEvents []NodeHealthEvent `json:"health_events"`
|
||||
}
|
||||
|
||||
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"`
|
||||
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"`
|
||||
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"`
|
||||
AgentToken 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"`
|
||||
MainConfig string `json:"main_config"`
|
||||
RouteConfig string `json:"route_config"`
|
||||
RenderedConfig string `json:"rendered_config"`
|
||||
SupportFiles []SupportFile `json:"support_files"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
}
|
||||
|
||||
type ActiveConfigMeta struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
type SupportFile struct {
|
||||
Path string `json:"path"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
@@ -0,0 +1,222 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strconv"
|
||||
"sync"
|
||||
|
||||
"openflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
const observabilityBufferWindowSeconds = 60
|
||||
|
||||
type ObservabilityBufferRecord struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
Snapshot *protocol.NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
TrafficReport *protocol.NodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []protocol.NodeAccessLog `json:"access_logs,omitempty"`
|
||||
QueuedAtUnix int64 `json:"queued_at_unix"`
|
||||
}
|
||||
|
||||
type ObservabilityBufferStore struct {
|
||||
path string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewObservabilityBufferStore(path string) *ObservabilityBufferStore {
|
||||
return &ObservabilityBufferStore{path: filepath.Clean(path)}
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) Upsert(record ObservabilityBufferRecord, retainAfterUnix int64) error {
|
||||
if s == nil || record.WindowStartedAtUnix <= 0 || (record.Snapshot == nil && record.TrafficReport == nil && len(record.AccessLogs) == 0) {
|
||||
return nil
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
records, err := s.loadUnlocked()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
records = pruneObservabilityBufferRecords(records, retainAfterUnix)
|
||||
replaced := false
|
||||
for index := range records {
|
||||
if records[index].WindowStartedAtUnix != record.WindowStartedAtUnix {
|
||||
continue
|
||||
}
|
||||
records[index] = mergeObservabilityBufferRecord(records[index], record)
|
||||
replaced = true
|
||||
break
|
||||
}
|
||||
if !replaced {
|
||||
records = append(records, record)
|
||||
}
|
||||
sort.Slice(records, func(i int, j int) bool {
|
||||
return records[i].WindowStartedAtUnix < records[j].WindowStartedAtUnix
|
||||
})
|
||||
return s.saveUnlocked(records)
|
||||
}
|
||||
|
||||
func mergeObservabilityBufferRecord(existing ObservabilityBufferRecord, incoming ObservabilityBufferRecord) ObservabilityBufferRecord {
|
||||
merged := existing
|
||||
if incoming.Snapshot != nil {
|
||||
merged.Snapshot = incoming.Snapshot
|
||||
}
|
||||
if incoming.TrafficReport != nil {
|
||||
merged.TrafficReport = incoming.TrafficReport
|
||||
}
|
||||
merged.AccessLogs = mergeAccessLogs(existing.AccessLogs, incoming.AccessLogs)
|
||||
if incoming.QueuedAtUnix > 0 {
|
||||
merged.QueuedAtUnix = incoming.QueuedAtUnix
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
func mergeAccessLogs(existing []protocol.NodeAccessLog, incoming []protocol.NodeAccessLog) []protocol.NodeAccessLog {
|
||||
if len(existing) == 0 && len(incoming) == 0 {
|
||||
return nil
|
||||
}
|
||||
merged := make([]protocol.NodeAccessLog, 0, len(existing)+len(incoming))
|
||||
seen := make(map[string]struct{}, len(existing)+len(incoming))
|
||||
appendIfNeeded := func(items []protocol.NodeAccessLog) {
|
||||
for _, item := range items {
|
||||
key := accessLogKey(item)
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
merged = append(merged, item)
|
||||
}
|
||||
}
|
||||
appendIfNeeded(existing)
|
||||
appendIfNeeded(incoming)
|
||||
sort.Slice(merged, func(i int, j int) bool {
|
||||
if merged[i].LoggedAtUnix == merged[j].LoggedAtUnix {
|
||||
return accessLogKey(merged[i]) < accessLogKey(merged[j])
|
||||
}
|
||||
return merged[i].LoggedAtUnix < merged[j].LoggedAtUnix
|
||||
})
|
||||
return merged
|
||||
}
|
||||
|
||||
func accessLogKey(item protocol.NodeAccessLog) string {
|
||||
return strconv.FormatInt(item.LoggedAtUnix, 10) + "|" + item.RemoteAddr + "|" + item.Host + "|" + item.Path + "|" + strconv.Itoa(item.StatusCode)
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) Replayable(currentWindowStartedAtUnix int64, retainAfterUnix int64) ([]ObservabilityBufferRecord, error) {
|
||||
if s == nil {
|
||||
return nil, nil
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
records, err := s.loadUnlocked()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
records = pruneObservabilityBufferRecords(records, retainAfterUnix)
|
||||
if err = s.saveUnlocked(records); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make([]ObservabilityBufferRecord, 0, len(records))
|
||||
for _, record := range records {
|
||||
if currentWindowStartedAtUnix > 0 && record.WindowStartedAtUnix >= currentWindowStartedAtUnix {
|
||||
continue
|
||||
}
|
||||
result = append(result, record)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) Ack(windowStartedAtUnix []int64, retainAfterUnix int64) error {
|
||||
if s == nil || len(windowStartedAtUnix) == 0 {
|
||||
return nil
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
records, err := s.loadUnlocked()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
acked := make(map[int64]struct{}, len(windowStartedAtUnix))
|
||||
for _, value := range windowStartedAtUnix {
|
||||
if value > 0 {
|
||||
acked[value] = struct{}{}
|
||||
}
|
||||
}
|
||||
filtered := make([]ObservabilityBufferRecord, 0, len(records))
|
||||
for _, record := range records {
|
||||
if _, ok := acked[record.WindowStartedAtUnix]; ok {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, record)
|
||||
}
|
||||
filtered = pruneObservabilityBufferRecords(filtered, retainAfterUnix)
|
||||
return s.saveUnlocked(filtered)
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) loadUnlocked() ([]ObservabilityBufferRecord, error) {
|
||||
data, err := os.ReadFile(s.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return []ObservabilityBufferRecord{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return []ObservabilityBufferRecord{}, nil
|
||||
}
|
||||
var records []ObservabilityBufferRecord
|
||||
if err = json.Unmarshal(data, &records); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return records, nil
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) saveUnlocked(records []ObservabilityBufferRecord) error {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(records, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(s.path, data, 0o644)
|
||||
}
|
||||
|
||||
func ObservabilityWindowStartedAt(snapshot *protocol.NodeMetricSnapshot, traffic *protocol.NodeTrafficReport) int64 {
|
||||
if traffic != nil && traffic.WindowStartedAtUnix > 0 {
|
||||
return traffic.WindowStartedAtUnix - (traffic.WindowStartedAtUnix % observabilityBufferWindowSeconds)
|
||||
}
|
||||
if snapshot == nil || snapshot.CapturedAtUnix <= 0 {
|
||||
return 0
|
||||
}
|
||||
return snapshot.CapturedAtUnix - (snapshot.CapturedAtUnix % observabilityBufferWindowSeconds)
|
||||
}
|
||||
|
||||
func pruneObservabilityBufferRecords(records []ObservabilityBufferRecord, retainAfterUnix int64) []ObservabilityBufferRecord {
|
||||
if len(records) == 0 {
|
||||
return []ObservabilityBufferRecord{}
|
||||
}
|
||||
filtered := make([]ObservabilityBufferRecord, 0, len(records))
|
||||
for _, record := range records {
|
||||
if record.WindowStartedAtUnix <= 0 {
|
||||
continue
|
||||
}
|
||||
if retainAfterUnix > 0 && record.WindowStartedAtUnix < retainAfterUnix {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, record)
|
||||
}
|
||||
sort.Slice(filtered, func(i int, j int) bool {
|
||||
return filtered[i].WindowStartedAtUnix < filtered[j].WindowStartedAtUnix
|
||||
})
|
||||
return filtered
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"openflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
func TestObservabilityBufferStoreUpsertReplayAndAck(t *testing.T) {
|
||||
store := NewObservabilityBufferStore(filepath.Join(t.TempDir(), "observability-buffer.json"))
|
||||
|
||||
if err := store.Upsert(ObservabilityBufferRecord{
|
||||
WindowStartedAtUnix: 1710403200,
|
||||
Snapshot: &protocol.NodeMetricSnapshot{CapturedAtUnix: 1710403205},
|
||||
TrafficReport: &protocol.NodeTrafficReport{WindowStartedAtUnix: 1710403200, WindowEndedAtUnix: 1710403260, RequestCount: 5},
|
||||
QueuedAtUnix: 1710403205,
|
||||
}, 1710403000); err != nil {
|
||||
t.Fatalf("first upsert failed: %v", err)
|
||||
}
|
||||
if err := store.Upsert(ObservabilityBufferRecord{
|
||||
WindowStartedAtUnix: 1710403200,
|
||||
Snapshot: &protocol.NodeMetricSnapshot{CapturedAtUnix: 1710403255},
|
||||
TrafficReport: &protocol.NodeTrafficReport{WindowStartedAtUnix: 1710403200, WindowEndedAtUnix: 1710403260, RequestCount: 12},
|
||||
QueuedAtUnix: 1710403255,
|
||||
}, 1710403000); err != nil {
|
||||
t.Fatalf("second upsert failed: %v", err)
|
||||
}
|
||||
if err := store.Upsert(ObservabilityBufferRecord{
|
||||
WindowStartedAtUnix: 1710403260,
|
||||
Snapshot: &protocol.NodeMetricSnapshot{CapturedAtUnix: 1710403265},
|
||||
TrafficReport: &protocol.NodeTrafficReport{WindowStartedAtUnix: 1710403260, WindowEndedAtUnix: 1710403320, RequestCount: 2},
|
||||
QueuedAtUnix: 1710403265,
|
||||
}, 1710403000); err != nil {
|
||||
t.Fatalf("third upsert failed: %v", err)
|
||||
}
|
||||
|
||||
records, err := store.Replayable(1710403260, 1710403000)
|
||||
if err != nil {
|
||||
t.Fatalf("Replayable failed: %v", err)
|
||||
}
|
||||
if len(records) != 1 {
|
||||
t.Fatalf("expected one replayable record before current window, got %d", len(records))
|
||||
}
|
||||
if records[0].TrafficReport == nil || records[0].TrafficReport.RequestCount != 12 {
|
||||
t.Fatalf("expected replayable record to keep latest upsert, got %+v", records[0])
|
||||
}
|
||||
|
||||
if err = store.Ack([]int64{1710403200}, 1710403000); err != nil {
|
||||
t.Fatalf("Ack failed: %v", err)
|
||||
}
|
||||
records, err = store.Replayable(0, 1710403000)
|
||||
if err != nil {
|
||||
t.Fatalf("Replayable after ack failed: %v", err)
|
||||
}
|
||||
if len(records) != 1 || records[0].WindowStartedAtUnix != 1710403260 {
|
||||
t.Fatalf("unexpected records after ack: %+v", records)
|
||||
}
|
||||
}
|
||||
|
||||
func TestObservabilityBufferStoreMergesAccessLogsWithinWindow(t *testing.T) {
|
||||
store := NewObservabilityBufferStore(filepath.Join(t.TempDir(), "observability-buffer.json"))
|
||||
|
||||
if err := store.Upsert(ObservabilityBufferRecord{
|
||||
WindowStartedAtUnix: 1710403200,
|
||||
AccessLogs: []protocol.NodeAccessLog{
|
||||
{LoggedAtUnix: 1710403201, RemoteAddr: "10.0.0.1", Host: "app.example.com", Path: "/a", StatusCode: 200},
|
||||
},
|
||||
}, 1710403000); err != nil {
|
||||
t.Fatalf("first upsert failed: %v", err)
|
||||
}
|
||||
if err := store.Upsert(ObservabilityBufferRecord{
|
||||
WindowStartedAtUnix: 1710403200,
|
||||
AccessLogs: []protocol.NodeAccessLog{
|
||||
{LoggedAtUnix: 1710403201, RemoteAddr: "10.0.0.1", Host: "app.example.com", Path: "/a", StatusCode: 200},
|
||||
{LoggedAtUnix: 1710403205, RemoteAddr: "10.0.0.2", Host: "app.example.com", Path: "/b", StatusCode: 502},
|
||||
},
|
||||
}, 1710403000); err != nil {
|
||||
t.Fatalf("second upsert failed: %v", err)
|
||||
}
|
||||
|
||||
records, err := store.Replayable(0, 1710403000)
|
||||
if err != nil {
|
||||
t.Fatalf("Replayable failed: %v", err)
|
||||
}
|
||||
if len(records) != 1 || len(records[0].AccessLogs) != 2 {
|
||||
t.Fatalf("expected merged access logs, got %+v", records)
|
||||
}
|
||||
}
|
||||
|
||||
func TestObservabilityWindowStartedAt(t *testing.T) {
|
||||
if value := ObservabilityWindowStartedAt(nil, &protocol.NodeTrafficReport{WindowStartedAtUnix: 1710403200}); value != 1710403200 {
|
||||
t.Fatalf("unexpected traffic window start: %d", value)
|
||||
}
|
||||
if value := ObservabilityWindowStartedAt(&protocol.NodeMetricSnapshot{CapturedAtUnix: 1710403259}, nil); value != 1710403200 {
|
||||
t.Fatalf("unexpected snapshot-derived window start: %d", value)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type Snapshot struct {
|
||||
NodeID string `json:"node_id"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
CurrentChecksum string `json:"current_checksum"`
|
||||
LastError string `json:"last_error"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
OpenrestyMessage string `json:"openresty_message"`
|
||||
LastProfileFingerprint string `json:"last_profile_fingerprint"`
|
||||
LastCPUStatTotal uint64 `json:"last_cpu_stat_total"`
|
||||
LastCPUStatIdle uint64 `json:"last_cpu_stat_idle"`
|
||||
LastMetricAtUnix int64 `json:"last_metric_at_unix"`
|
||||
AccessLogOffset int64 `json:"access_log_offset"`
|
||||
}
|
||||
|
||||
type Store struct {
|
||||
path string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewStore(path string) *Store {
|
||||
return &Store{path: filepath.Clean(path)}
|
||||
}
|
||||
|
||||
func (s *Store) Load() (*Snapshot, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.loadUnlocked()
|
||||
}
|
||||
|
||||
func (s *Store) EnsureNodeID() (string, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
snapshot, err := s.loadUnlocked()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if snapshot.NodeID != "" {
|
||||
return snapshot.NodeID, nil
|
||||
}
|
||||
snapshot.NodeID, err = newNodeID()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err = s.saveUnlocked(snapshot); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return snapshot.NodeID, nil
|
||||
}
|
||||
|
||||
func (s *Store) Save(snapshot *Snapshot) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.saveUnlocked(snapshot)
|
||||
}
|
||||
|
||||
func (s *Store) loadUnlocked() (*Snapshot, error) {
|
||||
data, err := os.ReadFile(s.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return &Snapshot{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
snapshot := &Snapshot{}
|
||||
if len(data) == 0 {
|
||||
return snapshot, nil
|
||||
}
|
||||
if err = json.Unmarshal(data, snapshot); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
|
||||
func (s *Store) saveUnlocked(snapshot *Snapshot) error {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(snapshot, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(s.path, data, 0o644)
|
||||
}
|
||||
|
||||
func newNodeID() (string, error) {
|
||||
buf := make([]byte, 8)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "node-" + hex.EncodeToString(buf), nil
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEnsureNodeIDPersists(t *testing.T) {
|
||||
store := NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
nodeID1, err := store.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
nodeID2, err := store.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID second call failed: %v", err)
|
||||
}
|
||||
if nodeID1 == "" || nodeID1 != nodeID2 {
|
||||
t.Fatal("expected node id to persist across calls")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,209 @@
|
||||
package sync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
)
|
||||
|
||||
const (
|
||||
ApplyResultSuccess = "success"
|
||||
ApplyResultFailed = "failed"
|
||||
)
|
||||
|
||||
type ConfigClient interface {
|
||||
GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error)
|
||||
ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error
|
||||
}
|
||||
|
||||
type NginxManager interface {
|
||||
Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) error
|
||||
EnsureRuntime(ctx context.Context, recreate bool) error
|
||||
CurrentChecksum() (string, error)
|
||||
}
|
||||
|
||||
type Service struct {
|
||||
client ConfigClient
|
||||
nginxManager NginxManager
|
||||
stateStore *state.Store
|
||||
}
|
||||
|
||||
func New(client ConfigClient, nginxManager NginxManager, stateStore *state.Store) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
nginxManager: nginxManager,
|
||||
stateStore: stateStore,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error {
|
||||
return s.sync(ctx, false, target)
|
||||
}
|
||||
|
||||
func (s *Service) SyncOnStartup(ctx context.Context, target *protocol.ActiveConfigMeta) error {
|
||||
return s.sync(ctx, true, target)
|
||||
}
|
||||
|
||||
func (s *Service) sync(ctx context.Context, startup bool, target *protocol.ActiveConfigMeta) error {
|
||||
mode := "periodic"
|
||||
if startup {
|
||||
mode = "startup"
|
||||
}
|
||||
snapshot, err := s.stateStore.Load()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
currentChecksum, err := s.nginxManager.CurrentChecksum()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if target != nil {
|
||||
target.Version = strings.TrimSpace(target.Version)
|
||||
target.Checksum = strings.TrimSpace(target.Checksum)
|
||||
}
|
||||
|
||||
if target == nil || target.Version == "" || target.Checksum == "" {
|
||||
if !startup {
|
||||
slog.Debug("skipping sync because heartbeat returned no active config summary", "mode", mode)
|
||||
return nil
|
||||
}
|
||||
slog.Debug("sync startup fallback: active config summary unavailable, fetching active config directly")
|
||||
config, fetchErr := s.client.GetActiveConfig(ctx)
|
||||
if fetchErr != nil {
|
||||
slog.Error("fetch active config failed", "mode", mode, "error", fetchErr)
|
||||
return fetchErr
|
||||
}
|
||||
target = &protocol.ActiveConfigMeta{
|
||||
Version: config.Version,
|
||||
Checksum: config.Checksum,
|
||||
}
|
||||
return s.applyIfNeeded(ctx, mode, startup, snapshot, currentChecksum, target, config)
|
||||
}
|
||||
|
||||
if currentChecksum == target.Checksum {
|
||||
slog.Debug("local openresty config already up to date", "mode", mode, "version", target.Version)
|
||||
if startup {
|
||||
slog.Debug("ensuring openresty runtime on startup", "version", target.Version)
|
||||
if err = s.nginxManager.EnsureRuntime(ctx, true); err != nil {
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
|
||||
snapshot.OpenrestyMessage = err.Error()
|
||||
_ = s.stateStore.Save(snapshot)
|
||||
return err
|
||||
}
|
||||
slog.Debug("openresty runtime ensured on startup", "version", target.Version)
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
|
||||
snapshot.OpenrestyMessage = ""
|
||||
}
|
||||
snapshot.CurrentVersion = target.Version
|
||||
snapshot.CurrentChecksum = target.Checksum
|
||||
snapshot.LastError = ""
|
||||
slog.Debug("sync finished without changes", "mode", mode, "version", target.Version)
|
||||
return s.stateStore.Save(snapshot)
|
||||
}
|
||||
if snapshot.CurrentVersion == target.Version && snapshot.CurrentChecksum == target.Checksum && !startup {
|
||||
slog.Debug("skipping config fetch because state already records target version/checksum", "version", target.Version, "checksum", target.Checksum)
|
||||
return nil
|
||||
}
|
||||
|
||||
config, err := s.client.GetActiveConfig(ctx)
|
||||
if err != nil {
|
||||
slog.Error("fetch active config failed", "mode", mode, "error", err)
|
||||
return err
|
||||
}
|
||||
return s.applyIfNeeded(ctx, mode, startup, snapshot, currentChecksum, target, config)
|
||||
}
|
||||
|
||||
func (s *Service) applyIfNeeded(ctx context.Context, mode string, startup bool, snapshot *state.Snapshot, currentChecksum string, target *protocol.ActiveConfigMeta, config *protocol.ActiveConfigResponse) error {
|
||||
if currentChecksum == config.Checksum {
|
||||
slog.Debug("local openresty config already up to date", "mode", mode, "version", config.Version)
|
||||
if startup {
|
||||
slog.Debug("ensuring openresty runtime on startup", "version", config.Version)
|
||||
if err := s.nginxManager.EnsureRuntime(ctx, true); err != nil {
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
|
||||
snapshot.OpenrestyMessage = err.Error()
|
||||
_ = s.stateStore.Save(snapshot)
|
||||
return err
|
||||
}
|
||||
slog.Debug("openresty runtime ensured on startup", "version", config.Version)
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
|
||||
snapshot.OpenrestyMessage = ""
|
||||
}
|
||||
snapshot.CurrentVersion = config.Version
|
||||
snapshot.CurrentChecksum = config.Checksum
|
||||
snapshot.LastError = ""
|
||||
slog.Debug("sync finished without changes", "mode", mode, "version", config.Version)
|
||||
return s.stateStore.Save(snapshot)
|
||||
}
|
||||
if target != nil && (target.Version != config.Version || target.Checksum != config.Checksum) {
|
||||
slog.Warn("active config changed between heartbeat and fetch", "heartbeat_version", target.Version, "heartbeat_checksum", target.Checksum, "fetched_version", config.Version, "fetched_checksum", config.Checksum)
|
||||
}
|
||||
if snapshot.CurrentVersion == config.Version && snapshot.CurrentChecksum == config.Checksum && !startup {
|
||||
slog.Debug("skipping apply because state already records target version/checksum", "version", config.Version, "checksum", config.Checksum)
|
||||
return nil
|
||||
}
|
||||
routeConfig := config.RouteConfig
|
||||
if routeConfig == "" {
|
||||
routeConfig = config.RenderedConfig
|
||||
}
|
||||
mainConfigChecksum := checksumString(config.MainConfig)
|
||||
routeConfigChecksum := checksumString(routeConfig)
|
||||
slog.Info("applying new openresty config", "mode", mode, "from_version", snapshot.CurrentVersion, "to_version", config.Version, "old_checksum", currentChecksum, "new_checksum", config.Checksum)
|
||||
if err := s.nginxManager.Apply(ctx, config.MainConfig, routeConfig, config.SupportFiles); err != nil {
|
||||
slog.Error("apply openresty config failed", "mode", mode, "version", config.Version, "error", err)
|
||||
snapshot.LastError = err.Error()
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
|
||||
snapshot.OpenrestyMessage = err.Error()
|
||||
_ = s.stateStore.Save(snapshot)
|
||||
reportErr := s.client.ReportApplyLog(ctx, protocol.ApplyLogPayload{
|
||||
NodeID: snapshot.NodeID,
|
||||
Version: config.Version,
|
||||
Result: ApplyResultFailed,
|
||||
Message: err.Error(),
|
||||
Checksum: config.Checksum,
|
||||
MainConfigChecksum: mainConfigChecksum,
|
||||
RouteConfigChecksum: routeConfigChecksum,
|
||||
SupportFileCount: len(config.SupportFiles),
|
||||
})
|
||||
if reportErr != nil {
|
||||
slog.Error("report failed apply log failed", "version", config.Version, "error", reportErr)
|
||||
return reportErr
|
||||
}
|
||||
slog.Warn("failed apply log reported", "version", config.Version)
|
||||
return err
|
||||
}
|
||||
slog.Info("openresty config applied successfully", "mode", mode, "version", config.Version)
|
||||
snapshot.CurrentVersion = config.Version
|
||||
snapshot.CurrentChecksum = config.Checksum
|
||||
snapshot.LastError = ""
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
|
||||
snapshot.OpenrestyMessage = ""
|
||||
if err := s.stateStore.Save(snapshot); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.client.ReportApplyLog(ctx, protocol.ApplyLogPayload{
|
||||
NodeID: snapshot.NodeID,
|
||||
Version: config.Version,
|
||||
Result: ApplyResultSuccess,
|
||||
Message: "apply success",
|
||||
Checksum: config.Checksum,
|
||||
MainConfigChecksum: mainConfigChecksum,
|
||||
RouteConfigChecksum: routeConfigChecksum,
|
||||
SupportFileCount: len(config.SupportFiles),
|
||||
}); err != nil {
|
||||
slog.Error("report successful apply log failed", "version", config.Version, "error", err)
|
||||
return err
|
||||
}
|
||||
slog.Debug("successful apply log reported", "version", config.Version)
|
||||
return nil
|
||||
}
|
||||
|
||||
func checksumString(content string) string {
|
||||
sum := sha256.Sum256([]byte(content))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
package sync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"openflare-agent/internal/nginx"
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
)
|
||||
|
||||
type fakeExecutor struct {
|
||||
testErr error
|
||||
reloadErr error
|
||||
}
|
||||
|
||||
type fakeClient struct {
|
||||
config protocol.ActiveConfigResponse
|
||||
reports []protocol.ApplyLogPayload
|
||||
fetchCalls int
|
||||
}
|
||||
|
||||
type fakeManager struct {
|
||||
applyErr error
|
||||
currentChecksum string
|
||||
currentChecksumErr error
|
||||
ensureErr error
|
||||
ensureCalls []bool
|
||||
applyMainContents []string
|
||||
applyRouteContents []string
|
||||
applyFiles [][]protocol.SupportFile
|
||||
}
|
||||
|
||||
func (f *fakeExecutor) Test(ctx context.Context) error {
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeExecutor) Reload(ctx context.Context) error {
|
||||
return f.reloadErr
|
||||
}
|
||||
|
||||
func (f *fakeExecutor) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeExecutor) CheckHealth(ctx context.Context) error {
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeExecutor) Restart(ctx context.Context) error {
|
||||
return f.reloadErr
|
||||
}
|
||||
|
||||
func (f *fakeClient) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) {
|
||||
f.fetchCalls++
|
||||
return &f.config, nil
|
||||
}
|
||||
|
||||
func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
|
||||
f.reports = append(f.reports, payload)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *fakeManager) Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) error {
|
||||
m.applyMainContents = append(m.applyMainContents, mainConfig)
|
||||
m.applyRouteContents = append(m.applyRouteContents, routeConfig)
|
||||
m.applyFiles = append(m.applyFiles, append([]protocol.SupportFile(nil), supportFiles...))
|
||||
return m.applyErr
|
||||
}
|
||||
|
||||
func (m *fakeManager) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
m.ensureCalls = append(m.ensureCalls, recreate)
|
||||
return m.ensureErr
|
||||
}
|
||||
|
||||
func (m *fakeManager) CurrentChecksum() (string, error) {
|
||||
return m.currentChecksum, m.currentChecksumErr
|
||||
}
|
||||
|
||||
func TestSyncOnceSuccess(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
config: protocol.ActiveConfigResponse{
|
||||
Version: "20260309-001",
|
||||
Checksum: "checksum-1",
|
||||
MainConfig: "worker_processes auto;",
|
||||
RouteConfig: "server { listen 80; }",
|
||||
RenderedConfig: "server { listen 80; }",
|
||||
SupportFiles: []protocol.SupportFile{{Path: "1.crt", Content: "cert"}},
|
||||
CreatedAt: time.Now().Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
nodeID, err := stateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
snapshot, _ := stateStore.Load()
|
||||
snapshot.NodeID = nodeID
|
||||
if err = stateStore.Save(snapshot); err != nil {
|
||||
t.Fatalf("failed to save initial state: %v", err)
|
||||
}
|
||||
|
||||
routePath := filepath.Join(t.TempDir(), "routes.conf")
|
||||
service := New(client, &nginx.Manager{
|
||||
MainConfigPath: filepath.Join(filepath.Dir(routePath), "nginx.conf"),
|
||||
RouteConfigPath: routePath,
|
||||
Executor: &fakeExecutor{},
|
||||
}, stateStore)
|
||||
|
||||
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
|
||||
Version: client.config.Version,
|
||||
Checksum: client.config.Checksum,
|
||||
}); err != nil {
|
||||
t.Fatalf("SyncOnce failed: %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(routePath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read route config: %v", err)
|
||||
}
|
||||
if string(data) != "server { listen 80; }" {
|
||||
t.Fatal("expected rendered config to be written to route file")
|
||||
}
|
||||
mainData, err := os.ReadFile(filepath.Join(filepath.Dir(routePath), "nginx.conf"))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read main config: %v", err)
|
||||
}
|
||||
if string(mainData) != "worker_processes auto;" {
|
||||
t.Fatal("expected main config to be written")
|
||||
}
|
||||
snapshot, err = stateStore.Load()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load state: %v", err)
|
||||
}
|
||||
if snapshot.CurrentVersion != "20260309-001" || snapshot.CurrentChecksum != "checksum-1" {
|
||||
t.Fatal("expected state store to persist current version and checksum")
|
||||
}
|
||||
if len(client.reports) != 1 || client.reports[0].Result != ApplyResultSuccess {
|
||||
t.Fatal("expected successful apply report to be sent")
|
||||
}
|
||||
if client.reports[0].Checksum != "checksum-1" {
|
||||
t.Fatalf("expected config checksum to be reported, got %q", client.reports[0].Checksum)
|
||||
}
|
||||
if client.reports[0].MainConfigChecksum == "" || client.reports[0].RouteConfigChecksum == "" {
|
||||
t.Fatal("expected main and route config checksums to be reported")
|
||||
}
|
||||
if client.reports[0].SupportFileCount != 1 {
|
||||
t.Fatalf("expected support file count to be reported, got %d", client.reports[0].SupportFileCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncOnceRollbackOnNginxFailure(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
config: protocol.ActiveConfigResponse{
|
||||
Version: "20260309-002",
|
||||
Checksum: "checksum-2",
|
||||
MainConfig: "worker_processes 2;",
|
||||
RouteConfig: "server { listen 81; }",
|
||||
RenderedConfig: "server { listen 81; }",
|
||||
SupportFiles: []protocol.SupportFile{{Path: "1.crt", Content: "cert"}},
|
||||
CreatedAt: time.Now().Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
|
||||
tempDir := t.TempDir()
|
||||
mainPath := filepath.Join(tempDir, "nginx.conf")
|
||||
routePath := filepath.Join(tempDir, "routes.conf")
|
||||
if err := os.WriteFile(mainPath, []byte("worker_processes auto;"), 0o644); err != nil {
|
||||
t.Fatalf("failed to seed main file: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(routePath, []byte("server { listen 80; }"), 0o644); err != nil {
|
||||
t.Fatalf("failed to seed route file: %v", err)
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
nodeID, err := stateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
if err = stateStore.Save(&state.Snapshot{
|
||||
NodeID: nodeID,
|
||||
CurrentVersion: "20260309-001",
|
||||
CurrentChecksum: "checksum-1",
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to seed state: %v", err)
|
||||
}
|
||||
|
||||
service := New(client, &nginx.Manager{
|
||||
MainConfigPath: mainPath,
|
||||
RouteConfigPath: routePath,
|
||||
Executor: &fakeExecutor{
|
||||
testErr: context.DeadlineExceeded,
|
||||
},
|
||||
}, stateStore)
|
||||
|
||||
err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
|
||||
Version: client.config.Version,
|
||||
Checksum: client.config.Checksum,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected SyncOnce to fail when nginx test fails")
|
||||
}
|
||||
|
||||
data, readErr := os.ReadFile(routePath)
|
||||
if readErr != nil {
|
||||
t.Fatalf("failed to read route file after rollback: %v", readErr)
|
||||
}
|
||||
if string(data) != "server { listen 80; }" {
|
||||
t.Fatal("expected original route config to be restored after rollback")
|
||||
}
|
||||
mainData, readErr := os.ReadFile(mainPath)
|
||||
if readErr != nil {
|
||||
t.Fatalf("failed to read main file after rollback: %v", readErr)
|
||||
}
|
||||
if string(mainData) != "worker_processes auto;" {
|
||||
t.Fatal("expected original main config to be restored after rollback")
|
||||
}
|
||||
snapshot, loadErr := stateStore.Load()
|
||||
if loadErr != nil {
|
||||
t.Fatalf("failed to load state: %v", loadErr)
|
||||
}
|
||||
if snapshot.CurrentVersion != "20260309-001" {
|
||||
t.Fatal("expected failed sync not to overwrite current version")
|
||||
}
|
||||
if len(client.reports) != 1 || client.reports[0].Result != ApplyResultFailed {
|
||||
t.Fatal("expected failed apply report to be sent")
|
||||
}
|
||||
if client.reports[0].Checksum != "checksum-2" {
|
||||
t.Fatalf("expected failed report to retain target checksum, got %q", client.reports[0].Checksum)
|
||||
}
|
||||
if client.reports[0].MainConfigChecksum == "" || client.reports[0].RouteConfigChecksum == "" {
|
||||
t.Fatal("expected failed report to include main and route config checksums")
|
||||
}
|
||||
if client.reports[0].SupportFileCount != 1 {
|
||||
t.Fatalf("expected failed report to include support file count, got %d", client.reports[0].SupportFileCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncOnStartupRecreatesRuntimeWhenChecksumMatches(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
config: protocol.ActiveConfigResponse{
|
||||
Version: "20260309-003",
|
||||
Checksum: "checksum-3",
|
||||
MainConfig: "worker_processes auto;",
|
||||
RouteConfig: "server { listen 82; }",
|
||||
RenderedConfig: "server { listen 82; }",
|
||||
SupportFiles: []protocol.SupportFile{{Path: "1.crt", Content: "cert"}},
|
||||
CreatedAt: time.Now().Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
nodeID, err := stateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
if err = stateStore.Save(&state.Snapshot{NodeID: nodeID}); err != nil {
|
||||
t.Fatalf("failed to seed state: %v", err)
|
||||
}
|
||||
|
||||
manager := &fakeManager{currentChecksum: "checksum-3"}
|
||||
service := New(client, manager, stateStore)
|
||||
if err = service.SyncOnStartup(context.Background(), &protocol.ActiveConfigMeta{
|
||||
Version: client.config.Version,
|
||||
Checksum: client.config.Checksum,
|
||||
}); err != nil {
|
||||
t.Fatalf("SyncOnStartup failed: %v", err)
|
||||
}
|
||||
if len(manager.ensureCalls) != 1 || !manager.ensureCalls[0] {
|
||||
t.Fatal("expected startup sync to recreate runtime")
|
||||
}
|
||||
if len(client.reports) != 0 {
|
||||
t.Fatal("expected no apply report when checksum already matches")
|
||||
}
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load state: %v", err)
|
||||
}
|
||||
if snapshot.CurrentChecksum != "checksum-3" || snapshot.CurrentVersion != "20260309-003" {
|
||||
t.Fatal("expected snapshot to be refreshed from active config")
|
||||
}
|
||||
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusHealthy || snapshot.OpenrestyMessage != "" {
|
||||
t.Fatal("expected startup sync to mark openresty healthy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncOnStartupRecordsRuntimeFailure(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
config: protocol.ActiveConfigResponse{
|
||||
Version: "20260309-004",
|
||||
Checksum: "checksum-4",
|
||||
MainConfig: "worker_processes 4;",
|
||||
RouteConfig: "server { listen 83; }",
|
||||
RenderedConfig: "server { listen 83; }",
|
||||
CreatedAt: time.Now().Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
nodeID, err := stateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
if err = stateStore.Save(&state.Snapshot{NodeID: nodeID}); err != nil {
|
||||
t.Fatalf("failed to seed state: %v", err)
|
||||
}
|
||||
|
||||
manager := &fakeManager{
|
||||
currentChecksum: "checksum-4",
|
||||
ensureErr: context.DeadlineExceeded,
|
||||
}
|
||||
service := New(client, manager, stateStore)
|
||||
if err = service.SyncOnStartup(context.Background(), &protocol.ActiveConfigMeta{
|
||||
Version: client.config.Version,
|
||||
Checksum: client.config.Checksum,
|
||||
}); err == nil {
|
||||
t.Fatal("expected SyncOnStartup to fail when runtime recreation fails")
|
||||
}
|
||||
snapshot, err := stateStore.Load()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load state: %v", err)
|
||||
}
|
||||
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusUnhealthy {
|
||||
t.Fatalf("expected unhealthy openresty status, got %q", snapshot.OpenrestyStatus)
|
||||
}
|
||||
if snapshot.OpenrestyMessage == "" {
|
||||
t.Fatal("expected runtime error message to be recorded")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncOnceSkipsFetchWhenHeartbeatChecksumMatches(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
config: protocol.ActiveConfigResponse{
|
||||
Version: "20260309-005",
|
||||
Checksum: "checksum-5",
|
||||
MainConfig: "worker_processes auto;",
|
||||
RouteConfig: "server { listen 84; }",
|
||||
RenderedConfig: "server { listen 84; }",
|
||||
CreatedAt: time.Now().Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
nodeID, err := stateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
if err = stateStore.Save(&state.Snapshot{
|
||||
NodeID: nodeID,
|
||||
CurrentVersion: client.config.Version,
|
||||
CurrentChecksum: client.config.Checksum,
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to seed state: %v", err)
|
||||
}
|
||||
|
||||
manager := &fakeManager{currentChecksum: client.config.Checksum}
|
||||
service := New(client, manager, stateStore)
|
||||
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
|
||||
Version: client.config.Version,
|
||||
Checksum: client.config.Checksum,
|
||||
}); err != nil {
|
||||
t.Fatalf("SyncOnce failed: %v", err)
|
||||
}
|
||||
if client.fetchCalls != 0 {
|
||||
t.Fatalf("expected no active config fetch when heartbeat checksum matches, got %d", client.fetchCalls)
|
||||
}
|
||||
if len(client.reports) != 0 {
|
||||
t.Fatal("expected no apply log when no config change is needed")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
//go:build !windows
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
os.Remove(backupPath)
|
||||
if err := os.Rename(execPath, backupPath); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("backup current binary: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpPath, execPath); err != nil {
|
||||
os.Rename(backupPath, execPath)
|
||||
return fmt.Errorf("replace binary: %w", err)
|
||||
}
|
||||
os.Remove(backupPath)
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil {
|
||||
return fmt.Errorf("exec restart: %w", err)
|
||||
}
|
||||
return fmt.Errorf("unreachable after exec")
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
//go:build windows
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
backupPath := execPath + ".bak"
|
||||
scriptPath := execPath + ".update.cmd"
|
||||
script := fmt.Sprintf(`@echo off
|
||||
setlocal
|
||||
:waitloop
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 (
|
||||
ping 127.0.0.1 -n 2 >nul
|
||||
goto waitloop
|
||||
)
|
||||
move /Y "%s" "%s" >nul 2>nul
|
||||
if errorlevel 1 exit /b 1
|
||||
start "" %s
|
||||
del /Q "%s" >nul 2>nul
|
||||
del /Q "%%~f0" >nul 2>nul
|
||||
`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath)
|
||||
if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("write restart script: %w", err)
|
||||
}
|
||||
cmd := exec.Command("cmd", "/C", "start", "", scriptPath)
|
||||
if err := cmd.Start(); err != nil {
|
||||
os.Remove(scriptPath)
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("schedule restart: %w", err)
|
||||
}
|
||||
os.Exit(0)
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildWindowsCommandLine(execPath string, args []string) string {
|
||||
parts := []string{quoteWindowsArg(execPath)}
|
||||
for _, arg := range args {
|
||||
parts = append(parts, quoteWindowsArg(arg))
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func quoteWindowsArg(value string) string {
|
||||
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
|
||||
}
|
||||
@@ -0,0 +1,397 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"openflare-agent/internal/agent"
|
||||
"openflare-agent/internal/config"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
httpClient *http.Client
|
||||
lastCheckKey string
|
||||
}
|
||||
|
||||
func New() *Service {
|
||||
return &Service{
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
type githubRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
Assets []githubAsset `json:"assets"`
|
||||
}
|
||||
|
||||
type githubAsset struct {
|
||||
Name string `json:"name"`
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
}
|
||||
|
||||
func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options agent.UpdateOptions) error {
|
||||
release, err := s.getRelease(ctx, repo, options)
|
||||
if err != nil {
|
||||
return fmt.Errorf("check latest release: %w", err)
|
||||
}
|
||||
if release == nil || release.TagName == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
remoteVersion := normalizeVersion(release.TagName)
|
||||
localVersion := normalizeVersion(config.AgentVersion)
|
||||
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)
|
||||
|
||||
var downloadURL string
|
||||
for _, asset := range release.Assets {
|
||||
if asset.Name == assetName {
|
||||
downloadURL = asset.BrowserDownloadURL
|
||||
break
|
||||
}
|
||||
}
|
||||
if downloadURL == "" {
|
||||
s.lastCheckKey = checkKey
|
||||
return fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName)
|
||||
}
|
||||
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get executable path: %w", err)
|
||||
}
|
||||
if err = s.downloadAndRestart(ctx, downloadURL, execPath); err != nil {
|
||||
return fmt.Errorf("download and restart: %w", err)
|
||||
}
|
||||
s.lastCheckKey = checkKey
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) getRelease(ctx context.Context, repo string, options agent.UpdateOptions) (*githubRelease, error) {
|
||||
tagName := strings.TrimSpace(options.TagName)
|
||||
if tagName != "" {
|
||||
return s.getReleaseByTag(ctx, repo, tagName)
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(options.Channel), "preview") {
|
||||
return s.getLatestPreviewRelease(ctx, repo)
|
||||
}
|
||||
return s.getLatestStableRelease(ctx, repo)
|
||||
}
|
||||
|
||||
func (s *Service) getLatestStableRelease(ctx context.Context, repo string) (*githubRelease, error) {
|
||||
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/latest", repo)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return nil, nil
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("github api returned %s", resp.Status)
|
||||
}
|
||||
|
||||
return decodeRelease(resp.Body)
|
||||
}
|
||||
|
||||
func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*githubRelease, error) {
|
||||
url := fmt.Sprintf("https://api.github.com/repos/%s/releases?per_page=20", repo)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("github api returned %s", resp.Status)
|
||||
}
|
||||
|
||||
var releases []githubRelease
|
||||
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, release := range releases {
|
||||
if release.Draft || !release.Prerelease {
|
||||
continue
|
||||
}
|
||||
releaseCopy := release
|
||||
return &releaseCopy, nil
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (s *Service) getReleaseByTag(ctx context.Context, repo string, tag string) (*githubRelease, error) {
|
||||
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/tags/%s", repo, strings.TrimSpace(tag))
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return nil, nil
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("github api returned %s", resp.Status)
|
||||
}
|
||||
|
||||
return decodeRelease(resp.Body)
|
||||
}
|
||||
|
||||
func decodeRelease(reader io.Reader) (*githubRelease, error) {
|
||||
var release githubRelease
|
||||
if err := json.NewDecoder(reader).Decode(&release); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
func (s *Service) downloadAndRestart(ctx context.Context, url string, targetPath 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("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, 0o755)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err = io.Copy(tmpFile, resp.Body); err != nil {
|
||||
tmpFile.Close()
|
||||
os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
tmpFile.Close()
|
||||
|
||||
slog.Info("agent binary updated, restarting")
|
||||
return replaceAndRestart(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
|
||||
}
|
||||
|
||||
type versionInfo struct {
|
||||
valid bool
|
||||
isDev bool
|
||||
numbers []int
|
||||
prerelease []string
|
||||
}
|
||||
|
||||
func parseVersionInfo(version string) versionInfo {
|
||||
normalized := normalizeVersion(version)
|
||||
if normalized == "" || strings.EqualFold(normalized, "dev") {
|
||||
return versionInfo{isDev: strings.EqualFold(normalized, "dev")}
|
||||
}
|
||||
base := normalized
|
||||
prerelease := ""
|
||||
if index := strings.IndexRune(normalized, '-'); index >= 0 {
|
||||
base = normalized[:index]
|
||||
prerelease = normalized[index+1:]
|
||||
}
|
||||
segments := strings.Split(base, ".")
|
||||
parts := make([]int, 0, len(segments))
|
||||
for _, segment := range segments {
|
||||
segment = strings.TrimSpace(segment)
|
||||
if segment == "" {
|
||||
parts = append(parts, 0)
|
||||
continue
|
||||
}
|
||||
numeric := strings.Builder{}
|
||||
for _, r := range segment {
|
||||
if r < '0' || r > '9' {
|
||||
break
|
||||
}
|
||||
numeric.WriteRune(r)
|
||||
}
|
||||
if numeric.Len() == 0 {
|
||||
return versionInfo{}
|
||||
}
|
||||
value, err := strconv.Atoi(numeric.String())
|
||||
if err != nil {
|
||||
return versionInfo{}
|
||||
}
|
||||
parts = append(parts, value)
|
||||
}
|
||||
info := versionInfo{valid: len(parts) > 0, numbers: parts}
|
||||
if prerelease != "" {
|
||||
info.prerelease = splitPrereleaseIdentifiers(prerelease)
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
func splitPrereleaseIdentifiers(value string) []string {
|
||||
parts := strings.FieldsFunc(strings.TrimSpace(value), func(r rune) bool {
|
||||
return r == '.' || r == '-'
|
||||
})
|
||||
filtered := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
if part != "" {
|
||||
filtered = append(filtered, part)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func compareVersions(local string, remote string) int {
|
||||
left := parseVersionInfo(local)
|
||||
right := parseVersionInfo(remote)
|
||||
if left.isDev {
|
||||
if right.valid {
|
||||
return -1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
if !left.valid || !right.valid {
|
||||
return 0
|
||||
}
|
||||
|
||||
maxLen := len(left.numbers)
|
||||
if len(right.numbers) > maxLen {
|
||||
maxLen = len(right.numbers)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
leftValue := 0
|
||||
rightValue := 0
|
||||
if index < len(left.numbers) {
|
||||
leftValue = left.numbers[index]
|
||||
}
|
||||
if index < len(right.numbers) {
|
||||
rightValue = right.numbers[index]
|
||||
}
|
||||
if leftValue < rightValue {
|
||||
return -1
|
||||
}
|
||||
if leftValue > rightValue {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
if len(left.prerelease) == 0 && len(right.prerelease) == 0 {
|
||||
return 0
|
||||
}
|
||||
if len(left.prerelease) == 0 {
|
||||
return 1
|
||||
}
|
||||
if len(right.prerelease) == 0 {
|
||||
return -1
|
||||
}
|
||||
maxLen = len(left.prerelease)
|
||||
if len(right.prerelease) > maxLen {
|
||||
maxLen = len(right.prerelease)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
if index >= len(left.prerelease) {
|
||||
return -1
|
||||
}
|
||||
if index >= len(right.prerelease) {
|
||||
return 1
|
||||
}
|
||||
leftPart := left.prerelease[index]
|
||||
rightPart := right.prerelease[index]
|
||||
leftNumber, leftErr := strconv.Atoi(leftPart)
|
||||
rightNumber, rightErr := strconv.Atoi(rightPart)
|
||||
switch {
|
||||
case leftErr == nil && rightErr == nil:
|
||||
if leftNumber < rightNumber {
|
||||
return -1
|
||||
}
|
||||
if leftNumber > rightNumber {
|
||||
return 1
|
||||
}
|
||||
case leftErr == nil && rightErr != nil:
|
||||
return -1
|
||||
case leftErr != nil && rightErr == nil:
|
||||
return 1
|
||||
default:
|
||||
if leftPart < rightPart {
|
||||
return -1
|
||||
}
|
||||
if leftPart > rightPart {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"openflare-agent/internal/agent"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type roundTripFunc func(req *http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
func TestGetLatestPreviewRelease(t *testing.T) {
|
||||
service := &Service{
|
||||
httpClient: &http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if req.URL.String() != "https://api.github.com/repos/Rain-kl/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 TestIsNewerSupportsPrerelease(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
local string
|
||||
remote string
|
||||
expected bool
|
||||
}{
|
||||
{name: "stable newer than prerelease", local: "1.2.3-rc.1", remote: "1.2.3", expected: true},
|
||||
{name: "same stable not newer", local: "1.2.3", remote: "1.2.3-rc.1", expected: false},
|
||||
{name: "higher prerelease sequence", local: "1.2.3-rc.1", remote: "1.2.3-rc.2", expected: true},
|
||||
{name: "higher minor", local: "1.2.3", remote: "1.3.0-rc.1", expected: true},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
if actual := isNewer(testCase.local, testCase.remote); actual != testCase.expected {
|
||||
t.Fatalf("unexpected compare result: local=%s remote=%s actual=%v expected=%v", testCase.local, testCase.remote, actual, testCase.expected)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user