mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
[新增] 添加 WAF 规则组及其绑定的 API 支持,更新前端页面以集成 WAF 功能
This commit is contained in:
@@ -53,6 +53,7 @@ type ConfigDiffResult struct {
|
||||
RemovedDomains []string `json:"removed_domains"`
|
||||
ModifiedDomains []string `json:"modified_domains"`
|
||||
MainConfigChanged bool `json:"main_config_changed"`
|
||||
WAFConfigChanged bool `json:"waf_config_changed"`
|
||||
ChangedOptionKeys []string `json:"changed_option_keys"`
|
||||
ChangedOptionDetails []ConfigOptionDiffItem `json:"changed_option_details"`
|
||||
CurrentWebsiteCount int `json:"current_website_count"`
|
||||
@@ -93,6 +94,32 @@ type snapshotRoute struct {
|
||||
Remark string `json:"remark,omitempty"`
|
||||
}
|
||||
|
||||
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"`
|
||||
CountryWhitelist []string `json:"country_whitelist,omitempty"`
|
||||
CountryBlacklist []string `json:"country_blacklist,omitempty"`
|
||||
RegionWhitelist []string `json:"region_whitelist,omitempty"`
|
||||
RegionBlacklist []string `json:"region_blacklist,omitempty"`
|
||||
}
|
||||
|
||||
type snapshotWAFBinding struct {
|
||||
RouteID uint `json:"route_id"`
|
||||
SiteName string `json:"site_name"`
|
||||
RuleGroupIDs []uint `json:"rule_group_ids"`
|
||||
}
|
||||
|
||||
type snapshotWAFDocument struct {
|
||||
RuleGroups []snapshotWAFRuleGroup `json:"rule_groups"`
|
||||
Bindings []snapshotWAFBinding `json:"bindings"`
|
||||
}
|
||||
|
||||
type routeCacheConfig struct {
|
||||
Enabled bool
|
||||
Policy string
|
||||
@@ -153,11 +180,13 @@ type openRestyConfigSnapshot struct {
|
||||
type snapshotDocument struct {
|
||||
Routes []snapshotRoute `json:"routes"`
|
||||
OpenRestyConfig openRestyConfigSnapshot `json:"openresty_config"`
|
||||
WAF snapshotWAFDocument `json:"waf"`
|
||||
}
|
||||
|
||||
type configBundle struct {
|
||||
Routes []*model.ProxyRoute
|
||||
SnapshotRoutes []snapshotRoute
|
||||
WAFSnapshot snapshotWAFDocument
|
||||
OpenRestyConfig openRestyConfigSnapshot
|
||||
SnapshotJSON string
|
||||
MainConfig string
|
||||
@@ -310,6 +339,7 @@ func DiffConfigVersion() (*ConfigDiffResult, error) {
|
||||
}
|
||||
}
|
||||
result.MainConfigChanged = activeVersion.MainConfig != bundle.MainConfig
|
||||
result.WAFConfigChanged = !snapshotWAFConfigEqual(activeSnapshot.WAF, bundle.WAFSnapshot)
|
||||
result.ChangedOptionDetails = diffOpenRestyOptionDetails(activeSnapshot.OpenRestyConfig, bundle.OpenRestyConfig)
|
||||
result.ChangedOptionKeys = extractOptionDiffKeys(result.ChangedOptionDetails)
|
||||
sort.Strings(result.AddedSites)
|
||||
@@ -460,10 +490,15 @@ func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
wafSnapshot, err := buildSnapshotWAFDocument(routes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
openRestyConfig := buildOpenRestyConfigSnapshot()
|
||||
snapshotDoc := snapshotDocument{
|
||||
Routes: snapshotRoutes,
|
||||
OpenRestyConfig: openRestyConfig,
|
||||
WAF: wafSnapshot,
|
||||
}
|
||||
snapshotJSON, err := json.Marshal(snapshotDoc)
|
||||
if err != nil {
|
||||
@@ -473,6 +508,10 @@ func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
wafConfigJSON, err := renderWAFConfigBundle(wafSnapshot)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
powConfigJSON, powSupportFiles, err := renderPowConfigBundle(routes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -480,9 +519,11 @@ func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) {
|
||||
supportFiles = append(supportFiles, powSupportFiles...)
|
||||
mainConfig := renderMainConfig(openRestyConfig)
|
||||
supportFiles = append(supportFiles, SupportFile{Path: "pow_config.json", Content: powConfigJSON})
|
||||
supportFiles = append(supportFiles, SupportFile{Path: "waf_config.json", Content: wafConfigJSON})
|
||||
return &configBundle{
|
||||
Routes: routes,
|
||||
SnapshotRoutes: snapshotRoutes,
|
||||
WAFSnapshot: wafSnapshot,
|
||||
OpenRestyConfig: openRestyConfig,
|
||||
SnapshotJSON: string(snapshotJSON),
|
||||
MainConfig: mainConfig,
|
||||
@@ -550,6 +591,74 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func buildSnapshotWAFDocument(routes []*model.ProxyRoute) (snapshotWAFDocument, error) {
|
||||
if err := EnsureDefaultWAFRuleGroup(); err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
views, err := ListWAFRuleGroups()
|
||||
if err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views))
|
||||
for _, view := range views {
|
||||
if !view.Enabled {
|
||||
continue
|
||||
}
|
||||
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,
|
||||
CountryWhitelist: view.CountryWhitelist,
|
||||
CountryBlacklist: view.CountryBlacklist,
|
||||
RegionWhitelist: view.RegionWhitelist,
|
||||
RegionBlacklist: view.RegionBlacklist,
|
||||
})
|
||||
}
|
||||
enabledRouteIDs := make(map[uint]string, len(routes))
|
||||
for _, route := range routes {
|
||||
if route == nil {
|
||||
continue
|
||||
}
|
||||
siteName := strings.TrimSpace(route.SiteName)
|
||||
if siteName == "" {
|
||||
siteName = route.Domain
|
||||
}
|
||||
enabledRouteIDs[route.ID] = siteName
|
||||
}
|
||||
var rawBindings []model.WAFRuleGroupBinding
|
||||
if err := model.DB.Order("proxy_route_id asc").Order("rule_group_id asc").Find(&rawBindings).Error; err != nil {
|
||||
return snapshotWAFDocument{}, err
|
||||
}
|
||||
groupIDsByRoute := make(map[uint][]uint, len(rawBindings))
|
||||
for _, binding := range rawBindings {
|
||||
if _, ok := enabledRouteIDs[binding.ProxyRouteID]; !ok {
|
||||
continue
|
||||
}
|
||||
groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
|
||||
}
|
||||
bindings := make([]snapshotWAFBinding, 0, len(groupIDsByRoute))
|
||||
for routeID, groupIDs := range groupIDsByRoute {
|
||||
sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] })
|
||||
bindings = append(bindings, snapshotWAFBinding{
|
||||
RouteID: routeID,
|
||||
SiteName: enabledRouteIDs[routeID],
|
||||
RuleGroupIDs: groupIDs,
|
||||
})
|
||||
}
|
||||
sort.Slice(bindings, func(i, j int) bool {
|
||||
if bindings[i].SiteName == bindings[j].SiteName {
|
||||
return bindings[i].RouteID < bindings[j].RouteID
|
||||
}
|
||||
return bindings[i].SiteName < bindings[j].SiteName
|
||||
})
|
||||
return snapshotWAFDocument{RuleGroups: ruleGroups, Bindings: bindings}, nil
|
||||
}
|
||||
|
||||
func mustDecodeSnapshotCertIDs(route *model.ProxyRoute) []uint {
|
||||
if route == nil {
|
||||
return []uint{}
|
||||
@@ -729,6 +838,18 @@ func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func snapshotWAFConfigEqual(left snapshotWAFDocument, right snapshotWAFDocument) bool {
|
||||
leftJSON, err := json.Marshal(left)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
rightJSON, err := json.Marshal(right)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return string(leftJSON) == string(rightJSON)
|
||||
}
|
||||
|
||||
func snapshotPoWConfigEqual(left *ProxyRoutePoWConfig, right *ProxyRoutePoWConfig) bool {
|
||||
if left == nil || right == nil {
|
||||
return left == nil && right == nil
|
||||
@@ -949,7 +1070,7 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
continue
|
||||
}
|
||||
certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID)
|
||||
@@ -1010,24 +1131,24 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot)
|
||||
|
||||
if route.RedirectHTTP {
|
||||
if len(httpOnlyDomains) > 0 {
|
||||
builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), displayName, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
|
||||
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains), displayName))
|
||||
}
|
||||
} else {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), route.OriginURL, route.OriginHost, certID, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), displayName, route.OriginURL, route.OriginHost, certID, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
}
|
||||
}
|
||||
return builder.String(), dedupeSupportFiles(supportFiles), nil
|
||||
@@ -1143,6 +1264,11 @@ func renderPowAccessBlock(powEnabled bool) string {
|
||||
return fmt.Sprintf(" access_by_lua_file %s/pow/check.lua;\n", nginxLuaDirPlaceholder)
|
||||
}
|
||||
|
||||
func renderWAFAccessBlock(siteName string) string {
|
||||
escapedSiteName := escapeNginxString(siteName)
|
||||
return fmt.Sprintf(" set $openflare_waf_site \"%s\";\n access_by_lua_file %s/waf/check.lua;\n", escapedSiteName, nginxLuaDirPlaceholder)
|
||||
}
|
||||
|
||||
func renderBasicAuthBlock(enabled bool, username, password string) string {
|
||||
if !enabled || username == "" || password == "" {
|
||||
return ""
|
||||
@@ -1299,18 +1425,19 @@ func nextVersionNumber(now time.Time) (string, error) {
|
||||
return fmt.Sprintf("%s-%03d", prefix, sequence+1), nil
|
||||
}
|
||||
|
||||
func renderHTTPProxyServer(serverNames string, originURL string, originHost string, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
|
||||
func renderHTTPProxyServer(serverNames string, siteName string, originURL string, originHost string, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string {
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, renderWAFAccessBlock(siteName), renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
|
||||
}
|
||||
|
||||
func renderHTTPRedirectServer(serverNames string) string {
|
||||
func renderHTTPRedirectServer(serverNames string, siteName string) string {
|
||||
_ = siteName
|
||||
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", serverNames)
|
||||
}
|
||||
|
||||
func renderHTTPSServer(serverNames string, originURL string, originHost string, certificateID uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string {
|
||||
func renderHTTPSServer(serverNames string, siteName string, originURL string, originHost string, certificateID uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string {
|
||||
certPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateCertFileName(certificateID))
|
||||
keyPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateKeyFileName(certificateID))
|
||||
return fmt.Sprintf("server {\n listen 443 ssl;\n http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, certPath, keyPath, renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
|
||||
return fmt.Sprintf("server {\n listen 443 ssl;\n http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, certPath, keyPath, renderWAFAccessBlock(siteName), renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
|
||||
}
|
||||
|
||||
func renderHTTPSServerWithCertificates(serverNames string, originURL string, originHost string, certificateIDs []uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string {
|
||||
@@ -1648,6 +1775,12 @@ func quoteNginxStringLiteral(value string) string {
|
||||
return fmt.Sprintf(`"%s"`, escaped)
|
||||
}
|
||||
|
||||
func escapeNginxString(value string) string {
|
||||
escaped := strings.ReplaceAll(value, `\`, `\\`)
|
||||
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
|
||||
return escaped
|
||||
}
|
||||
|
||||
func certificateCertFileName(id uint) string {
|
||||
return fmt.Sprintf("%d.crt", id)
|
||||
}
|
||||
@@ -1711,3 +1844,80 @@ func renderPowConfigBundle(routes []*model.ProxyRoute) (string, []SupportFile, e
|
||||
}
|
||||
return string(data), nil, nil
|
||||
}
|
||||
|
||||
func renderWAFConfigBundle(snapshot snapshotWAFDocument) (string, error) {
|
||||
type wafRuntimeRuleGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
IsGlobal bool `json:"is_global"`
|
||||
BlockStatusCode int `json:"block_status_code"`
|
||||
BlockResponseBody string `json:"block_response_body"`
|
||||
IPWhitelist []string `json:"ip_whitelist"`
|
||||
IPBlacklist []string `json:"ip_blacklist"`
|
||||
CountryWhitelist []string `json:"country_whitelist"`
|
||||
CountryBlacklist []string `json:"country_blacklist"`
|
||||
RegionWhitelist []string `json:"region_whitelist"`
|
||||
RegionBlacklist []string `json:"region_blacklist"`
|
||||
}
|
||||
type wafRuntimeConfig struct {
|
||||
DefaultBlockStatusCode int `json:"default_block_status_code"`
|
||||
RuleGroups []wafRuntimeRuleGroup `json:"rule_groups"`
|
||||
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
|
||||
}
|
||||
groups := make([]wafRuntimeRuleGroup, 0, len(snapshot.RuleGroups))
|
||||
globalGroupIDs := make([]uint, 0)
|
||||
enabledGroupIDs := make(map[uint]struct{}, len(snapshot.RuleGroups))
|
||||
for _, group := range snapshot.RuleGroups {
|
||||
if !group.Enabled {
|
||||
continue
|
||||
}
|
||||
statusCode := group.BlockStatusCode
|
||||
if statusCode == 0 {
|
||||
statusCode = defaultWAFBlockStatusCode
|
||||
}
|
||||
if group.IsGlobal {
|
||||
globalGroupIDs = append(globalGroupIDs, group.ID)
|
||||
}
|
||||
enabledGroupIDs[group.ID] = struct{}{}
|
||||
groups = append(groups, wafRuntimeRuleGroup{
|
||||
ID: group.ID,
|
||||
Name: group.Name,
|
||||
IsGlobal: group.IsGlobal,
|
||||
BlockStatusCode: statusCode,
|
||||
BlockResponseBody: group.BlockResponseBody,
|
||||
IPWhitelist: group.IPWhitelist,
|
||||
IPBlacklist: group.IPBlacklist,
|
||||
CountryWhitelist: group.CountryWhitelist,
|
||||
CountryBlacklist: group.CountryBlacklist,
|
||||
RegionWhitelist: group.RegionWhitelist,
|
||||
RegionBlacklist: group.RegionBlacklist,
|
||||
})
|
||||
}
|
||||
sort.Slice(groups, func(i, j int) bool {
|
||||
if groups[i].IsGlobal != groups[j].IsGlobal {
|
||||
return groups[i].IsGlobal
|
||||
}
|
||||
return groups[i].ID < groups[j].ID
|
||||
})
|
||||
sort.Slice(globalGroupIDs, func(i, j int) bool { return globalGroupIDs[i] < globalGroupIDs[j] })
|
||||
siteRuleGroups := make(map[string][]uint, len(snapshot.Bindings))
|
||||
for _, binding := range snapshot.Bindings {
|
||||
ids := append([]uint{}, globalGroupIDs...)
|
||||
for _, id := range binding.RuleGroupIDs {
|
||||
if _, ok := enabledGroupIDs[id]; ok {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
}
|
||||
siteRuleGroups[binding.SiteName] = uniqueUintIDs(ids)
|
||||
}
|
||||
runtimeConfig := wafRuntimeConfig{
|
||||
DefaultBlockStatusCode: defaultWAFBlockStatusCode,
|
||||
RuleGroups: groups,
|
||||
SiteRuleGroups: siteRuleGroups,
|
||||
}
|
||||
data, err := json.Marshal(runtimeConfig)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@ func renderOpenRestyObservabilityTemplateBlock() string {
|
||||
" lua_shared_dict openflare_pow_config 1m;",
|
||||
" lua_shared_dict openflare_pow_challenges 10m;",
|
||||
" lua_shared_dict openflare_pow_sessions 20m;",
|
||||
" lua_shared_dict openflare_waf_config 2m;",
|
||||
fmt.Sprintf(" init_worker_by_lua_file %s/%s;", nginxLuaDirPlaceholder, openRestyObservabilityInitLuaPath),
|
||||
fmt.Sprintf(" log_by_lua_file %s/%s;", nginxLuaDirPlaceholder, openRestyObservabilityLogLuaPath),
|
||||
"",
|
||||
|
||||
@@ -0,0 +1,514 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"openflare/model"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultWAFBlockStatusCode = 418
|
||||
maxWAFBlockBodyBytes = 16 * 1024
|
||||
)
|
||||
|
||||
type WAFRuleGroupInput struct {
|
||||
Name string `json:"name"`
|
||||
Enabled bool `json:"enabled"`
|
||||
BlockStatusCode int `json:"block_status_code"`
|
||||
BlockResponseBody string `json:"block_response_body"`
|
||||
IPWhitelist []string `json:"ip_whitelist"`
|
||||
IPBlacklist []string `json:"ip_blacklist"`
|
||||
CountryWhitelist []string `json:"country_whitelist"`
|
||||
CountryBlacklist []string `json:"country_blacklist"`
|
||||
RegionWhitelist []string `json:"region_whitelist"`
|
||||
RegionBlacklist []string `json:"region_blacklist"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
type WAFRuleGroupView 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"`
|
||||
IPWhitelist []string `json:"ip_whitelist"`
|
||||
IPBlacklist []string `json:"ip_blacklist"`
|
||||
CountryWhitelist []string `json:"country_whitelist"`
|
||||
CountryBlacklist []string `json:"country_blacklist"`
|
||||
RegionWhitelist []string `json:"region_whitelist"`
|
||||
RegionBlacklist []string `json:"region_blacklist"`
|
||||
Remark string `json:"remark"`
|
||||
AppliedSiteIDs []uint `json:"applied_site_ids"`
|
||||
AppliedSiteCount int `json:"applied_site_count"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
type WAFSiteRuleGroupsView struct {
|
||||
RouteID uint `json:"route_id"`
|
||||
GlobalRuleGroup *WAFRuleGroupView `json:"global_rule_group"`
|
||||
RuleGroups []WAFRuleGroupView `json:"rule_groups"`
|
||||
AppliedRuleGroups []WAFRuleGroupView `json:"applied_rule_groups"`
|
||||
AppliedIDs []uint `json:"applied_ids"`
|
||||
}
|
||||
|
||||
func ListWAFRuleGroups() ([]WAFRuleGroupView, error) {
|
||||
if err := EnsureDefaultWAFRuleGroup(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups, err := model.ListWAFRuleGroups()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := loadWAFBindings()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]WAFRuleGroupView, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
view, err := buildWAFRuleGroupView(group, bindings[group.ID])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views = append(views, view)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
func GetWAFRuleGroup(id uint) (*WAFRuleGroupView, error) {
|
||||
group, err := model.GetWAFRuleGroupByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := loadWAFBindings()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
view, err := buildWAFRuleGroupView(group, bindings[group.ID])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &view, nil
|
||||
}
|
||||
|
||||
func CreateWAFRuleGroup(input WAFRuleGroupInput) (*WAFRuleGroupView, error) {
|
||||
group, err := buildWAFRuleGroup(nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
group.IsGlobal = false
|
||||
if err := group.Insert(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFRuleGroup(group.ID)
|
||||
}
|
||||
|
||||
func UpdateWAFRuleGroup(id uint, input WAFRuleGroupInput) (*WAFRuleGroupView, error) {
|
||||
group, err := model.GetWAFRuleGroupByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
isGlobal := group.IsGlobal
|
||||
group, err = buildWAFRuleGroup(group, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
group.IsGlobal = isGlobal
|
||||
if isGlobal && strings.TrimSpace(group.Name) == "" {
|
||||
group.Name = "全局规则组"
|
||||
}
|
||||
if err := group.Update(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFRuleGroup(group.ID)
|
||||
}
|
||||
|
||||
func DeleteWAFRuleGroup(id uint) error {
|
||||
group, err := model.GetWAFRuleGroupByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if group.IsGlobal {
|
||||
return errors.New("全局 WAF 规则组不能删除")
|
||||
}
|
||||
return model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("rule_group_id = ?", group.ID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(group).Error
|
||||
})
|
||||
}
|
||||
|
||||
func ReplaceWAFRuleGroupSites(groupID uint, routeIDs []uint) (*WAFRuleGroupView, error) {
|
||||
group, err := model.GetWAFRuleGroupByID(groupID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if group.IsGlobal {
|
||||
return nil, errors.New("全局 WAF 规则组默认应用到所有网站,不能手动绑定")
|
||||
}
|
||||
normalized, err := normalizeWAFRouteIDs(routeIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("rule_group_id = ?", groupID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, routeID := range normalized {
|
||||
binding := model.WAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID}
|
||||
if err := tx.Create(&binding).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFRuleGroup(groupID)
|
||||
}
|
||||
|
||||
func GetWAFSiteRuleGroups(routeID uint) (*WAFSiteRuleGroupsView, error) {
|
||||
if _, err := model.GetProxyRouteByID(routeID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups, err := ListWAFRuleGroups()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
appliedIDs, err := ListWAFSiteRuleGroupIDs(routeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
appliedSet := make(map[uint]struct{}, len(appliedIDs))
|
||||
for _, id := range appliedIDs {
|
||||
appliedSet[id] = struct{}{}
|
||||
}
|
||||
var global *WAFRuleGroupView
|
||||
custom := make([]WAFRuleGroupView, 0, len(groups))
|
||||
applied := make([]WAFRuleGroupView, 0, len(appliedIDs))
|
||||
for index := range groups {
|
||||
group := groups[index]
|
||||
if group.IsGlobal {
|
||||
item := group
|
||||
global = &item
|
||||
continue
|
||||
}
|
||||
custom = append(custom, group)
|
||||
if _, ok := appliedSet[group.ID]; ok {
|
||||
applied = append(applied, group)
|
||||
}
|
||||
}
|
||||
return &WAFSiteRuleGroupsView{
|
||||
RouteID: routeID,
|
||||
GlobalRuleGroup: global,
|
||||
RuleGroups: custom,
|
||||
AppliedRuleGroups: applied,
|
||||
AppliedIDs: appliedIDs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func ReplaceWAFSiteRuleGroups(routeID uint, groupIDs []uint) (*WAFSiteRuleGroupsView, error) {
|
||||
if _, err := model.GetProxyRouteByID(routeID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
normalized, err := normalizeWAFRuleGroupIDs(groupIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = model.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("proxy_route_id = ?", routeID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, groupID := range normalized {
|
||||
binding := model.WAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID}
|
||||
if err := tx.Create(&binding).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return GetWAFSiteRuleGroups(routeID)
|
||||
}
|
||||
|
||||
func ListWAFSiteRuleGroupIDs(routeID uint) ([]uint, error) {
|
||||
var bindings []model.WAFRuleGroupBinding
|
||||
if err := model.DB.Where("proxy_route_id = ?", routeID).Order("rule_group_id asc").Find(&bindings).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ids := make([]uint, 0, len(bindings))
|
||||
for _, binding := range bindings {
|
||||
ids = append(ids, binding.RuleGroupID)
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func EnsureDefaultWAFRuleGroup() error {
|
||||
_, err := model.GetGlobalWAFRuleGroup()
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return err
|
||||
}
|
||||
group := &model.WAFRuleGroup{
|
||||
Name: "全局规则组",
|
||||
Enabled: true,
|
||||
IsGlobal: true,
|
||||
BlockStatusCode: defaultWAFBlockStatusCode,
|
||||
IPWhitelist: "[]",
|
||||
IPBlacklist: "[]",
|
||||
CountryWhitelist: "[]",
|
||||
CountryBlacklist: "[]",
|
||||
RegionWhitelist: "[]",
|
||||
RegionBlacklist: "[]",
|
||||
BlockResponseBody: "",
|
||||
}
|
||||
return group.Insert()
|
||||
}
|
||||
|
||||
func buildWAFRuleGroup(group *model.WAFRuleGroup, input WAFRuleGroupInput) (*model.WAFRuleGroup, error) {
|
||||
name := strings.TrimSpace(input.Name)
|
||||
if name == "" {
|
||||
return nil, errors.New("规则组名称不能为空")
|
||||
}
|
||||
statusCode := input.BlockStatusCode
|
||||
if statusCode == 0 {
|
||||
statusCode = defaultWAFBlockStatusCode
|
||||
}
|
||||
if statusCode < 400 || statusCode > 599 {
|
||||
return nil, errors.New("拦截状态码必须在 400-599 之间")
|
||||
}
|
||||
if len([]byte(input.BlockResponseBody)) > maxWAFBlockBodyBytes {
|
||||
return nil, fmt.Errorf("拦截页面内容不能超过 %d 字节", maxWAFBlockBodyBytes)
|
||||
}
|
||||
ipWhitelist, err := normalizeWAFIPList(input.IPWhitelist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("IP 白名单无效: %w", err)
|
||||
}
|
||||
ipBlacklist, err := normalizeWAFIPList(input.IPBlacklist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("IP 黑名单无效: %w", err)
|
||||
}
|
||||
countryWhitelist, err := normalizeWAFCountryList(input.CountryWhitelist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("地域白名单无效: %w", err)
|
||||
}
|
||||
countryBlacklist, err := normalizeWAFCountryList(input.CountryBlacklist)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("地域黑名单无效: %w", err)
|
||||
}
|
||||
regionWhitelist := normalizeStringList(input.RegionWhitelist)
|
||||
regionBlacklist := normalizeStringList(input.RegionBlacklist)
|
||||
|
||||
ipWhitelistJSON, _ := json.Marshal(ipWhitelist)
|
||||
ipBlacklistJSON, _ := json.Marshal(ipBlacklist)
|
||||
countryWhitelistJSON, _ := json.Marshal(countryWhitelist)
|
||||
countryBlacklistJSON, _ := json.Marshal(countryBlacklist)
|
||||
regionWhitelistJSON, _ := json.Marshal(regionWhitelist)
|
||||
regionBlacklistJSON, _ := json.Marshal(regionBlacklist)
|
||||
|
||||
if group == nil {
|
||||
group = &model.WAFRuleGroup{}
|
||||
}
|
||||
group.Name = name
|
||||
group.Enabled = input.Enabled
|
||||
group.BlockStatusCode = statusCode
|
||||
group.BlockResponseBody = input.BlockResponseBody
|
||||
group.IPWhitelist = string(ipWhitelistJSON)
|
||||
group.IPBlacklist = string(ipBlacklistJSON)
|
||||
group.CountryWhitelist = string(countryWhitelistJSON)
|
||||
group.CountryBlacklist = string(countryBlacklistJSON)
|
||||
group.RegionWhitelist = string(regionWhitelistJSON)
|
||||
group.RegionBlacklist = string(regionBlacklistJSON)
|
||||
group.Remark = strings.TrimSpace(input.Remark)
|
||||
return group, nil
|
||||
}
|
||||
|
||||
func buildWAFRuleGroupView(group *model.WAFRuleGroup, appliedSiteIDs []uint) (WAFRuleGroupView, error) {
|
||||
if group == nil {
|
||||
return WAFRuleGroupView{}, errors.New("waf rule group is nil")
|
||||
}
|
||||
sort.Slice(appliedSiteIDs, func(i, j int) bool { return appliedSiteIDs[i] < appliedSiteIDs[j] })
|
||||
view := WAFRuleGroupView{
|
||||
ID: group.ID,
|
||||
Name: group.Name,
|
||||
Enabled: group.Enabled,
|
||||
IsGlobal: group.IsGlobal,
|
||||
BlockStatusCode: group.BlockStatusCode,
|
||||
BlockResponseBody: group.BlockResponseBody,
|
||||
Remark: group.Remark,
|
||||
AppliedSiteIDs: appliedSiteIDs,
|
||||
AppliedSiteCount: len(appliedSiteIDs),
|
||||
CreatedAt: group.CreatedAt.Format(time.RFC3339),
|
||||
UpdatedAt: group.UpdatedAt.Format(time.RFC3339),
|
||||
}
|
||||
var err error
|
||||
if view.IPWhitelist, err = decodeStringList(group.IPWhitelist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.IPBlacklist, err = decodeStringList(group.IPBlacklist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.CountryWhitelist, err = decodeStringList(group.CountryWhitelist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.CountryBlacklist, err = decodeStringList(group.CountryBlacklist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.RegionWhitelist, err = decodeStringList(group.RegionWhitelist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if view.RegionBlacklist, err = decodeStringList(group.RegionBlacklist); err != nil {
|
||||
return view, err
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
func loadWAFBindings() (map[uint][]uint, error) {
|
||||
var bindings []model.WAFRuleGroupBinding
|
||||
if err := model.DB.Order("rule_group_id asc").Order("proxy_route_id asc").Find(&bindings).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make(map[uint][]uint, len(bindings))
|
||||
for _, binding := range bindings {
|
||||
result[binding.RuleGroupID] = append(result[binding.RuleGroupID], binding.ProxyRouteID)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func normalizeWAFIPList(items []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(items))
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
for _, raw := range items {
|
||||
item := strings.TrimSpace(raw)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(item, "/") {
|
||||
prefix, err := netip.ParsePrefix(item)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s 不是合法 IP 段", item)
|
||||
}
|
||||
item = prefix.Masked().String()
|
||||
} else {
|
||||
addr, err := netip.ParseAddr(item)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s 不是合法 IP", item)
|
||||
}
|
||||
item = addr.String()
|
||||
}
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
sort.Strings(normalized)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeWAFCountryList(items []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(items))
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
for _, raw := range items {
|
||||
item := strings.ToUpper(strings.TrimSpace(raw))
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
if len(item) != 2 || !unicode.IsLetter(rune(item[0])) || !unicode.IsLetter(rune(item[1])) {
|
||||
return nil, fmt.Errorf("%s 不是合法国家代码", item)
|
||||
}
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
sort.Strings(normalized)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeStringList(items []string) []string {
|
||||
normalized := make([]string, 0, len(items))
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
for _, raw := range items {
|
||||
item := strings.TrimSpace(raw)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
normalized = append(normalized, item)
|
||||
}
|
||||
sort.Strings(normalized)
|
||||
return normalized
|
||||
}
|
||||
|
||||
func decodeStringList(raw string) ([]string, error) {
|
||||
text := strings.TrimSpace(raw)
|
||||
if text == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(text), &items); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func normalizeWAFRouteIDs(routeIDs []uint) ([]uint, error) {
|
||||
normalized := uniqueUintIDs(routeIDs)
|
||||
for _, routeID := range normalized {
|
||||
if _, err := model.GetProxyRouteByID(routeID); err != nil {
|
||||
return nil, fmt.Errorf("网站 %d 不存在", routeID)
|
||||
}
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeWAFRuleGroupIDs(groupIDs []uint) ([]uint, error) {
|
||||
normalized := uniqueUintIDs(groupIDs)
|
||||
for _, groupID := range normalized {
|
||||
group, err := model.GetWAFRuleGroupByID(groupID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID)
|
||||
}
|
||||
if group.IsGlobal {
|
||||
return nil, errors.New("全局 WAF 规则组不需要手动绑定")
|
||||
}
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func uniqueUintIDs(ids []uint) []uint {
|
||||
seen := make(map[uint]struct{}, len(ids))
|
||||
normalized := make([]uint, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if id == 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
sort.Slice(normalized, func(i, j int) bool { return normalized[i] < normalized[j] })
|
||||
return normalized
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWAFRuleGroupValidationAndNormalization(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
group, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
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"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if len(group.IPWhitelist) != 2 || group.IPWhitelist[0] != "192.0.2.1" || group.IPWhitelist[1] != "198.51.100.0/24" {
|
||||
t.Fatalf("unexpected normalized ip whitelist: %#v", group.IPWhitelist)
|
||||
}
|
||||
if len(group.CountryBlacklist) != 2 || group.CountryBlacklist[0] != "CN" || group.CountryBlacklist[1] != "US" {
|
||||
t.Fatalf("unexpected normalized countries: %#v", group.CountryBlacklist)
|
||||
}
|
||||
|
||||
if _, err = CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "bad ip",
|
||||
Enabled: true,
|
||||
IPBlacklist: []string{"not-an-ip"},
|
||||
}); err == nil {
|
||||
t.Fatal("expected invalid IP to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWAFGlobalGroupAndBindings(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
groups, err := ListWAFRuleGroups()
|
||||
if err != nil {
|
||||
t.Fatalf("ListWAFRuleGroups failed: %v", err)
|
||||
}
|
||||
if len(groups) == 0 || !groups[0].IsGlobal {
|
||||
t.Fatalf("expected default global WAF rule group, got %#v", groups)
|
||||
}
|
||||
if err = DeleteWAFRuleGroup(groups[0].ID); err == nil {
|
||||
t.Fatal("expected global WAF rule group delete to be rejected")
|
||||
}
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "waf-site",
|
||||
Domains: []string{"waf.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
custom, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "custom",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 418,
|
||||
IPBlacklist: []string{"203.0.113.10"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if _, err = ReplaceWAFRuleGroupSites(custom.ID, []uint{route.ID}); err != nil {
|
||||
t.Fatalf("ReplaceWAFRuleGroupSites failed: %v", err)
|
||||
}
|
||||
siteGroups, err := GetWAFSiteRuleGroups(route.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetWAFSiteRuleGroups failed: %v", err)
|
||||
}
|
||||
if len(siteGroups.AppliedIDs) != 1 || siteGroups.AppliedIDs[0] != custom.ID {
|
||||
t.Fatalf("unexpected site WAF bindings: %#v", siteGroups.AppliedIDs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishConfigVersionIncludesWAFSnapshotAndRuntimeConfig(t *testing.T) {
|
||||
setupServiceTestDB(t)
|
||||
|
||||
route, err := CreateProxyRoute(ProxyRouteInput{
|
||||
SiteName: "waf-publish",
|
||||
Domains: []string{"waf-publish.example.com"},
|
||||
OriginURL: "https://origin.internal",
|
||||
Enabled: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateProxyRoute failed: %v", err)
|
||||
}
|
||||
group, err := CreateWAFRuleGroup(WAFRuleGroupInput{
|
||||
Name: "publish group",
|
||||
Enabled: true,
|
||||
BlockStatusCode: 451,
|
||||
IPBlacklist: []string{"203.0.113.0/24"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateWAFRuleGroup failed: %v", err)
|
||||
}
|
||||
if _, err = ReplaceWAFSiteRuleGroups(route.ID, []uint{group.ID}); err != nil {
|
||||
t.Fatalf("ReplaceWAFSiteRuleGroups failed: %v", err)
|
||||
}
|
||||
result, err := PublishConfigVersion("root", false)
|
||||
if err != nil {
|
||||
t.Fatalf("PublishConfigVersion failed: %v", err)
|
||||
}
|
||||
if !strings.Contains(result.Version.RenderedConfig, "access_by_lua_file __OPENFLARE_LUA_DIR__/waf/check.lua;") {
|
||||
t.Fatal("expected route config to include WAF lua access hook")
|
||||
}
|
||||
if !strings.Contains(result.Version.SnapshotJSON, `"waf"`) {
|
||||
t.Fatal("expected snapshot to include waf document")
|
||||
}
|
||||
var files []SupportFile
|
||||
if err = json.Unmarshal([]byte(result.Version.SupportFilesJSON), &files); err != nil {
|
||||
t.Fatalf("decode support files failed: %v", err)
|
||||
}
|
||||
found := false
|
||||
for _, file := range files {
|
||||
if file.Path == "waf_config.json" && strings.Contains(file.Content, "203.0.113.0/24") {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("expected waf_config.json support file, got %#v", files)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user