[新增] 添加 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
+138
View File
@@ -0,0 +1,138 @@
package controller
import (
"encoding/json"
"net/http"
"openflare/service"
"strconv"
"github.com/gin-gonic/gin"
)
type wafIDsRequest struct {
IDs []uint `json:"ids"`
}
func ListWAFRuleGroups(c *gin.Context) {
groups, err := service.ListWAFRuleGroups()
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": groups})
}
func GetWAFRuleGroup(c *gin.Context) {
id, ok := parseUintPathParam(c, "id")
if !ok {
return
}
group, err := service.GetWAFRuleGroup(id)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group})
}
func CreateWAFRuleGroup(c *gin.Context) {
var input service.WAFRuleGroupInput
if err := json.NewDecoder(c.Request.Body).Decode(&input); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"})
return
}
group, err := service.CreateWAFRuleGroup(input)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group})
}
func UpdateWAFRuleGroup(c *gin.Context) {
id, ok := parseUintPathParam(c, "id")
if !ok {
return
}
var input service.WAFRuleGroupInput
if err := json.NewDecoder(c.Request.Body).Decode(&input); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"})
return
}
group, err := service.UpdateWAFRuleGroup(id, input)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group})
}
func DeleteWAFRuleGroup(c *gin.Context) {
id, ok := parseUintPathParam(c, "id")
if !ok {
return
}
if err := service.DeleteWAFRuleGroup(id); err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": ""})
}
func ReplaceWAFRuleGroupSites(c *gin.Context) {
id, ok := parseUintPathParam(c, "id")
if !ok {
return
}
var request wafIDsRequest
if err := json.NewDecoder(c.Request.Body).Decode(&request); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"})
return
}
group, err := service.ReplaceWAFRuleGroupSites(id, request.IDs)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group})
}
func GetWAFSiteRuleGroups(c *gin.Context) {
routeID, ok := parseUintPathParam(c, "route_id")
if !ok {
return
}
view, err := service.GetWAFSiteRuleGroups(routeID)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": view})
}
func ReplaceWAFSiteRuleGroups(c *gin.Context) {
routeID, ok := parseUintPathParam(c, "route_id")
if !ok {
return
}
var request wafIDsRequest
if err := json.NewDecoder(c.Request.Body).Decode(&request); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"})
return
}
view, err := service.ReplaceWAFSiteRuleGroups(routeID, request.IDs)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": view})
}
func parseUintPathParam(c *gin.Context, name string) (uint, bool) {
id, err := strconv.ParseUint(c.Param(name), 10, 64)
if err != nil || id == 0 {
c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid id"})
return 0, false
}
return uint(id), true
}
@@ -4,7 +4,7 @@ import "time"
const (
legacyDatabaseSchemaVersion = 1
currentDatabaseSchemaVersion = 12
currentDatabaseSchemaVersion = 13
databaseSchemaVersionRowID = 1
)
+2
View File
@@ -43,6 +43,8 @@ func registeredModels() []any {
&ManagedDomain{},
&AcmeAccount{},
&DnsAccount{},
&WAFRuleGroup{},
&WAFRuleGroupBinding{},
}
}
+66 -1
View File
@@ -1388,6 +1388,67 @@ func validateDatabaseSchemaV12(db *gorm.DB, backend string) error {
return nil
}
func ensureDefaultWAFRuleGroup(db *gorm.DB) error {
if db == nil {
return fmt.Errorf("database handle is nil")
}
if !db.Migrator().HasTable(&WAFRuleGroup{}) {
return nil
}
var count int64
if err := db.Model(&WAFRuleGroup{}).Where("is_global = ?", true).Count(&count).Error; err != nil {
return fmt.Errorf("count global waf rule groups failed: %w", err)
}
if count > 0 {
return nil
}
group := WAFRuleGroup{
Name: "全局规则组",
Enabled: true,
IsGlobal: true,
BlockStatusCode: 418,
IPWhitelist: "[]",
IPBlacklist: "[]",
CountryWhitelist: "[]",
CountryBlacklist: "[]",
RegionWhitelist: "[]",
RegionBlacklist: "[]",
BlockResponseBody: "",
}
if err := db.Create(&group).Error; err != nil {
return fmt.Errorf("create default waf rule group failed: %w", err)
}
return nil
}
// migrateV13 adds WAF rule groups and website bindings.
func migrateV13(db *gorm.DB, backend string) error {
if err := applyCurrentSchema(db, backend); err != nil {
return err
}
return ensureDefaultWAFRuleGroup(db)
}
func validateDatabaseSchemaV13(db *gorm.DB, backend string) error {
if err := validateDatabaseSchemaV12(db, backend); err != nil {
return err
}
if !db.Migrator().HasTable(&WAFRuleGroup{}) {
return fmt.Errorf("table waf_rule_groups is missing")
}
if !db.Migrator().HasTable(&WAFRuleGroupBinding{}) {
return fmt.Errorf("table waf_rule_group_bindings is missing")
}
var count int64
if err := db.Model(&WAFRuleGroup{}).Where("is_global = ?", true).Count(&count).Error; err != nil {
return fmt.Errorf("count global waf rule groups failed: %w", err)
}
if count != 1 {
return fmt.Errorf("expected exactly one global waf rule group, got %d", count)
}
return nil
}
func databaseSchemaMigrations() []databaseSchemaMigration {
return []databaseSchemaMigration{
{fromVersion: 1, toVersion: 2, migrate: migrateV2, validate: validateDatabaseSchemaV2},
@@ -1401,6 +1462,7 @@ func databaseSchemaMigrations() []databaseSchemaMigration {
{fromVersion: 9, toVersion: 10, migrate: migrateV10, validate: validateDatabaseSchemaV10},
{fromVersion: 10, toVersion: 11, migrate: migrateV11, validate: validateDatabaseSchemaV11},
{fromVersion: 11, toVersion: 12, migrate: migrateV12, validate: validateDatabaseSchemaV12},
{fromVersion: 12, toVersion: 13, migrate: migrateV13, validate: validateDatabaseSchemaV13},
}
}
@@ -1486,7 +1548,10 @@ func initializeFreshDatabaseSchema(db *gorm.DB, backend string) error {
if err := ensureDefaultGitHubAuthSource(db); err != nil {
return err
}
if err := validateDatabaseSchemaV12(db, backend); err != nil {
if err := ensureDefaultWAFRuleGroup(db); err != nil {
return err
}
if err := validateDatabaseSchemaV13(db, backend); err != nil {
return err
}
return saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion)
+71
View File
@@ -0,0 +1,71 @@
package model
import "time"
type WAFRuleGroup struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"`
BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"`
BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"`
IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"`
IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"`
CountryWhitelist string `json:"country_whitelist" gorm:"type:text;not null;default:'[]'"`
CountryBlacklist string `json:"country_blacklist" gorm:"type:text;not null;default:'[]'"`
RegionWhitelist string `json:"region_whitelist" gorm:"type:text;not null;default:'[]'"`
RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type WAFRuleGroupBinding struct {
ID uint `json:"id" gorm:"primaryKey"`
RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_waf_group_route"`
ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_waf_group_route;index"`
CreatedAt time.Time `json:"created_at"`
}
func ListWAFRuleGroups() ([]*WAFRuleGroup, error) {
var groups []*WAFRuleGroup
err := DB.Order("is_global desc").Order("id asc").Find(&groups).Error
return groups, err
}
func GetWAFRuleGroupByID(id uint) (*WAFRuleGroup, error) {
group := &WAFRuleGroup{}
err := DB.First(group, id).Error
return group, err
}
func GetGlobalWAFRuleGroup() (*WAFRuleGroup, error) {
group := &WAFRuleGroup{}
err := DB.Where("is_global = ?", true).Order("id asc").First(group).Error
return group, err
}
func (group *WAFRuleGroup) Insert() error {
return DB.Create(group).Error
}
func (group *WAFRuleGroup) Update() error {
return DB.Model(&WAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"enabled": group.Enabled,
"is_global": group.IsGlobal,
"block_status_code": group.BlockStatusCode,
"block_response_body": group.BlockResponseBody,
"ip_whitelist": group.IPWhitelist,
"ip_blacklist": group.IPBlacklist,
"country_whitelist": group.CountryWhitelist,
"country_blacklist": group.CountryBlacklist,
"region_whitelist": group.RegionWhitelist,
"region_blacklist": group.RegionBlacklist,
"remark": group.Remark,
}).Error
}
func (group *WAFRuleGroup) Delete() error {
return DB.Delete(group).Error
}
+12
View File
@@ -94,6 +94,18 @@ func SetApiRouter(router *gin.Engine) {
proxyRoute.POST("/:id/update", controller.UpdateProxyRoute)
proxyRoute.POST("/:id/delete", controller.DeleteProxyRoute)
}
wafRoute := apiRouter.Group("/waf")
wafRoute.Use(middleware.AdminAuth())
{
wafRoute.GET("/rule-groups", controller.ListWAFRuleGroups)
wafRoute.GET("/rule-groups/:id", controller.GetWAFRuleGroup)
wafRoute.POST("/rule-groups", controller.CreateWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/update", controller.UpdateWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/delete", controller.DeleteWAFRuleGroup)
wafRoute.POST("/rule-groups/:id/sites", controller.ReplaceWAFRuleGroupSites)
wafRoute.GET("/sites/:route_id/rule-groups", controller.GetWAFSiteRuleGroups)
wafRoute.POST("/sites/:route_id/rule-groups", controller.ReplaceWAFSiteRuleGroups)
}
originRoute := apiRouter.Group("/origins")
originRoute.Use(middleware.AdminAuth())
{
+220 -10
View File
@@ -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),
"",
+514
View File
@@ -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
}
+130
View File
@@ -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)
}
}
+30 -8
View File
@@ -33,8 +33,18 @@ func (s *MaxMindGeoIPService) Name() string {
}
func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) {
return NewMaxMindGeoIPServiceWithConfig(GeoIpFilePath, GeoIpUrl)
}
func NewMaxMindGeoIPServiceWithConfig(dbFilePath string, downloadURL string) (*MaxMindGeoIPService, error) {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
service := &MaxMindGeoIPService{
dbFilePath: GeoIpFilePath,
dbFilePath: dbFilePath,
}
if err := os.MkdirAll(filepath.Dir(service.dbFilePath), os.ModePerm); err != nil {
@@ -42,7 +52,7 @@ func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) {
}
if _, err := os.Stat(service.dbFilePath); os.IsNotExist(err) {
if err := service.UpdateDatabase(); err != nil {
if err := DownloadMaxMindDatabase(service.dbFilePath, downloadURL); err != nil {
return nil, fmt.Errorf("failed to download initial MaxMind database: %w", err)
}
}
@@ -99,7 +109,20 @@ func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
}
func (s *MaxMindGeoIPService) UpdateDatabase() error {
resp, err := http.Get(GeoIpUrl)
if err := DownloadMaxMindDatabase(s.dbFilePath, GeoIpUrl); err != nil {
return err
}
return s.initialize()
}
func DownloadMaxMindDatabase(dbFilePath string, downloadURL string) error {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
resp, err := http.Get(downloadURL)
if err != nil {
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
}
@@ -109,11 +132,11 @@ func (s *MaxMindGeoIPService) UpdateDatabase() error {
return fmt.Errorf("failed to download MaxMind database: HTTP status %s", resp.Status)
}
if err := os.MkdirAll(filepath.Dir(s.dbFilePath), os.ModePerm); err != nil {
if err := os.MkdirAll(filepath.Dir(dbFilePath), os.ModePerm); err != nil {
return fmt.Errorf("failed to create data directory for MaxMind database update: %w", err)
}
tempPath := s.dbFilePath + ".download"
tempPath := dbFilePath + ".download"
out, err := os.Create(tempPath)
if err != nil {
return fmt.Errorf("failed to create MaxMind database file at %s: %w", tempPath, err)
@@ -128,11 +151,10 @@ func (s *MaxMindGeoIPService) UpdateDatabase() error {
if err = out.Close(); err != nil {
return fmt.Errorf("failed to close MaxMind database file: %w", err)
}
if err = os.Rename(tempPath, s.dbFilePath); err != nil {
if err = os.Rename(tempPath, dbFilePath); err != nil {
return fmt.Errorf("failed to move MaxMind database file into place: %w", err)
}
return s.initialize()
return nil
}
func (s *MaxMindGeoIPService) Close() error {
@@ -0,0 +1,5 @@
import { WAFPage } from '@/features/waf/components/waf-page';
export default function WAFRoute() {
return <WAFPage />;
}
@@ -3,6 +3,7 @@
import { useEffect } from 'react';
import Link from 'next/link';
import { usePathname } from 'next/navigation';
import { ShieldCheck } from 'lucide-react';
import { dashboardNavigation } from '@/lib/constants/navigation';
import { cn } from '@/lib/utils/cn';
@@ -79,6 +80,8 @@ function SidebarIcon({ icon }: { icon: NavigationIconKey }) {
<path d="m14 15 3 2-3 2" />
</svg>
);
case 'waf':
return <ShieldCheck className="h-[18px] w-[18px]" strokeWidth={1.8} />;
case 'release':
return (
<svg {...commonProps}>
@@ -39,6 +39,7 @@ export interface ConfigDiffResult {
removed_domains: string[];
modified_domains: string[];
main_config_changed: boolean;
waf_config_changed: boolean;
changed_option_keys: string[];
changed_option_details: ConfigOptionDiffItem[];
current_website_count: number;
@@ -28,7 +28,6 @@ import {
requestNodeForceSync,
requestNodeOpenrestyRestart,
requestNodeAgentUpdate,
rotateNodeBootstrapToken,
updateNode,
} from '@/features/nodes/api/nodes';
import { NodeEditorModal } from '@/features/nodes/components/node-editor-modal';
@@ -48,6 +48,7 @@ function hasConfigChanges(diff: {
removed_domains: string[];
modified_domains: string[];
main_config_changed: boolean;
waf_config_changed?: boolean;
changed_option_keys: string[];
}) {
return (
@@ -58,6 +59,7 @@ function hasConfigChanges(diff: {
diff.removed_domains.length > 0 ||
diff.modified_domains.length > 0 ||
diff.main_config_changed ||
Boolean(diff.waf_config_changed) ||
diff.changed_option_keys.length > 0 ||
!diff.active_version
);
@@ -0,0 +1,49 @@
import { apiRequest } from '@/lib/api/client';
import type {
WAFRuleGroup,
WAFRuleGroupPayload,
WAFSiteRuleGroups,
} from '@/features/waf/types';
export function getWAFRuleGroups() {
return apiRequest<WAFRuleGroup[]>('/waf/rule-groups');
}
export function createWAFRuleGroup(payload: WAFRuleGroupPayload) {
return apiRequest<WAFRuleGroup>('/waf/rule-groups', {
method: 'POST',
body: JSON.stringify(payload),
});
}
export function updateWAFRuleGroup(id: number, payload: WAFRuleGroupPayload) {
return apiRequest<WAFRuleGroup>(`/waf/rule-groups/${id}/update`, {
method: 'POST',
body: JSON.stringify(payload),
});
}
export function deleteWAFRuleGroup(id: number) {
return apiRequest<void>(`/waf/rule-groups/${id}/delete`, {
method: 'POST',
});
}
export function replaceWAFRuleGroupSites(id: number, ids: number[]) {
return apiRequest<WAFRuleGroup>(`/waf/rule-groups/${id}/sites`, {
method: 'POST',
body: JSON.stringify({ ids }),
});
}
export function getWAFSiteRuleGroups(routeId: number) {
return apiRequest<WAFSiteRuleGroups>(`/waf/sites/${routeId}/rule-groups`);
}
export function replaceWAFSiteRuleGroups(routeId: number, ids: number[]) {
return apiRequest<WAFSiteRuleGroups>(`/waf/sites/${routeId}/rule-groups`, {
method: 'POST',
body: JSON.stringify({ ids }),
});
}
@@ -0,0 +1,525 @@
'use client';
import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query';
import { Check, Globe2, Plus, Save, Search, ShieldCheck, Trash2 } from 'lucide-react';
import { useEffect, useMemo, useState } from 'react';
import { EmptyState } from '@/components/feedback/empty-state';
import { ErrorState } from '@/components/feedback/error-state';
import { InlineMessage } from '@/components/feedback/inline-message';
import { LoadingState } from '@/components/feedback/loading-state';
import { PageHeader } from '@/components/layout/page-header';
import { AppCard } from '@/components/ui/app-card';
import { Drawer } from '@/components/ui/drawer';
import { getProxyRoutes } from '@/features/proxy-routes/api/proxy-routes';
import type { ProxyRouteItem } from '@/features/proxy-routes/types';
import {
DangerButton,
PrimaryButton,
ResourceField,
ResourceInput,
ResourceTextarea,
SecondaryButton,
ToggleField,
} from '@/features/shared/components/resource-primitives';
import {
createWAFRuleGroup,
deleteWAFRuleGroup,
getWAFRuleGroups,
replaceWAFRuleGroupSites,
updateWAFRuleGroup,
} from '@/features/waf/api/waf';
import type { WAFRuleGroup, WAFRuleGroupPayload } from '@/features/waf/types';
import { cn } from '@/lib/utils/cn';
type FeedbackState = {
tone: 'success' | 'danger' | 'info';
message: string;
};
const emptyDraft: WAFRuleGroupPayload = {
name: '',
enabled: true,
block_status_code: 418,
block_response_body: '',
ip_whitelist: [],
ip_blacklist: [],
country_whitelist: [],
country_blacklist: [],
region_whitelist: [],
region_blacklist: [],
remark: '',
};
function getErrorMessage(error: unknown) {
return error instanceof Error ? error.message : '操作失败';
}
function listToText(items: string[]) {
return items.join('\n');
}
function textToList(text: string) {
return text
.split(/[\n,,\s]+/)
.map((item) => item.trim())
.filter(Boolean);
}
function buildDraft(group: WAFRuleGroup | null): WAFRuleGroupPayload {
if (!group) {
return { ...emptyDraft };
}
return {
name: group.name,
enabled: group.enabled,
block_status_code: group.block_status_code || 418,
block_response_body: group.block_response_body ?? '',
ip_whitelist: group.ip_whitelist ?? [],
ip_blacklist: group.ip_blacklist ?? [],
country_whitelist: group.country_whitelist ?? [],
country_blacklist: group.country_blacklist ?? [],
region_whitelist: group.region_whitelist ?? [],
region_blacklist: group.region_blacklist ?? [],
remark: group.remark ?? '',
};
}
function ruleCount(group: WAFRuleGroup) {
return (
group.ip_whitelist.length +
group.ip_blacklist.length +
group.country_whitelist.length +
group.country_blacklist.length
);
}
function SiteApplyDrawer({
group,
routes,
open,
onOpenChange,
onSave,
pending,
}: {
group: WAFRuleGroup | null;
routes: ProxyRouteItem[];
open: boolean;
onOpenChange: (open: boolean) => void;
onSave: (ids: number[]) => void;
pending: boolean;
}) {
const [keyword, setKeyword] = useState('');
const [selectedIDs, setSelectedIDs] = useState<number[]>([]);
useEffect(() => {
setSelectedIDs(group?.applied_site_ids ?? []);
setKeyword('');
}, [group, open]);
const filteredRoutes = useMemo(() => {
const normalized = keyword.trim().toLowerCase();
if (!normalized) {
return routes;
}
return routes.filter((route) =>
[route.site_name, route.primary_domain, ...route.domains]
.join(' ')
.toLowerCase()
.includes(normalized),
);
}, [keyword, routes]);
const selectedSet = useMemo(() => new Set(selectedIDs), [selectedIDs]);
const toggleID = (id: number) => {
setSelectedIDs((current) =>
current.includes(id)
? current.filter((item) => item !== id)
: [...current, id].sort((left, right) => left - right),
);
};
const selectFiltered = () => {
const next = new Set(selectedIDs);
filteredRoutes.forEach((route) => next.add(route.id));
setSelectedIDs([...next].sort((left, right) => left - right));
};
return (
<Drawer
open={open}
onOpenChange={onOpenChange}
direction="right"
title={group ? `应用 ${group.name}` : '应用规则组'}
description="选择这个自定义规则组要叠加到哪些网站。"
footer={
<div className="flex justify-end gap-3">
<SecondaryButton type="button" onClick={() => onOpenChange(false)}>
取消
</SecondaryButton>
<PrimaryButton
type="button"
disabled={!group || pending}
onClick={() => onSave(selectedIDs)}
>
{pending ? '保存中...' : '保存应用范围'}
</PrimaryButton>
</div>
}
>
<div className="space-y-4">
<div className="flex items-center gap-3 rounded-2xl border border-[var(--border-default)] bg-[var(--surface-elevated)] px-4 py-3">
<Search className="h-4 w-4 text-[var(--foreground-secondary)]" />
<input
value={keyword}
onChange={(event) => setKeyword(event.target.value)}
placeholder="搜索网站或域名"
className="min-w-0 flex-1 bg-transparent text-sm text-[var(--foreground-primary)] outline-none placeholder:text-[var(--foreground-muted)]"
/>
<button
type="button"
onClick={selectFiltered}
className="text-xs font-medium text-[var(--brand-primary)]"
>
全选当前
</button>
</div>
<div className="space-y-2">
{filteredRoutes.map((route) => (
<button
key={route.id}
type="button"
onClick={() => toggleID(route.id)}
className={cn(
'flex w-full items-center gap-3 rounded-2xl border px-4 py-3 text-left transition',
selectedSet.has(route.id)
? 'border-[var(--border-strong)] bg-[var(--accent-soft)]'
: 'border-[var(--border-default)] bg-[var(--surface-elevated)] hover:bg-[var(--surface-muted)]',
)}
>
<span
className={cn(
'flex h-5 w-5 items-center justify-center rounded-md border',
selectedSet.has(route.id)
? 'border-[var(--brand-primary)] bg-[var(--brand-primary)] text-[var(--foreground-inverse)]'
: 'border-[var(--border-default)]',
)}
>
{selectedSet.has(route.id) ? <Check className="h-3 w-3" /> : null}
</span>
<span className="min-w-0 flex-1">
<span className="block truncate text-sm font-medium text-[var(--foreground-primary)]">
{route.site_name}
</span>
<span className="block truncate text-xs text-[var(--foreground-secondary)]">
{route.domains.join(', ')}
</span>
</span>
</button>
))}
</div>
</div>
</Drawer>
);
}
export function WAFPage() {
const queryClient = useQueryClient();
const [selectedID, setSelectedID] = useState<number | null>(null);
const [draft, setDraft] = useState<WAFRuleGroupPayload>(emptyDraft);
const [feedback, setFeedback] = useState<FeedbackState | null>(null);
const [applyGroup, setApplyGroup] = useState<WAFRuleGroup | null>(null);
const groupsQuery = useQuery({
queryKey: ['waf', 'rule-groups'],
queryFn: getWAFRuleGroups,
});
const routesQuery = useQuery({
queryKey: ['proxy-routes'],
queryFn: getProxyRoutes,
});
const groups = useMemo(() => groupsQuery.data ?? [], [groupsQuery.data]);
const routes = useMemo(() => routesQuery.data ?? [], [routesQuery.data]);
const selectedGroup = useMemo(
() =>
selectedID === 0
? null
: (groups.find((group) => group.id === selectedID) ?? groups[0] ?? null),
[groups, selectedID],
);
useEffect(() => {
if (selectedGroup) {
setSelectedID(selectedGroup.id);
setDraft(buildDraft(selectedGroup));
}
}, [selectedGroup]);
const invalidate = async () => {
await Promise.all([
queryClient.invalidateQueries({ queryKey: ['waf', 'rule-groups'] }),
queryClient.invalidateQueries({ queryKey: ['config-versions', 'diff'] }),
]);
};
const saveMutation = useMutation({
mutationFn: (payload: WAFRuleGroupPayload) => {
if (selectedGroup) {
return updateWAFRuleGroup(selectedGroup.id, payload);
}
return createWAFRuleGroup(payload);
},
onSuccess: async (group) => {
setSelectedID(group.id);
setFeedback({ tone: 'success', message: 'WAF 规则组已保存。' });
await invalidate();
},
onError: (error) => {
setFeedback({ tone: 'danger', message: getErrorMessage(error) });
},
});
const deleteMutation = useMutation({
mutationFn: deleteWAFRuleGroup,
onSuccess: async () => {
setSelectedID(null);
setFeedback({ tone: 'success', message: 'WAF 规则组已删除。' });
await invalidate();
},
onError: (error) => {
setFeedback({ tone: 'danger', message: getErrorMessage(error) });
},
});
const applyMutation = useMutation({
mutationFn: ({ id, ids }: { id: number; ids: number[] }) =>
replaceWAFRuleGroupSites(id, ids),
onSuccess: async () => {
setApplyGroup(null);
setFeedback({ tone: 'success', message: '规则组应用范围已更新。' });
await invalidate();
},
onError: (error) => {
setFeedback({ tone: 'danger', message: getErrorMessage(error) });
},
});
if (groupsQuery.isLoading || routesQuery.isLoading) {
return <LoadingState />;
}
if (groupsQuery.isError) {
return <ErrorState title="WAF 加载失败" description={getErrorMessage(groupsQuery.error)} />;
}
if (routesQuery.isError) {
return <ErrorState title="网站列表加载失败" description={getErrorMessage(routesQuery.error)} />;
}
if (!selectedGroup && groups.length === 0) {
return <EmptyState title="WAF 尚未初始化" description="刷新页面后系统会自动创建全局规则组。" />;
}
const enabledCount = groups.filter((group) => group.enabled).length;
const protectedSites = new Set(groups.flatMap((group) => group.applied_site_ids));
const totalRules = groups.reduce((sum, group) => sum + ruleCount(group), 0);
return (
<>
<div className="space-y-6">
<PageHeader
title="WAF"
description="按规则组维护 IP 与地域黑白名单,全局规则始终应用到所有网站。"
action={
<PrimaryButton
type="button"
onClick={() => {
setSelectedID(0);
setDraft({ ...emptyDraft, name: '自定义规则组' });
}}
>
<Plus className="mr-2 h-4 w-4" />
新建规则组
</PrimaryButton>
}
/>
{feedback ? <InlineMessage tone={feedback.tone} message={feedback.message} /> : null}
<div className="grid gap-4 xl:grid-cols-3">
<AppCard>
<p className="text-sm text-[var(--foreground-secondary)]">启用规则组</p>
<p className="mt-2 text-3xl font-semibold text-[var(--foreground-primary)]">{enabledCount}</p>
</AppCard>
<AppCard>
<p className="text-sm text-[var(--foreground-secondary)]">自定义覆盖网站</p>
<p className="mt-2 text-3xl font-semibold text-[var(--foreground-primary)]">{protectedSites.size}</p>
</AppCard>
<AppCard>
<p className="text-sm text-[var(--foreground-secondary)]">黑白名单条目</p>
<p className="mt-2 text-3xl font-semibold text-[var(--foreground-primary)]">{totalRules}</p>
</AppCard>
</div>
<div className="grid gap-5 xl:grid-cols-[360px_minmax(0,1fr)]">
<AppCard title="规则组">
<div className="space-y-2">
{groups.map((group) => (
<button
key={group.id}
type="button"
onClick={() => setSelectedID(group.id)}
className={cn(
'w-full rounded-2xl border px-4 py-3 text-left transition',
selectedGroup?.id === group.id
? 'border-[var(--border-strong)] bg-[var(--accent-soft)]'
: 'border-[var(--border-default)] bg-[var(--surface-elevated)] hover:bg-[var(--surface-muted)]',
)}
>
<span className="flex items-center justify-between gap-3">
<span className="flex min-w-0 items-center gap-2">
{group.is_global ? <Globe2 className="h-4 w-4" /> : <ShieldCheck className="h-4 w-4" />}
<span className="truncate text-sm font-semibold text-[var(--foreground-primary)]">
{group.name}
</span>
</span>
<span className="text-xs text-[var(--foreground-secondary)]">
{group.enabled ? '启用' : '停用'}
</span>
</span>
<span className="mt-2 block text-xs text-[var(--foreground-secondary)]">
{group.is_global ? '应用全部网站' : `已应用 ${group.applied_site_count} 个网站`} · {ruleCount(group)} 条规则
</span>
</button>
))}
</div>
</AppCard>
<AppCard
title={selectedGroup ? selectedGroup.name : '新建规则组'}
description="白名单命中后直接放行;未命中白名单时继续判断黑名单。"
action={
selectedGroup && !selectedGroup.is_global ? (
<SecondaryButton type="button" onClick={() => setApplyGroup(selectedGroup)}>
一键应用
</SecondaryButton>
) : null
}
>
<div className="grid gap-5 xl:grid-cols-2">
<ResourceField label="规则组名称">
<ResourceInput
value={draft.name}
disabled={selectedGroup?.is_global}
onChange={(event) => setDraft((current) => ({ ...current, name: event.target.value }))}
/>
</ResourceField>
<ResourceField label="拦截状态码">
<ResourceInput
type="number"
min={400}
max={599}
value={draft.block_status_code}
onChange={(event) =>
setDraft((current) => ({ ...current, block_status_code: Number(event.target.value) }))
}
/>
</ResourceField>
<ToggleField
label="启用规则组"
checked={draft.enabled}
onChange={(checked) => setDraft((current) => ({ ...current, enabled: checked }))}
/>
<ResourceField label="备注">
<ResourceInput
value={draft.remark}
onChange={(event) => setDraft((current) => ({ ...current, remark: event.target.value }))}
/>
</ResourceField>
<ResourceField label="IP / IP 段白名单" hint="每行一个 IP 或 CIDR。">
<ResourceTextarea
value={listToText(draft.ip_whitelist)}
onChange={(event) =>
setDraft((current) => ({ ...current, ip_whitelist: textToList(event.target.value) }))
}
/>
</ResourceField>
<ResourceField label="IP / IP 段黑名单" hint="每行一个 IP 或 CIDR。">
<ResourceTextarea
value={listToText(draft.ip_blacklist)}
onChange={(event) =>
setDraft((current) => ({ ...current, ip_blacklist: textToList(event.target.value) }))
}
/>
</ResourceField>
<ResourceField label="国家白名单" hint="ISO 两位国家代码,例如 CN、US。">
<ResourceTextarea
value={listToText(draft.country_whitelist)}
onChange={(event) =>
setDraft((current) => ({ ...current, country_whitelist: textToList(event.target.value) }))
}
/>
</ResourceField>
<ResourceField label="国家黑名单" hint="ISO 两位国家代码,例如 CN、US。">
<ResourceTextarea
value={listToText(draft.country_blacklist)}
onChange={(event) =>
setDraft((current) => ({ ...current, country_blacklist: textToList(event.target.value) }))
}
/>
</ResourceField>
<ResourceField label="拦截页面" className="xl:col-span-2" hint="留空时只返回状态码。">
<ResourceTextarea
value={draft.block_response_body}
onChange={(event) =>
setDraft((current) => ({ ...current, block_response_body: event.target.value }))
}
/>
</ResourceField>
</div>
<div className="mt-6 flex flex-wrap justify-between gap-3">
<div>
{selectedGroup && !selectedGroup.is_global ? (
<DangerButton
type="button"
disabled={deleteMutation.isPending}
onClick={() => {
if (window.confirm(`确认删除 WAF 规则组 ${selectedGroup.name} 吗?`)) {
deleteMutation.mutate(selectedGroup.id);
}
}}
>
<Trash2 className="mr-2 h-4 w-4" />
删除
</DangerButton>
) : null}
</div>
<PrimaryButton
type="button"
disabled={saveMutation.isPending}
onClick={() => saveMutation.mutate(draft)}
>
<Save className="mr-2 h-4 w-4" />
{saveMutation.isPending ? '保存中...' : '保存规则组'}
</PrimaryButton>
</div>
</AppCard>
</div>
</div>
<SiteApplyDrawer
group={applyGroup}
routes={routes}
open={Boolean(applyGroup)}
pending={applyMutation.isPending}
onOpenChange={(open) => {
if (!open) {
setApplyGroup(null);
}
}}
onSave={(ids) => {
if (applyGroup) {
applyMutation.mutate({ id: applyGroup.id, ids });
}
}}
/>
</>
);
}
@@ -0,0 +1,41 @@
export interface WAFRuleGroup {
id: number;
name: string;
enabled: boolean;
is_global: boolean;
block_status_code: number;
block_response_body: string;
ip_whitelist: string[];
ip_blacklist: string[];
country_whitelist: string[];
country_blacklist: string[];
region_whitelist: string[];
region_blacklist: string[];
remark: string;
applied_site_ids: number[];
applied_site_count: number;
created_at: string;
updated_at: string;
}
export interface WAFRuleGroupPayload {
name: string;
enabled: boolean;
block_status_code: number;
block_response_body: string;
ip_whitelist: string[];
ip_blacklist: string[];
country_whitelist: string[];
country_blacklist: string[];
region_whitelist: string[];
region_blacklist: string[];
remark: string;
}
export interface WAFSiteRuleGroups {
route_id: number;
global_rule_group: WAFRuleGroup | null;
rule_groups: WAFRuleGroup[];
applied_rule_groups: WAFRuleGroup[];
applied_ids: number[];
}
@@ -3,7 +3,7 @@
import Link from 'next/link';
import { useRouter } from 'next/navigation';
import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query';
import { useMemo, useState } from 'react';
import { useEffect, useMemo, useState } from 'react';
import { EmptyState } from '@/components/feedback/empty-state';
import { ErrorState } from '@/components/feedback/error-state';
@@ -21,6 +21,10 @@ import {
deleteTlsCertificate,
getTlsCertificates,
} from '@/features/tls-certificates/api/tls-certificates';
import {
getWAFSiteRuleGroups,
replaceWAFSiteRuleGroups,
} from '@/features/waf/api/waf';
import { CertificateDetailModal } from '@/features/websites/components/certificate-detail-modal';
import { CertificateEditorModal } from '@/features/websites/components/certificate-editor-modal';
import { CertificateImportModal } from '@/features/websites/components/certificate-import-modal';
@@ -38,6 +42,7 @@ import {
DangerButton,
PrimaryButton,
SecondaryButton,
ToggleField,
} from '@/features/shared/components/resource-primitives';
import { formatDateTime } from '@/lib/utils/date';
@@ -54,6 +59,7 @@ export function WebsiteDetailPage({ websiteId }: { websiteId: string }) {
const [isCertificateImportOpen, setIsCertificateImportOpen] = useState(false);
const [isCertificateDetailOpen, setIsCertificateDetailOpen] = useState(false);
const [isCertificateEditorOpen, setIsCertificateEditorOpen] = useState(false);
const [wafSelectedIDs, setWafSelectedIDs] = useState<number[]>([]);
const [convertCertificate, setConvertCertificate] =
useState<TlsCertificateItem | null>(null);
const [preferredCertificateId, setPreferredCertificateId] = useState<
@@ -130,6 +136,32 @@ export function WebsiteDetailPage({ websiteId }: { websiteId: string }) {
const enabledRoutesCount = relatedRoutes.filter(
(route) => route.enabled,
).length;
const wafRouteID = relatedRoutes[0]?.id ?? null;
const wafQuery = useQuery({
queryKey: ['waf', 'site-rule-groups', wafRouteID],
queryFn: () => getWAFSiteRuleGroups(wafRouteID ?? 0),
enabled: Boolean(wafRouteID),
});
const wafMutation = useMutation({
mutationFn: (ids: number[]) => replaceWAFSiteRuleGroups(wafRouteID ?? 0, ids),
onSuccess: async (view) => {
setWafSelectedIDs(view.applied_ids);
setFeedback({ tone: 'success', message: '网站 WAF 规则组已更新。' });
await Promise.all([
queryClient.invalidateQueries({ queryKey: ['waf'] }),
queryClient.invalidateQueries({ queryKey: ['config-versions', 'diff'] }),
]);
},
onError: (error) => {
setFeedback({ tone: 'danger', message: getErrorMessage(error) });
},
});
useEffect(() => {
if (wafQuery.data) {
setWafSelectedIDs(wafQuery.data.applied_ids);
}
}, [wafQuery.data]);
const handleDeleteWebsite = () => {
if (!website) {
@@ -371,6 +403,89 @@ export function WebsiteDetailPage({ websiteId }: { websiteId: string }) {
</AppCard>
</div>
<AppCard
title="WAF"
description="全局规则组始终生效,可为当前网站叠加多个自定义规则组。"
action={
wafRouteID ? (
<PrimaryButton
type="button"
disabled={wafMutation.isPending}
onClick={() => wafMutation.mutate(wafSelectedIDs)}
>
{wafMutation.isPending ? '保存中...' : '保存 WAF'}
</PrimaryButton>
) : null
}
>
{!wafRouteID ? (
<EmptyState
title="暂无可绑定规则"
description="当前网站还没有关联代理规则,创建规则后即可配置 WAF。"
/>
) : wafQuery.isLoading ? (
<LoadingState />
) : wafQuery.isError ? (
<ErrorState
title="WAF 规则组加载失败"
description={getErrorMessage(wafQuery.error)}
/>
) : (
<div className="space-y-4">
{wafQuery.data?.global_rule_group ? (
<div className="rounded-2xl border border-[var(--border-default)] bg-[var(--surface-elevated)] px-4 py-3">
<div className="flex flex-wrap items-center justify-between gap-3">
<div>
<p className="text-sm font-semibold text-[var(--foreground-primary)]">
{wafQuery.data.global_rule_group.name}
</p>
<p className="mt-1 text-xs text-[var(--foreground-secondary)]">
全局规则组默认应用到所有网站,不能在单站关闭。
</p>
</div>
<StatusBadge
label={
wafQuery.data.global_rule_group.enabled
? '全局启用'
: '全局停用'
}
variant={
wafQuery.data.global_rule_group.enabled
? 'success'
: 'warning'
}
/>
</div>
</div>
) : null}
<div className="grid gap-3 md:grid-cols-2">
{(wafQuery.data?.rule_groups ?? []).map((group) => (
<ToggleField
key={group.id}
label={group.name}
description={`已应用 ${group.applied_site_count} 个网站,${group.enabled ? '启用中' : '已停用'}`}
checked={wafSelectedIDs.includes(group.id)}
onChange={(checked) => {
setWafSelectedIDs((current) =>
checked
? [...current, group.id].sort(
(left, right) => left - right,
)
: current.filter((id) => id !== group.id),
);
}}
/>
))}
</div>
{(wafQuery.data?.rule_groups ?? []).length === 0 ? (
<p className="text-sm text-[var(--foreground-secondary)]">
暂无自定义规则组,可在 WAF 页面创建后再绑定。
</p>
) : null}
</div>
)}
</AppCard>
<AppCard title="关联规则">
{relatedRoutes.length === 0 ? (
<EmptyState
@@ -21,6 +21,11 @@ export const dashboardNavigation: NavigationItem[] = [
label: '网站',
icon: 'website',
},
{
href: '/waf',
label: 'WAF',
icon: 'waf',
},
{
href: '/origin',
label: '源站',
+1
View File
@@ -6,6 +6,7 @@ export type NavigationIconKey =
| 'domain'
| 'certificate'
| 'proxy'
| 'waf'
| 'release'
| 'log'
| 'performance'