refactor(backend): rename OpenFlare directory to lowercase openflare

This commit is contained in:
ryan
2026-08-30 17:43:23 +08:00
parent 06d5fedbfc
commit c93ff6674f
543 changed files with 819 additions and 819 deletions
@@ -0,0 +1,90 @@
// 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 (
"context"
"log/slog"
"net"
"strings"
"sync"
pkggeoip "Wavelet/openflare/share/geoip"
)
// 共享一个 GeoIP 服务实例:mmdb 打开(mmap + 解析元数据)成本不低,缺文件时还会
// 同步下载,绝不能每个上报批次重建。maxminddb.Reader 并发安全,无需额外加锁。
// 初始化失败不锁存:下一批上报会重试(与旧行为一致)。
var (
sharedAccessLogGeoMu sync.Mutex
sharedAccessLogGeoInstance pkggeoip.Service
)
func sharedAccessLogGeoService(ctx context.Context) pkggeoip.Service {
sharedAccessLogGeoMu.Lock()
defer sharedAccessLogGeoMu.Unlock()
if sharedAccessLogGeoInstance == nil {
service, err := pkggeoip.NewMaxMindGeoIPServiceWithContext(ctx, "", "")
if err != nil {
slog.WarnContext(ctx, "initialize access log geo service failed", "error", err)
return nil
}
sharedAccessLogGeoInstance = service
}
return sharedAccessLogGeoInstance
}
// resolveAccessLogRegion resolves the region name for an access-log remote address.
// mmdb Lookup 本身是内存映射 trie 查找(微秒级),无需再建应用层 IP 缓存。
// resolveAccessLogRegion resolves the region name for an access-log remote address.
// mmdb Lookup 本身是内存映射 trie 查找(微秒级),无需再建应用层 IP 缓存。
func resolveAccessLogRegion(ctx context.Context, rawIP string) string {
normalizedIP := normalizeAccessLogIP(rawIP)
if normalizedIP == "" {
return ""
}
service := sharedAccessLogGeoService(ctx)
if service == nil {
return ""
}
info, err := service.GetGeoInfo(net.ParseIP(normalizedIP))
if err != nil || info == nil {
return ""
}
region := strings.TrimSpace(info.Name)
if region == "" {
region = strings.TrimSpace(info.ISOCode)
}
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 ""
}
@@ -0,0 +1,161 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"sync"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"gorm.io/gorm"
)
const (
agentTokenPositiveCacheTTL = 2 * time.Minute
agentTokenNegativeCacheTTL = 10 * time.Minute
// ponytail: 上限仅防未授权口伪造 token 撑爆内存;打满后放弃缓存(回退 DB 查询),行为不变
maxAgentTokenNegativeCacheEntries = 10_000
)
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: repository.GetOpenFlareNodeByAccessToken,
}
}
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)
if len(c.negative) >= maxAgentTokenNegativeCacheEntries {
c.evictExpiredMissingLocked(c.now())
if len(c.negative) >= maxAgentTokenNegativeCacheEntries {
return // 缓存满:放弃缓存该 token,认证仍走 DB,仅防内存无限增长
}
}
c.negative[token] = expiresAt
}
// evictExpiredMissingLocked 清理已过期的 negative 条目,须持写锁调用。
func (c *accessTokenAuthCache) evictExpiredMissingLocked(now time.Time) {
for token, expiresAt := range c.negative {
if now.After(expiresAt) {
delete(c.negative, token)
}
}
}
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)
}
@@ -0,0 +1,88 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"errors"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/pages"
openrestyrender "Wavelet/openflare/share/render/openresty"
"gorm.io/gorm"
)
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
version, err := repository.GetActiveConfigVersion(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 := repository.GetActiveConfigVersion(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
}
}
// Main config version history is independent of Pages deployment history.
// Agents always receive pages routes bound to each project's current active
// deployment so config rollback never depends on pruned packages.
sourceJSON := version.SnapshotJSON
if rebound, rebindErr := pages.RebindSnapshotPagesToCurrentActive(ctx, version.SnapshotJSON); rebindErr != nil {
return nil, rebindErr
} else if strings.TrimSpace(rebound) != "" {
sourceJSON = rebound
}
return &ConfigResponse{
Version: version.Version,
Checksum: version.Checksum,
SourceConfigJSON: sourceJSON,
SupportFiles: sourceSupportFiles(supportFiles),
CreatedAt: version.CreatedAt,
}, 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: // pow_config.json is legacy; waf_config.json is canonical
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 "Wavelet/openflare/share/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)
}
}
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
const (
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 = "节点不存在"
errNodeIDRequired = "node_id 不能为空"
errVersionRequired = "version 不能为空"
errInvalidApplyResult = "result 仅支持 success、warning 或 failed"
errIPRequired = "ip 不能为空"
errIPInvalid = "ip 格式无效"
errAgentVersionRequired = "version 不能为空"
errNodeIDConflict = "节点标识生成冲突,请重试"
)
@@ -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 "Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
)
const (
openrestyStatusHealthy = "healthy"
openrestyStatusUnhealthy = "unhealthy"
openrestyStatusUnknown = "unknown"
releaseChannelStable = "stable"
randomTokenBytes = 16
maxDatabaseTextLength = 16000
defaultAgentHeartbeatInterval = 3000 // 默认心跳间隔 3 秒(毫秒)
defaultAgentUpdateRepo = "Rain-kl/OpenFlare"
)
func newRandomToken() (string, error) {
buf := make([]byte, randomTokenBytes)
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, maxDatabaseTextLength)
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, maxDatabaseTextLength)
// Align L2 edge_health with top-level status/message (PG is latest-state authority).
if payload.EdgeHealth != nil {
if s := strings.TrimSpace(payload.EdgeHealth.Status); s != "" {
if payload.OpenrestyStatus == "" || payload.OpenrestyStatus == openrestyStatusUnknown {
payload.OpenrestyStatus = normalizeOpenrestyStatus(s)
}
}
if m := strings.TrimSpace(payload.EdgeHealth.Message); m != "" && payload.OpenrestyMessage == "" {
payload.OpenrestyMessage = truncateForDatabase(m, maxDatabaseTextLength)
}
// CH series status must match the same authority as PG after normalize.
payload.EdgeHealth.Status = payload.OpenrestyStatus
payload.EdgeHealth.Message = payload.OpenrestyMessage
}
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(ctx context.Context, 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, maxDatabaseTextLength)
node.Status = nodeStatusOnline
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
now := time.Now()
node.LastSeenAt = &now
node.LastError = truncateForDatabase(payload.LastError, maxDatabaseTextLength)
if !node.GeoManualOverride {
ofgeoip.ApplyNodeGeoFromIP(ctx, node, node.IP)
}
}
func truncateForDatabase(value string, maxVal int) string {
if maxVal <= 0 {
return ""
}
runes := []rune(strings.TrimSpace(value))
if len(runes) <= maxVal {
return string(runes)
}
return string(runes[:maxVal])
}
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(ctx context.Context, 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
}
// 从 SystemConfig 读取配置,使用默认值作为降级
heartbeatInterval, _ := repository.GetIntByKey(ctx, model.ConfigKeyAgentHeartbeatInterval)
if heartbeatInterval <= 0 {
heartbeatInterval = defaultAgentHeartbeatInterval
}
wsUpgradeEnabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyAgentWebsocketUpgradeEnabled)
updateRepo, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentUpdateRepo)
if strings.TrimSpace(updateRepo.Value) == "" {
updateRepo.Value = defaultAgentUpdateRepo
}
return &Settings{
HeartbeatInterval: heartbeatInterval,
WebsocketUpgradeEnabled: wsUpgradeEnabled,
AutoUpdate: autoUpdate,
UpdateRepo: updateRepo.Value,
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), maxDatabaseTextLength)
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(_ 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,133 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"net"
"testing"
ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
pkggeoip "Wavelet/openflare/share/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"}
ofgeoip.ApplyNodeGeoFromIP(context.Background(), 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),
}
ofgeoip.ApplyNodeGeoFromIP(context.Background(), 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(context.Background(), 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)
}
}
@@ -0,0 +1,223 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"time"
cf "Wavelet/openflare/plugins/server/domain/cloudflare"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/node"
ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/logger"
)
// RegisterWithAccessToken registers an agent on a reserved node token.
func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*RegistrationResponse, error) {
_ = ofgeoip.EnsureRuntimeProvider(ctx)
payload = normalizeNodePayload(payload)
if authNode == nil {
return nil, errors.New(errNodeNotFound)
}
if err := validateNodePayload(payload); err != nil {
return nil, err
}
applyNodeRuntime(ctx, authNode, payload, true)
if err := repository.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) {
_ = ofgeoip.EnsureRuntimeProvider(ctx)
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(ctx, record, payload, false)
if err = repository.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) {
_ = ofgeoip.EnsureRuntimeProvider(ctx)
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(ctx, 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 := repository.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil {
return nil, err
}
if previous.IP != authNode.IP {
if _, dispatchErr := cf.DispatchNodeSync(ctx, authNode.ID, "cloudflare_agent_ip_update"); dispatchErr != nil {
logger.ErrorF(ctx, "[Cloudflare] enqueue heartbeat node sync failed: node_id=%d error=%v", authNode.ID, dispatchErr)
}
}
}
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(ctx, 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)
}
latest, err := repository.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID)
if err != nil {
return nil, err
}
if model.IsRepeatSuccessApplyLog(latest, payload.Version, payload.Checksum, payload.Result) {
if err := repository.UpdateOpenFlareNodeFromApplyResult(ctx, payload.NodeID, payload.Result, payload.Version, payload.Message, now); err != nil {
return nil, err
}
return latest, nil
}
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,
}
if err := repository.CreateOpenFlareApplyLogAndUpdateNode(ctx, log, payload.Result, payload.Version, payload.Message); 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,60 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"strings"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
const (
agentTokenHeader = "X-Agent-Token" //nolint:gosec // HTTP header name, not a credential value
agentNodeContextKey = "agent_node"
)
// 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)
if err != nil {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
c.Set(agentNodeContextKey, node)
c.Next()
}
}
// 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 {
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()
}
}
// 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
}
@@ -0,0 +1,198 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"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.SystemConfig{},
))
db.SetDB(sqliteDB)
tokenCache.reset()
return func() {
db.SetDB(nil)
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", Auth(), func(c *gin.Context) {
node, ok := NodeFromContext(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, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyAgentDiscoveryToken, "discovery-token"))
router := testhelper.NewTestGinEngine()
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
}
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,234 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"go.uber.org/zap"
)
const (
accessLogPathMaxLength = 100
accessLogUserAgentMaxLength = 512
accessLogCacheStatusMaxLength = 32
)
// PersistHeartbeatObservability stores profile, host metrics, edge health, and access logs.
func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
if strings.TrimSpace(nodeID) == "" {
return
}
if payload.Profile == nil &&
payload.HostMetrics == nil &&
payload.EdgeHealth == nil &&
len(payload.AccessLogs) == 0 &&
len(payload.Buffered) == 0 &&
payload.HealthEvents == nil {
return
}
accessLogRecords, err := buildNodeAccessLogRecords(ctx, nodeID, payload.AccessLogs, payload.Buffered, reportedAt)
if err != nil {
zap.L().Error("build heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
return
}
profile := buildNodeSystemProfileModel(nodeID, payload.Profile, reportedAt)
healthEvents := healthEventInputs(payload.HealthEvents)
if err := repository.PersistOpenFlareNodePGObservability(
ctx,
profile,
nodeID,
healthEvents,
payload.HealthEvents != nil,
reportedAt,
nil,
); err != nil {
zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err))
return
}
if err := persistBufferedObservability(ctx, nodeID, payload.Buffered, reportedAt); err != nil {
zap.L().Error("persist buffered observability failed", zap.String("node_id", nodeID), zap.Error(err))
}
if err := persistNodeMetricSnapshot(ctx, nodeID, payload.HostMetrics, reportedAt); err != nil {
zap.L().Error("persist metric snapshot failed", zap.String("node_id", nodeID), zap.Error(err))
}
if err := persistNodeEdgeHealth(ctx, nodeID, payload.EdgeHealth, payload.OpenrestyStatus, reportedAt); err != nil {
zap.L().Error("persist edge health failed", zap.String("node_id", nodeID), zap.Error(err))
}
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(ctx context.Context, nodeID string, records []BufferedObservabilityRecord, reportedAt time.Time) error {
for _, record := range records {
if err := persistNodeMetricSnapshot(ctx, nodeID, record.HostMetrics, reportedAt); err != nil {
return err
}
if err := persistNodeEdgeHealth(ctx, nodeID, record.EdgeHealth, "", reportedAt); err != nil {
return err
}
}
return nil
}
func persistNodeEdgeHealth(ctx context.Context, nodeID string, health *NodeEdgeHealth, fallbackStatus string, reportedAt time.Time) error {
if health == nil {
return nil
}
status := strings.TrimSpace(health.Status)
if status == "" {
status = strings.TrimSpace(fallbackStatus)
}
if status == "" {
status = openrestyStatusUnknown
}
return repository.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{
NodeID: nodeID,
CapturedAt: timeFromUnix(health.CapturedAtUnix, reportedAt),
Status: status,
Connections: health.Connections,
})
}
func buildNodeSystemProfileModel(nodeID string, profile *NodeSystemProfile, reportedAt time.Time) *model.OpenFlareNodeSystemProfile {
if profile == nil {
return nil
}
return &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),
}
}
func healthEventInputs(events []NodeHealthEvent) []repository.OpenFlareHealthEventInput {
if events == nil {
return nil
}
out := make([]repository.OpenFlareHealthEventInput, 0, len(events))
for _, event := range events {
out = append(out, repository.OpenFlareHealthEventInput{
EventType: event.EventType,
Severity: event.Severity,
Message: event.Message,
TriggeredAtUnix: event.TriggeredAtUnix,
Metadata: event.Metadata,
})
}
return out
}
func persistNodeMetricSnapshot(ctx context.Context, 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,
// NetworkRx/Tx no longer collected from agents; CH columns remain 0.
}
return repository.InsertOpenFlareMetricSnapshot(ctx, record)
}
func buildNodeAccessLogRecords(ctx context.Context, 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
}
records := make([]*model.OpenFlareAccessLog, 0, total)
appendLogs := func(logs []NodeAccessLog) {
for _, item := range logs {
bytesSent := max(item.BytesSent, 0)
requestLength := max(item.RequestLength, 0)
requestTimeMs := max(item.RequestTimeMs, 0)
record := &model.OpenFlareAccessLog{
NodeID: nodeID,
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
Region: resolveAccessLogRegion(ctx, item.RemoteAddr),
Host: strings.TrimSpace(item.Host),
Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength),
UserAgent: truncateForDatabase(strings.TrimSpace(item.UserAgent), accessLogUserAgentMaxLength),
CacheStatus: truncateForDatabase(strings.TrimSpace(item.CacheStatus), accessLogCacheStatusMaxLength),
StatusCode: item.StatusCode,
BytesSent: bytesSent,
RequestLength: requestLength,
RequestTimeMs: requestTimeMs,
}
records = append(records, record)
}
}
appendLogs(direct)
for _, record := range buffered {
appendLogs(record.AccessLogs)
}
return records, nil
}
func persistNodeAccessLogs(ctx context.Context, _ string, records []*model.OpenFlareAccessLog, _ time.Time) error {
if len(records) == 0 {
return nil
}
return repository.InsertOpenFlareAccessLogsBatch(ctx, records)
}
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types.
func ReconcileScopedNodeHealthEvents(ctx context.Context, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
return repository.ReconcileOpenFlareHealthEvents(ctx, nodeID, healthEventInputs(events), reportedAt, managedEventTypes)
}
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)
}
@@ -0,0 +1,34 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"testing"
"time"
)
func TestBuildNodeAccessLogRecordsPreservesBytesSent(t *testing.T) {
reportedAt := time.Date(2026, 7, 12, 10, 0, 0, 0, time.UTC)
records, err := buildNodeAccessLogRecords(context.Background(), "node-a", []NodeAccessLog{
{
LoggedAtUnix: reportedAt.Unix(),
RemoteAddr: "203.0.113.10",
Host: "api.example.com",
Path: "/v1/ping",
StatusCode: 200,
BytesSent: 4096,
},
}, nil, reportedAt)
if err != nil {
t.Fatalf("buildNodeAccessLogRecords() error = %v", err)
}
if len(records) != 1 {
t.Fatalf("expected one access log record, got %d", len(records))
}
if records[0].BytesSent != 4096 {
t.Fatalf("BytesSent = %d, want 4096", records[0].BytesSent)
}
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import pkgprotocol "Wavelet/openflare/share/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
// NodeEdgeHealth is L2 OpenResty health + connections.
type NodeEdgeHealth = pkgprotocol.NodeEdgeHealth
// 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
// WAFIPGroupSyncResult is an alias for WAFIPGroupSyncResponse kept for backward compatibility.
type WAFIPGroupSyncResult = WAFIPGroupSyncResponse
@@ -0,0 +1,298 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"net/http"
"strconv"
"Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/domain/pages"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/openflare/share/protocol"
"Wavelet/pkg/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 := NodeFromContext(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 := NodeFromContext(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 := NodeFromContext(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 := NodeFromContext(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))
}
// GetPagesDeploymentHashHandler returns the upload SHA-256 hash for a Pages deployment package.
// @Summary 查询 Pages 部署包哈希
// @Description 返回 upload 框架记录的 SHA-256 哈希,供 Agent 对比本地缓存并按需拉取部署包(兼容旧路径)
// @Tags openflare-agent
// @Produce json
// @Security AgentTokenAuth
// @Param deployment_id path int true "部署 ID"
// @Success 200 {object} response.Any{data=protocol.PagesDeploymentHashResponse}
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/deployments/{deployment_id}/hash [get]
func GetPagesDeploymentHashHandler(c *gin.Context) {
deploymentID, ok := pagesUintParam(c, "deployment_id")
if !ok {
return
}
hash, err := pages.GetDeploymentPackageHash(c.Request.Context(), deploymentID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(protocol.PagesDeploymentHashResponse{
DeploymentID: deploymentID,
Hash: hash,
}))
}
// 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 := pagesUintParam(c, "deployment_id")
if !ok {
return
}
packageObj, err := pages.OpenDeploymentPackage(c.Request.Context(), deploymentID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
defer func() { _ = packageObj.Body.Close() }()
c.Header("Content-Disposition", "attachment; filename="+packageObj.FileName)
if packageObj.ContentType != "" {
c.Header("Content-Type", packageObj.ContentType)
}
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
}
// GetPagesProjectLatestHashHandler returns the hash of a project's currently active deployment.
// @Summary 查询 Pages 项目最新激活部署哈希
// @Description 按项目 ID 返回当前激活部署的包哈希(类似 latest 指针),Agent 无需关心具体部署 ID
// @Tags openflare-agent
// @Produce json
// @Security AgentTokenAuth
// @Param project_id path int true "Pages 项目 ID"
// @Success 200 {object} response.Any{data=protocol.PagesProjectLatestHashResponse}
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/projects/{project_id}/latest/hash [get]
func GetPagesProjectLatestHashHandler(c *gin.Context) {
projectID, ok := pagesUintParam(c, "project_id")
if !ok {
return
}
metadata, err := pages.GetProjectLatestPackageMetadata(c.Request.Context(), projectID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: metadata.DeploymentID,
Hash: metadata.Hash,
PackageSize: metadata.PackageSize,
FileCount: metadata.FileCount,
TotalSize: metadata.TotalSize,
}))
}
// DownloadPagesProjectLatestPackageHandler streams the active deployment package for a project.
// @Summary 下载 Pages 项目最新激活部署包
// @Description 按项目 ID 下载当前激活部署的压缩包,供 Agent 边缘分发
// @Tags openflare-agent
// @Produce application/octet-stream
// @Security AgentTokenAuth
// @Param project_id path int true "Pages 项目 ID"
// @Success 200 {file} binary "部署包文件"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/projects/{project_id}/latest/package [get]
func DownloadPagesProjectLatestPackageHandler(c *gin.Context) {
projectID, ok := pagesUintParam(c, "project_id")
if !ok {
return
}
packageObj, err := pages.OpenProjectLatestPackage(c.Request.Context(), projectID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
defer func() { _ = packageObj.Body.Close() }()
c.Header("Content-Disposition", "attachment; filename="+packageObj.FileName)
if packageObj.ContentType != "" {
c.Header("Content-Type", packageObj.ContentType)
}
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
}
func pagesUintParam(c *gin.Context, name string) (uint, bool) {
raw := c.Param(name)
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
}
// 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 WebSocketHandler(c *gin.Context) {
authNode, ok := NodeFromContext(c)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
}
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"time"
"Wavelet/openflare/plugins/server/kernel/model"
)
const (
nodeStatusOnline = "online"
applyResultOK = "success"
applyResultWarn = "warning"
applyResultFailed = "failed"
)
// RegistrationResponse is returned after agent registration.
// Server uses access_token; the agent client expects agent_token via RegisterNodeResponse.
type RegistrationResponse struct {
NodeID string `json:"node_id"`
AccessToken string `json:"access_token"`
Name string `json:"name"`
}
// ConfigResponse is the full active config payload for agents.
// Server uses time.Time for CreatedAt; the agent client uses string via ActiveConfigResponse.
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"`
}
// 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,263 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"slices"
"sort"
"strconv"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/share/protocol"
openrestyrender "Wavelet/openflare/share/render/openresty"
)
type activeConfigSnapshot struct {
WAF openrestyrender.WAFDocument `json:"waf"`
}
type runtimeIPMatchConfig struct {
IPs []string `json:"ips,omitempty"`
CIDRs []string `json:"cidrs,omitempty"`
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
}
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
return validatedAgentWAFIPGroups(ctx, ids, false)
}
// 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) {
groups, err := validatedAgentWAFIPGroups(ctx, ids, true)
if err != nil {
return nil, err
}
changed := make([]WAFIPGroup, 0, len(groups))
for _, group := range groups {
if strings.TrimSpace(checksums[strconv.FormatUint(uint64(group.ID), 10)]) == group.Checksum {
continue
}
changed = append(changed, group)
}
return changed, nil
}
func validatedAgentWAFIPGroups(ctx context.Context, ids []uint, fallbackToActive bool) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
if err != nil {
return nil, err
}
if len(targetIDs) == 0 && fallbackToActive {
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []WAFIPGroup{}, nil
}
validationIDs := uniqueUintIDs(append(append([]uint{}, activeIDs...), targetIDs...))
allGroups, err := buildAgentWAFIPGroups(ctx, validationIDs)
if err != nil {
return nil, err
}
runtimeGroups := make(map[string]protocol.WAFIPGroup, len(allGroups))
for _, group := range allGroups {
runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = group
}
if err = protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups); err != nil {
return nil, err
}
targetSet := make(map[uint]struct{}, len(targetIDs))
for _, id := range targetIDs {
targetSet[id] = struct{}{}
}
result := make([]WAFIPGroup, 0, len(targetIDs))
for _, group := range allGroups {
if _, ok := targetSet[group.ID]; ok {
result = append(result, group)
}
}
return result, nil
}
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
return []WAFIPGroup{}, nil
}
slices.Sort(ids)
groups, err := repository.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 := repository.GetActiveConfigVersion(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 {
// Retain legacy flattened references while older active snapshots may
// still exist during a rolling Server upgrade.
for _, id := range group.IPWhitelistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
for _, id := range group.IPBlacklistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
for nodeID, node := range group.Graph.Nodes {
if node.Type != "ip_match" {
continue
}
ids, err := runtimeIPMatchGroupIDs(node.Config)
if err != nil {
return nil, fmt.Errorf("活动配置 WAF 规则 %d 节点 %s 的 IP 匹配配置无效: %w", group.ID, nodeID, err)
}
for _, id := range ids {
if id > 0 {
idSet[id] = struct{}{}
}
}
}
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
ids = append(ids, id)
}
slices.Sort(ids)
return ids, nil
}
func runtimeIPMatchGroupIDs(raw json.RawMessage) ([]uint, error) {
var config runtimeIPMatchConfig
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&config); err != nil {
return nil, err
}
return config.IPGroupIDs, 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 = []openrestyrender.WAFRuleGroup{}
}
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,272 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"strconv"
"strings"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/share/protocol"
db "Wavelet/plugins/infra/database"
"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{},
&model.ConfigVersion{},
))
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(&model.ConfigVersion{
Version: "20260618-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "test-checksum",
IsActive: true,
}).Error)
}
func seedActiveConfigWithWAFGraphIPGroup(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": "graph refs",
"enabled": true,
"graph": map[string]any{
"entry": "start",
"nodes": map[string]any{
"start": map[string]any{
"type": "start",
"config": map[string]any{},
"next": map[string]string{"next": "match"},
},
"match": map[string]any{
"type": "ip_match",
"config": map[string]any{
"ip_group_ids": []uint{ipGroupID},
},
"next": map[string]string{"true": "allow", "false": "allow"},
},
"allow": map[string]any{
"type": "allow",
"config": map[string]any{},
},
},
},
},
},
"bindings": []any{},
},
}
snapshotJSON, err := json.Marshal(snapshot)
require.NoError(t, err)
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
Version: "20260713-graph-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "graph-test-checksum",
IsActive: true,
}).Error)
}
func TestChangedWAFIPGroupsForAgentDiscoversGraphReferences(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "graph runtime group",
Type: "manual",
Enabled: true,
IPList: `["192.0.2.88"]`,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFGraphIPGroup(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, []string{"192.0.2.88"}, groups[0].IPList)
}
func TestChangedWAFIPGroupsForAgentRejectsMalformedIPMatchConfig(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
Version: "20260713-malformed-001",
SnapshotJSON: `{"waf":{"rule_groups":[{"id":7,"graph":{"entry":"match","nodes":{` +
`"match":{"type":"ip_match","config":{"ip_group_ids":"not-an-array"}}}}}],"bindings":[]}}`,
Checksum: "malformed-test-checksum",
IsActive: true,
}).Error)
_, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.ErrorContains(t, err, "规则 7 节点 match")
require.ErrorContains(t, err, "IP 匹配配置无效")
}
func TestChangedWAFIPGroupsForAgentRejectsOversizedSnapshotBeforeChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: strings.Repeat("x", protocol.MaxWAFIPGroupSnapshotBytes),
Type: "manual",
Enabled: true,
IPList: `[]`,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
agentGroup, err := buildAgentWAFIPGroup(ipGroup)
require.NoError(t, err)
_, err = ChangedWAFIPGroupsForAgent(ctx, []uint{ipGroup.ID}, map[string]string{
strconv.FormatUint(uint64(ipGroup.ID), 10): agentGroup.Checksum,
})
require.ErrorContains(t, err, "WAF IP 组快照大小")
require.ErrorContains(t, err, "超过上限")
}
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, repository.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, repository.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, repository.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, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
ipGroup.Enabled = false
require.NoError(t, repository.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,58 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"log/slog"
"Wavelet/openflare/plugins/server/kernel/repository"
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
)
// 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 := repository.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,
)
}