mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-10 17:26:38 +08:00
refactor(backend): rename OpenFlare directory to lowercase openflare
This commit is contained in:
@@ -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())
|
||||
}
|
||||
Reference in New Issue
Block a user