refactor(backend): rename OpenFlare directory to lowercase openflare

This commit is contained in:
ryan
2026-08-30 17:43:23 +08:00
parent 06d5fedbfc
commit c93ff6674f
543 changed files with 819 additions and 819 deletions
@@ -0,0 +1,5 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package waf defines shared error messages for WAF management.
package waf
@@ -0,0 +1,156 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"errors"
"fmt"
"slices"
"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{}, errors.New("规则图包含多个开始节点")
}
runtime.Entry = node.ID
}
runtime.Nodes[node.ID] = RuntimeRuleNode{Type: node.Type, Config: config}
}
if runtime.Entry == "" {
return RuntimeRuleGraph{}, errors.New("规则图缺少开始节点")
}
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 RuleNodeUACheck:
var config UACheckConfig
if err := decodeStrictConfig(node.Config, &config); err != nil {
return nil, err
}
config.Browsers = sortedUniqueStrings(config.Browsers)
config.OperatingSystems = sortedUniqueStrings(config.OperatingSystems)
config.CustomUAPatterns = sortedUniqueStrings(config.CustomUAPatterns)
if config.MatchMode == "" {
config.MatchMode = UACheckMatchModeOr
}
return config, nil
case RuleNodeSecurityCheck:
var config SecurityCheckConfig
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...)
slices.Sort(result)
write := 0
for _, value := range result {
if write == 0 || result[write-1] != value {
result[write] = value
write++
}
}
return result[:write]
}
@@ -0,0 +1,182 @@
// 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 TestCompileSecurityCheckConfig(t *testing.T) {
graph := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
{ID: "start", Type: RuleNodeStart, Config: rawConfig(`{}`)},
{ID: "sec", Type: RuleNodeSecurityCheck, Config: rawConfig(`{"path_traversal":true,"file_inclusion":true,"sql_injection":false}`)},
{ID: "allow", Type: RuleNodeAllow, Config: rawConfig(`{}`)},
{ID: "block", Type: RuleNodeBlock, Config: rawConfig(`{"status_code":403}`)},
}, Edges: []RuleEdge{
{ID: "e1", Source: "start", SourceHandle: "next", Target: "sec"},
{ID: "e2", Source: "sec", SourceHandle: "true", Target: "allow"},
{ID: "e3", Source: "sec", SourceHandle: "false", Target: "block"},
}}
compiled, err := CompileRuleGraph(graph)
if err != nil {
t.Fatalf("CompileRuleGraph() error = %v", err)
}
cfg, ok := compiled.Nodes["sec"].Config.(SecurityCheckConfig)
if !ok {
t.Fatalf("config type = %T", compiled.Nodes["sec"].Config)
}
if !cfg.PathTraversal || !cfg.FileInclusion || cfg.SQLInjection {
t.Fatalf("unexpected config %#v", cfg)
}
}
func TestCompileUACheckConfigNormalizesListsAndMatchMode(t *testing.T) {
graph := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
{ID: "start", Type: RuleNodeStart, Config: rawConfig(`{}`)},
{ID: "ua", Type: RuleNodeUACheck, Config: rawConfig(`{"browsers":["Safari","Chrome","Chrome"],"operating_systems":["iOS","Android"],"require_ua":true,"block_common_bots":true}`)},
{ID: "allow", Type: RuleNodeAllow, Config: rawConfig(`{}`)},
{ID: "block", Type: RuleNodeBlock, Config: rawConfig(`{"status_code":403}`)},
}, Edges: []RuleEdge{
{ID: "e1", Source: "start", SourceHandle: "next", Target: "ua"},
{ID: "e2", Source: "ua", SourceHandle: "true", Target: "allow"},
{ID: "e3", Source: "ua", SourceHandle: "false", Target: "block"},
}}
compiled, err := CompileRuleGraph(graph)
if err != nil {
t.Fatalf("CompileRuleGraph() error = %v", err)
}
cfg, ok := compiled.Nodes["ua"].Config.(UACheckConfig)
if !ok {
t.Fatalf("config type = %T", compiled.Nodes["ua"].Config)
}
if !reflect.DeepEqual(cfg.Browsers, []string{"Chrome", "Safari"}) {
t.Fatalf("browsers = %#v", cfg.Browsers)
}
if !reflect.DeepEqual(cfg.OperatingSystems, []string{"Android", "iOS"}) {
t.Fatalf("os = %#v", cfg.OperatingSystems)
}
if cfg.MatchMode != UACheckMatchModeOr {
t.Fatalf("match_mode = %q, want or", cfg.MatchMode)
}
if !cfg.RequireUA || !cfg.BlockCommonBots || cfg.BlockAbnormalUA {
t.Fatalf("flags = %#v", cfg)
}
}
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,136 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import "encoding/json"
// RuleGraphSchemaVersion is the current persisted rule graph schema version.
const RuleGraphSchemaVersion = 1
// RuleNodeType identifies the behavior of a rule graph node.
type RuleNodeType string
const (
// RuleNodeStart begins graph execution.
RuleNodeStart RuleNodeType = "start"
// RuleNodeAllow terminates execution with an allow decision.
RuleNodeAllow RuleNodeType = "allow"
// RuleNodeBlock terminates execution with a blocking response.
RuleNodeBlock RuleNodeType = "block"
// RuleNodeIPMatch branches on an IP match.
RuleNodeIPMatch RuleNodeType = "ip_match"
// RuleNodeGeoMatch branches on a geographic match.
RuleNodeGeoMatch RuleNodeType = "geo_match"
// RuleNodePoW runs a proof-of-work challenge before continuing.
RuleNodePoW RuleNodeType = "pow"
// RuleNodeUACheck branches on User-Agent presence, classification, and lists.
RuleNodeUACheck RuleNodeType = "ua_check"
// RuleNodeSecurityCheck branches on basic request payload attack signatures.
RuleNodeSecurityCheck RuleNodeType = "security_check"
)
// RuleGraph is the editor-facing representation of an executable WAF graph.
type RuleGraph struct {
SchemaVersion int `json:"schema_version"`
Nodes []RuleNode `json:"nodes"`
Edges []RuleEdge `json:"edges"`
}
// RuleNode stores one editor node and its type-specific configuration.
type RuleNode struct {
ID string `json:"id"`
Type RuleNodeType `json:"type"`
Label string `json:"label,omitempty"`
Position RulePosition `json:"position"`
Config json.RawMessage `json:"config"`
}
// RulePosition stores a node's editor canvas coordinates.
type RulePosition struct {
X float64 `json:"x"`
Y float64 `json:"y"`
}
// RuleEdge connects one source handle to a target node.
type RuleEdge struct {
ID string `json:"id"`
Source string `json:"source"`
SourceHandle string `json:"source_handle"`
Target string `json:"target"`
}
// IPMatchConfig configures literal, CIDR, and managed-group IP matching.
type IPMatchConfig struct {
IPs []string `json:"ips,omitempty"`
CIDRs []string `json:"cidrs,omitempty"`
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
}
// GeoMatchConfig configures country and region matching.
type GeoMatchConfig struct {
Countries []string `json:"countries,omitempty"`
Regions []string `json:"regions,omitempty"`
}
// PoWNodeConfig configures a proof-of-work challenge node.
type PoWNodeConfig struct {
Algorithm string `json:"algorithm"`
Difficulty int `json:"difficulty"`
SessionTTL int `json:"session_ttl"`
ChallengeTTL int `json:"challenge_ttl"`
}
// BlockNodeConfig configures a terminal blocking response.
type BlockNodeConfig struct {
StatusCode int `json:"status_code"`
ResponseBody string `json:"response_body,omitempty"`
}
// UACheckConfig configures User-Agent presence, whitelist, and block switches.
type UACheckConfig struct {
RequireUA bool `json:"require_ua"`
Browsers []string `json:"browsers,omitempty"`
OperatingSystems []string `json:"operating_systems,omitempty"`
MatchMode string `json:"match_mode,omitempty"`
BlockCommonBots bool `json:"block_common_bots"`
BlockAbnormalUA bool `json:"block_abnormal_ua"`
BlockCustomUA bool `json:"block_custom_ua"`
CustomUAPatterns []string `json:"custom_ua_patterns,omitempty"`
}
// UA check match modes.
const (
UACheckMatchModeAnd = "and"
UACheckMatchModeOr = "or"
)
// SecurityCheckConfig toggles basic payload signature protections.
// Default graph nodes enable path_traversal and file_inclusion only.
type SecurityCheckConfig struct {
SQLInjection bool `json:"sql_injection"`
PathTraversal bool `json:"path_traversal"`
CommandInjection bool `json:"command_injection"`
XSS bool `json:"xss"`
SSRF bool `json:"ssrf"`
FileInclusion bool `json:"file_inclusion"`
MaliciousUpload bool `json:"malicious_upload"`
XXE bool `json:"xxe"`
CRLFInjection bool `json:"crlf_injection"`
}
// DefaultSecurityCheckConfig returns low false-positive defaults.
func DefaultSecurityCheckConfig() SecurityCheckConfig {
return SecurityCheckConfig{
PathTraversal: true,
FileInclusion: true,
}
}
// DefaultRuleGraph returns the minimal start-to-allow graph.
func DefaultRuleGraph() RuleGraph {
return RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
{ID: "start", Type: RuleNodeStart, Position: RulePosition{X: 0, Y: 0}, Config: json.RawMessage(`{}`)},
{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,458 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/netip"
"regexp"
"slices"
"strings"
)
const (
maxRuleGraphNodes = 128
maxRuleGraphEdges = 256
maxRuleGraphBytes = 256 * 1024
maxUACustomPatterns = 32
maxUACustomPatternBytes = 256
)
var (
countryCodePattern = regexp.MustCompile(`^[A-Z]{2}$`)
regionCodePattern = regexp.MustCompile(`^[A-Z]{2}-[A-Z0-9]{1,3}$`)
)
// ValidateRuleGraph validates graph structure, node configuration, references,
// reachability, and termination before compilation.
func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(context.Context, uint) (bool, error)) error {
if err := validateRuleGraphLimits(graph); err != nil {
return err
}
nodes, startID, err := validateRuleGraphNodes(ctx, graph.Nodes, ipGroupExists)
if err != nil {
return err
}
outgoing, incoming, handleTargets, err := validateRuleGraphEdges(nodes, graph.Edges)
if err != nil {
return err
}
if hasRuleGraphCycle(nodes, outgoing, incoming) {
return errors.New("规则图不能包含循环")
}
if err := validateRequiredHandles(graph.Nodes, handleTargets); err != nil {
return err
}
if err := validateRuleGraphConnectivity(graph.Nodes, startID, outgoing, incoming); err != nil {
return err
}
return validateTerminalPaths(graph.Nodes, outgoing)
}
func validateRuleGraphLimits(graph RuleGraph) error {
if graph.SchemaVersion != RuleGraphSchemaVersion {
return fmt.Errorf("规则图 schema_version 必须为 %d", RuleGraphSchemaVersion)
}
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 errors.New("规则图大小不能超过 256 KiB")
}
return nil
}
func validateRuleGraphNodes(ctx context.Context, graphNodes []RuleNode, ipGroupExists func(context.Context, uint) (bool, error)) (map[string]RuleNode, string, error) {
nodes := make(map[string]RuleNode, len(graphNodes))
startCount, allowCount, startID := 0, 0, ""
for _, node := range graphNodes {
if strings.TrimSpace(node.ID) == "" {
return nil, "", errors.New("节点 ID 不能为空")
}
if _, exists := nodes[node.ID]; exists {
return nil, "", 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, RuleNodeUACheck, RuleNodeSecurityCheck:
default:
return nil, "", fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
}
if err := validateRuleNodeConfig(ctx, node, ipGroupExists); err != nil {
return nil, "", err
}
}
if startCount != 1 {
return nil, "", errors.New("规则图必须恰好包含一个开始节点")
}
if allowCount != 1 {
return nil, "", errors.New("规则图必须恰好包含一个通过节点")
}
return nodes, startID, nil
}
func validateRuleGraphEdges(nodes map[string]RuleNode, graphEdges []RuleEdge) (map[string][]RuleEdge, map[string]int, map[string]int, error) {
edgeIDs := make(map[string]struct{}, len(graphEdges))
outgoing := make(map[string][]RuleEdge)
incoming := make(map[string]int)
handleTargets := make(map[string]int)
for _, edge := range graphEdges {
if strings.TrimSpace(edge.ID) == "" {
return nil, nil, nil, errors.New("边 ID 不能为空")
}
if _, exists := edgeIDs[edge.ID]; exists {
return nil, nil, nil, fmt.Errorf("边 ID %s 重复", edge.ID)
}
edgeIDs[edge.ID] = struct{}{}
source, ok := nodes[edge.Source]
if !ok {
return nil, nil, nil, fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
}
if _, ok := nodes[edge.Target]; !ok {
return nil, nil, nil, fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
}
if !validSourceHandle(source.Type, edge.SourceHandle) {
return nil, nil, nil, fmt.Errorf("边 %s 的源端口 %s 不适用于节点 %s", edge.ID, edge.SourceHandle, edge.Source)
}
key := edge.Source + "\x00" + edge.SourceHandle
handleTargets[key]++
if handleTargets[key] > 1 {
return nil, nil, nil, fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
}
outgoing[edge.Source] = append(outgoing[edge.Source], edge)
incoming[edge.Target]++
}
return outgoing, incoming, handleTargets, nil
}
func validateRequiredHandles(nodes []RuleNode, handleTargets map[string]int) error {
for _, node := range nodes {
for _, handle := range requiredHandles(node.Type) {
if handleTargets[node.ID+"\x00"+handle] == 0 {
return fmt.Errorf("节点 %s 的 %s 出口未连接", node.ID, handle)
}
}
}
return nil
}
func validateRuleGraphConnectivity(nodes []RuleNode, startID string, outgoing map[string][]RuleEdge, incoming map[string]int) error {
reachable := walkRuleGraph(startID, outgoing)
for _, node := range nodes {
if !reachable[node.ID] {
return fmt.Errorf("节点 %s 无法从开始节点到达", node.ID)
}
}
for _, node := range 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)
}
}
return nil
}
func validateRuleNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
switch node.Type {
case RuleNodeStart, RuleNodeAllow:
return validateEmptyNodeConfig(node)
case RuleNodeIPMatch:
return validateIPMatchNodeConfig(ctx, node, exists)
case RuleNodeGeoMatch:
return validateGeoMatchNodeConfig(node)
case RuleNodePoW:
return validatePoWNodeConfig(node)
case RuleNodeUACheck:
return validateUACheckNodeConfig(node)
case RuleNodeSecurityCheck:
return validateSecurityCheckNodeConfig(node)
case RuleNodeBlock:
return validateBlockNodeConfig(node)
}
return nil
}
func validateEmptyNodeConfig(node RuleNode) error {
var cfg struct{}
return decodeNodeConfig(node, &cfg)
}
func validateIPMatchNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
var cfg IPMatchConfig
if err := decodeNodeConfig(node, &cfg); err != nil {
return err
}
for _, raw := range cfg.IPs {
if _, err := netip.ParseAddr(raw); err != nil {
return fmt.Errorf("节点 %s 的 IP %s 无效", node.ID, raw)
}
}
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 err := validateIPGroupReference(ctx, node.ID, id, exists); err != nil {
return err
}
}
return nil
}
func validateIPGroupReference(ctx context.Context, nodeID string, id uint, exists func(context.Context, uint) (bool, error)) error {
if id == 0 {
return fmt.Errorf("节点 %s 引用的 IP 组 ID 无效", nodeID)
}
if exists == nil {
return fmt.Errorf("节点 %s 无法校验 IP 组 %d", nodeID, id)
}
ok, err := exists(ctx, id)
if err != nil {
return fmt.Errorf("节点 %s 校验 IP 组 %d 失败: %w", nodeID, id, err)
}
if !ok {
return fmt.Errorf("节点 %s 引用的 IP 组 %d 不存在", nodeID, id)
}
return nil
}
func validateGeoMatchNodeConfig(node RuleNode) error {
var cfg GeoMatchConfig
if err := decodeNodeConfig(node, &cfg); err != nil {
return err
}
for _, code := range cfg.Countries {
if !countryCodePattern.MatchString(code) {
return fmt.Errorf("节点 %s 的国家代码 %s 无效", node.ID, code)
}
}
for _, code := range cfg.Regions {
if !regionCodePattern.MatchString(code) {
return fmt.Errorf("节点 %s 的地区代码 %s 无效", node.ID, code)
}
}
return nil
}
func validatePoWNodeConfig(node RuleNode) error {
var cfg PoWNodeConfig
if err := decodeNodeConfig(node, &cfg); err != nil {
return err
}
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
return fmt.Errorf("节点 %s 的 PoW 难度必须在 1-16 之间", node.ID)
}
if cfg.Algorithm != powAlgorithmFast && cfg.Algorithm != powAlgorithmSlow {
return fmt.Errorf("节点 %s 的 PoW 算法必须为 fast 或 slow", node.ID)
}
if cfg.SessionTTL < minPoWSessionTTLSeconds {
return fmt.Errorf("节点 %s 的 PoW 会话 TTL 不能小于 60 秒", node.ID)
}
if cfg.ChallengeTTL < minPoWChallengeTTLSeconds {
return fmt.Errorf("节点 %s 的 PoW 挑战 TTL 不能小于 30 秒", node.ID)
}
return nil
}
func validateBlockNodeConfig(node RuleNode) error {
var cfg BlockNodeConfig
if err := decodeNodeConfig(node, &cfg); err != nil {
return err
}
if cfg.StatusCode < 400 || cfg.StatusCode > 599 {
return fmt.Errorf("节点 %s 的阻止状态码必须在 400-599 之间", node.ID)
}
if len([]byte(cfg.ResponseBody)) > maxWAFBlockBodyBytes {
return fmt.Errorf("节点 %s 的阻止响应体不能超过 %d 字节", node.ID, maxWAFBlockBodyBytes)
}
return nil
}
func validateUACheckNodeConfig(node RuleNode) error {
var cfg UACheckConfig
if err := decodeNodeConfig(node, &cfg); err != nil {
return err
}
mode := cfg.MatchMode
if mode == "" {
mode = UACheckMatchModeOr
}
if mode != UACheckMatchModeAnd && mode != UACheckMatchModeOr {
return fmt.Errorf("节点 %s 的匹配模式必须为 and 或 or", node.ID)
}
for _, label := range cfg.Browsers {
if !uaBrowserLabels[label] {
return fmt.Errorf("节点 %s 的浏览器标签 %s 无效", node.ID, label)
}
}
for _, label := range cfg.OperatingSystems {
if !uaOSLabels[label] {
return fmt.Errorf("节点 %s 的操作系统标签 %s 无效", node.ID, label)
}
}
if len(cfg.CustomUAPatterns) > maxUACustomPatterns {
return fmt.Errorf("节点 %s 的自定义 UA 正则不能超过 %d 条", node.ID, maxUACustomPatterns)
}
for _, pattern := range cfg.CustomUAPatterns {
if strings.TrimSpace(pattern) == "" {
return fmt.Errorf("节点 %s 的自定义 UA 正则不能为空", node.ID)
}
if len(pattern) > maxUACustomPatternBytes {
return fmt.Errorf("节点 %s 的自定义 UA 正则不能超过 %d 字节", node.ID, maxUACustomPatternBytes)
}
if _, err := regexp.Compile(pattern); err != nil {
return fmt.Errorf("节点 %s 的自定义 UA 正则无效: %s", node.ID, pattern)
}
}
if cfg.BlockCustomUA && len(cfg.CustomUAPatterns) == 0 {
return fmt.Errorf("节点 %s 开启屏蔽自定义 UA 时至少需要一条正则", node.ID)
}
return nil
}
func validateSecurityCheckNodeConfig(node RuleNode) error {
var cfg SecurityCheckConfig
return decodeNodeConfig(node, &cfg)
}
var uaBrowserLabels = map[string]bool{
"Chrome": true, "Safari": true, "Firefox": true, "Edge": true, "Opera": true,
"Chromium": true, "WeChat": true, "Postman": true, "CLI": true, "Bot": true,
"Unknown": true, "Other": true,
}
var uaOSLabels = map[string]bool{
"Android": true, "iOS": true, "Windows": true, "macOS": true, "Chrome OS": true,
"Linux": true, "Bot": true, "Unknown": true, "Other": true,
}
func decodeNodeConfig(node RuleNode, dst any) error {
if err := decodeStrictConfig(node.Config, dst); err != nil {
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
}
return nil
}
func decodeStrictConfig(raw json.RawMessage, dst any) error {
trimmed := bytes.TrimSpace(raw)
if bytes.Equal(trimmed, []byte("null")) {
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 {
return slices.Contains(requiredHandles(t), handle)
}
func requiredHandles(t RuleNodeType) []string {
switch t {
case RuleNodeStart, RuleNodePoW:
return []string{"next"}
case RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodeUACheck, RuleNodeSecurityCheck:
return []string{"true", "false"}
case RuleNodeAllow, RuleNodeBlock:
return nil
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,133 @@
// 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 之间"},
{"invalid ua browser", func(g *RuleGraph) {
g.Nodes[1].Type = RuleNodeUACheck
g.Nodes[1].Config = rawConfig(`{"browsers":["NotABrowser"],"match_mode":"or"}`)
}, "节点 match-1 的浏览器标签 NotABrowser 无效"},
{"invalid ua match mode", func(g *RuleGraph) {
g.Nodes[1].Type = RuleNodeUACheck
g.Nodes[1].Config = rawConfig(`{"match_mode":"xor"}`)
}, "节点 match-1 的匹配模式必须为 and 或 or"},
{"invalid ua custom regex", func(g *RuleGraph) {
g.Nodes[1].Type = RuleNodeUACheck
g.Nodes[1].Config = rawConfig(`{"block_custom_ua":true,"custom_ua_patterns":["("]}`)
}, "节点 match-1 的自定义 UA 正则无效"},
{"custom ua requires patterns", func(g *RuleGraph) {
g.Nodes[1].Type = RuleNodeUACheck
g.Nodes[1].Config = rawConfig(`{"block_custom_ua":true}`)
}, "节点 match-1 开启屏蔽自定义 UA 时至少需要一条正则"},
{"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,536 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"maps"
"net"
"net/http"
"net/netip"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/agent"
"Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/kernel/model"
exprlang "github.com/expr-lang/expr"
"github.com/expr-lang/expr/vm"
)
const maxWAFIPGroupSubscriptionBytes = 2 * 1024 * 1024
type ipGroupAutoRuleEnv struct {
IP string `expr:"ip"`
RequestCount int `expr:"request_count"`
Status404Count int `expr:"status_404_count"`
Status404Ratio float64 `expr:"status_404_ratio"`
IPHostCount int `expr:"ip_host_count"`
IPHostRatio float64 `expr:"ip_host_ratio"`
ClientErrorCount int `expr:"client_error_count"`
ServerErrorCount int `expr:"server_error_count"`
LastSeenUnix int64 `expr:"last_seen_unix"`
statusCounts map[int]int
}
func (env ipGroupAutoRuleEnv) StatusCount(code any) int {
if env.statusCounts == nil {
return 0
}
return countStatusMatches(env.statusCounts, code)
}
func (env ipGroupAutoRuleEnv) StatusRatio(code any) float64 {
if env.RequestCount <= 0 || env.statusCounts == nil {
return 0.0
}
return float64(countStatusMatches(env.statusCounts, code)) / float64(env.RequestCount)
}
const maxHTTPStatusCodeDigits = 999
// countStatusMatches sums status counts for an exact code or class token.
// Accepted forms:
// - int / int64 / float64: exact status code (e.g. 404)
// - string digits: exact status code (e.g. "404")
// - string class: "1xx".."5xx" (case-insensitive), matching that hundred range
func countStatusMatches(statusCounts map[int]int, code any) int {
if statusCounts == nil {
return 0
}
switch v := code.(type) {
case int:
return statusCounts[v]
case int64:
if v < 0 || v > int64(maxHTTPStatusCodeDigits) {
return 0
}
return statusCounts[int(v)]
case float64:
if v != float64(int64(v)) || v < 0 || v > float64(maxHTTPStatusCodeDigits) {
return 0
}
return statusCounts[int(v)]
case string:
return countStatusMatchesString(statusCounts, v)
default:
return 0
}
}
func countStatusMatchesString(statusCounts map[int]int, raw string) int {
token := strings.TrimSpace(strings.ToLower(raw))
if token == "" {
return 0
}
if len(token) == 3 && token[1] == 'x' && token[2] == 'x' {
classDigit := token[0]
if classDigit < '1' || classDigit > '5' {
return 0
}
base := int(classDigit-'0') * 100
total := 0
for code, count := range statusCounts {
if code >= base && code < base+100 {
total += count
}
}
return total
}
// exact numeric string, e.g. "404"
var code int
for _, ch := range token {
if ch < '0' || ch > '9' {
return 0
}
code = code*10 + int(ch-'0')
if code > maxHTTPStatusCodeDigits {
return 0
}
}
return statusCounts[code]
}
type ipGroupAutoAccumulator struct {
ip string
requestCount int
status404Count int
ipHostCount int
clientErrorCount int
serverErrorCount int
lastSeen time.Time
statusCounts map[int]int
}
// SyncDueWAFIPGroups syncs all enabled automatic/subscription IP groups that are due.
func SyncDueWAFIPGroups(ctx context.Context) error {
now := time.Now().UTC()
groups, err := repository.ListDueOpenFlareWAFIPGroups(ctx, now)
if err != nil {
return err
}
for _, group := range groups {
if _, err := syncOpenFlareWAFIPGroup(ctx, group, now); err != nil {
continue
}
}
return nil
}
func syncOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGroup, now time.Time) (*IPGroupSyncResult, error) {
if group == nil {
return nil, errors.New("IP 组不存在")
}
switch group.Type {
case wafIPGroupTypeSubscription:
return syncIPGroupSubscription(ctx, group, now)
case wafIPGroupTypeAutomatic:
return syncIPGroupAutomatic(ctx, group, now)
default:
return nil, &RuleValidationError{Err: errors.New("只有自动和订阅类型 IP 组支持同步")}
}
}
func syncIPGroupSubscription(ctx context.Context, group *model.OpenFlareWAFIPGroup, now time.Time) (*IPGroupSyncResult, error) {
content, err := downloadIPGroupSubscription(ctx, group.SubscriptionURL)
if err != nil {
recordIPGroupSyncFailure(ctx, group, now, err)
return nil, err
}
ips, err := parseIPGroupSubscription(content, group.SubscriptionFormat, group.SubscriptionMappingRule)
if err != nil {
recordIPGroupSyncFailure(ctx, group, now, err)
return nil, err
}
ipListJSON, _ := json.Marshal(ips)
nextSyncAt := now.Add(time.Duration(group.SyncIntervalMinutes) * time.Minute)
group.IPList = string(ipListJSON)
group.LastSyncedAt = &now
group.NextSyncAt = &nextSyncAt
group.LastSyncStatus = "success"
group.LastSyncMessage = fmt.Sprintf("同步成功,共 %d 条 IP/IP 段", len(ips))
if err := repository.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group); err != nil {
return nil, err
}
broadcastIPGroupToAgents(ctx, group.ID)
view, err := GetIPGroup(ctx, group.ID)
if err != nil {
return nil, err
}
return &IPGroupSyncResult{
Group: *view,
IPCount: len(ips),
SyncedAt: now.Format(time.RFC3339),
NextSyncAt: nextSyncAt.Format(time.RFC3339),
Status: group.LastSyncStatus,
Message: group.LastSyncMessage,
}, nil
}
func syncIPGroupAutomatic(ctx context.Context, group *model.OpenFlareWAFIPGroup, now time.Time) (*IPGroupSyncResult, error) {
config, err := parseIPGroupAutoConfig(json.RawMessage(group.AutoConfig))
if err != nil {
recordIPGroupSyncFailure(ctx, group, now, err)
return nil, err
}
var existingExtIPs []ipGroupExtIP
if group.ExtIPs != "" && group.ExtIPs != "[]" {
_ = json.Unmarshal([]byte(group.ExtIPs), &existingExtIPs)
}
activeExtIPs := make([]ipGroupExtIP, 0, len(existingExtIPs))
for _, extIP := range existingExtIPs {
if config.TTL > 0 {
expirationTime := extIP.CapturedAt.Add(time.Duration(config.TTL) * time.Second)
if expirationTime.Before(now) {
continue
}
}
activeExtIPs = append(activeExtIPs, extIP)
}
ips, err := evaluateParsedIPGroupAutoConfig(ctx, config, now)
if err != nil {
recordIPGroupSyncFailure(ctx, group, now, err)
return nil, err
}
extIPMap := make(map[string]int)
for idx, extIP := range activeExtIPs {
extIPMap[extIP.IP] = idx
}
for _, ip := range ips {
if idx, ok := extIPMap[ip]; ok {
activeExtIPs[idx].CapturedAt = now
} else {
activeExtIPs = append(activeExtIPs, ipGroupExtIP{
IP: ip,
CapturedAt: now,
})
}
}
finalIPs := make([]string, 0, len(activeExtIPs))
for _, extIP := range activeExtIPs {
finalIPs = append(finalIPs, extIP.IP)
}
finalIPs, err = normalizeIPList(finalIPs)
if err != nil {
recordIPGroupSyncFailure(ctx, group, now, err)
return nil, err
}
extIPsJSON, _ := json.Marshal(activeExtIPs)
ipListJSON, _ := json.Marshal(finalIPs)
nextSyncAt := now.Add(time.Duration(normalizeIPGroupSyncInterval(group.SyncIntervalMinutes)) * time.Minute)
group.IPList = string(ipListJSON)
group.ExtIPs = string(extIPsJSON)
group.LastSyncedAt = &now
group.NextSyncAt = &nextSyncAt
group.LastSyncStatus = "success"
group.LastSyncMessage = fmt.Sprintf("自动规则执行成功,共命中 %d 个 IP,当前生效 %d 个 IP", len(ips), len(finalIPs))
if err := repository.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group); err != nil {
return nil, err
}
broadcastIPGroupToAgents(ctx, group.ID)
view, err := GetIPGroup(ctx, group.ID)
if err != nil {
return nil, err
}
return &IPGroupSyncResult{
Group: *view,
IPCount: len(finalIPs),
SyncedAt: now.Format(time.RFC3339),
NextSyncAt: nextSyncAt.Format(time.RFC3339),
Status: group.LastSyncStatus,
Message: group.LastSyncMessage,
}, nil
}
func recordIPGroupSyncFailure(ctx context.Context, group *model.OpenFlareWAFIPGroup, now time.Time, syncErr error) {
nextSyncAt := now.Add(time.Duration(normalizeIPGroupSyncInterval(group.SyncIntervalMinutes)) * time.Minute)
group.LastSyncedAt = &now
group.NextSyncAt = &nextSyncAt
group.LastSyncStatus = "failed"
group.LastSyncMessage = syncErr.Error()
_ = repository.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group)
}
func evaluateParsedIPGroupAutoConfig(ctx context.Context, config ipGroupAutoConfig, now time.Time) ([]string, error) {
if len(config.Rules) == 0 {
return []string{}, nil
}
programs := make([]*vm.Program, 0, len(config.Rules))
for i, rule := range config.Rules {
program, err := exprlang.Compile(rule.Expr, exprlang.Env(ipGroupAutoRuleEnv{}), exprlang.AsBool())
if err != nil {
return nil, fmt.Errorf("自动规则 %s Expr 无效: %w", displayIPGroupAutoRuleName(rule, i), err)
}
programs = append(programs, program)
}
lookback := config.lookbackDuration
if lookback <= 0 {
lookback = defaultWAFIPGroupAutoLookbackDur
}
aggregates, err := repository.ListOpenFlareAccessLogWAFIPAggregates(ctx, model.OpenFlareAccessLogQuery{
Since: now.Add(-lookback),
Until: now,
})
if err != nil {
return nil, err
}
accumulators := make(map[string]*ipGroupAutoAccumulator, len(aggregates))
for _, item := range aggregates {
if item == nil {
continue
}
ip, ok := normalizeIPLiteral(item.RemoteAddr)
if !ok {
continue
}
lastSeen := time.Time{}
if item.LastSeenEpoch > 0 {
lastSeen = time.Unix(item.LastSeenEpoch, 0).UTC()
}
statusCounts := make(map[int]int, len(item.StatusCounts))
maps.Copy(statusCounts, item.StatusCounts)
accumulators[ip] = &ipGroupAutoAccumulator{
ip: ip,
requestCount: item.RequestCount,
status404Count: item.Status404Count,
ipHostCount: item.IPHostCount,
clientErrorCount: item.ClientErrorCount,
serverErrorCount: item.ServerErrorCount,
lastSeen: lastSeen,
statusCounts: statusCounts,
}
}
matched := make([]string, 0)
for _, acc := range accumulators {
env := acc.toExprEnv()
for _, program := range programs {
output, err := exprlang.Run(program, env)
if err != nil {
return nil, fmt.Errorf("执行自动规则失败: %w", err)
}
if matchedRule, ok := output.(bool); ok && matchedRule {
matched = append(matched, acc.ip)
break
}
}
}
return normalizeIPList(matched)
}
func (acc *ipGroupAutoAccumulator) toExprEnv() ipGroupAutoRuleEnv {
env := ipGroupAutoRuleEnv{
IP: acc.ip,
RequestCount: acc.requestCount,
Status404Count: acc.status404Count,
IPHostCount: acc.ipHostCount,
ClientErrorCount: acc.clientErrorCount,
ServerErrorCount: acc.serverErrorCount,
statusCounts: acc.statusCounts,
}
if acc.requestCount > 0 {
env.Status404Ratio = float64(acc.status404Count) / float64(acc.requestCount)
env.IPHostRatio = float64(acc.ipHostCount) / float64(acc.requestCount)
}
if !acc.lastSeen.IsZero() {
env.LastSeenUnix = acc.lastSeen.Unix()
}
return env
}
func displayIPGroupAutoRuleName(rule ipGroupAutoRule, index int) string {
if rule.Name != "" {
return rule.Name
}
return fmt.Sprintf("#%d", index+1)
}
func normalizeIPLiteral(value string) (string, bool) {
host := strings.TrimSpace(value)
if host == "" {
return "", false
}
if parsedHost, _, err := net.SplitHostPort(host); err == nil {
host = parsedHost
}
host = strings.Trim(host, "[]")
addr, err := netip.ParseAddr(host)
if err != nil {
return "", false
}
return addr.String(), true
}
func downloadIPGroupSubscription(ctx context.Context, rawURL string) ([]byte, error) {
if err := validateSubscriptionURL(rawURL); err != nil {
return nil, err
}
client := http.Client{Timeout: 15 * time.Second}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
if err != nil {
return nil, fmt.Errorf("下载订阅失败: %w", err)
}
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("下载订阅失败: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("订阅返回状态码 %d", resp.StatusCode)
}
var buffer bytes.Buffer
reader := io.LimitReader(resp.Body, maxWAFIPGroupSubscriptionBytes+1)
if _, err := buffer.ReadFrom(reader); err != nil {
return nil, fmt.Errorf("读取订阅内容失败: %w", err)
}
if buffer.Len() > maxWAFIPGroupSubscriptionBytes {
return nil, fmt.Errorf("订阅内容不能超过 %d 字节", maxWAFIPGroupSubscriptionBytes)
}
return buffer.Bytes(), nil
}
func parseIPGroupSubscription(content []byte, format string, mappingRule string) ([]string, error) {
switch normalizeIPGroupSubscriptionFormat(format) {
case wafIPGroupSubscriptionFormatJSON:
items, err := parseIPGroupJSONSubscription(content, mappingRule)
if err != nil {
return nil, err
}
return normalizeIPList(items)
default:
return normalizeIPList(parseIPGroupTextSubscription(string(content)))
}
}
func parseIPGroupTextSubscription(text string) []string {
lines := strings.Split(text, "\n")
items := make([]string, 0, len(lines))
for _, line := range lines {
item := strings.TrimSpace(line)
if item == "" || strings.HasPrefix(item, "#") {
continue
}
items = append(items, item)
}
return items
}
func parseIPGroupJSONSubscription(content []byte, mappingRule string) ([]string, error) {
var payload any
if err := json.Unmarshal(content, &payload); err != nil {
return nil, fmt.Errorf("JSON 订阅解析失败: %w", err)
}
nodes, err := selectJSONMappingNodes(payload, mappingRule)
if err != nil {
return nil, err
}
items := make([]string, 0, len(nodes))
for _, node := range nodes {
collectJSONStrings(node, &items)
}
if len(items) == 0 {
return nil, errors.New("JSON 订阅没有解析到 IP/IP 段")
}
return items, nil
}
func selectJSONMappingNodes(payload any, mappingRule string) ([]any, error) {
rule := strings.TrimSpace(mappingRule)
if rule == "" || rule == "$" {
return []any{payload}, nil
}
rule = strings.TrimPrefix(rule, "$.")
nodes := []any{payload}
for rawSegment := range strings.SplitSeq(rule, ".") {
segment := strings.TrimSpace(rawSegment)
if segment == "" {
continue
}
expandArray := strings.HasSuffix(segment, "[]")
segment = strings.TrimSuffix(segment, "[]")
next := make([]any, 0)
for _, node := range nodes {
object, ok := node.(map[string]any)
if !ok {
continue
}
value, ok := object[segment]
if !ok {
continue
}
if expandArray {
array, ok := value.([]any)
if !ok {
continue
}
next = append(next, array...)
} else {
next = append(next, value)
}
}
nodes = next
}
if len(nodes) == 0 {
return nil, fmt.Errorf("JSON 映射规则 %q 未匹配到内容", mappingRule)
}
return nodes, nil
}
func collectJSONStrings(node any, items *[]string) {
switch value := node.(type) {
case string:
*items = append(*items, value)
case []any:
for _, item := range value {
collectJSONStrings(item, items)
}
}
}
func broadcastIPGroupToAgents(ctx context.Context, id uint) {
groups, err := agent.WAFIPGroupsForAgent(ctx, []uint{id})
if err != nil || len(groups) == 0 {
if err != nil {
slog.Debug("build waf ip group broadcast payload failed", "id", id, "error", err)
}
return
}
websocket.BroadcastWAFIPGroups(groups)
}
@@ -0,0 +1,312 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupIPGroupSyncTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFIPGroup{},
))
db.SetDB(sqliteDB)
testhelper.SetupLogStoresForTest(t)
return func() {
db.SetDB(nil)
}
}
func TestParseIPGroupSubscriptionParsers(t *testing.T) {
textItems, err := parseIPGroupSubscription([]byte("# comment\n203.0.113.10\n\n198.51.100.0/24\n"), "text", "")
require.NoError(t, err)
require.Len(t, textItems, 2)
assert.Equal(t, "198.51.100.0/24", textItems[0])
assert.Equal(t, "203.0.113.10", textItems[1])
jsonItems, err := parseIPGroupSubscription([]byte(`{"data":{"items":[{"ip":"203.0.113.11"},{"ip":"203.0.113.12"}]}}`), "json", "data.items[].ip")
require.NoError(t, err)
require.Len(t, jsonItems, 2)
assert.Equal(t, "203.0.113.11", jsonItems[0])
assert.Equal(t, "203.0.113.12", jsonItems[1])
}
func TestSyncIPGroupDownloadsSubscription(t *testing.T) {
cleanup := setupIPGroupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("203.0.113.20\n"))
}))
defer server.Close()
group, err := CreateIPGroup(ctx, IPGroupInput{
Name: "subscription",
Type: wafIPGroupTypeSubscription,
Enabled: true,
SubscriptionURL: server.URL,
SubscriptionFormat: wafIPGroupSubscriptionFormatText,
SyncIntervalMinutes: 10,
})
require.NoError(t, err)
result, err := SyncIPGroup(ctx, group.ID)
require.NoError(t, err)
require.Equal(t, 1, result.IPCount)
assert.Equal(t, "203.0.113.20", result.Group.IPList[0])
assert.Equal(t, "success", result.Status)
}
func TestSyncIPGroupAutomaticExprRules(t *testing.T) {
cleanup := setupIPGroupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
seedWAFAccessLogs(t, ctx, now, "203.0.113.10", "app.example.com", 101, 81)
seedWAFAccessLogs(t, ctx, now, "203.0.113.11", "198.51.100.10", 60, 0)
seedWAFAccessLogs(t, ctx, now, "203.0.113.12", "app.example.com", 120, 10)
group, err := CreateIPGroup(ctx, IPGroupInput{
Name: "auto blacklist",
Type: wafIPGroupTypeAutomatic,
Enabled: true,
AutoConfig: json.RawMessage(`{
"lookback": "60m",
"rules": [
{"name":"单 IP 404 高频扫描","expr":"request_count > 100 && StatusRatio(404) >= 0.8"},
{"name":"单 IP 直连访问异常","expr":"ip_host_count > 50 && ip_host_ratio > 0.5"}
]
}`),
})
require.NoError(t, err)
result, err := SyncIPGroup(ctx, group.ID)
require.NoError(t, err)
require.Equal(t, 2, result.IPCount)
want := map[string]bool{"203.0.113.10": true, "203.0.113.11": true}
for _, item := range result.Group.IPList {
assert.True(t, want[item], "unexpected matched IP %s", item)
delete(want, item)
}
assert.Empty(t, want)
}
func TestTestIPGroupAutoConfigReturnsMatchedIPs(t *testing.T) {
cleanup := setupIPGroupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
seedWAFAccessLogs(t, ctx, now, "203.0.113.10", "app.example.com", 101, 81)
seedWAFAccessLogs(t, ctx, now, "203.0.113.11", "198.51.100.10", 60, 0)
seedWAFAccessLogs(t, ctx, now, "203.0.113.12", "app.example.com", 120, 10)
result, err := TestIPGroupAutoConfig(ctx, IPGroupAutoTestInput{
AutoConfig: json.RawMessage(`{
"lookback": "1h",
"rules": [
{"name":"单 IP 404 高频扫描","expr":"request_count > 100 && StatusRatio(404) >= 0.8"},
{"name":"单 IP 直连访问异常","expr":"ip_host_count > 50 && ip_host_ratio > 0.5"}
]
}`),
})
require.NoError(t, err)
assert.Equal(t, 2, result.MatchedCount)
assert.Equal(t, 2, result.RuleCount)
assert.Equal(t, "1h", result.Lookback)
want := map[string]bool{"203.0.113.10": true, "203.0.113.11": true}
for _, item := range result.MatchedIPs {
assert.True(t, want[item], "unexpected matched IP %s", item)
delete(want, item)
}
assert.Empty(t, want)
}
func TestListDueOpenFlareWAFIPGroups(t *testing.T) {
cleanup := setupIPGroupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
past := time.Now().UTC().Add(-time.Hour)
future := time.Now().UTC().Add(time.Hour)
dueAuto := &model.OpenFlareWAFIPGroup{
Name: "due auto", Type: wafIPGroupTypeAutomatic, Enabled: true,
IPList: "[]", AutoConfig: "{}", ExtIPs: "[]", NextSyncAt: &past,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, dueAuto))
futureAuto := &model.OpenFlareWAFIPGroup{
Name: "future auto", Type: wafIPGroupTypeAutomatic, Enabled: true,
IPList: "[]", AutoConfig: "{}", ExtIPs: "[]", NextSyncAt: &future,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, futureAuto))
dueSub := &model.OpenFlareWAFIPGroup{
Name: "due sub", Type: wafIPGroupTypeSubscription, Enabled: true,
IPList: "[]", AutoConfig: "{}", ExtIPs: "[]",
SubscriptionURL: "https://example.com/list", NextSyncAt: &past,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, dueSub))
manual := &model.OpenFlareWAFIPGroup{
Name: "manual", Type: wafIPGroupTypeManual, Enabled: true,
IPList: "[]", AutoConfig: "{}", ExtIPs: "[]", NextSyncAt: &past,
}
require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, manual))
groups, err := repository.ListDueOpenFlareWAFIPGroups(ctx, time.Now().UTC())
require.NoError(t, err)
require.Len(t, groups, 2)
ids := []uint{groups[0].ID, groups[1].ID}
assert.Contains(t, ids, dueAuto.ID)
assert.Contains(t, ids, dueSub.ID)
}
func seedWAFAccessLogs(t *testing.T, ctx context.Context, loggedAt time.Time, remoteAddr string, host string, total int, notFound int) {
t.Helper()
records := make([]*model.OpenFlareAccessLog, 0, total)
for i := 0; i < total; i++ {
statusCode := http.StatusOK
if i < notFound {
statusCode = http.StatusNotFound
}
records = append(records, &model.OpenFlareAccessLog{
NodeID: "node-waf-auto",
LoggedAt: loggedAt.Add(-time.Duration(i%30) * time.Second),
RemoteAddr: remoteAddr,
Host: host,
Path: "/probe",
StatusCode: statusCode,
})
}
require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, records))
}
func TestParseIPGroupAutoConfigLookback(t *testing.T) {
cases := []struct {
name string
raw string
want string
wantDur time.Duration
wantErr bool
}{
{name: "duration 60m", raw: `{"lookback":"60m","rules":[]}`, want: "1h", wantDur: time.Hour},
{name: "duration 1h", raw: `{"lookback":"1h","rules":[]}`, want: "1h", wantDur: time.Hour},
{name: "duration 30m", raw: `{"lookback":"30m","rules":[]}`, want: "30m", wantDur: 30 * time.Minute},
{name: "duration 1m no min floor", raw: `{"lookback":"1m","rules":[]}`, want: "1m", wantDur: time.Minute},
{name: "legacy minutes", raw: `{"lookback_minutes":45,"rules":[]}`, want: "45m", wantDur: 45 * time.Minute},
{name: "default empty", raw: `{"rules":[]}`, want: "1h", wantDur: time.Hour},
{name: "invalid", raw: `{"lookback":"abc","rules":[]}`, wantErr: true},
{name: "zero lookback uses default", raw: `{"lookback":"","rules":[]}`, want: "1h", wantDur: time.Hour},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
cfg, err := parseIPGroupAutoConfig(json.RawMessage(tc.raw))
if tc.wantErr {
if err == nil {
t.Fatalf("parseIPGroupAutoConfig(%s) error = nil, want error", tc.raw)
}
return
}
if err != nil {
t.Fatalf("parseIPGroupAutoConfig(%s) error = %v", tc.raw, err)
}
if cfg.Lookback != tc.want {
t.Errorf("Lookback = %q, want %q", cfg.Lookback, tc.want)
}
if cfg.lookbackDuration != tc.wantDur {
t.Errorf("lookbackDuration = %v, want %v", cfg.lookbackDuration, tc.wantDur)
}
})
}
}
func TestCountStatusMatchesSupportsClassTokens(t *testing.T) {
counts := map[int]int{
200: 10,
201: 5,
404: 20,
403: 10,
500: 4,
502: 1,
}
cases := []struct {
name string
code any
want int
}{
{name: "exact int", code: 404, want: 20},
{name: "exact string", code: "403", want: 10},
{name: "2xx class", code: "2xx", want: 15},
{name: "4xx class upper", code: "4XX", want: 30},
{name: "5xx class", code: "5xx", want: 5},
{name: "unknown class", code: "9xx", want: 0},
{name: "invalid token", code: "abc", want: 0},
{name: "float exact", code: float64(200), want: 10},
{name: "float non-int", code: 200.5, want: 0},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
got := countStatusMatches(counts, tc.code)
if got != tc.want {
t.Errorf("countStatusMatches(%v) = %d, want %d", tc.code, got, tc.want)
}
})
}
}
func TestStatusRatioClassTokenInExpr(t *testing.T) {
cleanup := setupIPGroupSyncTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
// 100 requests, 80 of which are 404 → 4xx ratio 0.8
seedWAFAccessLogs(t, ctx, now, "203.0.113.40", "app.example.com", 100, 80)
// mostly OK → should not match
seedWAFAccessLogs(t, ctx, now, "203.0.113.41", "app.example.com", 100, 10)
result, err := TestIPGroupAutoConfig(ctx, IPGroupAutoTestInput{
AutoConfig: json.RawMessage(`{
"lookback": "60m",
"rules": [
{"name":"高 4xx 占比","expr":"request_count >= 100 && StatusRatio(\"4xx\") >= 0.8"}
]
}`),
})
require.NoError(t, err)
assert.Equal(t, 1, result.MatchedCount)
require.Len(t, result.MatchedIPs, 1)
assert.Equal(t, "203.0.113.40", result.MatchedIPs[0])
}
@@ -0,0 +1,842 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/netip"
"net/url"
"sort"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
exprlang "github.com/expr-lang/expr"
"gorm.io/gorm"
)
const (
maxWAFBlockBodyBytes = 16 * 1024
wafIPGroupTypeManual = "manual"
wafIPGroupTypeAutomatic = "automatic"
wafIPGroupTypeSubscription = "subscription"
wafIPGroupSubscriptionFormatText = "text"
wafIPGroupSubscriptionFormatJSON = "json"
defaultWAFIPGroupSyncIntervalMinutes = 1440
defaultWAFIPGroupAutoLookback = "1h"
defaultWAFIPGroupAutoLookbackDur = time.Hour
minWAFIPGroupSyncIntervalMinutes = 1
maxWAFIPGroupSyncIntervalMinutes = 43200
maxWAFIPGroupAutoLookback = 30 * 24 * time.Hour
minPoWSessionTTLSeconds = 60
minPoWChallengeTTLSeconds = 30
powAlgorithmFast = "fast"
powAlgorithmSlow = "slow"
)
// SiteRuleGroupsView is the site-level WAF binding view.
type SiteRuleGroupsView struct {
RouteID uint `json:"route_id"`
GlobalRuleGroup *RuleView `json:"global_rule_group"`
RuleGroups []RuleView `json:"rule_groups"`
AppliedRuleGroups []RuleView `json:"applied_rule_groups"`
AppliedIDs []uint `json:"applied_ids"`
}
// IDsRequest carries a list of numeric ids.
type IDsRequest struct {
IDs []uint `json:"ids"`
}
// IPGroupInput is the create/update payload for WAF IP groups.
type IPGroupInput struct {
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
AutoConfig json.RawMessage `json:"auto_config"`
SubscriptionURL string `json:"subscription_url"`
SubscriptionFormat string `json:"subscription_format"`
SubscriptionMappingRule string `json:"subscription_mapping_rule"`
SyncIntervalMinutes int `json:"sync_interval_minutes"`
}
// IPGroupExtIPView is an external IP entry in API responses.
type IPGroupExtIPView struct {
IP string `json:"ip"`
CapturedAt string `json:"captured_at"`
}
// IPGroupView is the API view for a WAF IP group.
type IPGroupView struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
AutoConfig json.RawMessage `json:"auto_config"`
ExtIPs []IPGroupExtIPView `json:"ext_ips"`
SubscriptionURL string `json:"subscription_url"`
SubscriptionFormat string `json:"subscription_format"`
SubscriptionMappingRule string `json:"subscription_mapping_rule"`
SyncIntervalMinutes int `json:"sync_interval_minutes"`
LastSyncedAt string `json:"last_synced_at,omitempty"`
NextSyncAt string `json:"next_sync_at,omitempty"`
LastSyncStatus string `json:"last_sync_status"`
LastSyncMessage string `json:"last_sync_message"`
ReferencedByRuleCount int `json:"referenced_by_rule_count"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
// IPGroupSyncResult is the response for manual IP group sync.
type IPGroupSyncResult struct {
Group IPGroupView `json:"group"`
IPCount int `json:"ip_count"`
SyncedAt string `json:"synced_at"`
NextSyncAt string `json:"next_sync_at"`
Status string `json:"status"`
Message string `json:"message"`
}
// IPGroupAutoTestInput tests automatic IP group configuration.
type IPGroupAutoTestInput struct {
AutoConfig json.RawMessage `json:"auto_config"`
}
// IPGroupAutoTestResult is the response for automatic IP group test.
type IPGroupAutoTestResult struct {
MatchedIPs []string `json:"matched_ips"`
MatchedCount int `json:"matched_count"`
Lookback string `json:"lookback"`
RuleCount int `json:"rule_count"`
TestedAt string `json:"tested_at"`
}
type ipGroupAutoConfig struct {
Lookback string `json:"lookback"`
TTL int `json:"ttl"`
Rules []ipGroupAutoRule `json:"rules"`
// lookbackDuration is resolved from Lookback (and legacy lookback_minutes) for runtime queries.
lookbackDuration time.Duration `json:"-"`
}
type ipGroupAutoRule struct {
Name string `json:"name"`
Expr string `json:"expr"`
}
type ipGroupExtIP struct {
IP string `json:"ip"`
CapturedAt time.Time `json:"captured_at"`
}
// GetSiteRuleGroups returns WAF rule groups for a proxy route.
func GetSiteRuleGroups(ctx context.Context, routeID uint) (*SiteRuleGroupsView, error) {
if _, err := repository.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil {
return nil, err
}
groups, err := ListRules(ctx)
if err != nil {
return nil, err
}
appliedIDs, err := ListSiteRuleGroupIDs(ctx, routeID)
if err != nil {
return nil, err
}
var global *RuleView
custom := make([]RuleView, 0, len(groups))
applied := make([]RuleView, 0, len(appliedIDs))
groupByID := make(map[uint]RuleView, len(groups))
for index := range groups {
group := groups[index]
if group.IsGlobal {
item := group
global = &item
continue
}
custom = append(custom, group)
groupByID[group.ID] = group
}
for _, id := range appliedIDs {
if group, ok := groupByID[id]; ok {
applied = append(applied, group)
}
}
return &SiteRuleGroupsView{
RouteID: routeID,
GlobalRuleGroup: global,
RuleGroups: custom,
AppliedRuleGroups: applied,
AppliedIDs: appliedIDs,
}, nil
}
// ReplaceSiteRuleGroups replaces rule group bindings for a proxy route.
func ReplaceSiteRuleGroups(ctx context.Context, routeID uint, groupIDs []uint) (*SiteRuleGroupsView, error) {
if _, err := repository.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil {
return nil, err
}
normalized, err := normalizeRuleGroupIDs(ctx, groupIDs)
if err != nil {
return nil, &RuleValidationError{Err: err}
}
if err = repository.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, routeID, normalized); err != nil {
return nil, err
}
return GetSiteRuleGroups(ctx, routeID)
}
// ListSiteRuleGroupIDs returns rule group ids bound to a proxy route.
func ListSiteRuleGroupIDs(ctx context.Context, routeID uint) ([]uint, error) {
bindings, err := repository.ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx, routeID)
if err != nil {
return nil, err
}
ids := make([]uint, 0, len(bindings))
for _, binding := range bindings {
ids = append(ids, binding.RuleGroupID)
}
return ids, nil
}
// EnsureDefaultRuleGroup ensures the global WAF rule group exists.
func EnsureDefaultRuleGroup(ctx context.Context) error {
_, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx)
if err == nil {
return nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
graph, marshalErr := json.Marshal(DefaultRuleGraph())
if marshalErr != nil {
return marshalErr
}
group := &model.OpenFlareWAFRuleGroup{
Name: "全局规则组", Enabled: true, IsGlobal: true, Graph: string(graph), Revision: 1,
}
return repository.CreateOpenFlareWAFRuleGroup(ctx, group)
}
// ListIPGroups returns all WAF IP groups.
func ListIPGroups(ctx context.Context) ([]IPGroupView, error) {
groups, err := repository.ListOpenFlareWAFIPGroups(ctx)
if err != nil {
return nil, err
}
referenceCounts, err := loadIPGroupReferenceCounts(ctx)
if err != nil {
return nil, err
}
views := make([]IPGroupView, 0, len(groups))
for _, group := range groups {
view, buildErr := buildIPGroupView(group, referenceCounts[group.ID])
if buildErr != nil {
return nil, buildErr
}
views = append(views, view)
}
return views, nil
}
// GetIPGroup returns a WAF IP group by id.
func GetIPGroup(ctx context.Context, id uint) (*IPGroupView, error) {
group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if err != nil {
return nil, err
}
referenceCounts, err := loadIPGroupReferenceCounts(ctx)
if err != nil {
return nil, err
}
view, err := buildIPGroupView(group, referenceCounts[group.ID])
if err != nil {
return nil, err
}
return &view, nil
}
// CreateIPGroup creates a WAF IP group.
func CreateIPGroup(ctx context.Context, input IPGroupInput) (*IPGroupView, error) {
group, err := buildIPGroup(nil, input)
if err != nil {
return nil, &RuleValidationError{Err: err}
}
if err = repository.CreateOpenFlareWAFIPGroup(ctx, group); err != nil {
return nil, err
}
broadcastIPGroupToAgents(ctx, group.ID)
return GetIPGroup(ctx, group.ID)
}
// UpdateIPGroup updates a WAF IP group.
func UpdateIPGroup(ctx context.Context, id uint, input IPGroupInput) (*IPGroupView, error) {
group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if err != nil {
return nil, err
}
group, err = buildIPGroup(group, input)
if err != nil {
return nil, &RuleValidationError{Err: err}
}
if err = repository.UpdateOpenFlareWAFIPGroup(ctx, group); err != nil {
return nil, err
}
broadcastIPGroupToAgents(ctx, group.ID)
return GetIPGroup(ctx, group.ID)
}
// DeleteIPGroup deletes a WAF IP group when not referenced.
func DeleteIPGroup(ctx context.Context, id uint) error {
group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if err != nil {
return err
}
counts, err := loadIPGroupReferenceCounts(ctx)
if err != nil {
return err
}
if counts[group.ID] > 0 {
return &RuleValidationError{Err: errors.New("IP 组已被 WAF 规则引用,请先移除引用")}
}
return repository.DeleteOpenFlareWAFIPGroup(ctx, group.ID)
}
// SyncIPGroup synchronizes a subscription or automatic WAF IP group.
func SyncIPGroup(ctx context.Context, id uint) (*IPGroupSyncResult, error) {
group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if err != nil {
return nil, err
}
return syncOpenFlareWAFIPGroup(ctx, group, time.Now().UTC())
}
// TestIPGroupAutoConfig evaluates automatic IP group rules against recent access logs.
func TestIPGroupAutoConfig(ctx context.Context, input IPGroupAutoTestInput) (*IPGroupAutoTestResult, error) {
config, err := parseIPGroupAutoConfig(input.AutoConfig)
if err != nil {
return nil, &RuleValidationError{Err: err}
}
now := time.Now().UTC()
ips, err := evaluateParsedIPGroupAutoConfig(ctx, config, now)
if err != nil {
return nil, err
}
return &IPGroupAutoTestResult{
MatchedIPs: ips,
MatchedCount: len(ips),
Lookback: config.Lookback,
RuleCount: len(config.Rules),
TestedAt: now.Format(time.RFC3339),
}, nil
}
func loadRuleGroupBindings(ctx context.Context) (map[uint][]uint, error) {
bindings, err := repository.ListOpenFlareWAFRuleGroupBindings(ctx)
if err != nil {
return nil, err
}
result := make(map[uint][]uint, len(bindings))
for _, binding := range bindings {
result[binding.RuleGroupID] = append(result[binding.RuleGroupID], binding.ProxyRouteID)
}
return result, nil
}
func loadIPGroupReferenceCounts(ctx context.Context) (map[uint]int, error) {
groups, err := repository.ListOpenFlareWAFRuleGroups(ctx)
if err != nil {
return nil, err
}
counts := make(map[uint]int)
for _, group := range groups {
graph := DefaultRuleGraph()
if strings.TrimSpace(group.Graph) != "" {
if err = json.Unmarshal([]byte(group.Graph), &graph); err != nil {
return nil, fmt.Errorf("decode WAF rule %d graph: %w", group.ID, err)
}
}
for _, id := range ReferencedIPGroupIDs(graph) {
counts[id]++
}
}
return counts, nil
}
func pruneIPGroupExtIPs(group *model.OpenFlareWAFIPGroup, ipList []string) error {
if group == nil {
return nil
}
allowed := make(map[string]struct{}, len(ipList))
for _, ip := range ipList {
allowed[ip] = struct{}{}
}
var extIPs []ipGroupExtIP
if group.ExtIPs != "" && group.ExtIPs != "[]" {
if err := json.Unmarshal([]byte(group.ExtIPs), &extIPs); err != nil {
return err
}
}
pruned := make([]ipGroupExtIP, 0, len(extIPs))
for _, extIP := range extIPs {
if _, ok := allowed[extIP.IP]; ok {
pruned = append(pruned, extIP)
}
}
extIPsJSON, err := json.Marshal(pruned)
if err != nil {
return err
}
group.ExtIPs = string(extIPsJSON)
return nil
}
func normalizeIPList(items []string) ([]string, error) {
normalized := make([]string, 0, len(items))
for _, raw := range items {
item := strings.TrimSpace(raw)
if item == "" {
continue
}
if strings.Contains(item, "/") {
prefix, err := netip.ParsePrefix(item)
if err != nil {
return nil, fmt.Errorf("%s 不是合法 IP 段", item)
}
item = prefix.Masked().String()
} else {
addr, err := netip.ParseAddr(item)
if err != nil {
return nil, fmt.Errorf("%s 不是合法 IP", item)
}
item = addr.String()
}
normalized = append(normalized, item)
}
normalized = uniqueStrings(normalized)
sort.Strings(normalized)
return normalized, nil
}
func decodeStringList(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []string{}, nil
}
var items []string
if err := json.Unmarshal([]byte(text), &items); err != nil {
return nil, err
}
return items, nil
}
func normalizeRuleGroupIDs(ctx context.Context, groupIDs []uint) ([]uint, error) {
normalized := uniqueUintIDsInOrder(groupIDs)
for _, groupID := range normalized {
group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, groupID)
if err != nil {
return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID)
}
if group.IsGlobal {
return nil, errors.New("全局 WAF 规则组不需要手动绑定")
}
}
return normalized, nil
}
func uniqueUintIDsInOrder(ids []uint) []uint {
seen := make(map[uint]struct{}, len(ids))
result := make([]uint, 0, len(ids))
for _, id := range ids {
if id == 0 {
continue
}
if _, exists := seen[id]; exists {
continue
}
seen[id] = struct{}{}
result = append(result, id)
}
return result
}
func uniqueStrings(items []string) []string {
seen := make(map[string]struct{}, len(items))
result := make([]string, 0, len(items))
for _, item := range items {
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
result = append(result, item)
}
return result
}
func normalizeIPGroupAutoConfig(raw json.RawMessage) (string, error) {
text := strings.TrimSpace(string(raw))
if text == "" {
text = "{}"
}
config, err := parseIPGroupAutoConfig(json.RawMessage(text))
if err != nil {
return "", err
}
normalized, _ := json.Marshal(config)
return string(normalized), nil
}
func parseIPGroupAutoConfig(raw json.RawMessage) (ipGroupAutoConfig, error) {
text := strings.TrimSpace(string(raw))
if text == "" {
text = "{}"
}
var object map[string]any
if err := json.Unmarshal([]byte(text), &object); err != nil || object == nil {
return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象")
}
var config ipGroupAutoConfig
if err := json.Unmarshal([]byte(text), &config); err != nil {
return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象")
}
lookbackDur, lookbackText, err := resolveIPGroupAutoLookback(object)
if err != nil {
return ipGroupAutoConfig{}, err
}
config.Lookback = lookbackText
config.lookbackDuration = lookbackDur
if config.TTL == 0 {
config.TTL = -1
}
if config.Rules == nil {
config.Rules = []ipGroupAutoRule{}
}
for i, rule := range config.Rules {
rule.Name = strings.TrimSpace(rule.Name)
rule.Expr = strings.TrimSpace(rule.Expr)
if rule.Expr == "" {
return ipGroupAutoConfig{}, fmt.Errorf("自动规则 %d 的 Expr 表达式不能为空", i+1)
}
if _, err := exprlang.Compile(rule.Expr, exprlang.Env(ipGroupAutoRuleEnv{}), exprlang.AsBool()); err != nil {
return ipGroupAutoConfig{}, fmt.Errorf("自动规则 %s Expr 无效: %w", displayIPGroupAutoRuleName(rule, i), err)
}
config.Rules[i] = rule
}
return config, nil
}
// resolveIPGroupAutoLookback accepts lookback as duration string (60m/1h) or legacy lookback_minutes number.
func resolveIPGroupAutoLookback(object map[string]any) (time.Duration, string, error) {
if raw, ok := object["lookback"]; ok && raw != nil {
dur, text, err := parseIPGroupLookbackValue(raw)
if err != nil {
return 0, "", err
}
return dur, text, nil
}
if raw, ok := object["lookback_minutes"]; ok && raw != nil {
// legacy: minutes as number or numeric string
minutes, err := parsePositiveNumber(raw)
if err != nil {
return 0, "", fmt.Errorf("lookback_minutes 无效: %w", err)
}
if minutes <= 0 {
return defaultWAFIPGroupAutoLookbackDur, defaultWAFIPGroupAutoLookback, nil
}
dur := time.Duration(minutes) * time.Minute
if dur > maxWAFIPGroupAutoLookback {
return 0, "", fmt.Errorf("回看窗口不能超过 %s", formatLookbackDuration(maxWAFIPGroupAutoLookback))
}
return dur, formatLookbackDuration(dur), nil
}
return defaultWAFIPGroupAutoLookbackDur, defaultWAFIPGroupAutoLookback, nil
}
func parseIPGroupLookbackValue(raw any) (time.Duration, string, error) {
switch v := raw.(type) {
case string:
trimmed := strings.TrimSpace(v)
if trimmed == "" {
return defaultWAFIPGroupAutoLookbackDur, defaultWAFIPGroupAutoLookback, nil
}
// bare integer string → minutes
if isAllDigits(trimmed) {
minutes, err := parsePositiveNumber(trimmed)
if err != nil || minutes <= 0 {
return 0, "", errors.New("lookback 格式不合法,请使用 60m、1h 等时长")
}
dur := time.Duration(minutes) * time.Minute
if dur > maxWAFIPGroupAutoLookback {
return 0, "", fmt.Errorf("回看窗口不能超过 %s", formatLookbackDuration(maxWAFIPGroupAutoLookback))
}
return dur, formatLookbackDuration(dur), nil
}
dur, err := time.ParseDuration(strings.ToLower(trimmed))
if err != nil || dur <= 0 {
return 0, "", errors.New("lookback 格式不合法,请使用 60m、1h 等时长")
}
if dur > maxWAFIPGroupAutoLookback {
return 0, "", fmt.Errorf("回看窗口不能超过 %s", formatLookbackDuration(maxWAFIPGroupAutoLookback))
}
return dur, formatLookbackDuration(dur), nil
case float64:
if v <= 0 {
return defaultWAFIPGroupAutoLookbackDur, defaultWAFIPGroupAutoLookback, nil
}
if v != float64(int64(v)) {
return 0, "", errors.New("lookback 数值必须为整数分钟")
}
dur := time.Duration(int64(v)) * time.Minute
if dur > maxWAFIPGroupAutoLookback {
return 0, "", fmt.Errorf("回看窗口不能超过 %s", formatLookbackDuration(maxWAFIPGroupAutoLookback))
}
return dur, formatLookbackDuration(dur), nil
case json.Number:
return parseIPGroupLookbackValue(string(v))
default:
return 0, "", errors.New("lookback 格式不合法,请使用 60m、1h 等时长")
}
}
func parsePositiveNumber(raw any) (int, error) {
switch v := raw.(type) {
case float64:
if v != float64(int(v)) {
return 0, errors.New("必须为整数")
}
return int(v), nil
case int:
return v, nil
case int64:
return int(v), nil
case json.Number:
i, err := v.Int64()
if err != nil {
return 0, err
}
return int(i), nil
case string:
trimmed := strings.TrimSpace(v)
if trimmed == "" || !isAllDigits(trimmed) {
return 0, errors.New("必须为整数")
}
n := 0
for _, ch := range trimmed {
n = n*10 + int(ch-'0')
}
return n, nil
default:
return 0, errors.New("必须为整数")
}
}
func isAllDigits(s string) bool {
if s == "" {
return false
}
for _, ch := range s {
if ch < '0' || ch > '9' {
return false
}
}
return true
}
func formatLookbackDuration(d time.Duration) string {
if d <= 0 {
return defaultWAFIPGroupAutoLookback
}
// Prefer compact human units used in config examples.
if d%time.Hour == 0 {
return fmt.Sprintf("%dh", int(d/time.Hour))
}
if d%time.Minute == 0 {
return fmt.Sprintf("%dm", int(d/time.Minute))
}
if d%time.Second == 0 {
return fmt.Sprintf("%ds", int(d/time.Second))
}
return d.String()
}
func validateSubscriptionURL(rawURL string) error {
parsed, err := url.Parse(strings.TrimSpace(rawURL))
if err != nil || parsed.Host == "" {
return errors.New("订阅 URL 无效")
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
return errors.New("订阅 URL 仅支持 http 或 https")
}
return nil
}
func normalizeIPGroupType(value string) string {
switch strings.TrimSpace(value) {
case wafIPGroupTypeManual, "":
return wafIPGroupTypeManual
case wafIPGroupTypeAutomatic:
return wafIPGroupTypeAutomatic
case wafIPGroupTypeSubscription:
return wafIPGroupTypeSubscription
default:
return ""
}
}
func normalizeIPGroupSubscriptionFormat(value string) string {
switch strings.TrimSpace(value) {
case wafIPGroupSubscriptionFormatJSON:
return wafIPGroupSubscriptionFormatJSON
default:
return wafIPGroupSubscriptionFormatText
}
}
func normalizeIPGroupSyncInterval(value int) int {
if value <= 0 {
return defaultWAFIPGroupSyncIntervalMinutes
}
if value < minWAFIPGroupSyncIntervalMinutes {
return minWAFIPGroupSyncIntervalMinutes
}
if value > maxWAFIPGroupSyncIntervalMinutes {
return maxWAFIPGroupSyncIntervalMinutes
}
return value
}
func nextIPGroupSyncAt(groupType string, enabled bool, interval int, current *time.Time) *time.Time {
if (groupType != wafIPGroupTypeSubscription && groupType != wafIPGroupTypeAutomatic) || !enabled {
return nil
}
if current != nil && current.After(time.Now().UTC()) {
return current
}
next := time.Now().UTC().Add(time.Duration(normalizeIPGroupSyncInterval(interval)) * time.Minute)
return &next
}
func buildIPGroup(group *model.OpenFlareWAFIPGroup, input IPGroupInput) (*model.OpenFlareWAFIPGroup, error) {
name := strings.TrimSpace(input.Name)
if name == "" {
return nil, errors.New("IP 组名称不能为空")
}
groupType := normalizeIPGroupType(input.Type)
if groupType == "" {
return nil, errors.New("IP 组类型无效")
}
ipList := input.IPList
subscriptionURL := ""
subscriptionFormat := normalizeIPGroupSubscriptionFormat(input.SubscriptionFormat)
mappingRule := strings.TrimSpace(input.SubscriptionMappingRule)
syncInterval := normalizeIPGroupSyncInterval(input.SyncIntervalMinutes)
autoConfig := "{}"
switch groupType {
case wafIPGroupTypeManual:
subscriptionFormat = wafIPGroupSubscriptionFormatText
mappingRule = ""
case wafIPGroupTypeAutomatic:
normalizedConfig, err := normalizeIPGroupAutoConfig(input.AutoConfig)
if err != nil {
return nil, err
}
autoConfig = normalizedConfig
subscriptionFormat = wafIPGroupSubscriptionFormatText
mappingRule = ""
case wafIPGroupTypeSubscription:
subscriptionURL = strings.TrimSpace(input.SubscriptionURL)
if err := validateSubscriptionURL(subscriptionURL); err != nil {
return nil, err
}
if subscriptionFormat == "" {
subscriptionFormat = wafIPGroupSubscriptionFormatText
}
}
normalizedIPs, err := normalizeIPList(ipList)
if err != nil {
return nil, err
}
ipListJSON, _ := json.Marshal(normalizedIPs)
if group == nil {
group = &model.OpenFlareWAFIPGroup{}
group.ExtIPs = "[]"
}
group.Name = name
group.Type = groupType
group.Enabled = input.Enabled
group.IPList = string(ipListJSON)
if groupType == wafIPGroupTypeAutomatic {
if err := pruneIPGroupExtIPs(group, normalizedIPs); err != nil {
return nil, err
}
}
group.AutoConfig = autoConfig
group.SubscriptionURL = subscriptionURL
group.SubscriptionFormat = subscriptionFormat
group.SubscriptionMappingRule = mappingRule
group.SyncIntervalMinutes = syncInterval
group.NextSyncAt = nextIPGroupSyncAt(group.Type, group.Enabled, syncInterval, group.NextSyncAt)
return group, nil
}
func buildIPGroupView(group *model.OpenFlareWAFIPGroup, referenceCount int) (IPGroupView, error) {
if group == nil {
return IPGroupView{}, errors.New("waf ip group is nil")
}
ips, err := decodeStringList(group.IPList)
if err != nil {
return IPGroupView{}, err
}
autoConfig := json.RawMessage(strings.TrimSpace(group.AutoConfig))
if len(autoConfig) == 0 {
autoConfig = json.RawMessage("{}")
}
var extIPs []ipGroupExtIP
if group.ExtIPs != "" && group.ExtIPs != "[]" {
_ = json.Unmarshal([]byte(group.ExtIPs), &extIPs)
}
viewExtIPs := make([]IPGroupExtIPView, 0, len(extIPs))
for _, extIP := range extIPs {
viewExtIPs = append(viewExtIPs, IPGroupExtIPView{
IP: extIP.IP,
CapturedAt: extIP.CapturedAt.Format(time.RFC3339),
})
}
view := IPGroupView{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ips,
AutoConfig: autoConfig,
ExtIPs: viewExtIPs,
SubscriptionURL: group.SubscriptionURL,
SubscriptionFormat: group.SubscriptionFormat,
SubscriptionMappingRule: group.SubscriptionMappingRule,
SyncIntervalMinutes: group.SyncIntervalMinutes,
LastSyncStatus: group.LastSyncStatus,
LastSyncMessage: group.LastSyncMessage,
ReferencedByRuleCount: referenceCount,
CreatedAt: group.CreatedAt.Format(time.RFC3339),
UpdatedAt: group.UpdatedAt.Format(time.RFC3339),
}
if group.LastSyncedAt != nil {
view.LastSyncedAt = group.LastSyncedAt.Format(time.RFC3339)
}
if group.NextSyncAt != nil {
view.NextSyncAt = group.NextSyncAt.Format(time.RFC3339)
}
return view, nil
}
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"context"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupWAFTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFIPGroup{},
&model.OpenFlareWAFRuleGroupBinding{},
&model.OriginProxyRoute{},
))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func TestPruneIPGroupExtIPs(t *testing.T) {
group := &model.OpenFlareWAFIPGroup{
ExtIPs: `[{"ip":"203.0.113.10","captured_at":"2026-06-18T10:00:00Z"},{"ip":"203.0.113.11","captured_at":"2026-06-18T11:00:00Z"}]`,
}
err := pruneIPGroupExtIPs(group, []string{"203.0.113.10"})
require.NoError(t, err)
assert.JSONEq(t, `[{"ip":"203.0.113.10","captured_at":"2026-06-18T10:00:00Z"}]`, group.ExtIPs)
}
func TestUpdateIPGroupPrunesAutomaticExtIPs(t *testing.T) {
cleanup := setupWAFTestDB(t)
defer cleanup()
ctx := context.Background()
created, err := CreateIPGroup(ctx, IPGroupInput{
Name: "auto group",
Type: wafIPGroupTypeAutomatic,
Enabled: true,
AutoConfig: []byte(`{"lookback":"60m","ttl":-1,"rules":[{"name":"scan","expr":"request_count > 1"}]}`),
})
require.NoError(t, err)
group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, created.ID)
require.NoError(t, err)
group.IPList = `["203.0.113.10","203.0.113.11"]`
group.ExtIPs = `[{"ip":"203.0.113.10","captured_at":"2026-06-18T10:00:00Z"},{"ip":"203.0.113.11","captured_at":"2026-06-18T11:00:00Z"}]`
require.NoError(t, repository.UpdateOpenFlareWAFIPGroup(ctx, group))
updated, err := UpdateIPGroup(ctx, created.ID, IPGroupInput{
Name: created.Name,
Type: created.Type,
Enabled: created.Enabled,
IPList: []string{"203.0.113.10"},
AutoConfig: created.AutoConfig,
})
require.NoError(t, err)
require.Len(t, updated.IPList, 1)
assert.Equal(t, "203.0.113.10", updated.IPList[0])
require.Len(t, updated.ExtIPs, 1)
assert.Equal(t, "203.0.113.10", updated.ExtIPs[0].IP)
}
@@ -0,0 +1,267 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"net/http"
"strconv"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
func routeIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("route_id")
if raw == "" {
response.AbortBadRequest(c, "invalid id")
return 0, false
}
id64, err := strconv.ParseUint(raw, 10, 64)
if err != nil || id64 == 0 {
response.AbortBadRequest(c, "invalid id")
return 0, false
}
return uint(id64), true
}
// GetSiteRuleGroupsHandler 获取站点的 WAF 规则组绑定。
// @Summary 获取站点 WAF 规则组
// @Description 返回代理站点关联的 WAF 规则组绑定,需要管理员权限
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Param route_id path int true "代理路由 ID"
// @Success 200 {object} response.Any{data=waf.SiteRuleGroupsView} "站点规则组绑定"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/sites/{route_id}/rule-groups [get]
func GetSiteRuleGroupsHandler(c *gin.Context) {
routeID, ok := routeIDParam(c)
if !ok {
return
}
view, err := GetSiteRuleGroups(c.Request.Context(), routeID)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// ReplaceSiteRuleGroupsHandler 替换站点的 WAF 规则组绑定。
// @Summary 替换站点 WAF 规则组
// @Description 替换代理站点关联的 WAF 规则组列表,需要管理员权限
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param route_id path int true "代理路由 ID"
// @Param request body waf.IDsRequest true "规则组 ID 列表"
// @Success 200 {object} response.Any{data=waf.SiteRuleGroupsView} "更新后的站点规则组绑定"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/sites/{route_id}/rule-groups [post]
func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
routeID, ok := routeIDParam(c)
if !ok {
return
}
var request IDsRequest
if !apiutil.BindJSON(c, &request) {
return
}
view, err := ReplaceSiteRuleGroups(c.Request.Context(), routeID, request.IDs)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
}
// ListIPGroupsHandler 列出全部 WAF IP 组。
// @Summary 列出 WAF IP 组
// @Description 返回全部 WAF IP 组,需要管理员权限
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]waf.IPGroupView} "IP 组列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/ip-groups [get]
func ListIPGroupsHandler(c *gin.Context) {
groups, err := ListIPGroups(c.Request.Context())
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(groups))
}
// GetIPGroupHandler 获取 WAF IP 组详情。
// @Summary 获取 WAF IP 组详情
// @Description 按 ID 返回 WAF IP 组详情,需要管理员权限
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Param id path int true "IP 组 ID"
// @Success 200 {object} response.Any{data=waf.IPGroupView} "IP 组详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/ip-groups/{id} [get]
func GetIPGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
group, err := GetIPGroup(c.Request.Context(), id)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
}
// CreateIPGroupHandler 创建 WAF IP 组。
// @Summary 创建 WAF IP 组
// @Description 创建新的 WAF IP 组,需要管理员权限
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body waf.IPGroupInput true "IP 组参数"
// @Success 200 {object} response.Any{data=waf.IPGroupView} "创建成功的 IP 组"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/ip-groups [post]
func CreateIPGroupHandler(c *gin.Context) {
var input IPGroupInput
if !apiutil.BindJSON(c, &input) {
return
}
group, err := CreateIPGroup(c.Request.Context(), input)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
}
// UpdateIPGroupHandler 更新 WAF IP 组。
// @Summary 更新 WAF IP 组
// @Description 按 ID 更新 WAF IP 组,需要管理员权限
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "IP 组 ID"
// @Param request body waf.IPGroupInput true "IP 组参数"
// @Success 200 {object} response.Any{data=waf.IPGroupView} "更新后的 IP 组"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/ip-groups/{id}/update [post]
func UpdateIPGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input IPGroupInput
if !apiutil.BindJSON(c, &input) {
return
}
group, err := UpdateIPGroup(c.Request.Context(), id, input)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
}
// DeleteIPGroupHandler 删除 WAF IP 组。
// @Summary 删除 WAF IP 组
// @Description 按 ID 删除 WAF IP 组,需要管理员权限
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Param id path int true "IP 组 ID"
// @Success 200 {object} response.Any "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/ip-groups/{id}/delete [post]
func DeleteIPGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if err := DeleteIPGroup(c.Request.Context(), id); handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// SyncIPGroupHandler 触发 WAF IP 组同步。
// @Summary 同步 WAF IP 组
// @Description 手动触发 WAF IP 组外部 IP 同步,需要管理员权限
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Param id path int true "IP 组 ID"
// @Success 200 {object} response.Any{data=waf.IPGroupSyncResult} "同步结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/ip-groups/{id}/sync [post]
func SyncIPGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
result, err := SyncIPGroup(c.Request.Context(), id)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// TestIPGroupAutoConfigHandler 测试 WAF IP 组自动配置。
// @Summary 测试 WAF IP 组自动配置
// @Description 根据自动配置规则测试 IP 匹配结果(桩实现),需要管理员权限
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body waf.IPGroupAutoTestInput true "自动配置参数"
// @Success 200 {object} response.Any{data=waf.IPGroupAutoTestResult} "测试结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/ip-groups/test [post]
func TestIPGroupAutoConfigHandler(c *gin.Context) {
var input IPGroupAutoTestInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := TestIPGroupAutoConfig(c.Request.Context(), input)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
@@ -0,0 +1,188 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"gorm.io/gorm"
)
// CreateRuleInput is the minimal payload used to create an orchestrated rule.
type CreateRuleInput struct {
Name string `json:"name"`
}
// SaveRuleGraphInput atomically replaces a rule graph at the supplied revision.
type SaveRuleGraphInput struct {
Revision uint64 `json:"revision"`
Graph RuleGraph `json:"graph"`
}
// UpdateRuleMetaInput updates metadata without replacing the graph.
type UpdateRuleMetaInput struct {
Name string `json:"name"`
Enabled bool `json:"enabled"`
}
// RuleValidationError represents a safe user-facing validation failure.
type RuleValidationError struct{ Err error }
func (err *RuleValidationError) Error() string { return err.Err.Error() }
func (err *RuleValidationError) Unwrap() error { return err.Err }
// RuleView is the API representation of an orchestrated WAF rule.
type RuleView struct {
ID uint `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"`
Graph RuleGraph `json:"graph"`
Revision uint64 `json:"revision"`
AppliedSiteIDs []uint `json:"applied_site_ids"`
AppliedSiteCount int `json:"applied_site_count"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
// ListRules returns all orchestrated WAF rules.
func ListRules(ctx context.Context) ([]RuleView, error) {
if err := EnsureDefaultRuleGroup(ctx); err != nil {
return nil, err
}
groups, err := repository.ListOpenFlareWAFRuleGroups(ctx)
if err != nil {
return nil, err
}
bindings, err := loadRuleGroupBindings(ctx)
if err != nil {
return nil, err
}
views := make([]RuleView, 0, len(groups))
for _, group := range groups {
view, buildErr := buildRuleView(group, bindings[group.ID])
if buildErr != nil {
return nil, buildErr
}
views = append(views, view)
}
return views, nil
}
// GetRule returns one orchestrated WAF rule.
func GetRule(ctx context.Context, id uint) (*RuleView, error) {
group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id)
if err != nil {
return nil, err
}
bindings, err := loadRuleGroupBindings(ctx)
if err != nil {
return nil, err
}
view, err := buildRuleView(group, bindings[group.ID])
return &view, err
}
// CreateRule creates a disabled custom rule with the safe default graph.
func CreateRule(ctx context.Context, input CreateRuleInput) (*RuleView, error) {
name := strings.TrimSpace(input.Name)
if name == "" {
return nil, &RuleValidationError{Err: errors.New("WAF 规则名称不能为空")}
}
raw, err := json.Marshal(DefaultRuleGraph())
if err != nil {
return nil, err
}
group := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: false, IsGlobal: false, Graph: string(raw), Revision: 1}
if err = repository.CreateOpenFlareWAFRuleGroup(ctx, group); err != nil {
return nil, err
}
// GORM applies the model's database default to a false bool on Create, so
// explicitly persist the safe disabled state after the row has an ID.
group.Enabled = false
if err = repository.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil {
return nil, err
}
return GetRule(ctx, group.ID)
}
// UpdateRuleMeta updates rule metadata without touching its graph revision.
func UpdateRuleMeta(ctx context.Context, id uint, input UpdateRuleMetaInput) (*RuleView, error) {
group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id)
if err != nil {
return nil, err
}
name := strings.TrimSpace(input.Name)
if name == "" {
return nil, &RuleValidationError{Err: errors.New("WAF 规则名称不能为空")}
}
group.Name, group.Enabled = name, input.Enabled
if err = repository.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil {
return nil, err
}
return GetRule(ctx, id)
}
// DeleteRuleGroup deletes a non-global orchestrated WAF rule.
func DeleteRuleGroup(ctx context.Context, id uint) error {
group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id)
if err != nil {
return err
}
if group.IsGlobal {
return &RuleValidationError{Err: errors.New("全局 WAF 规则不能删除")}
}
return repository.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, id)
}
// SaveRuleGraph validates and atomically replaces a rule graph.
func SaveRuleGraph(ctx context.Context, id uint, input SaveRuleGraphInput) (*RuleView, error) {
if _, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id); err != nil {
return nil, err
}
if err := ValidateRuleGraph(ctx, input.Graph, ruleIPGroupExists); err != nil {
return nil, &RuleValidationError{Err: fmt.Errorf("规则图无效: %w", err)}
}
raw, err := json.Marshal(input.Graph)
if err != nil {
return nil, err
}
if _, err = repository.UpdateOpenFlareWAFRuleGraph(ctx, id, input.Revision, string(raw)); err != nil {
return nil, err
}
return GetRule(ctx, id)
}
func ruleIPGroupExists(ctx context.Context, id uint) (bool, error) {
_, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
return err == nil, err
}
func buildRuleView(group *model.OpenFlareWAFRuleGroup, appliedSiteIDs []uint) (RuleView, error) {
if group == nil {
return RuleView{}, errors.New("waf rule is nil")
}
graph := DefaultRuleGraph()
if strings.TrimSpace(group.Graph) != "" {
if err := json.Unmarshal([]byte(group.Graph), &graph); err != nil {
return RuleView{}, err
}
}
ids := append([]uint(nil), appliedSiteIDs...)
return RuleView{ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal,
Graph: graph, Revision: group.Revision, AppliedSiteIDs: ids, AppliedSiteCount: len(ids),
CreatedAt: group.CreatedAt.Format(time.RFC3339), UpdatedAt: group.UpdatedAt.Format(time.RFC3339)}, nil
}
@@ -0,0 +1,186 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestDeleteIPGroupRejectsGraphReference(t *testing.T) {
cleanup := setupWAFTestDB(t)
defer cleanup()
ctx := context.Background()
group, err := CreateIPGroup(ctx, IPGroupInput{Name: "trusted", Type: wafIPGroupTypeManual, Enabled: true})
require.NoError(t, err)
rule, err := CreateRule(ctx, CreateRuleInput{Name: "guard"})
require.NoError(t, err)
graph := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
{ID: "start", Type: RuleNodeStart, Config: json.RawMessage(`{}`)},
{ID: "match", Type: RuleNodeIPMatch, Config: json.RawMessage(`{"ip_group_ids":[` + strconv.FormatUint(uint64(group.ID), 10) + `]}`)},
{ID: "allow", Type: RuleNodeAllow, Config: json.RawMessage(`{}`)},
}, Edges: []RuleEdge{
{ID: "e1", Source: "start", SourceHandle: "next", Target: "match"},
{ID: "e2", Source: "match", SourceHandle: "true", Target: "allow"},
{ID: "e3", Source: "match", SourceHandle: "false", Target: "allow"},
}}
_, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: graph})
require.NoError(t, err)
view, err := GetIPGroup(ctx, group.ID)
require.NoError(t, err)
assert.Equal(t, 1, view.ReferencedByRuleCount)
require.ErrorContains(t, DeleteIPGroup(ctx, group.ID), "已被 WAF 规则引用")
}
func TestRuleHandlersMapFailures(t *testing.T) {
gin.SetMode(gin.TestMode)
tests := []struct {
name string
method string
path string
body string
setup func(t *testing.T) func()
want int
}{
{name: "invalid id", method: http.MethodGet, path: "/rules/nope", setup: setupWAFTestDB, want: http.StatusBadRequest},
{name: "malformed json", method: http.MethodPost, path: "/rules", body: `{`, setup: setupWAFTestDB, want: http.StatusBadRequest},
{name: "invalid graph", method: http.MethodPost, path: "/rules/1/graph", body: `{"revision":1,"graph":{"schema_version":1,"nodes":[],"edges":[]}}`, setup: func(t *testing.T) func() {
t.Helper()
cleanup := setupWAFTestDB(t)
_, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
require.NoError(t, err)
return cleanup
}, want: http.StatusBadRequest},
{name: "manual IP group sync", method: http.MethodPost, path: "/ip-groups/1/sync", setup: func(t *testing.T) func() {
t.Helper()
cleanup := setupWAFTestDB(t)
_, err := CreateIPGroup(context.Background(), IPGroupInput{Name: "manual", Type: wafIPGroupTypeManual, Enabled: true})
require.NoError(t, err)
return cleanup
}, want: http.StatusBadRequest},
{name: "missing", method: http.MethodGet, path: "/rules/999", setup: setupWAFTestDB, want: http.StatusNotFound},
{name: "conflict", method: http.MethodPost, path: "/rules/1/graph", body: mustGraphRequest(t, 0), setup: func(t *testing.T) func() {
t.Helper()
cleanup := setupWAFTestDB(t)
_, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
require.NoError(t, err)
return cleanup
}, want: http.StatusConflict},
{name: "database failure", method: http.MethodGet, path: "/rules", setup: func(t *testing.T) func() { t.Helper(); db.SetDB(nil); return func() {} }, want: http.StatusInternalServerError},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cleanup := tt.setup(t)
defer cleanup()
router := gin.New()
router.Use(response.ErrorHandlerMiddleware())
router.GET("/rules", ListRulesHandler)
router.POST("/rules", CreateRuleHandler)
router.GET("/rules/:id", GetRuleHandler)
router.POST("/rules/:id/graph", SaveRuleGraphHandler)
router.POST("/ip-groups/:id/sync", SyncIPGroupHandler)
rec := httptest.NewRecorder()
req := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(rec, req)
assert.Equal(t, tt.want, rec.Code, rec.Body.String())
})
}
}
func mustGraphRequest(t *testing.T, revision uint64) string {
t.Helper()
raw, err := json.Marshal(SaveRuleGraphInput{Revision: revision, Graph: DefaultRuleGraph()})
require.NoError(t, err)
return string(raw)
}
func TestCreateRuleCreatesDefaultGraph(t *testing.T) {
cleanup := setupWAFTestDB(t)
defer cleanup()
rule, err := CreateRule(context.Background(), CreateRuleInput{Name: " edge guard "})
require.NoError(t, err)
assert.Equal(t, "edge guard", rule.Name)
assert.False(t, rule.Enabled)
assert.Equal(t, uint64(1), rule.Revision)
assert.Equal(t, DefaultRuleGraph(), rule.Graph)
}
func TestCreateRuleRejectsEmptyName(t *testing.T) {
cleanup := setupWAFTestDB(t)
defer cleanup()
_, err := CreateRule(context.Background(), CreateRuleInput{Name: " "})
require.ErrorContains(t, err, "名称不能为空")
}
func TestSaveRuleGraphValidationAndRevisionConflict(t *testing.T) {
cleanup := setupWAFTestDB(t)
defer cleanup()
ctx := context.Background()
rule, err := CreateRule(ctx, CreateRuleInput{Name: "guard"})
require.NoError(t, err)
invalid := DefaultRuleGraph()
invalid.Edges = nil
_, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: invalid})
require.Error(t, err)
updated, err := SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: DefaultRuleGraph()})
require.NoError(t, err)
assert.Equal(t, uint64(2), updated.Revision)
_, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: DefaultRuleGraph()})
assert.ErrorIs(t, err, model.ErrWAFRuleRevisionConflict)
}
func TestReplaceSiteRuleGroupsPreservesOrderAndRejectsGlobal(t *testing.T) {
cleanup := setupWAFTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Create(&model.OriginProxyRoute{ID: 7, Domain: "example.com"}).Error)
first, err := CreateRule(ctx, CreateRuleInput{Name: "first"})
require.NoError(t, err)
second, err := CreateRule(ctx, CreateRuleInput{Name: "second"})
require.NoError(t, err)
third, err := CreateRule(ctx, CreateRuleInput{Name: "third"})
require.NoError(t, err)
view, err := ReplaceSiteRuleGroups(ctx, 7, []uint{third.ID, first.ID, second.ID, first.ID})
require.NoError(t, err)
assert.Equal(t, []uint{third.ID, first.ID, second.ID}, view.AppliedIDs)
require.NoError(t, EnsureDefaultRuleGroup(ctx))
global, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
_, err = ReplaceSiteRuleGroups(ctx, 7, []uint{global.ID, second.ID})
require.Error(t, err)
require.NotErrorIs(t, err, model.ErrWAFRuleRevisionConflict)
assert.Equal(t, []uint{third.ID, first.ID, second.ID}, mustListSiteRuleGroupIDs(t, ctx, 7))
}
func mustListSiteRuleGroupIDs(t *testing.T, ctx context.Context, routeID uint) []uint {
t.Helper()
ids, err := ListSiteRuleGroupIDs(ctx, routeID)
require.NoError(t, err)
return ids
}
@@ -0,0 +1,187 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"errors"
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
func handleRuleError(c *gin.Context, err error) bool {
if err == nil {
return false
}
var validation *RuleValidationError
switch {
case errors.As(err, &validation):
response.AbortBadRequest(c, validation.Error())
case errors.Is(err, model.ErrWAFRuleRevisionConflict):
response.AbortConflict(c, "规则已被其他操作更新,请重新加载")
case errors.Is(err, gorm.ErrRecordNotFound):
response.AbortNotFound(c, "WAF 规则不存在")
default:
logger.ErrorF(c.Request.Context(), "[OpenFlareWAF] rule API failed: %v", err)
response.AbortInternal(c, "WAF 规则操作失败")
}
return true
}
// ListRulesHandler lists orchestrated WAF rules.
// @Summary 列出 WAF 规则
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]waf.RuleView} "规则列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/rule-groups [get]
func ListRulesHandler(c *gin.Context) {
rules, err := ListRules(c.Request.Context())
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(rules))
}
// GetRuleHandler gets an orchestrated WAF rule.
// @Summary 获取 WAF 规则详情
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Param id path int true "规则 ID"
// @Success 200 {object} response.Any{data=waf.RuleView} "规则详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/rule-groups/{id} [get]
func GetRuleHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
rule, err := GetRule(c.Request.Context(), id)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(rule))
}
// CreateRuleHandler creates an orchestrated WAF rule from a name only.
// @Summary 创建 WAF 规则
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body waf.CreateRuleInput true "规则名称"
// @Success 200 {object} response.Any{data=waf.RuleView} "创建成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/rule-groups [post]
func CreateRuleHandler(c *gin.Context) {
var input CreateRuleInput
if !apiutil.BindJSON(c, &input) {
return
}
rule, err := CreateRule(c.Request.Context(), input)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(rule))
}
// UpdateRuleMetaHandler updates rule name and enabled state.
// @Summary 更新 WAF 规则元数据
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "规则 ID"
// @Param request body waf.UpdateRuleMetaInput true "规则元数据"
// @Success 200 {object} response.Any{data=waf.RuleView} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/rule-groups/{id}/meta [post]
func UpdateRuleMetaHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input UpdateRuleMetaInput
if !apiutil.BindJSON(c, &input) {
return
}
rule, err := UpdateRuleMeta(c.Request.Context(), id, input)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(rule))
}
// SaveRuleGraphHandler saves a complete versioned rule graph.
// @Summary 保存 WAF 规则图
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "规则 ID"
// @Param request body waf.SaveRuleGraphInput true "规则图和修订号"
// @Success 200 {object} response.Any{data=waf.RuleView} "保存成功"
// @Failure 400 {object} response.Any "参数或规则图错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 409 {object} response.Any "修订冲突"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/rule-groups/{id}/graph [post]
func SaveRuleGraphHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input SaveRuleGraphInput
if !apiutil.BindJSON(c, &input) {
return
}
rule, err := SaveRuleGraph(c.Request.Context(), id, input)
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(rule))
}
// DeleteRuleHandler deletes a non-global WAF rule.
// @Summary 删除 WAF 规则
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Param id path int true "规则 ID"
// @Success 200 {object} response.Any "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/rule-groups/{id}/delete [post]
func DeleteRuleHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if err := DeleteRuleGroup(c.Request.Context(), id); handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}