refactor(repo): consolidate openflare-server to root and move subprojects to internal/apps

- Merge all files inside openflare-server to the repository root directory.
- Relocate agent, relay, and flared subprojects from internal/ to internal/apps/.
- Combine docker-compose files and update build context paths to root.
- Update GitHub workflows and Dockerfiles to refer to new directories and package names.
- Rewrite Go package imports across all files.
- Resolve database renew test race condition and clean up docs.
This commit is contained in:
ryan
2026-06-19 14:23:29 +08:00
parent 19d476ed7f
commit 63cd906cfc
1064 changed files with 366 additions and 1397 deletions
@@ -0,0 +1,95 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"log/slog"
"net"
"strings"
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
)
var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) {
return pkggeoip.NewMaxMindGeoIPService()
}
type accessLogRegionResolver struct {
provider pkggeoip.GeoIPService
cache map[string]string
}
func newAccessLogRegionResolver() (*accessLogRegionResolver, error) {
provider, err := accessLogGeoProviderFactory()
if err != nil {
return nil, err
}
return &accessLogRegionResolver{
provider: provider,
cache: make(map[string]string),
}, nil
}
func (r *accessLogRegionResolver) Close() {
if r == nil || r.provider == nil {
return
}
if err := r.provider.Close(); err != nil {
slog.Warn("close access log geo provider failed", "error", err)
}
}
func (r *accessLogRegionResolver) Resolve(rawIP string) string {
if r == nil || r.provider == nil {
return ""
}
normalizedIP := normalizeAccessLogIP(rawIP)
if normalizedIP == "" {
return ""
}
if cached, ok := r.cache[normalizedIP]; ok {
return cached
}
info, err := r.provider.GetGeoInfo(net.ParseIP(normalizedIP))
if err != nil || info == nil {
r.cache[normalizedIP] = ""
return ""
}
region := strings.TrimSpace(info.Name)
if region == "" {
region = strings.TrimSpace(info.ISOCode)
}
r.cache[normalizedIP] = region
return region
}
func normalizeAccessLogIP(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
if ip := net.ParseIP(trimmed); ip != nil {
return ip.String()
}
trimmed = strings.TrimPrefix(trimmed, "[")
trimmed = strings.TrimSuffix(trimmed, "]")
if ip := net.ParseIP(trimmed); ip != nil {
return ip.String()
}
host, _, err := net.SplitHostPort(strings.TrimSpace(raw))
if err != nil {
return ""
}
host = strings.TrimPrefix(host, "[")
host = strings.TrimSuffix(host, "]")
if ip := net.ParseIP(host); ip != nil {
return ip.String()
}
return ""
}
+143
View File
@@ -0,0 +1,143 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
const (
agentTokenPositiveCacheTTL = 2 * time.Minute
agentTokenNegativeCacheTTL = 10 * time.Minute
)
type cachedAgentNode struct {
node *model.OpenFlareNode
expiresAt time.Time
}
type accessTokenAuthCache struct {
mu sync.RWMutex
positive map[string]cachedAgentNode
negative map[string]time.Time
now func() time.Time
loadNodeByToken func(context.Context, string) (*model.OpenFlareNode, error)
}
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)
},
}
}
func (c *accessTokenAuthCache) authenticate(ctx context.Context, token string) (*model.OpenFlareNode, error) {
now := c.now()
if node, ok := c.getNode(token, now); ok {
return node, nil
}
if c.isMissing(token, now) {
return nil, gorm.ErrRecordNotFound
}
node, err := c.loadNodeByToken(ctx, token)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL))
}
return nil, err
}
c.storeNode(token, node)
return cloneNode(node), nil
}
func (c *accessTokenAuthCache) getNode(token string, now time.Time) (*model.OpenFlareNode, bool) {
c.mu.RLock()
entry, ok := c.positive[token]
c.mu.RUnlock()
if !ok {
return nil, false
}
if now.After(entry.expiresAt) {
c.mu.Lock()
delete(c.positive, token)
c.mu.Unlock()
return nil, false
}
return cloneNode(entry.node), true
}
func (c *accessTokenAuthCache) isMissing(token string, now time.Time) bool {
c.mu.RLock()
expiresAt, ok := c.negative[token]
c.mu.RUnlock()
if !ok {
return false
}
if now.After(expiresAt) {
c.mu.Lock()
delete(c.negative, token)
c.mu.Unlock()
return false
}
return true
}
func (c *accessTokenAuthCache) storeNode(token string, node *model.OpenFlareNode) {
if token == "" || node == nil {
return
}
c.mu.Lock()
defer c.mu.Unlock()
delete(c.negative, token)
c.positive[token] = cachedAgentNode{
node: cloneNode(node),
expiresAt: c.now().Add(agentTokenPositiveCacheTTL),
}
}
func (c *accessTokenAuthCache) storeMissing(token string, expiresAt time.Time) {
if token == "" {
return
}
c.mu.Lock()
defer c.mu.Unlock()
delete(c.positive, token)
c.negative[token] = expiresAt
}
func (c *accessTokenAuthCache) reset() {
c.mu.Lock()
defer c.mu.Unlock()
c.positive = make(map[string]cachedAgentNode)
c.negative = make(map[string]time.Time)
}
// ResetAuthCacheForTest clears the in-memory access token cache for integration tests.
func ResetAuthCacheForTest() {
tokenCache.reset()
}
// AuthenticateAccessToken validates X-Agent-Token against of_nodes.access_token.
func AuthenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
token = strings.TrimSpace(token)
if token == "" {
return nil, errors.New(errMissingAgentToken)
}
return tokenCache.authenticate(ctx, token)
}
+103
View File
@@ -0,0 +1,103 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"gorm.io/gorm"
)
type configVersionRecord struct {
ID uint `gorm:"primaryKey"`
Version string `gorm:"column:version"`
SnapshotJSON string `gorm:"column:snapshot_json"`
SupportFilesJSON string `gorm:"column:support_files_json"`
Checksum string `gorm:"column:checksum"`
IsActive bool `gorm:"column:is_active"`
CreatedAt time.Time `gorm:"column:created_at"`
}
func (configVersionRecord) TableName() string {
return "of_config_versions"
}
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
version, err := loadActiveConfigVersion(ctx)
if err != nil {
return nil, err
}
return &ActiveConfigMeta{
Version: version.Version,
Checksum: version.Checksum,
}, nil
}
func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) {
version, err := loadActiveConfigVersion(ctx)
if err != nil {
return nil, err
}
var supportFiles []SupportFile
if strings.TrimSpace(version.SupportFilesJSON) != "" {
if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil {
return nil, err
}
}
return &ConfigResponse{
Version: version.Version,
Checksum: version.Checksum,
SourceConfigJSON: version.SnapshotJSON,
SupportFiles: sourceSupportFiles(supportFiles),
CreatedAt: version.CreatedAt,
}, nil
}
func loadActiveConfigVersion(ctx context.Context) (*configVersionRecord, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New("database not initialized")
}
version := &configVersionRecord{}
err := conn.Where("is_active = ?", true).Order("id desc").First(version).Error
if err != nil {
return nil, err
}
return version, nil
}
func sourceSupportFiles(files []SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
result := make([]SupportFile, 0, len(files))
for _, file := range files {
if isRuntimeGeneratedSupportFile(file.Path) {
continue
}
result = append(result, file)
}
return result
}
func isRuntimeGeneratedSupportFile(path string) bool {
switch strings.TrimSpace(path) {
case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName:
return true
default:
return false
}
}
func isActiveConfigNotFound(err error) bool {
return errors.Is(err, gorm.ErrRecordNotFound)
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"testing"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
)
func TestIsRuntimeGeneratedSupportFile(t *testing.T) {
tests := []struct {
path string
want bool
}{
{path: "pow_config.json", want: true},
{path: "waf_config.json", want: true},
{path: openrestyrender.SourceConfigFileName, want: true},
{path: "runtime/custom.json", want: false},
{path: "certs/example.pem", want: false},
}
for _, tc := range tests {
if got := isRuntimeGeneratedSupportFile(tc.path); got != tc.want {
t.Fatalf("isRuntimeGeneratedSupportFile(%q) = %v, want %v", tc.path, got, tc.want)
}
}
}
func TestSourceSupportFilesFiltersRuntimeGeneratedFiles(t *testing.T) {
files := []SupportFile{
{Path: "certs/example.pem", Content: "pem"},
{Path: "pow_config.json", Content: "{}"},
{Path: "waf_config.json", Content: "{}"},
{Path: openrestyrender.SourceConfigFileName, Content: "{}"},
{Path: "routes/extra.json", Content: "{}"},
}
filtered := sourceSupportFiles(files)
if len(filtered) != 2 {
t.Fatalf("expected 2 support files, got %d: %+v", len(filtered), filtered)
}
if filtered[0].Path != "certs/example.pem" || filtered[1].Path != "routes/extra.json" {
t.Fatalf("unexpected filtered files: %+v", filtered)
}
}
+20
View File
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
const (
errMissingAgentToken = "缺少 Agent Token"
errInvalidAgentToken = "无权进行此操作,Agent Token 无效"
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效"
errNodeMissingFromContext = "Node object missing from context"
errNoActiveConfig = "当前没有激活版本"
errNodeNotFound = "节点不存在"
errNodeIDRequired = "node_id 不能为空"
errVersionRequired = "version 不能为空"
errInvalidApplyResult = "result 仅支持 success、warning 或 failed"
errIPRequired = "ip 不能为空"
errIPInvalid = "ip 格式无效"
errAgentVersionRequired = "version 不能为空"
errNodeIDConflict = "节点标识生成冲突,请重试"
)
+308
View File
@@ -0,0 +1,308 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"crypto/rand"
"encoding/hex"
"net"
"strings"
"time"
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
openrestyStatusHealthy = "healthy"
openrestyStatusUnhealthy = "unhealthy"
openrestyStatusUnknown = "unknown"
releaseChannelStable = "stable"
)
func newRandomToken() (string, error) {
buf := make([]byte, 16)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return hex.EncodeToString(buf), nil
}
func newServerNodeID() (string, error) {
token, err := newRandomToken()
if err != nil {
return "", err
}
return "node-" + token, nil
}
func normalizeOpenrestyStatus(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case openrestyStatusHealthy:
return openrestyStatusHealthy
case openrestyStatusUnhealthy:
return openrestyStatusUnhealthy
default:
return openrestyStatusUnknown
}
}
func normalizeNodePayload(payload NodePayload) NodePayload {
payload.Name = strings.TrimSpace(payload.Name)
payload.IP = strings.TrimSpace(payload.IP)
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.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
return payload
}
func validateNodePayload(payload NodePayload) error {
if payload.IP == "" {
return errPayload(errIPRequired)
}
if net.ParseIP(payload.IP) == nil {
return errPayload(errIPInvalid)
}
if payload.Version == "" {
return errPayload(errAgentVersionRequired)
}
return nil
}
type payloadError string
func (e payloadError) Error() string { return string(e) }
func errPayload(message string) error { return payloadError(message) }
func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveName bool) {
if !preserveName || strings.TrimSpace(node.Name) == "" {
if strings.TrimSpace(payload.Name) != "" {
node.Name = strings.TrimSpace(payload.Name)
}
}
if !node.IPManualOverride {
node.IP = strings.TrimSpace(payload.IP)
}
node.Version = strings.TrimSpace(payload.Version)
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
node.Status = nodeStatusOnline
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
now := time.Now()
node.LastSeenAt = &now
node.LastError = truncateForDatabase(payload.LastError, 16000)
if !node.GeoManualOverride {
applyGeoInfoFromIP(node, node.IP)
}
}
func applyGeoInfoFromIP(node *model.OpenFlareNode, rawIP string) {
if node == nil {
return
}
node.GeoName = ""
node.GeoLatitude = nil
node.GeoLongitude = nil
ip := net.ParseIP(strings.TrimSpace(rawIP))
if ip == nil {
return
}
info, err := ofgeoip.GeoInfoFromIP(ip)
if err != nil || info == nil {
return
}
if strings.TrimSpace(info.Name) != "" {
node.GeoName = strings.TrimSpace(info.Name)
}
if info.Latitude != nil && info.Longitude != nil {
node.GeoLatitude = cloneCoordinate(info.Latitude)
node.GeoLongitude = cloneCoordinate(info.Longitude)
}
}
func cloneCoordinate(value *float64) *float64 {
if value == nil {
return nil
}
cloned := *value
return &cloned
}
func truncateForDatabase(value string, max int) string {
if max <= 0 {
return ""
}
runes := []rune(strings.TrimSpace(value))
if len(runes) <= max {
return string(runes)
}
return string(runes[:max])
}
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
reported := normalizeIP(reportedIP)
remote := normalizeRemoteAddr(remoteAddr)
if reported == "" {
return remote
}
if isPublicNodeIP(reported) {
return reported
}
if isPublicNodeIP(remote) {
return remote
}
return reported
}
func normalizeIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
host := raw
if strings.Contains(raw, ":") {
if h, _, err := net.SplitHostPort(raw); err == nil {
host = h
}
}
host = strings.TrimPrefix(host, "[")
host = strings.TrimSuffix(host, "]")
if ip := net.ParseIP(host); ip != nil {
return ip.String()
}
return ""
}
func normalizeRemoteAddr(remoteAddr string) string {
remoteAddr = strings.TrimSpace(remoteAddr)
if remoteAddr == "" {
return ""
}
host, _, err := net.SplitHostPort(remoteAddr)
if err != nil {
return normalizeIP(remoteAddr)
}
return normalizeIP(host)
}
func isPublicNodeIP(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
if ip == nil {
return false
}
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
return false
}
return true
}
func buildAgentSettings(node *model.OpenFlareNode, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *Settings {
autoUpdate := false
if node != nil {
autoUpdate = node.AutoUpdateEnabled
}
if strings.TrimSpace(updateChannel) == "" {
updateChannel = releaseChannelStable
}
return &Settings{
HeartbeatInterval: model.AgentHeartbeatInterval,
WebsocketUpgradeEnabled: model.AgentWebsocketUpgradeEnabled,
AutoUpdate: autoUpdate,
UpdateRepo: model.AgentUpdateRepo,
UpdateNow: updateNow,
UpdateChannel: updateChannel,
UpdateTag: strings.TrimSpace(updateTag),
RestartOpenrestyNow: restartOpenrestyNow,
}
}
func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenFlareNode) map[string]any {
if previous == nil || current == nil {
return map[string]any{}
}
changes := make(map[string]any)
appendIfChanged := func(key string, before any, after any) {
if before != after {
changes[key] = after
}
}
appendIfChanged("name", previous.Name, current.Name)
appendIfChanged("ip", previous.IP, current.IP)
appendIfChanged("geo_name", previous.GeoName, current.GeoName)
appendIfChanged("version", previous.Version, current.Version)
appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion)
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage)
appendIfChanged("status", previous.Status, current.Status)
appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion)
appendIfChanged("last_error", previous.LastError, current.LastError)
appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested)
appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel)
appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag)
appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested)
if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) {
changes["geo_latitude"] = current.GeoLatitude
}
if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) {
changes["geo_longitude"] = current.GeoLongitude
}
if !lastSeenAtEqual(previous.LastSeenAt, current.LastSeenAt) {
changes["last_seen_at"] = current.LastSeenAt
}
return changes
}
func coordinatesEqual(before *float64, after *float64) bool {
if before == nil || after == nil {
return before == after
}
return *before == *after
}
func lastSeenAtEqual(before *time.Time, after *time.Time) bool {
if before == nil || after == nil {
return before == after
}
return before.Equal(*after)
}
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.Checksum = strings.TrimSpace(payload.Checksum)
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
return payload
}
func isUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
if node == nil {
return
}
tokenCache.storeNode(node.AccessToken, cloneNode(node))
}
func cloneNode(node *model.OpenFlareNode) *model.OpenFlareNode {
if node == nil {
return nil
}
cloned := *node
return &cloned
}
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"net"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
)
type fakeGeoIPProvider struct {
info *pkggeoip.GeoInfo
}
func (f *fakeGeoIPProvider) Name() string { return "fake-geoip" }
func (f *fakeGeoIPProvider) GetGeoInfo(ip net.IP) (*pkggeoip.GeoInfo, error) {
return f.info, nil
}
func (f *fakeGeoIPProvider) UpdateDatabase() error { return nil }
func (f *fakeGeoIPProvider) Close() error { return nil }
func withFakeGeoIPProvider(t *testing.T, info *pkggeoip.GeoInfo) {
t.Helper()
previous := pkggeoip.CurrentProvider
pkggeoip.CurrentProvider = &fakeGeoIPProvider{info: info}
t.Cleanup(func() {
pkggeoip.CurrentProvider = previous
})
}
func geoipFloat(value float64) *float64 {
return &value
}
func TestApplyGeoInfoFromIP(t *testing.T) {
latitude := 31.2304
longitude := 121.4737
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
Name: "Shanghai",
Latitude: geoipFloat(latitude),
Longitude: geoipFloat(longitude),
})
node := &model.OpenFlareNode{IP: "203.0.113.10"}
applyGeoInfoFromIP(node, node.IP)
if node.GeoName != "Shanghai" {
t.Fatalf("expected geo_name Shanghai, got %q", node.GeoName)
}
if node.GeoLatitude == nil || *node.GeoLatitude != latitude {
t.Fatalf("unexpected geo_latitude: %+v", node.GeoLatitude)
}
if node.GeoLongitude == nil || *node.GeoLongitude != longitude {
t.Fatalf("unexpected geo_longitude: %+v", node.GeoLongitude)
}
}
func TestApplyGeoInfoFromIPSkipsInvalidIP(t *testing.T) {
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{Name: "Should Not Apply"})
node := &model.OpenFlareNode{
IP: "203.0.113.10",
GeoName: "Existing",
GeoLatitude: geoipFloat(1),
GeoLongitude: geoipFloat(2),
}
applyGeoInfoFromIP(node, "not-an-ip")
if node.GeoName != "" || node.GeoLatitude != nil || node.GeoLongitude != nil {
t.Fatalf("expected geo fields to be cleared on invalid IP, got %+v", node)
}
}
func TestApplyNodeRuntimeRespectsGeoManualOverride(t *testing.T) {
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
Name: "Shanghai",
Latitude: geoipFloat(31.2304),
Longitude: geoipFloat(121.4737),
})
node := &model.OpenFlareNode{
GeoManualOverride: true,
GeoName: "Manual",
GeoLatitude: geoipFloat(10),
GeoLongitude: geoipFloat(20),
}
applyNodeRuntime(node, NodePayload{
IP: "203.0.113.10",
Version: "1.0.0",
}, true)
if node.GeoName != "Manual" {
t.Fatalf("expected manual geo_name to be preserved, got %q", node.GeoName)
}
if node.GeoLatitude == nil || *node.GeoLatitude != 10 {
t.Fatalf("expected manual geo_latitude to be preserved, got %+v", node.GeoLatitude)
}
}
func TestCollectHeartbeatChangesTracksGeoFields(t *testing.T) {
before := &model.OpenFlareNode{
IP: "10.0.0.1",
GeoName: "Old Region",
}
after := &model.OpenFlareNode{
IP: "203.0.113.10",
GeoName: "New Region",
GeoLatitude: geoipFloat(31.2304),
GeoLongitude: geoipFloat(121.4737),
}
changes := collectHeartbeatChanges(before, after)
if changes["ip"] != after.IP {
t.Fatalf("expected ip change, got %+v", changes)
}
if changes["geo_name"] != after.GeoName {
t.Fatalf("expected geo_name change, got %+v", changes)
}
if changes["geo_latitude"] != after.GeoLatitude {
t.Fatalf("expected geo_latitude change, got %+v", changes)
}
if changes["geo_longitude"] != after.GeoLongitude {
t.Fatalf("expected geo_longitude change, got %+v", changes)
}
}
+224
View File
@@ -0,0 +1,224 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/node"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
// RegisterWithAccessToken registers an agent on a reserved node token.
func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*RegistrationResponse, error) {
payload = normalizeNodePayload(payload)
if authNode == nil {
return nil, errors.New(errNodeNotFound)
}
if err := validateNodePayload(payload); err != nil {
return nil, err
}
applyNodeRuntime(authNode, payload, true)
if err := model.SaveOpenFlareNode(ctx, authNode); err != nil {
return nil, err
}
RefreshAccessTokenCache(ctx, authNode)
return &RegistrationResponse{
NodeID: authNode.NodeID,
AccessToken: authNode.AccessToken,
Name: authNode.Name,
}, nil
}
// RegisterWithDiscovery registers a new node using the global discovery token.
func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*RegistrationResponse, error) {
payload = normalizeNodePayload(payload)
if err := validateNodePayload(payload); err != nil {
return nil, err
}
nodeID, err := newServerNodeID()
if err != nil {
return nil, err
}
accessToken, err := newRandomToken()
if err != nil {
return nil, err
}
nodeName := payload.Name
if nodeName == "" {
nodeName = nodeID
}
record := &model.OpenFlareNode{
NodeID: nodeID,
Name: nodeName,
AccessToken: accessToken,
Status: nodeStatusOnline,
NodeType: "edge_node",
CapabilitiesJSON: "[]",
UpdateChannel: releaseChannelStable,
}
applyNodeRuntime(record, payload, false)
if err = model.CreateOpenFlareNode(ctx, record); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errNodeIDConflict)
}
return nil, err
}
RefreshAccessTokenCache(ctx, record)
return &RegistrationResponse{
NodeID: record.NodeID,
AccessToken: record.AccessToken,
Name: record.Name,
}, nil
}
// HeartbeatNode updates runtime state and returns agent settings.
func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*HeartbeatResponse, error) {
if authNode == nil {
return nil, errors.New(errNodeNotFound)
}
payload.NodeID = authNode.NodeID
payload = normalizeNodePayload(payload)
if err := validateNodePayload(payload); err != nil {
return nil, err
}
previous := *authNode
updateNow := authNode.UpdateRequested
restartOpenrestyNow := authNode.RestartOpenrestyRequested
updateChannel := strings.TrimSpace(authNode.UpdateChannel)
updateTag := strings.TrimSpace(authNode.UpdateTag)
applyNodeRuntime(authNode, payload, true)
authNode.UpdateRequested = false
authNode.UpdateChannel = releaseChannelStable
authNode.UpdateTag = ""
authNode.RestartOpenrestyRequested = false
changes := collectHeartbeatChanges(&previous, authNode)
if len(changes) > 0 {
fields := make([]string, 0, len(changes))
for field := range changes {
fields = append(fields, field)
}
if err := model.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil {
return nil, err
}
}
RefreshAccessTokenCache(ctx, authNode)
reportedAt := time.Now()
if authNode.LastSeenAt != nil {
reportedAt = *authNode.LastSeenAt
}
PersistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt)
activeConfig, err := getActiveConfigMeta(ctx)
if err != nil && !isActiveConfigNotFound(err) {
return nil, err
}
wafIPGroups, err := ChangedWAFIPGroupsForAgent(ctx, nil, payload.WAFIPGroupChecksums)
if err != nil {
return nil, err
}
return &HeartbeatResponse{
Node: authNode,
AgentSettings: buildAgentSettings(authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow),
ActiveConfig: activeConfig,
WAFIPGroups: wafIPGroups,
}, nil
}
// GetActiveConfig returns the active configuration for an agent.
func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) {
config, err := getActiveConfigForAgent(ctx)
if err != nil {
if isActiveConfigNotFound(err) {
return nil, errors.New(errNoActiveConfig)
}
return nil, err
}
return config, nil
}
// SyncWAFIPGroups returns WAF IP groups whose checksums differ from the agent state.
func SyncWAFIPGroups(ctx context.Context, input WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) {
groups, err := ChangedWAFIPGroupsForAgent(ctx, input.IDs, input.Checksums)
if err != nil {
return nil, err
}
return &WAFIPGroupSyncResult{Groups: groups}, nil
}
// ReportApplyLog records an agent apply result.
func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) {
now := time.Now()
payload = normalizeApplyLogPayload(payload)
if payload.NodeID == "" {
return nil, errors.New(errNodeIDRequired)
}
if payload.Version == "" {
return nil, errors.New(errVersionRequired)
}
if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFailed {
return nil, errors.New(errInvalidApplyResult)
}
log := &model.OpenFlareApplyLog{
NodeID: payload.NodeID,
Version: payload.Version,
Result: payload.Result,
Message: payload.Message,
Checksum: payload.Checksum,
MainConfigChecksum: payload.MainConfigChecksum,
RouteConfigChecksum: payload.RouteConfigChecksum,
SupportFileCount: payload.SupportFileCount,
CreatedAt: now,
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New("database not initialized")
}
err := conn.Transaction(func(tx *gorm.DB) error {
record := &model.OpenFlareNode{}
if err := tx.Where("node_id = ?", payload.NodeID).First(record).Error; err != nil {
return err
}
record.Status = nodeStatusOnline
record.LastSeenAt = &now
if payload.Result == applyResultOK {
record.CurrentVersion = payload.Version
record.LastError = ""
} else {
record.LastError = payload.Message
}
if err := tx.Create(log).Error; err != nil {
return err
}
return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).Error
})
if err != nil {
return nil, err
}
return log, nil
}
// ValidateDiscoveryToken delegates to the node package discovery token helper.
func ValidateDiscoveryToken(ctx context.Context, token string) error {
return node.ValidateDiscoveryToken(ctx, token)
}
@@ -0,0 +1,59 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"strings"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
)
const (
agentTokenHeader = "X-Agent-Token"
agentNodeContextKey = "agent_node"
)
// AgentAuth validates X-Agent-Token against of_nodes.access_token.
func AgentAuth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
node, err := AuthenticateAccessToken(c.Request.Context(), token)
if err != nil {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
c.Set(agentNodeContextKey, node)
c.Next()
}
}
// AgentRegisterAuth accepts either a node access token or the global discovery token.
func AgentRegisterAuth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil {
c.Set(agentNodeContextKey, node)
c.Next()
return
}
if err := ValidateDiscoveryToken(c.Request.Context(), token); err != nil {
response.AbortUnauthorized(c, errInvalidDiscoveryToken)
return
}
c.Set("discovery_enabled", true)
c.Next()
}
}
// AgentNodeFromContext returns the authenticated agent node.
func AgentNodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
value, ok := c.Get(agentNodeContextKey)
if !ok {
return nil, false
}
node, ok := value.(*model.OpenFlareNode)
return node, ok
}
@@ -0,0 +1,199 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupAgentAuthTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareNode{},
&model.OpenFlareOption{},
))
db.SetDB(sqliteDB)
option.ResetInitializationForTest()
tokenCache.reset()
return func() {
db.SetDB(nil)
option.ResetInitializationForTest()
tokenCache.reset()
}
}
func TestAuthenticateAccessToken(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-auth-1",
Name: "edge",
AccessToken: "valid-agent-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
t.Run("valid token", func(t *testing.T) {
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
require.NoError(t, err)
assert.Equal(t, "node-auth-1", node.NodeID)
})
t.Run("cached token", func(t *testing.T) {
originalLoader := tokenCache.loadNodeByToken
t.Cleanup(func() {
tokenCache.loadNodeByToken = originalLoader
})
tokenCache.loadNodeByToken = func(context.Context, string) (*model.OpenFlareNode, error) {
t.Fatal("db should not be queried for cached token")
return nil, nil
}
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
require.NoError(t, err)
assert.Equal(t, "node-auth-1", node.NodeID)
})
t.Run("missing token", func(t *testing.T) {
_, err := AuthenticateAccessToken(ctx, "")
require.Error(t, err)
assert.Contains(t, err.Error(), errMissingAgentToken)
})
t.Run("invalid token", func(t *testing.T) {
_, err := AuthenticateAccessToken(ctx, "invalid-token")
require.Error(t, err)
})
}
func TestAgentAuthMiddleware(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-mw-1",
Name: "edge",
AccessToken: "middleware-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
router := testhelper.NewTestGinEngine()
router.GET("/protected", AgentAuth(), func(c *gin.Context) {
node, ok := AgentNodeFromContext(c)
if !ok {
c.Status(http.StatusInternalServerError)
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"node_id": node.NodeID}))
})
t.Run("authorized request", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set(agentTokenHeader, "middleware-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
assert.Empty(t, apiResp.ErrorMsg)
})
t.Run("unauthorized request", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set(agentTokenHeader, "bad-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusUnauthorized, resp.Code)
})
}
func TestAgentRegisterAuthMiddleware(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-register-1",
Name: "edge",
AccessToken: "existing-node-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
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 {
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "node", "node_id": node.NodeID}))
return
}
if _, ok := c.Get("discovery_enabled"); ok {
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "discovery"}))
return
}
c.Status(http.StatusInternalServerError)
})
t.Run("existing node token", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/register", nil)
req.Header.Set(agentTokenHeader, "existing-node-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
data, ok := apiResp.Data.(map[string]any)
require.True(t, ok)
assert.Equal(t, "node", data["mode"])
})
t.Run("discovery token", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/register", nil)
req.Header.Set(agentTokenHeader, "discovery-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
data, ok := apiResp.Data.(map[string]any)
require.True(t, ok)
assert.Equal(t, "discovery", data["mode"])
})
}
@@ -0,0 +1,499 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"errors"
"log/slog"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"go.uber.org/zap"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
healthEventStatusActive = "active"
healthEventStatusResolved = "resolved"
healthSeverityInfo = "info"
healthSeverityWarning = "warning"
healthSeverityCritical = "critical"
nodeAccessLogRetentionDays = 90
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
accessLogPathMaxLength = 100
)
// NodeSystemProfile is the agent-reported system profile.
type NodeSystemProfile struct {
Hostname string `json:"hostname"`
OSName string `json:"os_name"`
OSVersion string `json:"os_version"`
KernelVersion string `json:"kernel_version"`
Architecture string `json:"architecture"`
CPUModel string `json:"cpu_model"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAtUnix int64 `json:"reported_at_unix"`
}
// NodeMetricSnapshot is the agent-reported capacity snapshot.
type NodeMetricSnapshot struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
MemoryTotalBytes int64 `json:"memory_total_bytes"`
StorageUsedBytes int64 `json:"storage_used_bytes"`
StorageTotalBytes int64 `json:"storage_total_bytes"`
DiskReadBytes int64 `json:"disk_read_bytes"`
DiskWriteBytes int64 `json:"disk_write_bytes"`
NetworkRxBytes int64 `json:"network_rx_bytes"`
NetworkTxBytes int64 `json:"network_tx_bytes"`
}
// NodeOpenrestyObservation is the agent-reported openresty network observation.
type NodeOpenrestyObservation struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
OpenrestyConnections int64 `json:"openresty_connections"`
}
// NodeTrafficReport is the agent-reported traffic window.
type NodeTrafficReport struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
RequestCount int64 `json:"request_count"`
ErrorCount int64 `json:"error_count"`
UniqueVisitorCount int64 `json:"unique_visitor_count"`
StatusCodes map[string]int64 `json:"status_codes"`
TopDomains map[string]int64 `json:"top_domains"`
SourceCountries map[string]int64 `json:"source_countries"`
}
// NodeAccessLog is a single access log row from the agent.
type NodeAccessLog struct {
LoggedAtUnix int64 `json:"logged_at_unix"`
RemoteAddr string `json:"remote_addr"`
Host string `json:"host"`
Path string `json:"path"`
StatusCode int `json:"status_code"`
}
// BufferedObservabilityRecord is a buffered observability window from the agent.
type BufferedObservabilityRecord struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
}
// NodeHealthEvent is an agent-reported health event.
type NodeHealthEvent struct {
EventType string `json:"event_type"`
Severity string `json:"severity"`
Message string `json:"message"`
TriggeredAtUnix int64 `json:"triggered_at_unix"`
Metadata map[string]string `json:"metadata"`
}
// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events.
func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
if strings.TrimSpace(nodeID) == "" {
return
}
if payload.Profile == nil &&
payload.Snapshot == nil &&
payload.TrafficReport == nil &&
len(payload.AccessLogs) == 0 &&
len(payload.BufferedObservability) == 0 &&
payload.HealthEvents == nil {
return
}
conn := db.DB(ctx)
if conn == nil {
return
}
accessLogRecords, err := buildNodeAccessLogRecords(nodeID, payload.AccessLogs, payload.BufferedObservability, reportedAt)
if err != nil {
zap.L().Error("build heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
return
}
if err := conn.Transaction(func(tx *gorm.DB) error {
if err := persistNodeSystemProfile(tx, nodeID, payload.Profile, reportedAt); err != nil {
return err
}
if err := persistBufferedObservability(tx, nodeID, payload.BufferedObservability, reportedAt); err != nil {
return err
}
if err := persistNodeMetricSnapshot(tx, nodeID, payload.Snapshot, reportedAt); err != nil {
return err
}
if err := persistNodeOpenrestyObservation(tx, nodeID, payload.OpenrestyObservation, reportedAt); err != nil {
return err
}
if err := persistNodeTrafficReport(tx, nodeID, payload.TrafficReport, reportedAt); err != nil {
return err
}
if payload.HealthEvents != nil {
if err := reconcileNodeHealthEvents(tx, nodeID, payload.HealthEvents, reportedAt); err != nil {
return err
}
}
return nil
}); err != nil {
zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err))
return
}
if err := persistNodeAccessLogs(ctx, nodeID, accessLogRecords, reportedAt); err != nil {
zap.L().Error("persist heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
}
}
func persistBufferedObservability(tx *gorm.DB, nodeID string, records []BufferedObservabilityRecord, reportedAt time.Time) error {
for _, record := range records {
if err := persistNodeMetricSnapshot(tx, nodeID, record.Snapshot, reportedAt); err != nil {
return err
}
if err := persistNodeOpenrestyObservation(tx, nodeID, record.OpenrestyObservation, reportedAt); err != nil {
return err
}
if err := persistNodeTrafficReport(tx, nodeID, record.TrafficReport, reportedAt); err != nil {
return err
}
}
return nil
}
func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *NodeSystemProfile, reportedAt time.Time) error {
if profile == nil {
return nil
}
record := &model.OpenFlareNodeSystemProfile{
NodeID: nodeID,
Hostname: strings.TrimSpace(profile.Hostname),
OSName: strings.TrimSpace(profile.OSName),
OSVersion: strings.TrimSpace(profile.OSVersion),
KernelVersion: strings.TrimSpace(profile.KernelVersion),
Architecture: strings.TrimSpace(profile.Architecture),
CPUModel: strings.TrimSpace(profile.CPUModel),
CPUCores: profile.CPUCores,
TotalMemoryBytes: profile.TotalMemoryBytes,
TotalDiskBytes: profile.TotalDiskBytes,
UptimeSeconds: profile.UptimeSeconds,
ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt),
}
return tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "node_id"}},
DoUpdates: clause.AssignmentColumns([]string{
"hostname",
"os_name",
"os_version",
"kernel_version",
"architecture",
"cpu_model",
"cpu_cores",
"total_memory_bytes",
"total_disk_bytes",
"uptime_seconds",
"reported_at",
"updated_at",
}),
}).Create(record).Error
}
func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *NodeMetricSnapshot, reportedAt time.Time) error {
if snapshot == nil {
return nil
}
record := &model.OpenFlareMetricSnapshot{
NodeID: nodeID,
CapturedAt: timeFromUnix(snapshot.CapturedAtUnix, reportedAt),
CPUUsagePercent: snapshot.CPUUsagePercent,
MemoryUsedBytes: snapshot.MemoryUsedBytes,
MemoryTotalBytes: snapshot.MemoryTotalBytes,
StorageUsedBytes: snapshot.StorageUsedBytes,
StorageTotalBytes: snapshot.StorageTotalBytes,
DiskReadBytes: snapshot.DiskReadBytes,
DiskWriteBytes: snapshot.DiskWriteBytes,
NetworkRxBytes: snapshot.NetworkRxBytes,
NetworkTxBytes: snapshot.NetworkTxBytes,
}
exists, err := metricSnapshotExists(tx, nodeID, record.CapturedAt)
if err != nil {
return err
}
if exists {
return nil
}
return tx.Create(record).Error
}
func persistNodeOpenrestyObservation(tx *gorm.DB, nodeID string, obs *NodeOpenrestyObservation, reportedAt time.Time) error {
if obs == nil {
return nil
}
record := &model.OpenFlareNodeObservationOpenresty{
NodeID: nodeID,
CapturedAt: timeFromUnix(obs.CapturedAtUnix, reportedAt),
OpenrestyRxBytes: obs.OpenrestyRxBytes,
OpenrestyTxBytes: obs.OpenrestyTxBytes,
OpenrestyConnections: obs.OpenrestyConnections,
}
return tx.Create(record).Error
}
func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *NodeTrafficReport, reportedAt time.Time) error {
if report == nil {
return nil
}
if report.WindowEndedAtUnix > 0 && report.WindowStartedAtUnix > report.WindowEndedAtUnix {
return errors.New("traffic report window_started_at_unix 不能大于 window_ended_at_unix")
}
record := &model.OpenFlareRequestReport{
NodeID: nodeID,
WindowStartedAt: timeFromUnix(report.WindowStartedAtUnix, reportedAt),
WindowEndedAt: timeFromUnix(report.WindowEndedAtUnix, reportedAt),
RequestCount: report.RequestCount,
ErrorCount: report.ErrorCount,
UniqueVisitorCount: report.UniqueVisitorCount,
StatusCodesJSON: marshalJSON(report.StatusCodes),
TopDomainsJSON: marshalJSON(report.TopDomains),
SourceCountriesJSON: marshalJSON(report.SourceCountries),
}
exists, err := requestReportExists(tx, nodeID, record.WindowStartedAt, record.WindowEndedAt)
if err != nil {
return err
}
if exists {
return nil
}
return tx.Create(record).Error
}
func buildNodeAccessLogRecords(nodeID string, direct []NodeAccessLog, buffered []BufferedObservabilityRecord, reportedAt time.Time) ([]*model.OpenFlareAccessLog, error) {
total := len(direct)
for _, record := range buffered {
total += len(record.AccessLogs)
}
if total == 0 {
return nil, nil
}
resolver, err := newAccessLogRegionResolver()
if err != nil {
slog.Warn("initialize access log geo resolver failed", "node_id", nodeID, "error", err)
}
if resolver != nil {
defer resolver.Close()
}
records := make([]*model.OpenFlareAccessLog, 0, total)
appendLogs := func(logs []NodeAccessLog) {
for _, item := range logs {
record := &model.OpenFlareAccessLog{
NodeID: nodeID,
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
Region: "",
Host: strings.TrimSpace(item.Host),
Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength),
StatusCode: item.StatusCode,
}
if resolver != nil {
record.Region = resolver.Resolve(record.RemoteAddr)
}
records = append(records, record)
}
}
appendLogs(direct)
for _, record := range buffered {
appendLogs(record.AccessLogs)
}
return records, nil
}
func persistNodeAccessLogs(ctx context.Context, nodeID string, records []*model.OpenFlareAccessLog, reportedAt time.Time) error {
if len(records) == 0 {
return nil
}
if err := model.InsertOpenFlareAccessLogsBatch(ctx, records); err != nil {
return err
}
_, err := model.DeleteOpenFlareAccessLogsByNodeBefore(ctx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow))
return err
}
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error {
return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
}
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types.
func ReconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
activeTypes := make(map[string]NodeHealthEvent, len(events))
for _, event := range events {
eventType := normalizeHealthEventType(event.EventType)
if eventType == "" {
continue
}
if len(managedEventTypes) > 0 {
if _, ok := managedEventTypes[eventType]; !ok {
continue
}
}
event.EventType = eventType
event.Severity = normalizeHealthSeverity(event.Severity)
if event.TriggeredAtUnix <= 0 {
event.TriggeredAtUnix = reportedAt.Unix()
}
activeTypes[eventType] = event
}
var activeEvents []*model.OpenFlareHealthEvent
query := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive)
if len(managedEventTypes) > 0 {
scopedTypes := make([]string, 0, len(managedEventTypes))
for eventType := range managedEventTypes {
eventType = normalizeHealthEventType(eventType)
if eventType != "" {
scopedTypes = append(scopedTypes, eventType)
}
}
if len(scopedTypes) == 0 {
return nil
}
query = query.Where("event_type IN ?", scopedTypes)
}
if err := query.Find(&activeEvents).Error; err != nil {
return err
}
activeByType := make(map[string]*model.OpenFlareHealthEvent, len(activeEvents))
for _, event := range activeEvents {
activeByType[event.EventType] = event
}
for eventType, event := range activeTypes {
triggeredAt := timeFromUnix(event.TriggeredAtUnix, reportedAt)
if existing, ok := activeByType[eventType]; ok {
existing.Severity = event.Severity
existing.Message = normalizeHealthEventMessage(event.Message)
existing.LastTriggeredAt = triggeredAt
existing.ReportedAt = reportedAt
existing.MetadataJSON = marshalJSON(event.Metadata)
existing.ResolvedAt = nil
if err := tx.Save(existing).Error; err != nil {
return err
}
continue
}
record := &model.OpenFlareHealthEvent{
NodeID: nodeID,
EventType: eventType,
Severity: event.Severity,
Status: healthEventStatusActive,
Message: normalizeHealthEventMessage(event.Message),
FirstTriggeredAt: triggeredAt,
LastTriggeredAt: triggeredAt,
ReportedAt: reportedAt,
MetadataJSON: marshalJSON(event.Metadata),
}
if err := tx.Create(record).Error; err != nil {
return err
}
}
for _, existing := range activeEvents {
if _, ok := activeTypes[existing.EventType]; ok {
continue
}
resolvedAt := reportedAt
existing.Status = healthEventStatusResolved
existing.ReportedAt = reportedAt
existing.ResolvedAt = &resolvedAt
if err := tx.Save(existing).Error; err != nil {
return err
}
}
return nil
}
func metricSnapshotExists(tx *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) {
var count int64
if err := tx.Model(&model.OpenFlareMetricSnapshot{}).
Where("node_id = ? AND captured_at = ?", nodeID, capturedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
return count > 0, nil
}
func requestReportExists(tx *gorm.DB, nodeID string, windowStartedAt, windowEndedAt time.Time) (bool, error) {
var count int64
if err := tx.Model(&model.OpenFlareRequestReport{}).
Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, windowStartedAt, windowEndedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
return count > 0, nil
}
func normalizeHealthEventType(eventType string) string {
eventType = strings.TrimSpace(strings.ToLower(eventType))
eventType = strings.ReplaceAll(eventType, " ", "_")
return eventType
}
func normalizeHealthSeverity(severity string) string {
switch strings.ToLower(strings.TrimSpace(severity)) {
case healthSeverityCritical:
return healthSeverityCritical
case healthSeverityInfo:
return healthSeverityInfo
default:
return healthSeverityWarning
}
}
func normalizeHealthEventMessage(message string) string {
return truncateForDatabase(message, 4096)
}
func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
if unixSeconds <= 0 {
return fallback
}
return time.Unix(unixSeconds, 0).UTC()
}
// MarshalJSON serializes a value for database JSON columns.
func MarshalJSON(value any) string {
return marshalJSON(value)
}
func marshalJSON(value any) string {
if value == nil {
return ""
}
raw, err := json.Marshal(value)
if err != nil {
return ""
}
return string(raw)
}
+212
View File
@@ -0,0 +1,212 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/gin-gonic/gin"
)
// RegisterHandler registers or discovers an agent node.
// @Summary 注册或发现 Agent 节点
// @Description 使用节点 access token 重新注册,或使用全局 discovery token 发现新节点;请求头需携带 X-Agent-Token
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.NodePayload true "节点上报数据"
// @Success 200 {object} response.Any{data=agent.RegistrationResponse} "注册成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/nodes/register [post]
func RegisterHandler(c *gin.Context) {
var payload NodePayload
if !apiutil.BindJSON(c, &payload) {
return
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
var (
result *RegistrationResponse
err error
)
if authNode, ok := AgentNodeFromContext(c); ok {
result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload)
} else {
result, err = RegisterWithDiscovery(c.Request.Context(), payload)
}
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// HeartbeatHandler records agent heartbeat state.
// @Summary Agent 心跳上报
// @Description 上报节点状态、指标与健康事件,返回远程控制配置与活跃配置元信息
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.NodePayload true "心跳数据"
// @Success 200 {object} response.Any{data=agent.HeartbeatResponse} "心跳成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/nodes/heartbeat [post]
func HeartbeatHandler(c *gin.Context) {
var payload NodePayload
if !apiutil.BindJSON(c, &payload) {
return
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := AgentNodeFromContext(c)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
heartbeat, err := HeartbeatNode(c.Request.Context(), authNode, payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(heartbeat))
}
// GetActiveConfigHandler returns the active configuration version.
// @Summary 获取活跃配置版本
// @Description 返回当前生效的完整配置包,供 Agent 拉取并应用
// @Tags openflare-agent
// @Produce json
// @Security AgentTokenAuth
// @Success 200 {object} response.Any{data=agent.ConfigResponse} "活跃配置"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/config-versions/active [get]
func GetActiveConfigHandler(c *gin.Context) {
if _, ok := AgentNodeFromContext(c); !ok {
response.AbortUnauthorized(c, errNodeMissingFromContext)
return
}
config, err := GetActiveConfig(c.Request.Context())
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(config))
}
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent.
// @Summary 同步 WAF IP 组
// @Description 按 ID 与校验和增量同步 WAF IP 组定义
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.WAFIPGroupSyncInput true "同步请求"
// @Success 200 {object} response.Any{data=agent.WAFIPGroupSyncResult} "同步结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/waf/ip-groups/sync [post]
func SyncWAFIPGroupsHandler(c *gin.Context) {
var input WAFIPGroupSyncInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := SyncWAFIPGroups(c.Request.Context(), input)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// ReportApplyLogHandler records an agent apply log entry.
// @Summary 上报配置应用日志
// @Description 记录 Agent 配置下发与应用结果
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.ApplyLogPayload true "应用日志"
// @Success 200 {object} response.Any{data=model.OpenFlareApplyLog} "日志记录"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/apply-logs [post]
func ReportApplyLogHandler(c *gin.Context) {
var payload ApplyLogPayload
if !apiutil.BindJSON(c, &payload) {
return
}
if authNode, ok := AgentNodeFromContext(c); ok {
payload.NodeID = authNode.NodeID
}
log, err := ReportApplyLog(c.Request.Context(), payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(log))
}
// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent.
// @Summary 下载 Pages 部署包
// @Description 流式下载指定部署的静态资源压缩包,供 Agent 边缘分发
// @Tags openflare-agent
// @Produce application/octet-stream
// @Security AgentTokenAuth
// @Param deployment_id path int true "部署 ID"
// @Success 200 {file} binary "部署包文件"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/deployments/{deployment_id}/package [get]
func DownloadPagesPackageHandler(c *gin.Context) {
deploymentID, ok := pagesDeploymentIDParam(c)
if !ok {
return
}
packageObj, fileName, err := pages.OpenDeploymentPackage(c.Request.Context(), deploymentID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
defer packageObj.Body.Close()
c.Header("Content-Disposition", "attachment; filename="+fileName)
if packageObj.ContentType != "" {
c.Header("Content-Type", packageObj.ContentType)
}
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
}
func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("deployment_id")
if raw == "" {
response.AbortBadRequest(c, "无效的 ID")
return 0, false
}
id64, err := strconv.ParseUint(raw, 10, 64)
if err != nil || id64 == 0 {
response.AbortBadRequest(c, "无效的 ID")
return 0, false
}
return uint(id64), true
}
// AgentWebSocketHandler 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)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
}
+119
View File
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"time"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
nodeStatusOnline = "online"
applyResultOK = "success"
applyResultWarn = "warning"
applyResultFailed = "failed"
)
// NodePayload is the agent register/heartbeat payload.
type NodePayload struct {
NodeID string `json:"node_id"`
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
ExtVersion string `json:"ext_version"`
CurrentVersion string `json:"current_version"`
LastError string `json:"last_error"`
OpenrestyStatus string `json:"openresty_status"`
OpenrestyMessage string `json:"openresty_message"`
Profile *NodeSystemProfile `json:"profile,omitempty"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
BufferedObservability []BufferedObservabilityRecord `json:"buffered_observability,omitempty"`
HealthEvents []NodeHealthEvent `json:"health_events"`
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
}
// ApplyLogPayload is the agent apply log report payload.
type ApplyLogPayload struct {
NodeID string `json:"node_id"`
Version string `json:"version"`
Result string `json:"result"`
Message string `json:"message"`
Checksum string `json:"checksum"`
MainConfigChecksum string `json:"main_config_checksum"`
RouteConfigChecksum string `json:"route_config_checksum"`
SupportFileCount int `json:"support_file_count"`
}
// RegistrationResponse is returned after agent registration.
type RegistrationResponse struct {
NodeID string `json:"node_id"`
AccessToken string `json:"access_token"`
Name string `json:"name"`
}
// Settings carries remote agent control flags.
type Settings struct {
HeartbeatInterval int `json:"heartbeat_interval"`
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
AutoUpdate bool `json:"auto_update"`
UpdateRepo string `json:"update_repo"`
UpdateNow bool `json:"update_now"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
RestartOpenrestyNow bool `json:"restart_openresty_now"`
}
// ActiveConfigMeta summarizes the active configuration version.
type ActiveConfigMeta struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
// SupportFile is a configuration support artifact shipped to agents.
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
// ConfigResponse is the full active config payload for agents.
type ConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
SourceConfigJSON string `json:"source_config_json"`
SupportFiles []SupportFile `json:"support_files"`
CreatedAt time.Time `json:"created_at"`
}
// WAFIPGroup is a WAF IP group snapshot for agents.
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
Checksum string `json:"checksum"`
}
// WAFIPGroupSyncInput requests changed WAF IP groups.
type WAFIPGroupSyncInput struct {
IDs []uint `json:"ids"`
Checksums map[string]string `json:"checksums"`
}
// WAFIPGroupSyncResult returns synced WAF IP groups.
type WAFIPGroupSyncResult struct {
Groups []WAFIPGroup `json:"groups"`
}
// HeartbeatResponse is the heartbeat handler result.
type HeartbeatResponse struct {
Node *model.OpenFlareNode `json:"node"`
AgentSettings *Settings `json:"agent_settings"`
ActiveConfig *ActiveConfigMeta `json:"active_config"`
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
}
@@ -0,0 +1,205 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"sort"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
)
type snapshotWAFRuleGroupRef struct {
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
}
type snapshotWAFSection struct {
RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"`
}
type activeConfigSnapshot struct {
WAF snapshotWAFSection `json:"waf"`
}
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
return buildAgentWAFIPGroups(ctx, ids)
}
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
if len(targetIDs) == 0 {
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
if err != nil {
return nil, err
}
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []WAFIPGroup{}, nil
}
groups, err := buildAgentWAFIPGroups(ctx, targetIDs)
if err != nil {
return nil, err
}
changed := make([]WAFIPGroup, 0, len(groups))
for _, group := range groups {
if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum {
continue
}
changed = append(changed, group)
}
return changed, nil
}
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
return []WAFIPGroup{}, nil
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
groups, err := model.ListOpenFlareWAFIPGroupsByIDs(ctx, ids)
if err != nil {
return nil, err
}
groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups))
for _, group := range groups {
groupByID[group.ID] = group
}
result := make([]WAFIPGroup, 0, len(ids))
for _, id := range ids {
group := groupByID[id]
if group == nil {
continue
}
agentGroup, err := buildAgentWAFIPGroup(group)
if err != nil {
return nil, err
}
result = append(result, agentGroup)
}
return result, nil
}
func buildAgentWAFIPGroup(group *model.OpenFlareWAFIPGroup) (WAFIPGroup, error) {
if group == nil {
return WAFIPGroup{}, errors.New("IP 组不存在")
}
ips, err := decodeWAFIPGroupStringList(group.IPList)
if err != nil {
return WAFIPGroup{}, err
}
if !group.Enabled {
ips = []string{}
}
agentGroup := WAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ips,
}
agentGroup.Checksum = checksumAgentWAFIPGroup(agentGroup)
return agentGroup, nil
}
func checksumAgentWAFIPGroup(group WAFIPGroup) string {
payload := struct {
ID uint `json:"id"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
}{
ID: group.ID,
Enabled: group.Enabled,
IPList: append([]string{}, group.IPList...),
}
sort.Strings(payload.IPList)
data, _ := json.Marshal(payload)
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
version, err := loadActiveConfigVersion(ctx)
if err != nil {
if isActiveConfigNotFound(err) {
return []uint{}, nil
}
return nil, err
}
snapshot, err := parseActiveConfigSnapshot(version.SnapshotJSON)
if err != nil {
return nil, err
}
idSet := make(map[uint]struct{})
for _, group := range snapshot.WAF.RuleGroups {
for _, id := range group.IPWhitelistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
for _, id := range group.IPBlacklistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids, nil
}
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
text := strings.TrimSpace(snapshotJSON)
if text == "" {
return &activeConfigSnapshot{}, nil
}
var snapshot activeConfigSnapshot
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
return nil, err
}
if snapshot.WAF.RuleGroups == nil {
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
}
return &snapshot, nil
}
func decodeWAFIPGroupStringList(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []string{}, nil
}
var items []string
if err := json.Unmarshal([]byte(text), &items); err != nil {
return nil, err
}
return items, nil
}
func uniqueUintIDs(ids []uint) []uint {
normalized := make([]uint, 0, len(ids))
seen := make(map[uint]struct{}, len(ids))
for _, id := range ids {
if id == 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
}
return normalized
}
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"strconv"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupWAFIPGroupTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareWAFIPGroup{},
&configVersionRecord{},
))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
t.Helper()
snapshot := map[string]any{
"routes": []any{},
"waf": map[string]any{
"rule_groups": []map[string]any{
{
"id": 1,
"name": "agent refs",
"enabled": true,
"ip_blacklist_group_ids": []uint{ipGroupID},
},
},
"bindings": []any{},
},
}
snapshotJSON, err := json.Marshal(snapshot)
require.NoError(t, err)
require.NoError(t, db.DB(ctx).Create(&configVersionRecord{
Version: "20260618-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "test-checksum",
IsActive: true,
}).Error)
}
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "agent runtime group",
Type: "manual",
Enabled: true,
IPList: `["203.0.113.44"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.NoError(t, err)
require.Len(t, groups, 1)
assert.Equal(t, ipGroup.ID, groups[0].ID)
assert.Equal(t, "203.0.113.44", groups[0].IPList[0])
assert.NotEmpty(t, groups[0].Checksum)
groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10)
same, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
require.NoError(t, err)
assert.Empty(t, same)
ipGroup.IPList = `["203.0.113.45"]`
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
require.NoError(t, err)
require.Len(t, delta, 1)
assert.Equal(t, ipGroup.ID, delta[0].ID)
assert.Equal(t, "203.0.113.45", delta[0].IPList[0])
assert.NotEqual(t, groups[0].Checksum, delta[0].Checksum)
}
func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "sync group",
Type: "manual",
Enabled: true,
IPList: `["198.51.100.10"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
IDs: []uint{ipGroup.ID},
Checksums: map[string]string{},
})
require.NoError(t, err)
require.Len(t, result.Groups, 1)
assert.Equal(t, ipGroup.ID, result.Groups[0].ID)
assert.Equal(t, "198.51.100.10", result.Groups[0].IPList[0])
result, err = SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
IDs: []uint{ipGroup.ID},
Checksums: map[string]string{
strconv.FormatUint(uint64(ipGroup.ID), 10): result.Groups[0].Checksum,
},
})
require.NoError(t, err)
assert.Empty(t, result.Groups)
}
func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "disabled group",
Type: "manual",
Enabled: true,
IPList: `["203.0.113.10"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
ipGroup.Enabled = false
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.NoError(t, err)
require.Len(t, groups, 1)
assert.False(t, groups[0].Enabled)
assert.Empty(t, groups[0].IPList)
assert.NotEmpty(t, groups[0].Checksum)
}
@@ -0,0 +1,57 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"log/slog"
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/Rain-kl/Wavelet/internal/model"
)
// HandleWSStatus processes an agent websocket status payload (replaces HTTP heartbeat in WS mode).
func HandleWSStatus(ctx context.Context, nodeID, remoteAddr string, rawPayload json.RawMessage) {
var payload NodePayload
if err := json.Unmarshal(rawPayload, &payload); err != nil {
slog.Debug("agent ws status payload decode failed", "node_id", nodeID, "error", err)
return
}
authNode, err := model.GetOpenFlareNodeByNodeID(ctx, nodeID)
if err != nil {
slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err)
return
}
payload.IP = resolveReportedNodeIP(payload.IP, remoteAddr)
response, err := HeartbeatNode(ctx, authNode, payload)
if err != nil {
slog.Debug("agent ws status handling failed", "node_id", nodeID, "error", err)
return
}
settingsSent := false
if response.AgentSettings != nil {
settingsSent = ofws.SendAgentSettings(nodeID, response.AgentSettings)
}
activeConfigSent := false
if response.ActiveConfig != nil {
activeConfigSent = ofws.SendAgentActiveConfig(nodeID, response.ActiveConfig)
}
wafIPGroupsSent := false
if len(response.WAFIPGroups) > 0 {
wafIPGroupsSent = ofws.SendAgentWAFIPGroups(nodeID, response.WAFIPGroups)
}
slog.Debug("agent ws status processed",
"node_id", nodeID,
"current_version", payload.CurrentVersion,
"openresty_status", payload.OpenrestyStatus,
"settings_sent", settingsSent,
"active_config_sent", activeConfigSent,
"waf_ip_groups_sent", wafIPGroupsSent,
)
}