mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +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:'[]'"`
|
||||
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:'{}'"`
|
||||
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"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
@@ -70,9 +72,13 @@ type OpenFlareWAFRuleGroupBinding struct {
|
||||
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
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"`
|
||||
Sequence int `json:"sequence" gorm:"not null;default:0"`
|
||||
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.
|
||||
func (OpenFlareWAFRuleGroupBinding) TableName() string {
|
||||
return "of_waf_rule_group_bindings"
|
||||
@@ -159,6 +165,24 @@ func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGro
|
||||
}).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.
|
||||
func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error {
|
||||
conn, err := wafDB(ctx)
|
||||
@@ -289,7 +313,7 @@ func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]OpenFlareWAFRuleG
|
||||
return nil, err
|
||||
}
|
||||
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 bindings, nil
|
||||
@@ -302,7 +326,7 @@ func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uin
|
||||
return nil, err
|
||||
}
|
||||
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 bindings, nil
|
||||
@@ -342,10 +366,11 @@ func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, rou
|
||||
return err
|
||||
}
|
||||
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(routeIDs))
|
||||
for _, routeID := range routeIDs {
|
||||
for index, routeID := range routeIDs {
|
||||
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
|
||||
RuleGroupID: groupID,
|
||||
ProxyRouteID: routeID,
|
||||
Sequence: index,
|
||||
})
|
||||
}
|
||||
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
|
||||
@@ -363,10 +388,11 @@ func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint,
|
||||
return err
|
||||
}
|
||||
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(groupIDs))
|
||||
for _, groupID := range groupIDs {
|
||||
for index, groupID := range groupIDs {
|
||||
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
|
||||
RuleGroupID: groupID,
|
||||
ProxyRouteID: routeID,
|
||||
Sequence: index,
|
||||
})
|
||||
}
|
||||
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