[优化] 优化 WAF IP 组同步功能及相关文档更新

This commit is contained in:
ryan
2026-06-01 13:55:04 +08:00
parent a092935623
commit a8e8a940a0
27 changed files with 735 additions and 43 deletions
+25
View File
@@ -41,6 +41,7 @@ type AgentNodePayload struct {
AccessLogs []AgentNodeAccessLog `json:"access_logs,omitempty"`
BufferedObservability []AgentBufferedObservabilityRecord `json:"buffered_observability,omitempty"`
HealthEvents []AgentNodeHealthEvent `json:"health_events"`
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
}
type ApplyLogPayload struct {
@@ -107,6 +108,25 @@ type HeartbeatResponse struct {
Node *model.Node `json:"node"`
AgentSettings *AgentSettings `json:"agent_settings"`
ActiveConfig *ActiveConfigMeta `json:"active_config"`
WAFIPGroups []AgentWAFIPGroup `json:"waf_ip_groups,omitempty"`
}
type AgentWAFIPGroup 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"`
}
type AgentWAFIPGroupSyncInput struct {
IDs []uint `json:"ids"`
Checksums map[string]string `json:"checksums"`
}
type AgentWAFIPGroupSyncResult struct {
Groups []AgentWAFIPGroup `json:"groups"`
}
type NodeView struct {
@@ -183,10 +203,15 @@ func HeartbeatNode(node *model.Node, payload AgentNodePayload) (*HeartbeatRespon
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
wafIPGroups, err := ChangedWAFIPGroupsForAgent(nil, payload.WAFIPGroupChecksums)
if err != nil {
return nil, err
}
return &HeartbeatResponse{
Node: node,
AgentSettings: buildAgentSettings(node, updateNow, updateChannel.String(), updateTag, restartOpenrestyNow),
ActiveConfig: activeConfig,
WAFIPGroups: wafIPGroups,
}, nil
}
+70
View File
@@ -3,6 +3,7 @@ package service
import (
"errors"
"openflare/model"
"strconv"
"strings"
"testing"
"time"
@@ -74,6 +75,75 @@ func TestGetActiveConfigForAgentIncludesWAFConfig(t *testing.T) {
}
}
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
setupServiceTestDB(t)
route, err := CreateProxyRoute(ProxyRouteInput{
SiteName: "agent-waf-ip-group",
Domains: []string{"agent-waf-ip-group.example.com"},
OriginURL: "https://origin.internal",
Enabled: true,
})
if err != nil {
t.Fatalf("CreateProxyRoute failed: %v", err)
}
ipGroup, err := CreateWAFIPGroup(WAFIPGroupInput{
Name: "agent runtime group",
Type: WAFIPGroupTypeManual,
Enabled: true,
IPList: []string{"203.0.113.44"},
})
if err != nil {
t.Fatalf("CreateWAFIPGroup failed: %v", err)
}
ruleGroup, err := CreateWAFRuleGroup(WAFRuleGroupInput{
Name: "agent refs",
Enabled: true,
IPBlacklistGroups: []uint{ipGroup.ID},
})
if err != nil {
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
}
if _, err = ReplaceWAFSiteRuleGroups(route.ID, []uint{ruleGroup.ID}); err != nil {
t.Fatalf("ReplaceWAFSiteRuleGroups failed: %v", err)
}
if _, err = PublishConfigVersion("root", false); err != nil {
t.Fatalf("PublishConfigVersion failed: %v", err)
}
groups, err := ChangedWAFIPGroupsForAgent(nil, nil)
if err != nil {
t.Fatalf("ChangedWAFIPGroupsForAgent failed: %v", err)
}
if len(groups) != 1 || groups[0].ID != ipGroup.ID || groups[0].IPList[0] != "203.0.113.44" || groups[0].Checksum == "" {
t.Fatalf("unexpected changed groups: %#v", groups)
}
groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10)
same, err := ChangedWAFIPGroupsForAgent(nil, map[string]string{groupKey: groups[0].Checksum})
if err != nil {
t.Fatalf("ChangedWAFIPGroupsForAgent with checksum failed: %v", err)
}
if len(same) != 0 {
t.Fatalf("expected no delta for matching checksum, got %#v", same)
}
updated, err := UpdateWAFIPGroup(ipGroup.ID, WAFIPGroupInput{
Name: "agent runtime group",
Type: WAFIPGroupTypeManual,
Enabled: true,
IPList: []string{"203.0.113.45"},
})
if err != nil {
t.Fatalf("UpdateWAFIPGroup failed: %v", err)
}
delta, err := ChangedWAFIPGroupsForAgent(nil, map[string]string{groupKey: groups[0].Checksum})
if err != nil {
t.Fatalf("ChangedWAFIPGroupsForAgent after update failed: %v", err)
}
if len(delta) != 1 || delta[0].ID != updated.ID || delta[0].IPList[0] != "203.0.113.45" || delta[0].Checksum == groups[0].Checksum {
t.Fatalf("expected updated group delta, got %#v", delta)
}
}
func TestGetActiveConfigForAgentUsesTenMinutePoWSessionDefault(t *testing.T) {
setupServiceTestDB(t)
+28
View File
@@ -10,6 +10,7 @@ const (
AgentWSMessageTypeSettings = "settings"
AgentWSMessageTypeActiveConfig = "active_config"
AgentWSMessageTypeForceSyncConfig = "force_sync_config"
AgentWSMessageTypeWAFIPGroups = "waf_ip_groups"
AgentWSMessageTypePing = "ping"
AgentWSMessageTypePong = "pong"
@@ -77,6 +78,16 @@ func SendAgentWSForceSyncConfig(nodeID string, activeConfig *ActiveConfigMeta) b
})
}
func SendAgentWSWAFIPGroups(nodeID string, groups []AgentWAFIPGroup) bool {
if len(groups) == 0 {
return false
}
return DefaultAgentWSHub.SendMessage(nodeID, WSMessage{
Type: AgentWSMessageTypeWAFIPGroups,
Payload: groups,
})
}
func SendAgentWSPong(nodeID string) bool {
return DefaultAgentWSHub.SendMessage(nodeID, WSMessage{
Type: AgentWSMessageTypePong,
@@ -111,3 +122,20 @@ func BroadcastAgentWSActiveConfig(activeConfig *ActiveConfigMeta) AgentWSBroadca
)
return result
}
func BroadcastAgentWSWAFIPGroups(groups []AgentWAFIPGroup) WSBroadcastResult {
if len(groups) == 0 {
return WSBroadcastResult{}
}
result := DefaultAgentWSHub.Broadcast(WSMessage{
Type: AgentWSMessageTypeWAFIPGroups,
Payload: groups,
})
slog.Debug("agent ws broadcast waf ip groups",
"group_count", len(groups),
"client_count", result.ClientCount,
"success_count", result.SuccessCount,
"failed_nodes", result.FailedIDs,
)
return result
}
@@ -653,16 +653,11 @@ func buildSnapshotWAFIPGroups(ruleGroups []snapshotWAFRuleGroup) ([]snapshotWAFI
if group == nil {
return nil, fmt.Errorf("IP 组 %d 不存在", id)
}
ips, err := decodeStringList(group.IPList)
if err != nil {
return nil, fmt.Errorf("IP 组 %s 列表无效: %w", group.Name, err)
}
snapshots = append(snapshots, snapshotWAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ips,
})
}
return snapshots, nil
+160 -2
View File
@@ -2,10 +2,13 @@ package service
import (
"bytes"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"net/netip"
@@ -178,7 +181,11 @@ func CreateWAFIPGroup(input WAFIPGroupInput) (*WAFIPGroupView, error) {
if err := group.Insert(); err != nil {
return nil, err
}
return GetWAFIPGroup(group.ID)
view, err := GetWAFIPGroup(group.ID)
if err == nil {
broadcastWAFIPGroupToAgents(group.ID)
}
return view, err
}
func UpdateWAFIPGroup(id uint, input WAFIPGroupInput) (*WAFIPGroupView, error) {
@@ -193,7 +200,11 @@ func UpdateWAFIPGroup(id uint, input WAFIPGroupInput) (*WAFIPGroupView, error) {
if err := group.Update(); err != nil {
return nil, err
}
return GetWAFIPGroup(group.ID)
view, err := GetWAFIPGroup(group.ID)
if err == nil {
broadcastWAFIPGroupToAgents(group.ID)
}
return view, err
}
func DeleteWAFIPGroup(id uint) error {
@@ -364,6 +375,151 @@ func buildWAFIPGroupView(group *model.WAFIPGroup, referenceCount int) (WAFIPGrou
return view, nil
}
func ChangedWAFIPGroupsForAgent(ids []uint, checksums map[string]string) ([]AgentWAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
if len(targetIDs) == 0 {
activeIDs, err := activeConfigWAFIPGroupIDs()
if err != nil {
return nil, err
}
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []AgentWAFIPGroup{}, nil
}
groups, err := buildAgentWAFIPGroups(targetIDs)
if err != nil {
return nil, err
}
changed := make([]AgentWAFIPGroup, 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 SyncWAFIPGroupsForAgent(input AgentWAFIPGroupSyncInput) (*AgentWAFIPGroupSyncResult, error) {
groups, err := ChangedWAFIPGroupsForAgent(input.IDs, input.Checksums)
if err != nil {
return nil, err
}
return &AgentWAFIPGroupSyncResult{Groups: groups}, nil
}
func buildAgentWAFIPGroups(ids []uint) ([]AgentWAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
return []AgentWAFIPGroup{}, nil
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
groups, err := model.ListWAFIPGroupsByIDs(ids)
if err != nil {
return nil, err
}
groupByID := make(map[uint]*model.WAFIPGroup, len(groups))
for _, group := range groups {
groupByID[group.ID] = group
}
result := make([]AgentWAFIPGroup, 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.WAFIPGroup) (AgentWAFIPGroup, error) {
if group == nil {
return AgentWAFIPGroup{}, errors.New("IP 组不存在")
}
ips, err := decodeStringList(group.IPList)
if err != nil {
return AgentWAFIPGroup{}, err
}
if !group.Enabled {
ips = []string{}
}
agentGroup := AgentWAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ips,
}
agentGroup.Checksum = checksumAgentWAFIPGroup(agentGroup)
return agentGroup, nil
}
func checksumAgentWAFIPGroup(group AgentWAFIPGroup) 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() ([]uint, error) {
version, err := model.GetActiveConfigVersion()
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return []uint{}, nil
}
return nil, err
}
snapshot, err := parseSnapshotDocument(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 broadcastWAFIPGroupToAgents(id uint) {
groups, err := buildAgentWAFIPGroups([]uint{id})
if err != nil || len(groups) == 0 {
if err != nil {
slog.Debug("build waf ip group broadcast payload failed", "id", id, "error", err)
}
return
}
BroadcastAgentWSWAFIPGroups(groups)
}
func syncWAFIPGroup(group *model.WAFIPGroup, now time.Time) (*WAFIPGroupSyncResult, error) {
if group == nil {
return nil, errors.New("IP 组不存在")
@@ -399,6 +555,7 @@ func syncWAFIPGroupSubscription(group *model.WAFIPGroup, now time.Time) (*WAFIPG
if err := group.UpdateSyncResult(); err != nil {
return nil, err
}
broadcastWAFIPGroupToAgents(group.ID)
view, err := GetWAFIPGroup(group.ID)
if err != nil {
return nil, err
@@ -481,6 +638,7 @@ func syncWAFIPGroupAutomatic(group *model.WAFIPGroup, now time.Time) (*WAFIPGrou
if err := group.UpdateSyncResult(); err != nil {
return nil, err
}
broadcastWAFIPGroupToAgents(group.ID)
view, err := GetWAFIPGroup(group.ID)
if err != nil {
return nil, err
+14 -6
View File
@@ -305,7 +305,7 @@ func TestWAFIPGroupAutomaticRejectsInvalidExpr(t *testing.T) {
}
}
func TestPublishConfigVersionExpandsWAFIPGroupReferences(t *testing.T) {
func TestPublishConfigVersionKeepsWAFIPGroupReferences(t *testing.T) {
setupServiceTestDB(t)
route, err := CreateProxyRoute(ProxyRouteInput{
@@ -345,18 +345,26 @@ func TestPublishConfigVersionExpandsWAFIPGroupReferences(t *testing.T) {
if !strings.Contains(result.Version.SnapshotJSON, `"ip_groups"`) {
t.Fatal("expected snapshot to include waf ip groups")
}
if strings.Contains(result.Version.SnapshotJSON, "203.0.113.30") {
t.Fatal("expected snapshot to avoid embedding waf ip group members")
}
var files []SupportFile
if err = json.Unmarshal([]byte(result.Version.SupportFilesJSON), &files); err != nil {
t.Fatalf("decode support files failed: %v", err)
}
found := false
foundReference := false
for _, file := range files {
if file.Path == "waf_config.json" && strings.Contains(file.Content, "203.0.113.30") {
found = true
if file.Path == "waf_config.json" {
if strings.Contains(file.Content, "203.0.113.30") {
t.Fatalf("expected waf_config.json to avoid expanded IP group members, got %s", file.Content)
}
if strings.Contains(file.Content, `"ip_blacklist_group_ids":[`) {
foundReference = true
}
}
}
if !found {
t.Fatalf("expected expanded IP group in waf_config.json, got %#v", files)
if !foundReference {
t.Fatalf("expected IP group reference in waf_config.json, got %#v", files)
}
}