mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
migrate
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/openflare/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 ""
|
||||
}
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -89,8 +90,12 @@ func sourceSupportFiles(files []SupportFile) []SupportFile {
|
||||
}
|
||||
|
||||
func isRuntimeGeneratedSupportFile(path string) bool {
|
||||
path = strings.TrimSpace(path)
|
||||
return strings.HasPrefix(path, "runtime/")
|
||||
switch strings.TrimSpace(path) {
|
||||
case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isActiveConfigNotFound(err error) bool {
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
openrestyrender "github.com/rain-kl/openflare/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)
|
||||
}
|
||||
}
|
||||
@@ -17,5 +17,4 @@ const (
|
||||
errIPInvalid = "ip 格式无效"
|
||||
errAgentVersionRequired = "version 不能为空"
|
||||
errNodeIDConflict = "节点标识生成冲突,请重试"
|
||||
errPagesPackageNotFound = "Pages 部署包尚未实现"
|
||||
)
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
@@ -97,6 +98,41 @@ func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveNa
|
||||
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 {
|
||||
@@ -199,6 +235,7 @@ func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenF
|
||||
}
|
||||
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)
|
||||
@@ -210,12 +247,25 @@ func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenF
|
||||
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
|
||||
@@ -241,7 +291,8 @@ func isUniqueConstraintError(err error) bool {
|
||||
return strings.Contains(strings.ToLower(err.Error()), "unique")
|
||||
}
|
||||
|
||||
func refreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
|
||||
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
|
||||
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -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/openflare/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)
|
||||
}
|
||||
}
|
||||
@@ -28,7 +28,7 @@ func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode,
|
||||
if err := model.SaveOpenFlareNode(ctx, authNode); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(ctx, authNode)
|
||||
RefreshAccessTokenCache(ctx, authNode)
|
||||
return &RegistrationResponse{
|
||||
NodeID: authNode.NodeID,
|
||||
AccessToken: authNode.AccessToken,
|
||||
@@ -74,7 +74,7 @@ func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*Registrat
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
refreshAccessTokenCache(ctx, record)
|
||||
RefreshAccessTokenCache(ctx, record)
|
||||
return &RegistrationResponse{
|
||||
NodeID: record.NodeID,
|
||||
AccessToken: record.AccessToken,
|
||||
@@ -116,24 +116,29 @@ func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload N
|
||||
}
|
||||
}
|
||||
|
||||
refreshAccessTokenCache(ctx, authNode)
|
||||
RefreshAccessTokenCache(ctx, authNode)
|
||||
|
||||
reportedAt := time.Now()
|
||||
if authNode.LastSeenAt != nil {
|
||||
reportedAt = *authNode.LastSeenAt
|
||||
}
|
||||
persistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt)
|
||||
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: nil,
|
||||
WAFIPGroups: wafIPGroups,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -149,9 +154,13 @@ func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) {
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// SyncWAFIPGroups is a stub until full WAF agent sync is migrated.
|
||||
func SyncWAFIPGroups(_ context.Context, _ WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) {
|
||||
return &WAFIPGroupSyncResult{Groups: []WAFIPGroup{}}, 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.
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -18,12 +19,14 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
healthEventStatusActive = "active"
|
||||
healthEventStatusResolved = "resolved"
|
||||
healthSeverityInfo = "info"
|
||||
healthSeverityWarning = "warning"
|
||||
healthSeverityCritical = "critical"
|
||||
accessLogPathMaxLength = 100
|
||||
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.
|
||||
@@ -102,7 +105,8 @@ type NodeHealthEvent struct {
|
||||
Metadata map[string]string `json:"metadata"`
|
||||
}
|
||||
|
||||
func persistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
|
||||
// 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
|
||||
}
|
||||
@@ -275,15 +279,29 @@ func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *NodeTrafficRep
|
||||
}
|
||||
|
||||
func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, reportedAt time.Time) error {
|
||||
if len(logs) == 0 {
|
||||
return 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()
|
||||
}
|
||||
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)
|
||||
}
|
||||
exists, err := accessLogExists(tx, record)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -295,16 +313,27 @@ func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, rep
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
_, err = deleteAccessLogsByNodeBefore(tx, 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 {
|
||||
@@ -314,7 +343,21 @@ func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEv
|
||||
}
|
||||
|
||||
var activeEvents []*model.OpenFlareHealthEvent
|
||||
if err := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive).Find(&activeEvents).Error; err != nil {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -391,6 +434,11 @@ func requestReportExists(tx *gorm.DB, nodeID string, windowStartedAt, windowEnde
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func deleteAccessLogsByNodeBefore(tx *gorm.DB, nodeID string, before time.Time) (int64, error) {
|
||||
result := tx.Where("node_id = ? AND logged_at < ?", nodeID, before).Delete(&model.OpenFlareAccessLog{})
|
||||
return result.RowsAffected, result.Error
|
||||
}
|
||||
|
||||
func accessLogExists(tx *gorm.DB, record *model.OpenFlareAccessLog) (bool, error) {
|
||||
var count int64
|
||||
if err := tx.Model(&model.OpenFlareAccessLog{}).
|
||||
@@ -438,6 +486,11 @@ func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
|
||||
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 ""
|
||||
|
||||
@@ -5,8 +5,10 @@ package agent
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/compat"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
@@ -98,7 +100,7 @@ func GetActiveConfigHandler(c *gin.Context) {
|
||||
compat.OK(c, config)
|
||||
}
|
||||
|
||||
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent (stub).
|
||||
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent.
|
||||
func SyncWAFIPGroupsHandler(c *gin.Context) {
|
||||
var input WAFIPGroupSyncInput
|
||||
if !compat.BindJSON(c, &input) {
|
||||
@@ -129,13 +131,33 @@ func ReportApplyLogHandler(c *gin.Context) {
|
||||
compat.OK(c, log)
|
||||
}
|
||||
|
||||
// DownloadPagesPackageHandler is a stub until Pages agent packaging is migrated.
|
||||
// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent.
|
||||
func DownloadPagesPackageHandler(c *gin.Context) {
|
||||
c.JSON(http.StatusNotFound, compat.Envelope{
|
||||
Success: false,
|
||||
Message: errPagesPackageNotFound,
|
||||
Data: nil,
|
||||
})
|
||||
deploymentID, ok := pagesDeploymentIDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
filePath, fileName, err := pages.GetDeploymentPackagePath(c.Request.Context(), deploymentID)
|
||||
if err != nil {
|
||||
compat.Fail(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.Header("Content-Disposition", "attachment; filename="+fileName)
|
||||
c.File(filePath)
|
||||
}
|
||||
|
||||
func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
|
||||
raw := c.Param("deployment_id")
|
||||
if raw == "" {
|
||||
compat.Fail(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
id64, err := strconv.ParseUint(raw, 10, 64)
|
||||
if err != nil || id64 == 0 {
|
||||
compat.Fail(c, "无效的 ID")
|
||||
return 0, false
|
||||
}
|
||||
return uint(id64), true
|
||||
}
|
||||
|
||||
// AgentWebSocketHandler upgrades an authenticated agent websocket connection.
|
||||
|
||||
@@ -0,0 +1,200 @@
|
||||
// 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"`
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
Reference in New Issue
Block a user