mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 08:06:37 +08:00
refactor(repo): consolidate openflare-server to root and move subprojects to internal/apps
- Merge all files inside openflare-server to the repository root directory. - Relocate agent, relay, and flared subprojects from internal/ to internal/apps/. - Combine docker-compose files and update build context paths to root. - Update GitHub workflows and Dockerfiles to refer to new directories and package names. - Rewrite Go package imports across all files. - Resolve database renew test race condition and clean up docs.
This commit is contained in:
@@ -0,0 +1,95 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
|
||||
)
|
||||
|
||||
var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) {
|
||||
return pkggeoip.NewMaxMindGeoIPService()
|
||||
}
|
||||
|
||||
type accessLogRegionResolver struct {
|
||||
provider pkggeoip.GeoIPService
|
||||
cache map[string]string
|
||||
}
|
||||
|
||||
func newAccessLogRegionResolver() (*accessLogRegionResolver, error) {
|
||||
provider, err := accessLogGeoProviderFactory()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &accessLogRegionResolver{
|
||||
provider: provider,
|
||||
cache: make(map[string]string),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Close() {
|
||||
if r == nil || r.provider == nil {
|
||||
return
|
||||
}
|
||||
if err := r.provider.Close(); err != nil {
|
||||
slog.Warn("close access log geo provider failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *accessLogRegionResolver) Resolve(rawIP string) string {
|
||||
if r == nil || r.provider == nil {
|
||||
return ""
|
||||
}
|
||||
normalizedIP := normalizeAccessLogIP(rawIP)
|
||||
if normalizedIP == "" {
|
||||
return ""
|
||||
}
|
||||
if cached, ok := r.cache[normalizedIP]; ok {
|
||||
return cached
|
||||
}
|
||||
|
||||
info, err := r.provider.GetGeoInfo(net.ParseIP(normalizedIP))
|
||||
if err != nil || info == nil {
|
||||
r.cache[normalizedIP] = ""
|
||||
return ""
|
||||
}
|
||||
|
||||
region := strings.TrimSpace(info.Name)
|
||||
if region == "" {
|
||||
region = strings.TrimSpace(info.ISOCode)
|
||||
}
|
||||
r.cache[normalizedIP] = region
|
||||
return region
|
||||
}
|
||||
|
||||
func normalizeAccessLogIP(raw string) string {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
trimmed = strings.TrimPrefix(trimmed, "[")
|
||||
trimmed = strings.TrimSuffix(trimmed, "]")
|
||||
if ip := net.ParseIP(trimmed); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(strings.TrimSpace(raw))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
host = strings.TrimPrefix(host, "[")
|
||||
host = strings.TrimSuffix(host, "]")
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
agentTokenPositiveCacheTTL = 2 * time.Minute
|
||||
agentTokenNegativeCacheTTL = 10 * time.Minute
|
||||
)
|
||||
|
||||
type cachedAgentNode struct {
|
||||
node *model.OpenFlareNode
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
type accessTokenAuthCache struct {
|
||||
mu sync.RWMutex
|
||||
positive map[string]cachedAgentNode
|
||||
negative map[string]time.Time
|
||||
now func() time.Time
|
||||
loadNodeByToken func(context.Context, string) (*model.OpenFlareNode, error)
|
||||
}
|
||||
|
||||
var tokenCache = newAccessTokenAuthCache()
|
||||
|
||||
func newAccessTokenAuthCache() *accessTokenAuthCache {
|
||||
return &accessTokenAuthCache{
|
||||
positive: make(map[string]cachedAgentNode),
|
||||
negative: make(map[string]time.Time),
|
||||
now: time.Now,
|
||||
loadNodeByToken: func(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
return model.GetOpenFlareNodeByAccessToken(ctx, token)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) authenticate(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
now := c.now()
|
||||
if node, ok := c.getNode(token, now); ok {
|
||||
return node, nil
|
||||
}
|
||||
if c.isMissing(token, now) {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
node, err := c.loadNodeByToken(ctx, token)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL))
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
c.storeNode(token, node)
|
||||
return cloneNode(node), nil
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) getNode(token string, now time.Time) (*model.OpenFlareNode, bool) {
|
||||
c.mu.RLock()
|
||||
entry, ok := c.positive[token]
|
||||
c.mu.RUnlock()
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
if now.After(entry.expiresAt) {
|
||||
c.mu.Lock()
|
||||
delete(c.positive, token)
|
||||
c.mu.Unlock()
|
||||
return nil, false
|
||||
}
|
||||
return cloneNode(entry.node), true
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) isMissing(token string, now time.Time) bool {
|
||||
c.mu.RLock()
|
||||
expiresAt, ok := c.negative[token]
|
||||
c.mu.RUnlock()
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if now.After(expiresAt) {
|
||||
c.mu.Lock()
|
||||
delete(c.negative, token)
|
||||
c.mu.Unlock()
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) storeNode(token string, node *model.OpenFlareNode) {
|
||||
if token == "" || node == nil {
|
||||
return
|
||||
}
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
delete(c.negative, token)
|
||||
c.positive[token] = cachedAgentNode{
|
||||
node: cloneNode(node),
|
||||
expiresAt: c.now().Add(agentTokenPositiveCacheTTL),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) storeMissing(token string, expiresAt time.Time) {
|
||||
if token == "" {
|
||||
return
|
||||
}
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
delete(c.positive, token)
|
||||
c.negative[token] = expiresAt
|
||||
}
|
||||
|
||||
func (c *accessTokenAuthCache) reset() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.positive = make(map[string]cachedAgentNode)
|
||||
c.negative = make(map[string]time.Time)
|
||||
}
|
||||
|
||||
// ResetAuthCacheForTest clears the in-memory access token cache for integration tests.
|
||||
func ResetAuthCacheForTest() {
|
||||
tokenCache.reset()
|
||||
}
|
||||
|
||||
// AuthenticateAccessToken validates X-Agent-Token against of_nodes.access_token.
|
||||
func AuthenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return nil, errors.New(errMissingAgentToken)
|
||||
}
|
||||
return tokenCache.authenticate(ctx, token)
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type configVersionRecord struct {
|
||||
ID uint `gorm:"primaryKey"`
|
||||
Version string `gorm:"column:version"`
|
||||
SnapshotJSON string `gorm:"column:snapshot_json"`
|
||||
SupportFilesJSON string `gorm:"column:support_files_json"`
|
||||
Checksum string `gorm:"column:checksum"`
|
||||
IsActive bool `gorm:"column:is_active"`
|
||||
CreatedAt time.Time `gorm:"column:created_at"`
|
||||
}
|
||||
|
||||
func (configVersionRecord) TableName() string {
|
||||
return "of_config_versions"
|
||||
}
|
||||
|
||||
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
|
||||
version, err := loadActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ActiveConfigMeta{
|
||||
Version: version.Version,
|
||||
Checksum: version.Checksum,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) {
|
||||
version, err := loadActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var supportFiles []SupportFile
|
||||
if strings.TrimSpace(version.SupportFilesJSON) != "" {
|
||||
if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return &ConfigResponse{
|
||||
Version: version.Version,
|
||||
Checksum: version.Checksum,
|
||||
SourceConfigJSON: version.SnapshotJSON,
|
||||
SupportFiles: sourceSupportFiles(supportFiles),
|
||||
CreatedAt: version.CreatedAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func loadActiveConfigVersion(ctx context.Context) (*configVersionRecord, error) {
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return nil, errors.New("database not initialized")
|
||||
}
|
||||
version := &configVersionRecord{}
|
||||
err := conn.Where("is_active = ?", true).Order("id desc").First(version).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return version, nil
|
||||
}
|
||||
|
||||
func sourceSupportFiles(files []SupportFile) []SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]SupportFile, 0, len(files))
|
||||
for _, file := range files {
|
||||
if isRuntimeGeneratedSupportFile(file.Path) {
|
||||
continue
|
||||
}
|
||||
result = append(result, file)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func isRuntimeGeneratedSupportFile(path string) bool {
|
||||
switch strings.TrimSpace(path) {
|
||||
case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isActiveConfigNotFound(err error) bool {
|
||||
return errors.Is(err, gorm.ErrRecordNotFound)
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
)
|
||||
|
||||
func TestIsRuntimeGeneratedSupportFile(t *testing.T) {
|
||||
tests := []struct {
|
||||
path string
|
||||
want bool
|
||||
}{
|
||||
{path: "pow_config.json", want: true},
|
||||
{path: "waf_config.json", want: true},
|
||||
{path: openrestyrender.SourceConfigFileName, want: true},
|
||||
{path: "runtime/custom.json", want: false},
|
||||
{path: "certs/example.pem", want: false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
if got := isRuntimeGeneratedSupportFile(tc.path); got != tc.want {
|
||||
t.Fatalf("isRuntimeGeneratedSupportFile(%q) = %v, want %v", tc.path, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSourceSupportFilesFiltersRuntimeGeneratedFiles(t *testing.T) {
|
||||
files := []SupportFile{
|
||||
{Path: "certs/example.pem", Content: "pem"},
|
||||
{Path: "pow_config.json", Content: "{}"},
|
||||
{Path: "waf_config.json", Content: "{}"},
|
||||
{Path: openrestyrender.SourceConfigFileName, Content: "{}"},
|
||||
{Path: "routes/extra.json", Content: "{}"},
|
||||
}
|
||||
|
||||
filtered := sourceSupportFiles(files)
|
||||
if len(filtered) != 2 {
|
||||
t.Fatalf("expected 2 support files, got %d: %+v", len(filtered), filtered)
|
||||
}
|
||||
if filtered[0].Path != "certs/example.pem" || filtered[1].Path != "routes/extra.json" {
|
||||
t.Fatalf("unexpected filtered files: %+v", filtered)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
const (
|
||||
errMissingAgentToken = "缺少 Agent Token"
|
||||
errInvalidAgentToken = "无权进行此操作,Agent Token 无效"
|
||||
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效"
|
||||
errNodeMissingFromContext = "Node object missing from context"
|
||||
errNoActiveConfig = "当前没有激活版本"
|
||||
errNodeNotFound = "节点不存在"
|
||||
errNodeIDRequired = "node_id 不能为空"
|
||||
errVersionRequired = "version 不能为空"
|
||||
errInvalidApplyResult = "result 仅支持 success、warning 或 failed"
|
||||
errIPRequired = "ip 不能为空"
|
||||
errIPInvalid = "ip 格式无效"
|
||||
errAgentVersionRequired = "version 不能为空"
|
||||
errNodeIDConflict = "节点标识生成冲突,请重试"
|
||||
)
|
||||
@@ -0,0 +1,308 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
openrestyStatusHealthy = "healthy"
|
||||
openrestyStatusUnhealthy = "unhealthy"
|
||||
openrestyStatusUnknown = "unknown"
|
||||
releaseChannelStable = "stable"
|
||||
)
|
||||
|
||||
func newRandomToken() (string, error) {
|
||||
buf := make([]byte, 16)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
func newServerNodeID() (string, error) {
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "node-" + token, nil
|
||||
}
|
||||
|
||||
func normalizeOpenrestyStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case openrestyStatusHealthy:
|
||||
return openrestyStatusHealthy
|
||||
case openrestyStatusUnhealthy:
|
||||
return openrestyStatusUnhealthy
|
||||
default:
|
||||
return openrestyStatusUnknown
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeNodePayload(payload NodePayload) NodePayload {
|
||||
payload.Name = strings.TrimSpace(payload.Name)
|
||||
payload.IP = strings.TrimSpace(payload.IP)
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
payload.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
return payload
|
||||
}
|
||||
|
||||
func validateNodePayload(payload NodePayload) error {
|
||||
if payload.IP == "" {
|
||||
return errPayload(errIPRequired)
|
||||
}
|
||||
if net.ParseIP(payload.IP) == nil {
|
||||
return errPayload(errIPInvalid)
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return errPayload(errAgentVersionRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type payloadError string
|
||||
|
||||
func (e payloadError) Error() string { return string(e) }
|
||||
|
||||
func errPayload(message string) error { return payloadError(message) }
|
||||
|
||||
func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveName bool) {
|
||||
if !preserveName || strings.TrimSpace(node.Name) == "" {
|
||||
if strings.TrimSpace(payload.Name) != "" {
|
||||
node.Name = strings.TrimSpace(payload.Name)
|
||||
}
|
||||
}
|
||||
if !node.IPManualOverride {
|
||||
node.IP = strings.TrimSpace(payload.IP)
|
||||
}
|
||||
node.Version = strings.TrimSpace(payload.Version)
|
||||
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
|
||||
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
|
||||
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
|
||||
node.Status = nodeStatusOnline
|
||||
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
|
||||
now := time.Now()
|
||||
node.LastSeenAt = &now
|
||||
node.LastError = truncateForDatabase(payload.LastError, 16000)
|
||||
if !node.GeoManualOverride {
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
}
|
||||
}
|
||||
|
||||
func applyGeoInfoFromIP(node *model.OpenFlareNode, rawIP string) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
node.GeoName = ""
|
||||
node.GeoLatitude = nil
|
||||
node.GeoLongitude = nil
|
||||
ip := net.ParseIP(strings.TrimSpace(rawIP))
|
||||
if ip == nil {
|
||||
return
|
||||
}
|
||||
info, err := ofgeoip.GeoInfoFromIP(ip)
|
||||
if err != nil || info == nil {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(info.Name) != "" {
|
||||
node.GeoName = strings.TrimSpace(info.Name)
|
||||
}
|
||||
if info.Latitude != nil && info.Longitude != nil {
|
||||
node.GeoLatitude = cloneCoordinate(info.Latitude)
|
||||
node.GeoLongitude = cloneCoordinate(info.Longitude)
|
||||
}
|
||||
}
|
||||
|
||||
func cloneCoordinate(value *float64) *float64 {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *value
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func truncateForDatabase(value string, max int) string {
|
||||
if max <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(value))
|
||||
if len(runes) <= max {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:max])
|
||||
}
|
||||
|
||||
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
|
||||
reported := normalizeIP(reportedIP)
|
||||
remote := normalizeRemoteAddr(remoteAddr)
|
||||
if reported == "" {
|
||||
return remote
|
||||
}
|
||||
if isPublicNodeIP(reported) {
|
||||
return reported
|
||||
}
|
||||
if isPublicNodeIP(remote) {
|
||||
return remote
|
||||
}
|
||||
return reported
|
||||
}
|
||||
|
||||
func normalizeIP(raw string) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
host := raw
|
||||
if strings.Contains(raw, ":") {
|
||||
if h, _, err := net.SplitHostPort(raw); err == nil {
|
||||
host = h
|
||||
}
|
||||
}
|
||||
host = strings.TrimPrefix(host, "[")
|
||||
host = strings.TrimSuffix(host, "]")
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func normalizeRemoteAddr(remoteAddr string) string {
|
||||
remoteAddr = strings.TrimSpace(remoteAddr)
|
||||
if remoteAddr == "" {
|
||||
return ""
|
||||
}
|
||||
host, _, err := net.SplitHostPort(remoteAddr)
|
||||
if err != nil {
|
||||
return normalizeIP(remoteAddr)
|
||||
}
|
||||
return normalizeIP(host)
|
||||
}
|
||||
|
||||
func isPublicNodeIP(raw string) bool {
|
||||
ip := net.ParseIP(strings.TrimSpace(raw))
|
||||
if ip == nil {
|
||||
return false
|
||||
}
|
||||
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func buildAgentSettings(node *model.OpenFlareNode, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *Settings {
|
||||
autoUpdate := false
|
||||
if node != nil {
|
||||
autoUpdate = node.AutoUpdateEnabled
|
||||
}
|
||||
if strings.TrimSpace(updateChannel) == "" {
|
||||
updateChannel = releaseChannelStable
|
||||
}
|
||||
return &Settings{
|
||||
HeartbeatInterval: model.AgentHeartbeatInterval,
|
||||
WebsocketUpgradeEnabled: model.AgentWebsocketUpgradeEnabled,
|
||||
AutoUpdate: autoUpdate,
|
||||
UpdateRepo: model.AgentUpdateRepo,
|
||||
UpdateNow: updateNow,
|
||||
UpdateChannel: updateChannel,
|
||||
UpdateTag: strings.TrimSpace(updateTag),
|
||||
RestartOpenrestyNow: restartOpenrestyNow,
|
||||
}
|
||||
}
|
||||
|
||||
func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenFlareNode) map[string]any {
|
||||
if previous == nil || current == nil {
|
||||
return map[string]any{}
|
||||
}
|
||||
changes := make(map[string]any)
|
||||
appendIfChanged := func(key string, before any, after any) {
|
||||
if before != after {
|
||||
changes[key] = after
|
||||
}
|
||||
}
|
||||
appendIfChanged("name", previous.Name, current.Name)
|
||||
appendIfChanged("ip", previous.IP, current.IP)
|
||||
appendIfChanged("geo_name", previous.GeoName, current.GeoName)
|
||||
appendIfChanged("version", previous.Version, current.Version)
|
||||
appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion)
|
||||
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
|
||||
appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage)
|
||||
appendIfChanged("status", previous.Status, current.Status)
|
||||
appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion)
|
||||
appendIfChanged("last_error", previous.LastError, current.LastError)
|
||||
appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested)
|
||||
appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel)
|
||||
appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag)
|
||||
appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested)
|
||||
if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) {
|
||||
changes["geo_latitude"] = current.GeoLatitude
|
||||
}
|
||||
if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) {
|
||||
changes["geo_longitude"] = current.GeoLongitude
|
||||
}
|
||||
if !lastSeenAtEqual(previous.LastSeenAt, current.LastSeenAt) {
|
||||
changes["last_seen_at"] = current.LastSeenAt
|
||||
}
|
||||
return changes
|
||||
}
|
||||
|
||||
func coordinatesEqual(before *float64, after *float64) bool {
|
||||
if before == nil || after == nil {
|
||||
return before == after
|
||||
}
|
||||
return *before == *after
|
||||
}
|
||||
|
||||
func lastSeenAtEqual(before *time.Time, after *time.Time) bool {
|
||||
if before == nil || after == nil {
|
||||
return before == after
|
||||
}
|
||||
return before.Equal(*after)
|
||||
}
|
||||
|
||||
func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
|
||||
payload.NodeID = strings.TrimSpace(payload.NodeID)
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
payload.Result = strings.ToLower(strings.TrimSpace(payload.Result))
|
||||
payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), 16000)
|
||||
payload.Checksum = strings.TrimSpace(payload.Checksum)
|
||||
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
|
||||
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
|
||||
return payload
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
|
||||
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
tokenCache.storeNode(node.AccessToken, cloneNode(node))
|
||||
}
|
||||
|
||||
func cloneNode(node *model.OpenFlareNode) *model.OpenFlareNode {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *node
|
||||
return &cloned
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
|
||||
)
|
||||
|
||||
type fakeGeoIPProvider struct {
|
||||
info *pkggeoip.GeoInfo
|
||||
}
|
||||
|
||||
func (f *fakeGeoIPProvider) Name() string { return "fake-geoip" }
|
||||
|
||||
func (f *fakeGeoIPProvider) GetGeoInfo(ip net.IP) (*pkggeoip.GeoInfo, error) {
|
||||
return f.info, nil
|
||||
}
|
||||
|
||||
func (f *fakeGeoIPProvider) UpdateDatabase() error { return nil }
|
||||
|
||||
func (f *fakeGeoIPProvider) Close() error { return nil }
|
||||
|
||||
func withFakeGeoIPProvider(t *testing.T, info *pkggeoip.GeoInfo) {
|
||||
t.Helper()
|
||||
previous := pkggeoip.CurrentProvider
|
||||
pkggeoip.CurrentProvider = &fakeGeoIPProvider{info: info}
|
||||
t.Cleanup(func() {
|
||||
pkggeoip.CurrentProvider = previous
|
||||
})
|
||||
}
|
||||
|
||||
func geoipFloat(value float64) *float64 {
|
||||
return &value
|
||||
}
|
||||
|
||||
func TestApplyGeoInfoFromIP(t *testing.T) {
|
||||
latitude := 31.2304
|
||||
longitude := 121.4737
|
||||
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
|
||||
Name: "Shanghai",
|
||||
Latitude: geoipFloat(latitude),
|
||||
Longitude: geoipFloat(longitude),
|
||||
})
|
||||
|
||||
node := &model.OpenFlareNode{IP: "203.0.113.10"}
|
||||
applyGeoInfoFromIP(node, node.IP)
|
||||
|
||||
if node.GeoName != "Shanghai" {
|
||||
t.Fatalf("expected geo_name Shanghai, got %q", node.GeoName)
|
||||
}
|
||||
if node.GeoLatitude == nil || *node.GeoLatitude != latitude {
|
||||
t.Fatalf("unexpected geo_latitude: %+v", node.GeoLatitude)
|
||||
}
|
||||
if node.GeoLongitude == nil || *node.GeoLongitude != longitude {
|
||||
t.Fatalf("unexpected geo_longitude: %+v", node.GeoLongitude)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyGeoInfoFromIPSkipsInvalidIP(t *testing.T) {
|
||||
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{Name: "Should Not Apply"})
|
||||
|
||||
node := &model.OpenFlareNode{
|
||||
IP: "203.0.113.10",
|
||||
GeoName: "Existing",
|
||||
GeoLatitude: geoipFloat(1),
|
||||
GeoLongitude: geoipFloat(2),
|
||||
}
|
||||
applyGeoInfoFromIP(node, "not-an-ip")
|
||||
|
||||
if node.GeoName != "" || node.GeoLatitude != nil || node.GeoLongitude != nil {
|
||||
t.Fatalf("expected geo fields to be cleared on invalid IP, got %+v", node)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNodeRuntimeRespectsGeoManualOverride(t *testing.T) {
|
||||
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
|
||||
Name: "Shanghai",
|
||||
Latitude: geoipFloat(31.2304),
|
||||
Longitude: geoipFloat(121.4737),
|
||||
})
|
||||
|
||||
node := &model.OpenFlareNode{
|
||||
GeoManualOverride: true,
|
||||
GeoName: "Manual",
|
||||
GeoLatitude: geoipFloat(10),
|
||||
GeoLongitude: geoipFloat(20),
|
||||
}
|
||||
applyNodeRuntime(node, NodePayload{
|
||||
IP: "203.0.113.10",
|
||||
Version: "1.0.0",
|
||||
}, true)
|
||||
|
||||
if node.GeoName != "Manual" {
|
||||
t.Fatalf("expected manual geo_name to be preserved, got %q", node.GeoName)
|
||||
}
|
||||
if node.GeoLatitude == nil || *node.GeoLatitude != 10 {
|
||||
t.Fatalf("expected manual geo_latitude to be preserved, got %+v", node.GeoLatitude)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectHeartbeatChangesTracksGeoFields(t *testing.T) {
|
||||
before := &model.OpenFlareNode{
|
||||
IP: "10.0.0.1",
|
||||
GeoName: "Old Region",
|
||||
}
|
||||
after := &model.OpenFlareNode{
|
||||
IP: "203.0.113.10",
|
||||
GeoName: "New Region",
|
||||
GeoLatitude: geoipFloat(31.2304),
|
||||
GeoLongitude: geoipFloat(121.4737),
|
||||
}
|
||||
|
||||
changes := collectHeartbeatChanges(before, after)
|
||||
if changes["ip"] != after.IP {
|
||||
t.Fatalf("expected ip change, got %+v", changes)
|
||||
}
|
||||
if changes["geo_name"] != after.GeoName {
|
||||
t.Fatalf("expected geo_name change, got %+v", changes)
|
||||
}
|
||||
if changes["geo_latitude"] != after.GeoLatitude {
|
||||
t.Fatalf("expected geo_latitude change, got %+v", changes)
|
||||
}
|
||||
if changes["geo_longitude"] != after.GeoLongitude {
|
||||
t.Fatalf("expected geo_longitude change, got %+v", changes)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,224 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/node"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// RegisterWithAccessToken registers an agent on a reserved node token.
|
||||
func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*RegistrationResponse, error) {
|
||||
payload = normalizeNodePayload(payload)
|
||||
if authNode == nil {
|
||||
return nil, errors.New(errNodeNotFound)
|
||||
}
|
||||
if err := validateNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
applyNodeRuntime(authNode, payload, true)
|
||||
if err := model.SaveOpenFlareNode(ctx, authNode); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
RefreshAccessTokenCache(ctx, authNode)
|
||||
return &RegistrationResponse{
|
||||
NodeID: authNode.NodeID,
|
||||
AccessToken: authNode.AccessToken,
|
||||
Name: authNode.Name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// RegisterWithDiscovery registers a new node using the global discovery token.
|
||||
func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*RegistrationResponse, error) {
|
||||
payload = normalizeNodePayload(payload)
|
||||
if err := validateNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nodeID, err := newServerNodeID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessToken, err := newRandomToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nodeName := payload.Name
|
||||
if nodeName == "" {
|
||||
nodeName = nodeID
|
||||
}
|
||||
|
||||
record := &model.OpenFlareNode{
|
||||
NodeID: nodeID,
|
||||
Name: nodeName,
|
||||
AccessToken: accessToken,
|
||||
Status: nodeStatusOnline,
|
||||
NodeType: "edge_node",
|
||||
CapabilitiesJSON: "[]",
|
||||
UpdateChannel: releaseChannelStable,
|
||||
}
|
||||
applyNodeRuntime(record, payload, false)
|
||||
|
||||
if err = model.CreateOpenFlareNode(ctx, record); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errNodeIDConflict)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
RefreshAccessTokenCache(ctx, record)
|
||||
return &RegistrationResponse{
|
||||
NodeID: record.NodeID,
|
||||
AccessToken: record.AccessToken,
|
||||
Name: record.Name,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// HeartbeatNode updates runtime state and returns agent settings.
|
||||
func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*HeartbeatResponse, error) {
|
||||
if authNode == nil {
|
||||
return nil, errors.New(errNodeNotFound)
|
||||
}
|
||||
payload.NodeID = authNode.NodeID
|
||||
payload = normalizeNodePayload(payload)
|
||||
if err := validateNodePayload(payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
previous := *authNode
|
||||
updateNow := authNode.UpdateRequested
|
||||
restartOpenrestyNow := authNode.RestartOpenrestyRequested
|
||||
updateChannel := strings.TrimSpace(authNode.UpdateChannel)
|
||||
updateTag := strings.TrimSpace(authNode.UpdateTag)
|
||||
|
||||
applyNodeRuntime(authNode, payload, true)
|
||||
authNode.UpdateRequested = false
|
||||
authNode.UpdateChannel = releaseChannelStable
|
||||
authNode.UpdateTag = ""
|
||||
authNode.RestartOpenrestyRequested = false
|
||||
|
||||
changes := collectHeartbeatChanges(&previous, authNode)
|
||||
if len(changes) > 0 {
|
||||
fields := make([]string, 0, len(changes))
|
||||
for field := range changes {
|
||||
fields = append(fields, field)
|
||||
}
|
||||
if err := model.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
RefreshAccessTokenCache(ctx, authNode)
|
||||
|
||||
reportedAt := time.Now()
|
||||
if authNode.LastSeenAt != nil {
|
||||
reportedAt = *authNode.LastSeenAt
|
||||
}
|
||||
PersistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt)
|
||||
|
||||
activeConfig, err := getActiveConfigMeta(ctx)
|
||||
if err != nil && !isActiveConfigNotFound(err) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
wafIPGroups, err := ChangedWAFIPGroupsForAgent(ctx, nil, payload.WAFIPGroupChecksums)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &HeartbeatResponse{
|
||||
Node: authNode,
|
||||
AgentSettings: buildAgentSettings(authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow),
|
||||
ActiveConfig: activeConfig,
|
||||
WAFIPGroups: wafIPGroups,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetActiveConfig returns the active configuration for an agent.
|
||||
func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) {
|
||||
config, err := getActiveConfigForAgent(ctx)
|
||||
if err != nil {
|
||||
if isActiveConfigNotFound(err) {
|
||||
return nil, errors.New(errNoActiveConfig)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// SyncWAFIPGroups returns WAF IP groups whose checksums differ from the agent state.
|
||||
func SyncWAFIPGroups(ctx context.Context, input WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) {
|
||||
groups, err := ChangedWAFIPGroupsForAgent(ctx, input.IDs, input.Checksums)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &WAFIPGroupSyncResult{Groups: groups}, nil
|
||||
}
|
||||
|
||||
// ReportApplyLog records an agent apply result.
|
||||
func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) {
|
||||
now := time.Now()
|
||||
payload = normalizeApplyLogPayload(payload)
|
||||
if payload.NodeID == "" {
|
||||
return nil, errors.New(errNodeIDRequired)
|
||||
}
|
||||
if payload.Version == "" {
|
||||
return nil, errors.New(errVersionRequired)
|
||||
}
|
||||
if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFailed {
|
||||
return nil, errors.New(errInvalidApplyResult)
|
||||
}
|
||||
|
||||
log := &model.OpenFlareApplyLog{
|
||||
NodeID: payload.NodeID,
|
||||
Version: payload.Version,
|
||||
Result: payload.Result,
|
||||
Message: payload.Message,
|
||||
Checksum: payload.Checksum,
|
||||
MainConfigChecksum: payload.MainConfigChecksum,
|
||||
RouteConfigChecksum: payload.RouteConfigChecksum,
|
||||
SupportFileCount: payload.SupportFileCount,
|
||||
CreatedAt: now,
|
||||
}
|
||||
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return nil, errors.New("database not initialized")
|
||||
}
|
||||
|
||||
err := conn.Transaction(func(tx *gorm.DB) error {
|
||||
record := &model.OpenFlareNode{}
|
||||
if err := tx.Where("node_id = ?", payload.NodeID).First(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
record.Status = nodeStatusOnline
|
||||
record.LastSeenAt = &now
|
||||
if payload.Result == applyResultOK {
|
||||
record.CurrentVersion = payload.Version
|
||||
record.LastError = ""
|
||||
} else {
|
||||
record.LastError = payload.Message
|
||||
}
|
||||
if err := tx.Create(log).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).Error
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return log, nil
|
||||
}
|
||||
|
||||
// ValidateDiscoveryToken delegates to the node package discovery token helper.
|
||||
func ValidateDiscoveryToken(ctx context.Context, token string) error {
|
||||
return node.ValidateDiscoveryToken(ctx, token)
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const (
|
||||
agentTokenHeader = "X-Agent-Token"
|
||||
agentNodeContextKey = "agent_node"
|
||||
)
|
||||
|
||||
// AgentAuth validates X-Agent-Token against of_nodes.access_token.
|
||||
func AgentAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
|
||||
node, err := AuthenticateAccessToken(c.Request.Context(), token)
|
||||
if err != nil {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
}
|
||||
c.Set(agentNodeContextKey, node)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// AgentRegisterAuth accepts either a node access token or the global discovery token.
|
||||
func AgentRegisterAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
|
||||
if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil {
|
||||
c.Set(agentNodeContextKey, node)
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
if err := ValidateDiscoveryToken(c.Request.Context(), token); err != nil {
|
||||
response.AbortUnauthorized(c, errInvalidDiscoveryToken)
|
||||
return
|
||||
}
|
||||
c.Set("discovery_enabled", true)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// AgentNodeFromContext returns the authenticated agent node.
|
||||
func AgentNodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
|
||||
value, ok := c.Get(agentNodeContextKey)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
node, ok := value.(*model.OpenFlareNode)
|
||||
return node, ok
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupAgentAuthTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.OpenFlareNode{},
|
||||
&model.OpenFlareOption{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
tokenCache.reset()
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
tokenCache.reset()
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthenticateAccessToken(t *testing.T) {
|
||||
cleanup := setupAgentAuthTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-auth-1",
|
||||
Name: "edge",
|
||||
AccessToken: "valid-agent-token",
|
||||
Status: nodeStatusOnline,
|
||||
LastSeenAt: &now,
|
||||
NodeType: "edge_node",
|
||||
}).Error)
|
||||
|
||||
t.Run("valid token", func(t *testing.T) {
|
||||
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "node-auth-1", node.NodeID)
|
||||
})
|
||||
|
||||
t.Run("cached token", func(t *testing.T) {
|
||||
originalLoader := tokenCache.loadNodeByToken
|
||||
t.Cleanup(func() {
|
||||
tokenCache.loadNodeByToken = originalLoader
|
||||
})
|
||||
tokenCache.loadNodeByToken = func(context.Context, string) (*model.OpenFlareNode, error) {
|
||||
t.Fatal("db should not be queried for cached token")
|
||||
return nil, nil
|
||||
}
|
||||
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "node-auth-1", node.NodeID)
|
||||
})
|
||||
|
||||
t.Run("missing token", func(t *testing.T) {
|
||||
_, err := AuthenticateAccessToken(ctx, "")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), errMissingAgentToken)
|
||||
})
|
||||
|
||||
t.Run("invalid token", func(t *testing.T) {
|
||||
_, err := AuthenticateAccessToken(ctx, "invalid-token")
|
||||
require.Error(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentAuthMiddleware(t *testing.T) {
|
||||
cleanup := setupAgentAuthTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-mw-1",
|
||||
Name: "edge",
|
||||
AccessToken: "middleware-token",
|
||||
Status: nodeStatusOnline,
|
||||
LastSeenAt: &now,
|
||||
NodeType: "edge_node",
|
||||
}).Error)
|
||||
|
||||
router := testhelper.NewTestGinEngine()
|
||||
router.GET("/protected", AgentAuth(), func(c *gin.Context) {
|
||||
node, ok := AgentNodeFromContext(c)
|
||||
if !ok {
|
||||
c.Status(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"node_id": node.NodeID}))
|
||||
})
|
||||
|
||||
t.Run("authorized request", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set(agentTokenHeader, "middleware-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, resp.Code)
|
||||
var apiResp response.Any
|
||||
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
|
||||
assert.Empty(t, apiResp.ErrorMsg)
|
||||
})
|
||||
|
||||
t.Run("unauthorized request", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set(agentTokenHeader, "bad-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, resp.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentRegisterAuthMiddleware(t *testing.T) {
|
||||
cleanup := setupAgentAuthTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-register-1",
|
||||
Name: "edge",
|
||||
AccessToken: "existing-node-token",
|
||||
Status: nodeStatusOnline,
|
||||
LastSeenAt: &now,
|
||||
NodeType: "edge_node",
|
||||
}).Error)
|
||||
require.NoError(t, model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", "discovery-token"))
|
||||
|
||||
router := testhelper.NewTestGinEngine()
|
||||
router.POST("/register", AgentRegisterAuth(), func(c *gin.Context) {
|
||||
if node, ok := AgentNodeFromContext(c); ok {
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "node", "node_id": node.NodeID}))
|
||||
return
|
||||
}
|
||||
if _, ok := c.Get("discovery_enabled"); ok {
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "discovery"}))
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusInternalServerError)
|
||||
})
|
||||
|
||||
t.Run("existing node token", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/register", nil)
|
||||
req.Header.Set(agentTokenHeader, "existing-node-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, resp.Code)
|
||||
var apiResp response.Any
|
||||
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
|
||||
data, ok := apiResp.Data.(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "node", data["mode"])
|
||||
})
|
||||
|
||||
t.Run("discovery token", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/register", nil)
|
||||
req.Header.Set(agentTokenHeader, "discovery-token")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, resp.Code)
|
||||
var apiResp response.Any
|
||||
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
|
||||
data, ok := apiResp.Data.(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "discovery", data["mode"])
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,499 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const (
|
||||
healthEventStatusActive = "active"
|
||||
healthEventStatusResolved = "resolved"
|
||||
healthSeverityInfo = "info"
|
||||
healthSeverityWarning = "warning"
|
||||
healthSeverityCritical = "critical"
|
||||
nodeAccessLogRetentionDays = 90
|
||||
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
|
||||
accessLogPathMaxLength = 100
|
||||
)
|
||||
|
||||
// NodeSystemProfile is the agent-reported system profile.
|
||||
type NodeSystemProfile struct {
|
||||
Hostname string `json:"hostname"`
|
||||
OSName string `json:"os_name"`
|
||||
OSVersion string `json:"os_version"`
|
||||
KernelVersion string `json:"kernel_version"`
|
||||
Architecture string `json:"architecture"`
|
||||
CPUModel string `json:"cpu_model"`
|
||||
CPUCores int `json:"cpu_cores"`
|
||||
TotalMemoryBytes int64 `json:"total_memory_bytes"`
|
||||
TotalDiskBytes int64 `json:"total_disk_bytes"`
|
||||
UptimeSeconds int64 `json:"uptime_seconds"`
|
||||
ReportedAtUnix int64 `json:"reported_at_unix"`
|
||||
}
|
||||
|
||||
// NodeMetricSnapshot is the agent-reported capacity snapshot.
|
||||
type NodeMetricSnapshot struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
CPUUsagePercent float64 `json:"cpu_usage_percent"`
|
||||
MemoryUsedBytes int64 `json:"memory_used_bytes"`
|
||||
MemoryTotalBytes int64 `json:"memory_total_bytes"`
|
||||
StorageUsedBytes int64 `json:"storage_used_bytes"`
|
||||
StorageTotalBytes int64 `json:"storage_total_bytes"`
|
||||
DiskReadBytes int64 `json:"disk_read_bytes"`
|
||||
DiskWriteBytes int64 `json:"disk_write_bytes"`
|
||||
NetworkRxBytes int64 `json:"network_rx_bytes"`
|
||||
NetworkTxBytes int64 `json:"network_tx_bytes"`
|
||||
}
|
||||
|
||||
// NodeOpenrestyObservation is the agent-reported openresty network observation.
|
||||
type NodeOpenrestyObservation struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
|
||||
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
|
||||
OpenrestyConnections int64 `json:"openresty_connections"`
|
||||
}
|
||||
|
||||
// NodeTrafficReport is the agent-reported traffic window.
|
||||
type NodeTrafficReport struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
StatusCodes map[string]int64 `json:"status_codes"`
|
||||
TopDomains map[string]int64 `json:"top_domains"`
|
||||
SourceCountries map[string]int64 `json:"source_countries"`
|
||||
}
|
||||
|
||||
// NodeAccessLog is a single access log row from the agent.
|
||||
type NodeAccessLog struct {
|
||||
LoggedAtUnix int64 `json:"logged_at_unix"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
StatusCode int `json:"status_code"`
|
||||
}
|
||||
|
||||
// BufferedObservabilityRecord is a buffered observability window from the agent.
|
||||
type BufferedObservabilityRecord struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
|
||||
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
|
||||
}
|
||||
|
||||
// NodeHealthEvent is an agent-reported health event.
|
||||
type NodeHealthEvent struct {
|
||||
EventType string `json:"event_type"`
|
||||
Severity string `json:"severity"`
|
||||
Message string `json:"message"`
|
||||
TriggeredAtUnix int64 `json:"triggered_at_unix"`
|
||||
Metadata map[string]string `json:"metadata"`
|
||||
}
|
||||
|
||||
// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events.
|
||||
func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
|
||||
if strings.TrimSpace(nodeID) == "" {
|
||||
return
|
||||
}
|
||||
if payload.Profile == nil &&
|
||||
payload.Snapshot == nil &&
|
||||
payload.TrafficReport == nil &&
|
||||
len(payload.AccessLogs) == 0 &&
|
||||
len(payload.BufferedObservability) == 0 &&
|
||||
payload.HealthEvents == nil {
|
||||
return
|
||||
}
|
||||
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return
|
||||
}
|
||||
|
||||
accessLogRecords, err := buildNodeAccessLogRecords(nodeID, payload.AccessLogs, payload.BufferedObservability, reportedAt)
|
||||
if err != nil {
|
||||
zap.L().Error("build heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
if err := conn.Transaction(func(tx *gorm.DB) error {
|
||||
if err := persistNodeSystemProfile(tx, nodeID, payload.Profile, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistBufferedObservability(tx, nodeID, payload.BufferedObservability, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeMetricSnapshot(tx, nodeID, payload.Snapshot, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeOpenrestyObservation(tx, nodeID, payload.OpenrestyObservation, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeTrafficReport(tx, nodeID, payload.TrafficReport, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if payload.HealthEvents != nil {
|
||||
if err := reconcileNodeHealthEvents(tx, nodeID, payload.HealthEvents, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
if err := persistNodeAccessLogs(ctx, nodeID, accessLogRecords, reportedAt); err != nil {
|
||||
zap.L().Error("persist heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
func persistBufferedObservability(tx *gorm.DB, nodeID string, records []BufferedObservabilityRecord, reportedAt time.Time) error {
|
||||
for _, record := range records {
|
||||
if err := persistNodeMetricSnapshot(tx, nodeID, record.Snapshot, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeOpenrestyObservation(tx, nodeID, record.OpenrestyObservation, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := persistNodeTrafficReport(tx, nodeID, record.TrafficReport, reportedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *NodeSystemProfile, reportedAt time.Time) error {
|
||||
if profile == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.OpenFlareNodeSystemProfile{
|
||||
NodeID: nodeID,
|
||||
Hostname: strings.TrimSpace(profile.Hostname),
|
||||
OSName: strings.TrimSpace(profile.OSName),
|
||||
OSVersion: strings.TrimSpace(profile.OSVersion),
|
||||
KernelVersion: strings.TrimSpace(profile.KernelVersion),
|
||||
Architecture: strings.TrimSpace(profile.Architecture),
|
||||
CPUModel: strings.TrimSpace(profile.CPUModel),
|
||||
CPUCores: profile.CPUCores,
|
||||
TotalMemoryBytes: profile.TotalMemoryBytes,
|
||||
TotalDiskBytes: profile.TotalDiskBytes,
|
||||
UptimeSeconds: profile.UptimeSeconds,
|
||||
ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt),
|
||||
}
|
||||
return tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "node_id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"hostname",
|
||||
"os_name",
|
||||
"os_version",
|
||||
"kernel_version",
|
||||
"architecture",
|
||||
"cpu_model",
|
||||
"cpu_cores",
|
||||
"total_memory_bytes",
|
||||
"total_disk_bytes",
|
||||
"uptime_seconds",
|
||||
"reported_at",
|
||||
"updated_at",
|
||||
}),
|
||||
}).Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *NodeMetricSnapshot, reportedAt time.Time) error {
|
||||
if snapshot == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.OpenFlareMetricSnapshot{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: timeFromUnix(snapshot.CapturedAtUnix, reportedAt),
|
||||
CPUUsagePercent: snapshot.CPUUsagePercent,
|
||||
MemoryUsedBytes: snapshot.MemoryUsedBytes,
|
||||
MemoryTotalBytes: snapshot.MemoryTotalBytes,
|
||||
StorageUsedBytes: snapshot.StorageUsedBytes,
|
||||
StorageTotalBytes: snapshot.StorageTotalBytes,
|
||||
DiskReadBytes: snapshot.DiskReadBytes,
|
||||
DiskWriteBytes: snapshot.DiskWriteBytes,
|
||||
NetworkRxBytes: snapshot.NetworkRxBytes,
|
||||
NetworkTxBytes: snapshot.NetworkTxBytes,
|
||||
}
|
||||
exists, err := metricSnapshotExists(tx, nodeID, record.CapturedAt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeOpenrestyObservation(tx *gorm.DB, nodeID string, obs *NodeOpenrestyObservation, reportedAt time.Time) error {
|
||||
if obs == nil {
|
||||
return nil
|
||||
}
|
||||
record := &model.OpenFlareNodeObservationOpenresty{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: timeFromUnix(obs.CapturedAtUnix, reportedAt),
|
||||
OpenrestyRxBytes: obs.OpenrestyRxBytes,
|
||||
OpenrestyTxBytes: obs.OpenrestyTxBytes,
|
||||
OpenrestyConnections: obs.OpenrestyConnections,
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *NodeTrafficReport, reportedAt time.Time) error {
|
||||
if report == nil {
|
||||
return nil
|
||||
}
|
||||
if report.WindowEndedAtUnix > 0 && report.WindowStartedAtUnix > report.WindowEndedAtUnix {
|
||||
return errors.New("traffic report window_started_at_unix 不能大于 window_ended_at_unix")
|
||||
}
|
||||
record := &model.OpenFlareRequestReport{
|
||||
NodeID: nodeID,
|
||||
WindowStartedAt: timeFromUnix(report.WindowStartedAtUnix, reportedAt),
|
||||
WindowEndedAt: timeFromUnix(report.WindowEndedAtUnix, reportedAt),
|
||||
RequestCount: report.RequestCount,
|
||||
ErrorCount: report.ErrorCount,
|
||||
UniqueVisitorCount: report.UniqueVisitorCount,
|
||||
StatusCodesJSON: marshalJSON(report.StatusCodes),
|
||||
TopDomainsJSON: marshalJSON(report.TopDomains),
|
||||
SourceCountriesJSON: marshalJSON(report.SourceCountries),
|
||||
}
|
||||
exists, err := requestReportExists(tx, nodeID, record.WindowStartedAt, record.WindowEndedAt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
return tx.Create(record).Error
|
||||
}
|
||||
|
||||
func buildNodeAccessLogRecords(nodeID string, direct []NodeAccessLog, buffered []BufferedObservabilityRecord, reportedAt time.Time) ([]*model.OpenFlareAccessLog, error) {
|
||||
total := len(direct)
|
||||
for _, record := range buffered {
|
||||
total += len(record.AccessLogs)
|
||||
}
|
||||
if total == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
resolver, err := newAccessLogRegionResolver()
|
||||
if err != nil {
|
||||
slog.Warn("initialize access log geo resolver failed", "node_id", nodeID, "error", err)
|
||||
}
|
||||
if resolver != nil {
|
||||
defer resolver.Close()
|
||||
}
|
||||
|
||||
records := make([]*model.OpenFlareAccessLog, 0, total)
|
||||
appendLogs := func(logs []NodeAccessLog) {
|
||||
for _, item := range logs {
|
||||
record := &model.OpenFlareAccessLog{
|
||||
NodeID: nodeID,
|
||||
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
|
||||
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
|
||||
Region: "",
|
||||
Host: strings.TrimSpace(item.Host),
|
||||
Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength),
|
||||
StatusCode: item.StatusCode,
|
||||
}
|
||||
if resolver != nil {
|
||||
record.Region = resolver.Resolve(record.RemoteAddr)
|
||||
}
|
||||
records = append(records, record)
|
||||
}
|
||||
}
|
||||
appendLogs(direct)
|
||||
for _, record := range buffered {
|
||||
appendLogs(record.AccessLogs)
|
||||
}
|
||||
return records, nil
|
||||
}
|
||||
|
||||
func persistNodeAccessLogs(ctx context.Context, nodeID string, records []*model.OpenFlareAccessLog, reportedAt time.Time) error {
|
||||
if len(records) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := model.InsertOpenFlareAccessLogsBatch(ctx, records); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := model.DeleteOpenFlareAccessLogsByNodeBefore(ctx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow))
|
||||
return err
|
||||
}
|
||||
|
||||
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error {
|
||||
return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
|
||||
}
|
||||
|
||||
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types.
|
||||
func ReconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
|
||||
activeTypes := make(map[string]NodeHealthEvent, len(events))
|
||||
for _, event := range events {
|
||||
eventType := normalizeHealthEventType(event.EventType)
|
||||
if eventType == "" {
|
||||
continue
|
||||
}
|
||||
if len(managedEventTypes) > 0 {
|
||||
if _, ok := managedEventTypes[eventType]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
event.EventType = eventType
|
||||
event.Severity = normalizeHealthSeverity(event.Severity)
|
||||
if event.TriggeredAtUnix <= 0 {
|
||||
event.TriggeredAtUnix = reportedAt.Unix()
|
||||
}
|
||||
activeTypes[eventType] = event
|
||||
}
|
||||
|
||||
var activeEvents []*model.OpenFlareHealthEvent
|
||||
query := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive)
|
||||
if len(managedEventTypes) > 0 {
|
||||
scopedTypes := make([]string, 0, len(managedEventTypes))
|
||||
for eventType := range managedEventTypes {
|
||||
eventType = normalizeHealthEventType(eventType)
|
||||
if eventType != "" {
|
||||
scopedTypes = append(scopedTypes, eventType)
|
||||
}
|
||||
}
|
||||
if len(scopedTypes) == 0 {
|
||||
return nil
|
||||
}
|
||||
query = query.Where("event_type IN ?", scopedTypes)
|
||||
}
|
||||
if err := query.Find(&activeEvents).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
activeByType := make(map[string]*model.OpenFlareHealthEvent, len(activeEvents))
|
||||
for _, event := range activeEvents {
|
||||
activeByType[event.EventType] = event
|
||||
}
|
||||
|
||||
for eventType, event := range activeTypes {
|
||||
triggeredAt := timeFromUnix(event.TriggeredAtUnix, reportedAt)
|
||||
if existing, ok := activeByType[eventType]; ok {
|
||||
existing.Severity = event.Severity
|
||||
existing.Message = normalizeHealthEventMessage(event.Message)
|
||||
existing.LastTriggeredAt = triggeredAt
|
||||
existing.ReportedAt = reportedAt
|
||||
existing.MetadataJSON = marshalJSON(event.Metadata)
|
||||
existing.ResolvedAt = nil
|
||||
if err := tx.Save(existing).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
record := &model.OpenFlareHealthEvent{
|
||||
NodeID: nodeID,
|
||||
EventType: eventType,
|
||||
Severity: event.Severity,
|
||||
Status: healthEventStatusActive,
|
||||
Message: normalizeHealthEventMessage(event.Message),
|
||||
FirstTriggeredAt: triggeredAt,
|
||||
LastTriggeredAt: triggeredAt,
|
||||
ReportedAt: reportedAt,
|
||||
MetadataJSON: marshalJSON(event.Metadata),
|
||||
}
|
||||
if err := tx.Create(record).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
for _, existing := range activeEvents {
|
||||
if _, ok := activeTypes[existing.EventType]; ok {
|
||||
continue
|
||||
}
|
||||
resolvedAt := reportedAt
|
||||
existing.Status = healthEventStatusResolved
|
||||
existing.ReportedAt = reportedAt
|
||||
existing.ResolvedAt = &resolvedAt
|
||||
if err := tx.Save(existing).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func metricSnapshotExists(tx *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) {
|
||||
var count int64
|
||||
if err := tx.Model(&model.OpenFlareMetricSnapshot{}).
|
||||
Where("node_id = ? AND captured_at = ?", nodeID, capturedAt).
|
||||
Limit(1).
|
||||
Count(&count).Error; err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func requestReportExists(tx *gorm.DB, nodeID string, windowStartedAt, windowEndedAt time.Time) (bool, error) {
|
||||
var count int64
|
||||
if err := tx.Model(&model.OpenFlareRequestReport{}).
|
||||
Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, windowStartedAt, windowEndedAt).
|
||||
Limit(1).
|
||||
Count(&count).Error; err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func normalizeHealthEventType(eventType string) string {
|
||||
eventType = strings.TrimSpace(strings.ToLower(eventType))
|
||||
eventType = strings.ReplaceAll(eventType, " ", "_")
|
||||
return eventType
|
||||
}
|
||||
|
||||
func normalizeHealthSeverity(severity string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(severity)) {
|
||||
case healthSeverityCritical:
|
||||
return healthSeverityCritical
|
||||
case healthSeverityInfo:
|
||||
return healthSeverityInfo
|
||||
default:
|
||||
return healthSeverityWarning
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeHealthEventMessage(message string) string {
|
||||
return truncateForDatabase(message, 4096)
|
||||
}
|
||||
|
||||
func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
|
||||
if unixSeconds <= 0 {
|
||||
return fallback
|
||||
}
|
||||
return time.Unix(unixSeconds, 0).UTC()
|
||||
}
|
||||
|
||||
// MarshalJSON serializes a value for database JSON columns.
|
||||
func MarshalJSON(value any) string {
|
||||
return marshalJSON(value)
|
||||
}
|
||||
|
||||
func marshalJSON(value any) string {
|
||||
if value == nil {
|
||||
return ""
|
||||
}
|
||||
raw, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(raw)
|
||||
}
|
||||
@@ -0,0 +1,212 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// RegisterHandler registers or discovers an agent node.
|
||||
// @Summary 注册或发现 Agent 节点
|
||||
// @Description 使用节点 access token 重新注册,或使用全局 discovery token 发现新节点;请求头需携带 X-Agent-Token
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.NodePayload true "节点上报数据"
|
||||
// @Success 200 {object} response.Any{data=agent.RegistrationResponse} "注册成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/nodes/register [post]
|
||||
func RegisterHandler(c *gin.Context) {
|
||||
var payload NodePayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
|
||||
|
||||
var (
|
||||
result *RegistrationResponse
|
||||
err error
|
||||
)
|
||||
if authNode, ok := AgentNodeFromContext(c); ok {
|
||||
result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload)
|
||||
} else {
|
||||
result, err = RegisterWithDiscovery(c.Request.Context(), payload)
|
||||
}
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// HeartbeatHandler records agent heartbeat state.
|
||||
// @Summary Agent 心跳上报
|
||||
// @Description 上报节点状态、指标与健康事件,返回远程控制配置与活跃配置元信息
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.NodePayload true "心跳数据"
|
||||
// @Success 200 {object} response.Any{data=agent.HeartbeatResponse} "心跳成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/nodes/heartbeat [post]
|
||||
func HeartbeatHandler(c *gin.Context) {
|
||||
var payload NodePayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
|
||||
|
||||
authNode, ok := AgentNodeFromContext(c)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
}
|
||||
|
||||
heartbeat, err := HeartbeatNode(c.Request.Context(), authNode, payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(heartbeat))
|
||||
}
|
||||
|
||||
// GetActiveConfigHandler returns the active configuration version.
|
||||
// @Summary 获取活跃配置版本
|
||||
// @Description 返回当前生效的完整配置包,供 Agent 拉取并应用
|
||||
// @Tags openflare-agent
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Success 200 {object} response.Any{data=agent.ConfigResponse} "活跃配置"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/config-versions/active [get]
|
||||
func GetActiveConfigHandler(c *gin.Context) {
|
||||
if _, ok := AgentNodeFromContext(c); !ok {
|
||||
response.AbortUnauthorized(c, errNodeMissingFromContext)
|
||||
return
|
||||
}
|
||||
config, err := GetActiveConfig(c.Request.Context())
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(config))
|
||||
}
|
||||
|
||||
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent.
|
||||
// @Summary 同步 WAF IP 组
|
||||
// @Description 按 ID 与校验和增量同步 WAF IP 组定义
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.WAFIPGroupSyncInput true "同步请求"
|
||||
// @Success 200 {object} response.Any{data=agent.WAFIPGroupSyncResult} "同步结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/waf/ip-groups/sync [post]
|
||||
func SyncWAFIPGroupsHandler(c *gin.Context) {
|
||||
var input WAFIPGroupSyncInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
result, err := SyncWAFIPGroups(c.Request.Context(), input)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// ReportApplyLogHandler records an agent apply log entry.
|
||||
// @Summary 上报配置应用日志
|
||||
// @Description 记录 Agent 配置下发与应用结果
|
||||
// @Tags openflare-agent
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security AgentTokenAuth
|
||||
// @Param body body agent.ApplyLogPayload true "应用日志"
|
||||
// @Success 200 {object} response.Any{data=model.OpenFlareApplyLog} "日志记录"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/apply-logs [post]
|
||||
func ReportApplyLogHandler(c *gin.Context) {
|
||||
var payload ApplyLogPayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
if authNode, ok := AgentNodeFromContext(c); ok {
|
||||
payload.NodeID = authNode.NodeID
|
||||
}
|
||||
log, err := ReportApplyLog(c.Request.Context(), payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(log))
|
||||
}
|
||||
|
||||
// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent.
|
||||
// @Summary 下载 Pages 部署包
|
||||
// @Description 流式下载指定部署的静态资源压缩包,供 Agent 边缘分发
|
||||
// @Tags openflare-agent
|
||||
// @Produce application/octet-stream
|
||||
// @Security AgentTokenAuth
|
||||
// @Param deployment_id path int true "部署 ID"
|
||||
// @Success 200 {file} binary "部署包文件"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/pages/deployments/{deployment_id}/package [get]
|
||||
func DownloadPagesPackageHandler(c *gin.Context) {
|
||||
deploymentID, ok := pagesDeploymentIDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
packageObj, fileName, err := pages.OpenDeploymentPackage(c.Request.Context(), deploymentID)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
defer packageObj.Body.Close()
|
||||
c.Header("Content-Disposition", "attachment; filename="+fileName)
|
||||
if packageObj.ContentType != "" {
|
||||
c.Header("Content-Type", packageObj.ContentType)
|
||||
}
|
||||
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
|
||||
}
|
||||
|
||||
func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
|
||||
raw := c.Param("deployment_id")
|
||||
if raw == "" {
|
||||
response.AbortBadRequest(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
id64, err := strconv.ParseUint(raw, 10, 64)
|
||||
if err != nil || id64 == 0 {
|
||||
response.AbortBadRequest(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
return uint(id64), true
|
||||
}
|
||||
|
||||
// AgentWebSocketHandler upgrades an authenticated agent websocket connection.
|
||||
// @Summary Agent WebSocket 连接
|
||||
// @Description 升级为 WebSocket 长连接,用于实时推送配置同步、WAF IP 组等指令;需携带 X-Agent-Token
|
||||
// @Tags openflare-agent
|
||||
// @Security AgentTokenAuth
|
||||
// @Failure 401 {object} response.Any "Token 无效"
|
||||
// @Router /api/v1/agent/ws [get]
|
||||
func AgentWebSocketHandler(c *gin.Context) {
|
||||
authNode, ok := AgentNodeFromContext(c)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errInvalidAgentToken)
|
||||
return
|
||||
}
|
||||
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
applyResultOK = "success"
|
||||
applyResultWarn = "warning"
|
||||
applyResultFailed = "failed"
|
||||
)
|
||||
|
||||
// NodePayload is the agent register/heartbeat payload.
|
||||
type NodePayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"ext_version"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastError string `json:"last_error"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
OpenrestyMessage string `json:"openresty_message"`
|
||||
Profile *NodeSystemProfile `json:"profile,omitempty"`
|
||||
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
|
||||
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
|
||||
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
|
||||
BufferedObservability []BufferedObservabilityRecord `json:"buffered_observability,omitempty"`
|
||||
HealthEvents []NodeHealthEvent `json:"health_events"`
|
||||
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
|
||||
}
|
||||
|
||||
// ApplyLogPayload is the agent apply log report payload.
|
||||
type ApplyLogPayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Version string `json:"version"`
|
||||
Result string `json:"result"`
|
||||
Message string `json:"message"`
|
||||
Checksum string `json:"checksum"`
|
||||
MainConfigChecksum string `json:"main_config_checksum"`
|
||||
RouteConfigChecksum string `json:"route_config_checksum"`
|
||||
SupportFileCount int `json:"support_file_count"`
|
||||
}
|
||||
|
||||
// RegistrationResponse is returned after agent registration.
|
||||
type RegistrationResponse struct {
|
||||
NodeID string `json:"node_id"`
|
||||
AccessToken string `json:"access_token"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
// Settings carries remote agent control flags.
|
||||
type Settings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
UpdateRepo string `json:"update_repo"`
|
||||
UpdateNow bool `json:"update_now"`
|
||||
UpdateChannel string `json:"update_channel"`
|
||||
UpdateTag string `json:"update_tag"`
|
||||
RestartOpenrestyNow bool `json:"restart_openresty_now"`
|
||||
}
|
||||
|
||||
// ActiveConfigMeta summarizes the active configuration version.
|
||||
type ActiveConfigMeta struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// SupportFile is a configuration support artifact shipped to agents.
|
||||
type SupportFile struct {
|
||||
Path string `json:"path"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// ConfigResponse is the full active config payload for agents.
|
||||
type ConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
SourceConfigJSON string `json:"source_config_json"`
|
||||
SupportFiles []SupportFile `json:"support_files"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// WAFIPGroup is a WAF IP group snapshot for agents.
|
||||
type WAFIPGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IPList []string `json:"ip_list"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncInput requests changed WAF IP groups.
|
||||
type WAFIPGroupSyncInput struct {
|
||||
IDs []uint `json:"ids"`
|
||||
Checksums map[string]string `json:"checksums"`
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncResult returns synced WAF IP groups.
|
||||
type WAFIPGroupSyncResult struct {
|
||||
Groups []WAFIPGroup `json:"groups"`
|
||||
}
|
||||
|
||||
// HeartbeatResponse is the heartbeat handler result.
|
||||
type HeartbeatResponse struct {
|
||||
Node *model.OpenFlareNode `json:"node"`
|
||||
AgentSettings *Settings `json:"agent_settings"`
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config"`
|
||||
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
|
||||
}
|
||||
@@ -0,0 +1,205 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
type snapshotWAFRuleGroupRef struct {
|
||||
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
|
||||
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
|
||||
}
|
||||
|
||||
type snapshotWAFSection struct {
|
||||
RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"`
|
||||
}
|
||||
|
||||
type activeConfigSnapshot struct {
|
||||
WAF snapshotWAFSection `json:"waf"`
|
||||
}
|
||||
|
||||
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
|
||||
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
|
||||
return buildAgentWAFIPGroups(ctx, ids)
|
||||
}
|
||||
|
||||
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
|
||||
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
|
||||
targetIDs := uniqueUintIDs(ids)
|
||||
if len(targetIDs) == 0 {
|
||||
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
targetIDs = activeIDs
|
||||
}
|
||||
if len(targetIDs) == 0 {
|
||||
return []WAFIPGroup{}, nil
|
||||
}
|
||||
groups, err := buildAgentWAFIPGroups(ctx, targetIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
changed := make([]WAFIPGroup, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum {
|
||||
continue
|
||||
}
|
||||
changed = append(changed, group)
|
||||
}
|
||||
return changed, nil
|
||||
}
|
||||
|
||||
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
|
||||
ids = uniqueUintIDs(ids)
|
||||
if len(ids) == 0 {
|
||||
return []WAFIPGroup{}, nil
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
groups, err := model.ListOpenFlareWAFIPGroupsByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups))
|
||||
for _, group := range groups {
|
||||
groupByID[group.ID] = group
|
||||
}
|
||||
result := make([]WAFIPGroup, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
group := groupByID[id]
|
||||
if group == nil {
|
||||
continue
|
||||
}
|
||||
agentGroup, err := buildAgentWAFIPGroup(group)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, agentGroup)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func buildAgentWAFIPGroup(group *model.OpenFlareWAFIPGroup) (WAFIPGroup, error) {
|
||||
if group == nil {
|
||||
return WAFIPGroup{}, errors.New("IP 组不存在")
|
||||
}
|
||||
ips, err := decodeWAFIPGroupStringList(group.IPList)
|
||||
if err != nil {
|
||||
return WAFIPGroup{}, err
|
||||
}
|
||||
if !group.Enabled {
|
||||
ips = []string{}
|
||||
}
|
||||
agentGroup := WAFIPGroup{
|
||||
ID: group.ID,
|
||||
Name: group.Name,
|
||||
Type: group.Type,
|
||||
Enabled: group.Enabled,
|
||||
IPList: ips,
|
||||
}
|
||||
agentGroup.Checksum = checksumAgentWAFIPGroup(agentGroup)
|
||||
return agentGroup, nil
|
||||
}
|
||||
|
||||
func checksumAgentWAFIPGroup(group WAFIPGroup) string {
|
||||
payload := struct {
|
||||
ID uint `json:"id"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IPList []string `json:"ip_list"`
|
||||
}{
|
||||
ID: group.ID,
|
||||
Enabled: group.Enabled,
|
||||
IPList: append([]string{}, group.IPList...),
|
||||
}
|
||||
sort.Strings(payload.IPList)
|
||||
data, _ := json.Marshal(payload)
|
||||
sum := sha256.Sum256(data)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
|
||||
version, err := loadActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
if isActiveConfigNotFound(err) {
|
||||
return []uint{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
snapshot, err := parseActiveConfigSnapshot(version.SnapshotJSON)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
idSet := make(map[uint]struct{})
|
||||
for _, group := range snapshot.WAF.RuleGroups {
|
||||
for _, id := range group.IPWhitelistGroups {
|
||||
if id > 0 {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
for _, id := range group.IPBlacklistGroups {
|
||||
if id > 0 {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
ids := make([]uint, 0, len(idSet))
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
|
||||
text := strings.TrimSpace(snapshotJSON)
|
||||
if text == "" {
|
||||
return &activeConfigSnapshot{}, nil
|
||||
}
|
||||
var snapshot activeConfigSnapshot
|
||||
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if snapshot.WAF.RuleGroups == nil {
|
||||
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
|
||||
}
|
||||
return &snapshot, nil
|
||||
}
|
||||
|
||||
func decodeWAFIPGroupStringList(raw string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(text), &items); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func uniqueUintIDs(ids []uint) []uint {
|
||||
normalized := make([]uint, 0, len(ids))
|
||||
seen := make(map[uint]struct{}, len(ids))
|
||||
for _, id := range ids {
|
||||
if id == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupWAFIPGroupTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.OpenFlareWAFIPGroup{},
|
||||
&configVersionRecord{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
|
||||
t.Helper()
|
||||
|
||||
snapshot := map[string]any{
|
||||
"routes": []any{},
|
||||
"waf": map[string]any{
|
||||
"rule_groups": []map[string]any{
|
||||
{
|
||||
"id": 1,
|
||||
"name": "agent refs",
|
||||
"enabled": true,
|
||||
"ip_blacklist_group_ids": []uint{ipGroupID},
|
||||
},
|
||||
},
|
||||
"bindings": []any{},
|
||||
},
|
||||
}
|
||||
snapshotJSON, err := json.Marshal(snapshot)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, db.DB(ctx).Create(&configVersionRecord{
|
||||
Version: "20260618-001",
|
||||
SnapshotJSON: string(snapshotJSON),
|
||||
Checksum: "test-checksum",
|
||||
IsActive: true,
|
||||
}).Error)
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: "agent runtime group",
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `["203.0.113.44"]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
|
||||
|
||||
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, groups, 1)
|
||||
assert.Equal(t, ipGroup.ID, groups[0].ID)
|
||||
assert.Equal(t, "203.0.113.44", groups[0].IPList[0])
|
||||
assert.NotEmpty(t, groups[0].Checksum)
|
||||
|
||||
groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10)
|
||||
same, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, same)
|
||||
|
||||
ipGroup.IPList = `["203.0.113.45"]`
|
||||
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
|
||||
delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, delta, 1)
|
||||
assert.Equal(t, ipGroup.ID, delta[0].ID)
|
||||
assert.Equal(t, "203.0.113.45", delta[0].IPList[0])
|
||||
assert.NotEqual(t, groups[0].Checksum, delta[0].Checksum)
|
||||
}
|
||||
|
||||
func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: "sync group",
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `["198.51.100.10"]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
|
||||
|
||||
result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
|
||||
IDs: []uint{ipGroup.ID},
|
||||
Checksums: map[string]string{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, result.Groups, 1)
|
||||
assert.Equal(t, ipGroup.ID, result.Groups[0].ID)
|
||||
assert.Equal(t, "198.51.100.10", result.Groups[0].IPList[0])
|
||||
|
||||
result, err = SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
|
||||
IDs: []uint{ipGroup.ID},
|
||||
Checksums: map[string]string{
|
||||
strconv.FormatUint(uint64(ipGroup.ID), 10): result.Groups[0].Checksum,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, result.Groups)
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: "disabled group",
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `["203.0.113.10"]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
ipGroup.Enabled = false
|
||||
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
|
||||
|
||||
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, groups, 1)
|
||||
assert.False(t, groups[0].Enabled)
|
||||
assert.Empty(t, groups[0].IPList)
|
||||
assert.NotEmpty(t, groups[0].Checksum)
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// HandleWSStatus processes an agent websocket status payload (replaces HTTP heartbeat in WS mode).
|
||||
func HandleWSStatus(ctx context.Context, nodeID, remoteAddr string, rawPayload json.RawMessage) {
|
||||
var payload NodePayload
|
||||
if err := json.Unmarshal(rawPayload, &payload); err != nil {
|
||||
slog.Debug("agent ws status payload decode failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
authNode, err := model.GetOpenFlareNodeByNodeID(ctx, nodeID)
|
||||
if err != nil {
|
||||
slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
payload.IP = resolveReportedNodeIP(payload.IP, remoteAddr)
|
||||
response, err := HeartbeatNode(ctx, authNode, payload)
|
||||
if err != nil {
|
||||
slog.Debug("agent ws status handling failed", "node_id", nodeID, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
settingsSent := false
|
||||
if response.AgentSettings != nil {
|
||||
settingsSent = ofws.SendAgentSettings(nodeID, response.AgentSettings)
|
||||
}
|
||||
activeConfigSent := false
|
||||
if response.ActiveConfig != nil {
|
||||
activeConfigSent = ofws.SendAgentActiveConfig(nodeID, response.ActiveConfig)
|
||||
}
|
||||
wafIPGroupsSent := false
|
||||
if len(response.WAFIPGroups) > 0 {
|
||||
wafIPGroupsSent = ofws.SendAgentWAFIPGroups(nodeID, response.WAFIPGroups)
|
||||
}
|
||||
|
||||
slog.Debug("agent ws status processed",
|
||||
"node_id", nodeID,
|
||||
"current_version", payload.CurrentVersion,
|
||||
"openresty_status", payload.OpenrestyStatus,
|
||||
"settings_sent", settingsSent,
|
||||
"active_config_sent", activeConfigSent,
|
||||
"waf_ip_groups_sent", wafIPGroupsSent,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package apiutil provides HTTP helpers for OpenFlare v1 custom API handlers.
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const errInvalidParams = "参数错误"
|
||||
const errInvalidID = "无效的 ID"
|
||||
|
||||
// BindJSON binds JSON body; returns false after aborting with 400.
|
||||
func BindJSON(c *gin.Context, dst any) bool {
|
||||
if err := c.ShouldBindJSON(dst); err != nil {
|
||||
response.AbortBadRequest(c, errInvalidParams)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// IDParam parses :id from the URL path.
|
||||
func IDParam(c *gin.Context) (uint, bool) {
|
||||
raw := c.Param("id")
|
||||
if raw == "" {
|
||||
response.AbortBadRequest(c, errInvalidID)
|
||||
return 0, false
|
||||
}
|
||||
id64, err := strconv.ParseUint(raw, 10, 64)
|
||||
if err != nil || id64 == 0 {
|
||||
response.AbortBadRequest(c, errInvalidID)
|
||||
return 0, false
|
||||
}
|
||||
return uint(id64), true
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// AbortNotFoundIfMissing maps gorm.ErrRecordNotFound to 404; other errors to 400.
|
||||
func AbortNotFoundIfMissing(c *gin.Context, err error, notFoundMsg string) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, notFoundMsg)
|
||||
return true
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return true
|
||||
}
|
||||
|
||||
// AbortBadRequestOnError writes a 400 for any non-nil error.
|
||||
func AbortBadRequestOnError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// AdminMiddlewares returns Wavelet-standard middlewares for OpenFlare console routes.
|
||||
// OpenFlare no longer distinguishes Admin vs Root tiers; all management endpoints share
|
||||
// the same gate: user.IsAdmin for session users, token_admin for Access Token callers.
|
||||
func AdminMiddlewares() []gin.HandlerFunc {
|
||||
return []gin.HandlerFunc{oauth.LoginRequired(), admin.LoginAdminRequired()}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupAdminMiddlewareTest(t *testing.T) (*gin.Engine, *gorm.DB, func()) {
|
||||
t.Helper()
|
||||
|
||||
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, dbConn.AutoMigrate(&model.User{}, &model.AccessToken{}))
|
||||
db.SetDB(dbConn)
|
||||
|
||||
sessionCookieName := "test_admin_middleware_session"
|
||||
if config.Config.App.SessionCookieName != "" {
|
||||
sessionCookieName = config.Config.App.SessionCookieName
|
||||
}
|
||||
store := cookie.NewStore([]byte("test_admin_middleware_session_secret"))
|
||||
store.Options(oauth.GetSessionOptions(3600))
|
||||
engine := testhelper.NewTestGinEngine(sessions.Sessions(sessionCookieName, store))
|
||||
protected := engine.Group("/protected", AdminMiddlewares()...)
|
||||
protected.GET("", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"ok": true}))
|
||||
})
|
||||
|
||||
cleanup := func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
|
||||
return engine, dbConn, cleanup
|
||||
}
|
||||
|
||||
func seedUser(t *testing.T, dbConn *gorm.DB, username 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, dbConn.Create(user).Error)
|
||||
return user
|
||||
}
|
||||
|
||||
func seedAccessToken(t *testing.T, dbConn *gorm.DB, user *model.User, isAdmin bool) string {
|
||||
t.Helper()
|
||||
|
||||
token, err := model.GenerateTokenString()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, dbConn.Create(&model.AccessToken{
|
||||
UserID: user.ID,
|
||||
Name: user.Username + "-token",
|
||||
TokenHash: model.HashToken(token),
|
||||
MaskedToken: model.MaskTokenString(token),
|
||||
IsAdmin: isAdmin,
|
||||
}).Error)
|
||||
return token
|
||||
}
|
||||
|
||||
func decodeResponse(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 TestAdminRequiredUnauthenticated(t *testing.T) {
|
||||
engine, _, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.NotEmpty(t, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestAdminRequiredNonAdminToken(t *testing.T) {
|
||||
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
user := seedUser(t, dbConn, "regular", false)
|
||||
token := seedAccessToken(t, dbConn, user, false)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set("X-Access-Token", token)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.Equal(t, admin.TokenAdminRequired, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestAdminRequiredAdminWithoutTokenAdmin(t *testing.T) {
|
||||
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
user := seedUser(t, dbConn, "admin-no-token-admin", true)
|
||||
token := seedAccessToken(t, dbConn, user, false)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set("X-Access-Token", token)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.Equal(t, admin.TokenAdminRequired, resp.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestAdminRequiredAdminWithTokenAdmin(t *testing.T) {
|
||||
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
|
||||
defer cleanup()
|
||||
|
||||
user := seedUser(t, dbConn, "admin", true)
|
||||
token := seedAccessToken(t, dbConn, user, true)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
|
||||
req.Header.Set("X-Access-Token", token)
|
||||
engine.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
resp := decodeResponse(t, rec)
|
||||
assert.Empty(t, resp.ErrorMsg)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apiutil
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// RegisterCollection registers a collection endpoint on both "" and "/" so requests
|
||||
// work with or without a trailing slash.
|
||||
func RegisterCollection(route *gin.RouterGroup, method string, handlers ...gin.HandlerFunc) {
|
||||
route.Handle(method, "/", handlers...)
|
||||
if !strings.HasSuffix(route.BasePath(), "/") {
|
||||
route.Handle(method, "", handlers...)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
const (
|
||||
errRetentionDaysOutOfRange = "retention_days 必须在 1 到 3650 之间"
|
||||
)
|
||||
@@ -0,0 +1,128 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultApplyLogPageSize = 20
|
||||
maxApplyLogPageSize = 200
|
||||
maxApplyLogRetentionDays = 3650
|
||||
)
|
||||
|
||||
// ListQuery filters apply logs for paginated listing.
|
||||
type ListQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
PageNo int `json:"pageNo"`
|
||||
PageSize int `json:"pageSize"`
|
||||
}
|
||||
|
||||
// ListResult is the paginated apply log list response.
|
||||
type ListResult struct {
|
||||
Rows []*model.OpenFlareApplyLog `json:"rows"`
|
||||
Current int `json:"current"`
|
||||
Total int `json:"total"`
|
||||
TotalPage int `json:"totalPage"`
|
||||
}
|
||||
|
||||
// CleanupInput controls apply log cleanup behavior.
|
||||
type CleanupInput struct {
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays int `json:"retention_days"`
|
||||
}
|
||||
|
||||
// CleanupResult reports apply log cleanup outcome.
|
||||
type CleanupResult struct {
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays int `json:"retention_days"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
Cutoff *time.Time `json:"cutoff,omitempty"`
|
||||
}
|
||||
|
||||
// ListPage returns paginated apply logs with optional node_id filter.
|
||||
func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) {
|
||||
pageNo := normalizePageNo(input.PageNo)
|
||||
pageSize := normalizePageSize(input.PageSize)
|
||||
nodeID := strings.TrimSpace(input.NodeID)
|
||||
|
||||
rows, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
|
||||
NodeID: nodeID,
|
||||
PageNo: pageNo,
|
||||
PageSize: pageSize,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
total, err := model.CountOpenFlareApplyLogs(ctx, nodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
totalPage := 0
|
||||
if total > 0 {
|
||||
totalPage = int((total + int64(pageSize) - 1) / int64(pageSize))
|
||||
}
|
||||
|
||||
return &ListResult{
|
||||
Rows: rows,
|
||||
Current: pageNo,
|
||||
Total: int(total),
|
||||
TotalPage: totalPage,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Cleanup removes old apply logs or deletes all records.
|
||||
func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) {
|
||||
if input.DeleteAll {
|
||||
deleted, err := model.DeleteAllOpenFlareApplyLogs(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &CleanupResult{
|
||||
DeleteAll: true,
|
||||
DeletedCount: deleted,
|
||||
}, nil
|
||||
}
|
||||
|
||||
if input.RetentionDays <= 0 || input.RetentionDays > maxApplyLogRetentionDays {
|
||||
return nil, errors.New(errRetentionDaysOutOfRange)
|
||||
}
|
||||
|
||||
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour)
|
||||
deleted, err := model.DeleteOpenFlareApplyLogsBefore(ctx, cutoff)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &CleanupResult{
|
||||
RetentionDays: input.RetentionDays,
|
||||
DeletedCount: deleted,
|
||||
Cutoff: &cutoff,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizePageNo(pageNo int) int {
|
||||
if pageNo <= 0 {
|
||||
return 1
|
||||
}
|
||||
return pageNo
|
||||
}
|
||||
|
||||
func normalizePageSize(pageSize int) int {
|
||||
if pageSize <= 0 {
|
||||
return defaultApplyLogPageSize
|
||||
}
|
||||
if pageSize > maxApplyLogPageSize {
|
||||
return maxApplyLogPageSize
|
||||
}
|
||||
return pageSize
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupApplyLogTestDB(t *testing.T) func() {
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{})
|
||||
require.NoError(t, err)
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListPageAndCleanup(t *testing.T) {
|
||||
cleanup := setupApplyLogTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC()
|
||||
|
||||
logs := []model.OpenFlareApplyLog{
|
||||
{NodeID: "node-logs", Version: "v1", Result: "success", Message: "1", CreatedAt: now.Add(-10 * 24 * time.Hour)},
|
||||
{NodeID: "node-logs", Version: "v2", Result: "success", Message: "2", CreatedAt: now.Add(-5 * 24 * time.Hour)},
|
||||
{NodeID: "node-logs", Version: "v3", Result: "success", Message: "3", CreatedAt: now},
|
||||
}
|
||||
for i := range logs {
|
||||
require.NoError(t, db.DB(ctx).Create(&logs[i]).Error)
|
||||
}
|
||||
|
||||
pageResult, err := ListPage(ctx, ListQuery{
|
||||
NodeID: "node-logs",
|
||||
PageNo: 1,
|
||||
PageSize: 2,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 3, pageResult.Total)
|
||||
assert.Len(t, pageResult.Rows, 2)
|
||||
assert.Equal(t, 2, pageResult.TotalPage)
|
||||
assert.Equal(t, 1, pageResult.Current)
|
||||
|
||||
cleanupResult, err := Cleanup(ctx, CleanupInput{
|
||||
DeleteAll: false,
|
||||
RetentionDays: 7,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), cleanupResult.DeletedCount)
|
||||
assert.NotNil(t, cleanupResult.Cutoff)
|
||||
|
||||
remaining, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
|
||||
NodeID: "node-logs",
|
||||
PageNo: 1,
|
||||
PageSize: 10,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, remaining, 2)
|
||||
|
||||
cleanupAll, err := Cleanup(ctx, CleanupInput{DeleteAll: true})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), cleanupAll.DeletedCount)
|
||||
assert.True(t, cleanupAll.DeleteAll)
|
||||
|
||||
finalLogs, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
|
||||
NodeID: "node-logs",
|
||||
PageNo: 1,
|
||||
PageSize: 10,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, finalLogs)
|
||||
}
|
||||
|
||||
func TestCleanupInvalidRetentionDays(t *testing.T) {
|
||||
cleanup := setupApplyLogTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := Cleanup(ctx, CleanupInput{RetentionDays: 0})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
|
||||
|
||||
_, err = Cleanup(ctx, CleanupInput{RetentionDays: 4000})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package apply_log
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
// GetApplyLogs lists apply logs with pagination and optional node_id filter.
|
||||
// @Summary 获取配置下发日志
|
||||
// @Description 分页返回节点配置下发记录,支持按节点 ID 筛选,需要管理员权限
|
||||
// @Tags openflare-apply-log
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param node_id query string false "节点 ID 筛选"
|
||||
// @Param pageNo query int false "页码"
|
||||
// @Param page_no query int false "页码(别名)"
|
||||
// @Param pageSize query int false "每页数量"
|
||||
// @Param page_size query int false "每页数量(别名)"
|
||||
// @Success 200 {object} response.Any{data=apply_log.ListResult} "下发日志列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/apply-logs [get]
|
||||
func GetApplyLogs(c *gin.Context) {
|
||||
result, err := ListPage(c.Request.Context(), ListQuery{
|
||||
NodeID: c.Query("node_id"),
|
||||
PageNo: readIntQuery(c, "pageNo", "page_no"),
|
||||
PageSize: readIntQuery(c, "pageSize", "page_size"),
|
||||
})
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// CleanupApplyLogs removes old apply logs or deletes all records.
|
||||
// @Summary 清理配置下发日志
|
||||
// @Description 按保留天数清理历史下发记录,或删除全部记录,需要管理员权限
|
||||
// @Tags openflare-apply-log
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param body body apply_log.CleanupInput true "清理参数"
|
||||
// @Success 200 {object} response.Any{data=apply_log.CleanupResult} "清理结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/apply-logs/cleanup [post]
|
||||
func CleanupApplyLogs(c *gin.Context) {
|
||||
var input CleanupInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
|
||||
result, err := Cleanup(c.Request.Context(), input)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
func readIntQuery(c *gin.Context, primary, secondary string) int {
|
||||
value := c.Query(primary)
|
||||
if value == "" {
|
||||
value = c.Query(secondary)
|
||||
}
|
||||
parsed, _ := strconv.Atoi(value)
|
||||
return parsed
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package openflare
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tasks"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/uptimekuma"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/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"
|
||||
|
||||
// DatabaseAutoCleanupTask prunes observability tables by retention policy.
|
||||
DatabaseAutoCleanupTask = "openflare:database_auto_cleanup"
|
||||
// TaskTypeDatabaseAutoCleanup is the admin task type for observability cleanup.
|
||||
TaskTypeDatabaseAutoCleanup = "of_database_auto_cleanup"
|
||||
|
||||
// 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"
|
||||
)
|
||||
|
||||
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,
|
||||
}
|
||||
|
||||
// DatabaseAutoCleanupMeta describes the observability auto-cleanup task.
|
||||
var DatabaseAutoCleanupMeta = task.TaskMeta{
|
||||
Type: TaskTypeDatabaseAutoCleanup,
|
||||
AsynqTask: DatabaseAutoCleanupTask,
|
||||
Name: "OpenFlare 可观测数据自动清理",
|
||||
Description: "按保留天数清理访问日志、性能快照与请求聚合数据",
|
||||
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,
|
||||
}
|
||||
|
||||
// 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 := tasks.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
|
||||
}
|
||||
|
||||
// DatabaseAutoCleanupHandler prunes observability data when auto-cleanup is enabled.
|
||||
type DatabaseAutoCleanupHandler struct{}
|
||||
|
||||
// Execute runs retention-based cleanup for all observability targets.
|
||||
func (h *DatabaseAutoCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
if !model.DatabaseAutoCleanupEnabled {
|
||||
msg := "自动清理未启用,跳过执行"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始执行可观测数据自动清理,保留天数=%d", model.DatabaseAutoCleanupRetentionDays)
|
||||
summary, err := tasks.RunDatabaseAutoCleanupOnce(time.Now())
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "可观测数据自动清理失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
if summary == nil {
|
||||
msg := "自动清理未启用,跳过执行"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
var totalDeleted int64
|
||||
for _, item := range summary.Results {
|
||||
totalDeleted += item.DeletedCount
|
||||
task.AppendLog(ctx, "清理 %s:删除 %d 条", item.TargetLabel, item.DeletedCount)
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf(
|
||||
"可观测数据自动清理完成,保留 %d 天,共删除 %d 条",
|
||||
summary.RetentionDays,
|
||||
totalDeleted,
|
||||
)
|
||||
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) {
|
||||
if !model.UptimeKumaEnabled {
|
||||
msg := "Uptime Kuma 集成未启用,跳过执行"
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
interval := model.UptimeKumaSyncInterval
|
||||
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,84 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package openflare
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestDatabaseAutoCleanupHandlerSkipsWhenDisabled(t *testing.T) {
|
||||
previousEnabled := model.DatabaseAutoCleanupEnabled
|
||||
model.DatabaseAutoCleanupEnabled = false
|
||||
t.Cleanup(func() {
|
||||
model.DatabaseAutoCleanupEnabled = previousEnabled
|
||||
})
|
||||
|
||||
result, err := (&DatabaseAutoCleanupHandler{}).Execute(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "未启用")
|
||||
}
|
||||
|
||||
func TestDatabaseAutoCleanupHandlerDeletesRowsWhenEnabled(t *testing.T) {
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
db.SetDB(sqliteDB)
|
||||
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore())
|
||||
t.Cleanup(func() {
|
||||
resetAccessLogStore()
|
||||
db.SetDB(nil)
|
||||
})
|
||||
|
||||
now := time.Now().UTC()
|
||||
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(context.Background(), []*model.OpenFlareAccessLog{{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-48 * time.Hour),
|
||||
RemoteAddr: "203.0.113.10",
|
||||
Host: "example.com",
|
||||
Path: "/access",
|
||||
StatusCode: 200,
|
||||
}}))
|
||||
|
||||
previousEnabled := model.DatabaseAutoCleanupEnabled
|
||||
previousRetentionDays := model.DatabaseAutoCleanupRetentionDays
|
||||
model.DatabaseAutoCleanupEnabled = true
|
||||
model.DatabaseAutoCleanupRetentionDays = 1
|
||||
t.Cleanup(func() {
|
||||
model.DatabaseAutoCleanupEnabled = previousEnabled
|
||||
model.DatabaseAutoCleanupRetentionDays = previousRetentionDays
|
||||
})
|
||||
|
||||
result, err := (&DatabaseAutoCleanupHandler{}).Execute(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "共删除")
|
||||
|
||||
rows, err := model.ListOpenFlareAccessLogs(context.Background(), model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, rows)
|
||||
}
|
||||
|
||||
func TestUptimeKumaSyncHandlerSkipsWhenDisabled(t *testing.T) {
|
||||
previousEnabled := model.UptimeKumaEnabled
|
||||
model.UptimeKumaEnabled = false
|
||||
t.Cleanup(func() {
|
||||
model.UptimeKumaEnabled = previousEnabled
|
||||
})
|
||||
|
||||
result, err := (&UptimeKumaSyncHandler{}).Execute(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
assert.Contains(t, result.Message, "未启用")
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
const (
|
||||
errNoActiveVersion = "当前没有激活版本"
|
||||
errNoEnabledRoutes = "没有可发布的启用规则"
|
||||
errNoChangesToPublish = "当前规则没有变更,不能重复发布"
|
||||
errVersionConflict = "版本号生成冲突,请重试"
|
||||
errInvalidSnapshotFormat = "历史版本快照格式不合法"
|
||||
)
|
||||
@@ -0,0 +1,352 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
type customHeaderInput struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
func normalizeProxyRouteSiteName(route *model.ProxyRoute, raw, primaryDomain string) string {
|
||||
siteName := strings.TrimSpace(raw)
|
||||
if siteName != "" {
|
||||
return siteName
|
||||
}
|
||||
if route != nil && strings.TrimSpace(route.SiteName) != "" {
|
||||
return strings.TrimSpace(route.SiteName)
|
||||
}
|
||||
return primaryDomain
|
||||
}
|
||||
|
||||
func normalizeProxyRouteDomains(rawDomains []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(rawDomains))
|
||||
seen := make(map[string]struct{}, len(rawDomains))
|
||||
for _, rawDomain := range rawDomains {
|
||||
domain := strings.ToLower(strings.TrimSpace(rawDomain))
|
||||
if domain == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, fmt.Errorf("domain %q is invalid", rawDomain)
|
||||
}
|
||||
if _, ok := seen[domain]; ok {
|
||||
continue
|
||||
}
|
||||
seen[domain] = struct{}{}
|
||||
normalized = append(normalized, domain)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, fmt.Errorf("domain is required")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return normalizeProxyRouteDomains([]string{fallbackDomain})
|
||||
}
|
||||
var domains []string
|
||||
if err := json.Unmarshal([]byte(text), &domains); err != nil {
|
||||
return nil, fmt.Errorf("domains payload is invalid")
|
||||
}
|
||||
return normalizeProxyRouteDomains(domains)
|
||||
}
|
||||
|
||||
func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return normalizeUpstreams(fallbackOriginURL, nil)
|
||||
}
|
||||
var upstreams []string
|
||||
if err := json.Unmarshal([]byte(text), &upstreams); err != nil {
|
||||
return nil, fmt.Errorf("upstreams payload is invalid")
|
||||
}
|
||||
return normalizeUpstreams(fallbackOriginURL, upstreams)
|
||||
}
|
||||
|
||||
func normalizeUpstreams(originURL string, upstreams []string) ([]string, error) {
|
||||
candidates := upstreams
|
||||
if len(candidates) == 0 {
|
||||
candidates = []string{originURL}
|
||||
}
|
||||
normalized := make([]string, 0, len(candidates))
|
||||
seen := make(map[string]struct{}, len(candidates))
|
||||
for _, item := range candidates {
|
||||
value := strings.TrimSpace(item)
|
||||
if value == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[value]; ok {
|
||||
continue
|
||||
}
|
||||
seen[value] = struct{}{}
|
||||
normalized = append(normalized, value)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, fmt.Errorf("upstream is required")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func decodeStoredCustomHeaders(raw string) ([]customHeaderInput, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []customHeaderInput{}, nil
|
||||
}
|
||||
var headers []customHeaderInput
|
||||
if err := json.Unmarshal([]byte(text), &headers); err != nil {
|
||||
return nil, fmt.Errorf("custom_headers payload is invalid")
|
||||
}
|
||||
return headers, nil
|
||||
}
|
||||
|
||||
func decodeStoredCacheRules(raw string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
var rules []string
|
||||
if err := json.Unmarshal([]byte(text), &rules); err != nil {
|
||||
return nil, fmt.Errorf("cache_rules payload is invalid")
|
||||
}
|
||||
normalized := make([]string, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
item := strings.TrimSpace(rule)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func decodeStoredCertIDs(raw string, fallbackCertID *uint) ([]uint, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
if fallbackCertID == nil || *fallbackCertID == 0 {
|
||||
return []uint{}, nil
|
||||
}
|
||||
return []uint{*fallbackCertID}, nil
|
||||
}
|
||||
var certIDs []uint
|
||||
if err := json.Unmarshal([]byte(text), &certIDs); err != nil {
|
||||
return nil, fmt.Errorf("cert_ids payload is invalid")
|
||||
}
|
||||
normalized := make([]uint, 0, len(certIDs))
|
||||
seen := make(map[uint]struct{}, len(certIDs))
|
||||
for _, certID := range certIDs {
|
||||
if certID == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[certID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[certID] = struct{}{}
|
||||
normalized = append(normalized, certID)
|
||||
}
|
||||
if len(normalized) == 0 && fallbackCertID != nil && *fallbackCertID != 0 {
|
||||
return []uint{*fallbackCertID}, nil
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func resolveDomainCertIDs(domains []string, certIDs []uint, rawDomainCertIDs string) ([]uint, error) {
|
||||
text := strings.TrimSpace(rawDomainCertIDs)
|
||||
if text != "" {
|
||||
var domainCertIDs []uint
|
||||
if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil {
|
||||
return nil, fmt.Errorf("domain_cert_ids payload is invalid")
|
||||
}
|
||||
if len(domains) > 0 && len(domainCertIDs) != len(domains) {
|
||||
return nil, fmt.Errorf("domain_cert_ids length is invalid")
|
||||
}
|
||||
return domainCertIDs, nil
|
||||
}
|
||||
if len(certIDs) == 0 {
|
||||
return []uint{}, nil
|
||||
}
|
||||
if len(certIDs) == 1 {
|
||||
result := make([]uint, len(domains))
|
||||
for index := range result {
|
||||
result[index] = certIDs[0]
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
if len(certIDs) == len(domains) {
|
||||
result := make([]uint, len(certIDs))
|
||||
copy(result, certIDs)
|
||||
return result, nil
|
||||
}
|
||||
return []uint{}, nil
|
||||
}
|
||||
|
||||
func mustDecodeCertIDs(route *model.ProxyRoute) []uint {
|
||||
if route == nil {
|
||||
return []uint{}
|
||||
}
|
||||
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID)
|
||||
if err != nil {
|
||||
return []uint{}
|
||||
}
|
||||
return certIDs
|
||||
}
|
||||
|
||||
func mustDecodeDomainCertIDs(route *model.ProxyRoute, domains []string) []uint {
|
||||
if route == nil {
|
||||
return []uint{}
|
||||
}
|
||||
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID)
|
||||
if err != nil {
|
||||
return []uint{}
|
||||
}
|
||||
domainCertIDs, err := resolveDomainCertIDs(domains, certIDs, route.DomainCertIDs)
|
||||
if err != nil {
|
||||
return []uint{}
|
||||
}
|
||||
return domainCertIDs
|
||||
}
|
||||
|
||||
func normalizeUpstreamType(raw string) string {
|
||||
value := strings.ToLower(strings.TrimSpace(raw))
|
||||
switch value {
|
||||
case "tunnel", "pages":
|
||||
return value
|
||||
default:
|
||||
return "direct"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeTunnelTargetProtocol(raw string) string {
|
||||
value := strings.ToLower(strings.TrimSpace(raw))
|
||||
switch value {
|
||||
case "http", "https", "tcp":
|
||||
return value
|
||||
default:
|
||||
return "http"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizePEM(content string) string {
|
||||
return strings.TrimSpace(content) + "\n"
|
||||
}
|
||||
|
||||
func certificateCertFileName(id uint) string {
|
||||
return fmt.Sprintf("%d.crt", id)
|
||||
}
|
||||
|
||||
func certificateKeyFileName(id uint) string {
|
||||
return fmt.Sprintf("%d.key", id)
|
||||
}
|
||||
|
||||
func dedupeSupportFiles(files []SupportFile) []SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
unique := make(map[string]SupportFile, len(files))
|
||||
for _, file := range files {
|
||||
unique[file.Path] = file
|
||||
}
|
||||
result := make([]SupportFile, 0, len(unique))
|
||||
for _, file := range unique {
|
||||
result = append(result, file)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func uintPtrEqual(left *uint, right *uint) bool {
|
||||
if left == nil || right == nil {
|
||||
return left == nil && right == nil
|
||||
}
|
||||
return *left == *right
|
||||
}
|
||||
|
||||
func uintSliceEqual(left []uint, right []uint) bool {
|
||||
if len(left) != len(right) {
|
||||
return false
|
||||
}
|
||||
for index := range left {
|
||||
if left[index] != right[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func relayAgentAddress(node *model.OpenFlareNode) string {
|
||||
if node == nil {
|
||||
return ""
|
||||
}
|
||||
port := node.RelayVhostHTTPPort
|
||||
if port <= 0 {
|
||||
port = 8080
|
||||
}
|
||||
addr := strings.TrimSpace(node.RelayAgentAccessAddr)
|
||||
if addr == "" {
|
||||
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 resolveTunnelOpenRestyUpstreamURL(ctx context.Context) string {
|
||||
nodes, err := model.ListOpenFlareNodes(ctx)
|
||||
if err == nil {
|
||||
for index := range nodes {
|
||||
node := &nodes[index]
|
||||
if node.NodeType != "tunnel_relay" {
|
||||
continue
|
||||
}
|
||||
addr := relayAgentAddress(node)
|
||||
if addr != "" {
|
||||
return "http://" + addr
|
||||
}
|
||||
}
|
||||
}
|
||||
return "http://127.0.0.1:8080"
|
||||
}
|
||||
|
||||
func listWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*model.OpenFlareWAFIPGroup, error) {
|
||||
if len(ids) == 0 {
|
||||
return []*model.OpenFlareWAFIPGroup{}, nil
|
||||
}
|
||||
groups := make([]*model.OpenFlareWAFIPGroup, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups = append(groups, group)
|
||||
}
|
||||
return groups, nil
|
||||
}
|
||||
@@ -0,0 +1,559 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// ConfigPreviewResult is the preview response payload.
|
||||
type ConfigPreviewResult struct {
|
||||
SnapshotJSON string `json:"snapshot_json"`
|
||||
MainConfig string `json:"main_config"`
|
||||
RouteConfig string `json:"route_config"`
|
||||
RenderedConfig string `json:"rendered_config"`
|
||||
SupportFiles []SupportFile `json:"support_files"`
|
||||
Checksum string `json:"checksum"`
|
||||
RouteCount int `json:"route_count"`
|
||||
WebsiteCount int `json:"website_count"`
|
||||
}
|
||||
|
||||
// ConfigDiffResult is the diff response payload.
|
||||
type ConfigDiffResult struct {
|
||||
ActiveVersion string `json:"active_version,omitempty"`
|
||||
AddedSites []string `json:"added_sites"`
|
||||
RemovedSites []string `json:"removed_sites"`
|
||||
ModifiedSites []string `json:"modified_sites"`
|
||||
AddedDomains []string `json:"added_domains"`
|
||||
RemovedDomains []string `json:"removed_domains"`
|
||||
ModifiedDomains []string `json:"modified_domains"`
|
||||
MainConfigChanged bool `json:"main_config_changed"`
|
||||
WAFConfigChanged bool `json:"waf_config_changed"`
|
||||
ChangedOptionKeys []string `json:"changed_option_keys"`
|
||||
ChangedOptionDetails []ConfigOptionDiffItem `json:"changed_option_details"`
|
||||
CurrentWebsiteCount int `json:"current_website_count"`
|
||||
ActiveWebsiteCount int `json:"active_website_count"`
|
||||
}
|
||||
|
||||
// ConfigOptionDiffItem describes a changed OpenResty option.
|
||||
type ConfigOptionDiffItem struct {
|
||||
Key string `json:"key"`
|
||||
PreviousValue string `json:"previous_value"`
|
||||
CurrentValue string `json:"current_value"`
|
||||
}
|
||||
|
||||
// CleanupInput is the cleanup request payload.
|
||||
type CleanupInput struct {
|
||||
KeepCount int `json:"keep_count"`
|
||||
}
|
||||
|
||||
// CleanupResult is the cleanup response payload.
|
||||
type CleanupResult struct {
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// ListConfigVersions returns all config version summaries.
|
||||
func ListConfigVersions(ctx context.Context) ([]*model.ConfigVersionSummary, error) {
|
||||
return model.ListConfigVersionSummaries(ctx)
|
||||
}
|
||||
|
||||
// GetConfigVersionDetail returns a config version by id.
|
||||
func GetConfigVersionDetail(ctx context.Context, id uint) (*model.ConfigVersion, error) {
|
||||
return model.GetConfigVersionByID(ctx, id)
|
||||
}
|
||||
|
||||
// GetActiveConfigVersion returns the active config version.
|
||||
func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) {
|
||||
return model.GetActiveConfigVersion(ctx)
|
||||
}
|
||||
|
||||
// PreviewConfigVersion renders the current draft configuration.
|
||||
func PreviewConfigVersion(ctx context.Context) (*ConfigPreviewResult, error) {
|
||||
bundle, err := buildCurrentConfigBundle(ctx, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ConfigPreviewResult{
|
||||
SnapshotJSON: bundle.SnapshotJSON,
|
||||
MainConfig: bundle.MainConfig,
|
||||
RouteConfig: bundle.RouteConfig,
|
||||
RenderedConfig: bundle.RouteConfig,
|
||||
SupportFiles: bundle.SupportFiles,
|
||||
Checksum: bundle.Checksum,
|
||||
RouteCount: len(bundle.Routes),
|
||||
WebsiteCount: len(bundle.SnapshotRoutes),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// DiffConfigVersion compares the current draft against the active version.
|
||||
func DiffConfigVersion(ctx context.Context) (*ConfigDiffResult, error) {
|
||||
bundle, err := buildCurrentConfigBundle(ctx, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := &ConfigDiffResult{
|
||||
AddedSites: []string{},
|
||||
RemovedSites: []string{},
|
||||
ModifiedSites: []string{},
|
||||
AddedDomains: []string{},
|
||||
RemovedDomains: []string{},
|
||||
ModifiedDomains: []string{},
|
||||
ChangedOptionKeys: []string{},
|
||||
ChangedOptionDetails: []ConfigOptionDiffItem{},
|
||||
CurrentWebsiteCount: len(bundle.SnapshotRoutes),
|
||||
}
|
||||
activeVersion, err := model.GetActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
for _, route := range bundle.SnapshotRoutes {
|
||||
result.AddedSites = append(result.AddedSites, route.SiteName)
|
||||
result.AddedDomains = append(result.AddedDomains, route.Domains...)
|
||||
}
|
||||
result.MainConfigChanged = true
|
||||
result.ChangedOptionKeys = openRestyOptionKeys()
|
||||
result.ChangedOptionDetails = buildInitialOpenRestyOptionDiffs(bundle.OpenRestyConfig)
|
||||
sort.Strings(result.AddedSites)
|
||||
sort.Strings(result.AddedDomains)
|
||||
sort.Strings(result.ChangedOptionKeys)
|
||||
return result, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
result.ActiveVersion = activeVersion.Version
|
||||
activeSnapshot, err := parseSnapshotDocument(activeVersion.SnapshotJSON)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.ActiveWebsiteCount = len(activeSnapshot.Routes)
|
||||
currentSiteMap := flattenSnapshotRoutesBySite(bundle.SnapshotRoutes)
|
||||
activeSiteMap := flattenSnapshotRoutesBySite(activeSnapshot.Routes)
|
||||
for siteName, currentRoute := range currentSiteMap {
|
||||
activeRoute, ok := activeSiteMap[siteName]
|
||||
if !ok {
|
||||
result.AddedSites = append(result.AddedSites, siteName)
|
||||
continue
|
||||
}
|
||||
if !snapshotRouteConfigEqual(activeRoute, currentRoute) {
|
||||
result.ModifiedSites = append(result.ModifiedSites, siteName)
|
||||
}
|
||||
}
|
||||
for siteName := range activeSiteMap {
|
||||
if _, ok := currentSiteMap[siteName]; !ok {
|
||||
result.RemovedSites = append(result.RemovedSites, siteName)
|
||||
}
|
||||
}
|
||||
currentMap := flattenSnapshotRoutesByDomain(bundle.SnapshotRoutes)
|
||||
activeMap := flattenSnapshotRoutesByDomain(activeSnapshot.Routes)
|
||||
for domain, currentRoute := range currentMap {
|
||||
activeRoute, ok := activeMap[domain]
|
||||
if !ok {
|
||||
result.AddedDomains = append(result.AddedDomains, domain)
|
||||
continue
|
||||
}
|
||||
if !snapshotRouteConfigEqual(activeRoute, currentRoute) {
|
||||
result.ModifiedDomains = append(result.ModifiedDomains, domain)
|
||||
}
|
||||
}
|
||||
for domain := range activeMap {
|
||||
if _, ok := currentMap[domain]; !ok {
|
||||
result.RemovedDomains = append(result.RemovedDomains, domain)
|
||||
}
|
||||
}
|
||||
result.MainConfigChanged = activeVersion.MainConfig != bundle.MainConfig
|
||||
result.WAFConfigChanged = !snapshotWAFConfigEqual(activeSnapshot.WAF, bundle.WAFSnapshot)
|
||||
result.ChangedOptionDetails = diffOpenRestyOptionDetails(activeSnapshot.OpenRestyConfig, bundle.OpenRestyConfig)
|
||||
result.ChangedOptionKeys = extractOptionDiffKeys(result.ChangedOptionDetails)
|
||||
sort.Strings(result.AddedSites)
|
||||
sort.Strings(result.RemovedSites)
|
||||
sort.Strings(result.ModifiedSites)
|
||||
sort.Strings(result.AddedDomains)
|
||||
sort.Strings(result.RemovedDomains)
|
||||
sort.Strings(result.ModifiedDomains)
|
||||
sort.Strings(result.ChangedOptionKeys)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// PublishConfigVersion publishes the current draft as a new active version.
|
||||
func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*model.ConfigVersion, error) {
|
||||
bundle, err := buildCurrentConfigBundle(ctx, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(bundle.Routes) == 0 {
|
||||
return nil, errors.New(errNoEnabledRoutes)
|
||||
}
|
||||
activeVersion, err := model.GetActiveConfigVersion(ctx)
|
||||
if !force && err == nil && activeVersion.Checksum == bundle.Checksum {
|
||||
return nil, errors.New(errNoChangesToPublish)
|
||||
}
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
supportFilesJSON, err := json.Marshal(bundle.SupportFiles)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
version, err := nextVersionNumber(ctx, time.Now())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
record := &model.ConfigVersion{
|
||||
Version: version,
|
||||
SnapshotJSON: bundle.SnapshotJSON,
|
||||
MainConfig: bundle.MainConfig,
|
||||
RenderedConfig: bundle.RouteConfig,
|
||||
SupportFilesJSON: string(supportFilesJSON),
|
||||
Checksum: bundle.Checksum,
|
||||
IsActive: true,
|
||||
CreatedBy: createdBy,
|
||||
}
|
||||
if err = model.PublishConfigVersionTx(ctx, record); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errVersionConflict)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return record, nil
|
||||
}
|
||||
|
||||
// ActivateConfigVersion activates an existing config version.
|
||||
func ActivateConfigVersion(ctx context.Context, id uint) (*model.ConfigVersion, error) {
|
||||
version, err := model.GetConfigVersionByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.ActivateConfigVersionTx(ctx, id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
version.IsActive = true
|
||||
return version, nil
|
||||
}
|
||||
|
||||
// CleanupConfigVersions removes old inactive config versions.
|
||||
func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, error) {
|
||||
if keepCount < 3 {
|
||||
keepCount = 3
|
||||
}
|
||||
versions, err := model.ListConfigVersionSummaries(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(versions) <= keepCount {
|
||||
return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil
|
||||
}
|
||||
var deleteIDs []uint
|
||||
for index, version := range versions {
|
||||
if index < keepCount {
|
||||
continue
|
||||
}
|
||||
if version.IsActive {
|
||||
continue
|
||||
}
|
||||
deleteIDs = append(deleteIDs, version.ID)
|
||||
}
|
||||
if len(deleteIDs) == 0 {
|
||||
return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil
|
||||
}
|
||||
deletedCount, err := model.DeleteConfigVersionsByIDs(ctx, deleteIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &CleanupResult{DeletedCount: deletedCount, Message: "清理成功"}, nil
|
||||
}
|
||||
|
||||
func nextVersionNumber(ctx context.Context, now time.Time) (string, error) {
|
||||
prefix := now.Format("20060102")
|
||||
latest, err := model.GetLatestConfigVersionByPrefix(ctx, prefix)
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return fmt.Sprintf("%s-%03d", prefix, 1), nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
suffix := strings.TrimPrefix(latest, prefix+"-")
|
||||
sequence, err := strconv.Atoi(suffix)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("invalid config version sequence %q: %w", latest, err)
|
||||
}
|
||||
return fmt.Sprintf("%s-%03d", prefix, sequence+1), nil
|
||||
}
|
||||
|
||||
func parseSnapshotDocument(snapshotJSON string) (*snapshotDocument, error) {
|
||||
text := strings.TrimSpace(snapshotJSON)
|
||||
if text == "" {
|
||||
return &snapshotDocument{Routes: []snapshotRoute{}}, nil
|
||||
}
|
||||
if strings.HasPrefix(text, "[") {
|
||||
var routes []snapshotRoute
|
||||
if err := json.Unmarshal([]byte(text), &routes); err != nil {
|
||||
return nil, errors.New(errInvalidSnapshotFormat)
|
||||
}
|
||||
return &snapshotDocument{Routes: normalizeSnapshotRoutes(routes)}, nil
|
||||
}
|
||||
var snapshot snapshotDocument
|
||||
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
|
||||
return nil, errors.New(errInvalidSnapshotFormat)
|
||||
}
|
||||
snapshot.Routes = normalizeSnapshotRoutes(snapshot.Routes)
|
||||
return &snapshot, nil
|
||||
}
|
||||
|
||||
func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
|
||||
if len(routes) == 0 {
|
||||
return []snapshotRoute{}
|
||||
}
|
||||
for index := range routes {
|
||||
normalizedDomains, err := decodeStoredDomains("", routes[index].Domain)
|
||||
if len(routes[index].Domains) > 0 {
|
||||
normalizedDomains, err = normalizeProxyRouteDomains(routes[index].Domains)
|
||||
}
|
||||
if err == nil && len(normalizedDomains) > 0 {
|
||||
routes[index].Domains = normalizedDomains
|
||||
routes[index].Domain = normalizedDomains[0]
|
||||
routes[index].SiteName = normalizeProxyRouteSiteName(nil, routes[index].SiteName, normalizedDomains[0])
|
||||
}
|
||||
normalizedCertIDs, primaryCertID, certErr := normalizeSnapshotCertificateIDs(routes[index].CertID, routes[index].CertIDs)
|
||||
if certErr == nil {
|
||||
routes[index].CertID = primaryCertID
|
||||
routes[index].CertIDs = normalizedCertIDs
|
||||
}
|
||||
normalizedDomainCertIDs, domainCertErr := resolveDomainCertIDs(routes[index].Domains, routes[index].CertIDs, "")
|
||||
if domainCertErr == nil && len(routes[index].DomainCertIDs) == 0 {
|
||||
routes[index].DomainCertIDs = normalizedDomainCertIDs
|
||||
}
|
||||
normalizedUpstreams, upstreamErr := normalizeUpstreams(routes[index].OriginURL, routes[index].Upstreams)
|
||||
if upstreamErr == nil {
|
||||
routes[index].OriginURL = normalizedUpstreams[0]
|
||||
routes[index].Upstreams = normalizedUpstreams
|
||||
}
|
||||
if !routes[index].BasicAuthEnabled {
|
||||
routes[index].BasicAuthUsername = ""
|
||||
routes[index].BasicAuthPassword = ""
|
||||
}
|
||||
routes[index].UpstreamType = normalizeUpstreamType(routes[index].UpstreamType)
|
||||
}
|
||||
return routes
|
||||
}
|
||||
|
||||
func flattenSnapshotRoutesBySite(routes []snapshotRoute) map[string]snapshotRoute {
|
||||
siteMap := make(map[string]snapshotRoute)
|
||||
for _, route := range normalizeSnapshotRoutes(routes) {
|
||||
siteMap[route.SiteName] = route
|
||||
}
|
||||
return siteMap
|
||||
}
|
||||
|
||||
func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRoute {
|
||||
domainMap := make(map[string]snapshotRoute)
|
||||
for _, route := range normalizeSnapshotRoutes(routes) {
|
||||
for _, domain := range route.Domains {
|
||||
item := route
|
||||
item.Domain = domain
|
||||
domainMap[domain] = item
|
||||
}
|
||||
}
|
||||
return domainMap
|
||||
}
|
||||
|
||||
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
|
||||
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL ||
|
||||
left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP ||
|
||||
left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP ||
|
||||
left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy ||
|
||||
left.BasicAuthEnabled != right.BasicAuthEnabled || left.BasicAuthUsername != right.BasicAuthUsername ||
|
||||
left.BasicAuthPassword != right.BasicAuthPassword || left.UpstreamType != right.UpstreamType ||
|
||||
!uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) || left.TunnelTargetAddr != right.TunnelTargetAddr ||
|
||||
left.TunnelTargetProto != right.TunnelTargetProto || !uintPtrEqual(left.PagesProjectID, right.PagesProjectID) ||
|
||||
!uintSliceEqual(left.CertIDs, right.CertIDs) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) {
|
||||
return false
|
||||
}
|
||||
if len(left.Domains) != len(right.Domains) {
|
||||
return false
|
||||
}
|
||||
for index := range left.Domains {
|
||||
if left.Domains[index] != right.Domains[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.Upstreams) != len(right.Upstreams) {
|
||||
return false
|
||||
}
|
||||
for index := range left.Upstreams {
|
||||
if left.Upstreams[index] != right.Upstreams[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.CacheRules) != len(right.CacheRules) {
|
||||
return false
|
||||
}
|
||||
for index := range left.CacheRules {
|
||||
if left.CacheRules[index] != right.CacheRules[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if len(left.CustomHeaders) != len(right.CustomHeaders) {
|
||||
return false
|
||||
}
|
||||
for index := range left.CustomHeaders {
|
||||
if left.CustomHeaders[index] != right.CustomHeaders[index] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func snapshotWAFConfigEqual(left snapshotWAFDocument, right snapshotWAFDocument) bool {
|
||||
leftJSON, err := json.Marshal(left)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
rightJSON, err := json.Marshal(right)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return string(leftJSON) == string(rightJSON)
|
||||
}
|
||||
|
||||
func normalizeSnapshotCertificateIDs(primaryCertID *uint, certIDs []uint) ([]uint, *uint, error) {
|
||||
candidates := make([]uint, 0, len(certIDs)+1)
|
||||
if primaryCertID != nil && *primaryCertID != 0 {
|
||||
candidates = append(candidates, *primaryCertID)
|
||||
}
|
||||
candidates = append(candidates, certIDs...)
|
||||
normalized := make([]uint, 0, len(candidates))
|
||||
seen := make(map[uint]struct{}, len(candidates))
|
||||
for _, certID := range candidates {
|
||||
if certID == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[certID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[certID] = struct{}{}
|
||||
normalized = append(normalized, certID)
|
||||
}
|
||||
var normalizedPrimary *uint
|
||||
if len(normalized) > 0 {
|
||||
normalizedPrimary = &normalized[0]
|
||||
}
|
||||
return normalized, normalizedPrimary, nil
|
||||
}
|
||||
|
||||
func buildInitialOpenRestyOptionDiffs(current openRestyConfigSnapshot) []ConfigOptionDiffItem {
|
||||
details := diffOpenRestyOptionDetails(openRestyConfigSnapshot{}, current)
|
||||
for index := range details {
|
||||
details[index].PreviousValue = ""
|
||||
}
|
||||
return details
|
||||
}
|
||||
|
||||
func diffOpenRestyOptionDetails(left openRestyConfigSnapshot, right openRestyConfigSnapshot) []ConfigOptionDiffItem {
|
||||
changes := make([]ConfigOptionDiffItem, 0)
|
||||
appendIfChanged := func(key string, previous string, current string) {
|
||||
if previous == current {
|
||||
return
|
||||
}
|
||||
changes = append(changes, ConfigOptionDiffItem{
|
||||
Key: key,
|
||||
PreviousValue: previous,
|
||||
CurrentValue: current,
|
||||
})
|
||||
}
|
||||
appendIfChanged("OpenRestyDefaultServerReturnStatus", fmt.Sprintf("%d", left.DefaultServerReturnStatus), fmt.Sprintf("%d", right.DefaultServerReturnStatus))
|
||||
appendIfChanged("OpenRestyWorkerProcesses", left.WorkerProcesses, right.WorkerProcesses)
|
||||
appendIfChanged("OpenRestyWorkerConnections", fmt.Sprintf("%d", left.WorkerConnections), fmt.Sprintf("%d", right.WorkerConnections))
|
||||
appendIfChanged("OpenRestyWorkerRlimitNofile", fmt.Sprintf("%d", left.WorkerRlimitNofile), fmt.Sprintf("%d", right.WorkerRlimitNofile))
|
||||
appendIfChanged("OpenRestyEventsUse", left.EventsUse, right.EventsUse)
|
||||
appendIfChanged("OpenRestyEventsMultiAcceptEnabled", fmt.Sprintf("%t", left.EventsMultiAcceptEnabled), fmt.Sprintf("%t", right.EventsMultiAcceptEnabled))
|
||||
appendIfChanged("OpenRestyKeepaliveTimeout", fmt.Sprintf("%d", left.KeepaliveTimeout), fmt.Sprintf("%d", right.KeepaliveTimeout))
|
||||
appendIfChanged("OpenRestyKeepaliveRequests", fmt.Sprintf("%d", left.KeepaliveRequests), fmt.Sprintf("%d", right.KeepaliveRequests))
|
||||
appendIfChanged("OpenRestyClientHeaderTimeout", fmt.Sprintf("%d", left.ClientHeaderTimeout), fmt.Sprintf("%d", right.ClientHeaderTimeout))
|
||||
appendIfChanged("OpenRestyClientBodyTimeout", fmt.Sprintf("%d", left.ClientBodyTimeout), fmt.Sprintf("%d", right.ClientBodyTimeout))
|
||||
appendIfChanged("OpenRestyClientMaxBodySize", left.ClientMaxBodySize, right.ClientMaxBodySize)
|
||||
appendIfChanged("OpenRestyLargeClientHeaderBuffers", left.LargeClientHeaderBuffers, right.LargeClientHeaderBuffers)
|
||||
appendIfChanged("OpenRestySendTimeout", fmt.Sprintf("%d", left.SendTimeout), fmt.Sprintf("%d", right.SendTimeout))
|
||||
appendIfChanged("OpenRestyProxyConnectTimeout", fmt.Sprintf("%d", left.ProxyConnectTimeout), fmt.Sprintf("%d", right.ProxyConnectTimeout))
|
||||
appendIfChanged("OpenRestyProxySendTimeout", fmt.Sprintf("%d", left.ProxySendTimeout), fmt.Sprintf("%d", right.ProxySendTimeout))
|
||||
appendIfChanged("OpenRestyProxyReadTimeout", fmt.Sprintf("%d", left.ProxyReadTimeout), fmt.Sprintf("%d", right.ProxyReadTimeout))
|
||||
appendIfChanged("OpenRestyWebsocketEnabled", fmt.Sprintf("%t", left.WebsocketEnabled), fmt.Sprintf("%t", right.WebsocketEnabled))
|
||||
appendIfChanged("OpenRestyHTTP3Enabled", fmt.Sprintf("%t", left.HTTP3Enabled), fmt.Sprintf("%t", right.HTTP3Enabled))
|
||||
appendIfChanged("OpenRestyProxyRequestBufferingEnabled", fmt.Sprintf("%t", left.ProxyRequestBuffering), fmt.Sprintf("%t", right.ProxyRequestBuffering))
|
||||
appendIfChanged("OpenRestyProxyBufferingEnabled", fmt.Sprintf("%t", left.ProxyBufferingEnabled), fmt.Sprintf("%t", right.ProxyBufferingEnabled))
|
||||
appendIfChanged("OpenRestyProxyBuffers", left.ProxyBuffers, right.ProxyBuffers)
|
||||
appendIfChanged("OpenRestyProxyBufferSize", left.ProxyBufferSize, right.ProxyBufferSize)
|
||||
appendIfChanged("OpenRestyProxyBusyBuffersSize", left.ProxyBusyBuffersSize, right.ProxyBusyBuffersSize)
|
||||
appendIfChanged("OpenRestyGzipEnabled", fmt.Sprintf("%t", left.GzipEnabled), fmt.Sprintf("%t", right.GzipEnabled))
|
||||
appendIfChanged("OpenRestyGzipMinLength", fmt.Sprintf("%d", left.GzipMinLength), fmt.Sprintf("%d", right.GzipMinLength))
|
||||
appendIfChanged("OpenRestyGzipCompLevel", fmt.Sprintf("%d", left.GzipCompLevel), fmt.Sprintf("%d", right.GzipCompLevel))
|
||||
appendIfChanged("OpenRestyResolvers", left.Resolvers, right.Resolvers)
|
||||
appendIfChanged("OpenRestyCacheEnabled", fmt.Sprintf("%t", left.CacheEnabled), fmt.Sprintf("%t", right.CacheEnabled))
|
||||
appendIfChanged("OpenRestyCachePath", left.CachePath, right.CachePath)
|
||||
appendIfChanged("OpenRestyCacheLevels", left.CacheLevels, right.CacheLevels)
|
||||
appendIfChanged("OpenRestyCacheInactive", left.CacheInactive, right.CacheInactive)
|
||||
appendIfChanged("OpenRestyCacheMaxSize", left.CacheMaxSize, right.CacheMaxSize)
|
||||
appendIfChanged("OpenRestyCacheKeyTemplate", left.CacheKeyTemplate, right.CacheKeyTemplate)
|
||||
appendIfChanged("OpenRestyCacheLockEnabled", fmt.Sprintf("%t", left.CacheLockEnabled), fmt.Sprintf("%t", right.CacheLockEnabled))
|
||||
appendIfChanged("OpenRestyCacheLockTimeout", left.CacheLockTimeout, right.CacheLockTimeout)
|
||||
appendIfChanged("OpenRestyCacheUseStale", left.CacheUseStale, right.CacheUseStale)
|
||||
return changes
|
||||
}
|
||||
|
||||
func extractOptionDiffKeys(details []ConfigOptionDiffItem) []string {
|
||||
keys := make([]string, 0, len(details))
|
||||
for _, item := range details {
|
||||
keys = append(keys, item.Key)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func openRestyOptionKeys() []string {
|
||||
return []string{
|
||||
"OpenRestyDefaultServerReturnStatus",
|
||||
"OpenRestyWorkerProcesses",
|
||||
"OpenRestyWorkerConnections",
|
||||
"OpenRestyWorkerRlimitNofile",
|
||||
"OpenRestyEventsUse",
|
||||
"OpenRestyEventsMultiAcceptEnabled",
|
||||
"OpenRestyKeepaliveTimeout",
|
||||
"OpenRestyKeepaliveRequests",
|
||||
"OpenRestyClientHeaderTimeout",
|
||||
"OpenRestyClientBodyTimeout",
|
||||
"OpenRestyClientMaxBodySize",
|
||||
"OpenRestyLargeClientHeaderBuffers",
|
||||
"OpenRestySendTimeout",
|
||||
"OpenRestyProxyConnectTimeout",
|
||||
"OpenRestyProxySendTimeout",
|
||||
"OpenRestyProxyReadTimeout",
|
||||
"OpenRestyWebsocketEnabled",
|
||||
"OpenRestyHTTP3Enabled",
|
||||
"OpenRestyProxyRequestBufferingEnabled",
|
||||
"OpenRestyProxyBufferingEnabled",
|
||||
"OpenRestyProxyBuffers",
|
||||
"OpenRestyProxyBufferSize",
|
||||
"OpenRestyProxyBusyBuffersSize",
|
||||
"OpenRestyGzipEnabled",
|
||||
"OpenRestyGzipMinLength",
|
||||
"OpenRestyGzipCompLevel",
|
||||
"OpenRestyCacheEnabled",
|
||||
"OpenRestyCachePath",
|
||||
"OpenRestyCacheLevels",
|
||||
"OpenRestyCacheInactive",
|
||||
"OpenRestyCacheMaxSize",
|
||||
"OpenRestyCacheKeyTemplate",
|
||||
"OpenRestyCacheLockEnabled",
|
||||
"OpenRestyCacheLockTimeout",
|
||||
"OpenRestyCacheUseStale",
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupConfigVersionTestDB(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.ProxyRoute{},
|
||||
&model.ConfigVersion{},
|
||||
&model.OpenFlareWAFRuleGroup{},
|
||||
&model.OpenFlareWAFRuleGroupBinding{},
|
||||
&model.OpenFlareWAFIPGroup{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionCreatesVersion(t *testing.T) {
|
||||
cleanup := setupConfigVersionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
route := &model.ProxyRoute{
|
||||
SiteName: "publish-site",
|
||||
Domain: "publish.example.com",
|
||||
Domains: `["publish.example.com"]`,
|
||||
OriginURL: "http://origin.publish.example.com:8080",
|
||||
Upstreams: `["http://origin.publish.example.com:8080"]`,
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, model.CreateProxyRouteRecord(ctx, route))
|
||||
|
||||
version, err := PublishConfigVersion(ctx, "tester", false)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, version)
|
||||
assert.NotZero(t, version.ID)
|
||||
assert.True(t, version.IsActive)
|
||||
assert.Equal(t, "tester", version.CreatedBy)
|
||||
assert.NotEmpty(t, version.Version)
|
||||
assert.NotEmpty(t, version.Checksum)
|
||||
assert.NotEmpty(t, version.SnapshotJSON)
|
||||
assert.NotEmpty(t, version.RenderedConfig)
|
||||
|
||||
var snapshot snapshotDocument
|
||||
require.NoError(t, json.Unmarshal([]byte(version.SnapshotJSON), &snapshot))
|
||||
require.Len(t, snapshot.Routes, 1)
|
||||
assert.Equal(t, "publish-site", snapshot.Routes[0].SiteName)
|
||||
assert.Equal(t, "publish.example.com", snapshot.Routes[0].Domain)
|
||||
|
||||
active, err := GetActiveConfigVersion(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, version.ID, active.ID)
|
||||
|
||||
_, err = PublishConfigVersion(ctx, "tester", false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), errNoChangesToPublish)
|
||||
|
||||
forced, err := PublishConfigVersion(ctx, "tester", true)
|
||||
require.NoError(t, err)
|
||||
assert.NotEqual(t, version.ID, forced.ID)
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
import (
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
)
|
||||
|
||||
// SupportFile is a rendered configuration support artifact.
|
||||
type SupportFile struct {
|
||||
Path string `json:"path"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
func renderSnapshotConfig(sourceJSON string, certificateFiles []SupportFile) (*openrestyrender.Result, error) {
|
||||
return openrestyrender.RenderJSON(sourceJSON, toOpenRestySupportFiles(certificateFiles))
|
||||
}
|
||||
|
||||
func toOpenRestySupportFiles(files []SupportFile) []openrestyrender.SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]openrestyrender.SupportFile, 0, len(files))
|
||||
for _, file := range files {
|
||||
result = append(result, openrestyrender.SupportFile{
|
||||
Path: file.Path,
|
||||
Content: file.Content,
|
||||
})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func fromOpenRestySupportFiles(files []openrestyrender.SupportFile) []SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]SupportFile, 0, len(files))
|
||||
for _, file := range files {
|
||||
result = append(result, SupportFile{
|
||||
Path: file.Path,
|
||||
Content: file.Content,
|
||||
})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func renderPlaceholderConfig(snapshotJSON string) (mainConfig, routeConfig, checksum string) {
|
||||
mainConfig = `{"placeholder":"main_config"}`
|
||||
routeConfig = snapshotJSON
|
||||
checksum = openrestyrender.ChecksumBundle(mainConfig, routeConfig, nil)
|
||||
return mainConfig, routeConfig, checksum
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return apiutil.AbortNotFoundIfMissing(c, err, "记录不存在")
|
||||
}
|
||||
|
||||
// ListConfigVersionsHandler lists config versions.
|
||||
// @Summary 获取配置版本列表
|
||||
// @Description 返回所有已发布的 OpenResty 配置版本摘要,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.ConfigVersionSummary} "配置版本列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/config-versions [get]
|
||||
func ListConfigVersionsHandler(c *gin.Context) {
|
||||
versions, err := ListConfigVersions(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(versions))
|
||||
}
|
||||
|
||||
// GetConfigVersionHandler returns a config version by id.
|
||||
// @Summary 获取配置版本详情
|
||||
// @Description 返回指定配置版本的完整快照与渲染内容,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "配置版本 ID"
|
||||
// @Success 200 {object} response.Any{data=model.ConfigVersion} "配置版本详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或版本不存在"
|
||||
// @Router /api/v1/d/config-versions/{id} [get]
|
||||
func GetConfigVersionHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
version, err := GetConfigVersionDetail(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(version))
|
||||
}
|
||||
|
||||
// GetActiveConfigVersionHandler returns the active config version.
|
||||
// @Summary 获取当前活跃配置版本
|
||||
// @Description 返回当前正在使用的配置版本,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=model.ConfigVersion} "活跃配置版本"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限、不存在或无活跃版本"
|
||||
// @Router /api/v1/d/config-versions/active [get]
|
||||
func GetActiveConfigVersionHandler(c *gin.Context) {
|
||||
version, err := GetActiveConfigVersion(c.Request.Context())
|
||||
if apiutil.AbortNotFoundIfMissing(c, err, errNoActiveVersion) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(version))
|
||||
}
|
||||
|
||||
// PreviewConfigVersionHandler previews the current draft configuration.
|
||||
// @Summary 预览当前草稿配置
|
||||
// @Description 渲染并返回当前草稿配置的预览结果,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=config_version.ConfigPreviewResult} "配置预览"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/config-versions/preview [get]
|
||||
func PreviewConfigVersionHandler(c *gin.Context) {
|
||||
preview, err := PreviewConfigVersion(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(preview))
|
||||
}
|
||||
|
||||
// DiffConfigVersionHandler diffs the current draft against the active version.
|
||||
// @Summary 对比草稿与活跃配置
|
||||
// @Description 对比当前草稿配置与活跃版本之间的差异,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=config_version.ConfigDiffResult} "配置差异"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/config-versions/diff [get]
|
||||
func DiffConfigVersionHandler(c *gin.Context) {
|
||||
diff, err := DiffConfigVersion(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(diff))
|
||||
}
|
||||
|
||||
// PublishConfigVersionHandler publishes a new config version.
|
||||
// @Summary 发布配置版本
|
||||
// @Description 将当前草稿配置发布为新版本,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param force query bool false "是否强制发布"
|
||||
// @Success 200 {object} response.Any{data=model.ConfigVersion} "发布成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/config-versions/publish [post]
|
||||
func PublishConfigVersionHandler(c *gin.Context) {
|
||||
username := c.GetString("username")
|
||||
force := c.Query("force") == "true"
|
||||
version, err := PublishConfigVersion(c.Request.Context(), username, force)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(version))
|
||||
}
|
||||
|
||||
// ActivateConfigVersionHandler activates an existing config version.
|
||||
// @Summary 激活配置版本
|
||||
// @Description 将指定历史版本设为当前活跃配置,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "配置版本 ID"
|
||||
// @Success 200 {object} response.Any{data=model.ConfigVersion} "激活成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或版本不存在"
|
||||
// @Router /api/v1/d/config-versions/{id}/activate [post]
|
||||
func ActivateConfigVersionHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
version, err := ActivateConfigVersion(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(version))
|
||||
}
|
||||
|
||||
// CleanupConfigVersionsHandler removes old inactive config versions.
|
||||
// @Summary 清理历史配置版本
|
||||
// @Description 删除超出保留数量的非活跃配置版本,需要管理员权限
|
||||
// @Tags openflare-config-version
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param body body config_version.CleanupInput true "清理参数"
|
||||
// @Success 200 {object} response.Any{data=config_version.CleanupResult} "清理结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/config-versions/cleanup [post]
|
||||
func CleanupConfigVersionsHandler(c *gin.Context) {
|
||||
var input CleanupInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
result, err := CleanupConfigVersions(c.Request.Context(), input.KeepCount)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
@@ -0,0 +1,510 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
)
|
||||
|
||||
type snapshotRoute struct {
|
||||
ID uint `json:"id,omitempty"`
|
||||
SiteName string `json:"site_name,omitempty"`
|
||||
Domain string `json:"domain"`
|
||||
Domains []string `json:"domains,omitempty"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginHost string `json:"origin_host,omitempty"`
|
||||
Upstreams []string `json:"upstreams,omitempty"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id,omitempty"`
|
||||
CertIDs []uint `json:"cert_ids,omitempty"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids,omitempty"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server,omitempty"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"`
|
||||
LimitRate string `json:"limit_rate,omitempty"`
|
||||
CacheEnabled bool `json:"cache_enabled"`
|
||||
CachePolicy string `json:"cache_policy,omitempty"`
|
||||
CacheRules []string `json:"cache_rules,omitempty"`
|
||||
CustomHeaders []customHeaderInput `json:"custom_headers,omitempty"`
|
||||
BasicAuthEnabled bool `json:"basic_auth_enabled,omitempty"`
|
||||
BasicAuthUsername string `json:"basic_auth_username,omitempty"`
|
||||
BasicAuthPassword string `json:"basic_auth_password,omitempty"`
|
||||
Remark string `json:"remark,omitempty"`
|
||||
UpstreamType string `json:"upstream_type,omitempty"`
|
||||
TunnelNodeID *uint `json:"tunnel_node_id,omitempty"`
|
||||
TunnelTargetAddr string `json:"tunnel_target_addr,omitempty"`
|
||||
TunnelTargetProto string `json:"tunnel_target_protocol,omitempty"`
|
||||
PagesProjectID *uint `json:"pages_project_id,omitempty"`
|
||||
}
|
||||
|
||||
type snapshotWAFRuleGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IsGlobal bool `json:"is_global"`
|
||||
BlockStatusCode int `json:"block_status_code"`
|
||||
BlockResponseBody string `json:"block_response_body,omitempty"`
|
||||
IPWhitelist []string `json:"ip_whitelist,omitempty"`
|
||||
IPBlacklist []string `json:"ip_blacklist,omitempty"`
|
||||
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
|
||||
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
|
||||
CountryWhitelist []string `json:"country_whitelist,omitempty"`
|
||||
CountryBlacklist []string `json:"country_blacklist,omitempty"`
|
||||
RegionWhitelist []string `json:"region_whitelist,omitempty"`
|
||||
RegionBlacklist []string `json:"region_blacklist,omitempty"`
|
||||
PoWEnabled bool `json:"pow_enabled,omitempty"`
|
||||
PoWConfig *openrestyrender.PoWConfig `json:"pow_config,omitempty"`
|
||||
}
|
||||
|
||||
type snapshotWAFIPGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IPList []string `json:"ip_list,omitempty"`
|
||||
}
|
||||
|
||||
type snapshotWAFBinding struct {
|
||||
RouteID uint `json:"route_id"`
|
||||
SiteName string `json:"site_name"`
|
||||
RuleGroupIDs []uint `json:"rule_group_ids"`
|
||||
}
|
||||
|
||||
type snapshotWAFDocument struct {
|
||||
RuleGroups []snapshotWAFRuleGroup `json:"rule_groups"`
|
||||
IPGroups []snapshotWAFIPGroup `json:"ip_groups,omitempty"`
|
||||
Bindings []snapshotWAFBinding `json:"bindings"`
|
||||
}
|
||||
|
||||
type openRestyConfigSnapshot struct {
|
||||
DefaultServerReturnStatus int `json:"default_server_return_status"`
|
||||
WorkerProcesses string `json:"worker_processes"`
|
||||
WorkerConnections int `json:"worker_connections"`
|
||||
WorkerRlimitNofile int `json:"worker_rlimit_nofile"`
|
||||
EventsUse string `json:"events_use,omitempty"`
|
||||
EventsMultiAcceptEnabled bool `json:"events_multi_accept_enabled"`
|
||||
KeepaliveTimeout int `json:"keepalive_timeout"`
|
||||
KeepaliveRequests int `json:"keepalive_requests"`
|
||||
ClientHeaderTimeout int `json:"client_header_timeout"`
|
||||
ClientBodyTimeout int `json:"client_body_timeout"`
|
||||
ClientMaxBodySize string `json:"client_max_body_size"`
|
||||
LargeClientHeaderBuffers string `json:"large_client_header_buffers"`
|
||||
SendTimeout int `json:"send_timeout"`
|
||||
ProxyConnectTimeout int `json:"proxy_connect_timeout"`
|
||||
ProxySendTimeout int `json:"proxy_send_timeout"`
|
||||
ProxyReadTimeout int `json:"proxy_read_timeout"`
|
||||
WebsocketEnabled bool `json:"websocket_enabled"`
|
||||
HTTP3Enabled bool `json:"http3_enabled"`
|
||||
ProxyRequestBuffering bool `json:"proxy_request_buffering"`
|
||||
ProxyBufferingEnabled bool `json:"proxy_buffering_enabled"`
|
||||
ProxyBuffers string `json:"proxy_buffers"`
|
||||
ProxyBufferSize string `json:"proxy_buffer_size"`
|
||||
ProxyBusyBuffersSize string `json:"proxy_busy_buffers_size"`
|
||||
GzipEnabled bool `json:"gzip_enabled"`
|
||||
GzipMinLength int `json:"gzip_min_length"`
|
||||
GzipCompLevel int `json:"gzip_comp_level"`
|
||||
Resolvers string `json:"resolvers,omitempty"`
|
||||
CacheEnabled bool `json:"cache_enabled"`
|
||||
CachePath string `json:"cache_path,omitempty"`
|
||||
CacheLevels string `json:"cache_levels"`
|
||||
CacheInactive string `json:"cache_inactive"`
|
||||
CacheMaxSize string `json:"cache_max_size"`
|
||||
CacheKeyTemplate string `json:"cache_key_template"`
|
||||
CacheLockEnabled bool `json:"cache_lock_enabled"`
|
||||
CacheLockTimeout string `json:"cache_lock_timeout"`
|
||||
CacheUseStale string `json:"cache_use_stale"`
|
||||
MainConfigTemplate string `json:"main_config_template,omitempty"`
|
||||
}
|
||||
|
||||
type snapshotDocument struct {
|
||||
Routes []snapshotRoute `json:"routes"`
|
||||
OpenRestyConfig openRestyConfigSnapshot `json:"openresty_config"`
|
||||
WAF snapshotWAFDocument `json:"waf"`
|
||||
}
|
||||
|
||||
type configBundle struct {
|
||||
Routes []*model.ProxyRoute
|
||||
SnapshotRoutes []snapshotRoute
|
||||
WAFSnapshot snapshotWAFDocument
|
||||
OpenRestyConfig openRestyConfigSnapshot
|
||||
SnapshotJSON string
|
||||
MainConfig string
|
||||
RouteConfig string
|
||||
SupportFiles []SupportFile
|
||||
Checksum string
|
||||
ChangedOptionKeys []string
|
||||
}
|
||||
|
||||
func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configBundle, error) {
|
||||
routes, err := model.ListEnabledProxyRoutes(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if requireRoutes && len(routes) == 0 {
|
||||
return nil, errors.New(errNoEnabledRoutes)
|
||||
}
|
||||
snapshotRoutes, err := buildSnapshotRoutes(ctx, routes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
wafSnapshot, err := buildSnapshotWAFDocument(ctx, routes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
openRestyConfig := buildOpenRestyConfigSnapshot()
|
||||
snapshotDoc := snapshotDocument{
|
||||
Routes: snapshotRoutes,
|
||||
OpenRestyConfig: openRestyConfig,
|
||||
WAF: wafSnapshot,
|
||||
}
|
||||
snapshotJSON, err := json.Marshal(snapshotDoc)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certificateFiles, err := buildCertificateSupportFiles(ctx, snapshotRoutes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
mainConfig := ""
|
||||
routeConfig := ""
|
||||
checksum := ""
|
||||
supportFiles := []SupportFile(nil)
|
||||
|
||||
rendered, renderErr := renderSnapshotConfig(string(snapshotJSON), certificateFiles)
|
||||
if renderErr == nil {
|
||||
mainConfig = rendered.MainConfig
|
||||
routeConfig = rendered.RouteConfig
|
||||
checksum = rendered.Checksum
|
||||
supportFiles = fromOpenRestySupportFiles(rendered.SupportFiles)
|
||||
} else {
|
||||
mainConfig, routeConfig, checksum = renderPlaceholderConfig(string(snapshotJSON))
|
||||
}
|
||||
|
||||
return &configBundle{
|
||||
Routes: routes,
|
||||
SnapshotRoutes: snapshotRoutes,
|
||||
WAFSnapshot: wafSnapshot,
|
||||
OpenRestyConfig: openRestyConfig,
|
||||
SnapshotJSON: string(snapshotJSON),
|
||||
MainConfig: mainConfig,
|
||||
RouteConfig: routeConfig,
|
||||
SupportFiles: supportFiles,
|
||||
Checksum: checksum,
|
||||
ChangedOptionKeys: openRestyOptionKeys(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildSnapshotRoutes(ctx context.Context, routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
items := make([]snapshotRoute, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
domains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("route %s domains are invalid", route.Domain)
|
||||
}
|
||||
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
|
||||
}
|
||||
upstreamType := normalizeUpstreamType(route.UpstreamType)
|
||||
originURL := route.OriginURL
|
||||
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 上游配置无效", route.Domain)
|
||||
}
|
||||
var tunnelNodeID *uint
|
||||
var tunnelTargetAddr string
|
||||
var tunnelTargetProtocol string
|
||||
var pagesProjectID *uint
|
||||
if upstreamType == "tunnel" {
|
||||
originURL = resolveTunnelOpenRestyUpstreamURL(ctx)
|
||||
upstreams = []string{originURL}
|
||||
tunnelNodeID = route.TunnelNodeID
|
||||
tunnelTargetAddr = strings.TrimSpace(route.TunnelTargetAddr)
|
||||
tunnelTargetProtocol = normalizeTunnelTargetProtocol(route.TunnelTargetProtocol)
|
||||
} else if upstreamType == "pages" {
|
||||
return nil, fmt.Errorf("路由 %s Pages 配置无效: pages module is not available", route.Domain)
|
||||
}
|
||||
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("路由 %s 缓存规则无效", route.Domain)
|
||||
}
|
||||
items = append(items, snapshotRoute{
|
||||
ID: route.ID,
|
||||
SiteName: normalizeProxyRouteSiteName(route, route.SiteName, domains[0]),
|
||||
Domain: domains[0],
|
||||
Domains: domains,
|
||||
OriginURL: originURL,
|
||||
OriginHost: route.OriginHost,
|
||||
Upstreams: upstreams,
|
||||
Enabled: route.Enabled,
|
||||
EnableHTTPS: route.EnableHTTPS,
|
||||
CertID: route.CertID,
|
||||
CertIDs: mustDecodeCertIDs(route),
|
||||
DomainCertIDs: mustDecodeDomainCertIDs(route, domains),
|
||||
RedirectHTTP: route.RedirectHTTP,
|
||||
LimitConnPerServer: route.LimitConnPerServer,
|
||||
LimitConnPerIP: route.LimitConnPerIP,
|
||||
LimitRate: route.LimitRate,
|
||||
CacheEnabled: route.CacheEnabled,
|
||||
CachePolicy: route.CachePolicy,
|
||||
CacheRules: cacheRules,
|
||||
CustomHeaders: customHeaders,
|
||||
BasicAuthEnabled: route.BasicAuthEnabled,
|
||||
BasicAuthUsername: route.BasicAuthUsername,
|
||||
BasicAuthPassword: route.BasicAuthPassword,
|
||||
Remark: route.Remark,
|
||||
UpstreamType: upstreamType,
|
||||
TunnelNodeID: tunnelNodeID,
|
||||
TunnelTargetAddr: tunnelTargetAddr,
|
||||
TunnelTargetProto: tunnelTargetProtocol,
|
||||
PagesProjectID: pagesProjectID,
|
||||
})
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (snapshotWAFDocument, error) {
|
||||
if err := waf.EnsureDefaultRuleGroup(ctx); err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
views, err := waf.ListRuleGroups(ctx)
|
||||
if err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views))
|
||||
for _, view := range views {
|
||||
if !view.Enabled {
|
||||
continue
|
||||
}
|
||||
ruleGroups = append(ruleGroups, snapshotWAFRuleGroup{
|
||||
ID: view.ID,
|
||||
Name: view.Name,
|
||||
Enabled: view.Enabled,
|
||||
IsGlobal: view.IsGlobal,
|
||||
BlockStatusCode: view.BlockStatusCode,
|
||||
BlockResponseBody: view.BlockResponseBody,
|
||||
IPWhitelist: view.IPWhitelist,
|
||||
IPBlacklist: view.IPBlacklist,
|
||||
IPWhitelistGroups: view.IPWhitelistGroups,
|
||||
IPBlacklistGroups: view.IPBlacklistGroups,
|
||||
CountryWhitelist: view.CountryWhitelist,
|
||||
CountryBlacklist: view.CountryBlacklist,
|
||||
RegionWhitelist: view.RegionWhitelist,
|
||||
RegionBlacklist: view.RegionBlacklist,
|
||||
PoWEnabled: view.PoWEnabled,
|
||||
PoWConfig: convertPoWConfig(view.PoWConfig),
|
||||
})
|
||||
}
|
||||
ipGroups, err := buildSnapshotWAFIPGroups(ctx, ruleGroups)
|
||||
if err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
enabledRouteIDs := make(map[uint]string, len(routes))
|
||||
for _, route := range routes {
|
||||
if route == nil {
|
||||
continue
|
||||
}
|
||||
siteName := strings.TrimSpace(route.SiteName)
|
||||
if siteName == "" {
|
||||
siteName = route.Domain
|
||||
}
|
||||
enabledRouteIDs[route.ID] = siteName
|
||||
}
|
||||
rawBindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx)
|
||||
if err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
groupIDsByRoute := make(map[uint][]uint, len(rawBindings))
|
||||
for _, binding := range rawBindings {
|
||||
if _, ok := enabledRouteIDs[binding.ProxyRouteID]; !ok {
|
||||
continue
|
||||
}
|
||||
groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
|
||||
}
|
||||
bindings := make([]snapshotWAFBinding, 0, len(groupIDsByRoute))
|
||||
for routeID, groupIDs := range groupIDsByRoute {
|
||||
sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] })
|
||||
bindings = append(bindings, snapshotWAFBinding{
|
||||
RouteID: routeID,
|
||||
SiteName: enabledRouteIDs[routeID],
|
||||
RuleGroupIDs: groupIDs,
|
||||
})
|
||||
}
|
||||
sort.Slice(bindings, func(i, j int) bool {
|
||||
if bindings[i].SiteName == bindings[j].SiteName {
|
||||
return bindings[i].RouteID < bindings[j].RouteID
|
||||
}
|
||||
return bindings[i].SiteName < bindings[j].SiteName
|
||||
})
|
||||
return snapshotWAFDocument{RuleGroups: ruleGroups, IPGroups: ipGroups, Bindings: bindings}, nil
|
||||
}
|
||||
|
||||
func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleGroup) ([]snapshotWAFIPGroup, error) {
|
||||
idSet := make(map[uint]struct{})
|
||||
for _, group := range ruleGroups {
|
||||
for _, id := range group.IPWhitelistGroups {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
for _, id := range group.IPBlacklistGroups {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
if len(idSet) == 0 {
|
||||
return []snapshotWAFIPGroup{}, nil
|
||||
}
|
||||
ids := make([]uint, 0, len(idSet))
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
groups, err := listWAFIPGroupsByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups))
|
||||
for _, group := range groups {
|
||||
groupByID[group.ID] = group
|
||||
}
|
||||
snapshots := make([]snapshotWAFIPGroup, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
group := groupByID[id]
|
||||
if group == nil {
|
||||
return nil, fmt.Errorf("IP 组 %d 不存在", id)
|
||||
}
|
||||
ipList, decodeErr := decodeIPList(group.IPList)
|
||||
if decodeErr != nil {
|
||||
return nil, decodeErr
|
||||
}
|
||||
snapshots = append(snapshots, snapshotWAFIPGroup{
|
||||
ID: group.ID,
|
||||
Name: group.Name,
|
||||
Type: group.Type,
|
||||
Enabled: group.Enabled,
|
||||
IPList: ipList,
|
||||
})
|
||||
}
|
||||
return snapshots, nil
|
||||
}
|
||||
|
||||
func decodeIPList(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, fmt.Errorf("ip_list payload is invalid")
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func convertPoWConfig(config *waf.PoWConfig) *openrestyrender.PoWConfig {
|
||||
if config == nil {
|
||||
return nil
|
||||
}
|
||||
return &openrestyrender.PoWConfig{
|
||||
Difficulty: config.Difficulty,
|
||||
Algorithm: config.Algorithm,
|
||||
SessionTTL: config.SessionTTL,
|
||||
ChallengeTTL: config.ChallengeTTL,
|
||||
Whitelist: openrestyrender.PoWListConfig{
|
||||
IPs: config.Whitelist.IPs,
|
||||
IPCidrs: config.Whitelist.IPCidrs,
|
||||
Paths: config.Whitelist.Paths,
|
||||
PathRegexes: config.Whitelist.PathRegexes,
|
||||
UserAgents: config.Whitelist.UserAgents,
|
||||
},
|
||||
Blacklist: openrestyrender.PoWListConfig{
|
||||
IPs: config.Blacklist.IPs,
|
||||
IPCidrs: config.Blacklist.IPCidrs,
|
||||
Paths: config.Blacklist.Paths,
|
||||
PathRegexes: config.Blacklist.PathRegexes,
|
||||
UserAgents: config.Blacklist.UserAgents,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildOpenRestyConfigSnapshot() openRestyConfigSnapshot {
|
||||
return openRestyConfigSnapshot{
|
||||
DefaultServerReturnStatus: model.OpenRestyDefaultServerReturnStatus,
|
||||
WorkerProcesses: model.OpenRestyWorkerProcesses,
|
||||
WorkerConnections: model.OpenRestyWorkerConnections,
|
||||
WorkerRlimitNofile: model.OpenRestyWorkerRlimitNofile,
|
||||
EventsUse: model.OpenRestyEventsUse,
|
||||
EventsMultiAcceptEnabled: model.OpenRestyEventsMultiAcceptEnabled,
|
||||
KeepaliveTimeout: model.OpenRestyKeepaliveTimeout,
|
||||
KeepaliveRequests: model.OpenRestyKeepaliveRequests,
|
||||
ClientHeaderTimeout: model.OpenRestyClientHeaderTimeout,
|
||||
ClientBodyTimeout: model.OpenRestyClientBodyTimeout,
|
||||
ClientMaxBodySize: model.OpenRestyClientMaxBodySize,
|
||||
LargeClientHeaderBuffers: model.OpenRestyLargeClientHeaderBuffers,
|
||||
SendTimeout: model.OpenRestySendTimeout,
|
||||
ProxyConnectTimeout: model.OpenRestyProxyConnectTimeout,
|
||||
ProxySendTimeout: model.OpenRestyProxySendTimeout,
|
||||
ProxyReadTimeout: model.OpenRestyProxyReadTimeout,
|
||||
WebsocketEnabled: model.OpenRestyWebsocketEnabled,
|
||||
HTTP3Enabled: model.OpenRestyHTTP3Enabled,
|
||||
ProxyRequestBuffering: model.OpenRestyProxyRequestBufferingEnabled,
|
||||
ProxyBufferingEnabled: model.OpenRestyProxyBufferingEnabled,
|
||||
ProxyBuffers: model.OpenRestyProxyBuffers,
|
||||
ProxyBufferSize: model.OpenRestyProxyBufferSize,
|
||||
ProxyBusyBuffersSize: model.OpenRestyProxyBusyBuffersSize,
|
||||
GzipEnabled: model.OpenRestyGzipEnabled,
|
||||
GzipMinLength: model.OpenRestyGzipMinLength,
|
||||
GzipCompLevel: model.OpenRestyGzipCompLevel,
|
||||
Resolvers: model.OpenRestyResolvers,
|
||||
CacheEnabled: model.OpenRestyCacheEnabled,
|
||||
CachePath: model.OpenRestyCachePath,
|
||||
CacheLevels: model.OpenRestyCacheLevels,
|
||||
CacheInactive: model.OpenRestyCacheInactive,
|
||||
CacheMaxSize: model.OpenRestyCacheMaxSize,
|
||||
CacheKeyTemplate: model.OpenRestyCacheKeyTemplate,
|
||||
CacheLockEnabled: model.OpenRestyCacheLockEnabled,
|
||||
CacheLockTimeout: model.OpenRestyCacheLockTimeout,
|
||||
CacheUseStale: model.OpenRestyCacheUseStale,
|
||||
MainConfigTemplate: model.OpenRestyMainConfigTemplate,
|
||||
}
|
||||
}
|
||||
|
||||
func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) ([]SupportFile, error) {
|
||||
certIDSet := make(map[uint]struct{})
|
||||
for _, route := range routes {
|
||||
for _, certID := range route.CertIDs {
|
||||
if certID != 0 {
|
||||
certIDSet[certID] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(certIDSet) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
certIDs := make([]uint, 0, len(certIDSet))
|
||||
for certID := range certIDSet {
|
||||
certIDs = append(certIDs, certID)
|
||||
}
|
||||
sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] })
|
||||
files := make([]SupportFile, 0, len(certIDs)*2)
|
||||
for _, certID := range certIDs {
|
||||
certificate, err := model.GetTLSCertificateByID(ctx, certID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
files = append(files,
|
||||
SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)},
|
||||
SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(certificate.KeyPEM)},
|
||||
)
|
||||
}
|
||||
return dedupeSupportFiles(files), nil
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
nodeStatusOffline = "offline"
|
||||
nodeStatusPending = "pending"
|
||||
)
|
||||
|
||||
func computeNodeStatus(node *model.OpenFlareNode) string {
|
||||
if node == nil {
|
||||
return nodeStatusOffline
|
||||
}
|
||||
if node.LastSeenAt == nil || node.LastSeenAt.IsZero() {
|
||||
return nodeStatusPending
|
||||
}
|
||||
if time.Since(*node.LastSeenAt) > model.NodeOfflineThreshold {
|
||||
return nodeStatusOffline
|
||||
}
|
||||
return nodeStatusOnline
|
||||
}
|
||||
|
||||
func nodeViewLastSeenAt(node *model.OpenFlareNode) any {
|
||||
if node == nil {
|
||||
return time.Time{}
|
||||
}
|
||||
nodeType := strings.TrimSpace(node.NodeType)
|
||||
if nodeType == "" {
|
||||
nodeType = "edge_node"
|
||||
}
|
||||
if nodeType == "tunnel_relay" && ofws.IsRelayConnected(node.NodeID) {
|
||||
return ofws.RelayWSConnectedLastSeenValue
|
||||
}
|
||||
if nodeType == "tunnel_client" && 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
|
||||
}
|
||||
@@ -0,0 +1,366 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/observability"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// Summary is the dashboard node summary section.
|
||||
type Summary struct {
|
||||
TotalNodes int `json:"total_nodes"`
|
||||
OnlineNodes int `json:"online_nodes"`
|
||||
OfflineNodes int `json:"offline_nodes"`
|
||||
PendingNodes int `json:"pending_nodes"`
|
||||
UnhealthyNodes int `json:"unhealthy_nodes"`
|
||||
}
|
||||
|
||||
// Traffic is the dashboard traffic section.
|
||||
type Traffic struct {
|
||||
RequestCount int64 `json:"request_count"`
|
||||
UniqueVisitors int64 `json:"unique_visitors"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
EstimatedQPS float64 `json:"estimated_qps"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
// Capacity is the dashboard capacity section.
|
||||
type Capacity struct {
|
||||
AverageCPUUsagePercent float64 `json:"average_cpu_usage_percent"`
|
||||
AverageMemoryUsagePercent float64 `json:"average_memory_usage_percent"`
|
||||
HighCPUNodes int `json:"high_cpu_nodes"`
|
||||
HighMemoryNodes int `json:"high_memory_nodes"`
|
||||
HighStorageNodes int `json:"high_storage_nodes"`
|
||||
}
|
||||
|
||||
// NodeHealth is a dashboard node health row.
|
||||
type NodeHealth struct {
|
||||
ID uint `json:"id"`
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
GeoName string `json:"geo_name"`
|
||||
GeoLatitude *float64 `json:"geo_latitude"`
|
||||
GeoLongitude *float64 `json:"geo_longitude"`
|
||||
Status string `json:"status"`
|
||||
OpenrestyStatus string `json:"openresty_status"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
LastSeenAt any `json:"last_seen_at"`
|
||||
ActiveEventCount int `json:"active_event_count"`
|
||||
CPUUsagePercent float64 `json:"cpu_usage_percent"`
|
||||
MemoryUsagePercent float64 `json:"memory_usage_percent"`
|
||||
StorageUsagePercent float64 `json:"storage_usage_percent"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
}
|
||||
|
||||
// OverviewView is the expanded dashboard overview payload.
|
||||
type OverviewView struct {
|
||||
GeneratedAt time.Time `json:"generated_at"`
|
||||
Summary Summary `json:"summary"`
|
||||
Traffic Traffic `json:"traffic"`
|
||||
Capacity Capacity `json:"capacity"`
|
||||
Distributions observability.TrafficDistributions `json:"distributions"`
|
||||
Trends observability.NodeTrends `json:"trends"`
|
||||
Nodes []NodeHealth `json:"nodes"`
|
||||
}
|
||||
|
||||
// OverviewPayload is the compact legacy dashboard overview response.
|
||||
type OverviewPayload struct {
|
||||
GeneratedAt any `json:"generated_at"`
|
||||
Summary Summary `json:"summary"`
|
||||
Traffic Traffic `json:"traffic"`
|
||||
Capacity Capacity `json:"capacity"`
|
||||
Distributions distributionsPayload `json:"distributions"`
|
||||
Trends trendsPayload `json:"trends"`
|
||||
Nodes [][]any `json:"nodes"`
|
||||
}
|
||||
|
||||
type distributionsPayload struct {
|
||||
StatusCodes [][]any `json:"status_codes"`
|
||||
TopDomains [][]any `json:"top_domains"`
|
||||
SourceCountries [][]any `json:"source_countries"`
|
||||
}
|
||||
|
||||
type trendsPayload struct {
|
||||
Traffic24h [][]any `json:"traffic_24h"`
|
||||
Capacity24h [][]any `json:"capacity_24h"`
|
||||
Network24h [][]any `json:"network_24h"`
|
||||
DiskIO24h [][]any `json:"disk_io_24h"`
|
||||
}
|
||||
|
||||
// GetOverview aggregates dashboard overview data from nodes and observability tables.
|
||||
func GetOverview(ctx context.Context) (*OverviewPayload, error) {
|
||||
view, err := buildOverviewView(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return compressOverview(view), nil
|
||||
}
|
||||
|
||||
func buildOverviewView(ctx context.Context) (*OverviewView, error) {
|
||||
now := time.Now()
|
||||
since := now.Add(-24 * time.Hour)
|
||||
|
||||
nodes, err := model.ListOpenFlareNodes(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", since, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reports, err := model.ListOpenFlareRequestReportsSince(ctx, "", since, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, 8)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
activeEvents, err := model.ListOpenFlareActiveHealthEvents(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
openrestySnapshots, err := model.ListOpenFlareNodeObservationOpenresty(ctx, "", since, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
view := &OverviewView{
|
||||
GeneratedAt: now,
|
||||
Nodes: make([]NodeHealth, 0, len(nodes)),
|
||||
Distributions: observability.BuildTrafficDistributions(reports, accessLogRegions, 8),
|
||||
Trends: observability.NodeTrends{
|
||||
Traffic24h: observability.BuildTrafficTrendPoints(now, reports),
|
||||
Capacity24h: observability.BuildCapacityTrendPoints(now, snapshots),
|
||||
Network24h: observability.BuildNetworkTrendPoints(now, snapshots, openrestySnapshots),
|
||||
DiskIO24h: observability.BuildDiskIOTrendPoints(now, snapshots),
|
||||
},
|
||||
}
|
||||
|
||||
var cpuNodeCount int
|
||||
var memoryNodeCount int
|
||||
latestSnapshots := observability.LatestMetricSnapshotsByNode(snapshots)
|
||||
latestTrafficReports := observability.LatestTrafficReportsByNode(reports)
|
||||
activeEventsByNode := observability.ActiveHealthEventsByNode(activeEvents)
|
||||
|
||||
for _, node := range nodes {
|
||||
computedStatus := computeNodeStatus(&node)
|
||||
switch computedStatus {
|
||||
case nodeStatusOnline:
|
||||
view.Summary.OnlineNodes++
|
||||
case nodeStatusOffline:
|
||||
view.Summary.OfflineNodes++
|
||||
case nodeStatusPending:
|
||||
view.Summary.PendingNodes++
|
||||
}
|
||||
if node.OpenrestyStatus == "unhealthy" {
|
||||
view.Summary.UnhealthyNodes++
|
||||
}
|
||||
|
||||
latestSnapshot := latestSnapshots[node.NodeID]
|
||||
latestTraffic := latestTrafficReports[node.NodeID]
|
||||
nodeActiveEvents := activeEventsByNode[node.NodeID]
|
||||
|
||||
nodeHealth := NodeHealth{
|
||||
ID: node.ID,
|
||||
NodeID: node.NodeID,
|
||||
Name: node.Name,
|
||||
GeoName: node.GeoName,
|
||||
GeoLatitude: node.GeoLatitude,
|
||||
GeoLongitude: node.GeoLongitude,
|
||||
Status: computedStatus,
|
||||
OpenrestyStatus: node.OpenrestyStatus,
|
||||
CurrentVersion: node.CurrentVersion,
|
||||
LastSeenAt: nodeViewLastSeenAt(&node),
|
||||
ActiveEventCount: len(nodeActiveEvents),
|
||||
}
|
||||
|
||||
if latestSnapshot != nil {
|
||||
nodeHealth.CPUUsagePercent = latestSnapshot.CPUUsagePercent
|
||||
nodeHealth.MemoryUsagePercent = observability.Percentage(latestSnapshot.MemoryUsedBytes, latestSnapshot.MemoryTotalBytes)
|
||||
nodeHealth.StorageUsagePercent = observability.Percentage(latestSnapshot.StorageUsedBytes, latestSnapshot.StorageTotalBytes)
|
||||
if latestSnapshot.CPUUsagePercent > 0 {
|
||||
view.Capacity.AverageCPUUsagePercent += latestSnapshot.CPUUsagePercent
|
||||
cpuNodeCount++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent > 0 {
|
||||
view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent
|
||||
memoryNodeCount++
|
||||
}
|
||||
if latestSnapshot.CPUUsagePercent >= 80 {
|
||||
view.Capacity.HighCPUNodes++
|
||||
}
|
||||
if nodeHealth.MemoryUsagePercent >= 85 {
|
||||
view.Capacity.HighMemoryNodes++
|
||||
}
|
||||
if nodeHealth.StorageUsagePercent >= 85 {
|
||||
view.Capacity.HighStorageNodes++
|
||||
}
|
||||
}
|
||||
|
||||
if latestTraffic != nil {
|
||||
nodeHealth.RequestCount = latestTraffic.RequestCount
|
||||
nodeHealth.ErrorCount = latestTraffic.ErrorCount
|
||||
nodeHealth.UniqueVisitorCount = latestTraffic.UniqueVisitorCount
|
||||
view.Traffic.RequestCount += latestTraffic.RequestCount
|
||||
view.Traffic.UniqueVisitors += latestTraffic.UniqueVisitorCount
|
||||
view.Traffic.ErrorCount += latestTraffic.ErrorCount
|
||||
if duration := latestTraffic.WindowEndedAt.Sub(latestTraffic.WindowStartedAt).Seconds(); duration > 0 {
|
||||
view.Traffic.EstimatedQPS += float64(latestTraffic.RequestCount) / duration
|
||||
}
|
||||
view.Traffic.ReportedNodes++
|
||||
}
|
||||
|
||||
view.Nodes = append(view.Nodes, nodeHealth)
|
||||
}
|
||||
|
||||
view.Summary.TotalNodes = len(nodes)
|
||||
if cpuNodeCount > 0 {
|
||||
view.Capacity.AverageCPUUsagePercent /= float64(cpuNodeCount)
|
||||
}
|
||||
if memoryNodeCount > 0 {
|
||||
view.Capacity.AverageMemoryUsagePercent /= float64(memoryNodeCount)
|
||||
}
|
||||
|
||||
sort.Slice(view.Nodes, func(i int, j int) bool {
|
||||
if view.Nodes[i].ActiveEventCount == view.Nodes[j].ActiveEventCount {
|
||||
return view.Nodes[i].CPUUsagePercent > view.Nodes[j].CPUUsagePercent
|
||||
}
|
||||
return view.Nodes[i].ActiveEventCount > view.Nodes[j].ActiveEventCount
|
||||
})
|
||||
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func compressOverview(view *OverviewView) *OverviewPayload {
|
||||
if view == nil {
|
||||
return &OverviewPayload{
|
||||
Distributions: distributionsPayload{
|
||||
StatusCodes: [][]any{},
|
||||
TopDomains: [][]any{},
|
||||
SourceCountries: [][]any{},
|
||||
},
|
||||
Trends: trendsPayload{
|
||||
Traffic24h: [][]any{},
|
||||
Capacity24h: [][]any{},
|
||||
Network24h: [][]any{},
|
||||
DiskIO24h: [][]any{},
|
||||
},
|
||||
Nodes: [][]any{},
|
||||
}
|
||||
}
|
||||
return &OverviewPayload{
|
||||
GeneratedAt: view.GeneratedAt,
|
||||
Summary: view.Summary,
|
||||
Traffic: view.Traffic,
|
||||
Capacity: view.Capacity,
|
||||
Distributions: distributionsPayload{
|
||||
StatusCodes: compressDistributionItems(view.Distributions.StatusCodes),
|
||||
TopDomains: compressDistributionItems(view.Distributions.TopDomains),
|
||||
SourceCountries: compressDistributionItems(view.Distributions.SourceCountries),
|
||||
},
|
||||
Trends: trendsPayload{
|
||||
Traffic24h: compressTrafficTrendPoints(view.Trends.Traffic24h),
|
||||
Capacity24h: compressCapacityTrendPoints(view.Trends.Capacity24h),
|
||||
Network24h: compressNetworkTrendPoints(view.Trends.Network24h),
|
||||
DiskIO24h: compressDiskIOTrendPoints(view.Trends.DiskIO24h),
|
||||
},
|
||||
Nodes: compressDashboardNodes(view.Nodes),
|
||||
}
|
||||
}
|
||||
|
||||
func compressDistributionItems(items []observability.DistributionItem) [][]any {
|
||||
rows := make([][]any, 0, len(items))
|
||||
for _, item := range items {
|
||||
rows = append(rows, []any{item.Key, item.Value})
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func compressTrafficTrendPoints(points []observability.TrafficTrendPoint) [][]any {
|
||||
rows := make([][]any, 0, len(points))
|
||||
for _, point := range points {
|
||||
rows = append(rows, []any{
|
||||
point.BucketStartedAt,
|
||||
point.RequestCount,
|
||||
point.ErrorCount,
|
||||
point.UniqueVisitorCount,
|
||||
})
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func compressCapacityTrendPoints(points []observability.CapacityTrendPoint) [][]any {
|
||||
rows := make([][]any, 0, len(points))
|
||||
for _, point := range points {
|
||||
rows = append(rows, []any{
|
||||
point.BucketStartedAt,
|
||||
point.AverageCPUUsagePercent,
|
||||
point.AverageMemoryUsagePercent,
|
||||
point.ReportedNodes,
|
||||
})
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func compressNetworkTrendPoints(points []observability.NetworkTrendPoint) [][]any {
|
||||
rows := make([][]any, 0, len(points))
|
||||
for _, point := range points {
|
||||
rows = append(rows, []any{
|
||||
point.BucketStartedAt,
|
||||
point.NetworkRxBytes,
|
||||
point.NetworkTxBytes,
|
||||
point.OpenrestyRxBytes,
|
||||
point.OpenrestyTxBytes,
|
||||
point.ReportedNodes,
|
||||
})
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func compressDiskIOTrendPoints(points []observability.DiskIOTrendPoint) [][]any {
|
||||
rows := make([][]any, 0, len(points))
|
||||
for _, point := range points {
|
||||
rows = append(rows, []any{
|
||||
point.BucketStartedAt,
|
||||
point.DiskReadBytes,
|
||||
point.DiskWriteBytes,
|
||||
point.ReportedNodes,
|
||||
})
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func compressDashboardNodes(nodes []NodeHealth) [][]any {
|
||||
rows := make([][]any, 0, len(nodes))
|
||||
for _, node := range nodes {
|
||||
rows = append(rows, []any{
|
||||
node.ID,
|
||||
node.NodeID,
|
||||
node.Name,
|
||||
node.GeoName,
|
||||
node.GeoLatitude,
|
||||
node.GeoLongitude,
|
||||
node.Status,
|
||||
node.OpenrestyStatus,
|
||||
node.CurrentVersion,
|
||||
node.LastSeenAt,
|
||||
node.ActiveEventCount,
|
||||
node.CPUUsagePercent,
|
||||
node.MemoryUsagePercent,
|
||||
node.StorageUsagePercent,
|
||||
node.RequestCount,
|
||||
node.ErrorCount,
|
||||
node.UniqueVisitorCount,
|
||||
})
|
||||
}
|
||||
return rows
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupDashboardTestDB(t *testing.T) func() {
|
||||
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)
|
||||
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore())
|
||||
return func() {
|
||||
resetAccessLogStore()
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetOverviewStructure(t *testing.T) {
|
||||
cleanup := setupDashboardTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC()
|
||||
lastSeen := now.Add(-time.Minute)
|
||||
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-dashboard-1",
|
||||
Name: "Edge 1",
|
||||
IP: "10.0.0.1",
|
||||
Status: "online",
|
||||
OpenrestyStatus: "healthy",
|
||||
CurrentVersion: "v1.0.0",
|
||||
LastSeenAt: &lastSeen,
|
||||
}).Error)
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
|
||||
NodeID: "node-dashboard-2",
|
||||
Name: "Edge 2",
|
||||
IP: "10.0.0.2",
|
||||
Status: "pending",
|
||||
OpenrestyStatus: "unknown",
|
||||
}).Error)
|
||||
|
||||
overview, err := GetOverview(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, overview)
|
||||
|
||||
assert.False(t, overview.GeneratedAt.(time.Time).IsZero())
|
||||
assert.Equal(t, 2, overview.Summary.TotalNodes)
|
||||
assert.Equal(t, 1, overview.Summary.OnlineNodes)
|
||||
assert.Equal(t, 1, overview.Summary.PendingNodes)
|
||||
assert.Equal(t, 0, overview.Summary.OfflineNodes)
|
||||
assert.Equal(t, 0, overview.Summary.UnhealthyNodes)
|
||||
|
||||
assert.Equal(t, int64(0), overview.Traffic.RequestCount)
|
||||
assert.Equal(t, int64(0), overview.Traffic.UniqueVisitors)
|
||||
assert.Equal(t, int64(0), overview.Traffic.ErrorCount)
|
||||
assert.Equal(t, float64(0), overview.Traffic.EstimatedQPS)
|
||||
assert.Equal(t, 0, overview.Traffic.ReportedNodes)
|
||||
|
||||
assert.Equal(t, float64(0), overview.Capacity.AverageCPUUsagePercent)
|
||||
assert.Equal(t, float64(0), overview.Capacity.AverageMemoryUsagePercent)
|
||||
assert.Equal(t, 0, overview.Capacity.HighCPUNodes)
|
||||
assert.Equal(t, 0, overview.Capacity.HighMemoryNodes)
|
||||
assert.Equal(t, 0, overview.Capacity.HighStorageNodes)
|
||||
|
||||
require.NotNil(t, overview.Distributions.StatusCodes)
|
||||
require.NotNil(t, overview.Distributions.TopDomains)
|
||||
require.NotNil(t, overview.Distributions.SourceCountries)
|
||||
assert.Empty(t, overview.Distributions.StatusCodes)
|
||||
assert.Empty(t, overview.Distributions.TopDomains)
|
||||
assert.Empty(t, overview.Distributions.SourceCountries)
|
||||
|
||||
require.Len(t, overview.Trends.Traffic24h, 24)
|
||||
require.Len(t, overview.Trends.Capacity24h, 24)
|
||||
require.Len(t, overview.Trends.Network24h, 24)
|
||||
require.Len(t, overview.Trends.DiskIO24h, 24)
|
||||
for _, row := range overview.Trends.Traffic24h {
|
||||
require.Len(t, row, 4)
|
||||
}
|
||||
for _, row := range overview.Trends.Capacity24h {
|
||||
require.Len(t, row, 4)
|
||||
}
|
||||
for _, row := range overview.Trends.Network24h {
|
||||
require.Len(t, row, 6)
|
||||
}
|
||||
for _, row := range overview.Trends.DiskIO24h {
|
||||
require.Len(t, row, 4)
|
||||
}
|
||||
|
||||
require.Len(t, overview.Nodes, 2)
|
||||
for _, row := range overview.Nodes {
|
||||
require.Len(t, row, 17)
|
||||
}
|
||||
|
||||
nodeByID := make(map[string][]any, len(overview.Nodes))
|
||||
for _, row := range overview.Nodes {
|
||||
nodeByID[row[1].(string)] = row
|
||||
}
|
||||
|
||||
onlineNode := nodeByID["node-dashboard-1"]
|
||||
require.NotNil(t, onlineNode)
|
||||
assert.Equal(t, "Edge 1", onlineNode[2])
|
||||
assert.Equal(t, "online", onlineNode[6])
|
||||
assert.Equal(t, "healthy", onlineNode[7])
|
||||
|
||||
pendingNode := nodeByID["node-dashboard-2"]
|
||||
require.NotNil(t, pendingNode)
|
||||
assert.Equal(t, "Edge 2", pendingNode[2])
|
||||
assert.Equal(t, "pending", pendingNode[6])
|
||||
assert.Equal(t, "unknown", pendingNode[7])
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package dashboard
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GetOverviewHandler 获取仪表盘概览数据。
|
||||
// @Summary 获取仪表盘概览
|
||||
// @Description 聚合节点与可观测性数据,返回 OpenFlare 控制台仪表盘概览,需要管理员权限
|
||||
// @Tags openflare-dashboard
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=dashboard.OverviewPayload} "仪表盘概览"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/dashboard/overview [get]
|
||||
func GetOverviewHandler(c *gin.Context) {
|
||||
overview, err := GetOverview(c.Request.Context())
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(overview))
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
const (
|
||||
errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效"
|
||||
errTunnelNodeTypeMismatch = "此节点不是 TunnelClient 类型"
|
||||
)
|
||||
@@ -0,0 +1,176 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/relay"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type configVersionRow struct {
|
||||
Version string `gorm:"column:version"`
|
||||
Checksum string `gorm:"column:checksum"`
|
||||
}
|
||||
|
||||
func (configVersionRow) TableName() string {
|
||||
return "of_config_versions"
|
||||
}
|
||||
|
||||
func normalizeReleaseChannel(channel string) string {
|
||||
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
|
||||
return "preview"
|
||||
}
|
||||
return "stable"
|
||||
}
|
||||
|
||||
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) {
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return nil, errors.New("database not initialized")
|
||||
}
|
||||
if !conn.Migrator().HasTable(&configVersionRow{}) {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
var version configVersionRow
|
||||
err := conn.Where("is_active = ?", true).Order("id desc").First(&version).Error
|
||||
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 := model.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 decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
domain := strings.ToLower(strings.TrimSpace(fallbackDomain))
|
||||
if domain == "" {
|
||||
return nil, errors.New("domain is required")
|
||||
}
|
||||
return []string{domain}, nil
|
||||
}
|
||||
var domains []string
|
||||
if err := json.Unmarshal([]byte(text), &domains); err != nil {
|
||||
return nil, errors.New("domains payload is invalid")
|
||||
}
|
||||
normalized := make([]string, 0, len(domains))
|
||||
for _, item := range domains {
|
||||
domain := strings.ToLower(strings.TrimSpace(item))
|
||||
if domain == "" {
|
||||
continue
|
||||
}
|
||||
normalized = append(normalized, domain)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, errors.New("domain is required")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func parseTunnelTargetAddr(addr string) (string, int) {
|
||||
addr = strings.TrimSpace(addr)
|
||||
if addr == "" {
|
||||
return "127.0.0.1", 80
|
||||
}
|
||||
host, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
lastColon := strings.LastIndex(addr, ":")
|
||||
if lastColon < 0 {
|
||||
return addr, 80
|
||||
}
|
||||
host = addr[:lastColon]
|
||||
portStr = addr[lastColon+1:]
|
||||
}
|
||||
port := 80
|
||||
if _, scanErr := fmt.Sscanf(portStr, "%d", &port); scanErr != nil {
|
||||
port = 80
|
||||
}
|
||||
if host == "" {
|
||||
host = "127.0.0.1"
|
||||
}
|
||||
return host, port
|
||||
}
|
||||
|
||||
func sanitizeProxyName(domain string) string {
|
||||
return strings.ReplaceAll(strings.ReplaceAll(domain, ".", "-"), "*", "wildcard")
|
||||
}
|
||||
|
||||
func buildTunnelSettings(node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *relay.Settings {
|
||||
return relay.BuildSettings(node, updateNow, updateChannel, updateTag)
|
||||
}
|
||||
@@ -0,0 +1,303 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/relay"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
applyResultOK = "success"
|
||||
applyResultWarn = "warning"
|
||||
applyResultFail = "failed"
|
||||
)
|
||||
|
||||
// HeartbeatPayload is sent by OpenFlared on each heartbeat.
|
||||
type HeartbeatPayload struct {
|
||||
ClientVersion string `json:"client_version"`
|
||||
FrpVersion string `json:"frp_version"`
|
||||
IP string `json:"ip"`
|
||||
TunnelStatus string `json:"tunnel_status"`
|
||||
ConnectedRelays []ConnectedRelay `json:"connected_relays"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
CurrentChecksum string `json:"current_checksum"`
|
||||
}
|
||||
|
||||
// ConnectedRelay describes relay connection status from the client.
|
||||
type ConnectedRelay struct {
|
||||
RelayNodeID string `json:"relay_node_id"`
|
||||
Status string `json:"status"`
|
||||
ProxyCount int `json:"proxy_count"`
|
||||
}
|
||||
|
||||
// ActiveConfigMeta summarizes the active config version.
|
||||
type ActiveConfigMeta struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// HeartbeatResponse is returned to the OpenFlared client.
|
||||
type HeartbeatResponse struct {
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config"`
|
||||
TunnelSettings *relay.Settings `json:"tunnel_settings"`
|
||||
}
|
||||
|
||||
// TunnelConfigResponse is the full tunnel routing config sent to the client.
|
||||
type TunnelConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
Relays []RelayInfo `json:"relays"`
|
||||
Proxies []ProxyEntry `json:"proxies"`
|
||||
}
|
||||
|
||||
// RelayInfo describes a relay the client should connect to.
|
||||
type RelayInfo struct {
|
||||
RelayNodeID string `json:"relay_node_id"`
|
||||
Address string `json:"address"`
|
||||
AuthToken string `json:"auth_token"`
|
||||
ProxyURL string `json:"proxy_url"`
|
||||
}
|
||||
|
||||
// ProxyEntry describes one frpc proxy definition.
|
||||
type ProxyEntry struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
LocalAddr string `json:"local_addr"`
|
||||
LocalPort int `json:"local_port"`
|
||||
CustomDomains []string `json:"custom_domains"`
|
||||
}
|
||||
|
||||
// ApplyLogPayload is the apply result reported by OpenFlared.
|
||||
type ApplyLogPayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Version string `json:"version"`
|
||||
Result string `json:"result"`
|
||||
Message string `json:"message"`
|
||||
Checksum string `json:"checksum"`
|
||||
MainConfigChecksum string `json:"main_config_checksum"`
|
||||
RouteConfigChecksum string `json:"route_config_checksum"`
|
||||
SupportFileCount int `json:"support_file_count"`
|
||||
}
|
||||
|
||||
// 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, fmt.Errorf("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": "stable",
|
||||
"update_tag": "",
|
||||
}
|
||||
if !previous.UpdateRequested {
|
||||
delete(changes, "update_requested")
|
||||
}
|
||||
if previous.UpdateChannel == "stable" {
|
||||
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 = "stable"
|
||||
node.UpdateTag = ""
|
||||
lastSeen := now
|
||||
node.LastSeenAt = &lastSeen
|
||||
node.Status = nodeStatusOnline
|
||||
|
||||
if err := db.DB(ctx).Model(node).Updates(changes).Error; 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(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, fmt.Errorf("node is nil")
|
||||
}
|
||||
|
||||
activeVersion, err := getActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("no active config version: %w", err)
|
||||
}
|
||||
|
||||
routes, err := model.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
|
||||
}
|
||||
domains, decodeErr := decodeStoredDomains(route.Domains, route.Domain)
|
||||
if decodeErr != nil {
|
||||
continue
|
||||
}
|
||||
localAddr, localPort := parseTunnelTargetAddr(route.TunnelTargetAddr)
|
||||
proxies = append(proxies, ProxyEntry{
|
||||
Name: fmt.Sprintf("%s-%s", node.NodeID, sanitizeProxyName(domains[0])),
|
||||
Type: "http",
|
||||
LocalAddr: localAddr,
|
||||
LocalPort: localPort,
|
||||
CustomDomains: domains,
|
||||
})
|
||||
}
|
||||
|
||||
return &TunnelConfigResponse{
|
||||
Version: activeVersion.Version,
|
||||
Checksum: activeVersion.Checksum,
|
||||
Relays: relays,
|
||||
Proxies: proxies,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// 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")
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
|
||||
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var node model.OpenFlareNode
|
||||
if err := tx.Where("node_id = ?", payload.NodeID).First(&node).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
node.Status = nodeStatusOnline
|
||||
lastSeen := now
|
||||
node.LastSeenAt = &lastSeen
|
||||
if payload.Result == applyResultOK {
|
||||
node.CurrentVersion = payload.Version
|
||||
node.LastError = ""
|
||||
} else {
|
||||
node.LastError = payload.Message
|
||||
}
|
||||
if err := tx.Create(log).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&node).Select("status", "last_seen_at", "current_version", "last_error").Updates(&node).Error
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return log, nil
|
||||
}
|
||||
|
||||
func getActiveConfigVersion(ctx context.Context) (*configVersionRow, error) {
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return nil, errors.New("database not initialized")
|
||||
}
|
||||
if !conn.Migrator().HasTable(&configVersionRow{}) {
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
var version configVersionRow
|
||||
if err := conn.Where("is_active = ?", true).Order("id desc").First(&version).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &version, 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) > 16000 {
|
||||
payload.Message = payload.Message[:16000]
|
||||
}
|
||||
return payload
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
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 := 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()
|
||||
}
|
||||
}
|
||||
|
||||
func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
if token == "" {
|
||||
return nil, errors.New("missing tunnel token")
|
||||
}
|
||||
node, err := model.GetOpenFlareNodeByAccessToken(ctx, token)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, errors.New("invalid tunnel token")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return node, nil
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"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, model.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,50 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"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, ","),
|
||||
},
|
||||
})
|
||||
}
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return
|
||||
}
|
||||
if err := agent.ReconcileScopedNodeHealthEvents(conn, 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,71 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func 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{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
agent.ResetAuthCacheForTest()
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
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 := model.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,120 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package flared
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"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 := authNode.(*model.OpenFlareNode)
|
||||
|
||||
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 := authNode.(*model.OpenFlareNode)
|
||||
|
||||
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 {
|
||||
payload.NodeID = authNode.(*model.OpenFlareNode).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 := authNode.(*model.OpenFlareNode)
|
||||
ofws.ServeFlared(c, node.NodeID)
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package geoip provides OpenFlare-compatible GeoIP lookup helpers.
|
||||
package geoip
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
|
||||
)
|
||||
|
||||
// LookupView is the legacy OpenFlare GeoIP lookup response shape.
|
||||
type LookupView struct {
|
||||
Provider string `json:"provider"`
|
||||
IP string `json:"ip"`
|
||||
ISOCode string `json:"iso_code"`
|
||||
Name string `json:"name"`
|
||||
Latitude *float64 `json:"latitude,omitempty"`
|
||||
Longitude *float64 `json:"longitude,omitempty"`
|
||||
}
|
||||
|
||||
const (
|
||||
errProviderInvalid = "归属方式仅支持 disabled、mmdb、ip-api、geojs、ipinfo"
|
||||
errIPEmpty = "IP 不能为空"
|
||||
errIPInvalid = "IP 格式无效"
|
||||
errLookupEmpty = "未获取到 IP 归属结果"
|
||||
)
|
||||
|
||||
// IsValidProvider reports whether provider is a supported GeoIP backend.
|
||||
func IsValidProvider(provider string) bool {
|
||||
return pkggeoip.IsValidProvider(provider)
|
||||
}
|
||||
|
||||
// GeoInfoFromIP resolves geographic information using the configured default provider.
|
||||
func GeoInfoFromIP(ip net.IP) (*pkggeoip.GeoInfo, error) {
|
||||
return pkggeoip.GetGeoInfo(ip)
|
||||
}
|
||||
|
||||
// Lookup resolves geographic information for rawIP using the given provider.
|
||||
func Lookup(provider, rawIP string) (*LookupView, error) {
|
||||
trimmedProvider := strings.TrimSpace(provider)
|
||||
if !pkggeoip.IsValidProvider(trimmedProvider) {
|
||||
return nil, errors.New(errProviderInvalid)
|
||||
}
|
||||
|
||||
trimmedIP := strings.TrimSpace(rawIP)
|
||||
if trimmedIP == "" {
|
||||
return nil, errors.New(errIPEmpty)
|
||||
}
|
||||
parsedIP := net.ParseIP(trimmedIP)
|
||||
if parsedIP == nil {
|
||||
return nil, errors.New(errIPInvalid)
|
||||
}
|
||||
|
||||
if trimmedProvider == pkggeoip.ProviderDisabled {
|
||||
return &LookupView{
|
||||
Provider: trimmedProvider,
|
||||
IP: parsedIP.String(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
info, err := pkggeoip.LookupGeoInfoWithProvider(trimmedProvider, parsedIP)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if info == nil {
|
||||
return nil, errors.New(errLookupEmpty)
|
||||
}
|
||||
|
||||
return &LookupView{
|
||||
Provider: trimmedProvider,
|
||||
IP: parsedIP.String(),
|
||||
ISOCode: info.ISOCode,
|
||||
Name: info.Name,
|
||||
Latitude: info.Latitude,
|
||||
Longitude: info.Longitude,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package geoip
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
|
||||
)
|
||||
|
||||
type fakeLookupProvider struct{}
|
||||
|
||||
func (f *fakeLookupProvider) Name() string { return "fake-lookup" }
|
||||
|
||||
func (f *fakeLookupProvider) GetGeoInfo(_ net.IP) (*pkggeoip.GeoInfo, error) {
|
||||
lat := 37.7749
|
||||
lon := -122.4194
|
||||
return &pkggeoip.GeoInfo{
|
||||
ISOCode: "US",
|
||||
Name: "United States",
|
||||
Latitude: &lat,
|
||||
Longitude: &lon,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeLookupProvider) UpdateDatabase() error { return nil }
|
||||
|
||||
func (f *fakeLookupProvider) Close() error { return nil }
|
||||
|
||||
func TestLookupWithProvider(t *testing.T) {
|
||||
previousFactory := pkggeoip.ProviderFactoryForTest()
|
||||
pkggeoip.SetProviderFactoryForTest(func(provider string) (pkggeoip.GeoIPService, error) {
|
||||
return &fakeLookupProvider{}, nil
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
pkggeoip.SetProviderFactoryForTest(previousFactory)
|
||||
})
|
||||
|
||||
view, err := Lookup("ipinfo", "8.8.8.8")
|
||||
if err != nil {
|
||||
t.Fatalf("Lookup() error = %v", err)
|
||||
}
|
||||
if view.Provider != "ipinfo" || view.IP != "8.8.8.8" {
|
||||
t.Fatalf("unexpected lookup view: %+v", view)
|
||||
}
|
||||
if view.ISOCode != "US" || view.Name != "United States" {
|
||||
t.Fatalf("unexpected geo fields: %+v", view)
|
||||
}
|
||||
if view.Latitude == nil || view.Longitude == nil {
|
||||
t.Fatalf("expected coordinates, got %+v", view)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLookupRejectsInvalidInput(t *testing.T) {
|
||||
if _, err := Lookup("invalid", "8.8.8.8"); err == nil {
|
||||
t.Fatal("expected invalid provider to fail")
|
||||
}
|
||||
if _, err := Lookup("ipinfo", "not-an-ip"); err == nil {
|
||||
t.Fatal("expected invalid IP to fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLookupDisabledProvider(t *testing.T) {
|
||||
view, err := Lookup("disabled", "8.8.8.8")
|
||||
if err != nil {
|
||||
t.Fatalf("Lookup() error = %v", err)
|
||||
}
|
||||
if view.Provider != "disabled" || view.IP != "8.8.8.8" {
|
||||
t.Fatalf("unexpected disabled view: %+v", view)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
ofnode "github.com/Rain-kl/Wavelet/internal/apps/openflare/node"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type configVersionRecord struct {
|
||||
ID uint `gorm:"primaryKey"`
|
||||
Version string `gorm:"column:version"`
|
||||
SnapshotJSON string `gorm:"column:snapshot_json"`
|
||||
SupportFilesJSON string `gorm:"column:support_files_json"`
|
||||
Checksum string `gorm:"column:checksum"`
|
||||
IsActive bool `gorm:"column:is_active"`
|
||||
}
|
||||
|
||||
func (configVersionRecord) TableName() string {
|
||||
return "of_config_versions"
|
||||
}
|
||||
|
||||
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.OpenFlareOption{},
|
||||
&model.OpenFlareApplyLog{},
|
||||
&model.OpenFlareNodeSystemProfile{},
|
||||
&model.OpenFlareMetricSnapshot{},
|
||||
&model.OpenFlareHealthEvent{},
|
||||
&model.OpenFlareNodeObservationFrps{},
|
||||
&model.OpenFlareNodeObservationFrpc{},
|
||||
&configVersionRecord{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
agent.ResetAuthCacheForTest()
|
||||
|
||||
engine := testhelper.NewTestGinEngine()
|
||||
mountOpenFlareTestRoutes(engine)
|
||||
|
||||
cleanup := func() {
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
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 := model.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 := model.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 := model.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 := model.GetOpenFlareNodeByNodeID(ctx, edge.NodeID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "online", stored.Status)
|
||||
assert.Equal(t, "20260618-001", stored.CurrentVersion)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,211 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/cap"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"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 {
|
||||
SystemName string `json:"system_name"`
|
||||
}
|
||||
|
||||
func setupAuthOptionIntegration(t *testing.T) (*gorm.DB, *gin.Engine) {
|
||||
t.Helper()
|
||||
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
t.Cleanup(cleanup)
|
||||
|
||||
require.NoError(t, dbConn.AutoMigrate(&model.OpenFlareOption{}))
|
||||
option.ResetInitializationForTest()
|
||||
t.Cleanup(option.ResetInitializationForTest)
|
||||
|
||||
require.NoError(t, dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeyCapLoginEnabled).
|
||||
Update("value", "false").Error)
|
||||
require.NoError(t, repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled))
|
||||
cap.InvalidateRuntimeSettings()
|
||||
|
||||
oldCookieName := config.Config.App.SessionCookieName
|
||||
oldSecret := config.Config.App.SessionSecret
|
||||
oldDomain := config.Config.App.SessionDomain
|
||||
oldSecure := config.Config.App.SessionSecure
|
||||
oldHTTPOnly := config.Config.App.SessionHTTPOnly
|
||||
t.Cleanup(func() {
|
||||
config.Config.App.SessionCookieName = oldCookieName
|
||||
config.Config.App.SessionSecret = oldSecret
|
||||
config.Config.App.SessionDomain = oldDomain
|
||||
config.Config.App.SessionSecure = oldSecure
|
||||
config.Config.App.SessionHTTPOnly = oldHTTPOnly
|
||||
})
|
||||
|
||||
config.Config.App.SessionCookieName = "test_openflare_session"
|
||||
config.Config.App.SessionSecret = "test_openflare_session_secret"
|
||||
config.Config.App.SessionDomain = ""
|
||||
config.Config.App.SessionSecure = false
|
||||
config.Config.App.SessionHTTPOnly = true
|
||||
|
||||
store := cookie.NewStore([]byte(config.Config.App.SessionSecret))
|
||||
store.Options(oauth.GetSessionOptions(3600))
|
||||
r := testhelper.NewTestGinEngine(sessions.Sessions(config.Config.App.SessionCookieName, 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.SystemName)
|
||||
}
|
||||
|
||||
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) {
|
||||
w := performJSONRequest(t, r, http.MethodGet, apiPath("/option/"), nil, nil)
|
||||
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
resp := decodeAPIResponse(t, w)
|
||||
assert.NotEmpty(t, resp.ErrorMsg)
|
||||
})
|
||||
|
||||
t.Run("non-admin user forbidden", func(t *testing.T) {
|
||||
w := performJSONRequest(t, r, http.MethodGet, apiPath("/option/"), nil, adminAuthHeaders(commonToken))
|
||||
assert.Equal(t, http.StatusNotFound, w.Code)
|
||||
resp := decodeAPIResponse(t, w)
|
||||
assert.Equal(t, admin.TokenAdminRequired, resp.ErrorMsg)
|
||||
})
|
||||
|
||||
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 TestOptionHotReloadAfterUpdate(t *testing.T) {
|
||||
dbConn, r := setupAuthOptionIntegration(t)
|
||||
adminToken := seedUserWithAccessToken(t, dbConn, "admin", "password123", true)
|
||||
|
||||
statusBefore := getStatusSystemName(t, r, nil)
|
||||
assert.NotEmpty(t, statusBefore)
|
||||
|
||||
updateResp := performJSONRequest(t, r, http.MethodPost, apiPath("/option/update"), map[string]string{
|
||||
"key": "SystemName",
|
||||
"value": "HotReloadIntegration",
|
||||
}, adminAuthHeaders(adminToken))
|
||||
assert.Equal(t, http.StatusOK, updateResp.Code)
|
||||
requireAPIOK(t, updateResp)
|
||||
|
||||
statusAfter := getStatusSystemName(t, r, nil)
|
||||
assert.Equal(t, "HotReloadIntegration", statusAfter)
|
||||
assert.Equal(t, "HotReloadIntegration", model.SystemName)
|
||||
|
||||
ctx := context.Background()
|
||||
require.NoError(t, option.EnsureInitialized(ctx))
|
||||
assert.Equal(t, "HotReloadIntegration", model.OptionValue("SystemName"))
|
||||
}
|
||||
|
||||
func getStatusSystemName(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.SystemName
|
||||
}
|
||||
@@ -0,0 +1,275 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
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.OpenFlareOption{},
|
||||
&model.OpenFlareApplyLog{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
agent.ResetAuthCacheForTest()
|
||||
|
||||
seed, err := seedAdminWithAccessToken(sqliteDB)
|
||||
require.NoError(t, err)
|
||||
|
||||
engine := testhelper.NewTestGinEngine()
|
||||
mountOpenFlareTestRoutes(engine)
|
||||
|
||||
cleanup := func() {
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
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) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{
|
||||
"site_name": "core-chain-site",
|
||||
"domain": "core-chain.example.com",
|
||||
"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.Equal(t, "core-chain.example.com", data["domain"])
|
||||
assert.Equal(t, float64(originID), data["origin_id"])
|
||||
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.Equal(t, float64(1), listData["total"])
|
||||
|
||||
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.Equal(t, float64(nodeID), nodeView["id"])
|
||||
assert.Equal(t, nodePublicID, nodeView["node_id"])
|
||||
assert.Equal(t, "success", nodeView["latest_apply_result"])
|
||||
assert.Equal(t, configChecksum, nodeView["latest_apply_checksum"])
|
||||
assert.Equal(t, float64(2), nodeView["latest_support_file_count"])
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
v1 "github.com/Rain-kl/Wavelet/internal/router/v1"
|
||||
ofrouter "github.com/Rain-kl/Wavelet/internal/router/v1/openflare"
|
||||
"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
|
||||
}
|
||||
|
||||
func mountOpenFlareTestRoutes(engine *gin.Engine) {
|
||||
api := engine.Group("/api")
|
||||
apiV1 := api.Group("/v1")
|
||||
v1.RegisterV1Routes(apiV1, api)
|
||||
}
|
||||
|
||||
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,360 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package integration
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func 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.ManagedDomain{},
|
||||
&model.DNSAccount{},
|
||||
&model.AcmeAccount{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
|
||||
seed, err := seedAdminWithAccessToken(sqliteDB)
|
||||
require.NoError(t, err)
|
||||
|
||||
oldSecret := config.Config.App.SessionSecret
|
||||
config.Config.App.SessionSecret = "test_session_secret_for_security_integration"
|
||||
|
||||
engine := testhelper.NewTestGinEngine()
|
||||
mountOpenFlareTestRoutes(engine)
|
||||
|
||||
cleanup := func() {
|
||||
config.Config.App.SessionSecret = oldSecret
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
}
|
||||
|
||||
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",
|
||||
"enabled": true,
|
||||
"block_status_code": 403,
|
||||
"ip_whitelist": []string{"192.0.2.1"},
|
||||
"ip_blacklist": []string{"203.0.113.10"},
|
||||
"country_blacklist": []string{"CN"},
|
||||
"remark": "integration rule group",
|
||||
}, 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.Equal(t, float64(403), data["block_status_code"])
|
||||
})
|
||||
|
||||
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.Equal(t, float64(ruleGroupID), data["id"])
|
||||
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/update", apiPath(""), ruleGroupID),
|
||||
map[string]any{
|
||||
"name": "edge-security-updated",
|
||||
"enabled": true,
|
||||
"block_status_code": 451,
|
||||
"remark": "updated by integration test",
|
||||
},
|
||||
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, float64(451), data["block_status_code"])
|
||||
})
|
||||
|
||||
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"},
|
||||
"remark": "manual deny list",
|
||||
}, 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) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{
|
||||
"site_name": "security-site",
|
||||
"domain": "security.example.com",
|
||||
"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.Equal(t, "security.example.com", data["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.Equal(t, float64(proxyRouteID), data["route_id"])
|
||||
|
||||
appliedIDs, ok := data["applied_ids"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Len(t, appliedIDs, 1)
|
||||
assert.Equal(t, float64(ruleGroupID), appliedIDs[0])
|
||||
})
|
||||
|
||||
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.Equal(t, float64(ruleGroupID), group["id"])
|
||||
})
|
||||
|
||||
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 managed domain", func(t *testing.T) {
|
||||
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/managed-domains/"), map[string]any{
|
||||
"domain": "security.example.com",
|
||||
"cert_id": certID,
|
||||
"enabled": true,
|
||||
"remark": "primary security domain",
|
||||
}, 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.Equal(t, float64(certID), data["cert_id"])
|
||||
assert.Equal(t, true, data["enabled"])
|
||||
})
|
||||
|
||||
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,21 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
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,473 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeStatusOnline = "online"
|
||||
nodeStatusOffline = "offline"
|
||||
nodeStatusPending = "pending"
|
||||
openrestyStatusHealthy = "healthy"
|
||||
openrestyStatusUnhealthy = "unhealthy"
|
||||
openrestyStatusUnknown = "unknown"
|
||||
githubReleasesAPIBase = "https://api.github.com/repos/%s/releases"
|
||||
)
|
||||
|
||||
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, 16)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
func newServerNodeID() (string, error) {
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "node-" + token, nil
|
||||
}
|
||||
|
||||
func normalizeNodeType(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "tunnel_relay":
|
||||
return "tunnel_relay"
|
||||
case "tunnel_client":
|
||||
return "tunnel_client"
|
||||
default:
|
||||
return "edge_node"
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
manualOverride := input.GeoManualOverride || geoName != "" || input.GeoLatitude != nil || input.GeoLongitude != nil
|
||||
if len(ip) > 64 {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPTooLong)
|
||||
}
|
||||
if ip != "" && net.ParseIP(ip) == nil {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPInvalid)
|
||||
}
|
||||
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPManualRequired)
|
||||
}
|
||||
if len(geoName) > 128 {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoNameTooLong)
|
||||
}
|
||||
|
||||
geoLatitude := cloneCoordinate(input.GeoLatitude)
|
||||
geoLongitude := cloneCoordinate(input.GeoLongitude)
|
||||
if (geoLatitude == nil) != (geoLongitude == nil) {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoCoordinateMismatch)
|
||||
}
|
||||
if geoLatitude != nil && (*geoLatitude < -90 || *geoLatitude > 90) {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLatitudeInvalid)
|
||||
}
|
||||
if geoLongitude != nil && (*geoLongitude < -180 || *geoLongitude > 180) {
|
||||
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLongitudeInvalid)
|
||||
}
|
||||
|
||||
if !manualOverride {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
if geoLatitude == nil && geoLongitude == nil && geoName == "" {
|
||||
return name, ip, "", nil, nil, false, nil
|
||||
}
|
||||
|
||||
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
|
||||
}
|
||||
|
||||
func computeNodeStatus(node *model.OpenFlareNode) string {
|
||||
if node == nil {
|
||||
return nodeStatusOffline
|
||||
}
|
||||
if node.LastSeenAt == nil || node.LastSeenAt.IsZero() {
|
||||
return nodeStatusPending
|
||||
}
|
||||
if time.Since(*node.LastSeenAt) > model.NodeOfflineThreshold {
|
||||
return nodeStatusOffline
|
||||
}
|
||||
return nodeStatusOnline
|
||||
}
|
||||
|
||||
func nodeViewLastSeenAt(node *model.OpenFlareNode) any {
|
||||
if node == nil {
|
||||
return time.Time{}
|
||||
}
|
||||
nodeType := strings.TrimSpace(node.NodeType)
|
||||
if nodeType == "" {
|
||||
nodeType = "edge_node"
|
||||
}
|
||||
if nodeType == "tunnel_relay" && ofws.IsRelayConnected(node.NodeID) {
|
||||
return ofws.RelayWSConnectedLastSeenValue
|
||||
}
|
||||
if nodeType == "tunnel_client" && 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 = "edge_node"
|
||||
}
|
||||
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 compareVersions(local, remote string) int {
|
||||
left := parseVersionInfo(local)
|
||||
right := parseVersionInfo(remote)
|
||||
if left.isDev {
|
||||
if right.valid {
|
||||
return -1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
if !left.valid || !right.valid {
|
||||
return 0
|
||||
}
|
||||
|
||||
maxLen := len(left.numbers)
|
||||
if len(right.numbers) > maxLen {
|
||||
maxLen = len(right.numbers)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
leftValue := 0
|
||||
rightValue := 0
|
||||
if index < len(left.numbers) {
|
||||
leftValue = left.numbers[index]
|
||||
}
|
||||
if index < len(right.numbers) {
|
||||
rightValue = right.numbers[index]
|
||||
}
|
||||
if leftValue < rightValue {
|
||||
return -1
|
||||
}
|
||||
if leftValue > rightValue {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
type versionInfo struct {
|
||||
valid bool
|
||||
isDev bool
|
||||
numbers []int
|
||||
}
|
||||
|
||||
func parseVersionInfo(version string) versionInfo {
|
||||
normalized := strings.TrimSpace(strings.TrimPrefix(version, "v"))
|
||||
if normalized == "" || normalized == "dev" {
|
||||
return versionInfo{isDev: strings.EqualFold(normalized, "dev")}
|
||||
}
|
||||
base := normalized
|
||||
if separator := strings.IndexRune(normalized, '-'); separator >= 0 {
|
||||
base = normalized[:separator]
|
||||
}
|
||||
segments := strings.Split(base, ".")
|
||||
parts := make([]int, 0, len(segments))
|
||||
for _, segment := range segments {
|
||||
segment = strings.TrimSpace(segment)
|
||||
if segment == "" {
|
||||
parts = append(parts, 0)
|
||||
continue
|
||||
}
|
||||
numeric := strings.Builder{}
|
||||
for _, r := range segment {
|
||||
if r < '0' || r > '9' {
|
||||
break
|
||||
}
|
||||
numeric.WriteRune(r)
|
||||
}
|
||||
if numeric.Len() == 0 {
|
||||
parts = append(parts, 0)
|
||||
continue
|
||||
}
|
||||
value, err := strconv.Atoi(numeric.String())
|
||||
if err != nil {
|
||||
return versionInfo{}
|
||||
}
|
||||
parts = append(parts, value)
|
||||
}
|
||||
return versionInfo{valid: len(parts) > 0, numbers: parts}
|
||||
}
|
||||
|
||||
func fetchLatestGitHubRelease(ctx context.Context, repo string, channel releaseChannel) (*githubReleaseResponse, error) {
|
||||
switch normalizeReleaseChannel(string(channel)) {
|
||||
case releaseChannelPreview:
|
||||
return fetchLatestPreviewGitHubRelease(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, fmt.Errorf("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取最新版本失败: %v", err)
|
||||
}
|
||||
defer 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, fmt.Errorf("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取 preview 版本失败: %v", err)
|
||||
}
|
||||
defer 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, fmt.Errorf("解析 preview 版本信息失败")
|
||||
}
|
||||
for _, release := range releases {
|
||||
if release.Draft || !release.Prerelease {
|
||||
continue
|
||||
}
|
||||
releaseCopy := release
|
||||
return &releaseCopy, nil
|
||||
}
|
||||
return nil, fmt.Errorf("当前没有可用的 preview 发布")
|
||||
}
|
||||
|
||||
func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*githubReleaseResponse, error) {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
return nil, fmt.Errorf("缺少发布版本号")
|
||||
}
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"/tags/%s", strings.TrimSpace(repo), tag)
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取指定版本失败: %v", err)
|
||||
}
|
||||
defer 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, fmt.Errorf("解析版本信息失败")
|
||||
}
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
func setReleaseHTTPClientForTest(client *http.Client) *http.Client {
|
||||
previous := releaseHTTPClient
|
||||
if client == nil {
|
||||
releaseHTTPClient = &http.Client{Timeout: 30 * time.Second}
|
||||
} else {
|
||||
releaseHTTPClient = client
|
||||
}
|
||||
return previous
|
||||
}
|
||||
@@ -0,0 +1,405 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/observability"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// 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 := model.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 := model.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, 7000)
|
||||
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, 8080)
|
||||
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 = model.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 := model.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ipManualOverride := resolveNodeIPManualOverride(input, 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 = model.SaveOpenFlareNode(ctx, node); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildNodeView(node), nil
|
||||
}
|
||||
|
||||
// DeleteNode removes a node by id.
|
||||
func DeleteNode(ctx context.Context, id uint) error {
|
||||
if _, err := model.GetOpenFlareNodeByID(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return model.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 = model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", 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 := model.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
release, err := fetchLatestGitHubRelease(ctx, model.AgentUpdateRepo, 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 := model.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
channel := normalizeReleaseChannel(input.Channel)
|
||||
tagName := strings.TrimSpace(input.TagName)
|
||||
if tagName != "" {
|
||||
release, releaseErr := fetchGitHubReleaseByTag(ctx, model.AgentUpdateRepo, 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 = model.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 := model.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.RestartOpenrestyRequested = true
|
||||
if err = model.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 := model.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
activeConfig, err := model.GetActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, fmt.Errorf("无法获取当前激活的配置版本:%s", errNoActiveConfigVersion)
|
||||
}
|
||||
return nil, fmt.Errorf("无法获取当前激活的配置版本:%v", 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) {
|
||||
if err := option.EnsureInitialized(ctx); err != nil {
|
||||
return "", err
|
||||
}
|
||||
model.OptionMapRWMutex.RLock()
|
||||
token := strings.TrimSpace(model.AgentDiscoveryToken)
|
||||
model.OptionMapRWMutex.RUnlock()
|
||||
if token != "" {
|
||||
return token, nil
|
||||
}
|
||||
token, err := newRandomToken()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err = model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", 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 fmt.Errorf("缺少 Discovery Token")
|
||||
}
|
||||
discoveryToken, err := ensureGlobalDiscoveryToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if token != discoveryToken {
|
||||
return fmt.Errorf("Discovery Token 无效")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,325 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func 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.OpenFlareOption{},
|
||||
&model.OpenFlareApplyLog{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore())
|
||||
|
||||
return func() {
|
||||
resetAccessLogStore()
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
}
|
||||
}
|
||||
|
||||
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 := model.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 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 = model.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)
|
||||
assert.Equal(t, rotated.DiscoveryToken, model.OptionValue("AgentDiscoveryToken"))
|
||||
}
|
||||
|
||||
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"))
|
||||
}
|
||||
|
||||
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/" + model.AgentUpdateRepo + "/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))
|
||||
|
||||
offlineAt := now.Add(-model.NodeOfflineThreshold - time.Minute)
|
||||
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)
|
||||
}
|
||||
@@ -0,0 +1,324 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package node
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/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,636 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultAccessLogPageSize = 20
|
||||
maxAccessLogPageSize = 200
|
||||
defaultAccessLogSortBy = "logged_at"
|
||||
defaultAccessLogSortOrder = "desc"
|
||||
defaultAccessLogFoldMinute = 3
|
||||
defaultIPTrendHours = 24
|
||||
defaultIPTrendBucketMinute = 30
|
||||
maxIPTrendHours = 168
|
||||
nodeAccessLogRetentionDays = 90
|
||||
)
|
||||
|
||||
var nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
|
||||
|
||||
// AccessLogQuery filters access log list queries.
|
||||
type AccessLogQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
}
|
||||
|
||||
// AccessLogView is a single access log row.
|
||||
type AccessLogView struct {
|
||||
ID uint `json:"id"`
|
||||
NodeID string `json:"node_id"`
|
||||
NodeName string `json:"node_name"`
|
||||
LoggedAt time.Time `json:"logged_at"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Region string `json:"region"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
StatusCode int `json:"status_code"`
|
||||
}
|
||||
|
||||
// AccessLogList is a paginated access log response.
|
||||
type AccessLogList struct {
|
||||
Items []AccessLogView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalRecord int64 `json:"total_record"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
}
|
||||
|
||||
// FoldedAccessLogView is a folded access log bucket.
|
||||
type FoldedAccessLogView struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
UniqueIPCount int64 `json:"unique_ip_count"`
|
||||
UniqueHostCount int64 `json:"unique_host_count"`
|
||||
SuccessCount int64 `json:"success_count"`
|
||||
ClientErrorCount int64 `json:"client_error_count"`
|
||||
ServerErrorCount int64 `json:"server_error_count"`
|
||||
}
|
||||
|
||||
// FoldedAccessLogList is a paginated folded access log response.
|
||||
type FoldedAccessLogList struct {
|
||||
Items []FoldedAccessLogView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalBucket int64 `json:"total_bucket"`
|
||||
TotalRecord int64 `json:"total_record"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
}
|
||||
|
||||
// FoldedAccessLogIPQuery filters folded IP summary queries.
|
||||
type FoldedAccessLogIPQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
BucketStartedAt string `json:"bucket_started_at"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
// FoldedAccessLogIPView is a folded IP row.
|
||||
type FoldedAccessLogIPView struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
SuccessCount int64 `json:"success_count"`
|
||||
ClientErrorCount int64 `json:"client_error_count"`
|
||||
ServerErrorCount int64 `json:"server_error_count"`
|
||||
LastSeenAt time.Time `json:"last_seen_at"`
|
||||
}
|
||||
|
||||
// FoldedAccessLogIPList is a paginated folded IP response.
|
||||
type FoldedAccessLogIPList struct {
|
||||
Items []FoldedAccessLogIPView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
FoldMinutes int `json:"fold_minutes"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
// AccessLogIPSummaryQuery filters IP summary list queries.
|
||||
type AccessLogIPSummaryQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
// AccessLogIPSummaryView is an IP summary row.
|
||||
type AccessLogIPSummaryView struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
TotalRequests int64 `json:"total_requests"`
|
||||
RecentRequests int64 `json:"recent_requests"`
|
||||
LastSeenAt time.Time `json:"last_seen_at"`
|
||||
}
|
||||
|
||||
// AccessLogIPSummaryList is a paginated IP summary response.
|
||||
type AccessLogIPSummaryList struct {
|
||||
Items []AccessLogIPSummaryView `json:"items"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
HasMore bool `json:"has_more"`
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
// AccessLogIPTrendQuery filters IP trend queries.
|
||||
type AccessLogIPTrendQuery struct {
|
||||
NodeID string `json:"node_id"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Host string `json:"host"`
|
||||
Hours int `json:"hours"`
|
||||
BucketMinutes int `json:"bucket_minutes"`
|
||||
}
|
||||
|
||||
// AccessLogIPTrendPoint is an IP trend bucket.
|
||||
type AccessLogIPTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
}
|
||||
|
||||
// AccessLogIPTrendView is the IP trend response.
|
||||
type AccessLogIPTrendView struct {
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
Hours int `json:"hours"`
|
||||
BucketMinutes int `json:"bucket_minutes"`
|
||||
Points []AccessLogIPTrendPoint `json:"points"`
|
||||
}
|
||||
|
||||
// AccessLogCleanupInput is the cleanup request payload.
|
||||
type AccessLogCleanupInput struct {
|
||||
RetentionDays int `json:"retention_days"`
|
||||
}
|
||||
|
||||
// AccessLogCleanupResult is the cleanup response payload.
|
||||
type AccessLogCleanupResult struct {
|
||||
RetentionDays int `json:"retention_days"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
Cutoff time.Time `json:"cutoff"`
|
||||
}
|
||||
|
||||
// ListAccessLogs returns paginated access logs.
|
||||
func ListAccessLogs(ctx context.Context, input AccessLogQuery) (*AccessLogList, error) {
|
||||
normalized := normalizeAccessLogQuery(input)
|
||||
modelQuery := buildModelAccessLogQuery(normalized)
|
||||
logs, err := model.ListOpenFlareAccessLogs(ctx, modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalRecords, totalIPs, err := model.CountOpenFlareAccessLogs(ctx, modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodeNames, err := listNodeNameMap(ctx, logs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]AccessLogView, 0, len(logs))
|
||||
for _, item := range logs {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, AccessLogView{
|
||||
ID: item.ID,
|
||||
NodeID: item.NodeID,
|
||||
NodeName: nodeNames[item.NodeID],
|
||||
LoggedAt: item.LoggedAt,
|
||||
RemoteAddr: item.RemoteAddr,
|
||||
Region: item.Region,
|
||||
Host: item.Host,
|
||||
Path: item.Path,
|
||||
StatusCode: item.StatusCode,
|
||||
})
|
||||
}
|
||||
return &AccessLogList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalRecords,
|
||||
TotalRecord: totalRecords,
|
||||
TotalIP: totalIPs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ListFoldedAccessLogs returns paginated folded access logs.
|
||||
func ListFoldedAccessLogs(ctx context.Context, input AccessLogQuery) (*FoldedAccessLogList, error) {
|
||||
normalized := normalizeAccessLogQuery(input)
|
||||
foldMinutes, err := normalizeFoldMinutes(normalized.FoldMinutes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
modelQuery := buildModelAccessLogQuery(normalized)
|
||||
bucketQuery := model.OpenFlareAccessLogBucketQuery{
|
||||
NodeID: modelQuery.NodeID,
|
||||
RemoteAddr: modelQuery.RemoteAddr,
|
||||
Host: modelQuery.Host,
|
||||
Path: modelQuery.Path,
|
||||
Since: modelQuery.Since,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
SortBy: normalizeFoldSortBy(input.SortBy),
|
||||
SortOrder: normalized.SortOrder,
|
||||
FoldMinutes: foldMinutes,
|
||||
}
|
||||
items, err := model.ListOpenFlareAccessLogBuckets(ctx, bucketQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalBuckets, err := model.CountOpenFlareAccessLogBuckets(ctx, bucketQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalRecords, totalIPs, err := model.CountOpenFlareAccessLogs(ctx, modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]FoldedAccessLogView, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, FoldedAccessLogView{
|
||||
BucketStartedAt: time.Unix(item.BucketEpoch, 0).UTC(),
|
||||
RequestCount: item.RequestCount,
|
||||
UniqueIPCount: item.UniqueIPCount,
|
||||
UniqueHostCount: item.UniqueHostCount,
|
||||
SuccessCount: item.SuccessCount,
|
||||
ClientErrorCount: item.ClientErrorCount,
|
||||
ServerErrorCount: item.ServerErrorCount,
|
||||
})
|
||||
}
|
||||
return &FoldedAccessLogList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalBuckets,
|
||||
TotalBucket: totalBuckets,
|
||||
TotalRecord: totalRecords,
|
||||
TotalIP: totalIPs,
|
||||
FoldMinutes: foldMinutes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ListFoldedAccessLogIPs returns paginated folded IP summaries.
|
||||
func ListFoldedAccessLogIPs(ctx context.Context, input FoldedAccessLogIPQuery) (*FoldedAccessLogIPList, error) {
|
||||
normalized, bucketStartedAt, err := normalizeFoldedAccessLogIPQuery(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
modelQuery := model.OpenFlareAccessLogBucketIPQuery{
|
||||
NodeID: normalized.NodeID,
|
||||
RemoteAddr: normalized.RemoteAddr,
|
||||
Host: normalized.Host,
|
||||
Path: normalized.Path,
|
||||
BucketStartedAt: bucketStartedAt,
|
||||
FoldMinutes: normalized.FoldMinutes,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}
|
||||
items, err := model.ListOpenFlareAccessLogBucketIPs(ctx, modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalIP, err := model.CountOpenFlareAccessLogBucketIPs(ctx, modelQuery)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]FoldedAccessLogIPView, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, FoldedAccessLogIPView{
|
||||
RemoteAddr: item.RemoteAddr,
|
||||
RequestCount: item.RequestCount,
|
||||
SuccessCount: item.SuccessCount,
|
||||
ClientErrorCount: item.ClientErrorCount,
|
||||
ServerErrorCount: item.ServerErrorCount,
|
||||
LastSeenAt: time.Unix(item.LastSeenEpoch, 0).UTC(),
|
||||
})
|
||||
}
|
||||
return &FoldedAccessLogIPList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalIP,
|
||||
TotalIP: totalIP,
|
||||
BucketStartedAt: bucketStartedAt,
|
||||
FoldMinutes: normalized.FoldMinutes,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ListAccessLogIPSummaries returns paginated IP summaries.
|
||||
func ListAccessLogIPSummaries(ctx context.Context, input AccessLogIPSummaryQuery) (*AccessLogIPSummaryList, error) {
|
||||
normalized := normalizeAccessLogIPSummaryQuery(input)
|
||||
since := time.Now().UTC().Add(-nodeAccessLogRetentionWindow)
|
||||
recentSince := time.Now().UTC().Add(-3 * time.Hour)
|
||||
query := model.OpenFlareAccessLogIPSummaryQuery{
|
||||
NodeID: strings.TrimSpace(normalized.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(normalized.RemoteAddr),
|
||||
Host: strings.TrimSpace(normalized.Host),
|
||||
Since: since,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}
|
||||
items, err := model.ListOpenFlareAccessLogIPSummaries(ctx, query, recentSince)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
totalIP, err := model.CountOpenFlareAccessLogIPSummaries(ctx, query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]AccessLogIPSummaryView, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
views = append(views, AccessLogIPSummaryView{
|
||||
RemoteAddr: item.RemoteAddr,
|
||||
TotalRequests: item.TotalRequests,
|
||||
RecentRequests: item.RecentRequests,
|
||||
LastSeenAt: time.Unix(item.LastSeenEpoch, 0).UTC(),
|
||||
})
|
||||
}
|
||||
return &AccessLogIPSummaryList{
|
||||
Items: views,
|
||||
Page: normalized.Page,
|
||||
PageSize: normalized.PageSize,
|
||||
HasMore: int64((normalized.Page+1)*normalized.PageSize) < totalIP,
|
||||
TotalIP: totalIP,
|
||||
SortBy: normalized.SortBy,
|
||||
SortOrder: normalized.SortOrder,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetAccessLogIPTrend returns IP request trend points.
|
||||
func GetAccessLogIPTrend(ctx context.Context, input AccessLogIPTrendQuery) (*AccessLogIPTrendView, error) {
|
||||
normalized, err := normalizeAccessLogIPTrendQuery(input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
points, err := model.ListOpenFlareAccessLogIPTrend(ctx, model.OpenFlareAccessLogIPTrendQuery{
|
||||
NodeID: strings.TrimSpace(normalized.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(normalized.RemoteAddr),
|
||||
Host: strings.TrimSpace(normalized.Host),
|
||||
Since: time.Now().UTC().Add(-time.Duration(normalized.Hours) * time.Hour),
|
||||
BucketMinutes: normalized.BucketMinutes,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pointMap := make(map[int64]int64, len(points))
|
||||
for _, item := range points {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
pointMap[item.BucketEpoch] = item.RequestCount
|
||||
}
|
||||
bucketDuration := time.Duration(normalized.BucketMinutes) * time.Minute
|
||||
start := time.Now().UTC().Add(-time.Duration(normalized.Hours) * time.Hour).Truncate(bucketDuration)
|
||||
end := time.Now().UTC().Truncate(bucketDuration)
|
||||
views := make([]AccessLogIPTrendPoint, 0, int(end.Sub(start)/bucketDuration)+1)
|
||||
for cursor := start; !cursor.After(end); cursor = cursor.Add(bucketDuration) {
|
||||
views = append(views, AccessLogIPTrendPoint{
|
||||
BucketStartedAt: cursor,
|
||||
RequestCount: pointMap[cursor.Unix()],
|
||||
})
|
||||
}
|
||||
return &AccessLogIPTrendView{
|
||||
RemoteAddr: normalized.RemoteAddr,
|
||||
Hours: normalized.Hours,
|
||||
BucketMinutes: normalized.BucketMinutes,
|
||||
Points: views,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// CleanupAccessLogs removes access logs older than retention days.
|
||||
func CleanupAccessLogs(ctx context.Context, input AccessLogCleanupInput) (*AccessLogCleanupResult, error) {
|
||||
if input.RetentionDays <= 0 || input.RetentionDays > nodeAccessLogRetentionDays {
|
||||
return nil, errors.New("retention_days 必须在 1 到 90 之间")
|
||||
}
|
||||
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour)
|
||||
deleted, err := model.DeleteOpenFlareAccessLogsBefore(ctx, cutoff)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &AccessLogCleanupResult{
|
||||
RetentionDays: input.RetentionDays,
|
||||
DeletedCount: deleted,
|
||||
Cutoff: cutoff,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildModelAccessLogQuery(input AccessLogQuery) model.OpenFlareAccessLogQuery {
|
||||
return model.OpenFlareAccessLogQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
Since: time.Now().UTC().Add(-nodeAccessLogRetentionWindow),
|
||||
Page: input.Page,
|
||||
PageSize: input.PageSize,
|
||||
SortBy: input.SortBy,
|
||||
SortOrder: input.SortOrder,
|
||||
}
|
||||
}
|
||||
|
||||
func listNodeNameMap(ctx context.Context, logs []*model.OpenFlareAccessLog) (map[string]string, error) {
|
||||
nodeIDs := make([]string, 0, len(logs))
|
||||
seen := make(map[string]struct{}, len(logs))
|
||||
for _, item := range logs {
|
||||
if item == nil || item.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[item.NodeID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[item.NodeID] = struct{}{}
|
||||
nodeIDs = append(nodeIDs, item.NodeID)
|
||||
}
|
||||
if len(nodeIDs) == 0 {
|
||||
return map[string]string{}, nil
|
||||
}
|
||||
nodes, err := model.ListOpenFlareNodesByNodeIDs(ctx, nodeIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make(map[string]string, len(nodes))
|
||||
for _, node := range nodes {
|
||||
result[node.NodeID] = node.Name
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func normalizeAccessLogQuery(input AccessLogQuery) AccessLogQuery {
|
||||
return AccessLogQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
Page: normalizeAccessLogPage(input.Page),
|
||||
PageSize: normalizeAccessLogPageSize(input.PageSize),
|
||||
SortBy: normalizeAccessLogSortBy(input.SortBy),
|
||||
SortOrder: normalizeAccessLogSortOrder(input.SortOrder),
|
||||
FoldMinutes: input.FoldMinutes,
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeAccessLogIPSummaryQuery(input AccessLogIPSummaryQuery) AccessLogIPSummaryQuery {
|
||||
return AccessLogIPSummaryQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Page: normalizeAccessLogPage(input.Page),
|
||||
PageSize: normalizeAccessLogPageSize(input.PageSize),
|
||||
SortBy: normalizeIPSummarySortBy(input.SortBy),
|
||||
SortOrder: normalizeAccessLogSortOrder(input.SortOrder),
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeFoldedAccessLogIPQuery(input FoldedAccessLogIPQuery) (FoldedAccessLogIPQuery, time.Time, error) {
|
||||
foldMinutes, err := normalizeFoldMinutes(input.FoldMinutes)
|
||||
if err != nil {
|
||||
return FoldedAccessLogIPQuery{}, time.Time{}, err
|
||||
}
|
||||
bucketStartedAt, err := time.Parse(time.RFC3339, strings.TrimSpace(input.BucketStartedAt))
|
||||
if err != nil {
|
||||
return FoldedAccessLogIPQuery{}, time.Time{}, errors.New("bucket_started_at 必须为 RFC3339 时间")
|
||||
}
|
||||
normalizedSortBy := strings.TrimSpace(input.SortBy)
|
||||
switch normalizedSortBy {
|
||||
case "last_seen_at", "remote_addr":
|
||||
default:
|
||||
normalizedSortBy = "request_count"
|
||||
}
|
||||
return FoldedAccessLogIPQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: strings.TrimSpace(input.RemoteAddr),
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
BucketStartedAt: strings.TrimSpace(input.BucketStartedAt),
|
||||
FoldMinutes: foldMinutes,
|
||||
Page: normalizeAccessLogPage(input.Page),
|
||||
PageSize: normalizeAccessLogPageSize(input.PageSize),
|
||||
SortBy: normalizedSortBy,
|
||||
SortOrder: normalizeAccessLogSortOrder(input.SortOrder),
|
||||
}, bucketStartedAt.UTC(), nil
|
||||
}
|
||||
|
||||
func normalizeAccessLogIPTrendQuery(input AccessLogIPTrendQuery) (AccessLogIPTrendQuery, error) {
|
||||
remoteAddr := strings.TrimSpace(input.RemoteAddr)
|
||||
if remoteAddr == "" {
|
||||
return AccessLogIPTrendQuery{}, errors.New("remote_addr 不能为空")
|
||||
}
|
||||
hours := input.Hours
|
||||
if hours <= 0 {
|
||||
hours = defaultIPTrendHours
|
||||
}
|
||||
if hours > maxIPTrendHours {
|
||||
hours = maxIPTrendHours
|
||||
}
|
||||
bucketMinutes := input.BucketMinutes
|
||||
if bucketMinutes <= 0 {
|
||||
bucketMinutes = defaultIPTrendBucketMinute
|
||||
}
|
||||
switch bucketMinutes {
|
||||
case 5, 10, 15, 30, 60:
|
||||
default:
|
||||
return AccessLogIPTrendQuery{}, errors.New("bucket_minutes 仅支持 5、10、15、30、60")
|
||||
}
|
||||
return AccessLogIPTrendQuery{
|
||||
NodeID: strings.TrimSpace(input.NodeID),
|
||||
RemoteAddr: remoteAddr,
|
||||
Host: strings.TrimSpace(input.Host),
|
||||
Hours: hours,
|
||||
BucketMinutes: bucketMinutes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeAccessLogPage(page int) int {
|
||||
if page < 0 {
|
||||
return 0
|
||||
}
|
||||
return page
|
||||
}
|
||||
|
||||
func normalizeAccessLogPageSize(pageSize int) int {
|
||||
if pageSize <= 0 {
|
||||
return defaultAccessLogPageSize
|
||||
}
|
||||
if pageSize > maxAccessLogPageSize {
|
||||
return maxAccessLogPageSize
|
||||
}
|
||||
return pageSize
|
||||
}
|
||||
|
||||
func normalizeAccessLogSortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "status_code", "remote_addr", "host", "path":
|
||||
return strings.TrimSpace(sortBy)
|
||||
default:
|
||||
return defaultAccessLogSortBy
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeAccessLogSortOrder(sortOrder string) string {
|
||||
if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") {
|
||||
return "asc"
|
||||
}
|
||||
return defaultAccessLogSortOrder
|
||||
}
|
||||
|
||||
func normalizeFoldSortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "request_count":
|
||||
return "request_count"
|
||||
default:
|
||||
return "bucket_started_at"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeIPSummarySortBy(sortBy string) string {
|
||||
switch strings.TrimSpace(sortBy) {
|
||||
case "recent_requests", "last_seen_at", "remote_addr":
|
||||
return strings.TrimSpace(sortBy)
|
||||
default:
|
||||
return "total_requests"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeFoldMinutes(value int) (int, error) {
|
||||
if value <= 0 {
|
||||
return defaultAccessLogFoldMinute, nil
|
||||
}
|
||||
switch value {
|
||||
case 3, 5:
|
||||
return value, nil
|
||||
default:
|
||||
return 0, errors.New("fold_minutes 仅支持 3 或 5")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,482 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const observabilityTrendBuckets = 24
|
||||
|
||||
const (
|
||||
healthEventStatusActive = "active"
|
||||
healthEventStatusResolved = "resolved"
|
||||
healthSeverityCritical = "critical"
|
||||
healthSeverityWarning = "warning"
|
||||
)
|
||||
|
||||
// DistributionItem is a key/value distribution entry.
|
||||
type DistributionItem struct {
|
||||
Key string `json:"key"`
|
||||
Value int64 `json:"value"`
|
||||
}
|
||||
|
||||
// TrafficDistributions groups traffic distribution charts.
|
||||
type TrafficDistributions struct {
|
||||
StatusCodes []DistributionItem `json:"status_codes"`
|
||||
TopDomains []DistributionItem `json:"top_domains"`
|
||||
SourceCountries []DistributionItem `json:"source_countries"`
|
||||
}
|
||||
|
||||
// TrafficWindowSummary summarizes a traffic reporting window.
|
||||
type TrafficWindowSummary struct {
|
||||
WindowStartedAt time.Time `json:"window_started_at"`
|
||||
WindowEndedAt time.Time `json:"window_ended_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
EstimatedQPS float64 `json:"estimated_qps"`
|
||||
ErrorRatePercent float64 `json:"error_rate_percent"`
|
||||
}
|
||||
|
||||
// HealthSummary summarizes node health alerts and risks.
|
||||
type HealthSummary struct {
|
||||
ActiveAlerts int `json:"active_alerts"`
|
||||
CriticalAlerts int `json:"critical_alerts"`
|
||||
WarningAlerts int `json:"warning_alerts"`
|
||||
InfoAlerts int `json:"info_alerts"`
|
||||
ResolvedAlerts int `json:"resolved_alerts"`
|
||||
HasCapacityRisk bool `json:"has_capacity_risk"`
|
||||
HasTrafficRisk bool `json:"has_traffic_risk"`
|
||||
HasRuntimeRisk bool `json:"has_runtime_risk"`
|
||||
}
|
||||
|
||||
// TrafficTrendPoint is a traffic trend bucket.
|
||||
type TrafficTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
RequestCount int64 `json:"request_count"`
|
||||
ErrorCount int64 `json:"error_count"`
|
||||
UniqueVisitorCount int64 `json:"unique_visitor_count"`
|
||||
}
|
||||
|
||||
// CapacityTrendPoint is a capacity trend bucket.
|
||||
type CapacityTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
AverageCPUUsagePercent float64 `json:"average_cpu_usage_percent"`
|
||||
AverageMemoryUsagePercent float64 `json:"average_memory_usage_percent"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
// NetworkTrendPoint is a network trend bucket.
|
||||
type NetworkTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
NetworkRxBytes int64 `json:"network_rx_bytes"`
|
||||
NetworkTxBytes int64 `json:"network_tx_bytes"`
|
||||
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
|
||||
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
// DiskIOTrendPoint is a disk IO trend bucket.
|
||||
type DiskIOTrendPoint struct {
|
||||
BucketStartedAt time.Time `json:"bucket_started_at"`
|
||||
DiskReadBytes int64 `json:"disk_read_bytes"`
|
||||
DiskWriteBytes int64 `json:"disk_write_bytes"`
|
||||
ReportedNodes int `json:"reported_nodes"`
|
||||
}
|
||||
|
||||
type distributionAccumulator map[string]int64
|
||||
|
||||
type capacityTrendAccumulator struct {
|
||||
cpuSum float64
|
||||
cpuCount int
|
||||
memSum float64
|
||||
memCount int
|
||||
nodes map[string]struct{}
|
||||
}
|
||||
|
||||
type snapshotTrendAccumulator struct {
|
||||
nodes map[string]struct{}
|
||||
}
|
||||
|
||||
type diskCounterState struct {
|
||||
read int64
|
||||
write int64
|
||||
seen bool
|
||||
}
|
||||
|
||||
func buildTrafficWindowSummary(report *model.OpenFlareRequestReport) TrafficWindowSummary {
|
||||
if report == nil {
|
||||
return TrafficWindowSummary{}
|
||||
}
|
||||
summary := TrafficWindowSummary{
|
||||
WindowStartedAt: report.WindowStartedAt,
|
||||
WindowEndedAt: report.WindowEndedAt,
|
||||
RequestCount: report.RequestCount,
|
||||
UniqueVisitorCount: report.UniqueVisitorCount,
|
||||
ErrorCount: report.ErrorCount,
|
||||
}
|
||||
if duration := report.WindowEndedAt.Sub(report.WindowStartedAt).Seconds(); duration > 0 {
|
||||
summary.EstimatedQPS = float64(report.RequestCount) / duration
|
||||
}
|
||||
if report.RequestCount > 0 {
|
||||
summary.ErrorRatePercent = (float64(report.ErrorCount) / float64(report.RequestCount)) * 100
|
||||
}
|
||||
return summary
|
||||
}
|
||||
|
||||
// BuildTrafficDistributions aggregates traffic distribution charts.
|
||||
func BuildTrafficDistributions(
|
||||
reports []*model.OpenFlareRequestReport,
|
||||
accessLogRegions []*model.OpenFlareAccessLogRegionCount,
|
||||
limit int,
|
||||
) TrafficDistributions {
|
||||
statusCodes := make(distributionAccumulator)
|
||||
topDomains := make(distributionAccumulator)
|
||||
reportSourceCountries := make(distributionAccumulator)
|
||||
for _, report := range reports {
|
||||
mergeJSONCounts(statusCodes, report.StatusCodesJSON)
|
||||
mergeJSONCounts(topDomains, report.TopDomainsJSON)
|
||||
mergeJSONCounts(reportSourceCountries, report.SourceCountriesJSON)
|
||||
}
|
||||
sourceCountries := reportSourceCountries
|
||||
if len(accessLogRegions) > 0 {
|
||||
sourceCountries = make(distributionAccumulator, len(accessLogRegions))
|
||||
for _, item := range accessLogRegions {
|
||||
if item == nil || strings.TrimSpace(item.Region) == "" || item.Count <= 0 {
|
||||
continue
|
||||
}
|
||||
sourceCountries[item.Region] = item.Count
|
||||
}
|
||||
}
|
||||
return TrafficDistributions{
|
||||
StatusCodes: toDistributionItems(statusCodes, limit),
|
||||
TopDomains: toDistributionItems(topDomains, limit),
|
||||
SourceCountries: toDistributionItems(sourceCountries, limit),
|
||||
}
|
||||
}
|
||||
|
||||
func buildHealthSummary(
|
||||
snapshot *model.OpenFlareMetricSnapshot,
|
||||
report *model.OpenFlareRequestReport,
|
||||
events []*model.OpenFlareHealthEvent,
|
||||
) HealthSummary {
|
||||
summary := HealthSummary{}
|
||||
for _, event := range events {
|
||||
if event == nil {
|
||||
continue
|
||||
}
|
||||
if event.Status == healthEventStatusResolved {
|
||||
summary.ResolvedAlerts++
|
||||
continue
|
||||
}
|
||||
summary.ActiveAlerts++
|
||||
switch event.Severity {
|
||||
case healthSeverityCritical:
|
||||
summary.CriticalAlerts++
|
||||
case healthSeverityWarning:
|
||||
summary.WarningAlerts++
|
||||
default:
|
||||
summary.InfoAlerts++
|
||||
}
|
||||
}
|
||||
if snapshot != nil {
|
||||
memoryUsage := Percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes)
|
||||
storageUsage := Percentage(snapshot.StorageUsedBytes, snapshot.StorageTotalBytes)
|
||||
summary.HasCapacityRisk = snapshot.CPUUsagePercent >= 80 || memoryUsage >= 85 || storageUsage >= 85
|
||||
}
|
||||
if report != nil && report.RequestCount >= 100 {
|
||||
summary.HasTrafficRisk = (float64(report.ErrorCount) / float64(report.RequestCount)) >= 0.05
|
||||
}
|
||||
summary.HasRuntimeRisk = summary.ActiveAlerts > 0 || summary.HasCapacityRisk || summary.HasTrafficRisk
|
||||
return summary
|
||||
}
|
||||
|
||||
// BuildTrafficTrendPoints builds 24h traffic trend buckets.
|
||||
func BuildTrafficTrendPoints(now time.Time, reports []*model.OpenFlareRequestReport) []TrafficTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]TrafficTrendPoint, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
}
|
||||
for _, report := range reports {
|
||||
index, ok := trendBucketIndex(report.WindowEndedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
points[index].RequestCount += report.RequestCount
|
||||
points[index].ErrorCount += report.ErrorCount
|
||||
points[index].UniqueVisitorCount += report.UniqueVisitorCount
|
||||
}
|
||||
return points
|
||||
}
|
||||
|
||||
// BuildCapacityTrendPoints builds 24h capacity trend buckets.
|
||||
func BuildCapacityTrendPoints(now time.Time, snapshots []*model.OpenFlareMetricSnapshot) []CapacityTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]CapacityTrendPoint, observabilityTrendBuckets)
|
||||
accumulators := make([]capacityTrendAccumulator, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
accumulators[index].nodes = make(map[string]struct{})
|
||||
}
|
||||
for _, snapshot := range snapshots {
|
||||
index, ok := trendBucketIndex(snapshot.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if snapshot.CPUUsagePercent > 0 {
|
||||
accumulators[index].cpuSum += snapshot.CPUUsagePercent
|
||||
accumulators[index].cpuCount++
|
||||
}
|
||||
if memoryUsage := Percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes); memoryUsage > 0 {
|
||||
accumulators[index].memSum += memoryUsage
|
||||
accumulators[index].memCount++
|
||||
}
|
||||
if snapshot.NodeID != "" {
|
||||
accumulators[index].nodes[snapshot.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
for index := range points {
|
||||
if accumulators[index].cpuCount > 0 {
|
||||
points[index].AverageCPUUsagePercent = accumulators[index].cpuSum / float64(accumulators[index].cpuCount)
|
||||
}
|
||||
if accumulators[index].memCount > 0 {
|
||||
points[index].AverageMemoryUsagePercent = accumulators[index].memSum / float64(accumulators[index].memCount)
|
||||
}
|
||||
points[index].ReportedNodes = len(accumulators[index].nodes)
|
||||
}
|
||||
return points
|
||||
}
|
||||
|
||||
// BuildNetworkTrendPoints builds 24h network trend buckets.
|
||||
func BuildNetworkTrendPoints(
|
||||
now time.Time,
|
||||
snapshots []*model.OpenFlareMetricSnapshot,
|
||||
openrestyObs []*model.OpenFlareNodeObservationOpenresty,
|
||||
) []NetworkTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]NetworkTrendPoint, observabilityTrendBuckets)
|
||||
accumulators := make([]snapshotTrendAccumulator, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
accumulators[index].nodes = make(map[string]struct{})
|
||||
}
|
||||
for _, snapshot := range snapshots {
|
||||
index, ok := trendBucketIndex(snapshot.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
points[index].NetworkRxBytes += snapshot.NetworkRxBytes
|
||||
points[index].NetworkTxBytes += snapshot.NetworkTxBytes
|
||||
if snapshot.NodeID != "" {
|
||||
accumulators[index].nodes[snapshot.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
for _, obs := range openrestyObs {
|
||||
index, ok := trendBucketIndex(obs.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
points[index].OpenrestyRxBytes += obs.OpenrestyRxBytes
|
||||
points[index].OpenrestyTxBytes += obs.OpenrestyTxBytes
|
||||
if obs.NodeID != "" {
|
||||
accumulators[index].nodes[obs.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
for index := range points {
|
||||
points[index].ReportedNodes = len(accumulators[index].nodes)
|
||||
}
|
||||
return points
|
||||
}
|
||||
|
||||
// BuildDiskIOTrendPoints builds 24h disk IO trend buckets.
|
||||
func BuildDiskIOTrendPoints(now time.Time, snapshots []*model.OpenFlareMetricSnapshot) []DiskIOTrendPoint {
|
||||
start := trendWindowStart(now)
|
||||
points := make([]DiskIOTrendPoint, observabilityTrendBuckets)
|
||||
accumulators := make([]snapshotTrendAccumulator, observabilityTrendBuckets)
|
||||
for index := range points {
|
||||
points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour)
|
||||
accumulators[index].nodes = make(map[string]struct{})
|
||||
}
|
||||
sort.Slice(snapshots, func(i int, j int) bool {
|
||||
if snapshots[i].CapturedAt.Equal(snapshots[j].CapturedAt) {
|
||||
return snapshots[i].NodeID < snapshots[j].NodeID
|
||||
}
|
||||
return snapshots[i].CapturedAt.Before(snapshots[j].CapturedAt)
|
||||
})
|
||||
previousByNode := make(map[string]diskCounterState, len(snapshots))
|
||||
for _, snapshot := range snapshots {
|
||||
nodeKey := snapshot.NodeID
|
||||
if nodeKey == "" {
|
||||
nodeKey = "__unknown__"
|
||||
}
|
||||
previous := previousByNode[nodeKey]
|
||||
previousByNode[nodeKey] = diskCounterState{
|
||||
read: snapshot.DiskReadBytes,
|
||||
write: snapshot.DiskWriteBytes,
|
||||
seen: true,
|
||||
}
|
||||
if !previous.seen {
|
||||
continue
|
||||
}
|
||||
index, ok := trendBucketIndex(snapshot.CapturedAt, start)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
readDelta := snapshot.DiskReadBytes - previous.read
|
||||
writeDelta := snapshot.DiskWriteBytes - previous.write
|
||||
if readDelta < 0 {
|
||||
readDelta = 0
|
||||
}
|
||||
if writeDelta < 0 {
|
||||
writeDelta = 0
|
||||
}
|
||||
points[index].DiskReadBytes += readDelta
|
||||
points[index].DiskWriteBytes += writeDelta
|
||||
if snapshot.NodeID != "" {
|
||||
accumulators[index].nodes[snapshot.NodeID] = struct{}{}
|
||||
}
|
||||
}
|
||||
for index := range points {
|
||||
points[index].ReportedNodes = len(accumulators[index].nodes)
|
||||
}
|
||||
return points
|
||||
}
|
||||
|
||||
func latestMetricSnapshot(snapshots []*model.OpenFlareMetricSnapshot) *model.OpenFlareMetricSnapshot {
|
||||
for _, snapshot := range snapshots {
|
||||
if snapshot != nil {
|
||||
return snapshot
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func latestTrafficReport(reports []*model.OpenFlareRequestReport) *model.OpenFlareRequestReport {
|
||||
for _, report := range reports {
|
||||
if report != nil {
|
||||
return report
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// LatestMetricSnapshotsByNode returns the latest snapshot per node.
|
||||
func LatestMetricSnapshotsByNode(snapshots []*model.OpenFlareMetricSnapshot) map[string]*model.OpenFlareMetricSnapshot {
|
||||
result := make(map[string]*model.OpenFlareMetricSnapshot, len(snapshots))
|
||||
for _, snapshot := range snapshots {
|
||||
if snapshot == nil || snapshot.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
if existing, ok := result[snapshot.NodeID]; ok && !snapshot.CapturedAt.After(existing.CapturedAt) {
|
||||
continue
|
||||
}
|
||||
result[snapshot.NodeID] = snapshot
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// LatestTrafficReportsByNode returns the latest traffic report per node.
|
||||
func LatestTrafficReportsByNode(reports []*model.OpenFlareRequestReport) map[string]*model.OpenFlareRequestReport {
|
||||
result := make(map[string]*model.OpenFlareRequestReport, len(reports))
|
||||
for _, report := range reports {
|
||||
if report == nil || report.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
if existing, ok := result[report.NodeID]; ok && !report.WindowEndedAt.After(existing.WindowEndedAt) {
|
||||
continue
|
||||
}
|
||||
result[report.NodeID] = report
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// ActiveHealthEventsByNode groups active health events by node id.
|
||||
func ActiveHealthEventsByNode(events []*model.OpenFlareHealthEvent) map[string][]*model.OpenFlareHealthEvent {
|
||||
result := make(map[string][]*model.OpenFlareHealthEvent)
|
||||
for _, event := range events {
|
||||
if event == nil || event.NodeID == "" {
|
||||
continue
|
||||
}
|
||||
result[event.NodeID] = append(result[event.NodeID], event)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// Percentage returns used/total as a percentage.
|
||||
func Percentage(used int64, total int64) float64 {
|
||||
if used <= 0 || total <= 0 {
|
||||
return 0
|
||||
}
|
||||
return (float64(used) / float64(total)) * 100
|
||||
}
|
||||
|
||||
func mergeJSONCounts(target distributionAccumulator, raw string) {
|
||||
if len(target) == 0 && strings.TrimSpace(raw) == "" {
|
||||
return
|
||||
}
|
||||
values := parseJSONCounts(raw)
|
||||
for key, value := range values {
|
||||
if strings.TrimSpace(key) == "" || value <= 0 {
|
||||
continue
|
||||
}
|
||||
target[key] += value
|
||||
}
|
||||
}
|
||||
|
||||
func parseJSONCounts(raw string) map[string]int64 {
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
return nil
|
||||
}
|
||||
values := make(map[string]int64)
|
||||
if err := json.Unmarshal([]byte(raw), &values); err != nil {
|
||||
return nil
|
||||
}
|
||||
return values
|
||||
}
|
||||
|
||||
func toDistributionItems(values distributionAccumulator, limit int) []DistributionItem {
|
||||
if len(values) == 0 {
|
||||
return []DistributionItem{}
|
||||
}
|
||||
items := make([]DistributionItem, 0, len(values))
|
||||
for key, value := range values {
|
||||
if strings.TrimSpace(key) == "" || value <= 0 {
|
||||
continue
|
||||
}
|
||||
items = append(items, DistributionItem{Key: key, Value: value})
|
||||
}
|
||||
sort.Slice(items, func(i int, j int) bool {
|
||||
if items[i].Value == items[j].Value {
|
||||
return items[i].Key < items[j].Key
|
||||
}
|
||||
return items[i].Value > items[j].Value
|
||||
})
|
||||
if limit > 0 && len(items) > limit {
|
||||
items = items[:limit]
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func trendWindowStart(now time.Time) time.Time {
|
||||
return now.Truncate(time.Hour).Add(-(observabilityTrendBuckets - 1) * time.Hour)
|
||||
}
|
||||
|
||||
func trendBucketIndex(timestamp time.Time, start time.Time) (int, bool) {
|
||||
if timestamp.Before(start) {
|
||||
return 0, false
|
||||
}
|
||||
delta := timestamp.Sub(start)
|
||||
index := int(delta / time.Hour)
|
||||
if index < 0 || index >= observabilityTrendBuckets {
|
||||
return 0, false
|
||||
}
|
||||
return index, true
|
||||
}
|
||||
@@ -0,0 +1,246 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultObservabilityWindow = 24 * time.Hour
|
||||
defaultObservabilityLimit = 120
|
||||
maxObservabilityLimit = 500
|
||||
)
|
||||
|
||||
// NodeQuery filters node observability data.
|
||||
type NodeQuery struct {
|
||||
Hours int `json:"hours"`
|
||||
Limit int `json:"limit"`
|
||||
}
|
||||
|
||||
// NodeAnalytics groups node observability analytics.
|
||||
type NodeAnalytics struct {
|
||||
Traffic TrafficWindowSummary `json:"traffic"`
|
||||
Distributions TrafficDistributions `json:"distributions"`
|
||||
Health HealthSummary `json:"health"`
|
||||
}
|
||||
|
||||
// NodeTrends groups node observability trend series.
|
||||
type NodeTrends struct {
|
||||
Traffic24h []TrafficTrendPoint `json:"traffic_24h"`
|
||||
Capacity24h []CapacityTrendPoint `json:"capacity_24h"`
|
||||
Network24h []NetworkTrendPoint `json:"network_24h"`
|
||||
DiskIO24h []DiskIOTrendPoint `json:"disk_io_24h"`
|
||||
}
|
||||
|
||||
// RelayDashboardSnapshot summarizes tunnel relay status.
|
||||
type RelayDashboardSnapshot struct {
|
||||
TotalProxies int `json:"total_proxies"`
|
||||
OnlineProxies int `json:"online_proxies"`
|
||||
OfflineProxies int `json:"offline_proxies"`
|
||||
Proxies []RelayProxyStat `json:"proxies"`
|
||||
TotalConnections int `json:"total_connections"`
|
||||
ClientCounts int `json:"client_counts"`
|
||||
}
|
||||
|
||||
// RelayProxyStat is a single relay proxy entry.
|
||||
type RelayProxyStat struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Status string `json:"status"`
|
||||
ClientVersion string `json:"client_version"`
|
||||
LastStartTime string `json:"last_start_time"`
|
||||
LastCloseTime string `json:"last_close_time"`
|
||||
ClientAddr string `json:"client_addr"`
|
||||
}
|
||||
|
||||
// NodeView is the node observability API response.
|
||||
type NodeView struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Profile *model.OpenFlareNodeSystemProfile `json:"profile"`
|
||||
MetricSnapshots []*model.OpenFlareMetricSnapshot `json:"metric_snapshots"`
|
||||
TrafficReports []*model.OpenFlareRequestReport `json:"traffic_reports"`
|
||||
HealthEvents []*model.OpenFlareHealthEvent `json:"health_events"`
|
||||
Analytics NodeAnalytics `json:"analytics"`
|
||||
Trends NodeTrends `json:"trends"`
|
||||
RelayDashboard *RelayDashboardSnapshot `json:"relay_dashboard,omitempty"`
|
||||
}
|
||||
|
||||
// HealthEventCleanupResult reports health event cleanup outcome.
|
||||
type HealthEventCleanupResult struct {
|
||||
NodeID string `json:"node_id"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
}
|
||||
|
||||
// GetNodeObservability returns observability details for a node.
|
||||
func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeView, error) {
|
||||
now := time.Now()
|
||||
node, err := model.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
limit := normalizeObservabilityLimit(query.Limit)
|
||||
since := now.Add(-normalizeObservabilityWindow(query.Hours))
|
||||
|
||||
profile, err := model.GetOpenFlareNodeSystemProfile(ctx, node.NodeID)
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
profile = nil
|
||||
}
|
||||
|
||||
snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, since, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reports, err := model.ListOpenFlareRequestReportsSince(ctx, node.NodeID, since, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, 8)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trendSnapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-24*time.Hour), 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trendOpenresty, err := model.ListOpenFlareNodeObservationOpenresty(ctx, node.NodeID, now.Add(-24*time.Hour), 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trendReports, err := model.ListOpenFlareRequestReportsSince(ctx, node.NodeID, now.Add(-24*time.Hour), 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
view := &NodeView{
|
||||
NodeID: node.NodeID,
|
||||
Profile: profile,
|
||||
MetricSnapshots: snapshots,
|
||||
TrafficReports: reports,
|
||||
HealthEvents: events,
|
||||
Analytics: NodeAnalytics{
|
||||
Traffic: buildTrafficWindowSummary(latestTrafficReport(reports)),
|
||||
Distributions: BuildTrafficDistributions(reports, accessLogRegions, 8),
|
||||
Health: buildHealthSummary(latestMetricSnapshot(snapshots), latestTrafficReport(reports), events),
|
||||
},
|
||||
Trends: NodeTrends{
|
||||
Traffic24h: BuildTrafficTrendPoints(now, trendReports),
|
||||
Capacity24h: BuildCapacityTrendPoints(now, trendSnapshots),
|
||||
Network24h: BuildNetworkTrendPoints(now, trendSnapshots, trendOpenresty),
|
||||
DiskIO24h: BuildDiskIOTrendPoints(now, trendSnapshots),
|
||||
},
|
||||
}
|
||||
if node.NodeType == "tunnel_relay" {
|
||||
frpsObs, frpsErr := model.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1)
|
||||
if frpsErr != nil {
|
||||
return nil, frpsErr
|
||||
}
|
||||
var latestFrps *model.OpenFlareNodeObservationFrps
|
||||
if len(frpsObs) > 0 {
|
||||
latestFrps = frpsObs[0]
|
||||
}
|
||||
view.RelayDashboard = buildRelayDashboardSnapshot(node, latestFrps)
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
// CleanupHealthEvents removes all health events for a node.
|
||||
func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) {
|
||||
node, err := model.GetOpenFlareNodeByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deletedCount, err := model.DeleteOpenFlareHealthEventsByNodeID(ctx, node.NodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &HealthEventCleanupResult{
|
||||
NodeID: node.NodeID,
|
||||
DeletedCount: deletedCount,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildRelayDashboardSnapshot(node *model.OpenFlareNode, obs *model.OpenFlareNodeObservationFrps) *RelayDashboardSnapshot {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
totalProxies := 0
|
||||
totalConnections := 0
|
||||
clientCounts := 0
|
||||
proxies := []RelayProxyStat{}
|
||||
|
||||
if obs != nil {
|
||||
totalProxies = obs.FrpsProxyCount
|
||||
totalConnections = obs.FrpsConnections
|
||||
clientCounts = obs.FrpsClientCount
|
||||
if obs.FrpsProxies != "" {
|
||||
var decoded []RelayProxyStat
|
||||
if err := json.Unmarshal([]byte(obs.FrpsProxies), &decoded); err == nil {
|
||||
proxies = decoded
|
||||
}
|
||||
}
|
||||
}
|
||||
if totalProxies < 0 {
|
||||
totalProxies = 0
|
||||
}
|
||||
onlineProxies := 0
|
||||
for _, proxy := range proxies {
|
||||
if proxy.Status == "online" {
|
||||
onlineProxies++
|
||||
}
|
||||
}
|
||||
if len(proxies) == 0 {
|
||||
onlineProxies = totalProxies
|
||||
if node.RelayStatus != "healthy" {
|
||||
onlineProxies = 0
|
||||
}
|
||||
}
|
||||
|
||||
return &RelayDashboardSnapshot{
|
||||
TotalProxies: totalProxies,
|
||||
OnlineProxies: onlineProxies,
|
||||
OfflineProxies: totalProxies - onlineProxies,
|
||||
Proxies: proxies,
|
||||
TotalConnections: maxInt(totalConnections, 0),
|
||||
ClientCounts: maxInt(clientCounts, 0),
|
||||
}
|
||||
}
|
||||
|
||||
func maxInt(a int, b int) int {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func normalizeObservabilityLimit(limit int) int {
|
||||
if limit <= 0 {
|
||||
return defaultObservabilityLimit
|
||||
}
|
||||
if limit > maxObservabilityLimit {
|
||||
return maxObservabilityLimit
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func normalizeObservabilityWindow(hours int) time.Duration {
|
||||
if hours <= 0 {
|
||||
return defaultObservabilityWindow
|
||||
}
|
||||
return time.Duration(hours) * time.Hour
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GetAccessLogsHandler 分页列出访问日志。
|
||||
// @Summary 列出访问日志
|
||||
// @Description 分页返回 OpenFlare 访问日志,支持按节点、IP、主机与路径筛选,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param node_id query string false "节点 ID"
|
||||
// @Param remote_addr query string false "客户端 IP"
|
||||
// @Param host query string false "请求 Host"
|
||||
// @Param path query string false "请求路径"
|
||||
// @Param p query int false "页码"
|
||||
// @Param page_size query int false "每页条数"
|
||||
// @Param sort_by query string false "排序字段"
|
||||
// @Param sort_order query string false "排序方向"
|
||||
// @Success 200 {object} response.Any{data=observability.AccessLogList} "访问日志列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/access-logs [get]
|
||||
func GetAccessLogsHandler(c *gin.Context) {
|
||||
logs, err := ListAccessLogs(c.Request.Context(), readAccessLogQuery(c))
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(logs))
|
||||
}
|
||||
|
||||
// getFoldedAccessLogsHandler 分页列出折叠访问日志。
|
||||
// @Summary 列出折叠访问日志
|
||||
// @Description 按时间桶聚合访问日志并分页返回,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param node_id query string false "节点 ID"
|
||||
// @Param remote_addr query string false "客户端 IP"
|
||||
// @Param host query string false "请求 Host"
|
||||
// @Param path query string false "请求路径"
|
||||
// @Param fold_minutes query int false "折叠时间窗口(分钟)"
|
||||
// @Param p query int false "页码"
|
||||
// @Param page_size query int false "每页条数"
|
||||
// @Param sort_by query string false "排序字段"
|
||||
// @Param sort_order query string false "排序方向"
|
||||
// @Success 200 {object} response.Any{data=observability.FoldedAccessLogList} "折叠访问日志列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/access-logs/folds [get]
|
||||
func GetFoldedAccessLogsHandler(c *gin.Context) {
|
||||
query := readAccessLogQuery(c)
|
||||
query.FoldMinutes = readQueryInt(c, "fold_minutes")
|
||||
logs, err := ListFoldedAccessLogs(c.Request.Context(), query)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(logs))
|
||||
}
|
||||
|
||||
// getFoldedAccessLogIPsHandler 列出折叠桶内的 IP 汇总。
|
||||
// @Summary 列出折叠访问日志 IP 汇总
|
||||
// @Description 在指定时间桶内按 IP 聚合访问统计,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param node_id query string false "节点 ID"
|
||||
// @Param remote_addr query string false "客户端 IP"
|
||||
// @Param host query string false "请求 Host"
|
||||
// @Param path query string false "请求路径"
|
||||
// @Param bucket_started_at query string false "时间桶起始时间"
|
||||
// @Param fold_minutes query int false "折叠时间窗口(分钟)"
|
||||
// @Param p query int false "页码"
|
||||
// @Param page_size query int false "每页条数"
|
||||
// @Param sort_by query string false "排序字段"
|
||||
// @Param sort_order query string false "排序方向"
|
||||
// @Success 200 {object} response.Any{data=observability.FoldedAccessLogIPList} "折叠 IP 汇总列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/access-logs/folds/ip-summary [get]
|
||||
func GetFoldedAccessLogIPsHandler(c *gin.Context) {
|
||||
result, err := ListFoldedAccessLogIPs(c.Request.Context(), FoldedAccessLogIPQuery{
|
||||
NodeID: c.Query("node_id"),
|
||||
RemoteAddr: c.Query("remote_addr"),
|
||||
Host: c.Query("host"),
|
||||
Path: c.Query("path"),
|
||||
BucketStartedAt: c.Query("bucket_started_at"),
|
||||
FoldMinutes: readQueryInt(c, "fold_minutes"),
|
||||
Page: readQueryInt(c, "p"),
|
||||
PageSize: readQueryInt(c, "page_size"),
|
||||
SortBy: c.Query("sort_by"),
|
||||
SortOrder: c.Query("sort_order"),
|
||||
})
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// getAccessLogIPSummariesHandler 列出访问日志 IP 汇总。
|
||||
// @Summary 列出访问日志 IP 汇总
|
||||
// @Description 按 IP 聚合访问日志统计并分页返回,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param node_id query string false "节点 ID"
|
||||
// @Param remote_addr query string false "客户端 IP"
|
||||
// @Param host query string false "请求 Host"
|
||||
// @Param p query int false "页码"
|
||||
// @Param page_size query int false "每页条数"
|
||||
// @Param sort_by query string false "排序字段"
|
||||
// @Param sort_order query string false "排序方向"
|
||||
// @Success 200 {object} response.Any{data=observability.AccessLogIPSummaryList} "IP 汇总列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/access-logs/ip-summary [get]
|
||||
func GetAccessLogIPSummariesHandler(c *gin.Context) {
|
||||
result, err := ListAccessLogIPSummaries(c.Request.Context(), AccessLogIPSummaryQuery{
|
||||
NodeID: c.Query("node_id"),
|
||||
RemoteAddr: c.Query("remote_addr"),
|
||||
Host: c.Query("host"),
|
||||
Page: readQueryInt(c, "p"),
|
||||
PageSize: readQueryInt(c, "page_size"),
|
||||
SortBy: c.Query("sort_by"),
|
||||
SortOrder: c.Query("sort_order"),
|
||||
})
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// getAccessLogIPTrendHandler 获取 IP 访问趋势。
|
||||
// @Summary 获取访问日志 IP 趋势
|
||||
// @Description 返回指定 IP 在时间范围内的访问趋势数据,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param node_id query string false "节点 ID"
|
||||
// @Param remote_addr query string false "客户端 IP"
|
||||
// @Param host query string false "请求 Host"
|
||||
// @Param hours query int false "统计时间范围(小时)"
|
||||
// @Param bucket_minutes query int false "时间桶粒度(分钟)"
|
||||
// @Success 200 {object} response.Any{data=observability.AccessLogIPTrendView} "IP 访问趋势"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/access-logs/ip-summary/trend [get]
|
||||
func GetAccessLogIPTrendHandler(c *gin.Context) {
|
||||
result, err := GetAccessLogIPTrend(c.Request.Context(), AccessLogIPTrendQuery{
|
||||
NodeID: c.Query("node_id"),
|
||||
RemoteAddr: c.Query("remote_addr"),
|
||||
Host: c.Query("host"),
|
||||
Hours: readQueryInt(c, "hours"),
|
||||
BucketMinutes: readQueryInt(c, "bucket_minutes"),
|
||||
})
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// cleanupAccessLogsHandler 清理过期访问日志。
|
||||
// @Summary 清理访问日志
|
||||
// @Description 按保留天数清理过期访问日志记录,需要管理员权限
|
||||
// @Tags openflare-observability
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body observability.AccessLogCleanupInput true "清理参数"
|
||||
// @Success 200 {object} response.Any{data=observability.AccessLogCleanupResult} "清理结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/access-logs/cleanup [post]
|
||||
func CleanupAccessLogsHandler(c *gin.Context) {
|
||||
var input AccessLogCleanupInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
result, err := CleanupAccessLogs(c.Request.Context(), input)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
func readAccessLogQuery(c *gin.Context) AccessLogQuery {
|
||||
return AccessLogQuery{
|
||||
NodeID: c.Query("node_id"),
|
||||
RemoteAddr: c.Query("remote_addr"),
|
||||
Host: c.Query("host"),
|
||||
Path: c.Query("path"),
|
||||
Page: readQueryInt(c, "p"),
|
||||
PageSize: readQueryInt(c, "page_size"),
|
||||
SortBy: c.Query("sort_by"),
|
||||
SortOrder: c.Query("sort_order"),
|
||||
}
|
||||
}
|
||||
|
||||
func readQueryInt(c *gin.Context, key string) int {
|
||||
value, _ := strconv.Atoi(c.DefaultQuery(key, "0"))
|
||||
return value
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package option
|
||||
|
||||
const (
|
||||
errInvalidParams = "无效的参数"
|
||||
errOptionInitFailed = "系统选项初始化失败"
|
||||
errGeoIPProvider = "归属方式仅支持 disabled、mmdb、ip-api、geojs、ipinfo"
|
||||
errGeoIPIPEmpty = "IP 不能为空"
|
||||
errGeoIPIPInvalid = "IP 格式无效"
|
||||
errGeoIPLookupDisabled = "GeoIP 查询已禁用"
|
||||
)
|
||||
@@ -0,0 +1,243 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package option
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
|
||||
oftasks "github.com/Rain-kl/Wavelet/internal/apps/openflare/tasks"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/uptimekuma"
|
||||
"github.com/Rain-kl/Wavelet/internal/buildinfo"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
var (
|
||||
initOnce sync.Once
|
||||
initErr error
|
||||
)
|
||||
|
||||
// EnsureInitialized loads OptionMap from defaults and database once per process.
|
||||
func EnsureInitialized(ctx context.Context) error {
|
||||
initOnce.Do(func() {
|
||||
initErr = model.InitOptionMap(ctx)
|
||||
})
|
||||
return initErr
|
||||
}
|
||||
|
||||
// ResetInitializationForTest clears lazy-init state for unit tests.
|
||||
func ResetInitializationForTest() {
|
||||
initOnce = sync.Once{}
|
||||
initErr = nil
|
||||
model.ResetOptionMapForTest()
|
||||
}
|
||||
|
||||
type publicAuthSourceView struct {
|
||||
ID uint64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
DisplayName string `json:"display_name"`
|
||||
AuthorizeURL string `json:"authorize_url"`
|
||||
IconURL string `json:"icon_url"`
|
||||
}
|
||||
|
||||
type statusView struct {
|
||||
Version string `json:"version"`
|
||||
StartTime int64 `json:"start_time"`
|
||||
EmailVerification bool `json:"email_verification"`
|
||||
GitHubOAuth bool `json:"github_oauth"`
|
||||
GitHubClientID string `json:"github_client_id"`
|
||||
SystemName string `json:"system_name"`
|
||||
HomePageLink string `json:"home_page_link"`
|
||||
FooterHTML string `json:"footer_html"`
|
||||
WeChatQRCode string `json:"wechat_qrcode"`
|
||||
WeChatLogin bool `json:"wechat_login"`
|
||||
ServerAddress string `json:"server_address"`
|
||||
PasswordRegisterEnabled bool `json:"password_register_enabled"`
|
||||
CapLoginEnabled bool `json:"cap_login_enabled"`
|
||||
AuthSources []publicAuthSourceView `json:"auth_sources"`
|
||||
}
|
||||
|
||||
type geoIPLookupRequest struct {
|
||||
Provider string `json:"provider"`
|
||||
IP string `json:"ip"`
|
||||
}
|
||||
|
||||
type geoIPLookupView struct {
|
||||
Provider string `json:"provider"`
|
||||
IP string `json:"ip"`
|
||||
ISOCode string `json:"iso_code"`
|
||||
Name string `json:"name"`
|
||||
Latitude *float64 `json:"latitude,omitempty"`
|
||||
Longitude *float64 `json:"longitude,omitempty"`
|
||||
}
|
||||
|
||||
type databaseCleanupInput struct {
|
||||
Target string `json:"target"`
|
||||
RetentionDays *int `json:"retention_days"`
|
||||
}
|
||||
|
||||
type databaseCleanupResult struct {
|
||||
Target string `json:"target"`
|
||||
TargetLabel string `json:"target_label"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays *int `json:"retention_days,omitempty"`
|
||||
}
|
||||
|
||||
type optionBatchPayload struct {
|
||||
Options []model.OpenFlareOption `json:"options"`
|
||||
}
|
||||
|
||||
func listOptions(ctx context.Context) ([]model.OpenFlareOption, error) {
|
||||
if err := EnsureInitialized(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
model.OptionMapRWMutex.RLock()
|
||||
defer model.OptionMapRWMutex.RUnlock()
|
||||
|
||||
options := make([]model.OpenFlareOption, 0, len(model.OptionMap))
|
||||
for key, value := range model.OptionMap {
|
||||
if isSecretOptionKey(key) {
|
||||
continue
|
||||
}
|
||||
options = append(options, model.OpenFlareOption{
|
||||
Key: key,
|
||||
Value: value,
|
||||
})
|
||||
}
|
||||
return options, nil
|
||||
}
|
||||
|
||||
func updateOption(ctx context.Context, option model.OpenFlareOption) error {
|
||||
if err := EnsureInitialized(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return updateOptions(ctx, []model.OpenFlareOption{option})
|
||||
}
|
||||
|
||||
func updateOptionsBatch(ctx context.Context, payload optionBatchPayload) error {
|
||||
if err := EnsureInitialized(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
if len(payload.Options) == 0 {
|
||||
return errors.New(errInvalidParams)
|
||||
}
|
||||
return updateOptions(ctx, payload.Options)
|
||||
}
|
||||
|
||||
func updateOptions(ctx context.Context, options []model.OpenFlareOption) error {
|
||||
if err := validateOptions(options); err != nil {
|
||||
return err
|
||||
}
|
||||
return model.UpdateOpenFlareOptions(ctx, options)
|
||||
}
|
||||
|
||||
func getNotice(ctx context.Context) (string, error) {
|
||||
if err := EnsureInitialized(ctx); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return model.OptionValue("Notice"), nil
|
||||
}
|
||||
|
||||
func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) {
|
||||
if err := EnsureInitialized(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
authSources, err := publicAuthSources(ctx, baseAPIPath)
|
||||
if err != nil {
|
||||
authSources = []publicAuthSourceView{}
|
||||
}
|
||||
|
||||
return &statusView{
|
||||
Version: buildinfo.Version,
|
||||
StartTime: model.StartTime,
|
||||
EmailVerification: model.EmailVerificationEnabled,
|
||||
GitHubOAuth: model.GitHubOAuthEnabled,
|
||||
GitHubClientID: model.GitHubClientId,
|
||||
SystemName: model.SystemName,
|
||||
HomePageLink: model.HomePageLink,
|
||||
FooterHTML: model.Footer,
|
||||
WeChatQRCode: model.WeChatAccountQRCodeImageURL,
|
||||
WeChatLogin: model.WeChatAuthEnabled,
|
||||
ServerAddress: model.ServerAddress,
|
||||
PasswordRegisterEnabled: model.PasswordRegisterEnabled,
|
||||
CapLoginEnabled: model.CapLoginEnabled,
|
||||
AuthSources: authSources,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func publicAuthSources(ctx context.Context, baseAPIPath string) ([]publicAuthSourceView, error) {
|
||||
sources, err := model.GetActiveAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make([]publicAuthSourceView, 0, len(sources))
|
||||
base := strings.TrimRight(baseAPIPath, "/")
|
||||
for _, source := range sources {
|
||||
result = append(result, publicAuthSourceView{
|
||||
ID: source.ID,
|
||||
Name: source.Name,
|
||||
Type: source.Type,
|
||||
DisplayName: source.DisplayName,
|
||||
AuthorizeURL: fmt.Sprintf("%s/oauth/%s/authorize", base, source.Name),
|
||||
IconURL: source.IconURL,
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func lookupGeoIP(_ context.Context, provider, rawIP string) (*geoIPLookupView, error) {
|
||||
view, err := geoip.Lookup(provider, rawIP)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &geoIPLookupView{
|
||||
Provider: view.Provider,
|
||||
IP: view.IP,
|
||||
ISOCode: view.ISOCode,
|
||||
Name: view.Name,
|
||||
Latitude: view.Latitude,
|
||||
Longitude: view.Longitude,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func cleanupDatabaseObservability(ctx context.Context, input databaseCleanupInput) (*databaseCleanupResult, error) {
|
||||
target := strings.TrimSpace(input.Target)
|
||||
if target == "" {
|
||||
return nil, errors.New(errInvalidParams)
|
||||
}
|
||||
|
||||
result, err := oftasks.CleanupDatabaseObservability(ctx, oftasks.DatabaseCleanupInput{
|
||||
Target: target,
|
||||
RetentionDays: input.RetentionDays,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &databaseCleanupResult{
|
||||
Target: result.Target,
|
||||
TargetLabel: result.TargetLabel,
|
||||
DeletedCount: result.DeletedCount,
|
||||
DeleteAll: result.DeleteAll,
|
||||
RetentionDays: result.RetentionDays,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func syncUptimeKuma(ctx context.Context) error {
|
||||
return uptimekuma.SyncToUptimeKuma(ctx)
|
||||
}
|
||||
|
||||
func isSecretOptionKey(key string) bool {
|
||||
return strings.Contains(key, "Token") ||
|
||||
strings.Contains(key, "Secret") ||
|
||||
strings.Contains(key, "Password")
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package option
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupOptionTestDB(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.OpenFlareOption{}))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
ResetInitializationForTest()
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
ResetInitializationForTest()
|
||||
}
|
||||
}
|
||||
|
||||
func TestListOptionsFiltersSecretKeys(t *testing.T) {
|
||||
cleanup := setupOptionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
require.NoError(t, model.UpdateOpenFlareOptions(ctx, []model.OpenFlareOption{
|
||||
{Key: "SystemName", Value: "TestFlare"},
|
||||
{Key: "SMTPToken", Value: "secret-token"},
|
||||
{Key: "GitHubClientSecret", Value: "secret-id"},
|
||||
}))
|
||||
|
||||
options, err := listOptions(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
keys := make(map[string]string, len(options))
|
||||
for _, option := range options {
|
||||
keys[option.Key] = option.Value
|
||||
}
|
||||
|
||||
assert.Equal(t, "TestFlare", keys["SystemName"])
|
||||
assert.NotContains(t, keys, "SMTPToken")
|
||||
assert.NotContains(t, keys, "GitHubClientSecret")
|
||||
}
|
||||
|
||||
func TestUpdateOptionHotReloadsOptionMap(t *testing.T) {
|
||||
cleanup := setupOptionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
err := updateOption(ctx, model.OpenFlareOption{
|
||||
Key: "SystemName",
|
||||
Value: "HotReloaded",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "HotReloaded", model.OptionValue("SystemName"))
|
||||
assert.Equal(t, "HotReloaded", model.SystemName)
|
||||
}
|
||||
|
||||
func TestGetNotice(t *testing.T) {
|
||||
cleanup := setupOptionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
require.NoError(t, updateOption(ctx, model.OpenFlareOption{Key: "Notice", Value: "hello"}))
|
||||
|
||||
notice, err := getNotice(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "hello", notice)
|
||||
}
|
||||
|
||||
func TestLookupGeoIPDisabledProvider(t *testing.T) {
|
||||
cleanup := setupOptionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
view, err := lookupGeoIP(ctx, "disabled", "8.8.8.8")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "disabled", view.Provider)
|
||||
assert.Equal(t, "8.8.8.8", view.IP)
|
||||
}
|
||||
|
||||
func TestCleanupDatabaseObservabilityDeletesRows(t *testing.T) {
|
||||
cleanup := setupOptionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore())
|
||||
defer resetAccessLogStore()
|
||||
|
||||
now := time.Now().UTC()
|
||||
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{
|
||||
{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-10 * 24 * time.Hour),
|
||||
RemoteAddr: "203.0.113.1",
|
||||
Host: "example.com",
|
||||
Path: "/old",
|
||||
StatusCode: 200,
|
||||
},
|
||||
{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-2 * time.Hour),
|
||||
RemoteAddr: "203.0.113.2",
|
||||
Host: "example.com",
|
||||
Path: "/recent",
|
||||
StatusCode: 200,
|
||||
},
|
||||
}))
|
||||
|
||||
retention := 7
|
||||
result, err := cleanupDatabaseObservability(ctx, databaseCleanupInput{
|
||||
Target: "node_access_logs",
|
||||
RetentionDays: &retention,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "node_access_logs", result.Target)
|
||||
assert.Equal(t, "访问日志", result.TargetLabel)
|
||||
assert.Equal(t, int64(1), result.DeletedCount)
|
||||
assert.False(t, result.DeleteAll)
|
||||
require.NotNil(t, result.RetentionDays)
|
||||
assert.Equal(t, 7, *result.RetentionDays)
|
||||
|
||||
rows, err := model.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, rows, 1)
|
||||
assert.Equal(t, "/recent", rows[0].Path)
|
||||
}
|
||||
@@ -0,0 +1,207 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package option
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GetStatusHandler 获取公开运行状态。
|
||||
// @Summary 获取 OpenFlare 公开状态
|
||||
// @Description 返回版本、认证源与系统公开配置,无需登录
|
||||
// @Tags openflare-option
|
||||
// @Produce json
|
||||
// @Success 200 {object} response.Any{data=option.statusView} "公开状态"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/status [get]
|
||||
func GetStatusHandler(c *gin.Context) {
|
||||
view, err := getStatus(c.Request.Context(), "/api/v1/d")
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// getNoticeHandler 获取系统公告。
|
||||
// @Summary 获取系统公告
|
||||
// @Description 返回 OpenFlare 控制台公告文本,无需登录
|
||||
// @Tags openflare-option
|
||||
// @Produce json
|
||||
// @Success 200 {object} response.Any{data=string} "系统公告"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/notice [get]
|
||||
// GetNoticeHandler returns the notice content.
|
||||
func GetNoticeHandler(c *gin.Context) {
|
||||
notice, err := getNotice(c.Request.Context())
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(notice))
|
||||
}
|
||||
|
||||
// listOptionsHandler 列出全部配置项。
|
||||
// @Summary 列出 OpenFlare 配置项
|
||||
// @Description 返回全部非敏感 OpenFlare 配置项,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.OpenFlareOption} "配置项列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option [get]
|
||||
// ListOptionsHandler lists OpenFlare options.
|
||||
func ListOptionsHandler(c *gin.Context) {
|
||||
options, err := listOptions(c.Request.Context())
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(options))
|
||||
}
|
||||
|
||||
// updateOptionHandler 更新单个配置项。
|
||||
// @Summary 更新 OpenFlare 配置项
|
||||
// @Description 更新单个 OpenFlare 配置项,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body model.OpenFlareOption true "配置项"
|
||||
// @Success 200 {object} response.Any "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/update [post]
|
||||
// UpdateOptionHandler updates a single option.
|
||||
func UpdateOptionHandler(c *gin.Context) {
|
||||
var option model.OpenFlareOption
|
||||
if !apiutil.BindJSON(c, &option) {
|
||||
return
|
||||
}
|
||||
if apiutil.AbortBadRequestOnError(c, updateOption(c.Request.Context(), option)) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// updateOptionsBatchHandler 批量更新配置项。
|
||||
// @Summary 批量更新 OpenFlare 配置项
|
||||
// @Description 批量更新多个 OpenFlare 配置项,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body option.optionBatchPayload true "批量配置项"
|
||||
// @Success 200 {object} response.Any "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/update-batch [post]
|
||||
// UpdateOptionsBatchHandler updates options in batch.
|
||||
func UpdateOptionsBatchHandler(c *gin.Context) {
|
||||
var payload optionBatchPayload
|
||||
if !apiutil.BindJSON(c, &payload) {
|
||||
return
|
||||
}
|
||||
if apiutil.AbortBadRequestOnError(c, updateOptionsBatch(c.Request.Context(), payload)) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// lookupGeoIPHandler 查询 GeoIP 信息。
|
||||
// @Summary GeoIP 地址查询
|
||||
// @Description 按提供商与 IP 查询地理位置信息,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body option.geoIPLookupRequest true "查询参数"
|
||||
// @Success 200 {object} response.Any{data=option.geoIPLookupView} "GeoIP 查询结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/geoip/lookup [post]
|
||||
// LookupGeoIPHandler performs a GeoIP lookup.
|
||||
func LookupGeoIPHandler(c *gin.Context) {
|
||||
var request geoIPLookupRequest
|
||||
if !apiutil.BindJSON(c, &request) {
|
||||
return
|
||||
}
|
||||
view, err := lookupGeoIP(c.Request.Context(), request.Provider, request.IP)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
// cleanupDatabaseHandler 清理可观测性数据库数据。
|
||||
// @Summary 清理可观测性数据库
|
||||
// @Description 按目标与保留天数清理可观测性相关数据表,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body option.databaseCleanupInput false "清理参数"
|
||||
// @Success 200 {object} response.Any{data=option.databaseCleanupResult} "清理结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/option/database/cleanup [post]
|
||||
// CleanupDatabaseHandler cleans up observability data.
|
||||
func CleanupDatabaseHandler(c *gin.Context) {
|
||||
var input databaseCleanupInput
|
||||
if err := bindOptionalJSON(c.Request.Body, &input); err != nil {
|
||||
response.AbortBadRequest(c, errInvalidParams)
|
||||
return
|
||||
}
|
||||
result, err := cleanupDatabaseObservability(c.Request.Context(), input)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
|
||||
// syncUptimeKumaHandler 同步 Uptime Kuma 监控。
|
||||
// @Summary 同步 Uptime Kuma
|
||||
// @Description 将 OpenFlare 节点同步到 Uptime Kuma,需要管理员权限
|
||||
// @Tags openflare-option
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=string} "同步成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/uptimekuma/sync [post]
|
||||
// SyncUptimeKumaHandler triggers UptimeKuma sync.
|
||||
func SyncUptimeKumaHandler(c *gin.Context) {
|
||||
if apiutil.AbortBadRequestOnError(c, syncUptimeKuma(c.Request.Context())) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK("同步成功"))
|
||||
}
|
||||
|
||||
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,286 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package option
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
var (
|
||||
openRestySizePattern = regexp.MustCompile(`^\d+[kKmMgG]?$`)
|
||||
openRestyProxyBuffersPattern = regexp.MustCompile(`^\d+\s+\d+[kKmMgG]?$`)
|
||||
openRestyCacheLevelsPattern = regexp.MustCompile(`^\d{1,2}(?::\d{1,2}){0,2}$`)
|
||||
openRestyDurationTokenPattern = regexp.MustCompile(`^\d+[smhdwSMHDW]$`)
|
||||
)
|
||||
|
||||
func buildOptionValidationState(options []model.OpenFlareOption) map[string]string {
|
||||
model.OptionMapRWMutex.RLock()
|
||||
state := make(map[string]string, len(model.OptionMap)+len(options))
|
||||
for key, value := range model.OptionMap {
|
||||
state[key] = value
|
||||
}
|
||||
model.OptionMapRWMutex.RUnlock()
|
||||
|
||||
for _, option := range options {
|
||||
state[option.Key] = option.Value
|
||||
}
|
||||
return state
|
||||
}
|
||||
|
||||
func validateOptionWithState(option model.OpenFlareOption, state map[string]string) error {
|
||||
switch option.Key {
|
||||
case "GitHubOAuthEnabled":
|
||||
if option.Value == "true" && strings.TrimSpace(state["GitHubClientId"]) == "" {
|
||||
return fmt.Errorf("无法启用 GitHub OAuth,请先填入 GitHub Client ID 以及 GitHub Client Secret!")
|
||||
}
|
||||
case "WeChatAuthEnabled":
|
||||
if option.Value == "true" && strings.TrimSpace(state["WeChatServerAddress"]) == "" {
|
||||
return fmt.Errorf("无法启用微信登录,请先填入微信登录相关配置信息!")
|
||||
}
|
||||
}
|
||||
|
||||
if err := validateOpenRestyOption(option.Key, option.Value); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateGeoIPOption(option.Key, option.Value); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateDatabaseCleanupOption(option.Key, option.Value); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateAgentOption(option.Key, option.Value); err != nil {
|
||||
return err
|
||||
}
|
||||
return validateUptimeKumaOption(option.Key, option.Value, state)
|
||||
}
|
||||
|
||||
func validatePositiveIntegerOption(key, value string) error {
|
||||
intValue, err := strconv.Atoi(value)
|
||||
if err != nil || intValue <= 0 {
|
||||
return fmt.Errorf("%s 必须为大于 0 的整数", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateBooleanOption(key, value string) error {
|
||||
switch value {
|
||||
case "true", "false":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("%s 必须为 true 或 false", key)
|
||||
}
|
||||
}
|
||||
|
||||
func validateGeoIPOption(key, value string) error {
|
||||
if key != "GeoIPProvider" {
|
||||
return nil
|
||||
}
|
||||
if geoip.IsValidProvider(value) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 仅支持 disabled、mmdb、ip-api、geojs、ipinfo", key)
|
||||
}
|
||||
|
||||
func validateDatabaseCleanupOption(key, value string) error {
|
||||
switch key {
|
||||
case "DatabaseAutoCleanupEnabled":
|
||||
return validateBooleanOption(key, value)
|
||||
case "DatabaseAutoCleanupRetentionDays":
|
||||
intValue, err := strconv.Atoi(value)
|
||||
if err != nil || intValue < 1 {
|
||||
return fmt.Errorf("%s 必须为大于等于 1 的整数天", key)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateAgentOption(key, value string) error {
|
||||
if key == "AgentWebsocketUpgradeEnabled" {
|
||||
return validateBooleanOption(key, strings.TrimSpace(value))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUptimeKumaOption(key, value string, state map[string]string) error {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
switch key {
|
||||
case "UptimeKumaEnabled":
|
||||
if err := validateBooleanOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
if trimmed == "true" {
|
||||
url := strings.TrimSpace(state["UptimeKumaUrl"])
|
||||
username := strings.TrimSpace(state["UptimeKumaUsername"])
|
||||
password := strings.TrimSpace(state["UptimeKumaPassword"])
|
||||
if url == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时地址不能为空")
|
||||
}
|
||||
if username == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
if password == "" && model.UptimeKumaPassword == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时密码不能为空")
|
||||
}
|
||||
}
|
||||
case "UptimeKumaUsername":
|
||||
if trimmed == "" && state["UptimeKumaEnabled"] == "true" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
case "UptimeKumaUrl":
|
||||
if trimmed != "" && !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") {
|
||||
return fmt.Errorf("Uptime Kuma 地址必须以 http:// 或 https:// 开头")
|
||||
}
|
||||
case "UptimeKumaMonitorScope":
|
||||
if trimmed != "all" && trimmed != "selected" {
|
||||
return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)")
|
||||
}
|
||||
case "UptimeKumaSyncInterval", "UptimeKumaInterval", "UptimeKumaRetryInterval", "UptimeKumaTimeout":
|
||||
return validatePositiveIntegerOption(key, trimmed)
|
||||
case "UptimeKumaRetry":
|
||||
intValue, err := strconv.Atoi(trimmed)
|
||||
if err != nil || intValue < 0 {
|
||||
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOpenRestyOption(key, value string) error {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
|
||||
switch key {
|
||||
case "OpenRestyDefaultServerReturnStatus":
|
||||
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
statusCode, _ := strconv.Atoi(trimmed)
|
||||
if statusCode < 100 || statusCode > 999 {
|
||||
return fmt.Errorf("%s 必须在 100 到 999 之间", key)
|
||||
}
|
||||
case "OpenRestyWorkerProcesses":
|
||||
if trimmed == "auto" {
|
||||
return nil
|
||||
}
|
||||
return validatePositiveIntegerOption(key, trimmed)
|
||||
case "OpenRestyWorkerConnections",
|
||||
"OpenRestyWorkerRlimitNofile",
|
||||
"OpenRestyKeepaliveTimeout",
|
||||
"OpenRestyKeepaliveRequests",
|
||||
"OpenRestyClientHeaderTimeout",
|
||||
"OpenRestyClientBodyTimeout",
|
||||
"OpenRestySendTimeout",
|
||||
"OpenRestyProxyConnectTimeout",
|
||||
"OpenRestyProxySendTimeout",
|
||||
"OpenRestyProxyReadTimeout",
|
||||
"OpenRestyGzipMinLength":
|
||||
return validatePositiveIntegerOption(key, trimmed)
|
||||
case "OpenRestyGzipCompLevel":
|
||||
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
|
||||
return err
|
||||
}
|
||||
level, _ := strconv.Atoi(trimmed)
|
||||
if level > 9 {
|
||||
return fmt.Errorf("%s 不能大于 9", key)
|
||||
}
|
||||
case "OpenRestyEventsUse":
|
||||
if trimmed == "" {
|
||||
return nil
|
||||
}
|
||||
switch trimmed {
|
||||
case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key)
|
||||
}
|
||||
case "OpenRestyResolvers":
|
||||
if trimmed == "" {
|
||||
return nil
|
||||
}
|
||||
if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) {
|
||||
return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key)
|
||||
}
|
||||
case "OpenRestyEventsMultiAcceptEnabled",
|
||||
"OpenRestyWebsocketEnabled",
|
||||
"OpenRestyHTTP3Enabled",
|
||||
"OpenRestyProxyRequestBufferingEnabled",
|
||||
"OpenRestyProxyBufferingEnabled",
|
||||
"OpenRestyGzipEnabled",
|
||||
"OpenRestyCacheEnabled",
|
||||
"OpenRestyCacheLockEnabled":
|
||||
return validateBooleanOption(key, trimmed)
|
||||
case "OpenRestyProxyBuffers", "OpenRestyLargeClientHeaderBuffers":
|
||||
if openRestyProxyBuffersPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key)
|
||||
case "OpenRestyProxyBufferSize", "OpenRestyProxyBusyBuffersSize", "OpenRestyCacheMaxSize", "OpenRestyClientMaxBodySize":
|
||||
if openRestySizePattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key)
|
||||
case "OpenRestyCachePath":
|
||||
if strings.ContainsAny(trimmed, "\r\n\t") {
|
||||
return fmt.Errorf("%s 不能包含换行或制表符", key)
|
||||
}
|
||||
case "OpenRestyCacheLevels":
|
||||
if openRestyCacheLevelsPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key)
|
||||
case "OpenRestyCacheInactive", "OpenRestyCacheLockTimeout":
|
||||
if openRestyDurationTokenPattern.MatchString(trimmed) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key)
|
||||
case "OpenRestyCacheKeyTemplate":
|
||||
if trimmed == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
if strings.ContainsAny(trimmed, "\r\n") {
|
||||
return fmt.Errorf("%s 不能包含换行", key)
|
||||
}
|
||||
case "OpenRestyCacheUseStale":
|
||||
if trimmed == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
allowedTokens := map[string]struct{}{
|
||||
"error": {}, "timeout": {}, "invalid_header": {}, "updating": {},
|
||||
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
|
||||
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
|
||||
}
|
||||
for _, token := range strings.Fields(trimmed) {
|
||||
if _, ok := allowedTokens[token]; !ok {
|
||||
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
|
||||
}
|
||||
}
|
||||
case "OpenRestyMainConfigTemplate":
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOptions(options []model.OpenFlareOption) error {
|
||||
if len(options) == 0 {
|
||||
return fmt.Errorf(errInvalidParams)
|
||||
}
|
||||
|
||||
state := buildOptionValidationState(options)
|
||||
for _, option := range options {
|
||||
if strings.TrimSpace(option.Key) == "" {
|
||||
return fmt.Errorf(errInvalidParams)
|
||||
}
|
||||
if err := validateOptionWithState(option, state); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package origin
|
||||
|
||||
const (
|
||||
errOriginAddressRequired = "源站地址不能为空"
|
||||
errOriginAddressInvalid = "源站地址格式不合法"
|
||||
errOriginAddressExists = "源站地址已存在"
|
||||
errOriginDeleteReferenced = "该源站仍被规则引用,无法删除"
|
||||
errOriginMissingPort = "源站地址缺少端口"
|
||||
errOriginNotFound = "源站不存在"
|
||||
)
|
||||
@@ -0,0 +1,87 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package origin
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
func normalizeOriginAddress(raw string) string {
|
||||
return strings.ToLower(strings.TrimSpace(raw))
|
||||
}
|
||||
|
||||
func validateOriginAddress(address string) error {
|
||||
if address == "" {
|
||||
return errors.New(errOriginAddressRequired)
|
||||
}
|
||||
if strings.Contains(address, "://") || strings.ContainsAny(address, "/?#") {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
if strings.HasPrefix(address, "[") || strings.HasSuffix(address, "]") {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
if ip := net.ParseIP(address); ip != nil {
|
||||
return nil
|
||||
}
|
||||
if len(address) > 253 {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
for _, label := range labels {
|
||||
if len(label) == 0 || len(label) > 63 {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
if label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
for _, r := range label {
|
||||
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' {
|
||||
continue
|
||||
}
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeOriginName(name string, address string) string {
|
||||
normalized := strings.TrimSpace(name)
|
||||
if normalized != "" {
|
||||
return normalized
|
||||
}
|
||||
return address
|
||||
}
|
||||
|
||||
func formatOriginHost(address string, port string) string {
|
||||
return net.JoinHostPort(address, port)
|
||||
}
|
||||
|
||||
func rewriteOriginURLAddress(rawURL string, newAddress string) (string, error) {
|
||||
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", errOriginAddressInvalid, err)
|
||||
}
|
||||
address := normalizeOriginAddress(newAddress)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return "", err
|
||||
}
|
||||
port := parsed.Port()
|
||||
if port == "" {
|
||||
return "", errors.New(errOriginMissingPort)
|
||||
}
|
||||
parsed.Host = formatOriginHost(address, port)
|
||||
return parsed.String(), nil
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package origin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Input 源站创建/更新请求。
|
||||
type Input struct {
|
||||
Name string `json:"name"`
|
||||
Address string `json:"address"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
// RouteSummary 源站详情中的代理规则摘要。
|
||||
type RouteSummary struct {
|
||||
ID uint `json:"id"`
|
||||
Domain string `json:"domain"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
Enabled bool `json:"enabled"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
// View 源站列表项。
|
||||
type View struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Address string `json:"address"`
|
||||
Remark string `json:"remark"`
|
||||
RouteCount int64 `json:"route_count"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
// DetailView 源站详情。
|
||||
type DetailView struct {
|
||||
View
|
||||
Routes []RouteSummary `json:"routes"`
|
||||
}
|
||||
|
||||
// ListOrigins 列出全部源站。
|
||||
func ListOrigins(ctx context.Context) ([]View, error) {
|
||||
origins, err := model.ListOrigins(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildOriginViews(ctx, origins)
|
||||
}
|
||||
|
||||
// GetOriginDetail 获取源站详情。
|
||||
func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) {
|
||||
origin, err := model.GetOriginByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views, err := buildOriginViews(ctx, []model.Origin{*origin})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
routes, err := model.ListProxyRoutesByOriginID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]RouteSummary, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
items = append(items, RouteSummary{
|
||||
ID: route.ID,
|
||||
Domain: route.Domain,
|
||||
OriginURL: route.OriginURL,
|
||||
Enabled: route.Enabled,
|
||||
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
|
||||
})
|
||||
}
|
||||
sort.Slice(items, func(i, j int) bool {
|
||||
return items[i].Domain < items[j].Domain
|
||||
})
|
||||
return &DetailView{
|
||||
View: views[0],
|
||||
Routes: items,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// CreateOrigin 创建源站。
|
||||
func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) {
|
||||
origin, err := buildOrigin(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.CreateOriginRecord(ctx, origin); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errOriginAddressExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return origin, nil
|
||||
}
|
||||
|
||||
// UpdateOrigin 更新源站。
|
||||
func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, error) {
|
||||
origin, err := model.GetOriginByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
previousAddress := origin.Address
|
||||
nextOrigin, err := buildOrigin(origin, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Save(nextOrigin).Error; err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return errors.New(errOriginAddressExists)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if previousAddress == nextOrigin.Address {
|
||||
return nil
|
||||
}
|
||||
return updateRoutesForOriginAddress(ctx, tx, nextOrigin.ID, nextOrigin.Address)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nextOrigin, nil
|
||||
}
|
||||
|
||||
// DeleteOrigin 删除源站。
|
||||
func DeleteOrigin(ctx context.Context, id uint) error {
|
||||
count, err := model.CountProxyRoutesByOriginID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New(errOriginDeleteReferenced)
|
||||
}
|
||||
if _, err = model.GetOriginByID(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return model.DeleteOriginRecord(ctx, id)
|
||||
}
|
||||
|
||||
func buildOrigin(existing *model.Origin, input Input) (*model.Origin, error) {
|
||||
address := normalizeOriginAddress(input.Address)
|
||||
if err := validateOriginAddress(address); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
existing = &model.Origin{}
|
||||
}
|
||||
existing.Address = address
|
||||
existing.Name = normalizeOriginName(input.Name, address)
|
||||
existing.Remark = strings.TrimSpace(input.Remark)
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
func buildOriginViews(ctx context.Context, origins []model.Origin) ([]View, error) {
|
||||
countRows, err := model.ListOriginRouteCounts(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
countMap := make(map[uint]int64, len(countRows))
|
||||
for _, row := range countRows {
|
||||
countMap[row.OriginID] = row.RouteCount
|
||||
}
|
||||
views := make([]View, 0, len(origins))
|
||||
for _, origin := range origins {
|
||||
views = append(views, View{
|
||||
ID: origin.ID,
|
||||
Name: origin.Name,
|
||||
Address: origin.Address,
|
||||
Remark: origin.Remark,
|
||||
RouteCount: countMap[origin.ID],
|
||||
CreatedAt: origin.CreatedAt.Format("2006-01-02T15:04:05Z07:00"),
|
||||
UpdatedAt: origin.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uint, address string) error {
|
||||
if !model.HasProxyRoutesTable(ctx) {
|
||||
return nil
|
||||
}
|
||||
var routes []model.OriginProxyRoute
|
||||
if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; err != nil {
|
||||
return fmt.Errorf("query routes for origin update failed: %w", err)
|
||||
}
|
||||
for _, route := range routes {
|
||||
rewrittenOriginURL, err := rewriteOriginURLAddress(route.OriginURL, address)
|
||||
if err != nil {
|
||||
return fmt.Errorf("rewrite route %d origin failed: %w", route.ID, err)
|
||||
}
|
||||
upstreams := make([]string, 0)
|
||||
if strings.TrimSpace(route.Upstreams) != "" {
|
||||
if err := json.Unmarshal([]byte(route.Upstreams), &upstreams); err != nil {
|
||||
return fmt.Errorf("decode route %d upstreams failed: %w", route.ID, err)
|
||||
}
|
||||
}
|
||||
if len(upstreams) == 0 {
|
||||
upstreams = append(upstreams, rewrittenOriginURL)
|
||||
} else {
|
||||
upstreams[0] = rewrittenOriginURL
|
||||
}
|
||||
upstreamsJSON, err := json.Marshal(upstreams)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode route %d upstreams failed: %w", route.ID, err)
|
||||
}
|
||||
if err := tx.Model(&model.OriginProxyRoute{}).
|
||||
Where("id = ?", route.ID).
|
||||
Updates(map[string]any{
|
||||
"origin_url": rewrittenOriginURL,
|
||||
"upstreams": string(upstreamsJSON),
|
||||
}).Error; err != nil {
|
||||
return fmt.Errorf("update route %d origin address failed: %w", route.ID, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package origin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupOriginTestDB(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.Origin{}))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateOrigin(t *testing.T) {
|
||||
cleanup := setupOriginTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
origin, err := CreateOrigin(ctx, Input{
|
||||
Name: "Primary Origin",
|
||||
Address: "origin-a.internal",
|
||||
Remark: "main upstream",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, origin.ID)
|
||||
assert.Equal(t, "Primary Origin", origin.Name)
|
||||
assert.Equal(t, "origin-a.internal", origin.Address)
|
||||
assert.Equal(t, "main upstream", origin.Remark)
|
||||
|
||||
_, err = CreateOrigin(ctx, Input{
|
||||
Address: "origin-a.internal",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errOriginAddressExists, err.Error())
|
||||
}
|
||||
|
||||
func TestListOrigins(t *testing.T) {
|
||||
cleanup := setupOriginTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
first, err := CreateOrigin(ctx, Input{
|
||||
Name: "first-origin",
|
||||
Address: "origin-a.internal",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
second, err := CreateOrigin(ctx, Input{
|
||||
Name: "second-origin",
|
||||
Address: "origin-b.internal",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
origins, err := ListOrigins(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, origins, 2)
|
||||
assert.Equal(t, second.ID, origins[0].ID)
|
||||
assert.Equal(t, first.ID, origins[1].ID)
|
||||
assert.Equal(t, int64(0), origins[0].RouteCount)
|
||||
assert.Equal(t, int64(0), origins[1].RouteCount)
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package origin
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return apiutil.AbortNotFoundIfMissing(c, err, errOriginNotFound)
|
||||
}
|
||||
|
||||
// GetOrigins 列出全部源站。
|
||||
// @Summary 获取源站列表
|
||||
// @Description 返回所有源站及关联代理规则数量,需要管理员权限
|
||||
// @Tags openflare-origin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]origin.View} "源站列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/origins [get]
|
||||
func GetOrigins(c *gin.Context) {
|
||||
origins, err := ListOrigins(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(origins))
|
||||
}
|
||||
|
||||
// GetOrigin 获取源站详情。
|
||||
// @Summary 获取源站详情
|
||||
// @Description 返回指定源站信息及关联代理规则摘要,需要管理员权限
|
||||
// @Tags openflare-origin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "源站 ID"
|
||||
// @Success 200 {object} response.Any{data=origin.DetailView} "源站详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或源站不存在"
|
||||
// @Router /api/v1/d/origins/{id} [get]
|
||||
func GetOrigin(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
detail, err := GetOriginDetail(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(detail))
|
||||
}
|
||||
|
||||
// CreateOriginHandler 创建源站。
|
||||
// @Summary 创建源站
|
||||
// @Description 创建新的上游源站记录,需要管理员权限
|
||||
// @Tags openflare-origin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param body body origin.Input true "源站参数"
|
||||
// @Success 200 {object} response.Any{data=origin.View} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/origins [post]
|
||||
func CreateOriginHandler(c *gin.Context) {
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
origin, err := CreateOrigin(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(origin))
|
||||
}
|
||||
|
||||
// UpdateOriginHandler 更新源站。
|
||||
// @Summary 更新源站
|
||||
// @Description 更新指定源站的配置信息,需要管理员权限
|
||||
// @Tags openflare-origin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "源站 ID"
|
||||
// @Param body body origin.Input true "源站参数"
|
||||
// @Success 200 {object} response.Any{data=origin.View} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或源站不存在"
|
||||
// @Router /api/v1/d/origins/{id}/update [post]
|
||||
func UpdateOriginHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
origin, err := UpdateOrigin(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(origin))
|
||||
}
|
||||
|
||||
// DeleteOriginHandler 删除源站。
|
||||
// @Summary 删除源站
|
||||
// @Description 删除指定源站记录,需要管理员权限
|
||||
// @Tags openflare-origin
|
||||
// @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/origins/{id}/delete [post]
|
||||
func DeleteOriginHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteOrigin(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
BIN
Binary file not shown.
@@ -0,0 +1,27 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pages
|
||||
|
||||
const (
|
||||
errPagesProjectNotFound = "Pages 项目不存在"
|
||||
errPagesSlugExists = "Pages 项目标识已存在"
|
||||
errPagesNameRequired = "Pages 项目名称不能为空"
|
||||
errPagesSlugInvalid = "Pages 项目标识只能包含小写字母、数字和连字符"
|
||||
errPagesDeleteReferenced = "Pages 项目已被规则引用,不能删除"
|
||||
errPagesDeploymentNotFound = "Pages 部署不存在"
|
||||
errPagesDeploymentMismatch = "Pages 部署不属于该项目"
|
||||
errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署"
|
||||
errPagesPackageMissing = "缺少 Pages 部署包"
|
||||
errPagesPackageNotZip = "Pages 部署包必须是 .zip 文件"
|
||||
errPagesPackageInvalidZip = "Pages 部署包不是有效 zip 文件"
|
||||
errPagesPackageEmpty = "Pages 部署包不能为空"
|
||||
errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空"
|
||||
errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头"
|
||||
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空"
|
||||
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL"
|
||||
errPagesPackagePathEmpty = "Pages 部署包路径为空"
|
||||
errPagesPackageUploadMissing = "Pages 部署包上传记录不存在"
|
||||
errPagesPackageNotInActiveConfig = "Pages 部署尚未进入激活配置"
|
||||
errPagesInvalidSnapshotFormat = "配置快照格式无效"
|
||||
)
|
||||
@@ -0,0 +1,370 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pages
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/upload"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
)
|
||||
|
||||
const (
|
||||
pagesMaxDeploymentFiles = 1000
|
||||
pagesMaxDeploymentBytes = 100 * 1024 * 1024
|
||||
defaultPagesEntryFile = "index.html"
|
||||
defaultPagesFallbackPath = "/index.html"
|
||||
pagesDeploymentUploadType = "openflare_pages_deployment"
|
||||
)
|
||||
|
||||
var pagesSlugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,126}[a-z0-9]$|^[a-z0-9]$`)
|
||||
|
||||
type deploymentManifest struct {
|
||||
Files []model.PagesDeploymentFile
|
||||
FileCount int
|
||||
TotalSize int64
|
||||
EntryFile string
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
func normalizePagesSlug(raw string) string {
|
||||
value := strings.ToLower(strings.TrimSpace(raw))
|
||||
var builder strings.Builder
|
||||
lastDash := false
|
||||
for _, r := range value {
|
||||
valid := (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9')
|
||||
if valid {
|
||||
builder.WriteRune(r)
|
||||
lastDash = false
|
||||
continue
|
||||
}
|
||||
if !lastDash {
|
||||
builder.WriteByte('-')
|
||||
lastDash = true
|
||||
}
|
||||
}
|
||||
return strings.Trim(builder.String(), "-")
|
||||
}
|
||||
|
||||
func validateAndNormalizePagesRootDir(raw string) (string, error) {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" {
|
||||
return "", nil
|
||||
}
|
||||
if len(value) > 512 {
|
||||
return "", errors.New("Pages 根目录长度不能超过 512")
|
||||
}
|
||||
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
|
||||
return "", errors.New("Pages 根目录包含不支持的字符")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r <= 0x20 || r == 0x7f {
|
||||
return "", errors.New("Pages 根目录不能包含空白或控制字符")
|
||||
}
|
||||
}
|
||||
cleaned := path.Clean(filepath.ToSlash(value))
|
||||
if cleaned == "." || cleaned == "/" {
|
||||
return "", nil
|
||||
}
|
||||
for _, segment := range strings.Split(cleaned, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "", errors.New("Pages 根目录不能包含 . 或 .. 路径段")
|
||||
}
|
||||
}
|
||||
return strings.TrimPrefix(cleaned, "/"), nil
|
||||
}
|
||||
|
||||
func normalizePagesFallbackPath(raw string) (string, error) {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" {
|
||||
value = defaultPagesFallbackPath
|
||||
}
|
||||
if len(value) > 512 {
|
||||
return "", errors.New("SPA fallback 回退路径长度不能超过 512")
|
||||
}
|
||||
if !strings.HasPrefix(value, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须以 / 开头")
|
||||
}
|
||||
if value == "/" || strings.HasSuffix(value, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
|
||||
}
|
||||
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
|
||||
return "", errors.New("SPA fallback 回退路径包含不支持的字符")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r <= 0x20 || r == 0x7f {
|
||||
return "", errors.New("SPA fallback 回退路径不能包含空白或控制字符")
|
||||
}
|
||||
}
|
||||
for _, segment := range strings.Split(value, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "", errors.New("SPA fallback 回退路径不能包含 . 或 .. 路径段")
|
||||
}
|
||||
}
|
||||
cleaned := path.Clean(value)
|
||||
if cleaned == "." || !strings.HasPrefix(cleaned, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径不合法")
|
||||
}
|
||||
if cleaned == "/" || strings.HasSuffix(cleaned, "/") {
|
||||
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
|
||||
}
|
||||
return cleaned, nil
|
||||
}
|
||||
|
||||
func normalizeStoredPagesFallbackPath(value string) string {
|
||||
normalized, err := normalizePagesFallbackPath(value)
|
||||
if err != nil {
|
||||
return defaultPagesFallbackPath
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func normalizePagesEntryFile(raw string) string {
|
||||
value := path.Clean(strings.TrimSpace(filepath.ToSlash(raw)))
|
||||
if value == "." || value == "/" {
|
||||
return defaultPagesEntryFile
|
||||
}
|
||||
return strings.TrimPrefix(value, "/")
|
||||
}
|
||||
|
||||
func persistPagesUploadTemp(fileHeader *multipart.FileHeader) (string, string, int64, error) {
|
||||
file, err := fileHeader.Open()
|
||||
if err != nil {
|
||||
return "", "", 0, err
|
||||
}
|
||||
defer file.Close()
|
||||
temp, err := os.CreateTemp("", "openflare-pages-*.zip")
|
||||
if err != nil {
|
||||
return "", "", 0, err
|
||||
}
|
||||
defer temp.Close()
|
||||
hash := sha256.New()
|
||||
limited := io.LimitReader(file, pagesMaxDeploymentBytes+1)
|
||||
written, err := io.Copy(io.MultiWriter(temp, hash), limited)
|
||||
if err != nil {
|
||||
_ = os.Remove(temp.Name())
|
||||
return "", "", 0, err
|
||||
}
|
||||
if written > pagesMaxDeploymentBytes {
|
||||
_ = os.Remove(temp.Name())
|
||||
return "", "", 0, fmt.Errorf("Pages 部署包不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
|
||||
}
|
||||
return temp.Name(), hex.EncodeToString(hash.Sum(nil)), written, nil
|
||||
}
|
||||
|
||||
func ingestPagesDeploymentPackage(
|
||||
ctx context.Context,
|
||||
tempPath string,
|
||||
checksum string,
|
||||
size int64,
|
||||
projectSlug string,
|
||||
fileName string,
|
||||
) (upload.IngestResult, error) {
|
||||
file, err := os.Open(tempPath)
|
||||
if err != nil {
|
||||
return upload.IngestResult{}, err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
systemUser := repository.GetSystemUser(ctx)
|
||||
accessMode := 0
|
||||
return upload.Ingest(ctx, upload.IngestRequest{
|
||||
UserID: systemUser.ID,
|
||||
Reader: file,
|
||||
Size: size,
|
||||
FileName: fileName,
|
||||
MimeType: "application/zip",
|
||||
Extension: "zip",
|
||||
Hash: checksum,
|
||||
Type: pagesDeploymentUploadType,
|
||||
AccessMode: &accessMode,
|
||||
SkipExtensionCheck: true,
|
||||
Policy: upload.PolicyDedupNewRecord,
|
||||
Metadata: model.UploadMetadata{
|
||||
Extra: map[string]any{
|
||||
"project_slug": projectSlug,
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func removeDeploymentArtifact(ctx context.Context, deployment *model.PagesDeployment) {
|
||||
if deployment == nil {
|
||||
return
|
||||
}
|
||||
if deployment.UploadID > 0 {
|
||||
if _, err := upload.Remove(ctx, deployment.UploadID); err != nil {
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(deployment.ArtifactPath) != "" {
|
||||
_ = os.Remove(deployment.ArtifactPath)
|
||||
}
|
||||
}
|
||||
|
||||
func findCommonRootPrefix(files []*zip.File) (string, error) {
|
||||
var firstFilePath string
|
||||
hasMultipleFiles := false
|
||||
for _, item := range files {
|
||||
normalizedPath, skip, err := normalizePagesZipPath(item.Name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if skip {
|
||||
continue
|
||||
}
|
||||
if firstFilePath == "" {
|
||||
firstFilePath = normalizedPath
|
||||
} else {
|
||||
hasMultipleFiles = true
|
||||
}
|
||||
}
|
||||
if firstFilePath == "" {
|
||||
return "", nil
|
||||
}
|
||||
parts := strings.Split(firstFilePath, "/")
|
||||
if len(parts) <= 1 {
|
||||
return "", nil
|
||||
}
|
||||
commonPrefix := parts[0] + "/"
|
||||
if hasMultipleFiles {
|
||||
for _, item := range files {
|
||||
normalizedPath, skip, err := normalizePagesZipPath(item.Name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if skip {
|
||||
continue
|
||||
}
|
||||
if !strings.HasPrefix(normalizedPath, commonPrefix) {
|
||||
return "", nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return commonPrefix, nil
|
||||
}
|
||||
|
||||
func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploymentManifest, error) {
|
||||
reader, err := zip.OpenReader(zipPath)
|
||||
if err != nil {
|
||||
return nil, errors.New(errPagesPackageInvalidZip)
|
||||
}
|
||||
defer reader.Close()
|
||||
|
||||
commonPrefix, err := findCommonRootPrefix(reader.File)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
manifest := &deploymentManifest{
|
||||
Files: []model.PagesDeploymentFile{},
|
||||
EntryFile: entryFile,
|
||||
}
|
||||
targetEntryPath := entryFile
|
||||
if rootDir != "" {
|
||||
targetEntryPath = path.Join(rootDir, entryFile)
|
||||
}
|
||||
entrySeen := false
|
||||
for _, item := range reader.File {
|
||||
normalizedPath, skip, err := normalizePagesZipPath(item.Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if skip {
|
||||
continue
|
||||
}
|
||||
if commonPrefix != "" {
|
||||
normalizedPath = strings.TrimPrefix(normalizedPath, commonPrefix)
|
||||
}
|
||||
if item.FileInfo().Mode()&os.ModeSymlink != 0 {
|
||||
return nil, fmt.Errorf("Pages 部署包不支持符号链接: %s", normalizedPath)
|
||||
}
|
||||
if item.UncompressedSize64 > pagesMaxDeploymentBytes {
|
||||
return nil, fmt.Errorf("Pages 文件过大: %s", normalizedPath)
|
||||
}
|
||||
manifest.FileCount++
|
||||
if manifest.FileCount > pagesMaxDeploymentFiles {
|
||||
return nil, fmt.Errorf("Pages 部署文件数不能超过 %d", pagesMaxDeploymentFiles)
|
||||
}
|
||||
manifest.TotalSize += int64(item.UncompressedSize64)
|
||||
if manifest.TotalSize > pagesMaxDeploymentBytes {
|
||||
return nil, fmt.Errorf("Pages 部署展开后不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
|
||||
}
|
||||
checksum, err := checksumZipFile(item)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if normalizedPath == targetEntryPath {
|
||||
entrySeen = true
|
||||
}
|
||||
manifest.Files = append(manifest.Files, model.PagesDeploymentFile{
|
||||
Path: normalizedPath,
|
||||
Size: int64(item.UncompressedSize64),
|
||||
Checksum: checksum,
|
||||
})
|
||||
}
|
||||
if manifest.FileCount == 0 {
|
||||
return nil, errors.New(errPagesPackageEmpty)
|
||||
}
|
||||
if !entrySeen {
|
||||
return nil, fmt.Errorf("Pages 部署包缺少入口文件 %s", targetEntryPath)
|
||||
}
|
||||
return manifest, nil
|
||||
}
|
||||
|
||||
func normalizePagesZipPath(raw string) (string, bool, error) {
|
||||
name := strings.TrimSpace(filepath.ToSlash(raw))
|
||||
if name == "" {
|
||||
return "", true, nil
|
||||
}
|
||||
if strings.HasSuffix(name, "/") {
|
||||
return "", true, nil
|
||||
}
|
||||
if strings.HasPrefix(name, "/") || path.IsAbs(name) {
|
||||
return "", false, fmt.Errorf("Pages 部署包不能包含绝对路径: %s", raw)
|
||||
}
|
||||
cleaned := path.Clean(name)
|
||||
if cleaned == "." {
|
||||
return "", true, nil
|
||||
}
|
||||
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || strings.Contains(cleaned, "/../") {
|
||||
return "", false, fmt.Errorf("Pages 部署包路径不能逃逸目录: %s", raw)
|
||||
}
|
||||
return cleaned, false, nil
|
||||
}
|
||||
|
||||
func checksumZipFile(item *zip.File) (string, error) {
|
||||
file, err := item.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer file.Close()
|
||||
hash := sha256.New()
|
||||
if _, err = io.Copy(hash, file); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(hash.Sum(nil)), nil
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,593 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pages
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/upload"
|
||||
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/storage"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Input Pages 项目创建/更新请求。
|
||||
type Input struct {
|
||||
Name string `json:"name"`
|
||||
Slug string `json:"slug"`
|
||||
Description string `json:"description"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SPAFallbackEnabled bool `json:"spa_fallback_enabled"`
|
||||
SPAFallbackPath string `json:"spa_fallback_path"`
|
||||
APIProxyEnabled bool `json:"api_proxy_enabled"`
|
||||
APIProxyPath string `json:"api_proxy_path"`
|
||||
APIProxyPass string `json:"api_proxy_pass"`
|
||||
APIProxyRewrite string `json:"api_proxy_rewrite"`
|
||||
RootDir string `json:"root_dir"`
|
||||
EntryFile string `json:"entry_file"`
|
||||
}
|
||||
|
||||
// DeploymentView Pages 部署视图。
|
||||
type DeploymentView struct {
|
||||
ID uint `json:"id"`
|
||||
ProjectID uint `json:"project_id"`
|
||||
DeploymentNumber int `json:"deployment_number"`
|
||||
Checksum string `json:"checksum"`
|
||||
Status string `json:"status"`
|
||||
FileCount int `json:"file_count"`
|
||||
TotalSize int64 `json:"total_size"`
|
||||
CreatedBy string `json:"created_by"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
ActivatedAt *time.Time `json:"activated_at"`
|
||||
}
|
||||
|
||||
// DeploymentFileView Pages 部署文件视图。
|
||||
type DeploymentFileView struct {
|
||||
ID uint `json:"id"`
|
||||
DeploymentID uint `json:"deployment_id"`
|
||||
Path string `json:"path"`
|
||||
Size int64 `json:"size"`
|
||||
Checksum string `json:"checksum"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// View Pages 项目视图。
|
||||
type View struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Slug string `json:"slug"`
|
||||
Description string `json:"description"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SPAFallbackEnabled bool `json:"spa_fallback_enabled"`
|
||||
SPAFallbackPath string `json:"spa_fallback_path"`
|
||||
APIProxyEnabled bool `json:"api_proxy_enabled"`
|
||||
APIProxyPath string `json:"api_proxy_path"`
|
||||
APIProxyPass string `json:"api_proxy_pass"`
|
||||
APIProxyRewrite string `json:"api_proxy_rewrite"`
|
||||
RootDir string `json:"root_dir"`
|
||||
EntryFile string `json:"entry_file"`
|
||||
ActiveDeploymentID *uint `json:"active_deployment_id"`
|
||||
ActiveDeployment *DeploymentView `json:"active_deployment,omitempty"`
|
||||
DeploymentCount int64 `json:"deployment_count"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ListProjects 列出全部 Pages 项目。
|
||||
func ListProjects(ctx context.Context) ([]View, error) {
|
||||
projects, err := model.ListPagesProjects(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]View, 0, len(projects))
|
||||
for _, project := range projects {
|
||||
view, err := buildProjectView(ctx, &project)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views = append(views, *view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// GetProject 获取 Pages 项目详情。
|
||||
func GetProject(ctx context.Context, id uint) (*View, error) {
|
||||
project, err := model.GetPagesProjectByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildProjectView(ctx, project)
|
||||
}
|
||||
|
||||
// CreateProject 创建 Pages 项目。
|
||||
func CreateProject(ctx context.Context, input Input) (*View, error) {
|
||||
project, err := buildProject(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.CreatePagesProjectRecord(ctx, project); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errPagesSlugExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return buildProjectView(ctx, project)
|
||||
}
|
||||
|
||||
// UpdateProject 更新 Pages 项目。
|
||||
func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) {
|
||||
project, err := model.GetPagesProjectByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
project, err = buildProject(project, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = db.DB(ctx).Model(project).Updates(map[string]any{
|
||||
"name": project.Name,
|
||||
"slug": project.Slug,
|
||||
"description": project.Description,
|
||||
"enabled": project.Enabled,
|
||||
"spa_fallback_enabled": project.SPAFallbackEnabled,
|
||||
"spa_fallback_path": project.SPAFallbackPath,
|
||||
"api_proxy_enabled": project.APIProxyEnabled,
|
||||
"api_proxy_path": project.APIProxyPath,
|
||||
"api_proxy_pass": project.APIProxyPass,
|
||||
"api_proxy_rewrite": project.APIProxyRewrite,
|
||||
"root_dir": project.RootDir,
|
||||
"entry_file": project.EntryFile,
|
||||
}).Error; err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errPagesSlugExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return buildProjectView(ctx, project)
|
||||
}
|
||||
|
||||
// DeleteProject 删除 Pages 项目。
|
||||
func DeleteProject(ctx context.Context, id uint) error {
|
||||
project, err := model.GetPagesProjectByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
routeCount, err := model.CountProxyRoutesByPagesProjectID(ctx, project.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if routeCount > 0 {
|
||||
return errors.New(errPagesDeleteReferenced)
|
||||
}
|
||||
deployments, err := model.ListPagesDeployments(ctx, project.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where(
|
||||
"deployment_id IN (?)",
|
||||
tx.Model(&model.PagesDeployment{}).Select("id").Where("project_id = ?", project.ID),
|
||||
).Delete(&model.PagesDeploymentFile{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("project_id = ?", project.ID).Delete(&model.PagesDeployment{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Delete(project).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for index := range deployments {
|
||||
removeDeploymentArtifact(ctx, &deployments[index])
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// ListProjectDeployments 列出项目的全部部署。
|
||||
func ListProjectDeployments(ctx context.Context, projectID uint) ([]DeploymentView, error) {
|
||||
if _, err := model.GetPagesProjectByID(ctx, projectID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deployments, err := model.ListPagesDeployments(ctx, projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]DeploymentView, 0, len(deployments))
|
||||
for _, deployment := range deployments {
|
||||
views = append(views, buildDeploymentView(&deployment))
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// ListDeploymentFiles 列出部署文件清单。
|
||||
func ListDeploymentFiles(ctx context.Context, deploymentID uint) ([]DeploymentFileView, error) {
|
||||
if _, err := model.GetPagesDeploymentByID(ctx, deploymentID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
files, err := model.ListPagesDeploymentFiles(ctx, deploymentID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]DeploymentFileView, 0, len(files))
|
||||
for _, file := range files {
|
||||
views = append(views, DeploymentFileView{
|
||||
ID: file.ID,
|
||||
DeploymentID: file.DeploymentID,
|
||||
Path: file.Path,
|
||||
Size: file.Size,
|
||||
Checksum: file.Checksum,
|
||||
CreatedAt: file.CreatedAt,
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// UploadDeployment 上传 Pages 部署包。
|
||||
func UploadDeployment(ctx context.Context, projectID uint, fileHeader *multipart.FileHeader, createdBy string) (*DeploymentView, error) {
|
||||
project, err := model.GetPagesProjectByID(ctx, projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if fileHeader == nil {
|
||||
return nil, errors.New(errPagesPackageMissing)
|
||||
}
|
||||
if !strings.EqualFold(filepath.Ext(fileHeader.Filename), ".zip") {
|
||||
return nil, errors.New(errPagesPackageNotZip)
|
||||
}
|
||||
rootDir, err := validateAndNormalizePagesRootDir(project.RootDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
entryFile := normalizePagesEntryFile(project.EntryFile)
|
||||
tempPath, checksum, packageSize, err := persistPagesUploadTemp(fileHeader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer os.Remove(tempPath)
|
||||
manifest, err := inspectPagesZip(tempPath, rootDir, entryFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ingestResult, err := ingestPagesDeploymentPackage(
|
||||
ctx,
|
||||
tempPath,
|
||||
checksum,
|
||||
packageSize,
|
||||
project.Slug,
|
||||
fileHeader.Filename,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ingestCommitted := false
|
||||
defer func() {
|
||||
if !ingestCommitted && ingestResult.Created {
|
||||
_, _ = upload.Remove(ctx, ingestResult.Upload.ID)
|
||||
}
|
||||
}()
|
||||
deployment := &model.PagesDeployment{}
|
||||
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var maxNumber int
|
||||
if err := tx.Model(&model.PagesDeployment{}).
|
||||
Where("project_id = ?", project.ID).
|
||||
Select("COALESCE(MAX(deployment_number), 0)").
|
||||
Scan(&maxNumber).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
deployment = &model.PagesDeployment{
|
||||
ProjectID: project.ID,
|
||||
DeploymentNumber: maxNumber + 1,
|
||||
Checksum: checksum,
|
||||
Status: model.PagesDeploymentStatusUploaded,
|
||||
UploadID: ingestResult.Upload.ID,
|
||||
FileCount: manifest.FileCount,
|
||||
TotalSize: manifest.TotalSize,
|
||||
CreatedBy: strings.TrimSpace(createdBy),
|
||||
}
|
||||
if err := tx.Create(deployment).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for index := range manifest.Files {
|
||||
manifest.Files[index].DeploymentID = deployment.ID
|
||||
}
|
||||
if len(manifest.Files) > 0 {
|
||||
if err := tx.Create(&manifest.Files).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ingestCommitted = true
|
||||
view := buildDeploymentView(deployment)
|
||||
return &view, nil
|
||||
}
|
||||
|
||||
// ActivateDeployment 激活 Pages 部署。
|
||||
func ActivateDeployment(ctx context.Context, projectID uint, deploymentID uint) (*View, error) {
|
||||
project, err := model.GetPagesProjectByID(ctx, projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if deployment.ProjectID != project.ID {
|
||||
return nil, errors.New(errPagesDeploymentMismatch)
|
||||
}
|
||||
now := time.Now()
|
||||
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&model.PagesDeployment{}).
|
||||
Where("project_id = ?", project.ID).
|
||||
Update("status", model.PagesDeploymentStatusUploaded).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Model(deployment).Updates(map[string]any{
|
||||
"status": model.PagesDeploymentStatusActive,
|
||||
"activated_at": &now,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(project).Updates(map[string]any{
|
||||
"active_deployment_id": deployment.ID,
|
||||
}).Error
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetProject(ctx, project.ID)
|
||||
}
|
||||
|
||||
// OpenDeploymentPackage opens the deployment artifact from the upload storage framework.
|
||||
func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (*storage.Object, string, error) {
|
||||
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
if err = ensureDeploymentInActiveSnapshot(ctx, deployment.ID); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
fileName := fmt.Sprintf("pages-deployment-%d.zip", deployment.ID)
|
||||
if deployment.UploadID > 0 {
|
||||
uploadRecord, err := repository.GetActiveUploadByID(ctx, deployment.UploadID)
|
||||
if err != nil {
|
||||
return nil, "", errors.New(errPagesPackageUploadMissing)
|
||||
}
|
||||
obj, err := uploadstorage.OpenStoredObject(ctx, &uploadRecord)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
|
||||
}
|
||||
if obj.ContentType == "" {
|
||||
obj.ContentType = "application/zip"
|
||||
}
|
||||
return obj, fileName, nil
|
||||
}
|
||||
if strings.TrimSpace(deployment.ArtifactPath) == "" {
|
||||
return nil, "", errors.New(errPagesPackagePathEmpty)
|
||||
}
|
||||
file, err := os.Open(deployment.ArtifactPath)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
|
||||
}
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
|
||||
}
|
||||
return &storage.Object{
|
||||
Body: file,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "application/zip",
|
||||
}, fileName, nil
|
||||
}
|
||||
|
||||
func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) error {
|
||||
version, err := model.GetActiveConfigVersion(ctx)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return errors.New(errPagesPackageNotInActiveConfig)
|
||||
}
|
||||
return err
|
||||
}
|
||||
routes, err := parseSnapshotRoutes(version.SnapshotJSON)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, route := range routes {
|
||||
if route.UpstreamType != "pages" || route.PagesDeployment == nil {
|
||||
continue
|
||||
}
|
||||
if route.PagesDeployment.DeploymentID == deploymentID {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return errors.New(errPagesPackageNotInActiveConfig)
|
||||
}
|
||||
|
||||
type snapshotPagesDeployment struct {
|
||||
DeploymentID uint `json:"deployment_id"`
|
||||
}
|
||||
|
||||
type snapshotRouteRef struct {
|
||||
UpstreamType string `json:"upstream_type"`
|
||||
PagesDeployment *snapshotPagesDeployment `json:"pages_deployment"`
|
||||
}
|
||||
|
||||
func parseSnapshotRoutes(snapshotJSON string) ([]snapshotRouteRef, error) {
|
||||
text := strings.TrimSpace(snapshotJSON)
|
||||
if text == "" {
|
||||
return []snapshotRouteRef{}, nil
|
||||
}
|
||||
if strings.HasPrefix(text, "[") {
|
||||
var routes []snapshotRouteRef
|
||||
if err := json.Unmarshal([]byte(text), &routes); err != nil {
|
||||
return nil, errors.New(errPagesInvalidSnapshotFormat)
|
||||
}
|
||||
return routes, nil
|
||||
}
|
||||
var snapshot struct {
|
||||
Routes []snapshotRouteRef `json:"routes"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
|
||||
return nil, errors.New(errPagesInvalidSnapshotFormat)
|
||||
}
|
||||
return snapshot.Routes, nil
|
||||
}
|
||||
|
||||
// DeleteDeployment 删除 Pages 部署。
|
||||
func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error {
|
||||
project, err := model.GetPagesProjectByID(ctx, projectID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if deployment.ProjectID != project.ID {
|
||||
return errors.New(errPagesDeploymentMismatch)
|
||||
}
|
||||
if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID == deployment.ID {
|
||||
return errors.New(errPagesDeleteActiveDeploy)
|
||||
}
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("deployment_id = ?", deployment.ID).Delete(&model.PagesDeploymentFile{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Delete(deployment).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
removeDeploymentArtifact(ctx, deployment)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func buildProject(existing *model.PagesProject, input Input) (*model.PagesProject, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, errors.New(errPagesNameRequired)
|
||||
}
|
||||
slug := normalizePagesSlug(input.Slug)
|
||||
if slug == "" {
|
||||
slug = normalizePagesSlug(name)
|
||||
}
|
||||
if !pagesSlugPattern.MatchString(slug) {
|
||||
return nil, errors.New(errPagesSlugInvalid)
|
||||
}
|
||||
if existing == nil {
|
||||
existing = &model.PagesProject{}
|
||||
}
|
||||
existing.Name = name
|
||||
existing.Slug = slug
|
||||
existing.Description = strings.TrimSpace(input.Description)
|
||||
existing.Enabled = input.Enabled
|
||||
existing.SPAFallbackEnabled = input.SPAFallbackEnabled
|
||||
fallbackPath, err := normalizePagesFallbackPath(input.SPAFallbackPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
existing.SPAFallbackPath = fallbackPath
|
||||
|
||||
existing.APIProxyEnabled = input.APIProxyEnabled
|
||||
apiProxyPath := strings.TrimSpace(input.APIProxyPath)
|
||||
apiProxyPass := strings.TrimSpace(input.APIProxyPass)
|
||||
apiProxyRewrite := strings.TrimSpace(input.APIProxyRewrite)
|
||||
|
||||
if existing.APIProxyEnabled {
|
||||
if apiProxyPath == "" {
|
||||
return nil, errors.New(errPagesAPIProxyPathRequired)
|
||||
}
|
||||
if !strings.HasPrefix(apiProxyPath, "/") {
|
||||
return nil, errors.New(errPagesAPIProxyPathPrefix)
|
||||
}
|
||||
if apiProxyPass == "" {
|
||||
return nil, errors.New(errPagesAPIProxyPassRequired)
|
||||
}
|
||||
parsedURL, err := url.Parse(apiProxyPass)
|
||||
if err != nil || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") || parsedURL.Host == "" {
|
||||
return nil, errors.New(errPagesAPIProxyPassInvalid)
|
||||
}
|
||||
}
|
||||
existing.APIProxyPath = apiProxyPath
|
||||
existing.APIProxyPass = apiProxyPass
|
||||
existing.APIProxyRewrite = apiProxyRewrite
|
||||
|
||||
rootDir, err := validateAndNormalizePagesRootDir(input.RootDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
existing.RootDir = rootDir
|
||||
existing.EntryFile = normalizePagesEntryFile(input.EntryFile)
|
||||
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
func buildProjectView(ctx context.Context, project *model.PagesProject) (*View, error) {
|
||||
if project == nil {
|
||||
return nil, errors.New(errPagesProjectNotFound)
|
||||
}
|
||||
view := &View{
|
||||
ID: project.ID,
|
||||
Name: project.Name,
|
||||
Slug: project.Slug,
|
||||
Description: project.Description,
|
||||
Enabled: project.Enabled,
|
||||
SPAFallbackEnabled: project.SPAFallbackEnabled,
|
||||
SPAFallbackPath: normalizeStoredPagesFallbackPath(project.SPAFallbackPath),
|
||||
APIProxyEnabled: project.APIProxyEnabled,
|
||||
APIProxyPath: project.APIProxyPath,
|
||||
APIProxyPass: project.APIProxyPass,
|
||||
APIProxyRewrite: project.APIProxyRewrite,
|
||||
RootDir: project.RootDir,
|
||||
EntryFile: project.EntryFile,
|
||||
ActiveDeploymentID: project.ActiveDeploymentID,
|
||||
CreatedAt: project.CreatedAt,
|
||||
UpdatedAt: project.UpdatedAt,
|
||||
}
|
||||
count, err := model.CountPagesDeploymentsByProjectID(ctx, project.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
view.DeploymentCount = count
|
||||
if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID != 0 {
|
||||
deployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID)
|
||||
if err == nil {
|
||||
active := buildDeploymentView(deployment)
|
||||
view.ActiveDeployment = &active
|
||||
}
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func buildDeploymentView(deployment *model.PagesDeployment) DeploymentView {
|
||||
if deployment == nil {
|
||||
return DeploymentView{}
|
||||
}
|
||||
return DeploymentView{
|
||||
ID: deployment.ID,
|
||||
ProjectID: deployment.ProjectID,
|
||||
DeploymentNumber: deployment.DeploymentNumber,
|
||||
Checksum: deployment.Checksum,
|
||||
Status: deployment.Status,
|
||||
FileCount: deployment.FileCount,
|
||||
TotalSize: deployment.TotalSize,
|
||||
CreatedBy: deployment.CreatedBy,
|
||||
CreatedAt: deployment.CreatedAt,
|
||||
ActivatedAt: deployment.ActivatedAt,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,257 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pages
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/storage"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupPagesTestDB(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.User{},
|
||||
&model.Upload{},
|
||||
&model.UploadStat{},
|
||||
&model.TaskExecution{},
|
||||
&model.PagesProject{},
|
||||
&model.PagesDeployment{},
|
||||
&model.PagesDeploymentFile{},
|
||||
&model.ConfigVersion{},
|
||||
))
|
||||
require.NoError(t, sqliteDB.Create(&model.User{
|
||||
ID: 999,
|
||||
Username: "system",
|
||||
Password: "*",
|
||||
Nickname: "系统",
|
||||
IsActive: true,
|
||||
}).Error)
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func setupPagesStorageMock(t *testing.T) (restore func(), disable func()) {
|
||||
t.Helper()
|
||||
mockFiles := make(map[string][]byte)
|
||||
restore = storage.MockStorage(
|
||||
func(_ context.Context, key string, body io.Reader, _ int64, _ string) error {
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
mockFiles[key] = data
|
||||
return nil
|
||||
},
|
||||
func(_ context.Context, key string) (*storage.Object, error) {
|
||||
data, ok := mockFiles[key]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return &storage.Object{
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/zip",
|
||||
}, nil
|
||||
},
|
||||
func(_ context.Context, key string) error {
|
||||
delete(mockFiles, key)
|
||||
return nil
|
||||
},
|
||||
)
|
||||
storage.IsEnabledFunc = func() bool { return true }
|
||||
storage.ResetCache()
|
||||
disable = func() {
|
||||
storage.IsEnabledFunc = func() bool { return false }
|
||||
storage.ResetCache()
|
||||
restore()
|
||||
}
|
||||
return restore, disable
|
||||
}
|
||||
|
||||
func TestCreateProject(t *testing.T) {
|
||||
cleanup := setupPagesTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
project, err := CreateProject(ctx, Input{
|
||||
Name: "Marketing Site",
|
||||
Slug: "marketing-site",
|
||||
Description: "public site",
|
||||
Enabled: true,
|
||||
SPAFallbackEnabled: true,
|
||||
SPAFallbackPath: "/index.html",
|
||||
EntryFile: "index.html",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, project.ID)
|
||||
assert.Equal(t, "Marketing Site", project.Name)
|
||||
assert.Equal(t, "marketing-site", project.Slug)
|
||||
assert.Equal(t, "public site", project.Description)
|
||||
assert.True(t, project.Enabled)
|
||||
assert.True(t, project.SPAFallbackEnabled)
|
||||
assert.Equal(t, "/index.html", project.SPAFallbackPath)
|
||||
assert.Equal(t, "index.html", project.EntryFile)
|
||||
assert.Equal(t, int64(0), project.DeploymentCount)
|
||||
|
||||
_, err = CreateProject(ctx, Input{
|
||||
Name: "Duplicate Slug",
|
||||
Slug: "marketing-site",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, errPagesSlugExists, err.Error())
|
||||
}
|
||||
|
||||
func TestCreateProjectRejectsUnsafeFallbackPath(t *testing.T) {
|
||||
cleanup := setupPagesTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := CreateProject(ctx, Input{
|
||||
Name: "Unsafe Fallback",
|
||||
Slug: "unsafe-fallback",
|
||||
Enabled: true,
|
||||
SPAFallbackEnabled: true,
|
||||
SPAFallbackPath: "/index.html; proxy_pass http://evil",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "回退路径")
|
||||
}
|
||||
|
||||
func TestUploadDeploymentStoresPackageInUploadFramework(t *testing.T) {
|
||||
cleanup := setupPagesTestDB(t)
|
||||
defer cleanup()
|
||||
_, disableStorage := setupPagesStorageMock(t)
|
||||
defer disableStorage()
|
||||
ctx := context.Background()
|
||||
|
||||
project, err := CreateProject(ctx, Input{
|
||||
Name: "Upload Framework Site",
|
||||
Slug: "upload-framework-site",
|
||||
Enabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deployment, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "site.zip", testPagesZip(t, map[string]string{
|
||||
"index.html": "ok",
|
||||
})), "root")
|
||||
require.NoError(t, err)
|
||||
|
||||
storedDeployment, err := model.GetPagesDeploymentByID(ctx, deployment.ID)
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, storedDeployment.UploadID)
|
||||
assert.Empty(t, storedDeployment.ArtifactPath)
|
||||
|
||||
var uploadCount int64
|
||||
require.NoError(t, db.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error)
|
||||
assert.Equal(t, int64(1), uploadCount)
|
||||
}
|
||||
|
||||
func TestOpenDeploymentPackageRequiresActiveConfigSnapshot(t *testing.T) {
|
||||
cleanup := setupPagesTestDB(t)
|
||||
defer cleanup()
|
||||
_, disableStorage := setupPagesStorageMock(t)
|
||||
defer disableStorage()
|
||||
ctx := context.Background()
|
||||
|
||||
project, err := CreateProject(ctx, Input{
|
||||
Name: "Published Site",
|
||||
Slug: "published-site",
|
||||
Enabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
deployment, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "site.zip", testPagesZip(t, map[string]string{
|
||||
"index.html": "ok",
|
||||
})), "root")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ActivateDeployment(ctx, project.ID, deployment.ID)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, _, err = OpenDeploymentPackage(ctx, deployment.ID)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "激活配置")
|
||||
|
||||
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
|
||||
Version: "v2026-001",
|
||||
SnapshotJSON: fmt.Sprintf(`{"routes":[{"upstream_type":"pages","pages_deployment":{"deployment_id":%d}}]}`, deployment.ID),
|
||||
MainConfig: "",
|
||||
RenderedConfig: "",
|
||||
SupportFilesJSON: "[]",
|
||||
Checksum: "test-checksum",
|
||||
IsActive: true,
|
||||
CreatedBy: "test",
|
||||
}).Error)
|
||||
|
||||
packageObj, fileName, err := OpenDeploymentPackage(ctx, deployment.ID)
|
||||
require.NoError(t, err)
|
||||
defer packageObj.Body.Close()
|
||||
assert.Equal(t, fmt.Sprintf("pages-deployment-%d.zip", deployment.ID), fileName)
|
||||
|
||||
body, err := io.ReadAll(packageObj.Body)
|
||||
require.NoError(t, err)
|
||||
reader, err := zip.NewReader(bytes.NewReader(body), int64(len(body)))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, reader.File, 1)
|
||||
assert.Equal(t, "index.html", reader.File[0].Name)
|
||||
}
|
||||
|
||||
func testPagesZip(t *testing.T, files map[string]string) []byte {
|
||||
t.Helper()
|
||||
|
||||
var buffer bytes.Buffer
|
||||
writer := zip.NewWriter(&buffer)
|
||||
for name, content := range files {
|
||||
file, err := writer.Create(name)
|
||||
require.NoError(t, err)
|
||||
_, err = file.Write([]byte(content))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
require.NoError(t, writer.Close())
|
||||
return buffer.Bytes()
|
||||
}
|
||||
|
||||
func testPagesMultipartFile(t *testing.T, fileName string, content []byte) *multipart.FileHeader {
|
||||
t.Helper()
|
||||
|
||||
var body bytes.Buffer
|
||||
writer := multipart.NewWriter(&body)
|
||||
part, err := writer.CreateFormFile("package", fileName)
|
||||
require.NoError(t, err)
|
||||
_, err = part.Write(content)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, writer.Close())
|
||||
|
||||
req := httptest.NewRequest("POST", "/", &body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
require.NoError(t, req.ParseMultipartForm(int64(len(content))+1024))
|
||||
|
||||
file, header, err := req.FormFile("package")
|
||||
require.NoError(t, err)
|
||||
file.Close()
|
||||
return header
|
||||
}
|
||||
@@ -0,0 +1,310 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pages
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return apiutil.AbortNotFoundIfMissing(c, err, errPagesProjectNotFound)
|
||||
}
|
||||
|
||||
func deploymentIDParam(c *gin.Context) (uint, bool) {
|
||||
raw := c.Param("deployment_id")
|
||||
if raw == "" {
|
||||
response.AbortBadRequest(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
id64, err := strconv.ParseUint(raw, 10, 64)
|
||||
if err != nil || id64 == 0 {
|
||||
response.AbortBadRequest(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
return uint(id64), true
|
||||
}
|
||||
|
||||
// ListProjectsHandler 列出全部 Pages 项目。
|
||||
// @Summary 列出 Pages 项目
|
||||
// @Description 返回全部 OpenFlare Pages 项目,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]pages.View} "Pages 项目列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages [get]
|
||||
func ListProjectsHandler(c *gin.Context) {
|
||||
projects, err := ListProjects(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(projects))
|
||||
}
|
||||
|
||||
// GetProjectHandler 获取 Pages 项目详情。
|
||||
// @Summary 获取 Pages 项目详情
|
||||
// @Description 按 ID 返回 Pages 项目详情,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "项目 ID"
|
||||
// @Success 200 {object} response.Any{data=pages.View} "Pages 项目详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "项目不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/{id} [get]
|
||||
func GetProjectHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
project, err := GetProject(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(project))
|
||||
}
|
||||
|
||||
// CreateProjectHandler 创建 Pages 项目。
|
||||
// @Summary 创建 Pages 项目
|
||||
// @Description 创建新的 OpenFlare Pages 项目,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body pages.Input true "项目参数"
|
||||
// @Success 200 {object} response.Any{data=pages.View} "创建成功的项目"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages [post]
|
||||
func CreateProjectHandler(c *gin.Context) {
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
project, err := CreateProject(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(project))
|
||||
}
|
||||
|
||||
// UpdateProjectHandler 更新 Pages 项目。
|
||||
// @Summary 更新 Pages 项目
|
||||
// @Description 按 ID 更新 OpenFlare Pages 项目,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "项目 ID"
|
||||
// @Param request body pages.Input true "项目参数"
|
||||
// @Success 200 {object} response.Any{data=pages.View} "更新后的项目"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "项目不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/{id}/update [post]
|
||||
func UpdateProjectHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
project, err := UpdateProject(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(project))
|
||||
}
|
||||
|
||||
// DeleteProjectHandler 删除 Pages 项目。
|
||||
// @Summary 删除 Pages 项目
|
||||
// @Description 按 ID 删除 OpenFlare Pages 项目,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "项目 ID"
|
||||
// @Success 200 {object} response.Any "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "项目不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/{id}/delete [post]
|
||||
func DeleteProjectHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteProject(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// ListDeploymentsHandler 列出项目的全部部署。
|
||||
// @Summary 列出 Pages 部署
|
||||
// @Description 返回指定项目的全部部署记录,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "项目 ID"
|
||||
// @Success 200 {object} response.Any{data=[]pages.DeploymentView} "部署列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "项目不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/{id}/deployments [get]
|
||||
func ListDeploymentsHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
deployments, err := ListProjectDeployments(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(deployments))
|
||||
}
|
||||
|
||||
// UploadDeploymentHandler 上传 Pages 部署包。
|
||||
// @Summary 上传 Pages 部署包
|
||||
// @Description 为指定项目上传 ZIP 部署包,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Accept multipart/form-data
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "项目 ID"
|
||||
// @Param package formData file true "部署包 ZIP 文件"
|
||||
// @Success 200 {object} response.Any{data=pages.DeploymentView} "部署记录"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "项目不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/{id}/deployments/upload [post]
|
||||
func UploadDeploymentHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
file, err := c.FormFile("package")
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errPagesPackageMissing)
|
||||
return
|
||||
}
|
||||
deployment, err := UploadDeployment(c.Request.Context(), id, file, "")
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(deployment))
|
||||
}
|
||||
|
||||
// ActivateDeploymentHandler 激活 Pages 部署。
|
||||
// @Summary 激活 Pages 部署
|
||||
// @Description 将指定部署设为项目当前生效版本,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "项目 ID"
|
||||
// @Param deployment_id path int true "部署 ID"
|
||||
// @Success 200 {object} response.Any{data=pages.View} "激活后的项目"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "项目或部署不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/{id}/deployments/{deployment_id}/activate [post]
|
||||
func ActivateDeploymentHandler(c *gin.Context) {
|
||||
projectID, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
deploymentID, ok := deploymentIDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
project, err := ActivateDeployment(c.Request.Context(), projectID, deploymentID)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(project))
|
||||
}
|
||||
|
||||
// DeleteDeploymentHandler 删除 Pages 部署。
|
||||
// @Summary 删除 Pages 部署
|
||||
// @Description 删除指定项目的部署记录,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "项目 ID"
|
||||
// @Param deployment_id path int true "部署 ID"
|
||||
// @Success 200 {object} response.Any "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "项目或部署不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/{id}/deployments/{deployment_id}/delete [post]
|
||||
func DeleteDeploymentHandler(c *gin.Context) {
|
||||
projectID, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
deploymentID, ok := deploymentIDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteDeployment(c.Request.Context(), projectID, deploymentID); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// ListDeploymentFilesHandler 列出部署文件清单。
|
||||
// @Summary 列出 Pages 部署文件
|
||||
// @Description 返回指定部署包含的文件清单,需要管理员权限
|
||||
// @Tags openflare-pages
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param deployment_id path int true "部署 ID"
|
||||
// @Success 200 {object} response.Any{data=[]pages.DeploymentFileView} "部署文件列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 404 {object} response.Any "部署不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/pages/deployments/{deployment_id}/files [get]
|
||||
func ListDeploymentFilesHandler(c *gin.Context) {
|
||||
deploymentID, ok := deploymentIDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
files, err := ListDeploymentFiles(c.Request.Context(), deploymentID)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(files))
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
const (
|
||||
errProxyRouteNotFound = "proxy route not found"
|
||||
errProxyRouteIdentityExists = "proxy route identity already exists"
|
||||
errProxyRouteSiteNameExists = "site_name already exists"
|
||||
errProxyRouteDomainExists = "domain %s already exists"
|
||||
errProxyRouteSiteNameEmpty = "site_name cannot be empty"
|
||||
errProxyRouteDomainRequired = "at least one domain is required"
|
||||
errProxyRouteDomainInvalid = "domain format is invalid"
|
||||
errProxyRouteDomainMismatch = "domain must match domains[0]"
|
||||
errProxyRouteOriginEmpty = "origin_url cannot be empty"
|
||||
errProxyRouteOriginInvalid = "origin URL format is invalid"
|
||||
errProxyRouteOriginScheme = "origin URL must start with http:// or https://"
|
||||
errProxyRouteOriginHostInvalid = "origin_host format is invalid"
|
||||
errProxyRouteUpstreamRequired = "at least one upstream is required"
|
||||
errProxyRouteUpstreamScheme = "all upstreams must use the same scheme"
|
||||
errProxyRouteUpstreamPath = "multi-upstream mode does not support origin paths"
|
||||
errProxyRouteUpstreamQuery = "multi-upstream mode does not support origin query strings"
|
||||
errProxyRouteOriginNotFound = "selected origin does not exist"
|
||||
errProxyRouteCertNotFound = "selected certificate does not exist"
|
||||
errProxyRouteCertRequired = "must select a certificate when HTTPS is enabled"
|
||||
errProxyRouteCertDomainLength = "domain_cert_ids must match domains length"
|
||||
errProxyRouteRedirectHTTP = "redirect_http requires enable_https"
|
||||
errProxyRouteBasicAuth = "basic_auth_username and basic_auth_password cannot be empty when basic auth is enabled"
|
||||
errProxyRouteLimitRate = "limit_rate must be a number or use the 512k / 1m format"
|
||||
errProxyRouteCachePolicy = "cache policy is not supported"
|
||||
errProxyRouteCacheSuffix = "cache suffix format is invalid"
|
||||
errProxyRouteCachePath = "cache path rule format is invalid"
|
||||
errProxyRouteCacheSuffixReq = "at least one suffix is required"
|
||||
errProxyRouteCachePrefixReq = "at least one path prefix is required"
|
||||
errProxyRouteCacheExactReq = "at least one exact path is required"
|
||||
errProxyRouteHeaderKeyEmpty = "custom header key cannot be empty"
|
||||
errProxyRouteHeaderKeyInvalid = "custom header key format is invalid"
|
||||
errProxyRouteHeaderNewline = "custom headers cannot contain newlines"
|
||||
errProxyRouteTunnelNodeReq = "tunnel_node_id is required for tunnel upstream"
|
||||
errProxyRouteTunnelNodeMissing = "tunnel client node does not exist"
|
||||
errProxyRouteTunnelNodeType = "tunnel_node_id must reference a tunnel_client node"
|
||||
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
|
||||
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
|
||||
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
|
||||
errProxyRoutePagesNotFound = "Pages 项目不存在"
|
||||
errProxyRoutePagesDisabled = "Pages 项目未启用"
|
||||
errProxyRoutePagesNoDeploy = "Pages 项目没有激活部署"
|
||||
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
|
||||
errProxyRouteOriginPort = "端口格式不合法"
|
||||
errProxyRouteOriginPortEmpty = "端口不能为空"
|
||||
errProxyRouteOriginURI = "源站路径需以 / 或 ? 开头"
|
||||
errProxyRouteOriginURIProto = "源站路径不能包含协议"
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,430 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// CustomHeaderInput 自定义响应头。
|
||||
type CustomHeaderInput struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
|
||||
// Input 代理规则创建/更新请求。
|
||||
type Input struct {
|
||||
SiteName string `json:"site_name"`
|
||||
Domain string `json:"domain"`
|
||||
Domains []string `json:"domains"`
|
||||
OriginID *uint `json:"origin_id"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginScheme string `json:"origin_scheme"`
|
||||
OriginAddress string `json:"origin_address"`
|
||||
OriginPort string `json:"origin_port"`
|
||||
OriginURI string `json:"origin_uri"`
|
||||
OriginHost string `json:"origin_host"`
|
||||
Upstreams []string `json:"upstreams"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
CertIDs []uint `json:"cert_ids"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip"`
|
||||
LimitRate string `json:"limit_rate"`
|
||||
CacheEnabled bool `json:"cache_enabled"`
|
||||
CachePolicy string `json:"cache_policy"`
|
||||
CacheRules []string `json:"cache_rules"`
|
||||
CustomHeaders []CustomHeaderInput `json:"custom_headers"`
|
||||
BasicAuthEnabled bool `json:"basic_auth_enabled"`
|
||||
BasicAuthUsername string `json:"basic_auth_username"`
|
||||
BasicAuthPassword string `json:"basic_auth_password"`
|
||||
Remark string `json:"remark"`
|
||||
UpstreamType string `json:"upstream_type"`
|
||||
TunnelNodeID *uint `json:"tunnel_node_id"`
|
||||
TunnelID *uint `json:"tunnel_id"`
|
||||
TunnelTargetAddr string `json:"tunnel_target_addr"`
|
||||
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
|
||||
PagesProjectID *uint `json:"pages_project_id"`
|
||||
}
|
||||
|
||||
// View 代理规则视图。
|
||||
type View struct {
|
||||
ID uint `json:"id"`
|
||||
SiteName string `json:"site_name"`
|
||||
Domain string `json:"domain"`
|
||||
Domains []string `json:"domains"`
|
||||
PrimaryDomain string `json:"primary_domain"`
|
||||
DomainCount int `json:"domain_count"`
|
||||
OriginID *uint `json:"origin_id"`
|
||||
OriginURL string `json:"origin_url"`
|
||||
OriginHost string `json:"origin_host"`
|
||||
Upstreams string `json:"upstreams"`
|
||||
UpstreamList []string `json:"upstream_list"`
|
||||
Enabled bool `json:"enabled"`
|
||||
EnableHTTPS bool `json:"enable_https"`
|
||||
CertID *uint `json:"cert_id"`
|
||||
CertIDs []uint `json:"cert_ids"`
|
||||
DomainCertIDs []uint `json:"domain_cert_ids"`
|
||||
RedirectHTTP bool `json:"redirect_http"`
|
||||
LimitConnPerServer int `json:"limit_conn_per_server"`
|
||||
LimitConnPerIP int `json:"limit_conn_per_ip"`
|
||||
LimitRate string `json:"limit_rate"`
|
||||
CacheEnabled bool `json:"cache_enabled"`
|
||||
CachePolicy string `json:"cache_policy"`
|
||||
CacheRules string `json:"cache_rules"`
|
||||
CacheRuleList []string `json:"cache_rule_list"`
|
||||
CustomHeaders string `json:"custom_headers"`
|
||||
CustomHeaderList []CustomHeaderInput `json:"custom_header_list"`
|
||||
BasicAuthEnabled bool `json:"basic_auth_enabled"`
|
||||
BasicAuthUsername string `json:"basic_auth_username"`
|
||||
BasicAuthPassword string `json:"basic_auth_password"`
|
||||
Remark string `json:"remark"`
|
||||
UpstreamType string `json:"upstream_type"`
|
||||
TunnelNodeID *uint `json:"tunnel_node_id"`
|
||||
TunnelID *uint `json:"tunnel_id"`
|
||||
TunnelTargetAddr string `json:"tunnel_target_addr"`
|
||||
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
|
||||
PagesProjectID *uint `json:"pages_project_id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ListProxyRoutes 列出全部代理规则。
|
||||
func ListProxyRoutes(ctx context.Context) ([]*View, error) {
|
||||
routes, err := model.ListProxyRoutes(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildProxyRouteViews(ctx, routes)
|
||||
}
|
||||
|
||||
// GetProxyRoute 获取代理规则详情。
|
||||
func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
|
||||
route, err := model.GetProxyRouteByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildProxyRouteView(ctx, route)
|
||||
}
|
||||
|
||||
// CreateProxyRoute 创建代理规则。
|
||||
func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
|
||||
route, err := buildProxyRoute(ctx, nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.CreateProxyRouteRecord(ctx, route); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errProxyRouteIdentityExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return buildProxyRouteView(ctx, route)
|
||||
}
|
||||
|
||||
// UpdateProxyRoute 更新代理规则。
|
||||
func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error) {
|
||||
route, err := model.GetProxyRouteByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
route, err = buildProxyRoute(ctx, route, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = model.UpdateProxyRouteRecord(ctx, route); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
return nil, errors.New(errProxyRouteIdentityExists)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return buildProxyRouteView(ctx, route)
|
||||
}
|
||||
|
||||
// DeleteProxyRoute 删除代理规则。
|
||||
func DeleteProxyRoute(ctx context.Context, id uint) error {
|
||||
if _, err := model.GetProxyRouteByID(ctx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return model.DeleteProxyRouteRecord(ctx, id)
|
||||
}
|
||||
|
||||
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, error) {
|
||||
domains, err := normalizeProxyRouteDomainsInput(route, input.Domain, input.Domains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domain := domains[0]
|
||||
siteName := normalizeProxyRouteSiteNameInput(route, input.SiteName, domain)
|
||||
|
||||
upstreamType := normalizeUpstreamType(input.UpstreamType)
|
||||
var originURL string
|
||||
var originID *uint
|
||||
var upstreams []string
|
||||
|
||||
if upstreamType == "tunnel" {
|
||||
originURL = "http://127.0.0.1"
|
||||
upstreams = []string{originURL}
|
||||
} else if upstreamType == "pages" {
|
||||
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
originURL = "http://127.0.0.1"
|
||||
upstreams = []string{originURL}
|
||||
} else {
|
||||
originURL, originID, err = resolveProxyRoutePrimaryOrigin(ctx, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreams, err = normalizeUpstreams(originURL, input.Upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
originHost := strings.TrimSpace(input.OriginHost)
|
||||
remark := strings.TrimSpace(input.Remark)
|
||||
cachePolicy := strings.TrimSpace(input.CachePolicy)
|
||||
cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cacheRulesJSON, err := json.Marshal(cacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreamsJSON, err := json.Marshal(upstreams)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeadersJSON, err := json.Marshal(customHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !input.EnableHTTPS {
|
||||
input.RedirectHTTP = false
|
||||
input.CertID = nil
|
||||
input.CertIDs = nil
|
||||
input.DomainCertIDs = nil
|
||||
}
|
||||
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs(
|
||||
ctx,
|
||||
domains,
|
||||
input.EnableHTTPS,
|
||||
input.DomainCertIDs,
|
||||
input.CertID,
|
||||
input.CertIDs,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateProxyRouteDomainCertificateCoverage(ctx, domains, domainCertIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDsJSON, err := json.Marshal(certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainCertIDsJSON, err := json.Marshal(domainCertIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainsJSON, err := json.Marshal(domains)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := validateProxyRouteSiteName(siteName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateProxyRouteIdentityUniqueness(ctx, route, siteName, domains); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateOriginHost(originHost); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
input.DomainCertIDs = domainCertIDs
|
||||
input.CertIDs = certIDs
|
||||
input.CertID = primaryCertID
|
||||
if input.RedirectHTTP && !input.EnableHTTPS {
|
||||
return nil, errors.New(errProxyRouteRedirectHTTP)
|
||||
}
|
||||
|
||||
if input.BasicAuthEnabled {
|
||||
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
|
||||
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
|
||||
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
|
||||
return nil, errors.New(errProxyRouteBasicAuth)
|
||||
}
|
||||
} else {
|
||||
input.BasicAuthUsername = ""
|
||||
input.BasicAuthPassword = ""
|
||||
}
|
||||
|
||||
if route == nil {
|
||||
route = &model.ProxyRoute{}
|
||||
}
|
||||
route.SiteName = siteName
|
||||
route.Domain = domain
|
||||
route.Domains = string(domainsJSON)
|
||||
route.OriginID = originID
|
||||
route.OriginURL = upstreams[0]
|
||||
route.OriginHost = originHost
|
||||
route.Upstreams = string(upstreamsJSON)
|
||||
route.Enabled = input.Enabled
|
||||
route.EnableHTTPS = input.EnableHTTPS
|
||||
route.CertID = input.CertID
|
||||
route.CertIDs = string(certIDsJSON)
|
||||
route.DomainCertIDs = string(domainCertIDsJSON)
|
||||
route.RedirectHTTP = input.RedirectHTTP
|
||||
route.LimitConnPerServer = limitConnPerServer
|
||||
route.LimitConnPerIP = limitConnPerIP
|
||||
route.LimitRate = limitRate
|
||||
route.CacheEnabled = input.CacheEnabled
|
||||
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
|
||||
route.CacheRules = string(cacheRulesJSON)
|
||||
route.CustomHeaders = string(customHeadersJSON)
|
||||
route.BasicAuthEnabled = input.BasicAuthEnabled
|
||||
route.BasicAuthUsername = input.BasicAuthUsername
|
||||
route.BasicAuthPassword = input.BasicAuthPassword
|
||||
route.Remark = remark
|
||||
route.UpstreamType = upstreamType
|
||||
if upstreamType == "tunnel" {
|
||||
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
route.TunnelNodeID = tunnelNodeID
|
||||
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
|
||||
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
|
||||
route.PagesProjectID = nil
|
||||
} else if upstreamType == "pages" {
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = input.PagesProjectID
|
||||
} else {
|
||||
route.TunnelNodeID = nil
|
||||
route.TunnelTargetAddr = ""
|
||||
route.TunnelTargetProtocol = ""
|
||||
route.PagesProjectID = nil
|
||||
}
|
||||
return route, nil
|
||||
}
|
||||
|
||||
func buildProxyRouteViews(ctx context.Context, routes []*model.ProxyRoute) ([]*View, error) {
|
||||
views := make([]*View, 0, len(routes))
|
||||
for _, route := range routes {
|
||||
view, err := buildProxyRouteView(ctx, route)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views = append(views, view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, error) {
|
||||
if route == nil {
|
||||
return nil, errors.New("proxy route is nil")
|
||||
}
|
||||
domains, err := decodeStoredDomains(route.Domains, route.Domain)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domainCertIDs, err := resolveProxyRouteDomainCertIDs(ctx, route, domains, certIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var certID *uint
|
||||
if len(certIDs) > 0 {
|
||||
certID = &certIDs[0]
|
||||
}
|
||||
primaryDomain := domains[0]
|
||||
return &View{
|
||||
ID: route.ID,
|
||||
SiteName: normalizeProxyRouteSiteNameInput(route, route.SiteName, primaryDomain),
|
||||
Domain: primaryDomain,
|
||||
Domains: domains,
|
||||
PrimaryDomain: primaryDomain,
|
||||
DomainCount: len(domains),
|
||||
OriginID: route.OriginID,
|
||||
OriginURL: route.OriginURL,
|
||||
OriginHost: route.OriginHost,
|
||||
Upstreams: route.Upstreams,
|
||||
UpstreamList: upstreams,
|
||||
Enabled: route.Enabled,
|
||||
EnableHTTPS: route.EnableHTTPS,
|
||||
CertID: certID,
|
||||
CertIDs: certIDs,
|
||||
DomainCertIDs: domainCertIDs,
|
||||
RedirectHTTP: route.RedirectHTTP,
|
||||
LimitConnPerServer: route.LimitConnPerServer,
|
||||
LimitConnPerIP: route.LimitConnPerIP,
|
||||
LimitRate: route.LimitRate,
|
||||
CacheEnabled: route.CacheEnabled,
|
||||
CachePolicy: route.CachePolicy,
|
||||
CacheRules: route.CacheRules,
|
||||
CacheRuleList: cacheRules,
|
||||
CustomHeaders: route.CustomHeaders,
|
||||
CustomHeaderList: customHeaders,
|
||||
BasicAuthEnabled: route.BasicAuthEnabled,
|
||||
BasicAuthUsername: route.BasicAuthUsername,
|
||||
BasicAuthPassword: route.BasicAuthPassword,
|
||||
Remark: route.Remark,
|
||||
UpstreamType: route.UpstreamType,
|
||||
TunnelNodeID: route.TunnelNodeID,
|
||||
TunnelID: route.TunnelNodeID,
|
||||
TunnelTargetAddr: route.TunnelTargetAddr,
|
||||
TunnelTargetProtocol: route.TunnelTargetProtocol,
|
||||
PagesProjectID: route.PagesProjectID,
|
||||
CreatedAt: route.CreatedAt,
|
||||
UpdatedAt: route.UpdatedAt,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupProxyRouteTestDB(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.ProxyRoute{}, &model.Origin{}))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateProxyRoute(t *testing.T) {
|
||||
cleanup := setupProxyRouteTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
view, err := CreateProxyRoute(ctx, Input{
|
||||
SiteName: "example-site",
|
||||
Domain: "example.com",
|
||||
OriginURL: "http://origin.example.com:8080",
|
||||
Enabled: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, view.ID)
|
||||
assert.Equal(t, "example-site", view.SiteName)
|
||||
assert.Equal(t, "example.com", view.Domain)
|
||||
assert.Equal(t, []string{"example.com"}, view.Domains)
|
||||
assert.Equal(t, "http://origin.example.com:8080", view.OriginURL)
|
||||
assert.Equal(t, []string{"http://origin.example.com:8080"}, view.UpstreamList)
|
||||
assert.True(t, view.Enabled)
|
||||
|
||||
_, err = CreateProxyRoute(ctx, Input{
|
||||
SiteName: "duplicate-site",
|
||||
Domain: "example.com",
|
||||
OriginURL: "http://origin.example.com:8080",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
}
|
||||
|
||||
func TestListProxyRoutes(t *testing.T) {
|
||||
cleanup := setupProxyRouteTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
first, err := CreateProxyRoute(ctx, Input{
|
||||
SiteName: "first-site",
|
||||
Domain: "first.example.com",
|
||||
OriginURL: "http://origin-a.internal:80",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
second, err := CreateProxyRoute(ctx, Input{
|
||||
SiteName: "second-site",
|
||||
Domain: "second.example.com",
|
||||
OriginURL: "http://origin-b.internal:80",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
routes, err := ListProxyRoutes(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, routes, 2)
|
||||
assert.Equal(t, second.ID, routes[0].ID)
|
||||
assert.Equal(t, first.ID, routes[1].ID)
|
||||
assert.Equal(t, "second.example.com", routes[0].Domain)
|
||||
assert.Equal(t, "first.example.com", routes[1].Domain)
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package proxy_route
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return apiutil.AbortNotFoundIfMissing(c, err, errProxyRouteNotFound)
|
||||
}
|
||||
|
||||
// GetProxyRoutes 列出全部代理规则。
|
||||
// @Summary 获取代理规则列表
|
||||
// @Description 返回所有代理规则配置,需要管理员权限
|
||||
// @Tags openflare-proxy-route
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]proxy_route.View} "代理规则列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/proxy-routes [get]
|
||||
func GetProxyRoutes(c *gin.Context) {
|
||||
routes, err := ListProxyRoutes(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(routes))
|
||||
}
|
||||
|
||||
// GetProxyRouteHandler 获取代理规则详情。
|
||||
// @Summary 获取代理规则详情
|
||||
// @Description 返回指定代理规则的完整配置,需要管理员权限
|
||||
// @Tags openflare-proxy-route
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "代理规则 ID"
|
||||
// @Success 200 {object} response.Any{data=proxy_route.View} "代理规则详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或规则不存在"
|
||||
// @Router /api/v1/d/proxy-routes/{id} [get]
|
||||
func GetProxyRouteHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
route, err := GetProxyRoute(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(route))
|
||||
}
|
||||
|
||||
// CreateProxyRouteHandler 创建代理规则。
|
||||
// @Summary 创建代理规则
|
||||
// @Description 创建新的反向代理规则,需要管理员权限
|
||||
// @Tags openflare-proxy-route
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param body body proxy_route.Input true "代理规则参数"
|
||||
// @Success 200 {object} response.Any{data=proxy_route.View} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Router /api/v1/d/proxy-routes [post]
|
||||
func CreateProxyRouteHandler(c *gin.Context) {
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
route, err := CreateProxyRoute(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(route))
|
||||
}
|
||||
|
||||
// UpdateProxyRouteHandler 更新代理规则。
|
||||
// @Summary 更新代理规则
|
||||
// @Description 更新指定代理规则的配置,需要管理员权限
|
||||
// @Tags openflare-proxy-route
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "代理规则 ID"
|
||||
// @Param body body proxy_route.Input true "代理规则参数"
|
||||
// @Success 200 {object} response.Any{data=proxy_route.View} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或规则不存在"
|
||||
// @Router /api/v1/d/proxy-routes/{id}/update [post]
|
||||
func UpdateProxyRouteHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input Input
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
route, err := UpdateProxyRoute(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(route))
|
||||
}
|
||||
|
||||
// DeleteProxyRouteHandler 删除代理规则。
|
||||
// @Summary 删除代理规则
|
||||
// @Description 删除指定代理规则,需要管理员权限
|
||||
// @Tags openflare-proxy-route
|
||||
// @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/proxy-routes/{id}/delete [post]
|
||||
func DeleteProxyRouteHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteProxyRoute(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
const (
|
||||
errAgentTokenInvalid = "无权进行此操作,Agent Token 无效"
|
||||
errRelayNodeTypeMismatch = "此节点不是 TunnelRelay 类型"
|
||||
)
|
||||
@@ -0,0 +1,112 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
func normalizeRelayStatus(status string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(status)) {
|
||||
case "healthy":
|
||||
return "healthy"
|
||||
case "unhealthy":
|
||||
return "unhealthy"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeReleaseChannel(channel string) string {
|
||||
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
|
||||
return "preview"
|
||||
}
|
||||
return "stable"
|
||||
}
|
||||
|
||||
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(node *model.OpenFlareNode) *Config {
|
||||
if node == nil {
|
||||
return nil
|
||||
}
|
||||
return &Config{
|
||||
BindPort: node.RelayBindPort,
|
||||
VhostHTTPPort: node.RelayVhostHTTPPort,
|
||||
AuthToken: node.RelayAuthToken,
|
||||
LogLevel: "info",
|
||||
WebServerEnabled: node.RelayWebServerEnabled,
|
||||
}
|
||||
}
|
||||
|
||||
// BuildSettings returns runtime settings shared by relay and flared clients.
|
||||
func BuildSettings(node *model.OpenFlareNode, updateNow bool, updateChannel, updateTag string) *Settings {
|
||||
autoUpdate := false
|
||||
if node != nil {
|
||||
autoUpdate = node.AutoUpdateEnabled
|
||||
}
|
||||
if strings.TrimSpace(updateChannel) == "" {
|
||||
updateChannel = "stable"
|
||||
}
|
||||
return &Settings{
|
||||
HeartbeatInterval: model.AgentHeartbeatInterval,
|
||||
WebsocketUpgradeEnabled: model.AgentWebsocketUpgradeEnabled,
|
||||
AutoUpdate: autoUpdate,
|
||||
UpdateRepo: model.AgentUpdateRepo,
|
||||
UpdateNow: updateNow,
|
||||
UpdateChannel: updateChannel,
|
||||
UpdateTag: strings.TrimSpace(updateTag),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const nodeStatusOnline = "online"
|
||||
|
||||
// ProxyStat describes a single frps proxy reported by the relay.
|
||||
type ProxyStat struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Status string `json:"status"`
|
||||
ClientVersion string `json:"client_version"`
|
||||
LastStartTime string `json:"last_start_time"`
|
||||
LastCloseTime string `json:"last_close_time"`
|
||||
ClientAddr string `json:"client_addr"`
|
||||
}
|
||||
|
||||
// HeartbeatPayload is sent by OpenFlareRelay on each heartbeat.
|
||||
type HeartbeatPayload struct {
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"frp_version"`
|
||||
RelayStatus string `json:"relay_status"`
|
||||
FrpsConnCount int `json:"frps_connections"`
|
||||
FrpsProxyCount int `json:"frps_proxy_count"`
|
||||
FrpsClientCount int `json:"frps_client_count"`
|
||||
FrpsProxies []ProxyStat `json:"frps_proxies,omitempty"`
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
Profile *agent.NodeSystemProfile `json:"profile,omitempty"`
|
||||
Snapshot *agent.NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
HealthEvents []agent.NodeHealthEvent `json:"health_events,omitempty"`
|
||||
}
|
||||
|
||||
// Config is the frps configuration sent to the relay.
|
||||
type Config struct {
|
||||
BindPort int `json:"bind_port"`
|
||||
VhostHTTPPort int `json:"vhost_http_port"`
|
||||
AuthToken string `json:"auth_token"`
|
||||
LogLevel string `json:"log_level"`
|
||||
WebServerEnabled bool `json:"web_server_enabled"`
|
||||
}
|
||||
|
||||
// Settings contains runtime settings for relay and flared clients.
|
||||
type Settings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
UpdateRepo string `json:"update_repo"`
|
||||
UpdateNow bool `json:"update_now"`
|
||||
UpdateChannel string `json:"update_channel"`
|
||||
UpdateTag string `json:"update_tag"`
|
||||
}
|
||||
|
||||
// HeartbeatResponse is returned from a relay heartbeat.
|
||||
type HeartbeatResponse struct {
|
||||
RelayConfig *Config `json:"relay_config"`
|
||||
RelaySettings *Settings `json:"relay_settings"`
|
||||
}
|
||||
|
||||
// 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, fmt.Errorf("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": "stable",
|
||||
"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 !previous.UpdateRequested {
|
||||
delete(changes, "update_requested")
|
||||
}
|
||||
if previous.UpdateChannel == "stable" {
|
||||
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 = "stable"
|
||||
node.UpdateTag = ""
|
||||
lastSeen := now
|
||||
node.LastSeenAt = &lastSeen
|
||||
node.Status = nodeStatusOnline
|
||||
|
||||
if err := db.DB(ctx).Model(node).Updates(changes).Error; 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(node),
|
||||
RelaySettings: BuildSettings(node, updateNow, updateChannel, updateTag),
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,174 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func 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.OpenFlareOption{},
|
||||
&model.OpenFlareNodeSystemProfile{},
|
||||
&model.OpenFlareMetricSnapshot{},
|
||||
&model.OpenFlareHealthEvent{},
|
||||
&model.OpenFlareNodeObservationFrps{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
option.ResetInitializationForTest()
|
||||
agent.ResetAuthCacheForTest()
|
||||
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
option.ResetInitializationForTest()
|
||||
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,
|
||||
NetworkRxBytes: 1024,
|
||||
NetworkTxBytes: 2048,
|
||||
},
|
||||
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 := model.GetOpenFlareNodeSystemProfile(ctx, node.NodeID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "relay-runtime", profile.Hostname)
|
||||
assert.Equal(t, "Ubuntu", profile.OSName)
|
||||
|
||||
snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, snapshots, 1)
|
||||
assert.Equal(t, 12.5, snapshots[0].CPUUsagePercent)
|
||||
|
||||
frpsObs, err := model.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 := model.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 = model.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,49 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const ctxRelayNodeKey = "relay_node"
|
||||
|
||||
// RelayAuth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type.
|
||||
func RelayAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := strings.TrimSpace(c.GetHeader("X-Agent-Token"))
|
||||
node, err := 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()
|
||||
}
|
||||
}
|
||||
|
||||
func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
|
||||
if token == "" {
|
||||
return nil, errors.New("missing agent token")
|
||||
}
|
||||
node, err := model.GetOpenFlareNodeByAccessToken(ctx, token)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, errors.New("invalid agent token")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return node, nil
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"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, model.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", RelayAuth(), 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", RelayAuth(), 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", RelayAuth(), 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,69 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
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 == "unhealthy" {
|
||||
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,
|
||||
},
|
||||
})
|
||||
}
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return nil
|
||||
}
|
||||
return conn.Transaction(func(tx *gorm.DB) error {
|
||||
return agent.ReconcileScopedNodeHealthEvents(tx, 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,
|
||||
Snapshot: payload.Snapshot,
|
||||
HealthEvents: payload.HealthEvents,
|
||||
}, reportedAt)
|
||||
|
||||
conn := db.DB(ctx)
|
||||
if conn == nil {
|
||||
return
|
||||
}
|
||||
frpsObs := &model.OpenFlareNodeObservationFrps{
|
||||
NodeID: nodeID,
|
||||
CapturedAt: reportedAt,
|
||||
FrpsConnections: payload.FrpsConnCount,
|
||||
FrpsProxyCount: payload.FrpsProxyCount,
|
||||
FrpsClientCount: payload.FrpsClientCount,
|
||||
FrpsProxies: agent.MarshalJSON(payload.FrpsProxies),
|
||||
}
|
||||
if err := conn.Create(frpsObs).Error; err != nil {
|
||||
zap.L().Error("persist relay frps observation failed", zap.String("node_id", nodeID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package relay
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
|
||||
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"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 := authNode.(*model.OpenFlareNode)
|
||||
|
||||
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 := authNode.(*model.OpenFlareNode)
|
||||
ofws.ServeRelay(c, node.NodeID)
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
// DatabaseCleanupTargetAccessLogs is the API cleanup target for access logs.
|
||||
DatabaseCleanupTargetAccessLogs = "node_access_logs"
|
||||
// DatabaseCleanupTargetMetricSnapshots is the API cleanup target for metric snapshots.
|
||||
DatabaseCleanupTargetMetricSnapshots = "node_metric_snapshots"
|
||||
// DatabaseCleanupTargetRequestReports is the API cleanup target for request reports.
|
||||
DatabaseCleanupTargetRequestReports = "node_request_reports"
|
||||
)
|
||||
|
||||
var databaseCleanupTargets = map[string]string{
|
||||
DatabaseCleanupTargetAccessLogs: "访问日志",
|
||||
DatabaseCleanupTargetMetricSnapshots: "性能快照",
|
||||
DatabaseCleanupTargetRequestReports: "请求聚合",
|
||||
}
|
||||
|
||||
// DatabaseCleanupInput describes a manual observability cleanup request.
|
||||
type DatabaseCleanupInput struct {
|
||||
Target string `json:"target"`
|
||||
RetentionDays *int `json:"retention_days"`
|
||||
}
|
||||
|
||||
// DatabaseCleanupResult summarizes a manual observability cleanup run.
|
||||
type DatabaseCleanupResult struct {
|
||||
Target string `json:"target"`
|
||||
TargetLabel string `json:"target_label"`
|
||||
DeletedCount int64 `json:"deleted_count"`
|
||||
DeleteAll bool `json:"delete_all"`
|
||||
RetentionDays *int `json:"retention_days,omitempty"`
|
||||
Cutoff *time.Time `json:"cutoff,omitempty"`
|
||||
}
|
||||
|
||||
// DatabaseAutoCleanupSummary summarizes a scheduled auto-cleanup run.
|
||||
type DatabaseAutoCleanupSummary struct {
|
||||
RetentionDays int `json:"retention_days"`
|
||||
ExecutedAt time.Time `json:"executed_at"`
|
||||
Results []DatabaseCleanupResult `json:"results"`
|
||||
}
|
||||
|
||||
// CleanupDatabaseObservability deletes observability rows for the given target.
|
||||
func CleanupDatabaseObservability(ctx context.Context, input DatabaseCleanupInput) (*DatabaseCleanupResult, error) {
|
||||
target := strings.TrimSpace(input.Target)
|
||||
targetLabel, ok := databaseCleanupTargets[target]
|
||||
if !ok {
|
||||
return nil, errors.New("unsupported cleanup target")
|
||||
}
|
||||
if input.RetentionDays != nil && *input.RetentionDays <= 0 {
|
||||
return nil, errors.New("retention_days 必须为大于 0 的整数")
|
||||
}
|
||||
|
||||
result := &DatabaseCleanupResult{
|
||||
Target: target,
|
||||
TargetLabel: targetLabel,
|
||||
DeleteAll: input.RetentionDays == nil,
|
||||
}
|
||||
|
||||
if input.RetentionDays == nil {
|
||||
deleted, err := deleteAllObservabilityRows(ctx, target)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.DeletedCount = deleted
|
||||
return result, nil
|
||||
}
|
||||
|
||||
retentionDays := *input.RetentionDays
|
||||
cutoff := time.Now().UTC().Add(-time.Duration(retentionDays) * 24 * time.Hour)
|
||||
deleted, err := deleteObservabilityRowsBefore(ctx, target, cutoff)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.DeletedCount = deleted
|
||||
result.RetentionDays = &retentionDays
|
||||
result.Cutoff = &cutoff
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// RunDatabaseAutoCleanupOnce runs retention-based cleanup for all observability targets.
|
||||
func RunDatabaseAutoCleanupOnce(now time.Time) (*DatabaseAutoCleanupSummary, error) {
|
||||
if !model.DatabaseAutoCleanupEnabled {
|
||||
return nil, nil
|
||||
}
|
||||
if model.DatabaseAutoCleanupRetentionDays < 1 {
|
||||
return nil, fmt.Errorf("database auto cleanup retention_days must be at least 1")
|
||||
}
|
||||
|
||||
retentionDays := model.DatabaseAutoCleanupRetentionDays
|
||||
ctx := context.Background()
|
||||
results := make([]DatabaseCleanupResult, 0, len(databaseCleanupTargets))
|
||||
for _, target := range []string{
|
||||
DatabaseCleanupTargetAccessLogs,
|
||||
DatabaseCleanupTargetMetricSnapshots,
|
||||
DatabaseCleanupTargetRequestReports,
|
||||
} {
|
||||
result, err := CleanupDatabaseObservability(ctx, DatabaseCleanupInput{
|
||||
Target: target,
|
||||
RetentionDays: &retentionDays,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results = append(results, *result)
|
||||
}
|
||||
|
||||
return &DatabaseAutoCleanupSummary{
|
||||
RetentionDays: retentionDays,
|
||||
ExecutedAt: now.UTC(),
|
||||
Results: results,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func deleteAllObservabilityRows(ctx context.Context, target string) (int64, error) {
|
||||
switch target {
|
||||
case DatabaseCleanupTargetAccessLogs:
|
||||
return model.DeleteAllOpenFlareAccessLogs(ctx)
|
||||
case DatabaseCleanupTargetMetricSnapshots:
|
||||
return model.DeleteAllOpenFlareMetricSnapshots(ctx)
|
||||
case DatabaseCleanupTargetRequestReports:
|
||||
return model.DeleteAllOpenFlareRequestReports(ctx)
|
||||
default:
|
||||
return 0, errors.New("unsupported cleanup target")
|
||||
}
|
||||
}
|
||||
|
||||
func deleteObservabilityRowsBefore(ctx context.Context, target string, cutoff time.Time) (int64, error) {
|
||||
switch target {
|
||||
case DatabaseCleanupTargetAccessLogs:
|
||||
return model.DeleteOpenFlareAccessLogsBefore(ctx, cutoff)
|
||||
case DatabaseCleanupTargetMetricSnapshots:
|
||||
return model.DeleteOpenFlareMetricSnapshotsBefore(ctx, cutoff)
|
||||
case DatabaseCleanupTargetRequestReports:
|
||||
return model.DeleteOpenFlareRequestReportsBefore(ctx, cutoff)
|
||||
default:
|
||||
return 0, errors.New("unsupported cleanup target")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,153 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupDatabaseCleanupTestDB(t *testing.T) context.Context {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(
|
||||
&model.OpenFlareMetricSnapshot{},
|
||||
&model.OpenFlareRequestReport{},
|
||||
))
|
||||
db.SetDB(sqliteDB)
|
||||
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore())
|
||||
t.Cleanup(func() {
|
||||
resetAccessLogStore()
|
||||
db.SetDB(nil)
|
||||
})
|
||||
return context.Background()
|
||||
}
|
||||
|
||||
func TestCleanupDatabaseObservabilityDeletesTargetedRows(t *testing.T) {
|
||||
ctx := setupDatabaseCleanupTestDB(t)
|
||||
now := time.Now().UTC()
|
||||
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareMetricSnapshot{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: now.Add(-10 * 24 * time.Hour),
|
||||
CPUUsagePercent: 10,
|
||||
}).Error)
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareMetricSnapshot{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: now.Add(-12 * time.Hour),
|
||||
CPUUsagePercent: 20,
|
||||
}).Error)
|
||||
|
||||
retentionDays := 7
|
||||
result, err := CleanupDatabaseObservability(ctx, DatabaseCleanupInput{
|
||||
Target: DatabaseCleanupTargetMetricSnapshots,
|
||||
RetentionDays: &retentionDays,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, result.DeleteAll)
|
||||
assert.Equal(t, int64(1), result.DeletedCount)
|
||||
|
||||
rows, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, rows, 1)
|
||||
assert.Equal(t, float64(20), rows[0].CPUUsagePercent)
|
||||
}
|
||||
|
||||
func TestCleanupDatabaseObservabilityDeletesAllRowsWhenRetentionMissing(t *testing.T) {
|
||||
ctx := setupDatabaseCleanupTestDB(t)
|
||||
now := time.Now().UTC()
|
||||
|
||||
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{
|
||||
{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-3 * time.Hour),
|
||||
RemoteAddr: "203.0.113.1",
|
||||
Host: "example.com",
|
||||
Path: "/one",
|
||||
StatusCode: 200,
|
||||
},
|
||||
{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-2 * time.Hour),
|
||||
RemoteAddr: "203.0.113.2",
|
||||
Host: "example.com",
|
||||
Path: "/two",
|
||||
StatusCode: 502,
|
||||
},
|
||||
}))
|
||||
|
||||
result, err := CleanupDatabaseObservability(ctx, DatabaseCleanupInput{
|
||||
Target: DatabaseCleanupTargetAccessLogs,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, result.DeleteAll)
|
||||
assert.Equal(t, int64(2), result.DeletedCount)
|
||||
|
||||
rows, err := model.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, rows)
|
||||
}
|
||||
|
||||
func TestRunDatabaseAutoCleanupOnceDeletesAllObservabilityTargets(t *testing.T) {
|
||||
ctx := setupDatabaseCleanupTestDB(t)
|
||||
now := time.Now().UTC()
|
||||
|
||||
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{{
|
||||
NodeID: "node-a",
|
||||
LoggedAt: now.Add(-48 * time.Hour),
|
||||
RemoteAddr: "203.0.113.10",
|
||||
Host: "example.com",
|
||||
Path: "/access",
|
||||
StatusCode: 200,
|
||||
}}))
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareMetricSnapshot{
|
||||
NodeID: "node-a",
|
||||
CapturedAt: now.Add(-48 * time.Hour),
|
||||
CPUUsagePercent: 10,
|
||||
}).Error)
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareRequestReport{
|
||||
NodeID: "node-a",
|
||||
WindowStartedAt: now.Add(-49 * time.Hour),
|
||||
WindowEndedAt: now.Add(-48 * time.Hour),
|
||||
RequestCount: 15,
|
||||
}).Error)
|
||||
|
||||
previousEnabled := model.DatabaseAutoCleanupEnabled
|
||||
previousRetentionDays := model.DatabaseAutoCleanupRetentionDays
|
||||
model.DatabaseAutoCleanupEnabled = true
|
||||
model.DatabaseAutoCleanupRetentionDays = 1
|
||||
t.Cleanup(func() {
|
||||
model.DatabaseAutoCleanupEnabled = previousEnabled
|
||||
model.DatabaseAutoCleanupRetentionDays = previousRetentionDays
|
||||
})
|
||||
|
||||
summary, err := RunDatabaseAutoCleanupOnce(now)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, summary)
|
||||
require.Len(t, summary.Results, 3)
|
||||
|
||||
accessLogs, err := model.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, accessLogs)
|
||||
|
||||
metricSnapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, metricSnapshots)
|
||||
|
||||
requestReports, err := model.ListOpenFlareRequestReportsSince(ctx, "", time.Time{}, 0)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, requestReports)
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package tasks hosts shared OpenFlare background job business logic.
|
||||
//
|
||||
// Scheduled execution is handled by the Wavelet Asynq task framework; handlers
|
||||
// live in internal/apps/openflare/async_tasks.go and are registered via
|
||||
// bootstrap.RegisterTasks().
|
||||
package tasks
|
||||
@@ -0,0 +1,44 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
)
|
||||
|
||||
// RunSSLRenewJob renews all TLS certificates that are due for renewal.
|
||||
func RunSSLRenewJob(ctx context.Context) error {
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job started")
|
||||
|
||||
certificates, err := model.ListTLSCertificates(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "[OpenFlareTasks] list certificates failed: %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
due := tls.CertificatesDueForRenewal(certificates, now)
|
||||
if len(due) == 0 {
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job completed: no certificates due")
|
||||
return nil
|
||||
}
|
||||
|
||||
var triggered int
|
||||
for _, cert := range due {
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] renewing certificate id=%d domain=%s", cert.ID, cert.PrimaryDomain)
|
||||
if _, err := tls.RenewCertificate(ctx, cert.ID); err != nil {
|
||||
logger.ErrorF(ctx, "[OpenFlareTasks] renew certificate id=%d domain=%s failed: %v", cert.ID, cert.PrimaryDomain, err)
|
||||
continue
|
||||
}
|
||||
triggered++
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job completed: triggered=%d eligible=%d", triggered, len(due))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupSSLRenewTestDB(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.TLSCertificate{}))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
oldSecret := config.Config.App.SessionSecret
|
||||
config.Config.App.SessionSecret = "test_session_secret_for_ssl_renew"
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
config.Config.App.SessionSecret = oldSecret
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunSSLRenewJobTriggersDueCertificates(t *testing.T) {
|
||||
cleanup := setupSSLRenewTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
restore := tls.SetObtainCertificateFuncForTest(func(ctx context.Context, cert *model.TLSCertificate) error {
|
||||
return nil
|
||||
})
|
||||
defer restore()
|
||||
|
||||
now := time.Date(2026, 6, 18, 12, 0, 0, 0, time.UTC)
|
||||
due := &model.TLSCertificate{
|
||||
Name: "due-cert",
|
||||
Provider: "acme",
|
||||
AutoRenew: true,
|
||||
ApplyStatus: "ready",
|
||||
PrimaryDomain: "due.example.com",
|
||||
CertPEM: " ",
|
||||
KeyPEM: " ",
|
||||
NotAfter: now.Add(2 * 24 * time.Hour),
|
||||
}
|
||||
fresh := &model.TLSCertificate{
|
||||
Name: "fresh-cert",
|
||||
Provider: "acme",
|
||||
AutoRenew: true,
|
||||
ApplyStatus: "ready",
|
||||
PrimaryDomain: "fresh.example.com",
|
||||
CertPEM: " ",
|
||||
KeyPEM: " ",
|
||||
NotAfter: now.Add(30 * 24 * time.Hour),
|
||||
}
|
||||
require.NoError(t, model.CreateTLSCertificateRecord(ctx, due))
|
||||
require.NoError(t, model.CreateTLSCertificateRecord(ctx, fresh))
|
||||
|
||||
require.NoError(t, RunSSLRenewJob(ctx))
|
||||
|
||||
renewed, err := model.GetTLSCertificateByID(ctx, due.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "applying", renewed.ApplyStatus)
|
||||
|
||||
unchanged, err := model.GetTLSCertificateByID(ctx, fresh.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "ready", unchanged.ApplyStatus)
|
||||
}
|
||||
@@ -0,0 +1,262 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package acme
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/go-acme/lego/v4/acme"
|
||||
"github.com/go-acme/lego/v4/certcrypto"
|
||||
"github.com/go-acme/lego/v4/certificate"
|
||||
"github.com/go-acme/lego/v4/challenge/dns01"
|
||||
"github.com/go-acme/lego/v4/lego"
|
||||
"github.com/go-acme/lego/v4/providers/dns/cloudflare"
|
||||
"github.com/go-acme/lego/v4/registration"
|
||||
)
|
||||
|
||||
// AcmeUser implements lego's user interface.
|
||||
type AcmeUser struct {
|
||||
Email string
|
||||
Registration *registration.Resource
|
||||
key crypto.PrivateKey
|
||||
}
|
||||
|
||||
func (u *AcmeUser) GetEmail() string {
|
||||
return u.Email
|
||||
}
|
||||
|
||||
func (u *AcmeUser) GetRegistration() *registration.Resource {
|
||||
return u.Registration
|
||||
}
|
||||
|
||||
func (u *AcmeUser) GetPrivateKey() crypto.PrivateKey {
|
||||
return u.key
|
||||
}
|
||||
|
||||
// CertificateResult holds obtained certificate material.
|
||||
type CertificateResult struct {
|
||||
CertPEM string
|
||||
KeyPEM string
|
||||
NotBefore time.Time
|
||||
NotAfter time.Time
|
||||
}
|
||||
|
||||
func parsePrivateKey(pemData string) (crypto.PrivateKey, error) {
|
||||
block, _ := pem.Decode([]byte(pemData))
|
||||
if block == nil {
|
||||
return nil, errors.New("failed to parse PEM block containing the key")
|
||||
}
|
||||
|
||||
if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
|
||||
return key, nil
|
||||
}
|
||||
return nil, errors.New("failed to parse private key")
|
||||
}
|
||||
|
||||
func encodePrivateKey(key crypto.PrivateKey) (string, error) {
|
||||
var pemBlock *pem.Block
|
||||
switch k := key.(type) {
|
||||
case *rsa.PrivateKey:
|
||||
pemBlock = &pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(k)}
|
||||
case *ecdsa.PrivateKey:
|
||||
b, err := x509.MarshalECPrivateKey(k)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
pemBlock = &pem.Block{Type: "EC PRIVATE KEY", Bytes: b}
|
||||
default:
|
||||
return "", errors.New("unsupported key type")
|
||||
}
|
||||
return string(pem.EncodeToMemory(pemBlock)), nil
|
||||
}
|
||||
|
||||
// GetOrCreateLegoClient returns a configured lego client and optional new account credentials.
|
||||
func GetOrCreateLegoClient(acmeEmail, privateKeyPEM, accountURL string, keyAlgorithm string) (*lego.Client, *AcmeUser, string, string, error) {
|
||||
var privateKey crypto.PrivateKey
|
||||
var err error
|
||||
var newPrivateKeyPEM string
|
||||
var newAccountURL string
|
||||
|
||||
if privateKeyPEM == "" {
|
||||
privateKey, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
pemStr, err := encodePrivateKey(privateKey)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
newPrivateKeyPEM = pemStr
|
||||
} else {
|
||||
privateKey, err = parsePrivateKey(privateKeyPEM)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
}
|
||||
|
||||
user := &AcmeUser{
|
||||
Email: acmeEmail,
|
||||
key: privateKey,
|
||||
}
|
||||
|
||||
if accountURL != "" {
|
||||
user.Registration = ®istration.Resource{
|
||||
Body: acme.Account{
|
||||
Status: "valid",
|
||||
Contact: []string{"mailto:" + acmeEmail},
|
||||
},
|
||||
URI: accountURL,
|
||||
}
|
||||
}
|
||||
|
||||
config := lego.NewConfig(user)
|
||||
config.CADirURL = lego.LEDirectoryProduction
|
||||
|
||||
switch keyAlgorithm {
|
||||
case "RSA2048":
|
||||
config.Certificate.KeyType = certcrypto.RSA2048
|
||||
case "RSA4096":
|
||||
config.Certificate.KeyType = certcrypto.RSA4096
|
||||
case "EC256":
|
||||
config.Certificate.KeyType = certcrypto.EC256
|
||||
case "EC384":
|
||||
config.Certificate.KeyType = certcrypto.EC384
|
||||
default:
|
||||
config.Certificate.KeyType = certcrypto.RSA2048
|
||||
}
|
||||
|
||||
client, err := lego.NewClient(config)
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
|
||||
if accountURL == "" {
|
||||
reg, err := client.Registration.Register(registration.RegisterOptions{TermsOfServiceAgreed: true})
|
||||
if err != nil {
|
||||
return nil, nil, "", "", err
|
||||
}
|
||||
user.Registration = reg
|
||||
newAccountURL = reg.URI
|
||||
}
|
||||
|
||||
return client, user, newPrivateKeyPEM, newAccountURL, nil
|
||||
}
|
||||
|
||||
// SetupDNSProvider configures DNS-01 challenge for the lego client.
|
||||
func SetupDNSProvider(client *lego.Client, dnsType, dnsAuth string, dns1, dns2 string, disableCNAME, skipDNS bool) error {
|
||||
var provider challengeProvider
|
||||
|
||||
switch dnsType {
|
||||
case "cloudflare":
|
||||
var creds map[string]string
|
||||
if err := json.Unmarshal([]byte(dnsAuth), &creds); err != nil {
|
||||
return fmt.Errorf("failed to parse cloudflare credentials: %v", err)
|
||||
}
|
||||
|
||||
config := cloudflare.NewDefaultConfig()
|
||||
config.AuthToken = creds["api_token"]
|
||||
|
||||
p, err := cloudflare.NewDNSProviderConfig(config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
provider = p
|
||||
default:
|
||||
return fmt.Errorf("unsupported DNS provider: %s", dnsType)
|
||||
}
|
||||
|
||||
var resolvers []string
|
||||
if dns1 != "" {
|
||||
resolvers = append(resolvers, dns1+":53")
|
||||
}
|
||||
if dns2 != "" {
|
||||
resolvers = append(resolvers, dns2+":53")
|
||||
}
|
||||
|
||||
var opts []dns01.ChallengeOption
|
||||
|
||||
if len(resolvers) > 0 {
|
||||
opts = append(opts, dns01.AddRecursiveNameservers(resolvers))
|
||||
}
|
||||
|
||||
if disableCNAME {
|
||||
opts = append(opts, dns01.DisableCompletePropagationRequirement())
|
||||
}
|
||||
|
||||
if skipDNS {
|
||||
opts = append(opts, dns01.WrapPreCheck(func(domain, fqdn, value string, check dns01.PreCheckFunc) (bool, error) {
|
||||
time.Sleep(20 * time.Second)
|
||||
return true, nil
|
||||
}))
|
||||
}
|
||||
|
||||
return client.Challenge.SetDNS01Provider(provider, opts...)
|
||||
}
|
||||
|
||||
type challengeProvider interface {
|
||||
Present(domain, token, keyAuth string) error
|
||||
CleanUp(domain, token, keyAuth string) error
|
||||
}
|
||||
|
||||
// ObtainSSL obtains a certificate via ACME DNS-01 challenge.
|
||||
func ObtainSSL(
|
||||
acmeEmail, acmePrivateKeyPEM, acmeURL string,
|
||||
dnsType, dnsAuth string,
|
||||
dns1, dns2 string,
|
||||
disableCNAME, skipDNS bool,
|
||||
keyAlgorithm string,
|
||||
domains []string,
|
||||
) (string, string, *CertificateResult, error) {
|
||||
client, _, newPrivateKeyPEM, newAccountURL, err := GetOrCreateLegoClient(acmeEmail, acmePrivateKeyPEM, acmeURL, keyAlgorithm)
|
||||
if err != nil {
|
||||
return "", "", nil, fmt.Errorf("failed to create ACME client: %w", err)
|
||||
}
|
||||
|
||||
err = SetupDNSProvider(client, dnsType, dnsAuth, dns1, dns2, disableCNAME, skipDNS)
|
||||
if err != nil {
|
||||
return newAccountURL, newPrivateKeyPEM, nil, fmt.Errorf("failed to setup DNS provider: %w", err)
|
||||
}
|
||||
|
||||
request := certificate.ObtainRequest{
|
||||
Domains: domains,
|
||||
Bundle: true,
|
||||
}
|
||||
|
||||
certificates, err := client.Certificate.Obtain(request)
|
||||
if err != nil {
|
||||
return newAccountURL, newPrivateKeyPEM, nil, fmt.Errorf("failed to obtain certificate: %w", err)
|
||||
}
|
||||
|
||||
result := &CertificateResult{
|
||||
CertPEM: string(certificates.Certificate),
|
||||
KeyPEM: string(certificates.PrivateKey),
|
||||
}
|
||||
|
||||
certBlock, _ := pem.Decode(certificates.Certificate)
|
||||
if certBlock != nil {
|
||||
parsedCert, err := x509.ParseCertificate(certBlock.Bytes)
|
||||
if err == nil {
|
||||
result.NotBefore = parsedCert.NotBefore
|
||||
result.NotAfter = parsedCert.NotAfter
|
||||
}
|
||||
}
|
||||
|
||||
return newAccountURL, newPrivateKeyPEM, result, nil
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSplitAcmeDomains(t *testing.T) {
|
||||
assert.Equal(t, []string{"example.com"}, splitAcmeDomains("example.com", ""))
|
||||
assert.Equal(t, []string{"example.com", "*.example.com"}, splitAcmeDomains("example.com", "*.example.com"))
|
||||
assert.Equal(t, []string{"example.com", "www.example.com", "api.example.com"}, splitAcmeDomains("example.com", "www.example.com\napi.example.com"))
|
||||
assert.Equal(t, []string{"example.com", "www.example.com", "api.example.com"}, splitAcmeDomains("example.com", "www.example.com, api.example.com"))
|
||||
}
|
||||
|
||||
func TestCertificatesDueForRenewal(t *testing.T) {
|
||||
now := time.Date(2026, 6, 18, 12, 0, 0, 0, time.UTC)
|
||||
certificates := []model.TLSCertificate{
|
||||
{ID: 1, Provider: "acme", AutoRenew: true, ApplyStatus: "ready", PrimaryDomain: "due.example.com", NotAfter: now.Add(3 * 24 * time.Hour)},
|
||||
{ID: 2, Provider: "acme", AutoRenew: true, ApplyStatus: "ready", PrimaryDomain: "fresh.example.com", NotAfter: now.Add(30 * 24 * time.Hour)},
|
||||
{ID: 3, Provider: "upload", AutoRenew: true, ApplyStatus: "ready", PrimaryDomain: "upload.example.com", NotAfter: now.Add(24 * time.Hour)},
|
||||
{ID: 4, Provider: "acme", AutoRenew: false, ApplyStatus: "ready", PrimaryDomain: "manual.example.com", NotAfter: now.Add(24 * time.Hour)},
|
||||
{ID: 5, Provider: "acme", AutoRenew: true, ApplyStatus: "applying", PrimaryDomain: "busy.example.com", NotAfter: now.Add(24 * time.Hour)},
|
||||
}
|
||||
|
||||
due := CertificatesDueForRenewal(certificates, now)
|
||||
require.Len(t, due, 1)
|
||||
assert.Equal(t, uint(1), due[0].ID)
|
||||
assert.Equal(t, "due.example.com", due[0].PrimaryDomain)
|
||||
}
|
||||
|
||||
func TestApplyCertificateReturnsApplying(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
dnsAccount, err := CreateDNSAccount(ctx, DNSAccountInput{
|
||||
Name: "Test Cloudflare",
|
||||
Type: "cloudflare",
|
||||
Authorization: `{"api_token": "dummy_token"}`,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, cert *model.TLSCertificate) error {
|
||||
return updateCertError(ctx, cert, "dns challenge failed")
|
||||
})
|
||||
defer restore()
|
||||
|
||||
cert, err := ApplyCertificate(ctx, ApplyInput{
|
||||
Name: "Test ACME Cert",
|
||||
PrimaryDomain: "example.com",
|
||||
OtherDomains: "*.example.com",
|
||||
DnsAccountID: dnsAccount.ID,
|
||||
KeyAlgorithm: "RSA2048",
|
||||
AutoRenew: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "applying", cert.ApplyStatus)
|
||||
assert.Equal(t, "acme", cert.Provider)
|
||||
}
|
||||
|
||||
func TestRenewCertificateSetsApplying(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
cert := &model.TLSCertificate{
|
||||
Name: "renew-cert",
|
||||
Provider: "acme",
|
||||
AutoRenew: true,
|
||||
ApplyStatus: "ready",
|
||||
PrimaryDomain: "renew.example.com",
|
||||
CertPEM: " ",
|
||||
KeyPEM: " ",
|
||||
}
|
||||
require.NoError(t, model.CreateTLSCertificateRecord(ctx, cert))
|
||||
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error {
|
||||
return nil
|
||||
})
|
||||
defer restore()
|
||||
|
||||
renewed, err := RenewCertificate(ctx, cert.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "applying", renewed.ApplyStatus)
|
||||
}
|
||||
|
||||
func TestConvertCertificateToACMEPreservesUploadOnFailure(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
originalCertPEM, originalKeyPEM := generateTestCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateCertificate(ctx, CertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: originalCertPEM,
|
||||
KeyPEM: originalKeyPEM,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := model.GetTLSCertificateByID(ctx, cert.ID)
|
||||
require.NoError(t, err)
|
||||
originalStoredCertPEM := stored.CertPEM
|
||||
originalStoredKeyPEM := stored.KeyPEM
|
||||
|
||||
stored.ApplyStatus = "applying"
|
||||
stored.PrimaryDomain = "manual.example.com"
|
||||
require.NoError(t, model.SaveTLSCertificate(ctx, stored))
|
||||
|
||||
err = updateCertError(ctx, stored, "dns challenge failed")
|
||||
require.Error(t, err)
|
||||
|
||||
finalCert, err := model.GetTLSCertificateByID(ctx, cert.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "upload", finalCert.Provider)
|
||||
assert.Equal(t, "error", finalCert.ApplyStatus)
|
||||
assert.Equal(t, originalStoredCertPEM, finalCert.CertPEM)
|
||||
assert.Equal(t, originalStoredKeyPEM, finalCert.KeyPEM)
|
||||
assert.True(t, strings.Contains(finalCert.ApplyMessage, "dns challenge failed"))
|
||||
}
|
||||
|
||||
func TestConvertCertificateToACMERejectsInvalidStates(t *testing.T) {
|
||||
cleanup := setupTLSTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
certPEM, keyPEM := generateTestCertificatePair(t, []string{"manual.example.com"})
|
||||
cert, err := CreateCertificate(ctx, CertificateInput{
|
||||
Name: "manual-cert",
|
||||
CertPEM: certPEM,
|
||||
KeyPEM: keyPEM,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
cert.Provider = "acme"
|
||||
require.NoError(t, model.SaveTLSCertificate(ctx, cert))
|
||||
_, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "only uploaded")
|
||||
|
||||
cert.Provider = "upload"
|
||||
cert.ApplyStatus = "applying"
|
||||
require.NoError(t, model.SaveTLSCertificate(ctx, cert))
|
||||
_, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already applying")
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
const (
|
||||
errCertificateNameRequired = "certificate name cannot be empty"
|
||||
errCertificateNameExists = "certificate name already exists"
|
||||
errCertificateContentRequired = "certificate content and key content cannot be empty"
|
||||
errCertificateContentInvalid = "certificate or key format is invalid"
|
||||
errCertificateDeleteReferenced = "certificate is still referenced by proxy routes"
|
||||
errCertificateOnlyACME = "only acme certificates can be updated via this endpoint"
|
||||
errCertificateOnlyUploadConvert = "only uploaded certificates can be converted to acme"
|
||||
errCertificateAlreadyApplying = "certificate is already applying"
|
||||
errCertificateOnlyACMERenew = "only acme certificates can be renewed"
|
||||
errCertificateFilesRequired = "certificate file and key file cannot be empty"
|
||||
errCertificatePEMInvalid = "证书 PEM 内容不合法"
|
||||
|
||||
errManagedDomainRequired = "域名不能为空"
|
||||
errManagedDomainInvalid = "域名格式不合法"
|
||||
errManagedDomainWildcardInvalid = "通配符域名仅支持 *.example.com 格式"
|
||||
errManagedDomainExists = "域名已存在"
|
||||
errManagedDomainCertNotFound = "所选证书不存在"
|
||||
|
||||
errDNSAccountInUse = "该 DNS 账号已被证书使用,无法删除"
|
||||
)
|
||||
@@ -0,0 +1,62 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package tls
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func parseLeafCertificate(certPEM string) (*x509.Certificate, error) {
|
||||
certPEMBlock, _ := pem.Decode([]byte(certPEM))
|
||||
if certPEMBlock == nil {
|
||||
return nil, errors.New(errCertificatePEMInvalid)
|
||||
}
|
||||
leaf, err := x509.ParseCertificate(certPEMBlock.Bytes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return leaf, nil
|
||||
}
|
||||
|
||||
func readMultipartFile(fileHeader *multipart.FileHeader) (string, error) {
|
||||
file, err := fileHeader.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer file.Close()
|
||||
data, err := io.ReadAll(file)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
func decodeStoredDomainCertIDs(raw string, domainCount int) ([]uint, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return nil, nil
|
||||
}
|
||||
var domainCertIDs []uint
|
||||
if err := json.Unmarshal([]byte(text), &domainCertIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if domainCount > 0 && len(domainCertIDs) != domainCount {
|
||||
return nil, fmt.Errorf("domain_cert_ids length mismatch")
|
||||
}
|
||||
return domainCertIDs, nil
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user