mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 16:16:37 +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:
@@ -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"},
|
||||
}}
|
||||
}
|
||||
Reference in New Issue
Block a user