feat(waf): add composable rule graph core

This commit is contained in:
ryan
2026-07-13 11:57:36 +08:00
parent 2fff30e188
commit af20e2e838
11 changed files with 974 additions and 4 deletions
+30 -4
View File
@@ -30,6 +30,8 @@ type OpenFlareWAFRuleGroup struct {
RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"`
PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"`
Graph string `json:"graph" gorm:"type:text;not null;default:''"`
Revision uint64 `json:"revision" gorm:"not null;default:1"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
@@ -70,9 +72,13 @@ type OpenFlareWAFRuleGroupBinding struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route"`
ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route;index"`
Sequence int `json:"sequence" gorm:"not null;default:0"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// ErrWAFRuleRevisionConflict indicates that a rule graph was updated from a stale revision.
var ErrWAFRuleRevisionConflict = errors.New("waf rule revision conflict")
// TableName returns the GORM table name.
func (OpenFlareWAFRuleGroupBinding) TableName() string {
return "of_waf_rule_group_bindings"
@@ -159,6 +165,24 @@ func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGro
}).Error
}
// UpdateOpenFlareWAFRuleGraph atomically replaces a graph when revision is current.
func UpdateOpenFlareWAFRuleGraph(ctx context.Context, id uint, revision uint64, graph string) (uint64, error) {
conn, err := wafDB(ctx)
if err != nil {
return 0, err
}
result := conn.Model(&OpenFlareWAFRuleGroup{}).
Where("id = ? AND revision = ?", id, revision).
Updates(map[string]any{"graph": graph, "revision": gorm.Expr("revision + 1")})
if result.Error != nil {
return 0, result.Error
}
if result.RowsAffected != 1 {
return 0, ErrWAFRuleRevisionConflict
}
return revision + 1, nil
}
// DeleteOpenFlareWAFRuleGroup removes a rule group.
func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error {
conn, err := wafDB(ctx)
@@ -289,7 +313,7 @@ func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]OpenFlareWAFRuleG
return nil, err
}
var bindings []OpenFlareWAFRuleGroupBinding
if err = conn.Order("rule_group_id asc").Order("proxy_route_id asc").Find(&bindings).Error; err != nil {
if err = conn.Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil {
return nil, err
}
return bindings, nil
@@ -302,7 +326,7 @@ func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uin
return nil, err
}
var bindings []OpenFlareWAFRuleGroupBinding
if err = conn.Where("proxy_route_id = ?", routeID).Order("rule_group_id asc").Find(&bindings).Error; err != nil {
if err = conn.Where("proxy_route_id = ?", routeID).Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil {
return nil, err
}
return bindings, nil
@@ -342,10 +366,11 @@ func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, rou
return err
}
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(routeIDs))
for _, routeID := range routeIDs {
for index, routeID := range routeIDs {
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
RuleGroupID: groupID,
ProxyRouteID: routeID,
Sequence: index,
})
}
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
@@ -363,10 +388,11 @@ func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint,
return err
}
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(groupIDs))
for _, groupID := range groupIDs {
for index, groupID := range groupIDs {
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
RuleGroupID: groupID,
ProxyRouteID: routeID,
Sequence: index,
})
}
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
+107
View File
@@ -0,0 +1,107 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"io/fs"
"os"
"path/filepath"
"runtime"
"testing"
"testing/fstest"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/glebarez/sqlite"
"github.com/pressly/goose/v3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
const defaultWAFRuleGraph = `{"schema_version":1,"nodes":[{"id":"start","type":"start","position":{"x":0,"y":0},"config":{}},{"id":"allow","type":"allow","position":{"x":320,"y":0},"config":{}}],"edges":[{"id":"start-allow","source":"start","source_handle":"next","target":"allow"}]}`
func wafMigrationFS(t *testing.T) fs.FS {
t.Helper()
_, filename, _, ok := runtime.Caller(0)
require.True(t, ok)
dir := filepath.Join(filepath.Dir(filename), "..", "db", "migrator", "goose", "sqlite")
migrations := fstest.MapFS{}
for _, name := range []string{
"202607150001_orchestrate_waf_rules.sql",
"202607150002_reset_waf_rule_graphs.sql",
} {
contents, err := os.ReadFile(filepath.Join(dir, name))
require.NoError(t, err)
migrations[name] = &fstest.MapFile{Data: contents}
}
return migrations
}
func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) {
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, err := conn.DB()
require.NoError(t, err)
require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_groups (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL)`).Error)
require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_group_bindings (id INTEGER PRIMARY KEY AUTOINCREMENT, rule_group_id INTEGER NOT NULL, proxy_route_id INTEGER NOT NULL)`).Error)
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (id, name) VALUES (1, 'one'), (2, 'two')`).Error)
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_group_bindings (id, rule_group_id, proxy_route_id) VALUES (20, 2, 7), (10, 1, 7)`).Error)
goose.SetBaseFS(wafMigrationFS(t))
require.NoError(t, goose.SetDialect("sqlite3"))
require.NoError(t, goose.Up(sqlDB, "."))
var groups []OpenFlareWAFRuleGroup
require.NoError(t, conn.Order("id asc").Find(&groups).Error)
require.Len(t, groups, 2)
for _, group := range groups {
require.JSONEq(t, defaultWAFRuleGraph, group.Graph)
assert.Equal(t, uint64(1), group.Revision)
}
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (name) VALUES ('new')`).Error)
var newGroup OpenFlareWAFRuleGroup
require.NoError(t, conn.First(&newGroup, 3).Error)
assert.Empty(t, newGroup.Graph)
assert.Equal(t, uint64(1), newGroup.Revision)
var bindings []OpenFlareWAFRuleGroupBinding
require.NoError(t, conn.Where("proxy_route_id = ?", 7).Order("sequence asc").Order("id asc").Find(&bindings).Error)
require.Len(t, bindings, 2)
assert.Equal(t, []int{0, 1}, []int{bindings[0].Sequence, bindings[1].Sequence})
assert.Equal(t, []uint{10, 20}, []uint{bindings[0].ID, bindings[1].ID})
}
func TestOpenFlareWAFGraphOptimisticUpdate(t *testing.T) {
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, conn.AutoMigrate(&OpenFlareWAFRuleGroup{}))
db.SetDB(conn)
t.Cleanup(func() { db.SetDB(nil) })
group := OpenFlareWAFRuleGroup{Name: "rule", Graph: defaultWAFRuleGraph, Revision: 1}
require.NoError(t, conn.Create(&group).Error)
nextRevision, err := UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, `{"schema_version":1}`)
require.NoError(t, err)
assert.Equal(t, uint64(2), nextRevision)
_, err = UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, defaultWAFRuleGraph)
assert.ErrorIs(t, err, ErrWAFRuleRevisionConflict)
}
func TestReplaceOpenFlareWAFRuleGroupBindingsPreservesInputOrder(t *testing.T) {
cleanup := setupWAFBindingsTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, 7, []uint{30, 10, 20}))
bindings, err := ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx, 7)
require.NoError(t, err)
require.Len(t, bindings, 3)
assert.Equal(t, []uint{30, 10, 20}, []uint{bindings[0].RuleGroupID, bindings[1].RuleGroupID, bindings[2].RuleGroupID})
assert.Equal(t, []int{0, 1, 2}, []int{bindings[0].Sequence, bindings[1].Sequence, bindings[2].Sequence})
}