feat(waf): complete composable rule orchestration

Add the React Flow rule editor, ordered graph APIs and runtime DAG execution.\n\nPublish rules only on OpenResty reload and reconcile checksum-driven IP group snapshots in bounded shared memory.
This commit is contained in:
ryan
2026-07-13 14:16:55 +08:00
parent d36409fbf9
commit a1a997bcda
72 changed files with 5897 additions and 3080 deletions
+22 -5
View File
@@ -5,25 +5,35 @@ 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 RuleNodeType = "start"
RuleNodeAllow RuleNodeType = "allow"
RuleNodeBlock RuleNodeType = "block"
RuleNodeIPMatch RuleNodeType = "ip_match"
// 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 RuleNodeType = "pow"
// RuleNodePoW runs a proof-of-work challenge before continuing.
RuleNodePoW RuleNodeType = "pow"
)
// 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"`
@@ -32,11 +42,13 @@ type RuleNode struct {
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"`
@@ -44,17 +56,20 @@ type RuleEdge struct {
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"`
@@ -62,11 +77,13 @@ type PoWNodeConfig struct {
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"`
}
// 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(`{}`)},
+164 -95
View File
@@ -26,7 +26,33 @@ var (
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)
}
@@ -41,15 +67,18 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
} else if len(raw) > maxRuleGraphBytes {
return fmt.Errorf("规则图大小不能超过 256 KiB")
}
return nil
}
nodes := make(map[string]RuleNode, len(graph.Nodes))
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 graph.Nodes {
for _, node := range graphNodes {
if strings.TrimSpace(node.ID) == "" {
return errors.New("节点 ID 不能为空")
return nil, "", errors.New("节点 ID 不能为空")
}
if _, exists := nodes[node.ID]; exists {
return fmt.Errorf("节点 ID %s 重复", node.ID)
return nil, "", fmt.Errorf("节点 ID %s 重复", node.ID)
}
nodes[node.ID] = node
switch node.Type {
@@ -60,67 +89,74 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
allowCount++
case RuleNodeBlock, RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodePoW:
default:
return fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
return nil, "", fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
}
if err := validateRuleNodeConfig(ctx, node, ipGroupExists); err != nil {
return err
return nil, "", err
}
}
if startCount != 1 {
return errors.New("规则图必须恰好包含一个开始节点")
return nil, "", errors.New("规则图必须恰好包含一个开始节点")
}
if allowCount != 1 {
return errors.New("规则图必须恰好包含一个通过节点")
return nil, "", errors.New("规则图必须恰好包含一个通过节点")
}
return nodes, startID, nil
}
edgeIDs := make(map[string]struct{}, len(graph.Edges))
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 graph.Edges {
for _, edge := range graphEdges {
if strings.TrimSpace(edge.ID) == "" {
return errors.New("边 ID 不能为空")
return nil, nil, nil, errors.New("边 ID 不能为空")
}
if _, exists := edgeIDs[edge.ID]; exists {
return fmt.Errorf("边 ID %s 重复", edge.ID)
return nil, nil, nil, fmt.Errorf("边 ID %s 重复", edge.ID)
}
edgeIDs[edge.ID] = struct{}{}
source, ok := nodes[edge.Source]
if !ok {
return fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
return nil, nil, nil, fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
}
if _, ok := nodes[edge.Target]; !ok {
return fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
return nil, nil, nil, fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
}
if !validSourceHandle(source.Type, edge.SourceHandle) {
return fmt.Errorf("边 %s 的源端口 %s 不适用于节点 %s", edge.ID, edge.SourceHandle, edge.Source)
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 fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
return nil, nil, nil, fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
}
outgoing[edge.Source] = append(outgoing[edge.Source], edge)
incoming[edge.Target]++
}
if hasRuleGraphCycle(nodes, outgoing, incoming) {
return errors.New("规则图不能包含循环")
}
for _, node := range graph.Nodes {
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 graph.Nodes {
for _, node := range nodes {
if !reachable[node.ID] {
return fmt.Errorf("节点 %s 无法从开始节点到达", node.ID)
}
}
for _, node := range graph.Nodes {
for _, node := range nodes {
if node.Type == RuleNodeStart && incoming[node.ID] != 0 {
return fmt.Errorf("开始节点 %s 不能有入边", node.ID)
}
@@ -131,96 +167,129 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
return fmt.Errorf("终止节点 %s 不能有出口", node.ID)
}
}
if err := validateTerminalPaths(graph.Nodes, outgoing); err != nil {
return err
}
return nil
}
func validateRuleNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
switch node.Type {
case RuleNodeStart, RuleNodeAllow:
var cfg struct{}
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
}
return validateEmptyNodeConfig(node)
case RuleNodeIPMatch:
var cfg IPMatchConfig
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
}
for _, raw := range cfg.IPs {
if _, err := netip.ParseAddr(raw); err != nil {
return fmt.Errorf("节点 %s 的 IP %s 无效", node.ID, raw)
}
}
for _, raw := range cfg.CIDRs {
if _, err := netip.ParsePrefix(raw); err != nil {
return fmt.Errorf("节点 %s 的 CIDR %s 无效", node.ID, raw)
}
}
for _, id := range cfg.IPGroupIDs {
if id == 0 {
return fmt.Errorf("节点 %s 引用的 IP 组 ID 无效", node.ID)
}
if exists == nil {
return fmt.Errorf("节点 %s 无法校验 IP 组 %d", node.ID, id)
}
ok, err := exists(ctx, id)
if err != nil {
return fmt.Errorf("节点 %s 校验 IP 组 %d 失败: %w", node.ID, id, err)
}
if !ok {
return fmt.Errorf("节点 %s 引用的 IP 组 %d 不存在", node.ID, id)
}
}
return validateIPMatchNodeConfig(ctx, node, exists)
case RuleNodeGeoMatch:
var cfg GeoMatchConfig
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
}
for _, code := range cfg.Countries {
if !countryCodePattern.MatchString(code) {
return fmt.Errorf("节点 %s 的国家代码 %s 无效", node.ID, code)
}
}
for _, code := range cfg.Regions {
if !regionCodePattern.MatchString(code) {
return fmt.Errorf("节点 %s 的地区代码 %s 无效", node.ID, code)
}
}
return validateGeoMatchNodeConfig(node)
case RuleNodePoW:
var cfg PoWNodeConfig
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
}
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
return fmt.Errorf("节点 %s 的 PoW 难度必须在 1-16 之间", node.ID)
}
if cfg.Algorithm != "fast" && cfg.Algorithm != "slow" {
return fmt.Errorf("节点 %s 的 PoW 算法必须为 fast 或 slow", node.ID)
}
if cfg.SessionTTL < 60 {
return fmt.Errorf("节点 %s 的 PoW 会话 TTL 不能小于 60 秒", node.ID)
}
if cfg.ChallengeTTL < 30 {
return fmt.Errorf("节点 %s 的 PoW 挑战 TTL 不能小于 30 秒", node.ID)
}
return validatePoWNodeConfig(node)
case RuleNodeBlock:
var cfg BlockNodeConfig
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
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)
}
if cfg.StatusCode < 400 || cfg.StatusCode > 599 {
return fmt.Errorf("节点 %s 的阻止状态码必须在 400-599 之间", node.ID)
}
for _, raw := range cfg.CIDRs {
if _, err := netip.ParsePrefix(raw); err != nil {
return fmt.Errorf("节点 %s 的 CIDR %s 无效", node.ID, raw)
}
if len([]byte(cfg.ResponseBody)) > maxWAFBlockBodyBytes {
return fmt.Errorf("节点 %s 的阻止响应体不能超过 %d 字节", node.ID, maxWAFBlockBodyBytes)
}
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 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")) {
+1 -1
View File
@@ -90,7 +90,7 @@ func syncOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGro
case wafIPGroupTypeAutomatic:
return syncIPGroupAutomatic(ctx, group, now)
default:
return nil, errors.New("只有自动和订阅类型 IP 组支持同步")
return nil, &RuleValidationError{Err: errors.New("只有自动和订阅类型 IP 组支持同步")}
}
}
File diff suppressed because it is too large Load Diff
+1 -32
View File
@@ -26,6 +26,7 @@ func setupWAFTestDB(t *testing.T) func() {
&model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFIPGroup{},
&model.OpenFlareWAFRuleGroupBinding{},
&model.OriginProxyRoute{},
))
db.SetDB(sqliteDB)
@@ -34,38 +35,6 @@ func setupWAFTestDB(t *testing.T) func() {
}
}
func TestCreateRuleGroup(t *testing.T) {
cleanup := setupWAFTestDB(t)
defer cleanup()
ctx := context.Background()
group, err := CreateRuleGroup(ctx, RuleGroupInput{
Name: "edge guard",
Enabled: true,
BlockStatusCode: 451,
IPWhitelist: []string{" 192.0.2.1 ", "192.0.2.1", "198.51.100.0/24"},
IPBlacklist: []string{"203.0.113.10"},
CountryBlacklist: []string{" cn ", "CN", "us"},
})
require.NoError(t, err)
assert.NotZero(t, group.ID)
assert.False(t, group.IsGlobal)
assert.Equal(t, "edge guard", group.Name)
require.Len(t, group.IPWhitelist, 2)
assert.Equal(t, "192.0.2.1", group.IPWhitelist[0])
assert.Equal(t, "198.51.100.0/24", group.IPWhitelist[1])
require.Len(t, group.CountryBlacklist, 2)
assert.Equal(t, "CN", group.CountryBlacklist[0])
assert.Equal(t, "US", group.CountryBlacklist[1])
_, err = CreateRuleGroup(ctx, RuleGroupInput{
Name: "bad ip",
Enabled: true,
IPBlacklist: []string{"not-an-ip"},
})
require.Error(t, err)
}
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"}]`,
@@ -1,92 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"encoding/json"
"errors"
"fmt"
"net"
"regexp"
"strings"
)
func parsePoWConfigRaw(enabled bool, raw string) (PoWConfig, error) {
if !enabled {
return defaultPoWConfig(), nil
}
cfg := defaultPoWConfig()
text := strings.TrimSpace(raw)
if text == "" || text == "{}" {
return cfg, nil
}
if err := json.Unmarshal([]byte(text), &cfg); err != nil {
return cfg, errors.New("pow_config 格式无效")
}
return cfg, nil
}
func validatePoWCoreSettings(cfg PoWConfig) error {
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
return errors.New("pow_config.difficulty 必须在 1-16 之间")
}
if !powAlgorithmValues[cfg.Algorithm] {
return errors.New("pow_config.algorithm 必须为 fast 或 slow")
}
if cfg.SessionTTL < minPoWSessionTTLSeconds {
return errors.New("pow_config.session_ttl 不能小于 60 秒")
}
if cfg.ChallengeTTL < minPoWChallengeTTLSeconds {
return errors.New("pow_config.challenge_ttl 不能小于 30 秒")
}
return nil
}
func validatePoWCIDRs(cidrs []string, listName string) error {
for _, cidr := range cidrs {
if _, _, err := net.ParseCIDR(cidr); err != nil {
return fmt.Errorf("pow_config %s IP CIDR 格式无效: %s", listName, cidr)
}
}
return nil
}
func validatePoWPathRegexes(regexes []string, listName string) error {
for _, re := range regexes {
if _, err := regexp.Compile(re); err != nil {
return fmt.Errorf("pow_config %s路径正则格式无效: %s", listName, re)
}
}
return nil
}
func validatePoWIPs(ips []string, listName string) error {
for _, ip := range ips {
if net.ParseIP(ip) == nil {
return fmt.Errorf("pow_config %s IP 格式无效: %s", listName, ip)
}
}
return nil
}
func validatePoWListMutualExclusion(cfg PoWConfig) error {
type dimension struct {
name string
wl []string
bl []string
}
dimensions := []dimension{
{"IP", cfg.Whitelist.IPs, cfg.Blacklist.IPs},
{"IP CIDR", cfg.Whitelist.IPCidrs, cfg.Blacklist.IPCidrs},
{"路径", cfg.Whitelist.Paths, cfg.Blacklist.Paths},
{"路径正则", cfg.Whitelist.PathRegexes, cfg.Blacklist.PathRegexes},
{"User-Agent", cfg.Whitelist.UserAgents, cfg.Blacklist.UserAgents},
}
for _, dim := range dimensions {
if len(dim.wl) > 0 && len(dim.bl) > 0 {
return fmt.Errorf("pow_config %s 不能同时配置白名单和黑名单", dim.name)
}
}
return nil
}
+10 -179
View File
@@ -12,14 +12,6 @@ import (
"github.com/gin-gonic/gin"
)
func handleLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
return apiutil.AbortNotFoundIfMissing(c, err, "记录不存在")
}
func routeIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("route_id")
if raw == "" {
@@ -34,167 +26,6 @@ func routeIDParam(c *gin.Context) (uint, bool) {
return uint(id64), true
}
// ListRuleGroupsHandler 列出全部 WAF 规则组。
// @Summary 列出 WAF 规则组
// @Description 返回全部 WAF 规则组,需要管理员权限
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]waf.RuleGroupView} "规则组列表"
// @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 [get]
func ListRuleGroupsHandler(c *gin.Context) {
groups, err := ListRuleGroups(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(groups))
}
// GetRuleGroupHandler 获取 WAF 规则组详情。
// @Summary 获取 WAF 规则组详情
// @Description 按 ID 返回 WAF 规则组详情,需要管理员权限
// @Tags openflare-waf
// @Produce json
// @Security SessionCookie
// @Param id path int true "规则组 ID"
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "规则组详情"
// @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/rule-groups/{id} [get]
func GetRuleGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
group, err := GetRuleGroup(c.Request.Context(), id)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
}
// CreateRuleGroupHandler 创建 WAF 规则组。
// @Summary 创建 WAF 规则组
// @Description 创建新的 WAF 规则组,需要管理员权限
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body waf.RuleGroupInput true "规则组参数"
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "创建成功的规则组"
// @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 CreateRuleGroupHandler(c *gin.Context) {
var input RuleGroupInput
if !apiutil.BindJSON(c, &input) {
return
}
group, err := CreateRuleGroup(c.Request.Context(), input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
}
// UpdateRuleGroupHandler 更新 WAF 规则组。
// @Summary 更新 WAF 规则组
// @Description 按 ID 更新 WAF 规则组,需要管理员权限
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "规则组 ID"
// @Param request body waf.RuleGroupInput true "规则组参数"
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "更新后的规则组"
// @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/rule-groups/{id}/update [post]
func UpdateRuleGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input RuleGroupInput
if !apiutil.BindJSON(c, &input) {
return
}
group, err := UpdateRuleGroup(c.Request.Context(), id, input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
}
// DeleteRuleGroupHandler 删除 WAF 规则组。
// @Summary 删除 WAF 规则组
// @Description 按 ID 删除 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 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/waf/rule-groups/{id}/delete [post]
func DeleteRuleGroupHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if err := DeleteRuleGroup(c.Request.Context(), id); handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ReplaceRuleGroupSitesHandler 替换规则组绑定的站点。
// @Summary 替换规则组站点绑定
// @Description 替换 WAF 规则组关联的代理站点列表,需要管理员权限
// @Tags openflare-waf
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "规则组 ID"
// @Param request body waf.IDsRequest true "站点 ID 列表"
// @Success 200 {object} response.Any{data=waf.RuleGroupView} "更新后的规则组"
// @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/rule-groups/{id}/sites [post]
func ReplaceRuleGroupSitesHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var request IDsRequest
if !apiutil.BindJSON(c, &request) {
return
}
group, err := ReplaceRuleGroupSites(c.Request.Context(), id, request.IDs)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
}
// GetSiteRuleGroupsHandler 获取站点的 WAF 规则组绑定。
// @Summary 获取站点 WAF 规则组
// @Description 返回代理站点关联的 WAF 规则组绑定,需要管理员权限
@@ -215,7 +46,7 @@ func GetSiteRuleGroupsHandler(c *gin.Context) {
return
}
view, err := GetSiteRuleGroups(c.Request.Context(), routeID)
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
@@ -247,7 +78,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
return
}
view, err := ReplaceSiteRuleGroups(c.Request.Context(), routeID, request.IDs)
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
@@ -267,7 +98,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
// @Router /api/v1/d/waf/ip-groups [get]
func ListIPGroupsHandler(c *gin.Context) {
groups, err := ListIPGroups(c.Request.Context())
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(groups))
@@ -293,7 +124,7 @@ func GetIPGroupHandler(c *gin.Context) {
return
}
group, err := GetIPGroup(c.Request.Context(), id)
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
@@ -319,7 +150,7 @@ func CreateIPGroupHandler(c *gin.Context) {
return
}
group, err := CreateIPGroup(c.Request.Context(), input)
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
@@ -351,7 +182,7 @@ func UpdateIPGroupHandler(c *gin.Context) {
return
}
group, err := UpdateIPGroup(c.Request.Context(), id, input)
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
@@ -376,7 +207,7 @@ func DeleteIPGroupHandler(c *gin.Context) {
if !ok {
return
}
if err := DeleteIPGroup(c.Request.Context(), id); handleLogicError(c, err) {
if err := DeleteIPGroup(c.Request.Context(), id); handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -402,7 +233,7 @@ func SyncIPGroupHandler(c *gin.Context) {
return
}
result, err := SyncIPGroup(c.Request.Context(), id)
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
@@ -428,8 +259,8 @@ func TestIPGroupAutoConfigHandler(c *gin.Context) {
return
}
result, err := TestIPGroupAutoConfig(c.Request.Context(), input)
if handleLogicError(c, err) {
if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
}
+185
View File
@@ -0,0 +1,185 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/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 := model.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 := model.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 = model.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 = model.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 := model.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 = model.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 := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
if err != nil {
return err
}
if group.IsGlobal {
return &RuleValidationError{Err: errors.New("全局 WAF 规则不能删除")}
}
return model.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, id)
}
// SaveRuleGraph validates and atomically replaces a rule graph.
func SaveRuleGraph(ctx context.Context, id uint, input SaveRuleGraphInput) (*RuleView, error) {
if _, err := model.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 = model.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 := model.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,181 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"strconv"
"testing"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"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() {
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() {
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() {
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() { 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 := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
_, err = ReplaceSiteRuleGroups(ctx, 7, []uint{global.ID, second.ID})
require.Error(t, err)
assert.False(t, errors.Is(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
}
+186
View File
@@ -0,0 +1,186 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"errors"
"net/http"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"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())
}