[优化] 优化 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
+29
View File
@@ -78,9 +78,33 @@ func AgentHeartbeat(c *gin.Context) {
respondSuccessWithExtras(c, node.Node, gin.H{
"agent_settings": node.AgentSettings,
"active_config": node.ActiveConfig,
"waf_ip_groups": node.WAFIPGroups,
})
}
// AgentSyncWAFIPGroups godoc
// @Summary Sync WAF IP groups for agent
// @Tags Agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param payload body service.AgentWAFIPGroupSyncInput true "WAF IP group sync payload"
// @Success 200 {object} map[string]interface{}
// @Failure 400 {object} map[string]interface{}
// @Router /api/agent/waf/ip-groups/sync [post]
func AgentSyncWAFIPGroups(c *gin.Context) {
var input service.AgentWAFIPGroupSyncInput
if !bindJSON(c, &input) {
return
}
result, err := service.SyncWAFIPGroupsForAgent(input)
if err != nil {
respondFailure(c, err.Error())
return
}
respondSuccess(c, result)
}
// AgentGetActiveConfig godoc
// @Summary Get active config for agent
// @Tags Agent
@@ -237,12 +261,17 @@ func handleAgentWSStatus(c *gin.Context, node *model.Node, message service.Agent
if response.ActiveConfig != nil {
activeConfigSent = service.SendAgentWSActiveConfig(node.NodeID, response.ActiveConfig)
}
wafIPGroupsSent := false
if len(response.WAFIPGroups) > 0 {
wafIPGroupsSent = service.SendAgentWSWAFIPGroups(node.NodeID, response.WAFIPGroups)
}
slog.Debug("agent ws status processed",
"node_id", node.NodeID,
"current_version", payload.CurrentVersion,
"openresty_status", payload.OpenrestyStatus,
"settings_sent", settingsSent,
"active_config_sent", activeConfigSent,
"waf_ip_groups_sent", wafIPGroupsSent,
)
}
+1
View File
@@ -231,6 +231,7 @@ func SetApiRouter(router *gin.Engine) {
authorizedRoute.GET("/ws", controller.AgentWebSocket)
authorizedRoute.POST("/nodes/heartbeat", controller.AgentHeartbeat)
authorizedRoute.GET("/config-versions/active", controller.AgentGetActiveConfig)
authorizedRoute.POST("/waf/ip-groups/sync", controller.AgentSyncWAFIPGroups)
authorizedRoute.POST("/apply-logs", controller.AgentReportApplyLog)
}
}
+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)
}
}
@@ -186,6 +186,8 @@ func RenderWAFConfig(snapshot WAFDocument) (string, error) {
BlockResponseBody string `json:"block_response_body"`
IPWhitelist []string `json:"ip_whitelist"`
IPBlacklist []string `json:"ip_blacklist"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist"`
CountryBlacklist []string `json:"country_blacklist"`
RegionWhitelist []string `json:"region_whitelist"`
@@ -199,10 +201,6 @@ func RenderWAFConfig(snapshot WAFDocument) (string, error) {
groups := make([]wafRuntimeRuleGroup, 0, len(snapshot.RuleGroups))
globalGroupIDs := make([]uint, 0)
enabledGroupIDs := make(map[uint]struct{}, len(snapshot.RuleGroups))
ipGroupsByID := make(map[uint]WAFIPGroup, len(snapshot.IPGroups))
for _, group := range snapshot.IPGroups {
ipGroupsByID[group.ID] = group
}
for _, group := range snapshot.RuleGroups {
if !group.Enabled {
continue
@@ -221,8 +219,10 @@ func RenderWAFConfig(snapshot WAFDocument) (string, error) {
IsGlobal: group.IsGlobal,
BlockStatusCode: statusCode,
BlockResponseBody: group.BlockResponseBody,
IPWhitelist: expandWAFIPGroups(group.IPWhitelist, group.IPWhitelistGroups, ipGroupsByID),
IPBlacklist: expandWAFIPGroups(group.IPBlacklist, group.IPBlacklistGroups, ipGroupsByID),
IPWhitelist: sortedUniqueStrings(group.IPWhitelist),
IPBlacklist: sortedUniqueStrings(group.IPBlacklist),
IPWhitelistGroups: sortedUniqueUintIDs(group.IPWhitelistGroups),
IPBlacklistGroups: sortedUniqueUintIDs(group.IPBlacklistGroups),
CountryWhitelist: group.CountryWhitelist,
CountryBlacklist: group.CountryBlacklist,
RegionWhitelist: group.RegionWhitelist,
@@ -250,20 +250,19 @@ func RenderWAFConfig(snapshot WAFDocument) (string, error) {
return string(data), err
}
func expandWAFIPGroups(direct []string, groupIDs []uint, ipGroupsByID map[uint]WAFIPGroup) []string {
items := append([]string{}, direct...)
for _, id := range groupIDs {
group, ok := ipGroupsByID[id]
if !ok || !group.Enabled {
continue
}
items = append(items, group.IPList...)
}
func sortedUniqueStrings(values []string) []string {
items := append([]string{}, values...)
items = uniqueStrings(items)
sort.Strings(items)
return items
}
func sortedUniqueUintIDs(values []uint) []uint {
items := uniqueUintIDs(values)
sort.Slice(items, func(i, j int) bool { return items[i] < items[j] })
return items
}
func ChecksumBundle(mainConfig string, routeConfig string, supportFiles []SupportFile) string {
var builder strings.Builder
builder.WriteString(mainConfig)
@@ -188,7 +188,7 @@ export function RuleEntryModal({
选择 IP 组
</h3>
<p className="mt-1 text-xs leading-5 text-[var(--foreground-secondary)]">
被引用的 IP 组会在发布配置时展开到 WAF 运行时名单。
发布版本只保存引用 ID,IP 组成员由 Agent 按 checksum 差异同步。
</p>
</div>
<span className="rounded-full border border-[var(--border-default)] px-2.5 py-1 text-xs font-medium text-[var(--foreground-secondary)]">
@@ -253,8 +253,8 @@ describe('WAF IP groups', () => {
await userEvent.click(screen.getByRole('button', { name: /测试规则/ }));
expect(await screen.findByText('命中 2 个 IP。')).toBeInTheDocument();
expect(screen.getByText('203.0.113.10')).toBeInTheDocument();
expect(screen.getByText('203.0.113.11')).toBeInTheDocument();
expect(screen.getByText(/203\.0\.113\.10/)).toBeInTheDocument();
expect(screen.getByText(/203\.0\.113\.11/)).toBeInTheDocument();
expect(testMock).toHaveBeenCalledWith({
auto_config: expect.objectContaining({
lookback_minutes: 60,