mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 00:26:37 +08:00
feat(waf): add composable rule graph core
This commit is contained in:
@@ -0,0 +1,139 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package waf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"sort"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RuntimeRuleGraph is the compact graph loaded by the request runtime.
|
||||||
|
// Nodes are indexed by ID so execution does not scan the editor node list.
|
||||||
|
type RuntimeRuleGraph struct {
|
||||||
|
Entry string `json:"entry"`
|
||||||
|
Nodes map[string]RuntimeRuleNode `json:"nodes"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RuntimeRuleNode contains the typed runtime configuration and compiled exits
|
||||||
|
// for one graph node.
|
||||||
|
type RuntimeRuleNode struct {
|
||||||
|
Type RuleNodeType `json:"type"`
|
||||||
|
Config any `json:"config"`
|
||||||
|
Next map[string]string `json:"next,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompileRuleGraph removes editor-only fields and compiles edges into per-node
|
||||||
|
// handle lookups. Callers should validate the graph before compiling it.
|
||||||
|
func CompileRuleGraph(graph RuleGraph) (RuntimeRuleGraph, error) {
|
||||||
|
runtime := RuntimeRuleGraph{Nodes: make(map[string]RuntimeRuleNode, len(graph.Nodes))}
|
||||||
|
for _, node := range graph.Nodes {
|
||||||
|
config, err := compileRuleNodeConfig(node)
|
||||||
|
if err != nil {
|
||||||
|
return RuntimeRuleGraph{}, fmt.Errorf("节点 %s 的配置无法编译: %w", node.ID, err)
|
||||||
|
}
|
||||||
|
if node.Type == RuleNodeStart {
|
||||||
|
if runtime.Entry != "" {
|
||||||
|
return RuntimeRuleGraph{}, fmt.Errorf("规则图包含多个开始节点")
|
||||||
|
}
|
||||||
|
runtime.Entry = node.ID
|
||||||
|
}
|
||||||
|
runtime.Nodes[node.ID] = RuntimeRuleNode{Type: node.Type, Config: config}
|
||||||
|
}
|
||||||
|
if runtime.Entry == "" {
|
||||||
|
return RuntimeRuleGraph{}, fmt.Errorf("规则图缺少开始节点")
|
||||||
|
}
|
||||||
|
for _, edge := range graph.Edges {
|
||||||
|
node, ok := runtime.Nodes[edge.Source]
|
||||||
|
if !ok {
|
||||||
|
return RuntimeRuleGraph{}, fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
|
||||||
|
}
|
||||||
|
if _, ok := runtime.Nodes[edge.Target]; !ok {
|
||||||
|
return RuntimeRuleGraph{}, fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
|
||||||
|
}
|
||||||
|
if node.Next == nil {
|
||||||
|
node.Next = make(map[string]string)
|
||||||
|
}
|
||||||
|
if _, exists := node.Next[edge.SourceHandle]; exists {
|
||||||
|
return RuntimeRuleGraph{}, fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
|
||||||
|
}
|
||||||
|
node.Next[edge.SourceHandle] = edge.Target
|
||||||
|
runtime.Nodes[edge.Source] = node
|
||||||
|
}
|
||||||
|
return runtime, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func compileRuleNodeConfig(node RuleNode) (any, error) {
|
||||||
|
switch node.Type {
|
||||||
|
case RuleNodeStart, RuleNodeAllow:
|
||||||
|
var config struct{}
|
||||||
|
return config, decodeStrictConfig(node.Config, &config)
|
||||||
|
case RuleNodeIPMatch:
|
||||||
|
var config IPMatchConfig
|
||||||
|
if err := decodeStrictConfig(node.Config, &config); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
config.IPs = sortedUniqueStrings(config.IPs)
|
||||||
|
config.CIDRs = sortedUniqueStrings(config.CIDRs)
|
||||||
|
config.IPGroupIDs = sortedUniqueUints(config.IPGroupIDs)
|
||||||
|
return config, nil
|
||||||
|
case RuleNodeGeoMatch:
|
||||||
|
var config GeoMatchConfig
|
||||||
|
if err := decodeStrictConfig(node.Config, &config); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
config.Countries = sortedUniqueStrings(config.Countries)
|
||||||
|
config.Regions = sortedUniqueStrings(config.Regions)
|
||||||
|
return config, nil
|
||||||
|
case RuleNodePoW:
|
||||||
|
var config PoWNodeConfig
|
||||||
|
return config, decodeStrictConfig(node.Config, &config)
|
||||||
|
case RuleNodeBlock:
|
||||||
|
var config BlockNodeConfig
|
||||||
|
return config, decodeStrictConfig(node.Config, &config)
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("未知节点类型 %s", node.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReferencedIPGroupIDs returns the unique, sorted IP group IDs referenced by
|
||||||
|
// decodable IP match nodes. Graph validation reports malformed configurations.
|
||||||
|
func ReferencedIPGroupIDs(graph RuleGraph) []uint {
|
||||||
|
ids := make([]uint, 0)
|
||||||
|
for _, node := range graph.Nodes {
|
||||||
|
if node.Type != RuleNodeIPMatch {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
var config IPMatchConfig
|
||||||
|
if err := decodeStrictConfig(node.Config, &config); err == nil {
|
||||||
|
ids = append(ids, config.IPGroupIDs...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return sortedUniqueUints(ids)
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedUniqueStrings(values []string) []string {
|
||||||
|
result := append([]string(nil), values...)
|
||||||
|
sort.Strings(result)
|
||||||
|
write := 0
|
||||||
|
for _, value := range result {
|
||||||
|
if write == 0 || result[write-1] != value {
|
||||||
|
result[write] = value
|
||||||
|
write++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return result[:write]
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedUniqueUints(values []uint) []uint {
|
||||||
|
result := append([]uint(nil), values...)
|
||||||
|
sort.Slice(result, func(i, j int) bool { return result[i] < result[j] })
|
||||||
|
write := 0
|
||||||
|
for _, value := range result {
|
||||||
|
if write == 0 || result[write-1] != value {
|
||||||
|
result[write] = value
|
||||||
|
write++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return result[:write]
|
||||||
|
}
|
||||||
@@ -0,0 +1,125 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package waf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCompileRuleGraph(t *testing.T) {
|
||||||
|
graph := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
|
||||||
|
{ID: "allow", Type: RuleNodeAllow, Label: "通过", Position: RulePosition{X: 900, Y: 100}, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "match", Type: RuleNodeIPMatch, Label: "IP 匹配", Position: RulePosition{X: 320, Y: 100}, Config: rawConfig(`{"ips":["192.0.2.2","192.0.2.1"],"cidrs":["2001:db8:2::/48","2001:db8:1::/48"],"ip_group_ids":[7,2,7]}`)},
|
||||||
|
{ID: "start", Type: RuleNodeStart, Label: "开始", Position: RulePosition{X: 0, Y: 100}, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "block", Type: RuleNodeBlock, Label: "阻止", Position: RulePosition{X: 900, Y: 300}, Config: rawConfig(`{"status_code":403,"response_body":"denied"}`)},
|
||||||
|
}, Edges: []RuleEdge{
|
||||||
|
{ID: "edge-false", Source: "match", SourceHandle: "false", Target: "block"},
|
||||||
|
{ID: "edge-start", Source: "start", SourceHandle: "next", Target: "match"},
|
||||||
|
{ID: "edge-true", Source: "match", SourceHandle: "true", Target: "allow"},
|
||||||
|
}}
|
||||||
|
|
||||||
|
compiled, err := CompileRuleGraph(graph)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CompileRuleGraph() error = %v", err)
|
||||||
|
}
|
||||||
|
if compiled.Entry != "start" {
|
||||||
|
t.Fatalf("entry = %q, want start", compiled.Entry)
|
||||||
|
}
|
||||||
|
if len(compiled.Nodes) != 4 || compiled.Nodes["match"].Next["true"] != "allow" || compiled.Nodes["match"].Next["false"] != "block" {
|
||||||
|
t.Fatalf("unexpected compiled node index: %#v", compiled.Nodes)
|
||||||
|
}
|
||||||
|
|
||||||
|
raw, err := json.Marshal(compiled)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal compiled graph: %v", err)
|
||||||
|
}
|
||||||
|
want := `{"entry":"start","nodes":{"allow":{"type":"allow","config":{}},"block":{"type":"block","config":{"status_code":403,"response_body":"denied"}},"match":{"type":"ip_match","config":{"ips":["192.0.2.1","192.0.2.2"],"cidrs":["2001:db8:1::/48","2001:db8:2::/48"],"ip_group_ids":[2,7]},"next":{"false":"block","true":"allow"}},"start":{"type":"start","config":{},"next":{"next":"match"}}}}`
|
||||||
|
if string(raw) != want {
|
||||||
|
t.Fatalf("compiled JSON = %s\nwant = %s", raw, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompileRuleGraphIsDeterministicForNodeAndEdgeOrder(t *testing.T) {
|
||||||
|
first := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
|
||||||
|
{ID: "start", Type: RuleNodeStart, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "match", Type: RuleNodeIPMatch, Config: rawConfig(`{"ip_group_ids":[7,2]}`)},
|
||||||
|
{ID: "allow", Type: RuleNodeAllow, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "block", Type: RuleNodeBlock, Config: rawConfig(`{"status_code":403}`)},
|
||||||
|
}, Edges: []RuleEdge{
|
||||||
|
{ID: "start-match", Source: "start", SourceHandle: "next", Target: "match"},
|
||||||
|
{ID: "match-allow", Source: "match", SourceHandle: "true", Target: "allow"},
|
||||||
|
{ID: "match-block", Source: "match", SourceHandle: "false", Target: "block"},
|
||||||
|
}}
|
||||||
|
second := RuleGraph{SchemaVersion: first.SchemaVersion,
|
||||||
|
Nodes: []RuleNode{first.Nodes[3], first.Nodes[2], first.Nodes[1], first.Nodes[0]},
|
||||||
|
Edges: []RuleEdge{first.Edges[2], first.Edges[1], first.Edges[0]},
|
||||||
|
}
|
||||||
|
|
||||||
|
a, err := CompileRuleGraph(first)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := CompileRuleGraph(second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
aJSON, _ := json.Marshal(a)
|
||||||
|
bJSON, _ := json.Marshal(b)
|
||||||
|
if string(aJSON) != string(bJSON) {
|
||||||
|
t.Fatalf("equivalent graphs compiled differently:\n%s\n%s", aJSON, bJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompileRuleGraphDoesNotMutateInput(t *testing.T) {
|
||||||
|
graph := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
|
||||||
|
{ID: "start", Type: RuleNodeStart, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "match", Type: RuleNodeIPMatch, Config: rawConfig(`{"ips":["192.0.2.2","192.0.2.1"],"cidrs":["2001:db8:2::/48","2001:db8:1::/48"],"ip_group_ids":[7,2,7]}`)},
|
||||||
|
{ID: "allow", Type: RuleNodeAllow, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "block", Type: RuleNodeBlock, Config: rawConfig(`{"status_code":403}`)},
|
||||||
|
}, Edges: []RuleEdge{
|
||||||
|
{ID: "start-match", Source: "start", SourceHandle: "next", Target: "match"},
|
||||||
|
{ID: "match-allow", Source: "match", SourceHandle: "true", Target: "allow"},
|
||||||
|
{ID: "match-block", Source: "match", SourceHandle: "false", Target: "block"},
|
||||||
|
}}
|
||||||
|
before := cloneRuleGraph(t, graph)
|
||||||
|
|
||||||
|
if _, err := CompileRuleGraph(graph); err != nil {
|
||||||
|
t.Fatalf("CompileRuleGraph() error = %v", err)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(graph, before) {
|
||||||
|
t.Fatalf("CompileRuleGraph mutated input:\ngot %#v\nwant %#v", graph, before)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReferencedIPGroupIDs(t *testing.T) {
|
||||||
|
graph := RuleGraph{Nodes: []RuleNode{
|
||||||
|
{ID: "one", Type: RuleNodeIPMatch, Config: rawConfig(`{"ip_group_ids":[7,2,7]}`)},
|
||||||
|
{ID: "geo", Type: RuleNodeGeoMatch, Config: rawConfig(`{"countries":["US"]}`)},
|
||||||
|
{ID: "two", Type: RuleNodeIPMatch, Config: rawConfig(`{"ip_group_ids":[9,2]}`)},
|
||||||
|
{ID: "bad", Type: RuleNodeIPMatch, Config: rawConfig(`{"ip_group_ids":`)},
|
||||||
|
}}
|
||||||
|
|
||||||
|
before := cloneRuleGraph(t, graph)
|
||||||
|
if got, want := ReferencedIPGroupIDs(graph), []uint{2, 7, 9}; !reflect.DeepEqual(got, want) {
|
||||||
|
t.Fatalf("ReferencedIPGroupIDs() = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(graph, before) {
|
||||||
|
t.Fatalf("ReferencedIPGroupIDs mutated input:\ngot %#v\nwant %#v", graph, before)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cloneRuleGraph(t *testing.T, graph RuleGraph) RuleGraph {
|
||||||
|
t.Helper()
|
||||||
|
clone := RuleGraph{
|
||||||
|
SchemaVersion: graph.SchemaVersion,
|
||||||
|
Nodes: append([]RuleNode(nil), graph.Nodes...),
|
||||||
|
Edges: append([]RuleEdge(nil), graph.Edges...),
|
||||||
|
}
|
||||||
|
for i := range clone.Nodes {
|
||||||
|
clone.Nodes[i].Config = append(json.RawMessage(nil), graph.Nodes[i].Config...)
|
||||||
|
}
|
||||||
|
return clone
|
||||||
|
}
|
||||||
@@ -0,0 +1,75 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package waf
|
||||||
|
|
||||||
|
import "encoding/json"
|
||||||
|
|
||||||
|
const RuleGraphSchemaVersion = 1
|
||||||
|
|
||||||
|
type RuleNodeType string
|
||||||
|
|
||||||
|
const (
|
||||||
|
RuleNodeStart RuleNodeType = "start"
|
||||||
|
RuleNodeAllow RuleNodeType = "allow"
|
||||||
|
RuleNodeBlock RuleNodeType = "block"
|
||||||
|
RuleNodeIPMatch RuleNodeType = "ip_match"
|
||||||
|
RuleNodeGeoMatch RuleNodeType = "geo_match"
|
||||||
|
RuleNodePoW RuleNodeType = "pow"
|
||||||
|
)
|
||||||
|
|
||||||
|
type RuleGraph struct {
|
||||||
|
SchemaVersion int `json:"schema_version"`
|
||||||
|
Nodes []RuleNode `json:"nodes"`
|
||||||
|
Edges []RuleEdge `json:"edges"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type RuleNode struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Type RuleNodeType `json:"type"`
|
||||||
|
Label string `json:"label,omitempty"`
|
||||||
|
Position RulePosition `json:"position"`
|
||||||
|
Config json.RawMessage `json:"config"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type RulePosition struct {
|
||||||
|
X float64 `json:"x"`
|
||||||
|
Y float64 `json:"y"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type RuleEdge struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Source string `json:"source"`
|
||||||
|
SourceHandle string `json:"source_handle"`
|
||||||
|
Target string `json:"target"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type IPMatchConfig struct {
|
||||||
|
IPs []string `json:"ips,omitempty"`
|
||||||
|
CIDRs []string `json:"cidrs,omitempty"`
|
||||||
|
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type GeoMatchConfig struct {
|
||||||
|
Countries []string `json:"countries,omitempty"`
|
||||||
|
Regions []string `json:"regions,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type PoWNodeConfig struct {
|
||||||
|
Algorithm string `json:"algorithm"`
|
||||||
|
Difficulty int `json:"difficulty"`
|
||||||
|
SessionTTL int `json:"session_ttl"`
|
||||||
|
ChallengeTTL int `json:"challenge_ttl"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type BlockNodeConfig struct {
|
||||||
|
StatusCode int `json:"status_code"`
|
||||||
|
ResponseBody string `json:"response_body,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func DefaultRuleGraph() RuleGraph {
|
||||||
|
return RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
|
||||||
|
{ID: "start", Type: RuleNodeStart, Position: RulePosition{X: 0, Y: 0}, Config: json.RawMessage(`{}`)},
|
||||||
|
{ID: "allow", Type: RuleNodeAllow, Position: RulePosition{X: 320, Y: 0}, Config: json.RawMessage(`{}`)},
|
||||||
|
}, Edges: []RuleEdge{{ID: "start-allow", Source: "start", SourceHandle: "next", Target: "allow"}}}
|
||||||
|
}
|
||||||
@@ -0,0 +1,327 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package waf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/netip"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxRuleGraphNodes = 128
|
||||||
|
maxRuleGraphEdges = 256
|
||||||
|
maxRuleGraphBytes = 256 * 1024
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
countryCodePattern = regexp.MustCompile(`^[A-Z]{2}$`)
|
||||||
|
regionCodePattern = regexp.MustCompile(`^[A-Z]{2}-[A-Z0-9]{1,3}$`)
|
||||||
|
)
|
||||||
|
|
||||||
|
func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(context.Context, uint) (bool, error)) error {
|
||||||
|
if graph.SchemaVersion != RuleGraphSchemaVersion {
|
||||||
|
return fmt.Errorf("规则图 schema_version 必须为 %d", RuleGraphSchemaVersion)
|
||||||
|
}
|
||||||
|
if len(graph.Nodes) > maxRuleGraphNodes {
|
||||||
|
return fmt.Errorf("规则图节点数不能超过 %d", maxRuleGraphNodes)
|
||||||
|
}
|
||||||
|
if len(graph.Edges) > maxRuleGraphEdges {
|
||||||
|
return fmt.Errorf("规则图边数不能超过 %d", maxRuleGraphEdges)
|
||||||
|
}
|
||||||
|
if raw, err := json.Marshal(graph); err != nil {
|
||||||
|
return fmt.Errorf("规则图无法序列化: %w", err)
|
||||||
|
} else if len(raw) > maxRuleGraphBytes {
|
||||||
|
return fmt.Errorf("规则图大小不能超过 256 KiB")
|
||||||
|
}
|
||||||
|
|
||||||
|
nodes := make(map[string]RuleNode, len(graph.Nodes))
|
||||||
|
startCount, allowCount, startID := 0, 0, ""
|
||||||
|
for _, node := range graph.Nodes {
|
||||||
|
if strings.TrimSpace(node.ID) == "" {
|
||||||
|
return errors.New("节点 ID 不能为空")
|
||||||
|
}
|
||||||
|
if _, exists := nodes[node.ID]; exists {
|
||||||
|
return fmt.Errorf("节点 ID %s 重复", node.ID)
|
||||||
|
}
|
||||||
|
nodes[node.ID] = node
|
||||||
|
switch node.Type {
|
||||||
|
case RuleNodeStart:
|
||||||
|
startCount++
|
||||||
|
startID = node.ID
|
||||||
|
case RuleNodeAllow:
|
||||||
|
allowCount++
|
||||||
|
case RuleNodeBlock, RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodePoW:
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
|
||||||
|
}
|
||||||
|
if err := validateRuleNodeConfig(ctx, node, ipGroupExists); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if startCount != 1 {
|
||||||
|
return errors.New("规则图必须恰好包含一个开始节点")
|
||||||
|
}
|
||||||
|
if allowCount != 1 {
|
||||||
|
return errors.New("规则图必须恰好包含一个通过节点")
|
||||||
|
}
|
||||||
|
|
||||||
|
edgeIDs := make(map[string]struct{}, len(graph.Edges))
|
||||||
|
outgoing := make(map[string][]RuleEdge)
|
||||||
|
incoming := make(map[string]int)
|
||||||
|
handleTargets := make(map[string]int)
|
||||||
|
for _, edge := range graph.Edges {
|
||||||
|
if strings.TrimSpace(edge.ID) == "" {
|
||||||
|
return errors.New("边 ID 不能为空")
|
||||||
|
}
|
||||||
|
if _, exists := edgeIDs[edge.ID]; exists {
|
||||||
|
return 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)
|
||||||
|
}
|
||||||
|
if _, ok := nodes[edge.Target]; !ok {
|
||||||
|
return 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)
|
||||||
|
}
|
||||||
|
key := edge.Source + "\x00" + edge.SourceHandle
|
||||||
|
handleTargets[key]++
|
||||||
|
if handleTargets[key] > 1 {
|
||||||
|
return 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 {
|
||||||
|
for _, handle := range requiredHandles(node.Type) {
|
||||||
|
if handleTargets[node.ID+"\x00"+handle] == 0 {
|
||||||
|
return fmt.Errorf("节点 %s 的 %s 出口未连接", node.ID, handle)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
reachable := walkRuleGraph(startID, outgoing)
|
||||||
|
for _, node := range graph.Nodes {
|
||||||
|
if !reachable[node.ID] {
|
||||||
|
return fmt.Errorf("节点 %s 无法从开始节点到达", node.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, node := range graph.Nodes {
|
||||||
|
if node.Type == RuleNodeStart && incoming[node.ID] != 0 {
|
||||||
|
return fmt.Errorf("开始节点 %s 不能有入边", node.ID)
|
||||||
|
}
|
||||||
|
if node.Type != RuleNodeStart && incoming[node.ID] == 0 {
|
||||||
|
return fmt.Errorf("节点 %s 必须至少有一条入边", node.ID)
|
||||||
|
}
|
||||||
|
if (node.Type == RuleNodeAllow || node.Type == RuleNodeBlock) && len(outgoing[node.ID]) != 0 {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
case RuleNodeBlock:
|
||||||
|
var cfg BlockNodeConfig
|
||||||
|
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
|
||||||
|
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, 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 decodeStrictConfig(raw json.RawMessage, dst any) error {
|
||||||
|
trimmed := bytes.TrimSpace(raw)
|
||||||
|
if bytes.Equal(trimmed, []byte("null")) {
|
||||||
|
return errors.New("配置不能为 null")
|
||||||
|
}
|
||||||
|
if len(trimmed) == 0 {
|
||||||
|
raw = json.RawMessage(`{}`)
|
||||||
|
}
|
||||||
|
decoder := json.NewDecoder(bytes.NewReader(raw))
|
||||||
|
decoder.DisallowUnknownFields()
|
||||||
|
if err := decoder.Decode(dst); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) {
|
||||||
|
return errors.New("配置包含额外 JSON 值")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validSourceHandle(t RuleNodeType, handle string) bool {
|
||||||
|
for _, expected := range requiredHandles(t) {
|
||||||
|
if handle == expected {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
func requiredHandles(t RuleNodeType) []string {
|
||||||
|
switch t {
|
||||||
|
case RuleNodeStart, RuleNodePoW:
|
||||||
|
return []string{"next"}
|
||||||
|
case RuleNodeIPMatch, RuleNodeGeoMatch:
|
||||||
|
return []string{"true", "false"}
|
||||||
|
default:
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
func walkRuleGraph(start string, outgoing map[string][]RuleEdge) map[string]bool {
|
||||||
|
seen := map[string]bool{}
|
||||||
|
stack := []string{start}
|
||||||
|
for len(stack) > 0 {
|
||||||
|
id := stack[len(stack)-1]
|
||||||
|
stack = stack[:len(stack)-1]
|
||||||
|
if seen[id] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[id] = true
|
||||||
|
for _, edge := range outgoing[id] {
|
||||||
|
stack = append(stack, edge.Target)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return seen
|
||||||
|
}
|
||||||
|
func hasRuleGraphCycle(nodes map[string]RuleNode, outgoing map[string][]RuleEdge, incoming map[string]int) bool {
|
||||||
|
degree := make(map[string]int, len(nodes))
|
||||||
|
queue := []string{}
|
||||||
|
for id := range nodes {
|
||||||
|
degree[id] = incoming[id]
|
||||||
|
if degree[id] == 0 {
|
||||||
|
queue = append(queue, id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
visited := 0
|
||||||
|
for len(queue) > 0 {
|
||||||
|
id := queue[0]
|
||||||
|
queue = queue[1:]
|
||||||
|
visited++
|
||||||
|
for _, edge := range outgoing[id] {
|
||||||
|
degree[edge.Target]--
|
||||||
|
if degree[edge.Target] == 0 {
|
||||||
|
queue = append(queue, edge.Target)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return visited != len(nodes)
|
||||||
|
}
|
||||||
|
func validateTerminalPaths(nodes []RuleNode, outgoing map[string][]RuleEdge) error {
|
||||||
|
reverse := map[string][]string{}
|
||||||
|
stack := []string{}
|
||||||
|
for _, node := range nodes {
|
||||||
|
if node.Type == RuleNodeAllow || node.Type == RuleNodeBlock {
|
||||||
|
stack = append(stack, node.ID)
|
||||||
|
}
|
||||||
|
for _, edge := range outgoing[node.ID] {
|
||||||
|
reverse[edge.Target] = append(reverse[edge.Target], node.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
terminating := map[string]bool{}
|
||||||
|
for len(stack) > 0 {
|
||||||
|
id := stack[len(stack)-1]
|
||||||
|
stack = stack[:len(stack)-1]
|
||||||
|
if terminating[id] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
terminating[id] = true
|
||||||
|
stack = append(stack, reverse[id]...)
|
||||||
|
}
|
||||||
|
for _, node := range nodes {
|
||||||
|
if !terminating[node.ID] {
|
||||||
|
return fmt.Errorf("节点 %s 不存在通往终止节点的路径", node.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,117 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package waf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func rawConfig(value string) json.RawMessage { return json.RawMessage(value) }
|
||||||
|
|
||||||
|
func TestDefaultRuleGraph(t *testing.T) {
|
||||||
|
graph := DefaultRuleGraph()
|
||||||
|
if err := ValidateRuleGraph(context.Background(), graph, nil); err != nil {
|
||||||
|
t.Fatalf("default graph must be valid: %v", err)
|
||||||
|
}
|
||||||
|
if graph.SchemaVersion != 1 || len(graph.Nodes) != 2 || len(graph.Edges) != 1 {
|
||||||
|
t.Fatalf("unexpected default graph: %#v", graph)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateRuleGraph(t *testing.T) {
|
||||||
|
validBranch := func() RuleGraph {
|
||||||
|
return RuleGraph{SchemaVersion: 1, Nodes: []RuleNode{
|
||||||
|
{ID: "start", Type: RuleNodeStart, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "match-1", Type: RuleNodeIPMatch, Config: rawConfig(`{"ips":["192.0.2.1"],"cidrs":["2001:db8::/32"],"ip_group_ids":[7]}`)},
|
||||||
|
{ID: "allow", Type: RuleNodeAllow, Config: rawConfig(`{}`)},
|
||||||
|
{ID: "block", Type: RuleNodeBlock, Config: rawConfig(`{"status_code":403,"response_body":"denied"}`)},
|
||||||
|
}, Edges: []RuleEdge{
|
||||||
|
{ID: "e1", Source: "start", SourceHandle: "next", Target: "match-1"},
|
||||||
|
{ID: "e2", Source: "match-1", SourceHandle: "true", Target: "block"},
|
||||||
|
{ID: "e3", Source: "match-1", SourceHandle: "false", Target: "allow"},
|
||||||
|
}}
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
mutate func(*RuleGraph)
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"duplicate start", func(g *RuleGraph) {
|
||||||
|
g.Nodes = append(g.Nodes, RuleNode{ID: "start-2", Type: RuleNodeStart, Config: rawConfig(`{}`)})
|
||||||
|
}, "恰好包含一个开始节点"},
|
||||||
|
{"duplicate allow", func(g *RuleGraph) {
|
||||||
|
g.Nodes = append(g.Nodes, RuleNode{ID: "allow-2", Type: RuleNodeAllow, Config: rawConfig(`{}`)})
|
||||||
|
}, "恰好包含一个通过节点"},
|
||||||
|
{"cycle", func(g *RuleGraph) {
|
||||||
|
g.Edges[2].Target = "start"
|
||||||
|
}, "规则图不能包含循环"},
|
||||||
|
{"unreachable", func(g *RuleGraph) {
|
||||||
|
g.Nodes = append(g.Nodes, RuleNode{ID: "orphan", Type: RuleNodeBlock, Config: rawConfig(`{"status_code":403,"response_body":""}`)})
|
||||||
|
}, "节点 orphan 无法从开始节点到达"},
|
||||||
|
{"missing false edge", func(g *RuleGraph) { g.Edges = g.Edges[:2] }, "节点 match-1 的 false 出口未连接"},
|
||||||
|
{"wrong handle", func(g *RuleGraph) { g.Edges[0].SourceHandle = "true" }, "边 e1 的源端口 true 不适用于节点 start"},
|
||||||
|
{"same handle multiple targets", func(g *RuleGraph) {
|
||||||
|
g.Edges = append(g.Edges, RuleEdge{ID: "e4", Source: "match-1", SourceHandle: "true", Target: "allow"})
|
||||||
|
}, "节点 match-1 的 true 出口连接了多个目标"},
|
||||||
|
{"unknown type", func(g *RuleGraph) { g.Nodes[1].Type = RuleNodeType("script") }, "节点 match-1 的类型 script 未知"},
|
||||||
|
{"dangling target", func(g *RuleGraph) { g.Edges[1].Target = "missing" }, "边 e2 的目标节点 missing 不存在"},
|
||||||
|
{"invalid ip", func(g *RuleGraph) { g.Nodes[1].Config = rawConfig(`{"ips":["bad"]}`) }, "节点 match-1 的 IP bad 无效"},
|
||||||
|
{"invalid cidr", func(g *RuleGraph) { g.Nodes[1].Config = rawConfig(`{"cidrs":["bad"]}`) }, "节点 match-1 的 CIDR bad 无效"},
|
||||||
|
{"missing ip group", func(g *RuleGraph) { g.Nodes[1].Config = rawConfig(`{"ip_group_ids":[99]}`) }, "节点 match-1 引用的 IP 组 99 不存在"},
|
||||||
|
{"invalid country", func(g *RuleGraph) {
|
||||||
|
g.Nodes[1].Type = RuleNodeGeoMatch
|
||||||
|
g.Nodes[1].Config = rawConfig(`{"countries":["USA"]}`)
|
||||||
|
}, "节点 match-1 的国家代码 USA 无效"},
|
||||||
|
{"invalid region", func(g *RuleGraph) {
|
||||||
|
g.Nodes[1].Type = RuleNodeGeoMatch
|
||||||
|
g.Nodes[1].Config = rawConfig(`{"regions":["US-"]}`)
|
||||||
|
}, "节点 match-1 的地区代码 US- 无效"},
|
||||||
|
{"pow range", func(g *RuleGraph) {
|
||||||
|
g.Nodes[1].Type = RuleNodePoW
|
||||||
|
g.Nodes[1].Config = rawConfig(`{"algorithm":"fast","difficulty":17,"session_ttl":60,"challenge_ttl":30}`)
|
||||||
|
g.Edges[1].SourceHandle = "next"
|
||||||
|
g.Edges = g.Edges[:2]
|
||||||
|
}, "节点 match-1 的 PoW 难度必须在 1-16 之间"},
|
||||||
|
{"unknown config field", func(g *RuleGraph) { g.Nodes[1].Config = rawConfig(`{"ips":[],"surprise":true}`) }, "节点 match-1 的配置无效"},
|
||||||
|
{"null config", func(g *RuleGraph) { g.Nodes[1].Config = rawConfig(`null`) }, "节点 match-1 的配置无效"},
|
||||||
|
{"too many nodes", func(g *RuleGraph) {
|
||||||
|
for i := 0; i < 125; i++ {
|
||||||
|
g.Nodes = append(g.Nodes, RuleNode{ID: strings.Repeat("x", i+1), Type: RuleNodeBlock, Config: rawConfig(`{"status_code":403}`)})
|
||||||
|
}
|
||||||
|
}, "规则图节点数不能超过 128"},
|
||||||
|
{"too many edges", func(g *RuleGraph) {
|
||||||
|
for i := 0; i < 254; i++ {
|
||||||
|
g.Edges = append(g.Edges, RuleEdge{ID: strings.Repeat("e", i+5), Source: "start", SourceHandle: "next", Target: "allow"})
|
||||||
|
}
|
||||||
|
}, "规则图边数不能超过 256"},
|
||||||
|
{"graph too large", func(g *RuleGraph) {
|
||||||
|
g.Nodes[0].Label = strings.Repeat("x", 256*1024)
|
||||||
|
}, "规则图大小不能超过 256 KiB"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
graph := validBranch()
|
||||||
|
tt.mutate(&graph)
|
||||||
|
exists := func(_ context.Context, id uint) (bool, error) { return id == 7, nil }
|
||||||
|
err := ValidateRuleGraph(context.Background(), graph, exists)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tt.want) {
|
||||||
|
t.Fatalf("ValidateRuleGraph() error = %v, want substring %q", err, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateRuleGraphDetectsNodeWithoutTerminalPath(t *testing.T) {
|
||||||
|
nodes := []RuleNode{{ID: "start", Type: RuleNodeStart}, {ID: "match", Type: RuleNodeIPMatch}, {ID: "allow", Type: RuleNodeAllow}}
|
||||||
|
outgoing := map[string][]RuleEdge{"start": {{Source: "start", Target: "match"}}}
|
||||||
|
err := validateTerminalPaths(nodes, outgoing)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "节点 start 不存在通往终止节点的路径") {
|
||||||
|
t.Fatalf("validateTerminalPaths() error = %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
-- +goose Up
|
||||||
|
ALTER TABLE of_waf_rule_groups
|
||||||
|
ADD COLUMN graph TEXT NOT NULL DEFAULT '',
|
||||||
|
ADD COLUMN revision BIGINT NOT NULL DEFAULT 1;
|
||||||
|
|
||||||
|
ALTER TABLE of_waf_rule_group_bindings
|
||||||
|
ADD COLUMN sequence INTEGER NOT NULL DEFAULT 0;
|
||||||
|
|
||||||
|
-- +goose Down
|
||||||
|
ALTER TABLE of_waf_rule_group_bindings DROP COLUMN sequence;
|
||||||
|
ALTER TABLE of_waf_rule_groups DROP COLUMN revision, DROP COLUMN graph;
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
-- +goose Up
|
||||||
|
UPDATE of_waf_rule_groups
|
||||||
|
SET graph = '{"schema_version":1,"nodes":[{"id":"start","type":"start","position":{"x":0,"y":0},"config":{}},{"id":"allow","type":"allow","position":{"x":320,"y":0},"config":{}}],"edges":[{"id":"start-allow","source":"start","source_handle":"next","target":"allow"}]}',
|
||||||
|
revision = 1;
|
||||||
|
|
||||||
|
WITH ordered AS (
|
||||||
|
SELECT id, ROW_NUMBER() OVER (PARTITION BY proxy_route_id ORDER BY id) - 1 AS new_sequence
|
||||||
|
FROM of_waf_rule_group_bindings
|
||||||
|
)
|
||||||
|
UPDATE of_waf_rule_group_bindings AS binding
|
||||||
|
SET sequence = ordered.new_sequence
|
||||||
|
FROM ordered
|
||||||
|
WHERE binding.id = ordered.id;
|
||||||
|
|
||||||
|
-- +goose Down
|
||||||
|
UPDATE of_waf_rule_group_bindings SET sequence = 0;
|
||||||
|
UPDATE of_waf_rule_groups SET revision = 1;
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
-- +goose Up
|
||||||
|
ALTER TABLE of_waf_rule_groups
|
||||||
|
ADD COLUMN graph TEXT NOT NULL DEFAULT '';
|
||||||
|
ALTER TABLE of_waf_rule_groups ADD COLUMN revision INTEGER NOT NULL DEFAULT 1;
|
||||||
|
ALTER TABLE of_waf_rule_group_bindings ADD COLUMN sequence INTEGER NOT NULL DEFAULT 0;
|
||||||
|
|
||||||
|
-- +goose Down
|
||||||
|
ALTER TABLE of_waf_rule_group_bindings DROP COLUMN sequence;
|
||||||
|
ALTER TABLE of_waf_rule_groups DROP COLUMN revision;
|
||||||
|
ALTER TABLE of_waf_rule_groups DROP COLUMN graph;
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
-- +goose Up
|
||||||
|
UPDATE of_waf_rule_groups
|
||||||
|
SET graph = '{"schema_version":1,"nodes":[{"id":"start","type":"start","position":{"x":0,"y":0},"config":{}},{"id":"allow","type":"allow","position":{"x":320,"y":0},"config":{}}],"edges":[{"id":"start-allow","source":"start","source_handle":"next","target":"allow"}]}',
|
||||||
|
revision = 1;
|
||||||
|
|
||||||
|
UPDATE of_waf_rule_group_bindings
|
||||||
|
SET sequence = (
|
||||||
|
SELECT COUNT(*)
|
||||||
|
FROM of_waf_rule_group_bindings AS preceding
|
||||||
|
WHERE preceding.proxy_route_id = of_waf_rule_group_bindings.proxy_route_id
|
||||||
|
AND preceding.id < of_waf_rule_group_bindings.id
|
||||||
|
);
|
||||||
|
|
||||||
|
-- +goose Down
|
||||||
|
UPDATE of_waf_rule_group_bindings SET sequence = 0;
|
||||||
|
UPDATE of_waf_rule_groups SET revision = 1;
|
||||||
@@ -30,6 +30,8 @@ type OpenFlareWAFRuleGroup struct {
|
|||||||
RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
|
RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
|
||||||
PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"`
|
PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"`
|
||||||
PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"`
|
PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"`
|
||||||
|
Graph string `json:"graph" gorm:"type:text;not null;default:''"`
|
||||||
|
Revision uint64 `json:"revision" gorm:"not null;default:1"`
|
||||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||||
}
|
}
|
||||||
@@ -70,9 +72,13 @@ type OpenFlareWAFRuleGroupBinding struct {
|
|||||||
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
|
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||||
RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route"`
|
RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route"`
|
||||||
ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route;index"`
|
ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route;index"`
|
||||||
|
Sequence int `json:"sequence" gorm:"not null;default:0"`
|
||||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ErrWAFRuleRevisionConflict indicates that a rule graph was updated from a stale revision.
|
||||||
|
var ErrWAFRuleRevisionConflict = errors.New("waf rule revision conflict")
|
||||||
|
|
||||||
// TableName returns the GORM table name.
|
// TableName returns the GORM table name.
|
||||||
func (OpenFlareWAFRuleGroupBinding) TableName() string {
|
func (OpenFlareWAFRuleGroupBinding) TableName() string {
|
||||||
return "of_waf_rule_group_bindings"
|
return "of_waf_rule_group_bindings"
|
||||||
@@ -159,6 +165,24 @@ func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGro
|
|||||||
}).Error
|
}).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// UpdateOpenFlareWAFRuleGraph atomically replaces a graph when revision is current.
|
||||||
|
func UpdateOpenFlareWAFRuleGraph(ctx context.Context, id uint, revision uint64, graph string) (uint64, error) {
|
||||||
|
conn, err := wafDB(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
result := conn.Model(&OpenFlareWAFRuleGroup{}).
|
||||||
|
Where("id = ? AND revision = ?", id, revision).
|
||||||
|
Updates(map[string]any{"graph": graph, "revision": gorm.Expr("revision + 1")})
|
||||||
|
if result.Error != nil {
|
||||||
|
return 0, result.Error
|
||||||
|
}
|
||||||
|
if result.RowsAffected != 1 {
|
||||||
|
return 0, ErrWAFRuleRevisionConflict
|
||||||
|
}
|
||||||
|
return revision + 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
// DeleteOpenFlareWAFRuleGroup removes a rule group.
|
// DeleteOpenFlareWAFRuleGroup removes a rule group.
|
||||||
func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error {
|
func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error {
|
||||||
conn, err := wafDB(ctx)
|
conn, err := wafDB(ctx)
|
||||||
@@ -289,7 +313,7 @@ func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]OpenFlareWAFRuleG
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
var bindings []OpenFlareWAFRuleGroupBinding
|
var bindings []OpenFlareWAFRuleGroupBinding
|
||||||
if err = conn.Order("rule_group_id asc").Order("proxy_route_id asc").Find(&bindings).Error; err != nil {
|
if err = conn.Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return bindings, nil
|
return bindings, nil
|
||||||
@@ -302,7 +326,7 @@ func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uin
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
var bindings []OpenFlareWAFRuleGroupBinding
|
var bindings []OpenFlareWAFRuleGroupBinding
|
||||||
if err = conn.Where("proxy_route_id = ?", routeID).Order("rule_group_id asc").Find(&bindings).Error; err != nil {
|
if err = conn.Where("proxy_route_id = ?", routeID).Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return bindings, nil
|
return bindings, nil
|
||||||
@@ -342,10 +366,11 @@ func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, rou
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(routeIDs))
|
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(routeIDs))
|
||||||
for _, routeID := range routeIDs {
|
for index, routeID := range routeIDs {
|
||||||
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
|
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
|
||||||
RuleGroupID: groupID,
|
RuleGroupID: groupID,
|
||||||
ProxyRouteID: routeID,
|
ProxyRouteID: routeID,
|
||||||
|
Sequence: index,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
|
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
|
||||||
@@ -363,10 +388,11 @@ func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint,
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(groupIDs))
|
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(groupIDs))
|
||||||
for _, groupID := range groupIDs {
|
for index, groupID := range groupIDs {
|
||||||
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
|
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
|
||||||
RuleGroupID: groupID,
|
RuleGroupID: groupID,
|
||||||
ProxyRouteID: routeID,
|
ProxyRouteID: routeID,
|
||||||
|
Sequence: index,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
|
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
|
||||||
|
|||||||
@@ -0,0 +1,107 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package model
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
"testing/fstest"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/glebarez/sqlite"
|
||||||
|
"github.com/pressly/goose/v3"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
const defaultWAFRuleGraph = `{"schema_version":1,"nodes":[{"id":"start","type":"start","position":{"x":0,"y":0},"config":{}},{"id":"allow","type":"allow","position":{"x":320,"y":0},"config":{}}],"edges":[{"id":"start-allow","source":"start","source_handle":"next","target":"allow"}]}`
|
||||||
|
|
||||||
|
func wafMigrationFS(t *testing.T) fs.FS {
|
||||||
|
t.Helper()
|
||||||
|
_, filename, _, ok := runtime.Caller(0)
|
||||||
|
require.True(t, ok)
|
||||||
|
dir := filepath.Join(filepath.Dir(filename), "..", "db", "migrator", "goose", "sqlite")
|
||||||
|
migrations := fstest.MapFS{}
|
||||||
|
for _, name := range []string{
|
||||||
|
"202607150001_orchestrate_waf_rules.sql",
|
||||||
|
"202607150002_reset_waf_rule_graphs.sql",
|
||||||
|
} {
|
||||||
|
contents, err := os.ReadFile(filepath.Join(dir, name))
|
||||||
|
require.NoError(t, err)
|
||||||
|
migrations[name] = &fstest.MapFile{Data: contents}
|
||||||
|
}
|
||||||
|
return migrations
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) {
|
||||||
|
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := conn.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_groups (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL)`).Error)
|
||||||
|
require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_group_bindings (id INTEGER PRIMARY KEY AUTOINCREMENT, rule_group_id INTEGER NOT NULL, proxy_route_id INTEGER NOT NULL)`).Error)
|
||||||
|
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (id, name) VALUES (1, 'one'), (2, 'two')`).Error)
|
||||||
|
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_group_bindings (id, rule_group_id, proxy_route_id) VALUES (20, 2, 7), (10, 1, 7)`).Error)
|
||||||
|
|
||||||
|
goose.SetBaseFS(wafMigrationFS(t))
|
||||||
|
require.NoError(t, goose.SetDialect("sqlite3"))
|
||||||
|
require.NoError(t, goose.Up(sqlDB, "."))
|
||||||
|
|
||||||
|
var groups []OpenFlareWAFRuleGroup
|
||||||
|
require.NoError(t, conn.Order("id asc").Find(&groups).Error)
|
||||||
|
require.Len(t, groups, 2)
|
||||||
|
for _, group := range groups {
|
||||||
|
require.JSONEq(t, defaultWAFRuleGraph, group.Graph)
|
||||||
|
assert.Equal(t, uint64(1), group.Revision)
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (name) VALUES ('new')`).Error)
|
||||||
|
var newGroup OpenFlareWAFRuleGroup
|
||||||
|
require.NoError(t, conn.First(&newGroup, 3).Error)
|
||||||
|
assert.Empty(t, newGroup.Graph)
|
||||||
|
assert.Equal(t, uint64(1), newGroup.Revision)
|
||||||
|
|
||||||
|
var bindings []OpenFlareWAFRuleGroupBinding
|
||||||
|
require.NoError(t, conn.Where("proxy_route_id = ?", 7).Order("sequence asc").Order("id asc").Find(&bindings).Error)
|
||||||
|
require.Len(t, bindings, 2)
|
||||||
|
assert.Equal(t, []int{0, 1}, []int{bindings[0].Sequence, bindings[1].Sequence})
|
||||||
|
assert.Equal(t, []uint{10, 20}, []uint{bindings[0].ID, bindings[1].ID})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenFlareWAFGraphOptimisticUpdate(t *testing.T) {
|
||||||
|
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, conn.AutoMigrate(&OpenFlareWAFRuleGroup{}))
|
||||||
|
db.SetDB(conn)
|
||||||
|
t.Cleanup(func() { db.SetDB(nil) })
|
||||||
|
|
||||||
|
group := OpenFlareWAFRuleGroup{Name: "rule", Graph: defaultWAFRuleGraph, Revision: 1}
|
||||||
|
require.NoError(t, conn.Create(&group).Error)
|
||||||
|
|
||||||
|
nextRevision, err := UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, `{"schema_version":1}`)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint64(2), nextRevision)
|
||||||
|
|
||||||
|
_, err = UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, defaultWAFRuleGraph)
|
||||||
|
assert.ErrorIs(t, err, ErrWAFRuleRevisionConflict)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReplaceOpenFlareWAFRuleGroupBindingsPreservesInputOrder(t *testing.T) {
|
||||||
|
cleanup := setupWAFBindingsTestDB(t)
|
||||||
|
defer cleanup()
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
require.NoError(t, ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, 7, []uint{30, 10, 20}))
|
||||||
|
bindings, err := ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx, 7)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, bindings, 3)
|
||||||
|
assert.Equal(t, []uint{30, 10, 20}, []uint{bindings[0].RuleGroupID, bindings[1].RuleGroupID, bindings[2].RuleGroupID})
|
||||||
|
assert.Equal(t, []int{0, 1, 2}, []int{bindings[0].Sequence, bindings[1].Sequence, bindings[2].Sequence})
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user