mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 15:46:37 +08:00
fix lint
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
// Package agent implements the local OpenFlare agent runtime loop.
|
||||
package agent
|
||||
|
||||
import (
|
||||
@@ -16,12 +17,14 @@ import (
|
||||
edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat"
|
||||
)
|
||||
|
||||
// HeartbeatService handles node registration and periodic heartbeat reporting.
|
||||
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)
|
||||
}
|
||||
|
||||
// SyncService handles configuration synchronisation between the agent and the server.
|
||||
type SyncService interface {
|
||||
SyncOnStartup(ctx context.Context, target *protocol.ActiveConfigMeta) error
|
||||
SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error
|
||||
@@ -30,30 +33,36 @@ type SyncService interface {
|
||||
ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error
|
||||
}
|
||||
|
||||
// RuntimeManager manages the lifecycle and health checks of the OpenResty runtime.
|
||||
type RuntimeManager interface {
|
||||
CheckHealth(ctx context.Context) error
|
||||
Restart(ctx context.Context) error
|
||||
}
|
||||
|
||||
// WebSocketService manages the persistent WebSocket connection to the server.
|
||||
type WebSocketService interface {
|
||||
Connect(ctx context.Context) (protocol.WebSocketConnection, error)
|
||||
SetToken(token string)
|
||||
URL() string
|
||||
}
|
||||
|
||||
const websocketBackoffDefaultDelay = 30 * time.Second
|
||||
|
||||
// Runner coordinates the agent's heartbeat, configuration sync, and WebSocket upgrade lifecycle.
|
||||
type Runner struct {
|
||||
Config *config.Config
|
||||
StateStore *state.Store
|
||||
HeartbeatCycle *agentheartbeat.Cycle
|
||||
HeartbeatService HeartbeatService
|
||||
SyncService SyncService
|
||||
RuntimeManager RuntimeManager
|
||||
WebSocketService WebSocketService
|
||||
Config *config.Config
|
||||
StateStore *state.Store
|
||||
HeartbeatCycle *agentheartbeat.Cycle
|
||||
HeartbeatService HeartbeatService
|
||||
SyncService SyncService
|
||||
RuntimeManager RuntimeManager
|
||||
WebSocketService WebSocketService
|
||||
|
||||
restartOpenrestyNow bool
|
||||
websocketUpgradeEnabled bool
|
||||
}
|
||||
|
||||
// Run starts the agent's main loop, performing heartbeats and upgrading to WebSocket when available.
|
||||
func (r *Runner) Run(ctx context.Context) error {
|
||||
if r.HeartbeatCycle != nil {
|
||||
r.HeartbeatCycle.RecordSyncError = r.recordSyncError
|
||||
@@ -63,13 +72,7 @@ func (r *Runner) Run(ctx context.Context) error {
|
||||
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)
|
||||
}
|
||||
r.runStartupAuth(ctx, &nodeID)
|
||||
|
||||
heartbeatTicker := time.NewTicker(r.Config.HeartbeatInterval.Duration())
|
||||
defer heartbeatTicker.Stop()
|
||||
@@ -108,42 +111,66 @@ func (r *Runner) Run(ctx context.Context) error {
|
||||
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)
|
||||
}
|
||||
}
|
||||
r.handleWSDisconnect(ctx, nodeID)
|
||||
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()
|
||||
}
|
||||
r.handleHeartbeatTick(ctx, &nodeID, heartbeatTicker, tryStartWebSocket)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) runStartupAuth(ctx context.Context, nodeID *string) {
|
||||
if r.hasAccessToken() {
|
||||
if _, hbErr := r.performHeartbeatCycle(ctx, *nodeID, true); hbErr != nil {
|
||||
slog.Error("agent startup heartbeat failed", "error", hbErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err := r.tryRegister(ctx, nodeID); err != nil {
|
||||
slog.Error("agent initial discovery register failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) handleWSDisconnect(ctx context.Context, nodeID string) {
|
||||
if !r.hasAccessToken() {
|
||||
return
|
||||
}
|
||||
if _, hbErr := r.performHeartbeatCycle(ctx, nodeID, false); hbErr != nil {
|
||||
slog.Error("agent heartbeat after ws disconnect failed", "error", hbErr)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) handleHeartbeatTick(ctx context.Context, nodeID *string, heartbeatTicker *time.Ticker, tryStartWebSocket func()) {
|
||||
if !r.hasAccessToken() {
|
||||
if err := r.tryRegister(ctx, nodeID); err != nil {
|
||||
slog.Error("agent discovery register failed", "error", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
changed, hbErr := r.performHeartbeatCycle(ctx, *nodeID, false)
|
||||
if hbErr != nil {
|
||||
slog.Error("agent heartbeat failed", "error", hbErr)
|
||||
return
|
||||
}
|
||||
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)
|
||||
return r.HeartbeatCycle.Perform(ctx, nodeID, startup, r)
|
||||
}
|
||||
|
||||
// Apply applies the provided agent settings and reports whether the heartbeat interval changed.
|
||||
func (r *Runner) Apply(settings *protocol.AgentSettings) bool {
|
||||
return r.applySettings(settings)
|
||||
}
|
||||
|
||||
// RestartOpenrestyIfNeeded restarts OpenResty when a server-requested restart is pending.
|
||||
func (r *Runner) RestartOpenrestyIfNeeded(ctx context.Context) {
|
||||
r.tryRestartOpenresty(ctx)
|
||||
}
|
||||
@@ -251,7 +278,7 @@ func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol.
|
||||
|
||||
func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
|
||||
r.refreshOpenrestyHealth(ctx)
|
||||
payload, ackWindows := r.HeartbeatCycle.PrepareHeartbeatPayload(nodeID)
|
||||
payload, ackWindows := r.HeartbeatCycle.PrepareHeartbeatPayload(ctx, nodeID)
|
||||
if err := conn.SendStatus(payload); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -338,7 +365,7 @@ func newWebSocketBackoff() *webSocketBackoff {
|
||||
|
||||
func (backoff *webSocketBackoff) Next() time.Duration {
|
||||
if backoff == nil || len(backoff.delays) == 0 {
|
||||
return 30 * time.Second
|
||||
return websocketBackoffDefaultDelay
|
||||
}
|
||||
if backoff.index >= len(backoff.delays) {
|
||||
return backoff.delays[len(backoff.delays)-1]
|
||||
@@ -402,7 +429,7 @@ func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error {
|
||||
return errors.New("agent_token 为空且未配置 discovery_token")
|
||||
}
|
||||
slog.Info("agent discovery registration started")
|
||||
response, err := r.HeartbeatService.Register(ctx, r.HeartbeatCycle.NodePayload(*nodeID))
|
||||
response, err := r.HeartbeatService.Register(ctx, r.HeartbeatCycle.NodePayload(ctx, *nodeID))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -501,4 +528,4 @@ func (r *Runner) recordOpenrestyUnhealthy(err error, fallbackOnly bool) {
|
||||
if saveErr := r.StateStore.Save(snapshot); saveErr != nil {
|
||||
slog.Error("save state after recording openresty error failed", "error", saveErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -395,7 +395,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
|
||||
t.Fatalf("failed to prepare access log: %v", err)
|
||||
}
|
||||
|
||||
firstPayload := runner.HeartbeatCycle.NodePayload("node-observe")
|
||||
firstPayload := runner.HeartbeatCycle.NodePayload(context.Background(), "node-observe")
|
||||
if firstPayload.Profile == nil {
|
||||
t.Fatal("expected first heartbeat payload to include system profile")
|
||||
}
|
||||
@@ -412,7 +412,7 @@ func TestRunnerHeartbeatPayloadIncludesObservabilityExtensions(t *testing.T) {
|
||||
t.Fatalf("expected health events for openresty and sync error, got %+v", firstPayload.HealthEvents)
|
||||
}
|
||||
|
||||
secondPayload := runner.HeartbeatCycle.NodePayload("node-observe")
|
||||
secondPayload := runner.HeartbeatCycle.NodePayload(context.Background(), "node-observe")
|
||||
if secondPayload.Profile != nil {
|
||||
t.Fatal("expected unchanged profile to be omitted on subsequent heartbeat")
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package config loads and persists agent daemon configuration.
|
||||
package config
|
||||
|
||||
import (
|
||||
@@ -30,8 +31,12 @@ const (
|
||||
defaultObservabilityReplayMinutes = 15
|
||||
defaultMMDBUpdateInterval = 24 * time.Hour
|
||||
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
|
||||
defaultHeartbeatInterval = 10 * time.Second
|
||||
defaultRequestTimeout = 10 * time.Second
|
||||
configFilePerm = 0o600
|
||||
)
|
||||
|
||||
// Config holds the full runtime configuration for the OpenFlare agent.
|
||||
type Config struct {
|
||||
ServerURL string `json:"server_url"`
|
||||
AccessToken string `json:"agent_token"`
|
||||
@@ -61,7 +66,7 @@ type Config struct {
|
||||
StatePath string `json:"state_path"`
|
||||
HeartbeatInterval MillisecondDuration `json:"heartbeat_interval"`
|
||||
RequestTimeout MillisecondDuration `json:"request_timeout"`
|
||||
configPath string
|
||||
configPath string `json:"-"`
|
||||
}
|
||||
|
||||
type configFile struct {
|
||||
@@ -93,8 +98,17 @@ type configFile struct {
|
||||
RequestTimeout MillisecondDuration `json:"request_timeout"`
|
||||
}
|
||||
|
||||
func transferPersistedConfig(dst, src any) error {
|
||||
data, err := json.Marshal(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return json.Unmarshal(data, dst)
|
||||
}
|
||||
|
||||
// Load reads and parses the agent configuration file at the given path.
|
||||
func Load(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
data, err := os.ReadFile(path) //nolint:gosec // path is the configured agent config location
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return nil, err
|
||||
}
|
||||
@@ -107,33 +121,11 @@ func Load(path string) (*Config, error) {
|
||||
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 := &Config{}
|
||||
if err == nil {
|
||||
if err = transferPersistedConfig(cfg, file); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
cfg.configPath = path
|
||||
applyEnvOverrides(cfg)
|
||||
@@ -148,51 +140,58 @@ func applyDefaults(cfg *Config, baseDir string) {
|
||||
baseDir = filepath.Clean(baseDir)
|
||||
cfg.Version = Version
|
||||
cfg.OpenrestyResolvers = utils.UniqueAndCleanStringSlice(cfg.OpenrestyResolvers)
|
||||
applyAgentIdentityDefaults(cfg)
|
||||
applyAgentPathDefaults(cfg, baseDir)
|
||||
applyAgentTimingDefaults(cfg)
|
||||
normalizeManagedPaths(cfg)
|
||||
}
|
||||
|
||||
func applyAgentIdentityDefaults(cfg *Config) {
|
||||
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 = nodeip.Detect()
|
||||
}
|
||||
if cfg.MainConfigPath == "" {
|
||||
cfg.MainConfigPath = joinManagedPath(cfg.DataDir, defaultMainConfigRelativePath)
|
||||
}
|
||||
|
||||
func applyAgentPathDefaults(cfg *Config, baseDir string) {
|
||||
if cfg.DataDir == "" {
|
||||
cfg.DataDir = filepath.Join(baseDir, "data")
|
||||
}
|
||||
if cfg.RouteConfigPath == "" {
|
||||
cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultRouteConfigRelativePath)
|
||||
type managedPathDefault struct {
|
||||
target *string
|
||||
relative string
|
||||
}
|
||||
if cfg.AccessLogPath == "" {
|
||||
cfg.AccessLogPath = joinManagedPath(cfg.DataDir, defaultAccessLogRelativePath)
|
||||
pathDefaults := []managedPathDefault{
|
||||
{&cfg.MainConfigPath, defaultMainConfigRelativePath},
|
||||
{&cfg.RouteConfigPath, defaultRouteConfigRelativePath},
|
||||
{&cfg.AccessLogPath, defaultAccessLogRelativePath},
|
||||
{&cfg.StatePath, defaultStateRelativePath},
|
||||
{&cfg.CertDir, defaultCertDirRelativePath},
|
||||
{&cfg.LuaDir, defaultLuaDirRelativePath},
|
||||
{&cfg.RuntimeConfigDir, defaultRuntimeConfigDirRelativePath},
|
||||
{&cfg.PagesDir, defaultPagesDirRelativePath},
|
||||
{&cfg.MMDBPath, defaultMMDBRelativePath},
|
||||
{&cfg.ObservabilityBufferPath, defaultObservabilityBufferRelativePath},
|
||||
}
|
||||
if cfg.StatePath == "" {
|
||||
cfg.StatePath = joinManagedPath(cfg.DataDir, defaultStateRelativePath)
|
||||
}
|
||||
if cfg.CertDir == "" {
|
||||
cfg.CertDir = joinManagedPath(cfg.DataDir, defaultCertDirRelativePath)
|
||||
for _, item := range pathDefaults {
|
||||
if strings.TrimSpace(*item.target) == "" {
|
||||
*item.target = joinManagedPath(cfg.DataDir, item.relative)
|
||||
}
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
func applyAgentTimingDefaults(cfg *Config) {
|
||||
if cfg.MMDBUpdateInterval <= 0 {
|
||||
cfg.MMDBUpdateInterval = MillisecondDuration(defaultMMDBUpdateInterval)
|
||||
}
|
||||
@@ -202,19 +201,15 @@ func applyDefaults(cfg *Config, baseDir string) {
|
||||
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)
|
||||
cfg.HeartbeatInterval = MillisecondDuration(defaultHeartbeatInterval)
|
||||
}
|
||||
if cfg.RequestTimeout <= 0 {
|
||||
cfg.RequestTimeout = MillisecondDuration(10 * time.Second)
|
||||
cfg.RequestTimeout = MillisecondDuration(defaultRequestTimeout)
|
||||
}
|
||||
normalizeManagedPaths(cfg)
|
||||
}
|
||||
|
||||
func normalizeManagedPaths(cfg *Config) {
|
||||
@@ -360,6 +355,7 @@ func validate(cfg *Config) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// InitialAuthToken returns the agent access token, falling back to the discovery token if absent.
|
||||
func (cfg *Config) InitialAuthToken() string {
|
||||
if cfg == nil {
|
||||
return ""
|
||||
@@ -370,6 +366,15 @@ func (cfg *Config) InitialAuthToken() string {
|
||||
return strings.TrimSpace(cfg.DiscoveryToken)
|
||||
}
|
||||
|
||||
func (cfg *Config) toConfigFile() configFile {
|
||||
var file configFile
|
||||
if err := transferPersistedConfig(&file, cfg); err != nil {
|
||||
return configFile{}
|
||||
}
|
||||
return file
|
||||
}
|
||||
|
||||
// Save persists the current configuration back to its original file path.
|
||||
func (cfg *Config) Save() error {
|
||||
if cfg == nil {
|
||||
return errors.New("config 不能为空")
|
||||
@@ -377,11 +382,11 @@ func (cfg *Config) Save() error {
|
||||
if cfg.configPath == "" {
|
||||
return errors.New("config path 未初始化")
|
||||
}
|
||||
data, err := json.MarshalIndent(cfg, "", " ")
|
||||
data, err := json.MarshalIndent(cfg.toConfigFile(), "", " ") //nolint:gosec // agent token must be persisted in local config file
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(cfg.configPath, data, 0o644)
|
||||
return os.WriteFile(cfg.configPath, data, configFilePerm)
|
||||
}
|
||||
|
||||
func detectHostname() string {
|
||||
@@ -391,5 +396,3 @@ func detectHostname() string {
|
||||
}
|
||||
return strings.TrimSpace(host)
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -2,4 +2,5 @@ package config
|
||||
|
||||
import edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config"
|
||||
|
||||
type MillisecondDuration = edgeconfig.MillisecondDuration
|
||||
// MillisecondDuration is an alias for the edge config millisecond-precision duration type.
|
||||
type MillisecondDuration = edgeconfig.MillisecondDuration
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
package config
|
||||
|
||||
// Version is the current agent version string, overridden at build time.
|
||||
var Version = "dev"
|
||||
|
||||
@@ -1,8 +1,12 @@
|
||||
// Package geoipdata embeds the default MaxMind GeoLite2 country database.
|
||||
package geoipdata
|
||||
|
||||
import "embed"
|
||||
|
||||
// FS holds the embedded GeoLite2-Country.mmdb database.
|
||||
//
|
||||
//go:embed GeoLite2-Country.mmdb
|
||||
var FS embed.FS
|
||||
|
||||
// DefaultMMDBName is the filename of the embedded MaxMind country database.
|
||||
const DefaultMMDBName = "GeoLite2-Country.mmdb"
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package geoipupdate schedules local MaxMind GeoIP database updates for the agent.
|
||||
package geoipupdate
|
||||
|
||||
import (
|
||||
@@ -13,12 +14,20 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/pkg/geoip"
|
||||
)
|
||||
|
||||
const (
|
||||
mmdbDirPerm = 0o750
|
||||
mmdbFilePerm = 0o600
|
||||
)
|
||||
|
||||
// Updater periodically downloads a fresh GeoIP MMDB file and seeds the
|
||||
// initial embedded database when none is present on disk.
|
||||
type Updater struct {
|
||||
MMDBPath string
|
||||
DownloadURL string
|
||||
UpdateInterval time.Duration
|
||||
}
|
||||
|
||||
// EnsureInitialDatabase seeds the MMDB file from the embedded database if it does not exist on disk.
|
||||
func (u *Updater) EnsureInitialDatabase() error {
|
||||
path := filepath.Clean(u.MMDBPath)
|
||||
if path == "" || path == "." {
|
||||
@@ -33,16 +42,17 @@ func (u *Updater) EnsureInitialDatabase() error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("read embedded mmdb failed: %w", err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(path), mmdbDirPerm); err != nil {
|
||||
return fmt.Errorf("create mmdb directory failed: %w", err)
|
||||
}
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
if err := os.WriteFile(path, data, mmdbFilePerm); 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
|
||||
}
|
||||
|
||||
// Run starts the periodic GeoIP update loop and blocks until ctx is cancelled.
|
||||
func (u *Updater) Run(ctx context.Context) {
|
||||
if u == nil || u.MMDBPath == "" || u.UpdateInterval <= 0 {
|
||||
return
|
||||
@@ -57,7 +67,7 @@ func (u *Updater) Run(ctx context.Context) {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
if err := geoip.DownloadMaxMindDatabase(u.MMDBPath, u.DownloadURL); err != nil {
|
||||
if err := geoip.DownloadMaxMindDatabase(ctx, u.MMDBPath, u.DownloadURL); err != nil {
|
||||
slog.Warn("update GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
// Package heartbeat implements the periodic heartbeat cycle executed by the agent,
|
||||
// including payload preparation, config sync, WAF IP group application, and observability buffering.
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
@@ -14,10 +16,7 @@ import (
|
||||
edgeheartbeat "github.com/Rain-kl/Wavelet/internal/apps/edge/heartbeat"
|
||||
)
|
||||
|
||||
type HeartbeatClient interface {
|
||||
Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error)
|
||||
}
|
||||
|
||||
// SyncService is the interface used by Cycle to sync active configuration and WAF IP groups.
|
||||
type SyncService interface {
|
||||
SyncOnStartup(ctx context.Context, target *protocol.ActiveConfigMeta) error
|
||||
SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error
|
||||
@@ -25,23 +24,26 @@ type SyncService interface {
|
||||
ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error
|
||||
}
|
||||
|
||||
// SettingsApplier is the interface used by Cycle to apply agent settings received from the server.
|
||||
type SettingsApplier interface {
|
||||
Apply(settings *protocol.AgentSettings) (intervalChanged bool)
|
||||
RestartOpenrestyIfNeeded(ctx context.Context)
|
||||
}
|
||||
|
||||
// Cycle holds the dependencies required to execute a single agent heartbeat cycle.
|
||||
type Cycle struct {
|
||||
Config *config.Config
|
||||
StateStore *state.Store
|
||||
ObservabilityBuffer *state.ObservabilityBufferStore
|
||||
Heartbeat HeartbeatClient
|
||||
Heartbeat API
|
||||
Sync SyncService
|
||||
Updater *updater.Service
|
||||
RecordSyncError func(err error)
|
||||
}
|
||||
|
||||
// Perform executes one complete heartbeat cycle: sends the heartbeat, syncs config, and applies settings.
|
||||
func (c *Cycle) Perform(ctx context.Context, nodeID string, startup bool, settings SettingsApplier) (bool, error) {
|
||||
payload, ackWindows := c.PrepareHeartbeatPayload(nodeID)
|
||||
payload, ackWindows := c.PrepareHeartbeatPayload(ctx, nodeID)
|
||||
heartbeatResult, err := c.Heartbeat.Heartbeat(ctx, payload)
|
||||
if err != nil {
|
||||
return false, err
|
||||
@@ -79,14 +81,15 @@ func (c *Cycle) Perform(ctx context.Context, nodeID string, startup bool, settin
|
||||
return changed, nil
|
||||
}
|
||||
|
||||
func (c *Cycle) NodePayload(nodeID string) protocol.NodePayload {
|
||||
// NodePayload builds and returns the full NodePayload to be sent in a heartbeat request.
|
||||
func (c *Cycle) NodePayload(ctx context.Context, nodeID string) protocol.NodePayload {
|
||||
snapshot, _ := c.StateStore.Load()
|
||||
openrestyStatus := strings.TrimSpace(snapshot.OpenrestyStatus)
|
||||
if openrestyStatus == "" {
|
||||
openrestyStatus = protocol.OpenrestyStatusUnknown
|
||||
}
|
||||
profile := observability.BuildProfile(c.Config, c.StateStore)
|
||||
managedOpenRestyMetrics := observability.CollectManagedOpenRestyMetrics(c.Config)
|
||||
managedOpenRestyMetrics := observability.CollectManagedOpenRestyMetrics(ctx, c.Config)
|
||||
trafficReport, accessLogs, fallbackMetrics := observability.BuildTrafficObservability(c.Config, c.StateStore, managedOpenRestyMetrics)
|
||||
if managedOpenRestyMetrics == nil {
|
||||
managedOpenRestyMetrics = fallbackMetrics
|
||||
@@ -122,8 +125,9 @@ func (c *Cycle) NodePayload(nodeID string) protocol.NodePayload {
|
||||
return payload
|
||||
}
|
||||
|
||||
func (c *Cycle) PrepareHeartbeatPayload(nodeID string) (protocol.NodePayload, []int64) {
|
||||
payload := c.NodePayload(nodeID)
|
||||
// PrepareHeartbeatPayload constructs the heartbeat payload with buffered observability records and returns the window timestamps to acknowledge.
|
||||
func (c *Cycle) PrepareHeartbeatPayload(ctx context.Context, nodeID string) (protocol.NodePayload, []int64) {
|
||||
payload := c.NodePayload(ctx, nodeID)
|
||||
if c.ObservabilityBuffer == nil || (payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0) {
|
||||
return payload, nil
|
||||
}
|
||||
@@ -173,6 +177,7 @@ func (c *Cycle) PrepareHeartbeatPayload(nodeID string) (protocol.NodePayload, []
|
||||
return payload, ackWindows
|
||||
}
|
||||
|
||||
// AckObservabilityWindows acknowledges the given observability window timestamps in the buffer store.
|
||||
func (c *Cycle) AckObservabilityWindows(windowStartedAtUnix []int64) {
|
||||
if c.ObservabilityBuffer == nil || len(windowStartedAtUnix) == 0 {
|
||||
return
|
||||
@@ -183,6 +188,7 @@ func (c *Cycle) AckObservabilityWindows(windowStartedAtUnix []int64) {
|
||||
}
|
||||
}
|
||||
|
||||
// ApplyWAFIPGroups applies the WAF IP groups received from the server via the SyncService.
|
||||
func (c *Cycle) ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) {
|
||||
if len(groups) == 0 || c.Sync == nil {
|
||||
return
|
||||
@@ -199,6 +205,7 @@ func (c *Cycle) recordSyncError(err error) {
|
||||
}
|
||||
}
|
||||
|
||||
// AgentSettingsToAutoUpdate converts AgentSettings to an AutoUpdateSettings value used by the edge heartbeat updater.
|
||||
func AgentSettingsToAutoUpdate(settings *protocol.AgentSettings) *edgeheartbeat.AutoUpdateSettings {
|
||||
if settings == nil {
|
||||
return nil
|
||||
@@ -214,4 +221,4 @@ func AgentSettingsToAutoUpdate(settings *protocol.AgentSettings) *edgeheartbeat.
|
||||
|
||||
func agentSettingsToAutoUpdate(settings *protocol.AgentSettings) *edgeheartbeat.AutoUpdateSettings {
|
||||
return AgentSettingsToAutoUpdate(settings)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,28 +6,41 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
)
|
||||
|
||||
type Client interface {
|
||||
// RemoteClient is the interface that abstracts the remote API calls performed by Service.
|
||||
type RemoteClient 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
|
||||
// API abstracts registration and heartbeat operations used by Cycle.
|
||||
type API interface {
|
||||
Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error)
|
||||
Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error)
|
||||
SetToken(token string)
|
||||
}
|
||||
|
||||
func New(client Client) *Service {
|
||||
// Service wraps a RemoteClient to expose agent registration and heartbeat operations.
|
||||
type Service struct {
|
||||
client RemoteClient
|
||||
}
|
||||
|
||||
// New creates a new Service backed by the given RemoteClient.
|
||||
func New(client RemoteClient) *Service {
|
||||
return &Service{client: client}
|
||||
}
|
||||
|
||||
// Register sends a node registration request to the server.
|
||||
func (s *Service) Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
|
||||
return s.client.RegisterNode(ctx, payload)
|
||||
}
|
||||
|
||||
// Heartbeat sends a heartbeat to the server using the service client.
|
||||
func (s *Service) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) {
|
||||
return s.client.Heartbeat(ctx, payload)
|
||||
}
|
||||
|
||||
// SetToken sets the authentication token for the service client.
|
||||
func (s *Service) SetToken(token string) {
|
||||
s.client.SetToken(token)
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package httpclient provides an authenticated HTTP client for the agent.
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
@@ -8,20 +9,23 @@ import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
|
||||
)
|
||||
|
||||
// Client is a HTTP client used by the agent to communicate with the control plane server.
|
||||
type Client struct {
|
||||
base *edgehttp.Client
|
||||
}
|
||||
|
||||
// New creates a new Client instance with the specified base URL, token, and timeout.
|
||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
base: edgehttp.New(baseURL, token, timeout, "X-Agent-Token"),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterNode registers the agent node with the control plane server.
|
||||
func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
|
||||
resp := protocol.APIResponse[protocol.RegisterNodeResponse]{}
|
||||
if err := c.base.PostJSON(ctx, "/api/v1/agent/nodes/register", payload, &resp); err != nil {
|
||||
@@ -33,6 +37,7 @@ func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload)
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
// Heartbeat sends a heartbeat payload to the control plane and returns the response result.
|
||||
func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*protocol.HeartbeatResult, error) {
|
||||
resp := protocol.APIResponse[protocol.HeartbeatData]{}
|
||||
if err := c.base.PostJSON(ctx, "/api/v1/agent/nodes/heartbeat", payload, &resp); err != nil {
|
||||
@@ -48,6 +53,7 @@ func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) (*
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetActiveConfig retrieves the current active configuration from the control plane server.
|
||||
func (c *Client) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) {
|
||||
resp := protocol.APIResponse[protocol.ActiveConfigResponse]{}
|
||||
if err := c.base.GetJSON(ctx, "/api/v1/agent/config-versions/active", &resp); err != nil {
|
||||
@@ -59,6 +65,7 @@ func (c *Client) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigRes
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
// ReportApplyLog reports the configuration application logs back to the control plane.
|
||||
func (c *Client) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
|
||||
resp := protocol.APIResponse[json.RawMessage]{}
|
||||
if err := c.base.PostJSON(ctx, "/api/v1/agent/apply-logs", payload, &resp); err != nil {
|
||||
@@ -67,6 +74,7 @@ func (c *Client) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPa
|
||||
return edgehttp.APIError(resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// SyncWAFIPGroups synchronizes WAF IP groups with the control plane server.
|
||||
func (c *Client) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) {
|
||||
resp := protocol.APIResponse[protocol.WAFIPGroupSyncResponse]{}
|
||||
if err := c.base.PostJSON(ctx, "/api/v1/agent/waf/ip-groups/sync", payload, &resp); err != nil {
|
||||
@@ -78,18 +86,20 @@ func (c *Client) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGrou
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
// DownloadPagesDeploymentPackage downloads the deployment package for the given Pages deployment ID.
|
||||
func (c *Client) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
|
||||
res, err := c.base.DoRaw(ctx, http.MethodGet, fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
defer func() { _ = res.Body.Close() }()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return nil, edgehttp.ReadHTTPError(res)
|
||||
}
|
||||
return io.ReadAll(res.Body)
|
||||
}
|
||||
|
||||
// SetToken updates the authentication token used for API requests.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.base.SetToken(token)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
// Package logging configures structured logging for the agent process.
|
||||
package logging
|
||||
|
||||
import edgelogging "github.com/Rain-kl/Wavelet/internal/apps/edge/logging"
|
||||
|
||||
// Setup initialises structured logging for the agent process.
|
||||
func Setup() {
|
||||
edgelogging.Setup(edgelogging.Options{AddSource: true})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package nginx manages OpenResty configuration, runtime, and supporting assets.
|
||||
package nginx
|
||||
|
||||
import (
|
||||
@@ -26,10 +27,26 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
)
|
||||
|
||||
// RuntimeConfigDirPlaceholder is substituted into generated configs at apply time.
|
||||
const RuntimeConfigDirPlaceholder = "__OPENFLARE_RUNTIME_CONFIG_DIR__"
|
||||
const ResolverDirectivePlaceholder = "__OPENFLARE_RESOLVER_DIRECTIVE__"
|
||||
const WAFIPGroupsConfigFileName = "waf_ip_groups.json"
|
||||
|
||||
// ResolverDirectivePlaceholder is substituted into generated configs at apply time.
|
||||
const ResolverDirectivePlaceholder = "__OPENFLARE_RESOLVER_DIRECTIVE__"
|
||||
|
||||
// WAFIPGroupsConfigFileName is the runtime filename for synced WAF IP group data.
|
||||
const WAFIPGroupsConfigFileName = "waf_ip_groups.json"
|
||||
const powConfigFileName = "pow_config.json"
|
||||
|
||||
const (
|
||||
nginxConfigFilePerm = 0o644
|
||||
nginxPrivateKeyFilePerm = 0o600
|
||||
nginxDirPerm = 0o755
|
||||
stubStatusCheckTimeout = 1500 * time.Millisecond
|
||||
nginxVersionSubmatchCount = 2
|
||||
resolverAddressCapacity = 2
|
||||
)
|
||||
|
||||
// Executor controls OpenResty validation, reload, health, and lifecycle operations.
|
||||
type Executor interface {
|
||||
Test(ctx context.Context) error
|
||||
Reload(ctx context.Context) error
|
||||
@@ -38,44 +55,49 @@ type Executor interface {
|
||||
Restart(ctx context.Context) error
|
||||
}
|
||||
|
||||
// CommandRunner executes external commands on behalf of an Executor.
|
||||
type CommandRunner interface {
|
||||
Run(ctx context.Context, name string, args ...string) ([]byte, error)
|
||||
}
|
||||
|
||||
// OSCommandRunner runs commands using the host operating system.
|
||||
type OSCommandRunner struct{}
|
||||
|
||||
// Run executes the named command and returns its combined output.
|
||||
func (r *OSCommandRunner) Run(ctx context.Context, name string, args ...string) ([]byte, error) {
|
||||
slog.Debug("OSCommandRunner starting command", "name", name, "args", args)
|
||||
tmpFile, err := os.CreateTemp("", "openflare-cmd-*")
|
||||
if err != nil {
|
||||
slog.Error("OSCommandRunner failed to create temp file, falling back to CombinedOutput", "error", err)
|
||||
cmd := exec.CommandContext(ctx, name, args...)
|
||||
cmd := exec.CommandContext(ctx, name, args...) //nolint:gosec // command name and args come from trusted OpenResty management paths
|
||||
output, outErr := cmd.CombinedOutput()
|
||||
slog.Debug("OSCommandRunner finished CombinedOutput", "name", name, "error", outErr)
|
||||
return output, outErr
|
||||
}
|
||||
defer os.Remove(tmpFile.Name())
|
||||
defer func() { _ = os.Remove(tmpFile.Name()) }()
|
||||
|
||||
cmd := exec.CommandContext(ctx, name, args...)
|
||||
cmd := exec.CommandContext(ctx, name, args...) //nolint:gosec // command name and args come from trusted OpenResty management paths
|
||||
cmd.Stdout = tmpFile
|
||||
cmd.Stderr = tmpFile
|
||||
|
||||
slog.Debug("OSCommandRunner executing cmd.Run()", "name", name)
|
||||
runErr := cmd.Run()
|
||||
slog.Debug("OSCommandRunner cmd.Run() returned", "name", name, "error", runErr)
|
||||
tmpFile.Close()
|
||||
_ = tmpFile.Close()
|
||||
|
||||
output, _ := os.ReadFile(tmpFile.Name())
|
||||
slog.Debug("OSCommandRunner command complete", "name", name, "output_len", len(output))
|
||||
return output, runErr
|
||||
}
|
||||
|
||||
// PathExecutor runs OpenResty using a configured binary and config path.
|
||||
type PathExecutor struct {
|
||||
Path string
|
||||
ConfigPath string
|
||||
Runner CommandRunner
|
||||
}
|
||||
|
||||
// Test validates the current OpenResty configuration.
|
||||
func (e *PathExecutor) Test(ctx context.Context) error {
|
||||
slog.Debug("running openresty test with binary", "path", e.Path, "config", e.ConfigPath)
|
||||
output, err := e.Runner.Run(ctx, e.Path, "-t", "-c", e.ConfigPath)
|
||||
@@ -86,6 +108,7 @@ func (e *PathExecutor) Test(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Reload reloads OpenResty or starts it when no runtime process is running.
|
||||
func (e *PathExecutor) Reload(ctx context.Context) error {
|
||||
slog.Debug("running openresty reload with binary", "path", e.Path, "config", e.ConfigPath)
|
||||
output, err := e.Runner.Run(ctx, e.Path, "-s", "reload", "-c", e.ConfigPath)
|
||||
@@ -104,6 +127,7 @@ func (e *PathExecutor) Reload(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// EnsureRuntime validates configuration and reloads the OpenResty runtime.
|
||||
func (e *PathExecutor) EnsureRuntime(ctx context.Context, _ bool) error {
|
||||
if err := e.Test(ctx); err != nil {
|
||||
return err
|
||||
@@ -111,10 +135,12 @@ func (e *PathExecutor) EnsureRuntime(ctx context.Context, _ bool) error {
|
||||
return e.Reload(ctx)
|
||||
}
|
||||
|
||||
// CheckHealth reports whether the OpenResty configuration is valid.
|
||||
func (e *PathExecutor) CheckHealth(ctx context.Context) error {
|
||||
return e.Test(ctx)
|
||||
}
|
||||
|
||||
// Restart stops and starts the OpenResty runtime process.
|
||||
func (e *PathExecutor) Restart(ctx context.Context) error {
|
||||
slog.Info("restarting openresty with binary", "path", e.Path, "config", e.ConfigPath)
|
||||
output, err := e.Runner.Run(ctx, e.Path, "-s", "quit", "-c", e.ConfigPath)
|
||||
@@ -132,6 +158,7 @@ func (e *PathExecutor) Restart(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Manager applies OpenResty configuration and manages runtime assets.
|
||||
type Manager struct {
|
||||
MainConfigPath string
|
||||
RouteConfigPath string
|
||||
@@ -148,9 +175,12 @@ type Manager struct {
|
||||
Executor Executor
|
||||
}
|
||||
|
||||
// ApplyStatus reports the outcome of an OpenResty configuration apply.
|
||||
type ApplyStatus string
|
||||
|
||||
// Apply outcome status values.
|
||||
const (
|
||||
// ApplyStatusSuccess indicates the configuration was applied successfully.
|
||||
ApplyStatusSuccess ApplyStatus = "success"
|
||||
ApplyStatusWarning ApplyStatus = "warning"
|
||||
ApplyStatusFatal ApplyStatus = "fatal"
|
||||
@@ -188,6 +218,7 @@ const safeDefaultFallbackObservabilityServerBlock = `
|
||||
}
|
||||
`
|
||||
|
||||
// ApplyOutcome contains the status and message from a configuration apply.
|
||||
type ApplyOutcome struct {
|
||||
Status ApplyStatus
|
||||
Message string
|
||||
@@ -197,6 +228,7 @@ type wafIPGroupsRuntimeConfig struct {
|
||||
Groups map[string]protocol.WAFIPGroup `json:"groups"`
|
||||
}
|
||||
|
||||
// Apply writes, validates, and activates new OpenResty configuration files.
|
||||
func (m *Manager) Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) ApplyOutcome {
|
||||
slog.Info("openresty apply started", "main_config", m.MainConfigPath, "route_config", m.RouteConfigPath, "cert_files", len(supportFiles))
|
||||
backup, err := m.backup()
|
||||
@@ -236,11 +268,11 @@ func (m *Manager) writeTargetFiles(mainConfig string, routeConfig string, suppor
|
||||
slog.Warn("runtime-resolved hostname upstreams detected without available resolvers; hostname origin requests may fail until resolvers are configured")
|
||||
}
|
||||
renderedMainConfig := m.renderMainConfig(mainConfig)
|
||||
if err := os.WriteFile(m.MainConfigPath, []byte(renderedMainConfig), 0o644); err != nil {
|
||||
if err := os.WriteFile(m.MainConfigPath, []byte(renderedMainConfig), nginxConfigFilePerm); err != nil {
|
||||
return err
|
||||
}
|
||||
renderedRouteConfig := m.renderRouteConfig(routeConfig)
|
||||
if err := os.WriteFile(m.RouteConfigPath, []byte(renderedRouteConfig), 0o644); err != nil {
|
||||
if err := os.WriteFile(m.RouteConfigPath, []byte(renderedRouteConfig), nginxConfigFilePerm); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
@@ -293,6 +325,7 @@ func fatalApplyOutcome(err error) ApplyOutcome {
|
||||
}
|
||||
}
|
||||
|
||||
// EnsureLuaAssets synchronizes managed Lua and static assets to the runtime directory.
|
||||
func (m *Manager) EnsureLuaAssets() error {
|
||||
if strings.TrimSpace(m.LuaDir) == "" {
|
||||
return nil
|
||||
@@ -317,12 +350,13 @@ func (m *Manager) EnsureLuaAssets() error {
|
||||
files = append(files, managedFile{
|
||||
Path: filepath.ToSlash(relativePath),
|
||||
Content: []byte(file.Content),
|
||||
Mode: 0o644,
|
||||
Mode: nginxConfigFilePerm,
|
||||
})
|
||||
}
|
||||
return syncManagedFiles(m.LuaDir, files)
|
||||
}
|
||||
|
||||
// EnsureRuntime validates and reloads the current OpenResty runtime configuration.
|
||||
func (m *Manager) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
if m.Executor == nil {
|
||||
return errors.New("executor 未配置")
|
||||
@@ -331,6 +365,7 @@ func (m *Manager) EnsureRuntime(ctx context.Context, recreate bool) error {
|
||||
return m.Executor.EnsureRuntime(ctx, recreate)
|
||||
}
|
||||
|
||||
// EnsureSafeFallbackRuntime starts a minimal safe default OpenResty runtime.
|
||||
func (m *Manager) EnsureSafeFallbackRuntime(ctx context.Context, reason string) error {
|
||||
if m.Executor == nil {
|
||||
return errors.New("executor 未配置")
|
||||
@@ -350,6 +385,7 @@ func (m *Manager) EnsureSafeFallbackRuntime(ctx context.Context, reason string)
|
||||
return nil
|
||||
}
|
||||
|
||||
// CheckHealth verifies that OpenResty configuration and health endpoints are available.
|
||||
func (m *Manager) CheckHealth(ctx context.Context) error {
|
||||
if m.Executor == nil {
|
||||
return errors.New("executor 未配置")
|
||||
@@ -365,6 +401,7 @@ func (m *Manager) CheckHealth(ctx context.Context) error {
|
||||
return m.checkStubStatus(ctx)
|
||||
}
|
||||
|
||||
// Restart restarts the OpenResty runtime process.
|
||||
func (m *Manager) Restart(ctx context.Context) error {
|
||||
if m.Executor == nil {
|
||||
return errors.New("executor 未配置")
|
||||
@@ -373,6 +410,7 @@ func (m *Manager) Restart(ctx context.Context) error {
|
||||
return m.Executor.Restart(ctx)
|
||||
}
|
||||
|
||||
// CurrentChecksum returns a stable checksum for the active OpenResty configuration bundle.
|
||||
func (m *Manager) CurrentChecksum() (string, error) {
|
||||
if m.RouteConfigPath == "" {
|
||||
return "", errors.New("route config path 不能为空")
|
||||
@@ -435,6 +473,7 @@ func (m *Manager) CurrentChecksum() (string, error) {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// WAFIPGroupChecksums returns checksums for locally synced WAF IP groups.
|
||||
func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) {
|
||||
config, err := m.readWAFIPGroupsRuntimeConfig()
|
||||
if err != nil {
|
||||
@@ -449,6 +488,7 @@ func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// SyncWAFIPGroups writes WAF IP group definitions to the runtime config directory.
|
||||
func (m *Manager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error {
|
||||
if m.RuntimeConfigDir == "" || len(groups) == 0 {
|
||||
return nil
|
||||
@@ -470,11 +510,11 @@ func (m *Manager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(m.RuntimeConfigDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(m.RuntimeConfigDir, nginxDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
path := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsConfigFileName)
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
if err := os.WriteFile(path, data, nginxConfigFilePerm); err != nil {
|
||||
return fmt.Errorf("write %s: %w", WAFIPGroupsConfigFileName, err)
|
||||
}
|
||||
slog.Info("synced waf ip groups", "path", path, "group_count", len(groups))
|
||||
@@ -487,7 +527,7 @@ func (m *Manager) readWAFIPGroupsRuntimeConfig() (*wafIPGroupsRuntimeConfig, err
|
||||
return config, nil
|
||||
}
|
||||
path := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsConfigFileName)
|
||||
data, err := os.ReadFile(path)
|
||||
data, err := os.ReadFile(path) //nolint:gosec // path is under managed RuntimeConfigDir
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return config, nil
|
||||
@@ -506,6 +546,7 @@ func (m *Manager) readWAFIPGroupsRuntimeConfig() (*wafIPGroupsRuntimeConfig, err
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// ExecutorOptions configures construction of an OpenResty Executor.
|
||||
type ExecutorOptions struct {
|
||||
NginxPath string
|
||||
MainConfigPath string
|
||||
@@ -517,6 +558,7 @@ type ExecutorOptions struct {
|
||||
OpenrestyObservabilityPort int
|
||||
}
|
||||
|
||||
// NewExecutor creates an Executor backed by a configured OpenResty binary.
|
||||
func NewExecutor(options ExecutorOptions) Executor {
|
||||
runner := &OSCommandRunner{}
|
||||
return &PathExecutor{
|
||||
@@ -526,6 +568,7 @@ func NewExecutor(options ExecutorOptions) Executor {
|
||||
}
|
||||
}
|
||||
|
||||
// DetectVersion returns the OpenResty version reported by the configured binary.
|
||||
func DetectVersion(ctx context.Context, options ExecutorOptions) string {
|
||||
version, err := detectVersion(ctx, options, &OSCommandRunner{})
|
||||
if err != nil {
|
||||
@@ -556,7 +599,7 @@ func detectVersion(ctx context.Context, options ExecutorOptions, runner CommandR
|
||||
|
||||
func parseExtVersion(output string) string {
|
||||
matches := nginxVersionPattern.FindStringSubmatch(output)
|
||||
if len(matches) != 2 {
|
||||
if len(matches) != nginxVersionSubmatchCount {
|
||||
return ""
|
||||
}
|
||||
return matches[1]
|
||||
@@ -606,24 +649,24 @@ func (m *Manager) backup() (*backupState, error) {
|
||||
if m.RouteConfigPath == "" {
|
||||
return nil, errors.New("route config path 不能为空")
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(m.MainConfigPath), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(m.MainConfigPath), nginxDirPerm); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), nginxDirPerm); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if m.AccessLogPath != "" {
|
||||
if err := os.MkdirAll(filepath.Dir(m.AccessLogPath), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(m.AccessLogPath), nginxDirPerm); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if m.CertDir != "" {
|
||||
if err := os.MkdirAll(m.CertDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(m.CertDir, nginxDirPerm); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if m.RuntimeConfigDir != "" {
|
||||
if err := os.MkdirAll(m.RuntimeConfigDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(m.RuntimeConfigDir, nginxDirPerm); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
@@ -672,14 +715,14 @@ func (m *Manager) restore(state *backupState) error {
|
||||
}
|
||||
slog.Warn("restoring nginx backup", "main_existed", state.MainExisted, "route_existed", state.RouteExisted, "cert_files", len(state.Files))
|
||||
if state.MainExisted {
|
||||
if err := os.WriteFile(m.MainConfigPath, state.MainData, 0o644); err != nil {
|
||||
if err := os.WriteFile(m.MainConfigPath, state.MainData, nginxConfigFilePerm); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err := os.Remove(m.MainConfigPath); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if state.RouteExisted {
|
||||
if err := os.WriteFile(m.RouteConfigPath, state.RouteData, 0o644); err != nil {
|
||||
if err := os.WriteFile(m.RouteConfigPath, state.RouteData, nginxConfigFilePerm); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err := os.Remove(m.RouteConfigPath); err != nil && !os.IsNotExist(err) {
|
||||
@@ -690,7 +733,7 @@ func (m *Manager) restore(state *backupState) error {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := m.restoreRuntimeConfig(state.PowConfig, "pow_config.json"); err != nil {
|
||||
if err := m.restoreRuntimeConfig(state.PowConfig, powConfigFileName); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.restoreRuntimeConfig(state.WAFConfig, "waf_config.json"); err != nil {
|
||||
@@ -710,10 +753,10 @@ func (m *Manager) writePowConfig(supportFiles []protocol.SupportFile) error {
|
||||
if m.RuntimeConfigDir == "" {
|
||||
return nil
|
||||
}
|
||||
configPath := filepath.Join(m.RuntimeConfigDir, "pow_config.json")
|
||||
configPath := filepath.Join(m.RuntimeConfigDir, powConfigFileName)
|
||||
for _, file := range supportFiles {
|
||||
if file.Path == "pow_config.json" {
|
||||
if err := os.WriteFile(configPath, []byte(file.Content), 0o644); err != nil {
|
||||
if file.Path == powConfigFileName {
|
||||
if err := os.WriteFile(configPath, []byte(file.Content), nginxConfigFilePerm); err != nil {
|
||||
return fmt.Errorf("write pow_config.json: %w", err)
|
||||
}
|
||||
slog.Info("wrote pow config", "path", configPath, "size", len(file.Content))
|
||||
@@ -723,10 +766,10 @@ func (m *Manager) writePowConfig(supportFiles []protocol.SupportFile) error {
|
||||
if err := os.Remove(configPath); err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("remove pow_config.json: %w", err)
|
||||
}
|
||||
if err := removeLegacyPowConfig(filepath.Join(m.LuaDir, "pow_config.json")); err != nil {
|
||||
if err := removeLegacyPowConfig(filepath.Join(m.LuaDir, powConfigFileName)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := removeLegacyPowConfig(filepath.Join(m.CertDir, "pow_config.json")); err != nil {
|
||||
if err := removeLegacyPowConfig(filepath.Join(m.CertDir, powConfigFileName)); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
@@ -739,7 +782,7 @@ func (m *Manager) writeWAFConfig(supportFiles []protocol.SupportFile) error {
|
||||
configPath := filepath.Join(m.RuntimeConfigDir, "waf_config.json")
|
||||
for _, file := range supportFiles {
|
||||
if file.Path == "waf_config.json" {
|
||||
if err := os.WriteFile(configPath, []byte(file.Content), 0o644); err != nil {
|
||||
if err := os.WriteFile(configPath, []byte(file.Content), nginxConfigFilePerm); err != nil {
|
||||
return fmt.Errorf("write waf_config.json: %w", err)
|
||||
}
|
||||
slog.Info("wrote waf config", "path", configPath, "size", len(file.Content))
|
||||
@@ -759,7 +802,7 @@ func (m *Manager) writeSourceConfig(supportFiles []protocol.SupportFile) error {
|
||||
configPath := filepath.Join(m.RuntimeConfigDir, openrestyrender.SourceConfigFileName)
|
||||
for _, file := range supportFiles {
|
||||
if file.Path == openrestyrender.SourceConfigFileName {
|
||||
if err := os.WriteFile(configPath, []byte(file.Content), 0o644); err != nil {
|
||||
if err := os.WriteFile(configPath, []byte(file.Content), nginxConfigFilePerm); err != nil {
|
||||
return fmt.Errorf("write %s: %w", openrestyrender.SourceConfigFileName, err)
|
||||
}
|
||||
slog.Info("wrote openresty source config", "path", configPath, "size", len(file.Content))
|
||||
@@ -775,7 +818,7 @@ func (m *Manager) writeSourceConfig(supportFiles []protocol.SupportFile) error {
|
||||
func (m *Manager) writeManagedCertFiles(certFiles []protocol.SupportFile) error {
|
||||
files := make([]managedFile, 0, len(certFiles))
|
||||
for _, file := range certFiles {
|
||||
if file.Path == "pow_config.json" || file.Path == "waf_config.json" || file.Path == openrestyrender.SourceConfigFileName {
|
||||
if file.Path == powConfigFileName || file.Path == "waf_config.json" || file.Path == openrestyrender.SourceConfigFileName {
|
||||
continue
|
||||
}
|
||||
targetPath, err := m.certFileTargetPath(file.Path)
|
||||
@@ -817,10 +860,10 @@ func (m *Manager) readCertFiles() ([]protocol.SupportFile, error) {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if filepath.ToSlash(relativePath) == "pow_config.json" {
|
||||
if filepath.ToSlash(relativePath) == powConfigFileName {
|
||||
return nil
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
data, err := os.ReadFile(path) //nolint:gosec // path is under managed baseDir walk root
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -840,7 +883,7 @@ func (m *Manager) readCertFiles() ([]protocol.SupportFile, error) {
|
||||
}
|
||||
|
||||
func (m *Manager) readPowConfigFile() (*protocol.SupportFile, error) {
|
||||
return m.readRuntimeConfigFile("pow_config.json")
|
||||
return m.readRuntimeConfigFile(powConfigFileName)
|
||||
}
|
||||
|
||||
func (m *Manager) readRuntimeConfigFile(name string) (*protocol.SupportFile, error) {
|
||||
@@ -848,7 +891,7 @@ func (m *Manager) readRuntimeConfigFile(name string) (*protocol.SupportFile, err
|
||||
return nil, nil
|
||||
}
|
||||
configPath := filepath.Join(m.RuntimeConfigDir, name)
|
||||
data, err := os.ReadFile(configPath)
|
||||
data, err := os.ReadFile(configPath) //nolint:gosec // configPath is under managed RuntimeConfigDir
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, nil
|
||||
@@ -894,7 +937,7 @@ func (m *Manager) restoreRuntimeConfig(file *protocol.SupportFile, name string)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return os.WriteFile(configPath, []byte(file.Content), 0o644)
|
||||
return os.WriteFile(configPath, []byte(file.Content), nginxConfigFilePerm)
|
||||
}
|
||||
|
||||
func (m *Manager) writeSafeDefaultFallbackFiles() error {
|
||||
@@ -904,16 +947,16 @@ func (m *Manager) writeSafeDefaultFallbackFiles() error {
|
||||
if strings.TrimSpace(m.RouteConfigPath) == "" {
|
||||
return errors.New("route config path 不能为空")
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(m.MainConfigPath), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(m.MainConfigPath), nginxDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), nginxDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(m.RouteConfigPath, nil, 0o644); err != nil {
|
||||
if err := os.WriteFile(m.RouteConfigPath, nil, nginxConfigFilePerm); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(m.MainConfigPath, []byte(m.safeDefaultFallbackMainConfig()), 0o644); err != nil {
|
||||
if err := os.WriteFile(m.MainConfigPath, []byte(m.safeDefaultFallbackMainConfig()), nginxConfigFilePerm); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
@@ -928,10 +971,10 @@ func (m *Manager) safeDefaultFallbackMainConfig() string {
|
||||
}
|
||||
|
||||
func (m *Manager) checkStubStatus(ctx context.Context) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, 1500*time.Millisecond)
|
||||
ctx, cancel := context.WithTimeout(ctx, stubStatusCheckTimeout)
|
||||
defer cancel()
|
||||
openrestyStubUrl := fmt.Sprintf("http://127.0.0.1:%d/openflare/stub_status", m.OpenrestyObservabilityPort)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, openrestyStubUrl, nil)
|
||||
openrestyStubURL := fmt.Sprintf("http://127.0.0.1:%d/openflare/stub_status", m.OpenrestyObservabilityPort)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, openrestyStubURL, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -939,11 +982,11 @@ func (m *Manager) checkStubStatus(ctx context.Context) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("openresty health endpoint unreachable: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("openresty health endpoint returned %s", resp.Status)
|
||||
}
|
||||
slog.Debug("openresty health endpoint is healthy", "url", openrestyStubUrl)
|
||||
slog.Debug("openresty health endpoint is healthy", "url", openrestyStubURL)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -986,11 +1029,11 @@ func (m *Manager) certFileTargetPath(relativePath string) (string, error) {
|
||||
func certFileMode(relativePath string) fs.FileMode {
|
||||
switch strings.ToLower(filepath.Ext(strings.TrimSpace(relativePath))) {
|
||||
case ".crt", ".pem":
|
||||
return 0o644
|
||||
return nginxConfigFilePerm
|
||||
case ".key":
|
||||
return 0o600
|
||||
return nginxPrivateKeyFilePerm
|
||||
default:
|
||||
return 0o644
|
||||
return nginxConfigFilePerm
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1029,7 +1072,7 @@ func syncManagedFiles(baseDir string, files []managedFile) error {
|
||||
} else if err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(baseDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(baseDir, nginxDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -1060,14 +1103,14 @@ func syncManagedFiles(baseDir string, files []managedFile) error {
|
||||
if _, ok := desired[filepath.Clean(relativePath)]; ok {
|
||||
return nil
|
||||
}
|
||||
return os.Remove(path)
|
||||
return os.Remove(path) //nolint:gosec // path is resolved under the managed baseDir walk root
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, file := range desired {
|
||||
targetPath := filepath.Join(baseDir, file.Path)
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), nginxDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(targetPath, file.Content, file.Mode); err != nil {
|
||||
@@ -1119,10 +1162,10 @@ func (m *Manager) ensureMimeTypes() error {
|
||||
} else if !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(configDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(configDir, nginxDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(mimeTypesPath, []byte(DefaultMimeTypes), 0o644)
|
||||
return os.WriteFile(mimeTypesPath, []byte(DefaultMimeTypes), nginxConfigFilePerm)
|
||||
}
|
||||
|
||||
func (m *Manager) renderRouteConfig(content string) string {
|
||||
@@ -1183,6 +1226,7 @@ func (m *Manager) managedWAFLuaFiles() []protocol.SupportFile {
|
||||
return files
|
||||
}
|
||||
|
||||
// ObservabilityListenAddress returns the localhost listen address for stub_status.
|
||||
func ObservabilityListenAddress(port int) string {
|
||||
if port <= 0 {
|
||||
return ""
|
||||
@@ -1190,6 +1234,7 @@ func ObservabilityListenAddress(port int) string {
|
||||
return fmt.Sprintf("127.0.0.1:%d", port)
|
||||
}
|
||||
|
||||
// ResolverDirective renders the nginx resolver block for runtime upstream lookups.
|
||||
func ResolverDirective(explicitResolvers []string) string {
|
||||
resolvers := resolverAddresses(explicitResolvers)
|
||||
if len(resolvers) == 0 {
|
||||
@@ -1211,7 +1256,7 @@ func resolverAddresses(explicitResolvers []string) []string {
|
||||
|
||||
func parseResolverAddresses(content string, dockerMode bool) []string {
|
||||
lines := strings.Split(content, "\n")
|
||||
resolvers := make([]string, 0, 2)
|
||||
resolvers := make([]string, 0, resolverAddressCapacity)
|
||||
seen := make(map[string]struct{})
|
||||
for _, line := range lines {
|
||||
fields := strings.Fields(strings.TrimSpace(line))
|
||||
@@ -1242,6 +1287,7 @@ func isUsableDockerResolver(addr string) bool {
|
||||
return !ip.IsLoopback() && !ip.IsUnspecified()
|
||||
}
|
||||
|
||||
// RequiresRuntimeResolver reports whether originURL needs a runtime DNS resolver.
|
||||
func RequiresRuntimeResolver(originURL string) bool {
|
||||
parsed, err := url.Parse(strings.TrimSpace(originURL))
|
||||
if err != nil || parsed.Hostname() == "" {
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
package nginx
|
||||
|
||||
// DefaultMimeTypes is the embedded nginx mime.types map used by generated configs.
|
||||
const DefaultMimeTypes = `
|
||||
types {
|
||||
text/html html htm shtml;
|
||||
|
||||
@@ -146,6 +146,7 @@ ngx.header.content_type = "application/json"
|
||||
ngx.say(cjson.encode(payload))
|
||||
`
|
||||
|
||||
// ManagedObservabilityLuaFiles returns embedded Lua assets for OpenResty observability.
|
||||
func ManagedObservabilityLuaFiles() []protocol.SupportFile {
|
||||
return []protocol.SupportFile{
|
||||
{Path: "init.lua", Content: openRestyObservabilityInitLua},
|
||||
|
||||
@@ -542,6 +542,7 @@ end
|
||||
return M
|
||||
`
|
||||
|
||||
// ManagedPowLuaFiles returns embedded Lua assets for proof-of-work challenges.
|
||||
func ManagedPowLuaFiles() []protocol.SupportFile {
|
||||
return []protocol.SupportFile{
|
||||
{Path: "pow/runtime.lua", Content: openRestyPowRuntimeLua},
|
||||
@@ -552,6 +553,7 @@ func ManagedPowLuaFiles() []protocol.SupportFile {
|
||||
}
|
||||
}
|
||||
|
||||
// ManagedPowStaticFiles returns embedded static assets served by the PoW module.
|
||||
func ManagedPowStaticFiles() ([]protocol.SupportFile, error) {
|
||||
var files []protocol.SupportFile
|
||||
entries, err := powStaticFS.ReadDir("pow_static")
|
||||
|
||||
@@ -302,6 +302,7 @@ end
|
||||
return require("waf.runtime").check()
|
||||
`
|
||||
|
||||
// ManagedWAFLuaFiles returns the embedded Lua source files that must be deployed to the WAF runtime directory.
|
||||
func ManagedWAFLuaFiles() []protocol.SupportFile {
|
||||
return []protocol.SupportFile{
|
||||
{Path: "waf/runtime.lua", Content: openRestyWAFRuntimeLua},
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package observability provides system and service level observability data collection for the agent.
|
||||
package observability
|
||||
|
||||
import (
|
||||
@@ -15,6 +16,9 @@ import (
|
||||
edgeobs "github.com/Rain-kl/Wavelet/internal/apps/edge/observability"
|
||||
)
|
||||
|
||||
const nodeHealthEventInitialCapacity = 2
|
||||
|
||||
// BuildProfile collects the system profile and returns it only if the fingerprint has changed.
|
||||
func BuildProfile(cfg *config.Config, stateStore *state.Store) *protocol.NodeSystemProfile {
|
||||
profile := collectProfile(cfg)
|
||||
if profile == nil {
|
||||
@@ -38,6 +42,7 @@ func BuildProfile(cfg *config.Config, stateStore *state.Store) *protocol.NodeSys
|
||||
return profile
|
||||
}
|
||||
|
||||
// BuildSnapshot captures current system metrics and returns a metric snapshot.
|
||||
func BuildSnapshot(cfg *config.Config, stateStore *state.Store) *protocol.NodeMetricSnapshot {
|
||||
now := time.Now().UTC()
|
||||
metric := &protocol.NodeMetricSnapshot{
|
||||
@@ -79,6 +84,7 @@ func BuildSnapshot(cfg *config.Config, stateStore *state.Store) *protocol.NodeMe
|
||||
return metric
|
||||
}
|
||||
|
||||
// BuildOpenrestyObservation builds the OpenResty observation protocol model from the managed metrics.
|
||||
func BuildOpenrestyObservation(managed *ManagedOpenRestyMetrics) *protocol.NodeOpenrestyObservation {
|
||||
if managed == nil {
|
||||
return nil
|
||||
@@ -91,11 +97,12 @@ func BuildOpenrestyObservation(managed *ManagedOpenRestyMetrics) *protocol.NodeO
|
||||
}
|
||||
}
|
||||
|
||||
// BuildHealthEvents converts system snapshot health state into a list of health events.
|
||||
func BuildHealthEvents(snapshot *state.Snapshot) []protocol.NodeHealthEvent {
|
||||
if snapshot == nil {
|
||||
return []protocol.NodeHealthEvent{}
|
||||
}
|
||||
events := make([]protocol.NodeHealthEvent, 0, 2)
|
||||
events := make([]protocol.NodeHealthEvent, 0, nodeHealthEventInitialCapacity)
|
||||
nowUnix := time.Now().UTC().Unix()
|
||||
if strings.TrimSpace(snapshot.OpenrestyStatus) == protocol.OpenrestyStatusUnhealthy {
|
||||
events = append(events, protocol.NodeHealthEvent{
|
||||
@@ -147,4 +154,4 @@ func fingerprintProfile(profile *protocol.NodeSystemProfile) string {
|
||||
}
|
||||
sum := sha256.Sum256(raw)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -14,11 +15,15 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
)
|
||||
|
||||
const openRestyObservabilityPath = "/openflare/observability"
|
||||
const openRestyStubStatusPath = "/openflare/stub_status"
|
||||
const (
|
||||
openRestyObservabilityPath = "/openflare/observability"
|
||||
openRestyStubStatusPath = "/openflare/stub_status"
|
||||
stubStatusActiveMatchGroupCount = 2
|
||||
)
|
||||
|
||||
var stubStatusActivePattern = regexp.MustCompile(`Active connections:\s+(\d+)`)
|
||||
|
||||
// ManagedOpenRestyMetrics holds metrics collected from the local OpenResty instance.
|
||||
type ManagedOpenRestyMetrics struct {
|
||||
TrafficReport *protocol.NodeTrafficReport
|
||||
OpenrestyRxBytes int64
|
||||
@@ -39,7 +44,8 @@ type openRestyObservabilityResponse struct {
|
||||
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
|
||||
}
|
||||
|
||||
func CollectManagedOpenRestyMetrics(cfg *config.Config) *ManagedOpenRestyMetrics {
|
||||
// CollectManagedOpenRestyMetrics collects metrics from the local OpenResty observability endpoints.
|
||||
func CollectManagedOpenRestyMetrics(ctx context.Context, cfg *config.Config) *ManagedOpenRestyMetrics {
|
||||
if cfg == nil || cfg.OpenrestyObservabilityPort <= 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -48,7 +54,7 @@ func CollectManagedOpenRestyMetrics(cfg *config.Config) *ManagedOpenRestyMetrics
|
||||
client := &http.Client{Timeout: 1500 * time.Millisecond}
|
||||
|
||||
observabilityResp := openRestyObservabilityResponse{}
|
||||
if err := fetchLocalJSON(client, baseURL+openRestyObservabilityPath, &observabilityResp); err != nil {
|
||||
if err := fetchLocalJSON(ctx, client, baseURL+openRestyObservabilityPath, &observabilityResp); err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -67,31 +73,39 @@ func CollectManagedOpenRestyMetrics(cfg *config.Config) *ManagedOpenRestyMetrics
|
||||
OpenrestyTxBytes: observabilityResp.OpenrestyTxBytes,
|
||||
}
|
||||
|
||||
if text, err := fetchLocalText(client, baseURL+openRestyStubStatusPath); err == nil {
|
||||
if text, err := fetchLocalText(ctx, 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)
|
||||
func fetchLocalJSON(ctx context.Context, client *http.Client, url string, target any) error {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", url, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = 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)
|
||||
func fetchLocalText(ctx context.Context, client *http.Client, url string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", url, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("unexpected local stub status: %s", resp.Status)
|
||||
}
|
||||
@@ -104,7 +118,7 @@ func fetchLocalText(client *http.Client, url string) (string, error) {
|
||||
|
||||
func parseStubStatusActiveConnections(raw string) int64 {
|
||||
matches := stubStatusActivePattern.FindStringSubmatch(raw)
|
||||
if len(matches) != 2 {
|
||||
if len(matches) != stubStatusActiveMatchGroupCount {
|
||||
return 0
|
||||
}
|
||||
value, err := strconv.ParseInt(matches[1], 10, 64)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package observability
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -31,7 +32,7 @@ func TestCollectManagedOpenRestyMetrics(t *testing.T) {
|
||||
server.Start()
|
||||
defer server.Close()
|
||||
|
||||
metrics := CollectManagedOpenRestyMetrics(&config.Config{
|
||||
metrics := CollectManagedOpenRestyMetrics(context.Background(), &config.Config{
|
||||
OpenrestyObservabilityPort: port,
|
||||
})
|
||||
if metrics == nil || metrics.TrafficReport == nil {
|
||||
@@ -67,7 +68,7 @@ func TestNormalizeCountMapDropsEmptyKeys(t *testing.T) {
|
||||
|
||||
func TestCollectManagedOpenRestyMetricsHandlesUnavailableEndpoint(t *testing.T) {
|
||||
cfg := &config.Config{OpenrestyObservabilityPort: 1}
|
||||
if metrics := CollectManagedOpenRestyMetrics(cfg); metrics != nil {
|
||||
if metrics := CollectManagedOpenRestyMetrics(context.Background(), cfg); metrics != nil {
|
||||
t.Fatalf("expected nil metrics for unavailable endpoint, got %+v", metrics)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"regexp"
|
||||
"sort"
|
||||
@@ -28,6 +29,11 @@ type accessLogRecord struct {
|
||||
RequestLength int64 `json:"request_length"`
|
||||
}
|
||||
|
||||
const (
|
||||
combinedAccessLogMatchGroupCount = 5
|
||||
trafficTopDomainsLimit = 8
|
||||
)
|
||||
|
||||
var combinedAccessLogPattern = regexp.MustCompile(`^(\S+)\s+\S+\s+\S+\s+\[([^]]+)]\s+"\S+\s+(\S+)(?:\s+[^"]*)?"\s+(\d{3})\s+\S+`)
|
||||
|
||||
type trafficAggregate struct {
|
||||
@@ -43,11 +49,13 @@ type trafficAggregate struct {
|
||||
logs []protocol.NodeAccessLog
|
||||
}
|
||||
|
||||
// BuildTrafficReport generates a traffic report using access logs or falling back to managed metrics.
|
||||
func BuildTrafficReport(cfg *config.Config, stateStore *state.Store, managed *ManagedOpenRestyMetrics) *protocol.NodeTrafficReport {
|
||||
report, _, _ := BuildTrafficObservability(cfg, stateStore, managed)
|
||||
return report
|
||||
}
|
||||
|
||||
// BuildTrafficObservability returns the traffic report, parsed access logs, and managed metrics.
|
||||
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 {
|
||||
@@ -78,7 +86,7 @@ func readAccessLogDelta(cfg *config.Config, stateStore *state.Store) *trafficAgg
|
||||
}
|
||||
|
||||
logPath := managedAccessLogPath(cfg)
|
||||
file, err := os.Open(logPath)
|
||||
file, err := os.Open(logPath) //nolint:gosec // path is the configured managed access log location
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
if snapshot.AccessLogOffset != 0 {
|
||||
@@ -167,7 +175,7 @@ func (aggregate *trafficAggregate) consume(line []byte) {
|
||||
}
|
||||
|
||||
aggregate.requestCount++
|
||||
if record.Status >= 500 {
|
||||
if record.Status >= http.StatusInternalServerError {
|
||||
aggregate.errorCount++
|
||||
}
|
||||
if record.Status > 0 {
|
||||
@@ -234,7 +242,7 @@ func parseJSONAccessLogRecord(raw string) (parsedAccessLogRecord, bool) {
|
||||
|
||||
func parseCombinedAccessLogRecord(raw string) (parsedAccessLogRecord, bool) {
|
||||
matches := combinedAccessLogPattern.FindStringSubmatch(raw)
|
||||
if len(matches) != 5 {
|
||||
if len(matches) != combinedAccessLogMatchGroupCount {
|
||||
return parsedAccessLogRecord{}, false
|
||||
}
|
||||
timestamp, err := parseAccessLogTime(matches[2])
|
||||
@@ -265,7 +273,7 @@ func (aggregate *trafficAggregate) report() *protocol.NodeTrafficReport {
|
||||
ErrorCount: aggregate.errorCount,
|
||||
UniqueVisitorCount: int64(len(aggregate.visitors)),
|
||||
StatusCodes: cloneTrafficCounts(aggregate.statusCodes, 0),
|
||||
TopDomains: topCounts(aggregate.topDomains, 8),
|
||||
TopDomains: topCounts(aggregate.topDomains, trafficTopDomainsLimit),
|
||||
SourceCountries: map[string]int64{},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,43 +1,99 @@
|
||||
// Package protocol defines type aliases and constants for the agent protocol.
|
||||
package protocol
|
||||
|
||||
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
|
||||
// APIResponse is an alias for pkgprotocol.APIResponse.
|
||||
type APIResponse[T any] = pkgprotocol.APIResponse[T]
|
||||
|
||||
// HeartbeatData is an alias for pkgprotocol.HeartbeatData.
|
||||
type HeartbeatData = pkgprotocol.HeartbeatData
|
||||
|
||||
// HeartbeatResult is an alias for pkgprotocol.HeartbeatResult.
|
||||
type HeartbeatResult = pkgprotocol.HeartbeatResult
|
||||
|
||||
// AgentSettings is an alias for pkgprotocol.AgentSettings.
|
||||
type AgentSettings = pkgprotocol.AgentSettings
|
||||
|
||||
// WSMessage is an alias for pkgprotocol.WSMessage.
|
||||
type WSMessage = pkgprotocol.WSMessage
|
||||
|
||||
// WSOutboundMessage is an alias for pkgprotocol.WSOutboundMessage.
|
||||
type WSOutboundMessage = pkgprotocol.WSOutboundMessage
|
||||
|
||||
// WebSocketConnection is an alias for pkgprotocol.WebSocketConnection.
|
||||
type WebSocketConnection = pkgprotocol.WebSocketConnection
|
||||
|
||||
// NodePayload is an alias for pkgprotocol.NodePayload.
|
||||
type NodePayload = pkgprotocol.NodePayload
|
||||
|
||||
// NodeSystemProfile is an alias for pkgprotocol.NodeSystemProfile.
|
||||
type NodeSystemProfile = pkgprotocol.NodeSystemProfile
|
||||
|
||||
// NodeMetricSnapshot is an alias for pkgprotocol.NodeMetricSnapshot.
|
||||
type NodeMetricSnapshot = pkgprotocol.NodeMetricSnapshot
|
||||
|
||||
// NodeOpenrestyObservation is an alias for pkgprotocol.NodeOpenrestyObservation.
|
||||
type NodeOpenrestyObservation = pkgprotocol.NodeOpenrestyObservation
|
||||
|
||||
// NodeTrafficReport is an alias for pkgprotocol.NodeTrafficReport.
|
||||
type NodeTrafficReport = pkgprotocol.NodeTrafficReport
|
||||
|
||||
// NodeAccessLog is an alias for pkgprotocol.NodeAccessLog.
|
||||
type NodeAccessLog = pkgprotocol.NodeAccessLog
|
||||
|
||||
// BufferedObservabilityRecord is an alias for pkgprotocol.BufferedObservabilityRecord.
|
||||
type BufferedObservabilityRecord = pkgprotocol.BufferedObservabilityRecord
|
||||
|
||||
// NodeHealthEvent is an alias for pkgprotocol.NodeHealthEvent.
|
||||
type NodeHealthEvent = pkgprotocol.NodeHealthEvent
|
||||
|
||||
// RegisterNodeResponse is an alias for pkgprotocol.RegisterNodeResponse.
|
||||
type RegisterNodeResponse = pkgprotocol.RegisterNodeResponse
|
||||
|
||||
// ApplyLogPayload is an alias for pkgprotocol.ApplyLogPayload.
|
||||
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
|
||||
|
||||
// ActiveConfigResponse is an alias for pkgprotocol.ActiveConfigResponse.
|
||||
type ActiveConfigResponse = pkgprotocol.ActiveConfigResponse
|
||||
|
||||
// ActiveConfigMeta is an alias for pkgprotocol.ActiveConfigMeta.
|
||||
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
|
||||
|
||||
// WAFIPGroup is an alias for pkgprotocol.WAFIPGroup.
|
||||
type WAFIPGroup = pkgprotocol.WAFIPGroup
|
||||
|
||||
// WAFIPGroupSyncRequest is an alias for pkgprotocol.WAFIPGroupSyncRequest.
|
||||
type WAFIPGroupSyncRequest = pkgprotocol.WAFIPGroupSyncRequest
|
||||
|
||||
// WAFIPGroupSyncResponse is an alias for pkgprotocol.WAFIPGroupSyncResponse.
|
||||
type WAFIPGroupSyncResponse = pkgprotocol.WAFIPGroupSyncResponse
|
||||
|
||||
// SupportFile is an alias for pkgprotocol.SupportFile.
|
||||
type SupportFile = pkgprotocol.SupportFile
|
||||
|
||||
const (
|
||||
WSMessageTypeStatus = pkgprotocol.WSMessageTypeStatus
|
||||
WSMessageTypeSettings = pkgprotocol.WSMessageTypeSettings
|
||||
WSMessageTypeActiveConfig = pkgprotocol.WSMessageTypeActiveConfig
|
||||
// WSMessageTypeStatus is an alias for pkgprotocol.WSMessageTypeStatus.
|
||||
WSMessageTypeStatus = pkgprotocol.WSMessageTypeStatus
|
||||
// WSMessageTypeSettings is an alias for pkgprotocol.WSMessageTypeSettings.
|
||||
WSMessageTypeSettings = pkgprotocol.WSMessageTypeSettings
|
||||
// WSMessageTypeActiveConfig is an alias for pkgprotocol.WSMessageTypeActiveConfig.
|
||||
WSMessageTypeActiveConfig = pkgprotocol.WSMessageTypeActiveConfig
|
||||
// WSMessageTypeForceSyncConfig is an alias for pkgprotocol.WSMessageTypeForceSyncConfig.
|
||||
WSMessageTypeForceSyncConfig = pkgprotocol.WSMessageTypeForceSyncConfig
|
||||
WSMessageTypeWAFIPGroups = pkgprotocol.WSMessageTypeWAFIPGroups
|
||||
WSMessageTypePing = pkgprotocol.WSMessageTypePing
|
||||
WSMessageTypePong = pkgprotocol.WSMessageTypePong
|
||||
// WSMessageTypeWAFIPGroups is an alias for pkgprotocol.WSMessageTypeWAFIPGroups.
|
||||
WSMessageTypeWAFIPGroups = pkgprotocol.WSMessageTypeWAFIPGroups
|
||||
// WSMessageTypePing is an alias for pkgprotocol.WSMessageTypePing.
|
||||
WSMessageTypePing = pkgprotocol.WSMessageTypePing
|
||||
// WSMessageTypePong is an alias for pkgprotocol.WSMessageTypePong.
|
||||
WSMessageTypePong = pkgprotocol.WSMessageTypePong
|
||||
)
|
||||
|
||||
const (
|
||||
OpenrestyStatusHealthy = pkgprotocol.OpenrestyStatusHealthy
|
||||
// OpenrestyStatusHealthy is an alias for pkgprotocol.OpenrestyStatusHealthy.
|
||||
OpenrestyStatusHealthy = pkgprotocol.OpenrestyStatusHealthy
|
||||
// OpenrestyStatusUnhealthy is an alias for pkgprotocol.OpenrestyStatusUnhealthy.
|
||||
OpenrestyStatusUnhealthy = pkgprotocol.OpenrestyStatusUnhealthy
|
||||
OpenrestyStatusUnknown = pkgprotocol.OpenrestyStatusUnknown
|
||||
)
|
||||
// OpenrestyStatusUnknown is an alias for pkgprotocol.OpenrestyStatusUnknown.
|
||||
OpenrestyStatusUnknown = pkgprotocol.OpenrestyStatusUnknown
|
||||
)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package state persists agent runtime state and observability snapshots.
|
||||
package state
|
||||
|
||||
import (
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
|
||||
const observabilityBufferWindowSeconds = 60
|
||||
|
||||
// ObservabilityBufferRecord stores observability data for a single time window.
|
||||
type ObservabilityBufferRecord struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
Snapshot *protocol.NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
@@ -22,15 +24,18 @@ type ObservabilityBufferRecord struct {
|
||||
QueuedAtUnix int64 `json:"queued_at_unix"`
|
||||
}
|
||||
|
||||
// ObservabilityBufferStore persists observability records to disk for replay on heartbeat.
|
||||
type ObservabilityBufferStore struct {
|
||||
path string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// NewObservabilityBufferStore creates a store backed by the file at path.
|
||||
func NewObservabilityBufferStore(path string) *ObservabilityBufferStore {
|
||||
return &ObservabilityBufferStore{path: filepath.Clean(path)}
|
||||
}
|
||||
|
||||
// Upsert inserts or merges an observability record and prunes entries older than retainAfterUnix.
|
||||
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
|
||||
@@ -113,6 +118,7 @@ func accessLogKey(item protocol.NodeAccessLog) string {
|
||||
return strconv.FormatInt(item.LoggedAtUnix, 10) + "|" + item.RemoteAddr + "|" + item.Host + "|" + item.Path + "|" + strconv.Itoa(item.StatusCode)
|
||||
}
|
||||
|
||||
// Replayable returns buffered records from windows before currentWindowStartedAtUnix.
|
||||
func (s *ObservabilityBufferStore) Replayable(currentWindowStartedAtUnix int64, retainAfterUnix int64) ([]ObservabilityBufferRecord, error) {
|
||||
if s == nil {
|
||||
return nil, nil
|
||||
@@ -138,6 +144,7 @@ func (s *ObservabilityBufferStore) Replayable(currentWindowStartedAtUnix int64,
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// Ack removes acknowledged observability windows and prunes entries older than retainAfterUnix.
|
||||
func (s *ObservabilityBufferStore) Ack(windowStartedAtUnix []int64, retainAfterUnix int64) error {
|
||||
if s == nil || len(windowStartedAtUnix) == 0 {
|
||||
return nil
|
||||
@@ -185,16 +192,17 @@ func (s *ObservabilityBufferStore) loadUnlocked() ([]ObservabilityBufferRecord,
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) saveUnlocked(records []ObservabilityBufferRecord) error {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), stateDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(records, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(s.path, data, 0o644)
|
||||
return os.WriteFile(s.path, data, stateFilePerm)
|
||||
}
|
||||
|
||||
// ObservabilityWindowStartedAt calculates the start of the 60-second window for the given metrics, openresty observation, or traffic report.
|
||||
func ObservabilityWindowStartedAt(snapshot *protocol.NodeMetricSnapshot, openresty *protocol.NodeOpenrestyObservation, traffic *protocol.NodeTrafficReport) int64 {
|
||||
if traffic != nil && traffic.WindowStartedAtUnix > 0 {
|
||||
return traffic.WindowStartedAtUnix - (traffic.WindowStartedAtUnix % observabilityBufferWindowSeconds)
|
||||
|
||||
@@ -9,6 +9,13 @@ import (
|
||||
"sync"
|
||||
)
|
||||
|
||||
const (
|
||||
stateDirPerm = 0o750
|
||||
stateFilePerm = 0o600
|
||||
nodeIDRandomBytes = 8
|
||||
)
|
||||
|
||||
// Snapshot represents the state of the agent at a given point in time.
|
||||
type Snapshot struct {
|
||||
NodeID string `json:"node_id"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
@@ -26,21 +33,25 @@ type Snapshot struct {
|
||||
AccessLogOffset int64 `json:"access_log_offset"`
|
||||
}
|
||||
|
||||
// Store manages the storage and retrieval of the agent state snapshot.
|
||||
type Store struct {
|
||||
path string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// NewStore creates a new Store instance at the given path.
|
||||
func NewStore(path string) *Store {
|
||||
return &Store{path: filepath.Clean(path)}
|
||||
}
|
||||
|
||||
// Load loads the snapshot from the store.
|
||||
func (s *Store) Load() (*Snapshot, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.loadUnlocked()
|
||||
}
|
||||
|
||||
// EnsureNodeID returns the existing node ID, or generates and saves a new one if it does not exist.
|
||||
func (s *Store) EnsureNodeID() (string, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
@@ -62,6 +73,7 @@ func (s *Store) EnsureNodeID() (string, error) {
|
||||
return snapshot.NodeID, nil
|
||||
}
|
||||
|
||||
// Save saves the given snapshot to the store.
|
||||
func (s *Store) Save(snapshot *Snapshot) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
@@ -87,18 +99,18 @@ func (s *Store) loadUnlocked() (*Snapshot, error) {
|
||||
}
|
||||
|
||||
func (s *Store) saveUnlocked(snapshot *Snapshot) error {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), stateDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(snapshot, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(s.path, data, 0o644)
|
||||
return os.WriteFile(s.path, data, stateFilePerm)
|
||||
}
|
||||
|
||||
func newNodeID() (string, error) {
|
||||
buf := make([]byte, 8)
|
||||
buf := make([]byte, nodeIDRandomBytes)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package sync applies control-plane configuration to the local agent runtime.
|
||||
package sync
|
||||
|
||||
import (
|
||||
@@ -10,6 +11,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
@@ -18,6 +20,12 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
)
|
||||
|
||||
const (
|
||||
pagesMaxExtractedFileBytes = 100 * 1024 * 1024
|
||||
pagesDirPerm = 0o755
|
||||
pagesManifestFilePerm = 0o644
|
||||
)
|
||||
|
||||
type pagesSourceDocument struct {
|
||||
Routes []pagesSourceRoute `json:"routes"`
|
||||
}
|
||||
@@ -66,7 +74,7 @@ func (s *Service) ensurePagesDeployment(ctx context.Context, deployment pagesDep
|
||||
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)
|
||||
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) {
|
||||
@@ -83,7 +91,7 @@ func referencedPagesDeployments(config *protocol.ActiveConfigResponse) ([]pagesD
|
||||
}
|
||||
var doc pagesSourceDocument
|
||||
if err := json.Unmarshal([]byte(config.SourceConfigJSON), &doc); err != nil {
|
||||
return nil, fmt.Errorf("decode Pages references: %w", err)
|
||||
return nil, fmt.Errorf("decode pages references: %w", err)
|
||||
}
|
||||
seen := make(map[uint]struct{})
|
||||
result := make([]pagesDeploymentSource, 0)
|
||||
@@ -94,7 +102,7 @@ func referencedPagesDeployments(config *protocol.ActiveConfigResponse) ([]pagesD
|
||||
deploymentID := route.PagesDeployment.DeploymentID
|
||||
checksum := strings.TrimSpace(route.PagesDeployment.Checksum)
|
||||
if deploymentID == 0 || checksum == "" {
|
||||
return nil, errors.New("Pages deployment snapshot is incomplete")
|
||||
return nil, errors.New("pages deployment snapshot is incomplete")
|
||||
}
|
||||
if _, ok := seen[deploymentID]; ok {
|
||||
continue
|
||||
@@ -152,7 +160,7 @@ func findCommonRootPrefix(files []*zip.File) (string, error) {
|
||||
func extractPagesPackage(packageBytes []byte, releaseDir string, deployment pagesDeploymentSource) error {
|
||||
tmpDir := releaseDir + ".tmp"
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
if err := os.MkdirAll(tmpDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(tmpDir, pagesDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
reader, err := zip.NewReader(bytes.NewReader(packageBytes), int64(len(packageBytes)))
|
||||
@@ -182,7 +190,7 @@ func extractPagesPackage(packageBytes []byte, releaseDir string, deployment page
|
||||
}
|
||||
if item.FileInfo().Mode()&os.ModeSymlink != 0 {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
return fmt.Errorf("Pages package contains unsupported symlink: %s", relativePath)
|
||||
return fmt.Errorf("pages package contains unsupported symlink: %s", relativePath)
|
||||
}
|
||||
if err := extractPagesFile(item, filepath.Join(tmpDir, relativePath)); err != nil {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
@@ -197,21 +205,32 @@ func extractPagesPackage(packageBytes []byte, releaseDir string, deployment page
|
||||
return os.Rename(tmpDir, releaseDir)
|
||||
}
|
||||
|
||||
func pagesZipEntryCopyLimit(size uint64) (int64, error) {
|
||||
if size == 0 || size > pagesMaxExtractedFileBytes || size > uint64(math.MaxInt64) {
|
||||
return 0, errors.New("pages file size out of bounds")
|
||||
}
|
||||
return int64(size), nil //nolint:gosec // size is bounded to math.MaxInt64 above
|
||||
}
|
||||
|
||||
func extractPagesFile(item *zip.File, targetPath string) error {
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), pagesDirPerm); 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())
|
||||
defer func() { _ = source.Close() }()
|
||||
target, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, item.FileInfo().Mode().Perm()) //nolint:gosec // targetPath is under managed PagesDir from validated zip entry
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer target.Close()
|
||||
_, err = io.Copy(target, source)
|
||||
defer func() { _ = target.Close() }()
|
||||
limit, err := pagesZipEntryCopyLimit(item.UncompressedSize64)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %w", item.Name, err)
|
||||
}
|
||||
_, err = io.CopyN(target, source, limit)
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -219,7 +238,7 @@ func switchPagesCurrentDir(baseDir string, deploymentID uint, releaseDir string)
|
||||
currentDir := pagesCurrentDir(baseDir, deploymentID)
|
||||
previousDir := currentDir + ".previous"
|
||||
_ = os.RemoveAll(previousDir)
|
||||
if err := os.MkdirAll(filepath.Dir(currentDir), 0o755); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(currentDir), pagesDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := os.Stat(currentDir); err == nil {
|
||||
@@ -249,25 +268,25 @@ func copyPagesDir(sourceDir string, targetDir string) error {
|
||||
}
|
||||
targetPath := filepath.Join(targetDir, relativePath)
|
||||
if entry.IsDir() {
|
||||
return os.MkdirAll(targetPath, 0o755)
|
||||
return os.MkdirAll(targetPath, pagesDirPerm)
|
||||
}
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
input, err := os.Open(sourcePath)
|
||||
input, err := os.Open(sourcePath) //nolint:gosec // sourcePath is under managed PagesDir walk root
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer input.Close()
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
|
||||
defer func() { _ = input.Close() }()
|
||||
if err := os.MkdirAll(filepath.Dir(targetPath), pagesDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
output, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, info.Mode().Perm())
|
||||
output, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, info.Mode().Perm()) //nolint:gosec // targetPath is under managed PagesDir walk root
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer output.Close()
|
||||
defer func() { _ = output.Close() }()
|
||||
_, err = io.Copy(output, input)
|
||||
return err
|
||||
})
|
||||
@@ -279,20 +298,20 @@ func normalizePagesArchivePath(raw string) (string, bool, error) {
|
||||
return "", true, nil
|
||||
}
|
||||
if strings.HasPrefix(name, "/") {
|
||||
return "", false, fmt.Errorf("Pages package contains absolute path: %s", raw)
|
||||
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 "", 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"))
|
||||
data, err := os.ReadFile(filepath.Join(dir, ".openflare-pages.json")) //nolint:gosec // dir is managed PagesDir
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
@@ -304,14 +323,11 @@ func markerMatches(dir string, deployment pagesDeploymentSource) bool {
|
||||
}
|
||||
|
||||
func writePagesMarker(dir string, deployment pagesDeploymentSource) error {
|
||||
data, err := json.Marshal(pagesDeploymentMarker{
|
||||
DeploymentID: deployment.DeploymentID,
|
||||
Checksum: deployment.Checksum,
|
||||
})
|
||||
data, err := json.Marshal(pagesDeploymentMarker(deployment))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(filepath.Join(dir, ".openflare-pages.json"), data, 0o644)
|
||||
return os.WriteFile(filepath.Join(dir, ".openflare-pages.json"), data, pagesManifestFilePerm)
|
||||
}
|
||||
|
||||
func pagesCurrentDir(baseDir string, deploymentID uint) string {
|
||||
|
||||
@@ -18,12 +18,14 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
|
||||
)
|
||||
|
||||
// Apply result constants indicate the outcome reported back to the server.
|
||||
const (
|
||||
ApplyResultSuccess = "success"
|
||||
ApplyResultWarning = "warning"
|
||||
ApplyResultFailed = "failed"
|
||||
)
|
||||
|
||||
// ConfigClient is the interface for communicating with the server control plane.
|
||||
type ConfigClient interface {
|
||||
GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error)
|
||||
DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error)
|
||||
@@ -31,6 +33,7 @@ type ConfigClient interface {
|
||||
SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error)
|
||||
}
|
||||
|
||||
// NginxManager is the interface for managing the local OpenResty instance.
|
||||
type NginxManager interface {
|
||||
Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) nginx.ApplyOutcome
|
||||
EnsureRuntime(ctx context.Context, recreate bool) error
|
||||
@@ -40,6 +43,7 @@ type NginxManager interface {
|
||||
SyncWAFIPGroups(groups []protocol.WAFIPGroup) error
|
||||
}
|
||||
|
||||
// Service orchestrates configuration synchronisation between the server and the local OpenResty instance.
|
||||
type Service struct {
|
||||
client ConfigClient
|
||||
nginxManager NginxManager
|
||||
@@ -47,10 +51,12 @@ type Service struct {
|
||||
pagesDir string
|
||||
}
|
||||
|
||||
// SetPagesDir sets the local directory used for pages deployment packages.
|
||||
func (s *Service) SetPagesDir(path string) {
|
||||
s.pagesDir = strings.TrimSpace(path)
|
||||
}
|
||||
|
||||
// New creates a new Service with the given client, nginx manager, and state store.
|
||||
func New(client ConfigClient, nginxManager NginxManager, stateStore *state.Store) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
@@ -59,100 +65,34 @@ func New(client ConfigClient, nginxManager NginxManager, stateStore *state.Store
|
||||
}
|
||||
}
|
||||
|
||||
// SyncOnce performs a single periodic sync against the given active config summary.
|
||||
func (s *Service) SyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error {
|
||||
return s.sync(ctx, false, target)
|
||||
}
|
||||
|
||||
// SyncOnStartup performs an initial sync at agent startup, applying config even when checksums already match.
|
||||
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()
|
||||
mode := syncMode(startup)
|
||||
snapshot, currentChecksum, err := s.loadSyncState()
|
||||
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)
|
||||
}
|
||||
normalizeSyncTarget(target)
|
||||
|
||||
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)
|
||||
return s.syncWithoutTarget(ctx, mode, startup, snapshot, currentChecksum)
|
||||
}
|
||||
|
||||
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)
|
||||
return s.syncMatchingChecksum(ctx, mode, startup, snapshot, currentChecksum, target)
|
||||
}
|
||||
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)
|
||||
return s.syncMismatchedChecksum(ctx, mode, startup, snapshot, currentChecksum, target)
|
||||
}
|
||||
|
||||
// ForceSyncOnce clears any blocked target state then unconditionally fetches and applies the active config.
|
||||
func (s *Service) ForceSyncOnce(ctx context.Context, target *protocol.ActiveConfigMeta) error {
|
||||
snapshot, err := s.stateStore.Load()
|
||||
if err != nil {
|
||||
@@ -174,6 +114,7 @@ func (s *Service) ForceSyncOnce(ctx context.Context, target *protocol.ActiveConf
|
||||
return s.applyIfNeeded(ctx, "force", true, snapshot, currentChecksum, target, config)
|
||||
}
|
||||
|
||||
// WAFIPGroupChecksums returns the current per-group checksums held by the nginx manager.
|
||||
func (s *Service) WAFIPGroupChecksums() (map[string]string, error) {
|
||||
if s.nginxManager == nil {
|
||||
return map[string]string{}, nil
|
||||
@@ -181,7 +122,8 @@ func (s *Service) WAFIPGroupChecksums() (map[string]string, error) {
|
||||
return s.nginxManager.WAFIPGroupChecksums()
|
||||
}
|
||||
|
||||
func (s *Service) ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPGroup) error {
|
||||
// ApplyWAFIPGroups writes the given WAF IP groups to the nginx manager.
|
||||
func (s *Service) ApplyWAFIPGroups(_ context.Context, groups []protocol.WAFIPGroup) error {
|
||||
if len(groups) == 0 || s.nginxManager == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -190,36 +132,13 @@ func (s *Service) ApplyWAFIPGroups(ctx context.Context, groups []protocol.WAFIPG
|
||||
|
||||
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)
|
||||
return s.handleUpToDateConfig(ctx, mode, snapshot, config)
|
||||
}
|
||||
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 handled, err := s.handleBlockedConfigAfterFetch(ctx, mode, startup, snapshot, currentChecksum, config); handled {
|
||||
return err
|
||||
}
|
||||
if hasBlockedTarget(snapshot) {
|
||||
clearBlockedTarget(snapshot)
|
||||
@@ -228,6 +147,10 @@ func (s *Service) applyIfNeeded(ctx context.Context, mode string, startup bool,
|
||||
slog.Debug("skipping apply because state already records target version/checksum", "version", config.Version, "checksum", config.Checksum)
|
||||
return s.stateStore.Save(snapshot)
|
||||
}
|
||||
return s.applyRenderedConfig(ctx, mode, snapshot, currentChecksum, config)
|
||||
}
|
||||
|
||||
func (s *Service) applyRenderedConfig(ctx context.Context, mode string, snapshot *state.Snapshot, currentChecksum string, config *protocol.ActiveConfigResponse) error {
|
||||
rendered, err := renderActiveConfig(config)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -238,49 +161,8 @@ func (s *Service) applyIfNeeded(ctx context.Context, mode string, startup bool,
|
||||
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
|
||||
}
|
||||
outcome, message := normalizeApplyOutcome(s.nginxManager.Apply(ctx, rendered.mainConfig, rendered.routeConfig, rendered.supportFiles))
|
||||
applyResult := updateSnapshotFromApplyOutcome(mode, snapshot, config, outcome, message)
|
||||
|
||||
if err := s.stateStore.Save(snapshot); err != nil {
|
||||
return err
|
||||
@@ -288,25 +170,25 @@ func (s *Service) applyIfNeeded(ctx context.Context, mode string, startup bool,
|
||||
if err := s.client.ReportApplyLog(ctx, protocol.ApplyLogPayload{
|
||||
NodeID: snapshot.NodeID,
|
||||
Version: config.Version,
|
||||
Result: reportResult,
|
||||
Message: message,
|
||||
Result: applyResult.reportResult,
|
||||
Message: applyResult.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)
|
||||
slog.Error("report apply log failed", "version", config.Version, "result", applyResult.reportResult, "error", err)
|
||||
return err
|
||||
}
|
||||
if reportResult == ApplyResultFailed {
|
||||
if applyResult.reportResult == ApplyResultFailed {
|
||||
slog.Warn("failed apply log reported", "version", config.Version)
|
||||
return outcomeError(config.Version, message)
|
||||
return outcomeError(config.Version, applyResult.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)
|
||||
slog.Debug("apply log reported", "version", config.Version, "result", applyResult.reportResult)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -476,13 +358,13 @@ func (s *Service) ensureRuntimeForCurrentConfig(ctx context.Context, mode string
|
||||
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 {
|
||||
fallbackErr := s.nginxManager.EnsureSafeFallbackRuntime(ctx, reason)
|
||||
if 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)
|
||||
}
|
||||
err = fmt.Errorf("%v; fallback recovery failed: %w", err, fallbackErr)
|
||||
}
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
|
||||
snapshot.OpenrestyMessage = err.Error()
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
package sync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/nginx"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
|
||||
)
|
||||
|
||||
func syncMode(startup bool) string {
|
||||
if startup {
|
||||
return "startup"
|
||||
}
|
||||
return "periodic"
|
||||
}
|
||||
|
||||
func normalizeSyncTarget(target *protocol.ActiveConfigMeta) {
|
||||
if target == nil {
|
||||
return
|
||||
}
|
||||
target.Version = strings.TrimSpace(target.Version)
|
||||
target.Checksum = strings.TrimSpace(target.Checksum)
|
||||
}
|
||||
|
||||
func (s *Service) loadSyncState() (*state.Snapshot, string, error) {
|
||||
snapshot, err := s.stateStore.Load()
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
currentChecksum, err := s.nginxManager.CurrentChecksum()
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
return snapshot, currentChecksum, nil
|
||||
}
|
||||
|
||||
func (s *Service) syncWithoutTarget(ctx context.Context, mode string, startup bool, snapshot *state.Snapshot, currentChecksum string) error {
|
||||
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)
|
||||
}
|
||||
|
||||
func (s *Service) syncMatchingChecksum(ctx context.Context, mode string, startup bool, snapshot *state.Snapshot, currentChecksum string, target *protocol.ActiveConfigMeta) error {
|
||||
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)
|
||||
}
|
||||
return s.finishUpToDateSync(ctx, mode, snapshot, target)
|
||||
}
|
||||
|
||||
func (s *Service) finishUpToDateSync(ctx context.Context, mode string, snapshot *state.Snapshot, target *protocol.ActiveConfigMeta) error {
|
||||
slog.Debug("local openresty config already up to date", "mode", mode, "version", target.Version)
|
||||
if shouldReportNoopApply(snapshot, target.Version, target.Checksum) {
|
||||
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)
|
||||
}
|
||||
|
||||
func (s *Service) syncMismatchedChecksum(ctx context.Context, mode string, startup bool, snapshot *state.Snapshot, currentChecksum string, target *protocol.ActiveConfigMeta) error {
|
||||
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) handleUpToDateConfig(ctx context.Context, mode string, snapshot *state.Snapshot, config *protocol.ActiveConfigResponse) error {
|
||||
slog.Debug("local openresty config already up to date", "mode", mode, "version", config.Version)
|
||||
if shouldReportNoopApply(snapshot, config.Version, config.Checksum) {
|
||||
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)
|
||||
}
|
||||
|
||||
func (s *Service) handleBlockedConfigAfterFetch(ctx context.Context, mode string, startup bool, snapshot *state.Snapshot, currentChecksum string, config *protocol.ActiveConfigResponse) (bool, error) {
|
||||
if !isBlockedTarget(snapshot, config.Version, config.Checksum) {
|
||||
return false, nil
|
||||
}
|
||||
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 true, err
|
||||
}
|
||||
return true, s.stateStore.Save(snapshot)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
type applyOutcomeResult struct {
|
||||
reportResult string
|
||||
message string
|
||||
}
|
||||
|
||||
func normalizeApplyOutcome(outcome nginx.ApplyOutcome) (nginx.ApplyOutcome, string) {
|
||||
message := strings.TrimSpace(outcome.Message)
|
||||
if outcome.Status == "" {
|
||||
outcome.Status = nginx.ApplyStatusFatal
|
||||
if message == "" {
|
||||
message = "openresty apply returned empty outcome"
|
||||
}
|
||||
}
|
||||
return outcome, message
|
||||
}
|
||||
|
||||
func updateSnapshotFromApplyOutcome(mode string, snapshot *state.Snapshot, config *protocol.ActiveConfigResponse, outcome nginx.ApplyOutcome, message string) applyOutcomeResult {
|
||||
result := applyOutcomeResult{reportResult: ApplyResultFailed, message: message}
|
||||
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 = ""
|
||||
result.reportResult = ApplyResultSuccess
|
||||
if result.message == "" {
|
||||
result.message = "apply success"
|
||||
}
|
||||
case nginx.ApplyStatusWarning:
|
||||
if result.message == "" {
|
||||
result.message = "apply rolled back to previous config"
|
||||
}
|
||||
slog.Warn("openresty config apply rolled back", "mode", mode, "version", config.Version, "message", result.message)
|
||||
markBlockedTarget(snapshot, config.Version, config.Checksum, result.message)
|
||||
snapshot.LastError = result.message
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
|
||||
snapshot.OpenrestyMessage = result.message
|
||||
result.reportResult = ApplyResultWarning
|
||||
default:
|
||||
if result.message == "" {
|
||||
result.message = "openresty apply failed"
|
||||
}
|
||||
slog.Error("apply openresty config failed", "mode", mode, "version", config.Version, "message", result.message)
|
||||
markBlockedTarget(snapshot, config.Version, config.Checksum, result.message)
|
||||
snapshot.LastError = result.message
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
|
||||
snapshot.OpenrestyMessage = result.message
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -1,17 +1,22 @@
|
||||
// Package updater provides agent self-update integration with the edge updater.
|
||||
package updater
|
||||
|
||||
import (
|
||||
edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/config"
|
||||
edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
|
||||
)
|
||||
|
||||
// Service is an alias for the edge updater service type used by the agent.
|
||||
type Service = edgeupdater.Service
|
||||
|
||||
// UpdateOptions is an alias for the edge updater options type.
|
||||
type UpdateOptions = edgeupdater.UpdateOptions
|
||||
|
||||
// New creates and returns a new agent updater Service with the agent-specific configuration.
|
||||
func New() *Service {
|
||||
return edgeupdater.New(edgeupdater.Config{
|
||||
LocalVersion: config.Version,
|
||||
AssetPrefix: "openflare-agent",
|
||||
LogLabel: "agent",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package wsclient provides the agent-side WebSocket client for connecting to the OpenFlare server.
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
@@ -8,28 +9,38 @@ import (
|
||||
edgews "github.com/Rain-kl/Wavelet/internal/apps/edge/wsclient"
|
||||
)
|
||||
|
||||
// WSMessage is an alias for the WebSocket message type.
|
||||
type WSMessage = edgews.WSMessage
|
||||
|
||||
// MessageHandler is an alias for the WebSocket message handler function type.
|
||||
type MessageHandler = edgews.MessageHandler
|
||||
|
||||
// Connection is an alias for the AgentConnection interface.
|
||||
type Connection = edgews.AgentConnection
|
||||
|
||||
// Client wraps the connection client for agent WebSockets.
|
||||
type Client struct {
|
||||
inner *edgews.Client
|
||||
}
|
||||
|
||||
// New creates a new Client instance.
|
||||
func New(baseURL, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
inner: edgews.New(edgews.PresetAgent, baseURL, token, timeout),
|
||||
}
|
||||
}
|
||||
|
||||
// SetToken updates the client's token.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.inner.SetToken(token)
|
||||
}
|
||||
|
||||
// URL returns the WebSocket client's target URL.
|
||||
func (c *Client) URL() string {
|
||||
return c.inner.URL()
|
||||
}
|
||||
|
||||
// Connect establishes a WebSocket connection to the server and returns the connection handle.
|
||||
func (c *Client) Connect(ctx context.Context) (protocol.WebSocketConnection, error) {
|
||||
return c.inner.ConnectAgent(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package config provides shared configuration types for edge applications.
|
||||
package config
|
||||
|
||||
import (
|
||||
@@ -8,8 +9,11 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// MillisecondDuration is a time.Duration that marshals to/from JSON as an integer number of milliseconds
|
||||
// or as a Go duration string (e.g. "1s", "500ms").
|
||||
type MillisecondDuration time.Duration
|
||||
|
||||
// Duration returns the underlying time.Duration value.
|
||||
func (d MillisecondDuration) Duration() time.Duration {
|
||||
return time.Duration(d)
|
||||
}
|
||||
@@ -18,6 +22,7 @@ func (d MillisecondDuration) String() string {
|
||||
return time.Duration(d).String()
|
||||
}
|
||||
|
||||
// UnmarshalJSON decodes either a numeric millisecond value or a quoted Go duration string.
|
||||
func (d *MillisecondDuration) UnmarshalJSON(data []byte) error {
|
||||
raw := strings.TrimSpace(string(data))
|
||||
if raw == "" || raw == "null" {
|
||||
@@ -49,6 +54,7 @@ func (d *MillisecondDuration) UnmarshalJSON(data []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarshalJSON encodes the duration as an integer number of milliseconds.
|
||||
func (d MillisecondDuration) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(time.Duration(d).Milliseconds())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package heartbeat handles periodic heartbeat and update checks.
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
@@ -8,6 +9,7 @@ import (
|
||||
edgeupdater "github.com/Rain-kl/Wavelet/internal/apps/edge/updater"
|
||||
)
|
||||
|
||||
// AutoUpdateSettings defines the settings for automatic edge updates.
|
||||
type AutoUpdateSettings struct {
|
||||
AutoUpdate bool
|
||||
UpdateNow bool
|
||||
@@ -16,6 +18,7 @@ type AutoUpdateSettings struct {
|
||||
UpdateTag string
|
||||
}
|
||||
|
||||
// TryAutoUpdate attempts to check and apply auto updates for the edge service.
|
||||
func TryAutoUpdate(ctx context.Context, updater *edgeupdater.Service, settings *AutoUpdateSettings, logLabel string) {
|
||||
if settings == nil || updater == nil {
|
||||
return
|
||||
@@ -38,4 +41,4 @@ func TryAutoUpdate(ctx context.Context, updater *edgeupdater.Service, settings *
|
||||
if err != nil {
|
||||
slog.Error(logLabel+" update check failed", "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package httpclient provides an authenticated HTTP client for edge services.
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
@@ -12,6 +13,7 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// Client is an HTTP client wrapper for communicating with remote HTTP services.
|
||||
type Client struct {
|
||||
baseURL string
|
||||
token string
|
||||
@@ -19,6 +21,7 @@ type Client struct {
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
// New creates a new Client instance.
|
||||
func New(baseURL, token string, timeout time.Duration, authHeader string) *Client {
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
@@ -28,11 +31,13 @@ func New(baseURL, token string, timeout time.Duration, authHeader string) *Clien
|
||||
}
|
||||
}
|
||||
|
||||
// SetToken updates the client auth token.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.token = strings.TrimSpace(token)
|
||||
slog.Debug("http client token updated")
|
||||
}
|
||||
|
||||
// GetJSON sends a GET request and decodes the response body into target.
|
||||
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 {
|
||||
@@ -42,6 +47,7 @@ func (c *Client) GetJSON(ctx context.Context, path string, target any) error {
|
||||
return c.do(req, target)
|
||||
}
|
||||
|
||||
// PostJSON sends a POST request with JSON body and decodes the response body into target.
|
||||
func (c *Client) PostJSON(ctx context.Context, path string, body any, target any) error {
|
||||
data, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
@@ -56,6 +62,7 @@ func (c *Client) PostJSON(ctx context.Context, path string, body any, target any
|
||||
return c.do(req, target)
|
||||
}
|
||||
|
||||
// DoRaw performs an HTTP request with custom headers and returns the raw response.
|
||||
func (c *Client) DoRaw(ctx context.Context, method, path string, headers map[string]string) (*http.Response, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, nil)
|
||||
if err != nil {
|
||||
@@ -80,12 +87,11 @@ func (c *Client) do(req *http.Request, target any) error {
|
||||
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 {
|
||||
defer func() {
|
||||
if err := res.Body.Close(); err != nil {
|
||||
slog.Error("failed to close response body", "error", err)
|
||||
}
|
||||
}(res.Body)
|
||||
}()
|
||||
|
||||
body, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
@@ -106,6 +112,7 @@ func (c *Client) do(req *http.Request, target any) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// APIError creates a new API error with the given message if it is not empty.
|
||||
func APIError(msg string) error {
|
||||
if strings.TrimSpace(msg) == "" {
|
||||
return nil
|
||||
@@ -113,6 +120,7 @@ func APIError(msg string) error {
|
||||
return errors.New(msg)
|
||||
}
|
||||
|
||||
// ReadBodyError parses the error message from the response body, or returns the fallback message.
|
||||
func ReadBodyError(body []byte, fallback string) error {
|
||||
var errBody struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
@@ -123,10 +131,11 @@ func ReadBodyError(body []byte, fallback string) error {
|
||||
return errors.New(fallback)
|
||||
}
|
||||
|
||||
// ReadHTTPError reads the error message from the HTTP response.
|
||||
func ReadHTTPError(res *http.Response) error {
|
||||
body, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
return errors.New(res.Status)
|
||||
}
|
||||
return ReadBodyError(body, res.Status)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package logging configures structured logging for edge applications.
|
||||
package logging
|
||||
|
||||
import (
|
||||
@@ -6,10 +7,12 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Options holds configuration options for the structured logger.
|
||||
type Options struct {
|
||||
AddSource bool
|
||||
}
|
||||
|
||||
// Setup initialises the default slog handler using the given options and the LOG_LEVEL environment variable.
|
||||
func Setup(opts Options) {
|
||||
handlerOpts := &slog.HandlerOptions{
|
||||
AddSource: opts.AddSource,
|
||||
@@ -19,6 +22,7 @@ func Setup(opts Options) {
|
||||
slog.SetDefault(slog.New(handler))
|
||||
}
|
||||
|
||||
// ParseLevel converts a log-level string (e.g. "debug", "warn") to the corresponding slog.Level.
|
||||
func ParseLevel(value string) slog.Level {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "debug":
|
||||
@@ -30,4 +34,4 @@ func ParseLevel(value string) slog.Level {
|
||||
default:
|
||||
return slog.LevelInfo
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package nodeip detects the preferred public IP address for edge nodes.
|
||||
package nodeip
|
||||
|
||||
import (
|
||||
@@ -9,20 +10,36 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
|
||||
)
|
||||
|
||||
const (
|
||||
outboundIPLookupTimeout = 5 * time.Second
|
||||
publicIPPriorityScore = 2 // matches iputil.Score for public IPv4 addresses
|
||||
)
|
||||
|
||||
// LookupOutboundIP and LookupLocalIP are the provider functions used to detect the node's outbound/local IP.
|
||||
// They are package-level variables so they can be overridden in tests.
|
||||
var (
|
||||
LookupOutboundIP = geoip.GetOutboundIP
|
||||
LookupLocalIP = DetectLocal
|
||||
)
|
||||
|
||||
// Detect returns the best available outbound or local IPv4 address for this node.
|
||||
func Detect() string {
|
||||
if ip := detectOutbound(); ip != "" {
|
||||
if ip := detectOutbound(context.Background()); ip != "" {
|
||||
return ip
|
||||
}
|
||||
return LookupLocalIP()
|
||||
}
|
||||
|
||||
func detectOutbound() string {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
// DetectWithContext returns the best available outbound or local IPv4 address, respecting ctx for cancellation.
|
||||
func DetectWithContext(ctx context.Context) string {
|
||||
if ip := detectOutbound(ctx); ip != "" {
|
||||
return ip
|
||||
}
|
||||
return LookupLocalIP()
|
||||
}
|
||||
|
||||
func detectOutbound(ctx context.Context) string {
|
||||
ctx, cancel := context.WithTimeout(ctx, outboundIPLookupTimeout)
|
||||
defer cancel()
|
||||
ip, err := LookupOutboundIP(ctx)
|
||||
if err != nil || ip == nil {
|
||||
@@ -31,6 +48,7 @@ func detectOutbound() string {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
// DetectLocal returns the highest-priority non-loopback local IPv4 address found on system interfaces.
|
||||
func DetectLocal() string {
|
||||
interfaces, err := net.Interfaces()
|
||||
if err != nil {
|
||||
@@ -60,10 +78,10 @@ func DetectLocal() string {
|
||||
bestIP = ipv4.String()
|
||||
bestPriority = priority
|
||||
}
|
||||
if bestPriority == 2 {
|
||||
if bestPriority == publicIPPriorityScore {
|
||||
return bestIP
|
||||
}
|
||||
}
|
||||
}
|
||||
return bestIP
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
// Package observability provides helpers that read Linux /proc and /sys metrics for system monitoring.
|
||||
package observability
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
@@ -10,13 +12,20 @@ import (
|
||||
"syscall"
|
||||
)
|
||||
|
||||
const (
|
||||
memInfoMinFieldCount = 2
|
||||
cpuStatMinFieldCount = 5
|
||||
netDevMinFieldCount = 16
|
||||
diskStatsMinFieldCount = 14
|
||||
)
|
||||
|
||||
// ReadLinuxOSRelease returns the OS name and version from /etc/os-release.
|
||||
func ReadLinuxOSRelease() (string, string) {
|
||||
file, err := os.Open("/etc/os-release")
|
||||
if err != nil {
|
||||
return runtime.GOOS, ""
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
values := make(map[string]string)
|
||||
scanner := bufio.NewScanner(file)
|
||||
@@ -47,7 +56,7 @@ func ReadLinuxCPUModel() string {
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
@@ -68,7 +77,7 @@ func ReadMemInfo() (int64, int64) {
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
var memTotalKB int64
|
||||
var memAvailableKB int64
|
||||
@@ -96,7 +105,7 @@ func ReadMemInfo() (int64, int64) {
|
||||
|
||||
func parseMemInfoValue(line string) int64 {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) < 2 {
|
||||
if len(fields) < memInfoMinFieldCount {
|
||||
return 0
|
||||
}
|
||||
value, err := strconv.ParseInt(fields[1], 10, 64)
|
||||
@@ -135,7 +144,7 @@ func ReadLinuxCPUStat() (uint64, uint64) {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) < 5 {
|
||||
if len(fields) < cpuStatMinFieldCount {
|
||||
return 0, 0
|
||||
}
|
||||
var total uint64
|
||||
@@ -161,7 +170,7 @@ func ReadLinuxNetworkTotals() (int64, int64) {
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
var rx int64
|
||||
var tx int64
|
||||
@@ -179,7 +188,7 @@ func ReadLinuxNetworkTotals() (int64, int64) {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(data)
|
||||
if len(fields) < 16 {
|
||||
if len(fields) < netDevMinFieldCount {
|
||||
continue
|
||||
}
|
||||
rxValue, err := strconv.ParseInt(fields[0], 10, 64)
|
||||
@@ -200,14 +209,14 @@ func ReadLinuxDiskTotals() (int64, int64) {
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
var readBytes int64
|
||||
var writeBytes int64
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
fields := strings.Fields(scanner.Text())
|
||||
if len(fields) < 14 {
|
||||
if len(fields) < diskStatsMinFieldCount {
|
||||
continue
|
||||
}
|
||||
device := fields[2]
|
||||
@@ -249,8 +258,8 @@ func StatFilesystem(path string) (int64, int64) {
|
||||
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)
|
||||
total := multiplyUint64ToInt64(stat.Blocks, uint64(stat.Bsize))
|
||||
free := multiplyUint64ToInt64(stat.Bavail, uint64(stat.Bsize))
|
||||
used := total - free
|
||||
if used < 0 {
|
||||
used = 0
|
||||
@@ -258,11 +267,21 @@ func StatFilesystem(path string) (int64, int64) {
|
||||
return total, used
|
||||
}
|
||||
|
||||
func multiplyUint64ToInt64(a uint64, b uint64) int64 {
|
||||
if a == 0 || b == 0 {
|
||||
return 0
|
||||
}
|
||||
if a > math.MaxInt64/b {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return int64(a * b) //nolint:gosec // product is bounded to math.MaxInt64 above
|
||||
}
|
||||
|
||||
// ReadFirstLine reads and returns the trimmed first line of a file.
|
||||
func ReadFirstLine(path string) string {
|
||||
content, err := os.ReadFile(path)
|
||||
content, err := os.ReadFile(path) //nolint:gosec // path is a fixed /proc or /sys path from internal callers, not user input
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(string(content))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package runner provides shared WebSocket reconnect helpers for edge daemons.
|
||||
package runner
|
||||
|
||||
import (
|
||||
@@ -6,10 +7,12 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// WSConnection defines the minimum interface required for a WebSocket connection that can be closed.
|
||||
type WSConnection interface {
|
||||
Close() error
|
||||
}
|
||||
|
||||
// WSReconnectConfig specifies configuration parameters for the WebSocket reconnect loop.
|
||||
type WSReconnectConfig struct {
|
||||
ComponentName string
|
||||
ConnectBackoff time.Duration
|
||||
@@ -17,6 +20,7 @@ type WSReconnectConfig struct {
|
||||
OnShutdown func()
|
||||
}
|
||||
|
||||
// RunWSReconnectLoop runs a loop that attempts to keep a WebSocket connection active, automatically reconnecting when closed or failed.
|
||||
func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig,
|
||||
connect func(context.Context) (WSConnection, error),
|
||||
handle func(context.Context, WSConnection),
|
||||
@@ -40,6 +44,7 @@ func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig,
|
||||
}
|
||||
return ctx.Err()
|
||||
default:
|
||||
// Continue reconnect loop
|
||||
}
|
||||
|
||||
conn, err := connect(ctx)
|
||||
@@ -56,9 +61,10 @@ func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig,
|
||||
}
|
||||
}
|
||||
|
||||
// SleepContext pauses execution for the given duration or until the context is canceled.
|
||||
func SleepContext(ctx context.Context, d time.Duration) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(d):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
//go:build !windows
|
||||
|
||||
// Package updater provides capabilities to check for, download, and apply updates.
|
||||
package updater
|
||||
|
||||
import (
|
||||
@@ -33,7 +34,7 @@ func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
if err := removeBackupBinary(backupPath); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil {
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil { //nolint:gosec // execPath is the validated edge updater binary path
|
||||
return fmt.Errorf("exec restart: %w", err)
|
||||
}
|
||||
return fmt.Errorf("unreachable after exec")
|
||||
@@ -48,4 +49,4 @@ func removeBackupBinary(path string) error {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package updater provides capabilities to check for, download, and apply updates.
|
||||
package updater
|
||||
|
||||
import (
|
||||
@@ -17,16 +18,23 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/pkg/utils"
|
||||
)
|
||||
|
||||
const maxChecksumAssetSize = 64 * 1024
|
||||
const (
|
||||
maxChecksumAssetSize = 64 * 1024
|
||||
goosWindows = "windows"
|
||||
updateTmpFilePerm = 0o600
|
||||
updateBinaryFilePerm = 0o755
|
||||
)
|
||||
|
||||
var replaceAndRestartFunc = replaceAndRestart
|
||||
|
||||
// Config defines the configuration for the update service.
|
||||
type Config struct {
|
||||
LocalVersion string
|
||||
AssetPrefix string
|
||||
LogLabel string
|
||||
}
|
||||
|
||||
// Service handles checking and applying application binary updates.
|
||||
type Service struct {
|
||||
httpClient *http.Client
|
||||
lastCheckKey string
|
||||
@@ -35,6 +43,7 @@ type Service struct {
|
||||
logLabel string
|
||||
}
|
||||
|
||||
// New creates a new updater Service with the provided configuration.
|
||||
func New(cfg Config) *Service {
|
||||
return &Service{
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
@@ -44,6 +53,7 @@ func New(cfg Config) *Service {
|
||||
}
|
||||
}
|
||||
|
||||
// UpdateOptions specifies parameters for checking and applying updates.
|
||||
type UpdateOptions struct {
|
||||
Channel string
|
||||
TagName string
|
||||
@@ -62,6 +72,7 @@ type githubAsset struct {
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
}
|
||||
|
||||
// CheckAndUpdate checks for a newer release on GitHub and performs an update if available.
|
||||
func (s *Service) CheckAndUpdate(ctx context.Context, repo string, options UpdateOptions) error {
|
||||
release, err := s.getRelease(ctx, repo, options)
|
||||
if err != nil {
|
||||
@@ -152,7 +163,7 @@ func (s *Service) getLatestPreviewRelease(ctx context.Context, repo string) (*gi
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("github api returned %s", resp.Status)
|
||||
@@ -188,12 +199,11 @@ func (s *Service) fetchReleaseFromURL(ctx context.Context, url string) (*githubR
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
defer func() {
|
||||
if err := resp.Body.Close(); err != nil {
|
||||
slog.Error("failed to close response body", "error", err)
|
||||
}
|
||||
}(resp.Body)
|
||||
}()
|
||||
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return nil, nil
|
||||
@@ -222,7 +232,7 @@ func (s *Service) downloadChecksum(ctx context.Context, url string, assetName st
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("checksum download returned %s", resp.Status)
|
||||
@@ -309,37 +319,37 @@ func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedCh
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = 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") {
|
||||
if runtime.GOOS == goosWindows && !strings.HasSuffix(strings.ToLower(tmpPath), ".exe") {
|
||||
tmpPath += ".exe"
|
||||
}
|
||||
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600)
|
||||
tmpFile, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, updateTmpFilePerm) //nolint:gosec // tmpPath is derived from the configured updater binary location
|
||||
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)
|
||||
_ = tmpFile.Close()
|
||||
_ = os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
if err = tmpFile.Close(); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
_ = os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
actualChecksum := hex.EncodeToString(hasher.Sum(nil))
|
||||
if actualChecksum != expectedChecksum {
|
||||
os.Remove(tmpPath)
|
||||
_ = 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)
|
||||
if err = os.Chmod(tmpPath, updateBinaryFilePerm); err != nil && runtime.GOOS != goosWindows { //nolint:gosec // downloaded edge binary must remain executable
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("set executable permission: %w", err)
|
||||
}
|
||||
|
||||
@@ -349,7 +359,7 @@ func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedCh
|
||||
|
||||
func (s *Service) assetNameForGOOSGOARCH(goos string, goarch string) string {
|
||||
name := fmt.Sprintf("%s-%s-%s", s.assetPrefix, goos, goarch)
|
||||
if goos == "windows" {
|
||||
if goos == goosWindows {
|
||||
return name + ".exe"
|
||||
}
|
||||
return name
|
||||
@@ -378,4 +388,4 @@ func buildReleaseCheckKey(options UpdateOptions, remoteVersion string) string {
|
||||
|
||||
func compareVersions(local string, remote string) int {
|
||||
return utils.CompareVersions(local, remote)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package wsclient provides WebSocket client abstractions for edge node communication.
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
@@ -8,14 +9,21 @@ import (
|
||||
shared "github.com/Rain-kl/Wavelet/pkg/wsclient"
|
||||
)
|
||||
|
||||
// WSMessage is an alias for the shared WebSocket message type.
|
||||
type WSMessage = shared.WSMessage
|
||||
|
||||
// MessageHandler is an alias for the shared WebSocket message handler type.
|
||||
type MessageHandler = shared.MessageHandler
|
||||
|
||||
// Preset represents a predefined connection configuration for a specific edge role.
|
||||
type Preset int
|
||||
|
||||
const (
|
||||
// PresetAgent is the configuration preset for agent connections.
|
||||
PresetAgent Preset = iota
|
||||
// PresetRelay is the configuration preset for relay connections.
|
||||
PresetRelay
|
||||
// PresetFlared is the configuration preset for flared (tunnel) connections.
|
||||
PresetFlared
|
||||
)
|
||||
|
||||
@@ -30,18 +38,22 @@ var presets = map[Preset]presetConfig{
|
||||
PresetFlared: {HeaderKey: "X-Tunnel-Token", WSPath: "/api/v1/tunnel/ws"},
|
||||
}
|
||||
|
||||
// PresetHeaderKey returns the HTTP header key used for authentication with the given preset.
|
||||
func PresetHeaderKey(preset Preset) string {
|
||||
return presets[preset].HeaderKey
|
||||
}
|
||||
|
||||
// PresetWSPath returns the WebSocket path used for the given preset.
|
||||
func PresetWSPath(preset Preset) string {
|
||||
return presets[preset].WSPath
|
||||
}
|
||||
|
||||
// Client is a WebSocket client configured for a specific edge preset.
|
||||
type Client struct {
|
||||
sharedClient *shared.Client
|
||||
}
|
||||
|
||||
// New creates a new Client for the given preset, base URL, token, and timeout.
|
||||
func New(preset Preset, baseURL, token string, timeout time.Duration) *Client {
|
||||
cfg := presets[preset]
|
||||
return &Client{
|
||||
@@ -55,18 +67,22 @@ func New(preset Preset, baseURL, token string, timeout time.Duration) *Client {
|
||||
}
|
||||
}
|
||||
|
||||
// SetToken updates the authentication token used by the client.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.sharedClient.SetToken(token)
|
||||
}
|
||||
|
||||
// URL returns the fully resolved WebSocket URL for this client.
|
||||
func (c *Client) URL() string {
|
||||
return c.sharedClient.URL()
|
||||
}
|
||||
|
||||
// Connection represents an established WebSocket connection to an edge node.
|
||||
type Connection struct {
|
||||
sharedConn *shared.Connection
|
||||
}
|
||||
|
||||
// Connect establishes a WebSocket connection using the client configuration.
|
||||
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||
conn, err := c.sharedClient.Connect(ctx)
|
||||
if err != nil {
|
||||
@@ -75,10 +91,12 @@ func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||
return &Connection{sharedConn: conn}, nil
|
||||
}
|
||||
|
||||
// AgentConnection is a Connection specialized for agent node communication.
|
||||
type AgentConnection struct {
|
||||
Connection
|
||||
}
|
||||
|
||||
// ConnectAgent establishes a WebSocket connection and returns it as an AgentConnection.
|
||||
func (c *Client) ConnectAgent(ctx context.Context) (*AgentConnection, error) {
|
||||
conn, err := c.Connect(ctx)
|
||||
if err != nil {
|
||||
@@ -87,6 +105,7 @@ func (c *Client) ConnectAgent(ctx context.Context) (*AgentConnection, error) {
|
||||
return &AgentConnection{Connection: *conn}, nil
|
||||
}
|
||||
|
||||
// URL returns the resolved WebSocket URL of this connection.
|
||||
func (conn *Connection) URL() string {
|
||||
if conn == nil || conn.sharedConn == nil {
|
||||
return ""
|
||||
@@ -94,18 +113,22 @@ func (conn *Connection) URL() string {
|
||||
return conn.sharedConn.URL
|
||||
}
|
||||
|
||||
// SendPing sends a ping message over the connection.
|
||||
func (conn *Connection) SendPing() error {
|
||||
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypePing, nil)
|
||||
}
|
||||
|
||||
// SendPong sends a pong message over the connection.
|
||||
func (conn *Connection) SendPong() error {
|
||||
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypePong, nil)
|
||||
}
|
||||
|
||||
// SendMessage sends a typed message with an optional payload over the connection.
|
||||
func (conn *Connection) SendMessage(msgType string, payload any) error {
|
||||
return conn.sharedConn.SendMessage(msgType, payload)
|
||||
}
|
||||
|
||||
// Receive reads the next message from the connection.
|
||||
func (conn *Connection) Receive() (pkgprotocol.WSMessage, error) {
|
||||
var message pkgprotocol.WSMessage
|
||||
if err := conn.sharedConn.Receive(&message); err != nil {
|
||||
@@ -114,10 +137,12 @@ func (conn *Connection) Receive() (pkgprotocol.WSMessage, error) {
|
||||
return message, nil
|
||||
}
|
||||
|
||||
// RunReceiveLoop blocks and dispatches incoming messages to the handler until the context is canceled.
|
||||
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
|
||||
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||
}
|
||||
|
||||
// Close gracefully closes the WebSocket connection.
|
||||
func (conn *Connection) Close() error {
|
||||
if conn == nil || conn.sharedConn == nil {
|
||||
return nil
|
||||
@@ -125,6 +150,7 @@ func (conn *Connection) Close() error {
|
||||
return conn.sharedConn.Close()
|
||||
}
|
||||
|
||||
// SendStatus sends a node status payload over the agent connection.
|
||||
func (conn *AgentConnection) SendStatus(payload pkgprotocol.NodePayload) error {
|
||||
return conn.sharedConn.SendMessage(pkgprotocol.WSMessageTypeStatus, payload)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package config loads and persists flared daemon configuration.
|
||||
package config
|
||||
|
||||
import (
|
||||
@@ -11,8 +12,17 @@ import (
|
||||
edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultHeartbeatInterval = 10 * time.Second
|
||||
defaultSyncInterval = 30 * time.Second
|
||||
defaultRequestTimeout = 10 * time.Second
|
||||
configFilePerm = 0o644
|
||||
)
|
||||
|
||||
// MillisecondDuration is a JSON-friendly duration type shared with edge config.
|
||||
type MillisecondDuration = edgeconfig.MillisecondDuration
|
||||
|
||||
// Config holds flared daemon settings loaded from file and environment.
|
||||
type Config struct {
|
||||
ServerURL string `json:"server_url"`
|
||||
TunnelToken string `json:"tunnel_token"`
|
||||
@@ -25,8 +35,9 @@ type Config struct {
|
||||
configPath string
|
||||
}
|
||||
|
||||
// Load reads configuration from path, applying environment overrides and defaults.
|
||||
func Load(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
data, err := os.ReadFile(path) //nolint:gosec // path is the flared config file location from startup configuration
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return nil, err
|
||||
}
|
||||
@@ -89,13 +100,13 @@ func applyDefaults(cfg *Config, baseDir string) {
|
||||
cfg.StatePath = filepath.Join(cfg.DataDir, "flared-state.json")
|
||||
}
|
||||
if cfg.HeartbeatInterval <= 0 {
|
||||
cfg.HeartbeatInterval = MillisecondDuration(10 * time.Second)
|
||||
cfg.HeartbeatInterval = MillisecondDuration(defaultHeartbeatInterval)
|
||||
}
|
||||
if cfg.SyncInterval <= 0 {
|
||||
cfg.SyncInterval = MillisecondDuration(30 * time.Second)
|
||||
cfg.SyncInterval = MillisecondDuration(defaultSyncInterval)
|
||||
}
|
||||
if cfg.RequestTimeout <= 0 {
|
||||
cfg.RequestTimeout = MillisecondDuration(10 * time.Second)
|
||||
cfg.RequestTimeout = MillisecondDuration(defaultRequestTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -109,6 +120,7 @@ func validate(cfg *Config) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// InitialAuthToken returns the tunnel token used for initial authentication.
|
||||
func (cfg *Config) InitialAuthToken() string {
|
||||
if cfg == nil {
|
||||
return ""
|
||||
@@ -116,6 +128,7 @@ func (cfg *Config) InitialAuthToken() string {
|
||||
return strings.TrimSpace(cfg.TunnelToken)
|
||||
}
|
||||
|
||||
// Save writes the current configuration back to the loaded config path.
|
||||
func (cfg *Config) Save() error {
|
||||
if cfg == nil {
|
||||
return errors.New("config 不能为空")
|
||||
@@ -127,5 +140,5 @@ func (cfg *Config) Save() error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(cfg.configPath, data, 0o644)
|
||||
return os.WriteFile(cfg.configPath, data, configFilePerm)
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
package config
|
||||
|
||||
// Version is the flared daemon build version string.
|
||||
var Version = "dev"
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package flared implements the tunnel client daemon runtime loop.
|
||||
package flared
|
||||
|
||||
import (
|
||||
@@ -13,15 +14,19 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/wsclient"
|
||||
)
|
||||
|
||||
// Runner is the top-level orchestrator for the flared agent. It wires together
|
||||
// heartbeat, sync, frpc management, and the WebSocket control-plane connection.
|
||||
type Runner struct {
|
||||
Config *config.Config
|
||||
HeartbeatService *heartbeat.Service
|
||||
FrpcManager *frpc.Manager
|
||||
SyncService *sync.Service
|
||||
WebSocketService *wsclient.Client
|
||||
HttpClient *httpclient.Client
|
||||
HTTPClient *httpclient.Client
|
||||
}
|
||||
|
||||
// Run starts all background services and enters the WebSocket reconnect loop.
|
||||
// It blocks until ctx is cancelled or an unrecoverable error occurs.
|
||||
func (r *Runner) Run(ctx context.Context) error {
|
||||
go r.HeartbeatService.Run(ctx)
|
||||
go r.SyncService.Run(ctx)
|
||||
@@ -40,11 +45,11 @@ type flaredWSHandler struct {
|
||||
runner *Runner
|
||||
}
|
||||
|
||||
func (h *flaredWSHandler) OnConnect(ctx context.Context) error {
|
||||
func (h *flaredWSHandler) OnConnect(_ context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *flaredWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||
func (h *flaredWSHandler) HandleMessage(_ context.Context, msg wsclient.WSMessage) error {
|
||||
switch msg.Type {
|
||||
case "active_config":
|
||||
slog.Info("received config update notification from server")
|
||||
@@ -66,4 +71,4 @@ func (r *Runner) handleConnection(ctx context.Context, conn edgerunner.WSConnect
|
||||
return
|
||||
}
|
||||
_ = wsConn.RunReceiveLoop(ctx, &flaredWSHandler{runner: r})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package frpc manages frpc child processes for tunnel relay connections.
|
||||
package frpc
|
||||
|
||||
import (
|
||||
@@ -18,6 +19,13 @@ import (
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
const (
|
||||
dataDirPerm = 0o750
|
||||
frpcConfigFilePerm = 0o644
|
||||
orphanProcessKillDelay = 500 * time.Millisecond
|
||||
)
|
||||
|
||||
// Manager supervises frpc processes for each active relay node.
|
||||
type Manager struct {
|
||||
cfg *config.Config
|
||||
processes map[string]*Process
|
||||
@@ -27,6 +35,7 @@ type Manager struct {
|
||||
currentChecksum string
|
||||
}
|
||||
|
||||
// Process tracks a single frpc child process and its runtime state.
|
||||
type Process struct {
|
||||
RelayID string
|
||||
Cmd *exec.Cmd
|
||||
@@ -36,6 +45,7 @@ type Process struct {
|
||||
LastError string
|
||||
}
|
||||
|
||||
// NewManager creates a Manager using the given flared configuration.
|
||||
func NewManager(cfg *config.Config) *Manager {
|
||||
return &Manager{
|
||||
cfg: cfg,
|
||||
@@ -43,8 +53,9 @@ func NewManager(cfg *config.Config) *Manager {
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) GetVersion() string {
|
||||
cmd := exec.Command(m.cfg.FrpcPath, "-v")
|
||||
// GetVersion returns the installed frpc binary version string.
|
||||
func (m *Manager) GetVersion(ctx context.Context) string {
|
||||
cmd := exec.CommandContext(ctx, m.cfg.FrpcPath, "-v") //nolint:gosec // FrpcPath is the configured trusted frpc binary location
|
||||
out, err := cmd.Output()
|
||||
if err != nil {
|
||||
return "unknown"
|
||||
@@ -52,6 +63,7 @@ func (m *Manager) GetVersion() string {
|
||||
return strings.TrimSpace(string(out))
|
||||
}
|
||||
|
||||
// GetConnectedRelays reports the relay nodes with active or managed frpc processes.
|
||||
func (m *Manager) GetConnectedRelays() []service.FlaredConnectedRelay {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -66,18 +78,21 @@ func (m *Manager) GetConnectedRelays() []service.FlaredConnectedRelay {
|
||||
return result
|
||||
}
|
||||
|
||||
// GetCurrentConfigVersion returns the version of the applied tunnel configuration.
|
||||
func (m *Manager) GetCurrentConfigVersion() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.currentVersion
|
||||
}
|
||||
|
||||
// GetCurrentConfigChecksum returns the checksum of the applied tunnel configuration.
|
||||
func (m *Manager) GetCurrentConfigChecksum() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.currentChecksum
|
||||
}
|
||||
|
||||
// UpdateConfig reconciles running frpc processes with the latest tunnel configuration.
|
||||
func (m *Manager) UpdateConfig(ctx context.Context, newConfig *service.FlaredTunnelConfigResponse) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
@@ -93,7 +108,7 @@ func (m *Manager) UpdateConfig(ctx context.Context, newConfig *service.FlaredTun
|
||||
slog.Debug("tunnel config version unchanged, ensuring processes are running", "version", newConfig.Version)
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(m.cfg.DataDir, 0o755); err != nil {
|
||||
if err := os.MkdirAll(m.cfg.DataDir, dataDirPerm); err != nil {
|
||||
return fmt.Errorf("create data dir failed: %w", err)
|
||||
}
|
||||
|
||||
@@ -105,14 +120,14 @@ func (m *Manager) UpdateConfig(ctx context.Context, newConfig *service.FlaredTun
|
||||
configPath := filepath.Join(m.cfg.DataDir, fmt.Sprintf("frpc_%s.toml", relay.RelayNodeID))
|
||||
|
||||
needsRestart := false
|
||||
existingData, err := os.ReadFile(configPath)
|
||||
existingData, err := os.ReadFile(configPath) //nolint:gosec // configPath is under managed DataDir
|
||||
if err != nil || string(existingData) != tomlContent {
|
||||
// 配置文件不存在或内容有变化,需要写入并重启
|
||||
needsRestart = true
|
||||
}
|
||||
|
||||
if needsRestart {
|
||||
if err := os.WriteFile(configPath, []byte(tomlContent), 0o644); err != nil {
|
||||
if err := os.WriteFile(configPath, []byte(tomlContent), frpcConfigFilePerm); err != nil {
|
||||
slog.Error("failed to write frpc config", "relay_id", relay.RelayNodeID, "error", err)
|
||||
continue
|
||||
}
|
||||
@@ -150,10 +165,7 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
_ = os.Remove(pidPath)
|
||||
}
|
||||
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
procCtx, cancel := context.WithCancel(ctx)
|
||||
procCtx, cancel := context.WithCancel(context.WithoutCancel(ctx))
|
||||
proc := &Process{
|
||||
RelayID: relayID,
|
||||
Cancel: cancel,
|
||||
@@ -176,7 +188,7 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
|
||||
ensureNoOrphanProcess(pidPath)
|
||||
|
||||
cmd := exec.CommandContext(procCtx, m.cfg.FrpcPath, "-c", configPath)
|
||||
cmd := exec.CommandContext(procCtx, m.cfg.FrpcPath, "-c", configPath) //nolint:gosec // FrpcPath and configPath are managed trusted locations
|
||||
|
||||
m.mu.Lock()
|
||||
proc.Cmd = cmd
|
||||
@@ -186,7 +198,7 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
startedAt := time.Now()
|
||||
err := cmd.Start()
|
||||
if err == nil {
|
||||
_ = os.WriteFile(pidPath, []byte(fmt.Sprintf("%d", cmd.Process.Pid)), 0o644)
|
||||
_ = os.WriteFile(pidPath, []byte(fmt.Sprintf("%d", cmd.Process.Pid)), frpcConfigFilePerm)
|
||||
err = cmd.Wait()
|
||||
}
|
||||
_ = os.Remove(pidPath)
|
||||
@@ -217,7 +229,7 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
case <-procCtx.Done():
|
||||
return
|
||||
case <-time.After(backoff):
|
||||
backoff = backoff * 2
|
||||
backoff *= 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
@@ -226,6 +238,7 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
}()
|
||||
}
|
||||
|
||||
// Stop cancels and removes all managed frpc processes.
|
||||
func (m *Manager) Stop() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
@@ -245,28 +258,23 @@ func buildFrpcToml(relay service.FlaredRelayInfo, proxies []service.FlaredProxyE
|
||||
|
||||
host, port := parseAddr(relay.Address)
|
||||
|
||||
buf.WriteString(fmt.Sprintf(`serverAddr = "%s"
|
||||
serverPort = %s
|
||||
`, host, port))
|
||||
fmt.Fprintf(&buf, "serverAddr = \"%s\"\nserverPort = %s\n", host, port)
|
||||
|
||||
if relay.AuthToken != "" {
|
||||
buf.WriteString(fmt.Sprintf(`auth.method = "token"
|
||||
auth.token = "%s"
|
||||
`, relay.AuthToken))
|
||||
fmt.Fprintf(&buf, "auth.method = \"token\"\nauth.token = \"%s\"\n", relay.AuthToken)
|
||||
}
|
||||
|
||||
if relay.ProxyURL != "" {
|
||||
buf.WriteString(fmt.Sprintf(`transport.proxyURL = "%s"
|
||||
`, relay.ProxyURL))
|
||||
fmt.Fprintf(&buf, "transport.proxyURL = \"%s\"\n", relay.ProxyURL)
|
||||
}
|
||||
|
||||
buf.WriteString("\n")
|
||||
|
||||
for _, proxy := range proxies {
|
||||
buf.WriteString(fmt.Sprintf("[[proxies]]\nname = \"%s\"\ntype = \"%s\"\nlocalIP = \"%s\"\nlocalPort = %d\n",
|
||||
proxy.Name, proxy.Type, proxy.LocalAddr, proxy.LocalPort))
|
||||
fmt.Fprintf(&buf, "[[proxies]]\nname = \"%s\"\ntype = \"%s\"\nlocalIP = \"%s\"\nlocalPort = %d\n",
|
||||
proxy.Name, proxy.Type, proxy.LocalAddr, proxy.LocalPort)
|
||||
if len(proxy.CustomDomains) > 0 {
|
||||
buf.WriteString(fmt.Sprintf("customDomains = [\"%s\"]\n", strings.Join(proxy.CustomDomains, "\", \"")))
|
||||
fmt.Fprintf(&buf, "customDomains = [\"%s\"]\n", strings.Join(proxy.CustomDomains, "\", \""))
|
||||
}
|
||||
buf.WriteString("\n")
|
||||
}
|
||||
@@ -290,7 +298,7 @@ func parseAddr(addr string) (string, string) {
|
||||
return addr, "7000"
|
||||
}
|
||||
|
||||
// State persistence
|
||||
// ManagerState persists the last applied tunnel configuration version and checksum.
|
||||
type ManagerState struct {
|
||||
Version string
|
||||
Checksum string
|
||||
@@ -305,9 +313,10 @@ func (m *Manager) saveState() error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(m.cfg.StatePath, data, 0o644)
|
||||
return os.WriteFile(m.cfg.StatePath, data, frpcConfigFilePerm)
|
||||
}
|
||||
|
||||
// LoadState restores the last applied configuration version and checksum from disk.
|
||||
func (m *Manager) LoadState() error {
|
||||
data, err := os.ReadFile(m.cfg.StatePath)
|
||||
if err != nil {
|
||||
@@ -328,7 +337,7 @@ func (m *Manager) LoadState() error {
|
||||
}
|
||||
|
||||
func ensureNoOrphanProcess(pidPath string) {
|
||||
data, err := os.ReadFile(pidPath)
|
||||
data, err := os.ReadFile(pidPath) //nolint:gosec // pidPath is under managed DataDir
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
@@ -344,7 +353,7 @@ func ensureNoOrphanProcess(pidPath string) {
|
||||
slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath)
|
||||
_ = process.Kill()
|
||||
// Wait a little bit to ensure the OS has reclaimed ports
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
time.Sleep(orphanProcessKillDelay)
|
||||
}
|
||||
_ = os.Remove(pidPath)
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package heartbeat runs the periodic flared heartbeat loop against the control plane.
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
// Service sends periodic heartbeat payloads and applies tunnel settings from responses.
|
||||
type Service struct {
|
||||
client *httpclient.Client
|
||||
frpcManager *frpc.Manager
|
||||
@@ -20,6 +22,7 @@ type Service struct {
|
||||
updater *updater.Service
|
||||
}
|
||||
|
||||
// New creates a heartbeat service with the given client, frpc manager, and config.
|
||||
func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
@@ -29,6 +32,7 @@ func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) *
|
||||
}
|
||||
}
|
||||
|
||||
// Run starts the heartbeat loop until ctx is canceled.
|
||||
func (s *Service) Run(ctx context.Context) {
|
||||
edgeheartbeat.RunLoop(ctx, s.config.HeartbeatInterval.Duration(), s.doHeartbeat)
|
||||
}
|
||||
@@ -38,8 +42,8 @@ func (s *Service) doHeartbeat(ctx context.Context) {
|
||||
|
||||
payload := service.FlaredHeartbeatPayload{
|
||||
ClientVersion: config.Version,
|
||||
FrpVersion: s.frpcManager.GetVersion(),
|
||||
IP: nodeip.Detect(),
|
||||
FrpVersion: s.frpcManager.GetVersion(ctx),
|
||||
IP: nodeip.DetectWithContext(ctx),
|
||||
TunnelStatus: "running",
|
||||
ConnectedRelays: s.frpcManager.GetConnectedRelays(),
|
||||
CurrentVersion: s.frpcManager.GetCurrentConfigVersion(),
|
||||
@@ -69,4 +73,4 @@ func tunnelSettingsToAutoUpdate(settings *service.RelaySettings) *edgeheartbeat.
|
||||
UpdateChannel: settings.UpdateChannel,
|
||||
UpdateTag: settings.UpdateTag,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package httpclient provides the HTTP client used by the flared agent to communicate with the Wavelet server.
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
@@ -8,21 +9,25 @@ import (
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
// APIResponse is the standard JSON envelope returned by the Wavelet API.
|
||||
type APIResponse[T any] struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
Data T `json:"data"`
|
||||
}
|
||||
|
||||
// Client is the HTTP client for the flared tunnel API.
|
||||
type Client struct {
|
||||
base *edgehttp.Client
|
||||
}
|
||||
|
||||
// New creates a new Client configured with the given base URL, authentication token, and request timeout.
|
||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
base: edgehttp.New(baseURL, token, timeout, "X-Tunnel-Token"),
|
||||
}
|
||||
}
|
||||
|
||||
// Heartbeat sends a tunnel heartbeat payload and returns the server response.
|
||||
func (c *Client) Heartbeat(ctx context.Context, payload service.FlaredHeartbeatPayload) (*service.FlaredHeartbeatResponse, error) {
|
||||
resp := APIResponse[service.FlaredHeartbeatResponse]{}
|
||||
if err := c.base.PostJSON(ctx, "/api/v1/tunnel/heartbeat", payload, &resp); err != nil {
|
||||
@@ -34,6 +39,7 @@ func (c *Client) Heartbeat(ctx context.Context, payload service.FlaredHeartbeatP
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
// GetActiveConfig fetches the currently active tunnel configuration from the server.
|
||||
func (c *Client) GetActiveConfig(ctx context.Context) (*service.FlaredTunnelConfigResponse, error) {
|
||||
resp := APIResponse[service.FlaredTunnelConfigResponse]{}
|
||||
if err := c.base.GetJSON(ctx, "/api/v1/tunnel/config/active", &resp); err != nil {
|
||||
@@ -45,6 +51,7 @@ func (c *Client) GetActiveConfig(ctx context.Context) (*service.FlaredTunnelConf
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
// ReportApplyLog submits a configuration apply-log entry to the server.
|
||||
func (c *Client) ReportApplyLog(ctx context.Context, payload service.ApplyLogPayload) error {
|
||||
resp := APIResponse[any]{}
|
||||
if err := c.base.PostJSON(ctx, "/api/v1/tunnel/apply-log", payload, &resp); err != nil {
|
||||
@@ -53,6 +60,7 @@ func (c *Client) ReportApplyLog(ctx context.Context, payload service.ApplyLogPay
|
||||
return edgehttp.APIError(resp.ErrorMsg)
|
||||
}
|
||||
|
||||
// SetToken updates the authentication token used by the client.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.base.SetToken(token)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package sync periodically fetches and applies the active tunnel configuration.
|
||||
package sync
|
||||
|
||||
import (
|
||||
@@ -11,6 +12,7 @@ import (
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
)
|
||||
|
||||
// Service synchronizes tunnel configuration from the control plane to the local frpc manager.
|
||||
type Service struct {
|
||||
client *httpclient.Client
|
||||
frpcManager *frpc.Manager
|
||||
@@ -18,6 +20,7 @@ type Service struct {
|
||||
triggerCh chan struct{}
|
||||
}
|
||||
|
||||
// New creates a sync service with the given client, frpc manager, and config.
|
||||
func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
@@ -27,6 +30,7 @@ func New(client *httpclient.Client, manager *frpc.Manager, cfg *config.Config) *
|
||||
}
|
||||
}
|
||||
|
||||
// Trigger requests an immediate configuration sync without waiting for the next interval.
|
||||
func (s *Service) Trigger() {
|
||||
select {
|
||||
case s.triggerCh <- struct{}{}:
|
||||
@@ -34,6 +38,7 @@ func (s *Service) Trigger() {
|
||||
}
|
||||
}
|
||||
|
||||
// Run starts the sync loop until ctx is canceled.
|
||||
func (s *Service) Run(ctx context.Context) {
|
||||
ticker := time.NewTicker(s.config.SyncInterval.Duration())
|
||||
defer ticker.Stop()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package updater provides update service capabilities for flared.
|
||||
package updater
|
||||
|
||||
import (
|
||||
@@ -5,13 +6,17 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/config"
|
||||
)
|
||||
|
||||
// Service is an alias for the edge updater Service.
|
||||
type Service = edgeupdater.Service
|
||||
|
||||
// UpdateOptions is an alias for the edge updater UpdateOptions.
|
||||
type UpdateOptions = edgeupdater.UpdateOptions
|
||||
|
||||
// New creates a new updater Service instance.
|
||||
func New() *Service {
|
||||
return edgeupdater.New(edgeupdater.Config{
|
||||
LocalVersion: config.Version,
|
||||
AssetPrefix: "openflared",
|
||||
LogLabel: "flared",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package wsclient provides a WebSocket client for flared control-plane communication.
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
@@ -7,24 +8,33 @@ import (
|
||||
edgews "github.com/Rain-kl/Wavelet/internal/apps/edge/wsclient"
|
||||
)
|
||||
|
||||
// WSMessage is a WebSocket message exchanged with the control plane.
|
||||
type WSMessage = edgews.WSMessage
|
||||
|
||||
// MessageHandler processes incoming WebSocket messages.
|
||||
type MessageHandler = edgews.MessageHandler
|
||||
|
||||
// Connection represents an active WebSocket connection.
|
||||
type Connection = edgews.Connection
|
||||
|
||||
// Client connects to the flared WebSocket endpoint on the control plane.
|
||||
type Client struct {
|
||||
inner *edgews.Client
|
||||
}
|
||||
|
||||
// New creates a WebSocket client for the flared control-plane endpoint.
|
||||
func New(baseURL, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
inner: edgews.New(edgews.PresetFlared, baseURL, token, timeout),
|
||||
}
|
||||
}
|
||||
|
||||
// SetToken updates the authentication token used for the WebSocket connection.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.inner.SetToken(token)
|
||||
}
|
||||
|
||||
// Connect establishes a WebSocket connection to the control plane.
|
||||
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||
return c.inner.Connect(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package agent implements the OpenFlare agent protocol: node registration,
|
||||
// heartbeat processing, access-log ingestion, and related middleware.
|
||||
package agent
|
||||
|
||||
import (
|
||||
@@ -11,12 +13,12 @@ import (
|
||||
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
|
||||
)
|
||||
|
||||
var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) {
|
||||
var accessLogGeoProviderFactory = func() (pkggeoip.Service, error) {
|
||||
return pkggeoip.NewMaxMindGeoIPService()
|
||||
}
|
||||
|
||||
type accessLogRegionResolver struct {
|
||||
provider pkggeoip.GeoIPService
|
||||
provider pkggeoip.Service
|
||||
cache map[string]string
|
||||
}
|
||||
|
||||
|
||||
@@ -36,12 +36,10 @@ var tokenCache = newAccessTokenAuthCache()
|
||||
|
||||
func newAccessTokenAuthCache() *accessTokenAuthCache {
|
||||
return &accessTokenAuthCache{
|
||||
positive: make(map[string]cachedAgentNode),
|
||||
negative: make(map[string]time.Time),
|
||||
now: time.Now,
|
||||
loadNodeByToken: func(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
return model.GetOpenFlareNodeByAccessToken(ctx, token)
|
||||
},
|
||||
positive: make(map[string]cachedAgentNode),
|
||||
negative: make(map[string]time.Time),
|
||||
now: time.Now,
|
||||
loadNodeByToken: model.GetOpenFlareNodeByAccessToken,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -4,9 +4,9 @@
|
||||
package agent
|
||||
|
||||
const (
|
||||
errMissingAgentToken = "缺少 Agent Token"
|
||||
errInvalidAgentToken = "无权进行此操作,Agent Token 无效"
|
||||
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效"
|
||||
errMissingAgentToken = "缺少 Agent Token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errInvalidAgentToken = "无权进行此操作,Agent Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errNodeMissingFromContext = "Node object missing from context"
|
||||
errNoActiveConfig = "当前没有激活版本"
|
||||
errNodeNotFound = "节点不存在"
|
||||
|
||||
@@ -20,10 +20,12 @@ const (
|
||||
openrestyStatusUnhealthy = "unhealthy"
|
||||
openrestyStatusUnknown = "unknown"
|
||||
releaseChannelStable = "stable"
|
||||
randomTokenBytes = 16
|
||||
maxDatabaseTextLength = 16000
|
||||
)
|
||||
|
||||
func newRandomToken() (string, error) {
|
||||
buf := make([]byte, 16)
|
||||
buf := make([]byte, randomTokenBytes)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -55,9 +57,9 @@ func normalizeNodePayload(payload NodePayload) NodePayload {
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
payload.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
payload.LastError = truncateForDatabase(payload.LastError, maxDatabaseTextLength)
|
||||
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, maxDatabaseTextLength)
|
||||
return payload
|
||||
}
|
||||
|
||||
@@ -92,12 +94,12 @@ func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveNa
|
||||
node.Version = strings.TrimSpace(payload.Version)
|
||||
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, maxDatabaseTextLength)
|
||||
node.Status = nodeStatusOnline
|
||||
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
now := time.Now()
|
||||
node.LastSeenAt = &now
|
||||
node.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
node.LastError = truncateForDatabase(payload.LastError, maxDatabaseTextLength)
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
}
|
||||
@@ -135,15 +137,15 @@ func cloneCoordinate(value *float64) *float64 {
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func truncateForDatabase(value string, max int) string {
|
||||
if max <= 0 {
|
||||
func truncateForDatabase(value string, maxVal int) string {
|
||||
if maxVal <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(value))
|
||||
if len(runes) <= max {
|
||||
if len(runes) <= maxVal {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:max])
|
||||
return string(runes[:maxVal])
|
||||
}
|
||||
|
||||
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
|
||||
@@ -277,7 +279,7 @@ func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
|
||||
payload.NodeID = strings.TrimSpace(payload.NodeID)
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.Result = strings.ToLower(strings.TrimSpace(payload.Result))
|
||||
payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), 16000)
|
||||
payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), maxDatabaseTextLength)
|
||||
payload.Checksum = strings.TrimSpace(payload.Checksum)
|
||||
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
|
||||
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
|
||||
@@ -292,7 +294,7 @@ func isUniqueConstraintError(err error) bool {
|
||||
}
|
||||
|
||||
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
|
||||
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
|
||||
func RefreshAccessTokenCache(_ context.Context, node *model.OpenFlareNode) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -12,12 +12,12 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
agentTokenHeader = "X-Agent-Token"
|
||||
agentTokenHeader = "X-Agent-Token" //nolint:gosec // HTTP header name, not a credential value
|
||||
agentNodeContextKey = "agent_node"
|
||||
)
|
||||
|
||||
// AgentAuth validates X-Agent-Token against of_nodes.access_token.
|
||||
func AgentAuth() gin.HandlerFunc {
|
||||
// Auth validates X-Agent-Token against of_nodes.access_token.
|
||||
func Auth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
|
||||
node, err := AuthenticateAccessToken(c.Request.Context(), token)
|
||||
@@ -30,8 +30,8 @@ func AgentAuth() gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// AgentRegisterAuth accepts either a node access token or the global discovery token.
|
||||
func AgentRegisterAuth() gin.HandlerFunc {
|
||||
// RegisterAuth accepts either a node access token or the global discovery token.
|
||||
func RegisterAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
|
||||
if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil {
|
||||
@@ -48,12 +48,12 @@ func AgentRegisterAuth() gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// AgentNodeFromContext returns the authenticated agent node.
|
||||
func AgentNodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
|
||||
// NodeFromContext returns the authenticated agent node.
|
||||
func NodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
|
||||
value, ok := c.Get(agentNodeContextKey)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
node, ok := value.(*model.OpenFlareNode)
|
||||
return node, ok
|
||||
}
|
||||
}
|
||||
|
||||
@@ -109,8 +109,8 @@ func TestAgentAuthMiddleware(t *testing.T) {
|
||||
}).Error)
|
||||
|
||||
router := testhelper.NewTestGinEngine()
|
||||
router.GET("/protected", AgentAuth(), func(c *gin.Context) {
|
||||
node, ok := AgentNodeFromContext(c)
|
||||
router.GET("/protected", Auth(), func(c *gin.Context) {
|
||||
node, ok := NodeFromContext(c)
|
||||
if !ok {
|
||||
c.Status(http.StatusInternalServerError)
|
||||
return
|
||||
@@ -157,8 +157,8 @@ func TestAgentRegisterAuthMiddleware(t *testing.T) {
|
||||
require.NoError(t, model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", "discovery-token"))
|
||||
|
||||
router := testhelper.NewTestGinEngine()
|
||||
router.POST("/register", AgentRegisterAuth(), func(c *gin.Context) {
|
||||
if node, ok := AgentNodeFromContext(c); ok {
|
||||
router.POST("/register", RegisterAuth(), func(c *gin.Context) {
|
||||
if node, ok := NodeFromContext(c); ok {
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "node", "node_id": node.NodeID}))
|
||||
return
|
||||
}
|
||||
@@ -196,4 +196,4 @@ func TestAgentRegisterAuthMiddleware(t *testing.T) {
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "discovery", data["mode"])
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -27,6 +27,7 @@ const (
|
||||
nodeAccessLogRetentionDays = 90
|
||||
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
|
||||
accessLogPathMaxLength = 100
|
||||
healthEventMessageMaxLength = 4096
|
||||
)
|
||||
|
||||
// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events.
|
||||
@@ -396,7 +397,7 @@ func normalizeHealthSeverity(severity string) string {
|
||||
}
|
||||
|
||||
func normalizeHealthEventMessage(message string) string {
|
||||
return truncateForDatabase(message, 4096)
|
||||
return truncateForDatabase(message, healthEventMessageMaxLength)
|
||||
}
|
||||
|
||||
func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
|
||||
|
||||
@@ -5,22 +5,55 @@ package agent
|
||||
|
||||
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
|
||||
// NodePayload is the data sent by an agent on registration or heartbeat.
|
||||
type NodePayload = pkgprotocol.NodePayload
|
||||
|
||||
// NodeSystemProfile carries static host information reported by an agent.
|
||||
type NodeSystemProfile = pkgprotocol.NodeSystemProfile
|
||||
|
||||
// NodeMetricSnapshot holds a point-in-time resource-usage sample from an agent.
|
||||
type NodeMetricSnapshot = pkgprotocol.NodeMetricSnapshot
|
||||
|
||||
// NodeOpenrestyObservation reports the OpenResty process health observed by an agent.
|
||||
type NodeOpenrestyObservation = pkgprotocol.NodeOpenrestyObservation
|
||||
|
||||
// NodeTrafficReport aggregates traffic counters collected by an agent.
|
||||
type NodeTrafficReport = pkgprotocol.NodeTrafficReport
|
||||
|
||||
// NodeAccessLog is a single access-log record forwarded by an agent.
|
||||
type NodeAccessLog = pkgprotocol.NodeAccessLog
|
||||
|
||||
// BufferedObservabilityRecord bundles multiple observability payloads into one upload.
|
||||
type BufferedObservabilityRecord = pkgprotocol.BufferedObservabilityRecord
|
||||
|
||||
// NodeHealthEvent represents a discrete health-state change on an agent node.
|
||||
type NodeHealthEvent = pkgprotocol.NodeHealthEvent
|
||||
|
||||
// ApplyLogPayload carries the result of a configuration-apply attempt reported by an agent.
|
||||
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
|
||||
|
||||
// Settings contains remote-control directives sent from the server to an agent.
|
||||
type Settings = pkgprotocol.AgentSettings
|
||||
|
||||
// ActiveConfigMeta describes the currently active configuration version on the server.
|
||||
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
|
||||
|
||||
// SupportFile represents a supplementary file bundled with an agent configuration package.
|
||||
type SupportFile = pkgprotocol.SupportFile
|
||||
|
||||
// WAFIPGroup is a named IP-address group used in WAF allow/block rules.
|
||||
type WAFIPGroup = pkgprotocol.WAFIPGroup
|
||||
|
||||
// WAFIPGroupSyncRequest is sent by an agent to request an incremental WAF IP-group sync.
|
||||
type WAFIPGroupSyncRequest = pkgprotocol.WAFIPGroupSyncRequest
|
||||
|
||||
// WAFIPGroupSyncResponse carries the server's reply to a WAF IP-group sync request.
|
||||
type WAFIPGroupSyncResponse = pkgprotocol.WAFIPGroupSyncResponse
|
||||
|
||||
// Backward-compatible names used by server routers and handlers.
|
||||
|
||||
// WAFIPGroupSyncInput is an alias for WAFIPGroupSyncRequest kept for backward compatibility.
|
||||
type WAFIPGroupSyncInput = WAFIPGroupSyncRequest
|
||||
type WAFIPGroupSyncResult = WAFIPGroupSyncResponse
|
||||
|
||||
// WAFIPGroupSyncResult is an alias for WAFIPGroupSyncResponse kept for backward compatibility.
|
||||
type WAFIPGroupSyncResult = WAFIPGroupSyncResponse
|
||||
|
||||
@@ -37,7 +37,7 @@ func RegisterHandler(c *gin.Context) {
|
||||
result *RegistrationResponse
|
||||
err error
|
||||
)
|
||||
if authNode, ok := AgentNodeFromContext(c); ok {
|
||||
if authNode, ok := NodeFromContext(c); ok {
|
||||
result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload)
|
||||
} else {
|
||||
result, err = RegisterWithDiscovery(c.Request.Context(), payload)
|
||||
@@ -67,7 +67,7 @@ func HeartbeatHandler(c *gin.Context) {
|
||||
}
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
|
||||
|
||||
authNode, ok := AgentNodeFromContext(c)
|
||||
authNode, ok := NodeFromContext(c)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
@@ -91,7 +91,7 @@ func HeartbeatHandler(c *gin.Context) {
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/config-versions/active [get]
|
||||
func GetActiveConfigHandler(c *gin.Context) {
|
||||
if _, ok := AgentNodeFromContext(c); !ok {
|
||||
if _, ok := NodeFromContext(c); !ok {
|
||||
response.AbortUnauthorized(c, errNodeMissingFromContext)
|
||||
return
|
||||
}
|
||||
@@ -143,7 +143,7 @@ func ReportApplyLogHandler(c *gin.Context) {
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
if authNode, ok := AgentNodeFromContext(c); ok {
|
||||
if authNode, ok := NodeFromContext(c); ok {
|
||||
payload.NodeID = authNode.NodeID
|
||||
}
|
||||
log, err := ReportApplyLog(c.Request.Context(), payload)
|
||||
@@ -173,7 +173,7 @@ func DownloadPagesPackageHandler(c *gin.Context) {
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
defer packageObj.Body.Close()
|
||||
defer func() { _ = packageObj.Body.Close() }()
|
||||
c.Header("Content-Disposition", "attachment; filename="+fileName)
|
||||
if packageObj.ContentType != "" {
|
||||
c.Header("Content-Type", packageObj.ContentType)
|
||||
@@ -195,18 +195,18 @@ func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
|
||||
return uint(id64), true
|
||||
}
|
||||
|
||||
// AgentWebSocketHandler upgrades an authenticated agent websocket connection.
|
||||
// WebSocketHandler upgrades an authenticated agent websocket connection.
|
||||
// @Summary Agent WebSocket 连接
|
||||
// @Description 升级为 WebSocket 长连接,用于实时推送配置同步、WAF IP 组等指令;需携带 X-Agent-Token
|
||||
// @Tags openflare-agent
|
||||
// @Security AgentTokenAuth
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/ws [get]
|
||||
func AgentWebSocketHandler(c *gin.Context) {
|
||||
authNode, ok := AgentNodeFromContext(c)
|
||||
func WebSocketHandler(c *gin.Context) {
|
||||
authNode, ok := NodeFromContext(c)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
}
|
||||
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package apply_log manages the application of configuration change logs,
|
||||
// including validation and retention policy enforcement.
|
||||
package apply_log
|
||||
|
||||
const (
|
||||
|
||||
@@ -1,6 +1,4 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package openflare implements openflare configuration, service orchestration, and background tasks.
|
||||
package openflare
|
||||
|
||||
import (
|
||||
@@ -118,7 +116,7 @@ func (h *DatabaseAutoCleanupHandler) Execute(ctx context.Context, _ []byte) (*ta
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始执行可观测数据自动清理,保留天数=%d", model.DatabaseAutoCleanupRetentionDays)
|
||||
summary, err := tasks.RunDatabaseAutoCleanupOnce(time.Now())
|
||||
summary, err := tasks.RunDatabaseAutoCleanupOnce(ctx, time.Now())
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "可观测数据自动清理失败: %v", err)
|
||||
return nil, err
|
||||
@@ -197,4 +195,4 @@ func (h *UptimeKumaSyncHandler) Execute(ctx context.Context, _ []byte) (*task.Ta
|
||||
msg := "Uptime Kuma 同步完成"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package config_version defines shared error messages for configuration versions.
|
||||
package config_version
|
||||
|
||||
const (
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -17,6 +18,11 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
cleanupSuccessMessage = "清理成功"
|
||||
minConfigVersionKeepCount = 3
|
||||
)
|
||||
|
||||
// ConfigPreviewResult is the preview response payload.
|
||||
type ConfigPreviewResult struct {
|
||||
SnapshotJSON string `json:"snapshot_json"`
|
||||
@@ -243,15 +249,15 @@ func ActivateConfigVersion(ctx context.Context, id uint) (*model.ConfigVersion,
|
||||
|
||||
// CleanupConfigVersions removes old inactive config versions.
|
||||
func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, error) {
|
||||
if keepCount < 3 {
|
||||
keepCount = 3
|
||||
if keepCount < minConfigVersionKeepCount {
|
||||
keepCount = minConfigVersionKeepCount
|
||||
}
|
||||
versions, err := model.ListConfigVersionSummaries(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(versions) <= keepCount {
|
||||
return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil
|
||||
return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil
|
||||
}
|
||||
var deleteIDs []uint
|
||||
for index, version := range versions {
|
||||
@@ -264,13 +270,13 @@ func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult,
|
||||
deleteIDs = append(deleteIDs, version.ID)
|
||||
}
|
||||
if len(deleteIDs) == 0 {
|
||||
return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil
|
||||
return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil
|
||||
}
|
||||
deletedCount, err := model.DeleteConfigVersionsByIDs(ctx, deleteIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &CleanupResult{DeletedCount: deletedCount, Message: "清理成功"}, nil
|
||||
return &CleanupResult{DeletedCount: deletedCount, Message: cleanupSuccessMessage}, nil
|
||||
}
|
||||
|
||||
func nextVersionNumber(ctx context.Context, now time.Time) (string, error) {
|
||||
@@ -368,50 +374,50 @@ func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRo
|
||||
}
|
||||
|
||||
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
|
||||
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL ||
|
||||
left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP ||
|
||||
left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP ||
|
||||
left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy ||
|
||||
left.BasicAuthEnabled != right.BasicAuthEnabled || left.BasicAuthUsername != right.BasicAuthUsername ||
|
||||
left.BasicAuthPassword != right.BasicAuthPassword || left.UpstreamType != right.UpstreamType ||
|
||||
!uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) || left.TunnelTargetAddr != right.TunnelTargetAddr ||
|
||||
left.TunnelTargetProto != right.TunnelTargetProto || !uintPtrEqual(left.PagesProjectID, right.PagesProjectID) ||
|
||||
!uintSliceEqual(left.CertIDs, right.CertIDs) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) {
|
||||
return false
|
||||
}
|
||||
if len(left.Domains) != len(right.Domains) {
|
||||
return false
|
||||
}
|
||||
for index := range left.Domains {
|
||||
if left.Domains[index] != right.Domains[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.Upstreams) != len(right.Upstreams) {
|
||||
return false
|
||||
}
|
||||
for index := range left.Upstreams {
|
||||
if left.Upstreams[index] != right.Upstreams[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.CacheRules) != len(right.CacheRules) {
|
||||
return false
|
||||
}
|
||||
for index := range left.CacheRules {
|
||||
if left.CacheRules[index] != right.CacheRules[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.CustomHeaders) != len(right.CustomHeaders) {
|
||||
return false
|
||||
}
|
||||
for index := range left.CustomHeaders {
|
||||
if left.CustomHeaders[index] != right.CustomHeaders[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
return snapshotRouteScalarsEqual(left, right) &&
|
||||
slices.Equal(left.Domains, right.Domains) &&
|
||||
slices.Equal(left.Upstreams, right.Upstreams) &&
|
||||
slices.Equal(left.CacheRules, right.CacheRules) &&
|
||||
slices.Equal(left.CustomHeaders, right.CustomHeaders)
|
||||
}
|
||||
|
||||
func snapshotRouteScalarsEqual(left, right snapshotRoute) bool {
|
||||
return snapshotRouteIdentityEqual(left, right) &&
|
||||
snapshotRouteOriginEqual(left, right) &&
|
||||
snapshotRoutePolicyEqual(left, right) &&
|
||||
snapshotRouteTunnelEqual(left, right) &&
|
||||
uintSliceEqual(left.CertIDs, right.CertIDs) &&
|
||||
uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs)
|
||||
}
|
||||
|
||||
func snapshotRouteIdentityEqual(left, right snapshotRoute) bool {
|
||||
return left.SiteName == right.SiteName && left.Domain == right.Domain
|
||||
}
|
||||
|
||||
func snapshotRouteOriginEqual(left, right snapshotRoute) bool {
|
||||
return left.OriginURL == right.OriginURL &&
|
||||
left.OriginHost == right.OriginHost &&
|
||||
left.UpstreamType == right.UpstreamType
|
||||
}
|
||||
|
||||
func snapshotRoutePolicyEqual(left, right snapshotRoute) bool {
|
||||
return left.EnableHTTPS == right.EnableHTTPS &&
|
||||
left.RedirectHTTP == right.RedirectHTTP &&
|
||||
left.LimitConnPerServer == right.LimitConnPerServer &&
|
||||
left.LimitConnPerIP == right.LimitConnPerIP &&
|
||||
left.LimitRate == right.LimitRate &&
|
||||
left.CacheEnabled == right.CacheEnabled &&
|
||||
left.CachePolicy == right.CachePolicy &&
|
||||
left.BasicAuthEnabled == right.BasicAuthEnabled &&
|
||||
left.BasicAuthUsername == right.BasicAuthUsername &&
|
||||
left.BasicAuthPassword == right.BasicAuthPassword
|
||||
}
|
||||
|
||||
func snapshotRouteTunnelEqual(left, right snapshotRoute) bool {
|
||||
return left.TunnelTargetAddr == right.TunnelTargetAddr &&
|
||||
left.TunnelTargetProto == right.TunnelTargetProto &&
|
||||
uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) &&
|
||||
uintPtrEqual(left.PagesProjectID, right.PagesProjectID)
|
||||
}
|
||||
|
||||
func snapshotWAFConfigEqual(left snapshotWAFDocument, right snapshotWAFDocument) bool {
|
||||
|
||||
@@ -16,6 +16,8 @@ import (
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
)
|
||||
|
||||
const supportFilesPerCertificate = 2
|
||||
|
||||
type snapshotRoute struct {
|
||||
ID uint `json:"id,omitempty"`
|
||||
SiteName string `json:"site_name,omitempty"`
|
||||
@@ -227,13 +229,14 @@ func buildSnapshotRoutes(ctx context.Context, routes []*model.ProxyRoute) ([]sna
|
||||
var tunnelTargetAddr string
|
||||
var tunnelTargetProtocol string
|
||||
var pagesProjectID *uint
|
||||
if upstreamType == "tunnel" {
|
||||
switch upstreamType {
|
||||
case "tunnel":
|
||||
originURL = resolveTunnelOpenRestyUpstreamURL(ctx)
|
||||
upstreams = []string{originURL}
|
||||
tunnelNodeID = route.TunnelNodeID
|
||||
tunnelTargetAddr = strings.TrimSpace(route.TunnelTargetAddr)
|
||||
tunnelTargetProtocol = normalizeTunnelTargetProtocol(route.TunnelTargetProtocol)
|
||||
} else if upstreamType == "pages" {
|
||||
case "pages":
|
||||
return nil, fmt.Errorf("路由 %s Pages 配置无效: pages module is not available", route.Domain)
|
||||
}
|
||||
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
|
||||
@@ -495,7 +498,7 @@ func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) (
|
||||
certIDs = append(certIDs, certID)
|
||||
}
|
||||
sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] })
|
||||
files := make([]SupportFile, 0, len(certIDs)*2)
|
||||
files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate)
|
||||
for _, certID := range certIDs {
|
||||
certificate, err := model.GetTLSCertificateByID(ctx, certID)
|
||||
if err != nil {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package dashboard provides helper utilities for dashboard API handlers.
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
@@ -15,6 +16,11 @@ const (
|
||||
nodeStatusOnline = "online"
|
||||
nodeStatusOffline = "offline"
|
||||
nodeStatusPending = "pending"
|
||||
|
||||
dashboardDistributionLimit = 8
|
||||
highCPUUsagePercentThreshold = 80
|
||||
highMemoryUsagePercentThreshold = 85
|
||||
highStorageUsagePercentThreshold = 85
|
||||
)
|
||||
|
||||
func computeNodeStatus(node *model.OpenFlareNode) string {
|
||||
|
||||
@@ -120,7 +120,7 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, 8)
|
||||
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, dashboardDistributionLimit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -136,7 +136,7 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
|
||||
view := &OverviewView{
|
||||
GeneratedAt: now,
|
||||
Nodes: make([]NodeHealth, 0, len(nodes)),
|
||||
Distributions: observability.BuildTrafficDistributions(reports, accessLogRegions, 8),
|
||||
Distributions: observability.BuildTrafficDistributions(reports, accessLogRegions, dashboardDistributionLimit),
|
||||
Trends: observability.NodeTrends{
|
||||
Traffic24h: observability.BuildTrafficTrendPoints(now, reports),
|
||||
Capacity24h: observability.BuildCapacityTrendPoints(now, snapshots),
|
||||
@@ -183,41 +183,8 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
|
||||
ActiveEventCount: len(nodeActiveEvents),
|
||||
}
|
||||
|
||||
if latestSnapshot != nil {
|
||||
nodeHealth.CPUUsagePercent = latestSnapshot.CPUUsagePercent
|
||||
nodeHealth.MemoryUsagePercent = observability.Percentage(latestSnapshot.MemoryUsedBytes, latestSnapshot.MemoryTotalBytes)
|
||||
nodeHealth.StorageUsagePercent = observability.Percentage(latestSnapshot.StorageUsedBytes, latestSnapshot.StorageTotalBytes)
|
||||
if latestSnapshot.CPUUsagePercent > 0 {
|
||||
view.Capacity.AverageCPUUsagePercent += latestSnapshot.CPUUsagePercent
|
||||
cpuNodeCount++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent > 0 {
|
||||
view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent
|
||||
memoryNodeCount++
|
||||
}
|
||||
if latestSnapshot.CPUUsagePercent >= 80 {
|
||||
view.Capacity.HighCPUNodes++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent >= 85 {
|
||||
view.Capacity.HighMemoryNodes++
|
||||
}
|
||||
if nodeHealth.StorageUsagePercent >= 85 {
|
||||
view.Capacity.HighStorageNodes++
|
||||
}
|
||||
}
|
||||
|
||||
if latestTraffic != nil {
|
||||
nodeHealth.RequestCount = latestTraffic.RequestCount
|
||||
nodeHealth.ErrorCount = latestTraffic.ErrorCount
|
||||
nodeHealth.UniqueVisitorCount = latestTraffic.UniqueVisitorCount
|
||||
view.Traffic.RequestCount += latestTraffic.RequestCount
|
||||
view.Traffic.UniqueVisitors += latestTraffic.UniqueVisitorCount
|
||||
view.Traffic.ErrorCount += latestTraffic.ErrorCount
|
||||
if duration := latestTraffic.WindowEndedAt.Sub(latestTraffic.WindowStartedAt).Seconds(); duration > 0 {
|
||||
view.Traffic.EstimatedQPS += float64(latestTraffic.RequestCount) / duration
|
||||
}
|
||||
view.Traffic.ReportedNodes++
|
||||
}
|
||||
cpuNodeCount, memoryNodeCount = applyNodeSnapshotMetrics(&nodeHealth, latestSnapshot, view, cpuNodeCount, memoryNodeCount)
|
||||
applyNodeTrafficMetrics(&nodeHealth, latestTraffic, view)
|
||||
|
||||
view.Nodes = append(view.Nodes, nodeHealth)
|
||||
}
|
||||
@@ -240,6 +207,49 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func applyNodeSnapshotMetrics(nodeHealth *NodeHealth, snapshot *model.OpenFlareMetricSnapshot, view *OverviewView, cpuNodeCount, memoryNodeCount int) (int, int) {
|
||||
if snapshot == nil {
|
||||
return cpuNodeCount, memoryNodeCount
|
||||
}
|
||||
nodeHealth.CPUUsagePercent = snapshot.CPUUsagePercent
|
||||
nodeHealth.MemoryUsagePercent = observability.Percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes)
|
||||
nodeHealth.StorageUsagePercent = observability.Percentage(snapshot.StorageUsedBytes, snapshot.StorageTotalBytes)
|
||||
if snapshot.CPUUsagePercent > 0 {
|
||||
view.Capacity.AverageCPUUsagePercent += snapshot.CPUUsagePercent
|
||||
cpuNodeCount++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent > 0 {
|
||||
view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent
|
||||
memoryNodeCount++
|
||||
}
|
||||
if snapshot.CPUUsagePercent >= highCPUUsagePercentThreshold {
|
||||
view.Capacity.HighCPUNodes++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent >= highMemoryUsagePercentThreshold {
|
||||
view.Capacity.HighMemoryNodes++
|
||||
}
|
||||
if nodeHealth.StorageUsagePercent >= highStorageUsagePercentThreshold {
|
||||
view.Capacity.HighStorageNodes++
|
||||
}
|
||||
return cpuNodeCount, memoryNodeCount
|
||||
}
|
||||
|
||||
func applyNodeTrafficMetrics(nodeHealth *NodeHealth, traffic *model.OpenFlareRequestReport, view *OverviewView) {
|
||||
if traffic == nil {
|
||||
return
|
||||
}
|
||||
nodeHealth.RequestCount = traffic.RequestCount
|
||||
nodeHealth.ErrorCount = traffic.ErrorCount
|
||||
nodeHealth.UniqueVisitorCount = traffic.UniqueVisitorCount
|
||||
view.Traffic.RequestCount += traffic.RequestCount
|
||||
view.Traffic.UniqueVisitors += traffic.UniqueVisitorCount
|
||||
view.Traffic.ErrorCount += traffic.ErrorCount
|
||||
if duration := traffic.WindowEndedAt.Sub(traffic.WindowStartedAt).Seconds(); duration > 0 {
|
||||
view.Traffic.EstimatedQPS += float64(traffic.RequestCount) / duration
|
||||
}
|
||||
view.Traffic.ReportedNodes++
|
||||
}
|
||||
|
||||
func compressOverview(view *OverviewView) *OverviewPayload {
|
||||
if view == nil {
|
||||
return &OverviewPayload{
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package flared defines shared error messages for tunnel client operations.
|
||||
package flared
|
||||
|
||||
const (
|
||||
errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效"
|
||||
errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTunnelNodeTypeMismatch = "此节点不是 TunnelClient 类型"
|
||||
)
|
||||
|
||||
@@ -18,6 +18,11 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
updateChannelStable = "stable"
|
||||
defaultTunnelTargetPort = 80
|
||||
)
|
||||
|
||||
type configVersionRow struct {
|
||||
Version string `gorm:"column:version"`
|
||||
Checksum string `gorm:"column:checksum"`
|
||||
@@ -31,7 +36,7 @@ func normalizeReleaseChannel(channel string) string {
|
||||
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
|
||||
return "preview"
|
||||
}
|
||||
return "stable"
|
||||
return updateChannelStable
|
||||
}
|
||||
|
||||
func normalizeFlaredHeartbeatPayload(payload HeartbeatPayload) HeartbeatPayload {
|
||||
@@ -146,20 +151,20 @@ func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) {
|
||||
func parseTunnelTargetAddr(addr string) (string, int) {
|
||||
addr = strings.TrimSpace(addr)
|
||||
if addr == "" {
|
||||
return "127.0.0.1", 80
|
||||
return "127.0.0.1", defaultTunnelTargetPort
|
||||
}
|
||||
host, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
lastColon := strings.LastIndex(addr, ":")
|
||||
if lastColon < 0 {
|
||||
return addr, 80
|
||||
return addr, defaultTunnelTargetPort
|
||||
}
|
||||
host = addr[:lastColon]
|
||||
portStr = addr[lastColon+1:]
|
||||
}
|
||||
port := 80
|
||||
port := defaultTunnelTargetPort
|
||||
if _, scanErr := fmt.Sscanf(portStr, "%d", &port); scanErr != nil {
|
||||
port = 80
|
||||
port = defaultTunnelTargetPort
|
||||
}
|
||||
if host == "" {
|
||||
host = "127.0.0.1"
|
||||
|
||||
@@ -17,10 +17,11 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
applyResultOK = "success"
|
||||
applyResultWarn = "warning"
|
||||
applyResultFail = "failed"
|
||||
nodeStatusOnline = "online"
|
||||
applyResultOK = "success"
|
||||
applyResultWarn = "warning"
|
||||
applyResultFail = "failed"
|
||||
maxApplyLogMessageLength = 16000
|
||||
)
|
||||
|
||||
// Heartbeat processes an OpenFlared heartbeat and returns runtime settings.
|
||||
@@ -46,13 +47,13 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
|
||||
"last_seen_at": now,
|
||||
"status": nodeStatusOnline,
|
||||
"update_requested": false,
|
||||
"update_channel": "stable",
|
||||
"update_channel": updateChannelStable,
|
||||
"update_tag": "",
|
||||
}
|
||||
if !previous.UpdateRequested {
|
||||
delete(changes, "update_requested")
|
||||
}
|
||||
if previous.UpdateChannel == "stable" {
|
||||
if previous.UpdateChannel == updateChannelStable {
|
||||
delete(changes, "update_channel")
|
||||
}
|
||||
if previous.UpdateTag == "" {
|
||||
@@ -67,7 +68,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
|
||||
node.ExtVersion = payload.FrpVersion
|
||||
node.CurrentVersion = payload.CurrentVersion
|
||||
node.UpdateRequested = false
|
||||
node.UpdateChannel = "stable"
|
||||
node.UpdateChannel = updateChannelStable
|
||||
node.UpdateTag = ""
|
||||
lastSeen := now
|
||||
node.LastSeenAt = &lastSeen
|
||||
@@ -228,8 +229,8 @@ func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
|
||||
payload.Checksum = strings.TrimSpace(payload.Checksum)
|
||||
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
|
||||
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
|
||||
if len(payload.Message) > 16000 {
|
||||
payload.Message = payload.Message[:16000]
|
||||
if len(payload.Message) > maxApplyLogMessageLength {
|
||||
payload.Message = payload.Message[:maxApplyLogMessageLength]
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
@@ -5,11 +5,26 @@ package flared
|
||||
|
||||
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
|
||||
// HeartbeatPayload is an alias for FlaredHeartbeatPayload.
|
||||
type HeartbeatPayload = pkgprotocol.FlaredHeartbeatPayload
|
||||
|
||||
// ConnectedRelay is an alias for FlaredConnectedRelay.
|
||||
type ConnectedRelay = pkgprotocol.FlaredConnectedRelay
|
||||
|
||||
// ActiveConfigMeta is an alias for ActiveConfigMeta.
|
||||
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
|
||||
|
||||
// HeartbeatResponse is an alias for FlaredHeartbeatResponse.
|
||||
type HeartbeatResponse = pkgprotocol.FlaredHeartbeatResponse
|
||||
|
||||
// TunnelConfigResponse is an alias for FlaredTunnelConfigResponse.
|
||||
type TunnelConfigResponse = pkgprotocol.FlaredTunnelConfigResponse
|
||||
|
||||
// RelayInfo is an alias for FlaredRelayInfo.
|
||||
type RelayInfo = pkgprotocol.FlaredRelayInfo
|
||||
|
||||
// ProxyEntry is an alias for FlaredProxyEntry.
|
||||
type ProxyEntry = pkgprotocol.FlaredProxyEntry
|
||||
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
|
||||
|
||||
// ApplyLogPayload is an alias for ApplyLogPayload.
|
||||
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
|
||||
|
||||
@@ -31,7 +31,7 @@ func (f *fakeLookupProvider) Close() error { return nil }
|
||||
|
||||
func TestLookupWithProvider(t *testing.T) {
|
||||
previousFactory := pkggeoip.ProviderFactoryForTest()
|
||||
pkggeoip.SetProviderFactoryForTest(func(provider string) (pkggeoip.GeoIPService, error) {
|
||||
pkggeoip.SetProviderFactoryForTest(func(provider string) (pkggeoip.Service, error) {
|
||||
return &fakeLookupProvider{}, nil
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package node defines node validation and management error messages.
|
||||
package node
|
||||
|
||||
const (
|
||||
|
||||
@@ -28,6 +28,13 @@ const (
|
||||
openrestyStatusUnhealthy = "unhealthy"
|
||||
openrestyStatusUnknown = "unknown"
|
||||
githubReleasesAPIBase = "https://api.github.com/repos/%s/releases"
|
||||
nodeTypeTunnelRelay = "tunnel_relay"
|
||||
nodeTypeTunnelClient = "tunnel_client"
|
||||
nodeTypeEdgeNode = "edge_node"
|
||||
|
||||
nodeTokenByteLength = 16
|
||||
maxNodeIPLength = 64
|
||||
maxNodeGeoNameLength = 128
|
||||
)
|
||||
|
||||
type releaseChannel string
|
||||
@@ -49,7 +56,7 @@ type githubReleaseResponse struct {
|
||||
}
|
||||
|
||||
func newRandomToken() (string, error) {
|
||||
buf := make([]byte, 16)
|
||||
buf := make([]byte, nodeTokenByteLength)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -66,12 +73,12 @@ func newServerNodeID() (string, error) {
|
||||
|
||||
func normalizeNodeType(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "tunnel_relay":
|
||||
return "tunnel_relay"
|
||||
case "tunnel_client":
|
||||
return "tunnel_client"
|
||||
case nodeTypeTunnelRelay:
|
||||
return nodeTypeTunnelRelay
|
||||
case nodeTypeTunnelClient:
|
||||
return nodeTypeTunnelClient
|
||||
default:
|
||||
return "edge_node"
|
||||
return nodeTypeEdgeNode
|
||||
}
|
||||
}
|
||||
|
||||
@@ -134,40 +141,50 @@ func normalizeNodeInput(input Input) (string, string, string, *float64, *float64
|
||||
name := strings.TrimSpace(input.Name)
|
||||
ip := strings.TrimSpace(input.IP)
|
||||
geoName := strings.TrimSpace(input.GeoName)
|
||||
manualOverride := input.GeoManualOverride || geoName != "" || input.GeoLatitude != nil || input.GeoLongitude != nil
|
||||
if len(ip) > 64 {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPTooLong)
|
||||
if err := validateNodeIPInput(input, ip); err != nil {
|
||||
return "", "", "", nil, nil, false, err
|
||||
}
|
||||
if ip != "" && net.ParseIP(ip) == nil {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPInvalid)
|
||||
}
|
||||
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPManualRequired)
|
||||
}
|
||||
if len(geoName) > 128 {
|
||||
if len(geoName) > maxNodeGeoNameLength {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoNameTooLong)
|
||||
}
|
||||
|
||||
geoLatitude := cloneCoordinate(input.GeoLatitude)
|
||||
geoLongitude := cloneCoordinate(input.GeoLongitude)
|
||||
if err := validateNodeGeoCoordinates(geoLatitude, geoLongitude); err != nil {
|
||||
return "", "", "", nil, nil, false, err
|
||||
}
|
||||
|
||||
manualOverride := input.GeoManualOverride || geoName != "" || geoLatitude != nil || geoLongitude != nil
|
||||
if !manualOverride || (geoLatitude == nil && geoLongitude == nil && geoName == "") {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
|
||||
}
|
||||
|
||||
func validateNodeIPInput(input Input, ip string) error {
|
||||
if len(ip) > maxNodeIPLength {
|
||||
return fmt.Errorf("%s", errNodeIPTooLong)
|
||||
}
|
||||
if ip != "" && net.ParseIP(ip) == nil {
|
||||
return fmt.Errorf("%s", errNodeIPInvalid)
|
||||
}
|
||||
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
|
||||
return fmt.Errorf("%s", errNodeIPManualRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateNodeGeoCoordinates(geoLatitude, geoLongitude *float64) error {
|
||||
if (geoLatitude == nil) != (geoLongitude == nil) {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoCoordinateMismatch)
|
||||
return fmt.Errorf("%s", errNodeGeoCoordinateMismatch)
|
||||
}
|
||||
if geoLatitude != nil && (*geoLatitude < -90 || *geoLatitude > 90) {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLatitudeInvalid)
|
||||
return fmt.Errorf("%s", errNodeGeoLatitudeInvalid)
|
||||
}
|
||||
if geoLongitude != nil && (*geoLongitude < -180 || *geoLongitude > 180) {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLongitudeInvalid)
|
||||
return fmt.Errorf("%s", errNodeGeoLongitudeInvalid)
|
||||
}
|
||||
|
||||
if !manualOverride {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
if geoLatitude == nil && geoLongitude == nil && geoName == "" {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
|
||||
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func computeNodeStatus(node *model.OpenFlareNode) string {
|
||||
@@ -189,12 +206,12 @@ func nodeViewLastSeenAt(node *model.OpenFlareNode) any {
|
||||
}
|
||||
nodeType := strings.TrimSpace(node.NodeType)
|
||||
if nodeType == "" {
|
||||
nodeType = "edge_node"
|
||||
nodeType = nodeTypeEdgeNode
|
||||
}
|
||||
if nodeType == "tunnel_relay" && ofws.IsRelayConnected(node.NodeID) {
|
||||
if nodeType == nodeTypeTunnelRelay && ofws.IsRelayConnected(node.NodeID) {
|
||||
return ofws.RelayWSConnectedLastSeenValue
|
||||
}
|
||||
if nodeType == "tunnel_client" && ofws.IsFlaredConnected(node.NodeID) {
|
||||
if nodeType == nodeTypeTunnelClient && ofws.IsFlaredConnected(node.NodeID) {
|
||||
return ofws.FlaredWSConnectedLastSeenValue
|
||||
}
|
||||
if ofws.IsAgentConnected(node.NodeID) {
|
||||
@@ -250,7 +267,7 @@ func buildNodeView(node *model.OpenFlareNode) *View {
|
||||
view.UpdateChannel = releaseChannelStable.String()
|
||||
}
|
||||
if view.NodeType == "" {
|
||||
view.NodeType = "edge_node"
|
||||
view.NodeType = nodeTypeEdgeNode
|
||||
}
|
||||
return view
|
||||
}
|
||||
@@ -378,7 +395,7 @@ func fetchLatestStableGitHubRelease(ctx context.Context, repo string) (*githubRe
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取最新版本失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
@@ -395,7 +412,7 @@ func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubR
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取 preview 版本失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
@@ -427,7 +444,7 @@ func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*git
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取指定版本失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return nil, fmt.Errorf("未找到指定版本: %s", tag)
|
||||
}
|
||||
@@ -461,13 +478,3 @@ func isUniqueConstraintError(err error) bool {
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
func setReleaseHTTPClientForTest(client *http.Client) *http.Client {
|
||||
previous := releaseHTTPClient
|
||||
if client == nil {
|
||||
releaseHTTPClient = &http.Client{Timeout: 30 * time.Second}
|
||||
} else {
|
||||
releaseHTTPClient = client
|
||||
}
|
||||
return previous
|
||||
}
|
||||
|
||||
@@ -17,6 +17,11 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultRelayBindPort = 7000
|
||||
defaultRelayVhostHTTPPort = 8080
|
||||
)
|
||||
|
||||
// Input is the create/update node payload.
|
||||
type Input struct {
|
||||
Name string `json:"name"`
|
||||
@@ -184,8 +189,8 @@ func CreateNode(ctx context.Context, input Input) (*View, error) {
|
||||
return nil, err
|
||||
}
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, 7000)
|
||||
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, 8080)
|
||||
node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, defaultRelayBindPort)
|
||||
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, defaultRelayVhostHTTPPort)
|
||||
node.RelayAuthToken, err = newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -399,7 +404,7 @@ func ValidateDiscoveryToken(ctx context.Context, token string) error {
|
||||
return err
|
||||
}
|
||||
if token != discoveryToken {
|
||||
return fmt.Errorf("Discovery Token 无效")
|
||||
return fmt.Errorf("discovery Token 无效") // error 消息首字母小写
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,6 +1,4 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package observability provides monitoring, metrics, and access log analysis for OpenFlare.
|
||||
package observability
|
||||
|
||||
import (
|
||||
@@ -22,6 +20,8 @@ const (
|
||||
defaultIPTrendBucketMinute = 30
|
||||
maxIPTrendHours = 168
|
||||
nodeAccessLogRetentionDays = 90
|
||||
accessLogFieldRemoteAddr = "remote_addr"
|
||||
accessLogFieldRequestCount = "request_count"
|
||||
)
|
||||
|
||||
var nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
|
||||
@@ -524,9 +524,9 @@ func normalizeFoldedAccessLogIPQuery(input FoldedAccessLogIPQuery) (FoldedAccess
|
||||
}
|
||||
normalizedSortBy := strings.TrimSpace(input.SortBy)
|
||||
switch normalizedSortBy {
|
||||
case "last_seen_at", "remote_addr":
|
||||
case "last_seen_at", accessLogFieldRemoteAddr:
|
||||
default:
|
||||
normalizedSortBy = "request_count"
|
||||
normalizedSortBy = accessLogFieldRequestCount
|
||||
}
|
||||
return FoldedAccessLogIPQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
@@ -591,7 +591,7 @@ func normalizeAccessLogPageSize(pageSize int) int {
|
||||
|
||||
func normalizeAccessLogSortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "status_code", "remote_addr", "host", "path":
|
||||
case "status_code", accessLogFieldRemoteAddr, "host", "path":
|
||||
return strings.TrimSpace(sortBy)
|
||||
default:
|
||||
return defaultAccessLogSortBy
|
||||
@@ -607,8 +607,8 @@ func normalizeAccessLogSortOrder(sortOrder string) string {
|
||||
|
||||
func normalizeFoldSortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "request_count":
|
||||
return "request_count"
|
||||
case accessLogFieldRequestCount:
|
||||
return accessLogFieldRequestCount
|
||||
default:
|
||||
return "bucket_started_at"
|
||||
}
|
||||
@@ -616,7 +616,7 @@ func normalizeFoldSortBy(sortBy string) string {
|
||||
|
||||
func normalizeIPSummarySortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "recent_requests", "last_seen_at", "remote_addr":
|
||||
case "recent_requests", "last_seen_at", accessLogFieldRemoteAddr:
|
||||
return strings.TrimSpace(sortBy)
|
||||
default:
|
||||
return "total_requests"
|
||||
|
||||
@@ -19,6 +19,7 @@ const (
|
||||
healthEventStatusResolved = "resolved"
|
||||
healthSeverityCritical = "critical"
|
||||
healthSeverityWarning = "warning"
|
||||
percentageMultiplier = 100
|
||||
)
|
||||
|
||||
// DistributionItem is a key/value distribution entry.
|
||||
@@ -415,7 +416,7 @@ func Percentage(used int64, total int64) float64 {
|
||||
if used <= 0 || total <= 0 {
|
||||
return 0
|
||||
}
|
||||
return (float64(used) / float64(total)) * 100
|
||||
return (float64(used) / float64(total)) * percentageMultiplier
|
||||
}
|
||||
|
||||
func mergeJSONCounts(target distributionAccumulator, raw string) {
|
||||
|
||||
@@ -14,9 +14,10 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
defaultObservabilityWindow = 24 * time.Hour
|
||||
defaultObservabilityLimit = 120
|
||||
maxObservabilityLimit = 500
|
||||
defaultObservabilityWindow = 24 * time.Hour
|
||||
defaultObservabilityLimit = 120
|
||||
maxObservabilityLimit = 500
|
||||
defaultTrafficDistributionLimit = 8
|
||||
)
|
||||
|
||||
// NodeQuery filters node observability data.
|
||||
@@ -106,7 +107,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, 8)
|
||||
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, defaultTrafficDistributionLimit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -135,7 +136,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV
|
||||
HealthEvents: events,
|
||||
Analytics: NodeAnalytics{
|
||||
Traffic: buildTrafficWindowSummary(latestTrafficReport(reports)),
|
||||
Distributions: BuildTrafficDistributions(reports, accessLogRegions, 8),
|
||||
Distributions: BuildTrafficDistributions(reports, accessLogRegions, defaultTrafficDistributionLimit),
|
||||
Health: buildHealthSummary(latestMetricSnapshot(snapshots), latestTrafficReport(reports), events),
|
||||
},
|
||||
Trends: NodeTrends{
|
||||
|
||||
@@ -40,7 +40,7 @@ func GetAccessLogsHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(logs))
|
||||
}
|
||||
|
||||
// getFoldedAccessLogsHandler 分页列出折叠访问日志。
|
||||
// GetFoldedAccessLogsHandler 分页列出折叠访问日志。
|
||||
// @Summary 列出折叠访问日志
|
||||
// @Description 按时间桶聚合访问日志并分页返回,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
@@ -71,7 +71,7 @@ func GetFoldedAccessLogsHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(logs))
|
||||
}
|
||||
|
||||
// getFoldedAccessLogIPsHandler 列出折叠桶内的 IP 汇总。
|
||||
// GetFoldedAccessLogIPsHandler 列出折叠桶内的 IP 汇总。
|
||||
// @Summary 列出折叠访问日志 IP 汇总
|
||||
// @Description 在指定时间桶内按 IP 聚合访问统计,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
@@ -112,7 +112,7 @@ func GetFoldedAccessLogIPsHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// getAccessLogIPSummariesHandler 列出访问日志 IP 汇总。
|
||||
// GetAccessLogIPSummariesHandler 列出访问日志 IP 汇总。
|
||||
// @Summary 列出访问日志 IP 汇总
|
||||
// @Description 按 IP 聚合访问日志统计并分页返回,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
@@ -147,7 +147,7 @@ func GetAccessLogIPSummariesHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// getAccessLogIPTrendHandler 获取 IP 访问趋势。
|
||||
// GetAccessLogIPTrendHandler 获取 IP 访问趋势。
|
||||
// @Summary 获取访问日志 IP 趋势
|
||||
// @Description 返回指定 IP 在时间范围内的访问趋势数据,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
@@ -178,7 +178,7 @@ func GetAccessLogIPTrendHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// cleanupAccessLogsHandler 清理过期访问日志。
|
||||
// CleanupAccessLogsHandler 清理过期访问日志。
|
||||
// @Summary 清理访问日志
|
||||
// @Description 按保留天数清理过期访问日志记录,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
@@ -220,4 +220,4 @@ func readAccessLogQuery(c *gin.Context) AccessLogQuery {
|
||||
func readQueryInt(c *gin.Context, key string) int {
|
||||
value, _ := strconv.Atoi(c.DefaultQuery(key, "0"))
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package option provides handler-level error message constants for the openflare option module.
|
||||
package option
|
||||
|
||||
const (
|
||||
|
||||
@@ -161,7 +161,7 @@ func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) {
|
||||
StartTime: model.StartTime,
|
||||
EmailVerification: model.EmailVerificationEnabled,
|
||||
GitHubOAuth: model.GitHubOAuthEnabled,
|
||||
GitHubClientID: model.GitHubClientId,
|
||||
GitHubClientID: model.GitHubClientID,
|
||||
SystemName: model.SystemName,
|
||||
HomePageLink: model.HomePageLink,
|
||||
FooterHTML: model.Footer,
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package option
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var openRestyOptionValidators = map[string]func(key, value string) error{
|
||||
"OpenRestyDefaultServerReturnStatus": validateOpenRestyDefaultServerReturnStatus,
|
||||
"OpenRestyWorkerProcesses": validateOpenRestyWorkerProcesses,
|
||||
"OpenRestyWorkerConnections": validatePositiveIntegerOption,
|
||||
"OpenRestyWorkerRlimitNofile": validatePositiveIntegerOption,
|
||||
"OpenRestyKeepaliveTimeout": validatePositiveIntegerOption,
|
||||
"OpenRestyKeepaliveRequests": validatePositiveIntegerOption,
|
||||
"OpenRestyClientHeaderTimeout": validatePositiveIntegerOption,
|
||||
"OpenRestyClientBodyTimeout": validatePositiveIntegerOption,
|
||||
"OpenRestySendTimeout": validatePositiveIntegerOption,
|
||||
"OpenRestyProxyConnectTimeout": validatePositiveIntegerOption,
|
||||
"OpenRestyProxySendTimeout": validatePositiveIntegerOption,
|
||||
"OpenRestyProxyReadTimeout": validatePositiveIntegerOption,
|
||||
"OpenRestyGzipMinLength": validatePositiveIntegerOption,
|
||||
"OpenRestyGzipCompLevel": validateOpenRestyGzipCompLevel,
|
||||
"OpenRestyEventsUse": validateOpenRestyEventsUse,
|
||||
"OpenRestyResolvers": validateOpenRestyResolvers,
|
||||
"OpenRestyEventsMultiAcceptEnabled": validateBooleanOption,
|
||||
"OpenRestyWebsocketEnabled": validateBooleanOption,
|
||||
"OpenRestyHTTP3Enabled": validateBooleanOption,
|
||||
"OpenRestyProxyRequestBufferingEnabled": validateBooleanOption,
|
||||
"OpenRestyProxyBufferingEnabled": validateBooleanOption,
|
||||
"OpenRestyGzipEnabled": validateBooleanOption,
|
||||
"OpenRestyCacheEnabled": validateBooleanOption,
|
||||
"OpenRestyCacheLockEnabled": validateBooleanOption,
|
||||
"OpenRestyProxyBuffers": validateOpenRestyProxyBuffers,
|
||||
"OpenRestyLargeClientHeaderBuffers": validateOpenRestyProxyBuffers,
|
||||
"OpenRestyProxyBufferSize": validateOpenRestySizeValue,
|
||||
"OpenRestyProxyBusyBuffersSize": validateOpenRestySizeValue,
|
||||
"OpenRestyCacheMaxSize": validateOpenRestySizeValue,
|
||||
"OpenRestyClientMaxBodySize": validateOpenRestySizeValue,
|
||||
"OpenRestyCachePath": validateOpenRestyCachePath,
|
||||
"OpenRestyCacheLevels": validateOpenRestyCacheLevels,
|
||||
"OpenRestyCacheInactive": validateOpenRestyDurationToken,
|
||||
"OpenRestyCacheLockTimeout": validateOpenRestyDurationToken,
|
||||
"OpenRestyCacheKeyTemplate": validateOpenRestyCacheKeyTemplate,
|
||||
"OpenRestyCacheUseStale": validateOpenRestyCacheUseStale,
|
||||
"OpenRestyMainConfigTemplate": validateOpenRestyMainConfigTemplate,
|
||||
}
|
||||
|
||||
func validateOpenRestyOption(key, value string) error {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if validator, ok := openRestyOptionValidators[key]; ok {
|
||||
return validator(key, trimmed)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyDefaultServerReturnStatus(key, trimmed string) error {
|
||||
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
statusCode, _ := strconv.Atoi(trimmed)
|
||||
if statusCode < 100 || statusCode > 999 {
|
||||
return fmt.Errorf("%s 必须在 100 到 999 之间", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyWorkerProcesses(key, trimmed string) error {
|
||||
if trimmed == "auto" {
|
||||
return nil
|
||||
}
|
||||
return validatePositiveIntegerOption(key, trimmed)
|
||||
}
|
||||
|
||||
func validateOpenRestyGzipCompLevel(key, trimmed string) error {
|
||||
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
level, _ := strconv.Atoi(trimmed)
|
||||
if level > maxOpenRestyGzipCompLevel {
|
||||
return fmt.Errorf("%s 不能大于 %d", key, maxOpenRestyGzipCompLevel)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyEventsUse(key, trimmed string) error {
|
||||
if trimmed == "" {
|
||||
return nil
|
||||
}
|
||||
switch trimmed {
|
||||
case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key)
|
||||
}
|
||||
}
|
||||
|
||||
func validateOpenRestyResolvers(key, trimmed string) error {
|
||||
if trimmed == "" {
|
||||
return nil
|
||||
}
|
||||
if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) {
|
||||
return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyProxyBuffers(key, trimmed string) error {
|
||||
if openRestyProxyBuffersPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key)
|
||||
}
|
||||
|
||||
func validateOpenRestySizeValue(key, trimmed string) error {
|
||||
if openRestySizePattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key)
|
||||
}
|
||||
|
||||
func validateOpenRestyCachePath(key, trimmed string) error {
|
||||
if strings.ContainsAny(trimmed, "\r\n\t") {
|
||||
return fmt.Errorf("%s 不能包含换行或制表符", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyCacheLevels(key, trimmed string) error {
|
||||
if openRestyCacheLevelsPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key)
|
||||
}
|
||||
|
||||
func validateOpenRestyDurationToken(key, trimmed string) error {
|
||||
if openRestyDurationTokenPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key)
|
||||
}
|
||||
|
||||
func validateOpenRestyCacheKeyTemplate(key, trimmed string) error {
|
||||
if trimmed == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
if strings.ContainsAny(trimmed, "\r\n") {
|
||||
return fmt.Errorf("%s 不能包含换行", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyCacheUseStale(key, trimmed string) error {
|
||||
if trimmed == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
allowedTokens := map[string]struct{}{
|
||||
"error": {}, "timeout": {}, "invalid_header": {}, "updating": {},
|
||||
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
|
||||
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
|
||||
}
|
||||
for _, token := range strings.Fields(trimmed) {
|
||||
if _, ok := allowedTokens[token]; !ok {
|
||||
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyMainConfigTemplate(key, value string) error {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -32,7 +32,7 @@ func GetStatusHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// getNoticeHandler 获取系统公告。
|
||||
// GetNoticeHandler 获取系统公告。
|
||||
// @Summary 获取系统公告
|
||||
// @Description 返回 OpenFlare 控制台公告文本,无需登录
|
||||
// @Tags openflare-option
|
||||
@@ -41,7 +41,6 @@ func GetStatusHandler(c *gin.Context) {
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/notice [get]
|
||||
// GetNoticeHandler returns the notice content.
|
||||
func GetNoticeHandler(c *gin.Context) {
|
||||
notice, err := getNotice(c.Request.Context())
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
@@ -50,7 +49,7 @@ func GetNoticeHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(notice))
|
||||
}
|
||||
|
||||
// listOptionsHandler 列出全部配置项。
|
||||
// ListOptionsHandler 列出全部配置项。
|
||||
// @Summary 列出 OpenFlare 配置项
|
||||
// @Description 返回全部非敏感 OpenFlare 配置项,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
@@ -62,7 +61,6 @@ func GetNoticeHandler(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option [get]
|
||||
// ListOptionsHandler lists OpenFlare options.
|
||||
func ListOptionsHandler(c *gin.Context) {
|
||||
options, err := listOptions(c.Request.Context())
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
@@ -71,7 +69,7 @@ func ListOptionsHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(options))
|
||||
}
|
||||
|
||||
// updateOptionHandler 更新单个配置项。
|
||||
// UpdateOptionHandler 更新单个配置项。
|
||||
// @Summary 更新 OpenFlare 配置项
|
||||
// @Description 更新单个 OpenFlare 配置项,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
@@ -85,7 +83,6 @@ func ListOptionsHandler(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/update [post]
|
||||
// UpdateOptionHandler updates a single option.
|
||||
func UpdateOptionHandler(c *gin.Context) {
|
||||
var option model.OpenFlareOption
|
||||
if !apiutil.BindJSON(c, &option) {
|
||||
@@ -97,7 +94,7 @@ func UpdateOptionHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// updateOptionsBatchHandler 批量更新配置项。
|
||||
// UpdateOptionsBatchHandler 批量更新配置项。
|
||||
// @Summary 批量更新 OpenFlare 配置项
|
||||
// @Description 批量更新多个 OpenFlare 配置项,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
@@ -111,7 +108,6 @@ func UpdateOptionHandler(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/update-batch [post]
|
||||
// UpdateOptionsBatchHandler updates options in batch.
|
||||
func UpdateOptionsBatchHandler(c *gin.Context) {
|
||||
var payload optionBatchPayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
@@ -123,7 +119,7 @@ func UpdateOptionsBatchHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// lookupGeoIPHandler 查询 GeoIP 信息。
|
||||
// LookupGeoIPHandler 查询 GeoIP 信息。
|
||||
// @Summary GeoIP 地址查询
|
||||
// @Description 按提供商与 IP 查询地理位置信息,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
@@ -137,7 +133,6 @@ func UpdateOptionsBatchHandler(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/geoip/lookup [post]
|
||||
// LookupGeoIPHandler performs a GeoIP lookup.
|
||||
func LookupGeoIPHandler(c *gin.Context) {
|
||||
var request geoIPLookupRequest
|
||||
if !apiutil.BindJSON(c, &request) {
|
||||
@@ -150,7 +145,7 @@ func LookupGeoIPHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// cleanupDatabaseHandler 清理可观测性数据库数据。
|
||||
// CleanupDatabaseHandler 清理可观测性数据库数据。
|
||||
// @Summary 清理可观测性数据库
|
||||
// @Description 按目标与保留天数清理可观测性相关数据表,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
@@ -164,7 +159,6 @@ func LookupGeoIPHandler(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/database/cleanup [post]
|
||||
// CleanupDatabaseHandler cleans up observability data.
|
||||
func CleanupDatabaseHandler(c *gin.Context) {
|
||||
var input databaseCleanupInput
|
||||
if err := bindOptionalJSON(c.Request.Body, &input); err != nil {
|
||||
@@ -178,7 +172,7 @@ func CleanupDatabaseHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// syncUptimeKumaHandler 同步 Uptime Kuma 监控。
|
||||
// SyncUptimeKumaHandler 同步 Uptime Kuma 监控。
|
||||
// @Summary 同步 Uptime Kuma
|
||||
// @Description 将 OpenFlare 节点同步到 Uptime Kuma,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
@@ -191,7 +185,6 @@ func CleanupDatabaseHandler(c *gin.Context) {
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/uptimekuma/sync [post]
|
||||
// SyncUptimeKumaHandler triggers UptimeKuma sync.
|
||||
func SyncUptimeKumaHandler(c *gin.Context) {
|
||||
if apiutil.AbortBadRequestOnError(c, syncUptimeKuma(c.Request.Context())) {
|
||||
return
|
||||
@@ -204,4 +197,4 @@ func bindOptionalJSON(body io.Reader, target any) error {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
package option
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strconv"
|
||||
@@ -13,6 +14,8 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const maxOpenRestyGzipCompLevel = 9
|
||||
|
||||
var (
|
||||
openRestySizePattern = regexp.MustCompile(`^\d+[kKmMgG]?$`)
|
||||
openRestyProxyBuffersPattern = regexp.MustCompile(`^\d+\s+\d+[kKmMgG]?$`)
|
||||
@@ -20,6 +23,8 @@ var (
|
||||
openRestyDurationTokenPattern = regexp.MustCompile(`^\d+[smhdwSMHDW]$`)
|
||||
)
|
||||
|
||||
const optionValueTrue = "true"
|
||||
|
||||
func buildOptionValidationState(options []model.OpenFlareOption) map[string]string {
|
||||
model.OptionMapRWMutex.RLock()
|
||||
state := make(map[string]string, len(model.OptionMap)+len(options))
|
||||
@@ -37,11 +42,11 @@ func buildOptionValidationState(options []model.OpenFlareOption) map[string]stri
|
||||
func validateOptionWithState(option model.OpenFlareOption, state map[string]string) error {
|
||||
switch option.Key {
|
||||
case "GitHubOAuthEnabled":
|
||||
if option.Value == "true" && strings.TrimSpace(state["GitHubClientId"]) == "" {
|
||||
if option.Value == optionValueTrue && strings.TrimSpace(state["GitHubClientId"]) == "" {
|
||||
return fmt.Errorf("无法启用 GitHub OAuth,请先填入 GitHub Client ID 以及 GitHub Client Secret!")
|
||||
}
|
||||
case "WeChatAuthEnabled":
|
||||
if option.Value == "true" && strings.TrimSpace(state["WeChatServerAddress"]) == "" {
|
||||
if option.Value == optionValueTrue && strings.TrimSpace(state["WeChatServerAddress"]) == "" {
|
||||
return fmt.Errorf("无法启用微信登录,请先填入微信登录相关配置信息!")
|
||||
}
|
||||
}
|
||||
@@ -71,7 +76,7 @@ func validatePositiveIntegerOption(key, value string) error {
|
||||
|
||||
func validateBooleanOption(key, value string) error {
|
||||
switch value {
|
||||
case "true", "false":
|
||||
case optionValueTrue, "false":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("%s 必须为 true 或 false", key)
|
||||
@@ -112,171 +117,81 @@ func validateUptimeKumaOption(key, value string, state map[string]string) error
|
||||
trimmed := strings.TrimSpace(value)
|
||||
switch key {
|
||||
case "UptimeKumaEnabled":
|
||||
if err := validateBooleanOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
if trimmed == "true" {
|
||||
url := strings.TrimSpace(state["UptimeKumaUrl"])
|
||||
username := strings.TrimSpace(state["UptimeKumaUsername"])
|
||||
password := strings.TrimSpace(state["UptimeKumaPassword"])
|
||||
if url == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时地址不能为空")
|
||||
}
|
||||
if username == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
if password == "" && model.UptimeKumaPassword == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时密码不能为空")
|
||||
}
|
||||
}
|
||||
return validateUptimeKumaEnabled(key, trimmed, state)
|
||||
case "UptimeKumaUsername":
|
||||
if trimmed == "" && state["UptimeKumaEnabled"] == "true" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
return validateUptimeKumaUsername(trimmed, state)
|
||||
case "UptimeKumaUrl":
|
||||
if trimmed != "" && !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") {
|
||||
return fmt.Errorf("Uptime Kuma 地址必须以 http:// 或 https:// 开头")
|
||||
}
|
||||
return validateUptimeKumaURL(trimmed)
|
||||
case "UptimeKumaMonitorScope":
|
||||
if trimmed != "all" && trimmed != "selected" {
|
||||
return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)")
|
||||
}
|
||||
return validateUptimeKumaMonitorScope(trimmed)
|
||||
case "UptimeKumaSyncInterval", "UptimeKumaInterval", "UptimeKumaRetryInterval", "UptimeKumaTimeout":
|
||||
return validatePositiveIntegerOption(key, trimmed)
|
||||
case "UptimeKumaRetry":
|
||||
intValue, err := strconv.Atoi(trimmed)
|
||||
if err != nil || intValue < 0 {
|
||||
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
|
||||
}
|
||||
return validateUptimeKumaRetry(key, trimmed)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyOption(key, value string) error {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
func validateUptimeKumaEnabled(key, trimmed string, state map[string]string) error {
|
||||
if err := validateBooleanOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
if trimmed != optionValueTrue {
|
||||
return nil
|
||||
}
|
||||
url := strings.TrimSpace(state["UptimeKumaUrl"])
|
||||
username := strings.TrimSpace(state["UptimeKumaUsername"])
|
||||
password := strings.TrimSpace(state["UptimeKumaPassword"])
|
||||
if url == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时地址不能为空")
|
||||
}
|
||||
if username == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
if password == "" && model.UptimeKumaPassword == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时密码不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
switch key {
|
||||
case "OpenRestyDefaultServerReturnStatus":
|
||||
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
statusCode, _ := strconv.Atoi(trimmed)
|
||||
if statusCode < 100 || statusCode > 999 {
|
||||
return fmt.Errorf("%s 必须在 100 到 999 之间", key)
|
||||
}
|
||||
case "OpenRestyWorkerProcesses":
|
||||
if trimmed == "auto" {
|
||||
return nil
|
||||
}
|
||||
return validatePositiveIntegerOption(key, trimmed)
|
||||
case "OpenRestyWorkerConnections",
|
||||
"OpenRestyWorkerRlimitNofile",
|
||||
"OpenRestyKeepaliveTimeout",
|
||||
"OpenRestyKeepaliveRequests",
|
||||
"OpenRestyClientHeaderTimeout",
|
||||
"OpenRestyClientBodyTimeout",
|
||||
"OpenRestySendTimeout",
|
||||
"OpenRestyProxyConnectTimeout",
|
||||
"OpenRestyProxySendTimeout",
|
||||
"OpenRestyProxyReadTimeout",
|
||||
"OpenRestyGzipMinLength":
|
||||
return validatePositiveIntegerOption(key, trimmed)
|
||||
case "OpenRestyGzipCompLevel":
|
||||
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
level, _ := strconv.Atoi(trimmed)
|
||||
if level > 9 {
|
||||
return fmt.Errorf("%s 不能大于 9", key)
|
||||
}
|
||||
case "OpenRestyEventsUse":
|
||||
if trimmed == "" {
|
||||
return nil
|
||||
}
|
||||
switch trimmed {
|
||||
case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key)
|
||||
}
|
||||
case "OpenRestyResolvers":
|
||||
if trimmed == "" {
|
||||
return nil
|
||||
}
|
||||
if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) {
|
||||
return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key)
|
||||
}
|
||||
case "OpenRestyEventsMultiAcceptEnabled",
|
||||
"OpenRestyWebsocketEnabled",
|
||||
"OpenRestyHTTP3Enabled",
|
||||
"OpenRestyProxyRequestBufferingEnabled",
|
||||
"OpenRestyProxyBufferingEnabled",
|
||||
"OpenRestyGzipEnabled",
|
||||
"OpenRestyCacheEnabled",
|
||||
"OpenRestyCacheLockEnabled":
|
||||
return validateBooleanOption(key, trimmed)
|
||||
case "OpenRestyProxyBuffers", "OpenRestyLargeClientHeaderBuffers":
|
||||
if openRestyProxyBuffersPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key)
|
||||
case "OpenRestyProxyBufferSize", "OpenRestyProxyBusyBuffersSize", "OpenRestyCacheMaxSize", "OpenRestyClientMaxBodySize":
|
||||
if openRestySizePattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key)
|
||||
case "OpenRestyCachePath":
|
||||
if strings.ContainsAny(trimmed, "\r\n\t") {
|
||||
return fmt.Errorf("%s 不能包含换行或制表符", key)
|
||||
}
|
||||
case "OpenRestyCacheLevels":
|
||||
if openRestyCacheLevelsPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key)
|
||||
case "OpenRestyCacheInactive", "OpenRestyCacheLockTimeout":
|
||||
if openRestyDurationTokenPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key)
|
||||
case "OpenRestyCacheKeyTemplate":
|
||||
if trimmed == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
if strings.ContainsAny(trimmed, "\r\n") {
|
||||
return fmt.Errorf("%s 不能包含换行", key)
|
||||
}
|
||||
case "OpenRestyCacheUseStale":
|
||||
if trimmed == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
allowedTokens := map[string]struct{}{
|
||||
"error": {}, "timeout": {}, "invalid_header": {}, "updating": {},
|
||||
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
|
||||
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
|
||||
}
|
||||
for _, token := range strings.Fields(trimmed) {
|
||||
if _, ok := allowedTokens[token]; !ok {
|
||||
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
|
||||
}
|
||||
}
|
||||
case "OpenRestyMainConfigTemplate":
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
func validateUptimeKumaUsername(trimmed string, state map[string]string) error {
|
||||
if trimmed == "" && state["UptimeKumaEnabled"] == optionValueTrue {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUptimeKumaURL(trimmed string) error {
|
||||
if trimmed != "" && !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") {
|
||||
return fmt.Errorf("uptime Kuma 地址必须以 http:// 或 https:// 开头")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUptimeKumaMonitorScope(trimmed string) error {
|
||||
if trimmed != "all" && trimmed != "selected" {
|
||||
return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUptimeKumaRetry(key, trimmed string) error {
|
||||
intValue, err := strconv.Atoi(trimmed)
|
||||
if err != nil || intValue < 0 {
|
||||
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOptions(options []model.OpenFlareOption) error {
|
||||
if len(options) == 0 {
|
||||
return fmt.Errorf(errInvalidParams)
|
||||
return errors.New(errInvalidParams)
|
||||
}
|
||||
|
||||
state := buildOptionValidationState(options)
|
||||
for _, option := range options {
|
||||
if strings.TrimSpace(option.Key) == "" {
|
||||
return fmt.Errorf(errInvalidParams)
|
||||
return errors.New(errInvalidParams)
|
||||
}
|
||||
if err := validateOptionWithState(option, state); err != nil {
|
||||
return err
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package origin defines shared error messages for origin management.
|
||||
package origin
|
||||
|
||||
const (
|
||||
|
||||
@@ -12,6 +12,8 @@ import (
|
||||
"unicode"
|
||||
)
|
||||
|
||||
const maxOriginHostnameLength = 253
|
||||
|
||||
func normalizeOriginAddress(raw string) string {
|
||||
return strings.ToLower(strings.TrimSpace(raw))
|
||||
}
|
||||
@@ -29,7 +31,7 @@ func validateOriginAddress(address string) error {
|
||||
if ip := net.ParseIP(address); ip != nil {
|
||||
return nil
|
||||
}
|
||||
if len(address) > 253 {
|
||||
if len(address) > maxOriginHostnameLength {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
|
||||
@@ -1,27 +1,28 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package pages provides logics and management for OpenFlare static page deployments.
|
||||
package pages
|
||||
|
||||
const (
|
||||
errPagesProjectNotFound = "Pages 项目不存在"
|
||||
errPagesSlugExists = "Pages 项目标识已存在"
|
||||
errPagesNameRequired = "Pages 项目名称不能为空"
|
||||
errPagesSlugInvalid = "Pages 项目标识只能包含小写字母、数字和连字符"
|
||||
errPagesDeleteReferenced = "Pages 项目已被规则引用,不能删除"
|
||||
errPagesDeploymentNotFound = "Pages 部署不存在"
|
||||
errPagesDeploymentMismatch = "Pages 部署不属于该项目"
|
||||
errPagesProjectNotFound = "pages 项目不存在"
|
||||
errPagesSlugExists = "pages 项目标识已存在"
|
||||
errPagesNameRequired = "pages 项目名称不能为空"
|
||||
errPagesSlugInvalid = "pages 项目标识只能包含小写字母、数字和连字符"
|
||||
errPagesDeleteReferenced = "pages 项目已被规则引用,不能删除"
|
||||
errPagesDeploymentNotFound = "pages 部署不存在"
|
||||
errPagesDeploymentMismatch = "pages 部署不属于该项目"
|
||||
errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署"
|
||||
errPagesPackageMissing = "缺少 Pages 部署包"
|
||||
errPagesPackageNotZip = "Pages 部署包必须是 .zip 文件"
|
||||
errPagesPackageInvalidZip = "Pages 部署包不是有效 zip 文件"
|
||||
errPagesPackageEmpty = "Pages 部署包不能为空"
|
||||
errPagesPackageNotZip = "pages 部署包必须是 .zip 文件"
|
||||
errPagesPackageInvalidZip = "pages 部署包不是有效 zip 文件"
|
||||
errPagesPackageEmpty = "pages 部署包不能为空"
|
||||
errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空"
|
||||
errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头"
|
||||
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空"
|
||||
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL"
|
||||
errPagesPackagePathEmpty = "Pages 部署包路径为空"
|
||||
errPagesPackageUploadMissing = "Pages 部署包上传记录不存在"
|
||||
errPagesPackageNotInActiveConfig = "Pages 部署尚未进入激活配置"
|
||||
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errPagesPackagePathEmpty = "pages 部署包路径为空"
|
||||
errPagesPackageUploadMissing = "pages 部署包上传记录不存在"
|
||||
errPagesPackageNotInActiveConfig = "pages 部署尚未进入激活配置"
|
||||
errPagesInvalidSnapshotFormat = "配置快照格式无效"
|
||||
)
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"mime/multipart"
|
||||
"os"
|
||||
"path"
|
||||
@@ -24,11 +25,14 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
pagesMaxDeploymentFiles = 1000
|
||||
pagesMaxDeploymentBytes = 100 * 1024 * 1024
|
||||
defaultPagesEntryFile = "index.html"
|
||||
defaultPagesFallbackPath = "/index.html"
|
||||
pagesMaxDeploymentFiles = 1000
|
||||
pagesMaxDeploymentBytes = 100 * 1024 * 1024
|
||||
defaultPagesEntryFile = "index.html"
|
||||
defaultPagesFallbackPath = "/index.html"
|
||||
pagesDeploymentUploadType = "openflare_pages_deployment"
|
||||
mimeTypeApplicationZip = "application/zip"
|
||||
pagesMaxPathLength = 512
|
||||
bytesPerKiB = 1024
|
||||
)
|
||||
|
||||
var pagesSlugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,126}[a-z0-9]$|^[a-z0-9]$`)
|
||||
@@ -71,15 +75,15 @@ func validateAndNormalizePagesRootDir(raw string) (string, error) {
|
||||
if value == "" {
|
||||
return "", nil
|
||||
}
|
||||
if len(value) > 512 {
|
||||
return "", errors.New("Pages 根目录长度不能超过 512")
|
||||
if len(value) > pagesMaxPathLength {
|
||||
return "", errors.New("pages 根目录长度不能超过 512") // error 消息首字母小写
|
||||
}
|
||||
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
|
||||
return "", errors.New("Pages 根目录包含不支持的字符")
|
||||
return "", errors.New("pages 根目录包含不支持的字符")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r <= 0x20 || r == 0x7f {
|
||||
return "", errors.New("Pages 根目录不能包含空白或控制字符")
|
||||
return "", errors.New("pages 根目录不能包含空白或控制字符")
|
||||
}
|
||||
}
|
||||
cleaned := path.Clean(filepath.ToSlash(value))
|
||||
@@ -88,7 +92,7 @@ func validateAndNormalizePagesRootDir(raw string) (string, error) {
|
||||
}
|
||||
for _, segment := range strings.Split(cleaned, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "", errors.New("Pages 根目录不能包含 . 或 .. 路径段")
|
||||
return "", errors.New("pages 根目录不能包含 . 或 .. 路径段")
|
||||
}
|
||||
}
|
||||
return strings.TrimPrefix(cleaned, "/"), nil
|
||||
@@ -99,34 +103,34 @@ func normalizePagesFallbackPath(raw string) (string, error) {
|
||||
if value == "" {
|
||||
value = defaultPagesFallbackPath
|
||||
}
|
||||
if len(value) > 512 {
|
||||
return "", errors.New("SPA fallback 回退路径长度不能超过 512")
|
||||
if len(value) > pagesMaxPathLength {
|
||||
return "", errors.New("spa fallback 回退路径长度不能超过 512")
|
||||
}
|
||||
if !strings.HasPrefix(value, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须以 / 开头")
|
||||
return "", errors.New("spa fallback 回退路径必须以 / 开头")
|
||||
}
|
||||
if value == "/" || strings.HasSuffix(value, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
|
||||
return "", errors.New("spa fallback 回退路径必须指向具体文件")
|
||||
}
|
||||
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
|
||||
return "", errors.New("SPA fallback 回退路径包含不支持的字符")
|
||||
return "", errors.New("spa fallback 回退路径包含不支持的字符")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r <= 0x20 || r == 0x7f {
|
||||
return "", errors.New("SPA fallback 回退路径不能包含空白或控制字符")
|
||||
return "", errors.New("spa fallback 回退路径不能包含空白或控制字符")
|
||||
}
|
||||
}
|
||||
for _, segment := range strings.Split(value, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "", errors.New("SPA fallback 回退路径不能包含 . 或 .. 路径段")
|
||||
return "", errors.New("spa fallback 回退路径不能包含 . 或 .. 路径段")
|
||||
}
|
||||
}
|
||||
cleaned := path.Clean(value)
|
||||
if cleaned == "." || !strings.HasPrefix(cleaned, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径不合法")
|
||||
return "", errors.New("spa fallback 回退路径不合法")
|
||||
}
|
||||
if cleaned == "/" || strings.HasSuffix(cleaned, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
|
||||
return "", errors.New("spa fallback 回退路径必须指向具体文件")
|
||||
}
|
||||
return cleaned, nil
|
||||
}
|
||||
@@ -152,12 +156,12 @@ func persistPagesUploadTemp(fileHeader *multipart.FileHeader) (string, string, i
|
||||
if err != nil {
|
||||
return "", "", 0, err
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
temp, err := os.CreateTemp("", "openflare-pages-*.zip")
|
||||
if err != nil {
|
||||
return "", "", 0, err
|
||||
}
|
||||
defer temp.Close()
|
||||
defer func() { _ = temp.Close() }()
|
||||
hash := sha256.New()
|
||||
limited := io.LimitReader(file, pagesMaxDeploymentBytes+1)
|
||||
written, err := io.Copy(io.MultiWriter(temp, hash), limited)
|
||||
@@ -167,7 +171,7 @@ func persistPagesUploadTemp(fileHeader *multipart.FileHeader) (string, string, i
|
||||
}
|
||||
if written > pagesMaxDeploymentBytes {
|
||||
_ = os.Remove(temp.Name())
|
||||
return "", "", 0, fmt.Errorf("Pages 部署包不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
|
||||
return "", "", 0, fmt.Errorf("pages 部署包不能超过 %d MiB", pagesMaxDeploymentBytes/bytesPerKiB/bytesPerKiB)
|
||||
}
|
||||
return temp.Name(), hex.EncodeToString(hash.Sum(nil)), written, nil
|
||||
}
|
||||
@@ -180,11 +184,11 @@ func ingestPagesDeploymentPackage(
|
||||
projectSlug string,
|
||||
fileName string,
|
||||
) (upload.IngestResult, error) {
|
||||
file, err := os.Open(tempPath)
|
||||
file, err := os.Open(tempPath) //nolint:gosec // tempPath is a validated pages deployment staging file
|
||||
if err != nil {
|
||||
return upload.IngestResult{}, err
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
systemUser := repository.GetSystemUser(ctx)
|
||||
accessMode := 0
|
||||
@@ -193,7 +197,7 @@ func ingestPagesDeploymentPackage(
|
||||
Reader: file,
|
||||
Size: size,
|
||||
FileName: fileName,
|
||||
MimeType: "application/zip",
|
||||
MimeType: mimeTypeApplicationZip,
|
||||
Extension: "zip",
|
||||
Hash: checksum,
|
||||
Type: pagesDeploymentUploadType,
|
||||
@@ -270,7 +274,7 @@ func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploym
|
||||
if err != nil {
|
||||
return nil, errors.New(errPagesPackageInvalidZip)
|
||||
}
|
||||
defer reader.Close()
|
||||
defer func() { _ = reader.Close() }()
|
||||
|
||||
commonPrefix, err := findCommonRootPrefix(reader.File)
|
||||
if err != nil {
|
||||
@@ -298,18 +302,18 @@ func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploym
|
||||
normalizedPath = strings.TrimPrefix(normalizedPath, commonPrefix)
|
||||
}
|
||||
if item.FileInfo().Mode()&os.ModeSymlink != 0 {
|
||||
return nil, fmt.Errorf("Pages 部署包不支持符号链接: %s", normalizedPath)
|
||||
return nil, fmt.Errorf("pages 部署包不支持符号链接: %s", normalizedPath)
|
||||
}
|
||||
if item.UncompressedSize64 > pagesMaxDeploymentBytes {
|
||||
return nil, fmt.Errorf("Pages 文件过大: %s", normalizedPath)
|
||||
return nil, fmt.Errorf("pages 文件过大: %s", normalizedPath)
|
||||
}
|
||||
manifest.FileCount++
|
||||
if manifest.FileCount > pagesMaxDeploymentFiles {
|
||||
return nil, fmt.Errorf("Pages 部署文件数不能超过 %d", pagesMaxDeploymentFiles)
|
||||
return nil, fmt.Errorf("pages 部署文件数不能超过 %d", pagesMaxDeploymentFiles)
|
||||
}
|
||||
manifest.TotalSize += int64(item.UncompressedSize64)
|
||||
if manifest.TotalSize > pagesMaxDeploymentBytes {
|
||||
return nil, fmt.Errorf("Pages 部署展开后不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
|
||||
return nil, fmt.Errorf("pages 部署展开后不能超过 %d MiB", pagesMaxDeploymentBytes/bytesPerKiB/bytesPerKiB)
|
||||
}
|
||||
checksum, err := checksumZipFile(item)
|
||||
if err != nil {
|
||||
@@ -328,7 +332,7 @@ func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploym
|
||||
return nil, errors.New(errPagesPackageEmpty)
|
||||
}
|
||||
if !entrySeen {
|
||||
return nil, fmt.Errorf("Pages 部署包缺少入口文件 %s", targetEntryPath)
|
||||
return nil, fmt.Errorf("pages 部署包缺少入口文件 %s", targetEntryPath)
|
||||
}
|
||||
return manifest, nil
|
||||
}
|
||||
@@ -342,29 +346,38 @@ func normalizePagesZipPath(raw string) (string, bool, error) {
|
||||
return "", true, nil
|
||||
}
|
||||
if strings.HasPrefix(name, "/") || path.IsAbs(name) {
|
||||
return "", false, fmt.Errorf("Pages 部署包不能包含绝对路径: %s", raw)
|
||||
return "", false, fmt.Errorf("pages 部署包不能包含绝对路径: %s", raw)
|
||||
}
|
||||
cleaned := path.Clean(name)
|
||||
if cleaned == "." {
|
||||
return "", true, nil
|
||||
}
|
||||
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || strings.Contains(cleaned, "/../") {
|
||||
return "", false, fmt.Errorf("Pages 部署包路径不能逃逸目录: %s", raw)
|
||||
return "", false, fmt.Errorf("pages 部署包路径不能逃逸目录: %s", raw)
|
||||
}
|
||||
return cleaned, false, nil
|
||||
}
|
||||
|
||||
func pagesZipEntryCopyLimit(size uint64) (int64, error) {
|
||||
if size == 0 || size > pagesMaxDeploymentBytes || size > uint64(math.MaxInt64) {
|
||||
return 0, errors.New("pages file size out of bounds")
|
||||
}
|
||||
return int64(size), nil //nolint:gosec // size is bounded to math.MaxInt64 above
|
||||
}
|
||||
|
||||
func checksumZipFile(item *zip.File) (string, error) {
|
||||
file, err := item.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer file.Close()
|
||||
defer func() { _ = file.Close() }()
|
||||
hash := sha256.New()
|
||||
if _, err = io.Copy(hash, file); err != nil {
|
||||
limit, err := pagesZipEntryCopyLimit(item.UncompressedSize64)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, err = io.CopyN(hash, file, limit); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(hash.Sum(nil)), nil
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -256,7 +256,7 @@ func UploadDeployment(ctx context.Context, projectID uint, fileHeader *multipart
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer os.Remove(tempPath)
|
||||
defer func() { _ = os.Remove(tempPath) }()
|
||||
manifest, err := inspectPagesZip(tempPath, rootDir, entryFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -370,10 +370,10 @@ func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (*storage.Obj
|
||||
}
|
||||
obj, err := uploadstorage.OpenStoredObject(ctx, &uploadRecord)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
|
||||
return nil, "", fmt.Errorf("pages 部署包不存在: %w", err)
|
||||
}
|
||||
if obj.ContentType == "" {
|
||||
obj.ContentType = "application/zip"
|
||||
obj.ContentType = mimeTypeApplicationZip
|
||||
}
|
||||
return obj, fileName, nil
|
||||
}
|
||||
@@ -382,17 +382,17 @@ func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (*storage.Obj
|
||||
}
|
||||
file, err := os.Open(deployment.ArtifactPath)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
|
||||
return nil, "", fmt.Errorf("pages 部署包不存在: %w", err)
|
||||
}
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
|
||||
return nil, "", fmt.Errorf("pages 部署包不存在: %w", err)
|
||||
}
|
||||
return &storage.Object{
|
||||
Body: file,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "application/zip",
|
||||
ContentType: mimeTypeApplicationZip,
|
||||
}, fileName, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,178 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package proxy_route provides helpers for building proxy route configurations.
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
type proxyRouteJSONFields struct {
|
||||
cacheRulesJSON string
|
||||
upstreamsJSON string
|
||||
customHeadersJSON string
|
||||
certIDsJSON string
|
||||
domainCertIDsJSON string
|
||||
domainsJSON string
|
||||
}
|
||||
|
||||
func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input Input) (string, *uint, []string, error) {
|
||||
switch upstreamType {
|
||||
case proxyRouteUpstreamTypeTunnel, proxyRouteUpstreamTypePages:
|
||||
if upstreamType == proxyRouteUpstreamTypePages {
|
||||
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
|
||||
return "", nil, nil, err
|
||||
}
|
||||
}
|
||||
originURL := "http://127.0.0.1"
|
||||
return originURL, nil, []string{originURL}, nil
|
||||
default:
|
||||
originURL, originID, err := resolveProxyRoutePrimaryOrigin(ctx, input)
|
||||
if err != nil {
|
||||
return "", nil, nil, err
|
||||
}
|
||||
upstreams, err := normalizeUpstreams(originURL, input.Upstreams)
|
||||
if err != nil {
|
||||
return "", nil, nil, err
|
||||
}
|
||||
return originURL, originID, upstreams, nil
|
||||
}
|
||||
}
|
||||
|
||||
func marshalProxyRouteJSONFields(
|
||||
domains []string,
|
||||
upstreams []string,
|
||||
cacheRules []string,
|
||||
customHeaders []CustomHeaderInput,
|
||||
certIDs []uint,
|
||||
domainCertIDs []uint,
|
||||
) (*proxyRouteJSONFields, error) {
|
||||
cacheRulesJSON, err := json.Marshal(cacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreamsJSON, err := json.Marshal(upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeadersJSON, err := json.Marshal(customHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDsJSON, err := json.Marshal(certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainCertIDsJSON, err := json.Marshal(domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainsJSON, err := json.Marshal(domains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &proxyRouteJSONFields{
|
||||
cacheRulesJSON: string(cacheRulesJSON),
|
||||
upstreamsJSON: string(upstreamsJSON),
|
||||
customHeadersJSON: string(customHeadersJSON),
|
||||
certIDsJSON: string(certIDsJSON),
|
||||
domainCertIDsJSON: string(domainCertIDsJSON),
|
||||
domainsJSON: string(domainsJSON),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteHTTPSInput(input *Input) {
|
||||
if input.EnableHTTPS {
|
||||
return
|
||||
}
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
input.DomainCertIDs = nil
|
||||
}
|
||||
|
||||
func normalizeProxyRouteBasicAuth(input *Input) error {
|
||||
if !input.BasicAuthEnabled {
|
||||
input.BasicAuthUsername = ""
|
||||
input.BasicAuthPassword = ""
|
||||
return nil
|
||||
}
|
||||
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
|
||||
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
|
||||
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
|
||||
return errors.New(errProxyRouteBasicAuth)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func populateProxyRouteFields(
|
||||
route *model.ProxyRoute,
|
||||
input Input,
|
||||
siteName, domain string,
|
||||
jsonFields *proxyRouteJSONFields,
|
||||
originID *uint,
|
||||
upstreams []string,
|
||||
originHost, remark, cachePolicy string,
|
||||
limitConnPerServer, limitConnPerIP int,
|
||||
limitRate, upstreamType string,
|
||||
) {
|
||||
route.SiteName = siteName
|
||||
route.Domain = domain
|
||||
route.Domains = jsonFields.domainsJSON
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = jsonFields.upstreamsJSON
|
||||
route.Enabled = input.Enabled
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
route.CertIDs = jsonFields.certIDsJSON
|
||||
route.DomainCertIDs = jsonFields.domainCertIDsJSON
|
||||
route.RedirectHTTP = input.RedirectHTTP
|
||||
route.LimitConnPerServer = limitConnPerServer
|
||||
route.LimitConnPerIP = limitConnPerIP
|
||||
route.LimitRate = limitRate
|
||||
route.CacheEnabled = input.CacheEnabled
|
||||
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
|
||||
route.CacheRules = jsonFields.cacheRulesJSON
|
||||
route.CustomHeaders = jsonFields.customHeadersJSON
|
||||
route.BasicAuthEnabled = input.BasicAuthEnabled
|
||||
route.BasicAuthUsername = input.BasicAuthUsername
|
||||
route.BasicAuthPassword = input.BasicAuthPassword
|
||||
route.Remark = remark
|
||||
route.UpstreamType = upstreamType
|
||||
}
|
||||
|
||||
func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, upstreamType string, input Input) error {
|
||||
switch upstreamType {
|
||||
case proxyRouteUpstreamTypeTunnel:
|
||||
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
|
||||
return err
|
||||
}
|
||||
route.TunnelNodeID = tunnelNodeID
|
||||
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
|
||||
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
|
||||
route.PagesProjectID = nil
|
||||
case proxyRouteUpstreamTypePages:
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = input.PagesProjectID
|
||||
default:
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
func normalizeExplicitDomainCertIDs(ctx context.Context, domains []string, rawDomainCertIDs []uint) ([]uint, []uint, *uint, error) {
|
||||
if len(rawDomainCertIDs) != len(domains) {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
|
||||
}
|
||||
|
||||
normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs))
|
||||
uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs))
|
||||
seen := make(map[uint]struct{}, len(rawDomainCertIDs))
|
||||
hasAssignedCertificate := false
|
||||
for index, item := range rawDomainCertIDs {
|
||||
if item == 0 {
|
||||
continue
|
||||
}
|
||||
if _, err := lookupTLSCertificateByID(ctx, item); err != nil {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
|
||||
}
|
||||
normalizedDomainCertIDs[index] = item
|
||||
hasAssignedCertificate = true
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
uniqueCertIDs = append(uniqueCertIDs, item)
|
||||
}
|
||||
if !hasAssignedCertificate {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
}
|
||||
|
||||
primaryCertID := &uniqueCertIDs[0]
|
||||
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
|
||||
}
|
||||
|
||||
func normalizeDerivedDomainCertIDs(
|
||||
ctx context.Context,
|
||||
domains []string,
|
||||
normalizedCertIDs []uint,
|
||||
) ([]uint, []uint, *uint, error) {
|
||||
switch {
|
||||
case len(normalizedCertIDs) == 0:
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
case len(normalizedCertIDs) == 1:
|
||||
domainCertIDs := make([]uint, len(domains))
|
||||
for index := range domainCertIDs {
|
||||
domainCertIDs[index] = normalizedCertIDs[0]
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
case len(normalizedCertIDs) == len(domains):
|
||||
domainCertIDs := make([]uint, len(normalizedCertIDs))
|
||||
copy(domainCertIDs, normalizedCertIDs)
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
default:
|
||||
domainCertIDs, err := deriveDomainCertIDsFromCertificateSet(ctx, domains, normalizedCertIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
}
|
||||
}
|
||||
@@ -42,9 +42,9 @@ const (
|
||||
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
|
||||
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
|
||||
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
|
||||
errProxyRoutePagesNotFound = "Pages 项目不存在"
|
||||
errProxyRoutePagesDisabled = "Pages 项目未启用"
|
||||
errProxyRoutePagesNoDeploy = "Pages 项目没有激活部署"
|
||||
errProxyRoutePagesNotFound = "pages 项目不存在"
|
||||
errProxyRoutePagesDisabled = "pages 项目未启用"
|
||||
errProxyRoutePagesNoDeploy = "pages 项目没有激活部署"
|
||||
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
|
||||
errProxyRouteOriginPort = "端口格式不合法"
|
||||
errProxyRouteOriginPortEmpty = "端口不能为空"
|
||||
|
||||
@@ -30,6 +30,13 @@ const (
|
||||
proxyRouteCachePolicySuffix = "suffix"
|
||||
proxyRouteCachePolicyPathPrefix = "path_prefix"
|
||||
proxyRouteCachePolicyPathExact = "path_exact"
|
||||
proxyRouteSchemeHTTP = "http"
|
||||
proxyRouteSchemeHTTPS = "https"
|
||||
proxyRouteUpstreamTypeTunnel = "tunnel"
|
||||
proxyRouteUpstreamTypePages = "pages"
|
||||
|
||||
maxOriginHostnameLength = 253
|
||||
originURIPathQueryParts = 2
|
||||
)
|
||||
|
||||
type tlsCertificateRow struct {
|
||||
@@ -100,7 +107,7 @@ func validateOriginAddress(address string) error {
|
||||
if ip := net.ParseIP(address); ip != nil {
|
||||
return nil
|
||||
}
|
||||
if len(address) > 253 {
|
||||
if len(address) > maxOriginHostnameLength {
|
||||
return errors.New(errProxyRouteOriginInvalid)
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
@@ -136,7 +143,7 @@ func normalizeOriginPort(raw string) (string, error) {
|
||||
func normalizeOriginScheme(raw string) (string, error) {
|
||||
scheme := strings.ToLower(strings.TrimSpace(raw))
|
||||
switch scheme {
|
||||
case "http", "https":
|
||||
case proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
|
||||
return scheme, nil
|
||||
default:
|
||||
return "", errors.New(errProxyRouteOriginSchemeOnly)
|
||||
@@ -187,7 +194,7 @@ func buildOriginURLFromParts(scheme, address, port, uri string) (string, error)
|
||||
if strings.HasPrefix(normalizedURI, "?") {
|
||||
parsed.RawQuery = strings.TrimPrefix(normalizedURI, "?")
|
||||
} else {
|
||||
pathQuery := strings.SplitN(normalizedURI, "?", 2)
|
||||
pathQuery := strings.SplitN(normalizedURI, "?", originURIPathQueryParts)
|
||||
parsed.Path = pathQuery[0]
|
||||
if len(pathQuery) > 1 {
|
||||
parsed.RawQuery = pathQuery[1]
|
||||
@@ -465,65 +472,14 @@ func normalizeProxyRouteDomainCertificateIDs(
|
||||
}
|
||||
|
||||
if len(rawDomainCertIDs) > 0 {
|
||||
if len(rawDomainCertIDs) != len(domains) {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
|
||||
}
|
||||
|
||||
normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs))
|
||||
uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs))
|
||||
seen := make(map[uint]struct{}, len(rawDomainCertIDs))
|
||||
hasAssignedCertificate := false
|
||||
for index, item := range rawDomainCertIDs {
|
||||
if item == 0 {
|
||||
continue
|
||||
}
|
||||
if _, err := lookupTLSCertificateByID(ctx, item); err != nil {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
|
||||
}
|
||||
normalizedDomainCertIDs[index] = item
|
||||
hasAssignedCertificate = true
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
uniqueCertIDs = append(uniqueCertIDs, item)
|
||||
}
|
||||
if !hasAssignedCertificate {
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
}
|
||||
|
||||
primaryCertID := &uniqueCertIDs[0]
|
||||
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
|
||||
return normalizeExplicitDomainCertIDs(ctx, domains, rawDomainCertIDs)
|
||||
}
|
||||
|
||||
normalizedCertIDs, err := normalizeProxyRouteCertificateIDs(ctx, enableHTTPS, certID, certIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
switch {
|
||||
case len(normalizedCertIDs) == 0:
|
||||
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
|
||||
case len(normalizedCertIDs) == 1:
|
||||
domainCertIDs := make([]uint, len(domains))
|
||||
for index := range domainCertIDs {
|
||||
domainCertIDs[index] = normalizedCertIDs[0]
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
case len(normalizedCertIDs) == len(domains):
|
||||
domainCertIDs := make([]uint, len(normalizedCertIDs))
|
||||
copy(domainCertIDs, normalizedCertIDs)
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
default:
|
||||
domainCertIDs, err := deriveDomainCertIDsFromCertificateSet(ctx, domains, normalizedCertIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
primaryCertID := &normalizedCertIDs[0]
|
||||
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
|
||||
}
|
||||
return normalizeDerivedDomainCertIDs(ctx, domains, normalizedCertIDs)
|
||||
}
|
||||
|
||||
func validateProxyRouteDomainCertificateCoverage(ctx context.Context, domains []string, domainCertIDs []uint) error {
|
||||
@@ -637,65 +593,6 @@ func hasStructuredOriginInput(input Input) bool {
|
||||
strings.TrimSpace(input.OriginURI) != ""
|
||||
}
|
||||
|
||||
func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) {
|
||||
if hasStructuredOriginInput(input) {
|
||||
scheme, err := normalizeOriginScheme(input.OriginScheme)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
port, err := normalizeOriginPort(input.OriginPort)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
uri, err := normalizeOriginURI(input.OriginURI)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
if input.OriginID != nil && *input.OriginID != 0 {
|
||||
origin, err := model.GetOriginByID(ctx, *input.OriginID)
|
||||
if err != nil {
|
||||
return "", nil, errors.New(errProxyRouteOriginNotFound)
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
address := normalizeOriginAddress(input.OriginAddress)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, err := getOrCreateOriginByAddress(ctx, address)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
originURL := strings.TrimSpace(input.OriginURL)
|
||||
if originURL == "" {
|
||||
return "", nil, errors.New(errProxyRouteOriginEmpty)
|
||||
}
|
||||
address, err := extractOriginAddress(originURL)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, findErr := model.GetOriginByAddress(ctx, address)
|
||||
if findErr == nil {
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
|
||||
return "", nil, findErr
|
||||
}
|
||||
return originURL, nil, nil
|
||||
}
|
||||
|
||||
func normalizeCustomHeaders(headers []CustomHeaderInput) ([]CustomHeaderInput, error) {
|
||||
if len(headers) == 0 {
|
||||
return []CustomHeaderInput{}, nil
|
||||
@@ -942,7 +839,7 @@ func validateOriginURL(raw string) error {
|
||||
if err != nil {
|
||||
return errors.New(errProxyRouteOriginInvalid)
|
||||
}
|
||||
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
||||
if parsed.Scheme != proxyRouteSchemeHTTP && parsed.Scheme != proxyRouteSchemeHTTPS {
|
||||
return errors.New(errProxyRouteOriginScheme)
|
||||
}
|
||||
if parsed.Host == "" {
|
||||
@@ -996,7 +893,7 @@ func validateTunnelRouteInput(ctx context.Context, tunnelNodeID *uint, targetAdd
|
||||
return errors.New(errProxyRouteTunnelAddrReq)
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(targetProtocol)) {
|
||||
case "", "http", "https":
|
||||
case "", proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
|
||||
return nil
|
||||
default:
|
||||
return errors.New(errProxyRouteTunnelProtocol)
|
||||
@@ -1025,10 +922,10 @@ func validatePagesRouteInput(ctx context.Context, projectID *uint) error {
|
||||
|
||||
func normalizeUpstreamType(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "tunnel":
|
||||
return "tunnel"
|
||||
case "pages":
|
||||
return "pages"
|
||||
case proxyRouteUpstreamTypeTunnel:
|
||||
return proxyRouteUpstreamTypeTunnel
|
||||
case proxyRouteUpstreamTypePages:
|
||||
return proxyRouteUpstreamTypePages
|
||||
default:
|
||||
return "direct"
|
||||
}
|
||||
@@ -1036,9 +933,9 @@ func normalizeUpstreamType(raw string) string {
|
||||
|
||||
func normalizeTunnelTargetProtocol(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "https":
|
||||
return "https"
|
||||
case proxyRouteSchemeHTTPS:
|
||||
return proxyRouteSchemeHTTPS
|
||||
default:
|
||||
return "http"
|
||||
return proxyRouteSchemeHTTP
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -168,28 +167,9 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
siteName := normalizeProxyRouteSiteNameInput(route, input.SiteName, domain)
|
||||
|
||||
upstreamType := normalizeUpstreamType(input.UpstreamType)
|
||||
var originURL string
|
||||
var originID *uint
|
||||
var upstreams []string
|
||||
|
||||
if upstreamType == "tunnel" {
|
||||
originURL = "http://127.0.0.1"
|
||||
upstreams = []string{originURL}
|
||||
} else if upstreamType == "pages" {
|
||||
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
originURL = "http://127.0.0.1"
|
||||
upstreams = []string{originURL}
|
||||
} else {
|
||||
originURL, originID, err = resolveProxyRoutePrimaryOrigin(ctx, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreams, err = normalizeUpstreams(originURL, input.Upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
originHost := strings.TrimSpace(input.OriginHost)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
@@ -215,25 +195,7 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cacheRulesJSON, err := json.Marshal(cacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreamsJSON, err := json.Marshal(upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeadersJSON, err := json.Marshal(customHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !input.EnableHTTPS {
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
input.DomainCertIDs = nil
|
||||
}
|
||||
normalizeProxyRouteHTTPSInput(&input)
|
||||
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs(
|
||||
ctx,
|
||||
domains,
|
||||
@@ -248,15 +210,7 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
if err := validateProxyRouteDomainCertificateCoverage(ctx, domains, domainCertIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDsJSON, err := json.Marshal(certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainCertIDsJSON, err := json.Marshal(domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainsJSON, err := json.Marshal(domains)
|
||||
jsonFields, err := marshalProxyRouteJSONFields(domains, upstreams, cacheRules, customHeaders, certIDs, domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -277,67 +231,31 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
return nil, errors.New(errProxyRouteRedirectHTTP)
|
||||
}
|
||||
|
||||
if input.BasicAuthEnabled {
|
||||
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
|
||||
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
|
||||
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
|
||||
return nil, errors.New(errProxyRouteBasicAuth)
|
||||
}
|
||||
} else {
|
||||
input.BasicAuthUsername = ""
|
||||
input.BasicAuthPassword = ""
|
||||
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if route == nil {
|
||||
route = &model.ProxyRoute{}
|
||||
}
|
||||
route.SiteName = siteName
|
||||
route.Domain = domain
|
||||
route.Domains = string(domainsJSON)
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = string(upstreamsJSON)
|
||||
route.Enabled = input.Enabled
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
route.CertIDs = string(certIDsJSON)
|
||||
route.DomainCertIDs = string(domainCertIDsJSON)
|
||||
route.RedirectHTTP = input.RedirectHTTP
|
||||
route.LimitConnPerServer = limitConnPerServer
|
||||
route.LimitConnPerIP = limitConnPerIP
|
||||
route.LimitRate = limitRate
|
||||
route.CacheEnabled = input.CacheEnabled
|
||||
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
|
||||
route.CacheRules = string(cacheRulesJSON)
|
||||
route.CustomHeaders = string(customHeadersJSON)
|
||||
route.BasicAuthEnabled = input.BasicAuthEnabled
|
||||
route.BasicAuthUsername = input.BasicAuthUsername
|
||||
route.BasicAuthPassword = input.BasicAuthPassword
|
||||
route.Remark = remark
|
||||
route.UpstreamType = upstreamType
|
||||
if upstreamType == "tunnel" {
|
||||
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
route.TunnelNodeID = tunnelNodeID
|
||||
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
|
||||
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
|
||||
route.PagesProjectID = nil
|
||||
} else if upstreamType == "pages" {
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = input.PagesProjectID
|
||||
} else {
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = nil
|
||||
populateProxyRouteFields(
|
||||
route,
|
||||
input,
|
||||
siteName,
|
||||
domain,
|
||||
jsonFields,
|
||||
originID,
|
||||
upstreams,
|
||||
originHost,
|
||||
remark,
|
||||
cachePolicy,
|
||||
limitConnPerServer,
|
||||
limitConnPerIP,
|
||||
limitRate,
|
||||
upstreamType,
|
||||
)
|
||||
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return route, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func resolveStructuredOriginInput(ctx context.Context, input Input) (string, *uint, error) {
|
||||
scheme, err := normalizeOriginScheme(input.OriginScheme)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
port, err := normalizeOriginPort(input.OriginPort)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
uri, err := normalizeOriginURI(input.OriginURI)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
if input.OriginID != nil && *input.OriginID != 0 {
|
||||
return resolveOriginByID(ctx, scheme, port, uri, *input.OriginID)
|
||||
}
|
||||
return resolveOriginByAddress(ctx, scheme, port, uri, input.OriginAddress)
|
||||
}
|
||||
|
||||
func resolveOriginByID(ctx context.Context, scheme, port, uri string, originID uint) (string, *uint, error) {
|
||||
origin, err := model.GetOriginByID(ctx, originID)
|
||||
if err != nil {
|
||||
return "", nil, errors.New(errProxyRouteOriginNotFound)
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
func resolveOriginByAddress(ctx context.Context, scheme, port, uri, rawAddress string) (string, *uint, error) {
|
||||
address := normalizeOriginAddress(rawAddress)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
originURL, err := buildOriginURLFromParts(scheme, address, port, uri)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, err := getOrCreateOriginByAddress(ctx, address)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
|
||||
func resolveLegacyOriginInput(ctx context.Context, originURL string) (string, *uint, error) {
|
||||
if originURL == "" {
|
||||
return "", nil, errors.New(errProxyRouteOriginEmpty)
|
||||
}
|
||||
address, err := extractOriginAddress(originURL)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
origin, findErr := model.GetOriginByAddress(ctx, address)
|
||||
if findErr == nil {
|
||||
return originURL, &origin.ID, nil
|
||||
}
|
||||
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
|
||||
return "", nil, findErr
|
||||
}
|
||||
return originURL, nil, nil
|
||||
}
|
||||
|
||||
func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) {
|
||||
if hasStructuredOriginInput(input) {
|
||||
return resolveStructuredOriginInput(ctx, input)
|
||||
}
|
||||
return resolveLegacyOriginInput(ctx, strings.TrimSpace(input.OriginURL))
|
||||
}
|
||||
@@ -1,9 +1,11 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package relay provides relay node management and authentication for the OpenFlare platform.
|
||||
package relay
|
||||
|
||||
const (
|
||||
//nolint:gosec // error message text, not a credential
|
||||
errAgentTokenInvalid = "无权进行此操作,Agent Token 无效"
|
||||
errRelayNodeTypeMismatch = "此节点不是 TunnelRelay 类型"
|
||||
)
|
||||
|
||||
@@ -10,12 +10,17 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
relayStatusUnhealthy = "unhealthy"
|
||||
releaseChannelStable = "stable"
|
||||
)
|
||||
|
||||
func normalizeRelayStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case "healthy":
|
||||
return "healthy"
|
||||
case "unhealthy":
|
||||
return "unhealthy"
|
||||
case relayStatusUnhealthy:
|
||||
return relayStatusUnhealthy
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
@@ -25,7 +30,7 @@ func normalizeReleaseChannel(channel string) string {
|
||||
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
|
||||
return "preview"
|
||||
}
|
||||
return "stable"
|
||||
return releaseChannelStable
|
||||
}
|
||||
|
||||
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
|
||||
@@ -98,7 +103,7 @@ func BuildSettings(node *model.OpenFlareNode, updateNow bool, updateChannel, upd
|
||||
autoUpdate = node.AutoUpdateEnabled
|
||||
}
|
||||
if strings.TrimSpace(updateChannel) == "" {
|
||||
updateChannel = "stable"
|
||||
updateChannel = releaseChannelStable
|
||||
}
|
||||
return &Settings{
|
||||
HeartbeatInterval: model.AgentHeartbeatInterval,
|
||||
|
||||
@@ -41,7 +41,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
|
||||
"last_seen_at": now,
|
||||
"status": nodeStatusOnline,
|
||||
"update_requested": false,
|
||||
"update_channel": "stable",
|
||||
"update_channel": releaseChannelStable,
|
||||
"update_tag": "",
|
||||
}
|
||||
if payload.Name != "" && strings.TrimSpace(node.Name) == "" {
|
||||
@@ -55,7 +55,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
|
||||
if !previous.UpdateRequested {
|
||||
delete(changes, "update_requested")
|
||||
}
|
||||
if previous.UpdateChannel == "stable" {
|
||||
if previous.UpdateChannel == releaseChannelStable {
|
||||
delete(changes, "update_channel")
|
||||
}
|
||||
if previous.UpdateTag == "" {
|
||||
@@ -66,7 +66,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
|
||||
node.ExtVersion = payload.ExtVersion
|
||||
node.RelayStatus = payload.RelayStatus
|
||||
node.UpdateRequested = false
|
||||
node.UpdateChannel = "stable"
|
||||
node.UpdateChannel = releaseChannelStable
|
||||
node.UpdateTag = ""
|
||||
lastSeen := now
|
||||
node.LastSeenAt = &lastSeen
|
||||
|
||||
@@ -16,8 +16,8 @@ import (
|
||||
|
||||
const ctxRelayNodeKey = "relay_node"
|
||||
|
||||
// RelayAuth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type.
|
||||
func RelayAuth() gin.HandlerFunc {
|
||||
// Auth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type.
|
||||
func Auth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader("X-Agent-Token"))
|
||||
node, err := authenticateAccessToken(c.Request.Context(), token)
|
||||
@@ -46,4 +46,4 @@ func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlar
|
||||
return nil, err
|
||||
}
|
||||
return node, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -55,7 +55,7 @@ func TestRelayAuthMissingToken(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) {
|
||||
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
@@ -74,7 +74,7 @@ func TestRelayAuthRejectsWrongNodeType(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) {
|
||||
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
@@ -93,7 +93,7 @@ func TestRelayAuthAcceptsTunnelRelay(t *testing.T) {
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) {
|
||||
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
|
||||
authNode, ok := c.Get(ctxRelayNodeKey)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID)
|
||||
|
||||
@@ -24,7 +24,7 @@ func reconcileRelayHealthEvents(ctx context.Context, nodeID string, relayStatus
|
||||
relayFrpsUnhealthyEventType: {},
|
||||
}
|
||||
events := []agent.NodeHealthEvent{}
|
||||
if relayStatus == "unhealthy" {
|
||||
if relayStatus == relayStatusUnhealthy {
|
||||
events = append(events, agent.NodeHealthEvent{
|
||||
EventType: relayFrpsUnhealthyEventType,
|
||||
Severity: "critical",
|
||||
|
||||
@@ -5,8 +5,17 @@ package relay
|
||||
|
||||
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
|
||||
// ProxyStat is an alias for protocol.RelayProxyStat.
|
||||
type ProxyStat = pkgprotocol.RelayProxyStat
|
||||
|
||||
// HeartbeatPayload is an alias for protocol.RelayHeartbeatPayload.
|
||||
type HeartbeatPayload = pkgprotocol.RelayHeartbeatPayload
|
||||
|
||||
// Config is an alias for protocol.RelayConfig.
|
||||
type Config = pkgprotocol.RelayConfig
|
||||
|
||||
// Settings is an alias for protocol.RelaySettings.
|
||||
type Settings = pkgprotocol.RelaySettings
|
||||
type HeartbeatResponse = pkgprotocol.RelayHeartbeatResponse
|
||||
|
||||
// HeartbeatResponse is an alias for protocol.RelayHeartbeatResponse.
|
||||
type HeartbeatResponse = pkgprotocol.RelayHeartbeatResponse
|
||||
|
||||
@@ -90,7 +90,7 @@ func CleanupDatabaseObservability(ctx context.Context, input DatabaseCleanupInpu
|
||||
}
|
||||
|
||||
// RunDatabaseAutoCleanupOnce runs retention-based cleanup for all observability targets.
|
||||
func RunDatabaseAutoCleanupOnce(now time.Time) (*DatabaseAutoCleanupSummary, error) {
|
||||
func RunDatabaseAutoCleanupOnce(ctx context.Context, now time.Time) (*DatabaseAutoCleanupSummary, error) {
|
||||
if !model.DatabaseAutoCleanupEnabled {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -99,7 +99,6 @@ func RunDatabaseAutoCleanupOnce(now time.Time) (*DatabaseAutoCleanupSummary, err
|
||||
}
|
||||
|
||||
retentionDays := model.DatabaseAutoCleanupRetentionDays
|
||||
ctx := context.Background()
|
||||
results := make([]DatabaseCleanupResult, 0, len(databaseCleanupTargets))
|
||||
for _, target := range []string{
|
||||
DatabaseCleanupTargetAccessLogs,
|
||||
|
||||
@@ -134,7 +134,7 @@ func TestRunDatabaseAutoCleanupOnceDeletesAllObservabilityTargets(t *testing.T)
|
||||
model.DatabaseAutoCleanupRetentionDays = previousRetentionDays
|
||||
})
|
||||
|
||||
summary, err := RunDatabaseAutoCleanupOnce(now)
|
||||
summary, err := RunDatabaseAutoCleanupOnce(ctx, now)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, summary)
|
||||
require.Len(t, summary.Results, 3)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user