[新增] 添加 WAF 规则组及其绑定的 API 支持,更新前端页面以集成 WAF 功能

This commit is contained in:
ryan
2026-05-30 12:16:28 +08:00
parent 290ddd7b51
commit 8300d3ec1c
39 changed files with 2574 additions and 38 deletions
+8
View File
@@ -10,6 +10,7 @@ import (
"openflare-agent/internal/agent"
"openflare-agent/internal/config"
"openflare-agent/internal/geoipupdate"
"openflare-agent/internal/heartbeat"
"openflare-agent/internal/httpclient"
"openflare-agent/internal/logging"
@@ -54,6 +55,7 @@ func main() {
"cert_dir", cfg.CertDir,
"lua_dir", cfg.LuaDir,
"runtime_config_dir", cfg.RuntimeConfigDir,
"mmdb_path", cfg.MMDBPath,
)
client := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
@@ -100,6 +102,12 @@ func main() {
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
geoIPUpdater := &geoipupdate.Updater{
MMDBPath: cfg.MMDBPath,
DownloadURL: cfg.MMDBDownloadURL,
UpdateInterval: cfg.MMDBUpdateInterval.Duration(),
}
go geoIPUpdater.Run(ctx)
slog.Info("agent process started")
if err = runner.Run(ctx); err != nil && err != context.Canceled {
+37
View File
@@ -22,11 +22,14 @@ const (
defaultCertDirRelativePath = "etc/nginx/certs"
defaultLuaDirRelativePath = "etc/nginx/lua"
defaultRuntimeConfigDirRelativePath = "etc/openflare"
defaultMMDBRelativePath = "etc/openflare/GeoLite2-Country.mmdb"
defaultAccessLogRelativePath = "var/log/openflare/access.log"
defaultStateRelativePath = "var/lib/openflare/agent-state.json"
defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json"
defaultOpenRestyObservabilityPort = 18081
defaultObservabilityReplayMinutes = 15
defaultMMDBUpdateInterval = 24 * time.Hour
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
)
var (
@@ -56,6 +59,9 @@ type Config struct {
LuaDir string `json:"lua_dir"`
OpenrestyLuaDir string `json:"openresty_lua_dir"`
RuntimeConfigDir string `json:"runtime_config_dir"`
MMDBPath string `json:"mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -85,6 +91,9 @@ type configFile struct {
LuaDir string `json:"lua_dir"`
OpenrestyLuaDir string `json:"openresty_lua_dir"`
RuntimeConfigDir string `json:"runtime_config_dir"`
MMDBPath string `json:"mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -127,6 +136,9 @@ func Load(path string) (*Config, error) {
LuaDir: file.LuaDir,
OpenrestyLuaDir: file.OpenrestyLuaDir,
RuntimeConfigDir: file.RuntimeConfigDir,
MMDBPath: file.MMDBPath,
MMDBUpdateInterval: file.MMDBUpdateInterval,
MMDBDownloadURL: file.MMDBDownloadURL,
OpenrestyObservabilityPort: file.OpenrestyObservabilityPort,
ObservabilityBufferPath: file.ObservabilityBufferPath,
ObservabilityReplayMinutes: file.ObservabilityReplayMinutes,
@@ -186,6 +198,15 @@ func applyDefaults(cfg *Config, baseDir string) {
if cfg.RuntimeConfigDir == "" {
cfg.RuntimeConfigDir = joinManagedPath(cfg.DataDir, defaultRuntimeConfigDirRelativePath)
}
if cfg.MMDBPath == "" {
cfg.MMDBPath = joinManagedPath(cfg.DataDir, defaultMMDBRelativePath)
}
if cfg.MMDBUpdateInterval <= 0 {
cfg.MMDBUpdateInterval = MillisecondDuration(defaultMMDBUpdateInterval)
}
if cfg.MMDBDownloadURL == "" {
cfg.MMDBDownloadURL = defaultMMDBDownloadURL
}
if cfg.OpenrestyObservabilityPort <= 0 {
cfg.OpenrestyObservabilityPort = defaultOpenRestyObservabilityPort
}
@@ -241,6 +262,9 @@ func normalizeManagedPaths(cfg *Config) {
if usesSlashPath(cfg.ObservabilityBufferPath) {
cfg.ObservabilityBufferPath = filepath.ToSlash(cfg.ObservabilityBufferPath)
}
if usesSlashPath(cfg.MMDBPath) {
cfg.MMDBPath = filepath.ToSlash(cfg.MMDBPath)
}
}
func hasEnvConfig() bool {
@@ -255,6 +279,9 @@ func hasEnvConfig() bool {
"OPENFLARE_HEARTBEAT_INTERVAL",
"OPENFLARE_REQUEST_TIMEOUT",
"OPENFLARE_OPENRESTY_OBSERVABILITY_PORT",
"OPENFLARE_MMDB_PATH",
"OPENFLARE_MMDB_UPDATE_INTERVAL",
"OPENFLARE_MMDB_DOWNLOAD_URL",
} {
if strings.TrimSpace(os.Getenv(key)) != "" {
return true
@@ -279,6 +306,8 @@ func applyEnvOverrides(cfg *Config) {
overrideString("OPENFLARE_NODE_IP", &cfg.NodeIP)
overrideString("OPENFLARE_DATA_DIR", &cfg.DataDir)
overrideString("OPENFLARE_OPENRESTY_PATH", &cfg.OpenrestyPath)
overrideString("OPENFLARE_MMDB_PATH", &cfg.MMDBPath)
overrideString("OPENFLARE_MMDB_DOWNLOAD_URL", &cfg.MMDBDownloadURL)
if value := strings.TrimSpace(os.Getenv("OPENFLARE_HEARTBEAT_INTERVAL")); value != "" {
if duration, err := parseDurationValue(value); err == nil {
cfg.HeartbeatInterval = duration
@@ -289,6 +318,11 @@ func applyEnvOverrides(cfg *Config) {
cfg.RequestTimeout = duration
}
}
if value := strings.TrimSpace(os.Getenv("OPENFLARE_MMDB_UPDATE_INTERVAL")); value != "" {
if duration, err := parseDurationValue(value); err == nil {
cfg.MMDBUpdateInterval = duration
}
}
if value := strings.TrimSpace(os.Getenv("OPENFLARE_OPENRESTY_OBSERVABILITY_PORT")); value != "" {
var port int
if _, err := fmt.Sscanf(value, "%d", &port); err == nil {
@@ -342,6 +376,9 @@ func validate(cfg *Config) error {
if cfg.ObservabilityReplayMinutes <= 0 {
return errors.New("observability_replay_minutes 必须大于 0")
}
if cfg.MMDBUpdateInterval <= 0 {
return errors.New("mmdb_update_interval 必须大于 0")
}
return nil
}
@@ -0,0 +1,8 @@
package geoipdata
import "embed"
//go:embed GeoLite2-Country.mmdb
var FS embed.FS
const DefaultMMDBName = "GeoLite2-Country.mmdb"
@@ -0,0 +1,67 @@
package geoipupdate
import (
"context"
"fmt"
"io/fs"
"log/slog"
"os"
"path/filepath"
"time"
"openflare-agent/internal/geoipdata"
"openflare/utils/geoip"
)
type Updater struct {
MMDBPath string
DownloadURL string
UpdateInterval time.Duration
}
func (u *Updater) EnsureInitialDatabase() error {
path := filepath.Clean(u.MMDBPath)
if path == "" || path == "." {
return nil
}
if _, err := os.Stat(path); err == nil {
return nil
} else if !os.IsNotExist(err) {
return fmt.Errorf("stat mmdb file failed: %w", err)
}
data, err := fs.ReadFile(geoipdata.FS, geoipdata.DefaultMMDBName)
if err != nil {
return fmt.Errorf("read embedded mmdb failed: %w", err)
}
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return fmt.Errorf("create mmdb directory failed: %w", err)
}
if err := os.WriteFile(path, data, 0o644); err != nil {
return fmt.Errorf("write initial mmdb failed: %w", err)
}
slog.Info("initialized GeoIP mmdb from embedded database", "path", path, "size", len(data))
return nil
}
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)
}
ticker := time.NewTicker(u.UpdateInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if err := geoip.DownloadMaxMindDatabase(u.MMDBPath, u.DownloadURL); err != nil {
slog.Warn("update GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
continue
}
slog.Info("GeoIP mmdb updated", "path", u.MMDBPath)
}
}
}
@@ -0,0 +1,24 @@
package geoipupdate
import (
"os"
"path/filepath"
"testing"
)
func TestEnsureInitialDatabaseCopiesEmbeddedMMDB(t *testing.T) {
tempDir := t.TempDir()
path := filepath.Join(tempDir, "GeoLite2-Country.mmdb")
updater := &Updater{MMDBPath: path}
if err := updater.EnsureInitialDatabase(); err != nil {
t.Fatalf("EnsureInitialDatabase failed: %v", err)
}
info, err := os.Stat(path)
if err != nil {
t.Fatalf("expected mmdb to exist: %v", err)
}
if info.Size() == 0 {
t.Fatal("expected copied mmdb to be non-empty")
}
}
+34 -1
View File
@@ -220,6 +220,9 @@ func (m *Manager) writeTargetFiles(mainConfig string, routeConfig string, suppor
if err := m.writePowConfig(supportFiles); err != nil {
return err
}
if err := m.writeWAFConfig(supportFiles); err != nil {
return err
}
if err := m.ensureMimeTypes(); err != nil {
return err
}
@@ -289,6 +292,7 @@ func (m *Manager) EnsureLuaAssets() error {
return nil
}
allSupportFiles := append(ManagedObservabilityLuaFiles(), m.managedPowLuaFiles()...)
allSupportFiles = append(allSupportFiles, m.managedWAFLuaFiles()...)
powStaticFiles, err := ManagedPowStaticFiles()
if err != nil {
return fmt.Errorf("load pow static files: %w", err)
@@ -628,10 +632,30 @@ func (m *Manager) writePowConfig(supportFiles []protocol.SupportFile) error {
return nil
}
func (m *Manager) writeWAFConfig(supportFiles []protocol.SupportFile) error {
if m.RuntimeConfigDir == "" {
return nil
}
configPath := filepath.Join(m.RuntimeConfigDir, "waf_config.json")
for _, file := range supportFiles {
if file.Path == "waf_config.json" {
if err := os.WriteFile(configPath, []byte(file.Content), 0o644); err != nil {
return fmt.Errorf("write waf_config.json: %w", err)
}
slog.Info("wrote waf config", "path", configPath, "size", len(file.Content))
return nil
}
}
if err := os.Remove(configPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove waf_config.json: %w", err)
}
return nil
}
func (m *Manager) writeManagedCertFiles(certFiles []protocol.SupportFile) error {
files := make([]managedFile, 0, len(certFiles))
for _, file := range certFiles {
if file.Path == "pow_config.json" {
if file.Path == "pow_config.json" || file.Path == "waf_config.json" {
continue
}
targetPath, err := m.certFileTargetPath(file.Path)
@@ -1014,6 +1038,15 @@ func (m *Manager) managedPowLuaFiles() []protocol.SupportFile {
return files
}
func (m *Manager) managedWAFLuaFiles() []protocol.SupportFile {
files := ManagedWAFLuaFiles()
runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir))
for index := range files {
files[index].Content = strings.ReplaceAll(files[index].Content, RuntimeConfigDirPlaceholder, runtimeConfigDir)
}
return files
}
func ObservabilityListenAddress(openrestyPath string, port int) string {
if port <= 0 {
return ""
@@ -0,0 +1,214 @@
package nginx
import "openflare-agent/internal/protocol"
const openRestyWAFCheckLua = `local cjson = require "cjson.safe"
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 list_contains(items, value)
if not items 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 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 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 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.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
return
end
local ip = ngx.var.remote_addr or ""
local groups = active_groups(config)
for _, group in ipairs(groups) do
if ip_matches(group.ip_whitelist, ip) then
return
end
end
local country = nil
for _, group in ipairs(groups) do
if group.country_whitelist 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) then
return exit_with_group(group)
end
end
for _, group in ipairs(groups) do
if group.country_blacklist 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
`
func ManagedWAFLuaFiles() []protocol.SupportFile {
return []protocol.SupportFile{
{Path: "waf/check.lua", Content: openRestyWAFCheckLua},
}
}