mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 00:26:37 +08:00
refactor(backend): rename OpenFlare directory to lowercase openflare
This commit is contained in:
@@ -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,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package fleet contains shared edge-fleet task metadata used by plugin registration.
|
||||
package fleet
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/option/uptimekuma"
|
||||
"Wavelet/openflare/plugins/server/domain/tls"
|
||||
"Wavelet/openflare/plugins/server/domain/waf"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
"Wavelet/openflare/plugins/server/kernel/task"
|
||||
)
|
||||
|
||||
const (
|
||||
// SSLRenewTask renews due ACME TLS certificates.
|
||||
SSLRenewTask = "openflare:ssl_renew"
|
||||
// TaskTypeSSLRenew is the admin task type for SSL renewal.
|
||||
TaskTypeSSLRenew = "of_ssl_renew"
|
||||
|
||||
// WAFIPGroupSyncTask syncs due automatic/subscription WAF IP groups.
|
||||
WAFIPGroupSyncTask = "openflare:waf_ip_group_sync"
|
||||
// TaskTypeWAFIPGroupSync is the admin task type for WAF IP group sync.
|
||||
TaskTypeWAFIPGroupSync = "of_waf_ip_group_sync"
|
||||
|
||||
// UptimeKumaSyncTask synchronizes proxy routes to Uptime Kuma monitors.
|
||||
UptimeKumaSyncTask = "openflare:uptime_kuma_sync"
|
||||
// TaskTypeUptimeKumaSync is the admin task type for Uptime Kuma sync.
|
||||
TaskTypeUptimeKumaSync = "of_uptime_kuma_sync"
|
||||
|
||||
// LogDBSwitchTask 切换日志数据库任务标识。
|
||||
LogDBSwitchTask = "openflare:log_db_switch"
|
||||
// TaskTypeLogDBSwitch is the admin task type for log database switch.
|
||||
TaskTypeLogDBSwitch = "of_log_db_switch"
|
||||
)
|
||||
|
||||
var (
|
||||
lastUptimeKumaSyncTime time.Time
|
||||
uptimeKumaSyncMutex sync.Mutex
|
||||
)
|
||||
|
||||
// SSLRenewMeta describes the SSL renewal task.
|
||||
var SSLRenewMeta = task.TaskMeta{
|
||||
Type: TaskTypeSSLRenew,
|
||||
AsynqTask: SSLRenewTask,
|
||||
Name: "OpenFlare SSL 自动续期",
|
||||
Description: "扫描即将到期的 ACME 证书并触发自动续期",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncMeta describes the WAF IP group sync task.
|
||||
var WAFIPGroupSyncMeta = task.TaskMeta{
|
||||
Type: TaskTypeWAFIPGroupSync,
|
||||
AsynqTask: WAFIPGroupSyncTask,
|
||||
Name: "OpenFlare WAF IP 组同步",
|
||||
Description: "同步到期的自动规则与订阅类型 WAF IP 组",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// UptimeKumaSyncMeta describes the Uptime Kuma sync task.
|
||||
var UptimeKumaSyncMeta = task.TaskMeta{
|
||||
Type: TaskTypeUptimeKumaSync,
|
||||
AsynqTask: UptimeKumaSyncTask,
|
||||
Name: "OpenFlare Uptime Kuma 同步",
|
||||
Description: "将启用的代理规则同步到 Uptime Kuma 监控",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
// LogDBSwitchMeta 描述切换日志数据库任务。
|
||||
var LogDBSwitchMeta = task.TaskMeta{
|
||||
Type: TaskTypeLogDBSwitch,
|
||||
AsynqTask: LogDBSwitchTask,
|
||||
Name: "切换日志数据库",
|
||||
Description: "复制迁移日志数据并在成功后切换日志主库(期间禁止日志写入)",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
|
||||
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
|
||||
},
|
||||
}
|
||||
|
||||
// SSLRenewHandler renews due TLS certificates.
|
||||
type SSLRenewHandler struct{}
|
||||
|
||||
// Execute runs SSL certificate renewal for all due certificates.
|
||||
func (h *SSLRenewHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
task.AppendLog(ctx, "开始扫描待续期证书")
|
||||
if err := tls.RunSSLRenewJob(ctx); err != nil {
|
||||
task.AppendLog(ctx, "SSL 自动续期失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
msg := "SSL 自动续期任务完成"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncHandler syncs due WAF IP groups to agents.
|
||||
type WAFIPGroupSyncHandler struct{}
|
||||
|
||||
// Execute syncs all due automatic/subscription WAF IP groups.
|
||||
func (h *WAFIPGroupSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
task.AppendLog(ctx, "开始同步到期的 WAF IP 组")
|
||||
if err := waf.SyncDueWAFIPGroups(ctx); err != nil {
|
||||
task.AppendLog(ctx, "WAF IP 组同步失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
msg := "WAF IP 组同步完成"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
// UptimeKumaSyncHandler synchronizes proxy routes to Uptime Kuma.
|
||||
type UptimeKumaSyncHandler struct{}
|
||||
|
||||
// Execute runs Uptime Kuma sync when integration is enabled and the interval has elapsed.
|
||||
func (h *UptimeKumaSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
// 从 SystemConfig 读取 UptimeKuma 配置
|
||||
enabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyUptimeKumaEnabled)
|
||||
if !enabled {
|
||||
msg := "Uptime Kuma 集成未启用,跳过执行"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
interval, _ := repository.GetIntByKey(ctx, model.ConfigKeyUptimeKumaSyncInterval)
|
||||
if interval <= 0 {
|
||||
interval = 5
|
||||
}
|
||||
if time.Since(lastUptimeKumaSyncTime) < time.Duration(interval)*time.Minute {
|
||||
msg := fmt.Sprintf("距上次同步不足 %d 分钟,跳过执行", interval)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
if !uptimeKumaSyncMutex.TryLock() {
|
||||
msg := "Uptime Kuma 同步任务正在执行,跳过本次调度"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
defer uptimeKumaSyncMutex.Unlock()
|
||||
|
||||
task.AppendLog(ctx, "开始同步代理规则到 Uptime Kuma")
|
||||
if err := uptimekuma.SyncToUptimeKuma(ctx); err != nil {
|
||||
task.AppendLog(ctx, "Uptime Kuma 同步失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
lastUptimeKumaSyncTime = time.Now()
|
||||
msg := "Uptime Kuma 同步完成"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package fleet
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestUptimeKumaSyncHandlerSkipsWhenDisabled(t *testing.T) {
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{}))
|
||||
db.SetDB(sqliteDB)
|
||||
t.Cleanup(func() { db.SetDB(nil) })
|
||||
|
||||
ctx := context.Background()
|
||||
require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyUptimeKumaEnabled, "false"))
|
||||
|
||||
result, err := (&UptimeKumaSyncHandler{}).Execute(ctx, nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "未启用")
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package flared defines shared error messages for tunnel client operations.
|
||||
package flared
|
||||
|
||||
const (
|
||||
errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errTunnelNodeTypeMismatch = "此节点不是 TunnelClient 类型"
|
||||
)
|
||||
@@ -0,0 +1,134 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/relay"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
)
|
||||
|
||||
const (
|
||||
updateChannelStable = "stable"
|
||||
defaultTunnelTargetPort = 80
|
||||
)
|
||||
|
||||
func normalizeReleaseChannel(channel string) string {
|
||||
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
|
||||
return "preview"
|
||||
}
|
||||
return updateChannelStable
|
||||
}
|
||||
|
||||
func normalizeFlaredHeartbeatPayload(payload HeartbeatPayload) HeartbeatPayload {
|
||||
payload.ClientVersion = strings.TrimSpace(payload.ClientVersion)
|
||||
payload.FrpVersion = strings.TrimSpace(payload.FrpVersion)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
payload.TunnelStatus = strings.ToLower(strings.TrimSpace(payload.TunnelStatus))
|
||||
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
payload.CurrentChecksum = strings.TrimSpace(payload.CurrentChecksum)
|
||||
|
||||
cleaned := make([]ConnectedRelay, 0, len(payload.ConnectedRelays))
|
||||
for _, item := range payload.ConnectedRelays {
|
||||
item.RelayNodeID = strings.TrimSpace(item.RelayNodeID)
|
||||
item.Status = strings.ToLower(strings.TrimSpace(item.Status))
|
||||
if item.RelayNodeID == "" {
|
||||
continue
|
||||
}
|
||||
if item.Status == "" {
|
||||
item.Status = "unknown"
|
||||
}
|
||||
cleaned = append(cleaned, item)
|
||||
}
|
||||
payload.ConnectedRelays = cleaned
|
||||
return payload
|
||||
}
|
||||
|
||||
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 listTunnelRelayNodes(ctx context.Context) ([]model.OpenFlareNode, error) {
|
||||
nodes, err := repository.ListOpenFlareNodes(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
relays := make([]model.OpenFlareNode, 0)
|
||||
for _, node := range nodes {
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
relays = append(relays, node)
|
||||
}
|
||||
}
|
||||
return relays, nil
|
||||
}
|
||||
|
||||
func relayClientAddress(node *model.OpenFlareNode) string {
|
||||
if node == nil {
|
||||
return ""
|
||||
}
|
||||
port := node.RelayBindPort
|
||||
if port <= 0 {
|
||||
port = 7000
|
||||
}
|
||||
addr := strings.TrimSpace(node.RelayClientAccessAddr)
|
||||
if addr == "" {
|
||||
addr = strings.TrimSpace(node.IP)
|
||||
}
|
||||
if addr == "" {
|
||||
return fmt.Sprintf("127.0.0.1:%d", port)
|
||||
}
|
||||
if _, _, err := net.SplitHostPort(addr); err == nil {
|
||||
return addr
|
||||
}
|
||||
if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 {
|
||||
return net.JoinHostPort(addr, strconv.Itoa(port))
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", addr, port)
|
||||
}
|
||||
|
||||
func parseTunnelTargetAddr(addr string) (string, int) {
|
||||
addr = strings.TrimSpace(addr)
|
||||
if addr == "" {
|
||||
return "127.0.0.1", defaultTunnelTargetPort
|
||||
}
|
||||
host, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
lastColon := strings.LastIndex(addr, ":")
|
||||
if lastColon < 0 {
|
||||
return addr, defaultTunnelTargetPort
|
||||
}
|
||||
host = addr[:lastColon]
|
||||
portStr = addr[lastColon+1:]
|
||||
}
|
||||
port := defaultTunnelTargetPort
|
||||
if _, scanErr := fmt.Sscanf(portStr, "%d", &port); scanErr != nil {
|
||||
port = defaultTunnelTargetPort
|
||||
}
|
||||
if host == "" {
|
||||
host = "127.0.0.1"
|
||||
}
|
||||
return host, port
|
||||
}
|
||||
|
||||
func sanitizeProxyName(domain string) string {
|
||||
return strings.ReplaceAll(strings.ReplaceAll(domain, ".", "-"), "*", "wildcard")
|
||||
}
|
||||
|
||||
func buildTunnelSettings(ctx context.Context, node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *relay.Settings {
|
||||
return relay.BuildSettings(ctx, node, updateNow, updateChannel, updateTag)
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
applyResultOK = "success"
|
||||
applyResultWarn = "warning"
|
||||
applyResultFail = "failed"
|
||||
maxApplyLogMessageLength = 16000
|
||||
)
|
||||
|
||||
// Heartbeat processes an OpenFlared heartbeat and returns runtime settings.
|
||||
func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) {
|
||||
if node == nil {
|
||||
return nil, errors.New("tunnel client node is nil")
|
||||
}
|
||||
if node.NodeType != "tunnel_client" {
|
||||
return nil, fmt.Errorf("node %s is not a tunnel_client", node.NodeID)
|
||||
}
|
||||
|
||||
payload = normalizeFlaredHeartbeatPayload(payload)
|
||||
previous := *node
|
||||
updateNow := node.UpdateRequested
|
||||
updateChannel := normalizeReleaseChannel(node.UpdateChannel)
|
||||
updateTag := strings.TrimSpace(node.UpdateTag)
|
||||
|
||||
now := time.Now().UTC()
|
||||
changes := map[string]any{
|
||||
"version": payload.ClientVersion,
|
||||
"ext_version": payload.FrpVersion,
|
||||
"current_version": payload.CurrentVersion,
|
||||
"last_seen_at": now,
|
||||
"status": nodeStatusOnline,
|
||||
"update_requested": false,
|
||||
"update_channel": updateChannelStable,
|
||||
"update_tag": "",
|
||||
}
|
||||
if !previous.UpdateRequested {
|
||||
delete(changes, "update_requested")
|
||||
}
|
||||
if previous.UpdateChannel == updateChannelStable {
|
||||
delete(changes, "update_channel")
|
||||
}
|
||||
if previous.UpdateTag == "" {
|
||||
delete(changes, "update_tag")
|
||||
}
|
||||
if !node.IPManualOverride && payload.IP != "" && previous.IP != payload.IP {
|
||||
changes["ip"] = payload.IP
|
||||
node.IP = payload.IP
|
||||
}
|
||||
|
||||
node.Version = payload.ClientVersion
|
||||
node.ExtVersion = payload.FrpVersion
|
||||
node.CurrentVersion = payload.CurrentVersion
|
||||
node.UpdateRequested = false
|
||||
node.UpdateChannel = updateChannelStable
|
||||
node.UpdateTag = ""
|
||||
lastSeen := now
|
||||
node.LastSeenAt = &lastSeen
|
||||
node.Status = nodeStatusOnline
|
||||
|
||||
if err := repository.UpdateOpenFlareNodeColumns(ctx, node, changes); err != nil {
|
||||
return nil, fmt.Errorf("update flared heartbeat: %w", err)
|
||||
}
|
||||
agent.RefreshAccessTokenCache(ctx, node)
|
||||
persistFlaredObservability(ctx, node.NodeID, payload, now)
|
||||
|
||||
activeConfig, err := getActiveConfigMeta(ctx)
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &HeartbeatResponse{
|
||||
ActiveConfig: activeConfig,
|
||||
TunnelSettings: buildTunnelSettings(ctx, node, updateNow, updateChannel, updateTag),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetTunnelConfig builds the full tunnel routing config for an OpenFlared client.
|
||||
func GetTunnelConfig(ctx context.Context, node *model.OpenFlareNode) (*TunnelConfigResponse, error) {
|
||||
if node == nil {
|
||||
return nil, errors.New("node is nil")
|
||||
}
|
||||
|
||||
activeVersion, err := getActiveConfigMeta(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("no active config version: %w", err)
|
||||
}
|
||||
|
||||
routes, err := repository.ListProxyRoutes(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get proxy routes: %w", err)
|
||||
}
|
||||
|
||||
relayNodes, err := listTunnelRelayNodes(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get relay nodes: %w", err)
|
||||
}
|
||||
|
||||
relays := make([]RelayInfo, 0, len(relayNodes))
|
||||
for i := range relayNodes {
|
||||
relayNode := relayNodes[i]
|
||||
if relayNode.RelayStatus == "healthy" || relayNode.Status == nodeStatusOnline {
|
||||
relays = append(relays, RelayInfo{
|
||||
RelayNodeID: relayNode.NodeID,
|
||||
Address: relayClientAddress(&relayNode),
|
||||
AuthToken: relayNode.RelayAuthToken,
|
||||
ProxyURL: strings.TrimSpace(relayNode.RelayClientProxyURL),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
proxies := make([]ProxyEntry, 0)
|
||||
for _, route := range routes {
|
||||
if route == nil || route.UpstreamType != "tunnel" || route.TunnelNodeID == nil || *route.TunnelNodeID != node.ID {
|
||||
continue
|
||||
}
|
||||
if !route.Enabled {
|
||||
continue
|
||||
}
|
||||
zoneDomains, domainErr := repository.ListZoneDomainsByRouteID(ctx, route.ID)
|
||||
if domainErr != nil || len(zoneDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
localAddr, localPort := parseTunnelTargetAddr(route.TunnelTargetAddr)
|
||||
proxies = append(proxies, ProxyEntry{
|
||||
Name: fmt.Sprintf("%s-%s", node.NodeID, sanitizeProxyName(zoneDomains[0].Domain)),
|
||||
Type: "http",
|
||||
LocalAddr: localAddr,
|
||||
LocalPort: localPort,
|
||||
CustomDomains: zoneDomainNames(zoneDomains),
|
||||
})
|
||||
}
|
||||
|
||||
return &TunnelConfigResponse{
|
||||
Version: activeVersion.Version,
|
||||
Checksum: activeVersion.Checksum,
|
||||
Relays: relays,
|
||||
Proxies: proxies,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func zoneDomainNames(domains []model.ZoneDomain) []string {
|
||||
names := make([]string, 0, len(domains))
|
||||
for _, domain := range domains {
|
||||
names = append(names, domain.Domain)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// ReportApplyLog records an apply result from OpenFlared.
|
||||
func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) {
|
||||
now := time.Now().UTC()
|
||||
payload = normalizeApplyLogPayload(payload)
|
||||
if payload.NodeID == "" {
|
||||
return nil, errors.New("node_id 不能为空")
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return nil, errors.New("version 不能为空")
|
||||
}
|
||||
if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFail {
|
||||
return nil, errors.New("result 仅支持 success、warning 或 failed")
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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 = strings.TrimSpace(payload.Message)
|
||||
payload.Checksum = strings.TrimSpace(payload.Checksum)
|
||||
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
|
||||
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
|
||||
if len(payload.Message) > maxApplyLogMessageLength {
|
||||
payload.Message = payload.Message[:maxApplyLogMessageLength]
|
||||
}
|
||||
return payload
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const ctxFlaredNodeKey = "flared_node"
|
||||
|
||||
// TunnelAuth authenticates flared requests using X-Tunnel-Token and verifies tunnel_client type.
|
||||
func TunnelAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader("X-Tunnel-Token"))
|
||||
node, err := agent.AuthenticateAccessToken(c.Request.Context(), token)
|
||||
if err != nil {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
if node.NodeType != "tunnel_client" {
|
||||
response.AbortForbidden(c, errTunnelNodeTypeMismatch)
|
||||
return
|
||||
}
|
||||
c.Set(ctxFlaredNodeKey, node)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"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 setupFlaredMiddlewareTestDB(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{}))
|
||||
db.SetDB(sqliteDB)
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func seedFlaredNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareNode {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
node := &model.OpenFlareNode{
|
||||
NodeID: "flared-test-node",
|
||||
Name: "flared-test",
|
||||
Status: "pending",
|
||||
NodeType: nodeType,
|
||||
AccessToken: accessToken,
|
||||
}
|
||||
require.NoError(t, repository.CreateOpenFlareNode(ctx, node))
|
||||
return node
|
||||
}
|
||||
|
||||
func TestTunnelAuthMissingToken(t *testing.T) {
|
||||
cleanup := setupFlaredMiddlewareTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/flared/test", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
}
|
||||
|
||||
func TestTunnelAuthRejectsWrongNodeType(t *testing.T) {
|
||||
cleanup := setupFlaredMiddlewareTestDB(t)
|
||||
defer cleanup()
|
||||
seedFlaredNode(t, "edge_node", "edge-token-flared")
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/flared/test", nil)
|
||||
req.Header.Set("X-Tunnel-Token", "edge-token-flared")
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code)
|
||||
}
|
||||
|
||||
func TestTunnelAuthAcceptsTunnelClient(t *testing.T) {
|
||||
cleanup := setupFlaredMiddlewareTestDB(t)
|
||||
defer cleanup()
|
||||
node := seedFlaredNode(t, "tunnel_client", "tunnel-token-valid")
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.GET("/flared/test", TunnelAuth(), func(c *gin.Context) {
|
||||
authNode, ok := c.Get(ctxFlaredNodeKey)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID)
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/flared/test", nil)
|
||||
req.Header.Set("X-Tunnel-Token", "tunnel-token-valid")
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const flaredRuntimeUnhealthyEventType = "flared_runtime_unhealthy"
|
||||
|
||||
func persistFlaredObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) {
|
||||
connected := make([]string, 0, len(payload.ConnectedRelays))
|
||||
for _, relay := range payload.ConnectedRelays {
|
||||
connected = append(connected, fmt.Sprintf("%s:%s", relay.RelayNodeID, relay.Status))
|
||||
}
|
||||
managedTypes := map[string]struct{}{
|
||||
flaredRuntimeUnhealthyEventType: {},
|
||||
}
|
||||
var events []agent.NodeHealthEvent
|
||||
if payload.TunnelStatus == "unhealthy" {
|
||||
events = append(events, agent.NodeHealthEvent{
|
||||
EventType: flaredRuntimeUnhealthyEventType,
|
||||
Severity: "critical",
|
||||
Message: "openflared runtime is not healthy",
|
||||
TriggeredAtUnix: reportedAt.Unix(),
|
||||
Metadata: map[string]string{
|
||||
"tunnel_status": payload.TunnelStatus,
|
||||
"client_version": payload.ClientVersion,
|
||||
"current_version": payload.CurrentVersion,
|
||||
"current_checksum": payload.CurrentChecksum,
|
||||
"connected_relays": strings.Join(connected, ","),
|
||||
},
|
||||
})
|
||||
}
|
||||
if err := agent.ReconcileScopedNodeHealthEvents(ctx, nodeID, events, reportedAt, managedTypes); err != nil {
|
||||
zap.L().Error("persist flared health events failed", zap.String("node_id", nodeID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupFlaredObservabilityTestDB(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.OpenFlareHealthEvent{},
|
||||
&model.SystemConfig{},
|
||||
&model.ConfigVersion{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
agent.ResetAuthCacheForTest()
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
agent.ResetAuthCacheForTest()
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatFlaredEmitsHealthEventOnUnhealthy(t *testing.T) {
|
||||
cleanup := setupFlaredObservabilityTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
node := &model.OpenFlareNode{
|
||||
NodeID: "node-flared-unhealthy",
|
||||
Name: "flared-unhealthy",
|
||||
AccessToken: "tunnel-token-unhealthy",
|
||||
Status: "pending",
|
||||
NodeType: "tunnel_client",
|
||||
}
|
||||
require.NoError(t, db.DB(ctx).Create(node).Error)
|
||||
|
||||
_, err := Heartbeat(ctx, node, HeartbeatPayload{
|
||||
ClientVersion: "v0.2.0",
|
||||
FrpVersion: "0.61.0",
|
||||
TunnelStatus: "unhealthy",
|
||||
CurrentVersion: "v1",
|
||||
CurrentChecksum: "checksum-1",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
events, err := repository.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 20)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, events, 1)
|
||||
assert.Equal(t, flaredRuntimeUnhealthyEventType, events[0].EventType)
|
||||
assert.Equal(t, "active", events[0].Status)
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import pkgprotocol "Wavelet/openflare/share/protocol"
|
||||
|
||||
// HeartbeatPayload is an alias for FlaredHeartbeatPayload.
|
||||
type HeartbeatPayload = pkgprotocol.FlaredHeartbeatPayload
|
||||
|
||||
// ConnectedRelay is an alias for FlaredConnectedRelay.
|
||||
type ConnectedRelay = pkgprotocol.FlaredConnectedRelay
|
||||
|
||||
// ActiveConfigMeta is an alias for ActiveConfigMeta.
|
||||
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
|
||||
|
||||
// HeartbeatResponse is an alias for FlaredHeartbeatResponse.
|
||||
type HeartbeatResponse = pkgprotocol.FlaredHeartbeatResponse
|
||||
|
||||
// TunnelConfigResponse is an alias for FlaredTunnelConfigResponse.
|
||||
type TunnelConfigResponse = pkgprotocol.FlaredTunnelConfigResponse
|
||||
|
||||
// RelayInfo is an alias for FlaredRelayInfo.
|
||||
type RelayInfo = pkgprotocol.FlaredRelayInfo
|
||||
|
||||
// ProxyEntry is an alias for FlaredProxyEntry.
|
||||
type ProxyEntry = pkgprotocol.FlaredProxyEntry
|
||||
|
||||
// ApplyLogPayload is an alias for ApplyLogPayload.
|
||||
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
|
||||
@@ -0,0 +1,135 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
|
||||
"Wavelet/openflare/plugins/server/kernel/apiutil"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// PostHeartbeat handles POST /tunnel/heartbeat.
|
||||
// @Summary 上报 Tunnel 心跳
|
||||
// @Description Tunnel 客户端定期上报运行状态与中继连接信息,返回活跃配置元数据与隧道设置
|
||||
// @Tags openflare-tunnel
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security TunnelTokenAuth
|
||||
// @Param body body flared.HeartbeatPayload true "心跳载荷"
|
||||
// @Success 200 {object} response.Any{data=flared.HeartbeatResponse} "心跳响应"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Tunnel Token 无效"
|
||||
// @Failure 403 {object} response.Any "节点类型不匹配"
|
||||
// @Router /api/v1/tunnel/heartbeat [post]
|
||||
func PostHeartbeat(c *gin.Context) {
|
||||
var payload HeartbeatPayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
|
||||
authNode, ok := c.Get(ctxFlaredNodeKey)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := Heartbeat(c.Request.Context(), node, payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// GetActiveConfig handles GET /tunnel/config/active.
|
||||
// @Summary 获取活跃隧道配置
|
||||
// @Description 返回 Tunnel 客户端当前应应用的完整路由配置(含中继列表与代理定义)
|
||||
// @Tags openflare-tunnel
|
||||
// @Produce json
|
||||
// @Security TunnelTokenAuth
|
||||
// @Success 200 {object} response.Any{data=flared.TunnelConfigResponse} "隧道配置"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Tunnel Token 无效"
|
||||
// @Failure 403 {object} response.Any "节点类型不匹配"
|
||||
// @Router /api/v1/tunnel/config/active [get]
|
||||
func GetActiveConfig(c *gin.Context) {
|
||||
authNode, ok := c.Get(ctxFlaredNodeKey)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
config, err := GetTunnelConfig(c.Request.Context(), node)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(config))
|
||||
}
|
||||
|
||||
// PostApplyLog handles POST /tunnel/apply-log.
|
||||
// @Summary 上报 Tunnel 配置下发结果
|
||||
// @Description Tunnel 客户端上报配置应用结果,服务端记录下发日志
|
||||
// @Tags openflare-tunnel
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security TunnelTokenAuth
|
||||
// @Param body body flared.ApplyLogPayload true "下发结果载荷"
|
||||
// @Success 200 {object} response.Any{data=model.OpenFlareApplyLog} "下发日志记录"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Tunnel Token 无效"
|
||||
// @Failure 403 {object} response.Any "节点类型不匹配"
|
||||
// @Router /api/v1/tunnel/apply-log [post]
|
||||
func PostApplyLog(c *gin.Context) {
|
||||
var payload ApplyLogPayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
if authNode, ok := c.Get(ctxFlaredNodeKey); ok {
|
||||
if node, ok := authNode.(*model.OpenFlareNode); ok {
|
||||
payload.NodeID = node.NodeID
|
||||
}
|
||||
}
|
||||
|
||||
log, err := ReportApplyLog(c.Request.Context(), payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(log))
|
||||
}
|
||||
|
||||
// GetWebSocket handles GET /tunnel/ws.
|
||||
// @Summary 升级 Tunnel WebSocket 连接
|
||||
// @Description 将已认证的 Tunnel 客户端连接升级为 WebSocket 长连接,用于配置推送
|
||||
// @Tags openflare-tunnel
|
||||
// @Security TunnelTokenAuth
|
||||
// @Failure 401 {object} response.Any "Tunnel Token 无效"
|
||||
// @Failure 403 {object} response.Any "节点类型不匹配"
|
||||
// @Router /api/v1/tunnel/ws [get]
|
||||
func GetWebSocket(c *gin.Context) {
|
||||
authNode, ok := c.Get(ctxFlaredNodeKey)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
ofws.ServeFlared(c, node.NodeID)
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
ofnode "Wavelet/openflare/plugins/server/domain/fleet/node"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
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 setupProtocolTestEnv(t *testing.T) (*gin.Engine, 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{},
|
||||
&model.OpenFlareApplyLog{},
|
||||
&model.OpenFlareNodeSystemProfile{},
|
||||
&model.OpenFlareHealthEvent{},
|
||||
&model.ConfigVersion{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
agent.ResetAuthCacheForTest()
|
||||
testhelper.SetupLogStoresForTest(t)
|
||||
|
||||
engine := testhelper.NewTestGinEngine()
|
||||
mountOpenFlareTestRoutes(engine)
|
||||
|
||||
cleanup := func() {
|
||||
db.SetDB(nil)
|
||||
agent.ResetAuthCacheForTest()
|
||||
}
|
||||
return engine, cleanup
|
||||
}
|
||||
|
||||
func TestAgentRelayFlaredProtocol(t *testing.T) {
|
||||
engine, cleanup := setupProtocolTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("create edge node and heartbeat with X-Agent-Token", func(t *testing.T) {
|
||||
edge, err := ofnode.CreateNode(ctx, ofnode.Input{
|
||||
Name: "edge-1",
|
||||
IP: "10.0.0.1",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, edge.AccessToken)
|
||||
assert.Equal(t, "edge_node", edge.NodeType)
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/nodes/heartbeat", map[string]any{
|
||||
"name": "edge-1",
|
||||
"ip": "203.0.113.10",
|
||||
"version": "0.1.0",
|
||||
}, map[string]string{
|
||||
"X-Agent-Token": edge.AccessToken,
|
||||
})
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.NotNil(t, data["agent_settings"])
|
||||
})
|
||||
|
||||
t.Run("create tunnel_relay node and relay heartbeat", func(t *testing.T) {
|
||||
relayNode, err := ofnode.CreateNode(ctx, ofnode.Input{
|
||||
Name: "relay-1",
|
||||
NodeType: "tunnel_relay",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, relayNode.AccessToken)
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/relay/heartbeat", map[string]any{
|
||||
"version": "v0.1.0",
|
||||
"frp_version": "0.61.0",
|
||||
"relay_status": "healthy",
|
||||
"name": "relay-1",
|
||||
"ip": "203.0.113.20",
|
||||
}, map[string]string{
|
||||
"X-Agent-Token": relayNode.AccessToken,
|
||||
})
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
var heartbeatData struct {
|
||||
RelayConfig map[string]any `json:"relay_config"`
|
||||
RelaySettings map[string]any `json:"relay_settings"`
|
||||
}
|
||||
unmarshalAPIData(t, resp.Data, &heartbeatData)
|
||||
assert.NotNil(t, heartbeatData.RelayConfig)
|
||||
assert.NotNil(t, heartbeatData.RelaySettings)
|
||||
|
||||
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, relayNode.NodeID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "online", stored.Status)
|
||||
assert.Equal(t, "healthy", stored.RelayStatus)
|
||||
})
|
||||
|
||||
t.Run("create tunnel_client node and flared heartbeat with X-Tunnel-Token", func(t *testing.T) {
|
||||
clientNode, err := ofnode.CreateNode(ctx, ofnode.Input{
|
||||
Name: "client-1",
|
||||
NodeType: "tunnel_client",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, clientNode.AccessToken)
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/tunnel/heartbeat", map[string]any{
|
||||
"client_version": "v0.2.0",
|
||||
"frp_version": "0.61.0",
|
||||
"tunnel_status": "running",
|
||||
}, map[string]string{
|
||||
"X-Tunnel-Token": clientNode.AccessToken,
|
||||
})
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
requireAPIOK(t, rec)
|
||||
|
||||
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, clientNode.NodeID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "online", stored.Status)
|
||||
assert.Equal(t, "v0.2.0", stored.Version)
|
||||
})
|
||||
|
||||
t.Run("agent register with discovery token from options", func(t *testing.T) {
|
||||
bootstrap, err := ofnode.GetBootstrapToken(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, bootstrap.DiscoveryToken)
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/nodes/register", map[string]any{
|
||||
"name": "discovered-edge",
|
||||
"ip": "203.0.113.30",
|
||||
"version": "0.2.0",
|
||||
}, map[string]string{
|
||||
"X-Agent-Token": bootstrap.DiscoveryToken,
|
||||
})
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
var registration agent.RegistrationResponse
|
||||
unmarshalAPIData(t, resp.Data, ®istration)
|
||||
assert.NotEmpty(t, registration.NodeID)
|
||||
assert.NotEmpty(t, registration.AccessToken)
|
||||
assert.Equal(t, "discovered-edge", registration.Name)
|
||||
|
||||
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, registration.NodeID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "online", stored.Status)
|
||||
assert.Equal(t, registration.AccessToken, stored.AccessToken)
|
||||
})
|
||||
|
||||
t.Run("POST agent apply-logs", func(t *testing.T) {
|
||||
edge, err := ofnode.CreateNode(ctx, ofnode.Input{
|
||||
Name: "edge-apply",
|
||||
IP: "10.0.0.2",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/apply-logs", map[string]any{
|
||||
"version": "20260618-001",
|
||||
"result": "success",
|
||||
"message": "apply ok",
|
||||
}, map[string]string{
|
||||
"X-Agent-Token": edge.AccessToken,
|
||||
})
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
var applyLog model.OpenFlareApplyLog
|
||||
unmarshalAPIData(t, resp.Data, &applyLog)
|
||||
assert.Equal(t, edge.NodeID, applyLog.NodeID)
|
||||
assert.Equal(t, "success", applyLog.Result)
|
||||
assert.Equal(t, "20260618-001", applyLog.Version)
|
||||
|
||||
stored, err := repository.GetOpenFlareNodeByNodeID(ctx, edge.NodeID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "online", stored.Status)
|
||||
assert.Equal(t, "20260618-001", stored.CurrentVersion)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,178 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
"Wavelet/pkg/idgen"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type statusPayload struct {
|
||||
Version string `json:"version"`
|
||||
ServerAddress string `json:"server_address"`
|
||||
}
|
||||
|
||||
func setupAuthOptionIntegration(t *testing.T) (*gorm.DB, *gin.Engine) {
|
||||
t.Helper()
|
||||
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
t.Cleanup(cleanup)
|
||||
|
||||
require.NoError(t, dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeyCapLoginEnabled).
|
||||
Update("value", "false").Error)
|
||||
require.NoError(t, repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled))
|
||||
runtimeconfig.SetSessionSecret("test_openflare_session_secret")
|
||||
|
||||
store := cookie.NewStore([]byte("test_openflare_session_secret"))
|
||||
r := testhelper.NewTestGinEngine(sessions.Sessions("test_openflare_session", store))
|
||||
mountOpenFlareTestRoutes(r)
|
||||
|
||||
return dbConn, r
|
||||
}
|
||||
|
||||
func seedUser(t *testing.T, dbConn *gorm.DB, username, password string, isAdmin bool) *model.User {
|
||||
t.Helper()
|
||||
|
||||
user := &model.User{
|
||||
ID: idgen.NextUint64ID(),
|
||||
Username: username,
|
||||
Nickname: username,
|
||||
Email: username + "@openflare.test",
|
||||
IsActive: true,
|
||||
IsAdmin: isAdmin,
|
||||
}
|
||||
require.NoError(t, user.SetEncryptedPassword(password))
|
||||
require.NoError(t, dbConn.Create(user).Error)
|
||||
return user
|
||||
}
|
||||
|
||||
func seedUserWithAccessToken(t *testing.T, dbConn *gorm.DB, username, password string, isAdmin bool) string {
|
||||
t.Helper()
|
||||
|
||||
user := seedUser(t, dbConn, username, password, isAdmin)
|
||||
|
||||
token, err := model.GenerateTokenString()
|
||||
require.NoError(t, err)
|
||||
|
||||
tokenRecord := model.AccessToken{
|
||||
UserID: user.ID,
|
||||
Name: username + "-integration-token",
|
||||
TokenHash: model.HashToken(token),
|
||||
MaskedToken: model.MaskTokenString(token),
|
||||
IsAdmin: isAdmin,
|
||||
}
|
||||
require.NoError(t, dbConn.Create(&tokenRecord).Error)
|
||||
return token
|
||||
}
|
||||
|
||||
func TestGETStatusReturnsSuccessEnvelope(t *testing.T) {
|
||||
_, r := setupAuthOptionIntegration(t)
|
||||
|
||||
w := performJSONRequest(t, r, http.MethodGet, apiPath("/status"), nil, nil)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
resp := requireAPIOK(t, w)
|
||||
|
||||
var status statusPayload
|
||||
unmarshalAPIData(t, resp.Data, &status)
|
||||
assert.NotEmpty(t, status.Version)
|
||||
}
|
||||
|
||||
func TestGETOptionRequiresAdminAuth(t *testing.T) {
|
||||
dbConn, r := setupAuthOptionIntegration(t)
|
||||
commonToken := seedUserWithAccessToken(t, dbConn, "commonuser", "password123", false)
|
||||
adminToken := seedUserWithAccessToken(t, dbConn, "adminuser", "password123", true)
|
||||
|
||||
t.Run("unauthenticated", func(t *testing.T) {
|
||||
t.Skip("console auth is owned by Wavelet auth plugin")
|
||||
})
|
||||
|
||||
t.Run("non-admin user forbidden", func(t *testing.T) {
|
||||
t.Skip("console auth is owned by Wavelet auth plugin")
|
||||
_ = commonToken
|
||||
})
|
||||
|
||||
t.Run("admin user allowed", func(t *testing.T) {
|
||||
w := performJSONRequest(t, r, http.MethodGet, apiPath("/option/"), nil, adminAuthHeaders(adminToken))
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
requireAPIOK(t, w)
|
||||
})
|
||||
}
|
||||
|
||||
func TestPOSTOptionUpdateRejectsInvalidParams(t *testing.T) {
|
||||
dbConn, r := setupAuthOptionIntegration(t)
|
||||
adminToken := seedUserWithAccessToken(t, dbConn, "adminuser", "password123", true)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, apiPath("/option/update"), bytes.NewReader([]byte("{invalid")))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Access-Token", adminToken)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
resp := decodeAPIResponse(t, w)
|
||||
assert.NotEmpty(t, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestGETNodesWithAccessToken(t *testing.T) {
|
||||
dbConn, r := setupAuthOptionIntegration(t)
|
||||
require.NoError(t, dbConn.AutoMigrate(&model.OpenFlareNode{}))
|
||||
adminToken := seedUserWithAccessToken(t, dbConn, "admin", "password123", true)
|
||||
|
||||
w := performJSONRequest(t, r, http.MethodGet, apiPath("/nodes/"), nil, adminAuthHeaders(adminToken))
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
requireAPIOK(t, w)
|
||||
}
|
||||
|
||||
func TestOptionUpdatePersistsAndReflectsInStatus(t *testing.T) {
|
||||
dbConn, r := setupAuthOptionIntegration(t)
|
||||
adminToken := seedUserWithAccessToken(t, dbConn, "admin", "password123", true)
|
||||
|
||||
updateResp := performJSONRequest(t, r, http.MethodPost, apiPath("/option/update"), map[string]string{
|
||||
"key": model.ConfigKeyServerAddress,
|
||||
"value": "https://hotreload.openflare.test",
|
||||
}, adminAuthHeaders(adminToken))
|
||||
assert.Equal(t, http.StatusOK, updateResp.Code)
|
||||
requireAPIOK(t, updateResp)
|
||||
|
||||
statusAfter := getStatusServerAddress(t, r, nil)
|
||||
assert.Equal(t, "https://hotreload.openflare.test", statusAfter)
|
||||
|
||||
// 验证已持久化到 SystemConfig
|
||||
ctx := context.Background()
|
||||
saved, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "https://hotreload.openflare.test", saved.Value)
|
||||
}
|
||||
|
||||
func getStatusServerAddress(t *testing.T, r http.Handler, headers map[string]string) string {
|
||||
t.Helper()
|
||||
|
||||
w := performJSONRequest(t, r, http.MethodGet, apiPath("/status"), nil, headers)
|
||||
require.Equal(t, http.StatusOK, w.Code)
|
||||
resp := requireAPIOK(t, w)
|
||||
|
||||
var status statusPayload
|
||||
unmarshalAPIData(t, resp.Data, &status)
|
||||
return status.ServerAddress
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
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"
|
||||
)
|
||||
|
||||
const (
|
||||
adminUserID = uint64(1001)
|
||||
adminUsername = "openflare-admin"
|
||||
)
|
||||
|
||||
type adminSeed struct {
|
||||
User model.User
|
||||
Token string
|
||||
TokenHash string
|
||||
}
|
||||
|
||||
func setupCoreChainTest(t *testing.T) (*gin.Engine, adminSeed, func()) {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.User{},
|
||||
&model.AccessToken{},
|
||||
&model.Origin{},
|
||||
&model.ProxyRoute{},
|
||||
&model.ConfigVersion{},
|
||||
&model.OpenFlareWAFRuleGroup{},
|
||||
&model.OpenFlareWAFRuleGroupBinding{},
|
||||
&model.OpenFlareWAFIPGroup{},
|
||||
&model.OpenFlareNode{},
|
||||
&model.SystemConfig{},
|
||||
&model.OpenFlareApplyLog{},
|
||||
&model.Zone{},
|
||||
&model.ZoneDomain{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
agent.ResetAuthCacheForTest()
|
||||
|
||||
seed, err := seedAdminWithAccessToken(sqliteDB)
|
||||
require.NoError(t, err)
|
||||
|
||||
engine := testhelper.NewTestGinEngine()
|
||||
mountOpenFlareTestRoutes(engine)
|
||||
|
||||
cleanup := func() {
|
||||
db.SetDB(nil)
|
||||
agent.ResetAuthCacheForTest()
|
||||
}
|
||||
|
||||
return engine, seed, cleanup
|
||||
}
|
||||
|
||||
func seedAdminWithAccessToken(conn *gorm.DB) (adminSeed, error) {
|
||||
now := time.Now().UTC()
|
||||
admin := model.User{
|
||||
ID: adminUserID,
|
||||
Username: adminUsername,
|
||||
Nickname: "OpenFlare Admin",
|
||||
IsActive: true,
|
||||
IsAdmin: true,
|
||||
LastLoginAt: now,
|
||||
}
|
||||
if err := conn.Create(&admin).Error; err != nil {
|
||||
return adminSeed{}, err
|
||||
}
|
||||
|
||||
token, err := model.GenerateTokenString()
|
||||
if err != nil {
|
||||
return adminSeed{}, err
|
||||
}
|
||||
tokenHash := model.HashToken(token)
|
||||
tokenRecord := model.AccessToken{
|
||||
UserID: adminUserID,
|
||||
Name: "integration-admin-token",
|
||||
TokenHash: tokenHash,
|
||||
MaskedToken: model.MaskTokenString(token),
|
||||
IsAdmin: true,
|
||||
}
|
||||
if err := conn.Create(&tokenRecord).Error; err != nil {
|
||||
return adminSeed{}, err
|
||||
}
|
||||
|
||||
return adminSeed{
|
||||
User: admin,
|
||||
Token: token,
|
||||
TokenHash: tokenHash,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func TestCoreChainMigrationFlow(t *testing.T) {
|
||||
engine, seed, cleanup := setupCoreChainTest(t)
|
||||
defer cleanup()
|
||||
|
||||
var (
|
||||
originID uint
|
||||
proxyRouteID uint
|
||||
configVersion string
|
||||
configChecksum string
|
||||
nodeID uint
|
||||
nodePublicID string
|
||||
agentToken string
|
||||
)
|
||||
|
||||
t.Run("create origin", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/origins/"), map[string]any{
|
||||
"name": "Primary Origin",
|
||||
"address": "origin.core-chain.internal",
|
||||
"remark": "integration upstream",
|
||||
}, map[string]string{
|
||||
"X-Access-Token": seed.Token,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
originID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, originID)
|
||||
assert.Equal(t, "Primary Origin", data["name"])
|
||||
assert.Equal(t, "origin.core-chain.internal", data["address"])
|
||||
})
|
||||
|
||||
t.Run("create proxy route linked to origin", func(t *testing.T) {
|
||||
// Create Zone and ZoneDomain directly in the DB
|
||||
zone := model.Zone{Domain: "example.com"}
|
||||
require.NoError(t, db.DB(context.Background()).Create(&zone).Error)
|
||||
zoneDomain := model.ZoneDomain{
|
||||
ZoneID: zone.ID,
|
||||
Domain: "core-chain.example.com",
|
||||
}
|
||||
require.NoError(t, db.DB(context.Background()).Create(&zoneDomain).Error)
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{
|
||||
"site_name": "core-chain-site",
|
||||
"zone_domain_ids": []uint{zoneDomain.ID},
|
||||
"origin_id": originID,
|
||||
"origin_scheme": "http",
|
||||
"origin_port": "8080",
|
||||
"enabled": true,
|
||||
}, map[string]string{
|
||||
"X-Access-Token": seed.Token,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
proxyRouteID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, proxyRouteID)
|
||||
assert.Equal(t, "core-chain-site", data["site_name"])
|
||||
assert.NotEmpty(t, data["zone_domains"])
|
||||
zoneDomains := data["zone_domains"].([]any)
|
||||
assert.Len(t, zoneDomains, 1)
|
||||
assert.Equal(t, "core-chain.example.com", zoneDomains[0].(map[string]any)["domain"])
|
||||
assert.InDelta(t, float64(originID), data["origin_id"], 1e-9)
|
||||
assert.Equal(t, "http://origin.core-chain.internal:8080", data["origin_url"])
|
||||
})
|
||||
|
||||
t.Run("publish config version", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/config-versions/publish"), nil, map[string]string{
|
||||
"X-Access-Token": seed.Token,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
configVersion, _ = data["version"].(string)
|
||||
configChecksum, _ = data["checksum"].(string)
|
||||
assert.NotEmpty(t, configVersion)
|
||||
assert.NotEmpty(t, configChecksum)
|
||||
assert.Equal(t, true, data["is_active"])
|
||||
|
||||
activeRec := performJSONRequest(t, engine, http.MethodGet, apiPath("/config-versions/active"), nil, map[string]string{
|
||||
"X-Access-Token": seed.Token,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, activeRec.Code)
|
||||
activeResp := requireAPIOK(t, activeRec)
|
||||
|
||||
activeData := unmarshalAPIMap(t, activeResp.Data)
|
||||
assert.Equal(t, configVersion, activeData["version"])
|
||||
assert.Equal(t, configChecksum, activeData["checksum"])
|
||||
})
|
||||
|
||||
t.Run("create node", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/nodes/"), map[string]any{
|
||||
"name": "edge-core-chain",
|
||||
"ip": "10.10.0.1",
|
||||
"auto_update_enabled": true,
|
||||
}, map[string]string{
|
||||
"X-Access-Token": seed.Token,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
nodeID = uint(data["id"].(float64))
|
||||
nodePublicID, _ = data["node_id"].(string)
|
||||
agentToken, _ = data["access_token"].(string)
|
||||
assert.NotZero(t, nodeID)
|
||||
assert.NotEmpty(t, nodePublicID)
|
||||
assert.Len(t, agentToken, 32)
|
||||
})
|
||||
|
||||
t.Run("create apply log for node", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, "/api/v1/agent/apply-logs", map[string]any{
|
||||
"version": configVersion,
|
||||
"result": "success",
|
||||
"message": "config applied",
|
||||
"checksum": configChecksum,
|
||||
"main_config_checksum": "main-checksum",
|
||||
"route_config_checksum": "route-checksum",
|
||||
"support_file_count": 2,
|
||||
}, map[string]string{
|
||||
"X-Agent-Token": agentToken,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.Equal(t, nodePublicID, data["node_id"])
|
||||
assert.Equal(t, configVersion, data["version"])
|
||||
assert.Equal(t, "success", data["result"])
|
||||
assert.Equal(t, configChecksum, data["checksum"])
|
||||
})
|
||||
|
||||
t.Run("verify apply log listing and node metadata", func(t *testing.T) {
|
||||
listRec := performJSONRequest(
|
||||
t,
|
||||
engine,
|
||||
http.MethodGet,
|
||||
apiPath("/apply-logs/?node_id="+nodePublicID+"&pageNo=1&pageSize=10"),
|
||||
nil,
|
||||
map[string]string{
|
||||
"X-Access-Token": seed.Token,
|
||||
},
|
||||
)
|
||||
require.Equal(t, http.StatusOK, listRec.Code)
|
||||
|
||||
listResp := requireAPIOK(t, listRec)
|
||||
listData := unmarshalAPIMap(t, listResp.Data)
|
||||
assert.InDelta(t, float64(1), listData["total"], 1e-9)
|
||||
|
||||
rows, ok := listData["rows"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Len(t, rows, 1)
|
||||
row, ok := rows[0].(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, nodePublicID, row["node_id"])
|
||||
assert.Equal(t, configVersion, row["version"])
|
||||
assert.Equal(t, "success", row["result"])
|
||||
|
||||
nodeRec := performJSONRequest(t, engine, http.MethodGet, apiPath("/nodes/"), nil, map[string]string{
|
||||
"X-Access-Token": seed.Token,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, nodeRec.Code)
|
||||
nodeResp := requireAPIOK(t, nodeRec)
|
||||
|
||||
nodes := unmarshalAPISlice(t, nodeResp.Data)
|
||||
require.Len(t, nodes, 1)
|
||||
nodeView, ok := nodes[0].(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, float64(nodeID), nodeView["id"], 1e-9)
|
||||
assert.Equal(t, nodePublicID, nodeView["node_id"])
|
||||
assert.Equal(t, "success", nodeView["latest_apply_result"])
|
||||
assert.Equal(t, configChecksum, nodeView["latest_apply_checksum"])
|
||||
assert.InDelta(t, float64(2), nodeView["latest_support_file_count"], 1e-9)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/openflare/plugins/server"
|
||||
ofrouter "Wavelet/openflare/plugins/server/httpapi"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func decodeAPIResponse(t *testing.T, rec *httptest.ResponseRecorder) response.Any {
|
||||
t.Helper()
|
||||
|
||||
var resp response.Any
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
|
||||
return resp
|
||||
}
|
||||
|
||||
func requireAPIOK(t *testing.T, rec *httptest.ResponseRecorder) response.Any {
|
||||
t.Helper()
|
||||
|
||||
resp := decodeAPIResponse(t, rec)
|
||||
require.Empty(t, resp.ErrorMsg, "unexpected API error: %s", resp.ErrorMsg)
|
||||
return resp
|
||||
}
|
||||
|
||||
func unmarshalAPIData(t *testing.T, data any, target any) {
|
||||
t.Helper()
|
||||
|
||||
payload, err := json.Marshal(data)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, json.Unmarshal(payload, target))
|
||||
}
|
||||
|
||||
func unmarshalAPIMap(t *testing.T, data any) map[string]any {
|
||||
t.Helper()
|
||||
|
||||
var result map[string]any
|
||||
unmarshalAPIData(t, data, &result)
|
||||
return result
|
||||
}
|
||||
|
||||
func unmarshalAPISlice(t *testing.T, data any) []any {
|
||||
t.Helper()
|
||||
|
||||
var result []any
|
||||
unmarshalAPIData(t, data, &result)
|
||||
return result
|
||||
}
|
||||
|
||||
// mountOpenFlareTestRoutes 复刻 driver_http 的挂载方式:先由 server 插件经内核
|
||||
// 路由注册表声明路由,再把每条 (方法, 路径, 中间件+处理链) 挂到测试引擎上。
|
||||
func mountOpenFlareTestRoutes(engine *gin.Engine) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
core.Provide[contracts.AuthService](ctx, testhelper.StubAuth{})
|
||||
if err := server.New().Apply(ctx); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
for _, rd := range ctx.Router().Routes() {
|
||||
chain := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers))
|
||||
for _, item := range append(append([]any{}, rd.Middlewares...), rd.Handlers...) {
|
||||
chain = append(chain, toGinHandler(item))
|
||||
}
|
||||
engine.Handle(rd.Method, rd.Path, chain...)
|
||||
}
|
||||
}
|
||||
|
||||
// toGinHandler 与 driver_http 接受的处理函数形态保持一致。
|
||||
func toGinHandler(item any) gin.HandlerFunc {
|
||||
switch fn := item.(type) {
|
||||
case gin.HandlerFunc:
|
||||
return fn
|
||||
case func(*gin.Context):
|
||||
return gin.HandlerFunc(fn)
|
||||
default:
|
||||
panic("unexpected handler type")
|
||||
}
|
||||
}
|
||||
|
||||
func apiPath(subpath string) string {
|
||||
return ofrouter.V1BasePath + subpath
|
||||
}
|
||||
|
||||
func performJSONRequest(
|
||||
t *testing.T,
|
||||
engine http.Handler,
|
||||
method, path string,
|
||||
body any,
|
||||
headers map[string]string,
|
||||
) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
|
||||
var payload []byte
|
||||
if body != nil {
|
||||
var err error
|
||||
payload, err = json.Marshal(body)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(method, path, bytes.NewReader(payload))
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
for key, value := range headers {
|
||||
req.Header.Set(key, value)
|
||||
}
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
func adminAuthHeaders(token string) map[string]string {
|
||||
return map[string]string{
|
||||
"X-Access-Token": token,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,374 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
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 setupSecurityTest(t *testing.T) (*gin.Engine, adminSeed, func()) {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.User{},
|
||||
&model.AccessToken{},
|
||||
&model.Origin{},
|
||||
&model.ProxyRoute{},
|
||||
&model.OpenFlareWAFRuleGroup{},
|
||||
&model.OpenFlareWAFRuleGroupBinding{},
|
||||
&model.OpenFlareWAFIPGroup{},
|
||||
&model.TLSCertificate{},
|
||||
&model.Zone{},
|
||||
&model.ZoneDomain{},
|
||||
&model.DNSAccount{},
|
||||
&model.AcmeAccount{},
|
||||
&model.SystemConfig{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
|
||||
seed, err := seedAdminWithAccessToken(sqliteDB)
|
||||
require.NoError(t, err)
|
||||
|
||||
previous := runtimeconfig.Get()
|
||||
runtimeconfig.SetSessionSecret("test_session_secret_for_security_integration")
|
||||
|
||||
engine := testhelper.NewTestGinEngine()
|
||||
mountOpenFlareTestRoutes(engine)
|
||||
|
||||
cleanup := func() {
|
||||
runtimeconfig.Set(previous)
|
||||
db.SetDB(nil)
|
||||
}
|
||||
|
||||
return engine, seed, cleanup
|
||||
}
|
||||
|
||||
func generateSelfSignedCertificatePair(t *testing.T, dnsNames []string) (string, string) {
|
||||
t.Helper()
|
||||
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
require.NoError(t, err)
|
||||
|
||||
template := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(time.Now().UnixNano()),
|
||||
Subject: pkix.Name{
|
||||
CommonName: dnsNames[0],
|
||||
},
|
||||
DNSNames: dnsNames,
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(24 * time.Hour),
|
||||
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||
}
|
||||
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
|
||||
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
|
||||
return string(certPEM), string(keyPEM)
|
||||
}
|
||||
|
||||
func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
engine, seed, cleanup := setupSecurityTest(t)
|
||||
defer cleanup()
|
||||
|
||||
var (
|
||||
ruleGroupID uint
|
||||
ipGroupID uint
|
||||
proxyRouteID uint
|
||||
certID uint
|
||||
domainID uint
|
||||
dnsAccountID uint
|
||||
)
|
||||
|
||||
t.Run("WAF rule group create", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/waf/rule-groups"), map[string]any{
|
||||
"name": "edge-security",
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
ruleGroupID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, ruleGroupID)
|
||||
assert.Equal(t, "edge-security", data["name"])
|
||||
assert.Equal(t, false, data["is_global"])
|
||||
assert.InDelta(t, float64(1), data["revision"], 1e-9)
|
||||
assert.NotNil(t, data["graph"])
|
||||
})
|
||||
|
||||
t.Run("WAF rule group list includes global and custom groups", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodGet, apiPath("/waf/rule-groups"), nil, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
groups := unmarshalAPISlice(t, resp.Data)
|
||||
require.GreaterOrEqual(t, len(groups), 2)
|
||||
|
||||
foundCustom := false
|
||||
foundGlobal := false
|
||||
for _, item := range groups {
|
||||
group, ok := item.(map[string]any)
|
||||
require.True(t, ok)
|
||||
if group["is_global"] == true {
|
||||
foundGlobal = true
|
||||
}
|
||||
if uint(group["id"].(float64)) == ruleGroupID {
|
||||
foundCustom = true
|
||||
assert.Equal(t, "edge-security", group["name"])
|
||||
}
|
||||
}
|
||||
assert.True(t, foundGlobal)
|
||||
assert.True(t, foundCustom)
|
||||
})
|
||||
|
||||
t.Run("WAF rule group get detail", func(t *testing.T) {
|
||||
rec := performJSONRequest(
|
||||
t,
|
||||
engine,
|
||||
http.MethodGet,
|
||||
fmt.Sprintf("%s/waf/rule-groups/%d", apiPath(""), ruleGroupID),
|
||||
nil,
|
||||
adminAuthHeaders(seed.Token),
|
||||
)
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.InDelta(t, float64(ruleGroupID), data["id"], 1e-9)
|
||||
assert.Equal(t, "edge-security", data["name"])
|
||||
})
|
||||
|
||||
t.Run("WAF rule group update", func(t *testing.T) {
|
||||
rec := performJSONRequest(
|
||||
t,
|
||||
engine,
|
||||
http.MethodPost,
|
||||
fmt.Sprintf("%s/waf/rule-groups/%d/meta", apiPath(""), ruleGroupID),
|
||||
map[string]any{
|
||||
"name": "edge-security-updated", "enabled": true,
|
||||
},
|
||||
adminAuthHeaders(seed.Token),
|
||||
)
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.Equal(t, "edge-security-updated", data["name"])
|
||||
assert.Equal(t, true, data["enabled"])
|
||||
})
|
||||
|
||||
t.Run("WAF IP group create", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/waf/ip-groups"), map[string]any{
|
||||
"name": "blocked-ips",
|
||||
"type": "manual",
|
||||
"enabled": true,
|
||||
"ip_list": []string{"203.0.113.0/24", "198.51.100.10"},
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
ipGroupID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, ipGroupID)
|
||||
assert.Equal(t, "blocked-ips", data["name"])
|
||||
assert.Equal(t, "manual", data["type"])
|
||||
})
|
||||
|
||||
t.Run("create proxy route for WAF binding", func(t *testing.T) {
|
||||
// Create Zone and ZoneDomain directly in the DB
|
||||
routeZone := model.Zone{Domain: "example-route.com"}
|
||||
require.NoError(t, db.DB(context.Background()).Create(&routeZone).Error)
|
||||
routeZoneDomain := model.ZoneDomain{
|
||||
ZoneID: routeZone.ID,
|
||||
Domain: "route.example-route.com",
|
||||
}
|
||||
require.NoError(t, db.DB(context.Background()).Create(&routeZoneDomain).Error)
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{
|
||||
"site_name": "security-site",
|
||||
"zone_domain_ids": []uint{routeZoneDomain.ID},
|
||||
"origin_url": "http://origin.security.internal:8080",
|
||||
"enabled": true,
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
proxyRouteID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, proxyRouteID)
|
||||
assert.NotEmpty(t, data["zone_domains"])
|
||||
zoneDomains := data["zone_domains"].([]any)
|
||||
assert.Len(t, zoneDomains, 1)
|
||||
assert.Equal(t, "route.example-route.com", zoneDomains[0].(map[string]any)["domain"])
|
||||
})
|
||||
|
||||
t.Run("bind WAF rule group to proxy route", func(t *testing.T) {
|
||||
rec := performJSONRequest(
|
||||
t,
|
||||
engine,
|
||||
http.MethodPost,
|
||||
fmt.Sprintf("%s/waf/sites/%d/rule-groups", apiPath(""), proxyRouteID),
|
||||
map[string]any{
|
||||
"ids": []uint{ruleGroupID},
|
||||
},
|
||||
adminAuthHeaders(seed.Token),
|
||||
)
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.InDelta(t, float64(proxyRouteID), data["route_id"], 1e-9)
|
||||
|
||||
appliedIDs, ok := data["applied_ids"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Len(t, appliedIDs, 1)
|
||||
assert.InDelta(t, float64(ruleGroupID), appliedIDs[0], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("verify site rule groups binding", func(t *testing.T) {
|
||||
rec := performJSONRequest(
|
||||
t,
|
||||
engine,
|
||||
http.MethodGet,
|
||||
fmt.Sprintf("%s/waf/sites/%d/rule-groups", apiPath(""), proxyRouteID),
|
||||
nil,
|
||||
adminAuthHeaders(seed.Token),
|
||||
)
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.NotNil(t, data["global_rule_group"])
|
||||
|
||||
appliedGroups, ok := data["applied_rule_groups"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Len(t, appliedGroups, 1)
|
||||
group, ok := appliedGroups[0].(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, float64(ruleGroupID), group["id"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("create TLS certificate with PEM", func(t *testing.T) {
|
||||
certPEM, keyPEM := generateSelfSignedCertificatePair(t, []string{"security.example.com"})
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/tls-certificates/"), map[string]any{
|
||||
"name": "security-cert",
|
||||
"cert_pem": certPEM,
|
||||
"key_pem": keyPEM,
|
||||
"remark": "self-signed integration cert",
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
certID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, certID)
|
||||
assert.Equal(t, "security-cert", data["name"])
|
||||
assert.Equal(t, "upload", data["provider"])
|
||||
})
|
||||
|
||||
t.Run("create Zone domain", func(t *testing.T) {
|
||||
zoneRec := performJSONRequest(t, engine, http.MethodPost, apiPath("/zones/"), map[string]any{
|
||||
"domain": "example.com",
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, zoneRec.Code)
|
||||
zoneData := unmarshalAPIMap(t, requireAPIOK(t, zoneRec).Data)
|
||||
zoneID := uint(zoneData["id"].(float64))
|
||||
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, fmt.Sprintf("%s/zones/%d/domains", apiPath(""), zoneID), map[string]any{
|
||||
"domain": "*.example.com",
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusBadRequest, rec.Code)
|
||||
errResp := decodeAPIResponse(t, rec)
|
||||
assert.NotEmpty(t, errResp.ErrorMsg)
|
||||
|
||||
rec = performJSONRequest(t, engine, http.MethodPost, fmt.Sprintf("%s/zones/%d/domains", apiPath(""), zoneID), map[string]any{
|
||||
"domain": "security.example.com", "cert_id": certID,
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
domainID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, domainID)
|
||||
assert.Equal(t, "security.example.com", data["domain"])
|
||||
assert.InDelta(t, float64(certID), data["cert_id"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("create DNS account", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/dns-accounts/"), map[string]any{
|
||||
"name": "cloudflare-dns",
|
||||
"type": "cloudflare",
|
||||
"authorization": "test-api-token-value",
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
dnsAccountID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, dnsAccountID)
|
||||
assert.Equal(t, "cloudflare-dns", data["name"])
|
||||
assert.Equal(t, "cloudflare", data["type"])
|
||||
// API 响应会脱敏 authorization,不应回显明文凭证。
|
||||
if auth, ok := data["authorization"]; ok {
|
||||
assert.NotEqual(t, "test-api-token-value", auth)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("WAF rule group delete", func(t *testing.T) {
|
||||
rec := performJSONRequest(
|
||||
t,
|
||||
engine,
|
||||
http.MethodPost,
|
||||
fmt.Sprintf("%s/waf/rule-groups/%d/delete", apiPath(""), ruleGroupID),
|
||||
nil,
|
||||
adminAuthHeaders(seed.Token),
|
||||
)
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
requireAPIOK(t, rec)
|
||||
|
||||
detailRec := performJSONRequest(
|
||||
t,
|
||||
engine,
|
||||
http.MethodGet,
|
||||
fmt.Sprintf("%s/waf/rule-groups/%d", apiPath(""), ruleGroupID),
|
||||
nil,
|
||||
adminAuthHeaders(seed.Token),
|
||||
)
|
||||
require.Equal(t, http.StatusNotFound, detailRec.Code)
|
||||
detailResp := decodeAPIResponse(t, detailRec)
|
||||
assert.NotEmpty(t, detailResp.ErrorMsg)
|
||||
})
|
||||
|
||||
_ = ipGroupID
|
||||
_ = domainID
|
||||
_ = dnsAccountID
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package node defines node validation and management error messages.
|
||||
package node
|
||||
|
||||
const (
|
||||
errNodeNameRequired = "节点名不能为空"
|
||||
errNodeIPTooLong = "节点 IP 不能超过 64 个字符"
|
||||
errNodeIPInvalid = "节点 IP 格式无效"
|
||||
errNodeIPManualRequired = "锁定节点 IP 时必须填写节点 IP"
|
||||
errNodeGeoNameTooLong = "节点位置名不能超过 128 个字符"
|
||||
errNodeGeoCoordinateMismatch = "地图坐标必须同时填写纬度和经度"
|
||||
errNodeGeoLatitudeInvalid = "纬度必须在 -90 到 90 之间"
|
||||
errNodeGeoLongitudeInvalid = "经度必须在 -180 到 180 之间"
|
||||
errNodeIDConflict = "节点标识生成冲突,请重试"
|
||||
errNodeNotFound = "节点不存在"
|
||||
errNodeForceSyncFailed = "节点不在线或通过 WebSocket 发送同步指令失败"
|
||||
errNoActiveConfigVersion = "当前没有激活版本"
|
||||
errAgentPreviewTagInvalid = "指定版本不是 preview 发布"
|
||||
errAgentStableTagInvalid = "正式版更新不能选择 preview 发布"
|
||||
)
|
||||
@@ -0,0 +1,414 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
nodeStatusOffline = "offline"
|
||||
nodeStatusPending = "pending"
|
||||
openrestyStatusHealthy = "healthy"
|
||||
openrestyStatusUnhealthy = "unhealthy"
|
||||
openrestyStatusUnknown = "unknown"
|
||||
githubReleasesAPIBase = "https://api.github.com/repos/%s/releases"
|
||||
nodeTypeTunnelRelay = "tunnel_relay"
|
||||
nodeTypeTunnelClient = "tunnel_client"
|
||||
nodeTypeEdgeNode = "edge_node"
|
||||
|
||||
nodeTokenByteLength = 16
|
||||
maxNodeIPLength = 64
|
||||
maxNodeGeoNameLength = 128
|
||||
)
|
||||
|
||||
type releaseChannel string
|
||||
|
||||
const (
|
||||
releaseChannelStable releaseChannel = "stable"
|
||||
releaseChannelPreview releaseChannel = "preview"
|
||||
)
|
||||
|
||||
var releaseHTTPClient = &http.Client{Timeout: 30 * time.Second}
|
||||
|
||||
type githubReleaseResponse struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Body string `json:"body"`
|
||||
HTMLURL string `json:"html_url"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
}
|
||||
|
||||
func newRandomToken() (string, error) {
|
||||
buf := make([]byte, nodeTokenByteLength)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
func tokenEqual(got, want string) bool {
|
||||
sumGot := sha256.Sum256([]byte(got))
|
||||
sumWant := sha256.Sum256([]byte(want))
|
||||
return subtle.ConstantTimeCompare(sumGot[:], sumWant[:]) == 1
|
||||
}
|
||||
|
||||
func newServerNodeID() (string, error) {
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "node-" + token, nil
|
||||
}
|
||||
|
||||
func normalizeNodeType(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case nodeTypeTunnelRelay:
|
||||
return nodeTypeTunnelRelay
|
||||
case nodeTypeTunnelClient:
|
||||
return nodeTypeTunnelClient
|
||||
default:
|
||||
return nodeTypeEdgeNode
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeRelayPort(port int, defaultPort int) int {
|
||||
if port <= 0 || port > 65535 {
|
||||
return defaultPort
|
||||
}
|
||||
return port
|
||||
}
|
||||
|
||||
func normalizeReleaseChannel(channel string) releaseChannel {
|
||||
switch strings.ToLower(strings.TrimSpace(channel)) {
|
||||
case string(releaseChannelPreview):
|
||||
return releaseChannelPreview
|
||||
default:
|
||||
return releaseChannelStable
|
||||
}
|
||||
}
|
||||
|
||||
func (channel releaseChannel) String() string {
|
||||
if channel == releaseChannelPreview {
|
||||
return string(releaseChannelPreview)
|
||||
}
|
||||
return string(releaseChannelStable)
|
||||
}
|
||||
|
||||
func normalizeOpenrestyStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case openrestyStatusHealthy:
|
||||
return openrestyStatusHealthy
|
||||
case openrestyStatusUnhealthy:
|
||||
return openrestyStatusUnhealthy
|
||||
default:
|
||||
return openrestyStatusUnknown
|
||||
}
|
||||
}
|
||||
|
||||
func cloneCoordinate(value *float64) *float64 {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *value
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func resolveNodeIPManualOverride(input Input, existing *model.OpenFlareNode, normalizedIP string) bool {
|
||||
if input.IPManualOverride != nil {
|
||||
return *input.IPManualOverride
|
||||
}
|
||||
if existing == nil {
|
||||
return strings.TrimSpace(normalizedIP) != ""
|
||||
}
|
||||
if existing.IPManualOverride {
|
||||
return true
|
||||
}
|
||||
return strings.TrimSpace(normalizedIP) != "" && strings.TrimSpace(normalizedIP) != strings.TrimSpace(existing.IP)
|
||||
}
|
||||
|
||||
func normalizeNodeInput(input Input) (string, string, string, *float64, *float64, bool, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
ip := strings.TrimSpace(input.IP)
|
||||
geoName := strings.TrimSpace(input.GeoName)
|
||||
if err := validateNodeIPInput(input, ip); err != nil {
|
||||
return "", "", "", nil, nil, false, err
|
||||
}
|
||||
if len(geoName) > maxNodeGeoNameLength {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoNameTooLong)
|
||||
}
|
||||
|
||||
geoLatitude := cloneCoordinate(input.GeoLatitude)
|
||||
geoLongitude := cloneCoordinate(input.GeoLongitude)
|
||||
if err := validateNodeGeoCoordinates(geoLatitude, geoLongitude); err != nil {
|
||||
return "", "", "", nil, nil, false, err
|
||||
}
|
||||
|
||||
manualOverride := input.GeoManualOverride || geoName != "" || geoLatitude != nil || geoLongitude != nil
|
||||
if !manualOverride || (geoLatitude == nil && geoLongitude == nil && geoName == "") {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
|
||||
}
|
||||
|
||||
func validateNodeIPInput(input Input, ip string) error {
|
||||
if len(ip) > maxNodeIPLength {
|
||||
return fmt.Errorf("%s", errNodeIPTooLong)
|
||||
}
|
||||
if ip != "" && net.ParseIP(ip) == nil {
|
||||
return fmt.Errorf("%s", errNodeIPInvalid)
|
||||
}
|
||||
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
|
||||
return fmt.Errorf("%s", errNodeIPManualRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateNodeGeoCoordinates(geoLatitude, geoLongitude *float64) error {
|
||||
if (geoLatitude == nil) != (geoLongitude == nil) {
|
||||
return fmt.Errorf("%s", errNodeGeoCoordinateMismatch)
|
||||
}
|
||||
if geoLatitude != nil && (*geoLatitude < -90 || *geoLatitude > 90) {
|
||||
return fmt.Errorf("%s", errNodeGeoLatitudeInvalid)
|
||||
}
|
||||
if geoLongitude != nil && (*geoLongitude < -180 || *geoLongitude > 180) {
|
||||
return fmt.Errorf("%s", errNodeGeoLongitudeInvalid)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func computeNodeStatus(node *model.OpenFlareNode) string {
|
||||
if node == nil {
|
||||
return nodeStatusOffline
|
||||
}
|
||||
if node.LastSeenAt == nil || node.LastSeenAt.IsZero() {
|
||||
return nodeStatusPending
|
||||
}
|
||||
// 默认离线阈值 60 秒(与 node_offline_threshold 默认一致),避免在这里读取配置
|
||||
// 实际阈值会在需要精确判断的地方通过 getNodeOfflineThreshold 读取
|
||||
threshold := 60 * time.Second
|
||||
if time.Since(*node.LastSeenAt) > threshold {
|
||||
return nodeStatusOffline
|
||||
}
|
||||
return nodeStatusOnline
|
||||
}
|
||||
|
||||
func nodeViewLastSeenAt(node *model.OpenFlareNode) any {
|
||||
if node == nil {
|
||||
return time.Time{}
|
||||
}
|
||||
nodeType := strings.TrimSpace(node.NodeType)
|
||||
if nodeType == "" {
|
||||
nodeType = nodeTypeEdgeNode
|
||||
}
|
||||
if nodeType == nodeTypeTunnelRelay && ofws.IsRelayConnected(node.NodeID) {
|
||||
return ofws.RelayWSConnectedLastSeenValue
|
||||
}
|
||||
if nodeType == nodeTypeTunnelClient && ofws.IsFlaredConnected(node.NodeID) {
|
||||
return ofws.FlaredWSConnectedLastSeenValue
|
||||
}
|
||||
if ofws.IsAgentConnected(node.NodeID) {
|
||||
return ofws.AgentWSConnectedLastSeenValue
|
||||
}
|
||||
if node.LastSeenAt == nil {
|
||||
return time.Time{}
|
||||
}
|
||||
return *node.LastSeenAt
|
||||
}
|
||||
|
||||
func buildNodeView(node *model.OpenFlareNode) *View {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
status := computeNodeStatus(node)
|
||||
view := &View{
|
||||
ID: node.ID,
|
||||
NodeID: node.NodeID,
|
||||
Name: node.Name,
|
||||
IP: node.IP,
|
||||
IPManualOverride: node.IPManualOverride,
|
||||
GeoName: strings.TrimSpace(node.GeoName),
|
||||
GeoLatitude: node.GeoLatitude,
|
||||
GeoLongitude: node.GeoLongitude,
|
||||
GeoManualOverride: node.GeoManualOverride,
|
||||
AccessToken: node.AccessToken,
|
||||
UpdateChannel: strings.TrimSpace(node.UpdateChannel),
|
||||
UpdateTag: strings.TrimSpace(node.UpdateTag),
|
||||
RestartOpenrestyRequested: node.RestartOpenrestyRequested,
|
||||
Version: node.Version,
|
||||
ExtVersion: node.ExtVersion,
|
||||
OpenrestyStatus: normalizeOpenrestyStatus(node.OpenrestyStatus),
|
||||
OpenrestyMessage: strings.TrimSpace(node.OpenrestyMessage),
|
||||
Status: status,
|
||||
CurrentVersion: node.CurrentVersion,
|
||||
LastSeenAt: nodeViewLastSeenAt(node),
|
||||
LastError: node.LastError,
|
||||
CreatedAt: node.CreatedAt,
|
||||
UpdatedAt: node.UpdatedAt,
|
||||
AutoUpdateEnabled: node.AutoUpdateEnabled,
|
||||
UpdateRequested: node.UpdateRequested,
|
||||
NodeType: node.NodeType,
|
||||
RelayBindPort: node.RelayBindPort,
|
||||
RelayVhostHTTPPort: node.RelayVhostHTTPPort,
|
||||
RelayAgentAccessAddr: node.RelayAgentAccessAddr,
|
||||
RelayClientAccessAddr: node.RelayClientAccessAddr,
|
||||
RelayClientProxyURL: node.RelayClientProxyURL,
|
||||
RelayStatus: node.RelayStatus,
|
||||
RelayWebServerEnabled: node.RelayWebServerEnabled,
|
||||
}
|
||||
if view.UpdateChannel == "" {
|
||||
view.UpdateChannel = releaseChannelStable.String()
|
||||
}
|
||||
if view.NodeType == "" {
|
||||
view.NodeType = nodeTypeEdgeNode
|
||||
}
|
||||
return view
|
||||
}
|
||||
|
||||
func buildNodeAgentReleaseView(node *model.OpenFlareNode, release *githubReleaseResponse, channel releaseChannel) *AgentReleaseInfo {
|
||||
currentVersion := strings.TrimSpace(node.Version)
|
||||
view := &AgentReleaseInfo{
|
||||
CurrentVersion: currentVersion,
|
||||
Channel: channel.String(),
|
||||
UpdateRequested: node.UpdateRequested,
|
||||
RequestedChannel: normalizeReleaseChannel(node.UpdateChannel).String(),
|
||||
RequestedTag: strings.TrimSpace(node.UpdateTag),
|
||||
}
|
||||
if release == nil {
|
||||
return view
|
||||
}
|
||||
view.TagName = release.TagName
|
||||
view.Body = release.Body
|
||||
view.HTMLURL = release.HTMLURL
|
||||
view.PublishedAt = release.PublishedAt
|
||||
view.Prerelease = release.Prerelease
|
||||
view.HasUpdate = isVersionNewer(currentVersion, release.TagName)
|
||||
return view
|
||||
}
|
||||
|
||||
func isVersionNewer(current string, latest string) bool {
|
||||
return compareVersions(current, latest) < 0
|
||||
}
|
||||
|
||||
func fetchLatestGitHubRelease(ctx context.Context, repo string, channel releaseChannel) (*githubReleaseResponse, error) {
|
||||
switch normalizeReleaseChannel(string(channel)) {
|
||||
case releaseChannelPreview:
|
||||
return fetchLatestPreviewGitHubRelease(ctx, repo)
|
||||
case releaseChannelStable:
|
||||
return fetchLatestStableGitHubRelease(ctx, repo)
|
||||
default:
|
||||
return fetchLatestStableGitHubRelease(ctx, repo)
|
||||
}
|
||||
}
|
||||
|
||||
func fetchLatestStableGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) {
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"/latest", strings.TrimSpace(repo))
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, errors.New("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取最新版本失败: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
return decodeGitHubRelease(resp.Body)
|
||||
}
|
||||
|
||||
func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubReleaseResponse, error) {
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"?per_page=20", strings.TrimSpace(repo))
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, errors.New("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取 preview 版本失败: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
var releases []githubReleaseResponse
|
||||
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
return nil, errors.New("解析 preview 版本信息失败")
|
||||
}
|
||||
for _, release := range releases {
|
||||
if release.Draft || !release.Prerelease {
|
||||
continue
|
||||
}
|
||||
releaseCopy := release
|
||||
return &releaseCopy, nil
|
||||
}
|
||||
return nil, errors.New("当前没有可用的 preview 发布")
|
||||
}
|
||||
|
||||
func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*githubReleaseResponse, error) {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
return nil, errors.New("缺少发布版本号")
|
||||
}
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"/tags/%s", strings.TrimSpace(repo), tag)
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, errors.New("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取指定版本失败: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return nil, fmt.Errorf("未找到指定版本: %s", tag)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
|
||||
}
|
||||
return decodeGitHubRelease(resp.Body)
|
||||
}
|
||||
|
||||
func newGitHubReleaseRequest(ctx context.Context, url string) (*http.Request, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
req.Header.Set("User-Agent", "OpenFlare-Server")
|
||||
return req, nil
|
||||
}
|
||||
|
||||
func decodeGitHubRelease(reader io.Reader) (*githubReleaseResponse, error) {
|
||||
var release githubReleaseResponse
|
||||
if err := json.NewDecoder(reader).Decode(&release); err != nil {
|
||||
return nil, errors.New("解析版本信息失败")
|
||||
}
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
@@ -0,0 +1,428 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
cf "Wavelet/openflare/plugins/server/domain/cloudflare"
|
||||
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
|
||||
"Wavelet/openflare/plugins/server/domain/observability"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
"Wavelet/pkg/logger"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultRelayBindPort = 7000
|
||||
defaultRelayVhostHTTPPort = 8080
|
||||
)
|
||||
|
||||
// getAgentUpdateRepo 从 SystemConfig 读取 Agent 更新仓库配置
|
||||
func getAgentUpdateRepo(ctx context.Context) string {
|
||||
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentUpdateRepo)
|
||||
if err != nil || strings.TrimSpace(config.Value) == "" {
|
||||
return "Rain-kl/OpenFlare" // 默认值
|
||||
}
|
||||
return strings.TrimSpace(config.Value)
|
||||
}
|
||||
|
||||
// Input is the create/update node payload.
|
||||
type Input struct {
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
IPManualOverride *bool `json:"ip_manual_override"`
|
||||
AutoUpdateEnabled bool `json:"auto_update_enabled"`
|
||||
GeoName string `json:"geo_name"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
GeoManualOverride bool `json:"geo_manual_override"`
|
||||
NodeType string `json:"node_type"`
|
||||
RelayBindPort int `json:"relay_bind_port"`
|
||||
RelayVhostHTTPPort int `json:"relay_vhost_http_port"`
|
||||
RelayAgentAccessAddr string `json:"relay_agent_access_addr"`
|
||||
RelayClientAccessAddr string `json:"relay_client_access_addr"`
|
||||
RelayClientProxyURL string `json:"relay_client_proxy_url"`
|
||||
RelayWebServerEnabled bool `json:"relay_web_server_enabled"`
|
||||
}
|
||||
|
||||
// AgentUpdateInput requests an agent self-update on a node.
|
||||
type AgentUpdateInput struct {
|
||||
Channel string `json:"channel"`
|
||||
TagName string `json:"tag_name"`
|
||||
}
|
||||
|
||||
// AgentReleaseInfo describes the latest agent release for a node.
|
||||
type AgentReleaseInfo struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Body string `json:"body"`
|
||||
HTMLURL string `json:"html_url"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
HasUpdate bool `json:"has_update"`
|
||||
Channel string `json:"channel"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
UpdateRequested bool `json:"update_requested"`
|
||||
RequestedChannel string `json:"requested_channel"`
|
||||
RequestedTag string `json:"requested_tag"`
|
||||
}
|
||||
|
||||
// BootstrapView exposes the global discovery token.
|
||||
type BootstrapView struct {
|
||||
DiscoveryToken string `json:"discovery_token"`
|
||||
}
|
||||
|
||||
// View is the admin-facing node representation.
|
||||
type View struct {
|
||||
ID uint `json:"id"`
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
IPManualOverride bool `json:"ip_manual_override"`
|
||||
GeoName string `json:"geo_name"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
GeoManualOverride bool `json:"geo_manual_override"`
|
||||
AccessToken string `json:"access_token"`
|
||||
AutoUpdateEnabled bool `json:"auto_update_enabled"`
|
||||
UpdateRequested bool `json:"update_requested"`
|
||||
UpdateChannel string `json:"update_channel"`
|
||||
UpdateTag string `json:"update_tag"`
|
||||
RestartOpenrestyRequested bool `json:"restart_openresty_requested"`
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"ext_version"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
OpenrestyMessage string `json:"openresty_message"`
|
||||
Status string `json:"status"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastSeenAt any `json:"last_seen_at"`
|
||||
LastError string `json:"last_error"`
|
||||
LatestApplyResult string `json:"latest_apply_result"`
|
||||
LatestApplyMessage string `json:"latest_apply_message"`
|
||||
LatestApplyChecksum string `json:"latest_apply_checksum"`
|
||||
LatestMainConfigChecksum string `json:"latest_main_config_checksum"`
|
||||
LatestRouteConfigChecksum string `json:"latest_route_config_checksum"`
|
||||
LatestSupportFileCount int `json:"latest_support_file_count"`
|
||||
LatestApplyAt *time.Time `json:"latest_apply_at"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
NodeType string `json:"node_type"`
|
||||
RelayBindPort int `json:"relay_bind_port"`
|
||||
RelayVhostHTTPPort int `json:"relay_vhost_http_port"`
|
||||
RelayAgentAccessAddr string `json:"relay_agent_access_addr"`
|
||||
RelayClientAccessAddr string `json:"relay_client_access_addr"`
|
||||
RelayClientProxyURL string `json:"relay_client_proxy_url"`
|
||||
RelayStatus string `json:"relay_status"`
|
||||
RelayWebServerEnabled bool `json:"relay_web_server_enabled"`
|
||||
}
|
||||
|
||||
// ObservabilityQuery filters node observability data.
|
||||
type ObservabilityQuery struct {
|
||||
Hours int `json:"hours"`
|
||||
Limit int `json:"limit"`
|
||||
}
|
||||
|
||||
// ObservabilityView is the node observability API response.
|
||||
type ObservabilityView = observability.NodeView
|
||||
|
||||
// HealthEventCleanupResult reports health event cleanup outcome.
|
||||
type HealthEventCleanupResult = observability.HealthEventCleanupResult
|
||||
|
||||
// ListNodes returns all node views with latest apply log metadata.
|
||||
func ListNodes(ctx context.Context) ([]*View, error) {
|
||||
nodes, err := repository.ListOpenFlareNodes(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeIDs := make([]string, 0, len(nodes))
|
||||
for _, node := range nodes {
|
||||
nodeIDs = append(nodeIDs, node.NodeID)
|
||||
}
|
||||
latestLogs, err := repository.GetLatestOpenFlareApplyLogsByNodeIDs(ctx, nodeIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]*View, 0, len(nodes))
|
||||
for _, node := range nodes {
|
||||
view := buildNodeView(&node)
|
||||
view.Status = computeNodeStatus(&node)
|
||||
if log, ok := latestLogs[node.NodeID]; ok {
|
||||
view.LatestApplyResult = log.Result
|
||||
view.LatestApplyMessage = log.Message
|
||||
view.LatestApplyChecksum = log.Checksum
|
||||
view.LatestMainConfigChecksum = log.MainConfigChecksum
|
||||
view.LatestRouteConfigChecksum = log.RouteConfigChecksum
|
||||
view.LatestSupportFileCount = log.SupportFileCount
|
||||
view.LatestApplyAt = &log.CreatedAt
|
||||
}
|
||||
views = append(views, view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// CreateNode creates a reserved node with generated node_id and access_token.
|
||||
func CreateNode(ctx context.Context, input Input) (*View, error) {
|
||||
name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if name == "" {
|
||||
return nil, errors.New(errNodeNameRequired)
|
||||
}
|
||||
ipManualOverride := resolveNodeIPManualOverride(input, nil, ip)
|
||||
node := &model.OpenFlareNode{
|
||||
Name: name,
|
||||
IP: ip,
|
||||
IPManualOverride: ipManualOverride,
|
||||
GeoName: geoName,
|
||||
GeoLatitude: geoLatitude,
|
||||
GeoLongitude: geoLongitude,
|
||||
GeoManualOverride: geoManualOverride,
|
||||
Version: "",
|
||||
ExtVersion: "",
|
||||
Status: nodeStatusPending,
|
||||
AutoUpdateEnabled: input.AutoUpdateEnabled,
|
||||
NodeType: normalizeNodeType(input.NodeType),
|
||||
CapabilitiesJSON: "[]",
|
||||
}
|
||||
node.NodeID, err = newServerNodeID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.AccessToken, err = newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, defaultRelayBindPort)
|
||||
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, defaultRelayVhostHTTPPort)
|
||||
node.RelayAuthToken, err = newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr)
|
||||
node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr)
|
||||
node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL)
|
||||
node.RelayWebServerEnabled = input.RelayWebServerEnabled
|
||||
}
|
||||
if err = repository.CreateOpenFlareNode(ctx, node); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errNodeIDConflict)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
// UpdateNode updates an existing node.
|
||||
func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) {
|
||||
name, ip, geoName, geoLatitude, geoLongitude, geoManualOverride, err := normalizeNodeInput(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if name == "" {
|
||||
return nil, errors.New(errNodeNameRequired)
|
||||
}
|
||||
node, err := repository.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ipManualOverride := resolveNodeIPManualOverride(input, node, ip)
|
||||
previousIP := node.IP
|
||||
node.Name = name
|
||||
node.IP = ip
|
||||
node.IPManualOverride = ipManualOverride
|
||||
node.GeoName = geoName
|
||||
node.GeoLatitude = geoLatitude
|
||||
node.GeoLongitude = geoLongitude
|
||||
node.GeoManualOverride = geoManualOverride
|
||||
node.AutoUpdateEnabled = input.AutoUpdateEnabled
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
node.RelayAgentAccessAddr = strings.TrimSpace(input.RelayAgentAccessAddr)
|
||||
node.RelayClientAccessAddr = strings.TrimSpace(input.RelayClientAccessAddr)
|
||||
node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL)
|
||||
node.RelayWebServerEnabled = input.RelayWebServerEnabled
|
||||
if input.RelayBindPort > 0 {
|
||||
node.RelayBindPort = input.RelayBindPort
|
||||
}
|
||||
if input.RelayVhostHTTPPort > 0 {
|
||||
node.RelayVhostHTTPPort = input.RelayVhostHTTPPort
|
||||
}
|
||||
}
|
||||
if err = repository.SaveOpenFlareNode(ctx, node); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.TrimSpace(previousIP) != strings.TrimSpace(node.IP) {
|
||||
if _, dispatchErr := cf.DispatchNodeSync(ctx, node.ID, "cloudflare_node_ip_update"); dispatchErr != nil {
|
||||
logger.ErrorF(ctx, "[Cloudflare] enqueue node sync failed: node_id=%d error=%v", node.ID, dispatchErr)
|
||||
}
|
||||
}
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
// DeleteNode removes a node by id.
|
||||
func DeleteNode(ctx context.Context, id uint) error {
|
||||
if _, err := repository.GetOpenFlareNodeByID(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return repository.DeleteOpenFlareNode(ctx, id)
|
||||
}
|
||||
|
||||
// GetBootstrapToken returns the global discovery token, creating one if missing.
|
||||
func GetBootstrapToken(ctx context.Context) (*BootstrapView, error) {
|
||||
token, err := ensureGlobalDiscoveryToken(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &BootstrapView{DiscoveryToken: token}, nil
|
||||
}
|
||||
|
||||
// RotateBootstrapToken rotates the global discovery token.
|
||||
func RotateBootstrapToken(ctx context.Context) (*BootstrapView, error) {
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyAgentDiscoveryToken, token); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &BootstrapView{DiscoveryToken: token}, nil
|
||||
}
|
||||
|
||||
// GetAgentRelease checks the latest agent release for a node.
|
||||
func GetAgentRelease(ctx context.Context, id uint, channel string) (*AgentReleaseInfo, error) {
|
||||
node, err := repository.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
release, err := fetchLatestGitHubRelease(ctx, getAgentUpdateRepo(ctx), normalizeReleaseChannel(channel))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildNodeAgentReleaseView(node, release, normalizeReleaseChannel(channel)), nil
|
||||
}
|
||||
|
||||
// RequestAgentUpdate marks a node for manual agent update.
|
||||
func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (*View, error) {
|
||||
node, err := repository.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
channel := normalizeReleaseChannel(input.Channel)
|
||||
tagName := strings.TrimSpace(input.TagName)
|
||||
if tagName != "" {
|
||||
release, releaseErr := fetchGitHubReleaseByTag(ctx, getAgentUpdateRepo(ctx), tagName)
|
||||
if releaseErr != nil {
|
||||
return nil, releaseErr
|
||||
}
|
||||
if channel == releaseChannelPreview && !release.Prerelease {
|
||||
return nil, errors.New(errAgentPreviewTagInvalid)
|
||||
}
|
||||
if channel == releaseChannelStable && release.Prerelease {
|
||||
return nil, errors.New(errAgentStableTagInvalid)
|
||||
}
|
||||
}
|
||||
node.UpdateRequested = true
|
||||
node.UpdateChannel = channel.String()
|
||||
node.UpdateTag = tagName
|
||||
if err = repository.UpdateOpenFlareNodeFields(ctx, node, "update_requested", "update_channel", "update_tag"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
// RequestOpenrestyRestart marks a node for openresty restart.
|
||||
func RequestOpenrestyRestart(ctx context.Context, id uint) (*View, error) {
|
||||
node, err := repository.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.RestartOpenrestyRequested = true
|
||||
if err = repository.UpdateOpenFlareNodeFields(ctx, node, "restart_openresty_requested"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
// RequestForceSync pushes force_sync_config to a connected agent websocket.
|
||||
func RequestForceSync(ctx context.Context, id uint) (*View, error) {
|
||||
node, err := repository.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
activeConfig, err := repository.GetActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, fmt.Errorf("无法获取当前激活的配置版本:%s", errNoActiveConfigVersion)
|
||||
}
|
||||
return nil, fmt.Errorf("无法获取当前激活的配置版本:%w", err)
|
||||
}
|
||||
if !ofws.SendForceSyncConfig(node.NodeID, forceSyncConfigPayload{
|
||||
Version: activeConfig.Version,
|
||||
Checksum: activeConfig.Checksum,
|
||||
}) {
|
||||
return nil, errors.New(errNodeForceSyncFailed)
|
||||
}
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
// GetObservability returns observability details for a node.
|
||||
func GetObservability(ctx context.Context, id uint, query ObservabilityQuery) (*ObservabilityView, error) {
|
||||
return observability.GetNodeObservability(ctx, id, observability.NodeQuery{
|
||||
Hours: query.Hours,
|
||||
Limit: query.Limit,
|
||||
})
|
||||
}
|
||||
|
||||
// CleanupHealthEvents removes all health events for a node.
|
||||
func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) {
|
||||
return observability.CleanupHealthEvents(ctx, id)
|
||||
}
|
||||
|
||||
func ensureGlobalDiscoveryToken(ctx context.Context) (string, error) {
|
||||
// 从 SystemConfig 读取 Agent 发现令牌
|
||||
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentDiscoveryToken)
|
||||
if err == nil && strings.TrimSpace(config.Value) != "" {
|
||||
return strings.TrimSpace(config.Value), nil
|
||||
}
|
||||
|
||||
// 如果不存在,生成新令牌并保存
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 更新到 SystemConfig
|
||||
if err = repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyAgentDiscoveryToken, token); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
type forceSyncConfigPayload struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// ValidateDiscoveryToken validates the global discovery token.
|
||||
func ValidateDiscoveryToken(ctx context.Context, token string) error {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return errors.New("缺少 Discovery Token")
|
||||
}
|
||||
discoveryToken, err := ensureGlobalDiscoveryToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !tokenEqual(token, discoveryToken) {
|
||||
return errors.New("discovery Token 无效") // error 消息首字母小写
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,383 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
cf "Wavelet/openflare/plugins/server/domain/cloudflare"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setReleaseHTTPClientForTest(client *http.Client) *http.Client {
|
||||
previous := releaseHTTPClient
|
||||
releaseHTTPClient = client
|
||||
return previous
|
||||
}
|
||||
|
||||
func setupNodeTestDB(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{},
|
||||
&model.OpenFlareApplyLog{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
testhelper.SetupLogStoresForTest(t)
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateEdgeNode(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
view, err := CreateNode(ctx, Input{
|
||||
Name: "edge-1",
|
||||
IP: "10.0.0.1",
|
||||
AutoUpdateEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, view.ID)
|
||||
assert.True(t, strings.HasPrefix(view.NodeID, "node-"))
|
||||
assert.Len(t, view.AccessToken, 32)
|
||||
assert.Equal(t, "edge_node", view.NodeType)
|
||||
assert.Equal(t, nodeStatusPending, view.Status)
|
||||
assert.True(t, view.AutoUpdateEnabled)
|
||||
}
|
||||
|
||||
func TestCreateTunnelRelayNode(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
view, err := CreateNode(ctx, Input{
|
||||
Name: "relay-1",
|
||||
NodeType: "tunnel_relay",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "tunnel_relay", view.NodeType)
|
||||
assert.Equal(t, 7000, view.RelayBindPort)
|
||||
assert.Equal(t, 8080, view.RelayVhostHTTPPort)
|
||||
|
||||
stored, err := repository.GetOpenFlareNodeByID(ctx, view.ID)
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, stored.RelayAuthToken)
|
||||
}
|
||||
|
||||
func TestCreateTunnelClientNode(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
view, err := CreateNode(ctx, Input{
|
||||
Name: "client-1",
|
||||
NodeType: "tunnel_client",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "tunnel_client", view.NodeType)
|
||||
}
|
||||
|
||||
func TestCreateNodeRequiresName(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := CreateNode(ctx, Input{IP: "10.0.0.2"})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errNodeNameRequired, err.Error())
|
||||
}
|
||||
|
||||
func TestUpdateNode(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-update"})
|
||||
require.NoError(t, err)
|
||||
|
||||
updated, err := UpdateNode(ctx, created.ID, Input{
|
||||
Name: "edge-updated",
|
||||
IP: "192.168.1.10",
|
||||
AutoUpdateEnabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "edge-updated", updated.Name)
|
||||
assert.Equal(t, "192.168.1.10", updated.IP)
|
||||
assert.True(t, updated.AutoUpdateEnabled)
|
||||
}
|
||||
|
||||
func TestUpdateNodeDispatchesCloudflareSyncWhenIPChanges(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-update", IP: "192.0.2.10"})
|
||||
require.NoError(t, err)
|
||||
|
||||
var dispatchedNodeID uint
|
||||
restore := cf.SetDispatchTaskForTest(func(_ context.Context, taskType string, payload []byte, _ string) (string, error) {
|
||||
assert.Equal(t, cf.TaskTypeSyncByNode, taskType)
|
||||
assert.Contains(t, string(payload), `"node_id":`)
|
||||
dispatchedNodeID = created.ID
|
||||
return "task-1", nil
|
||||
})
|
||||
defer restore()
|
||||
|
||||
_, err = UpdateNode(ctx, created.ID, Input{Name: "edge-update", IP: "192.0.2.11"})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, created.ID, dispatchedNodeID)
|
||||
}
|
||||
|
||||
func TestDeleteNode(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-delete"})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, DeleteNode(ctx, created.ID))
|
||||
_, err = repository.GetOpenFlareNodeByID(ctx, created.ID)
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, gorm.ErrRecordNotFound)
|
||||
}
|
||||
|
||||
func TestListNodesWithApplyLogMetadata(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-list"})
|
||||
require.NoError(t, err)
|
||||
|
||||
applyAt := time.Now().UTC().Truncate(time.Second)
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{
|
||||
NodeID: created.NodeID,
|
||||
Version: "20260618-001",
|
||||
Result: "success",
|
||||
Message: "ok",
|
||||
Checksum: "checksum-1",
|
||||
MainConfigChecksum: "main-1",
|
||||
RouteConfigChecksum: "route-1",
|
||||
SupportFileCount: 3,
|
||||
CreatedAt: applyAt,
|
||||
}).Error)
|
||||
|
||||
views, err := ListNodes(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, views, 1)
|
||||
assert.Equal(t, "success", views[0].LatestApplyResult)
|
||||
assert.Equal(t, "checksum-1", views[0].LatestApplyChecksum)
|
||||
assert.Equal(t, 3, views[0].LatestSupportFileCount)
|
||||
require.NotNil(t, views[0].LatestApplyAt)
|
||||
assert.Equal(t, applyAt, views[0].LatestApplyAt.UTC())
|
||||
}
|
||||
|
||||
func TestBootstrapTokenLifecycle(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
first, err := GetBootstrapToken(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, first.DiscoveryToken, 32)
|
||||
|
||||
second, err := GetBootstrapToken(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, first.DiscoveryToken, second.DiscoveryToken)
|
||||
|
||||
rotated, err := RotateBootstrapToken(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.NotEqual(t, first.DiscoveryToken, rotated.DiscoveryToken)
|
||||
// 验证令牌已保存到 SystemConfig
|
||||
savedToken, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyAgentDiscoveryToken)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, rotated.DiscoveryToken, savedToken.Value)
|
||||
}
|
||||
|
||||
func TestValidateDiscoveryToken(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
bootstrap, err := GetBootstrapToken(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, ValidateDiscoveryToken(ctx, bootstrap.DiscoveryToken))
|
||||
require.Error(t, ValidateDiscoveryToken(ctx, "invalid-token"))
|
||||
require.Error(t, ValidateDiscoveryToken(ctx, ""))
|
||||
require.Error(t, ValidateDiscoveryToken(ctx, bootstrap.DiscoveryToken[:len(bootstrap.DiscoveryToken)-1]+"x"))
|
||||
}
|
||||
|
||||
func TestRequestAgentUpdateWithPreviewTag(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-update-agent"})
|
||||
require.NoError(t, err)
|
||||
|
||||
originalClient := setReleaseHTTPClientForTest(&http.Client{
|
||||
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
expected := "https://api.github.com/repos/Rain-kl/OpenFlare/releases/tags/v0.5.0-rc.1"
|
||||
require.Equal(t, expected, req.URL.String())
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{"tag_name":"v0.5.0-rc.1","prerelease":true}`)),
|
||||
}, nil
|
||||
}),
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
setReleaseHTTPClientForTest(originalClient)
|
||||
})
|
||||
|
||||
updated, err := RequestAgentUpdate(ctx, created.ID, AgentUpdateInput{
|
||||
Channel: "preview",
|
||||
TagName: "v0.5.0-rc.1",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, updated.UpdateRequested)
|
||||
assert.Equal(t, "preview", updated.UpdateChannel)
|
||||
assert.Equal(t, "v0.5.0-rc.1", updated.UpdateTag)
|
||||
}
|
||||
|
||||
func TestRequestOpenrestyRestart(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-restart"})
|
||||
require.NoError(t, err)
|
||||
|
||||
updated, err := RequestOpenrestyRestart(ctx, created.ID)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, updated.RestartOpenrestyRequested)
|
||||
}
|
||||
|
||||
func seedActiveConfigVersion(t *testing.T, ctx context.Context) {
|
||||
t.Helper()
|
||||
conn := db.DB(ctx)
|
||||
require.NotNil(t, conn)
|
||||
require.NoError(t, conn.AutoMigrate(&model.ConfigVersion{}))
|
||||
require.NoError(t, conn.Create(&model.ConfigVersion{
|
||||
Version: "20260618-001",
|
||||
SnapshotJSON: `{}`,
|
||||
RenderedConfig: `server {}`,
|
||||
Checksum: "abc123",
|
||||
IsActive: true,
|
||||
CreatedBy: "test",
|
||||
}).Error)
|
||||
}
|
||||
|
||||
func TestRequestForceSyncRequiresWebSocket(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
seedActiveConfigVersion(t, ctx)
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-sync"})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = RequestForceSync(ctx, created.ID)
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errNodeForceSyncFailed, err.Error())
|
||||
}
|
||||
|
||||
func TestRequestForceSyncRequiresActiveConfig(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
conn := db.DB(ctx)
|
||||
require.NotNil(t, conn)
|
||||
require.NoError(t, conn.AutoMigrate(&model.ConfigVersion{}))
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-sync-active"})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = RequestForceSync(ctx, created.ID)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), errNoActiveConfigVersion)
|
||||
}
|
||||
|
||||
func TestGetObservabilityStub(t *testing.T) {
|
||||
cleanup := setupNodeTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := CreateNode(ctx, Input{Name: "edge-obs"})
|
||||
require.NoError(t, err)
|
||||
|
||||
view, err := GetObservability(ctx, created.ID, ObservabilityQuery{Hours: 24, Limit: 50})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, created.NodeID, view.NodeID)
|
||||
assert.Empty(t, view.MetricSnapshots)
|
||||
}
|
||||
|
||||
func TestComputeNodeStatus(t *testing.T) {
|
||||
now := time.Now()
|
||||
pending := &model.OpenFlareNode{}
|
||||
assert.Equal(t, nodeStatusPending, computeNodeStatus(pending))
|
||||
|
||||
online := &model.OpenFlareNode{LastSeenAt: &now}
|
||||
assert.Equal(t, nodeStatusOnline, computeNodeStatus(online))
|
||||
|
||||
// computeNodeStatus 使用默认阈值 60 秒
|
||||
offlineAt := now.Add(-61 * time.Second)
|
||||
offline := &model.OpenFlareNode{LastSeenAt: &offlineAt}
|
||||
assert.Equal(t, nodeStatusOffline, computeNodeStatus(offline))
|
||||
}
|
||||
|
||||
type roundTripFunc func(req *http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
func TestCompareVersions(t *testing.T) {
|
||||
tests := []struct {
|
||||
local string
|
||||
remote string
|
||||
expected int
|
||||
}{
|
||||
{"v3.0.0-beta", "v3.0.0-beta.1", -1},
|
||||
{"v3.0.0-beta", "v3.0.0", -1},
|
||||
{"v3.0.0-beta.1", "v3.0.0", -1},
|
||||
{"dev", "v3.0.0", -1},
|
||||
{"v3.0.0", "v3.0.0", 0},
|
||||
{"v3.0.0", "v2.9.9", 1},
|
||||
{"v3.0.0", "v3.0.1", -1},
|
||||
{"v3.0.0-beta.1", "v3.0.0-beta.2", -1},
|
||||
{"v3.0.0-beta.11", "v3.0.0-beta.2", 1},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.local+"_vs_"+tt.remote, func(t *testing.T) {
|
||||
res := compareVersions(tt.local, tt.remote)
|
||||
assert.Equal(t, tt.expected, res)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,324 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/apiutil"
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return apiutil.AbortNotFoundIfMissing(c, err, errNodeNotFound)
|
||||
}
|
||||
|
||||
// ListNodesHandler lists all nodes.
|
||||
// @Summary 获取节点列表
|
||||
// @Description 返回所有节点及最新配置下发记录,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]node.View} "节点列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/nodes [get]
|
||||
func ListNodesHandler(c *gin.Context) {
|
||||
nodes, err := ListNodes(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(nodes))
|
||||
}
|
||||
|
||||
// CreateNodeHandler creates a node.
|
||||
// @Summary 创建节点
|
||||
// @Description 创建新的边缘节点记录,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param body body node.Input true "节点参数"
|
||||
// @Success 200 {object} response.Any{data=node.View} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/nodes [post]
|
||||
func CreateNodeHandler(c *gin.Context) {
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
view, err := CreateNode(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// UpdateNodeHandler updates a node.
|
||||
// @Summary 更新节点
|
||||
// @Description 更新指定节点的配置信息,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Param body body node.Input true "节点参数"
|
||||
// @Success 200 {object} response.Any{data=node.View} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/update [post]
|
||||
func UpdateNodeHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
view, err := UpdateNode(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// DeleteNodeHandler deletes a node.
|
||||
// @Summary 删除节点
|
||||
// @Description 删除指定节点记录,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/delete [post]
|
||||
func DeleteNodeHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteNode(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// GetBootstrapTokenHandler returns the global discovery token.
|
||||
// @Summary 获取引导令牌
|
||||
// @Description 返回全局节点发现引导令牌,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=node.BootstrapView} "引导令牌"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/nodes/bootstrap-token [get]
|
||||
func GetBootstrapTokenHandler(c *gin.Context) {
|
||||
view, err := GetBootstrapToken(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// RotateBootstrapTokenHandler rotates the global discovery token.
|
||||
// @Summary 轮换引导令牌
|
||||
// @Description 重新生成全局节点发现引导令牌,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=node.BootstrapView} "新引导令牌"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/nodes/bootstrap-token/rotate [post]
|
||||
func RotateBootstrapTokenHandler(c *gin.Context) {
|
||||
view, err := RotateBootstrapToken(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// GetAgentReleaseHandler returns the latest agent release for a node.
|
||||
// @Summary 获取 Agent 发布信息
|
||||
// @Description 返回指定节点可用的最新 Agent 版本信息,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Param channel query string false "发布渠道"
|
||||
// @Success 200 {object} response.Any{data=node.AgentReleaseInfo} "Agent 发布信息"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/agent-release [get]
|
||||
func GetAgentReleaseHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
release, err := GetAgentRelease(c.Request.Context(), id, c.Query("channel"))
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(release))
|
||||
}
|
||||
|
||||
// RequestAgentUpdateHandler requests agent self-update on a node.
|
||||
// @Summary 请求 Agent 更新
|
||||
// @Description 向指定节点下发 Agent 自更新指令,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Param body body node.AgentUpdateInput false "更新参数(可选)"
|
||||
// @Success 200 {object} response.Any{data=node.View} "更新请求已下发"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/agent-update [post]
|
||||
func RequestAgentUpdateHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var request AgentUpdateInput
|
||||
if c.Request.ContentLength > 0 {
|
||||
if err := bindOptionalJSON(c.Request.Body, &request); err != nil {
|
||||
response.AbortBadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
}
|
||||
view, err := RequestAgentUpdate(c.Request.Context(), id, request)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// RequestOpenrestyRestartHandler requests openresty restart on a node.
|
||||
// @Summary 请求重启 OpenResty
|
||||
// @Description 向指定节点下发 OpenResty 重启指令,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Success 200 {object} response.Any{data=node.View} "重启请求已下发"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/openresty-restart [post]
|
||||
func RequestOpenrestyRestartHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
view, err := RequestOpenrestyRestart(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// RequestForceSyncHandler requests force sync on a node.
|
||||
// @Summary 请求强制同步配置
|
||||
// @Description 向指定节点下发强制同步当前活跃配置的指令,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Success 200 {object} response.Any{data=node.View} "同步请求已下发"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/force-sync [post]
|
||||
func RequestForceSyncHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
view, err := RequestForceSync(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// GetObservabilityHandler returns node observability details.
|
||||
// @Summary 获取节点可观测性数据
|
||||
// @Description 返回指定节点的指标、健康事件与流量分析数据,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Param hours query int false "统计时间范围(小时)"
|
||||
// @Param limit query int false "返回记录数量上限"
|
||||
// @Success 200 {object} response.Any{data=node.ObservabilityView} "可观测性数据"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/observability [get]
|
||||
func GetObservabilityHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var query ObservabilityQuery
|
||||
if err := c.ShouldBindQuery(&query); err != nil {
|
||||
response.AbortBadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
view, err := GetObservability(c.Request.Context(), id, query)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// CleanupHealthEventsHandler cleans up node health events.
|
||||
// @Summary 清理节点健康事件
|
||||
// @Description 清理指定节点的历史健康事件记录,需要管理员权限
|
||||
// @Tags openflare-node
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "节点 ID"
|
||||
// @Success 200 {object} response.Any{data=node.HealthEventCleanupResult} "清理结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或节点不存在"
|
||||
// @Router /api/v1/d/nodes/{id}/observability/cleanup [post]
|
||||
func CleanupHealthEventsHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
result, err := CleanupHealthEvents(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
func bindOptionalJSON(body io.Reader, target any) error {
|
||||
if err := json.NewDecoder(body).Decode(target); err != nil && !errors.Is(err, io.EOF) {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import "Wavelet/openflare/share/ofutil"
|
||||
|
||||
func compareVersions(local, remote string) int {
|
||||
return ofutil.CompareVersions(local, remote)
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package relay provides relay node management and authentication for the OpenFlare platform.
|
||||
package relay
|
||||
|
||||
const (
|
||||
//nolint:gosec // error message text, not a credential
|
||||
errAgentTokenInvalid = "无权进行此操作,Agent Token 无效"
|
||||
errRelayNodeTypeMismatch = "此节点不是 TunnelRelay 类型"
|
||||
)
|
||||
@@ -0,0 +1,139 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
)
|
||||
|
||||
const (
|
||||
relayStatusUnhealthy = "unhealthy"
|
||||
releaseChannelStable = "stable"
|
||||
|
||||
defaultAgentHeartbeatInterval = 3000 // 默认心跳间隔 3 秒(毫秒)
|
||||
defaultAgentUpdateRepo = "Rain-kl/OpenFlare"
|
||||
)
|
||||
|
||||
func normalizeRelayStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case "healthy":
|
||||
return "healthy"
|
||||
case relayStatusUnhealthy:
|
||||
return relayStatusUnhealthy
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeReleaseChannel(channel string) string {
|
||||
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
|
||||
return "preview"
|
||||
}
|
||||
return releaseChannelStable
|
||||
}
|
||||
|
||||
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
|
||||
reported := normalizeNodeIP(reportedIP)
|
||||
remote := normalizeRemoteAddr(remoteAddr)
|
||||
if reported == "" {
|
||||
return remote
|
||||
}
|
||||
if isPublicNodeIP(reported) {
|
||||
return reported
|
||||
}
|
||||
if isPublicNodeIP(remote) {
|
||||
return remote
|
||||
}
|
||||
return reported
|
||||
}
|
||||
|
||||
func normalizeNodeIP(raw string) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
if host, _, err := net.SplitHostPort(raw); err == nil {
|
||||
raw = host
|
||||
}
|
||||
raw = strings.Trim(raw, "[]")
|
||||
return raw
|
||||
}
|
||||
|
||||
func normalizeRemoteAddr(remoteAddr string) string {
|
||||
remoteAddr = strings.TrimSpace(remoteAddr)
|
||||
if remoteAddr == "" {
|
||||
return ""
|
||||
}
|
||||
host, _, err := net.SplitHostPort(remoteAddr)
|
||||
if err != nil {
|
||||
return normalizeNodeIP(remoteAddr)
|
||||
}
|
||||
return normalizeNodeIP(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 buildRelayConfig(ctx context.Context, node *model.OpenFlareNode) *Config {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
webServerPort, err := repository.GetIntByKey(ctx, model.ConfigKeyRelayFRPSWebUIPort)
|
||||
if err != nil || webServerPort <= 0 {
|
||||
webServerPort = node.RelayBindPort + 500
|
||||
}
|
||||
return &Config{
|
||||
BindPort: node.RelayBindPort,
|
||||
VhostHTTPPort: node.RelayVhostHTTPPort,
|
||||
AuthToken: node.RelayAuthToken,
|
||||
LogLevel: "info",
|
||||
WebServerEnabled: node.RelayWebServerEnabled,
|
||||
WebServerPort: webServerPort,
|
||||
}
|
||||
}
|
||||
|
||||
// BuildSettings returns runtime settings shared by relay and flared clients.
|
||||
func BuildSettings(ctx context.Context, node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *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),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
)
|
||||
|
||||
const nodeStatusOnline = "online"
|
||||
|
||||
// Heartbeat processes a relay heartbeat, updates node status, and returns config.
|
||||
func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) {
|
||||
if node == nil {
|
||||
return nil, errors.New("relay node is nil")
|
||||
}
|
||||
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
payload.RelayStatus = normalizeRelayStatus(payload.RelayStatus)
|
||||
payload.Name = strings.TrimSpace(payload.Name)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
|
||||
previous := *node
|
||||
updateNow := node.UpdateRequested
|
||||
updateChannel := normalizeReleaseChannel(node.UpdateChannel)
|
||||
updateTag := strings.TrimSpace(node.UpdateTag)
|
||||
|
||||
now := time.Now().UTC()
|
||||
changes := map[string]any{
|
||||
"version": payload.Version,
|
||||
"ext_version": payload.ExtVersion,
|
||||
"relay_status": payload.RelayStatus,
|
||||
"last_seen_at": now,
|
||||
"status": nodeStatusOnline,
|
||||
"update_requested": false,
|
||||
"update_channel": releaseChannelStable,
|
||||
"update_tag": "",
|
||||
}
|
||||
if payload.Name != "" && strings.TrimSpace(node.Name) == "" {
|
||||
changes["name"] = payload.Name
|
||||
node.Name = payload.Name
|
||||
}
|
||||
if payload.IP != "" && !node.IPManualOverride {
|
||||
changes["ip"] = payload.IP
|
||||
node.IP = payload.IP
|
||||
}
|
||||
if !node.GeoManualOverride {
|
||||
beforeGeo := node.GeoName
|
||||
beforeLat := node.GeoLatitude
|
||||
beforeLon := node.GeoLongitude
|
||||
ofgeoip.ApplyNodeGeoFromIP(ctx, node, node.IP)
|
||||
if node.GeoName != beforeGeo {
|
||||
changes["geo_name"] = node.GeoName
|
||||
}
|
||||
if !coordinatesEqual(beforeLat, node.GeoLatitude) {
|
||||
changes["geo_latitude"] = node.GeoLatitude
|
||||
}
|
||||
if !coordinatesEqual(beforeLon, node.GeoLongitude) {
|
||||
changes["geo_longitude"] = node.GeoLongitude
|
||||
}
|
||||
}
|
||||
if !previous.UpdateRequested {
|
||||
delete(changes, "update_requested")
|
||||
}
|
||||
if previous.UpdateChannel == releaseChannelStable {
|
||||
delete(changes, "update_channel")
|
||||
}
|
||||
if previous.UpdateTag == "" {
|
||||
delete(changes, "update_tag")
|
||||
}
|
||||
|
||||
node.Version = payload.Version
|
||||
node.ExtVersion = payload.ExtVersion
|
||||
node.RelayStatus = payload.RelayStatus
|
||||
node.UpdateRequested = false
|
||||
node.UpdateChannel = releaseChannelStable
|
||||
node.UpdateTag = ""
|
||||
lastSeen := now
|
||||
node.LastSeenAt = &lastSeen
|
||||
node.Status = nodeStatusOnline
|
||||
|
||||
if err := repository.UpdateOpenFlareNodeColumns(ctx, node, changes); err != nil {
|
||||
return nil, fmt.Errorf("update relay heartbeat: %w", err)
|
||||
}
|
||||
if err := reconcileRelayHealthEvents(ctx, node.NodeID, payload.RelayStatus, now); err != nil {
|
||||
return nil, fmt.Errorf("reconcile relay health events: %w", err)
|
||||
}
|
||||
agent.RefreshAccessTokenCache(ctx, node)
|
||||
persistRelayHeartbeatObservability(ctx, node.NodeID, payload, now)
|
||||
|
||||
return &HeartbeatResponse{
|
||||
RelayConfig: buildRelayConfig(ctx, node),
|
||||
RelaySettings: BuildSettings(ctx, node, updateNow, updateChannel, updateTag),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func coordinatesEqual(before *float64, after *float64) bool {
|
||||
if before == nil || after == nil {
|
||||
return before == after
|
||||
}
|
||||
return *before == *after
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
"Wavelet/openflare/plugins/server/kernel/testhelper"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupRelayTestDB(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{},
|
||||
&model.OpenFlareNodeSystemProfile{},
|
||||
&model.OpenFlareMetricSnapshot{},
|
||||
&model.OpenFlareHealthEvent{},
|
||||
&model.OpenFlareNodeObservationFrps{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
agent.ResetAuthCacheForTest()
|
||||
testhelper.SetupLogStoresForTest(t)
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
agent.ResetAuthCacheForTest()
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) {
|
||||
cleanup := setupRelayTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC().Truncate(time.Second)
|
||||
|
||||
node := &model.OpenFlareNode{
|
||||
NodeID: "node-relay-observe",
|
||||
Name: "relay-1",
|
||||
AccessToken: "relay-token",
|
||||
Status: "pending",
|
||||
NodeType: "tunnel_relay",
|
||||
RelayStatus: "unknown",
|
||||
}
|
||||
require.NoError(t, db.DB(ctx).Create(node).Error)
|
||||
|
||||
proxies := []ProxyStat{
|
||||
{
|
||||
Name: "proxy-a",
|
||||
Type: "http",
|
||||
Status: "online",
|
||||
ClientVersion: "0.61.0",
|
||||
ClientAddr: "10.0.0.2:12345",
|
||||
},
|
||||
}
|
||||
_, err := Heartbeat(ctx, node, HeartbeatPayload{
|
||||
Version: "v0.1.0",
|
||||
ExtVersion: "0.61.0",
|
||||
RelayStatus: "healthy",
|
||||
FrpsConnCount: 7,
|
||||
FrpsProxyCount: 3,
|
||||
FrpsClientCount: 2,
|
||||
FrpsProxies: proxies,
|
||||
Name: "relay-runtime",
|
||||
IP: "203.0.113.9",
|
||||
Profile: &agent.NodeSystemProfile{
|
||||
Hostname: "relay-runtime",
|
||||
OSName: "Ubuntu",
|
||||
OSVersion: "24.04",
|
||||
Architecture: "amd64",
|
||||
CPUCores: 4,
|
||||
ReportedAtUnix: now.Unix(),
|
||||
},
|
||||
Snapshot: &agent.NodeMetricSnapshot{
|
||||
CapturedAtUnix: now.Unix(),
|
||||
CPUUsagePercent: 12.5,
|
||||
DiskReadBytes: 100,
|
||||
DiskWriteBytes: 200,
|
||||
},
|
||||
HealthEvents: []agent.NodeHealthEvent{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
var stored model.OpenFlareNode
|
||||
require.NoError(t, db.DB(ctx).Where("node_id = ?", node.NodeID).First(&stored).Error)
|
||||
assert.Equal(t, "online", stored.Status)
|
||||
assert.Equal(t, "healthy", stored.RelayStatus)
|
||||
assert.Equal(t, "203.0.113.9", stored.IP)
|
||||
assert.Equal(t, "v0.1.0", stored.Version)
|
||||
assert.Equal(t, "0.61.0", stored.ExtVersion)
|
||||
|
||||
profile, err := repository.GetOpenFlareNodeSystemProfile(ctx, node.NodeID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "relay-runtime", profile.Hostname)
|
||||
assert.Equal(t, "Ubuntu", profile.OSName)
|
||||
|
||||
snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, snapshots, 1)
|
||||
assert.InDelta(t, 12.5, snapshots[0].CPUUsagePercent, 1e-9)
|
||||
|
||||
frpsObs, err := repository.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, frpsObs, 1)
|
||||
assert.Equal(t, 7, frpsObs[0].FrpsConnections)
|
||||
assert.Equal(t, 3, frpsObs[0].FrpsProxyCount)
|
||||
assert.Equal(t, 2, frpsObs[0].FrpsClientCount)
|
||||
|
||||
var decoded []ProxyStat
|
||||
require.NoError(t, json.Unmarshal([]byte(frpsObs[0].FrpsProxies), &decoded))
|
||||
require.Len(t, decoded, 1)
|
||||
assert.Equal(t, "proxy-a", decoded[0].Name)
|
||||
assert.Equal(t, "online", decoded[0].Status)
|
||||
}
|
||||
|
||||
func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) {
|
||||
cleanup := setupRelayTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
node := &model.OpenFlareNode{
|
||||
NodeID: "node-relay-unhealthy",
|
||||
Name: "relay-unhealthy",
|
||||
AccessToken: "relay-token-unhealthy",
|
||||
Status: "pending",
|
||||
NodeType: "tunnel_relay",
|
||||
RelayStatus: "healthy",
|
||||
}
|
||||
require.NoError(t, db.DB(ctx).Create(node).Error)
|
||||
|
||||
_, err := Heartbeat(ctx, node, HeartbeatPayload{
|
||||
Version: "v0.1.0",
|
||||
ExtVersion: "0.61.0",
|
||||
RelayStatus: "unhealthy",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
events, err := repository.ListOpenFlareHealthEvents(ctx, node.NodeID, true, 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, events, 1)
|
||||
assert.Equal(t, relayFrpsUnhealthyEventType, events[0].EventType)
|
||||
assert.Equal(t, "active", events[0].Status)
|
||||
|
||||
_, err = Heartbeat(ctx, node, HeartbeatPayload{
|
||||
Version: "v0.1.0",
|
||||
ExtVersion: "0.61.0",
|
||||
RelayStatus: "healthy",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
events, err = repository.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, events, 1)
|
||||
assert.Equal(t, "resolved", events[0].Status)
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const ctxRelayNodeKey = "relay_node"
|
||||
|
||||
// Auth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type.
|
||||
func Auth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader("X-Agent-Token"))
|
||||
node, err := agent.AuthenticateAccessToken(c.Request.Context(), token)
|
||||
if err != nil {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
if node.NodeType != "tunnel_relay" {
|
||||
response.AbortForbidden(c, errRelayNodeTypeMismatch)
|
||||
return
|
||||
}
|
||||
c.Set(ctxRelayNodeKey, node)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"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 setupRelayMiddlewareTestDB(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{}))
|
||||
db.SetDB(sqliteDB)
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func seedRelayNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareNode {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
node := &model.OpenFlareNode{
|
||||
NodeID: "relay-test-node",
|
||||
Name: "relay-test",
|
||||
Status: "pending",
|
||||
NodeType: nodeType,
|
||||
AccessToken: accessToken,
|
||||
}
|
||||
require.NoError(t, repository.CreateOpenFlareNode(ctx, node))
|
||||
return node
|
||||
}
|
||||
|
||||
func TestRelayAuthMissingToken(t *testing.T) {
|
||||
cleanup := setupRelayMiddlewareTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/relay/test", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
}
|
||||
|
||||
func TestRelayAuthRejectsWrongNodeType(t *testing.T) {
|
||||
cleanup := setupRelayMiddlewareTestDB(t)
|
||||
defer cleanup()
|
||||
seedRelayNode(t, "edge_node", "edge-token-relay")
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.Use(response.ErrorHandlerMiddleware())
|
||||
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/relay/test", nil)
|
||||
req.Header.Set("X-Agent-Token", "edge-token-relay")
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code)
|
||||
}
|
||||
|
||||
func TestRelayAuthAcceptsTunnelRelay(t *testing.T) {
|
||||
cleanup := setupRelayMiddlewareTestDB(t)
|
||||
defer cleanup()
|
||||
node := seedRelayNode(t, "tunnel_relay", "relay-token-valid")
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := gin.New()
|
||||
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
|
||||
authNode, ok := c.Get(ctxRelayNodeKey)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID)
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/relay/test", nil)
|
||||
req.Header.Set("X-Agent-Token", "relay-token-valid")
|
||||
rec := httptest.NewRecorder()
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"Wavelet/openflare/plugins/server/kernel/repository"
|
||||
|
||||
"Wavelet/openflare/plugins/server/domain/fleet/agent"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const relayFrpsUnhealthyEventType = "frps_unhealthy"
|
||||
|
||||
func reconcileRelayHealthEvents(ctx context.Context, nodeID string, relayStatus string, reportedAt time.Time) error {
|
||||
if relayStatus == "unknown" {
|
||||
return nil
|
||||
}
|
||||
managedTypes := map[string]struct{}{
|
||||
relayFrpsUnhealthyEventType: {},
|
||||
}
|
||||
events := []agent.NodeHealthEvent{}
|
||||
if relayStatus == relayStatusUnhealthy {
|
||||
events = append(events, agent.NodeHealthEvent{
|
||||
EventType: relayFrpsUnhealthyEventType,
|
||||
Severity: "critical",
|
||||
Message: "frps runtime is not healthy",
|
||||
TriggeredAtUnix: reportedAt.Unix(),
|
||||
Metadata: map[string]string{
|
||||
"relay_status": relayStatus,
|
||||
},
|
||||
})
|
||||
}
|
||||
return agent.ReconcileScopedNodeHealthEvents(ctx, nodeID, events, reportedAt, managedTypes)
|
||||
}
|
||||
|
||||
func persistRelayHeartbeatObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) {
|
||||
agent.PersistHeartbeatObservability(ctx, nodeID, agent.NodePayload{
|
||||
Profile: payload.Profile,
|
||||
HostMetrics: payload.Snapshot,
|
||||
HealthEvents: payload.HealthEvents,
|
||||
}, reportedAt)
|
||||
|
||||
frpsObs := &model.OpenFlareNodeObservationFrps{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: reportedAt,
|
||||
FrpsConnections: payload.FrpsConnCount,
|
||||
FrpsProxyCount: payload.FrpsProxyCount,
|
||||
FrpsClientCount: payload.FrpsClientCount,
|
||||
FrpsProxies: agent.MarshalJSON(payload.FrpsProxies),
|
||||
}
|
||||
if err := repository.InsertOpenFlareNodeObservationFrps(ctx, frpsObs); err != nil {
|
||||
zap.L().Error("persist relay frps observation failed", zap.String("node_id", nodeID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import pkgprotocol "Wavelet/openflare/share/protocol"
|
||||
|
||||
// ProxyStat is an alias for protocol.RelayProxyStat.
|
||||
type ProxyStat = pkgprotocol.RelayProxyStat
|
||||
|
||||
// HeartbeatPayload is an alias for protocol.RelayHeartbeatPayload.
|
||||
type HeartbeatPayload = pkgprotocol.RelayHeartbeatPayload
|
||||
|
||||
// Config is an alias for protocol.RelayConfig.
|
||||
type Config = pkgprotocol.RelayConfig
|
||||
|
||||
// Settings is an alias for protocol.RelaySettings.
|
||||
type Settings = pkgprotocol.RelaySettings
|
||||
|
||||
// HeartbeatResponse is an alias for protocol.RelayHeartbeatResponse.
|
||||
type HeartbeatResponse = pkgprotocol.RelayHeartbeatResponse
|
||||
@@ -0,0 +1,75 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
ofws "Wavelet/openflare/plugins/server/domain/fleet/websocket"
|
||||
"Wavelet/openflare/plugins/server/kernel/apiutil"
|
||||
"Wavelet/openflare/plugins/server/kernel/model"
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// PostHeartbeat handles POST /relay/heartbeat.
|
||||
// @Summary 上报 Relay 心跳
|
||||
// @Description Relay 节点定期上报运行状态与 frps 观测数据,返回运行时配置
|
||||
// @Tags openflare-relay
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body relay.HeartbeatPayload true "心跳载荷"
|
||||
// @Success 200 {object} response.Any{data=relay.HeartbeatResponse} "心跳响应"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Agent Token 无效"
|
||||
// @Failure 403 {object} response.Any "节点类型不匹配"
|
||||
// @Router /api/v1/relay/heartbeat [post]
|
||||
func PostHeartbeat(c *gin.Context) {
|
||||
var payload HeartbeatPayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
|
||||
|
||||
authNode, ok := c.Get(ctxRelayNodeKey)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := Heartbeat(c.Request.Context(), node, payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// GetWebSocket handles GET /relay/ws.
|
||||
// @Summary 升级 Relay WebSocket 连接
|
||||
// @Description 将已认证的 Relay 连接升级为 WebSocket 长连接,用于配置推送
|
||||
// @Tags openflare-relay
|
||||
// @Security AgentTokenAuth
|
||||
// @Failure 401 {object} response.Any "Agent Token 无效"
|
||||
// @Failure 403 {object} response.Any "节点类型不匹配"
|
||||
// @Router /api/v1/relay/ws [get]
|
||||
func GetWebSocket(c *gin.Context) {
|
||||
authNode, ok := c.Get(ctxRelayNodeKey)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
ofws.ServeRelay(c, node.NodeID)
|
||||
}
|
||||
@@ -0,0 +1,211 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package websocket manages persistent WebSocket connections between the OpenFlare server and its agents.
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const (
|
||||
// AgentWSConnectedLastSeenValue is the sentinel last_seen_at value when agent WS is connected.
|
||||
AgentWSConnectedLastSeenValue = "__OPENFLARE_WS_CONNECTED__"
|
||||
|
||||
agentMessageTypeStatus = "status"
|
||||
agentMessageTypeSettings = "settings"
|
||||
agentMessageTypeActiveConfig = "active_config"
|
||||
agentMessageTypeForceSyncConfig = "force_sync_config"
|
||||
agentMessageTypeWAFIPGroups = "waf_ip_groups"
|
||||
)
|
||||
|
||||
// AgentStatusHandler processes inbound agent websocket status payloads.
|
||||
type AgentStatusHandler func(ctx context.Context, nodeID, remoteAddr string, payload json.RawMessage)
|
||||
|
||||
type agentClient struct {
|
||||
wsClientCore
|
||||
remoteAddr string
|
||||
onStatus AgentStatusHandler
|
||||
}
|
||||
|
||||
type agentHub struct {
|
||||
mu sync.RWMutex
|
||||
clients map[string]*agentClient
|
||||
}
|
||||
|
||||
var defaultAgentHub = &agentHub{clients: make(map[string]*agentClient)}
|
||||
|
||||
// ServeAgent handles an upgraded agent websocket connection.
|
||||
func ServeAgent(c *gin.Context, nodeID string, onStatus AgentStatusHandler) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
slog.Debug("agent ws upgrade failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
client := &agentClient{
|
||||
wsClientCore: wsClientCore{
|
||||
nodeID: nodeID,
|
||||
conn: conn,
|
||||
send: make(chan Message, wsChannelBuf),
|
||||
done: make(chan struct{}),
|
||||
},
|
||||
remoteAddr: c.Request.RemoteAddr,
|
||||
onStatus: onStatus,
|
||||
}
|
||||
defaultAgentHub.register(client)
|
||||
defer defaultAgentHub.unregister(client)
|
||||
|
||||
slog.Debug("agent ws connected", "node_id", nodeID, "remote", client.remoteAddr)
|
||||
|
||||
go client.writePump()
|
||||
client.readPump()
|
||||
}
|
||||
|
||||
func (h *agentHub) register(client *agentClient) {
|
||||
h.mu.Lock()
|
||||
if existing := h.clients[client.nodeID]; existing != nil {
|
||||
existing.close()
|
||||
}
|
||||
h.clients[client.nodeID] = client
|
||||
h.mu.Unlock()
|
||||
}
|
||||
|
||||
func (h *agentHub) unregister(client *agentClient) {
|
||||
h.mu.Lock()
|
||||
if current := h.clients[client.nodeID]; current == client {
|
||||
delete(h.clients, client.nodeID)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
client.close()
|
||||
}
|
||||
|
||||
// IsAgentConnected reports whether an agent websocket is active.
|
||||
func IsAgentConnected(nodeID string) bool {
|
||||
defaultAgentHub.mu.RLock()
|
||||
client := defaultAgentHub.clients[nodeID]
|
||||
defaultAgentHub.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-client.done:
|
||||
return false
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// SendAgentSettings pushes agent settings to a connected agent.
|
||||
func SendAgentSettings(nodeID string, payload any) bool {
|
||||
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeSettings, Payload: payload})
|
||||
}
|
||||
|
||||
// SendAgentActiveConfig pushes active config metadata to a connected agent.
|
||||
func SendAgentActiveConfig(nodeID string, payload any) bool {
|
||||
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeActiveConfig, Payload: payload})
|
||||
}
|
||||
|
||||
// SendAgentWAFIPGroups pushes WAF IP group updates to a connected agent.
|
||||
func SendAgentWAFIPGroups(nodeID string, payload any) bool {
|
||||
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeWAFIPGroups, Payload: payload})
|
||||
}
|
||||
|
||||
// BroadcastWAFIPGroups pushes changed WAF IP groups to all connected agents.
|
||||
func BroadcastWAFIPGroups(payload any) int {
|
||||
return broadcastAgent(agentMessageTypeWAFIPGroups, payload)
|
||||
}
|
||||
|
||||
// BroadcastActiveConfig pushes active config metadata to all connected agents.
|
||||
func BroadcastActiveConfig(payload any) int {
|
||||
return broadcastAgent(agentMessageTypeActiveConfig, payload)
|
||||
}
|
||||
|
||||
func broadcastAgent(messageType string, payload any) int {
|
||||
if payload == nil {
|
||||
return 0
|
||||
}
|
||||
message := Message{Type: messageType, Payload: payload}
|
||||
defaultAgentHub.mu.RLock()
|
||||
clients := make([]*agentClient, 0, len(defaultAgentHub.clients))
|
||||
for _, client := range defaultAgentHub.clients {
|
||||
clients = append(clients, client)
|
||||
}
|
||||
defaultAgentHub.mu.RUnlock()
|
||||
|
||||
success := 0
|
||||
for _, client := range clients {
|
||||
if client.enqueue(message) {
|
||||
success++
|
||||
}
|
||||
}
|
||||
return success
|
||||
}
|
||||
|
||||
// SendForceSyncConfig notifies an agent to force sync configuration.
|
||||
func SendForceSyncConfig(nodeID string, payload any) bool {
|
||||
return sendAgentMessage(nodeID, Message{Type: agentMessageTypeForceSyncConfig, Payload: payload})
|
||||
}
|
||||
|
||||
func sendAgentMessage(nodeID string, message Message) bool {
|
||||
defaultAgentHub.mu.RLock()
|
||||
client := defaultAgentHub.clients[nodeID]
|
||||
defaultAgentHub.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
return client.enqueue(message)
|
||||
}
|
||||
|
||||
func (c *agentClient) readPump() {
|
||||
defer c.close()
|
||||
|
||||
for {
|
||||
_ = c.conn.SetReadDeadline(time.Now().Add(agentWSReadTimeout()))
|
||||
_, data, err := c.conn.ReadMessage()
|
||||
if err != nil {
|
||||
slog.Debug("agent ws read closed", "node_id", c.nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
var inbound struct {
|
||||
Type string `json:"type"`
|
||||
Payload json.RawMessage `json:"payload,omitempty"`
|
||||
}
|
||||
if err = json.Unmarshal(data, &inbound); err != nil {
|
||||
slog.Debug("agent ws invalid message", "node_id", c.nodeID, "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
slog.Debug("agent ws message received", "node_id", c.nodeID, "type", inbound.Type)
|
||||
switch inbound.Type {
|
||||
case agentMessageTypeStatus:
|
||||
if c.onStatus != nil {
|
||||
c.onStatus(context.Background(), c.nodeID, c.remoteAddr, inbound.Payload)
|
||||
}
|
||||
case messageTypePing:
|
||||
_ = c.enqueue(Message{Type: messageTypePong})
|
||||
case messageTypePong:
|
||||
default:
|
||||
slog.Debug("agent ws unsupported message type", "node_id", c.nodeID, "type", inbound.Type)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func agentWSReadTimeout() time.Duration {
|
||||
timeout := wsReadDeadline
|
||||
if timeout < minAgentWSReadTimeout {
|
||||
return minAgentWSReadTimeout
|
||||
}
|
||||
return timeout
|
||||
}
|
||||
|
||||
func (c *agentClient) writePump() {
|
||||
runWritePump(c.nodeID, c.conn, c.done, c.send, c.close, "agent ws")
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// wsClientCore holds the state and lifecycle shared by all WebSocket client
|
||||
// variants (agent/relay/flared). Embed it; call close exactly-once semantics
|
||||
// are guaranteed via once.
|
||||
type wsClientCore struct {
|
||||
nodeID string
|
||||
conn *websocket.Conn
|
||||
send chan Message
|
||||
done chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
// close tears down the connection at most once.
|
||||
func (c *wsClientCore) close() {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
c.once.Do(func() {
|
||||
close(c.done)
|
||||
if c.conn != nil {
|
||||
_ = c.conn.Close()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// enqueue best-effort delivers message; it never blocks and fails fast when
|
||||
// the client is closed or its send buffer is full.
|
||||
func (c *wsClientCore) enqueue(message Message) bool {
|
||||
// 先确定性检查 closed:若与发送合并在同一个 select,两个 case 同时就绪时
|
||||
// Go 会随机选择,close 后仍可能投递成功。
|
||||
select {
|
||||
case <-c.done:
|
||||
return false
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case c.send <- message:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWSClientCoreCloseIsIdempotent(t *testing.T) {
|
||||
core := &wsClientCore{
|
||||
send: make(chan Message, 1),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
var wg sync.WaitGroup
|
||||
for range 8 {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
core.close()
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
select {
|
||||
case <-core.done:
|
||||
default:
|
||||
t.Fatal("close did not signal done")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWSClientCoreEnqueueFailsAfterClose(t *testing.T) {
|
||||
core := &wsClientCore{
|
||||
send: make(chan Message, 1),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
core.close()
|
||||
// 循环多次:若 close 检查与发送合并在同一个 select,两 case 同时就绪时
|
||||
// Go 随机选择,单次调用可能碰巧通过。
|
||||
for range 50 {
|
||||
if core.enqueue(Message{Type: messageTypePing}) {
|
||||
t.Fatal("enqueue must fail after close")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWSClientCoreEnqueueNeverBlocks(t *testing.T) {
|
||||
core := &wsClientCore{
|
||||
send: make(chan Message, 1), // 缓冲小于消息数,验证不阻塞
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
defer core.close()
|
||||
for range 3 {
|
||||
if !core.enqueue(Message{Type: messageTypePing}) && len(core.send) == 0 {
|
||||
t.Fatal("enqueue failed with empty buffer")
|
||||
}
|
||||
}
|
||||
if core.enqueue(Message{Type: messageTypePing}) {
|
||||
t.Fatal("enqueue must fail when buffer full")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
const (
|
||||
messageTypePing = "ping"
|
||||
messageTypePong = "pong"
|
||||
messageTypeNotify = "notify"
|
||||
)
|
||||
|
||||
// Message is a JSON-framed websocket payload.
|
||||
type Message struct {
|
||||
Type string `json:"type"`
|
||||
Payload any `json:"payload,omitempty"`
|
||||
}
|
||||
|
||||
var upgrader = websocket.Upgrader{
|
||||
CheckOrigin: func(_ *http.Request) bool { return true },
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import "time"
|
||||
|
||||
const (
|
||||
wsChannelBuf = 16
|
||||
wsPingInterval = 30 * time.Second
|
||||
wsReadDeadline = 90 * time.Second
|
||||
wsWriteDeadline = 10 * time.Second
|
||||
minAgentWSReadTimeout = 30 * time.Second
|
||||
)
|
||||
@@ -0,0 +1,124 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const (
|
||||
// FlaredWSConnectedLastSeenValue is the sentinel last_seen_at value when flared WS is connected.
|
||||
FlaredWSConnectedLastSeenValue = "__OPENFLARE_FLARED_WS_CONNECTED__"
|
||||
|
||||
flaredMessageTypeActiveConfig = "active_config"
|
||||
flaredMessageTypeForceSync = "force_sync"
|
||||
flaredMessageTypePong = "pong"
|
||||
)
|
||||
|
||||
type flaredClient struct {
|
||||
wsClientCore
|
||||
}
|
||||
|
||||
type flaredHub struct {
|
||||
mu sync.RWMutex
|
||||
clients map[string]*flaredClient
|
||||
}
|
||||
|
||||
var defaultFlaredHub = &flaredHub{clients: make(map[string]*flaredClient)}
|
||||
|
||||
// ServeFlared handles an upgraded flared websocket connection.
|
||||
func ServeFlared(c *gin.Context, nodeID string) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
slog.Debug("flared ws upgrade failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
client := &flaredClient{
|
||||
wsClientCore: wsClientCore{
|
||||
nodeID: nodeID,
|
||||
conn: conn,
|
||||
send: make(chan Message, wsChannelBuf),
|
||||
done: make(chan struct{}),
|
||||
},
|
||||
}
|
||||
defaultFlaredHub.register(client)
|
||||
defer defaultFlaredHub.unregister(client)
|
||||
|
||||
slog.Debug("flared ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr)
|
||||
|
||||
go client.writePump()
|
||||
client.readPump()
|
||||
}
|
||||
|
||||
func (h *flaredHub) register(client *flaredClient) {
|
||||
h.mu.Lock()
|
||||
if existing := h.clients[client.nodeID]; existing != nil {
|
||||
existing.close()
|
||||
}
|
||||
h.clients[client.nodeID] = client
|
||||
h.mu.Unlock()
|
||||
}
|
||||
|
||||
func (h *flaredHub) unregister(client *flaredClient) {
|
||||
h.mu.Lock()
|
||||
if current := h.clients[client.nodeID]; current == client {
|
||||
delete(h.clients, client.nodeID)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
client.close()
|
||||
}
|
||||
|
||||
// DisconnectFlaredClient forcefully disconnects a flared websocket client.
|
||||
func DisconnectFlaredClient(nodeID string) {
|
||||
defaultFlaredHub.mu.Lock()
|
||||
client := defaultFlaredHub.clients[nodeID]
|
||||
if client != nil {
|
||||
delete(defaultFlaredHub.clients, nodeID)
|
||||
}
|
||||
defaultFlaredHub.mu.Unlock()
|
||||
if client != nil {
|
||||
client.close()
|
||||
}
|
||||
}
|
||||
|
||||
// IsFlaredConnected reports whether a flared websocket is active.
|
||||
func IsFlaredConnected(nodeID string) bool {
|
||||
defaultFlaredHub.mu.RLock()
|
||||
client := defaultFlaredHub.clients[nodeID]
|
||||
defaultFlaredHub.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-client.done:
|
||||
return false
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// SendFlaredPong enqueues a pong message for the flared node.
|
||||
func SendFlaredPong(nodeID string) bool {
|
||||
defaultFlaredHub.mu.RLock()
|
||||
client := defaultFlaredHub.clients[nodeID]
|
||||
defaultFlaredHub.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
// 委托 enqueue:closed 检查与发送不能合并在同一个 select(两 case 同时
|
||||
// 就绪时 Go 随机选择,close 后仍可能投递成功)。
|
||||
return client.enqueue(Message{Type: flaredMessageTypePong})
|
||||
}
|
||||
|
||||
func (c *flaredClient) readPump() {
|
||||
runReadPump(c.nodeID, c.conn, c.close, "flared ws", SendFlaredPong, flaredMessageTypePong)
|
||||
}
|
||||
|
||||
func (c *flaredClient) writePump() {
|
||||
runWritePump(c.nodeID, c.conn, c.done, c.send, c.close, "flared ws")
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
func runReadPump(
|
||||
nodeID string,
|
||||
conn *websocket.Conn,
|
||||
closeFn func(),
|
||||
logLabel string,
|
||||
sendPong func(string) bool,
|
||||
clientPongType string,
|
||||
) {
|
||||
defer closeFn()
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
return conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
|
||||
})
|
||||
|
||||
for {
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
slog.Debug(logLabel+" read closed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
var message Message
|
||||
if err = json.Unmarshal(data, &message); err != nil {
|
||||
slog.Debug(logLabel+" invalid message", "node_id", nodeID, "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
switch message.Type {
|
||||
case messageTypePing:
|
||||
_ = sendPong(nodeID)
|
||||
case clientPongType:
|
||||
// Refresh read deadline when the client replies with a JSON pong.
|
||||
// This keeps the connection alive when the WebSocket is proxied
|
||||
// through Cloudflare, which enforces a 100-second idle timeout on
|
||||
// the TCP stream. Without this refresh, the server's 90-second read
|
||||
// deadline expires and terminates the connection even though the
|
||||
// client is actively responding to pings.
|
||||
_ = conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
|
||||
default:
|
||||
slog.Debug(logLabel+" unsupported message", "node_id", nodeID, "type", message.Type)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// RelayWSConnectedLastSeenValue is the sentinel last_seen_at value when relay WS is connected.
|
||||
const RelayWSConnectedLastSeenValue = "__OPENFLARE_WS_CONNECTED__"
|
||||
|
||||
type relayClient struct {
|
||||
wsClientCore
|
||||
}
|
||||
|
||||
type relayHub struct {
|
||||
mu sync.RWMutex
|
||||
clients map[string]*relayClient
|
||||
}
|
||||
|
||||
var defaultRelayHub = &relayHub{clients: make(map[string]*relayClient)}
|
||||
|
||||
// ServeRelay handles an upgraded relay websocket connection.
|
||||
func ServeRelay(c *gin.Context, nodeID string) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
slog.Debug("relay ws upgrade failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
client := &relayClient{
|
||||
wsClientCore: wsClientCore{
|
||||
nodeID: nodeID,
|
||||
conn: conn,
|
||||
send: make(chan Message, wsChannelBuf),
|
||||
done: make(chan struct{}),
|
||||
},
|
||||
}
|
||||
defaultRelayHub.register(client)
|
||||
defer defaultRelayHub.unregister(client)
|
||||
|
||||
slog.Debug("relay ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr)
|
||||
|
||||
go client.writePump()
|
||||
client.readPump()
|
||||
}
|
||||
|
||||
func (h *relayHub) register(client *relayClient) {
|
||||
h.mu.Lock()
|
||||
if existing := h.clients[client.nodeID]; existing != nil {
|
||||
existing.close()
|
||||
}
|
||||
h.clients[client.nodeID] = client
|
||||
h.mu.Unlock()
|
||||
}
|
||||
|
||||
func (h *relayHub) unregister(client *relayClient) {
|
||||
h.mu.Lock()
|
||||
if current := h.clients[client.nodeID]; current == client {
|
||||
delete(h.clients, client.nodeID)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
client.close()
|
||||
}
|
||||
|
||||
// IsRelayConnected reports whether a relay websocket is active.
|
||||
func IsRelayConnected(nodeID string) bool {
|
||||
defaultRelayHub.mu.RLock()
|
||||
client := defaultRelayHub.clients[nodeID]
|
||||
defaultRelayHub.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-client.done:
|
||||
return false
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// SendRelayPong enqueues a pong message for the relay node.
|
||||
func SendRelayPong(nodeID string) bool {
|
||||
defaultRelayHub.mu.RLock()
|
||||
client := defaultRelayHub.clients[nodeID]
|
||||
defaultRelayHub.mu.RUnlock()
|
||||
if client == nil {
|
||||
return false
|
||||
}
|
||||
// 委托 enqueue:closed 检查与发送不能合并在同一个 select(两 case 同时
|
||||
// 就绪时 Go 随机选择,close 后仍可能投递成功)。
|
||||
return client.enqueue(Message{Type: messageTypePong})
|
||||
}
|
||||
|
||||
func (c *relayClient) readPump() {
|
||||
runReadPump(c.nodeID, c.conn, c.close, "relay ws", SendRelayPong, messageTypePong)
|
||||
}
|
||||
|
||||
func (c *relayClient) writePump() {
|
||||
runWritePump(c.nodeID, c.conn, c.done, c.send, c.close, "relay ws")
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// runWritePump drains send onto conn until done is closed, emitting
|
||||
// JSON pings at wsPingInterval. Shared by agent/relay/flared clients;
|
||||
// closeFn must be idempotent.
|
||||
func runWritePump(
|
||||
nodeID string,
|
||||
conn *websocket.Conn,
|
||||
done <-chan struct{},
|
||||
send chan Message,
|
||||
closeFn func(),
|
||||
logLabel string,
|
||||
) {
|
||||
ticker := time.NewTicker(wsPingInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-done:
|
||||
return
|
||||
case message := <-send:
|
||||
_ = conn.SetWriteDeadline(time.Now().Add(wsWriteDeadline))
|
||||
if err := conn.WriteJSON(message); err != nil {
|
||||
slog.Debug(logLabel+" write failed", "node_id", nodeID, "error", err)
|
||||
closeFn()
|
||||
return
|
||||
}
|
||||
case <-ticker.C:
|
||||
select {
|
||||
case <-done:
|
||||
return
|
||||
case send <- Message{Type: messageTypePing}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user