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
+15
View File
@@ -24,6 +24,7 @@ const (
defaultRuntimeConfigDirRelativePath = "etc/openflare"
defaultPagesDirRelativePath = "var/lib/openflare/pages"
defaultMMDBRelativePath = "etc/openflare/GeoLite2-Country.mmdb"
defaultCityMMDBRelativePath = "etc/openflare/GeoLite2-City.mmdb"
defaultAccessLogRelativePath = "var/log/openflare/access.log"
defaultStateRelativePath = "var/lib/openflare/agent-state.json"
defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json"
@@ -31,6 +32,7 @@ const (
defaultObservabilityReplayMinutes = 15
defaultMMDBUpdateInterval = 24 * time.Hour
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
defaultCityMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-City.mmdb"
defaultHeartbeatInterval = 10 * time.Second
defaultRequestTimeout = 10 * time.Second
configFilePerm = 0o600
@@ -58,8 +60,10 @@ type Config struct {
RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"`
CityMMDBPath string `json:"city_mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
CityMMDBDownloadURL string `json:"city_mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -89,8 +93,10 @@ type configFile struct {
RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"`
CityMMDBPath string `json:"city_mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
CityMMDBDownloadURL string `json:"city_mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -178,6 +184,7 @@ func applyAgentPathDefaults(cfg *Config, baseDir string) {
{&cfg.RuntimeConfigDir, defaultRuntimeConfigDirRelativePath},
{&cfg.PagesDir, defaultPagesDirRelativePath},
{&cfg.MMDBPath, defaultMMDBRelativePath},
{&cfg.CityMMDBPath, defaultCityMMDBRelativePath},
{&cfg.ObservabilityBufferPath, defaultObservabilityBufferRelativePath},
}
for _, item := range pathDefaults {
@@ -200,6 +207,9 @@ func applyAgentTimingDefaults(cfg *Config) {
if cfg.MMDBDownloadURL == "" {
cfg.MMDBDownloadURL = defaultMMDBDownloadURL
}
if cfg.CityMMDBDownloadURL == "" {
cfg.CityMMDBDownloadURL = defaultCityMMDBDownloadURL
}
if cfg.OpenrestyObservabilityPort <= 0 {
cfg.OpenrestyObservabilityPort = defaultOpenRestyObservabilityPort
}
@@ -232,6 +242,7 @@ func normalizeManagedPaths(cfg *Config) {
&cfg.StatePath,
&cfg.ObservabilityBufferPath,
&cfg.MMDBPath,
&cfg.CityMMDBPath,
}
for _, p := range paths {
if usesSlashPath(*p) {
@@ -256,6 +267,8 @@ func hasEnvConfig() bool {
"OPENFLARE_MMDB_PATH",
"OPENFLARE_MMDB_UPDATE_INTERVAL",
"OPENFLARE_MMDB_DOWNLOAD_URL",
"OPENFLARE_CITY_MMDB_PATH",
"OPENFLARE_CITY_MMDB_DOWNLOAD_URL",
} {
if strings.TrimSpace(os.Getenv(key)) != "" {
return true
@@ -283,6 +296,8 @@ func applyEnvOverrides(cfg *Config) {
overrideString("OPENFLARE_PAGES_DIR", &cfg.PagesDir)
overrideString("OPENFLARE_MMDB_PATH", &cfg.MMDBPath)
overrideString("OPENFLARE_MMDB_DOWNLOAD_URL", &cfg.MMDBDownloadURL)
overrideString("OPENFLARE_CITY_MMDB_PATH", &cfg.CityMMDBPath)
overrideString("OPENFLARE_CITY_MMDB_DOWNLOAD_URL", &cfg.CityMMDBDownloadURL)
if value := strings.TrimSpace(os.Getenv("OPENFLARE_HEARTBEAT_INTERVAL")); value != "" {
if duration, err := parseDurationValue(value); err == nil {
cfg.HeartbeatInterval = duration
+30
View File
@@ -61,6 +61,12 @@ func TestLoadDefaultsToManagedBinaryPaths(t *testing.T) {
if cfg.RuntimeConfigDir != filepath.Join(dir, "data", defaultRuntimeConfigDirRelativePath) {
t.Fatalf("unexpected runtime config dir: %s", cfg.RuntimeConfigDir)
}
if cfg.CityMMDBPath != filepath.Join(dir, "data", defaultCityMMDBRelativePath) {
t.Fatalf("unexpected city mmdb path: %s", cfg.CityMMDBPath)
}
if cfg.CityMMDBDownloadURL != defaultCityMMDBDownloadURL {
t.Fatalf("unexpected city mmdb download URL: %s", cfg.CityMMDBDownloadURL)
}
if cfg.OpenrestyCertDir != cfg.CertDir {
t.Fatalf("unexpected openresty cert dir: %s", cfg.OpenrestyCertDir)
}
@@ -333,6 +339,8 @@ func TestLoadEnvOverridesConfigFile(t *testing.T) {
t.Setenv("OPENFLARE_SERVER_URL", "http://new:3000")
t.Setenv("OPENFLARE_AGENT_TOKEN", "new-token")
t.Setenv("OPENFLARE_OPENRESTY_PATH", "/new/openresty")
t.Setenv("OPENFLARE_CITY_MMDB_PATH", "/new/GeoLite2-City.mmdb")
t.Setenv("OPENFLARE_CITY_MMDB_DOWNLOAD_URL", "https://geo.example/GeoLite2-City.mmdb")
cfg, err := Load(configPath)
if err != nil {
@@ -347,6 +355,25 @@ func TestLoadEnvOverridesConfigFile(t *testing.T) {
if cfg.OpenrestyPath != "/new/openresty" {
t.Fatalf("expected openresty path from env, got %s", cfg.OpenrestyPath)
}
if cfg.CityMMDBPath != "/new/GeoLite2-City.mmdb" || cfg.CityMMDBDownloadURL != "https://geo.example/GeoLite2-City.mmdb" {
t.Fatalf("unexpected City MMDB env overrides: %s / %s", cfg.CityMMDBPath, cfg.CityMMDBDownloadURL)
}
}
func TestLoadKeepsExplicitCityMMDBConfig(t *testing.T) {
dir := t.TempDir()
configPath := filepath.Join(dir, "agent.json")
payload := `{"server_url":"http://127.0.0.1:3000","agent_token":"token","node_name":"edge-01","node_ip":"10.0.0.8","city_mmdb_path":"/custom/GeoLite2-City.mmdb","city_mmdb_download_url":"https://custom.example/GeoLite2-City.mmdb"}`
if err := os.WriteFile(configPath, []byte(payload), 0o644); err != nil {
t.Fatalf("failed to write config: %v", err)
}
cfg, err := Load(configPath)
if err != nil {
t.Fatalf("Load failed: %v", err)
}
if cfg.CityMMDBPath != "/custom/GeoLite2-City.mmdb" || cfg.CityMMDBDownloadURL != "https://custom.example/GeoLite2-City.mmdb" {
t.Fatalf("explicit City MMDB config changed: %s / %s", cfg.CityMMDBPath, cfg.CityMMDBDownloadURL)
}
}
func TestLoadUsesMillisecondsForIntervals(t *testing.T) {
@@ -431,6 +458,9 @@ func TestSavePersistsMillisecondsAndOmitsRuntimeVersions(t *testing.T) {
if decoded["observability_replay_minutes"] != float64(defaultObservabilityReplayMinutes) {
t.Fatalf("unexpected observability replay minutes: %#v", decoded["observability_replay_minutes"])
}
if decoded["city_mmdb_path"] != cfg.CityMMDBPath || decoded["city_mmdb_download_url"] != cfg.CityMMDBDownloadURL {
t.Fatalf("City MMDB config was not persisted: %#v", decoded)
}
if _, ok := decoded["nginx_path"]; ok {
t.Fatal("legacy nginx_path should not be persisted")
}
+66 -9
View File
@@ -3,6 +3,7 @@ package geoipupdate
import (
"context"
"errors"
"fmt"
"io/fs"
"log/slog"
@@ -22,9 +23,12 @@ const (
// Updater periodically downloads a fresh GeoIP MMDB file and seeds the
// initial embedded database when none is present on disk.
type Updater struct {
MMDBPath string
DownloadURL string
UpdateInterval time.Duration
MMDBPath string
DownloadURL string
CityMMDBPath string
CityDownloadURL string
UpdateInterval time.Duration
downloadDatabase func(context.Context, string, string) error
}
// EnsureInitialDatabase seeds the MMDB file from the embedded database if it does not exist on disk.
@@ -52,13 +56,68 @@ func (u *Updater) EnsureInitialDatabase() error {
return nil
}
// EnsureInitialDatabases retains the embedded Country seed and immediately
// downloads City when it is absent so subdivision rules work before the first ticker interval.
func (u *Updater) EnsureInitialDatabases(ctx context.Context) error {
var errs []error
if err := u.EnsureInitialDatabase(); err != nil {
errs = append(errs, err)
}
cityPath := filepath.Clean(u.CityMMDBPath)
if cityPath == "" || cityPath == "." || u.CityDownloadURL == "" {
return errors.Join(errs...)
}
if _, err := os.Stat(cityPath); err == nil {
return errors.Join(errs...)
} else if !os.IsNotExist(err) {
errs = append(errs, fmt.Errorf("stat City mmdb file failed: %w", err))
return errors.Join(errs...)
}
if err := u.download(ctx, cityPath, u.CityDownloadURL); err != nil {
errs = append(errs, fmt.Errorf("download initial City mmdb failed: %w", err))
} else {
slog.Info("initialized GeoIP City mmdb from provider", "path", cityPath)
}
return errors.Join(errs...)
}
func (u *Updater) download(ctx context.Context, path string, downloadURL string) error {
if u.downloadDatabase != nil {
return u.downloadDatabase(ctx, path, downloadURL)
}
return geoip.DownloadMaxMindDatabase(ctx, path, downloadURL)
}
func (u *Updater) updateDatabases(ctx context.Context) error {
databases := []struct {
name string
path string
downloadURL string
}{
{name: "Country", path: u.MMDBPath, downloadURL: u.DownloadURL},
{name: "City", path: u.CityMMDBPath, downloadURL: u.CityDownloadURL},
}
var errs []error
for _, database := range databases {
if database.path == "" || (database.name == "City" && database.downloadURL == "") {
continue
}
if err := u.download(ctx, database.path, database.downloadURL); err != nil {
errs = append(errs, fmt.Errorf("update GeoIP %s mmdb failed: %w", database.name, err))
continue
}
slog.Info("GeoIP mmdb updated", "database", database.name, "path", database.path)
}
return errors.Join(errs...)
}
// Run starts the periodic GeoIP update loop and blocks until ctx is cancelled.
func (u *Updater) Run(ctx context.Context) {
if u == nil || u.MMDBPath == "" || u.UpdateInterval <= 0 {
return
}
if err := u.EnsureInitialDatabase(); err != nil {
slog.Warn("initialize GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
if err := u.EnsureInitialDatabases(ctx); err != nil {
slog.Warn("initialize GeoIP databases failed", "country_path", u.MMDBPath, "city_path", u.CityMMDBPath, "error", err)
}
ticker := time.NewTicker(u.UpdateInterval)
defer ticker.Stop()
@@ -67,11 +126,9 @@ func (u *Updater) Run(ctx context.Context) {
case <-ctx.Done():
return
case <-ticker.C:
if err := geoip.DownloadMaxMindDatabase(ctx, u.MMDBPath, u.DownloadURL); err != nil {
slog.Warn("update GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
continue
if err := u.updateDatabases(ctx); err != nil {
slog.Warn("update GeoIP databases failed", "error", err)
}
slog.Info("GeoIP mmdb updated", "path", u.MMDBPath)
}
}
}
@@ -1,8 +1,11 @@
package geoipupdate
import (
"context"
"errors"
"os"
"path/filepath"
"slices"
"testing"
)
@@ -22,3 +25,77 @@ func TestEnsureInitialDatabaseCopiesEmbeddedMMDB(t *testing.T) {
t.Fatal("expected copied mmdb to be non-empty")
}
}
func TestEnsureInitialDatabasesDownloadsMissingCity(t *testing.T) {
tempDir := t.TempDir()
countryPath := filepath.Join(tempDir, "GeoLite2-Country.mmdb")
cityPath := filepath.Join(tempDir, "GeoLite2-City.mmdb")
updater := &Updater{
MMDBPath: countryPath,
CityMMDBPath: cityPath,
CityDownloadURL: "https://geo.example/GeoLite2-City.mmdb",
downloadDatabase: func(_ context.Context, path, downloadURL string) error {
if path != cityPath || downloadURL != "https://geo.example/GeoLite2-City.mmdb" {
t.Fatalf("unexpected initial download: %s / %s", path, downloadURL)
}
return os.WriteFile(path, []byte("city-mmdb"), 0o600)
},
}
if err := updater.EnsureInitialDatabases(context.Background()); err != nil {
t.Fatalf("EnsureInitialDatabases failed: %v", err)
}
if _, err := os.Stat(countryPath); err != nil {
t.Fatalf("expected embedded Country database: %v", err)
}
data, err := os.ReadFile(cityPath)
if err != nil || string(data) != "city-mmdb" {
t.Fatalf("expected downloaded City database, data=%q err=%v", data, err)
}
}
func TestEnsureInitialDatabasesKeepsCountryFallbackWhenCityDownloadFails(t *testing.T) {
tempDir := t.TempDir()
countryPath := filepath.Join(tempDir, "GeoLite2-Country.mmdb")
cityPath := filepath.Join(tempDir, "GeoLite2-City.mmdb")
updater := &Updater{
MMDBPath: countryPath,
CityMMDBPath: cityPath,
CityDownloadURL: "https://geo.example/GeoLite2-City.mmdb",
downloadDatabase: func(_ context.Context, _, _ string) error {
return errors.New("city unavailable")
},
}
if err := updater.EnsureInitialDatabases(context.Background()); err == nil {
t.Fatal("expected City download error to be reported")
}
if _, err := os.Stat(countryPath); err != nil {
t.Fatalf("expected Country fallback to remain available: %v", err)
}
if _, err := os.Stat(cityPath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("expected failed City download not to create a database, err=%v", err)
}
}
func TestUpdateDatabasesAttemptsCityAfterCountryFailure(t *testing.T) {
var paths []string
updater := &Updater{
MMDBPath: "/data/GeoLite2-Country.mmdb",
DownloadURL: "https://geo.example/GeoLite2-Country.mmdb",
CityMMDBPath: "/data/GeoLite2-City.mmdb",
CityDownloadURL: "https://geo.example/GeoLite2-City.mmdb",
downloadDatabase: func(_ context.Context, path, _ string) error {
paths = append(paths, path)
if path == "/data/GeoLite2-Country.mmdb" {
return errors.New("country unavailable")
}
return nil
},
}
err := updater.updateDatabases(context.Background())
if err == nil || !slices.Equal(paths, []string{"/data/GeoLite2-Country.mmdb", "/data/GeoLite2-City.mmdb"}) {
t.Fatalf("expected independent Country then City attempts, paths=%#v err=%v", paths, err)
}
}
+192 -11
View File
@@ -18,9 +18,12 @@ import (
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"sync"
"time"
sharedprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"github.com/Rain-kl/Wavelet/pkg/utils"
@@ -31,11 +34,23 @@ import (
// RuntimeConfigDirPlaceholder is substituted into generated configs at apply time.
const RuntimeConfigDirPlaceholder = "__OPENFLARE_RUNTIME_CONFIG_DIR__"
// CountryMMDBPathPlaceholder is substituted with the configured Country database path.
const CountryMMDBPathPlaceholder = "__OPENFLARE_COUNTRY_MMDB_PATH__"
// CityMMDBPathPlaceholder is substituted with the configured City database path.
const CityMMDBPathPlaceholder = "__OPENFLARE_CITY_MMDB_PATH__"
// WAFIPGroupsMaxSnapshotBytesPlaceholder is replaced with the shared protocol aggregate snapshot limit.
const WAFIPGroupsMaxSnapshotBytesPlaceholder = "__OPENFLARE_WAF_IP_GROUPS_MAX_SNAPSHOT_BYTES__"
// ResolverDirectivePlaceholder is substituted into generated configs at apply time.
const ResolverDirectivePlaceholder = "__OPENFLARE_RESOLVER_DIRECTIVE__"
// WAFIPGroupsConfigFileName is the runtime filename for synced WAF IP group data.
const WAFIPGroupsConfigFileName = "waf_ip_groups.json"
// WAFIPGroupsChecksumFileName is atomically published after the IP group JSON snapshot.
const WAFIPGroupsChecksumFileName = WAFIPGroupsConfigFileName + ".checksum"
const powConfigFileName = "pow_config.json"
const (
@@ -45,6 +60,7 @@ const (
stubStatusCheckTimeout = 1500 * time.Millisecond
nginxVersionSubmatchCount = 2
resolverAddressCapacity = 2
workerInitSubmatchCount = 2
)
// Executor controls OpenResty validation, reload, health, and lifecycle operations.
@@ -169,11 +185,15 @@ type Manager struct {
LuaDir string
NginxLuaDir string
RuntimeConfigDir string
MMDBPath string
CityMMDBPath string
PagesDir string
OpenrestyObservabilityListen string
OpenrestyObservabilityPort int
OpenrestyResolverDirective string
Executor Executor
atomicFileWriter func(path string, data []byte, perm os.FileMode) error
wafIPGroupsMu sync.Mutex
}
// ApplyStatus reports the outcome of an OpenResty configuration apply.
@@ -522,12 +542,21 @@ func (m *Manager) CurrentChecksum() (string, error) {
// WAFIPGroupChecksums returns checksums for locally synced WAF IP groups.
func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) {
m.wafIPGroupsMu.Lock()
defer m.wafIPGroupsMu.Unlock()
config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil {
return nil, err
}
if err = m.ensureWAFIPGroupsChecksum(); err != nil {
return nil, err
}
result := make(map[string]string, len(config.Groups))
for id, group := range config.Groups {
if id != strconv.FormatUint(uint64(group.ID), 10) {
continue
}
if strings.TrimSpace(group.Checksum) != "" {
result[id] = strings.TrimSpace(group.Checksum)
}
@@ -535,39 +564,164 @@ func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) {
return result, nil
}
// SyncWAFIPGroups writes WAF IP group definitions to the runtime config directory.
func (m *Manager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error {
if m.RuntimeConfigDir == "" || len(groups) == 0 {
// ReconcileWAFIPGroups atomically replaces the runtime snapshot with exactly the
// authoritative target IDs, retaining local definitions that did not change.
func (m *Manager) ReconcileWAFIPGroups(targetIDs []uint, changed []protocol.WAFIPGroup) error {
if m.RuntimeConfigDir == "" {
return nil
}
m.wafIPGroupsMu.Lock()
defer m.wafIPGroupsMu.Unlock()
config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil {
return err
}
if config.Groups == nil {
config.Groups = make(map[string]protocol.WAFIPGroup)
}
for _, group := range groups {
if group.ID == 0 {
target := make(map[string]protocol.WAFIPGroup, len(targetIDs))
targetSet := make(map[uint]struct{}, len(targetIDs))
for _, id := range targetIDs {
if id == 0 {
continue
}
config.Groups[fmt.Sprintf("%d", group.ID)] = group
targetSet[id] = struct{}{}
key := strconv.FormatUint(uint64(id), 10)
if group, ok := config.Groups[key]; ok && group.ID == id {
target[key] = group
}
}
data, err := json.Marshal(config)
for _, group := range changed {
if _, ok := targetSet[group.ID]; !ok {
continue
}
target[strconv.FormatUint(uint64(group.ID), 10)] = group
}
for _, id := range targetIDs {
if id == 0 {
continue
}
if _, ok := target[strconv.FormatUint(uint64(id), 10)]; !ok {
return fmt.Errorf("missing referenced WAF IP group %d after synchronization", id)
}
}
return m.publishWAFIPGroups(target)
}
// UpdateExistingWAFIPGroups applies real-time changes only to groups already
// present in the authoritative local snapshot.
func (m *Manager) UpdateExistingWAFIPGroups(changed []protocol.WAFIPGroup) error {
if m.RuntimeConfigDir == "" || len(changed) == 0 {
return nil
}
m.wafIPGroupsMu.Lock()
defer m.wafIPGroupsMu.Unlock()
config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil {
return err
}
updated := false
for _, group := range changed {
key := strconv.FormatUint(uint64(group.ID), 10)
if _, ok := config.Groups[key]; !ok || group.ID == 0 {
continue
}
config.Groups[key] = group
updated = true
}
if !updated {
return nil
}
return m.publishWAFIPGroups(config.Groups)
}
func (m *Manager) publishWAFIPGroups(groups map[string]protocol.WAFIPGroup) error {
data, err := sharedprotocol.MarshalWAFIPGroupSnapshot(groups)
if err != nil {
return err
}
if len(data) > sharedprotocol.MaxWAFIPGroupSnapshotBytes {
return fmt.Errorf("WAF IP group snapshot size %d exceeds maximum %d bytes", len(data), sharedprotocol.MaxWAFIPGroupSnapshotBytes)
}
if err := os.MkdirAll(m.RuntimeConfigDir, nginxDirPerm); err != nil {
return err
}
path := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsConfigFileName)
if err := os.WriteFile(path, data, nginxConfigFilePerm); err != nil {
if err := m.writeAtomicFile(path, data, nginxConfigFilePerm); err != nil {
return fmt.Errorf("write %s: %w", WAFIPGroupsConfigFileName, err)
}
checksumPath := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsChecksumFileName)
if err := m.writeAtomicFile(checksumPath, []byte(checksum(string(data))+"\n"), nginxConfigFilePerm); err != nil {
return fmt.Errorf("write %s: %w", WAFIPGroupsChecksumFileName, err)
}
slog.Info("synced waf ip groups", "path", path, "group_count", len(groups))
return nil
}
func (m *Manager) ensureWAFIPGroupsChecksum() error {
if m.RuntimeConfigDir == "" {
return nil
}
jsonPath := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsConfigFileName)
data, err := os.ReadFile(jsonPath) //nolint:gosec // path is under managed RuntimeConfigDir
if err != nil {
if os.IsNotExist(err) {
return nil
}
return err
}
expected := checksum(string(data))
checksumPath := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsChecksumFileName)
current, err := os.ReadFile(checksumPath) //nolint:gosec // path is under managed RuntimeConfigDir
if err == nil && strings.TrimSpace(string(current)) == expected {
return nil
}
if err != nil && !os.IsNotExist(err) {
return err
}
return m.writeAtomicFile(checksumPath, []byte(expected+"\n"), nginxConfigFilePerm)
}
func (m *Manager) writeAtomicFile(path string, data []byte, perm os.FileMode) error {
if m.atomicFileWriter != nil {
return m.atomicFileWriter(path, data, perm)
}
return writeAtomicFile(path, data, perm)
}
func writeAtomicFile(path string, data []byte, perm os.FileMode) (resultErr error) {
tempFile, err := os.CreateTemp(filepath.Dir(path), "."+filepath.Base(path)+".tmp-*")
if err != nil {
return err
}
tempPath := tempFile.Name()
closed := false
defer func() {
if !closed {
if closeErr := tempFile.Close(); resultErr == nil && closeErr != nil {
resultErr = closeErr
}
}
_ = os.Remove(tempPath)
}()
if err = tempFile.Chmod(perm); err != nil {
return err
}
if _, err = tempFile.Write(data); err != nil {
return err
}
if err = tempFile.Sync(); err != nil {
return err
}
if err = tempFile.Close(); err != nil {
return err
}
closed = true
if err = os.Rename(tempPath, path); err != nil {
return err
}
return nil
}
func (m *Manager) readWAFIPGroupsRuntimeConfig() (*wafIPGroupsRuntimeConfig, error) {
config := &wafIPGroupsRuntimeConfig{Groups: map[string]protocol.WAFIPGroup{}}
if m.RuntimeConfigDir == "" {
@@ -1276,6 +1430,9 @@ func (m *Manager) renderMainConfig(content string) string {
}
if luaDir := m.luaRuntimePath(); luaDir != "" {
rendered = strings.ReplaceAll(rendered, openrestyrender.LuaDirPlaceholder, luaDir)
if strings.Contains(rendered, "lua_shared_dict openflare_waf_config") {
rendered = injectWAFWorkerInit(rendered, luaDir)
}
}
if listen := strings.TrimSpace(m.OpenrestyObservabilityListen); listen != "" {
rendered = strings.ReplaceAll(rendered, openrestyrender.ObservabilityListenPlaceholder, listen)
@@ -1289,6 +1446,19 @@ func (m *Manager) renderMainConfig(content string) string {
return rendered
}
func injectWAFWorkerInit(content string, luaDir string) string {
packagePath := fmt.Sprintf(" lua_package_path \"%s/?.lua;%s/?/init.lua;;\";\n", luaDir, luaDir)
if !strings.Contains(content, "lua_package_path ") {
content = strings.Replace(content, "http {", "http {\n"+packagePath, 1)
}
initPattern := regexp.MustCompile(`(?m)^[ \t]*init_worker_by_lua_file[ \t]+([^;]+);[ \t]*$`)
if match := initPattern.FindStringSubmatch(content); len(match) == workerInitSubmatchCount {
block := fmt.Sprintf(" init_worker_by_lua_block {\n require(\"waf.runtime\").init()\n dofile(%q)\n }", strings.TrimSpace(match[1]))
return initPattern.ReplaceAllString(content, block)
}
return strings.Replace(content, "http {", "http {\n init_worker_by_lua_block { require(\"waf.runtime\").init() }", 1)
}
func (m *Manager) managedPowLuaFiles() []protocol.SupportFile {
files := ManagedPowLuaFiles()
runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir))
@@ -1301,8 +1471,19 @@ func (m *Manager) managedPowLuaFiles() []protocol.SupportFile {
func (m *Manager) managedWAFLuaFiles() []protocol.SupportFile {
files := ManagedWAFLuaFiles()
runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir))
countryMMDBPath := filepath.ToSlash(strings.TrimSpace(m.MMDBPath))
if countryMMDBPath == "" {
countryMMDBPath = filepath.ToSlash(filepath.Join(runtimeConfigDir, "GeoLite2-Country.mmdb"))
}
cityMMDBPath := filepath.ToSlash(strings.TrimSpace(m.CityMMDBPath))
if cityMMDBPath == "" {
cityMMDBPath = filepath.ToSlash(filepath.Join(runtimeConfigDir, "GeoLite2-City.mmdb"))
}
for index := range files {
files[index].Content = strings.ReplaceAll(files[index].Content, RuntimeConfigDirPlaceholder, runtimeConfigDir)
files[index].Content = strings.ReplaceAll(files[index].Content, CountryMMDBPathPlaceholder, countryMMDBPath)
files[index].Content = strings.ReplaceAll(files[index].Content, CityMMDBPathPlaceholder, cityMMDBPath)
files[index].Content = strings.ReplaceAll(files[index].Content, WAFIPGroupsMaxSnapshotBytesPlaceholder, strconv.Itoa(sharedprotocol.MaxWAFIPGroupSnapshotBytes))
}
return files
}
+340 -14
View File
@@ -3,6 +3,7 @@ package nginx
import (
"context"
"errors"
"fmt"
"net"
"net/http"
"os"
@@ -13,6 +14,7 @@ import (
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
sharedprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
)
type runCall struct {
@@ -330,6 +332,29 @@ func TestManagerApplyWritesSupportFilesAndReplacesPlaceholder(t *testing.T) {
}
}
func TestManagerRenderMainConfigInitializesWAFRuntimeInWorker(t *testing.T) {
manager := &Manager{NginxLuaDir: "/etc/nginx/openflare-lua"}
rendered := manager.renderMainConfig("events {}\nhttp {\n lua_shared_dict openflare_waf_config 1m;\n server {}\n}\n")
want := `init_worker_by_lua_block { require("waf.runtime").init() }`
if !strings.Contains(rendered, want) {
t.Fatalf("expected worker-time WAF initialization %q, got:\n%s", want, rendered)
}
}
func TestManagerRenderMainConfigMergesExistingWorkerInitializer(t *testing.T) {
manager := &Manager{NginxLuaDir: "/etc/nginx/openflare-lua"}
rendered := manager.renderMainConfig("http {\n lua_shared_dict openflare_waf_config 1m;\n init_worker_by_lua_file /etc/nginx/openflare-lua/observability/init.lua;\n}\n")
if strings.Count(rendered, "init_worker_by_lua_") != 1 {
t.Fatalf("expected one merged worker initializer, got:\n%s", rendered)
}
if !strings.Contains(rendered, `require("waf.runtime").init()`) {
t.Fatalf("expected WAF initialization in merged block, got:\n%s", rendered)
}
if !strings.Contains(rendered, `dofile("/etc/nginx/openflare-lua/observability/init.lua")`) {
t.Fatalf("expected existing worker initializer to be preserved, got:\n%s", rendered)
}
}
func TestManagerCheckHealthUsesStubStatusInsteadOfConfigTest(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
@@ -545,12 +570,51 @@ func TestManagerEnsureLuaAssetsWritesReadableFiles(t *testing.T) {
if _, err := os.Stat(filepath.Join(manager.LuaDir, "pow", "check.lua")); err != nil {
t.Fatalf("failed to stat pow lua file: %v", err)
}
data, err := os.ReadFile(filepath.Join(manager.LuaDir, "pow", "runtime.lua"))
data, err := os.ReadFile(filepath.Join(manager.LuaDir, "waf", "runtime.lua"))
if err != nil {
t.Fatalf("failed to read pow lua file: %v", err)
}
if !strings.Contains(string(data), filepath.ToSlash(manager.RuntimeConfigDir)+"/waf_config.json") {
t.Fatalf("expected pow lua to read runtime config dir, got %s", string(data))
if !strings.Contains(string(data), filepath.ToSlash(manager.RuntimeConfigDir)) || !strings.Contains(string(data), `runtime_dir .. "/waf_config.json"`) {
t.Fatalf("expected WAF runtime to load its worker snapshot from the runtime config dir, got %s", string(data))
}
ipGroupsData, err := os.ReadFile(filepath.Join(manager.LuaDir, "waf", "ip_groups.lua"))
if err != nil {
t.Fatalf("failed to read WAF IP group refresh module: %v", err)
}
if !strings.Contains(string(ipGroupsData), filepath.ToSlash(manager.RuntimeConfigDir)) || !strings.Contains(string(ipGroupsData), "waf_ip_groups.json.checksum") {
t.Fatalf("expected IP group module to use the managed runtime checksum path, got %s", string(ipGroupsData))
}
if !strings.Contains(string(ipGroupsData), "openflare_waf_ip_groups") ||
!strings.Contains(string(ipGroupsData), fmt.Sprintf("max_snapshot_bytes = options.max_snapshot_bytes or tonumber(\"%d\")", sharedprotocol.MaxWAFIPGroupSnapshotBytes)) {
t.Fatalf("expected dedicated dictionary and shared protocol size limit in deployed IP group module, got %s", string(ipGroupsData))
}
powData, err := os.ReadFile(filepath.Join(manager.LuaDir, "pow", "runtime.lua"))
if err != nil {
t.Fatalf("failed to read pow lua file: %v", err)
}
if strings.Contains(string(powData), "io.open") {
t.Fatalf("expected PoW node evaluation not to read configuration files, got %s", string(powData))
}
}
func TestManagerEnsureLuaAssetsUsesConfiguredGeoIPDatabasePaths(t *testing.T) {
tempDir := t.TempDir()
manager := &Manager{
LuaDir: filepath.Join(tempDir, "lua"),
MMDBPath: "/custom/GeoLite2-Country.mmdb",
CityMMDBPath: "/custom/GeoLite2-City.mmdb",
}
if err := manager.EnsureLuaAssets(); err != nil {
t.Fatalf("EnsureLuaAssets failed: %v", err)
}
data, err := os.ReadFile(filepath.Join(manager.LuaDir, "waf", "runtime.lua"))
if err != nil {
t.Fatalf("read WAF runtime: %v", err)
}
for _, path := range []string{manager.MMDBPath, manager.CityMMDBPath} {
if !strings.Contains(string(data), path) {
t.Fatalf("expected configured GeoIP path %q in WAF runtime", path)
}
}
}
@@ -667,7 +731,7 @@ func TestManagerCurrentChecksumIncludesPowConfig(t *testing.T) {
}
func TestManagedPowLuaFilesUseInternalChallengeFlow(t *testing.T) {
if !strings.Contains(openRestyPowRuntimeLua, `return ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")`) {
if !strings.Contains(openRestyPowRuntimeLua, `ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")`) {
t.Fatal("expected pow runtime lua to internally execute make-challenge instead of issuing a 302 redirect")
}
if strings.Contains(openRestyPowRuntimeLua, "ngx.redirect(") {
@@ -699,12 +763,28 @@ func TestManagedPowLuaFilesUseInternalChallengeFlow(t *testing.T) {
}
}
func TestManagedWAFLuaTreatsWhitelistAsBypass(t *testing.T) {
if !strings.Contains(openRestyWAFRuntimeLua, "if ip_matches(group.ip_whitelist, ip)") {
t.Fatal("expected waf runtime to bypass request when ip matches whitelist")
func TestManagedPowLuaFilesPreserveConfigAcrossInternalRedirect(t *testing.T) {
for _, expected := range []string{
`pow_config_dict:set(config_key, cjson.encode(config)`,
`openflare_pow_config_key = config_key`,
`return false`,
} {
if !strings.Contains(openRestyPowRuntimeLua, expected) {
t.Fatalf("expected PoW runtime to contain %q", expected)
}
}
if strings.Contains(openRestyWAFRuntimeLua, "first_allowlist_group") {
t.Fatal("expected waf runtime not to block requests that miss configured whitelists")
if !strings.Contains(openRestyPowChallengeLua, `pow_config_dict:get(config_key)`) {
t.Fatal("expected challenge handler to restore reached PoW node config after internal redirect")
}
}
func TestManagedWAFLuaExecutesCompiledGraphWithoutRequestIO(t *testing.T) {
if !strings.Contains(openRestyWAFRuntimeLua, `node.type == "ip_match"`) {
t.Fatal("expected WAF runtime to execute compiled IP match nodes")
}
checkStart := strings.Index(openRestyWAFRuntimeLua, "function _M.check()")
if checkStart < 0 || strings.Contains(openRestyWAFRuntimeLua[checkStart:], "io.open") {
t.Fatal("expected WAF request path not to perform file I/O")
}
}
@@ -924,18 +1004,18 @@ func TestManagerApplyRejectsCertFilePathTraversal(t *testing.T) {
}
}
func TestManagerSyncWAFIPGroupsWritesDeltaRuntimeFile(t *testing.T) {
func TestManagerReconcileWAFIPGroupsRetainsUnchangedDeltaRuntimeFile(t *testing.T) {
manager := &Manager{RuntimeConfigDir: t.TempDir()}
if err := manager.SyncWAFIPGroups([]protocol.WAFIPGroup{
if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{
{ID: 1, Enabled: true, IPList: []string{"203.0.113.10"}, Checksum: "sum-1"},
}); err != nil {
t.Fatalf("SyncWAFIPGroups failed: %v", err)
t.Fatalf("ReconcileWAFIPGroups failed: %v", err)
}
if err := manager.SyncWAFIPGroups([]protocol.WAFIPGroup{
if err := manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{
{ID: 2, Enabled: true, IPList: []string{"198.51.100.10"}, Checksum: "sum-2"},
}); err != nil {
t.Fatalf("SyncWAFIPGroups second delta failed: %v", err)
t.Fatalf("ReconcileWAFIPGroups second delta failed: %v", err)
}
checksums, err := manager.WAFIPGroupChecksums()
@@ -955,6 +1035,252 @@ func TestManagerSyncWAFIPGroupsWritesDeltaRuntimeFile(t *testing.T) {
}
}
func TestManagerReconcileWAFIPGroupsConvergesToAuthoritativeTarget(t *testing.T) {
runtimeDir := t.TempDir()
manager := &Manager{RuntimeConfigDir: runtimeDir}
initial, err := sharedprotocol.MarshalWAFIPGroupSnapshot(map[string]protocol.WAFIPGroup{
"1": {ID: 1, Name: "unchanged", Enabled: true, Checksum: "sum-1"},
"2": {ID: 2, Name: "old", Enabled: true, Checksum: "old-2"},
"99": {ID: 99, Name: strings.Repeat("x", sharedprotocol.MaxWAFIPGroupSnapshotBytes-1024), Enabled: true, Checksum: "stale"},
})
if err != nil {
t.Fatal(err)
}
if err = os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), initial, 0o644); err != nil {
t.Fatal(err)
}
if err = manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{{
ID: 2, Name: "changed", Enabled: true, Checksum: "sum-2",
}}); err != nil {
t.Fatalf("ReconcileWAFIPGroups failed after pruning oversized stale data: %v", err)
}
config, err := manager.readWAFIPGroupsRuntimeConfig()
if err != nil {
t.Fatal(err)
}
if len(config.Groups) != 2 {
t.Fatalf("authoritative group count = %d, want 2: %#v", len(config.Groups), config.Groups)
}
if got := config.Groups["1"].Name; got != "unchanged" {
t.Fatalf("unchanged referenced group was not retained: %q", got)
}
if got := config.Groups["2"].Name; got != "changed" {
t.Fatalf("changed referenced group was not merged: %q", got)
}
if _, exists := config.Groups["99"]; exists {
t.Fatal("historical unreferenced group was not pruned")
}
}
func TestManagerUpdateExistingWAFIPGroupsIgnoresUnrelatedBroadcast(t *testing.T) {
runtimeDir := t.TempDir()
manager := &Manager{RuntimeConfigDir: runtimeDir}
if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{{ID: 1, Name: "old", Checksum: "old"}}); err != nil {
t.Fatal(err)
}
if err := manager.UpdateExistingWAFIPGroups([]protocol.WAFIPGroup{
{ID: 1, Name: "new", Checksum: "new"},
{ID: 2, Name: "unrelated", Checksum: "sum-2"},
}); err != nil {
t.Fatal(err)
}
config, err := manager.readWAFIPGroupsRuntimeConfig()
if err != nil {
t.Fatal(err)
}
if len(config.Groups) != 1 || config.Groups["1"].Name != "new" {
t.Fatalf("broadcast update escaped existing target: %#v", config.Groups)
}
}
func TestManagerReconcileWAFIPGroupsPublishesRemovalOnlyAndEmptyTargets(t *testing.T) {
runtimeDir := t.TempDir()
manager := &Manager{RuntimeConfigDir: runtimeDir}
if err := manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{
{ID: 1, Checksum: "sum-1"}, {ID: 2, Checksum: "sum-2"},
}); err != nil {
t.Fatal(err)
}
before, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
if err != nil {
t.Fatal(err)
}
if err = manager.ReconcileWAFIPGroups([]uint{1}, nil); err != nil {
t.Fatalf("removal-only reconcile failed: %v", err)
}
after, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
if err != nil {
t.Fatal(err)
}
if string(before) == string(after) {
t.Fatal("removal-only reconcile did not publish a new checksum")
}
if err = manager.ReconcileWAFIPGroups(nil, nil); err != nil {
t.Fatalf("empty authoritative reconcile failed: %v", err)
}
config, err := manager.readWAFIPGroupsRuntimeConfig()
if err != nil {
t.Fatal(err)
}
if len(config.Groups) != 0 {
t.Fatalf("empty authoritative target retained groups: %#v", config.Groups)
}
}
func TestManagerReconcileWAFIPGroupsRejectsMissingReferencedGroup(t *testing.T) {
runtimeDir := t.TempDir()
data := []byte(`{"groups":{"7":{"id":8,"checksum":"mistaken-match"}}}`)
if err := os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), data, 0o644); err != nil {
t.Fatal(err)
}
manager := &Manager{RuntimeConfigDir: runtimeDir}
checksums, err := manager.WAFIPGroupChecksums()
if err != nil {
t.Fatal(err)
}
if _, ok := checksums["7"]; ok {
t.Fatalf("invalid local group was mistakenly reported matched: %#v", checksums)
}
err = manager.ReconcileWAFIPGroups([]uint{7}, nil)
if err == nil || !strings.Contains(err.Error(), "missing referenced WAF IP group 7") {
t.Fatalf("expected clear missing referenced group error, got %v", err)
}
}
func TestWAFIPGroupChecksumPublishesJSONBeforeSidecar(t *testing.T) {
runtimeDir := t.TempDir()
var writes []string
manager := &Manager{
RuntimeConfigDir: runtimeDir,
atomicFileWriter: func(path string, data []byte, perm os.FileMode) error {
writes = append(writes, filepath.Base(path))
return os.WriteFile(path, data, perm)
},
}
if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{{
ID: 1, Enabled: true, IPList: []string{"203.0.113.10"}, Checksum: "sum-1",
}}); err != nil {
t.Fatalf("SyncWAFIPGroups failed: %v", err)
}
if got, want := strings.Join(writes, ","), WAFIPGroupsConfigFileName+","+WAFIPGroupsChecksumFileName; got != want {
t.Fatalf("expected JSON then checksum publication, got %s", got)
}
jsonData, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName))
if err != nil {
t.Fatalf("read JSON snapshot: %v", err)
}
checksumData, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
if err != nil {
t.Fatalf("read checksum sidecar: %v", err)
}
if got, want := strings.TrimSpace(string(checksumData)), checksum(string(jsonData)); got != want {
t.Fatalf("checksum mismatch: got %q want %q", got, want)
}
}
func TestWAFIPGroupChecksumJSONFailurePreservesOldSidecar(t *testing.T) {
runtimeDir := t.TempDir()
checksumPath := filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName)
if err := os.WriteFile(checksumPath, []byte("old-checksum\n"), 0o644); err != nil {
t.Fatal(err)
}
manager := &Manager{
RuntimeConfigDir: runtimeDir,
atomicFileWriter: func(path string, _ []byte, _ os.FileMode) error {
if filepath.Base(path) == WAFIPGroupsConfigFileName {
return errors.New("json rename failed")
}
return errors.New("checksum must not be written")
},
}
if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{{ID: 1, Checksum: "sum-1"}}); err == nil {
t.Fatal("expected JSON publication failure")
}
data, err := os.ReadFile(checksumPath)
if err != nil || string(data) != "old-checksum\n" {
t.Fatalf("old checksum must remain unchanged, data=%q err=%v", data, err)
}
}
func TestWAFIPGroupChecksumBootstrapsLegacySnapshot(t *testing.T) {
runtimeDir := t.TempDir()
jsonData := []byte(`{"groups":{"7":{"id":7,"enabled":true,"checksum":"sum-7"}}}`)
if err := os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), jsonData, 0o644); err != nil {
t.Fatal(err)
}
manager := &Manager{RuntimeConfigDir: runtimeDir}
checksums, err := manager.WAFIPGroupChecksums()
if err != nil {
t.Fatalf("WAFIPGroupChecksums failed: %v", err)
}
if checksums["7"] != "sum-7" {
t.Fatalf("unexpected group checksums: %#v", checksums)
}
sidecar, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
if err != nil {
t.Fatalf("legacy checksum sidecar was not created: %v", err)
}
if got, want := strings.TrimSpace(string(sidecar)), checksum(string(jsonData)); got != want {
t.Fatalf("legacy checksum mismatch: got %q want %q", got, want)
}
}
func TestWAFIPGroupChecksumRepairsStaleSidecar(t *testing.T) {
runtimeDir := t.TempDir()
jsonData := []byte(`{"groups":{"9":{"id":9,"enabled":true,"checksum":"sum-9"}}}`)
if err := os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), jsonData, 0o644); err != nil {
t.Fatal(err)
}
checksumPath := filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName)
if err := os.WriteFile(checksumPath, []byte("stale-checksum\n"), 0o644); err != nil {
t.Fatal(err)
}
manager := &Manager{RuntimeConfigDir: runtimeDir}
if _, err := manager.WAFIPGroupChecksums(); err != nil {
t.Fatalf("WAFIPGroupChecksums failed: %v", err)
}
sidecar, err := os.ReadFile(checksumPath)
if err != nil {
t.Fatalf("read repaired checksum: %v", err)
}
if got, want := strings.TrimSpace(string(sidecar)), checksum(string(jsonData)); got != want {
t.Fatalf("stale checksum was not repaired: got %q want %q", got, want)
}
}
func TestWAFIPGroupSnapshotRejectsOversizeBeforePublication(t *testing.T) {
runtimeDir := t.TempDir()
jsonPath := filepath.Join(runtimeDir, WAFIPGroupsConfigFileName)
checksumPath := filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName)
oldJSON := []byte(`{"groups":{"1":{"id":1,"enabled":true,"checksum":"old"}}}`)
oldChecksum := []byte("old-checksum\n")
if err := os.WriteFile(jsonPath, oldJSON, 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(checksumPath, oldChecksum, 0o644); err != nil {
t.Fatal(err)
}
manager := &Manager{RuntimeConfigDir: runtimeDir}
err := manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{{
ID: 2, Name: strings.Repeat("x", sharedprotocol.MaxWAFIPGroupSnapshotBytes), Enabled: true, Checksum: "new",
}})
if err == nil || !strings.Contains(err.Error(), "exceeds maximum") {
t.Fatalf("expected clear oversized snapshot error, got %v", err)
}
if got, readErr := os.ReadFile(jsonPath); readErr != nil || string(got) != string(oldJSON) {
t.Fatalf("oversized snapshot touched committed JSON, got=%q err=%v", got, readErr)
}
if got, readErr := os.ReadFile(checksumPath); readErr != nil || string(got) != string(oldChecksum) {
t.Fatalf("oversized snapshot touched committed checksum, got=%q err=%v", got, readErr)
}
}
func TestObservabilityListenAddress(t *testing.T) {
if got := ObservabilityListenAddress(18081); got != "127.0.0.1:18081" {
t.Fatalf("unexpected default observability listen address: %s", got)
+87 -21
View File
@@ -13,7 +13,81 @@ var powStaticFS embed.FS
const openRestyPowRuntimeLua = `local _M = {}
local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
local script_path = string.sub(source, 2)
local base_dir = string.match(script_path, "^(.*)/pow/[^/]+%.lua$")
if base_dir and base_dir ~= "" and not string.find(package.path, base_dir, 1, true) then
package.path = base_dir .. "/?.lua;" .. base_dir .. "/?/init.lua;" .. package.path
end
end
local policy = require "pow.policy"
local pow_sessions = ngx.shared.openflare_pow_sessions
local pow_config_dict = ngx.shared.openflare_pow_config
local cjson = require "cjson.safe"
local function session_cookie(value, ttl)
local cookie = "__openflare_pow=" .. value .. "; Path=/; HttpOnly; SameSite=Lax; Max-Age=" .. tostring(ttl)
if ngx.var.scheme == "https" then cookie = cookie .. "; Secure" end
return cookie
end
-- evaluate is called by a DAG pow node. true continues along its next edge;
-- false means the challenge flow has taken ownership of the request.
function _M.evaluate(config)
config = config or {}
ngx.ctx.openflare_pow_config = config
local host = ngx.var.host
if not host or host == "" then return true end
local session_ttl = config.session_ttl or 600
local uri = ngx.var.uri or ""
local ua = ngx.var.http_user_agent or ""
local remote_ip = ngx.var.remote_addr or ""
if policy.match_any(remote_ip, ua, uri, config.whitelist or {}) then return true end
local blacklist = config.blacklist or {}
if policy.has_entries(blacklist) and not policy.match_any(remote_ip, ua, uri, blacklist) then return true end
local cookie_val = ngx.var["cookie___openflare_pow"]
if cookie_val and cookie_val ~= "" then
local session_key = host .. ":" .. cookie_val
if pow_sessions:get(session_key) then
pow_sessions:set(session_key, "1", session_ttl)
ngx.header["Set-Cookie"] = session_cookie(cookie_val, session_ttl)
return true
end
end
local api_prefix = "/.within.website/x/cmd/anubis/api/"
local static_prefix = "/.within.website/x/cmd/anubis/static/"
if string.sub(uri, 1, #api_prefix) == api_prefix or string.sub(uri, 1, #static_prefix) == static_prefix then
return false
end
local config_key = "_request_config:" .. (ngx.var.request_id or ngx.md5(host .. uri .. tostring(ngx.now())))
pow_config_dict:set(config_key, cjson.encode(config), config.challenge_ttl or 300)
ngx.req.set_uri_args({
redir = ngx.var.scheme .. "://" .. host .. uri .. (ngx.var.args and ("?" .. ngx.var.args) or ""),
host = host,
openflare_pow_config_key = config_key,
})
ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")
return false
end
-- Compatibility entrypoint for old rendered routes. PoW selection now belongs
-- exclusively to WAF graph nodes, so this function intentionally does nothing.
function _M.check()
return true
end
return _M
`
/* Removed legacy request-time configuration scanner. Graph execution now calls
evaluate(config) with the reached node.
local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
local script_path = string.sub(source, 2)
@@ -198,7 +272,7 @@ return ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")
end
return _M
`
*/
const openRestyPowCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
@@ -214,8 +288,8 @@ return require("pow.runtime").check()
const openRestyPowChallengeLua = `local cjson = require "cjson.safe"
local pow_config_dict = ngx.shared.openflare_pow_config
local pow_challenges = ngx.shared.openflare_pow_challenges
local pow_config_dict = ngx.shared.openflare_pow_config
local function generate_entropy()
local pieces = {
@@ -233,28 +307,20 @@ local args = ngx.req.get_uri_args()
local host = args["host"] or ngx.var.host or ""
local redir = args["redir"] or ""
local site = ngx.var.openflare_waf_site or ""
if site == "" then
local config = ngx.ctx.openflare_pow_config
local config_key = args["openflare_pow_config_key"] or ""
if type(config) ~= "table" and config_key ~= "" then
local config_raw = pow_config_dict:get(config_key)
if config_raw then
config = cjson.decode(config_raw)
end
end
if config_key ~= "" then pow_config_dict:delete(config_key) end
if type(config) ~= "table" then
ngx.status = 403
ngx.say("PoW site not resolved; openflare_waf_site is required")
ngx.say("PoW graph node was not evaluated for this request")
return
end
local config_raw = pow_config_dict:get(site)
if not config_raw then
ngx.status = 403
ngx.say("PoW not configured for this site")
return
end
local ok, route_config = pcall(cjson.decode, config_raw)
if not ok or not route_config or not route_config.enabled then
ngx.status = 403
ngx.say("PoW not enabled for this site")
return
end
local config = route_config.config or {}
local difficulty = config.difficulty or 4
local algorithm = config.algorithm or "fast"
local challenge_ttl = config.challenge_ttl or 300
+9 -270
View File
@@ -1,278 +1,16 @@
package nginx
import "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
import (
_ "embed"
const openRestyWAFRuntimeLua = `local _M = {}
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
)
function _M.check()
local cjson = require "cjson.safe"
//go:embed waf_runtime.lua
var openRestyWAFRuntimeLua string
local config_dict = ngx.shared.openflare_waf_config
local function read_file(path)
local f = io.open(path, "r")
if not f then
return nil
end
local content = f:read("*a")
f:close()
return content
end
local function load_config()
local paths = {
"__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_config.json",
"/etc/nginx/openflare-lua/waf_config.json",
"/usr/local/openresty/nginx/conf/waf_config.json"
}
for _, path in ipairs(paths) do
local content = read_file(path)
if content and content ~= "" then
local hash = ngx.md5(content)
if config_dict:get("_config_hash") == hash then
local cached = config_dict:get("_config_json")
if cached then
local decoded = cjson.decode(cached)
if decoded then
return decoded
end
end
end
local decoded = cjson.decode(content)
if decoded then
config_dict:set("_config_hash", hash, 0)
config_dict:set("_config_json", content, 0)
return decoded
end
end
end
return nil
end
local function load_ip_groups()
local paths = {
"__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_ip_groups.json",
"/etc/nginx/openflare-lua/waf_ip_groups.json",
"/usr/local/openresty/nginx/conf/waf_ip_groups.json"
}
for _, path in ipairs(paths) do
local content = read_file(path)
if content and content ~= "" then
local hash = ngx.md5(content)
if config_dict:get("_ip_groups_hash") == hash then
local cached = config_dict:get("_ip_groups_json")
if cached then
local decoded = cjson.decode(cached)
if decoded then
return decoded
end
end
end
local decoded = cjson.decode(content)
if decoded then
config_dict:set("_ip_groups_hash", hash, 0)
config_dict:set("_ip_groups_json", content, 0)
return decoded
end
end
end
return { groups = {} }
end
local function list_contains(items, value)
if not items or type(items) ~= "table" or not value or value == "" then
return false
end
for _, item in ipairs(items) do
if item == value then
return true
end
end
return false
end
local function table_has_items(items)
return type(items) == "table" and #items > 0
end
local function parse_ipv4(value)
local a, b, c, d = string.match(value or "", "^(%d+)%.(%d+)%.(%d+)%.(%d+)$")
if not a then
return nil
end
a, b, c, d = tonumber(a), tonumber(b), tonumber(c), tonumber(d)
if a > 255 or b > 255 or c > 255 or d > 255 then
return nil
end
return ((a * 256 + b) * 256 + c) * 256 + d
end
local function ipv4_in_cidr(ip, cidr)
local base, bits = string.match(cidr or "", "^([^/]+)/(%d+)$")
if not base then
return false
end
bits = tonumber(bits)
if not bits or bits < 0 or bits > 32 then
return false
end
local ip_num = parse_ipv4(ip)
local base_num = parse_ipv4(base)
if not ip_num or not base_num then
return false
end
if bits == 0 then
return true
end
local mask = 4294967295 - (2 ^ (32 - bits) - 1)
return (ip_num - (ip_num % (2 ^ (32 - bits)))) == (base_num - (base_num % (2 ^ (32 - bits))))
end
local function ip_matches(items, ip)
if not items or type(items) ~= "table" or not ip or ip == "" then
return false
end
for _, item in ipairs(items) do
if item == ip then
return true
end
if string.find(item, "/", 1, true) and ipv4_in_cidr(ip, item) then
return true
end
end
return false
end
local function ip_matches_group_ids(group_ids, ip, ip_groups_config)
if not group_ids or type(group_ids) ~= "table" or not ip or ip == "" then
return false
end
local groups = (ip_groups_config or {}).groups or {}
for _, id in ipairs(group_ids) do
local group = groups[tostring(id)]
if group and group.enabled and ip_matches(group.ip_list, ip) then
return true
end
end
return false
end
local function lookup_country(ip)
local ok, maxminddb = pcall(require, "resty.maxminddb")
if not ok or not maxminddb then
return nil
end
local paths = {
"__OPENFLARE_RUNTIME_CONFIG_DIR__/GeoLite2-Country.mmdb",
"/etc/openflare/GeoLite2-Country.mmdb",
"/usr/local/share/openflare/GeoLite2-Country.mmdb"
}
for _, path in ipairs(paths) do
local opened = pcall(maxminddb.init, path)
if opened then
local res, err = maxminddb.lookup(ip)
if res and res.country and res.country.iso_code then
return string.upper(res.country.iso_code)
end
end
end
return nil
end
local function group_by_id(config)
local result = {}
for _, group in ipairs(config.rule_groups or {}) do
result[tostring(group.id)] = group
end
return result
end
local function active_groups(config, groups)
local site = ngx.var.openflare_waf_site or ""
local ids = (config.site_rule_groups or {})[site]
local result = {}
for _, group in ipairs(config.rule_groups or {}) do
if group.is_global then
result[#result + 1] = group
end
end
if ids then
local by_id = group_by_id(config)
for _, id in ipairs(ids) do
local group = by_id[tostring(id)]
if group and not group.is_global then
result[#result + 1] = group
end
end
end
return result
end
local function exit_with_group(group)
ngx.ctx.openflare_waf_blocked = true
ngx.status = tonumber(group.block_status_code) or 418
local body = group.block_response_body or ""
if body ~= "" then
ngx.header["Content-Type"] = "text/html; charset=utf-8"
ngx.say(body)
end
return ngx.exit(ngx.status)
end
local config = load_config()
if not config then
if config_dict:add("_missing_config_logged", true, 60) then
ngx.log(ngx.WARN, "openflare waf config is missing or invalid; requests will be allowed")
end
return
end
local ip = ngx.var.remote_addr or ""
local groups = active_groups(config)
local ip_groups_config = load_ip_groups()
if #groups == 0 then
if config_dict:add("_empty_groups_logged", true, 60) then
ngx.log(ngx.WARN, "openflare waf has no active rule group for site: ", ngx.var.openflare_waf_site or "")
end
return
end
for _, group in ipairs(groups) do
if ip_matches(group.ip_whitelist, ip) or ip_matches_group_ids(group.ip_whitelist_group_ids, ip, ip_groups_config) then
return
end
end
local country = nil
for _, group in ipairs(groups) do
if type(group.country_whitelist) == "table" and #group.country_whitelist > 0 then
country = country or lookup_country(ip)
if list_contains(group.country_whitelist, country) then
return
end
end
end
for _, group in ipairs(groups) do
if ip_matches(group.ip_blacklist, ip) or ip_matches_group_ids(group.ip_blacklist_group_ids, ip, ip_groups_config) then
return exit_with_group(group)
end
end
for _, group in ipairs(groups) do
if type(group.country_blacklist) == "table" and #group.country_blacklist > 0 then
country = country or lookup_country(ip)
if list_contains(group.country_blacklist, country) then
return exit_with_group(group)
end
end
end
return "ok"
end
return _M
`
//go:embed waf_ip_groups.lua
var openRestyWAFIPGroupsLua string
const openRestyWAFCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
@@ -290,6 +28,7 @@ return require("waf.runtime").check()
func ManagedWAFLuaFiles() []protocol.SupportFile {
return []protocol.SupportFile{
{Path: "waf/runtime.lua", Content: openRestyWAFRuntimeLua},
{Path: "waf/ip_groups.lua", Content: openRestyWAFIPGroupsLua},
{Path: "waf/check.lua", Content: openRestyWAFCheckLua},
}
}
@@ -0,0 +1,44 @@
package nginx
import (
"path/filepath"
"testing"
lua "github.com/yuin/gopher-lua"
)
func TestWAFRuntime(t *testing.T) {
state := lua.NewState()
defer state.Close()
runtimePath, err := filepath.Abs("waf_runtime.lua")
if err != nil {
t.Fatal(err)
}
specPath, err := filepath.Abs("waf_runtime_spec.lua")
if err != nil {
t.Fatal(err)
}
state.SetGlobal("WAF_RUNTIME_PATH", lua.LString(runtimePath))
if err := state.DoFile(specPath); err != nil {
t.Fatalf("WAF runtime specification failed: %v", err)
}
}
func TestWAFIPGroupRefresh(t *testing.T) {
state := lua.NewState()
defer state.Close()
modulePath, err := filepath.Abs("waf_ip_groups.lua")
if err != nil {
t.Fatal(err)
}
specPath, err := filepath.Abs("waf_ip_groups_spec.lua")
if err != nil {
t.Fatal(err)
}
state.SetGlobal("WAF_IP_GROUPS_PATH", lua.LString(modulePath))
if err := state.DoFile(specPath); err != nil {
t.Fatalf("WAF IP group refresh specification failed: %v", err)
}
}
+150
View File
@@ -0,0 +1,150 @@
local _M = {}
local current_groups = { groups = {} }
local current_version
local initialized = false
local shared
local read_checksum
local read_json
local decode
local log_warning
local max_snapshot_bytes
local refresh_lock_key = "ip_groups_refresh_lock"
local raw_snapshot_prefix = "ip_groups_raw:"
local version_key = "ip_groups_version"
local previous_version_key = "ip_groups_previous_version"
local function warn(message, err, forcible)
local suffix = err and (": " .. tostring(err)) or ""
if forcible then suffix = suffix .. " (forcible eviction refused)" end
pcall(log_warning, "openflare WAF IP group refresh " .. message .. suffix)
end
local function safe_set(key, value, description)
local ok, err, forcible = shared:safe_set(key, value)
if ok ~= true or forcible == true then
warn(description, err, forcible)
return false
end
return true
end
local function read_file(path)
local file, err = io.open(path, "rb")
if not file then return nil, err end
local content = file:read("*a")
file:close()
return content
end
local function valid_snapshot(snapshot)
return type(snapshot) == "table" and type(snapshot.groups) == "table"
end
local function decode_snapshot(raw)
if type(raw) ~= "string" or raw == "" then return nil end
local called, snapshot = pcall(decode, raw)
if not called or not valid_snapshot(snapshot) then return nil end
return snapshot
end
local function refresh_from_checksum()
local called, checksum = pcall(read_checksum)
if not called or type(checksum) ~= "string" then return end
checksum = string.match(checksum, "^%s*(.-)%s*$")
local committed_version = shared:get(version_key)
if checksum == "" or checksum == committed_version then return end
local json_called, raw = pcall(read_json)
if not json_called then
warn("JSON read failed", raw)
return
end
if type(raw) ~= "string" or #raw > max_snapshot_bytes then
warn("snapshot exceeds maximum " .. tostring(max_snapshot_bytes) .. " bytes")
return
end
if not decode_snapshot(raw) then return end
local raw_key = raw_snapshot_prefix .. checksum
local existing_raw = shared:get(raw_key)
local published_new_raw = false
if existing_raw == nil then
if not safe_set(raw_key, raw, "raw publication failed") then return end
published_new_raw = true
elseif existing_raw ~= raw then
return
end
if not safe_set(version_key, checksum, "commit pointer publication failed") then
if published_new_raw then shared:delete(raw_key) end
return
end
local previous_version = shared:get(previous_version_key)
if type(committed_version) == "string" and committed_version ~= "" and committed_version ~= checksum then
if not safe_set(previous_version_key, committed_version, "previous version metadata publication failed") then return end
if type(previous_version) == "string" and previous_version ~= "" and
previous_version ~= committed_version and previous_version ~= checksum then
shared:delete(raw_snapshot_prefix .. previous_version)
end
end
end
local function adopt_shared_snapshot_if_changed()
local version = shared:get(version_key)
if type(version) ~= "string" or version == "" or version == current_version then return end
local snapshot = decode_snapshot(shared:get(raw_snapshot_prefix .. version))
if not snapshot then return end
current_groups = snapshot
current_version = version
end
local function tick(premature)
if premature then return end
local locked, lock_error, forcible = shared:safe_add(refresh_lock_key, true, 4)
if forcible == true then
warn("coordination lock refused forcible eviction", lock_error, true)
locked = false
elseif not locked and lock_error and lock_error ~= "exists" then
warn("coordination lock failed", lock_error)
end
if locked then refresh_from_checksum() end
adopt_shared_snapshot_if_changed()
end
function _M.init(options)
if initialized then return true end
options = options or {}
local runtime_dir = options.runtime_dir or "__OPENFLARE_RUNTIME_CONFIG_DIR__"
shared = options.shared or (ngx.shared and ngx.shared.openflare_waf_ip_groups)
assert(shared, "openflare_waf_ip_groups shared dictionary is required")
max_snapshot_bytes = options.max_snapshot_bytes or tonumber("__OPENFLARE_WAF_IP_GROUPS_MAX_SNAPSHOT_BYTES__")
assert(max_snapshot_bytes and max_snapshot_bytes > 0, "WAF IP group maximum snapshot size is required")
log_warning = options.log_warning or function(message)
if ngx and ngx.log then ngx.log(ngx.WARN, message) end
end
read_checksum = options.read_checksum or function()
return read_file(runtime_dir .. "/waf_ip_groups.json.checksum")
end
read_json = options.read_json or function()
return read_file(runtime_dir .. "/waf_ip_groups.json")
end
if options.decode then
decode = options.decode
else
local cjson = require("cjson.safe")
decode = cjson.decode
end
local timer_every = options.timer_every or ngx.timer.every
local ok, err = timer_every(5, tick)
if not ok then return nil, err end
initialized = true
tick(false)
return true
end
function _M.current()
return current_groups
end
return _M
@@ -0,0 +1,356 @@
local module_path = assert(WAF_IP_GROUPS_PATH, "WAF_IP_GROUPS_PATH is required")
local function assert_equal(actual, expected, message)
if actual ~= expected then
error((message or "values differ") .. ": expected " .. tostring(expected) .. ", got " .. tostring(actual), 2)
end
end
local shared_data = {}
local locks = {}
local shared = {}
function shared:get(key) return shared_data[key] end
function shared:set(key, value) shared_data[key] = value return true end
function shared:delete(key) shared_data[key] = nil return true end
function shared:safe_set(key, value) return shared:set(key, value) end
function shared:add(key, value, ttl)
assert_equal(ttl, 4, "coordination lock TTL")
if locks[key] then return false end
locks[key] = value
return true
end
function shared:safe_add(key, value, ttl) return shared:add(key, value, ttl) end
local function advance_time() locks = {} end
local disk_checksum = "v1"
local disk_json = "valid-v1"
local checksum_reads = 0
local json_reads = 0
local timer_callbacks = {}
local function decode(raw)
if raw == "valid-v1" then
return { groups = { ["1"] = { enabled = true, ip_list = { "192.0.2.1" } } } }
end
if raw == "valid-v2" then
return { groups = { ["2"] = { enabled = true, ip_list = { "198.51.100.2" } } } }
end
if raw == "valid-v3" then
return { groups = { ["3"] = { enabled = true, ip_list = { "203.0.113.3" } } } }
end
return nil, "invalid json"
end
local function load_worker()
local worker = assert(loadfile(module_path))()
worker.init({
shared = shared,
timer_every = function(interval, callback)
assert_equal(interval, 5, "refresh interval")
timer_callbacks[#timer_callbacks + 1] = callback
return true
end,
read_checksum = function()
checksum_reads = checksum_reads + 1
return disk_checksum
end,
read_json = function()
json_reads = json_reads + 1
return disk_json
end,
decode = decode,
max_snapshot_bytes = 20 * 1024 * 1024,
})
return worker
end
local first = load_worker()
local second = load_worker()
assert_equal(#timer_callbacks, 2, "each worker schedules a refresh timer")
assert_equal(checksum_reads, 1, "one worker coordinates initial checksum read")
assert_equal(json_reads, 1, "one worker reads initial JSON")
assert_equal(first.current().groups["1"].ip_list[1], "192.0.2.1", "first worker adopts initial snapshot")
assert_equal(second.current().groups["1"].ip_list[1], "192.0.2.1", "second worker adopts initial snapshot")
local function tick_all()
advance_time()
for _, callback in ipairs(timer_callbacks) do callback(false) end
end
checksum_reads = 0
json_reads = 0
for _ = 1, 3 do tick_all() end
assert_equal(checksum_reads, 3, "stable 15 seconds reads checksum once per interval")
assert_equal(json_reads, 0, "unchanged checksum never reads JSON")
disk_checksum = "v2"
disk_json = "valid-v2"
tick_all()
assert_equal(json_reads, 1, "changed snapshot JSON is read once across workers")
assert_equal(first.current().groups["2"].ip_list[1], "198.51.100.2", "first worker adopts v2")
assert_equal(second.current().groups["2"].ip_list[1], "198.51.100.2", "second worker adopts v2")
disk_checksum = "v3"
disk_json = "valid-v3"
tick_all()
assert_equal(shared_data.ip_groups_previous_version, "v2", "previous pointer follows committed version")
assert_equal(shared_data["ip_groups_raw:v1"], nil, "snapshot older than previous is cleaned")
assert_equal(shared_data["ip_groups_raw:v2"], "valid-v2", "previous committed raw is retained")
assert_equal(shared_data["ip_groups_raw:v3"], "valid-v3", "current committed raw is retained")
disk_checksum = "v2"
disk_json = "valid-v2"
tick_all()
assert_equal(shared_data.ip_groups_version, "v2", "rollback checksum becomes current commit")
assert_equal(shared_data.ip_groups_previous_version, "v3", "rollback retains former current as previous")
assert_equal(shared_data["ip_groups_raw:v2"], "valid-v2", "rollback must not clean its new current raw")
assert_equal(shared_data["ip_groups_raw:v3"], "valid-v3", "rollback retains previous raw")
disk_checksum = "v4"
disk_json = "invalid-v4"
tick_all()
assert_equal(shared_data.ip_groups_version, "v2", "invalid update preserves shared version")
assert_equal(first.current().groups["2"].ip_list[1], "198.51.100.2", "invalid update preserves first worker")
assert_equal(second.current().groups["2"].ip_list[1], "198.51.100.2", "invalid update preserves second worker")
local reads_before_requests = checksum_reads + json_reads
for _ = 1, 20 do
assert_equal(first.current().groups["2"].enabled, true, "request reads worker-local object")
end
assert_equal(checksum_reads + json_reads, reads_before_requests, "current() performs zero file I/O")
timer_callbacks[1](true)
assert_equal(checksum_reads + json_reads, reads_before_requests, "premature timer performs zero file I/O")
local function test_failed_commit_never_exposes_unpublished_raw_to_new_worker()
local data = {}
local held_locks = {}
local callbacks = {}
local checksum = "v1"
local raw = "valid-v1"
local reads = 0
local fail_commit = false
local interleaved_worker
local load_regression_worker
local regression_shared = {}
function regression_shared:get(key) return data[key] end
function regression_shared:add(key, value)
if held_locks[key] then return false end
held_locks[key] = value
return true
end
function regression_shared:delete(key) data[key] = nil return true end
local function set_regression_value(key, value)
if key == "ip_groups_version" and fail_commit then
return false, "shared dictionary full"
end
data[key] = value
if fail_commit and string.sub(key, 1, #"ip_groups_raw") == "ip_groups_raw" and not interleaved_worker then
interleaved_worker = load_regression_worker()
end
return true
end
function regression_shared:set(key, value) return set_regression_value(key, value) end
function regression_shared:safe_set(key, value) return set_regression_value(key, value) end
function regression_shared:safe_add(key, value) return regression_shared:add(key, value) end
load_regression_worker = function()
local worker = assert(loadfile(module_path))()
assert(worker.init({
shared = regression_shared,
timer_every = function(_, callback) callbacks[#callbacks + 1] = callback return true end,
read_checksum = function() return checksum end,
read_json = function() reads = reads + 1 return raw end,
decode = decode,
max_snapshot_bytes = 20 * 1024 * 1024,
}))
return worker
end
local established_worker = load_regression_worker()
assert_equal(established_worker.current().groups["1"].ip_list[1], "192.0.2.1", "v1 is committed before failure")
held_locks = {}
reads = 0
checksum = "v2"
raw = "valid-v2"
fail_commit = true
callbacks[1](false)
assert_equal(reads, 1, "failed commit still reads changed JSON only once")
assert_equal(data.ip_groups_version, "v1", "failed pointer write preserves committed version")
assert_equal(data["ip_groups_raw:v2"], nil, "failed commit cleans only unpublished v2 raw")
assert_equal(established_worker.current().groups["1"].ip_list[1], "192.0.2.1", "existing worker preserves committed v1")
assert(interleaved_worker, "raw publication must interleave a newly initialized worker")
assert_equal(interleaved_worker.current().groups["2"], nil, "new worker must not expose unpublished v2")
assert_equal(interleaved_worker.current().groups["1"].ip_list[1], "192.0.2.1", "new worker must never adopt unpublished v2 raw")
end
test_failed_commit_never_exposes_unpublished_raw_to_new_worker()
local function test_capacity_failure_never_evicts_committed_snapshot()
local data = {
ip_groups_version = "v1",
ip_groups_previous_version = "v0",
["ip_groups_raw:v1"] = "valid-v1",
["ip_groups_raw:v0"] = "valid-v0",
}
local locks = {}
local callbacks = {}
local disk_checksum = "v1"
local disk_raw = "valid-v1"
local json_reads = 0
local ordinary_writes = 0
local warnings = {}
local dict = {}
function dict:get(key) return data[key] end
function dict:delete(key) data[key] = nil return true end
function dict:add(key, value)
if locks[key] then return false end
locks[key] = value
return true
end
function dict:safe_add(key, value) return dict:add(key, value) end
function dict:set(key, value)
ordinary_writes = ordinary_writes + 1
if key == "ip_groups_raw:v2" then
data = { [key] = value }
return true, nil, true
end
data[key] = value
return true, nil, false
end
function dict:safe_set(key, value)
if key == "ip_groups_raw:v2" then return nil, "no memory", false end
data[key] = value
return true, nil, false
end
local worker = assert(loadfile(module_path))()
assert(worker.init({
shared = dict,
timer_every = function(_, callback) callbacks[1] = callback return true end,
read_checksum = function() return disk_checksum end,
read_json = function() json_reads = json_reads + 1 return disk_raw end,
decode = decode,
max_snapshot_bytes = 20 * 1024 * 1024,
log_warning = function(message) warnings[#warnings + 1] = message end,
}))
assert_equal(worker.current().groups["1"].ip_list[1], "192.0.2.1", "worker starts from committed v1")
locks = {}
disk_checksum = "v2"
disk_raw = "valid-v2"
callbacks[1](false)
assert_equal(ordinary_writes, 0, "snapshot publication must never use evicting set")
assert_equal(json_reads, 1, "capacity failure reads changed JSON once")
assert_equal(data.ip_groups_version, "v1", "capacity failure preserves commit pointer")
assert_equal(data.ip_groups_previous_version, "v0", "capacity failure preserves previous metadata")
assert_equal(data["ip_groups_raw:v1"], "valid-v1", "capacity failure preserves current raw")
assert_equal(data["ip_groups_raw:v0"], "valid-v0", "capacity failure preserves previous raw")
assert_equal(data["ip_groups_raw:v2"], nil, "capacity failure does not publish new raw")
assert_equal(worker.current().groups["1"].ip_list[1], "192.0.2.1", "capacity failure preserves worker-local snapshot")
assert_equal(#warnings, 1, "capacity failure is logged")
end
local function test_previous_metadata_failure_keeps_committed_snapshot_without_cleanup()
local data = {
ip_groups_version = "v1",
ip_groups_previous_version = "v0",
["ip_groups_raw:v1"] = "valid-v1",
["ip_groups_raw:v0"] = "valid-v0",
}
local locks = {}
local callback
local checksum = "v1"
local raw = "valid-v1"
local deletes = 0
local warnings = {}
local dict = {}
function dict:get(key) return data[key] end
function dict:delete(key) deletes = deletes + 1 data[key] = nil return true end
function dict:add(key, value)
if locks[key] then return false end
locks[key] = value
return true
end
function dict:safe_add(key, value) return dict:add(key, value) end
function dict:set(key, value) data[key] = value return true end
function dict:safe_set(key, value)
if key == "ip_groups_previous_version" then return nil, "no memory", false end
data[key] = value
return true, nil, false
end
local worker = assert(loadfile(module_path))()
assert(worker.init({
shared = dict,
timer_every = function(_, value) callback = value return true end,
read_checksum = function() return checksum end,
read_json = function() return raw end,
decode = decode,
max_snapshot_bytes = 20 * 1024 * 1024,
log_warning = function(message) warnings[#warnings + 1] = message end,
}))
locks = {}
checksum = "v2"
raw = "valid-v2"
callback(false)
assert_equal(data.ip_groups_version, "v2", "successful commit pointer remains authoritative")
assert_equal(data.ip_groups_previous_version, "v0", "failed previous metadata write is not forced")
assert_equal(data["ip_groups_raw:v2"], "valid-v2", "new committed raw remains")
assert_equal(data["ip_groups_raw:v1"], "valid-v1", "old current raw remains when cleanup is skipped")
assert_equal(data["ip_groups_raw:v0"], "valid-v0", "old previous raw remains when cleanup is skipped")
assert_equal(deletes, 0, "previous metadata failure skips all cleanup")
assert_equal(worker.current().groups["2"].ip_list[1], "198.51.100.2", "worker adopts valid committed v2")
assert_equal(#warnings, 1, "previous metadata failure is logged")
end
local function test_oversized_raw_is_rejected_before_shared_publication()
local data = { ip_groups_version = "v1", ["ip_groups_raw:v1"] = "valid-v1" }
local locks = {}
local callback
local checksum = "v1"
local raw = "valid-v1"
local shared_writes = 0
local warnings = {}
local dict = {}
function dict:get(key) return data[key] end
function dict:delete(key) data[key] = nil return true end
function dict:add(key, value) if locks[key] then return false end locks[key] = value return true end
function dict:safe_add(key, value) return dict:add(key, value) end
function dict:set(key, value) shared_writes = shared_writes + 1 data[key] = value return true end
function dict:safe_set(key, value) shared_writes = shared_writes + 1 data[key] = value return true, nil, false end
local worker = assert(loadfile(module_path))()
assert(worker.init({
shared = dict,
timer_every = function(_, value) callback = value return true end,
read_checksum = function() return checksum end,
read_json = function() return raw end,
decode = decode,
max_snapshot_bytes = 4,
log_warning = function(message) warnings[#warnings + 1] = message end,
}))
locks = {}
checksum = "v2"
raw = "valid-v2"
callback(false)
assert_equal(shared_writes, 0, "oversized raw is rejected before shared writes")
assert_equal(data.ip_groups_version, "v1", "oversized raw preserves commit pointer")
assert_equal(data["ip_groups_raw:v1"], "valid-v1", "oversized raw preserves committed data")
assert_equal(worker.current().groups["1"].ip_list[1], "192.0.2.1", "oversized raw preserves worker-local snapshot")
assert_equal(#warnings, 1, "oversized raw rejection is logged")
end
test_capacity_failure_never_evicts_committed_snapshot()
test_previous_metadata_failure_keeps_committed_snapshot_without_cleanup()
test_oversized_raw_is_rejected_before_shared_publication()
return true
+404
View File
@@ -0,0 +1,404 @@
local _M = {}
local rules_config
local ip_groups_config
local ip_groups_runtime
local pow_runtime
local geo_lookup
local geo_module
local geo_profiles = { city = false, country = false }
local function read_file(path)
local file, err = io.open(path, "r")
if not file then
return nil, err
end
local content = file:read("*a")
file:close()
return content
end
local function load_json(path)
local content, err = read_file(path)
if not content or content == "" then
return nil, err or "empty file"
end
local decoded, decode_err = require("cjson.safe").decode(content)
if not decoded then
return nil, decode_err or "invalid JSON"
end
return decoded
end
local function warn_rate_limited(key, ...)
local dict = ngx.shared and ngx.shared.openflare_waf_config
if not dict or not dict.add or dict:add(key, true, 60) then
ngx.log(ngx.WARN, ...)
end
end
local function file_exists(path)
local file = io.open(path, "rb")
if not file then return false end
file:close()
return true
end
local function init_geo_databases(country_path, city_path, path_exists, region_required)
local ok, module_or_error = pcall(require, "resty.maxminddb")
if not ok or not module_or_error then
warn_rate_limited("_geo_module_unavailable", "openflare waf GeoIP module unavailable: ", module_or_error)
return
end
geo_module = module_or_error
local profiles = {}
if path_exists(city_path) then profiles.city = city_path end
if path_exists(country_path) then profiles.country = country_path end
if not profiles.city and region_required then
warn_rate_limited("_geo_city_unavailable", "openflare waf GeoLite2 City database unavailable; region match takes false branch")
end
if not profiles.country and not profiles.city then
warn_rate_limited("_geo_database_unavailable", "openflare waf GeoIP databases unavailable")
return
end
local function initialize_profile(profile, path)
local called, init_result, init_error = pcall(geo_module.init, { [profile] = path })
if not called or init_result ~= true then
return false, init_error or init_result
end
geo_profiles[profile] = true
return true
end
local city_initialized, city_error = false, nil
if profiles.city then
city_initialized, city_error = initialize_profile("city", profiles.city)
if not city_initialized and region_required then
warn_rate_limited("_geo_city_unavailable", "openflare waf GeoLite2 City database initialization failed; region match takes false branch: ", city_error)
end
end
local country_initialized, country_error = false, nil
if profiles.country then
country_initialized, country_error = initialize_profile("country", profiles.country)
end
if not city_initialized and not country_initialized then
warn_rate_limited("_geo_database_unavailable", "openflare waf GeoIP database initialization failed: ", country_error or city_error)
end
end
local function lookup_geo_profile(ip, profile)
if not geo_module or not geo_profiles[profile] then return nil end
local ok, result, lookup_error = pcall(geo_module.lookup, ip, nil, profile)
if not ok or not result then
warn_rate_limited("_geo_lookup_failed_" .. profile, "openflare waf GeoIP ", profile, " lookup failed: ", lookup_error or result)
return nil
end
return result
end
local function default_geo_lookup(ip, region_required)
local result = lookup_geo_profile(ip, "city")
local from_city = result ~= nil
if not result then result = lookup_geo_profile(ip, "country") end
if not result then return nil, nil end
local country = result.country and result.country.iso_code or nil
local subdivision
if from_city then
subdivision = result.most_specific_subdivision and result.most_specific_subdivision.iso_code or nil
if not subdivision and result.subdivisions and result.subdivisions[1] then
subdivision = result.subdivisions[1].iso_code
end
elseif region_required then
warn_rate_limited("_geo_city_unavailable", "openflare waf GeoLite2 City database unavailable; region match takes false branch")
end
country = country and string.upper(country) or nil
subdivision = subdivision and string.upper(subdivision) or nil
local region = subdivision
if country and subdivision and not string.match(subdivision, "^[A-Z][A-Z]%-") then
region = country .. "-" .. subdivision
end
return country, region
end
local function config_geo_requirements(config)
local uses_geo, uses_region = false, false
for _, rule in ipairs(config.rule_groups or {}) do
for _, node in pairs((rule.graph or {}).nodes or {}) do
if node.type == "geo_match" then
uses_geo = true
local node_config = node.config or {}
if type(node_config.regions) == "table" and #node_config.regions > 0 then uses_region = true end
end
end
end
return uses_geo, uses_region
end
function _M.init(options)
if rules_config then
return true
end
options = options or {}
local runtime_dir = options.runtime_dir or "__OPENFLARE_RUNTIME_CONFIG_DIR__"
if options.config then
rules_config = options.config
else
local err
rules_config, err = load_json(runtime_dir .. "/waf_config.json")
assert(rules_config, "load waf_config.json failed: " .. tostring(err))
end
if options.ip_groups then
ip_groups_config = options.ip_groups
else
ip_groups_runtime = options.ip_groups_runtime or require("waf.ip_groups")
local initialized, init_error = ip_groups_runtime.init({ runtime_dir = runtime_dir })
assert(initialized, "initialize WAF IP groups failed: " .. tostring(init_error))
end
pow_runtime = options.pow or require("pow.runtime")
if options.geo_lookup then
geo_lookup = options.geo_lookup
else
local uses_geo, uses_region = config_geo_requirements(rules_config)
if uses_geo then
init_geo_databases(
options.country_mmdb_path or "__OPENFLARE_COUNTRY_MMDB_PATH__",
options.city_mmdb_path or "__OPENFLARE_CITY_MMDB_PATH__",
options.geo_file_exists or file_exists,
uses_region
)
end
geo_lookup = default_geo_lookup
end
return true
end
-- Task 7 can atomically replace the worker-local IP group snapshot through this seam.
function _M.replace_ip_groups(snapshot)
ip_groups_config = snapshot or { groups = {} }
end
local function list_contains(items, value)
if type(items) ~= "table" or not value then return false end
value = string.upper(value)
for _, item in ipairs(items) do
if string.upper(tostring(item)) == value then return true end
end
return false
end
local function parse_ipv4(value)
local a, b, c, d = string.match(value or "", "^(%d+)%.(%d+)%.(%d+)%.(%d+)$")
if not a then return nil end
a, b, c, d = tonumber(a), tonumber(b), tonumber(c), tonumber(d)
if a > 255 or b > 255 or c > 255 or d > 255 then return nil end
return ((a * 256 + b) * 256 + c) * 256 + d
end
local function split_ipv6_side(value)
local result = {}
if value == "" then return result end
for part in string.gmatch(value, "[^:]+") do
if string.find(part, ".", 1, true) then
local ipv4 = parse_ipv4(part)
if not ipv4 then return nil end
result[#result + 1] = math.floor(ipv4 / 65536)
result[#result + 1] = ipv4 % 65536
else
if #part > 4 or not string.match(part, "^[%x]+$") then return nil end
local number = tonumber(part, 16)
if not number or number > 65535 then return nil end
result[#result + 1] = number
end
end
return result
end
local function parse_ipv6(value)
value = string.lower(value or "")
local compressed_at = string.find(value, "::", 1, true)
if compressed_at and string.find(value, "::", compressed_at + 2, true) then return nil end
local left, right
if compressed_at then
left = split_ipv6_side(string.sub(value, 1, compressed_at - 1))
right = split_ipv6_side(string.sub(value, compressed_at + 2))
else
if string.sub(value, 1, 1) == ":" or string.sub(value, -1) == ":" then return nil end
left, right = split_ipv6_side(value), {}
end
if not left or not right then return nil end
local missing = 8 - #left - #right
if (compressed_at and missing < 1) or (not compressed_at and missing ~= 0) then return nil end
local result = {}
for _, number in ipairs(left) do result[#result + 1] = number end
for _ = 1, missing do result[#result + 1] = 0 end
for _, number in ipairs(right) do result[#result + 1] = number end
if #result ~= 8 then return nil end
return result
end
local function ipv6_equal(left, right)
left, right = parse_ipv6(left), parse_ipv6(right)
if not left or not right then return false end
for index = 1, 8 do
if left[index] ~= right[index] then return false end
end
return true
end
local function ip_in_cidr(ip, cidr)
local base, bits = string.match(cidr or "", "^([^/]+)/(%d+)$")
bits = tonumber(bits)
if not base or not bits then return false end
local ip_number, base_number = parse_ipv4(ip), parse_ipv4(base)
if ip_number and base_number then
if bits < 0 or bits > 32 then return false end
if bits == 0 then return true end
local size = 2 ^ (32 - bits)
return ip_number - (ip_number % size) == base_number - (base_number % size)
end
local ip_groups, base_groups = parse_ipv6(ip), parse_ipv6(base)
if not ip_groups or not base_groups or bits < 0 or bits > 128 then return false end
local full_groups, remaining_bits = math.floor(bits / 16), bits % 16
for index = 1, full_groups do
if ip_groups[index] ~= base_groups[index] then return false end
end
if remaining_bits > 0 then
local size = 2 ^ (16 - remaining_bits)
local index = full_groups + 1
if math.floor(ip_groups[index] / size) ~= math.floor(base_groups[index] / size) then return false end
end
return true
end
local function matches_ip_values(config, ip)
for _, item in ipairs(config.ips or {}) do
if item == ip or ipv6_equal(item, ip) then return true end
end
for _, cidr in ipairs(config.cidrs or {}) do
if ip_in_cidr(ip, cidr) then return true end
end
local snapshot = ip_groups_config or ip_groups_runtime.current()
local groups = (snapshot or {}).groups or {}
for _, id in ipairs(config.ip_group_ids or {}) do
local group = groups[tostring(id)]
if group and group.enabled then
for _, item in ipairs(group.ip_list or {}) do
if item == ip or ipv6_equal(item, ip) or ip_in_cidr(ip, item) then return true end
end
end
end
return false
end
local function fail_closed(reason)
local dict = ngx.shared and ngx.shared.openflare_waf_config
if not dict or not dict.add or dict:add("_damaged_graph_logged", true, 60) then
ngx.log(ngx.ERR, "openflare waf damaged runtime graph: ", reason)
end
ngx.ctx.openflare_waf_blocked = true
ngx.status = 500
ngx.header["Content-Type"] = "text/plain; charset=utf-8"
ngx.say("OpenFlare WAF runtime error")
return ngx.exit(500)
end
local function render_block(config)
config = config or {}
local status = tonumber(config.status_code) or 403
ngx.ctx.openflare_waf_blocked = true
ngx.status = status
local body = config.response_body or ""
if body ~= "" then
ngx.header["Content-Type"] = "text/html; charset=utf-8"
ngx.say(body)
end
return ngx.exit(status)
end
local function execute_graph(graph)
if type(graph) ~= "table" or type(graph.nodes) ~= "table" or type(graph.entry) ~= "string" then
return nil, "invalid graph"
end
local node_count = 0
for _ in pairs(graph.nodes) do node_count = node_count + 1 end
local current = graph.entry
for _ = 1, node_count do
local node = graph.nodes[current]
if type(node) ~= "table" or type(node.type) ~= "string" then
return nil, "missing node " .. tostring(current)
end
if node.type == "allow" then
return { kind = "allow" }
end
if node.type == "block" then
return { kind = "block", config = node.config }
end
local handle
if node.type == "start" then
handle = "next"
elseif node.type == "ip_match" then
handle = matches_ip_values(node.config or {}, ngx.var.remote_addr or "") and "true" or "false"
elseif node.type == "geo_match" then
local config = node.config or {}
local region_required = type(config.regions) == "table" and #config.regions > 0
local country, region = geo_lookup(ngx.var.remote_addr or "", region_required)
handle = (list_contains(config.countries, country) or list_contains(config.regions, region)) and "true" or "false"
elseif node.type == "pow" then
if pow_runtime.evaluate(node.config or {}) ~= true then
return { kind = "takeover" }
end
handle = "next"
else
return nil, "unknown node type " .. node.type
end
if type(node.next) ~= "table" or type(node.next[handle]) ~= "string" then
return nil, "missing " .. handle .. " edge from " .. current
end
current = node.next[handle]
end
return nil, "graph exceeded maximum steps"
end
local function active_rules(site)
local by_id, result = {}, {}
for _, rule in ipairs(rules_config.rule_groups or {}) do
by_id[tostring(rule.id)] = rule
if rule.enabled and rule.is_global then result[#result + 1] = rule end
end
for _, binding in ipairs(rules_config.bindings or {}) do
if binding.site_name == site then
for _, id in ipairs(binding.rule_group_ids or {}) do
local rule = by_id[tostring(id)]
if rule and rule.enabled and not rule.is_global then result[#result + 1] = rule end
end
break
end
end
return result
end
local function is_internal_pow_continuation()
if not ngx.req or not ngx.req.is_internal or not ngx.req.is_internal() then return false end
local uri = ngx.var.uri or ""
local api_prefix = "/.within.website/x/cmd/anubis/api/"
local static_prefix = "/.within.website/x/cmd/anubis/static/"
return string.sub(uri, 1, #api_prefix) == api_prefix or string.sub(uri, 1, #static_prefix) == static_prefix
end
function _M.check()
if not rules_config then
return fail_closed("runtime not initialized")
end
if is_internal_pow_continuation() then
ngx.ctx.openflare_pow_takeover = true
return
end
for _, rule in ipairs(active_rules(ngx.var.openflare_waf_site or "")) do
local decision, err = execute_graph(rule.graph)
if not decision then return fail_closed(err) end
if decision.kind == "block" then return render_block(decision.config) end
if decision.kind == "takeover" then return end
end
return "ok"
end
return _M
@@ -0,0 +1,535 @@
local runtime_path = assert(WAF_RUNTIME_PATH, "WAF_RUNTIME_PATH is required")
local function assert_equal(actual, expected, message)
if actual ~= expected then
error((message or "values differ") .. ": expected " .. tostring(expected) .. ", got " .. tostring(actual), 2)
end
end
local output
local pow_calls
local pow_results
local shared_keys = {}
local logs = {}
ngx = {
WARN = "WARN",
ERR = "ERR",
var = {},
ctx = {},
header = {},
shared = {
openflare_waf_config = {
add = function(_, key)
if shared_keys[key] then return false end
shared_keys[key] = true
return true
end,
},
},
req = { is_internal = function() return ngx.var.openflare_internal == true end },
say = function(body) output.body = body end,
exit = function(status) output.exit = status return status end,
log = function(_, ...)
local parts = { ... }
for index, value in ipairs(parts) do parts[index] = tostring(value) end
output.log = table.concat(parts)
logs[#logs + 1] = output.log
end,
}
local pow_stub = {}
function pow_stub.evaluate(config)
pow_calls[#pow_calls + 1] = config.difficulty
local result = pow_results[1]
table.remove(pow_results, 1)
return result
end
local function node(node_type, config, next_nodes)
return { type = node_type, config = config or {}, next = next_nodes }
end
local function graph(nodes, entry)
return { entry = entry or "start", nodes = nodes }
end
local function rule(id, is_global, rule_graph)
return { id = id, enabled = true, is_global = is_global or false, graph = rule_graph }
end
local function start_to(target)
return node("start", {}, { next = target })
end
local function load_runtime(config, options)
local chunk = assert(loadfile(runtime_path))
local runtime = chunk()
options = options or {}
runtime.init({
config = config,
ip_groups = options.ip_groups or { groups = {} },
pow = pow_stub,
geo_lookup = options.geo_lookup,
runtime_dir = options.runtime_dir,
geo_file_exists = options.geo_file_exists,
country_mmdb_path = options.country_mmdb_path or (options.runtime_dir and (options.runtime_dir .. "/GeoLite2-Country.mmdb") or nil),
city_mmdb_path = options.city_mmdb_path or (options.runtime_dir and (options.runtime_dir .. "/GeoLite2-City.mmdb") or nil),
})
return runtime
end
local function reset_request(site, ip, uri, is_internal)
ngx.var = { openflare_waf_site = site, remote_addr = ip or "192.0.2.1", uri = uri or "/", request_id = "request-1", openflare_internal = is_internal == true }
ngx.ctx = {}
ngx.header = {}
output = {}
pow_calls = {}
pow_results = {}
end
local function binding(site, ids)
return { site_name = site, rule_group_ids = ids }
end
local function test_ip_true_and_false()
local config = {
rule_groups = { rule(1, false, graph({
start = start_to("match"),
match = node("ip_match", { ips = { "192.0.2.1" }, cidrs = { "198.51.100.0/24" }, ip_group_ids = { 7 } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 451, response_body = "ip blocked" }),
allow = node("allow"),
})) },
bindings = { binding("ip-site", { 1 }) },
}
local runtime = load_runtime(config, { ip_groups = { groups = { ["7"] = { enabled = true, ip_list = { "203.0.113.7" } } } } })
reset_request("ip-site", "192.0.2.1")
runtime.check()
assert_equal(output.exit, 451, "exact IP true branch")
reset_request("ip-site", "198.51.100.8")
runtime.check()
assert_equal(output.exit, 451, "CIDR true branch")
reset_request("ip-site", "203.0.113.7")
runtime.check()
assert_equal(output.exit, 451, "IP group true branch")
reset_request("ip-site", "203.0.113.8")
runtime.check()
assert_equal(output.exit, nil, "IP false branch")
end
local function test_ipv6_exact_cidr_and_group()
local config = {
rule_groups = { rule(8, false, graph({
start = start_to("match"),
match = node("ip_match", { ips = { "2001:db8::1" }, cidrs = { "2001:db8:abcd::/48" }, ip_group_ids = { 9 } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 451, response_body = "ipv6 blocked" }),
allow = node("allow"),
})) },
bindings = { binding("ipv6-site", { 8 }) },
}
local runtime = load_runtime(config, { ip_groups = { groups = { ["9"] = { enabled = true, ip_list = { "2001:db8:ffff::/48" } } } } })
reset_request("ipv6-site", "2001:0db8:0:0:0:0:0:1")
runtime.check()
assert_equal(output.exit, 451, "canonical-equivalent IPv6 exact match")
reset_request("ipv6-site", "2001:db8:abcd:12::9")
runtime.check()
assert_equal(output.exit, 451, "IPv6 CIDR true branch")
reset_request("ipv6-site", "2001:db8:ffff:beef::9")
runtime.check()
assert_equal(output.exit, 451, "IP group IPv6 CIDR true branch")
reset_request("ipv6-site", "2001:db9::1")
runtime.check()
assert_equal(output.exit, nil, "IPv6 false branch")
end
local function test_geo_true_and_false()
local config = {
rule_groups = { rule(2, false, graph({
start = start_to("geo"),
geo = node("geo_match", { countries = { "US" }, regions = { "DE-BE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 403, response_body = "geo blocked" }),
allow = node("allow"),
})) },
bindings = { binding("geo-site", { 2 }) },
}
local country, region = "US", "NY"
local runtime = load_runtime(config, { geo_lookup = function() return country, region end })
reset_request("geo-site")
runtime.check()
assert_equal(output.exit, 403, "country true branch")
country, region = "DE", "DE-BE"
reset_request("geo-site")
runtime.check()
assert_equal(output.exit, 403, "region true branch")
country, region = "DE", "BE"
reset_request("geo-site")
runtime.check()
assert_equal(output.exit, nil, "geo false branch")
end
local function test_geo_module_is_initialized_once_and_composes_region()
local init_calls, lookup_calls = 0, 0
local initialized_profiles = {}
package.loaded["resty.maxminddb"] = nil
package.preload["resty.maxminddb"] = function()
return {
init = function(profiles)
init_calls = init_calls + 1
for profile, path in pairs(profiles) do initialized_profiles[profile] = path end
return true
end,
has_profile = function(profile) return initialized_profiles[profile] ~= nil end,
lookup = function(_, _, profile)
lookup_calls = lookup_calls + 1
assert_equal(profile, "city", "subdivision lookup uses City profile")
return { country = { iso_code = "US" }, subdivisions = { { iso_code = "CA" } } }
end,
}
end
local config = {
rule_groups = { rule(12, false, graph({
start = start_to("geo"),
geo = node("geo_match", { regions = { "US-CA" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 403 }),
allow = node("allow"),
})) },
bindings = { binding("geo-cache", { 12 }) },
}
local runtime = load_runtime(config, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
assert_equal(init_calls, 2, "each MaxMind profile initializes independently during worker init")
assert_equal(initialized_profiles.city, "/runtime/GeoLite2-City.mmdb", "City profile path")
assert_equal(initialized_profiles.country, "/runtime/GeoLite2-Country.mmdb", "Country profile path")
for _ = 1, 3 do
reset_request("geo-cache")
runtime.check()
assert_equal(output.exit, 403, "MaxMind subdivision composes validator-compatible region")
end
assert_equal(init_calls, 2, "MaxMind database is not initialized on requests")
assert_equal(lookup_calls, 3, "requests only perform lookup")
end
local function test_geo_country_fallback_does_not_fake_region()
local profiles
package.loaded["resty.maxminddb"] = nil
package.preload["resty.maxminddb"] = function()
return {
init = function(value) profiles = value return true end,
has_profile = function(profile) return profiles[profile] ~= nil end,
lookup = function(_, _, profile)
assert_equal(profile, "country", "fallback lookup uses Country profile")
return { country = { iso_code = "US" }, subdivisions = { { iso_code = "CA" } } }
end,
}
end
shared_keys = {}
logs = {}
local country_graph = graph({
start = start_to("geo"),
geo = node("geo_match", { countries = { "US" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 403 }), allow = node("allow"),
})
local region_graph = graph({
start = start_to("geo"),
geo = node("geo_match", { regions = { "US-CA" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 451 }), allow = node("allow"),
})
local runtime = load_runtime({
rule_groups = { rule(15, false, country_graph), rule(16, false, region_graph) },
bindings = { binding("country-only", { 15 }), binding("region-without-city", { 16 }) },
}, {
runtime_dir = "/runtime",
geo_file_exists = function(path) return string.find(path, "Country", 1, true) ~= nil end,
})
reset_request("country-only")
runtime.check()
assert_equal(output.exit, 403, "Country fallback remains available")
reset_request("region-without-city")
runtime.check()
assert_equal(output.exit, nil, "Country subdivisions must not satisfy region")
runtime.check()
assert_equal(#logs, 1, "missing City warning is rate limited")
end
local function test_geo_city_init_failure_retries_country_profile()
local init_calls = {}
local profiles = {}
package.loaded["resty.maxminddb"] = nil
package.preload["resty.maxminddb"] = function()
return {
init = function(value)
init_calls[#init_calls + 1] = value
if value.city then return false end
profiles = value
return true
end,
has_profile = function(profile) return profiles[profile] ~= nil end,
lookup = function(_, _, profile)
assert_equal(profile, "country", "corrupt City fallback uses Country")
return { country = { iso_code = "DE" } }
end,
}
end
shared_keys = {}
logs = {}
local runtime = load_runtime({
rule_groups = { rule(17, false, graph({
start = start_to("geo"),
geo = node("geo_match", { countries = { "DE" }, regions = { "DE-BE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 403 }), allow = node("allow"),
})) },
bindings = { binding("corrupt-city", { 17 }) },
}, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
reset_request("corrupt-city")
runtime.check()
assert_equal(#init_calls, 2, "Country profile is retried after City profile init failure")
assert_equal(output.exit, 403, "Country remains available after corrupt City init")
end
local function test_geo_partial_init_never_looks_up_corrupt_city()
local opened = {}
local lookups = {}
package.loaded["resty.maxminddb"] = nil
package.preload["resty.maxminddb"] = function()
return {
init = function(profiles)
if profiles.country then opened.country = true end
if profiles.city then return nil, "corrupt City" end
return true
end,
initted = function() return next(opened) ~= nil end,
lookup = function(_, _, profile)
lookups[#lookups + 1] = profile
assert_equal(opened[profile], true, "lookup must only use an opened profile")
return { country = { iso_code = "DE" } }
end,
}
end
local runtime = load_runtime({
rule_groups = { rule(18, false, graph({
start = start_to("geo"),
geo = node("geo_match", { countries = { "DE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 403 }), allow = node("allow"),
})) },
bindings = { binding("partial-corrupt-city", { 18 }) },
}, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
reset_request("partial-corrupt-city")
runtime.check()
assert_equal(table.concat(lookups, ","), "country", "corrupt City is never looked up")
assert_equal(output.exit, 403, "valid Country remains available")
end
local function test_geo_partial_init_never_looks_up_corrupt_country()
local opened = {}
local lookups = {}
package.loaded["resty.maxminddb"] = nil
package.preload["resty.maxminddb"] = function()
return {
init = function(profiles)
if profiles.city then opened.city = true end
if profiles.country then return nil, "corrupt Country" end
return true
end,
initted = function() return next(opened) ~= nil end,
lookup = function(_, _, profile)
lookups[#lookups + 1] = profile
assert_equal(opened[profile], true, "lookup must only use an opened profile")
return nil, "address absent"
end,
}
end
local runtime = load_runtime({
rule_groups = { rule(19, false, graph({
start = start_to("geo"),
geo = node("geo_match", { countries = { "DE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 403 }), allow = node("allow"),
})) },
bindings = { binding("partial-corrupt-country", { 19 }) },
}, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
reset_request("partial-corrupt-country")
runtime.check()
assert_equal(table.concat(lookups, ","), "city", "corrupt Country is never used as fallback")
assert_equal(output.exit, nil, "missing City result takes false branch without corrupt fallback")
end
local function test_geo_unavailable_warning_is_rate_limited()
package.loaded["resty.maxminddb"] = nil
package.preload["resty.maxminddb"] = function() error("module unavailable") end
shared_keys = {}
logs = {}
local config = {
rule_groups = { rule(13, false, graph({
start = start_to("geo"),
geo = node("geo_match", { countries = { "US" } }, { ["true"] = "blocked", ["false"] = "allow" }),
blocked = node("block", { status_code = 403 }),
allow = node("allow"),
})) },
bindings = { binding("geo-missing", { 13 }) },
}
local first = load_runtime(config)
local second = load_runtime(config)
reset_request("geo-missing")
first.check()
second.check()
assert_equal(#logs, 1, "missing MaxMind warning is rate limited across workers")
end
local function test_pow_takeover_and_completion()
local config = {
rule_groups = { rule(3, false, graph({
start = start_to("pow"),
pow = node("pow", { algorithm = "fast", difficulty = 5, session_ttl = 600, challenge_ttl = 300 }, { next = "blocked" }),
blocked = node("block", { status_code = 429, response_body = "after pow" }),
allow = node("allow"),
})) },
bindings = { binding("pow-site", { 3 }) },
}
local runtime = load_runtime(config)
reset_request("pow-site")
pow_results = { false }
runtime.check()
assert_equal(output.exit, nil, "PoW takeover must stop graph execution")
assert_equal(#pow_calls, 1, "PoW evaluated once")
reset_request("pow-site")
pow_results = { true }
runtime.check()
assert_equal(output.exit, 429, "completed PoW follows next edge")
end
local function test_pow_internal_redirect_bypasses_graph_as_takeover()
local config = {
rule_groups = { rule(14, false, graph({
start = start_to("pow"),
pow = node("pow", { difficulty = 4 }, { next = "blocked" }),
blocked = node("block", { status_code = 429 }),
allow = node("allow"),
})) },
bindings = { binding("pow-internal", { 14 }) },
}
local runtime = load_runtime(config)
reset_request("pow-internal", "192.0.2.1", "/.within.website/x/cmd/anubis/api/make-challenge", true)
pow_results = { true }
runtime.check()
assert_equal(#pow_calls, 0, "internal challenge continuation must not re-enter DAG")
assert_equal(output.exit, nil, "internal challenge continuation must not follow pow next")
end
local function test_block_config_and_rule_order()
local function pow_allow(difficulty)
return graph({
start = start_to("pow"),
pow = node("pow", { algorithm = "fast", difficulty = difficulty, session_ttl = 600, challenge_ttl = 300 }, { next = "allow" }),
allow = node("allow"),
})
end
local config = {
rule_groups = {
rule(10, true, pow_allow(10)),
rule(20, false, pow_allow(20)),
rule(30, false, pow_allow(30)),
rule(40, false, graph({
start = start_to("blocked"),
blocked = node("block", { status_code = 418, response_body = "custom block" }),
allow = node("allow"),
})),
},
bindings = { binding("ordered-site", { 30, 20, 40 }) },
}
local runtime = load_runtime(config)
reset_request("ordered-site")
pow_results = { true, true, true }
runtime.check()
assert_equal(table.concat(pow_calls, ","), "10,30,20", "global rule precedes binding order")
assert_equal(output.exit, 418, "block status comes from reached block node")
assert_equal(output.body, "custom block", "block body comes from reached block node")
assert_equal(ngx.header["Content-Type"], "text/html; charset=utf-8", "block content type")
end
local function test_damaged_graphs_fail_closed()
local configs = {
graph({ start = start_to("unknown"), unknown = node("future_node"), allow = node("allow") }),
graph({ start = start_to("missing"), allow = node("allow") }),
graph({ start = start_to("loop"), loop = node("start", {}, { next = "loop" }), allow = node("allow") }),
}
for index, damaged in ipairs(configs) do
local runtime = load_runtime({ rule_groups = { rule(index, false, damaged) }, bindings = { binding("damaged", { index }) } })
reset_request("damaged")
runtime.check()
assert_equal(output.exit, 500, "damaged graph " .. index .. " must fail closed")
end
end
local function test_request_path_has_no_file_io()
local opens = 0
local original_open = io.open
io.open = function(path, mode)
opens = opens + 1
local value = path:match("waf_ip_groups%.json$") and "IP_GROUPS" or "CONFIG"
return {
read = function() return value end,
close = function() end,
}
end
package.loaded["cjson.safe"] = nil
package.preload["cjson.safe"] = function()
return { decode = function(value)
if value == "IP_GROUPS" then return { groups = {} } end
return {
rule_groups = { rule(1, false, graph({ start = start_to("allow"), allow = node("allow") })) },
bindings = { binding("io-site", { 1 }) },
}
end }
end
local chunk = assert(loadfile(runtime_path))
local runtime = chunk()
runtime.init({
runtime_dir = "/runtime",
pow = pow_stub,
ip_groups_runtime = {
init = function() return true end,
current = function() return { groups = {} } end,
},
})
local init_opens = opens
assert_equal(init_opens, 1, "WAF graph initializes once; IP groups are owned by refresh module")
reset_request("io-site")
for _ = 1, 3 do runtime.check() end
assert_equal(opens, init_opens, "request execution performs no file I/O")
io.open = original_open
end
test_ip_true_and_false()
test_ipv6_exact_cidr_and_group()
test_geo_true_and_false()
test_geo_module_is_initialized_once_and_composes_region()
test_geo_country_fallback_does_not_fake_region()
test_geo_city_init_failure_retries_country_profile()
test_geo_partial_init_never_looks_up_corrupt_city()
test_geo_partial_init_never_looks_up_corrupt_country()
test_geo_unavailable_warning_is_rate_limited()
test_pow_takeover_and_completion()
test_pow_internal_redirect_bypasses_graph_as_takeover()
test_block_config_and_rule_order()
test_damaged_graphs_fail_closed()
test_request_path_has_no_file_io()
return true
+54 -25
View File
@@ -9,6 +9,7 @@ import (
"fmt"
"log/slog"
"sort"
"strconv"
"strings"
"sync"
@@ -42,7 +43,8 @@ type NginxManager interface {
EnsureSafeFallbackRuntime(ctx context.Context, reason string) error
CurrentChecksum() (string, error)
WAFIPGroupChecksums() (map[string]string, error)
SyncWAFIPGroups(groups []protocol.WAFIPGroup) error
ReconcileWAFIPGroups(targetIDs []uint, changed []protocol.WAFIPGroup) error
UpdateExistingWAFIPGroups(changed []protocol.WAFIPGroup) error
EnsureWorkerReadAccess() error
}
@@ -132,12 +134,12 @@ func (s *Service) WAFIPGroupChecksums() (map[string]string, error) {
return s.nginxManager.WAFIPGroupChecksums()
}
// ApplyWAFIPGroups writes the given WAF IP groups to the nginx manager.
// ApplyWAFIPGroups applies real-time changes only to groups already in the local authoritative snapshot.
func (s *Service) ApplyWAFIPGroups(_ context.Context, groups []protocol.WAFIPGroup) error {
if len(groups) == 0 || s.nginxManager == nil {
return nil
}
return s.nginxManager.SyncWAFIPGroups(groups)
return s.nginxManager.UpdateExistingWAFIPGroups(groups)
}
func (s *Service) applyIfNeeded(ctx context.Context, mode string, startup bool, snapshot *state.Snapshot, currentChecksum string, target *protocol.ActiveConfigMeta, config *protocol.ActiveConfigResponse) error {
@@ -218,25 +220,42 @@ func (s *Service) applyRenderedConfig(ctx context.Context, mode string, snapshot
}
func (s *Service) syncReferencedWAFIPGroups(ctx context.Context, supportFiles []protocol.SupportFile) error {
ids := referencedWAFIPGroupIDs(supportFiles)
ids, err := referencedWAFIPGroupIDs(supportFiles)
if err != nil {
return err
}
if len(ids) == 0 {
return nil
if s.nginxManager == nil {
return nil
}
return s.nginxManager.ReconcileWAFIPGroups([]uint{}, nil)
}
checksums, err := s.WAFIPGroupChecksums()
if err != nil {
return err
}
targetChecksums := make(map[string]string, len(ids))
for _, id := range ids {
key := strconv.FormatUint(uint64(id), 10)
if value := strings.TrimSpace(checksums[key]); value != "" {
targetChecksums[key] = value
}
}
response, err := s.client.SyncWAFIPGroups(ctx, protocol.WAFIPGroupSyncRequest{
IDs: ids,
Checksums: checksums,
Checksums: targetChecksums,
})
if err != nil {
return err
}
if response == nil || len(response.Groups) == 0 {
if s.nginxManager == nil {
return nil
}
return s.ApplyWAFIPGroups(ctx, response.Groups)
var changed []protocol.WAFIPGroup
if response != nil {
changed = response.Groups
}
return s.nginxManager.ReconcileWAFIPGroups(ids, changed)
}
type renderedActiveConfig struct {
@@ -288,7 +307,7 @@ func fromOpenRestySupportFiles(files []openrestyrender.SupportFile) []protocol.S
return result
}
func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) ([]uint, error) {
var content string
for _, file := range supportFiles {
if file.Path == "waf_config.json" {
@@ -297,28 +316,38 @@ func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
}
}
if content == "" {
return []uint{}
}
var payload struct {
RuleGroups []struct {
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
} `json:"rule_groups"`
return []uint{}, nil
}
var payload openrestyrender.WAFDocument
if err := json.Unmarshal([]byte(content), &payload); err != nil {
slog.Debug("decode waf_config.json for ip group references failed", "error", err)
return []uint{}
return nil, fmt.Errorf("decode waf_config.json for ip group references: %w", err)
}
seen := make(map[uint]struct{})
for _, group := range payload.RuleGroups {
for _, id := range group.IPWhitelistGroups {
if id > 0 {
seen[id] = struct{}{}
for _, legacyIDs := range [][]uint{group.IPWhitelistGroups, group.IPBlacklistGroups} {
for _, id := range legacyIDs {
if id > 0 {
seen[id] = struct{}{}
}
}
}
for _, id := range group.IPBlacklistGroups {
if id > 0 {
seen[id] = struct{}{}
for nodeID, node := range group.Graph.Nodes {
if node.Type != "ip_match" {
continue
}
var config *struct {
IPGroupIDs []uint `json:"ip_group_ids"`
}
if err := json.Unmarshal(node.Config, &config); err != nil || config == nil {
if err == nil {
err = errors.New("config must be a JSON object")
}
return nil, fmt.Errorf("decode ip_match config for rule group %d node %s: %w", group.ID, nodeID, err)
}
for _, id := range config.IPGroupIDs {
if id > 0 {
seen[id] = struct{}{}
}
}
}
}
@@ -327,7 +356,7 @@ func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids
return ids, nil
}
func shouldReportNoopApply(snapshot *state.Snapshot, version string, checksum string) bool {
+158 -3
View File
@@ -30,8 +30,10 @@ func testPagesSourceConfigJSON(deploymentID uint, checksum string) string {
type fakeClient struct {
config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload
wafSyncCalls []protocol.WAFIPGroupSyncRequest
pagesPackages map[uint][]byte
pagesHashes map[uint]string
wafSyncResult protocol.WAFIPGroupSyncResponse
fetchCalls int
hashCalls int
}
@@ -47,6 +49,12 @@ type fakeManager struct {
applyMainContents []string
applyRouteContents []string
applyFiles [][]protocol.SupportFile
wafChecksums map[string]string
wafReconcileIDs []uint
wafReconcileGroups []protocol.WAFIPGroup
wafUpdatedGroups []protocol.WAFIPGroup
wafReconcileErr error
wafReconcileCalls int
}
func testSourceConfigJSON(workerProcesses string, listen int) string {
@@ -106,7 +114,9 @@ func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyL
}
func (f *fakeClient) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) {
return &protocol.WAFIPGroupSyncResponse{}, nil
f.wafSyncCalls = append(f.wafSyncCalls, payload)
result := f.wafSyncResult
return &result, nil
}
func (m *fakeManager) Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) nginx.ApplyOutcome {
@@ -134,10 +144,18 @@ func (m *fakeManager) CurrentChecksum() (string, error) {
}
func (m *fakeManager) WAFIPGroupChecksums() (map[string]string, error) {
return map[string]string{}, nil
return m.wafChecksums, nil
}
func (m *fakeManager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error {
func (m *fakeManager) ReconcileWAFIPGroups(ids []uint, groups []protocol.WAFIPGroup) error {
m.wafReconcileCalls++
m.wafReconcileIDs = append([]uint(nil), ids...)
m.wafReconcileGroups = append([]protocol.WAFIPGroup(nil), groups...)
return m.wafReconcileErr
}
func (m *fakeManager) UpdateExistingWAFIPGroups(groups []protocol.WAFIPGroup) error {
m.wafUpdatedGroups = append([]protocol.WAFIPGroup(nil), groups...)
return nil
}
@@ -145,6 +163,143 @@ func (m *fakeManager) EnsureWorkerReadAccess() error {
return nil
}
func TestReferencedWAFIPGroupIDsSyncsCompiledDAGReferences(t *testing.T) {
client := &fakeClient{wafSyncResult: protocol.WAFIPGroupSyncResponse{Groups: []protocol.WAFIPGroup{{ID: 7, Checksum: "new-7"}}}}
manager := &fakeManager{wafChecksums: map[string]string{"2": "sum-2", "7": "old-7", "99": "stale"}}
service := New(client, manager, nil)
supportFiles := []protocol.SupportFile{{Path: "waf_config.json", Content: `{
"rule_groups":[
{"id":1,"ip_whitelist_group_ids":[0,11,7],"graph":{"entry":"start","nodes":{
"start":{"type":"start","config":{}},
"first":{"type":"ip_match","config":{"ip_group_ids":[7,0,2,7]}},
"geo":{"type":"geo_match","config":{"countries":["US"],"ip_group_ids":[700]}},
"pow":{"type":"pow","config":{"difficulty":4,"ip_group_ids":[800]}},
"second":{"type":"ip_match","config":{"ip_group_ids":[9,2]}}
}}}
],
"ip_groups":[{"id":2},{"id":7},{"id":9},{"id":11},{"id":404}],
"bindings":[]
}`}}
if err := service.syncReferencedWAFIPGroups(context.Background(), supportFiles); err != nil {
t.Fatalf("syncReferencedWAFIPGroups failed: %v", err)
}
if len(client.wafSyncCalls) != 1 {
t.Fatalf("expected one WAF IP group sync request, got %d", len(client.wafSyncCalls))
}
if got, want := fmt.Sprint(client.wafSyncCalls[0].IDs), "[2 7 9 11]"; got != want {
t.Fatalf("referenced IDs = %s, want %s", got, want)
}
if got, want := fmt.Sprint(client.wafSyncCalls[0].Checksums), "map[2:sum-2 7:old-7]"; got != want {
t.Fatalf("request checksums = %s, want target-only %s", got, want)
}
if got, want := fmt.Sprint(manager.wafReconcileIDs), "[2 7 9 11]"; got != want {
t.Fatalf("reconcile IDs = %s, want %s", got, want)
}
if len(manager.wafReconcileGroups) != 1 || manager.wafReconcileGroups[0].ID != 7 {
t.Fatalf("changed groups not passed to reconcile: %#v", manager.wafReconcileGroups)
}
}
func TestReferencedWAFIPGroupIDsReconcilesEmptyResponseAndSurfacesMissingLocal(t *testing.T) {
client := &fakeClient{}
manager := &fakeManager{
wafChecksums: map[string]string{"7": "mistaken-match"},
wafReconcileErr: fmt.Errorf("missing referenced WAF IP group 7"),
}
service := New(client, manager, nil)
err := service.syncReferencedWAFIPGroups(context.Background(), []protocol.SupportFile{{
Path: "waf_config.json", Content: `{"rule_groups":[{"graph":{"nodes":{"match":{"type":"ip_match","config":{"ip_group_ids":[7]}}}}}]}`,
}})
if err == nil || !strings.Contains(err.Error(), "missing referenced WAF IP group 7") {
t.Fatalf("expected missing local group error after empty delta, got %v", err)
}
if len(client.wafSyncCalls) != 1 || len(manager.wafReconcileIDs) != 1 || manager.wafReconcileIDs[0] != 7 {
t.Fatalf("empty response did not reach authoritative reconcile: calls=%#v ids=%#v", client.wafSyncCalls, manager.wafReconcileIDs)
}
}
func TestReferencedWAFIPGroupIDsEmptyTargetClearsWithoutRequest(t *testing.T) {
client := &fakeClient{}
manager := &fakeManager{}
service := New(client, manager, nil)
if err := service.syncReferencedWAFIPGroups(context.Background(), nil); err != nil {
t.Fatal(err)
}
if len(client.wafSyncCalls) != 0 {
t.Fatalf("empty target performed request I/O: %#v", client.wafSyncCalls)
}
if manager.wafReconcileCalls != 1 {
t.Fatal("empty authoritative target was not reconciled")
}
}
func TestApplyWAFIPGroupsUsesExistingOnlyUpdatePath(t *testing.T) {
manager := &fakeManager{}
service := New(nil, manager, nil)
groups := []protocol.WAFIPGroup{{ID: 99, Checksum: "broadcast"}}
if err := service.ApplyWAFIPGroups(context.Background(), groups); err != nil {
t.Fatal(err)
}
if len(manager.wafUpdatedGroups) != 1 || manager.wafUpdatedGroups[0].ID != 99 {
t.Fatalf("broadcast did not use existing-only update path: %#v", manager.wafUpdatedGroups)
}
if manager.wafReconcileCalls != 0 {
t.Fatal("broadcast must not use authoritative reconciliation")
}
}
func TestReferencedWAFIPGroupIDsRejectsMalformedRuntimeConfig(t *testing.T) {
tests := []struct {
name string
content string
wantErr string
}{
{name: "document", content: `{`, wantErr: "decode waf_config.json"},
{
name: "ip match config",
content: `{"rule_groups":[{"id":1,"graph":{"nodes":{"match":{"type":"ip_match","config":{"ip_group_ids":"bad"}}}}}],"bindings":[]}`,
wantErr: "decode ip_match config",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
client := &fakeClient{}
service := New(client, &fakeManager{}, nil)
err := service.syncReferencedWAFIPGroups(context.Background(), []protocol.SupportFile{{
Path: "waf_config.json", Content: test.content,
}})
if err == nil || !strings.Contains(err.Error(), test.wantErr) {
t.Fatalf("expected %q error, got %v", test.wantErr, err)
}
if len(client.wafSyncCalls) != 0 {
t.Fatalf("malformed WAF config must not send sync request, got %#v", client.wafSyncCalls)
}
})
}
}
func TestWAFIPGroupChecksumServicePublishesSidecar(t *testing.T) {
runtimeDir := t.TempDir()
manager := &nginx.Manager{RuntimeConfigDir: runtimeDir}
if err := manager.ReconcileWAFIPGroups([]uint{3}, []protocol.WAFIPGroup{{
ID: 3, Enabled: true, IPList: []string{"203.0.113.3"}, Checksum: "sum-3",
}}); err != nil {
t.Fatalf("ReconcileWAFIPGroups failed: %v", err)
}
jsonData, err := os.ReadFile(filepath.Join(runtimeDir, nginx.WAFIPGroupsConfigFileName))
if err != nil {
t.Fatalf("read IP group JSON: %v", err)
}
checksumData, err := os.ReadFile(filepath.Join(runtimeDir, nginx.WAFIPGroupsChecksumFileName))
if err != nil {
t.Fatalf("read IP group checksum: %v", err)
}
if got, want := strings.TrimSpace(string(checksumData)), testBytesChecksum(jsonData); got != want {
t.Fatalf("published checksum mismatch: got %q want %q", got, want)
}
}
func TestSyncOnceSuccess(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
+79 -24
View File
@@ -4,6 +4,7 @@
package agent
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
@@ -11,43 +12,32 @@ import (
"errors"
"fmt"
"sort"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
)
type snapshotWAFRuleGroupRef struct {
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
}
type snapshotWAFSection struct {
RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"`
}
type activeConfigSnapshot struct {
WAF snapshotWAFSection `json:"waf"`
WAF openrestyrender.WAFDocument `json:"waf"`
}
type runtimeIPMatchConfig struct {
IPs []string `json:"ips,omitempty"`
CIDRs []string `json:"cidrs,omitempty"`
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
}
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
return buildAgentWAFIPGroups(ctx, ids)
return validatedAgentWAFIPGroups(ctx, ids, false)
}
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
if len(targetIDs) == 0 {
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
if err != nil {
return nil, err
}
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []WAFIPGroup{}, nil
}
groups, err := buildAgentWAFIPGroups(ctx, targetIDs)
groups, err := validatedAgentWAFIPGroups(ctx, ids, true)
if err != nil {
return nil, err
}
@@ -61,6 +51,45 @@ func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[s
return changed, nil
}
func validatedAgentWAFIPGroups(ctx context.Context, ids []uint, fallbackToActive bool) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
if err != nil {
return nil, err
}
if len(targetIDs) == 0 && fallbackToActive {
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []WAFIPGroup{}, nil
}
validationIDs := uniqueUintIDs(append(append([]uint{}, activeIDs...), targetIDs...))
allGroups, err := buildAgentWAFIPGroups(ctx, validationIDs)
if err != nil {
return nil, err
}
runtimeGroups := make(map[string]protocol.WAFIPGroup, len(allGroups))
for _, group := range allGroups {
runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = group
}
if err = protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups); err != nil {
return nil, err
}
targetSet := make(map[uint]struct{}, len(targetIDs))
for _, id := range targetIDs {
targetSet[id] = struct{}{}
}
result := make([]WAFIPGroup, 0, len(targetIDs))
for _, group := range allGroups {
if _, ok := targetSet[group.ID]; ok {
result = append(result, group)
}
}
return result, nil
}
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
@@ -142,6 +171,8 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
}
idSet := make(map[uint]struct{})
for _, group := range snapshot.WAF.RuleGroups {
// Retain legacy flattened references while older active snapshots may
// still exist during a rolling Server upgrade.
for _, id := range group.IPWhitelistGroups {
if id > 0 {
idSet[id] = struct{}{}
@@ -152,6 +183,20 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
idSet[id] = struct{}{}
}
}
for nodeID, node := range group.Graph.Nodes {
if node.Type != "ip_match" {
continue
}
ids, err := runtimeIPMatchGroupIDs(node.Config)
if err != nil {
return nil, fmt.Errorf("活动配置 WAF 规则 %d 节点 %s 的 IP 匹配配置无效: %w", group.ID, nodeID, err)
}
for _, id := range ids {
if id > 0 {
idSet[id] = struct{}{}
}
}
}
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
@@ -161,6 +206,16 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
return ids, nil
}
func runtimeIPMatchGroupIDs(raw json.RawMessage) ([]uint, error) {
var config runtimeIPMatchConfig
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&config); err != nil {
return nil, err
}
return config.IPGroupIDs, nil
}
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
text := strings.TrimSpace(snapshotJSON)
if text == "" {
@@ -171,7 +226,7 @@ func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, erro
return nil, err
}
if snapshot.WAF.RuleGroups == nil {
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
snapshot.WAF.RuleGroups = []openrestyrender.WAFRuleGroup{}
}
return &snapshot, nil
}
@@ -7,10 +7,12 @@ import (
"context"
"encoding/json"
"strconv"
"strings"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/protocol"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -63,6 +65,115 @@ func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID
}).Error)
}
func seedActiveConfigWithWAFGraphIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
t.Helper()
snapshot := map[string]any{
"routes": []any{},
"waf": map[string]any{
"rule_groups": []map[string]any{
{
"id": 1,
"name": "graph refs",
"enabled": true,
"graph": map[string]any{
"entry": "start",
"nodes": map[string]any{
"start": map[string]any{
"type": "start",
"config": map[string]any{},
"next": map[string]string{"next": "match"},
},
"match": map[string]any{
"type": "ip_match",
"config": map[string]any{
"ip_group_ids": []uint{ipGroupID},
},
"next": map[string]string{"true": "allow", "false": "allow"},
},
"allow": map[string]any{
"type": "allow",
"config": map[string]any{},
},
},
},
},
},
"bindings": []any{},
},
}
snapshotJSON, err := json.Marshal(snapshot)
require.NoError(t, err)
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
Version: "20260713-graph-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "graph-test-checksum",
IsActive: true,
}).Error)
}
func TestChangedWAFIPGroupsForAgentDiscoversGraphReferences(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "graph runtime group",
Type: "manual",
Enabled: true,
IPList: `["192.0.2.88"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFGraphIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.NoError(t, err)
require.Len(t, groups, 1)
assert.Equal(t, ipGroup.ID, groups[0].ID)
assert.Equal(t, []string{"192.0.2.88"}, groups[0].IPList)
}
func TestChangedWAFIPGroupsForAgentRejectsMalformedIPMatchConfig(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
Version: "20260713-malformed-001",
SnapshotJSON: `{"waf":{"rule_groups":[{"id":7,"graph":{"entry":"match","nodes":{` +
`"match":{"type":"ip_match","config":{"ip_group_ids":"not-an-array"}}}}}],"bindings":[]}}`,
Checksum: "malformed-test-checksum",
IsActive: true,
}).Error)
_, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.ErrorContains(t, err, "规则 7 节点 match")
require.ErrorContains(t, err, "IP 匹配配置无效")
}
func TestChangedWAFIPGroupsForAgentRejectsOversizedSnapshotBeforeChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: strings.Repeat("x", protocol.MaxWAFIPGroupSnapshotBytes),
Type: "manual",
Enabled: true,
IPList: `[]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
agentGroup, err := buildAgentWAFIPGroup(ipGroup)
require.NoError(t, err)
_, err = ChangedWAFIPGroupsForAgent(ctx, []uint{ipGroup.ID}, map[string]string{
strconv.FormatUint(uint64(ipGroup.ID), 10): agentGroup.Checksum,
})
require.ErrorContains(t, err, "WAF IP 组快照大小")
require.ErrorContains(t, err, "超过上限")
}
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
@@ -13,6 +13,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -151,13 +152,7 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
customGroup := &model.OpenFlareWAFRuleGroup{
Name: "pow-group",
Enabled: true,
PoWEnabled: true,
PoWConfig: `{"difficulty":4,"algorithm":"fast","session_ttl":600,"challenge_ttl":300}`,
}
require.NoError(t, model.CreateOpenFlareWAFRuleGroup(ctx, customGroup))
customGroup := createSnapshotRule(t, ctx, "pow-group", waf.DefaultRuleGraph())
require.NoError(t, model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID}))
bundle, err := buildCurrentConfigBundle(ctx, true)
@@ -177,20 +172,24 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
}
assert.True(t, found, "expected WAF binding for enabled route")
var wafRuntime struct {
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
}
var wafRuntime openrestyrender.WAFDocument
foundWAFConfig := false
for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" {
continue
}
foundWAFConfig = true
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
}
require.Contains(t, wafRuntime.SiteRuleGroups, "example.com")
require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], customGroup.ID)
require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], globalGroup.ID)
require.True(t, foundWAFConfig, "expected rendered WAF support file")
require.NotEmpty(t, wafRuntime.RuleGroups)
assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
require.Len(t, wafRuntime.Bindings, 1)
assert.Equal(t, route.ID, wafRuntime.Bindings[0].RouteID)
assert.Equal(t, "example.com", wafRuntime.Bindings[0].SiteName)
assert.Equal(t, []uint{customGroup.ID}, wafRuntime.Bindings[0].RuleGroupIDs)
assert.Contains(t, bundle.RouteConfig, `set $openflare_waf_site "example.com"`)
assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`)
}
func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testing.T) {
@@ -210,34 +209,30 @@ func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testi
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
globalGroup.PoWEnabled = true
globalGroup.PoWConfig = `{"difficulty":4,"algorithm":"fast","session_ttl":600,"challenge_ttl":300}`
require.NoError(t, model.UpdateOpenFlareWAFRuleGroup(ctx, globalGroup))
graphJSON, err := json.Marshal(snapshotPoWGraph())
require.NoError(t, err)
globalGroup.Graph = string(graphJSON)
require.NoError(t, db.DB(ctx).Model(globalGroup).Update("graph", globalGroup.Graph).Error)
bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err)
assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`)
var wafRuntime struct {
RuleGroups []struct {
ID uint `json:"id"`
PoWEnabled bool `json:"pow_enabled"`
PoWConfig *struct {
Difficulty int `json:"difficulty"`
} `json:"pow_config"`
} `json:"rule_groups"`
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
}
var wafRuntime openrestyrender.WAFDocument
foundWAFConfig := false
for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" {
continue
}
foundWAFConfig = true
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
}
require.Contains(t, wafRuntime.SiteRuleGroups, "pow-global.example.com")
require.Contains(t, wafRuntime.SiteRuleGroups["pow-global.example.com"], globalGroup.ID)
require.True(t, foundWAFConfig, "expected rendered WAF support file")
require.NotEmpty(t, wafRuntime.RuleGroups)
assert.True(t, wafRuntime.RuleGroups[0].PoWEnabled)
require.NotNil(t, wafRuntime.RuleGroups[0].PoWConfig)
assert.Equal(t, 4, wafRuntime.RuleGroups[0].PoWConfig.Difficulty)
assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
assert.Equal(t, string(waf.RuleNodePoW), wafRuntime.RuleGroups[0].Graph.Nodes["pow"].Type)
require.Len(t, wafRuntime.Bindings, 1)
assert.Equal(t, "pow-global.example.com", wafRuntime.Bindings[0].SiteName)
assert.Empty(t, wafRuntime.Bindings[0].RuleGroupIDs)
require.NotEmpty(t, bundle.WAFSnapshot.RuleGroups)
assert.Equal(t, waf.RuleNodePoW, bundle.WAFSnapshot.RuleGroups[0].Graph.Nodes["pow"].Type)
}
@@ -9,17 +9,21 @@ import (
"errors"
"fmt"
"sort"
"strconv"
"strings"
oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"gorm.io/gorm"
)
const (
supportFilesPerCertificate = 2
supportFilesPerCertificate = 2
wafIPGroupChecksumHexLength = 64
// OpenResty 默认配置值
defaultOpenRestyReturnStatus = 421
@@ -66,22 +70,11 @@ type snapshotRoute struct {
}
type snapshotWAFRuleGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"`
BlockResponseBody string `json:"block_response_body,omitempty"`
IPWhitelist []string `json:"ip_whitelist,omitempty"`
IPBlacklist []string `json:"ip_blacklist,omitempty"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist,omitempty"`
CountryBlacklist []string `json:"country_blacklist,omitempty"`
RegionWhitelist []string `json:"region_whitelist,omitempty"`
RegionBlacklist []string `json:"region_blacklist,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *openrestyrender.PoWConfig `json:"pow_config,omitempty"`
ID uint `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"`
Graph waf.RuntimeRuleGraph `json:"graph"`
}
type snapshotWAFIPGroup struct {
@@ -311,35 +304,37 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if err := waf.EnsureDefaultRuleGroup(ctx); err != nil {
return snapshotWAFDocument{}, err
}
views, err := waf.ListRuleGroups(ctx)
groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
if err != nil {
return snapshotWAFDocument{}, err
}
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views))
for _, view := range views {
if !view.Enabled {
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(groups))
referencedIPGroupIDs := make(map[uint]struct{})
enabledRuleIDs := make(map[uint]struct{})
for _, group := range groups {
if !group.Enabled {
continue
}
var editorGraph waf.RuleGraph
if err = json.Unmarshal([]byte(group.Graph), &editorGraph); err != nil {
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图数据无效: %w", group.Name, err)
}
if err = waf.ValidateRuleGraph(ctx, editorGraph, snapshotWAFIPGroupExists); err != nil {
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图无效: %w", group.Name, err)
}
runtimeGraph, compileErr := waf.CompileRuleGraph(editorGraph)
if compileErr != nil {
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 编译失败: %w", group.Name, compileErr)
}
ruleGroups = append(ruleGroups, snapshotWAFRuleGroup{
ID: view.ID,
Name: view.Name,
Enabled: view.Enabled,
IsGlobal: view.IsGlobal,
BlockStatusCode: view.BlockStatusCode,
BlockResponseBody: view.BlockResponseBody,
IPWhitelist: view.IPWhitelist,
IPBlacklist: view.IPBlacklist,
IPWhitelistGroups: view.IPWhitelistGroups,
IPBlacklistGroups: view.IPBlacklistGroups,
CountryWhitelist: view.CountryWhitelist,
CountryBlacklist: view.CountryBlacklist,
RegionWhitelist: view.RegionWhitelist,
RegionBlacklist: view.RegionBlacklist,
PoWEnabled: view.PoWEnabled,
PoWConfig: convertPoWConfig(view.PoWEnabled, view.PoWConfig),
ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal, Graph: runtimeGraph,
})
enabledRuleIDs[group.ID] = struct{}{}
for _, id := range waf.ReferencedIPGroupIDs(editorGraph) {
referencedIPGroupIDs[id] = struct{}{}
}
}
ipGroups, err := buildSnapshotWAFIPGroups(ctx, ruleGroups)
ipGroups, err := buildSnapshotWAFIPGroups(ctx, referencedIPGroupIDs)
if err != nil {
return snapshotWAFDocument{}, err
}
@@ -366,16 +361,16 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if _, ok := enabledRouteSiteNames[binding.ProxyRouteID]; !ok {
continue
}
groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
if _, enabled := enabledRuleIDs[binding.RuleGroupID]; enabled {
groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
}
}
bindings := make([]snapshotWAFBinding, 0, len(enabledRouteSiteNames))
for routeID, siteName := range enabledRouteSiteNames {
groupIDs := groupIDsByRoute[routeID]
sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] })
bindings = append(bindings, snapshotWAFBinding{
RouteID: routeID,
SiteName: siteName,
RuleGroupIDs: groupIDs,
RuleGroupIDs: groupIDsByRoute[routeID],
})
}
sort.Slice(bindings, func(i, j int) bool {
@@ -387,16 +382,26 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
return snapshotWAFDocument{RuleGroups: ruleGroups, IPGroups: ipGroups, Bindings: bindings}, nil
}
func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleGroup) ([]snapshotWAFIPGroup, error) {
idSet := make(map[uint]struct{})
for _, group := range ruleGroups {
for _, id := range group.IPWhitelistGroups {
idSet[id] = struct{}{}
func validateSnapshotWAFIPGroupSize(groups []snapshotWAFIPGroup) error {
runtimeGroups := make(map[string]protocol.WAFIPGroup, len(groups))
for _, group := range groups {
ipList := group.IPList
if !group.Enabled {
ipList = []string{}
}
for _, id := range group.IPBlacklistGroups {
idSet[id] = struct{}{}
runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = protocol.WAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ipList,
Checksum: strings.Repeat("0", wafIPGroupChecksumHexLength),
}
}
return protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups)
}
func buildSnapshotWAFIPGroups(ctx context.Context, idSet map[uint]struct{}) ([]snapshotWAFIPGroup, error) {
if len(idSet) == 0 {
return []snapshotWAFIPGroup{}, nil
}
@@ -431,9 +436,20 @@ func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleG
IPList: ipList,
})
}
if err = validateSnapshotWAFIPGroupSize(snapshots); err != nil {
return nil, err
}
return snapshots, nil
}
func snapshotWAFIPGroupExists(ctx context.Context, id uint) (bool, error) {
group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id)
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
return group != nil, err
}
func decodeIPList(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
@@ -446,36 +462,6 @@ func decodeIPList(raw string) ([]string, error) {
return items, nil
}
func convertPoWConfig(enabled bool, config *waf.PoWConfig) *openrestyrender.PoWConfig {
if !enabled {
return nil
}
if config == nil {
defaultConfig := openrestyrender.DefaultPoWConfig()
return &defaultConfig
}
return &openrestyrender.PoWConfig{
Difficulty: config.Difficulty,
Algorithm: config.Algorithm,
SessionTTL: config.SessionTTL,
ChallengeTTL: config.ChallengeTTL,
Whitelist: openrestyrender.PoWListConfig{
IPs: config.Whitelist.IPs,
IPCidrs: config.Whitelist.IPCidrs,
Paths: config.Whitelist.Paths,
PathRegexes: config.Whitelist.PathRegexes,
UserAgents: config.Whitelist.UserAgents,
},
Blacklist: openrestyrender.PoWListConfig{
IPs: config.Blacklist.IPs,
IPCidrs: config.Blacklist.IPCidrs,
Paths: config.Blacklist.Paths,
PathRegexes: config.Blacklist.PathRegexes,
UserAgents: config.Blacklist.UserAgents,
},
}
}
func buildOpenRestyConfigSnapshot(ctx context.Context) openRestyConfigSnapshot {
// 读取所有 OpenResty 配置,使用默认值作为降级
getIntConfig := func(key string, defaultVal int) int {
@@ -0,0 +1,135 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"strings"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestBuildSnapshotRejectsOversizedAggregateWAFIPGroups(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
// Each group remains below the existing 2 MiB per-subscription ceiling,
// while the complete Agent runtime document exceeds the aggregate limit.
ipList, err := json.Marshal(strings.Fields(strings.Repeat("192.0.2.1 ", 165000)))
require.NoError(t, err)
require.Less(t, len(ipList), 2<<20)
groupIDs := make([]uint, 0, 12)
for index := 0; index < 12; index++ {
group := &model.OpenFlareWAFIPGroup{
Name: "aggregate-" + strings.Repeat("x", index),
Type: "manual",
Enabled: true,
IPList: string(ipList),
}
require.NoError(t, db.DB(ctx).Create(group).Error)
groupIDs = append(groupIDs, group.ID)
}
createSnapshotRule(t, ctx, "oversized-aggregate", snapshotIPMatchGraphForGroups(groupIDs))
_, err = buildSnapshotWAFDocument(ctx, nil)
require.ErrorContains(t, err, "WAF IP 组快照大小")
require.ErrorContains(t, err, "超过上限")
}
func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
route := &model.ProxyRoute{SiteName: "ordered.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true}
require.NoError(t, model.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, route.SiteName)
referenced := &model.OpenFlareWAFIPGroup{Name: "referenced", Type: "manual", Enabled: true, IPList: `["192.0.2.1"]`}
unused := &model.OpenFlareWAFIPGroup{Name: "unused", Type: "manual", Enabled: true, IPList: `["198.51.100.1"]`}
require.NoError(t, db.DB(ctx).Create(referenced).Error)
require.NoError(t, db.DB(ctx).Create(unused).Error)
customA := createSnapshotRule(t, ctx, "custom-a", waf.DefaultRuleGraph())
customB := createSnapshotRule(t, ctx, "custom-b", snapshotIPMatchGraph(referenced.ID))
require.NoError(t, model.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, route.ID, []uint{customB.ID, customA.ID}))
snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route})
require.NoError(t, err)
require.Len(t, snapshot.Bindings, 1)
assert.Equal(t, []uint{customB.ID, customA.ID}, snapshot.Bindings[0].RuleGroupIDs)
require.Len(t, snapshot.IPGroups, 1)
assert.Equal(t, referenced.ID, snapshot.IPGroups[0].ID)
var customBSnapshot *snapshotWAFRuleGroup
for index := range snapshot.RuleGroups {
if snapshot.RuleGroups[index].ID == customB.ID {
customBSnapshot = &snapshot.RuleGroups[index]
}
}
require.NotNil(t, customBSnapshot)
assert.Equal(t, "start", customBSnapshot.Graph.Entry)
assert.Equal(t, waf.RuleNodeIPMatch, customBSnapshot.Graph.Nodes["match"].Type)
raw, err := json.Marshal(customBSnapshot)
require.NoError(t, err)
assert.NotContains(t, string(raw), "position")
assert.NotContains(t, string(raw), "ip_whitelist")
}
func TestBuildSnapshotRejectsInvalidWAFGraph(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
invalid := &model.OpenFlareWAFRuleGroup{Name: "invalid", Enabled: true, Graph: `{"schema_version":1,"nodes":[],"edges":[]}`, Revision: 1}
require.NoError(t, db.DB(ctx).Create(invalid).Error)
_, err := buildSnapshotWAFDocument(ctx, nil)
require.ErrorContains(t, err, "invalid")
}
func createSnapshotRule(t *testing.T, ctx context.Context, name string, graph waf.RuleGraph) *model.OpenFlareWAFRuleGroup {
t.Helper()
raw, err := json.Marshal(graph)
require.NoError(t, err)
rule := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: true, Graph: string(raw), Revision: 1}
require.NoError(t, db.DB(ctx).Create(rule).Error)
return rule
}
func snapshotIPMatchGraph(ipGroupID uint) waf.RuleGraph {
return snapshotIPMatchGraphForGroups([]uint{ipGroupID})
}
func snapshotIPMatchGraphForGroups(ipGroupIDs []uint) waf.RuleGraph {
config, _ := json.Marshal(waf.IPMatchConfig{IPGroupIDs: ipGroupIDs})
return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
{ID: "start", Type: waf.RuleNodeStart, Position: waf.RulePosition{X: 1, Y: 2}, Config: json.RawMessage(`{}`)},
{ID: "match", Type: waf.RuleNodeIPMatch, Position: waf.RulePosition{X: 3, Y: 4}, Config: config},
{ID: "allow", Type: waf.RuleNodeAllow, Position: waf.RulePosition{X: 5, Y: 6}, Config: json.RawMessage(`{}`)},
}, Edges: []waf.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"},
}}
}
func snapshotPoWGraph() waf.RuleGraph {
config, _ := json.Marshal(waf.PoWNodeConfig{Algorithm: "fast", Difficulty: 4, SessionTTL: 600, ChallengeTTL: 300})
return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
{ID: "start", Type: waf.RuleNodeStart, Config: json.RawMessage(`{}`)},
{ID: "pow", Type: waf.RuleNodePoW, Config: config},
{ID: "allow", Type: waf.RuleNodeAllow, Config: json.RawMessage(`{}`)},
}, Edges: []waf.RuleEdge{
{ID: "e1", Source: "start", SourceHandle: "next", Target: "pow"},
{ID: "e2", Source: "pow", SourceHandle: "next", Target: "allow"},
}}
}
@@ -109,12 +109,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
t.Run("WAF rule group create", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/waf/rule-groups"), map[string]any{
"name": "edge-security",
"enabled": true,
"block_status_code": 403,
"ip_whitelist": []string{"192.0.2.1"},
"ip_blacklist": []string{"203.0.113.10"},
"country_blacklist": []string{"CN"},
"name": "edge-security",
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
@@ -124,7 +119,8 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
assert.NotZero(t, ruleGroupID)
assert.Equal(t, "edge-security", data["name"])
assert.Equal(t, false, data["is_global"])
assert.Equal(t, float64(403), data["block_status_code"])
assert.Equal(t, float64(1), data["revision"])
assert.NotNil(t, data["graph"])
})
t.Run("WAF rule group list includes global and custom groups", func(t *testing.T) {
@@ -174,11 +170,9 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
t,
engine,
http.MethodPost,
fmt.Sprintf("%s/waf/rule-groups/%d/update", apiPath(""), ruleGroupID),
fmt.Sprintf("%s/waf/rule-groups/%d/meta", apiPath(""), ruleGroupID),
map[string]any{
"name": "edge-security-updated",
"enabled": true,
"block_status_code": 451,
"name": "edge-security-updated", "enabled": true,
},
adminAuthHeaders(seed.Token),
)
@@ -187,7 +181,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.Equal(t, "edge-security-updated", data["name"])
assert.Equal(t, float64(451), data["block_status_code"])
assert.Equal(t, true, data["enabled"])
})
t.Run("WAF IP group create", func(t *testing.T) {
+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())
}