mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 14:46:36 +08:00
[优化] go 引用调整
This commit is contained in:
@@ -0,0 +1 @@
|
||||
data
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
Binary file not shown.
@@ -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
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
@@ -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");
|
||||
}
|
||||
Binary file not shown.
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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[:])
|
||||
}
|
||||
@@ -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, `"`, `""`) + `"`
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user