mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-05 15:26:36 +08:00
feat(waf): complete composable rule orchestration
Add the React Flow rule editor, ordered graph APIs and runtime DAG execution.\n\nPublish rules only on OpenResty reload and reconcile checksum-driven IP group snapshots in bounded shared memory.
This commit is contained in:
@@ -4,6 +4,7 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
@@ -11,43 +12,32 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
)
|
||||
|
||||
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"`
|
||||
WAF openrestyrender.WAFDocument `json:"waf"`
|
||||
}
|
||||
|
||||
type runtimeIPMatchConfig struct {
|
||||
IPs []string `json:"ips,omitempty"`
|
||||
CIDRs []string `json:"cidrs,omitempty"`
|
||||
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
|
||||
}
|
||||
|
||||
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
|
||||
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
|
||||
return buildAgentWAFIPGroups(ctx, ids)
|
||||
return validatedAgentWAFIPGroups(ctx, ids, false)
|
||||
}
|
||||
|
||||
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
|
||||
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
|
||||
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)
|
||||
groups, err := validatedAgentWAFIPGroups(ctx, ids, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -61,6 +51,45 @@ func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[s
|
||||
return changed, nil
|
||||
}
|
||||
|
||||
func validatedAgentWAFIPGroups(ctx context.Context, ids []uint, fallbackToActive bool) ([]WAFIPGroup, error) {
|
||||
targetIDs := uniqueUintIDs(ids)
|
||||
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(targetIDs) == 0 && fallbackToActive {
|
||||
targetIDs = activeIDs
|
||||
}
|
||||
if len(targetIDs) == 0 {
|
||||
return []WAFIPGroup{}, nil
|
||||
}
|
||||
|
||||
validationIDs := uniqueUintIDs(append(append([]uint{}, activeIDs...), targetIDs...))
|
||||
allGroups, err := buildAgentWAFIPGroups(ctx, validationIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
runtimeGroups := make(map[string]protocol.WAFIPGroup, len(allGroups))
|
||||
for _, group := range allGroups {
|
||||
runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = group
|
||||
}
|
||||
if err = protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
targetSet := make(map[uint]struct{}, len(targetIDs))
|
||||
for _, id := range targetIDs {
|
||||
targetSet[id] = struct{}{}
|
||||
}
|
||||
result := make([]WAFIPGroup, 0, len(targetIDs))
|
||||
for _, group := range allGroups {
|
||||
if _, ok := targetSet[group.ID]; ok {
|
||||
result = append(result, group)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
|
||||
ids = uniqueUintIDs(ids)
|
||||
if len(ids) == 0 {
|
||||
@@ -142,6 +171,8 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
|
||||
}
|
||||
idSet := make(map[uint]struct{})
|
||||
for _, group := range snapshot.WAF.RuleGroups {
|
||||
// Retain legacy flattened references while older active snapshots may
|
||||
// still exist during a rolling Server upgrade.
|
||||
for _, id := range group.IPWhitelistGroups {
|
||||
if id > 0 {
|
||||
idSet[id] = struct{}{}
|
||||
@@ -152,6 +183,20 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
for nodeID, node := range group.Graph.Nodes {
|
||||
if node.Type != "ip_match" {
|
||||
continue
|
||||
}
|
||||
ids, err := runtimeIPMatchGroupIDs(node.Config)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("活动配置 WAF 规则 %d 节点 %s 的 IP 匹配配置无效: %w", group.ID, nodeID, err)
|
||||
}
|
||||
for _, id := range ids {
|
||||
if id > 0 {
|
||||
idSet[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
ids := make([]uint, 0, len(idSet))
|
||||
for id := range idSet {
|
||||
@@ -161,6 +206,16 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func runtimeIPMatchGroupIDs(raw json.RawMessage) ([]uint, error) {
|
||||
var config runtimeIPMatchConfig
|
||||
decoder := json.NewDecoder(bytes.NewReader(raw))
|
||||
decoder.DisallowUnknownFields()
|
||||
if err := decoder.Decode(&config); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return config.IPGroupIDs, nil
|
||||
}
|
||||
|
||||
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
|
||||
text := strings.TrimSpace(snapshotJSON)
|
||||
if text == "" {
|
||||
@@ -171,7 +226,7 @@ func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, erro
|
||||
return nil, err
|
||||
}
|
||||
if snapshot.WAF.RuleGroups == nil {
|
||||
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
|
||||
snapshot.WAF.RuleGroups = []openrestyrender.WAFRuleGroup{}
|
||||
}
|
||||
return &snapshot, nil
|
||||
}
|
||||
|
||||
@@ -7,10 +7,12 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -63,6 +65,115 @@ func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID
|
||||
}).Error)
|
||||
}
|
||||
|
||||
func seedActiveConfigWithWAFGraphIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
|
||||
t.Helper()
|
||||
|
||||
snapshot := map[string]any{
|
||||
"routes": []any{},
|
||||
"waf": map[string]any{
|
||||
"rule_groups": []map[string]any{
|
||||
{
|
||||
"id": 1,
|
||||
"name": "graph refs",
|
||||
"enabled": true,
|
||||
"graph": map[string]any{
|
||||
"entry": "start",
|
||||
"nodes": map[string]any{
|
||||
"start": map[string]any{
|
||||
"type": "start",
|
||||
"config": map[string]any{},
|
||||
"next": map[string]string{"next": "match"},
|
||||
},
|
||||
"match": map[string]any{
|
||||
"type": "ip_match",
|
||||
"config": map[string]any{
|
||||
"ip_group_ids": []uint{ipGroupID},
|
||||
},
|
||||
"next": map[string]string{"true": "allow", "false": "allow"},
|
||||
},
|
||||
"allow": map[string]any{
|
||||
"type": "allow",
|
||||
"config": map[string]any{},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"bindings": []any{},
|
||||
},
|
||||
}
|
||||
snapshotJSON, err := json.Marshal(snapshot)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
|
||||
Version: "20260713-graph-001",
|
||||
SnapshotJSON: string(snapshotJSON),
|
||||
Checksum: "graph-test-checksum",
|
||||
IsActive: true,
|
||||
}).Error)
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentDiscoversGraphReferences(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: "graph runtime group",
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `["192.0.2.88"]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
seedActiveConfigWithWAFGraphIPGroup(t, ctx, ipGroup.ID)
|
||||
|
||||
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, groups, 1)
|
||||
assert.Equal(t, ipGroup.ID, groups[0].ID)
|
||||
assert.Equal(t, []string{"192.0.2.88"}, groups[0].IPList)
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentRejectsMalformedIPMatchConfig(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
|
||||
Version: "20260713-malformed-001",
|
||||
SnapshotJSON: `{"waf":{"rule_groups":[{"id":7,"graph":{"entry":"match","nodes":{` +
|
||||
`"match":{"type":"ip_match","config":{"ip_group_ids":"not-an-array"}}}}}],"bindings":[]}}`,
|
||||
Checksum: "malformed-test-checksum",
|
||||
IsActive: true,
|
||||
}).Error)
|
||||
|
||||
_, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
|
||||
require.ErrorContains(t, err, "规则 7 节点 match")
|
||||
require.ErrorContains(t, err, "IP 匹配配置无效")
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentRejectsOversizedSnapshotBeforeChecksumDelta(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
ipGroup := &model.OpenFlareWAFIPGroup{
|
||||
Name: strings.Repeat("x", protocol.MaxWAFIPGroupSnapshotBytes),
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: `[]`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
|
||||
agentGroup, err := buildAgentWAFIPGroup(ipGroup)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = ChangedWAFIPGroupsForAgent(ctx, []uint{ipGroup.ID}, map[string]string{
|
||||
strconv.FormatUint(uint64(ipGroup.ID), 10): agentGroup.Checksum,
|
||||
})
|
||||
require.ErrorContains(t, err, "WAF IP 组快照大小")
|
||||
require.ErrorContains(t, err, "超过上限")
|
||||
}
|
||||
|
||||
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
|
||||
cleanup := setupWAFIPGroupTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -151,13 +152,7 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
|
||||
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
customGroup := &model.OpenFlareWAFRuleGroup{
|
||||
Name: "pow-group",
|
||||
Enabled: true,
|
||||
PoWEnabled: true,
|
||||
PoWConfig: `{"difficulty":4,"algorithm":"fast","session_ttl":600,"challenge_ttl":300}`,
|
||||
}
|
||||
require.NoError(t, model.CreateOpenFlareWAFRuleGroup(ctx, customGroup))
|
||||
customGroup := createSnapshotRule(t, ctx, "pow-group", waf.DefaultRuleGraph())
|
||||
require.NoError(t, model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID}))
|
||||
|
||||
bundle, err := buildCurrentConfigBundle(ctx, true)
|
||||
@@ -177,20 +172,24 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
|
||||
}
|
||||
assert.True(t, found, "expected WAF binding for enabled route")
|
||||
|
||||
var wafRuntime struct {
|
||||
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
|
||||
}
|
||||
var wafRuntime openrestyrender.WAFDocument
|
||||
foundWAFConfig := false
|
||||
for _, file := range bundle.SupportFiles {
|
||||
if file.Path != "waf_config.json" {
|
||||
continue
|
||||
}
|
||||
foundWAFConfig = true
|
||||
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
|
||||
}
|
||||
require.Contains(t, wafRuntime.SiteRuleGroups, "example.com")
|
||||
require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], customGroup.ID)
|
||||
require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], globalGroup.ID)
|
||||
require.True(t, foundWAFConfig, "expected rendered WAF support file")
|
||||
require.NotEmpty(t, wafRuntime.RuleGroups)
|
||||
assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
|
||||
assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
|
||||
require.Len(t, wafRuntime.Bindings, 1)
|
||||
assert.Equal(t, route.ID, wafRuntime.Bindings[0].RouteID)
|
||||
assert.Equal(t, "example.com", wafRuntime.Bindings[0].SiteName)
|
||||
assert.Equal(t, []uint{customGroup.ID}, wafRuntime.Bindings[0].RuleGroupIDs)
|
||||
assert.Contains(t, bundle.RouteConfig, `set $openflare_waf_site "example.com"`)
|
||||
assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`)
|
||||
}
|
||||
|
||||
func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testing.T) {
|
||||
@@ -210,34 +209,30 @@ func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testi
|
||||
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
|
||||
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
|
||||
require.NoError(t, err)
|
||||
globalGroup.PoWEnabled = true
|
||||
globalGroup.PoWConfig = `{"difficulty":4,"algorithm":"fast","session_ttl":600,"challenge_ttl":300}`
|
||||
require.NoError(t, model.UpdateOpenFlareWAFRuleGroup(ctx, globalGroup))
|
||||
graphJSON, err := json.Marshal(snapshotPoWGraph())
|
||||
require.NoError(t, err)
|
||||
globalGroup.Graph = string(graphJSON)
|
||||
require.NoError(t, db.DB(ctx).Model(globalGroup).Update("graph", globalGroup.Graph).Error)
|
||||
|
||||
bundle, err := buildCurrentConfigBundle(ctx, true)
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`)
|
||||
|
||||
var wafRuntime struct {
|
||||
RuleGroups []struct {
|
||||
ID uint `json:"id"`
|
||||
PoWEnabled bool `json:"pow_enabled"`
|
||||
PoWConfig *struct {
|
||||
Difficulty int `json:"difficulty"`
|
||||
} `json:"pow_config"`
|
||||
} `json:"rule_groups"`
|
||||
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
|
||||
}
|
||||
var wafRuntime openrestyrender.WAFDocument
|
||||
foundWAFConfig := false
|
||||
for _, file := range bundle.SupportFiles {
|
||||
if file.Path != "waf_config.json" {
|
||||
continue
|
||||
}
|
||||
foundWAFConfig = true
|
||||
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
|
||||
}
|
||||
require.Contains(t, wafRuntime.SiteRuleGroups, "pow-global.example.com")
|
||||
require.Contains(t, wafRuntime.SiteRuleGroups["pow-global.example.com"], globalGroup.ID)
|
||||
require.True(t, foundWAFConfig, "expected rendered WAF support file")
|
||||
require.NotEmpty(t, wafRuntime.RuleGroups)
|
||||
assert.True(t, wafRuntime.RuleGroups[0].PoWEnabled)
|
||||
require.NotNil(t, wafRuntime.RuleGroups[0].PoWConfig)
|
||||
assert.Equal(t, 4, wafRuntime.RuleGroups[0].PoWConfig.Difficulty)
|
||||
assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
|
||||
assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
|
||||
assert.Equal(t, string(waf.RuleNodePoW), wafRuntime.RuleGroups[0].Graph.Nodes["pow"].Type)
|
||||
require.Len(t, wafRuntime.Bindings, 1)
|
||||
assert.Equal(t, "pow-global.example.com", wafRuntime.Bindings[0].SiteName)
|
||||
assert.Empty(t, wafRuntime.Bindings[0].RuleGroupIDs)
|
||||
require.NotEmpty(t, bundle.WAFSnapshot.RuleGroups)
|
||||
assert.Equal(t, waf.RuleNodePoW, bundle.WAFSnapshot.RuleGroups[0].Graph.Nodes["pow"].Type)
|
||||
}
|
||||
|
||||
@@ -9,17 +9,21 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
supportFilesPerCertificate = 2
|
||||
supportFilesPerCertificate = 2
|
||||
wafIPGroupChecksumHexLength = 64
|
||||
|
||||
// OpenResty 默认配置值
|
||||
defaultOpenRestyReturnStatus = 421
|
||||
@@ -66,22 +70,11 @@ type snapshotRoute struct {
|
||||
}
|
||||
|
||||
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"`
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IsGlobal bool `json:"is_global"`
|
||||
Graph waf.RuntimeRuleGraph `json:"graph"`
|
||||
}
|
||||
|
||||
type snapshotWAFIPGroup struct {
|
||||
@@ -311,35 +304,37 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
|
||||
if err := waf.EnsureDefaultRuleGroup(ctx); err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
views, err := waf.ListRuleGroups(ctx)
|
||||
groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
|
||||
if err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views))
|
||||
for _, view := range views {
|
||||
if !view.Enabled {
|
||||
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(groups))
|
||||
referencedIPGroupIDs := make(map[uint]struct{})
|
||||
enabledRuleIDs := make(map[uint]struct{})
|
||||
for _, group := range groups {
|
||||
if !group.Enabled {
|
||||
continue
|
||||
}
|
||||
var editorGraph waf.RuleGraph
|
||||
if err = json.Unmarshal([]byte(group.Graph), &editorGraph); err != nil {
|
||||
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图数据无效: %w", group.Name, err)
|
||||
}
|
||||
if err = waf.ValidateRuleGraph(ctx, editorGraph, snapshotWAFIPGroupExists); err != nil {
|
||||
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图无效: %w", group.Name, err)
|
||||
}
|
||||
runtimeGraph, compileErr := waf.CompileRuleGraph(editorGraph)
|
||||
if compileErr != nil {
|
||||
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 编译失败: %w", group.Name, compileErr)
|
||||
}
|
||||
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.PoWEnabled, view.PoWConfig),
|
||||
ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal, Graph: runtimeGraph,
|
||||
})
|
||||
enabledRuleIDs[group.ID] = struct{}{}
|
||||
for _, id := range waf.ReferencedIPGroupIDs(editorGraph) {
|
||||
referencedIPGroupIDs[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
ipGroups, err := buildSnapshotWAFIPGroups(ctx, ruleGroups)
|
||||
ipGroups, err := buildSnapshotWAFIPGroups(ctx, referencedIPGroupIDs)
|
||||
if err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
@@ -366,16 +361,16 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
|
||||
if _, ok := enabledRouteSiteNames[binding.ProxyRouteID]; !ok {
|
||||
continue
|
||||
}
|
||||
groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
|
||||
if _, enabled := enabledRuleIDs[binding.RuleGroupID]; enabled {
|
||||
groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
|
||||
}
|
||||
}
|
||||
bindings := make([]snapshotWAFBinding, 0, len(enabledRouteSiteNames))
|
||||
for routeID, siteName := range enabledRouteSiteNames {
|
||||
groupIDs := groupIDsByRoute[routeID]
|
||||
sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] })
|
||||
bindings = append(bindings, snapshotWAFBinding{
|
||||
RouteID: routeID,
|
||||
SiteName: siteName,
|
||||
RuleGroupIDs: groupIDs,
|
||||
RuleGroupIDs: groupIDsByRoute[routeID],
|
||||
})
|
||||
}
|
||||
sort.Slice(bindings, func(i, j int) bool {
|
||||
@@ -387,16 +382,26 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
|
||||
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{}{}
|
||||
func validateSnapshotWAFIPGroupSize(groups []snapshotWAFIPGroup) error {
|
||||
runtimeGroups := make(map[string]protocol.WAFIPGroup, len(groups))
|
||||
for _, group := range groups {
|
||||
ipList := group.IPList
|
||||
if !group.Enabled {
|
||||
ipList = []string{}
|
||||
}
|
||||
for _, id := range group.IPBlacklistGroups {
|
||||
idSet[id] = struct{}{}
|
||||
runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = protocol.WAFIPGroup{
|
||||
ID: group.ID,
|
||||
Name: group.Name,
|
||||
Type: group.Type,
|
||||
Enabled: group.Enabled,
|
||||
IPList: ipList,
|
||||
Checksum: strings.Repeat("0", wafIPGroupChecksumHexLength),
|
||||
}
|
||||
}
|
||||
return protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups)
|
||||
}
|
||||
|
||||
func buildSnapshotWAFIPGroups(ctx context.Context, idSet map[uint]struct{}) ([]snapshotWAFIPGroup, error) {
|
||||
if len(idSet) == 0 {
|
||||
return []snapshotWAFIPGroup{}, nil
|
||||
}
|
||||
@@ -431,9 +436,20 @@ func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleG
|
||||
IPList: ipList,
|
||||
})
|
||||
}
|
||||
if err = validateSnapshotWAFIPGroupSize(snapshots); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return snapshots, nil
|
||||
}
|
||||
|
||||
func snapshotWAFIPGroupExists(ctx context.Context, id uint) (bool, error) {
|
||||
group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id)
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
return group != nil, err
|
||||
}
|
||||
|
||||
func decodeIPList(raw string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
@@ -446,36 +462,6 @@ func decodeIPList(raw string) ([]string, error) {
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func convertPoWConfig(enabled bool, config *waf.PoWConfig) *openrestyrender.PoWConfig {
|
||||
if !enabled {
|
||||
return nil
|
||||
}
|
||||
if config == nil {
|
||||
defaultConfig := openrestyrender.DefaultPoWConfig()
|
||||
return &defaultConfig
|
||||
}
|
||||
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(ctx context.Context) openRestyConfigSnapshot {
|
||||
// 读取所有 OpenResty 配置,使用默认值作为降级
|
||||
getIntConfig := func(key string, defaultVal int) int {
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config_version
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestBuildSnapshotRejectsOversizedAggregateWAFIPGroups(t *testing.T) {
|
||||
cleanup := setupConfigVersionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// Each group remains below the existing 2 MiB per-subscription ceiling,
|
||||
// while the complete Agent runtime document exceeds the aggregate limit.
|
||||
ipList, err := json.Marshal(strings.Fields(strings.Repeat("192.0.2.1 ", 165000)))
|
||||
require.NoError(t, err)
|
||||
require.Less(t, len(ipList), 2<<20)
|
||||
|
||||
groupIDs := make([]uint, 0, 12)
|
||||
for index := 0; index < 12; index++ {
|
||||
group := &model.OpenFlareWAFIPGroup{
|
||||
Name: "aggregate-" + strings.Repeat("x", index),
|
||||
Type: "manual",
|
||||
Enabled: true,
|
||||
IPList: string(ipList),
|
||||
}
|
||||
require.NoError(t, db.DB(ctx).Create(group).Error)
|
||||
groupIDs = append(groupIDs, group.ID)
|
||||
}
|
||||
createSnapshotRule(t, ctx, "oversized-aggregate", snapshotIPMatchGraphForGroups(groupIDs))
|
||||
|
||||
_, err = buildSnapshotWAFDocument(ctx, nil)
|
||||
require.ErrorContains(t, err, "WAF IP 组快照大小")
|
||||
require.ErrorContains(t, err, "超过上限")
|
||||
}
|
||||
|
||||
func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) {
|
||||
cleanup := setupConfigVersionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
route := &model.ProxyRoute{SiteName: "ordered.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true}
|
||||
require.NoError(t, model.CreateProxyRouteRecord(ctx, route))
|
||||
createSnapshotZoneDomains(t, ctx, route, route.SiteName)
|
||||
|
||||
referenced := &model.OpenFlareWAFIPGroup{Name: "referenced", Type: "manual", Enabled: true, IPList: `["192.0.2.1"]`}
|
||||
unused := &model.OpenFlareWAFIPGroup{Name: "unused", Type: "manual", Enabled: true, IPList: `["198.51.100.1"]`}
|
||||
require.NoError(t, db.DB(ctx).Create(referenced).Error)
|
||||
require.NoError(t, db.DB(ctx).Create(unused).Error)
|
||||
|
||||
customA := createSnapshotRule(t, ctx, "custom-a", waf.DefaultRuleGraph())
|
||||
customB := createSnapshotRule(t, ctx, "custom-b", snapshotIPMatchGraph(referenced.ID))
|
||||
require.NoError(t, model.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, route.ID, []uint{customB.ID, customA.ID}))
|
||||
|
||||
snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, snapshot.Bindings, 1)
|
||||
assert.Equal(t, []uint{customB.ID, customA.ID}, snapshot.Bindings[0].RuleGroupIDs)
|
||||
require.Len(t, snapshot.IPGroups, 1)
|
||||
assert.Equal(t, referenced.ID, snapshot.IPGroups[0].ID)
|
||||
|
||||
var customBSnapshot *snapshotWAFRuleGroup
|
||||
for index := range snapshot.RuleGroups {
|
||||
if snapshot.RuleGroups[index].ID == customB.ID {
|
||||
customBSnapshot = &snapshot.RuleGroups[index]
|
||||
}
|
||||
}
|
||||
require.NotNil(t, customBSnapshot)
|
||||
assert.Equal(t, "start", customBSnapshot.Graph.Entry)
|
||||
assert.Equal(t, waf.RuleNodeIPMatch, customBSnapshot.Graph.Nodes["match"].Type)
|
||||
raw, err := json.Marshal(customBSnapshot)
|
||||
require.NoError(t, err)
|
||||
assert.NotContains(t, string(raw), "position")
|
||||
assert.NotContains(t, string(raw), "ip_whitelist")
|
||||
}
|
||||
|
||||
func TestBuildSnapshotRejectsInvalidWAFGraph(t *testing.T) {
|
||||
cleanup := setupConfigVersionTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
invalid := &model.OpenFlareWAFRuleGroup{Name: "invalid", Enabled: true, Graph: `{"schema_version":1,"nodes":[],"edges":[]}`, Revision: 1}
|
||||
require.NoError(t, db.DB(ctx).Create(invalid).Error)
|
||||
_, err := buildSnapshotWAFDocument(ctx, nil)
|
||||
require.ErrorContains(t, err, "invalid")
|
||||
}
|
||||
|
||||
func createSnapshotRule(t *testing.T, ctx context.Context, name string, graph waf.RuleGraph) *model.OpenFlareWAFRuleGroup {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(graph)
|
||||
require.NoError(t, err)
|
||||
rule := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: true, Graph: string(raw), Revision: 1}
|
||||
require.NoError(t, db.DB(ctx).Create(rule).Error)
|
||||
return rule
|
||||
}
|
||||
|
||||
func snapshotIPMatchGraph(ipGroupID uint) waf.RuleGraph {
|
||||
return snapshotIPMatchGraphForGroups([]uint{ipGroupID})
|
||||
}
|
||||
|
||||
func snapshotIPMatchGraphForGroups(ipGroupIDs []uint) waf.RuleGraph {
|
||||
config, _ := json.Marshal(waf.IPMatchConfig{IPGroupIDs: ipGroupIDs})
|
||||
return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
|
||||
{ID: "start", Type: waf.RuleNodeStart, Position: waf.RulePosition{X: 1, Y: 2}, Config: json.RawMessage(`{}`)},
|
||||
{ID: "match", Type: waf.RuleNodeIPMatch, Position: waf.RulePosition{X: 3, Y: 4}, Config: config},
|
||||
{ID: "allow", Type: waf.RuleNodeAllow, Position: waf.RulePosition{X: 5, Y: 6}, Config: json.RawMessage(`{}`)},
|
||||
}, Edges: []waf.RuleEdge{
|
||||
{ID: "e1", Source: "start", SourceHandle: "next", Target: "match"},
|
||||
{ID: "e2", Source: "match", SourceHandle: "true", Target: "allow"},
|
||||
{ID: "e3", Source: "match", SourceHandle: "false", Target: "allow"},
|
||||
}}
|
||||
}
|
||||
|
||||
func snapshotPoWGraph() waf.RuleGraph {
|
||||
config, _ := json.Marshal(waf.PoWNodeConfig{Algorithm: "fast", Difficulty: 4, SessionTTL: 600, ChallengeTTL: 300})
|
||||
return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
|
||||
{ID: "start", Type: waf.RuleNodeStart, Config: json.RawMessage(`{}`)},
|
||||
{ID: "pow", Type: waf.RuleNodePoW, Config: config},
|
||||
{ID: "allow", Type: waf.RuleNodeAllow, Config: json.RawMessage(`{}`)},
|
||||
}, Edges: []waf.RuleEdge{
|
||||
{ID: "e1", Source: "start", SourceHandle: "next", Target: "pow"},
|
||||
{ID: "e2", Source: "pow", SourceHandle: "next", Target: "allow"},
|
||||
}}
|
||||
}
|
||||
@@ -109,12 +109,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
|
||||
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"},
|
||||
"name": "edge-security",
|
||||
}, adminAuthHeaders(seed.Token))
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
|
||||
@@ -124,7 +119,8 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
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"])
|
||||
assert.Equal(t, float64(1), data["revision"])
|
||||
assert.NotNil(t, data["graph"])
|
||||
})
|
||||
|
||||
t.Run("WAF rule group list includes global and custom groups", func(t *testing.T) {
|
||||
@@ -174,11 +170,9 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
t,
|
||||
engine,
|
||||
http.MethodPost,
|
||||
fmt.Sprintf("%s/waf/rule-groups/%d/update", apiPath(""), ruleGroupID),
|
||||
fmt.Sprintf("%s/waf/rule-groups/%d/meta", apiPath(""), ruleGroupID),
|
||||
map[string]any{
|
||||
"name": "edge-security-updated",
|
||||
"enabled": true,
|
||||
"block_status_code": 451,
|
||||
"name": "edge-security-updated", "enabled": true,
|
||||
},
|
||||
adminAuthHeaders(seed.Token),
|
||||
)
|
||||
@@ -187,7 +181,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
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"])
|
||||
assert.Equal(t, true, data["enabled"])
|
||||
})
|
||||
|
||||
t.Run("WAF IP group create", func(t *testing.T) {
|
||||
|
||||
@@ -5,25 +5,35 @@ package waf
|
||||
|
||||
import "encoding/json"
|
||||
|
||||
// RuleGraphSchemaVersion is the current persisted rule graph schema version.
|
||||
const RuleGraphSchemaVersion = 1
|
||||
|
||||
// RuleNodeType identifies the behavior of a rule graph node.
|
||||
type RuleNodeType string
|
||||
|
||||
const (
|
||||
RuleNodeStart RuleNodeType = "start"
|
||||
RuleNodeAllow RuleNodeType = "allow"
|
||||
RuleNodeBlock RuleNodeType = "block"
|
||||
RuleNodeIPMatch RuleNodeType = "ip_match"
|
||||
// RuleNodeStart begins graph execution.
|
||||
RuleNodeStart RuleNodeType = "start"
|
||||
// RuleNodeAllow terminates execution with an allow decision.
|
||||
RuleNodeAllow RuleNodeType = "allow"
|
||||
// RuleNodeBlock terminates execution with a blocking response.
|
||||
RuleNodeBlock RuleNodeType = "block"
|
||||
// RuleNodeIPMatch branches on an IP match.
|
||||
RuleNodeIPMatch RuleNodeType = "ip_match"
|
||||
// RuleNodeGeoMatch branches on a geographic match.
|
||||
RuleNodeGeoMatch RuleNodeType = "geo_match"
|
||||
RuleNodePoW RuleNodeType = "pow"
|
||||
// RuleNodePoW runs a proof-of-work challenge before continuing.
|
||||
RuleNodePoW RuleNodeType = "pow"
|
||||
)
|
||||
|
||||
// RuleGraph is the editor-facing representation of an executable WAF graph.
|
||||
type RuleGraph struct {
|
||||
SchemaVersion int `json:"schema_version"`
|
||||
Nodes []RuleNode `json:"nodes"`
|
||||
Edges []RuleEdge `json:"edges"`
|
||||
}
|
||||
|
||||
// RuleNode stores one editor node and its type-specific configuration.
|
||||
type RuleNode struct {
|
||||
ID string `json:"id"`
|
||||
Type RuleNodeType `json:"type"`
|
||||
@@ -32,11 +42,13 @@ type RuleNode struct {
|
||||
Config json.RawMessage `json:"config"`
|
||||
}
|
||||
|
||||
// RulePosition stores a node's editor canvas coordinates.
|
||||
type RulePosition struct {
|
||||
X float64 `json:"x"`
|
||||
Y float64 `json:"y"`
|
||||
}
|
||||
|
||||
// RuleEdge connects one source handle to a target node.
|
||||
type RuleEdge struct {
|
||||
ID string `json:"id"`
|
||||
Source string `json:"source"`
|
||||
@@ -44,17 +56,20 @@ type RuleEdge struct {
|
||||
Target string `json:"target"`
|
||||
}
|
||||
|
||||
// IPMatchConfig configures literal, CIDR, and managed-group IP matching.
|
||||
type IPMatchConfig struct {
|
||||
IPs []string `json:"ips,omitempty"`
|
||||
CIDRs []string `json:"cidrs,omitempty"`
|
||||
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
|
||||
}
|
||||
|
||||
// GeoMatchConfig configures country and region matching.
|
||||
type GeoMatchConfig struct {
|
||||
Countries []string `json:"countries,omitempty"`
|
||||
Regions []string `json:"regions,omitempty"`
|
||||
}
|
||||
|
||||
// PoWNodeConfig configures a proof-of-work challenge node.
|
||||
type PoWNodeConfig struct {
|
||||
Algorithm string `json:"algorithm"`
|
||||
Difficulty int `json:"difficulty"`
|
||||
@@ -62,11 +77,13 @@ type PoWNodeConfig struct {
|
||||
ChallengeTTL int `json:"challenge_ttl"`
|
||||
}
|
||||
|
||||
// BlockNodeConfig configures a terminal blocking response.
|
||||
type BlockNodeConfig struct {
|
||||
StatusCode int `json:"status_code"`
|
||||
ResponseBody string `json:"response_body,omitempty"`
|
||||
}
|
||||
|
||||
// DefaultRuleGraph returns the minimal start-to-allow graph.
|
||||
func DefaultRuleGraph() RuleGraph {
|
||||
return RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
|
||||
{ID: "start", Type: RuleNodeStart, Position: RulePosition{X: 0, Y: 0}, Config: json.RawMessage(`{}`)},
|
||||
|
||||
@@ -26,7 +26,33 @@ var (
|
||||
regionCodePattern = regexp.MustCompile(`^[A-Z]{2}-[A-Z0-9]{1,3}$`)
|
||||
)
|
||||
|
||||
// ValidateRuleGraph validates graph structure, node configuration, references,
|
||||
// reachability, and termination before compilation.
|
||||
func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(context.Context, uint) (bool, error)) error {
|
||||
if err := validateRuleGraphLimits(graph); err != nil {
|
||||
return err
|
||||
}
|
||||
nodes, startID, err := validateRuleGraphNodes(ctx, graph.Nodes, ipGroupExists)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
outgoing, incoming, handleTargets, err := validateRuleGraphEdges(nodes, graph.Edges)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if hasRuleGraphCycle(nodes, outgoing, incoming) {
|
||||
return errors.New("规则图不能包含循环")
|
||||
}
|
||||
if err := validateRequiredHandles(graph.Nodes, handleTargets); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateRuleGraphConnectivity(graph.Nodes, startID, outgoing, incoming); err != nil {
|
||||
return err
|
||||
}
|
||||
return validateTerminalPaths(graph.Nodes, outgoing)
|
||||
}
|
||||
|
||||
func validateRuleGraphLimits(graph RuleGraph) error {
|
||||
if graph.SchemaVersion != RuleGraphSchemaVersion {
|
||||
return fmt.Errorf("规则图 schema_version 必须为 %d", RuleGraphSchemaVersion)
|
||||
}
|
||||
@@ -41,15 +67,18 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
|
||||
} else if len(raw) > maxRuleGraphBytes {
|
||||
return fmt.Errorf("规则图大小不能超过 256 KiB")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
nodes := make(map[string]RuleNode, len(graph.Nodes))
|
||||
func validateRuleGraphNodes(ctx context.Context, graphNodes []RuleNode, ipGroupExists func(context.Context, uint) (bool, error)) (map[string]RuleNode, string, error) {
|
||||
nodes := make(map[string]RuleNode, len(graphNodes))
|
||||
startCount, allowCount, startID := 0, 0, ""
|
||||
for _, node := range graph.Nodes {
|
||||
for _, node := range graphNodes {
|
||||
if strings.TrimSpace(node.ID) == "" {
|
||||
return errors.New("节点 ID 不能为空")
|
||||
return nil, "", errors.New("节点 ID 不能为空")
|
||||
}
|
||||
if _, exists := nodes[node.ID]; exists {
|
||||
return fmt.Errorf("节点 ID %s 重复", node.ID)
|
||||
return nil, "", fmt.Errorf("节点 ID %s 重复", node.ID)
|
||||
}
|
||||
nodes[node.ID] = node
|
||||
switch node.Type {
|
||||
@@ -60,67 +89,74 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
|
||||
allowCount++
|
||||
case RuleNodeBlock, RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodePoW:
|
||||
default:
|
||||
return fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
|
||||
return nil, "", fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
|
||||
}
|
||||
if err := validateRuleNodeConfig(ctx, node, ipGroupExists); err != nil {
|
||||
return err
|
||||
return nil, "", err
|
||||
}
|
||||
}
|
||||
if startCount != 1 {
|
||||
return errors.New("规则图必须恰好包含一个开始节点")
|
||||
return nil, "", errors.New("规则图必须恰好包含一个开始节点")
|
||||
}
|
||||
if allowCount != 1 {
|
||||
return errors.New("规则图必须恰好包含一个通过节点")
|
||||
return nil, "", errors.New("规则图必须恰好包含一个通过节点")
|
||||
}
|
||||
return nodes, startID, nil
|
||||
}
|
||||
|
||||
edgeIDs := make(map[string]struct{}, len(graph.Edges))
|
||||
func validateRuleGraphEdges(nodes map[string]RuleNode, graphEdges []RuleEdge) (map[string][]RuleEdge, map[string]int, map[string]int, error) {
|
||||
edgeIDs := make(map[string]struct{}, len(graphEdges))
|
||||
outgoing := make(map[string][]RuleEdge)
|
||||
incoming := make(map[string]int)
|
||||
handleTargets := make(map[string]int)
|
||||
for _, edge := range graph.Edges {
|
||||
for _, edge := range graphEdges {
|
||||
if strings.TrimSpace(edge.ID) == "" {
|
||||
return errors.New("边 ID 不能为空")
|
||||
return nil, nil, nil, errors.New("边 ID 不能为空")
|
||||
}
|
||||
if _, exists := edgeIDs[edge.ID]; exists {
|
||||
return fmt.Errorf("边 ID %s 重复", edge.ID)
|
||||
return nil, nil, nil, fmt.Errorf("边 ID %s 重复", edge.ID)
|
||||
}
|
||||
edgeIDs[edge.ID] = struct{}{}
|
||||
source, ok := nodes[edge.Source]
|
||||
if !ok {
|
||||
return fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
|
||||
return nil, nil, nil, fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
|
||||
}
|
||||
if _, ok := nodes[edge.Target]; !ok {
|
||||
return fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
|
||||
return nil, nil, nil, fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
|
||||
}
|
||||
if !validSourceHandle(source.Type, edge.SourceHandle) {
|
||||
return fmt.Errorf("边 %s 的源端口 %s 不适用于节点 %s", edge.ID, edge.SourceHandle, edge.Source)
|
||||
return nil, nil, nil, fmt.Errorf("边 %s 的源端口 %s 不适用于节点 %s", edge.ID, edge.SourceHandle, edge.Source)
|
||||
}
|
||||
key := edge.Source + "\x00" + edge.SourceHandle
|
||||
handleTargets[key]++
|
||||
if handleTargets[key] > 1 {
|
||||
return fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
|
||||
return nil, nil, nil, fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
|
||||
}
|
||||
outgoing[edge.Source] = append(outgoing[edge.Source], edge)
|
||||
incoming[edge.Target]++
|
||||
}
|
||||
if hasRuleGraphCycle(nodes, outgoing, incoming) {
|
||||
return errors.New("规则图不能包含循环")
|
||||
}
|
||||
for _, node := range graph.Nodes {
|
||||
return outgoing, incoming, handleTargets, nil
|
||||
}
|
||||
|
||||
func validateRequiredHandles(nodes []RuleNode, handleTargets map[string]int) error {
|
||||
for _, node := range nodes {
|
||||
for _, handle := range requiredHandles(node.Type) {
|
||||
if handleTargets[node.ID+"\x00"+handle] == 0 {
|
||||
return fmt.Errorf("节点 %s 的 %s 出口未连接", node.ID, handle)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateRuleGraphConnectivity(nodes []RuleNode, startID string, outgoing map[string][]RuleEdge, incoming map[string]int) error {
|
||||
reachable := walkRuleGraph(startID, outgoing)
|
||||
for _, node := range graph.Nodes {
|
||||
for _, node := range nodes {
|
||||
if !reachable[node.ID] {
|
||||
return fmt.Errorf("节点 %s 无法从开始节点到达", node.ID)
|
||||
}
|
||||
}
|
||||
for _, node := range graph.Nodes {
|
||||
for _, node := range nodes {
|
||||
if node.Type == RuleNodeStart && incoming[node.ID] != 0 {
|
||||
return fmt.Errorf("开始节点 %s 不能有入边", node.ID)
|
||||
}
|
||||
@@ -131,96 +167,129 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
|
||||
return fmt.Errorf("终止节点 %s 不能有出口", node.ID)
|
||||
}
|
||||
}
|
||||
if err := validateTerminalPaths(graph.Nodes, outgoing); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateRuleNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
|
||||
switch node.Type {
|
||||
case RuleNodeStart, RuleNodeAllow:
|
||||
var cfg struct{}
|
||||
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
|
||||
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
|
||||
}
|
||||
return validateEmptyNodeConfig(node)
|
||||
case RuleNodeIPMatch:
|
||||
var cfg IPMatchConfig
|
||||
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
|
||||
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
|
||||
}
|
||||
for _, raw := range cfg.IPs {
|
||||
if _, err := netip.ParseAddr(raw); err != nil {
|
||||
return fmt.Errorf("节点 %s 的 IP %s 无效", node.ID, raw)
|
||||
}
|
||||
}
|
||||
for _, raw := range cfg.CIDRs {
|
||||
if _, err := netip.ParsePrefix(raw); err != nil {
|
||||
return fmt.Errorf("节点 %s 的 CIDR %s 无效", node.ID, raw)
|
||||
}
|
||||
}
|
||||
for _, id := range cfg.IPGroupIDs {
|
||||
if id == 0 {
|
||||
return fmt.Errorf("节点 %s 引用的 IP 组 ID 无效", node.ID)
|
||||
}
|
||||
if exists == nil {
|
||||
return fmt.Errorf("节点 %s 无法校验 IP 组 %d", node.ID, id)
|
||||
}
|
||||
ok, err := exists(ctx, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("节点 %s 校验 IP 组 %d 失败: %w", node.ID, id, err)
|
||||
}
|
||||
if !ok {
|
||||
return fmt.Errorf("节点 %s 引用的 IP 组 %d 不存在", node.ID, id)
|
||||
}
|
||||
}
|
||||
return validateIPMatchNodeConfig(ctx, node, exists)
|
||||
case RuleNodeGeoMatch:
|
||||
var cfg GeoMatchConfig
|
||||
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
|
||||
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
|
||||
}
|
||||
for _, code := range cfg.Countries {
|
||||
if !countryCodePattern.MatchString(code) {
|
||||
return fmt.Errorf("节点 %s 的国家代码 %s 无效", node.ID, code)
|
||||
}
|
||||
}
|
||||
for _, code := range cfg.Regions {
|
||||
if !regionCodePattern.MatchString(code) {
|
||||
return fmt.Errorf("节点 %s 的地区代码 %s 无效", node.ID, code)
|
||||
}
|
||||
}
|
||||
return validateGeoMatchNodeConfig(node)
|
||||
case RuleNodePoW:
|
||||
var cfg PoWNodeConfig
|
||||
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
|
||||
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
|
||||
}
|
||||
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
|
||||
return fmt.Errorf("节点 %s 的 PoW 难度必须在 1-16 之间", node.ID)
|
||||
}
|
||||
if cfg.Algorithm != "fast" && cfg.Algorithm != "slow" {
|
||||
return fmt.Errorf("节点 %s 的 PoW 算法必须为 fast 或 slow", node.ID)
|
||||
}
|
||||
if cfg.SessionTTL < 60 {
|
||||
return fmt.Errorf("节点 %s 的 PoW 会话 TTL 不能小于 60 秒", node.ID)
|
||||
}
|
||||
if cfg.ChallengeTTL < 30 {
|
||||
return fmt.Errorf("节点 %s 的 PoW 挑战 TTL 不能小于 30 秒", node.ID)
|
||||
}
|
||||
return validatePoWNodeConfig(node)
|
||||
case RuleNodeBlock:
|
||||
var cfg BlockNodeConfig
|
||||
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
|
||||
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
|
||||
return validateBlockNodeConfig(node)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateEmptyNodeConfig(node RuleNode) error {
|
||||
var cfg struct{}
|
||||
return decodeNodeConfig(node, &cfg)
|
||||
}
|
||||
|
||||
func validateIPMatchNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
|
||||
var cfg IPMatchConfig
|
||||
if err := decodeNodeConfig(node, &cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, raw := range cfg.IPs {
|
||||
if _, err := netip.ParseAddr(raw); err != nil {
|
||||
return fmt.Errorf("节点 %s 的 IP %s 无效", node.ID, raw)
|
||||
}
|
||||
if cfg.StatusCode < 400 || cfg.StatusCode > 599 {
|
||||
return fmt.Errorf("节点 %s 的阻止状态码必须在 400-599 之间", node.ID)
|
||||
}
|
||||
for _, raw := range cfg.CIDRs {
|
||||
if _, err := netip.ParsePrefix(raw); err != nil {
|
||||
return fmt.Errorf("节点 %s 的 CIDR %s 无效", node.ID, raw)
|
||||
}
|
||||
if len([]byte(cfg.ResponseBody)) > maxWAFBlockBodyBytes {
|
||||
return fmt.Errorf("节点 %s 的阻止响应体不能超过 %d 字节", node.ID, maxWAFBlockBodyBytes)
|
||||
}
|
||||
for _, id := range cfg.IPGroupIDs {
|
||||
if err := validateIPGroupReference(ctx, node.ID, id, exists); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateIPGroupReference(ctx context.Context, nodeID string, id uint, exists func(context.Context, uint) (bool, error)) error {
|
||||
if id == 0 {
|
||||
return fmt.Errorf("节点 %s 引用的 IP 组 ID 无效", nodeID)
|
||||
}
|
||||
if exists == nil {
|
||||
return fmt.Errorf("节点 %s 无法校验 IP 组 %d", nodeID, id)
|
||||
}
|
||||
ok, err := exists(ctx, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("节点 %s 校验 IP 组 %d 失败: %w", nodeID, id, err)
|
||||
}
|
||||
if !ok {
|
||||
return fmt.Errorf("节点 %s 引用的 IP 组 %d 不存在", nodeID, id)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateGeoMatchNodeConfig(node RuleNode) error {
|
||||
var cfg GeoMatchConfig
|
||||
if err := decodeNodeConfig(node, &cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, code := range cfg.Countries {
|
||||
if !countryCodePattern.MatchString(code) {
|
||||
return fmt.Errorf("节点 %s 的国家代码 %s 无效", node.ID, code)
|
||||
}
|
||||
}
|
||||
for _, code := range cfg.Regions {
|
||||
if !regionCodePattern.MatchString(code) {
|
||||
return fmt.Errorf("节点 %s 的地区代码 %s 无效", node.ID, code)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validatePoWNodeConfig(node RuleNode) error {
|
||||
var cfg PoWNodeConfig
|
||||
if err := decodeNodeConfig(node, &cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
|
||||
return fmt.Errorf("节点 %s 的 PoW 难度必须在 1-16 之间", node.ID)
|
||||
}
|
||||
if cfg.Algorithm != powAlgorithmFast && cfg.Algorithm != powAlgorithmSlow {
|
||||
return fmt.Errorf("节点 %s 的 PoW 算法必须为 fast 或 slow", node.ID)
|
||||
}
|
||||
if cfg.SessionTTL < minPoWSessionTTLSeconds {
|
||||
return fmt.Errorf("节点 %s 的 PoW 会话 TTL 不能小于 60 秒", node.ID)
|
||||
}
|
||||
if cfg.ChallengeTTL < minPoWChallengeTTLSeconds {
|
||||
return fmt.Errorf("节点 %s 的 PoW 挑战 TTL 不能小于 30 秒", node.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateBlockNodeConfig(node RuleNode) error {
|
||||
var cfg BlockNodeConfig
|
||||
if err := decodeNodeConfig(node, &cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg.StatusCode < 400 || cfg.StatusCode > 599 {
|
||||
return fmt.Errorf("节点 %s 的阻止状态码必须在 400-599 之间", node.ID)
|
||||
}
|
||||
if len([]byte(cfg.ResponseBody)) > maxWAFBlockBodyBytes {
|
||||
return fmt.Errorf("节点 %s 的阻止响应体不能超过 %d 字节", node.ID, maxWAFBlockBodyBytes)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func decodeNodeConfig(node RuleNode, dst any) error {
|
||||
if err := decodeStrictConfig(node.Config, dst); err != nil {
|
||||
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func decodeStrictConfig(raw json.RawMessage, dst any) error {
|
||||
trimmed := bytes.TrimSpace(raw)
|
||||
if bytes.Equal(trimmed, []byte("null")) {
|
||||
|
||||
@@ -90,7 +90,7 @@ func syncOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGro
|
||||
case wafIPGroupTypeAutomatic:
|
||||
return syncIPGroupAutomatic(ctx, group, now)
|
||||
default:
|
||||
return nil, errors.New("只有自动和订阅类型 IP 组支持同步")
|
||||
return nil, &RuleValidationError{Err: errors.New("只有自动和订阅类型 IP 组支持同步")}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -26,6 +26,7 @@ func setupWAFTestDB(t *testing.T) func() {
|
||||
&model.OpenFlareWAFRuleGroup{},
|
||||
&model.OpenFlareWAFIPGroup{},
|
||||
&model.OpenFlareWAFRuleGroupBinding{},
|
||||
&model.OriginProxyRoute{},
|
||||
))
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
@@ -34,38 +35,6 @@ func setupWAFTestDB(t *testing.T) func() {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateRuleGroup(t *testing.T) {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
group, err := CreateRuleGroup(ctx, RuleGroupInput{
|
||||
Name: "edge guard",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 451,
|
||||
IPWhitelist: []string{" 192.0.2.1 ", "192.0.2.1", "198.51.100.0/24"},
|
||||
IPBlacklist: []string{"203.0.113.10"},
|
||||
CountryBlacklist: []string{" cn ", "CN", "us"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, group.ID)
|
||||
assert.False(t, group.IsGlobal)
|
||||
assert.Equal(t, "edge guard", group.Name)
|
||||
require.Len(t, group.IPWhitelist, 2)
|
||||
assert.Equal(t, "192.0.2.1", group.IPWhitelist[0])
|
||||
assert.Equal(t, "198.51.100.0/24", group.IPWhitelist[1])
|
||||
require.Len(t, group.CountryBlacklist, 2)
|
||||
assert.Equal(t, "CN", group.CountryBlacklist[0])
|
||||
assert.Equal(t, "US", group.CountryBlacklist[1])
|
||||
|
||||
_, err = CreateRuleGroup(ctx, RuleGroupInput{
|
||||
Name: "bad ip",
|
||||
Enabled: true,
|
||||
IPBlacklist: []string{"not-an-ip"},
|
||||
})
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestPruneIPGroupExtIPs(t *testing.T) {
|
||||
group := &model.OpenFlareWAFIPGroup{
|
||||
ExtIPs: `[{"ip":"203.0.113.10","captured_at":"2026-06-18T10:00:00Z"},{"ip":"203.0.113.11","captured_at":"2026-06-18T11:00:00Z"}]`,
|
||||
|
||||
@@ -1,92 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package waf
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func parsePoWConfigRaw(enabled bool, raw string) (PoWConfig, error) {
|
||||
if !enabled {
|
||||
return defaultPoWConfig(), nil
|
||||
}
|
||||
cfg := defaultPoWConfig()
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" || text == "{}" {
|
||||
return cfg, nil
|
||||
}
|
||||
if err := json.Unmarshal([]byte(text), &cfg); err != nil {
|
||||
return cfg, errors.New("pow_config 格式无效")
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func validatePoWCoreSettings(cfg PoWConfig) error {
|
||||
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
|
||||
return errors.New("pow_config.difficulty 必须在 1-16 之间")
|
||||
}
|
||||
if !powAlgorithmValues[cfg.Algorithm] {
|
||||
return errors.New("pow_config.algorithm 必须为 fast 或 slow")
|
||||
}
|
||||
if cfg.SessionTTL < minPoWSessionTTLSeconds {
|
||||
return errors.New("pow_config.session_ttl 不能小于 60 秒")
|
||||
}
|
||||
if cfg.ChallengeTTL < minPoWChallengeTTLSeconds {
|
||||
return errors.New("pow_config.challenge_ttl 不能小于 30 秒")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validatePoWCIDRs(cidrs []string, listName string) error {
|
||||
for _, cidr := range cidrs {
|
||||
if _, _, err := net.ParseCIDR(cidr); err != nil {
|
||||
return fmt.Errorf("pow_config %s IP CIDR 格式无效: %s", listName, cidr)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validatePoWPathRegexes(regexes []string, listName string) error {
|
||||
for _, re := range regexes {
|
||||
if _, err := regexp.Compile(re); err != nil {
|
||||
return fmt.Errorf("pow_config %s路径正则格式无效: %s", listName, re)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validatePoWIPs(ips []string, listName string) error {
|
||||
for _, ip := range ips {
|
||||
if net.ParseIP(ip) == nil {
|
||||
return fmt.Errorf("pow_config %s IP 格式无效: %s", listName, ip)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validatePoWListMutualExclusion(cfg PoWConfig) error {
|
||||
type dimension struct {
|
||||
name string
|
||||
wl []string
|
||||
bl []string
|
||||
}
|
||||
dimensions := []dimension{
|
||||
{"IP", cfg.Whitelist.IPs, cfg.Blacklist.IPs},
|
||||
{"IP CIDR", cfg.Whitelist.IPCidrs, cfg.Blacklist.IPCidrs},
|
||||
{"路径", cfg.Whitelist.Paths, cfg.Blacklist.Paths},
|
||||
{"路径正则", cfg.Whitelist.PathRegexes, cfg.Blacklist.PathRegexes},
|
||||
{"User-Agent", cfg.Whitelist.UserAgents, cfg.Blacklist.UserAgents},
|
||||
}
|
||||
for _, dim := range dimensions {
|
||||
if len(dim.wl) > 0 && len(dim.bl) > 0 {
|
||||
return fmt.Errorf("pow_config %s 不能同时配置白名单和黑名单", dim.name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -12,14 +12,6 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
func handleLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return apiutil.AbortNotFoundIfMissing(c, err, "记录不存在")
|
||||
}
|
||||
|
||||
func routeIDParam(c *gin.Context) (uint, bool) {
|
||||
raw := c.Param("route_id")
|
||||
if raw == "" {
|
||||
@@ -34,167 +26,6 @@ func routeIDParam(c *gin.Context) (uint, bool) {
|
||||
return uint(id64), true
|
||||
}
|
||||
|
||||
// ListRuleGroupsHandler 列出全部 WAF 规则组。
|
||||
// @Summary 列出 WAF 规则组
|
||||
// @Description 返回全部 WAF 规则组,需要管理员权限
|
||||
// @Tags openflare-waf
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]waf.RuleGroupView} "规则组列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups [get]
|
||||
func ListRuleGroupsHandler(c *gin.Context) {
|
||||
groups, err := ListRuleGroups(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(groups))
|
||||
}
|
||||
|
||||
// GetRuleGroupHandler 获取 WAF 规则组详情。
|
||||
// @Summary 获取 WAF 规则组详情
|
||||
// @Description 按 ID 返回 WAF 规则组详情,需要管理员权限
|
||||
// @Tags openflare-waf
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "规则组 ID"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "规则组详情"
|
||||
// @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/waf/rule-groups/{id} [get]
|
||||
func GetRuleGroupHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
group, err := GetRuleGroup(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(group))
|
||||
}
|
||||
|
||||
// CreateRuleGroupHandler 创建 WAF 规则组。
|
||||
// @Summary 创建 WAF 规则组
|
||||
// @Description 创建新的 WAF 规则组,需要管理员权限
|
||||
// @Tags openflare-waf
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body waf.RuleGroupInput true "规则组参数"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "创建成功的规则组"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups [post]
|
||||
func CreateRuleGroupHandler(c *gin.Context) {
|
||||
var input RuleGroupInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
group, err := CreateRuleGroup(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(group))
|
||||
}
|
||||
|
||||
// UpdateRuleGroupHandler 更新 WAF 规则组。
|
||||
// @Summary 更新 WAF 规则组
|
||||
// @Description 按 ID 更新 WAF 规则组,需要管理员权限
|
||||
// @Tags openflare-waf
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "规则组 ID"
|
||||
// @Param request body waf.RuleGroupInput true "规则组参数"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "更新后的规则组"
|
||||
// @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/waf/rule-groups/{id}/update [post]
|
||||
func UpdateRuleGroupHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input RuleGroupInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
group, err := UpdateRuleGroup(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(group))
|
||||
}
|
||||
|
||||
// DeleteRuleGroupHandler 删除 WAF 规则组。
|
||||
// @Summary 删除 WAF 规则组
|
||||
// @Description 按 ID 删除 WAF 规则组,需要管理员权限
|
||||
// @Tags openflare-waf
|
||||
// @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/waf/rule-groups/{id}/delete [post]
|
||||
func DeleteRuleGroupHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteRuleGroup(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// ReplaceRuleGroupSitesHandler 替换规则组绑定的站点。
|
||||
// @Summary 替换规则组站点绑定
|
||||
// @Description 替换 WAF 规则组关联的代理站点列表,需要管理员权限
|
||||
// @Tags openflare-waf
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "规则组 ID"
|
||||
// @Param request body waf.IDsRequest true "站点 ID 列表"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "更新后的规则组"
|
||||
// @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/waf/rule-groups/{id}/sites [post]
|
||||
func ReplaceRuleGroupSitesHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var request IDsRequest
|
||||
if !apiutil.BindJSON(c, &request) {
|
||||
return
|
||||
}
|
||||
group, err := ReplaceRuleGroupSites(c.Request.Context(), id, request.IDs)
|
||||
if handleLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(group))
|
||||
}
|
||||
|
||||
// GetSiteRuleGroupsHandler 获取站点的 WAF 规则组绑定。
|
||||
// @Summary 获取站点 WAF 规则组
|
||||
// @Description 返回代理站点关联的 WAF 规则组绑定,需要管理员权限
|
||||
@@ -215,7 +46,7 @@ func GetSiteRuleGroupsHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
view, err := GetSiteRuleGroups(c.Request.Context(), routeID)
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
@@ -247,7 +78,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
view, err := ReplaceSiteRuleGroups(c.Request.Context(), routeID, request.IDs)
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
@@ -267,7 +98,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
|
||||
// @Router /api/v1/d/waf/ip-groups [get]
|
||||
func ListIPGroupsHandler(c *gin.Context) {
|
||||
groups, err := ListIPGroups(c.Request.Context())
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(groups))
|
||||
@@ -293,7 +124,7 @@ func GetIPGroupHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
group, err := GetIPGroup(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(group))
|
||||
@@ -319,7 +150,7 @@ func CreateIPGroupHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
group, err := CreateIPGroup(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(group))
|
||||
@@ -351,7 +182,7 @@ func UpdateIPGroupHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
group, err := UpdateIPGroup(c.Request.Context(), id, input)
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(group))
|
||||
@@ -376,7 +207,7 @@ func DeleteIPGroupHandler(c *gin.Context) {
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteIPGroup(c.Request.Context(), id); handleLogicError(c, err) {
|
||||
if err := DeleteIPGroup(c.Request.Context(), id); handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
@@ -402,7 +233,7 @@ func SyncIPGroupHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
result, err := SyncIPGroup(c.Request.Context(), id)
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
@@ -428,8 +259,8 @@ func TestIPGroupAutoConfigHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
result, err := TestIPGroupAutoConfig(c.Request.Context(), input)
|
||||
if handleLogicError(c, err) {
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(result))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,185 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package waf
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// CreateRuleInput is the minimal payload used to create an orchestrated rule.
|
||||
type CreateRuleInput struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
// SaveRuleGraphInput atomically replaces a rule graph at the supplied revision.
|
||||
type SaveRuleGraphInput struct {
|
||||
Revision uint64 `json:"revision"`
|
||||
Graph RuleGraph `json:"graph"`
|
||||
}
|
||||
|
||||
// UpdateRuleMetaInput updates metadata without replacing the graph.
|
||||
type UpdateRuleMetaInput struct {
|
||||
Name string `json:"name"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// RuleValidationError represents a safe user-facing validation failure.
|
||||
type RuleValidationError struct{ Err error }
|
||||
|
||||
func (err *RuleValidationError) Error() string { return err.Err.Error() }
|
||||
func (err *RuleValidationError) Unwrap() error { return err.Err }
|
||||
|
||||
// RuleView is the API representation of an orchestrated WAF rule.
|
||||
type RuleView struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Enabled bool `json:"enabled"`
|
||||
IsGlobal bool `json:"is_global"`
|
||||
Graph RuleGraph `json:"graph"`
|
||||
Revision uint64 `json:"revision"`
|
||||
AppliedSiteIDs []uint `json:"applied_site_ids"`
|
||||
AppliedSiteCount int `json:"applied_site_count"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ListRules returns all orchestrated WAF rules.
|
||||
func ListRules(ctx context.Context) ([]RuleView, error) {
|
||||
if err := EnsureDefaultRuleGroup(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := loadRuleGroupBindings(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]RuleView, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
view, buildErr := buildRuleView(group, bindings[group.ID])
|
||||
if buildErr != nil {
|
||||
return nil, buildErr
|
||||
}
|
||||
views = append(views, view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// GetRule returns one orchestrated WAF rule.
|
||||
func GetRule(ctx context.Context, id uint) (*RuleView, error) {
|
||||
group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := loadRuleGroupBindings(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
view, err := buildRuleView(group, bindings[group.ID])
|
||||
return &view, err
|
||||
}
|
||||
|
||||
// CreateRule creates a disabled custom rule with the safe default graph.
|
||||
func CreateRule(ctx context.Context, input CreateRuleInput) (*RuleView, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, &RuleValidationError{Err: errors.New("WAF 规则名称不能为空")}
|
||||
}
|
||||
raw, err := json.Marshal(DefaultRuleGraph())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
group := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: false, IsGlobal: false, Graph: string(raw), Revision: 1}
|
||||
if err = model.CreateOpenFlareWAFRuleGroup(ctx, group); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// GORM applies the model's database default to a false bool on Create, so
|
||||
// explicitly persist the safe disabled state after the row has an ID.
|
||||
group.Enabled = false
|
||||
if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetRule(ctx, group.ID)
|
||||
}
|
||||
|
||||
// UpdateRuleMeta updates rule metadata without touching its graph revision.
|
||||
func UpdateRuleMeta(ctx context.Context, id uint, input UpdateRuleMetaInput) (*RuleView, error) {
|
||||
group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, &RuleValidationError{Err: errors.New("WAF 规则名称不能为空")}
|
||||
}
|
||||
group.Name, group.Enabled = name, input.Enabled
|
||||
if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetRule(ctx, id)
|
||||
}
|
||||
|
||||
// DeleteRuleGroup deletes a non-global orchestrated WAF rule.
|
||||
func DeleteRuleGroup(ctx context.Context, id uint) error {
|
||||
group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if group.IsGlobal {
|
||||
return &RuleValidationError{Err: errors.New("全局 WAF 规则不能删除")}
|
||||
}
|
||||
return model.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, id)
|
||||
}
|
||||
|
||||
// SaveRuleGraph validates and atomically replaces a rule graph.
|
||||
func SaveRuleGraph(ctx context.Context, id uint, input SaveRuleGraphInput) (*RuleView, error) {
|
||||
if _, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := ValidateRuleGraph(ctx, input.Graph, ruleIPGroupExists); err != nil {
|
||||
return nil, &RuleValidationError{Err: fmt.Errorf("规则图无效: %w", err)}
|
||||
}
|
||||
raw, err := json.Marshal(input.Graph)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err = model.UpdateOpenFlareWAFRuleGraph(ctx, id, input.Revision, string(raw)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetRule(ctx, id)
|
||||
}
|
||||
|
||||
func ruleIPGroupExists(ctx context.Context, id uint) (bool, error) {
|
||||
_, err := model.GetOpenFlareWAFIPGroupByID(ctx, id)
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
return err == nil, err
|
||||
}
|
||||
|
||||
func buildRuleView(group *model.OpenFlareWAFRuleGroup, appliedSiteIDs []uint) (RuleView, error) {
|
||||
if group == nil {
|
||||
return RuleView{}, errors.New("waf rule is nil")
|
||||
}
|
||||
graph := DefaultRuleGraph()
|
||||
if strings.TrimSpace(group.Graph) != "" {
|
||||
if err := json.Unmarshal([]byte(group.Graph), &graph); err != nil {
|
||||
return RuleView{}, err
|
||||
}
|
||||
}
|
||||
ids := append([]uint(nil), appliedSiteIDs...)
|
||||
return RuleView{ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal,
|
||||
Graph: graph, Revision: group.Revision, AppliedSiteIDs: ids, AppliedSiteCount: len(ids),
|
||||
CreatedAt: group.CreatedAt.Format(time.RFC3339), UpdatedAt: group.UpdatedAt.Format(time.RFC3339)}, nil
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package waf
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"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/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestDeleteIPGroupRejectsGraphReference(t *testing.T) {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
group, err := CreateIPGroup(ctx, IPGroupInput{Name: "trusted", Type: wafIPGroupTypeManual, Enabled: true})
|
||||
require.NoError(t, err)
|
||||
rule, err := CreateRule(ctx, CreateRuleInput{Name: "guard"})
|
||||
require.NoError(t, err)
|
||||
graph := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
|
||||
{ID: "start", Type: RuleNodeStart, Config: json.RawMessage(`{}`)},
|
||||
{ID: "match", Type: RuleNodeIPMatch, Config: json.RawMessage(`{"ip_group_ids":[` + strconv.FormatUint(uint64(group.ID), 10) + `]}`)},
|
||||
{ID: "allow", Type: RuleNodeAllow, Config: json.RawMessage(`{}`)},
|
||||
}, Edges: []RuleEdge{
|
||||
{ID: "e1", Source: "start", SourceHandle: "next", Target: "match"},
|
||||
{ID: "e2", Source: "match", SourceHandle: "true", Target: "allow"},
|
||||
{ID: "e3", Source: "match", SourceHandle: "false", Target: "allow"},
|
||||
}}
|
||||
_, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: graph})
|
||||
require.NoError(t, err)
|
||||
|
||||
view, err := GetIPGroup(ctx, group.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 1, view.ReferencedByRuleCount)
|
||||
require.ErrorContains(t, DeleteIPGroup(ctx, group.ID), "已被 WAF 规则引用")
|
||||
}
|
||||
|
||||
func TestRuleHandlersMapFailures(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
tests := []struct {
|
||||
name string
|
||||
method string
|
||||
path string
|
||||
body string
|
||||
setup func(t *testing.T) func()
|
||||
want int
|
||||
}{
|
||||
{name: "invalid id", method: http.MethodGet, path: "/rules/nope", setup: setupWAFTestDB, want: http.StatusBadRequest},
|
||||
{name: "malformed json", method: http.MethodPost, path: "/rules", body: `{`, setup: setupWAFTestDB, want: http.StatusBadRequest},
|
||||
{name: "invalid graph", method: http.MethodPost, path: "/rules/1/graph", body: `{"revision":1,"graph":{"schema_version":1,"nodes":[],"edges":[]}}`, setup: func(t *testing.T) func() {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
_, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
|
||||
require.NoError(t, err)
|
||||
return cleanup
|
||||
}, want: http.StatusBadRequest},
|
||||
{name: "manual IP group sync", method: http.MethodPost, path: "/ip-groups/1/sync", setup: func(t *testing.T) func() {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
_, err := CreateIPGroup(context.Background(), IPGroupInput{Name: "manual", Type: wafIPGroupTypeManual, Enabled: true})
|
||||
require.NoError(t, err)
|
||||
return cleanup
|
||||
}, want: http.StatusBadRequest},
|
||||
{name: "missing", method: http.MethodGet, path: "/rules/999", setup: setupWAFTestDB, want: http.StatusNotFound},
|
||||
{name: "conflict", method: http.MethodPost, path: "/rules/1/graph", body: mustGraphRequest(t, 0), setup: func(t *testing.T) func() {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
_, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
|
||||
require.NoError(t, err)
|
||||
return cleanup
|
||||
}, want: http.StatusConflict},
|
||||
{name: "database failure", method: http.MethodGet, path: "/rules", setup: func(t *testing.T) func() { db.SetDB(nil); return func() {} }, want: http.StatusInternalServerError},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cleanup := tt.setup(t)
|
||||
defer cleanup()
|
||||
router := gin.New()
|
||||
router.Use(response.ErrorHandlerMiddleware())
|
||||
router.GET("/rules", ListRulesHandler)
|
||||
router.POST("/rules", CreateRuleHandler)
|
||||
router.GET("/rules/:id", GetRuleHandler)
|
||||
router.POST("/rules/:id/graph", SaveRuleGraphHandler)
|
||||
router.POST("/ip-groups/:id/sync", SyncIPGroupHandler)
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
router.ServeHTTP(rec, req)
|
||||
assert.Equal(t, tt.want, rec.Code, rec.Body.String())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func mustGraphRequest(t *testing.T, revision uint64) string {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(SaveRuleGraphInput{Revision: revision, Graph: DefaultRuleGraph()})
|
||||
require.NoError(t, err)
|
||||
return string(raw)
|
||||
}
|
||||
|
||||
func TestCreateRuleCreatesDefaultGraph(t *testing.T) {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
rule, err := CreateRule(context.Background(), CreateRuleInput{Name: " edge guard "})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "edge guard", rule.Name)
|
||||
assert.False(t, rule.Enabled)
|
||||
assert.Equal(t, uint64(1), rule.Revision)
|
||||
assert.Equal(t, DefaultRuleGraph(), rule.Graph)
|
||||
}
|
||||
|
||||
func TestCreateRuleRejectsEmptyName(t *testing.T) {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
defer cleanup()
|
||||
|
||||
_, err := CreateRule(context.Background(), CreateRuleInput{Name: " "})
|
||||
require.ErrorContains(t, err, "名称不能为空")
|
||||
}
|
||||
|
||||
func TestSaveRuleGraphValidationAndRevisionConflict(t *testing.T) {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
rule, err := CreateRule(ctx, CreateRuleInput{Name: "guard"})
|
||||
require.NoError(t, err)
|
||||
invalid := DefaultRuleGraph()
|
||||
invalid.Edges = nil
|
||||
_, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: invalid})
|
||||
require.Error(t, err)
|
||||
|
||||
updated, err := SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: DefaultRuleGraph()})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(2), updated.Revision)
|
||||
_, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: DefaultRuleGraph()})
|
||||
assert.ErrorIs(t, err, model.ErrWAFRuleRevisionConflict)
|
||||
}
|
||||
|
||||
func TestReplaceSiteRuleGroupsPreservesOrderAndRejectsGlobal(t *testing.T) {
|
||||
cleanup := setupWAFTestDB(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
require.NoError(t, db.DB(ctx).Create(&model.OriginProxyRoute{ID: 7, Domain: "example.com"}).Error)
|
||||
first, err := CreateRule(ctx, CreateRuleInput{Name: "first"})
|
||||
require.NoError(t, err)
|
||||
second, err := CreateRule(ctx, CreateRuleInput{Name: "second"})
|
||||
require.NoError(t, err)
|
||||
third, err := CreateRule(ctx, CreateRuleInput{Name: "third"})
|
||||
require.NoError(t, err)
|
||||
|
||||
view, err := ReplaceSiteRuleGroups(ctx, 7, []uint{third.ID, first.ID, second.ID, first.ID})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []uint{third.ID, first.ID, second.ID}, view.AppliedIDs)
|
||||
|
||||
require.NoError(t, EnsureDefaultRuleGroup(ctx))
|
||||
global, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
|
||||
require.NoError(t, err)
|
||||
_, err = ReplaceSiteRuleGroups(ctx, 7, []uint{global.ID, second.ID})
|
||||
require.Error(t, err)
|
||||
assert.False(t, errors.Is(err, model.ErrWAFRuleRevisionConflict))
|
||||
assert.Equal(t, []uint{third.ID, first.ID, second.ID}, mustListSiteRuleGroupIDs(t, ctx, 7))
|
||||
}
|
||||
|
||||
func mustListSiteRuleGroupIDs(t *testing.T, ctx context.Context, routeID uint) []uint {
|
||||
t.Helper()
|
||||
ids, err := ListSiteRuleGroupIDs(ctx, routeID)
|
||||
require.NoError(t, err)
|
||||
return ids
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package waf
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"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/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func handleRuleError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
var validation *RuleValidationError
|
||||
switch {
|
||||
case errors.As(err, &validation):
|
||||
response.AbortBadRequest(c, validation.Error())
|
||||
case errors.Is(err, model.ErrWAFRuleRevisionConflict):
|
||||
response.AbortConflict(c, "规则已被其他操作更新,请重新加载")
|
||||
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||
response.AbortNotFound(c, "WAF 规则不存在")
|
||||
default:
|
||||
logger.ErrorF(c.Request.Context(), "[OpenFlareWAF] rule API failed: %v", err)
|
||||
response.AbortInternal(c, "WAF 规则操作失败")
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// ListRulesHandler lists orchestrated WAF rules.
|
||||
// @Summary 列出 WAF 规则
|
||||
// @Tags openflare-waf
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]waf.RuleView} "规则列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups [get]
|
||||
func ListRulesHandler(c *gin.Context) {
|
||||
rules, err := ListRules(c.Request.Context())
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(rules))
|
||||
}
|
||||
|
||||
// GetRuleHandler gets an orchestrated WAF rule.
|
||||
// @Summary 获取 WAF 规则详情
|
||||
// @Tags openflare-waf
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "规则 ID"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleView} "规则详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups/{id} [get]
|
||||
func GetRuleHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
rule, err := GetRule(c.Request.Context(), id)
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(rule))
|
||||
}
|
||||
|
||||
// CreateRuleHandler creates an orchestrated WAF rule from a name only.
|
||||
// @Summary 创建 WAF 规则
|
||||
// @Tags openflare-waf
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body waf.CreateRuleInput true "规则名称"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleView} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups [post]
|
||||
func CreateRuleHandler(c *gin.Context) {
|
||||
var input CreateRuleInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
rule, err := CreateRule(c.Request.Context(), input)
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(rule))
|
||||
}
|
||||
|
||||
// UpdateRuleMetaHandler updates rule name and enabled state.
|
||||
// @Summary 更新 WAF 规则元数据
|
||||
// @Tags openflare-waf
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "规则 ID"
|
||||
// @Param request body waf.UpdateRuleMetaInput true "规则元数据"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleView} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups/{id}/meta [post]
|
||||
func UpdateRuleMetaHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input UpdateRuleMetaInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
rule, err := UpdateRuleMeta(c.Request.Context(), id, input)
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(rule))
|
||||
}
|
||||
|
||||
// SaveRuleGraphHandler saves a complete versioned rule graph.
|
||||
// @Summary 保存 WAF 规则图
|
||||
// @Tags openflare-waf
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "规则 ID"
|
||||
// @Param request body waf.SaveRuleGraphInput true "规则图和修订号"
|
||||
// @Success 200 {object} response.Any{data=waf.RuleView} "保存成功"
|
||||
// @Failure 400 {object} response.Any "参数或规则图错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 404 {object} response.Any "无权限或不存在"
|
||||
// @Failure 409 {object} response.Any "修订冲突"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups/{id}/graph [post]
|
||||
func SaveRuleGraphHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var input SaveRuleGraphInput
|
||||
if !apiutil.BindJSON(c, &input) {
|
||||
return
|
||||
}
|
||||
rule, err := SaveRuleGraph(c.Request.Context(), id, input)
|
||||
if handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(rule))
|
||||
}
|
||||
|
||||
// DeleteRuleHandler deletes a non-global WAF rule.
|
||||
// @Summary 删除 WAF 规则
|
||||
// @Tags openflare-waf
|
||||
// @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 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/waf/rule-groups/{id}/delete [post]
|
||||
func DeleteRuleHandler(c *gin.Context) {
|
||||
id, ok := apiutil.IDParam(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := DeleteRuleGroup(c.Request.Context(), id); handleRuleError(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
Reference in New Issue
Block a user