This commit is contained in:
ryan
2026-06-19 15:13:24 +08:00
parent 0b34792709
commit 32861c5db9
376 changed files with 3648 additions and 19957 deletions
+64 -37
View File
@@ -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)
}
}
}
+2 -2
View File
@@ -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")
}
+68 -65
View File
@@ -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 -1
View File
@@ -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
View File
@@ -1,3 +1,4 @@
package config
// Version is the current agent version string, overridden at build time.
var Version = "dev"
+4
View File
@@ -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"
+13 -3
View File
@@ -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
}
+18 -11
View File
@@ -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)
}
}
+17 -4
View File
@@ -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)
}
+13 -3
View File
@@ -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)
}
}
+3 -1
View File
@@ -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})
}
}
+98 -52
View File
@@ -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
View File
@@ -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},
+2
View File
@@ -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")
+1
View File
@@ -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)
}
}
+12 -4
View File
@@ -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{},
}
}
+65 -9
View File
@@ -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)
+15 -3
View File
@@ -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
}
+41 -25
View File
@@ -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 {
+36 -154
View File
@@ -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()
+205
View File
@@ -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
}
+7 -2
View File
@@ -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",
})
}
}
+12 -1
View File
@@ -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)
}
}
+7 -1
View File
@@ -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())
}
}
+4 -1
View File
@@ -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)
}
}
}
+14 -5
View File
@@ -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)
}
}
+5 -1
View File
@@ -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
}
}
}
+23 -5
View File
@@ -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
}
}
+32 -13
View File
@@ -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))
}
}
+7 -1
View File
@@ -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):
}
}
}
+3 -2
View File
@@ -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
}
}
+28 -18
View File
@@ -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)
}
}
+27 -1
View File
@@ -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)
}
}
+18 -5
View File
@@ -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
View File
@@ -1,3 +1,4 @@
package config
// Version is the flared daemon build version string.
var Version = "dev"
+9 -4
View File
@@ -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})
}
}
+36 -27
View File
@@ -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)
}
+7 -3
View File
@@ -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,
}
}
}
+9 -1
View File
@@ -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)
}
}
+5
View File
@@ -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()
+6 -1
View File
@@ -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",
})
}
}
+11 -1
View File
@@ -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
}
+4 -6
View File
@@ -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,
}
}
+3 -3
View File
@@ -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 = "节点不存在"
+13 -11
View File
@@ -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
}
+8 -8
View File
@@ -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
+9 -9
View File
@@ -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 (
+3 -5
View File
@@ -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 {
+47 -37
View File
@@ -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{
+2 -1
View File
@@ -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 类型"
)
+10 -5
View File
@@ -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"
+10 -9
View File
@@ -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
+1 -1
View File
@@ -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
View File
@@ -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 (
+52 -45
View File
@@ -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
}
+8 -3
View File
@@ -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
View File
@@ -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 (
+1 -1
View File
@@ -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
}
+8 -15
View File
@@ -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
}
}
+61 -146
View File
@@ -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
View File
@@ -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 (
+3 -1
View File
@@ -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, ".")
+16 -15
View File
@@ -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 = "配置快照格式无效"
)
+49 -36
View File
@@ -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
}
+6 -6
View File
@@ -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
}
}
+3 -3
View File
@@ -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 = "端口不能为空"
+21 -124
View File
@@ -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
}
}
+25 -107
View File
@@ -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))
}
+2
View File
@@ -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 类型"
)
+9 -4
View File
@@ -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,
+3 -3
View File
@@ -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
+3 -3
View File
@@ -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