[优化] go 引用调整

This commit is contained in:
ryan
2026-06-06 10:26:20 +08:00
parent ee1110b752
commit 3cfefb4367
552 changed files with 1642 additions and 2185 deletions
+1
View File
@@ -0,0 +1 @@
data
+40
View File
@@ -0,0 +1,40 @@
# syntax=docker/dockerfile:1.7
ARG VERSION=dev
FROM golang:1.25-alpine AS builder
ARG VERSION
ARG TARGETOS=linux
ARG TARGETARCH
ENV CGO_ENABLED=0 \
GOOS=${TARGETOS} \
GOARCH=${TARGETARCH}
WORKDIR /build
COPY go.mod go.sum ./
RUN --mount=type=cache,target=/go/pkg/mod \
go mod download
COPY openflare-server ./openflare-server
COPY openflare-agent ./openflare-agent
RUN --mount=type=cache,target=/go/pkg/mod \
--mount=type=cache,target=/root/.cache/go-build \
go build -trimpath -ldflags "-s -w -X 'github.com/rain-kl/openflare/openflare-agent/internal/config.Version=$VERSION'" -o /build/bin/openflare-agent ./openflare-agent/cmd/agent
FROM openresty/openresty:alpine
RUN apk add --no-cache ca-certificates tzdata perl libmaxminddb \
&& ln -sf /usr/lib/libmaxminddb.so.0 /usr/lib/libmaxminddb.so \
&& opm get anjia0532/lua-resty-maxminddb \
&& mkdir -p /etc/openflare /data
ENV OPENFLARE_OPENRESTY_PATH=openresty \
OPENFLARE_DATA_DIR=/data
COPY --from=builder /build/bin/openflare-agent /usr/local/bin/openflare-agent
EXPOSE 80 443 18081
ENTRYPOINT ["/usr/local/bin/openflare-agent"]
CMD ["-config", "/etc/openflare/agent.json"]
+8
View File
@@ -0,0 +1,8 @@
{
"server_url": "http://127.0.0.1:3000",
"agent_token": "373956188ddead1df6dd7c86cd330b73",
"data_dir": "./data",
"openresty_path": "openresty",
"heartbeat_interval": 10000,
"request_timeout": 10000
}
+121
View File
@@ -0,0 +1,121 @@
package main
import (
"context"
"flag"
"log/slog"
"os"
"os/signal"
"syscall"
"github.com/rain-kl/openflare/openflare-agent/internal/agent"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
"github.com/rain-kl/openflare/openflare-agent/internal/geoipupdate"
"github.com/rain-kl/openflare/openflare-agent/internal/heartbeat"
"github.com/rain-kl/openflare/openflare-agent/internal/httpclient"
"github.com/rain-kl/openflare/openflare-agent/internal/logging"
"github.com/rain-kl/openflare/openflare-agent/internal/nginx"
"github.com/rain-kl/openflare/openflare-agent/internal/state"
syncservice "github.com/rain-kl/openflare/openflare-agent/internal/sync"
"github.com/rain-kl/openflare/openflare-agent/internal/updater"
"github.com/rain-kl/openflare/openflare-agent/internal/wsclient"
)
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.ExtVersion = nginx.DetectVersion(
context.Background(),
nginx.ExecutorOptions{
NginxPath: cfg.OpenrestyPath,
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,
"access_log", cfg.AccessLogPath,
"cert_dir", cfg.CertDir,
"lua_dir", cfg.LuaDir,
"runtime_config_dir", cfg.RuntimeConfigDir,
"mmdb_path", cfg.MMDBPath,
)
client := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
wsClient := wsclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
stateStore := state.NewStore(cfg.StatePath)
observabilityBuffer := state.NewObservabilityBufferStore(cfg.ObservabilityBufferPath)
runtimeManager := &nginx.Manager{
MainConfigPath: cfg.MainConfigPath,
RouteConfigPath: cfg.RouteConfigPath,
AccessLogPath: cfg.AccessLogPath,
CertDir: cfg.CertDir,
NginxCertDir: cfg.OpenrestyCertDir,
LuaDir: cfg.LuaDir,
NginxLuaDir: cfg.OpenrestyLuaDir,
RuntimeConfigDir: cfg.RuntimeConfigDir,
PagesDir: cfg.PagesDir,
OpenrestyObservabilityListen: nginx.ObservabilityListenAddress(cfg.OpenrestyObservabilityPort),
OpenrestyObservabilityPort: cfg.OpenrestyObservabilityPort,
OpenrestyResolverDirective: "",
Executor: nginx.NewExecutor(nginx.ExecutorOptions{
NginxPath: cfg.OpenrestyPath,
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)
}
syncService := syncservice.New(client, runtimeManager, stateStore)
syncService.SetPagesDir(cfg.PagesDir)
runner := &agent.Runner{
Config: cfg,
StateStore: stateStore,
ObservabilityBuffer: observabilityBuffer,
HeartbeatService: heartbeat.New(client),
SyncService: syncService,
Updater: updater.New(),
RuntimeManager: runtimeManager,
WebSocketService: wsClient,
}
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
geoIPUpdater := &geoipupdate.Updater{
MMDBPath: cfg.MMDBPath,
DownloadURL: cfg.MMDBDownloadURL,
UpdateInterval: cfg.MMDBUpdateInterval.Duration(),
}
go geoIPUpdater.Run(ctx)
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")
}
+699
View File
@@ -0,0 +1,699 @@
package agent
import (
"context"
"encoding/json"
"errors"
"log/slog"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
"github.com/rain-kl/openflare/openflare-agent/internal/observability"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
"github.com/rain-kl/openflare/openflare-agent/internal/state"
"github.com/rain-kl/openflare/openflare-agent/internal/wsclient"
)
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
ForceSyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error
WAFIPGroupChecksums() (map[string]string, error)
ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error
}
type Updater interface {
CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error
}
type RuntimeManager interface {
CheckHealth(ctx context.Context) error
Restart(ctx context.Context) error
}
type WebSocketService interface {
Connect(ctx context.Context) (protocol.WebSocketConnection, error)
SetToken(token string)
URL() string
}
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
WebSocketService WebSocketService
autoUpdate bool
updateNow bool
updateRepo string
updateChan string
updateTag string
restartOpenrestyNow bool
websocketUpgradeEnabled 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.hasAccessToken() {
if _, hbErr := r.performHeartbeatCycle(ctx, nodeID, true); hbErr != nil {
slog.Error("agent startup heartbeat failed", "error", hbErr)
}
} 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()
var wsDone <-chan error
wsBackoff := newWebSocketBackoff()
nextWSAttempt := time.Now()
tryStartWebSocket := func() {
if wsDone != nil || !r.shouldUseWebSocket() || time.Now().Before(nextWSAttempt) {
return
}
done, startErr := r.startWebSocket(ctx, nodeID)
if startErr != nil {
delay := wsBackoff.Next()
nextWSAttempt = time.Now().Add(delay)
slog.Debug("agent ws upgrade failed; falling back to http heartbeat",
"enabled", r.websocketUpgradeEnabled,
"url", r.websocketURL(),
"retry_after", delay,
"error", startErr,
)
return
}
wsBackoff.Reset()
wsDone = done
slog.Debug("agent switched to websocket mode", "url", r.websocketURL())
}
tryStartWebSocket()
for {
select {
case <-ctx.Done():
slog.Info("agent runner shutting down", "error", ctx.Err())
return ctx.Err()
case wsErr := <-wsDone:
wsDone = nil
delay := wsBackoff.Next()
nextWSAttempt = time.Now().Add(delay)
slog.Debug("agent ws disconnected; resuming http heartbeat", "retry_after", delay, "error", wsErr)
if r.hasAccessToken() {
if _, hbErr := r.performHeartbeatCycle(ctx, nodeID, false); hbErr != nil {
slog.Error("agent heartbeat after ws disconnect failed", "error", hbErr)
}
}
case <-heartbeatTicker.C:
if wsDone != nil {
continue
}
if !r.hasAccessToken() {
if err = r.tryRegister(ctx, &nodeID); err != nil {
slog.Error("agent discovery register failed", "error", err)
}
continue
}
if changed, hbErr := r.performHeartbeatCycle(ctx, nodeID, false); hbErr != nil {
slog.Error("agent heartbeat failed", "error", hbErr)
} else {
if changed {
heartbeatTicker.Reset(r.Config.HeartbeatInterval.Duration())
}
tryStartWebSocket()
}
}
}
}
func (r *Runner) performHeartbeatCycle(ctx context.Context, nodeID string, startup bool) (bool, error) {
r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(nodeID)
heartbeatResult, err := r.HeartbeatService.Heartbeat(ctx, payload)
if err != nil {
return false, err
}
r.ackObservabilityWindows(ackWindows)
if heartbeatResult == nil {
heartbeatResult = &protocol.HeartbeatResult{}
}
mode := "periodic"
if startup {
mode = "startup"
}
slog.Debug("agent heartbeat succeeded", "mode", mode, "node_id", nodeID)
changed := r.applySettings(heartbeatResult.AgentSettings)
r.applyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups)
if startup {
if err = r.SyncService.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil {
r.recordSyncError(err)
slog.Error("agent startup sync failed", "error", err)
} else {
slog.Debug("agent startup sync completed")
}
} else if err = r.SyncService.SyncOnce(ctx, heartbeatResult.ActiveConfig); err != nil {
r.recordSyncError(err)
slog.Error("agent sync failed", "error", err)
}
r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx)
return changed, nil
}
func (r *Runner) shouldUseWebSocket() bool {
enabled := r.WebSocketService != nil && r.websocketUpgradeEnabled && r.hasAccessToken()
slog.Debug("agent ws upgrade eligibility checked", "enabled", enabled, "server_enabled", r.websocketUpgradeEnabled, "url", r.websocketURL())
return enabled
}
func (r *Runner) websocketURL() string {
if r.WebSocketService == nil {
return ""
}
return r.WebSocketService.URL()
}
func (r *Runner) startWebSocket(ctx context.Context, nodeID string) (<-chan error, error) {
if r.WebSocketService == nil {
return nil, errors.New("websocket service is not configured")
}
conn, err := r.WebSocketService.Connect(ctx)
if err != nil {
return nil, err
}
done := make(chan error, 1)
go func() {
defer func() {
_ = conn.Close()
}()
done <- r.runWebSocket(ctx, nodeID, conn)
}()
return done, nil
}
type agentWSHandler struct {
runner *Runner
conn protocol.WebSocketConnection
nodeID string
statusTicker *time.Ticker
}
func (h *agentWSHandler) OnConnect(ctx context.Context) error {
return h.runner.sendWebSocketStatus(ctx, h.nodeID, h.conn)
}
func (h *agentWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
var payloadBytes []byte
if msg.Payload != nil {
payloadBytes = []byte(msg.Payload)
}
protoMsg := protocol.WSMessage{
Type: msg.Type,
Payload: payloadBytes,
}
changed, err := h.runner.handleWebSocketMessage(ctx, protoMsg, h.conn)
if err != nil {
return err
}
if changed {
h.statusTicker.Reset(h.runner.Config.HeartbeatInterval.Duration())
}
return nil
}
func (h *agentWSHandler) OnClose(err error) {
slog.Error("agent ws receive failed", "error", err)
}
func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
slog.Debug("agent ws connected", "url", conn.URL(), "node_id", nodeID)
statusTicker := time.NewTicker(r.Config.HeartbeatInterval.Duration())
defer statusTicker.Stop()
childCtx, cancel := context.WithCancel(ctx)
defer cancel()
// Start status ticker sender in background
go func() {
for {
select {
case <-childCtx.Done():
return
case <-statusTicker.C:
if err := r.sendWebSocketStatus(childCtx, nodeID, conn); err != nil {
slog.Error("agent ws send status failed", "error", err)
_ = conn.Close()
return
}
}
}
}()
wsConn, ok := conn.(*wsclient.Connection)
if !ok {
return errors.New("invalid websocket connection type")
}
return wsConn.RunReceiveLoop(childCtx, &agentWSHandler{
runner: r,
conn: conn,
nodeID: nodeID,
statusTicker: statusTicker,
})
}
func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
r.refreshOpenrestyHealth(ctx)
payload, ackWindows := r.prepareHeartbeatPayload(nodeID)
if err := conn.SendStatus(payload); err != nil {
return err
}
r.ackObservabilityWindows(ackWindows)
return nil
}
func (r *Runner) handleWebSocketMessage(ctx context.Context, message protocol.WSMessage, conn protocol.WebSocketConnection) (bool, error) {
switch message.Type {
case protocol.WSMessageTypeSettings:
var settings protocol.AgentSettings
if err := json.Unmarshal(message.Payload, &settings); err != nil {
slog.Debug("agent ws settings decode failed", "error", err)
return false, nil
}
changed := r.applySettings(&settings)
r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx)
if !r.websocketUpgradeEnabled {
slog.Debug("agent ws disabled by server settings; falling back to http heartbeat")
return changed, errors.New("websocket upgrade disabled by server")
}
return changed, nil
case protocol.WSMessageTypeActiveConfig:
var target protocol.ActiveConfigMeta
if err := json.Unmarshal(message.Payload, &target); err != nil {
slog.Debug("agent ws active config decode failed", "error", err)
return false, nil
}
slog.Debug("agent ws active config received", "version", target.Version, "checksum", target.Checksum, "trigger_sync", true)
if err := r.SyncService.SyncOnce(ctx, &target); err != nil {
r.recordSyncError(err)
slog.Error("agent ws triggered sync failed", "version", target.Version, "error", err)
}
return false, nil
case protocol.WSMessageTypeForceSyncConfig:
var target protocol.ActiveConfigMeta
if err := json.Unmarshal(message.Payload, &target); err != nil {
slog.Debug("agent ws force sync config decode failed", "error", err)
return false, nil
}
slog.Debug("agent ws force sync config received", "version", target.Version, "checksum", target.Checksum, "trigger_sync", true)
if err := r.SyncService.ForceSyncOnce(ctx, &target); err != nil {
r.recordSyncError(err)
slog.Error("agent ws triggered force sync failed", "version", target.Version, "error", err)
}
return false, nil
case protocol.WSMessageTypeWAFIPGroups:
var groups []protocol.WAFIPGroup
if err := json.Unmarshal(message.Payload, &groups); err != nil {
slog.Debug("agent ws waf ip groups decode failed", "error", err)
return false, nil
}
r.applyWAFIPGroups(ctx, groups)
return false, nil
case protocol.WSMessageTypePing:
slog.Debug("agent ws ping received")
return false, conn.SendPong()
case protocol.WSMessageTypePong:
slog.Debug("agent ws pong received")
return false, nil
default:
slog.Debug("agent ws unsupported message type", "type", message.Type)
return false, nil
}
}
type webSocketBackoff struct {
delays []time.Duration
index int
}
func newWebSocketBackoff() *webSocketBackoff {
return &webSocketBackoff{
delays: []time.Duration{
time.Second,
2 * time.Second,
5 * time.Second,
10 * time.Second,
30 * time.Second,
},
}
}
func (backoff *webSocketBackoff) Next() time.Duration {
if backoff == nil || len(backoff.delays) == 0 {
return 30 * time.Second
}
if backoff.index >= len(backoff.delays) {
return backoff.delays[len(backoff.delays)-1]
}
delay := backoff.delays[backoff.index]
backoff.index++
return delay
}
func (backoff *webSocketBackoff) Reset() {
if backoff != nil {
backoff.index = 0
}
}
func (r *Runner) hasAccessToken() bool {
return strings.TrimSpace(r.Config.AccessToken) != ""
}
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
}
}
if settings.WebsocketUpgradeEnabled != r.websocketUpgradeEnabled {
slog.Debug("agent websocket upgrade setting updated", "from", r.websocketUpgradeEnabled, "to", settings.WebsocketUpgradeEnabled)
}
r.websocketUpgradeEnabled = settings.WebsocketUpgradeEnabled
r.autoUpdate = settings.AutoUpdate
r.updateNow = settings.UpdateNow
r.updateRepo = strings.TrimSpace(settings.UpdateRepo)
r.updateChan = strings.TrimSpace(settings.UpdateChannel)
r.updateTag = strings.TrimSpace(settings.UpdateTag)
r.restartOpenrestyNow = settings.RestartOpenrestyNow
return changed
}
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.AccessToken) == "" || 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.AccessToken = response.AccessToken
r.Config.DiscoveryToken = ""
if err = r.Config.Save(); err != nil {
return err
}
r.HeartbeatService.SetToken(response.AccessToken)
if r.WebSocketService != nil {
r.WebSocketService.SetToken(response.AccessToken)
}
*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)
r.applyWAFIPGroups(ctx, heartbeatResult.WAFIPGroups)
if err = r.SyncService.SyncOnStartup(ctx, heartbeatResult.ActiveConfig); err != nil {
r.recordSyncError(err)
slog.Error("agent post-register startup sync failed", "error", err)
} else {
slog.Debug("agent post-register startup sync completed")
}
r.tryRestartOpenresty(ctx)
r.tryAutoUpdate(ctx)
return nil
}
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 {
if strings.Contains(err.Error(), "openresty config not exists") {
return
}
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)
openrestyObservation := observability.BuildOpenrestyObservation(managedOpenRestyMetrics)
healthEvents := observability.BuildHealthEvents(snapshot)
payload := protocol.NodePayload{
NodeID: nodeID,
Name: r.Config.NodeName,
IP: r.Config.NodeIP,
Version: r.Config.Version,
ExtVersion: r.Config.ExtVersion,
CurrentVersion: snapshot.CurrentVersion,
LastError: snapshot.LastError,
OpenrestyStatus: openrestyStatus,
OpenrestyMessage: snapshot.OpenrestyMessage,
Profile: profile,
Snapshot: metricSnapshot,
OpenrestyObservation: openrestyObservation,
TrafficReport: trafficReport,
AccessLogs: accessLogs,
HealthEvents: healthEvents,
}
if r.SyncService != nil {
checksums, err := r.SyncService.WAFIPGroupChecksums()
if err != nil {
slog.Debug("load local waf ip group checksums failed", "error", err)
} else if len(checksums) > 0 {
payload.WAFIPGroupChecksums = checksums
}
}
return payload
}
func (r *Runner) applyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) {
if len(groups) == 0 || r.SyncService == nil {
return
}
if err := r.SyncService.ApplyWAFIPGroups(ctx, groups); err != nil {
r.recordSyncError(err)
slog.Error("agent apply waf ip groups failed", "error", err)
}
}
func (r *Runner) prepareHeartbeatPayload(nodeID string) (protocol.NodePayload, []int64) {
payload := r.nodePayload(nodeID)
if r.ObservabilityBuffer == nil || (payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0) {
return payload, nil
}
now := time.Now().UTC()
retainAfterUnix := now.Add(-time.Duration(r.Config.ObservabilityReplayMinutes) * time.Minute).Unix()
windowStartedAtUnix := state.ObservabilityWindowStartedAt(payload.Snapshot, payload.OpenrestyObservation, payload.TrafficReport)
if windowStartedAtUnix <= 0 {
return payload, nil
}
record := state.ObservabilityBufferRecord{
WindowStartedAtUnix: windowStartedAtUnix,
Snapshot: payload.Snapshot,
OpenrestyObservation: payload.OpenrestyObservation,
TrafficReport: payload.TrafficReport,
AccessLogs: payload.AccessLogs,
QueuedAtUnix: now.Unix(),
}
if err := r.ObservabilityBuffer.Upsert(record, retainAfterUnix); err != nil {
slog.Error("upsert observability buffer failed", "error", err)
return payload, nil
}
records, err := r.ObservabilityBuffer.Replayable(windowStartedAtUnix, retainAfterUnix)
if err != nil {
slog.Error("load replayable observability buffer failed", "error", err)
return payload, []int64{windowStartedAtUnix}
}
ackWindows := make([]int64, 0, len(records)+1)
buffered := make([]protocol.BufferedObservabilityRecord, 0, len(records))
for _, item := range records {
if item.WindowStartedAtUnix <= 0 {
continue
}
buffered = append(buffered, protocol.BufferedObservabilityRecord{
WindowStartedAtUnix: item.WindowStartedAtUnix,
Snapshot: item.Snapshot,
OpenrestyObservation: item.OpenrestyObservation,
TrafficReport: item.TrafficReport,
AccessLogs: item.AccessLogs,
})
ackWindows = append(ackWindows, item.WindowStartedAtUnix)
}
payload.BufferedObservability = buffered
ackWindows = append(ackWindows, windowStartedAtUnix)
return payload, ackWindows
}
func (r *Runner) ackObservabilityWindows(windowStartedAtUnix []int64) {
if r.ObservabilityBuffer == nil || len(windowStartedAtUnix) == 0 {
return
}
retainAfterUnix := time.Now().UTC().Add(-time.Duration(r.Config.ObservabilityReplayMinutes) * time.Minute).Unix()
if err := r.ObservabilityBuffer.Ack(windowStartedAtUnix, retainAfterUnix); err != nil {
slog.Error("ack observability buffer failed", "error", err)
}
}
@@ -0,0 +1,640 @@
package agent
import (
"context"
"encoding/json"
"errors"
"os"
"path/filepath"
"sync"
"testing"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
"github.com/rain-kl/openflare/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
lastTarget *protocol.ActiveConfigMeta
onSyncOnceCall func(int)
wafChecksums map[string]string
wafGroups []protocol.WAFIPGroup
}
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++
if target != nil {
copied := *target
f.lastTarget = &copied
}
callIndex := f.syncOnceCalls
callback := f.onSyncOnceCall
f.mu.Unlock()
if callback != nil {
callback(callIndex)
}
return f.syncOnceErr
}
func (f *fakeSyncService) ForceSyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error {
f.mu.Lock()
f.syncOnceCalls++
if target != nil {
copied := *target
f.lastTarget = &copied
}
callIndex := f.syncOnceCalls
callback := f.onSyncOnceCall
f.mu.Unlock()
if callback != nil {
callback(callIndex)
}
return f.syncOnceErr
}
func (f *fakeSyncService) WAFIPGroupChecksums() (map[string]string, error) {
if f.wafChecksums == nil {
return map[string]string{}, nil
}
return f.wafChecksums, nil
}
func (f *fakeSyncService) ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error {
f.mu.Lock()
defer f.mu.Unlock()
f.wafGroups = append(f.wafGroups, groups...)
return nil
}
type fakeWebSocketConnection struct {
pongCalls int
}
func (f *fakeWebSocketConnection) URL() string {
return "ws://127.0.0.1/api/agent/ws"
}
func (f *fakeWebSocketConnection) SendStatus(payload protocol.NodePayload) error {
return nil
}
func (f *fakeWebSocketConnection) SendPong() error {
f.pongCalls++
return nil
}
func (f *fakeWebSocketConnection) Receive() (protocol.WSMessage, error) {
return protocol.WSMessage{}, errors.New("not implemented")
}
func (f *fakeWebSocketConnection) Close() error {
return nil
}
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{
AccessToken: "agent-token",
NodeName: "edge-01",
NodeIP: "10.0.0.8",
Version: config.Version,
ExtVersion: "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{
AccessToken: "agent-token",
NodeName: "edge-01",
NodeIP: "10.0.0.8",
Version: config.Version,
ExtVersion: "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{
AccessToken: "agent-token",
NodeName: "edge-01",
NodeIP: "10.0.0.8",
Version: config.Version,
ExtVersion: "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",
Version: config.Version,
ExtVersion: "1.27.1.2",
DataDir: tempDir,
RouteConfigPath: filepath.Join(tempDir, "conf.d", "openflare_routes.conf"),
AccessLogPath: filepath.Join(tempDir, "var", "log", "openflare", "access.log"),
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
},
StateStore: stateStore,
}
if err := os.MkdirAll(filepath.Dir(runner.Config.AccessLogPath), 0o755); err != nil {
t.Fatalf("failed to prepare access log dir: %v", err)
}
if err := os.WriteFile(
runner.Config.AccessLogPath,
[]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{
AccessToken: "agent-token",
NodeName: "edge-buffer-01",
NodeIP: "10.0.0.52",
Version: config.Version,
ExtVersion: "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",
AccessToken: "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,
Version: config.Version,
ExtVersion: "1.27.1.2",
HeartbeatInterval: config.MillisecondDuration(10 * time.Millisecond),
},
StateStore: stateStore,
HeartbeatService: heartbeatService,
SyncService: syncService,
}
runner.Config = cfg
runner.Config.Version = config.Version
runner.Config.ExtVersion = "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.AccessToken != "agent-token-issued" || runner.Config.DiscoveryToken != "" {
t.Fatal("expected config token rotation to complete")
}
}
func TestRunnerHandlesWebSocketActiveConfigMessage(t *testing.T) {
syncService := &fakeSyncService{}
runner := &Runner{SyncService: syncService}
payload, err := json.Marshal(protocol.ActiveConfigMeta{
Version: "20260529-001",
Checksum: "checksum-ws",
})
if err != nil {
t.Fatalf("marshal active config: %v", err)
}
changed, err := runner.handleWebSocketMessage(context.Background(), protocol.WSMessage{
Type: protocol.WSMessageTypeActiveConfig,
Payload: payload,
}, &fakeWebSocketConnection{})
if err != nil {
t.Fatalf("handle websocket active config: %v", err)
}
if changed {
t.Fatal("active config message should not change heartbeat interval")
}
if syncService.syncOnceCalls != 1 {
t.Fatalf("expected one sync call, got %d", syncService.syncOnceCalls)
}
if syncService.lastTarget == nil || syncService.lastTarget.Version != "20260529-001" || syncService.lastTarget.Checksum != "checksum-ws" {
t.Fatalf("unexpected sync target: %+v", syncService.lastTarget)
}
}
func TestRunnerHandlesWebSocketSettingsDisabled(t *testing.T) {
runner := &Runner{
Config: &config.Config{
HeartbeatInterval: config.MillisecondDuration(10 * time.Second),
},
websocketUpgradeEnabled: true,
}
payload, err := json.Marshal(protocol.AgentSettings{
HeartbeatInterval: 15000,
WebsocketUpgradeEnabled: false,
})
if err != nil {
t.Fatalf("marshal settings: %v", err)
}
changed, err := runner.handleWebSocketMessage(context.Background(), protocol.WSMessage{
Type: protocol.WSMessageTypeSettings,
Payload: payload,
}, &fakeWebSocketConnection{})
if err == nil {
t.Fatal("expected disabled websocket setting to request fallback")
}
if !changed {
t.Fatal("expected heartbeat interval change to be reported")
}
if runner.websocketUpgradeEnabled {
t.Fatal("expected websocket upgrade to be disabled")
}
}
func TestWebSocketBackoffSequence(t *testing.T) {
backoff := newWebSocketBackoff()
expected := []time.Duration{
time.Second,
2 * time.Second,
5 * time.Second,
10 * time.Second,
30 * time.Second,
30 * time.Second,
}
for _, want := range expected {
if got := backoff.Next(); got != want {
t.Fatalf("unexpected backoff: got %s want %s", got, want)
}
}
backoff.Reset()
if got := backoff.Next(); got != time.Second {
t.Fatalf("expected reset backoff to return 1s, got %s", got)
}
}
+463
View File
@@ -0,0 +1,463 @@
package config
import (
"context"
"encoding/json"
"errors"
"fmt"
"net"
"os"
pathpkg "path"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/utils"
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
"github.com/rain-kl/openflare/openflare-server/utils/geoip/iputil"
)
const (
defaultMainConfigRelativePath = "etc/nginx/nginx.conf"
defaultRouteConfigRelativePath = "etc/nginx/conf.d/openflare_routes.conf"
defaultCertDirRelativePath = "etc/nginx/certs"
defaultLuaDirRelativePath = "etc/nginx/lua"
defaultRuntimeConfigDirRelativePath = "etc/openflare"
defaultPagesDirRelativePath = "var/lib/openflare/pages"
defaultMMDBRelativePath = "etc/openflare/GeoLite2-Country.mmdb"
defaultAccessLogRelativePath = "var/log/openflare/access.log"
defaultStateRelativePath = "var/lib/openflare/agent-state.json"
defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json"
defaultOpenRestyObservabilityPort = 18081
defaultObservabilityReplayMinutes = 15
defaultMMDBUpdateInterval = 24 * time.Hour
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
)
var (
lookupOutboundIP = geoip.GetOutboundIP
lookupLocalIP = detectLocalNodeIP
)
type Config struct {
ServerURL string `json:"server_url"`
AccessToken string `json:"agent_token"`
DiscoveryToken string `json:"discovery_token"`
NodeName string `json:"node_name"`
NodeIP string `json:"node_ip"`
Version string `json:"-"`
ExtVersion string `json:"-"`
OpenrestyPath string `json:"openresty_path"`
OpenrestyResolvers []string `json:"openresty_resolvers,omitempty"`
DataDir string `json:"data_dir"`
MainConfigPath string `json:"main_config_path"`
RouteConfigPath string `json:"route_config_path"`
AccessLogPath string `json:"access_log_path"`
CertDir string `json:"cert_dir"`
OpenrestyCertDir string `json:"openresty_cert_dir"`
LuaDir string `json:"lua_dir"`
OpenrestyLuaDir string `json:"openresty_lua_dir"`
RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
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"`
AccessToken string `json:"agent_token"`
DiscoveryToken string `json:"discovery_token"`
NodeName string `json:"node_name"`
NodeIP string `json:"node_ip"`
OpenrestyPath string `json:"openresty_path"`
OpenrestyResolvers []string `json:"openresty_resolvers"`
DataDir string `json:"data_dir"`
MainConfigPath string `json:"main_config_path"`
RouteConfigPath string `json:"route_config_path"`
AccessLogPath string `json:"access_log_path"`
CertDir string `json:"cert_dir"`
OpenrestyCertDir string `json:"openresty_cert_dir"`
LuaDir string `json:"lua_dir"`
OpenrestyLuaDir string `json:"openresty_lua_dir"`
RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
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 && !os.IsNotExist(err) {
return nil, err
}
file := &configFile{}
if err == nil {
if err = json.Unmarshal(data, file); err != nil {
return nil, err
}
}
if err != nil && !hasEnvConfig() {
return nil, err
}
cfg := &Config{
ServerURL: file.ServerURL,
AccessToken: file.AccessToken,
DiscoveryToken: file.DiscoveryToken,
NodeName: file.NodeName,
NodeIP: file.NodeIP,
OpenrestyPath: file.OpenrestyPath,
OpenrestyResolvers: append([]string{}, file.OpenrestyResolvers...),
DataDir: file.DataDir,
MainConfigPath: file.MainConfigPath,
RouteConfigPath: file.RouteConfigPath,
AccessLogPath: file.AccessLogPath,
CertDir: file.CertDir,
OpenrestyCertDir: file.OpenrestyCertDir,
LuaDir: file.LuaDir,
OpenrestyLuaDir: file.OpenrestyLuaDir,
RuntimeConfigDir: file.RuntimeConfigDir,
PagesDir: file.PagesDir,
MMDBPath: file.MMDBPath,
MMDBUpdateInterval: file.MMDBUpdateInterval,
MMDBDownloadURL: file.MMDBDownloadURL,
OpenrestyObservabilityPort: file.OpenrestyObservabilityPort,
ObservabilityBufferPath: file.ObservabilityBufferPath,
ObservabilityReplayMinutes: file.ObservabilityReplayMinutes,
StatePath: file.StatePath,
HeartbeatInterval: file.HeartbeatInterval,
RequestTimeout: file.RequestTimeout,
}
cfg.configPath = path
applyEnvOverrides(cfg)
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.Version = Version
cfg.OpenrestyResolvers = utils.UniqueAndCleanStringSlice(cfg.OpenrestyResolvers)
if cfg.OpenrestyPath == "" {
cfg.OpenrestyPath = "openresty"
}
if cfg.DataDir == "" {
cfg.DataDir = filepath.Join(baseDir, "data")
}
if cfg.NodeName == "" {
cfg.NodeName = detectHostname()
}
if cfg.NodeIP == "" {
cfg.NodeIP = detectNodeIP()
}
if cfg.MainConfigPath == "" {
cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultMainConfigRelativePath)
}
if cfg.RouteConfigPath == "" {
cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultRouteConfigRelativePath)
}
if cfg.AccessLogPath == "" {
cfg.AccessLogPath = joinManagedPath(cfg.DataDir, defaultAccessLogRelativePath)
}
if cfg.StatePath == "" {
cfg.StatePath = joinManagedPath(cfg.DataDir, defaultStateRelativePath)
}
if cfg.CertDir == "" {
cfg.CertDir = joinManagedPath(cfg.DataDir, defaultCertDirRelativePath)
}
if cfg.OpenrestyCertDir == "" {
cfg.OpenrestyCertDir = cfg.CertDir
}
if cfg.LuaDir == "" {
cfg.LuaDir = joinManagedPath(cfg.DataDir, defaultLuaDirRelativePath)
}
if cfg.OpenrestyLuaDir == "" {
cfg.OpenrestyLuaDir = cfg.LuaDir
}
if cfg.RuntimeConfigDir == "" {
cfg.RuntimeConfigDir = joinManagedPath(cfg.DataDir, defaultRuntimeConfigDirRelativePath)
}
if cfg.PagesDir == "" {
cfg.PagesDir = joinManagedPath(cfg.DataDir, defaultPagesDirRelativePath)
}
if cfg.MMDBPath == "" {
cfg.MMDBPath = joinManagedPath(cfg.DataDir, defaultMMDBRelativePath)
}
if cfg.MMDBUpdateInterval <= 0 {
cfg.MMDBUpdateInterval = MillisecondDuration(defaultMMDBUpdateInterval)
}
if cfg.MMDBDownloadURL == "" {
cfg.MMDBDownloadURL = defaultMMDBDownloadURL
}
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
}
paths := []*string{
&cfg.DataDir,
&cfg.MainConfigPath,
&cfg.RouteConfigPath,
&cfg.AccessLogPath,
&cfg.CertDir,
&cfg.OpenrestyCertDir,
&cfg.LuaDir,
&cfg.OpenrestyLuaDir,
&cfg.RuntimeConfigDir,
&cfg.PagesDir,
&cfg.StatePath,
&cfg.ObservabilityBufferPath,
&cfg.MMDBPath,
}
for _, p := range paths {
if usesSlashPath(*p) {
*p = filepath.ToSlash(*p)
}
}
}
func hasEnvConfig() bool {
for _, key := range []string{
"OPENFLARE_SERVER_URL",
"OPENFLARE_AGENT_TOKEN",
"OPENFLARE_DISCOVERY_TOKEN",
"OPENFLARE_NODE_NAME",
"OPENFLARE_NODE_IP",
"OPENFLARE_DATA_DIR",
"OPENFLARE_OPENRESTY_PATH",
"OPENFLARE_PAGES_DIR",
"OPENFLARE_HEARTBEAT_INTERVAL",
"OPENFLARE_REQUEST_TIMEOUT",
"OPENFLARE_OPENRESTY_OBSERVABILITY_PORT",
"OPENFLARE_MMDB_PATH",
"OPENFLARE_MMDB_UPDATE_INTERVAL",
"OPENFLARE_MMDB_DOWNLOAD_URL",
} {
if strings.TrimSpace(os.Getenv(key)) != "" {
return true
}
}
return false
}
func applyEnvOverrides(cfg *Config) {
if cfg == nil {
return
}
overrideString := func(key string, target *string) {
if value := strings.TrimSpace(os.Getenv(key)); value != "" {
*target = value
}
}
overrideString("OPENFLARE_SERVER_URL", &cfg.ServerURL)
overrideString("OPENFLARE_AGENT_TOKEN", &cfg.AccessToken)
overrideString("OPENFLARE_DISCOVERY_TOKEN", &cfg.DiscoveryToken)
overrideString("OPENFLARE_NODE_NAME", &cfg.NodeName)
overrideString("OPENFLARE_NODE_IP", &cfg.NodeIP)
overrideString("OPENFLARE_DATA_DIR", &cfg.DataDir)
overrideString("OPENFLARE_OPENRESTY_PATH", &cfg.OpenrestyPath)
overrideString("OPENFLARE_PAGES_DIR", &cfg.PagesDir)
overrideString("OPENFLARE_MMDB_PATH", &cfg.MMDBPath)
overrideString("OPENFLARE_MMDB_DOWNLOAD_URL", &cfg.MMDBDownloadURL)
if value := strings.TrimSpace(os.Getenv("OPENFLARE_HEARTBEAT_INTERVAL")); value != "" {
if duration, err := parseDurationValue(value); err == nil {
cfg.HeartbeatInterval = duration
}
}
if value := strings.TrimSpace(os.Getenv("OPENFLARE_REQUEST_TIMEOUT")); value != "" {
if duration, err := parseDurationValue(value); err == nil {
cfg.RequestTimeout = duration
}
}
if value := strings.TrimSpace(os.Getenv("OPENFLARE_MMDB_UPDATE_INTERVAL")); value != "" {
if duration, err := parseDurationValue(value); err == nil {
cfg.MMDBUpdateInterval = duration
}
}
if value := strings.TrimSpace(os.Getenv("OPENFLARE_OPENRESTY_OBSERVABILITY_PORT")); value != "" {
var port int
if _, err := fmt.Sscanf(value, "%d", &port); err == nil {
cfg.OpenrestyObservabilityPort = port
}
}
}
func parseDurationValue(value string) (MillisecondDuration, error) {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return 0, nil
}
if parsed, err := time.ParseDuration(trimmed); err == nil {
return MillisecondDuration(parsed), nil
}
ms, err := strconv.ParseInt(trimmed, 10, 64)
if err != nil {
return 0, err
}
return MillisecondDuration(time.Duration(ms) * time.Millisecond), nil
}
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.AccessToken) == "" && 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")
}
if cfg.MMDBUpdateInterval <= 0 {
return errors.New("mmdb_update_interval 必须大于 0")
}
return nil
}
func (cfg *Config) InitialAuthToken() string {
if cfg == nil {
return ""
}
if token := strings.TrimSpace(cfg.AccessToken); 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 detectNodeIP() string {
if ip := detectOutboundNodeIP(); ip != "" {
return ip
}
return lookupLocalIP()
}
func detectOutboundNodeIP() string {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
ip, err := lookupOutboundIP(ctx)
if err != nil || ip == nil {
return ""
}
return ip.String()
}
func detectLocalNodeIP() string {
interfaces, err := net.Interfaces()
if err != nil {
return ""
}
bestIP := ""
bestPriority := -1
for _, iface := range interfaces {
if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 {
continue
}
addrs, err := iface.Addrs()
if err != nil {
continue
}
for _, addr := range addrs {
ipNet, ok := addr.(*net.IPNet)
if !ok || ipNet.IP == nil || ipNet.IP.IsLoopback() {
continue
}
ipv4 := normalizeIPv4(ipNet.IP)
priority := nodeIPPriority(ipv4)
if priority > bestPriority {
bestIP = ipv4.String()
bestPriority = priority
}
if bestPriority == 2 {
return bestIP
}
}
}
return bestIP
}
func normalizeIPv4(ip net.IP) net.IP {
if ip == nil {
return nil
}
return ip.To4()
}
func nodeIPPriority(ip net.IP) int {
return iputil.Score(ip)
}
@@ -0,0 +1,520 @@
package config
import (
"context"
"encoding/json"
"errors"
"net"
"os"
"path/filepath"
"testing"
"time"
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
)
func TestLoadDefaultsToManagedBinaryPaths(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.OpenrestyPath != "openresty" {
t.Fatalf("unexpected openresty path: %s", cfg.OpenrestyPath)
}
if cfg.MainConfigPath != filepath.Join(dir, "data", defaultMainConfigRelativePath) {
t.Fatalf("unexpected main config path: %s", cfg.MainConfigPath)
}
if cfg.RouteConfigPath != filepath.Join(dir, "data", defaultRouteConfigRelativePath) {
t.Fatalf("unexpected route config path: %s", cfg.RouteConfigPath)
}
if cfg.AccessLogPath != filepath.Join(dir, "data", defaultAccessLogRelativePath) {
t.Fatalf("unexpected access log path: %s", cfg.AccessLogPath)
}
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.RuntimeConfigDir != filepath.Join(dir, "data", defaultRuntimeConfigDirRelativePath) {
t.Fatalf("unexpected runtime config dir: %s", cfg.RuntimeConfigDir)
}
if cfg.OpenrestyCertDir != cfg.CertDir {
t.Fatalf("unexpected openresty cert dir: %s", cfg.OpenrestyCertDir)
}
if cfg.OpenrestyLuaDir != cfg.LuaDir {
t.Fatalf("unexpected openresty lua dir: %s", cfg.OpenrestyLuaDir)
}
if cfg.StatePath != filepath.Join(dir, "data", defaultStateRelativePath) {
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 TestLoadNormalizesExplicitResolvers(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_resolvers": []string{" 10.0.0.2 ", "10.0.0.2", "", "1.1.1.1"},
}
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)
}
expected := []string{"10.0.0.2", "1.1.1.1"}
if len(cfg.OpenrestyResolvers) != len(expected) {
t.Fatalf("unexpected resolver count: %#v", cfg.OpenrestyResolvers)
}
for index, value := range expected {
if cfg.OpenrestyResolvers[index] != value {
t.Fatalf("unexpected resolver at %d: got %q want %q", index, cfg.OpenrestyResolvers[index], value)
}
}
}
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/"+defaultRouteConfigRelativePath {
t.Fatalf("unexpected route config path: %s", cfg.RouteConfigPath)
}
if cfg.MainConfigPath != "/srv/openflare/"+defaultMainConfigRelativePath {
t.Fatalf("unexpected main config path: %s", cfg.MainConfigPath)
}
if cfg.AccessLogPath != "/srv/openflare/"+defaultAccessLogRelativePath {
t.Fatalf("unexpected access log path: %s", cfg.AccessLogPath)
}
if cfg.StatePath != "/srv/openflare/"+defaultStateRelativePath {
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)
}
if cfg.RuntimeConfigDir != "/srv/openflare/"+defaultRuntimeConfigDirRelativePath {
t.Fatalf("unexpected runtime config dir: %s", cfg.RuntimeConfigDir)
}
}
func TestLoadUsesEnvConfigWhenFileIsMissing(t *testing.T) {
dir := t.TempDir()
t.Setenv("OPENFLARE_SERVER_URL", "http://127.0.0.1:3000")
t.Setenv("OPENFLARE_AGENT_TOKEN", "token")
t.Setenv("OPENFLARE_NODE_NAME", "edge-env")
t.Setenv("OPENFLARE_NODE_IP", "10.0.0.9")
t.Setenv("OPENFLARE_DATA_DIR", "/srv/openflare-env")
t.Setenv("OPENFLARE_OPENRESTY_PATH", "/usr/bin/openresty")
t.Setenv("OPENFLARE_HEARTBEAT_INTERVAL", "45s")
t.Setenv("OPENFLARE_REQUEST_TIMEOUT", "2500")
t.Setenv("OPENFLARE_OPENRESTY_OBSERVABILITY_PORT", "19091")
cfg, err := Load(filepath.Join(dir, "missing-agent.json"))
if err != nil {
t.Fatalf("Load failed: %v", err)
}
if cfg.ServerURL != "http://127.0.0.1:3000" || cfg.AccessToken != "token" {
t.Fatalf("unexpected env auth config: %#v", cfg)
}
if cfg.OpenrestyPath != "/usr/bin/openresty" {
t.Fatalf("unexpected openresty path: %s", cfg.OpenrestyPath)
}
if cfg.DataDir != "/srv/openflare-env" {
t.Fatalf("unexpected data dir: %s", cfg.DataDir)
}
if cfg.HeartbeatInterval.Duration() != 45*time.Second {
t.Fatalf("unexpected heartbeat interval: %s", cfg.HeartbeatInterval)
}
if cfg.RequestTimeout.Duration() != 2500*time.Millisecond {
t.Fatalf("unexpected request timeout: %s", cfg.RequestTimeout)
}
if cfg.OpenrestyObservabilityPort != 19091 {
t.Fatalf("unexpected observability port: %d", cfg.OpenrestyObservabilityPort)
}
}
func TestLoadDetectsOutboundIPWhenNodeIPMissing(t *testing.T) {
previousLookup := lookupOutboundIP
lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
return net.ParseIP("8.8.8.8"), nil
}
defer func() {
lookupOutboundIP = previousLookup
}()
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",
}
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.NodeIP != "8.8.8.8" {
t.Fatalf("expected outbound IP, got %s", cfg.NodeIP)
}
}
func TestLoadFallsBackToLocalIPWhenOutboundLookupFails(t *testing.T) {
previousOutboundLookup := lookupOutboundIP
previousLocalLookup := lookupLocalIP
lookupOutboundIP = func(ctx context.Context, strategies ...geoip.OutboundIPStrategy) (net.IP, error) {
return nil, errors.New("realip.cc unavailable")
}
lookupLocalIP = func() string {
return "9.9.9.9"
}
defer func() {
lookupOutboundIP = previousOutboundLookup
lookupLocalIP = previousLocalLookup
}()
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",
}
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.NodeIP != "9.9.9.9" {
t.Fatalf("expected local fallback IP, got %s", cfg.NodeIP)
}
}
func TestLoadEnvOverridesConfigFile(t *testing.T) {
dir := t.TempDir()
configPath := filepath.Join(dir, "agent.json")
if err := os.WriteFile(configPath, []byte(`{"server_url":"http://old:3000","agent_token":"old","node_name":"edge-01","node_ip":"10.0.0.8","openresty_path":"/old/openresty"}`), 0o644); err != nil {
t.Fatalf("failed to write config: %v", err)
}
t.Setenv("OPENFLARE_SERVER_URL", "http://new:3000")
t.Setenv("OPENFLARE_AGENT_TOKEN", "new-token")
t.Setenv("OPENFLARE_OPENRESTY_PATH", "/new/openresty")
cfg, err := Load(configPath)
if err != nil {
t.Fatalf("Load failed: %v", err)
}
if cfg.ServerURL != "http://new:3000" {
t.Fatalf("expected server url from env, got %s", cfg.ServerURL)
}
if cfg.AccessToken != "new-token" {
t.Fatalf("expected token from env, got %s", cfg.AccessToken)
}
if cfg.OpenrestyPath != "/new/openresty" {
t.Fatalf("expected openresty path from env, got %s", cfg.OpenrestyPath)
}
}
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.ExtVersion = "1.27.1.2"
cfg.HeartbeatInterval = MillisecondDuration(5 * time.Second)
cfg.RequestTimeout = MillisecondDuration(7 * time.Second)
cfg.OpenrestyResolvers = []string{"10.0.0.2", "1.1.1.1"}
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"])
}
resolvers, ok := decoded["openresty_resolvers"].([]any)
if !ok || len(resolvers) != 2 || resolvers[0] != "10.0.0.2" || resolvers[1] != "1.1.1.1" {
t.Fatalf("unexpected resolvers: %#v", decoded["openresty_resolvers"])
}
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{
AccessToken: tt.agentToken,
DiscoveryToken: tt.discoveryToken,
}
}
if token := cfg.InitialAuthToken(); token != tt.expected {
t.Fatalf("unexpected initial auth token: %q", token)
}
})
}
}
func TestNodeIPPriority(t *testing.T) {
tests := []struct {
name string
ip string
expected int
}{
{
name: "public ipv4 preferred",
ip: "8.8.8.8",
expected: 2,
},
{
name: "private ipv4 fallback",
ip: "10.0.0.8",
expected: 1,
},
{
name: "link local ignored",
ip: "169.254.1.10",
expected: -1,
},
{
name: "loopback ignored",
ip: "127.0.0.1",
expected: -1,
},
{
name: "nil ignored",
ip: "",
expected: -1,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var parsed net.IP
if tt.ip != "" {
parsed = net.ParseIP(tt.ip)
}
if got := nodeIPPriority(parsed); got != tt.expected {
t.Fatalf("unexpected priority for %q: got %d want %d", tt.ip, got, tt.expected)
}
})
}
}
@@ -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 Version = "dev"
@@ -0,0 +1,8 @@
package geoipdata
import "embed"
//go:embed GeoLite2-Country.mmdb
var FS embed.FS
const DefaultMMDBName = "GeoLite2-Country.mmdb"
@@ -0,0 +1,67 @@
package geoipupdate
import (
"context"
"fmt"
"io/fs"
"log/slog"
"os"
"path/filepath"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/geoipdata"
"github.com/rain-kl/openflare/openflare-server/utils/geoip"
)
type Updater struct {
MMDBPath string
DownloadURL string
UpdateInterval time.Duration
}
func (u *Updater) EnsureInitialDatabase() error {
path := filepath.Clean(u.MMDBPath)
if path == "" || path == "." {
return nil
}
if _, err := os.Stat(path); err == nil {
return nil
} else if !os.IsNotExist(err) {
return fmt.Errorf("stat mmdb file failed: %w", err)
}
data, err := fs.ReadFile(geoipdata.FS, geoipdata.DefaultMMDBName)
if err != nil {
return fmt.Errorf("read embedded mmdb failed: %w", err)
}
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return fmt.Errorf("create mmdb directory failed: %w", err)
}
if err := os.WriteFile(path, data, 0o644); err != nil {
return fmt.Errorf("write initial mmdb failed: %w", err)
}
slog.Info("initialized GeoIP mmdb from embedded database", "path", path, "size", len(data))
return nil
}
func (u *Updater) Run(ctx context.Context) {
if u == nil || u.MMDBPath == "" || u.UpdateInterval <= 0 {
return
}
if err := u.EnsureInitialDatabase(); err != nil {
slog.Warn("initialize GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
}
ticker := time.NewTicker(u.UpdateInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if err := geoip.DownloadMaxMindDatabase(u.MMDBPath, u.DownloadURL); err != nil {
slog.Warn("update GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
continue
}
slog.Info("GeoIP mmdb updated", "path", u.MMDBPath)
}
}
}
@@ -0,0 +1,24 @@
package geoipupdate
import (
"os"
"path/filepath"
"testing"
)
func TestEnsureInitialDatabaseCopiesEmbeddedMMDB(t *testing.T) {
tempDir := t.TempDir()
path := filepath.Join(tempDir, "GeoLite2-Country.mmdb")
updater := &Updater{MMDBPath: path}
if err := updater.EnsureInitialDatabase(); err != nil {
t.Fatalf("EnsureInitialDatabase failed: %v", err)
}
info, err := os.Stat(path)
if err != nil {
t.Fatalf("expected mmdb to exist: %v", err)
}
if info.Size() == 0 {
t.Fatal("expected copied mmdb to be non-empty")
}
}
@@ -0,0 +1,33 @@
package heartbeat
import (
"context"
"github.com/rain-kl/openflare/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,168 @@
package httpclient
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"time"
"github.com/rain-kl/openflare/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,
WAFIPGroups: resp.WAFIPGroups,
}, 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) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) {
resp := protocol.APIResponse[protocol.WAFIPGroupSyncResponse]{}
if err := c.postJSON(ctx, "/api/agent/waf/ip-groups/sync", payload, &resp); err != nil {
return nil, err
}
if !resp.Success {
return nil, errors.New(resp.Message)
}
return &resp.Data, nil
}
func (c *Client) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+fmt.Sprintf("/api/agent/pages/deployments/%d/package", deploymentID), nil)
if err != nil {
return nil, err
}
req.Header.Set("X-Agent-Token", c.token)
res, err := c.httpClient.Do(req)
if err != nil {
return nil, err
}
defer res.Body.Close()
if res.StatusCode != http.StatusOK {
return nil, errors.New(res.Status)
}
return io.ReadAll(res.Body)
}
func (c *Client) SetToken(token string) {
c.token = strings.TrimSpace(token)
slog.Debug("http client token updated")
}
func (c *Client) getJSON(ctx context.Context, path string, target any) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil)
if err != nil {
return err
}
req.Header.Set("X-Agent-Token", c.token)
return c.do(req, target)
}
func (c *Client) postJSON(ctx context.Context, path string, body any, target any) error {
data, err := json.Marshal(body)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+path, bytes.NewReader(data))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Agent-Token", c.token)
return c.do(req, target)
}
func (c *Client) do(req *http.Request, target any) error {
res, err := c.httpClient.Do(req)
if err != nil {
slog.Error("http request failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
defer func(Body io.ReadCloser) {
err := Body.Close()
if err != nil {
slog.Error("failed to close response body", "error", err)
}
}(res.Body)
if res.StatusCode != http.StatusOK {
slog.Warn("http request returned non-200", "method", req.Method, "path", req.URL.Path, "status", res.Status)
return errors.New(res.Status)
}
if target == nil {
var wrapper protocol.APIResponse[json.RawMessage]
if err = json.NewDecoder(res.Body).Decode(&wrapper); err != nil {
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
if !wrapper.Success {
slog.Warn("http api response failed", "method", req.Method, "path", req.URL.Path, "message", wrapper.Message)
return errors.New(wrapper.Message)
}
return nil
}
if err = json.NewDecoder(res.Body).Decode(target); err != nil {
slog.Error("http response decode failed", "method", req.Method, "path", req.URL.Path, "error", err)
return err
}
return nil
}
@@ -0,0 +1,29 @@
package logging
import (
"log/slog"
"os"
"strings"
)
func Setup() {
opts := &slog.HandlerOptions{
AddSource: true,
Level: parseLevel(os.Getenv("LOG_LEVEL")),
}
handler := slog.NewTextHandler(os.Stdout, opts)
slog.SetDefault(slog.New(handler))
}
func parseLevel(value string) slog.Level {
switch strings.ToLower(strings.TrimSpace(value)) {
case "debug":
return slog.LevelDebug
case "warn", "warning":
return slog.LevelWarn
case "error":
return slog.LevelError
default:
return slog.LevelInfo
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,968 @@
package nginx
import (
"context"
"errors"
"net"
"net/http"
"os"
"path/filepath"
"reflect"
"runtime"
"strings"
"testing"
"github.com/rain-kl/openflare/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
}
type scriptedExecutor struct {
testErrors []error
testCalls int
reloadErrors []error
reloadCalls int
}
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 (e *scriptedExecutor) Test(ctx context.Context) error {
index := e.testCalls
e.testCalls++
if index >= len(e.testErrors) {
return nil
}
return e.testErrors[index]
}
func (e *scriptedExecutor) Reload(ctx context.Context) error {
index := e.reloadCalls
e.reloadCalls++
if index >= len(e.reloadErrors) {
return nil
}
return e.reloadErrors[index]
}
func (e *scriptedExecutor) EnsureRuntime(ctx context.Context, recreate bool) error {
return nil
}
func (e *scriptedExecutor) CheckHealth(ctx context.Context) error {
return nil
}
func (e *scriptedExecutor) Restart(ctx context.Context) error {
return nil
}
func TestPathExecutorCommands(t *testing.T) {
runner := &fakeRunner{}
executor := &PathExecutor{
Path: "/usr/local/openresty/nginx/sbin/openresty",
ConfigPath: "/data/etc/nginx/nginx.conf",
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", "-c", "/data/etc/nginx/nginx.conf"}},
{name: "/usr/local/openresty/nginx/sbin/openresty", args: []string{"-s", "reload", "-c", "/data/etc/nginx/nginx.conf"}},
}
if !reflect.DeepEqual(runner.calls, expected) {
t.Fatalf("unexpected calls: %#v", runner.calls)
}
}
func TestPathExecutorEnsureRuntimeNoop(t *testing.T) {
runner := &fakeRunner{}
executor := &PathExecutor{
Path: "/usr/local/openresty/nginx/sbin/openresty",
ConfigPath: "/data/etc/nginx/nginx.conf",
Runner: runner,
}
if err := executor.EnsureRuntime(context.Background(), true); err != nil {
t.Fatalf("EnsureRuntime failed: %v", err)
}
if len(runner.calls) != 2 {
t.Fatalf("expected test and reload calls, got %d", len(runner.calls))
}
}
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",
ConfigPath: "/data/etc/nginx/nginx.conf",
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 TestPathExecutorReloadStartsWhenRuntimeIsNotRunning(t *testing.T) {
runner := &fakeRunner{
runFn: func(name string, args ...string) ([]byte, error) {
if len(args) >= 2 && args[0] == "-s" && args[1] == "reload" {
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",
ConfigPath: "/data/etc/nginx/nginx.conf",
Runner: runner,
}
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{"-s", "reload", "-c", "/data/etc/nginx/nginx.conf"}},
{name: "/usr/local/openresty/nginx/sbin/openresty", args: []string{"-c", "/data/etc/nginx/nginx.conf"}},
}
if !reflect.DeepEqual(runner.calls, expected) {
t.Fatalf("unexpected calls: %#v", runner.calls)
}
}
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")
accessLogPath := filepath.Join(tempDir, "var", "log", "openflare", "access.log")
manager := &Manager{
MainConfigPath: mainPath,
RouteConfigPath: routePath,
AccessLogPath: accessLogPath,
CertDir: certDir,
NginxCertDir: "/etc/nginx/openflare-certs",
LuaDir: filepath.Join(tempDir, "lua"),
NginxLuaDir: "/etc/nginx/openflare-lua",
Executor: &fakeExecutor{},
}
outcome := 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 outcome.Status != ApplyStatusSuccess {
t.Fatalf("Apply failed: %#v", outcome)
}
mainData, err := os.ReadFile(mainPath)
if err != nil {
t.Fatalf("failed to read main config: %v", err)
}
expectedMain := "include " + routePath + ";\naccess_log " + filepath.ToSlash(accessLogPath) + " 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 TestParseExtVersionIgnoresDockerEntrypointPaths(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 := parseExtVersion(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",
OpenrestyResolverDirective: " resolver 127.0.0.11 valid=30s ipv6=off;\n resolver_timeout 5s;\n",
Executor: &fakeExecutor{},
}
outcome := manager.Apply(context.Background(), "include __OPENFLARE_ROUTE_CONFIG__;\n__OPENFLARE_RESOLVER_DIRECTIVE__server { 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 outcome.Status != ApplyStatusSuccess {
t.Fatalf("Apply failed: %#v", outcome)
}
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))
}
renderedRoute := manager.renderRouteConfig("access_by_lua_file __OPENFLARE_LUA_DIR__/pow/check.lua;\nlocation /.within.website/x/cmd/anubis/static/ { alias __OPENFLARE_POW_STATIC_DIR__/; }\n")
if !strings.Contains(renderedRoute, "access_by_lua_file /etc/nginx/openflare-lua/pow/check.lua;") {
t.Fatalf("expected lua dir placeholder replacement in route config, got %s", renderedRoute)
}
if !strings.Contains(renderedRoute, "alias /etc/nginx/openflare-lua/pow/static/;") {
t.Fatalf("expected pow static dir placeholder replacement in route config, got %s", renderedRoute)
}
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))
}
if !strings.Contains(string(mainData), "resolver 127.0.0.11 valid=30s ipv6=off;") {
t.Fatalf("expected resolver directive 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 runtime.GOOS != "windows" && luaInfo.Mode().Perm() != 0o644 {
t.Fatalf("unexpected lua mode: %o", luaInfo.Mode().Perm())
}
}
func TestManagerCheckHealthUsesStubStatusInsteadOfConfigTest(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
server := &http.Server{
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/openflare/stub_status" {
http.NotFound(w, r)
return
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("Active connections: 1\n"))
}),
}
go func() {
_ = server.Serve(listener)
}()
defer server.Shutdown(context.Background())
mainPath := filepath.Join(t.TempDir(), "nginx.conf")
if err := os.WriteFile(mainPath, []byte("main"), 0o644); err != nil {
t.Fatalf("WriteFile failed: %v", err)
}
manager := &Manager{
MainConfigPath: mainPath,
OpenrestyObservabilityPort: port,
Executor: &fakeExecutor{
testErr: errors.New("openresty -t should not be called"),
},
}
if err := manager.CheckHealth(context.Background()); err != nil {
t.Fatalf("CheckHealth failed: %v", err)
}
}
func TestManagerCheckHealthFailsWhenStubStatusUnavailable(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
if err := listener.Close(); err != nil {
t.Fatalf("listener close failed: %v", err)
}
mainPath := filepath.Join(t.TempDir(), "nginx.conf")
if err := os.WriteFile(mainPath, []byte("main"), 0o644); err != nil {
t.Fatalf("WriteFile failed: %v", err)
}
manager := &Manager{
MainConfigPath: mainPath,
OpenrestyObservabilityPort: port,
Executor: &fakeExecutor{},
}
if err := manager.CheckHealth(context.Background()); err == nil {
t.Fatal("expected CheckHealth to fail when stub_status is unavailable")
}
}
func TestResolverDirectiveUsesExplicitResolvers(t *testing.T) {
got := ResolverDirective([]string{"10.0.0.2", "1.1.1.1"})
if !strings.Contains(got, "resolver 10.0.0.2 1.1.1.1") {
t.Fatalf("expected explicit resolver directive, got %q", got)
}
}
func TestParseResolverAddressesFiltersLoopbackForDocker(t *testing.T) {
content := strings.Join([]string{
"nameserver 127.0.0.53",
"nameserver 10.0.0.2",
"nameserver ::1",
"nameserver 1.1.1.1",
}, "\n")
got := parseResolverAddresses(content, true)
expected := []string{"10.0.0.2", "1.1.1.1"}
if !reflect.DeepEqual(got, expected) {
t.Fatalf("unexpected docker resolvers: got %#v want %#v", got, expected)
}
}
func TestParseResolverAddressesKeepsLoopbackForLocalBinary(t *testing.T) {
content := strings.Join([]string{
"nameserver 127.0.0.53",
"nameserver 10.0.0.2",
}, "\n")
got := parseResolverAddresses(content, false)
expected := []string{"127.0.0.53", "10.0.0.2"}
if !reflect.DeepEqual(got, expected) {
t.Fatalf("unexpected local resolvers: got %#v want %#v", got, expected)
}
}
func TestRequiresRuntimeResolver(t *testing.T) {
testCases := []struct {
name string
originURL string
want bool
}{
{name: "hostname", originURL: "https://origin.internal", want: true},
{name: "ipv4", originURL: "https://10.0.0.8", want: false},
{name: "ipv6", originURL: "https://[2001:db8::1]", want: false},
{name: "invalid", originURL: "://bad", want: false},
}
for _, testCase := range testCases {
if got := RequiresRuntimeResolver(testCase.originURL); got != testCase.want {
t.Fatalf("%s: got %v want %v", testCase.name, got, testCase.want)
}
}
}
func TestWriteCertFilesKeepsBaseDirAndRemovesStaleFiles(t *testing.T) {
tempDir := t.TempDir()
certDir := filepath.Join(tempDir, "certs")
if err := os.MkdirAll(filepath.Join(certDir, "stale"), 0o755); err != nil {
t.Fatalf("MkdirAll failed: %v", err)
}
if err := os.WriteFile(filepath.Join(certDir, "stale", "old.crt"), []byte("old"), 0o644); err != nil {
t.Fatalf("WriteFile failed: %v", err)
}
manager := &Manager{CertDir: certDir}
if err := manager.writeCertFiles([]protocol.SupportFile{
{Path: "1.crt", Content: "cert"},
{Path: "1.key", Content: "key"},
}); err != nil {
t.Fatalf("writeCertFiles failed: %v", err)
}
if _, err := os.Stat(certDir); err != nil {
t.Fatalf("expected cert dir to persist, stat err = %v", err)
}
if _, err := os.Stat(filepath.Join(certDir, "stale", "old.crt")); !os.IsNotExist(err) {
t.Fatalf("expected stale cert file to be removed, stat err = %v", err)
}
if _, err := os.Stat(filepath.Join(certDir, "1.crt")); err != nil {
t.Fatalf("expected new cert file to exist, stat err = %v", err)
}
}
func TestEnsureLuaAssetsKeepsBaseDirAndRemovesStaleFiles(t *testing.T) {
tempDir := t.TempDir()
luaDir := filepath.Join(tempDir, "lua")
if err := os.MkdirAll(filepath.Join(luaDir, "stale"), 0o755); err != nil {
t.Fatalf("MkdirAll failed: %v", err)
}
if err := os.WriteFile(filepath.Join(luaDir, "stale", "old.lua"), []byte("old"), 0o644); err != nil {
t.Fatalf("WriteFile failed: %v", err)
}
manager := &Manager{LuaDir: luaDir}
if err := manager.EnsureLuaAssets(); err != nil {
t.Fatalf("EnsureLuaAssets failed: %v", err)
}
if _, err := os.Stat(luaDir); err != nil {
t.Fatalf("expected lua dir to persist, stat err = %v", err)
}
if _, err := os.Stat(filepath.Join(luaDir, "stale", "old.lua")); !os.IsNotExist(err) {
t.Fatalf("expected stale lua file to be removed, stat err = %v", err)
}
if _, err := os.Stat(filepath.Join(luaDir, "log.lua")); err != nil {
t.Fatalf("expected managed lua file to exist, stat err = %v", err)
}
if _, err := os.Stat(filepath.Join(luaDir, "pow", "check.lua")); err != nil {
t.Fatalf("expected managed pow lua file to exist, stat err = %v", err)
}
if _, err := os.Stat(filepath.Join(luaDir, "pow", "static", "js", "main.mjs")); err != nil {
t.Fatalf("expected managed pow static asset to exist, stat err = %v", err)
}
}
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",
RuntimeConfigDir: filepath.Join(tempDir, "runtime"),
}
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())
}
if _, err := os.Stat(filepath.Join(manager.LuaDir, "pow", "check.lua")); err != nil {
t.Fatalf("failed to stat pow lua file: %v", err)
}
data, err := os.ReadFile(filepath.Join(manager.LuaDir, "pow", "runtime.lua"))
if err != nil {
t.Fatalf("failed to read pow lua file: %v", err)
}
if !strings.Contains(string(data), filepath.ToSlash(manager.RuntimeConfigDir)+"/waf_config.json") {
t.Fatalf("expected pow lua to read runtime config dir, got %s", string(data))
}
}
func TestEnsureLuaAssetsLeavesRuntimePowConfigOutsideLuaDir(t *testing.T) {
tempDir := t.TempDir()
luaDir := filepath.Join(tempDir, "lua")
runtimeConfigDir := filepath.Join(tempDir, "runtime")
if err := os.MkdirAll(runtimeConfigDir, 0o755); err != nil {
t.Fatalf("MkdirAll failed: %v", err)
}
powConfigPath := filepath.Join(runtimeConfigDir, "pow_config.json")
want := `[{"domains":["pow.example.com"],"enabled":true}]`
if err := os.WriteFile(powConfigPath, []byte(want), 0o644); err != nil {
t.Fatalf("WriteFile failed: %v", err)
}
manager := &Manager{LuaDir: luaDir, RuntimeConfigDir: runtimeConfigDir}
if err := manager.EnsureLuaAssets(); err != nil {
t.Fatalf("EnsureLuaAssets failed: %v", err)
}
got, err := os.ReadFile(powConfigPath)
if err != nil {
t.Fatalf("expected pow_config.json to remain after EnsureLuaAssets: %v", err)
}
if string(got) != want {
t.Fatalf("unexpected pow_config.json content: got %s want %s", string(got), want)
}
if _, err := os.Stat(filepath.Join(luaDir, "pow_config.json")); !os.IsNotExist(err) {
t.Fatalf("expected lua pow_config.json to stay absent, stat err = %v", err)
}
}
func TestManagerApplyWritesPowConfigToRuntimeDirAndCleansLegacyCopies(t *testing.T) {
tempDir := t.TempDir()
certDir := filepath.Join(tempDir, "certs")
luaDir := filepath.Join(tempDir, "lua")
runtimeConfigDir := filepath.Join(tempDir, "runtime")
for _, dir := range []string{certDir, luaDir, runtimeConfigDir} {
if err := os.MkdirAll(dir, 0o755); err != nil {
t.Fatalf("MkdirAll failed: %v", err)
}
}
for _, path := range []string{filepath.Join(certDir, "pow_config.json"), filepath.Join(luaDir, "pow_config.json")} {
if err := os.WriteFile(path, []byte("stale"), 0o644); err != nil {
t.Fatalf("WriteFile failed: %v", err)
}
}
manager := &Manager{
MainConfigPath: filepath.Join(tempDir, "nginx.conf"),
RouteConfigPath: filepath.Join(tempDir, "routes.conf"),
CertDir: certDir,
LuaDir: luaDir,
RuntimeConfigDir: runtimeConfigDir,
Executor: &fakeExecutor{},
}
outcome := manager.Apply(context.Background(), "main", "route", []protocol.SupportFile{
{Path: "pow_config.json", Content: "runtime"},
})
if outcome.Status != ApplyStatusSuccess {
t.Fatalf("Apply failed: %#v", outcome)
}
data, err := os.ReadFile(filepath.Join(runtimeConfigDir, "pow_config.json"))
if err != nil {
t.Fatalf("failed to read runtime pow config: %v", err)
}
if string(data) != "runtime" {
t.Fatalf("unexpected runtime pow config: %s", string(data))
}
for _, path := range []string{filepath.Join(certDir, "pow_config.json"), filepath.Join(luaDir, "pow_config.json")} {
if _, err := os.Stat(path); !os.IsNotExist(err) {
t.Fatalf("expected legacy pow config to be removed from %s, stat err = %v", path, err)
}
}
}
func TestManagerCurrentChecksumIncludesPowConfig(t *testing.T) {
tempDir := t.TempDir()
mainPath := filepath.Join(tempDir, "nginx.conf")
routePath := filepath.Join(tempDir, "routes.conf")
luaDir := filepath.Join(tempDir, "lua")
runtimeConfigDir := filepath.Join(tempDir, "runtime")
manager := &Manager{
MainConfigPath: mainPath,
RouteConfigPath: routePath,
LuaDir: luaDir,
NginxLuaDir: "/etc/nginx/openflare-lua",
RuntimeConfigDir: runtimeConfigDir,
Executor: &fakeExecutor{},
}
outcome := manager.Apply(
context.Background(),
"access_log __OPENFLARE_ACCESS_LOG__ openflare_json;\n",
"location /.within.website/x/cmd/anubis/static/ { alias __OPENFLARE_POW_STATIC_DIR__/; }\n",
[]protocol.SupportFile{{Path: "pow_config.json", Content: `[{"domains":["pow.example.com"],"enabled":true}]`}},
)
if outcome.Status != ApplyStatusSuccess {
t.Fatalf("Apply failed: %#v", outcome)
}
value, err := manager.CurrentChecksum()
if err != nil {
t.Fatalf("CurrentChecksum failed: %v", err)
}
expected := bundleChecksum(
"access_log __OPENFLARE_ACCESS_LOG__ openflare_json;\n",
"location /.within.website/x/cmd/anubis/static/ { alias __OPENFLARE_POW_STATIC_DIR__/; }\n",
[]protocol.SupportFile{{Path: "pow_config.json", Content: `[{"domains":["pow.example.com"],"enabled":true}]`}},
)
if value != expected {
t.Fatalf("unexpected checksum with pow config: got %s want %s", value, expected)
}
}
func TestManagedPowLuaFilesUseInternalChallengeFlow(t *testing.T) {
if !strings.Contains(openRestyPowRuntimeLua, `return ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")`) {
t.Fatal("expected pow runtime lua to internally execute make-challenge instead of issuing a 302 redirect")
}
if strings.Contains(openRestyPowRuntimeLua, "ngx.redirect(") {
t.Fatal("expected pow runtime lua to avoid external redirects for challenge rendering")
}
if !strings.Contains(openRestyPowChallengeLua, `<h1 id="title" class="centered-div">`) {
t.Fatal("expected challenge html to include Anubis-compatible title node")
}
if !strings.Contains(openRestyPowChallengeLua, `<div id="progress" role="progressbar" aria-labelledby="status"><div class="bar-inner"></div></div>`) {
t.Fatal("expected challenge html to include Anubis-compatible progress markup")
}
if !strings.Contains(openRestyPowChallengeLua, `<script id="anubis_public_url" type="application/json">"__openflare_internal__"</script>`) {
t.Fatal("expected challenge html to force Anubis frontend to reuse the current URL as redir target")
}
if !strings.Contains(openRestyPowRuntimeLua, `pow_sessions:set(session_key, "1", session_ttl)`) {
t.Fatal("expected pow runtime lua to refresh the PoW session TTL on each valid request")
}
if !strings.Contains(openRestyPowRuntimeLua, `ngx.header["Set-Cookie"] = session_cookie(cookie_val, session_ttl)`) {
t.Fatal("expected pow runtime lua to refresh the browser session cookie on each valid request")
}
if !strings.Contains(openRestyPowChallengeLua, `local session_ttl = config.session_ttl or 600`) {
t.Fatal("expected challenge.lua to default session TTL to 10 minutes")
}
if !strings.Contains(openRestyPowVerifyLua, `local session_ttl = challenge_info.session_ttl or 600`) {
t.Fatal("expected verify.lua to default session TTL to 10 minutes")
}
if !strings.Contains(openRestyPowVerifyLua, `if ngx.var.scheme == "https" then`) {
t.Fatal("expected verify.lua to only mark the session cookie as Secure for HTTPS requests")
}
}
func TestManagedWAFLuaTreatsWhitelistAsAllowlist(t *testing.T) {
if !strings.Contains(openRestyWAFRuntimeLua, "local function first_allowlist_group(groups)") {
t.Fatal("expected waf runtime to detect allowlist rule groups")
}
if !strings.Contains(openRestyWAFRuntimeLua, "local allowlist_group = first_allowlist_group(groups)") {
t.Fatal("expected waf runtime to enter allowlist mode when whitelist rules exist")
}
if !strings.Contains(openRestyWAFRuntimeLua, "return exit_with_group(allowlist_group)") {
t.Fatal("expected waf runtime to block requests that miss configured whitelists")
}
}
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{
reloadErr: errors.New("openresty reload failed"),
},
}
outcome := manager.Apply(context.Background(), "new-main", "new-route", []protocol.SupportFile{
{Path: "1.crt", Content: "new-cert"},
})
if outcome.Status != ApplyStatusFatal {
t.Fatalf("expected fatal apply outcome, got %#v", outcome)
}
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 TestManagerApplyReturnsWarningWhenRollbackRecoversRuntime(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: &scriptedExecutor{
reloadErrors: []error{errors.New("target config failed"), nil},
},
}
outcome := manager.Apply(context.Background(), "new-main", "new-route", []protocol.SupportFile{
{Path: "1.crt", Content: "new-cert"},
})
if outcome.Status != ApplyStatusWarning {
t.Fatalf("expected warning apply outcome, got %#v", outcome)
}
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))
}
}
func TestManagerApplyStartsSafeFallbackWhenNoRollbackConfigExists(t *testing.T) {
tempDir := t.TempDir()
routePath := filepath.Join(tempDir, "routes.conf")
mainPath := filepath.Join(tempDir, "nginx.conf")
executor := &scriptedExecutor{
testErrors: []error{errors.New("target config failed"), errors.New("rollback config missing"), nil},
}
manager := &Manager{
MainConfigPath: mainPath,
RouteConfigPath: routePath,
OpenrestyObservabilityListen: "127.0.0.1:18081",
Executor: executor,
}
outcome := manager.Apply(context.Background(), "bad-main", "bad-route", nil)
if outcome.Status != ApplyStatusWarning {
t.Fatalf("expected warning apply outcome, got %#v", outcome)
}
if !strings.Contains(outcome.Message, "fallback runtime started") {
t.Fatalf("expected fallback message, got %q", outcome.Message)
}
if executor.testCalls != 3 {
t.Fatalf("expected target, rollback, and fallback tests, got %d", executor.testCalls)
}
mainData, err := os.ReadFile(mainPath)
if err != nil {
t.Fatalf("failed to read main config: %v", err)
}
if !strings.Contains(string(mainData), "OpenFlare: No Valid Configuration") {
t.Fatalf("expected safe fallback main config, got %s", string(mainData))
}
if !strings.Contains(string(mainData), "listen 80 default_server") {
t.Fatalf("expected fallback to listen on port 80, got %s", string(mainData))
}
if !strings.Contains(string(mainData), "listen 127.0.0.1:18081") {
t.Fatalf("expected fallback to expose local stub_status port, got %s", string(mainData))
}
if !strings.Contains(string(mainData), "stub_status;") {
t.Fatalf("expected fallback to expose stub_status, got %s", string(mainData))
}
routeData, err := os.ReadFile(routePath)
if err != nil {
t.Fatalf("failed to read route config: %v", err)
}
if len(routeData) != 0 {
t.Fatalf("expected fallback route config to be empty, got %q", string(routeData))
}
}
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{},
}
outcome := manager.Apply(context.Background(), "main", "route", []protocol.SupportFile{
{Path: "../escape.crt", Content: "bad"},
})
if outcome.Status != ApplyStatusWarning {
t.Fatalf("expected warning apply outcome, got %#v", outcome)
}
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 TestManagerSyncWAFIPGroupsWritesDeltaRuntimeFile(t *testing.T) {
manager := &Manager{RuntimeConfigDir: t.TempDir()}
if err := manager.SyncWAFIPGroups([]protocol.WAFIPGroup{
{ID: 1, Enabled: true, IPList: []string{"203.0.113.10"}, Checksum: "sum-1"},
}); err != nil {
t.Fatalf("SyncWAFIPGroups failed: %v", err)
}
if err := manager.SyncWAFIPGroups([]protocol.WAFIPGroup{
{ID: 2, Enabled: true, IPList: []string{"198.51.100.10"}, Checksum: "sum-2"},
}); err != nil {
t.Fatalf("SyncWAFIPGroups second delta failed: %v", err)
}
checksums, err := manager.WAFIPGroupChecksums()
if err != nil {
t.Fatalf("WAFIPGroupChecksums failed: %v", err)
}
if checksums["1"] != "sum-1" || checksums["2"] != "sum-2" {
t.Fatalf("expected merged checksums, got %#v", checksums)
}
data, err := os.ReadFile(filepath.Join(manager.RuntimeConfigDir, WAFIPGroupsConfigFileName))
if err != nil {
t.Fatalf("failed to read runtime ip group file: %v", err)
}
text := string(data)
if !strings.Contains(text, "203.0.113.10") || !strings.Contains(text, "198.51.100.10") {
t.Fatalf("expected runtime file to keep both groups, got %s", text)
}
}
func TestObservabilityListenAddress(t *testing.T) {
if got := ObservabilityListenAddress(18081); got != "127.0.0.1:18081" {
t.Fatalf("unexpected default observability listen address: %s", got)
}
if got := ObservabilityListenAddress(18081); got != "127.0.0.1:18081" {
t.Fatalf("unexpected path observability listen address: %s", got)
}
}
@@ -0,0 +1,98 @@
package nginx
const DefaultMimeTypes = `
types {
text/html html htm shtml;
text/css css;
text/xml xml;
image/gif gif;
image/jpeg jpeg jpg;
application/javascript js;
application/atom+xml atom;
application/rss+xml rss;
text/mathml mml;
text/plain txt;
text/vnd.sun.j2me.app-descriptor jad;
text/vnd.wap.wml wml;
text/x-component htc;
image/png png;
image/svg+xml svg svgz;
image/tiff tif tiff;
image/vnd.wap.wbmp wbmp;
image/webp webp;
image/x-icon ico;
image/x-jng jng;
image/x-ms-bmp bmp;
application/font-woff woff;
application/java-archive jar war ear;
application/json json;
application/mac-binhex40 hqx;
application/msword doc;
application/pdf pdf;
application/postscript ps eps ai;
application/rtf rtf;
application/vnd.apple.mpegurl m3u8;
application/vnd.google-earth.kml+xml kml;
application/vnd.google-earth.kmz kmz;
application/vnd.ms-excel xls;
application/vnd.ms-fontobject eot;
application/vnd.ms-powerpoint ppt;
application/vnd.oasis.opendocument.graphics odg;
application/vnd.oasis.opendocument.presentation odp;
application/vnd.oasis.opendocument.spreadsheet ods;
application/vnd.oasis.opendocument.text odt;
application/vnd.openxmlformats-officedocument.presentationml.presentation
pptx;
application/vnd.openxmlformats-officedocument.spreadsheetml.sheet
xlsx;
application/vnd.openxmlformats-officedocument.wordprocessingml.document
docx;
application/vnd.wap.wmlc wmlc;
application/x-7z-compressed 7z;
application/x-cocoa cco;
application/x-java-archive-diff jardiff;
application/x-java-jnlp-file jnlp;
application/x-makeself run;
application/x-perl pl pm;
application/x-pilot prc pdb;
application/x-rar-compressed rar;
application/x-redhat-package-manager rpm;
application/x-sea sea;
application/x-shockwave-flash swf;
application/x-stuffit sit;
application/x-tcl tcl tk;
application/x-x509-ca-cert der pem crt;
application/x-xpinstall xpi;
application/xhtml+xml xhtml;
application/xspf+xml xspf;
application/zip zip;
application/octet-stream bin exe dll;
application/octet-stream deb;
application/octet-stream dmg;
application/octet-stream iso img;
application/octet-stream msi msp msm;
audio/midi mid midi kar;
audio/mpeg mp3;
audio/ogg ogg;
audio/x-m4a m4a;
audio/x-realaudio ra;
video/3gpp 3gpp 3gp;
video/mp2t ts;
video/mp4 mp4;
video/mpeg mpeg mpg;
video/quicktime mov;
video/webm webm;
video/x-flv flv;
video/x-m4v m4v;
video/x-mng mng;
video/x-ms-asf asx asf;
video/x-ms-wmv wmv;
video/x-msvideo avi;
}
`
@@ -0,0 +1,158 @@
package nginx
import "github.com/rain-kl/openflare/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 = ` + openRestyObservabilityWindowTTL + `
local now = ngx.time()
local window_size = ` + openRestyObservabilityWindowSize + `
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 = ` + openRestyObservabilityWindowSize + `
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,607 @@
package nginx
import (
"embed"
"path/filepath"
"strings"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
)
//go:embed pow_static
var powStaticFS embed.FS
const openRestyPowRuntimeLua = `local _M = {}
function _M.check()
local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
local script_path = string.sub(source, 2)
local base_dir = string.match(script_path, "^(.*)/pow/[^/]+%.lua$")
if base_dir and base_dir ~= "" and not string.find(package.path, base_dir, 1, true) then
package.path = base_dir .. "/?.lua;" .. base_dir .. "/?/init.lua;" .. package.path
end
end
local cjson = require "cjson.safe"
local policy = require "pow.policy"
local pow_config_dict = ngx.shared.openflare_pow_config
local pow_sessions = ngx.shared.openflare_pow_sessions
local function session_cookie(value, ttl)
local cookie = "__openflare_pow=" .. value .. "; Path=/; HttpOnly; SameSite=Lax; Max-Age=" .. tostring(ttl)
if ngx.var.scheme == "https" then
cookie = cookie .. "; Secure"
end
return cookie
end
-- Lazy-load pow_config from file; reload when content changes
local function load_pow_config()
local config_paths = {
"__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_config.json",
"/etc/nginx/openflare-lua/waf_config.json",
"/usr/local/openresty/nginx/conf/waf_config.json"
}
for _, config_path in ipairs(config_paths) do
local f = io.open(config_path, "r")
if f then
local content = f:read("*a")
f:close()
local current_hash = ngx.md5(content or "")
if current_hash == pow_config_dict:get("_config_hash") then
return
end
-- Clear old domain/site entries
local old_keys = pow_config_dict:get("_domain_keys")
if old_keys then
for domain in string.gmatch(old_keys, "[^\n]+") do
pow_config_dict:delete(domain)
end
end
local domain_keys = {}
if content and content ~= "" and content ~= "{}" then
local ok, decoded = pcall(cjson.decode, content)
if ok and decoded and decoded.rule_groups and decoded.site_rule_groups then
-- Build rule groups map (group ID -> PoWConfig)
local groups = {}
for _, group in ipairs(decoded.rule_groups) do
if group.pow_enabled then
groups[tostring(group.id)] = group.pow_config
end
end
-- Build site name to pow_config map
for site, group_ids in pairs(decoded.site_rule_groups) do
local pow_config = nil
-- Check custom group IDs first
for _, id in ipairs(group_ids) do
pow_config = groups[tostring(id)]
if pow_config then
break
end
end
-- If not found, check global group IDs
if not pow_config then
for _, group in ipairs(decoded.rule_groups) do
if group.is_global and group.pow_enabled then
pow_config = group.pow_config
break
end
end
end
if pow_config then
pow_config_dict:set(site, cjson.encode({enabled = true, config = pow_config}), 0)
domain_keys[#domain_keys+1] = site
end
end
end
end
pow_config_dict:set("_domain_keys", table.concat(domain_keys, "\n"), 0)
pow_config_dict:set("_config_hash", current_hash, 0)
return
end
end
end
load_pow_config()
local host = ngx.var.host
if not host or host == "" then
return
end
local site = ngx.var.openflare_waf_site or ""
if site == "" then
site = host
end
local config_raw = pow_config_dict:get(site)
if not config_raw then
return
end
local ok, route_config = pcall(cjson.decode, config_raw)
if not ok or not route_config then
return
end
if not route_config.enabled then
return
end
local config = route_config.config or {}
local session_ttl = config.session_ttl or 600
local uri = ngx.var.uri or ""
local ua = ngx.var.http_user_agent or ""
local remote_ip = ngx.var.remote_addr or ""
-- Check whitelist: if matched, skip PoW
local whitelist = config.whitelist or {}
if policy.match_any(remote_ip, ua, uri, whitelist) then
return
end
-- Check blacklist: if matched, require PoW
local blacklist = config.blacklist or {}
local has_blacklist = policy.has_entries(blacklist)
local need_pow = false
if has_blacklist then
need_pow = policy.match_any(remote_ip, ua, uri, blacklist)
else
-- No blacklist means all non-whitelisted need PoW
need_pow = true
end
if not need_pow then
return
end
-- Check valid session cookie
local cookie_val = ngx.var["cookie___openflare_pow"]
if cookie_val and cookie_val ~= "" then
local session_key = host .. ":" .. cookie_val
local session_data = pow_sessions:get(session_key)
if session_data then
pow_sessions:set(session_key, "1", session_ttl)
ngx.header["Set-Cookie"] = session_cookie(cookie_val, session_ttl)
return
end
end
-- If requesting the challenge API endpoints, let them through (handled by content_by_lua)
local anubis_api_prefix = "/.within.website/x/cmd/anubis/api/"
local anubis_static_prefix = "/.within.website/x/cmd/anubis/static/"
if string.sub(uri, 1, #anubis_api_prefix) == anubis_api_prefix then
return
end
if string.sub(uri, 1, #anubis_static_prefix) == anubis_static_prefix then
return
end
-- Render the challenge page through an internal redirect so the browser stays
-- on the originally requested URL instead of seeing a 302 hop.
ngx.req.set_uri_args({
redir = ngx.var.scheme .. "://" .. host .. uri .. (ngx.var.args and ("?" .. ngx.var.args) or ""),
host = host
})
return ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")
end
return _M
`
const openRestyPowCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
local script_path = string.sub(source, 2)
local base_dir = string.match(script_path, "^(.*)/pow/[^/]+%.lua$")
if base_dir and base_dir ~= "" and not string.find(package.path, base_dir, 1, true) then
package.path = base_dir .. "/?.lua;" .. base_dir .. "/?/init.lua;" .. package.path
end
end
return require("pow.runtime").check()
`
const openRestyPowChallengeLua = `local cjson = require "cjson.safe"
local pow_config_dict = ngx.shared.openflare_pow_config
local pow_challenges = ngx.shared.openflare_pow_challenges
local function generate_entropy()
local pieces = {
tostring(ngx.now()),
tostring(ngx.worker.pid()),
tostring(math.random()),
ngx.var.remote_addr or "",
ngx.var.http_user_agent or "",
ngx.var.request_id or "",
}
return table.concat(pieces, ":")
end
local args = ngx.req.get_uri_args()
local host = args["host"] or ngx.var.host or ""
local redir = args["redir"] or ""
local site = ngx.var.openflare_waf_site or ""
if site == "" then
site = host
end
local config_raw = pow_config_dict:get(site)
if not config_raw then
ngx.status = 403
ngx.say("PoW not configured for this site")
return
end
local ok, route_config = pcall(cjson.decode, config_raw)
if not ok or not route_config or not route_config.enabled then
ngx.status = 403
ngx.say("PoW not enabled for this site")
return
end
local config = route_config.config or {}
local difficulty = config.difficulty or 4
local algorithm = config.algorithm or "fast"
local challenge_ttl = config.challenge_ttl or 300
local session_ttl = config.session_ttl or 600
-- Generate challenge data without depending on ngx.random_bytes, which is not
-- available in every OpenResty runtime build.
local entropy = generate_entropy()
local challenge_id = ngx.md5(entropy .. ":id")
local challenge_data = ngx.md5(entropy .. ":data-a") .. ngx.md5(entropy .. ":data-b")
-- Store challenge
local challenge_info = cjson.encode({
data = challenge_data,
difficulty = difficulty,
host = host,
redir = redir,
session_ttl = session_ttl
})
pow_challenges:set(challenge_id, challenge_info, challenge_ttl)
local static_prefix = "/.within.website/x/cmd/anubis/static/"
local accept_lang = ngx.var.http_accept_language or ""
local lang = "en"
if string.find(accept_lang, "zh") then
lang = "zh-CN"
end
local t_title = "Making sure you're not a bot!"
local t_status = "Loading..."
local t_protected = "This site is protected by a Proof-of-Work challenge. Your browser will solve a small puzzle before the upstream response is shown."
local t_why = "Why am I seeing this?"
local t_why_desc = "OpenFlare is asking your browser to complete a lightweight computation to distinguish normal browser traffic from automated abuse. This should finish automatically."
local t_noscript = "JavaScript is required to pass this verification. Please enable JavaScript and reload."
if lang == "zh-CN" then
t_title = "正在确认你是不是机器人!"
t_status = "加载中..."
t_protected = "本网站受工作量证明(Proof-of-Work)挑战保护。在显示源站响应之前,您的浏览器将解决一个微型谜题。"
t_why = "为什么我会看到这个?"
t_why_desc = "OpenFlare 正在要求您的浏览器完成一项轻量级计算,以区分正常的浏览器流量和自动化的恶意请求。这应该会自动完成。"
t_noscript = "很遗憾,您必须启用 JavaScript 才能通过这项验证。请开启 JavaScript 并刷新页面。"
end
ngx.header.content_type = "text/html; charset=utf-8"
ngx.say([[<!DOCTYPE html>
<html lang="]] .. lang .. [[">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<meta name="robots" content="noindex,nofollow">
<title>]] .. t_title .. [[</title>
<link rel="stylesheet" href="]] .. static_prefix .. [[css/xess.css">
<style>
body,html{height:100%;display:flex;justify-content:center;align-items:center;margin-left:auto;margin-right:auto}
.centered-div{text-align:center}
#status{font-variant-numeric:tabular-nums}
#progress{display:none;width:min(20rem,90%);height:2rem;border-radius:1rem;overflow:hidden;margin:1rem 0 2rem;outline-offset:2px;outline:#b16286 solid 4px}
.bar-inner{background-color:#b16286;height:100%;width:0;transition:width .25s ease-in}
</style>
<script id="anubis_version" type="application/json">"openflare-pow"</script>
<script id="anubis_challenge" type="application/json">]] .. cjson.encode({
challenge = {
id = challenge_id,
randomData = challenge_data,
method = algorithm
},
rules = {
difficulty = difficulty,
algorithm = algorithm
}
}) .. [[</script>
<script id="anubis_base_prefix" type="application/json">""</script>
<script id="anubis_public_url" type="application/json">"__openflare_internal__"</script>
</head>
<body id="top">
<main>
<h1 id="title" class="centered-div">]] .. t_title .. [[</h1>
<div class="centered-div">
<img id="image" style="width:100%;max-width:256px;" src="]] .. static_prefix .. [[img/pensive.webp?cacheBuster=openflare-pow">
<p id="status">]] .. t_status .. [[</p>
<p>]] .. t_protected .. [[</p>
<div id="progress" role="progressbar" aria-labelledby="status"><div class="bar-inner"></div></div>
<details>
<summary>]] .. t_why .. [[</summary>
<p>]] .. t_why_desc .. [[</p>
</details>
<noscript><p>]] .. t_noscript .. [[</p></noscript>
</div>
</main>
<script type="module" src="]] .. static_prefix .. [[js/main.mjs"></script>
</body>
</html>]])
`
const openRestyPowVerifyLua = `local cjson = require "cjson.safe"
local pow_challenges = ngx.shared.openflare_pow_challenges
local pow_sessions = ngx.shared.openflare_pow_sessions
local args = ngx.req.get_uri_args()
local challenge_id = args["id"] or ""
local response = args["response"] or ""
local nonce_str = args["nonce"] or ""
local redir = args["redir"] or ""
local elapsed = args["elapsedTime"] or ""
if challenge_id == "" or response == "" or nonce_str == "" then
ngx.status = 400
ngx.header.content_type = "application/json"
ngx.say(cjson.encode({error = "missing parameters"}))
return
end
local nonce = tonumber(nonce_str)
if not nonce then
ngx.status = 400
ngx.header.content_type = "application/json"
ngx.say(cjson.encode({error = "invalid nonce"}))
return
end
-- Get stored challenge
local challenge_raw = pow_challenges:get(challenge_id)
if not challenge_raw then
ngx.status = 410
ngx.header.content_type = "application/json"
ngx.say(cjson.encode({error = "challenge expired or not found"}))
return
end
local ok, challenge_info = pcall(cjson.decode, challenge_raw)
if not ok or not challenge_info then
ngx.status = 500
ngx.header.content_type = "application/json"
ngx.say(cjson.encode({error = "invalid challenge data"}))
return
end
local challenge_data = challenge_info.data or ""
local difficulty = challenge_info.difficulty or 4
local host = challenge_info.host or ngx.var.host or ""
local session_ttl = challenge_info.session_ttl or 600
-- Compute SHA-256(challenge_data + nonce)
local calc_string = challenge_data .. tostring(math.floor(nonce))
local calculated = ngx.sha1_bin ~= nil and "" or ""
-- Use resty.sha256 for proper SHA-256
local sha256 = require "resty.sha256"
local str = require "resty.string"
local hasher = sha256:new()
hasher:update(calc_string)
local hash_bytes = hasher:final()
local hash_hex = str.to_hex(hash_bytes)
-- Verify hash matches response
if hash_hex ~= string.lower(response) then
ngx.status = 403
ngx.header.content_type = "application/json"
ngx.say(cjson.encode({error = "hash mismatch"}))
return
end
-- Verify difficulty (leading zeros in hex)
local prefix = string.rep("0", difficulty)
if string.sub(hash_hex, 1, difficulty) ~= prefix then
ngx.status = 403
ngx.header.content_type = "application/json"
ngx.say(cjson.encode({error = "insufficient difficulty"}))
return
end
-- Invalidate challenge (prevent replay)
pow_challenges:delete(challenge_id)
-- Generate session token
local session_token = str.to_hex(ngx.sha1_bin(challenge_id .. ngx.now() .. tostring(ngx.worker.pid())))
-- Store session
pow_sessions:set(host .. ":" .. session_token, "1", session_ttl)
-- Set cookie. Secure cookies are not sent over HTTP, so only add Secure when
-- the current request itself is HTTPS.
local cookie = "__openflare_pow=" .. session_token .. "; Path=/; HttpOnly; SameSite=Lax; Max-Age=" .. tostring(session_ttl)
if ngx.var.scheme == "https" then
cookie = cookie .. "; Secure"
end
ngx.header["Set-Cookie"] = cookie
if redir ~= "" then
return ngx.redirect(redir)
end
ngx.header.content_type = "application/json"
ngx.say(cjson.encode({ok = true}))
`
const openRestyPowPolicyLua = `local M = {}
local function match_ip(remote_ip, ips)
if not ips or #ips == 0 then return false end
for _, ip in ipairs(ips) do
if ip == remote_ip then
return true
end
end
return false
end
local function match_cidr(remote_ip, cidrs)
if not cidrs or #cidrs == 0 then return false end
for _, cidr in ipairs(cidrs) do
local m, err = ngx.re.match(cidr, "^(\\\\d{1,3}\\\\.\\\\d{1,3}\\\\.\\\\d{1,3}\\\\.\\\\d{1,3})/(\\\\d{1,2})$")
if m then
local mask_bits = tonumber(m[2])
if mask_bits and mask_bits >= 0 and mask_bits <= 32 then
local function ip_to_num(ip_str)
local parts = {}
for part in string.gmatch(ip_str, "%d+") do
parts[#parts+1] = tonumber(part) or 0
end
if #parts ~= 4 then return 0 end
return parts[1]*16777216 + parts[2]*65536 + parts[3]*256 + parts[4]
end
local remote_num = ip_to_num(remote_ip)
local net_num = ip_to_num(m[1])
if mask_bits == 0 then
return true
end
local mask = math.floor(2^(32 - mask_bits))
mask = 4294967296 - mask
if bit.band(remote_num, mask) == bit.band(net_num, mask) then
return true
end
end
end
end
return false
end
local function match_path(uri, patterns)
if not patterns or #patterns == 0 then return false end
for _, pattern in ipairs(patterns) do
local ok, match = pcall(ngx.re.match, uri, "^" .. ngx.re.gsub(pattern, "([%^%$%(%)%%%.%[%]%+%-%?])", function(c)
if c == "*" then return ".*" end
return "%" .. c
end) .. "$", "i")
if ok and match then
return true
end
end
return false
end
local function match_path_regex(uri, patterns)
if not patterns or #patterns == 0 then return false end
for _, pattern in ipairs(patterns) do
local ok, match = pcall(ngx.re.match, uri, pattern)
if ok and match then
return true
end
end
return false
end
local function match_ua(ua, patterns)
if not patterns or #patterns == 0 then return false end
for _, pattern in ipairs(patterns) do
if ua and string.find(ua, pattern, 1, true) then
return true
end
end
return false
end
function M.match_any(remote_ip, ua, uri, list)
if not list then return false end
if match_ip(remote_ip, list.ips) then return true end
if match_cidr(remote_ip, list.ip_cidrs) then return true end
if match_path(uri, list.paths) then return true end
if match_path_regex(uri, list.path_regexes) then return true end
if match_ua(ua, list.user_agents) then return true end
return false
end
function M.has_entries(list)
if not list then return false end
return (#(list.ips or {}) + #(list.ip_cidrs or {}) + #(list.paths or {}) + #(list.path_regexes or {}) + #(list.user_agents or {})) > 0
end
return M
`
func ManagedPowLuaFiles() []protocol.SupportFile {
return []protocol.SupportFile{
{Path: "pow/runtime.lua", Content: openRestyPowRuntimeLua},
{Path: "pow/check.lua", Content: openRestyPowCheckLua},
{Path: "pow/challenge.lua", Content: openRestyPowChallengeLua},
{Path: "pow/verify.lua", Content: openRestyPowVerifyLua},
{Path: "pow/policy.lua", Content: openRestyPowPolicyLua},
}
}
func ManagedPowStaticFiles() ([]protocol.SupportFile, error) {
var files []protocol.SupportFile
entries, err := powStaticFS.ReadDir("pow_static")
if err != nil {
return nil, err
}
var walk func(dir string) error
walk = func(dir string) error {
entries, err := powStaticFS.ReadDir(dir)
if err != nil {
return err
}
for _, entry := range entries {
fullPath := filepath.Join(dir, entry.Name())
if entry.IsDir() {
if err := walk(fullPath); err != nil {
return err
}
continue
}
data, err := powStaticFS.ReadFile(fullPath)
if err != nil {
return err
}
// Convert pow_static/css/xess.css -> pow/static/css/xess.css
relPath := strings.TrimPrefix(fullPath, "pow_static/")
files = append(files, protocol.SupportFile{
Path: "pow/static/" + relPath,
Content: string(data),
})
}
return nil
}
for _, entry := range entries {
fullPath := filepath.Join("pow_static", entry.Name())
if entry.IsDir() {
if err := walk(fullPath); err != nil {
return nil, err
}
} else {
data, err := powStaticFS.ReadFile(fullPath)
if err != nil {
return nil, err
}
relPath := strings.TrimPrefix(fullPath, "pow_static/")
files = append(files, protocol.SupportFile{
Path: "pow/static/" + relPath,
Content: string(data),
})
}
}
return files, nil
}
@@ -0,0 +1,7 @@
@font-face {
font-family: "Podkova";
font-style: normal;
font-weight: 400 800;
font-display: swap;
src: url("podkova.woff2") format("woff2");
}
@@ -0,0 +1,149 @@
:root {
--body-sans-font: Geist, sans-serif;
--body-preformatted-font: Iosevka Curly Iaso, monospace;
--body-title-font: Podkova, serif;
--background: #1d2021;
--text: #f9f5d7;
--text-selection: #d3869b;
--preformatted-background: #3c3836;
--link-foreground: #b16286;
--link-background: #282828;
--blockquote-border-left: 1px solid #bdae93;
--progress-bar-outline: #b16286 solid 4px;
--progress-bar-fill: #b16286;
}
@media (prefers-color-scheme: light) {
:root {
--background: #f9f5d7;
--text: #1d2021;
--text-selection: #d3869b;
--preformatted-background: #ebdbb2;
--link-foreground: #b16286;
--link-background: #fbf1c7;
--blockquote-border-left: 1px solid #655c54;
}
}
@font-face {
font-family: "Geist";
font-style: normal;
font-weight: 100 900;
font-display: swap;
src: url("./static/geist.woff2") format("woff2");
}
@font-face {
font-family: "Podkova";
font-style: normal;
font-weight: 400 800;
font-display: swap;
src: url("./static/podkova.woff2") format("woff2");
}
@font-face {
font-family: "Iosevka Curly";
font-style: monospace;
font-display: swap;
src: url("./static/iosevka-curly.woff2") format("woff2");
}
main {
font-family: var(--body-sans-font);
max-width: 50rem;
padding: 2rem;
margin: auto;
}
::selection {
background: var(--text-selection);
}
body {
background: var(--background);
color: var(--text);
}
body,
html {
height: 100%;
display: flex;
justify-content: center;
align-items: center;
margin-left: auto;
margin-right: auto;
}
.centered-div {
text-align: center;
}
#status {
font-variant-numeric: tabular-nums;
}
.centered-div {
text-align: center;
}
#status {
font-variant-numeric: tabular-nums;
}
#progress {
display: none;
width: min(20rem, 90%);
height: 2rem;
border-radius: 1rem;
overflow: hidden;
margin: 1rem 0 2rem;
outline-offset: 2px;
outline: var(--progress-bar-outline);
}
.bar-inner {
background-color: var(--progress-bar-fill);
height: 100%;
width: 0;
transition: width 0.25s ease-in;
}
@media (prefers-reduced-motion: no-preference) {
.bar-inner {
transition: width 0.25s ease-in;
}
}
pre {
background-color: var(--preformatted-background);
padding: 1em;
border: 0;
font-family: var(--body-preformatted-font);
}
a,
a:active,
a:visited {
color: var(--link-foreground);
background-color: var(--link-background);
}
h1,
h2,
h3,
h4,
h5 {
margin-bottom: 0.1rem;
font-family: var(--body-title-font);
}
blockquote {
border-left: var(--blockquote-border-left);
margin: 0.5em 10px;
padding: 0.5em 10px;
}
footer {
text-align: center;
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 30 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 28 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 26 KiB

@@ -0,0 +1,32 @@
/*
@licstart The following is the entire license notice for the
JavaScript code in this page.
Copyright (c) 2025 Xe Iaso <xe.iaso@techaro.lol>
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in
all copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
THE SOFTWARE.
Includes code from https://github.com/aws/aws-sdk-js-crypto-helpers which is
used under the terms of the Apache 2 license.
@licend The above is the entire license notice
for the JavaScript code in this page.
*/
(()=>{var k=()=>navigator.hardwareConcurrency!==void 0?navigator.hardwareConcurrency:1;function n(c,b,w=5,e=null,g,u=Math.trunc(Math.max(k()/2,1))){console.debug("fast algo");let s="purejs";return window.isSecureContext&&(s="webcrypto"),(navigator.userAgent.includes("Firefox")||navigator.userAgent.includes("Goanna"))&&(console.log("Firefox detected, using pure-JS fallback"),s="purejs"),new Promise((p,l)=>{let m=`${c.basePrefix}/.within.website/x/cmd/anubis/static/js/worker/sha256-${s}.mjs?cacheBuster=${c.version}`,f=[],d=!1,a=()=>{console.log("PoW aborted"),i(),l(new DOMException("Aborted","AbortError"))},i=()=>{d||(d=!0,f.forEach(r=>r.terminate()),e?.removeEventListener("abort",a))};if(e!=null){if(e.aborted)return a();e.addEventListener("abort",a,{once:!0})}for(let r=0;r<u;r++){let t=new Worker(m);t.onmessage=o=>{typeof o.data=="number"?g?.(o.data):(i(),p(o.data))},t.onerror=o=>{i(),l(o)},t.postMessage({data:b,difficulty:w,nonce:r,threads:u}),f.push(t)}})}var P={fast:n,slow:n};})();
//# sourceMappingURL=index.mjs.map
@@ -0,0 +1,32 @@
/*
@licstart The following is the entire license notice for the
JavaScript code in this page.
Copyright (c) 2025 Xe Iaso <xe.iaso@techaro.lol>
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in
all copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
THE SOFTWARE.
Includes code from https://github.com/aws/aws-sdk-js-crypto-helpers which is
used under the terms of the Apache 2 license.
@licend The above is the entire license notice
for the JavaScript code in this page.
*/
(()=>{var I=()=>navigator.hardwareConcurrency!==void 0?navigator.hardwareConcurrency:1;function _(e,n,s=5,o=null,i,u=Math.trunc(Math.max(I()/2,1))){console.debug("fast algo");let a="purejs";return window.isSecureContext&&(a="webcrypto"),(navigator.userAgent.includes("Firefox")||navigator.userAgent.includes("Goanna"))&&(console.log("Firefox detected, using pure-JS fallback"),a="purejs"),new Promise((E,x)=>{let M=`${e.basePrefix}/.within.website/x/cmd/anubis/static/js/worker/sha256-${a}.mjs?cacheBuster=${e.version}`,p=[],d=!1,b=()=>{console.log("PoW aborted"),h(),x(new DOMException("Aborted","AbortError"))},h=()=>{d||(d=!0,p.forEach(c=>c.terminate()),o?.removeEventListener("abort",b))};if(o!=null){if(o.aborted)return b();o.addEventListener("abort",b,{once:!0})}for(let c=0;c<u;c++){let g=new Worker(M);g.onmessage=m=>{typeof m.data=="number"?i?.(m.data):(h(),E(m.data))},g.onerror=m=>{h(),x(m)},g.postMessage({data:n,difficulty:s,nonce:c,threads:u}),p.push(g)}})}var j={fast:_,slow:_};var v=(e="",n={})=>{let s=new URL(e,window.location.href);return Object.entries(n).forEach(([o,i])=>s.searchParams.set(o,i)),s.toString()},L=e=>{let n=document.getElementById(e);return n===null?null:JSON.parse(n.textContent)},k=(e,n,s)=>v(`${s}/.within.website/x/cmd/anubis/static/img/${e}.webp`,{cacheBuster:n});var W=async()=>document.documentElement.lang,S=async e=>{let n=L("anubis_base_prefix");if(n!==null)try{return await(await fetch(`${n}/.within.website/x/cmd/anubis/static/locales/${e}.json`)).json()}catch(s){if(console.warn(`Failed to load translations for ${e}, falling back to English`),e!=="en")return await S("en");throw s}},C=()=>{let e=L("anubis_public_url");if(e!==null)return e&&window.location.href.startsWith(e)?new URLSearchParams(window.location.search).get("redir"):window.location.href},$={},D,A=async()=>{D=await W(),$=await S(D)},r=e=>$[`js_${e}`]||$[e]||e;(async()=>{await A();let e=[{name:"Web Workers",msg:r("web_workers_error"),value:window.Worker},{name:"Cookies",msg:r("cookies_error"),value:navigator.cookieEnabled}],n=document.getElementById("status"),s=document.getElementById("image"),o=document.getElementById("title"),i=document.getElementById("progress"),u=L("anubis_version"),a=L("anubis_base_prefix"),E=document.querySelector("details"),x=!1;E&&E.addEventListener("toggle",()=>{E.open&&(x=!0)});let M=({titleMsg:l,statusMsg:f,imageSrc:w})=>{o.innerHTML=l,n.innerHTML=f,s.src=w,i.style.display="none"};n.innerHTML=r("calculating");for(let{value:l,name:f,msg:w}of e)if(!l){M({titleMsg:`${r("missing_feature")} ${f}`,statusMsg:w,imageSrc:k("reject",u,a)});return}let{challenge:p,rules:d}=L("anubis_challenge"),b=j[d.algorithm];if(!b){M({titleMsg:r("challenge_error"),statusMsg:r("challenge_error_msg"),imageSrc:k("reject",u,a)});return}n.innerHTML=`${r("calculating_difficulty")} ${d.difficulty}, `,i.style.display="inline-block";let h=document.createTextNode(`${r("speed")} 0kH/s`);n.appendChild(h);let c=0,g=!1,m=Math.pow(16,-d.difficulty);try{let l=Date.now(),{hash:f,nonce:w}=await b({basePrefix:a,version:u},p.randomData,d.difficulty,null,t=>{let y=Date.now()-l;y-c>1e3&&(c=y,h.data=`${r("speed")} ${(t/y).toFixed(3)}kH/s`);let T=Math.pow(1-m,t),P=(1-Math.pow(T,2))*100;i["aria-valuenow"]=P,i.firstElementChild!==null&&(i.firstElementChild.style.width=`${P}%`),T<.1&&!g&&(n.append(document.createElement("br"),document.createTextNode(r("verification_longer"))),g=!0)}),H=Date.now();if(console.log({hash:f,nonce:w}),x){let y=function(){let T=C();window.location.replace(v(`${a}/.within.website/x/cmd/anubis/api/pass-challenge`,{id:p.id,response:f,nonce:w,redir:T,elapsedTime:H-l}))},t=document.getElementById("progress");t.style.display="flex",t.style.alignItems="center",t.style.justifyContent="center",t.style.height="2rem",t.style.borderRadius="1rem",t.style.cursor="pointer",t.style.background="#b16286",t.style.color="white",t.style.fontWeight="bold",t.style.outline="4px solid #b16286",t.style.outlineOffset="2px",t.style.width="min(20rem, 90%)",t.style.margin="1rem auto 2rem",t.innerHTML=r("finished_reading"),t.onclick=y,setTimeout(y,3e4)}else{let t=C();window.location.replace(v(`${a}/.within.website/x/cmd/anubis/api/pass-challenge`,{id:p.id,response:f,nonce:w,redir:t,elapsedTime:H-l}))}}catch(l){M({titleMsg:r("calculation_error"),statusMsg:`${r("calculation_error_msg")} ${l.message}`,imageSrc:k("reject",u,a)})}})();})();
//# sourceMappingURL=main.mjs.map
File diff suppressed because one or more lines are too long
@@ -0,0 +1,32 @@
/*
@licstart The following is the entire license notice for the
JavaScript code in this page.
Copyright (c) 2025 Xe Iaso <xe.iaso@techaro.lol>
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in
all copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
THE SOFTWARE.
Includes code from https://github.com/aws/aws-sdk-js-crypto-helpers which is
used under the terms of the Apache 2 license.
@licend The above is the entire license notice
for the JavaScript code in this page.
*/
(()=>{var h=new TextEncoder,y=async e=>{let s=h.encode(e);return await crypto.subtle.digest("SHA-256",s)},g=e=>e.reduce((s,a)=>s+a.toString(16).padStart(2,"0"),"");addEventListener("message",async({data:e})=>{let{data:s,difficulty:a,threads:d}=e,t=e.nonce,f=t===0,o=0,c=Math.floor(a/2),l=a%2!==0;for(;;){let u=await y(s+t),i=new Uint8Array(u),r=!0;for(let n=0;n<c;n++)if(i[n]!==0){r=!1;break}if(r&&l&&i[c]>>4!==0&&(r=!1),r){let n=g(i);postMessage({hash:n,data:s,difficulty:a,nonce:t});return}t+=d,o++,t%1!==0&&(t=Math.trunc(t)),f&&(o&1023)===0&&postMessage(t)}});})();
//# sourceMappingURL=sha256-webcrypto.mjs.map
@@ -0,0 +1,66 @@
{
"loading": "Loading...",
"why_am_i_seeing": "Why am I seeing this?",
"protected_by": "Protected by",
"protected_from": "From",
"made_with": "Made with ❤️ in 🇨🇦",
"mascot_design": "Mascot design by",
"ai_companies_explanation": "You are seeing this because the administrator of this website has set up Anubis to protect the server against the scourge of AI companies aggressively scraping websites. This can and does cause downtime for the websites, which makes their resources inaccessible for everyone.",
"anubis_compromise": "Anubis is a compromise. Anubis uses a Proof-of-Work scheme in the vein of Hashcash, a proposed proof-of-work scheme for reducing email spam. The idea is that at individual scales the additional load is ignorable, but at mass scraper levels it adds up and makes scraping much more expensive.",
"hack_purpose": "Ultimately, this is a placeholder solution so that more time can be spent on fingerprinting and identifying headless browsers (EG: via how they do font rendering) so that the challenge proof of work page doesn't need to be presented to users that are much more likely to be legitimate.",
"simplified_explanation": "This is a measure against bots and malicious requests similar to a CAPTCHA. However, instead of having to do work yourself, your browser is given a calculation task that it has to solve to ensure that it is a valid client. This concept is called <a href=\"https://en.wikipedia.org/wiki/Proof_of_work\">Proof of Work</a>. The task is calculated in a few seconds and you are granted access to the website. Thank you for your understanding and patience.",
"jshelter_note": "Please note that Anubis requires the use of modern JavaScript features that plugins like JShelter will disable. Please disable JShelter or other such plugins for this domain.",
"version_info": "This website is running Anubis version",
"try_again": "Try again",
"go_home": "Go home",
"contact_webmaster": "or if you believe you should not be blocked, please contact the webmaster at",
"connection_security": "Please wait a moment while we ensure the security of your connection.",
"javascript_required": "Sadly, you must enable JavaScript to get past this challenge. This is required because AI companies have changed the social contract around how website hosting works. A no-JS solution is a work-in-progress.",
"benchmark_requires_js": "Running the benchmark tool requires JavaScript to be enabled.",
"difficulty": "Difficulty:",
"algorithm": "Algorithm:",
"compare": "Compare:",
"time": "Time",
"iters": "Iters",
"time_a": "Time A",
"iters_a": "Iters A",
"time_b": "Time B",
"iters_b": "Iters B",
"static_check_endpoint": "This is just a check endpoint for your reverse proxy to use.",
"authorization_required": "Authorization required",
"cookies_disabled": "Your browser is configured to disable cookies. Anubis requires cookies for the legitimate interest of making sure you are a valid client. Please enable cookies for this domain",
"access_denied": "Access Denied: error code",
"dronebl_entry": "DroneBL reported an entry",
"see_dronebl_lookup": "see",
"internal_server_error": "Internal Server Error: administrator has misconfigured Anubis. Please contact the administrator and ask them to look for the logs around",
"invalid_redirect": "Invalid redirect",
"redirect_not_parseable": "Redirect URL not parseable",
"redirect_domain_not_allowed": "Redirect domain not allowed",
"missing_required_forwarded_headers": "Missing required X-Forwarded-* headers",
"failed_to_sign_jwt": "failed to sign JWT",
"invalid_invocation": "Invalid invocation of MakeChallenge",
"client_error_browser": "Client Error: Please ensure your browser is up to date and try again later.",
"oh_noes": "Oh noes!",
"benchmarking_anubis": "Benchmarking Anubis!",
"you_are_not_a_bot": "You are not a bot!",
"making_sure_not_bot": "Making sure you're not a bot!",
"celphase": "CELPHASE",
"js_web_crypto_error": "Your browser doesn't have a functioning web.crypto element. Are you viewing this over a secure context?",
"js_web_workers_error": "Your browser doesn't support web workers (Anubis uses this to avoid freezing your browser). Do you have a plugin like JShelter installed?",
"js_cookies_error": "Your browser doesn't store cookies. Anubis uses cookies to determine which clients have passed challenges by storing a signed token in a cookie. Please enable storing cookies for this domain. The names of the cookies Anubis stores may vary without notice. Cookie names and values are not part of the public API.",
"js_context_not_secure": "Your context is not secure!",
"js_context_not_secure_msg": "Try connecting over HTTPS or let the admin know to set up HTTPS. For more information, see <a href=\"https://developer.mozilla.org/en-US/docs/Web/Security/Secure_Contexts#when_is_a_context_considered_secure\">MDN</a>.",
"js_calculating": "Calculating...",
"js_missing_feature": "Missing feature",
"js_challenge_error": "Challenge error!",
"js_challenge_error_msg": "Failed to resolve check algorithm. You may want to reload the page.",
"js_calculating_difficulty": "Calculating...<br/>Difficulty:",
"js_speed": "Speed:",
"js_verification_longer": "Verification is taking longer than expected. Please do not refresh the page.",
"js_success": "Success!",
"js_done_took": "Done! Took",
"js_iterations": "iterations",
"js_finished_reading": "I've finished reading, continue →",
"js_calculation_error": "Calculation error!",
"js_calculation_error_msg": "Failed to calculate challenge:"
}
@@ -0,0 +1,66 @@
{
"loading": "加载中...",
"why_am_i_seeing": "为什么我会看到这个?",
"protected_by": "本网站由",
"protected_from": "保护,来自",
"made_with": "在 🇨🇦 用 ❤️ 制作",
"mascot_design": "吉祥物由",
"ai_companies_explanation": "您会看到这个画面,是因为网站管理员启用了 Anubis 来保护服务器,避免 AI 公司大量爬取网站内容。这类行为会导致网站崩溃,让所有用户都无法正常访问资源。",
"anubis_compromise": "Anubis 是一种折中做法。它采用了类似 Hashcash 的工作量证明机制(Proof-of-Work),该机制最初是为了减少垃圾邮件而提出。其核心概念是:对个别用户而言,额外的计算负担可以忽略,但对大规模爬虫来说,累积起来的成本将大幅增加,从而让爬取行为变得更困难。",
"hack_purpose": "最终,这是一个占位符解决方案,以便将更多时间用于指纹识别和识别无头浏览器(例如:通过它们如何进行字体渲染),从而无需向更可能是合法用户的用户呈现挑战工作量证明页面。",
"jshelter_note": "请注意,Anubis 需要使用现代 JavaScript 功能,而像 JShelter 这类插件可能会阻挡这些功能。请为此域名停用 JShelter 或类似的插件。",
"version_info": "这个网站正在运行的 Anubis 版本为",
"try_again": "再试一次",
"go_home": "返回首页",
"contact_webmaster": "或者您觉得您不应该被封锁,请联系网站管理员于",
"connection_security": "请稍等,我们需要在继续之前检查您的连接安全性。",
"javascript_required": "很遗憾,您必须启用 JavaScript 才能通过这项验证。这是因为 AI 公司已经改变了网站托管的社会契约,因此我们必须采取这样的保护机制。无需 JavaScript 的解决方案仍在开发中。",
"benchmark_requires_js": "运行基准测试工具需要启用 JavaScript。",
"difficulty": "难度:",
"algorithm": "算法:",
"compare": "比较:",
"time": "时间",
"iters": "迭代",
"time_a": "时间 A",
"iters_a": "迭代 A",
"time_b": "时间 B",
"iters_b": "迭代 B",
"static_check_endpoint": "这是提供给您的反向代理服务器使用的检查端点。",
"authorization_required": "需要认证",
"cookies_disabled": "您的浏览器目前已禁用 Cookie,为了确认您是合法用户,Anubis 需要启用 Cookie。 请您为此域名启用 Cookie",
"access_denied": "拒绝访问:错误代码",
"dronebl_entry": "DroneBL 报告了一条记录",
"see_dronebl_lookup": "见",
"internal_server_error": "内部服务器错误:管理员错误地配置了 Anubis。 请联系管理员要求他们检查日志",
"invalid_redirect": "无效的重定向",
"redirect_not_parseable": "重定向 URL 无法解析",
"redirect_domain_not_allowed": "重定向的域名并不允许",
"failed_to_sign_jwt": "签署 JWT 失败",
"invalid_invocation": "无效的 MakeChallenge 调用",
"client_error_browser": "客户端错误:请确保您的浏览器是最新版本并稍候再试。",
"oh_noes": "哎呀糟糕了!",
"benchmarking_anubis": "正在进行 Anubis 性能测试!",
"you_are_not_a_bot": "你不是机器人!",
"making_sure_not_bot": "正在确认你是不是机器人!",
"celphase": "CELPHASE 设计",
"js_web_crypto_error": "您的浏览器无法正常使用 web.crypto 组件。您是否通过安全连接(HTTPS)查看此网站?",
"js_web_workers_error": "您的浏览器并不支持 Web workers (Anubis 使用这个来避免冻结您的浏览器 )您有安装像是 JShelter 之类的插件吗?",
"js_cookies_error": "您的浏览器无法存储 Cookie。 Anubis 会使用 Cookie 存储签署的凭证,以判断用户是否已通过验证。请为此域名启用 Cookie 存储功能。 请注意,Anubis 存储的 Cookie 名称可能会变动,且其名称与内容不属于公开 API 的一部分。",
"js_context_not_secure": "您的内容并不安全",
"js_context_not_secure_msg": "请尝试使用 HTTPS 连接,或联系网站管理员设置 HTTPS。更多信息请参见 <a href=\"https://developer.mozilla.org/en-US/docs/Web/Security/Secure_Contexts#when_is_a_context_considered_secure\">MDN</a>。",
"js_calculating": "计算中...",
"js_missing_feature": "缺少功能",
"js_challenge_error": "挑战错误!",
"js_challenge_error_msg": "解决检查算法失败。 您可能会想要刷新页面。",
"js_calculating_difficulty": "计算中...<br/>难度:",
"js_speed": "速度:",
"js_verification_longer": "验证所花的时间高于预期。 请不要刷新页面。",
"js_success": "成功!",
"js_done_took": "完成! 花费",
"js_iterations": "迭代",
"js_finished_reading": "我读完了,继续 →",
"js_calculation_error": "计算错误!",
"js_calculation_error_msg": "计算挑战失败:",
"missing_required_forwarded_headers": "缺少必要的 X-Forwarded-* 头",
"simplified_explanation": "这是一种类似于验证码的措施,用于防止机器人和恶意请求。但是,您无需自己动手,您的浏览器会收到一个计算任务,必须解决该任务以确保它是有效的客户端。这个概念称为<a href=\"https://en.wikipedia.org/wiki/Proof_of_work\">工作量证明</a>。该任务在几秒钟内计算完毕,您将被授予访问网站的权限。感谢您的理解和耐心。"
}
@@ -0,0 +1,310 @@
package nginx
import "github.com/rain-kl/openflare/openflare-agent/internal/protocol"
const openRestyWAFRuntimeLua = `local _M = {}
function _M.check()
local cjson = require "cjson.safe"
local config_dict = ngx.shared.openflare_waf_config
local function read_file(path)
local f = io.open(path, "r")
if not f then
return nil
end
local content = f:read("*a")
f:close()
return content
end
local function load_config()
local paths = {
"__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_config.json",
"/etc/nginx/openflare-lua/waf_config.json",
"/usr/local/openresty/nginx/conf/waf_config.json"
}
for _, path in ipairs(paths) do
local content = read_file(path)
if content and content ~= "" then
local hash = ngx.md5(content)
if config_dict:get("_config_hash") == hash then
local cached = config_dict:get("_config_json")
if cached then
local decoded = cjson.decode(cached)
if decoded then
return decoded
end
end
end
local decoded = cjson.decode(content)
if decoded then
config_dict:set("_config_hash", hash, 0)
config_dict:set("_config_json", content, 0)
return decoded
end
end
end
return nil
end
local function load_ip_groups()
local paths = {
"__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_ip_groups.json",
"/etc/nginx/openflare-lua/waf_ip_groups.json",
"/usr/local/openresty/nginx/conf/waf_ip_groups.json"
}
for _, path in ipairs(paths) do
local content = read_file(path)
if content and content ~= "" then
local hash = ngx.md5(content)
if config_dict:get("_ip_groups_hash") == hash then
local cached = config_dict:get("_ip_groups_json")
if cached then
local decoded = cjson.decode(cached)
if decoded then
return decoded
end
end
end
local decoded = cjson.decode(content)
if decoded then
config_dict:set("_ip_groups_hash", hash, 0)
config_dict:set("_ip_groups_json", content, 0)
return decoded
end
end
end
return { groups = {} }
end
local function list_contains(items, value)
if not items or type(items) ~= "table" or not value or value == "" then
return false
end
for _, item in ipairs(items) do
if item == value then
return true
end
end
return false
end
local function table_has_items(items)
return type(items) == "table" and #items > 0
end
local function parse_ipv4(value)
local a, b, c, d = string.match(value or "", "^(%d+)%.(%d+)%.(%d+)%.(%d+)$")
if not a then
return nil
end
a, b, c, d = tonumber(a), tonumber(b), tonumber(c), tonumber(d)
if a > 255 or b > 255 or c > 255 or d > 255 then
return nil
end
return ((a * 256 + b) * 256 + c) * 256 + d
end
local function ipv4_in_cidr(ip, cidr)
local base, bits = string.match(cidr or "", "^([^/]+)/(%d+)$")
if not base then
return false
end
bits = tonumber(bits)
if not bits or bits < 0 or bits > 32 then
return false
end
local ip_num = parse_ipv4(ip)
local base_num = parse_ipv4(base)
if not ip_num or not base_num then
return false
end
if bits == 0 then
return true
end
local mask = 4294967295 - (2 ^ (32 - bits) - 1)
return (ip_num - (ip_num % (2 ^ (32 - bits)))) == (base_num - (base_num % (2 ^ (32 - bits))))
end
local function ip_matches(items, ip)
if not items or type(items) ~= "table" or not ip or ip == "" then
return false
end
for _, item in ipairs(items) do
if item == ip then
return true
end
if string.find(item, "/", 1, true) and ipv4_in_cidr(ip, item) then
return true
end
end
return false
end
local function ip_matches_group_ids(group_ids, ip, ip_groups_config)
if not group_ids or type(group_ids) ~= "table" or not ip or ip == "" then
return false
end
local groups = (ip_groups_config or {}).groups or {}
for _, id in ipairs(group_ids) do
local group = groups[tostring(id)]
if group and group.enabled and ip_matches(group.ip_list, ip) then
return true
end
end
return false
end
local function lookup_country(ip)
local ok, maxminddb = pcall(require, "resty.maxminddb")
if not ok or not maxminddb then
return nil
end
local paths = {
"__OPENFLARE_RUNTIME_CONFIG_DIR__/GeoLite2-Country.mmdb",
"/etc/openflare/GeoLite2-Country.mmdb",
"/usr/local/share/openflare/GeoLite2-Country.mmdb"
}
for _, path in ipairs(paths) do
local opened = pcall(maxminddb.init, path)
if opened then
local res, err = maxminddb.lookup(ip)
if res and res.country and res.country.iso_code then
return string.upper(res.country.iso_code)
end
end
end
return nil
end
local function group_by_id(config)
local result = {}
for _, group in ipairs(config.rule_groups or {}) do
result[tostring(group.id)] = group
end
return result
end
local function active_groups(config, groups)
local site = ngx.var.openflare_waf_site or ""
local ids = (config.site_rule_groups or {})[site]
local result = {}
for _, group in ipairs(config.rule_groups or {}) do
if group.is_global then
result[#result + 1] = group
end
end
if ids then
local by_id = group_by_id(config)
for _, id in ipairs(ids) do
local group = by_id[tostring(id)]
if group and not group.is_global then
result[#result + 1] = group
end
end
end
return result
end
local function exit_with_group(group)
ngx.ctx.openflare_waf_blocked = true
ngx.status = tonumber(group.block_status_code) or 418
local body = group.block_response_body or ""
if body ~= "" then
ngx.header["Content-Type"] = "text/html; charset=utf-8"
ngx.say(body)
end
return ngx.exit(ngx.status)
end
local function first_allowlist_group(groups)
for _, group in ipairs(groups) do
if table_has_items(group.ip_whitelist)
or table_has_items(group.ip_whitelist_group_ids)
or table_has_items(group.country_whitelist) then
return group
end
end
return nil
end
local config = load_config()
if not config then
if config_dict:add("_missing_config_logged", true, 60) then
ngx.log(ngx.WARN, "openflare waf config is missing or invalid; requests will be allowed")
end
return
end
local ip = ngx.var.remote_addr or ""
local groups = active_groups(config)
local ip_groups_config = load_ip_groups()
if #groups == 0 then
if config_dict:add("_empty_groups_logged", true, 60) then
ngx.log(ngx.WARN, "openflare waf has no active rule group for site: ", ngx.var.openflare_waf_site or "")
end
return
end
for _, group in ipairs(groups) do
if ip_matches(group.ip_whitelist, ip) or ip_matches_group_ids(group.ip_whitelist_group_ids, ip, ip_groups_config) then
return
end
end
local country = nil
for _, group in ipairs(groups) do
if type(group.country_whitelist) == "table" and #group.country_whitelist > 0 then
country = country or lookup_country(ip)
if list_contains(group.country_whitelist, country) then
return
end
end
end
local allowlist_group = first_allowlist_group(groups)
if allowlist_group then
return exit_with_group(allowlist_group)
end
for _, group in ipairs(groups) do
if ip_matches(group.ip_blacklist, ip) or ip_matches_group_ids(group.ip_blacklist_group_ids, ip, ip_groups_config) then
return exit_with_group(group)
end
end
for _, group in ipairs(groups) do
if type(group.country_blacklist) == "table" and #group.country_blacklist > 0 then
country = country or lookup_country(ip)
if list_contains(group.country_blacklist, country) then
return exit_with_group(group)
end
end
end
return "ok"
end
return _M
`
const openRestyWAFCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
local script_path = string.sub(source, 2)
local base_dir = string.match(script_path, "^(.*)/waf/[^/]+%.lua$")
if base_dir and base_dir ~= "" and not string.find(package.path, base_dir, 1, true) then
package.path = base_dir .. "/?.lua;" .. base_dir .. "/?/init.lua;" .. package.path
end
end
return require("waf.runtime").check()
`
func ManagedWAFLuaFiles() []protocol.SupportFile {
return []protocol.SupportFile{
{Path: "waf/runtime.lua", Content: openRestyWAFRuntimeLua},
{Path: "waf/check.lua", Content: openRestyWAFCheckLua},
}
}
@@ -0,0 +1,404 @@
package observability
import (
"bufio"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"syscall"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
"github.com/rain-kl/openflare/openflare-agent/internal/state"
)
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) *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 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 BuildOpenrestyObservation(managed *ManagedOpenRestyMetrics) *protocol.NodeOpenrestyObservation {
if managed == nil {
return nil
}
return &protocol.NodeOpenrestyObservation{
CapturedAtUnix: time.Now().UTC().Unix(),
OpenrestyRxBytes: managed.OpenrestyRxBytes,
OpenrestyTxBytes: managed.OpenrestyTxBytes,
OpenrestyConnections: managed.OpenrestyConnections,
}
}
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,130 @@
package observability
import (
"encoding/json"
"fmt"
"io"
"net/http"
"regexp"
"strconv"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
)
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"
"github.com/rain-kl/openflare/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,362 @@
package observability
import (
"bufio"
"encoding/json"
"errors"
"io"
"log/slog"
"os"
"regexp"
"sort"
"strconv"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
"github.com/rain-kl/openflare/openflare-agent/internal/state"
)
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)
var 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 func(file *os.File) {
err := file.Close()
if err != nil {
slog.Error("failed to close access log file", "error", err)
}
}(file)
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.AccessLogPath) == "" {
return ""
}
return cfg.AccessLogPath
}
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
}
const accessLogPathMaxRunes = 100
func normalizeAccessLogPath(value string) string {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return ""
}
if strings.HasPrefix(trimmed, "http://") || strings.HasPrefix(trimmed, "https://") {
return truncateAccessLogPath(trimmed)
}
if strings.HasPrefix(trimmed, "/") {
return truncateAccessLogPath(trimmed)
}
return truncateAccessLogPath("/" + trimmed)
}
func truncateAccessLogPath(value string) string {
runes := []rune(value)
if len(runes) <= accessLogPathMaxRunes {
return value
}
return string(runes[:accessLogPathMaxRunes])
}
func topCounts(values map[string]int64, limit int) map[string]int64 {
return cloneTrafficCounts(values, limit)
}
@@ -0,0 +1,188 @@
package observability
import (
"os"
"path/filepath"
"strings"
"testing"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
"github.com/rain-kl/openflare/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{AccessLogPath: logPath}, 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{AccessLogPath: logPath}, 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{AccessLogPath: logPath}, 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{AccessLogPath: logPath}, 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 TestBuildTrafficObservabilityTruncatesLongAccessLogPath(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")
longPath := "/" + strings.Repeat("a", 140)
content := []byte(
"{\"ts\":\"2026-03-14T08:00:00Z\",\"host\":\"app.example.com\",\"path\":\"" + longPath + "\",\"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"))
_, accessLogs, _ := BuildTrafficObservability(&config.Config{AccessLogPath: logPath}, stateStore, nil)
if len(accessLogs) != 1 {
t.Fatalf("expected one access log, got %+v", accessLogs)
}
if got := len([]rune(accessLogs[0].Path)); got != accessLogPathMaxRunes {
t.Fatalf("expected truncated path length %d, got %d (%q)", accessLogPathMaxRunes, got, accessLogs[0].Path)
}
}
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{AccessLogPath: logPath}, 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,211 @@
package protocol
import "encoding/json"
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"`
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
}
type HeartbeatResult struct {
AgentSettings *AgentSettings
ActiveConfig *ActiveConfigMeta
WAFIPGroups []WAFIPGroup
}
type AgentSettings struct {
HeartbeatInterval int `json:"heartbeat_interval"`
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
AutoUpdate bool `json:"auto_update"`
UpdateRepo string `json:"update_repo"`
UpdateNow bool `json:"update_now"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
RestartOpenrestyNow bool `json:"restart_openresty_now"`
}
const (
WSMessageTypeStatus = "status"
WSMessageTypeSettings = "settings"
WSMessageTypeActiveConfig = "active_config"
WSMessageTypeForceSyncConfig = "force_sync_config"
WSMessageTypeWAFIPGroups = "waf_ip_groups"
WSMessageTypePing = "ping"
WSMessageTypePong = "pong"
)
type WSMessage struct {
Type string `json:"type"`
Payload json.RawMessage `json:"payload,omitempty"`
}
type WSOutboundMessage struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}
type WebSocketConnection interface {
URL() string
SendStatus(payload NodePayload) error
SendPong() error
Receive() (WSMessage, error)
Close() error
}
const (
OpenrestyStatusHealthy = "healthy"
OpenrestyStatusUnhealthy = "unhealthy"
OpenrestyStatusUnknown = "unknown"
)
type NodePayload struct {
NodeID string `json:"node_id"`
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
ExtVersion string `json:"ext_version"`
CurrentVersion string `json:"current_version"`
LastError string `json:"last_error"`
OpenrestyStatus string `json:"openresty_status"`
OpenrestyMessage string `json:"openresty_message"`
Profile *NodeSystemProfile `json:"profile,omitempty"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
BufferedObservability []BufferedObservabilityRecord `json:"buffered_observability,omitempty"`
HealthEvents []NodeHealthEvent `json:"health_events"`
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
}
type NodeSystemProfile struct {
Hostname string `json:"hostname"`
OSName string `json:"os_name"`
OSVersion string `json:"os_version"`
KernelVersion string `json:"kernel_version"`
Architecture string `json:"architecture"`
CPUModel string `json:"cpu_model"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAtUnix int64 `json:"reported_at_unix"`
}
type NodeMetricSnapshot struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
MemoryTotalBytes int64 `json:"memory_total_bytes"`
StorageUsedBytes int64 `json:"storage_used_bytes"`
StorageTotalBytes int64 `json:"storage_total_bytes"`
DiskReadBytes int64 `json:"disk_read_bytes"`
DiskWriteBytes int64 `json:"disk_write_bytes"`
NetworkRxBytes int64 `json:"network_rx_bytes"`
NetworkTxBytes int64 `json:"network_tx_bytes"`
}
type NodeOpenrestyObservation struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
OpenrestyConnections int64 `json:"openresty_connections"`
}
type NodeTrafficReport struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
RequestCount int64 `json:"request_count"`
ErrorCount int64 `json:"error_count"`
UniqueVisitorCount int64 `json:"unique_visitor_count"`
StatusCodes map[string]int64 `json:"status_codes"`
TopDomains map[string]int64 `json:"top_domains"`
SourceCountries map[string]int64 `json:"source_countries"`
}
type NodeAccessLog struct {
LoggedAtUnix int64 `json:"logged_at_unix"`
RemoteAddr string `json:"remote_addr"`
Host string `json:"host"`
Path string `json:"path"`
StatusCode int `json:"status_code"`
}
type BufferedObservabilityRecord struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
}
type NodeHealthEvent struct {
EventType string `json:"event_type"`
Severity string `json:"severity"`
Message string `json:"message"`
TriggeredAtUnix int64 `json:"triggered_at_unix"`
Metadata map[string]string `json:"metadata,omitempty"`
}
type RegisterNodeResponse struct {
NodeID string `json:"node_id"`
AccessToken string `json:"agent_token"`
Name string `json:"name"`
}
type ApplyLogPayload struct {
NodeID string `json:"node_id"`
Version string `json:"version"`
Result string `json:"result"`
Message string `json:"message"`
Checksum string `json:"checksum"`
MainConfigChecksum string `json:"main_config_checksum"`
RouteConfigChecksum string `json:"route_config_checksum"`
SupportFileCount int `json:"support_file_count"`
}
type ActiveConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
SourceConfigJSON string `json:"source_config_json"`
SupportFiles []SupportFile `json:"support_files"`
CreatedAt string `json:"created_at"`
}
type ActiveConfigMeta struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
Checksum string `json:"checksum"`
}
type WAFIPGroupSyncRequest struct {
IDs []uint `json:"ids,omitempty"`
Checksums map[string]string `json:"checksums,omitempty"`
}
type WAFIPGroupSyncResponse struct {
Groups []WAFIPGroup `json:"groups"`
}
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
@@ -0,0 +1,229 @@
package state
import (
"encoding/json"
"os"
"path/filepath"
"sort"
"strconv"
"sync"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
)
const observabilityBufferWindowSeconds = 60
type ObservabilityBufferRecord struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
Snapshot *protocol.NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *protocol.NodeOpenrestyObservation `json:"openresty_observation,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.OpenrestyObservation == 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.OpenrestyObservation != nil {
merged.OpenrestyObservation = incoming.OpenrestyObservation
}
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, openresty *protocol.NodeOpenrestyObservation, traffic *protocol.NodeTrafficReport) int64 {
if traffic != nil && traffic.WindowStartedAtUnix > 0 {
return traffic.WindowStartedAtUnix - (traffic.WindowStartedAtUnix % observabilityBufferWindowSeconds)
}
if openresty != nil && openresty.CapturedAtUnix > 0 {
return openresty.CapturedAtUnix - (openresty.CapturedAtUnix % 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"
"github.com/rain-kl/openflare/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, nil, &protocol.NodeTrafficReport{WindowStartedAtUnix: 1710403200}); value != 1710403200 {
t.Fatalf("unexpected traffic window start: %d", value)
}
if value := ObservabilityWindowStartedAt(&protocol.NodeMetricSnapshot{CapturedAtUnix: 1710403259}, nil, nil); value != 1710403200 {
t.Fatalf("unexpected snapshot-derived window start: %d", value)
}
}
+106
View File
@@ -0,0 +1,106 @@
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"`
BlockedVersion string `json:"blocked_version"`
BlockedChecksum string `json:"blocked_checksum"`
BlockedReason string `json:"blocked_reason"`
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,150 @@
package state
import (
"fmt"
"os"
"path/filepath"
"sync"
"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")
}
}
func TestStore_Load_NonExistentFile(t *testing.T) {
// Loading from a non-existent path should succeed and return an empty Snapshot
tempFile := filepath.Join(t.TempDir(), "nonexistent.json")
store := NewStore(tempFile)
snap, err := store.Load()
if err != nil {
t.Fatalf("expected Load to succeed for non-existent file, got err: %v", err)
}
if snap == nil {
t.Fatal("expected non-nil snapshot")
}
if snap.NodeID != "" || snap.CurrentVersion != "" {
t.Errorf("expected empty snapshot, got: %+v", snap)
}
}
func TestStore_Load_EmptyFile(t *testing.T) {
// Loading from an empty file should succeed and return an empty Snapshot
tempFile := filepath.Join(t.TempDir(), "empty.json")
if err := os.WriteFile(tempFile, []byte(""), 0644); err != nil {
t.Fatalf("failed to create empty file: %v", err)
}
store := NewStore(tempFile)
snap, err := store.Load()
if err != nil {
t.Fatalf("expected Load to succeed for empty file, got err: %v", err)
}
if snap == nil {
t.Fatal("expected non-nil snapshot")
}
if snap.NodeID != "" {
t.Errorf("expected empty snapshot, got: %+v", snap)
}
}
func TestStore_Load_InvalidJSON(t *testing.T) {
// Loading from a corrupted file with invalid JSON should fail with parsing error
tempFile := filepath.Join(t.TempDir(), "corrupted.json")
if err := os.WriteFile(tempFile, []byte("{invalid-json"), 0644); err != nil {
t.Fatalf("failed to create corrupted file: %v", err)
}
store := NewStore(tempFile)
_, err := store.Load()
if err == nil {
t.Fatal("expected Load to fail for corrupted JSON file")
}
}
func TestStore_SaveAndLoad(t *testing.T) {
tempFile := filepath.Join(t.TempDir(), "state.json")
store := NewStore(tempFile)
original := &Snapshot{
NodeID: "node-test-123",
CurrentVersion: "20260531-001",
CurrentChecksum: "chk-active-xyz",
BlockedVersion: "20260531-002",
BlockedChecksum: "chk-blocked-abc",
BlockedReason: "invalid upstream domain name",
LastError: "configuration reload timeout",
OpenrestyStatus: "unhealthy",
}
if err := store.Save(original); err != nil {
t.Fatalf("expected Save to succeed, got: %v", err)
}
loaded, err := store.Load()
if err != nil {
t.Fatalf("expected Load to succeed, got: %v", err)
}
if loaded.NodeID != original.NodeID ||
loaded.CurrentVersion != original.CurrentVersion ||
loaded.CurrentChecksum != original.CurrentChecksum ||
loaded.BlockedVersion != original.BlockedVersion ||
loaded.BlockedChecksum != original.BlockedChecksum ||
loaded.BlockedReason != original.BlockedReason ||
loaded.LastError != original.LastError ||
loaded.OpenrestyStatus != original.OpenrestyStatus {
t.Errorf("loaded snapshot does not match original: %+v vs %+v", loaded, original)
}
}
func TestStore_ConcurrencySafety(t *testing.T) {
tempFile := filepath.Join(t.TempDir(), "state.json")
store := NewStore(tempFile)
var wg sync.WaitGroup
workers := 20
iterations := 50
// Run concurrent writers and readers
for i := 0; i < workers; i++ {
wg.Add(1)
go func(workerID int) {
defer wg.Done()
for j := 0; j < iterations; j++ {
// Concurrently save
snap := &Snapshot{
NodeID: fmt.Sprintf("node-%d", workerID),
CurrentVersion: fmt.Sprintf("v-%d", j),
}
if err := store.Save(snap); err != nil {
t.Errorf("Save failed under concurrency: %v", err)
}
// Concurrently load
if _, err := store.Load(); err != nil {
t.Errorf("Load failed under concurrency: %v", err)
}
// Concurrently ensure ID
if _, err := store.EnsureNodeID(); err != nil {
t.Errorf("EnsureNodeID failed under concurrency: %v", err)
}
}
}(i)
}
wg.Wait()
}
+328
View File
@@ -0,0 +1,328 @@
package sync
import (
"archive/zip"
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path"
"path/filepath"
"strings"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
)
type pagesSourceDocument struct {
Routes []pagesSourceRoute `json:"routes"`
}
type pagesSourceRoute struct {
UpstreamType string `json:"upstream_type"`
PagesDeployment *pagesDeploymentSource `json:"pages_deployment"`
}
type pagesDeploymentSource struct {
DeploymentID uint `json:"deployment_id"`
Checksum string `json:"checksum"`
}
type pagesDeploymentMarker struct {
DeploymentID uint `json:"deployment_id"`
Checksum string `json:"checksum"`
}
func (s *Service) syncPagesDeployments(ctx context.Context, config *protocol.ActiveConfigResponse) error {
deployments, err := referencedPagesDeployments(config)
if err != nil {
return err
}
if len(deployments) == 0 {
return nil
}
if strings.TrimSpace(s.pagesDir) == "" {
return errors.New("pages_dir is required when active config references Pages deployments")
}
for _, deployment := range deployments {
if err := s.ensurePagesDeployment(ctx, deployment); err != nil {
return err
}
}
return nil
}
func (s *Service) ensurePagesDeployment(ctx context.Context, deployment pagesDeploymentSource) error {
currentDir := pagesCurrentDir(s.pagesDir, deployment.DeploymentID)
if markerMatches(currentDir, deployment) {
return nil
}
packageBytes, err := s.client.DownloadPagesDeploymentPackage(ctx, deployment.DeploymentID)
if err != nil {
return fmt.Errorf("download Pages deployment %d: %w", deployment.DeploymentID, err)
}
if got := checksumBytes(packageBytes); got != deployment.Checksum {
return fmt.Errorf("Pages deployment %d checksum mismatch: expected %s, got %s", deployment.DeploymentID, deployment.Checksum, got)
}
releaseDir := pagesReleaseDir(s.pagesDir, deployment.DeploymentID, deployment.Checksum)
if !markerMatches(releaseDir, deployment) {
if err := extractPagesPackage(packageBytes, releaseDir, deployment); err != nil {
return err
}
}
return switchPagesCurrentDir(s.pagesDir, deployment.DeploymentID, releaseDir)
}
func referencedPagesDeployments(config *protocol.ActiveConfigResponse) ([]pagesDeploymentSource, error) {
if config == nil || strings.TrimSpace(config.SourceConfigJSON) == "" {
return nil, nil
}
var doc pagesSourceDocument
if err := json.Unmarshal([]byte(config.SourceConfigJSON), &doc); err != nil {
return nil, fmt.Errorf("decode Pages references: %w", err)
}
seen := make(map[uint]struct{})
result := make([]pagesDeploymentSource, 0)
for _, route := range doc.Routes {
if strings.ToLower(strings.TrimSpace(route.UpstreamType)) != "pages" || route.PagesDeployment == nil {
continue
}
deploymentID := route.PagesDeployment.DeploymentID
checksum := strings.TrimSpace(route.PagesDeployment.Checksum)
if deploymentID == 0 || checksum == "" {
return nil, errors.New("Pages deployment snapshot is incomplete")
}
if _, ok := seen[deploymentID]; ok {
continue
}
seen[deploymentID] = struct{}{}
result = append(result, pagesDeploymentSource{DeploymentID: deploymentID, Checksum: checksum})
}
return result, nil
}
func findCommonRootPrefix(files []*zip.File) (string, error) {
var firstFilePath string
hasMultipleFiles := false
for _, item := range files {
relativePath, skip, err := normalizePagesArchivePath(item.Name)
if err != nil {
return "", err
}
if skip {
continue
}
normalizedPath := filepath.ToSlash(relativePath)
if firstFilePath == "" {
firstFilePath = normalizedPath
} else {
hasMultipleFiles = true
}
}
if firstFilePath == "" {
return "", nil
}
parts := strings.Split(firstFilePath, "/")
if len(parts) <= 1 {
return "", nil
}
commonPrefix := parts[0] + "/"
if hasMultipleFiles {
for _, item := range files {
relativePath, skip, err := normalizePagesArchivePath(item.Name)
if err != nil {
return "", err
}
if skip {
continue
}
normalizedPath := filepath.ToSlash(relativePath)
if !strings.HasPrefix(normalizedPath, commonPrefix) {
return "", nil
}
}
}
return commonPrefix, nil
}
func extractPagesPackage(packageBytes []byte, releaseDir string, deployment pagesDeploymentSource) error {
tmpDir := releaseDir + ".tmp"
_ = os.RemoveAll(tmpDir)
if err := os.MkdirAll(tmpDir, 0o755); err != nil {
return err
}
reader, err := zip.NewReader(bytes.NewReader(packageBytes), int64(len(packageBytes)))
if err != nil {
_ = os.RemoveAll(tmpDir)
return fmt.Errorf("open Pages zip: %w", err)
}
commonPrefix, err := findCommonRootPrefix(reader.File)
if err != nil {
_ = os.RemoveAll(tmpDir)
return err
}
for _, item := range reader.File {
relativePath, skip, err := normalizePagesArchivePath(item.Name)
if err != nil {
_ = os.RemoveAll(tmpDir)
return err
}
if skip {
continue
}
if commonPrefix != "" {
slashPath := filepath.ToSlash(relativePath)
if strings.HasPrefix(slashPath, commonPrefix) {
relativePath = filepath.FromSlash(strings.TrimPrefix(slashPath, commonPrefix))
}
}
if item.FileInfo().Mode()&os.ModeSymlink != 0 {
_ = os.RemoveAll(tmpDir)
return fmt.Errorf("Pages package contains unsupported symlink: %s", relativePath)
}
if err := extractPagesFile(item, filepath.Join(tmpDir, relativePath)); err != nil {
_ = os.RemoveAll(tmpDir)
return err
}
}
if err := writePagesMarker(tmpDir, deployment); err != nil {
_ = os.RemoveAll(tmpDir)
return err
}
_ = os.RemoveAll(releaseDir)
return os.Rename(tmpDir, releaseDir)
}
func extractPagesFile(item *zip.File, targetPath string) error {
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
return err
}
source, err := item.Open()
if err != nil {
return err
}
defer source.Close()
target, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, item.FileInfo().Mode().Perm())
if err != nil {
return err
}
defer target.Close()
_, err = io.Copy(target, source)
return err
}
func switchPagesCurrentDir(baseDir string, deploymentID uint, releaseDir string) error {
currentDir := pagesCurrentDir(baseDir, deploymentID)
previousDir := currentDir + ".previous"
_ = os.RemoveAll(previousDir)
if err := os.MkdirAll(filepath.Dir(currentDir), 0o755); err != nil {
return err
}
if _, err := os.Stat(currentDir); err == nil {
if err := os.Rename(currentDir, previousDir); err != nil {
return err
}
}
if err := copyPagesDir(releaseDir, currentDir); err != nil {
_ = os.RemoveAll(currentDir)
if _, restoreErr := os.Stat(previousDir); restoreErr == nil {
_ = os.Rename(previousDir, currentDir)
}
return err
}
_ = os.RemoveAll(previousDir)
return nil
}
func copyPagesDir(sourceDir string, targetDir string) error {
return filepath.WalkDir(sourceDir, func(sourcePath string, entry os.DirEntry, err error) error {
if err != nil {
return err
}
relativePath, err := filepath.Rel(sourceDir, sourcePath)
if err != nil || relativePath == "." {
return err
}
targetPath := filepath.Join(targetDir, relativePath)
if entry.IsDir() {
return os.MkdirAll(targetPath, 0o755)
}
info, err := entry.Info()
if err != nil {
return err
}
input, err := os.Open(sourcePath)
if err != nil {
return err
}
defer input.Close()
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
return err
}
output, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, info.Mode().Perm())
if err != nil {
return err
}
defer output.Close()
_, err = io.Copy(output, input)
return err
})
}
func normalizePagesArchivePath(raw string) (string, bool, error) {
name := strings.TrimSpace(filepath.ToSlash(raw))
if name == "" || strings.HasSuffix(name, "/") {
return "", true, nil
}
if strings.HasPrefix(name, "/") {
return "", false, fmt.Errorf("Pages package contains absolute path: %s", raw)
}
cleaned := path.Clean(name)
if cleaned == "." {
return "", true, nil
}
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || strings.Contains(cleaned, "/../") {
return "", false, fmt.Errorf("Pages package path escapes deployment root: %s", raw)
}
return filepath.FromSlash(cleaned), false, nil
}
func markerMatches(dir string, deployment pagesDeploymentSource) bool {
data, err := os.ReadFile(filepath.Join(dir, ".openflare-pages.json"))
if err != nil {
return false
}
var marker pagesDeploymentMarker
if err := json.Unmarshal(data, &marker); err != nil {
return false
}
return marker.DeploymentID == deployment.DeploymentID && marker.Checksum == deployment.Checksum
}
func writePagesMarker(dir string, deployment pagesDeploymentSource) error {
data, err := json.Marshal(pagesDeploymentMarker{
DeploymentID: deployment.DeploymentID,
Checksum: deployment.Checksum,
})
if err != nil {
return err
}
return os.WriteFile(filepath.Join(dir, ".openflare-pages.json"), data, 0o644)
}
func pagesCurrentDir(baseDir string, deploymentID uint) string {
return filepath.Join(baseDir, "deployments", fmt.Sprintf("%d", deploymentID), "current")
}
func pagesReleaseDir(baseDir string, deploymentID uint, checksum string) string {
return filepath.Join(baseDir, "deployments", fmt.Sprintf("%d", deploymentID), "releases", checksum)
}
func checksumBytes(data []byte) string {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
+533
View File
@@ -0,0 +1,533 @@
package sync
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"log/slog"
"sort"
"strings"
openrestyrender "github.com/rain-kl/openflare/openflare-server/utils/render/openresty"
"github.com/rain-kl/openflare/openflare-agent/internal/nginx"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
"github.com/rain-kl/openflare/openflare-agent/internal/state"
)
const (
ApplyResultSuccess = "success"
ApplyResultWarning = "warning"
ApplyResultFailed = "failed"
)
type ConfigClient interface {
GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error)
DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error)
ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error
SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error)
}
type NginxManager interface {
Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) nginx.ApplyOutcome
EnsureRuntime(ctx context.Context, recreate bool) error
EnsureSafeFallbackRuntime(ctx context.Context, reason string) error
CurrentChecksum() (string, error)
WAFIPGroupChecksums() (map[string]string, error)
SyncWAFIPGroups(groups []protocol.WAFIPGroup) error
}
type Service struct {
client ConfigClient
nginxManager NginxManager
stateStore *state.Store
pagesDir string
}
func (s *Service) SetPagesDir(path string) {
s.pagesDir = strings.TrimSpace(path)
}
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 {
if startup {
config, fetchErr := s.client.GetActiveConfig(ctx)
if fetchErr != nil {
slog.Error("fetch active config failed", "mode", mode, "error", fetchErr)
return fetchErr
}
return s.applyIfNeeded(ctx, mode, startup, snapshot, currentChecksum, target, config)
}
slog.Debug("local openresty config already up to date", "mode", mode, "version", target.Version)
shouldReport := shouldReportNoopApply(snapshot, target.Version, target.Checksum)
if shouldReport {
if err = s.reportNoopApply(ctx, snapshot.NodeID, target.Version, target.Checksum, "", "", 0); err != nil {
return err
}
}
snapshot.CurrentVersion = target.Version
snapshot.CurrentChecksum = target.Checksum
clearBlockedTarget(snapshot)
snapshot.LastError = ""
slog.Debug("sync finished without changes", "mode", mode, "version", target.Version)
return s.stateStore.Save(snapshot)
}
if isBlockedTarget(snapshot, target.Version, target.Checksum) {
slog.Warn("skipping blocked config version after previous failed apply", "mode", mode, "version", target.Version, "checksum", target.Checksum)
if startup {
if err = s.ensureRuntimeForCurrentConfig(ctx, mode, snapshot, currentChecksum); err != nil {
return err
}
return s.stateStore.Save(snapshot)
}
return nil
}
if hasBlockedTarget(snapshot) {
clearBlockedTarget(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 s.stateStore.Save(snapshot)
}
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) ForceSyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error {
snapshot, err := s.stateStore.Load()
if err != nil {
return err
}
if hasBlockedTarget(snapshot) {
clearBlockedTarget(snapshot)
_ = s.stateStore.Save(snapshot)
}
currentChecksum, err := s.nginxManager.CurrentChecksum()
if err != nil {
return err
}
config, err := s.client.GetActiveConfig(ctx)
if err != nil {
slog.Error("fetch active config failed", "mode", "force", "error", err)
return err
}
return s.applyIfNeeded(ctx, "force", true, snapshot, currentChecksum, target, config)
}
func (s *Service) WAFIPGroupChecksums() (map[string]string, error) {
if s.nginxManager == nil {
return map[string]string{}, nil
}
return s.nginxManager.WAFIPGroupChecksums()
}
func (s *Service) ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error {
if len(groups) == 0 || s.nginxManager == nil {
return nil
}
return s.nginxManager.SyncWAFIPGroups(groups)
}
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 && !startup {
slog.Debug("local openresty config already up to date", "mode", mode, "version", config.Version)
shouldReport := shouldReportNoopApply(snapshot, config.Version, config.Checksum)
if shouldReport {
rendered, renderErr := renderActiveConfig(config)
if renderErr != nil {
return renderErr
}
if err := s.reportNoopApply(ctx, snapshot.NodeID, config.Version, config.Checksum, checksumString(rendered.mainConfig), checksumString(rendered.routeConfig), len(rendered.supportFiles)); err != nil {
return err
}
}
snapshot.CurrentVersion = config.Version
snapshot.CurrentChecksum = config.Checksum
clearBlockedTarget(snapshot)
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 isBlockedTarget(snapshot, config.Version, config.Checksum) {
slog.Warn("skipping blocked config after fetch because the same version previously failed", "mode", mode, "version", config.Version, "checksum", config.Checksum)
if startup {
if err := s.ensureRuntimeForCurrentConfig(ctx, mode, snapshot, currentChecksum); err != nil {
return err
}
return s.stateStore.Save(snapshot)
}
return nil
}
if hasBlockedTarget(snapshot) {
clearBlockedTarget(snapshot)
}
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 s.stateStore.Save(snapshot)
}
rendered, err := renderActiveConfig(config)
if err != nil {
return err
}
if err := s.syncPagesDeployments(ctx, config); err != nil {
return err
}
mainConfigChecksum := checksumString(rendered.mainConfig)
routeConfigChecksum := checksumString(rendered.routeConfig)
slog.Info("applying new openresty config", "mode", mode, "from_version", snapshot.CurrentVersion, "to_version", config.Version, "old_checksum", currentChecksum, "new_checksum", config.Checksum)
outcome := s.nginxManager.Apply(ctx, rendered.mainConfig, rendered.routeConfig, rendered.supportFiles)
message := strings.TrimSpace(outcome.Message)
if outcome.Status == "" {
outcome.Status = nginx.ApplyStatusFatal
if message == "" {
message = "openresty apply returned empty outcome"
}
}
reportResult := ApplyResultFailed
switch outcome.Status {
case nginx.ApplyStatusSuccess:
slog.Info("openresty config applied successfully", "mode", mode, "version", config.Version)
snapshot.CurrentVersion = config.Version
snapshot.CurrentChecksum = config.Checksum
clearBlockedTarget(snapshot)
snapshot.LastError = ""
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
snapshot.OpenrestyMessage = ""
reportResult = ApplyResultSuccess
if message == "" {
message = "apply success"
}
case nginx.ApplyStatusWarning:
if message == "" {
message = "apply rolled back to previous config"
}
slog.Warn("openresty config apply rolled back", "mode", mode, "version", config.Version, "message", message)
markBlockedTarget(snapshot, config.Version, config.Checksum, message)
snapshot.LastError = message
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
snapshot.OpenrestyMessage = message
reportResult = ApplyResultWarning
default:
if message == "" {
message = "openresty apply failed"
}
slog.Error("apply openresty config failed", "mode", mode, "version", config.Version, "message", message)
markBlockedTarget(snapshot, config.Version, config.Checksum, message)
snapshot.LastError = message
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
snapshot.OpenrestyMessage = message
}
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: reportResult,
Message: message,
Checksum: config.Checksum,
MainConfigChecksum: mainConfigChecksum,
RouteConfigChecksum: routeConfigChecksum,
SupportFileCount: len(rendered.supportFiles),
}); err != nil {
slog.Error("report apply log failed", "version", config.Version, "result", reportResult, "error", err)
return err
}
if reportResult == ApplyResultFailed {
slog.Warn("failed apply log reported", "version", config.Version)
return outcomeError(config.Version, message)
}
if err := s.syncReferencedWAFIPGroups(ctx, rendered.supportFiles); err != nil {
slog.Error("sync referenced waf ip groups failed", "version", config.Version, "error", err)
return err
}
slog.Debug("apply log reported", "version", config.Version, "result", reportResult)
return nil
}
func (s *Service) syncReferencedWAFIPGroups(ctx context.Context, supportFiles []protocol.SupportFile) error {
ids := referencedWAFIPGroupIDs(supportFiles)
if len(ids) == 0 {
return nil
}
checksums, err := s.WAFIPGroupChecksums()
if err != nil {
return err
}
response, err := s.client.SyncWAFIPGroups(ctx, protocol.WAFIPGroupSyncRequest{
IDs: ids,
Checksums: checksums,
})
if err != nil {
return err
}
if response == nil || len(response.Groups) == 0 {
return nil
}
return s.ApplyWAFIPGroups(ctx, response.Groups)
}
type renderedActiveConfig struct {
mainConfig string
routeConfig string
supportFiles []protocol.SupportFile
}
func renderActiveConfig(config *protocol.ActiveConfigResponse) (*renderedActiveConfig, error) {
if config == nil {
return nil, errors.New("active config is nil")
}
sourceJSON := strings.TrimSpace(config.SourceConfigJSON)
if sourceJSON == "" {
return nil, errors.New("active config source_config_json is empty")
}
rendered, err := openrestyrender.RenderJSON(sourceJSON, toOpenRestySupportFiles(config.SupportFiles))
if err != nil {
return nil, err
}
files := fromOpenRestySupportFiles(rendered.SupportFiles)
files = append(files, protocol.SupportFile{Path: openrestyrender.SourceConfigFileName, Content: sourceJSON})
return &renderedActiveConfig{
mainConfig: rendered.MainConfig,
routeConfig: rendered.RouteConfig,
supportFiles: files,
}, nil
}
func toOpenRestySupportFiles(files []protocol.SupportFile) []openrestyrender.SupportFile {
if len(files) == 0 {
return nil
}
result := make([]openrestyrender.SupportFile, 0, len(files))
for _, file := range files {
result = append(result, openrestyrender.SupportFile{Path: file.Path, Content: file.Content})
}
return result
}
func fromOpenRestySupportFiles(files []openrestyrender.SupportFile) []protocol.SupportFile {
if len(files) == 0 {
return nil
}
result := make([]protocol.SupportFile, 0, len(files))
for _, file := range files {
result = append(result, protocol.SupportFile{Path: file.Path, Content: file.Content})
}
return result
}
func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
var content string
for _, file := range supportFiles {
if file.Path == "waf_config.json" {
content = strings.TrimSpace(file.Content)
break
}
}
if content == "" {
return []uint{}
}
var payload struct {
RuleGroups []struct {
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
} `json:"rule_groups"`
}
if err := json.Unmarshal([]byte(content), &payload); err != nil {
slog.Debug("decode waf_config.json for ip group references failed", "error", err)
return []uint{}
}
seen := make(map[uint]struct{})
for _, group := range payload.RuleGroups {
for _, id := range group.IPWhitelistGroups {
if id > 0 {
seen[id] = struct{}{}
}
}
for _, id := range group.IPBlacklistGroups {
if id > 0 {
seen[id] = struct{}{}
}
}
}
ids := make([]uint, 0, len(seen))
for id := range seen {
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids
}
func shouldReportNoopApply(snapshot *state.Snapshot, version string, checksum string) bool {
if snapshot == nil {
return false
}
return strings.TrimSpace(snapshot.CurrentVersion) != strings.TrimSpace(version) ||
strings.TrimSpace(snapshot.CurrentChecksum) != strings.TrimSpace(checksum)
}
func (s *Service) reportNoopApply(ctx context.Context, nodeID string, version string, checksum string, mainConfigChecksum string, routeConfigChecksum string, supportFileCount int) error {
message := "local config already matches active version; apply skipped"
if err := s.client.ReportApplyLog(ctx, protocol.ApplyLogPayload{
NodeID: nodeID,
Version: strings.TrimSpace(version),
Result: ApplyResultSuccess,
Message: message,
Checksum: strings.TrimSpace(checksum),
MainConfigChecksum: strings.TrimSpace(mainConfigChecksum),
RouteConfigChecksum: strings.TrimSpace(routeConfigChecksum),
SupportFileCount: supportFileCount,
}); err != nil {
slog.Error("report noop apply log failed", "version", version, "error", err)
return err
}
slog.Debug("noop apply log reported", "version", version)
return nil
}
func outcomeError(version string, message string) error {
trimmed := strings.TrimSpace(message)
if trimmed == "" {
trimmed = "openresty apply failed"
}
return fmt.Errorf("apply version %s failed: %s", version, trimmed)
}
func (s *Service) ensureRuntimeForCurrentConfig(ctx context.Context, mode string, snapshot *state.Snapshot, currentChecksum string) error {
if strings.TrimSpace(currentChecksum) == "" {
slog.Warn("blocked config cannot be retried and no local checksum is available for runtime recovery", "mode", mode, "blocked_version", snapshot.BlockedVersion)
reason := fmt.Sprintf("blocked config %s has no valid local config available for runtime recovery", strings.TrimSpace(snapshot.BlockedVersion))
if err := s.nginxManager.EnsureSafeFallbackRuntime(ctx, reason); err != nil {
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
snapshot.OpenrestyMessage = err.Error()
_ = s.stateStore.Save(snapshot)
return err
}
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
snapshot.OpenrestyMessage = "safe default fallback runtime started"
return nil
}
slog.Info("ensuring runtime with current local config while active target remains blocked", "mode", mode, "current_version", snapshot.CurrentVersion, "current_checksum", currentChecksum, "blocked_version", snapshot.BlockedVersion)
if err := s.nginxManager.EnsureRuntime(ctx, true); err != nil {
if strings.TrimSpace(snapshot.CurrentChecksum) == "" {
reason := fmt.Sprintf("blocked config %s has no historical config and current local config cannot start: %v", strings.TrimSpace(snapshot.BlockedVersion), err)
if fallbackErr := s.nginxManager.EnsureSafeFallbackRuntime(ctx, reason); fallbackErr == nil {
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
snapshot.OpenrestyMessage = "safe default fallback runtime started"
return nil
} else {
err = fmt.Errorf("%v; fallback recovery failed: %w", err, fallbackErr)
}
}
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
snapshot.OpenrestyMessage = err.Error()
_ = s.stateStore.Save(snapshot)
return err
}
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
if strings.TrimSpace(snapshot.OpenrestyMessage) == strings.TrimSpace(snapshot.BlockedReason) {
snapshot.OpenrestyMessage = ""
}
return nil
}
func markBlockedTarget(snapshot *state.Snapshot, version string, checksum string, reason string) {
if snapshot == nil {
return
}
snapshot.BlockedVersion = strings.TrimSpace(version)
snapshot.BlockedChecksum = strings.TrimSpace(checksum)
snapshot.BlockedReason = strings.TrimSpace(reason)
}
func clearBlockedTarget(snapshot *state.Snapshot) {
if snapshot == nil {
return
}
snapshot.BlockedVersion = ""
snapshot.BlockedChecksum = ""
snapshot.BlockedReason = ""
}
func hasBlockedTarget(snapshot *state.Snapshot) bool {
return snapshot != nil && (strings.TrimSpace(snapshot.BlockedVersion) != "" || strings.TrimSpace(snapshot.BlockedChecksum) != "")
}
func isBlockedTarget(snapshot *state.Snapshot, version string, checksum string) bool {
if snapshot == nil {
return false
}
return strings.TrimSpace(snapshot.BlockedVersion) == strings.TrimSpace(version) &&
strings.TrimSpace(snapshot.BlockedChecksum) == strings.TrimSpace(checksum) &&
(strings.TrimSpace(version) != "" || strings.TrimSpace(checksum) != "")
}
func checksumString(content string) string {
sum := sha256.Sum256([]byte(content))
return hex.EncodeToString(sum[:])
}
@@ -0,0 +1,923 @@
package sync
import (
"archive/zip"
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/nginx"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
"github.com/rain-kl/openflare/openflare-agent/internal/state"
)
type fakeExecutor struct {
testErr error
reloadErr error
}
func testPagesSourceConfigJSON(deploymentID uint, checksum string) string {
return fmt.Sprintf(`{"routes":[{"id":1,"site_name":"pages","domain":"pages.example.com","domains":["pages.example.com"],"origin_url":"openflare-pages://project/1","upstreams":["openflare-pages://project/1"],"enabled":true,"upstream_type":"pages","pages_deployment":{"project_id":1,"project_slug":"pages","deployment_id":%d,"deployment_number":1,"checksum":"%s","entry_file":"index.html","spa_fallback_enabled":true,"local_root":"__OPENFLARE_PAGES_DIR__/deployments/%d/current"}}],"openresty_config":{"worker_processes":"auto","worker_connections":1024,"worker_rlimit_nofile":65535,"events_multi_accept_enabled":true,"keepalive_timeout":20,"keepalive_requests":1000,"client_header_timeout":15,"client_body_timeout":15,"client_max_body_size":"64m","large_client_header_buffers":"4 16k","send_timeout":30,"proxy_connect_timeout":3,"proxy_send_timeout":60,"proxy_read_timeout":60,"websocket_enabled":true,"proxy_request_buffering":false,"proxy_buffering_enabled":true,"proxy_buffers":"16 16k","proxy_buffer_size":"8k","proxy_busy_buffers_size":"64k","gzip_enabled":true,"gzip_min_length":1024,"gzip_comp_level":5,"cache_enabled":false,"cache_levels":"1:2","cache_inactive":"30m","cache_max_size":"1g","cache_key_template":"$scheme$host$request_uri","cache_lock_enabled":true,"cache_lock_timeout":"5s","cache_use_stale":"error timeout updating http_500 http_502 http_503 http_504","main_config_template":"worker_processes {{OpenRestyWorkerProcesses}};"},"waf":{"rule_groups":[],"bindings":[]}}`, deploymentID, checksum, deploymentID)
}
type fakeClient struct {
config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload
pagesPackages map[uint][]byte
fetchCalls int
}
type fakeManager struct {
applyOutcome nginx.ApplyOutcome
currentChecksum string
currentChecksumErr error
ensureErr error
fallbackErr error
ensureCalls []bool
fallbackReasons []string
applyMainContents []string
applyRouteContents []string
applyFiles [][]protocol.SupportFile
}
func testSourceConfigJSON(workerProcesses string, listen int) string {
return fmt.Sprintf(`{"routes":[{"id":1,"site_name":"example","domain":"example.com","domains":["example.com"],"origin_url":"http://127.0.0.1:%d","upstreams":["http://127.0.0.1:%d"],"enabled":true}],"openresty_config":{"worker_processes":"%s","worker_connections":1024,"worker_rlimit_nofile":65535,"events_multi_accept_enabled":true,"keepalive_timeout":20,"keepalive_requests":1000,"client_header_timeout":15,"client_body_timeout":15,"client_max_body_size":"64m","large_client_header_buffers":"4 16k","send_timeout":30,"proxy_connect_timeout":3,"proxy_send_timeout":60,"proxy_read_timeout":60,"websocket_enabled":true,"proxy_request_buffering":false,"proxy_buffering_enabled":true,"proxy_buffers":"16 16k","proxy_buffer_size":"8k","proxy_busy_buffers_size":"64k","gzip_enabled":true,"gzip_min_length":1024,"gzip_comp_level":5,"cache_enabled":false,"cache_levels":"1:2","cache_inactive":"30m","cache_max_size":"1g","cache_key_template":"$scheme$host$request_uri","cache_lock_enabled":true,"cache_lock_timeout":"5s","cache_use_stale":"error timeout updating http_500 http_502 http_503 http_504","main_config_template":"worker_processes {{OpenRestyWorkerProcesses}};"},"waf":{"rule_groups":[],"bindings":[]}}`, listen, listen, workerProcesses)
}
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) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
if f.pagesPackages == nil {
return nil, fmt.Errorf("missing Pages package %d", deploymentID)
}
return f.pagesPackages[deploymentID], nil
}
func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
f.reports = append(f.reports, payload)
return nil
}
func (f *fakeClient) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) {
return &protocol.WAFIPGroupSyncResponse{}, nil
}
func (m *fakeManager) Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) nginx.ApplyOutcome {
m.applyMainContents = append(m.applyMainContents, mainConfig)
m.applyRouteContents = append(m.applyRouteContents, routeConfig)
m.applyFiles = append(m.applyFiles, append([]protocol.SupportFile(nil), supportFiles...))
if m.applyOutcome.Status == "" {
return nginx.ApplyOutcome{Status: nginx.ApplyStatusSuccess}
}
return m.applyOutcome
}
func (m *fakeManager) EnsureRuntime(ctx context.Context, recreate bool) error {
m.ensureCalls = append(m.ensureCalls, recreate)
return m.ensureErr
}
func (m *fakeManager) EnsureSafeFallbackRuntime(ctx context.Context, reason string) error {
m.fallbackReasons = append(m.fallbackReasons, reason)
return m.fallbackErr
}
func (m *fakeManager) CurrentChecksum() (string, error) {
return m.currentChecksum, m.currentChecksumErr
}
func (m *fakeManager) WAFIPGroupChecksums() (map[string]string, error) {
return map[string]string{}, nil
}
func (m *fakeManager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error {
return nil
}
func TestSyncOnceSuccess(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-001",
Checksum: "checksum-1",
SourceConfigJSON: testSourceConfigJSON("auto", 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 !strings.Contains(string(data), "listen 80;") || !strings.Contains(string(data), "server_name example.com;") {
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 != 3 {
t.Fatalf("expected support file count to be reported, got %d", client.reports[0].SupportFileCount)
}
}
func TestSyncOnceDownloadsPagesDeploymentBeforeApply(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{"index.html": "hello"})
checksum := testBytesChecksum(packageBytes)
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-101",
Checksum: "pages-config-checksum",
SourceConfigJSON: testPagesSourceConfigJSON(7, checksum),
CreatedAt: time.Now().Format(time.RFC3339),
},
pagesPackages: map[uint][]byte{7: packageBytes},
}
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("save state failed: %v", err)
}
manager := &fakeManager{currentChecksum: "old-checksum"}
service := New(client, manager, stateStore)
pagesDir := t.TempDir()
service.SetPagesDir(pagesDir)
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{Version: "20260309-101", Checksum: "pages-config-checksum"}); err != nil {
t.Fatalf("SyncOnce failed: %v", err)
}
data, err := os.ReadFile(filepath.Join(pagesDir, "deployments", "7", "current", "index.html"))
if err != nil {
t.Fatalf("expected Pages file to be extracted: %v", err)
}
if string(data) != "hello" {
t.Fatalf("unexpected Pages file content: %s", string(data))
}
if len(manager.applyRouteContents) != 1 || !strings.Contains(manager.applyRouteContents[0], "__OPENFLARE_PAGES_DIR__/deployments/7/current") {
t.Fatalf("expected Pages placeholder in rendered route config, got %#v", manager.applyRouteContents)
}
}
func TestSyncOnceRejectsPagesZipSlipBeforeApply(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{"../escape.html": "bad", "index.html": "ok"})
checksum := testBytesChecksum(packageBytes)
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-102",
Checksum: "pages-config-checksum",
SourceConfigJSON: testPagesSourceConfigJSON(8, checksum),
CreatedAt: time.Now().Format(time.RFC3339),
},
pagesPackages: map[uint][]byte{8: packageBytes},
}
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
if _, err := stateStore.EnsureNodeID(); err != nil {
t.Fatalf("EnsureNodeID failed: %v", err)
}
manager := &fakeManager{currentChecksum: "old-checksum"}
service := New(client, manager, stateStore)
service.SetPagesDir(t.TempDir())
err := service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{Version: "20260309-102", Checksum: "pages-config-checksum"})
if err == nil || !strings.Contains(err.Error(), "escapes deployment root") {
t.Fatalf("expected zip-slip rejection, got %v", err)
}
if len(manager.applyRouteContents) != 0 {
t.Fatalf("OpenResty apply must not run after Pages package rejection")
}
}
func TestSyncOnceRollbackOnNginxFailure(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-002",
Checksum: "checksum-2",
SourceConfigJSON: testSourceConfigJSON("2", 81),
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,
CurrentVersion: "20260309-001",
CurrentChecksum: "checksum-1",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
service := New(client, &fakeManager{
applyOutcome: nginx.ApplyOutcome{
Status: nginx.ApplyStatusFatal,
Message: "openresty failed after rollback",
},
}, 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 apply outcome is fatal")
}
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 snapshot.BlockedVersion != "20260309-002" || snapshot.BlockedChecksum != "checksum-2" {
t.Fatalf("expected failed target version to be blocked, got %+v", snapshot)
}
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusUnhealthy {
t.Fatalf("expected unhealthy openresty status, got %q", snapshot.OpenrestyStatus)
}
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 != 3 {
t.Fatalf("expected failed report to include support file count, got %d", client.reports[0].SupportFileCount)
}
}
func TestSyncOnceReportsWarningWhenRollbackKeepsOpenrestyHealthy(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-002",
Checksum: "checksum-2",
SourceConfigJSON: testSourceConfigJSON("2", 81),
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,
CurrentVersion: "20260309-001",
CurrentChecksum: "checksum-1",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
service := New(client, &fakeManager{
applyOutcome: nginx.ApplyOutcome{
Status: nginx.ApplyStatusWarning,
Message: "apply failed, rolled back to previous config",
},
}, stateStore)
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
Version: client.config.Version,
Checksum: client.config.Checksum,
}); err != nil {
t.Fatalf("expected warning outcome to keep sync successful, got %v", err)
}
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 warning apply to keep previous version state")
}
if snapshot.BlockedVersion != "20260309-002" || snapshot.BlockedChecksum != "checksum-2" {
t.Fatalf("expected rolled-back target version to be blocked, got %+v", snapshot)
}
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusHealthy {
t.Fatalf("expected healthy openresty after rollback, got %q", snapshot.OpenrestyStatus)
}
if snapshot.LastError == "" {
t.Fatal("expected rollback warning to be recorded")
}
if len(client.reports) != 1 || client.reports[0].Result != ApplyResultWarning {
t.Fatal("expected warning apply report to be sent")
}
}
func TestSyncOnStartupRecreatesRuntimeWhenChecksumMatches(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-003",
Checksum: "checksum-3",
SourceConfigJSON: testSourceConfigJSON("auto", 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.applyMainContents) != 1 {
t.Fatal("expected startup sync to re-render and apply local config")
}
if len(client.reports) != 1 || client.reports[0].Result != ApplyResultSuccess {
t.Fatal("expected startup sync to report apply success when state is refreshed")
}
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 TestSyncOnceReportsNoopWhenVersionChangesButChecksumMatches(t *testing.T) {
client := &fakeClient{}
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: "20260309-002",
CurrentChecksum: "checksum-3",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
manager := &fakeManager{currentChecksum: "checksum-3"}
service := New(client, manager, stateStore)
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
Version: "20260309-003",
Checksum: "checksum-3",
}); err != nil {
t.Fatalf("SyncOnce failed: %v", err)
}
if client.fetchCalls != 0 {
t.Fatalf("expected checksum match to skip config fetch, got %d", client.fetchCalls)
}
if len(manager.applyMainContents) != 0 {
t.Fatal("expected checksum match to skip apply")
}
if len(client.reports) != 1 || client.reports[0].Result != ApplyResultSuccess {
t.Fatalf("expected noop apply success report, got %+v", client.reports)
}
if client.reports[0].Version != "20260309-003" || client.reports[0].Checksum != "checksum-3" {
t.Fatalf("unexpected noop apply report: %+v", client.reports[0])
}
snapshot, err := stateStore.Load()
if err != nil {
t.Fatalf("failed to load state: %v", err)
}
if snapshot.CurrentVersion != "20260309-003" || snapshot.CurrentChecksum != "checksum-3" {
t.Fatalf("expected state to refresh active version, got %+v", snapshot)
}
}
func TestSyncOnceDoesNotRepeatNoopReportWhenStateAlreadyMatches(t *testing.T) {
client := &fakeClient{}
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: "20260309-003",
CurrentChecksum: "checksum-3",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
manager := &fakeManager{currentChecksum: "checksum-3"}
service := New(client, manager, stateStore)
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
Version: "20260309-003",
Checksum: "checksum-3",
}); err != nil {
t.Fatalf("SyncOnce failed: %v", err)
}
if len(client.reports) != 0 {
t.Fatalf("expected matching state to skip duplicate noop report, got %+v", client.reports)
}
}
func TestSyncOnStartupRecordsRuntimeFailure(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-004",
Checksum: "checksum-4",
SourceConfigJSON: testSourceConfigJSON("4", 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",
applyOutcome: nginx.ApplyOutcome{Status: nginx.ApplyStatusFatal, Message: context.DeadlineExceeded.Error()},
}
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 TestSyncOnceSkipsPreviouslyBlockedVersion(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-006",
Checksum: "checksum-6",
SourceConfigJSON: testSourceConfigJSON("6", 86),
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: "20260309-005",
CurrentChecksum: "checksum-5",
BlockedVersion: "20260309-006",
BlockedChecksum: "checksum-6",
BlockedReason: "apply failed, rolled back to previous config",
LastError: "apply failed, rolled back to previous config",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
manager := &fakeManager{currentChecksum: "checksum-5"}
service := New(client, manager, stateStore)
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
Version: "20260309-006",
Checksum: "checksum-6",
}); err != nil {
t.Fatalf("expected blocked version to be skipped, got %v", err)
}
if client.fetchCalls != 0 {
t.Fatalf("expected blocked version to skip fetch, got %d", client.fetchCalls)
}
if len(manager.applyMainContents) != 0 {
t.Fatal("expected blocked version to skip apply")
}
if len(client.reports) != 0 {
t.Fatal("expected blocked version to skip reporting duplicate apply result")
}
}
func TestSyncOnStartupKeepsBlockedVersionSuppressedUntilNewTargetArrives(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-007",
Checksum: "checksum-7",
SourceConfigJSON: testSourceConfigJSON("7", 87),
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: "20260309-005",
CurrentChecksum: "checksum-5",
BlockedVersion: "20260309-007",
BlockedChecksum: "checksum-7",
BlockedReason: "apply failed, rolled back to previous config",
OpenrestyStatus: protocol.OpenrestyStatusUnhealthy,
OpenrestyMessage: "apply failed, rolled back to previous config",
LastError: "apply failed, rolled back to previous config",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
manager := &fakeManager{currentChecksum: "checksum-5"}
service := New(client, manager, stateStore)
if err = service.SyncOnStartup(context.Background(), &protocol.ActiveConfigMeta{
Version: "20260309-007",
Checksum: "checksum-7",
}); err != nil {
t.Fatalf("expected blocked startup target to be skipped, got %v", err)
}
if len(manager.ensureCalls) != 1 || !manager.ensureCalls[0] {
t.Fatal("expected startup skip to ensure runtime with current local config")
}
if client.fetchCalls != 0 {
t.Fatalf("expected blocked startup target to skip fetch, got %d", client.fetchCalls)
}
if len(client.reports) != 0 {
t.Fatal("expected blocked startup target to skip duplicate apply report")
}
snapshot, err := stateStore.Load()
if err != nil {
t.Fatalf("failed to load state: %v", err)
}
if snapshot.BlockedVersion != "20260309-007" || snapshot.BlockedChecksum != "checksum-7" {
t.Fatalf("expected blocked target to remain recorded, got %+v", snapshot)
}
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusHealthy {
t.Fatalf("expected startup runtime recovery to mark openresty healthy, got %q", snapshot.OpenrestyStatus)
}
}
func TestSyncOnStartupStartsFallbackWhenBlockedVersionHasNoLocalConfig(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-007",
Checksum: "checksum-7",
SourceConfigJSON: testSourceConfigJSON("7", 87),
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,
BlockedVersion: "20260309-007",
BlockedChecksum: "checksum-7",
BlockedReason: "apply failed, but fallback runtime started",
OpenrestyStatus: protocol.OpenrestyStatusUnhealthy,
OpenrestyMessage: "apply failed, but fallback runtime started",
LastError: "apply failed, but fallback runtime started",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
manager := &fakeManager{}
service := New(client, manager, stateStore)
if err = service.SyncOnStartup(context.Background(), &protocol.ActiveConfigMeta{
Version: "20260309-007",
Checksum: "checksum-7",
}); err != nil {
t.Fatalf("expected blocked startup target to start fallback, got %v", err)
}
if len(manager.fallbackReasons) != 1 {
t.Fatalf("expected fallback runtime to be started once, got %d", len(manager.fallbackReasons))
}
if client.fetchCalls != 0 {
t.Fatalf("expected blocked startup target to skip fetch, got %d", client.fetchCalls)
}
if len(client.reports) != 0 {
t.Fatal("expected blocked startup target to skip duplicate apply report")
}
snapshot, err := stateStore.Load()
if err != nil {
t.Fatalf("failed to load state: %v", err)
}
if snapshot.BlockedVersion != "20260309-007" || snapshot.BlockedChecksum != "checksum-7" {
t.Fatalf("expected blocked target to remain recorded, got %+v", snapshot)
}
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusHealthy {
t.Fatalf("expected fallback startup recovery to mark openresty healthy, got %q", snapshot.OpenrestyStatus)
}
if snapshot.OpenrestyMessage != "safe default fallback runtime started" {
t.Fatalf("expected fallback status message, got %q", snapshot.OpenrestyMessage)
}
}
func TestSyncOnStartupStartsFallbackWhenResidualConfigCannotRecover(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-007",
Checksum: "checksum-7",
SourceConfigJSON: testSourceConfigJSON("7", 87),
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,
BlockedVersion: "20260309-007",
BlockedChecksum: "checksum-7",
BlockedReason: "apply failed, but fallback runtime started",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
manager := &fakeManager{
currentChecksum: "residual-checksum",
ensureErr: context.DeadlineExceeded,
}
service := New(client, manager, stateStore)
if err = service.SyncOnStartup(context.Background(), &protocol.ActiveConfigMeta{
Version: "20260309-007",
Checksum: "checksum-7",
}); err != nil {
t.Fatalf("expected residual config failure to start fallback, got %v", err)
}
if len(manager.ensureCalls) != 1 {
t.Fatalf("expected residual config to be tested once, got %d", len(manager.ensureCalls))
}
if len(manager.fallbackReasons) != 1 {
t.Fatalf("expected fallback runtime to be started once, got %d", len(manager.fallbackReasons))
}
snapshot, err := stateStore.Load()
if err != nil {
t.Fatalf("failed to load state: %v", err)
}
if snapshot.OpenrestyStatus != protocol.OpenrestyStatusHealthy {
t.Fatalf("expected fallback startup recovery to mark openresty healthy, got %q", snapshot.OpenrestyStatus)
}
if snapshot.BlockedVersion != "20260309-007" || snapshot.BlockedChecksum != "checksum-7" {
t.Fatalf("expected blocked target to remain recorded, got %+v", snapshot)
}
}
func TestSyncOnceClearsBlockedTargetWhenNewVersionArrives(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-008",
Checksum: "checksum-8",
SourceConfigJSON: testSourceConfigJSON("8", 88),
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: "20260309-005",
CurrentChecksum: "checksum-5",
BlockedVersion: "20260309-007",
BlockedChecksum: "checksum-7",
BlockedReason: "apply failed, rolled back to previous config",
}); err != nil {
t.Fatalf("failed to seed state: %v", err)
}
manager := &fakeManager{}
service := New(client, manager, stateStore)
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{
Version: "20260309-008",
Checksum: "checksum-8",
}); err != nil {
t.Fatalf("expected new target version to be applied, got %v", err)
}
if client.fetchCalls != 1 {
t.Fatalf("expected new target to trigger fetch, got %d", client.fetchCalls)
}
if len(manager.applyMainContents) != 1 {
t.Fatal("expected new target to trigger apply")
}
snapshot, err := stateStore.Load()
if err != nil {
t.Fatalf("failed to load state: %v", err)
}
if snapshot.BlockedVersion != "" || snapshot.BlockedChecksum != "" {
t.Fatalf("expected blocked target to be cleared after new version succeeds, got %+v", snapshot)
}
if snapshot.CurrentVersion != "20260309-008" || snapshot.CurrentChecksum != "checksum-8" {
t.Fatalf("expected current version to move to new target, got %+v", snapshot)
}
}
func TestSyncOnceSkipsFetchWhenHeartbeatChecksumMatches(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-005",
Checksum: "checksum-5",
SourceConfigJSON: testSourceConfigJSON("auto", 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")
}
}
func testPagesPackage(t *testing.T, files map[string]string) []byte {
t.Helper()
var buffer bytes.Buffer
writer := zip.NewWriter(&buffer)
for name, content := range files {
file, err := writer.Create(name)
if err != nil {
t.Fatalf("create zip file failed: %v", err)
}
if _, err := file.Write([]byte(content)); err != nil {
t.Fatalf("write zip file failed: %v", err)
}
}
if err := writer.Close(); err != nil {
t.Fatalf("close zip failed: %v", err)
}
return buffer.Bytes()
}
func testBytesChecksum(data []byte) string {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
func TestSyncOnceDownloadsPagesDeploymentWithTopLevelFolder(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{
"Speed-Test-source/index.html": "hello html",
"Speed-Test-source/assets/app.js": "hello js",
})
checksum := testBytesChecksum(packageBytes)
client := &fakeClient{
config: protocol.ActiveConfigResponse{
Version: "20260309-105",
Checksum: "pages-config-checksum",
SourceConfigJSON: testPagesSourceConfigJSON(77, checksum),
CreatedAt: time.Now().Format(time.RFC3339),
},
pagesPackages: map[uint][]byte{77: packageBytes},
}
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("save state failed: %v", err)
}
manager := &fakeManager{currentChecksum: "old-checksum"}
service := New(client, manager, stateStore)
pagesDir := t.TempDir()
service.SetPagesDir(pagesDir)
if err = service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{Version: "20260309-105", Checksum: "pages-config-checksum"}); err != nil {
t.Fatalf("SyncOnce failed: %v", err)
}
data, err := os.ReadFile(filepath.Join(pagesDir, "deployments", "77", "current", "index.html"))
if err != nil {
t.Fatalf("expected Pages index.html file to be extracted: %v", err)
}
if string(data) != "hello html" {
t.Fatalf("unexpected Pages index.html content: %s", string(data))
}
jsData, err := os.ReadFile(filepath.Join(pagesDir, "deployments", "77", "current", "assets", "app.js"))
if err != nil {
t.Fatalf("expected Pages assets/app.js file to be extracted: %v", err)
}
if string(jsData) != "hello js" {
t.Fatalf("unexpected Pages assets/app.js content: %s", string(jsData))
}
}
@@ -0,0 +1,51 @@
//go:build !windows
package updater
import (
"fmt"
"log/slog"
"os"
"syscall"
)
func replaceAndRestart(execPath string, tmpPath string) error {
backupPath := execPath + ".bak"
if err := removeBackupBinary(backupPath); err != nil {
return err
}
if err := os.Rename(execPath, backupPath); err != nil {
renameErr := err
if err := os.Remove(tmpPath); err != nil && !os.IsNotExist(err) {
slog.Error("remove tmp binary failed", "path", tmpPath, "error", err)
return fmt.Errorf("backup current binary: %w; remove tmp binary: %v", renameErr, err)
}
return fmt.Errorf("backup current binary: %w", renameErr)
}
if err := os.Rename(tmpPath, execPath); err != nil {
replaceErr := err
if err := os.Rename(backupPath, execPath); err != nil {
slog.Error("restore backup binary failed", "path", backupPath, "error", err)
return fmt.Errorf("replace binary: %w; restore backup binary: %v", replaceErr, err)
}
return fmt.Errorf("replace binary: %w", replaceErr)
}
if err := removeBackupBinary(backupPath); err != nil {
return err
}
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil {
return fmt.Errorf("exec restart: %w", err)
}
return fmt.Errorf("unreachable after exec")
}
func removeBackupBinary(path string) error {
if err := os.Remove(path); err != nil {
if os.IsNotExist(err) {
return nil
}
slog.Error("remove backup binary failed", "path", path, "error", err)
return err
}
return nil
}
@@ -0,0 +1,15 @@
//go:build !windows
package updater
import (
"path/filepath"
"testing"
)
func TestRemoveBackupBinaryIgnoresMissingFile(t *testing.T) {
backupPath := filepath.Join(t.TempDir(), "openflare-agent.bak")
if err := removeBackupBinary(backupPath); err != nil {
t.Fatalf("expected missing backup cleanup to be ignored: %v", err)
}
}
@@ -0,0 +1,53 @@
//go:build windows
package updater
import (
"fmt"
"os"
"os/exec"
"strings"
)
func replaceAndRestart(execPath string, tmpPath string) error {
backupPath := execPath + ".bak"
scriptPath := execPath + ".update.cmd"
script := fmt.Sprintf(`@echo off
setlocal
:waitloop
move /Y "%s" "%s" >nul 2>nul
if errorlevel 1 (
ping 127.0.0.1 -n 2 >nul
goto waitloop
)
move /Y "%s" "%s" >nul 2>nul
if errorlevel 1 exit /b 1
start "" %s
del /Q "%s" >nul 2>nul
del /Q "%%~f0" >nul 2>nul
`, execPath, backupPath, tmpPath, execPath, buildWindowsCommandLine(execPath, os.Args[1:]), backupPath)
if err := os.WriteFile(scriptPath, []byte(script), 0o700); err != nil {
os.Remove(tmpPath)
return fmt.Errorf("write restart script: %w", err)
}
cmd := exec.Command("cmd", "/C", "start", "", scriptPath)
if err := cmd.Start(); err != nil {
os.Remove(scriptPath)
os.Remove(tmpPath)
return fmt.Errorf("schedule restart: %w", err)
}
os.Exit(0)
return nil
}
func buildWindowsCommandLine(execPath string, args []string) string {
parts := []string{quoteWindowsArg(execPath)}
for _, arg := range args {
parts = append(parts, quoteWindowsArg(arg))
}
return strings.Join(parts, " ")
}
func quoteWindowsArg(value string) string {
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
}
+366
View File
@@ -0,0 +1,366 @@
package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"os"
"runtime"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/utils"
"github.com/rain-kl/openflare/openflare-agent/internal/agent"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
)
const maxChecksumAssetSize = 64 * 1024
var replaceAndRestartFunc = replaceAndRestart
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.Version)
checkKey := buildReleaseCheckKey(options, remoteVersion)
if remoteVersion == localVersion {
return nil
}
if !options.Force && checkKey != "" && checkKey == s.lastCheckKey {
return nil
}
if !isNewer(localVersion, remoteVersion) {
s.lastCheckKey = checkKey
return nil
}
slog.Info("agent update available", "from", localVersion, "to", remoteVersion)
assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
checksumAssetName := assetName + ".sha256"
var downloadURL string
var checksumURL string
for _, asset := range release.Assets {
switch asset.Name {
case assetName:
downloadURL = asset.BrowserDownloadURL
case checksumAssetName:
checksumURL = asset.BrowserDownloadURL
}
}
if downloadURL == "" {
s.lastCheckKey = checkKey
return fmt.Errorf("no matching asset %q in release %s", assetName, release.TagName)
}
if checksumURL == "" {
return fmt.Errorf("no matching checksum asset %q in release %s", checksumAssetName, release.TagName)
}
expectedChecksum, err := s.downloadChecksum(ctx, checksumURL, assetName)
if err != nil {
return fmt.Errorf("download checksum: %w", err)
}
execPath, err := os.Executable()
if err != nil {
return fmt.Errorf("get executable path: %w", err)
}
if err = s.downloadAndRestart(ctx, downloadURL, expectedChecksum, execPath); err != nil {
return fmt.Errorf("download and restart: %w", err)
}
s.lastCheckKey = checkKey
return nil
}
func (s *Service) getRelease(ctx context.Context, repo string, options agent.UpdateOptions) (*githubRelease, error) {
tagName := strings.TrimSpace(options.TagName)
if tagName != "" {
return s.getReleaseByTag(ctx, repo, tagName)
}
if strings.EqualFold(strings.TrimSpace(options.Channel), "preview") {
return s.getLatestPreviewRelease(ctx, repo)
}
return s.getLatestStableRelease(ctx, repo)
}
func (s *Service) getLatestStableRelease(ctx context.Context, repo string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/latest", repo)
return s.fetchReleaseFromURL(ctx, url)
}
func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases?per_page=20", repo)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/vnd.github+json")
resp, err := s.httpClient.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("github api returned %s", resp.Status)
}
var releases []githubRelease
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
return nil, err
}
for _, release := range releases {
if release.Draft || !release.Prerelease {
continue
}
releaseCopy := release
return &releaseCopy, nil
}
return nil, nil
}
func (s *Service) getReleaseByTag(ctx context.Context, repo string, tag string) (*githubRelease, error) {
url := fmt.Sprintf("https://api.github.com/repos/%s/releases/tags/%s", repo, strings.TrimSpace(tag))
return s.fetchReleaseFromURL(ctx, url)
}
func (s *Service) fetchReleaseFromURL(ctx context.Context, url string) (*githubRelease, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
req.Header.Set("Accept", "application/vnd.github+json")
resp, err := s.httpClient.Do(req)
if err != nil {
return nil, err
}
defer func(Body io.ReadCloser) {
err := Body.Close()
if err != nil {
slog.Error("failed to close response body", "error", err)
}
}(resp.Body)
if resp.StatusCode == http.StatusNotFound {
return nil, nil
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("github api returned %s", resp.Status)
}
return decodeRelease(resp.Body)
}
func decodeRelease(reader io.Reader) (*githubRelease, error) {
var release githubRelease
if err := json.NewDecoder(reader).Decode(&release); err != nil {
return nil, err
}
return &release, nil
}
func (s *Service) downloadChecksum(ctx context.Context, url string, assetName string) (string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return "", err
}
resp, err := s.httpClient.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("checksum download returned %s", resp.Status)
}
content, err := io.ReadAll(io.LimitReader(resp.Body, maxChecksumAssetSize+1))
if err != nil {
return "", err
}
if len(content) > maxChecksumAssetSize {
return "", fmt.Errorf("checksum asset exceeds %d bytes", maxChecksumAssetSize)
}
checksum, err := parseSHA256Checksum(string(content), assetName)
if err != nil {
return "", err
}
return checksum, nil
}
func parseSHA256Checksum(content string, assetName string) (string, error) {
assetName = strings.TrimSpace(assetName)
for _, line := range strings.Split(content, "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if checksum, ok := parseSHA256Line(line, assetName); ok {
return checksum, nil
}
}
if assetName == "" {
return "", fmt.Errorf("checksum asset does not contain a valid sha256 digest")
}
return "", fmt.Errorf("checksum asset does not contain a sha256 digest for %q", assetName)
}
func parseSHA256Line(line string, assetName string) (string, bool) {
fields := strings.Fields(line)
if len(fields) == 1 && isSHA256Hex(fields[0]) {
return strings.ToLower(fields[0]), true
}
if len(fields) >= 2 && isSHA256Hex(fields[0]) {
fileName := strings.TrimPrefix(strings.TrimSpace(fields[1]), "*")
if assetName == "" || fileName == assetName {
return strings.ToLower(fields[0]), true
}
}
prefix := "SHA256("
if strings.HasPrefix(line, prefix) {
closing := strings.Index(line, ")")
if closing > len(prefix) && closing+1 < len(line) {
fileName := strings.TrimSpace(line[len(prefix):closing])
rest := strings.TrimSpace(line[closing+1:])
rest = strings.TrimPrefix(rest, "=")
rest = strings.TrimSpace(rest)
if isSHA256Hex(rest) && (assetName == "" || fileName == assetName) {
return strings.ToLower(rest), true
}
}
}
return "", false
}
func isSHA256Hex(value string) bool {
value = strings.TrimSpace(value)
if len(value) != sha256.Size*2 {
return false
}
_, err := hex.DecodeString(value)
return err == nil
}
func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedChecksum string, targetPath string) error {
expectedChecksum = strings.ToLower(strings.TrimSpace(expectedChecksum))
if !isSHA256Hex(expectedChecksum) {
return fmt.Errorf("invalid expected sha256 checksum")
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
resp, err := s.httpClient.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("download returned %s", resp.Status)
}
tmpPath := targetPath + ".update"
if runtime.GOOS == "windows" && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") {
tmpPath += ".exe"
}
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600)
if err != nil {
return err
}
hasher := sha256.New()
if _, err = io.Copy(io.MultiWriter(tmpFile, hasher), resp.Body); err != nil {
tmpFile.Close()
os.Remove(tmpPath)
return err
}
if err = tmpFile.Close(); err != nil {
os.Remove(tmpPath)
return err
}
actualChecksum := hex.EncodeToString(hasher.Sum(nil))
if actualChecksum != expectedChecksum {
os.Remove(tmpPath)
return fmt.Errorf("sha256 checksum mismatch: expected %s, got %s", expectedChecksum, actualChecksum)
}
if err = os.Chmod(tmpPath, 0o755); err != nil && runtime.GOOS != "windows" {
os.Remove(tmpPath)
return fmt.Errorf("set executable permission: %w", err)
}
slog.Info("agent binary updated, restarting")
return replaceAndRestartFunc(targetPath, tmpPath)
}
func assetNameForGOOSGOARCH(goos string, goarch string) string {
name := fmt.Sprintf("openflare-agent-%s-%s", goos, goarch)
if goos == "windows" {
return name + ".exe"
}
return name
}
func normalizeVersion(v string) string {
v = strings.TrimSpace(v)
v = strings.TrimPrefix(v, "v")
return v
}
func isNewer(local, remote string) bool {
return compareVersions(local, remote) < 0
}
func buildReleaseCheckKey(options agent.UpdateOptions, remoteVersion string) string {
channel := strings.TrimSpace(options.Channel)
if channel == "" {
channel = "stable"
}
if tagName := strings.TrimSpace(options.TagName); tagName != "" {
return channel + ":" + tagName
}
return channel + ":" + remoteVersion
}
func compareVersions(local string, remote string) int {
return utils.CompareVersions(local, remote)
}
@@ -0,0 +1,242 @@
package updater
import (
"context"
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"github.com/rain-kl/openflare/openflare-agent/internal/agent"
"github.com/rain-kl/openflare/openflare-agent/internal/config"
)
type roundTripFunc func(req *http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestGetLatestPreviewRelease(t *testing.T) {
service := &Service{
httpClient: &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases?per_page=20" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`[
{"tag_name":"v1.0.0","prerelease":false},
{"tag_name":"v1.1.0-rc.1","prerelease":true}
]`)),
}, nil
}),
},
}
release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", agent.UpdateOptions{Channel: "preview"})
if err != nil {
t.Fatalf("expected preview release query to succeed: %v", err)
}
if release == nil || release.TagName != "v1.1.0-rc.1" {
t.Fatalf("unexpected preview release: %#v", release)
}
}
func TestGetReleaseByTag(t *testing.T) {
service := &Service{
httpClient: &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/tags/v1.1.0-rc.1" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{"tag_name":"v1.1.0-rc.1","prerelease":true}`)),
}, nil
}),
},
}
release, err := service.getRelease(context.Background(), "Rain-kl/OpenFlare", agent.UpdateOptions{Channel: "preview", TagName: "v1.1.0-rc.1", Force: true})
if err != nil {
t.Fatalf("expected tag release query to succeed: %v", err)
}
if release == nil || release.TagName != "v1.1.0-rc.1" {
t.Fatalf("unexpected tag release: %#v", release)
}
}
func TestCheckAndUpdateRequiresChecksumAsset(t *testing.T) {
originalVersion := config.Version
config.Version = "v1.0.0"
t.Cleanup(func() {
config.Version = originalVersion
})
assetName := assetNameForGOOSGOARCH(runtime.GOOS, runtime.GOARCH)
service := &Service{
httpClient: &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
if req.URL.String() != "https://api.github.com/repos/Rain-kl/OpenFlare/releases/latest" {
t.Fatalf("unexpected request url: %s", req.URL.String())
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{
"tag_name":"v1.0.1",
"assets":[
{"name":"` + assetName + `","browser_download_url":"https://example.test/agent"}
]
}`)),
}, nil
}),
},
}
err := service.CheckAndUpdate(context.Background(), "Rain-kl/OpenFlare", agent.UpdateOptions{})
if err == nil || !strings.Contains(err.Error(), "no matching checksum asset") {
t.Fatalf("expected missing checksum asset error, got %v", err)
}
}
func TestParseSHA256Checksum(t *testing.T) {
checksum := strings.Repeat("a", sha256.Size*2)
testCases := []struct {
name string
content string
asset string
want string
}{
{name: "single digest", content: checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
{name: "sha256sum format", content: checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
{name: "bsd format", content: "SHA256(openflare-agent-linux-amd64)= " + checksum + "\n", asset: "openflare-agent-linux-amd64", want: checksum},
{name: "selects matching file", content: strings.Repeat("b", sha256.Size*2) + " other\n" + checksum + " openflare-agent-linux-amd64\n", asset: "openflare-agent-linux-amd64", want: checksum},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
got, err := parseSHA256Checksum(testCase.content, testCase.asset)
if err != nil {
t.Fatalf("expected checksum parse to succeed: %v", err)
}
if got != testCase.want {
t.Fatalf("unexpected checksum: got %s want %s", got, testCase.want)
}
})
}
}
func TestDownloadAndRestartVerifiesChecksum(t *testing.T) {
payload := []byte("new-agent-binary")
sum := sha256.Sum256(payload)
expectedChecksum := hex.EncodeToString(sum[:])
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
t.Fatalf("write target: %v", err)
}
var replacedTarget string
var replacedTemp string
originalReplace := replaceAndRestartFunc
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
replacedTarget = execPath
replacedTemp = tmpPath
return nil
}
t.Cleanup(func() {
replaceAndRestartFunc = originalReplace
})
service := &Service{
httpClient: &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(string(payload))),
}, nil
}),
},
}
if err := service.downloadAndRestart(context.Background(), "https://example.test/agent", expectedChecksum, targetPath); err != nil {
t.Fatalf("expected verified download to succeed: %v", err)
}
if replacedTarget != targetPath {
t.Fatalf("unexpected replace target: %s", replacedTarget)
}
if replacedTemp == "" {
t.Fatal("expected replacement temp path to be recorded")
}
if _, err := os.Stat(replacedTemp); err != nil {
t.Fatalf("expected verified temp binary to remain for replacement: %v", err)
}
}
func TestDownloadAndRestartRejectsChecksumMismatch(t *testing.T) {
targetPath := filepath.Join(t.TempDir(), "openflare-agent")
if err := os.WriteFile(targetPath, []byte("old-agent-binary"), 0o755); err != nil {
t.Fatalf("write target: %v", err)
}
originalReplace := replaceAndRestartFunc
replaceAndRestartFunc = func(execPath string, tmpPath string) error {
t.Fatal("replace should not run on checksum mismatch")
return nil
}
t.Cleanup(func() {
replaceAndRestartFunc = originalReplace
})
service := &Service{
httpClient: &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader("tampered")),
}, nil
}),
},
}
err := service.downloadAndRestart(context.Background(), "https://example.test/agent", strings.Repeat("0", sha256.Size*2), targetPath)
if err == nil || !strings.Contains(err.Error(), "sha256 checksum mismatch") {
t.Fatalf("expected checksum mismatch error, got %v", err)
}
if _, err = os.Stat(targetPath + ".update"); !os.IsNotExist(err) {
t.Fatalf("expected temp update file to be removed, stat err=%v", err)
}
}
func TestIsNewerSupportsPrerelease(t *testing.T) {
testCases := []struct {
name string
local string
remote string
expected bool
}{
{name: "stable newer than prerelease", local: "1.2.3-rc.1", remote: "1.2.3", expected: true},
{name: "same stable not newer", local: "1.2.3", remote: "1.2.3-rc.1", expected: false},
{name: "higher prerelease sequence", local: "1.2.3-rc.1", remote: "1.2.3-rc.2", expected: true},
{name: "higher minor", local: "1.2.3", remote: "1.3.0-rc.1", expected: true},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
if actual := isNewer(testCase.local, testCase.remote); actual != testCase.expected {
t.Fatalf("unexpected compare result: local=%s remote=%s actual=%v expected=%v", testCase.local, testCase.remote, actual, testCase.expected)
}
})
}
}
@@ -0,0 +1,82 @@
package wsclient
import (
"context"
"time"
"github.com/rain-kl/openflare/openflare-agent/internal/protocol"
shared "github.com/rain-kl/openflare/openflare-server/utils/wsclient"
)
type WSMessage = shared.WSMessage
type MessageHandler = shared.MessageHandler
type Client struct {
sharedClient *shared.Client
}
type Connection struct {
sharedConn *shared.Connection
}
func New(baseURL string, token string, timeout time.Duration) *Client {
return &Client{
sharedClient: shared.New(shared.Config{
BaseURL: baseURL,
Token: token,
Timeout: timeout,
HeaderKey: "X-Agent-Token",
WSPath: "/api/agent/ws",
}),
}
}
func (c *Client) SetToken(token string) {
c.sharedClient.SetToken(token)
}
func (c *Client) URL() string {
return c.sharedClient.URL()
}
func (c *Client) Connect(ctx context.Context) (protocol.WebSocketConnection, error) {
conn, err := c.sharedClient.Connect(ctx)
if err != nil {
return nil, err
}
return &Connection{sharedConn: conn}, nil
}
func (conn *Connection) URL() string {
if conn == nil || conn.sharedConn == nil {
return ""
}
return conn.sharedConn.URL
}
func (conn *Connection) SendStatus(payload protocol.NodePayload) error {
return conn.sharedConn.SendMessage(protocol.WSMessageTypeStatus, payload)
}
func (conn *Connection) SendPong() error {
return conn.sharedConn.SendMessage(protocol.WSMessageTypePong, nil)
}
func (conn *Connection) Receive() (protocol.WSMessage, error) {
var message protocol.WSMessage
if err := conn.sharedConn.Receive(&message); err != nil {
return message, err
}
return message, nil
}
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error {
return conn.sharedConn.RunReceiveLoop(ctx, handler)
}
func (conn *Connection) Close() error {
if conn == nil || conn.sharedConn == nil {
return nil
}
return conn.sharedConn.Close()
}