[优化] 改名

This commit is contained in:
ryan
2026-03-15 16:04:54 +08:00
parent d68773c554
commit 32d90ba641
304 changed files with 629 additions and 969 deletions
+1
View File
@@ -0,0 +1 @@
data
+9
View File
@@ -0,0 +1,9 @@
{
"server_url": "http://127.0.0.1:3000",
"agent_token": "89d0efbf7bcc8fa53fe48889dce4d045",
"data_dir": "./data",
"openresty_container_name": "openflare-openresty",
"openresty_docker_image": "openresty/openresty:alpine",
"heartbeat_interval": 10000,
"request_timeout": 10000
}
+113
View File
@@ -0,0 +1,113 @@
package main
import (
"context"
"flag"
"log/slog"
"os"
"os/signal"
"syscall"
"openflare-agent/internal/agent"
"openflare-agent/internal/config"
"openflare-agent/internal/heartbeat"
"openflare-agent/internal/httpclient"
"openflare-agent/internal/logging"
"openflare-agent/internal/nginx"
"openflare-agent/internal/state"
syncservice "openflare-agent/internal/sync"
"openflare-agent/internal/updater"
)
func main() {
logging.Setup()
configPath := flag.String("config", "./agent.json", "agent config path")
flag.Parse()
cfg, err := config.Load(*configPath)
if err != nil {
slog.Error("load agent config failed", "error", err)
os.Exit(1)
}
cfg.NginxVersion = nginx.DetectVersion(
context.Background(),
nginx.ExecutorOptions{
NginxPath: cfg.OpenrestyPath,
DockerBinary: cfg.DockerBinary,
ContainerName: cfg.OpenrestyContainerName,
Image: cfg.OpenrestyDockerImage,
MainConfigPath: cfg.MainConfigPath,
RouteConfigPath: cfg.RouteConfigPath,
CertDir: cfg.CertDir,
NginxCertDir: cfg.OpenrestyCertDir,
LuaDir: cfg.LuaDir,
NginxLuaDir: cfg.OpenrestyLuaDir,
OpenrestyObservabilityPort: cfg.OpenrestyObservabilityPort,
},
)
slog.Info("agent config loaded",
"server", cfg.ServerURL,
"node", cfg.NodeName,
"ip", cfg.NodeIP,
"heartbeat_interval", cfg.HeartbeatInterval,
"route_config", cfg.RouteConfigPath,
"cert_dir", cfg.CertDir,
"lua_dir", cfg.LuaDir,
)
client := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
stateStore := state.NewStore(cfg.StatePath)
observabilityBuffer := state.NewObservabilityBufferStore(cfg.ObservabilityBufferPath)
runtimeRouteConfigPath := cfg.RouteConfigPath
if cfg.OpenrestyPath == "" {
runtimeRouteConfigPath = nginx.DockerRouteConfigPath
}
runtimeManager := &nginx.Manager{
MainConfigPath: cfg.MainConfigPath,
RouteConfigPath: cfg.RouteConfigPath,
RuntimeRouteConfigPath: runtimeRouteConfigPath,
CertDir: cfg.CertDir,
NginxCertDir: cfg.OpenrestyCertDir,
LuaDir: cfg.LuaDir,
NginxLuaDir: cfg.OpenrestyLuaDir,
OpenrestyObservabilityListen: nginx.ObservabilityListenAddress(cfg.OpenrestyPath, cfg.OpenrestyObservabilityPort),
OpenrestyObservabilityPort: cfg.OpenrestyObservabilityPort,
Executor: nginx.NewExecutor(nginx.ExecutorOptions{
NginxPath: cfg.OpenrestyPath,
DockerBinary: cfg.DockerBinary,
ContainerName: cfg.OpenrestyContainerName,
Image: cfg.OpenrestyDockerImage,
MainConfigPath: cfg.MainConfigPath,
RouteConfigPath: cfg.RouteConfigPath,
CertDir: cfg.CertDir,
NginxCertDir: cfg.OpenrestyCertDir,
LuaDir: cfg.LuaDir,
NginxLuaDir: cfg.OpenrestyLuaDir,
OpenrestyObservabilityPort: cfg.OpenrestyObservabilityPort,
}),
}
if err = runtimeManager.EnsureLuaAssets(); err != nil {
slog.Error("ensure managed lua assets failed", "error", err)
os.Exit(1)
}
runner := &agent.Runner{
Config: cfg,
StateStore: stateStore,
ObservabilityBuffer: observabilityBuffer,
HeartbeatService: heartbeat.New(client),
SyncService: syncservice.New(client, runtimeManager, stateStore),
Updater: updater.New(),
RuntimeManager: runtimeManager,
}
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
slog.Info("agent process started")
if err = runner.Run(ctx); err != nil && err != context.Canceled {
slog.Error("agent process exited with error", "error", err)
os.Exit(1)
}
slog.Info("agent process stopped")
}
+3
View File
@@ -0,0 +1,3 @@
module openflare-agent
go 1.23.0
+404
View File
@@ -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")
}
}
+327
View File
@@ -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
}
+133
View File
@@ -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, " ")
}
+803
View File
@@ -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)
}
}
+103
View File
@@ -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")
}
}
+209
View File
@@ -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, `"`, `""`) + `"`
}
+397
View File
@@ -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)
}
})
}
}