mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 00:26:37 +08:00
[新增] 添加 WAF 规则组及其绑定的 API 支持,更新前端页面以集成 WAF 功能
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Binary file not shown.
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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},
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user