refactor(backend): rename OpenFlare directory to lowercase openflare

This commit is contained in:
ryan
2026-08-30 17:43:23 +08:00
parent 06d5fedbfc
commit c93ff6674f
543 changed files with 819 additions and 819 deletions
@@ -0,0 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package apply_log manages the application of configuration change logs,
// including validation and retention policy enforcement.
package apply_log
const (
errRetentionDaysOutOfRange = "retention_days 必须在 1 到 3650 之间"
)
@@ -0,0 +1,130 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apply_log
import (
"context"
"errors"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
)
const (
defaultApplyLogPageSize = 20
maxApplyLogPageSize = 200
maxApplyLogRetentionDays = 3650
)
// ListQuery filters apply logs for paginated listing.
type ListQuery struct {
NodeID string `json:"node_id"`
PageNo int `json:"pageNo"`
PageSize int `json:"pageSize"`
}
// ListResult is the paginated apply log list response.
type ListResult struct {
Rows []*model.OpenFlareApplyLog `json:"rows"`
Current int `json:"current"`
Total int `json:"total"`
TotalPage int `json:"totalPage"`
}
// CleanupInput controls apply log cleanup behavior.
type CleanupInput struct {
DeleteAll bool `json:"delete_all"`
RetentionDays int `json:"retention_days"`
}
// CleanupResult reports apply log cleanup outcome.
type CleanupResult struct {
DeleteAll bool `json:"delete_all"`
RetentionDays int `json:"retention_days"`
DeletedCount int64 `json:"deleted_count"`
Cutoff *time.Time `json:"cutoff,omitempty"`
}
// ListPage returns paginated apply logs with optional node_id filter.
func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) {
pageNo := normalizePageNo(input.PageNo)
pageSize := normalizePageSize(input.PageSize)
nodeID := strings.TrimSpace(input.NodeID)
rows, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: nodeID,
PageNo: pageNo,
PageSize: pageSize,
})
if err != nil {
return nil, err
}
total, err := repository.CountOpenFlareApplyLogs(ctx, nodeID)
if err != nil {
return nil, err
}
totalPage := 0
if total > 0 {
totalPage = int((total + int64(pageSize) - 1) / int64(pageSize))
}
return &ListResult{
Rows: rows,
Current: pageNo,
Total: int(total),
TotalPage: totalPage,
}, nil
}
// Cleanup removes old apply logs or deletes all records.
func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) {
if input.DeleteAll {
deleted, err := repository.DeleteAllOpenFlareApplyLogs(ctx)
if err != nil {
return nil, err
}
return &CleanupResult{
DeleteAll: true,
DeletedCount: deleted,
}, nil
}
if input.RetentionDays <= 0 || input.RetentionDays > maxApplyLogRetentionDays {
return nil, errors.New(errRetentionDaysOutOfRange)
}
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour)
deleted, err := repository.DeleteOpenFlareApplyLogsBefore(ctx, cutoff)
if err != nil {
return nil, err
}
return &CleanupResult{
RetentionDays: input.RetentionDays,
DeletedCount: deleted,
Cutoff: &cutoff,
}, nil
}
func normalizePageNo(pageNo int) int {
if pageNo <= 0 {
return 1
}
return pageNo
}
func normalizePageSize(pageSize int) int {
if pageSize <= 0 {
return defaultApplyLogPageSize
}
if pageSize > maxApplyLogPageSize {
return maxApplyLogPageSize
}
return pageSize
}
@@ -0,0 +1,109 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apply_log
import (
"context"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupApplyLogTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
err = sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{})
require.NoError(t, err)
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func TestListPageAndCleanup(t *testing.T) {
cleanup := setupApplyLogTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
logs := []model.OpenFlareApplyLog{
{NodeID: "node-logs", Version: "v1", Result: "success", Message: "1", CreatedAt: now.Add(-10 * 24 * time.Hour)},
{NodeID: "node-logs", Version: "v2", Result: "success", Message: "2", CreatedAt: now.Add(-5 * 24 * time.Hour)},
{NodeID: "node-logs", Version: "v3", Result: "success", Message: "3", CreatedAt: now},
}
for i := range logs {
require.NoError(t, db.DB(ctx).Create(&logs[i]).Error)
}
pageResult, err := ListPage(ctx, ListQuery{
NodeID: "node-logs",
PageNo: 1,
PageSize: 2,
})
require.NoError(t, err)
assert.Equal(t, 3, pageResult.Total)
assert.Len(t, pageResult.Rows, 2)
assert.Equal(t, 2, pageResult.TotalPage)
assert.Equal(t, 1, pageResult.Current)
cleanupResult, err := Cleanup(ctx, CleanupInput{
DeleteAll: false,
RetentionDays: 7,
})
require.NoError(t, err)
assert.Equal(t, int64(1), cleanupResult.DeletedCount)
assert.NotNil(t, cleanupResult.Cutoff)
remaining, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: "node-logs",
PageNo: 1,
PageSize: 10,
})
require.NoError(t, err)
assert.Len(t, remaining, 2)
cleanupAll, err := Cleanup(ctx, CleanupInput{DeleteAll: true})
require.NoError(t, err)
assert.Equal(t, int64(2), cleanupAll.DeletedCount)
assert.True(t, cleanupAll.DeleteAll)
finalLogs, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: "node-logs",
PageNo: 1,
PageSize: 10,
})
require.NoError(t, err)
assert.Empty(t, finalLogs)
}
func TestCleanupInvalidRetentionDays(t *testing.T) {
cleanup := setupApplyLogTestDB(t)
defer cleanup()
ctx := context.Background()
_, err := Cleanup(ctx, CleanupInput{RetentionDays: 0})
require.Error(t, err)
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
_, err = Cleanup(ctx, CleanupInput{RetentionDays: 4000})
require.Error(t, err)
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
}
@@ -0,0 +1,77 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apply_log
import (
"net/http"
"strconv"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// GetApplyLogs lists apply logs with pagination and optional node_id filter.
// @Summary 获取配置下发日志
// @Description 分页返回节点配置下发记录,支持按节点 ID 筛选,需要管理员权限
// @Tags openflare-apply-log
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID 筛选"
// @Param pageNo query int false "页码"
// @Param page_no query int false "页码(别名)"
// @Param pageSize query int false "每页数量"
// @Param page_size query int false "每页数量(别名)"
// @Success 200 {object} response.Any{data=apply_log.ListResult} "下发日志列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/apply-logs [get]
func GetApplyLogs(c *gin.Context) {
result, err := ListPage(c.Request.Context(), ListQuery{
NodeID: c.Query("node_id"),
PageNo: readIntQuery(c, "pageNo", "page_no"),
PageSize: readIntQuery(c, "pageSize", "page_size"),
})
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// CleanupApplyLogs removes old apply logs or deletes all records.
// @Summary 清理配置下发日志
// @Description 按保留天数清理历史下发记录,或删除全部记录,需要管理员权限
// @Tags openflare-apply-log
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body apply_log.CleanupInput true "清理参数"
// @Success 200 {object} response.Any{data=apply_log.CleanupResult} "清理结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/apply-logs/cleanup [post]
func CleanupApplyLogs(c *gin.Context) {
var input CleanupInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := Cleanup(c.Request.Context(), input)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
func readIntQuery(c *gin.Context, primary, secondary string) int {
value := c.Query(primary)
if value == "" {
value = c.Query(secondary)
}
parsed, _ := strconv.Atoi(value)
return parsed
}
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"strings"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
oftls "Wavelet/openflare/plugins/server/domain/tls"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
db "Wavelet/plugins/infra/database"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestBuildCertificateSupportFilesDecryptsSealedPrivateKey(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
require.NoError(t, db.DB(context.Background()).AutoMigrate(&model.TLSCertificate{}))
previous := runtimeconfig.Get()
runtimeconfig.SetSessionSecret("test-session-secret-for-tls-seal")
t.Cleanup(func() { runtimeconfig.Set(previous) })
ctx := context.Background()
certPEM, keyPEM := generateTestCertKeyPairForSnapshot(t)
certificate, err := oftls.CreateCertificate(ctx, oftls.CertificateInput{
Name: "publish-cert",
CertPEM: certPEM,
KeyPEM: keyPEM,
})
require.NoError(t, err)
files, err := buildCertificateSupportFiles(ctx, []snapshotRoute{
{DomainCertIDs: []uint{certificate.ID}},
})
require.NoError(t, err)
require.Len(t, files, 2)
var keyContent string
for _, file := range files {
if file.Path == certificateKeyFileName(certificate.ID) {
keyContent = file.Content
}
assert.NotContains(t, file.Content, "enc:v1:")
}
assert.Contains(t, keyContent, "BEGIN")
assert.Equal(t, normalizePEM(strings.TrimSpace(keyPEM)), keyContent)
}
func TestBuildSnapshotReadsZoneDomainCertificates(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, db.DB(ctx).AutoMigrate(&model.TLSCertificate{}))
previous := runtimeconfig.Get()
runtimeconfig.SetSessionSecret("test-session-secret-for-zone-domain-snapshots")
t.Cleanup(func() { runtimeconfig.Set(previous) })
firstCertPEM, firstKeyPEM := generateTestCertKeyPairForSnapshotForDomain(t, "one.example.com")
first, err := oftls.CreateCertificate(ctx, oftls.CertificateInput{Name: "first", CertPEM: firstCertPEM, KeyPEM: firstKeyPEM})
require.NoError(t, err)
secondCertPEM, secondKeyPEM := generateTestCertKeyPairForSnapshotForDomain(t, "two.example.com")
second, err := oftls.CreateCertificate(ctx, oftls.CertificateInput{Name: "second", CertPEM: secondCertPEM, KeyPEM: secondKeyPEM})
require.NoError(t, err)
route := &model.ProxyRoute{SiteName: "tls-site", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true, EnableHTTPS: true}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
zone := &model.Zone{Domain: "example.com"}
require.NoError(t, db.DB(ctx).Create(zone).Error)
require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: "one.example.com", CertID: &first.ID}).Error)
require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: "two.example.com", CertID: &second.ID}).Error)
bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err)
require.Len(t, bundle.SnapshotRoutes, 1)
assert.Equal(t, []string{"one.example.com", "two.example.com"}, bundle.SnapshotRoutes[0].Domains)
assert.Equal(t, []uint{first.ID, second.ID}, bundle.SnapshotRoutes[0].DomainCertIDs)
assert.Contains(t, bundle.RouteConfig, "server_name one.example.com;")
assert.Contains(t, bundle.RouteConfig, "server_name two.example.com;")
}
func generateTestCertKeyPairForSnapshot(t *testing.T) (certPEM string, keyPEM string) {
t.Helper()
return generateTestCertKeyPairForSnapshotForDomain(t, "test.example.com")
}
func generateTestCertKeyPairForSnapshotForDomain(t *testing.T, domain string) (certPEM string, keyPEM string) {
t.Helper()
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
template := x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: domain},
DNSNames: []string{domain},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(24 * time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
}
der, err := x509.CreateCertificate(rand.Reader, &template, &template, &privateKey.PublicKey, privateKey)
require.NoError(t, err)
certPEM = string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}))
keyPEM = string(pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)}))
return certPEM, keyPEM
}
@@ -0,0 +1,13 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package config_version defines shared error messages for configuration versions.
package config_version
const (
errNoActiveVersion = "当前没有激活版本"
errNoEnabledRoutes = "没有可发布的启用规则"
errNoChangesToPublish = "当前规则没有变更,不能重复发布"
errVersionConflict = "版本号生成冲突,请重试"
errInvalidSnapshotFormat = "历史版本快照格式不合法"
)
@@ -0,0 +1,243 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"errors"
"fmt"
"net"
"strconv"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
)
type customHeaderInput struct {
Key string `json:"key"`
Value string `json:"value"`
}
func normalizeSnapshotDomains(domains []string) ([]string, error) {
normalized := make([]string, 0, len(domains))
seen := make(map[string]struct{}, len(domains))
for _, raw := range domains {
domain := strings.ToLower(strings.TrimSpace(raw))
if domain == "" || strings.Contains(domain, "://") || strings.Contains(domain, "/") {
return nil, errors.New("domains payload is invalid")
}
if _, ok := seen[domain]; ok {
continue
}
seen[domain] = struct{}{}
normalized = append(normalized, domain)
}
if len(normalized) == 0 {
return nil, errors.New("domain is required")
}
return normalized, nil
}
func isUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return normalizeUpstreams(fallbackOriginURL, nil)
}
var upstreams []string
if err := json.Unmarshal([]byte(text), &upstreams); err != nil {
return nil, errors.New("upstreams payload is invalid")
}
return normalizeUpstreams(fallbackOriginURL, upstreams)
}
func normalizeUpstreams(originURL string, upstreams []string) ([]string, error) {
candidates := upstreams
if len(candidates) == 0 {
candidates = []string{originURL}
}
normalized := make([]string, 0, len(candidates))
seen := make(map[string]struct{}, len(candidates))
for _, item := range candidates {
value := strings.TrimSpace(item)
if value == "" {
continue
}
if _, ok := seen[value]; ok {
continue
}
seen[value] = struct{}{}
normalized = append(normalized, value)
}
if len(normalized) == 0 {
return nil, errors.New("upstream is required")
}
return normalized, nil
}
func decodeStoredCustomHeaders(raw string) ([]customHeaderInput, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []customHeaderInput{}, nil
}
var headers []customHeaderInput
if err := json.Unmarshal([]byte(text), &headers); err != nil {
return nil, errors.New("custom_headers payload is invalid")
}
return headers, nil
}
func decodeStoredCacheRules(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []string{}, nil
}
var rules []string
if err := json.Unmarshal([]byte(text), &rules); err != nil {
return nil, errors.New("cache_rules payload is invalid")
}
normalized := make([]string, 0, len(rules))
for _, rule := range rules {
item := strings.TrimSpace(rule)
if item == "" {
continue
}
normalized = append(normalized, item)
}
return normalized, nil
}
func normalizeUpstreamType(raw string) string {
value := strings.ToLower(strings.TrimSpace(raw))
switch value {
case "tunnel", "pages":
return value
default:
return "direct"
}
}
func normalizeTunnelTargetProtocol(raw string) string {
value := strings.ToLower(strings.TrimSpace(raw))
switch value {
case "http", "https", "tcp":
return value
default:
return "http"
}
}
func normalizePEM(content string) string {
return strings.TrimSpace(content) + "\n"
}
func certificateCertFileName(id uint) string {
return fmt.Sprintf("%d.crt", id)
}
func certificateKeyFileName(id uint) string {
return fmt.Sprintf("%d.key", id)
}
func dedupeSupportFiles(files []SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
unique := make(map[string]SupportFile, len(files))
for _, file := range files {
unique[file.Path] = file
}
result := make([]SupportFile, 0, len(unique))
for _, file := range unique {
result = append(result, file)
}
return result
}
func uintPtrEqual(left *uint, right *uint) bool {
if left == nil || right == nil {
return left == nil && right == nil
}
return *left == *right
}
func uintSliceEqual(left []uint, right []uint) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] != right[index] {
return false
}
}
return true
}
func relayAgentAddress(node *model.OpenFlareNode) string {
if node == nil {
return ""
}
port := node.RelayVhostHTTPPort
if port <= 0 {
port = 8080
}
addr := strings.TrimSpace(node.RelayAgentAccessAddr)
if addr == "" {
addr = strings.TrimSpace(node.RelayClientAccessAddr)
}
if addr == "" {
addr = strings.TrimSpace(node.IP)
}
if addr == "" {
return fmt.Sprintf("127.0.0.1:%d", port)
}
if _, _, err := net.SplitHostPort(addr); err == nil {
return addr
}
if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 {
return net.JoinHostPort(addr, strconv.Itoa(port))
}
return fmt.Sprintf("%s:%d", addr, port)
}
func resolveTunnelOpenRestyUpstreamURL(ctx context.Context) string {
nodes, err := repository.ListOpenFlareNodes(ctx)
if err == nil {
for index := range nodes {
node := &nodes[index]
if node.NodeType != "tunnel_relay" {
continue
}
addr := relayAgentAddress(node)
if addr != "" {
return "http://" + addr
}
}
}
return "http://127.0.0.1:8080"
}
func listWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*model.OpenFlareWAFIPGroup, error) {
if len(ids) == 0 {
return []*model.OpenFlareWAFIPGroup{}, nil
}
groups := make([]*model.OpenFlareWAFIPGroup, 0, len(ids))
for _, id := range ids {
group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if err != nil {
return nil, err
}
groups = append(groups, group)
}
return groups, nil
}
@@ -0,0 +1,627 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"errors"
"fmt"
"slices"
"sort"
"strconv"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/fleet/websocket"
"Wavelet/openflare/plugins/server/kernel/model"
pkgprotocol "Wavelet/openflare/share/protocol"
openrestyrender "Wavelet/openflare/share/render/openresty"
"gorm.io/gorm"
)
const (
cleanupSuccessMessage = "清理成功"
minConfigVersionKeepCount = 3
)
// ConfigPreviewResult is the preview response payload.
type ConfigPreviewResult struct {
SnapshotJSON string `json:"snapshot_json"`
MainConfig string `json:"main_config"`
RouteConfig string `json:"route_config"`
RenderedConfig string `json:"rendered_config"`
SupportFiles []SupportFile `json:"support_files"`
Checksum string `json:"checksum"`
RouteCount int `json:"route_count"`
WebsiteCount int `json:"website_count"`
}
// ConfigDiffResult is the diff response payload.
type ConfigDiffResult struct {
ActiveVersion string `json:"active_version,omitempty"`
AddedSites []string `json:"added_sites"`
RemovedSites []string `json:"removed_sites"`
ModifiedSites []string `json:"modified_sites"`
AddedDomains []string `json:"added_domains"`
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"`
ActiveWebsiteCount int `json:"active_website_count"`
}
// ConfigOptionDiffItem describes a changed OpenResty option.
type ConfigOptionDiffItem struct {
Key string `json:"key"`
PreviousValue string `json:"previous_value"`
CurrentValue string `json:"current_value"`
}
// CleanupInput is the cleanup request payload.
type CleanupInput struct {
KeepCount int `json:"keep_count"`
}
// CleanupResult is the cleanup response payload.
type CleanupResult struct {
DeletedCount int64 `json:"deleted_count"`
Message string `json:"message"`
}
// ListConfigVersions returns all config version summaries.
func ListConfigVersions(ctx context.Context) ([]*model.ConfigVersionSummary, error) {
return repository.ListConfigVersionSummaries(ctx)
}
// GetConfigVersionDetail returns a config version by version.
func GetConfigVersionDetail(ctx context.Context, version string) (*model.ConfigVersion, error) {
return repository.GetConfigVersionByVersion(ctx, version)
}
// GetActiveConfigVersion returns the active config version.
func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) {
return repository.GetActiveConfigVersion(ctx)
}
// PreviewConfigVersion renders the current draft configuration.
func PreviewConfigVersion(ctx context.Context) (*ConfigPreviewResult, error) {
bundle, err := buildCurrentConfigBundle(ctx, false)
if err != nil {
return nil, err
}
return &ConfigPreviewResult{
SnapshotJSON: bundle.SnapshotJSON,
MainConfig: bundle.MainConfig,
RouteConfig: bundle.RouteConfig,
RenderedConfig: bundle.RouteConfig,
SupportFiles: bundle.SupportFiles,
Checksum: bundle.Checksum,
RouteCount: len(bundle.Routes),
WebsiteCount: len(bundle.SnapshotRoutes),
}, nil
}
// DiffConfigVersion compares the current draft against the active version.
func DiffConfigVersion(ctx context.Context) (*ConfigDiffResult, error) {
bundle, err := buildCurrentConfigBundle(ctx, false)
if err != nil {
return nil, err
}
result := &ConfigDiffResult{
AddedSites: []string{},
RemovedSites: []string{},
ModifiedSites: []string{},
AddedDomains: []string{},
RemovedDomains: []string{},
ModifiedDomains: []string{},
ChangedOptionKeys: []string{},
ChangedOptionDetails: []ConfigOptionDiffItem{},
CurrentWebsiteCount: len(bundle.SnapshotRoutes),
}
activeVersion, err := repository.GetActiveConfigVersion(ctx)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
for _, route := range bundle.SnapshotRoutes {
result.AddedSites = append(result.AddedSites, route.SiteName)
result.AddedDomains = append(result.AddedDomains, route.Domains...)
}
result.MainConfigChanged = true
result.ChangedOptionKeys = openRestyOptionKeys()
result.ChangedOptionDetails = buildInitialOpenRestyOptionDiffs(bundle.OpenRestyConfig)
sort.Strings(result.AddedSites)
sort.Strings(result.AddedDomains)
sort.Strings(result.ChangedOptionKeys)
return result, nil
}
return nil, err
}
result.ActiveVersion = activeVersion.Version
activeSnapshot, err := parseSnapshotDocument(activeVersion.SnapshotJSON)
if err != nil {
return nil, err
}
result.ActiveWebsiteCount = len(activeSnapshot.Routes)
currentSiteMap := flattenSnapshotRoutesBySite(bundle.SnapshotRoutes)
activeSiteMap := flattenSnapshotRoutesBySite(activeSnapshot.Routes)
for siteName, currentRoute := range currentSiteMap {
activeRoute, ok := activeSiteMap[siteName]
if !ok {
result.AddedSites = append(result.AddedSites, siteName)
continue
}
if !snapshotRouteConfigEqual(activeRoute, currentRoute) {
result.ModifiedSites = append(result.ModifiedSites, siteName)
}
}
for siteName := range activeSiteMap {
if _, ok := currentSiteMap[siteName]; !ok {
result.RemovedSites = append(result.RemovedSites, siteName)
}
}
currentMap := flattenSnapshotRoutesByDomain(bundle.SnapshotRoutes)
activeMap := flattenSnapshotRoutesByDomain(activeSnapshot.Routes)
for domain, currentRoute := range currentMap {
activeRoute, ok := activeMap[domain]
if !ok {
result.AddedDomains = append(result.AddedDomains, domain)
continue
}
if !snapshotRouteConfigEqual(activeRoute, currentRoute) {
result.ModifiedDomains = append(result.ModifiedDomains, domain)
}
}
for domain := range activeMap {
if _, ok := currentMap[domain]; !ok {
result.RemovedDomains = append(result.RemovedDomains, domain)
}
}
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)
sort.Strings(result.RemovedSites)
sort.Strings(result.ModifiedSites)
sort.Strings(result.AddedDomains)
sort.Strings(result.RemovedDomains)
sort.Strings(result.ModifiedDomains)
sort.Strings(result.ChangedOptionKeys)
return result, nil
}
// PublishConfigVersion publishes the current draft as a new active version.
func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*model.ConfigVersion, error) {
bundle, err := buildCurrentConfigBundle(ctx, true)
if err != nil {
return nil, err
}
if len(bundle.Routes) == 0 {
return nil, errors.New(errNoEnabledRoutes)
}
activeVersion, err := repository.GetActiveConfigVersion(ctx)
if !force && err == nil && activeVersion.Checksum == bundle.Checksum {
return nil, errors.New(errNoChangesToPublish)
}
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
supportFilesJSON, err := json.Marshal(bundle.SupportFiles)
if err != nil {
return nil, err
}
version, err := nextVersionNumber(ctx, time.Now())
if err != nil {
return nil, err
}
record := &model.ConfigVersion{
Version: version,
SnapshotJSON: bundle.SnapshotJSON,
MainConfig: bundle.MainConfig,
RenderedConfig: bundle.RouteConfig,
SupportFilesJSON: string(supportFilesJSON),
Checksum: bundle.Checksum,
IsActive: true,
CreatedBy: createdBy,
}
if err = repository.PublishConfigVersionTx(ctx, record); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errVersionConflict)
}
return nil, err
}
websocket.BroadcastActiveConfig(pkgprotocol.ActiveConfigMeta{
Version: record.Version,
Checksum: record.Checksum,
})
return record, nil
}
// ActivateConfigVersion activates an existing config version.
func ActivateConfigVersion(ctx context.Context, versionStr string) (*model.ConfigVersion, error) {
version, err := repository.GetConfigVersionByVersion(ctx, versionStr)
if err != nil {
return nil, err
}
if err = repository.ActivateConfigVersionTx(ctx, versionStr); err != nil {
return nil, err
}
version.IsActive = true
websocket.BroadcastActiveConfig(pkgprotocol.ActiveConfigMeta{
Version: version.Version,
Checksum: version.Checksum,
})
return version, nil
}
// CleanupConfigVersions removes old inactive config versions.
func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, error) {
if keepCount < minConfigVersionKeepCount {
keepCount = minConfigVersionKeepCount
}
versions, err := repository.ListConfigVersionSummaries(ctx)
if err != nil {
return nil, err
}
if len(versions) <= keepCount {
return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil
}
var deleteVersions []string
for index, version := range versions {
if index < keepCount {
continue
}
if version.IsActive {
continue
}
deleteVersions = append(deleteVersions, version.Version)
}
if len(deleteVersions) == 0 {
return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil
}
deletedCount, err := repository.DeleteConfigVersionsByVersions(ctx, deleteVersions)
if err != nil {
return nil, err
}
return &CleanupResult{DeletedCount: deletedCount, Message: cleanupSuccessMessage}, nil
}
func nextVersionNumber(ctx context.Context, now time.Time) (string, error) {
prefix := now.Format("20060102")
latest, err := repository.GetLatestConfigVersionByPrefix(ctx, prefix)
if errors.Is(err, gorm.ErrRecordNotFound) {
return fmt.Sprintf("%s-%03d", prefix, 1), nil
}
if err != nil {
return "", err
}
suffix := strings.TrimPrefix(latest, prefix+"-")
sequence, err := strconv.Atoi(suffix)
if err != nil {
return "", fmt.Errorf("invalid config version sequence %q: %w", latest, err)
}
return fmt.Sprintf("%s-%03d", prefix, sequence+1), nil
}
func parseSnapshotDocument(snapshotJSON string) (*snapshotDocument, error) {
text := strings.TrimSpace(snapshotJSON)
if text == "" {
return &snapshotDocument{Routes: []snapshotRoute{}}, nil
}
if strings.HasPrefix(text, "[") {
var routes []snapshotRoute
if err := json.Unmarshal([]byte(text), &routes); err != nil {
return nil, errors.New(errInvalidSnapshotFormat)
}
return &snapshotDocument{Routes: normalizeSnapshotRoutes(routes)}, nil
}
var snapshot snapshotDocument
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
return nil, errors.New(errInvalidSnapshotFormat)
}
snapshot.Routes = normalizeSnapshotRoutes(snapshot.Routes)
return &snapshot, nil
}
func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
if len(routes) == 0 {
return []snapshotRoute{}
}
for index := range routes {
normalizedDomains, err := normalizeSnapshotDomains(routes[index].Domains)
if err == nil && len(normalizedDomains) > 0 {
routes[index].Domains = normalizedDomains
routes[index].SiteName = strings.TrimSpace(routes[index].SiteName)
}
normalizedUpstreams, upstreamErr := normalizeUpstreams(routes[index].OriginURL, routes[index].Upstreams)
if upstreamErr == nil {
routes[index].OriginURL = normalizedUpstreams[0]
routes[index].Upstreams = normalizedUpstreams
}
if !routes[index].BasicAuthEnabled {
routes[index].BasicAuthUsername = ""
routes[index].BasicAuthPassword = ""
}
routes[index].UpstreamType = normalizeUpstreamType(routes[index].UpstreamType)
routes[index].CachePolicy = normalizeSnapshotCachePolicy(
routes[index].CacheEnabled,
routes[index].CachePolicy,
)
if !routes[index].CacheEnabled {
routes[index].CacheRules = nil
}
}
return routes
}
// normalizeSnapshotCachePolicy aligns published policy with edge-cache-design:
// legacy empty/url → all; disabled → empty; static/suffix/... kept.
func normalizeSnapshotCachePolicy(enabled bool, raw string) string {
if !enabled {
return ""
}
policy := strings.TrimSpace(strings.ToLower(raw))
switch policy {
case "", "url", "all":
return "all"
case "static", "suffix", "path_prefix", "path_exact":
return policy
default:
// Unknown: prefer static over caching everything.
return "static"
}
}
func flattenSnapshotRoutesBySite(routes []snapshotRoute) map[string]snapshotRoute {
siteMap := make(map[string]snapshotRoute)
for _, route := range normalizeSnapshotRoutes(routes) {
siteMap[route.SiteName] = route
}
return siteMap
}
func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRoute {
domainMap := make(map[string]snapshotRoute)
for _, route := range normalizeSnapshotRoutes(routes) {
for _, domain := range route.Domains {
item := route
domainMap[domain] = item
}
}
return domainMap
}
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
return snapshotRouteScalarsEqual(left, right) &&
slices.Equal(left.Domains, right.Domains) &&
slices.Equal(left.Upstreams, right.Upstreams) &&
slices.Equal(left.CacheRules, right.CacheRules) &&
slices.Equal(left.CustomHeaders, right.CustomHeaders)
}
func snapshotRouteScalarsEqual(left, right snapshotRoute) bool {
return snapshotRouteIdentityEqual(left, right) &&
snapshotRouteOriginEqual(left, right) &&
snapshotRoutePolicyEqual(left, right) &&
snapshotRouteTunnelEqual(left, right) &&
uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs)
}
func snapshotRouteIdentityEqual(left, right snapshotRoute) bool {
return left.SiteName == right.SiteName
}
func snapshotRouteOriginEqual(left, right snapshotRoute) bool {
return left.OriginURL == right.OriginURL &&
left.OriginHost == right.OriginHost &&
left.UpstreamType == right.UpstreamType &&
snapshotPagesDeploymentEqual(left.PagesDeployment, right.PagesDeployment)
}
func snapshotPagesDeploymentEqual(left, right *openrestyrender.PagesDeployment) bool {
if left == nil && right == nil {
return true
}
if left == nil || right == nil {
return false
}
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 snapshotRoutePolicyEqual(left, right snapshotRoute) bool {
return left.EnableHTTPS == right.EnableHTTPS &&
left.RedirectHTTP == right.RedirectHTTP &&
left.LimitConnPerServer == right.LimitConnPerServer &&
left.LimitConnPerIP == right.LimitConnPerIP &&
left.LimitRate == right.LimitRate &&
left.CacheEnabled == right.CacheEnabled &&
left.CachePolicy == right.CachePolicy &&
left.BasicAuthEnabled == right.BasicAuthEnabled &&
left.BasicAuthUsername == right.BasicAuthUsername &&
left.BasicAuthPassword == right.BasicAuthPassword
}
func snapshotRouteTunnelEqual(left, right snapshotRoute) bool {
return left.TunnelTargetAddr == right.TunnelTargetAddr &&
left.TunnelTargetProto == right.TunnelTargetProto &&
uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) &&
uintPtrEqual(left.PagesProjectID, right.PagesProjectID)
}
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 buildInitialOpenRestyOptionDiffs(current openRestyConfigSnapshot) []ConfigOptionDiffItem {
details := diffOpenRestyOptionDetails(openRestyConfigSnapshot{}, current)
for index := range details {
details[index].PreviousValue = ""
}
return details
}
func diffOpenRestyOptionDetails(left openRestyConfigSnapshot, right openRestyConfigSnapshot) []ConfigOptionDiffItem {
changes := make([]ConfigOptionDiffItem, 0)
appendIfChanged := func(key string, previous string, current string) {
if previous == current {
return
}
changes = append(changes, ConfigOptionDiffItem{
Key: key,
PreviousValue: previous,
CurrentValue: current,
})
}
appendIfChanged("OpenRestyDefaultServerReturnStatus", strconv.Itoa(left.DefaultServerReturnStatus), strconv.Itoa(right.DefaultServerReturnStatus))
appendIfChanged("OpenRestyWorkerProcesses", left.WorkerProcesses, right.WorkerProcesses)
appendIfChanged("OpenRestyWorkerConnections", strconv.Itoa(left.WorkerConnections), strconv.Itoa(right.WorkerConnections))
appendIfChanged("OpenRestyWorkerRlimitNofile", strconv.Itoa(left.WorkerRlimitNofile), strconv.Itoa(right.WorkerRlimitNofile))
appendIfChanged("OpenRestyEventsUse", left.EventsUse, right.EventsUse)
appendIfChanged("OpenRestyEventsMultiAcceptEnabled", strconv.FormatBool(left.EventsMultiAcceptEnabled), strconv.FormatBool(right.EventsMultiAcceptEnabled))
appendIfChanged("OpenRestyKeepaliveTimeout", strconv.Itoa(left.KeepaliveTimeout), strconv.Itoa(right.KeepaliveTimeout))
appendIfChanged("OpenRestyKeepaliveRequests", strconv.Itoa(left.KeepaliveRequests), strconv.Itoa(right.KeepaliveRequests))
appendIfChanged("OpenRestyClientHeaderTimeout", strconv.Itoa(left.ClientHeaderTimeout), strconv.Itoa(right.ClientHeaderTimeout))
appendIfChanged("OpenRestyClientBodyTimeout", strconv.Itoa(left.ClientBodyTimeout), strconv.Itoa(right.ClientBodyTimeout))
appendIfChanged("OpenRestyClientMaxBodySize", left.ClientMaxBodySize, right.ClientMaxBodySize)
appendIfChanged("OpenRestyLargeClientHeaderBuffers", left.LargeClientHeaderBuffers, right.LargeClientHeaderBuffers)
appendIfChanged("OpenRestySendTimeout", strconv.Itoa(left.SendTimeout), strconv.Itoa(right.SendTimeout))
appendIfChanged("OpenRestyProxyConnectTimeout", strconv.Itoa(left.ProxyConnectTimeout), strconv.Itoa(right.ProxyConnectTimeout))
appendIfChanged("OpenRestyProxySendTimeout", strconv.Itoa(left.ProxySendTimeout), strconv.Itoa(right.ProxySendTimeout))
appendIfChanged("OpenRestyProxyReadTimeout", strconv.Itoa(left.ProxyReadTimeout), strconv.Itoa(right.ProxyReadTimeout))
appendIfChanged("OpenRestyWebsocketEnabled", strconv.FormatBool(left.WebsocketEnabled), strconv.FormatBool(right.WebsocketEnabled))
appendIfChanged("OpenRestyHTTP3Enabled", strconv.FormatBool(left.HTTP3Enabled), strconv.FormatBool(right.HTTP3Enabled))
appendIfChanged("OpenRestyProxyRequestBufferingEnabled", strconv.FormatBool(left.ProxyRequestBuffering), strconv.FormatBool(right.ProxyRequestBuffering))
appendIfChanged("OpenRestyProxyBufferingEnabled", strconv.FormatBool(left.ProxyBufferingEnabled), strconv.FormatBool(right.ProxyBufferingEnabled))
appendIfChanged("OpenRestyProxyBuffers", left.ProxyBuffers, right.ProxyBuffers)
appendIfChanged("OpenRestyProxyBufferSize", left.ProxyBufferSize, right.ProxyBufferSize)
appendIfChanged("OpenRestyProxyBusyBuffersSize", left.ProxyBusyBuffersSize, right.ProxyBusyBuffersSize)
appendIfChanged("OpenRestyGzipEnabled", strconv.FormatBool(left.GzipEnabled), strconv.FormatBool(right.GzipEnabled))
appendIfChanged("OpenRestyGzipMinLength", strconv.Itoa(left.GzipMinLength), strconv.Itoa(right.GzipMinLength))
appendIfChanged("OpenRestyGzipCompLevel", strconv.Itoa(left.GzipCompLevel), strconv.Itoa(right.GzipCompLevel))
appendIfChanged("OpenRestyResolvers", left.Resolvers, right.Resolvers)
appendIfChanged("OpenRestyCacheEnabled", strconv.FormatBool(left.CacheEnabled), strconv.FormatBool(right.CacheEnabled))
appendIfChanged("OpenRestyCachePath", left.CachePath, right.CachePath)
appendIfChanged("OpenRestyCacheLevels", left.CacheLevels, right.CacheLevels)
appendIfChanged("OpenRestyCacheInactive", left.CacheInactive, right.CacheInactive)
appendIfChanged("OpenRestyCacheMaxSize", left.CacheMaxSize, right.CacheMaxSize)
appendIfChanged("OpenRestyCacheKeyTemplate", left.CacheKeyTemplate, right.CacheKeyTemplate)
appendIfChanged("OpenRestyCacheLockEnabled", strconv.FormatBool(left.CacheLockEnabled), strconv.FormatBool(right.CacheLockEnabled))
appendIfChanged("OpenRestyCacheLockTimeout", left.CacheLockTimeout, right.CacheLockTimeout)
appendIfChanged("OpenRestyCacheUseStale", left.CacheUseStale, right.CacheUseStale)
appendIfChanged("OpenRestyDefaultLimitConnPerServer", strconv.Itoa(left.DefaultLimitConnPerServer), strconv.Itoa(right.DefaultLimitConnPerServer))
appendIfChanged("OpenRestyDefaultLimitConnPerIP", strconv.Itoa(left.DefaultLimitConnPerIP), strconv.Itoa(right.DefaultLimitConnPerIP))
appendIfChanged("OpenRestyDefaultLimitRate", left.DefaultLimitRate, right.DefaultLimitRate)
appendIfChanged("OpenRestyDefaultLimitReqPerIP", left.DefaultLimitReqPerIP, right.DefaultLimitReqPerIP)
appendIfChanged("OriginErrorPageEnabled", strconv.FormatBool(left.OriginErrorPageEnabled), strconv.FormatBool(right.OriginErrorPageEnabled))
appendIfChanged("OriginErrorPageStatusCodes", encodeOriginErrorPageStatusCodes(left.OriginErrorPageStatusCodes), encodeOriginErrorPageStatusCodes(right.OriginErrorPageStatusCodes))
appendIfChanged("OriginErrorPageHTML", left.OriginErrorPageHTML, right.OriginErrorPageHTML)
appendIfChanged("OriginErrorPageGetOnly", strconv.FormatBool(left.OriginErrorPageGetOnly), strconv.FormatBool(right.OriginErrorPageGetOnly))
appendIfChanged("SWOfflineEnabled", strconv.FormatBool(left.SWOfflineEnabled), strconv.FormatBool(right.SWOfflineEnabled))
appendIfChanged("SWOfflineHTML", left.SWOfflineHTML, right.SWOfflineHTML)
appendIfChanged("SWOfflineDomains", encodeSWOfflineDomains(left.SWOfflineDomains), encodeSWOfflineDomains(right.SWOfflineDomains))
return changes
}
func encodeOriginErrorPageStatusCodes(tags []string) string {
if len(tags) == 0 {
return ""
}
payload, err := json.Marshal(tags)
if err != nil {
return strings.Join(tags, ",")
}
return string(payload)
}
func encodeSWOfflineDomains(domains []string) string {
if len(domains) == 0 {
return ""
}
payload, err := json.Marshal(domains)
if err != nil {
return strings.Join(domains, ",")
}
return string(payload)
}
func extractOptionDiffKeys(details []ConfigOptionDiffItem) []string {
keys := make([]string, 0, len(details))
for _, item := range details {
keys = append(keys, item.Key)
}
return keys
}
func openRestyOptionKeys() []string {
return []string{
"OpenRestyDefaultServerReturnStatus",
"OpenRestyWorkerProcesses",
"OpenRestyWorkerConnections",
"OpenRestyWorkerRlimitNofile",
"OpenRestyEventsUse",
"OpenRestyEventsMultiAcceptEnabled",
"OpenRestyKeepaliveTimeout",
"OpenRestyKeepaliveRequests",
"OpenRestyClientHeaderTimeout",
"OpenRestyClientBodyTimeout",
"OpenRestyClientMaxBodySize",
"OpenRestyLargeClientHeaderBuffers",
"OpenRestySendTimeout",
"OpenRestyProxyConnectTimeout",
"OpenRestyProxySendTimeout",
"OpenRestyProxyReadTimeout",
"OpenRestyWebsocketEnabled",
"OpenRestyHTTP3Enabled",
"OpenRestyProxyRequestBufferingEnabled",
"OpenRestyProxyBufferingEnabled",
"OpenRestyProxyBuffers",
"OpenRestyProxyBufferSize",
"OpenRestyProxyBusyBuffersSize",
"OpenRestyGzipEnabled",
"OpenRestyGzipMinLength",
"OpenRestyGzipCompLevel",
"OpenRestyCacheEnabled",
"OpenRestyCachePath",
"OpenRestyCacheLevels",
"OpenRestyCacheInactive",
"OpenRestyCacheMaxSize",
"OpenRestyCacheKeyTemplate",
"OpenRestyCacheLockEnabled",
"OpenRestyCacheLockTimeout",
"OpenRestyCacheUseStale",
"OpenRestyDefaultLimitConnPerServer",
"OpenRestyDefaultLimitConnPerIP",
"OpenRestyDefaultLimitRate",
"OpenRestyDefaultLimitReqPerIP",
"OriginErrorPageEnabled",
"OriginErrorPageStatusCodes",
"OriginErrorPageHTML",
"OriginErrorPageGetOnly",
"SWOfflineEnabled",
"SWOfflineHTML",
"SWOfflineDomains",
}
}
@@ -0,0 +1,248 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"fmt"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/waf"
"Wavelet/openflare/plugins/server/kernel/model"
openrestyrender "Wavelet/openflare/share/render/openresty"
"Wavelet/pkg/cache/ram"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupConfigVersionTestDB(t *testing.T) func() {
t.Helper()
// 同 origin_error_page_snapshot_test.go:换 DB 前后重置进程级 RAM 配置缓存。
ram.ResetForTest()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.ProxyRoute{},
&model.Zone{},
&model.ZoneDomain{},
&model.ConfigVersion{},
&model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFRuleGroupBinding{},
&model.OpenFlareWAFIPGroup{},
&model.SystemConfig{},
))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
ram.ResetForTest()
}
}
func createSnapshotZoneDomains(t *testing.T, ctx context.Context, route *model.ProxyRoute, domains ...string) {
t.Helper()
zone := &model.Zone{Domain: fmt.Sprintf("zone-%d.example", route.ID)}
require.NoError(t, db.DB(ctx).Create(zone).Error)
for _, domain := range domains {
require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{
ZoneID: zone.ID,
ProxyRouteID: &route.ID,
Domain: domain,
}).Error)
}
}
func TestListConfigVersionsOrdersByCreatedAtDesc(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
conn := db.DB(ctx)
require.NotNil(t, conn)
newer := &model.ConfigVersion{
Version: "20260102-001",
SnapshotJSON: "{}",
RenderedConfig: "route {}",
Checksum: "checksum-newer",
CreatedBy: "tester",
CreatedAt: time.Date(2026, 1, 2, 12, 0, 0, 0, time.UTC),
}
older := &model.ConfigVersion{
Version: "20260101-001",
SnapshotJSON: "{}",
RenderedConfig: "route {}",
Checksum: "checksum-older",
CreatedBy: "tester",
CreatedAt: time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC),
}
require.NoError(t, conn.Create(newer).Error)
require.NoError(t, conn.Create(older).Error)
versions, err := ListConfigVersions(ctx)
require.NoError(t, err)
require.Len(t, versions, 2)
assert.Equal(t, newer.Version, versions[0].Version)
assert.Equal(t, older.Version, versions[1].Version)
}
func TestPublishConfigVersionCreatesVersion(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
route := &model.ProxyRoute{
SiteName: "publish-site",
OriginURL: "http://origin.publish.example.com:8080",
Upstreams: `["http://origin.publish.example.com:8080"]`,
Enabled: true,
}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, "publish.example.com")
version, err := PublishConfigVersion(ctx, "tester", false)
require.NoError(t, err)
require.NotNil(t, version)
assert.NotEmpty(t, version.ID)
assert.True(t, version.IsActive)
assert.Equal(t, "tester", version.CreatedBy)
assert.NotEmpty(t, version.Version)
assert.NotEmpty(t, version.Checksum)
assert.NotEmpty(t, version.SnapshotJSON)
assert.NotEmpty(t, version.RenderedConfig)
var snapshot snapshotDocument
require.NoError(t, json.Unmarshal([]byte(version.SnapshotJSON), &snapshot))
require.Len(t, snapshot.Routes, 1)
assert.Equal(t, "publish-site", snapshot.Routes[0].SiteName)
assert.Equal(t, []string{"publish.example.com"}, snapshot.Routes[0].Domains)
active, err := GetActiveConfigVersion(ctx)
require.NoError(t, err)
assert.Equal(t, version.ID, active.ID)
_, err = PublishConfigVersion(ctx, "tester", false)
require.Error(t, err)
assert.Contains(t, err.Error(), errNoChangesToPublish)
forced, err := PublishConfigVersion(ctx, "tester", true)
require.NoError(t, err)
assert.NotEqual(t, version.ID, forced.ID)
}
func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
route := &model.ProxyRoute{
SiteName: "example.com",
OriginURL: "http://origin.example.com:8080",
Upstreams: `["http://origin.example.com:8080"]`,
Enabled: true,
}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, "example.com", "www.example.com")
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
globalGroup, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
customGroup := createSnapshotRule(t, ctx, "pow-group", waf.DefaultRuleGraph())
require.NoError(t, repository.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID}))
bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err)
require.Len(t, bundle.SnapshotRoutes, 1)
assert.Equal(t, "example.com", bundle.SnapshotRoutes[0].SiteName)
require.NotEmpty(t, bundle.WAFSnapshot.Bindings)
found := false
for _, binding := range bundle.WAFSnapshot.Bindings {
if binding.RouteID != route.ID {
continue
}
found = true
assert.Equal(t, "example.com", binding.SiteName)
assert.Contains(t, binding.RuleGroupIDs, customGroup.ID)
}
assert.True(t, found, "expected WAF binding for enabled route")
var wafRuntime openrestyrender.WAFDocument
foundWAFConfig := false
for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" {
continue
}
foundWAFConfig = true
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
}
require.True(t, foundWAFConfig, "expected rendered WAF support file")
require.NotEmpty(t, wafRuntime.RuleGroups)
assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
require.Len(t, wafRuntime.Bindings, 1)
assert.Equal(t, route.ID, wafRuntime.Bindings[0].RouteID)
assert.Equal(t, "example.com", wafRuntime.Bindings[0].SiteName)
assert.Equal(t, []uint{customGroup.ID}, wafRuntime.Bindings[0].RuleGroupIDs)
assert.Contains(t, bundle.RouteConfig, `set $openflare_waf_site "example.com"`)
}
func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
route := &model.ProxyRoute{
SiteName: "pow-global.example.com",
OriginURL: "http://origin.example.com:8080",
Upstreams: `["http://origin.example.com:8080"]`,
Enabled: true,
}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, "pow-global.example.com")
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
globalGroup, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
graphJSON, err := json.Marshal(snapshotPoWGraph())
require.NoError(t, err)
globalGroup.Graph = string(graphJSON)
require.NoError(t, db.DB(ctx).Model(globalGroup).Update("graph", globalGroup.Graph).Error)
bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err)
var wafRuntime openrestyrender.WAFDocument
foundWAFConfig := false
for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" {
continue
}
foundWAFConfig = true
assert.Contains(t, file.Content, `"rule_group_ids":[]`)
assert.NotContains(t, file.Content, `"rule_group_ids":null`)
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
}
require.True(t, foundWAFConfig, "expected rendered WAF support file")
require.NotEmpty(t, wafRuntime.RuleGroups)
assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
assert.Equal(t, string(waf.RuleNodePoW), wafRuntime.RuleGroups[0].Graph.Nodes["pow"].Type)
require.Len(t, wafRuntime.Bindings, 1)
assert.Equal(t, "pow-global.example.com", wafRuntime.Bindings[0].SiteName)
assert.Empty(t, wafRuntime.Bindings[0].RuleGroupIDs)
require.NotEmpty(t, bundle.WAFSnapshot.RuleGroups)
assert.Equal(t, waf.RuleNodePoW, bundle.WAFSnapshot.RuleGroups[0].Graph.Nodes["pow"].Type)
}
@@ -0,0 +1,128 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/pkg/cache/ram"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupOriginErrorPageSnapshotDB(t *testing.T) func() {
t.Helper()
// repository 读配置会写进程级 RAM 缓存(跨测试存活),换 DB 前后必须
// 重置,否则 shuffle 下先跑的用例会污染后跑的用例。
ram.ResetForTest()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
ram.ResetForTest()
}
}
func TestBuildOpenRestyConfigSnapshotOriginErrorPageDefaults(t *testing.T) {
cleanup := setupOriginErrorPageSnapshotDB(t)
defer cleanup()
snapshot := buildOpenRestyConfigSnapshot(context.Background())
assert.True(t, snapshot.OriginErrorPageEnabled)
assert.Equal(t, []string{"500-599"}, snapshot.OriginErrorPageStatusCodes)
assert.Empty(t, snapshot.OriginErrorPageHTML)
payload, err := json.Marshal(snapshot)
require.NoError(t, err)
assert.Contains(t, string(payload), `"origin_error_page_enabled":true`)
assert.Contains(t, string(payload), `"origin_error_page_status_codes":["500-599"]`)
}
func TestBuildOpenRestyConfigSnapshotOriginErrorPageCustom(t *testing.T) {
cleanup := setupOriginErrorPageSnapshotDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{
Key: model.ConfigKeyOriginErrorPageEnabled, Value: "false", Type: "business",
}).Error)
require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{
Key: model.ConfigKeyOriginErrorPageStatusCodes, Value: `["522","500-502"]`, Type: "business",
}).Error)
require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{
Key: model.ConfigKeyOriginErrorPageHTML, Value: "<h1>{{status}}</h1>", Type: "business",
}).Error)
snapshot := buildOpenRestyConfigSnapshot(ctx)
assert.False(t, snapshot.OriginErrorPageEnabled)
assert.Equal(t, []string{"522", "500-502"}, snapshot.OriginErrorPageStatusCodes)
assert.Equal(t, "<h1>{{status}}</h1>", snapshot.OriginErrorPageHTML)
}
func TestParseOriginErrorPageStatusCodesFallback(t *testing.T) {
t.Parallel()
assert.Equal(t, []string{"500-599"}, parseOriginErrorPageStatusCodes(""))
assert.Equal(t, []string{"500-599"}, parseOriginErrorPageStatusCodes("not-json"))
assert.Equal(t, []string{"500-599"}, parseOriginErrorPageStatusCodes("[]"))
assert.Equal(t, []string{"502"}, parseOriginErrorPageStatusCodes(`["502"]`))
}
func TestDiffOpenRestyOptionDetailsOriginErrorPage(t *testing.T) {
t.Parallel()
left := openRestyConfigSnapshot{
OriginErrorPageEnabled: true,
OriginErrorPageStatusCodes: []string{"500-599"},
OriginErrorPageHTML: "",
}
right := openRestyConfigSnapshot{
OriginErrorPageEnabled: false,
OriginErrorPageStatusCodes: []string{"522"},
OriginErrorPageHTML: "<p>x</p>",
}
details := diffOpenRestyOptionDetails(left, right)
keys := make(map[string]ConfigOptionDiffItem, len(details))
for _, item := range details {
keys[item.Key] = item
}
assert.Equal(t, "true", keys["OriginErrorPageEnabled"].PreviousValue)
assert.Equal(t, "false", keys["OriginErrorPageEnabled"].CurrentValue)
assert.Equal(t, `["500-599"]`, keys["OriginErrorPageStatusCodes"].PreviousValue)
assert.Equal(t, `["522"]`, keys["OriginErrorPageStatusCodes"].CurrentValue)
assert.Empty(t, keys["OriginErrorPageHTML"].PreviousValue)
assert.Equal(t, "<p>x</p>", keys["OriginErrorPageHTML"].CurrentValue)
}
func TestDiffOpenRestyOptionDetailsSWOfflineDomains(t *testing.T) {
t.Parallel()
left := openRestyConfigSnapshot{
SWOfflineDomains: []string{"a.com,b.com"},
}
right := openRestyConfigSnapshot{
SWOfflineDomains: []string{"a.com", "b.com"},
}
details := diffOpenRestyOptionDetails(left, right)
keys := make(map[string]ConfigOptionDiffItem, len(details))
for _, item := range details {
keys[item.Key] = item
}
assert.Equal(t, `["a.com,b.com"]`, keys["SWOfflineDomains"].PreviousValue)
assert.Equal(t, `["a.com","b.com"]`, keys["SWOfflineDomains"].CurrentValue)
}
@@ -0,0 +1,118 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"errors"
"fmt"
"path"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/share/pagesarchive"
openrestyrender "Wavelet/openflare/share/render/openresty"
"gorm.io/gorm"
)
const defaultPagesSnapshotEntryFile = "index.html"
const defaultPagesSnapshotFallbackPath = "/index.html"
func buildPagesRouteSnapshot(
ctx context.Context,
route *model.ProxyRoute,
) (originURL string, upstreams []string, pagesProjectID *uint, deployment *openrestyrender.PagesDeployment, err error) {
if route == nil {
return "", nil, nil, nil, errors.New("pages 路由配置无效")
}
if !repository.HasPagesProjectsTable(ctx) {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 模块不可用", route.SiteName)
}
if route.PagesProjectID == nil || *route.PagesProjectID == 0 {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: 未绑定 Pages 项目", route.SiteName)
}
project, err := repository.GetPagesProjectByID(ctx, *route.PagesProjectID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目不存在", route.SiteName)
}
return "", nil, nil, nil, err
}
if !project.Enabled {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目未启用", route.SiteName)
}
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目没有激活部署", route.SiteName)
}
activeDeployment, err := repository.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不存在", route.SiteName)
}
return "", nil, nil, nil, err
}
if activeDeployment.ProjectID != project.ID {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不匹配", route.SiteName)
}
if strings.TrimSpace(activeDeployment.Checksum) == "" {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 部署校验和缺失", route.SiteName)
}
pagesProjectID = route.PagesProjectID
deployment, err = buildSnapshotPagesDeployment(project, activeDeployment)
if err != nil {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: %w", route.SiteName, err)
}
originURL = fmt.Sprintf("openflare-pages://project/%d", project.ID)
return originURL, []string{originURL}, pagesProjectID, deployment, nil
}
func buildSnapshotPagesDeployment(
project *model.PagesProject,
activeDeployment *model.PagesDeployment,
) (*openrestyrender.PagesDeployment, error) {
if project == nil || activeDeployment == nil {
return nil, errors.New("pages 项目或部署为空")
}
rootDir, err := pagesarchive.NormalizeLogicalPath(strings.TrimSpace(project.RootDir), true)
if err != nil {
return nil, fmt.Errorf("pages 根目录不合法: %w", err)
}
entryFile := strings.TrimSpace(project.EntryFile)
if entryFile == "" {
entryFile = defaultPagesSnapshotEntryFile
}
entryFile, err = pagesarchive.NormalizeLogicalPath(entryFile, false)
if err != nil {
return nil, fmt.Errorf("pages 入口文件不合法: %w", err)
}
fallbackPath := strings.TrimSpace(project.SPAFallbackPath)
if fallbackPath == "" {
fallbackPath = defaultPagesSnapshotFallbackPath
}
localRoot := openrestyrender.PagesProjectLocalRoot(project.ID)
if rootDir != "" {
localRoot = path.Join(localRoot, rootDir)
}
return &openrestyrender.PagesDeployment{
ProjectID: project.ID,
ProjectSlug: strings.TrimSpace(project.Slug),
DeploymentID: activeDeployment.ID,
DeploymentNumber: activeDeployment.DeploymentNumber,
Checksum: strings.TrimSpace(activeDeployment.Checksum),
EntryFile: entryFile,
SPAFallbackEnabled: project.SPAFallbackEnabled,
SPAFallbackPath: fallbackPath,
APIProxyEnabled: project.APIProxyEnabled,
APIProxyPath: strings.TrimSpace(project.APIProxyPath),
APIProxyPass: strings.TrimSpace(project.APIProxyPass),
APIProxyRewrite: strings.TrimSpace(project.APIProxyRewrite),
// Root is project-scoped so Agents can swap active packages without
// re-publishing main config (nginx root stays stable).
LocalRoot: localRoot,
}, nil
}
@@ -0,0 +1,105 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
openrestyrender "Wavelet/openflare/share/render/openresty"
db "Wavelet/plugins/infra/database"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func TestBuildSnapshotRoutesPages(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
conn := requireDB(t, ctx)
project := &model.PagesProject{
Name: "Speed Test",
Slug: "speedtest",
Enabled: true,
SPAFallbackEnabled: true,
SPAFallbackPath: "/index.html",
RootDir: "public/site",
EntryFile: "index.html",
}
require.NoError(t, conn.Create(project).Error)
deployment := &model.PagesDeployment{
ProjectID: project.ID,
DeploymentNumber: 1,
Checksum: "abc123checksum",
Status: model.PagesDeploymentStatusActive,
FileCount: 1,
TotalSize: 12,
}
require.NoError(t, conn.Create(deployment).Error)
require.NoError(t, conn.Model(project).Update("active_deployment_id", deployment.ID).Error)
route := &model.ProxyRoute{
SiteName: "speedtest",
OriginURL: "openflare-pages://project/1",
Upstreams: `["openflare-pages://project/1"]`,
Enabled: true,
UpstreamType: "pages",
PagesProjectID: &project.ID,
}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, "speedtest.arctel.net")
bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err)
require.Len(t, bundle.SnapshotRoutes, 1)
snapshotRoute := bundle.SnapshotRoutes[0]
assert.Equal(t, "pages", snapshotRoute.UpstreamType)
assert.Equal(t, "openflare-pages://project/1", snapshotRoute.OriginURL)
require.NotNil(t, snapshotRoute.PagesDeployment)
assert.Equal(t, deployment.ID, snapshotRoute.PagesDeployment.DeploymentID)
assert.Equal(t, deployment.Checksum, snapshotRoute.PagesDeployment.Checksum)
assert.Equal(t, "__OPENFLARE_PAGES_DIR__/projects/1/current/public/site", snapshotRoute.PagesDeployment.LocalRoot)
_, err = renderSnapshotConfig(bundle.SnapshotJSON, nil)
require.NoError(t, err)
var decoded struct {
Routes []struct {
PagesDeployment *openrestyrender.PagesDeployment `json:"pages_deployment"`
} `json:"routes"`
}
require.NoError(t, json.Unmarshal([]byte(bundle.SnapshotJSON), &decoded))
require.NotNil(t, decoded.Routes[0].PagesDeployment)
}
func TestBuildSnapshotPagesDeploymentRejectsUnsafeStoredPaths(t *testing.T) {
deployment := &model.PagesDeployment{ID: 1, ProjectID: 1, Checksum: "checksum"}
for _, project := range []*model.PagesProject{
{ID: 1, RootDir: "../escape", EntryFile: "index.html"},
{ID: 1, RootDir: "public", EntryFile: "/index.html"},
} {
_, err := buildSnapshotPagesDeployment(project, deployment)
require.Error(t, err)
}
}
func requireDB(t *testing.T, ctx context.Context) *gorm.DB {
t.Helper()
conn := db.DB(ctx)
require.NotNil(t, conn)
require.NoError(t, conn.AutoMigrate(
&model.PagesProject{},
&model.PagesDeployment{},
))
return conn
}
@@ -0,0 +1,19 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"testing"
openrestyrender "Wavelet/openflare/share/render/openresty"
"github.com/stretchr/testify/assert"
)
func TestNormalizeProxyCachePathForSnapshot(t *testing.T) {
assert.Equal(t, "/var/cache/openresty", normalizeProxyCachePathForSnapshot(false, "/var/cache/openresty"))
assert.Equal(t, openrestyrender.ProxyCachePathPlaceholder, normalizeProxyCachePathForSnapshot(true, "/var/cache/openresty"))
assert.Equal(t, openrestyrender.ProxyCachePathPlaceholder, normalizeProxyCachePathForSnapshot(true, ""))
assert.Equal(t, "/data/var/cache/custom", normalizeProxyCachePathForSnapshot(true, "/data/var/cache/custom"))
}
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
openrestyrender "Wavelet/openflare/share/render/openresty"
)
// SupportFile is a rendered configuration support artifact.
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
func renderSnapshotConfig(sourceJSON string, certificateFiles []SupportFile) (*openrestyrender.Result, error) {
return openrestyrender.RenderJSON(sourceJSON, toOpenRestySupportFiles(certificateFiles))
}
func toOpenRestySupportFiles(files []SupportFile) []openrestyrender.SupportFile {
if len(files) == 0 {
return nil
}
result := make([]openrestyrender.SupportFile, 0, len(files))
for _, file := range files {
result = append(result, openrestyrender.SupportFile{
Path: file.Path,
Content: file.Content,
})
}
return result
}
func fromOpenRestySupportFiles(files []openrestyrender.SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
result := make([]SupportFile, 0, len(files))
for _, file := range files {
result = append(result, SupportFile{
Path: file.Path,
Content: file.Content,
})
}
return result
}
func renderPlaceholderConfig(snapshotJSON string) (mainConfig, routeConfig, checksum string) {
mainConfig = `{"placeholder":"main_config"}`
routeConfig = snapshotJSON
checksum = openrestyrender.ChecksumBundle(mainConfig, routeConfig, nil)
return mainConfig, routeConfig, checksum
}
@@ -0,0 +1,199 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
func handleLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
return apiutil.AbortNotFoundIfMissing(c, err, "记录不存在")
}
func versionParam(c *gin.Context) (string, bool) {
version := c.Param("id")
if version == "" {
response.AbortBadRequest(c, "无效的版本号")
return "", false
}
return version, true
}
// ListConfigVersionsHandler lists config versions.
// @Summary 获取配置版本列表
// @Description 返回所有已发布的 OpenResty 配置版本摘要,按创建时间倒序排列,需要管理员权限
// @Tags openflare-config-version
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.ConfigVersionSummary} "配置版本列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/config-versions [get]
func ListConfigVersionsHandler(c *gin.Context) {
versions, err := ListConfigVersions(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(versions))
}
// GetConfigVersionHandler returns a config version by id.
// @Summary 获取配置版本详情
// @Description 返回指定配置版本的完整快照与渲染内容,需要管理员权限
// @Tags openflare-config-version
// @Produce json
// @Security SessionCookie
// @Param id path int true "配置版本 ID"
// @Success 200 {object} response.Any{data=model.ConfigVersion} "配置版本详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或版本不存在"
// @Router /api/v1/d/config-versions/{id} [get]
func GetConfigVersionHandler(c *gin.Context) {
versionStr, ok := versionParam(c)
if !ok {
return
}
version, err := GetConfigVersionDetail(c.Request.Context(), versionStr)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(version))
}
// GetActiveConfigVersionHandler returns the active config version.
// @Summary 获取当前活跃配置版本
// @Description 返回当前正在使用的配置版本,需要管理员权限
// @Tags openflare-config-version
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.ConfigVersion} "活跃配置版本"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限、不存在或无活跃版本"
// @Router /api/v1/d/config-versions/active [get]
func GetActiveConfigVersionHandler(c *gin.Context) {
version, err := GetActiveConfigVersion(c.Request.Context())
if apiutil.AbortNotFoundIfMissing(c, err, errNoActiveVersion) {
return
}
c.JSON(http.StatusOK, response.OK(version))
}
// PreviewConfigVersionHandler previews the current draft configuration.
// @Summary 预览当前草稿配置
// @Description 渲染并返回当前草稿配置的预览结果,需要管理员权限
// @Tags openflare-config-version
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=config_version.ConfigPreviewResult} "配置预览"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/config-versions/preview [get]
func PreviewConfigVersionHandler(c *gin.Context) {
preview, err := PreviewConfigVersion(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(preview))
}
// DiffConfigVersionHandler diffs the current draft against the active version.
// @Summary 对比草稿与活跃配置
// @Description 对比当前草稿配置与活跃版本之间的差异,需要管理员权限
// @Tags openflare-config-version
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=config_version.ConfigDiffResult} "配置差异"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/config-versions/diff [get]
func DiffConfigVersionHandler(c *gin.Context) {
diff, err := DiffConfigVersion(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(diff))
}
// PublishConfigVersionHandler publishes a new config version.
// @Summary 发布配置版本
// @Description 将当前草稿配置发布为新版本,需要管理员权限
// @Tags openflare-config-version
// @Produce json
// @Security SessionCookie
// @Param force query bool false "是否强制发布"
// @Success 200 {object} response.Any{data=model.ConfigVersion} "发布成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/config-versions/publish [post]
func PublishConfigVersionHandler(c *gin.Context) {
username := c.GetString("username")
force := c.Query("force") == "true"
version, err := PublishConfigVersion(c.Request.Context(), username, force)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(version))
}
// ActivateConfigVersionHandler activates an existing config version.
// @Summary 激活配置版本
// @Description 将指定历史版本设为当前活跃配置,需要管理员权限
// @Tags openflare-config-version
// @Produce json
// @Security SessionCookie
// @Param id path int true "配置版本 ID"
// @Success 200 {object} response.Any{data=model.ConfigVersion} "激活成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或版本不存在"
// @Router /api/v1/d/config-versions/{id}/activate [post]
func ActivateConfigVersionHandler(c *gin.Context) {
versionStr, ok := versionParam(c)
if !ok {
return
}
version, err := ActivateConfigVersion(c.Request.Context(), versionStr)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(version))
}
// CleanupConfigVersionsHandler removes old inactive config versions.
// @Summary 清理历史配置版本
// @Description 删除超出保留数量的非活跃配置版本,需要管理员权限
// @Tags openflare-config-version
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body config_version.CleanupInput true "清理参数"
// @Success 200 {object} response.Any{data=config_version.CleanupResult} "清理结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/config-versions/cleanup [post]
func CleanupConfigVersionsHandler(c *gin.Context) {
var input CleanupInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := CleanupConfigVersions(c.Request.Context(), input.KeepCount)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
@@ -0,0 +1,652 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"errors"
"fmt"
"slices"
"sort"
"strconv"
"strings"
oftls "Wavelet/openflare/plugins/server/domain/tls"
"Wavelet/openflare/plugins/server/domain/waf"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/share/protocol"
openrestyrender "Wavelet/openflare/share/render/openresty"
"gorm.io/gorm"
)
const (
supportFilesPerCertificate = 2
wafIPGroupChecksumHexLength = 64
// OpenResty 默认配置值
defaultOpenRestyReturnStatus = 421
defaultOpenRestyWorkerConns = 4096
defaultOpenRestyRlimitNofile = 65535
defaultOpenRestyKeepaliveTimeout = 20
defaultOpenRestyKeepaliveReqs = 1000
defaultOpenRestyHeaderTimeout = 15
defaultOpenRestyBodyTimeout = 15
defaultOpenRestySendTimeout = 30
defaultOpenRestyConnectTimeout = 3
defaultOpenRestyProxyTimeout = 60
defaultOpenRestyGzipMinLen = 1024
defaultOpenRestyGzipLevel = 5
)
type snapshotRoute struct {
ID uint `json:"id,omitempty"`
SiteName string `json:"site_name,omitempty"`
Domains []string `json:"domains,omitempty"`
OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host,omitempty"`
Upstreams []string `json:"upstreams,omitempty"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
DomainCertIDs []uint `json:"domain_cert_ids,omitempty"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server,omitempty"`
LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"`
LimitRate string `json:"limit_rate,omitempty"`
LimitReqPerIP string `json:"limit_req_per_ip,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy,omitempty"`
CacheRules []string `json:"cache_rules,omitempty"`
CustomHeaders []customHeaderInput `json:"custom_headers,omitempty"`
BasicAuthEnabled bool `json:"basic_auth_enabled,omitempty"`
BasicAuthUsername string `json:"basic_auth_username,omitempty"`
BasicAuthPassword string `json:"basic_auth_password,omitempty"`
UpstreamType string `json:"upstream_type,omitempty"`
TunnelNodeID *uint `json:"tunnel_node_id,omitempty"`
TunnelTargetAddr string `json:"tunnel_target_addr,omitempty"`
TunnelTargetProto string `json:"tunnel_target_protocol,omitempty"`
PagesProjectID *uint `json:"pages_project_id,omitempty"`
PagesDeployment *openrestyrender.PagesDeployment `json:"pages_deployment,omitempty"`
}
type snapshotWAFRuleGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"`
Graph waf.RuntimeRuleGraph `json:"graph"`
}
type snapshotWAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list,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"`
IPGroups []snapshotWAFIPGroup `json:"ip_groups,omitempty"`
Bindings []snapshotWAFBinding `json:"bindings"`
}
type openRestyConfigSnapshot struct {
DefaultServerReturnStatus int `json:"default_server_return_status"`
WorkerProcesses string `json:"worker_processes"`
WorkerConnections int `json:"worker_connections"`
WorkerRlimitNofile int `json:"worker_rlimit_nofile"`
EventsUse string `json:"events_use,omitempty"`
EventsMultiAcceptEnabled bool `json:"events_multi_accept_enabled"`
KeepaliveTimeout int `json:"keepalive_timeout"`
KeepaliveRequests int `json:"keepalive_requests"`
ClientHeaderTimeout int `json:"client_header_timeout"`
ClientBodyTimeout int `json:"client_body_timeout"`
ClientMaxBodySize string `json:"client_max_body_size"`
LargeClientHeaderBuffers string `json:"large_client_header_buffers"`
SendTimeout int `json:"send_timeout"`
ProxyConnectTimeout int `json:"proxy_connect_timeout"`
ProxySendTimeout int `json:"proxy_send_timeout"`
ProxyReadTimeout int `json:"proxy_read_timeout"`
WebsocketEnabled bool `json:"websocket_enabled"`
HTTP3Enabled bool `json:"http3_enabled"`
ProxyRequestBuffering bool `json:"proxy_request_buffering"`
ProxyBufferingEnabled bool `json:"proxy_buffering_enabled"`
ProxyBuffers string `json:"proxy_buffers"`
ProxyBufferSize string `json:"proxy_buffer_size"`
ProxyBusyBuffersSize string `json:"proxy_busy_buffers_size"`
GzipEnabled bool `json:"gzip_enabled"`
GzipMinLength int `json:"gzip_min_length"`
GzipCompLevel int `json:"gzip_comp_level"`
Resolvers string `json:"resolvers,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePath string `json:"cache_path,omitempty"`
CacheLevels string `json:"cache_levels"`
CacheInactive string `json:"cache_inactive"`
CacheMaxSize string `json:"cache_max_size"`
CacheKeyTemplate string `json:"cache_key_template"`
CacheLockEnabled bool `json:"cache_lock_enabled"`
CacheLockTimeout string `json:"cache_lock_timeout"`
CacheUseStale string `json:"cache_use_stale"`
MainConfigTemplate string `json:"main_config_template,omitempty"`
DefaultLimitConnPerServer int `json:"default_limit_conn_per_server,omitempty"`
DefaultLimitConnPerIP int `json:"default_limit_conn_per_ip,omitempty"`
DefaultLimitRate string `json:"default_limit_rate,omitempty"`
DefaultLimitReqPerIP string `json:"default_limit_req_per_ip,omitempty"`
OriginErrorPageEnabled bool `json:"origin_error_page_enabled"`
OriginErrorPageStatusCodes []string `json:"origin_error_page_status_codes,omitempty"`
OriginErrorPageHTML string `json:"origin_error_page_html,omitempty"`
OriginErrorPageGetOnly bool `json:"origin_error_page_get_only,omitempty"`
SWOfflineEnabled bool `json:"sw_offline_enabled,omitempty"`
SWOfflineHTML string `json:"sw_offline_html,omitempty"`
SWOfflineDomains []string `json:"sw_offline_domains,omitempty"`
}
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
RouteConfig string
SupportFiles []SupportFile
Checksum string
ChangedOptionKeys []string
}
func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configBundle, error) {
routes, err := repository.ListEnabledProxyRoutes(ctx)
if err != nil {
return nil, err
}
if requireRoutes && len(routes) == 0 {
return nil, errors.New(errNoEnabledRoutes)
}
snapshotRoutes, err := buildSnapshotRoutes(ctx, routes)
if err != nil {
return nil, err
}
wafSnapshot, err := buildSnapshotWAFDocument(ctx, routes)
if err != nil {
return nil, err
}
openRestyConfig := buildOpenRestyConfigSnapshot(ctx)
snapshotDoc := snapshotDocument{
Routes: snapshotRoutes,
OpenRestyConfig: openRestyConfig,
WAF: wafSnapshot,
}
snapshotJSON, err := json.Marshal(snapshotDoc)
if err != nil {
return nil, err
}
certificateFiles, err := buildCertificateSupportFiles(ctx, snapshotRoutes)
if err != nil {
return nil, err
}
var mainConfig, routeConfig, checksum string
supportFiles := []SupportFile(nil)
rendered, renderErr := renderSnapshotConfig(string(snapshotJSON), certificateFiles)
if renderErr == nil {
mainConfig = rendered.MainConfig
routeConfig = rendered.RouteConfig
checksum = rendered.Checksum
supportFiles = fromOpenRestySupportFiles(rendered.SupportFiles)
} else {
mainConfig, routeConfig, checksum = renderPlaceholderConfig(string(snapshotJSON))
}
return &configBundle{
Routes: routes,
SnapshotRoutes: snapshotRoutes,
WAFSnapshot: wafSnapshot,
OpenRestyConfig: openRestyConfig,
SnapshotJSON: string(snapshotJSON),
MainConfig: mainConfig,
RouteConfig: routeConfig,
SupportFiles: supportFiles,
Checksum: checksum,
ChangedOptionKeys: openRestyOptionKeys(),
}, nil
}
func buildSnapshotRoutes(ctx context.Context, routes []*model.ProxyRoute) ([]snapshotRoute, error) {
items := make([]snapshotRoute, 0, len(routes))
for _, route := range routes {
zoneDomains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return nil, err
}
if len(zoneDomains) == 0 {
return nil, fmt.Errorf("route %s has no zone domains", route.SiteName)
}
domains := make([]string, 0, len(zoneDomains))
domainCertIDs := make([]uint, 0, len(zoneDomains))
for _, zoneDomain := range zoneDomains {
domains = append(domains, zoneDomain.Domain)
if zoneDomain.CertID == nil {
domainCertIDs = append(domainCertIDs, 0)
continue
}
domainCertIDs = append(domainCertIDs, *zoneDomain.CertID)
}
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
if err != nil {
return nil, fmt.Errorf("路由 %s 自定义请求头无效", route.SiteName)
}
upstreamType := normalizeUpstreamType(route.UpstreamType)
originURL := route.OriginURL
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
if err != nil {
return nil, fmt.Errorf("路由 %s 上游配置无效", route.SiteName)
}
var tunnelNodeID *uint
var tunnelTargetAddr string
var tunnelTargetProtocol string
var pagesProjectID *uint
var pagesDeployment *openrestyrender.PagesDeployment
switch upstreamType {
case "tunnel":
originURL = resolveTunnelOpenRestyUpstreamURL(ctx)
upstreams = []string{originURL}
tunnelNodeID = route.TunnelNodeID
tunnelTargetAddr = strings.TrimSpace(route.TunnelTargetAddr)
tunnelTargetProtocol = normalizeTunnelTargetProtocol(route.TunnelTargetProtocol)
case "pages":
originURL, upstreams, pagesProjectID, pagesDeployment, err = buildPagesRouteSnapshot(ctx, route)
if err != nil {
return nil, err
}
}
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
if err != nil {
return nil, fmt.Errorf("路由 %s 缓存规则无效", route.SiteName)
}
items = append(items, snapshotRoute{
ID: route.ID,
SiteName: route.SiteName,
Domains: domains,
OriginURL: originURL,
OriginHost: route.OriginHost,
Upstreams: upstreams,
Enabled: route.Enabled,
EnableHTTPS: route.EnableHTTPS,
DomainCertIDs: domainCertIDs,
RedirectHTTP: route.RedirectHTTP,
LimitConnPerServer: route.LimitConnPerServer,
LimitConnPerIP: route.LimitConnPerIP,
LimitRate: route.LimitRate,
LimitReqPerIP: route.LimitReqPerIP,
CacheEnabled: route.CacheEnabled,
CachePolicy: route.CachePolicy,
CacheRules: cacheRules,
CustomHeaders: customHeaders,
BasicAuthEnabled: route.BasicAuthEnabled,
BasicAuthUsername: route.BasicAuthUsername,
BasicAuthPassword: route.BasicAuthPassword,
UpstreamType: upstreamType,
TunnelNodeID: tunnelNodeID,
TunnelTargetAddr: tunnelTargetAddr,
TunnelTargetProto: tunnelTargetProtocol,
PagesProjectID: pagesProjectID,
PagesDeployment: pagesDeployment,
})
}
return items, nil
}
func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (snapshotWAFDocument, error) {
if err := waf.EnsureDefaultRuleGroup(ctx); err != nil {
return snapshotWAFDocument{}, err
}
groups, err := repository.ListOpenFlareWAFRuleGroups(ctx)
if err != nil {
return snapshotWAFDocument{}, err
}
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(groups))
referencedIPGroupIDs := make(map[uint]struct{})
enabledRuleIDs := make(map[uint]struct{})
for _, group := range groups {
if !group.Enabled {
continue
}
var editorGraph waf.RuleGraph
if err = json.Unmarshal([]byte(group.Graph), &editorGraph); err != nil {
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图数据无效: %w", group.Name, err)
}
if err = waf.ValidateRuleGraph(ctx, editorGraph, snapshotWAFIPGroupExists); err != nil {
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图无效: %w", group.Name, err)
}
runtimeGraph, compileErr := waf.CompileRuleGraph(editorGraph)
if compileErr != nil {
return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 编译失败: %w", group.Name, compileErr)
}
ruleGroups = append(ruleGroups, snapshotWAFRuleGroup{
ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal, Graph: runtimeGraph,
})
enabledRuleIDs[group.ID] = struct{}{}
for _, id := range waf.ReferencedIPGroupIDs(editorGraph) {
referencedIPGroupIDs[id] = struct{}{}
}
}
ipGroups, err := buildSnapshotWAFIPGroups(ctx, referencedIPGroupIDs)
if err != nil {
return snapshotWAFDocument{}, err
}
enabledRouteSiteNames := make(map[uint]string, len(routes))
for _, route := range routes {
if route == nil {
continue
}
domains, domainErr := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if domainErr != nil {
return snapshotWAFDocument{}, domainErr
}
if len(domains) == 0 {
return snapshotWAFDocument{}, fmt.Errorf("route %s has no zone domains", route.SiteName)
}
enabledRouteSiteNames[route.ID] = route.SiteName
}
rawBindings, err := repository.ListOpenFlareWAFRuleGroupBindings(ctx)
if err != nil {
return snapshotWAFDocument{}, err
}
groupIDsByRoute := make(map[uint][]uint, len(rawBindings))
for _, binding := range rawBindings {
if _, ok := enabledRouteSiteNames[binding.ProxyRouteID]; !ok {
continue
}
if _, enabled := enabledRuleIDs[binding.RuleGroupID]; enabled {
groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
}
}
bindings := make([]snapshotWAFBinding, 0, len(enabledRouteSiteNames))
for routeID, siteName := range enabledRouteSiteNames {
bindings = append(bindings, snapshotWAFBinding{
RouteID: routeID,
SiteName: siteName,
RuleGroupIDs: nonNilUintSlice(groupIDsByRoute[routeID]),
})
}
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, IPGroups: ipGroups, Bindings: bindings}, nil
}
func nonNilUintSlice(values []uint) []uint {
if values == nil {
return make([]uint, 0)
}
return values
}
func validateSnapshotWAFIPGroupSize(groups []snapshotWAFIPGroup) error {
runtimeGroups := make(map[string]protocol.WAFIPGroup, len(groups))
for _, group := range groups {
ipList := group.IPList
if !group.Enabled {
ipList = []string{}
}
runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = protocol.WAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ipList,
Checksum: strings.Repeat("0", wafIPGroupChecksumHexLength),
}
}
return protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups)
}
func buildSnapshotWAFIPGroups(ctx context.Context, idSet map[uint]struct{}) ([]snapshotWAFIPGroup, error) {
if len(idSet) == 0 {
return []snapshotWAFIPGroup{}, nil
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
ids = append(ids, id)
}
slices.Sort(ids)
groups, err := listWAFIPGroupsByIDs(ctx, ids)
if err != nil {
return nil, err
}
groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups))
for _, group := range groups {
groupByID[group.ID] = group
}
snapshots := make([]snapshotWAFIPGroup, 0, len(ids))
for _, id := range ids {
group := groupByID[id]
if group == nil {
return nil, fmt.Errorf("IP 组 %d 不存在", id)
}
ipList, decodeErr := decodeIPList(group.IPList)
if decodeErr != nil {
return nil, decodeErr
}
snapshots = append(snapshots, snapshotWAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ipList,
})
}
if err = validateSnapshotWAFIPGroupSize(snapshots); err != nil {
return nil, err
}
return snapshots, nil
}
func snapshotWAFIPGroupExists(ctx context.Context, id uint) (bool, error) {
group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
return group != nil, err
}
func decodeIPList(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []string{}, nil
}
var items []string
if err := json.Unmarshal([]byte(text), &items); err != nil {
return nil, errors.New("ip_list payload is invalid")
}
return items, nil
}
func buildOpenRestyConfigSnapshot(ctx context.Context) openRestyConfigSnapshot {
// 读取所有 OpenResty 配置,使用默认值作为降级
getIntConfig := func(key string, defaultVal int) int {
val, err := repository.GetIntByKey(ctx, key)
if err != nil || val <= 0 {
return defaultVal
}
return val
}
// 0 为合法关闭值,不能与 getIntConfig 的 val<=0 语义混用
getNonNegIntConfig := func(key string, defaultVal int) int {
val, err := repository.GetIntByKey(ctx, key)
if err != nil || val < 0 {
return defaultVal
}
return val
}
getBoolConfig := func(key string, defaultVal bool) bool {
val, err := repository.GetBoolByKey(ctx, key)
if err != nil {
return defaultVal
}
return val
}
getStringConfig := func(key string, defaultVal string) string {
config, err := repository.GetSystemConfigByKey(ctx, key)
if err != nil {
return defaultVal
}
return config.Value
}
getStringSliceConfig := func(key string, defaultVal []string) []string {
config, err := repository.GetSystemConfigByKey(ctx, key)
if err != nil {
return defaultVal
}
var values []string
if err := json.Unmarshal([]byte(config.Value), &values); err != nil {
return defaultVal
}
return values
}
snapshot := openRestyConfigSnapshot{
DefaultServerReturnStatus: getIntConfig(model.ConfigKeyOpenRestyDefaultServerReturnStatus, defaultOpenRestyReturnStatus),
WorkerProcesses: getStringConfig(model.ConfigKeyOpenRestyWorkerProcesses, "auto"),
WorkerConnections: getIntConfig(model.ConfigKeyOpenRestyWorkerConnections, defaultOpenRestyWorkerConns),
WorkerRlimitNofile: getIntConfig(model.ConfigKeyOpenRestyWorkerRlimitNofile, defaultOpenRestyRlimitNofile),
EventsUse: getStringConfig(model.ConfigKeyOpenRestyEventsUse, "epoll"),
EventsMultiAcceptEnabled: getBoolConfig(model.ConfigKeyOpenRestyEventsMultiAcceptEnabled, true),
KeepaliveTimeout: getIntConfig(model.ConfigKeyOpenRestyKeepaliveTimeout, defaultOpenRestyKeepaliveTimeout),
KeepaliveRequests: getIntConfig(model.ConfigKeyOpenRestyKeepaliveRequests, defaultOpenRestyKeepaliveReqs),
ClientHeaderTimeout: getIntConfig(model.ConfigKeyOpenRestyClientHeaderTimeout, defaultOpenRestyHeaderTimeout),
ClientBodyTimeout: getIntConfig(model.ConfigKeyOpenRestyClientBodyTimeout, defaultOpenRestyBodyTimeout),
ClientMaxBodySize: getStringConfig(model.ConfigKeyOpenRestyClientMaxBodySize, "64m"),
LargeClientHeaderBuffers: getStringConfig(model.ConfigKeyOpenRestyLargeClientHeaderBuffers, "4 16k"),
SendTimeout: getIntConfig(model.ConfigKeyOpenRestySendTimeout, defaultOpenRestySendTimeout),
ProxyConnectTimeout: getIntConfig(model.ConfigKeyOpenRestyProxyConnectTimeout, defaultOpenRestyConnectTimeout),
ProxySendTimeout: getIntConfig(model.ConfigKeyOpenRestyProxySendTimeout, defaultOpenRestyProxyTimeout),
ProxyReadTimeout: getIntConfig(model.ConfigKeyOpenRestyProxyReadTimeout, defaultOpenRestyProxyTimeout),
WebsocketEnabled: getBoolConfig(model.ConfigKeyOpenRestyWebsocketEnabled, true),
HTTP3Enabled: getBoolConfig(model.ConfigKeyOpenRestyHTTP3Enabled, true),
ProxyRequestBuffering: getBoolConfig(model.ConfigKeyOpenRestyProxyRequestBufferingEnabled, false),
ProxyBufferingEnabled: getBoolConfig(model.ConfigKeyOpenRestyProxyBufferingEnabled, true),
ProxyBuffers: getStringConfig(model.ConfigKeyOpenRestyProxyBuffers, "16 16k"),
ProxyBufferSize: getStringConfig(model.ConfigKeyOpenRestyProxyBufferSize, "8k"),
ProxyBusyBuffersSize: getStringConfig(model.ConfigKeyOpenRestyProxyBusyBuffersSize, "64k"),
GzipEnabled: getBoolConfig(model.ConfigKeyOpenRestyGzipEnabled, true),
GzipMinLength: getIntConfig(model.ConfigKeyOpenRestyGzipMinLength, defaultOpenRestyGzipMinLen),
GzipCompLevel: getIntConfig(model.ConfigKeyOpenRestyGzipCompLevel, defaultOpenRestyGzipLevel),
Resolvers: getStringConfig(model.ConfigKeyOpenRestyResolvers, ""),
CacheEnabled: getBoolConfig(model.ConfigKeyOpenRestyCacheEnabled, false),
CachePath: getStringConfig(model.ConfigKeyOpenRestyCachePath, ""),
CacheLevels: getStringConfig(model.ConfigKeyOpenRestyCacheLevels, "1:2"),
CacheInactive: getStringConfig(model.ConfigKeyOpenRestyCacheInactive, "30m"),
CacheMaxSize: getStringConfig(model.ConfigKeyOpenRestyCacheMaxSize, "1g"),
CacheKeyTemplate: getStringConfig(model.ConfigKeyOpenRestyCacheKeyTemplate, "$scheme$host$request_uri"),
CacheLockEnabled: getBoolConfig(model.ConfigKeyOpenRestyCacheLockEnabled, true),
CacheLockTimeout: getStringConfig(model.ConfigKeyOpenRestyCacheLockTimeout, "5s"),
CacheUseStale: getStringConfig(model.ConfigKeyOpenRestyCacheUseStale, "error timeout updating http_500 http_502 http_503 http_504"),
MainConfigTemplate: getStringConfig(model.ConfigKeyOpenRestyMainConfigTemplate, model.DefaultOpenRestyMainConfigTemplate),
DefaultLimitConnPerServer: getNonNegIntConfig(model.ConfigKeyOpenRestyDefaultLimitConnPerServer, 0),
DefaultLimitConnPerIP: getNonNegIntConfig(model.ConfigKeyOpenRestyDefaultLimitConnPerIP, 0),
DefaultLimitRate: strings.ToLower(strings.TrimSpace(getStringConfig(model.ConfigKeyOpenRestyDefaultLimitRate, ""))),
DefaultLimitReqPerIP: strings.ToLower(strings.TrimSpace(getStringConfig(model.ConfigKeyOpenRestyDefaultLimitReqPerIP, ""))),
OriginErrorPageEnabled: getBoolConfig(model.ConfigKeyOriginErrorPageEnabled, true),
OriginErrorPageStatusCodes: parseOriginErrorPageStatusCodes(getStringConfig(model.ConfigKeyOriginErrorPageStatusCodes, `["500-599"]`)),
OriginErrorPageHTML: getStringConfig(model.ConfigKeyOriginErrorPageHTML, ""),
OriginErrorPageGetOnly: getBoolConfig(model.ConfigKeyOriginErrorPageGetOnly, false),
SWOfflineEnabled: getBoolConfig(model.ConfigKeySWOfflineEnabled, false),
SWOfflineHTML: getStringConfig(model.ConfigKeySWOfflineHTML, ""),
SWOfflineDomains: getStringSliceConfig(model.ConfigKeySWOfflineDomains, nil),
}
if snapshot.DefaultLimitRate == "0" {
snapshot.DefaultLimitRate = ""
}
if snapshot.DefaultLimitReqPerIP == "0" {
snapshot.DefaultLimitReqPerIP = ""
}
snapshot.CachePath = normalizeProxyCachePathForSnapshot(snapshot.CacheEnabled, snapshot.CachePath)
return snapshot
}
func parseOriginErrorPageStatusCodes(raw string) []string {
const defaultTag = "500-599"
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return []string{defaultTag}
}
var tags []string
if err := json.Unmarshal([]byte(trimmed), &tags); err != nil || len(tags) == 0 {
return []string{defaultTag}
}
return tags
}
func normalizeProxyCachePathForSnapshot(cacheEnabled bool, cachePath string) string {
if !cacheEnabled {
return strings.TrimSpace(cachePath)
}
trimmed := strings.TrimSpace(cachePath)
if trimmed == "" || strings.HasPrefix(trimmed, "/var/") {
return openrestyrender.ProxyCachePathPlaceholder
}
return trimmed
}
func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) ([]SupportFile, error) {
certIDSet := make(map[uint]struct{})
for _, route := range routes {
for _, certID := range route.DomainCertIDs {
if certID != 0 {
certIDSet[certID] = struct{}{}
}
}
}
if len(certIDSet) == 0 {
return nil, nil
}
certIDs := make([]uint, 0, len(certIDSet))
for certID := range certIDSet {
certIDs = append(certIDs, certID)
}
slices.Sort(certIDs)
files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate)
for _, certID := range certIDs {
certificate, err := repository.GetTLSCertificateByID(ctx, certID)
if err != nil {
return nil, err
}
keyPEM, err := oftls.OpenKeyPEM(certificate.KeyPEM)
if err != nil {
return nil, fmt.Errorf("certificate %d private key: %w", certificate.ID, err)
}
if strings.TrimSpace(keyPEM) == "" {
return nil, fmt.Errorf("certificate %d has no private key", certificate.ID)
}
files = append(files,
SupportFile{Path: certificateCertFileName(certificate.ID), Content: normalizePEM(certificate.CertPEM)},
SupportFile{Path: certificateKeyFileName(certificate.ID), Content: normalizePEM(keyPEM)},
)
}
return dedupeSupportFiles(files), nil
}
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config_version
import (
"context"
"encoding/json"
"strings"
"testing"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/domain/waf"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestBuildSnapshotRejectsOversizedAggregateWAFIPGroups(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
// Each group remains below the existing 2 MiB per-subscription ceiling,
// while the complete Agent runtime document exceeds the aggregate limit.
ipList, err := json.Marshal(strings.Fields(strings.Repeat("192.0.2.1 ", 165000)))
require.NoError(t, err)
require.Less(t, len(ipList), 2<<20)
groupIDs := make([]uint, 0, 12)
for index := 0; index < 12; index++ {
group := &model.OpenFlareWAFIPGroup{
Name: "aggregate-" + strings.Repeat("x", index),
Type: "manual",
Enabled: true,
IPList: string(ipList),
}
require.NoError(t, db.DB(ctx).Create(group).Error)
groupIDs = append(groupIDs, group.ID)
}
createSnapshotRule(t, ctx, "oversized-aggregate", snapshotIPMatchGraphForGroups(groupIDs))
_, err = buildSnapshotWAFDocument(ctx, nil)
require.ErrorContains(t, err, "WAF IP 组快照大小")
require.ErrorContains(t, err, "超过上限")
}
func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
route := &model.ProxyRoute{SiteName: "ordered.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, route.SiteName)
referenced := &model.OpenFlareWAFIPGroup{Name: "referenced", Type: "manual", Enabled: true, IPList: `["192.0.2.1"]`}
unused := &model.OpenFlareWAFIPGroup{Name: "unused", Type: "manual", Enabled: true, IPList: `["198.51.100.1"]`}
require.NoError(t, db.DB(ctx).Create(referenced).Error)
require.NoError(t, db.DB(ctx).Create(unused).Error)
customA := createSnapshotRule(t, ctx, "custom-a", waf.DefaultRuleGraph())
customB := createSnapshotRule(t, ctx, "custom-b", snapshotIPMatchGraph(referenced.ID))
require.NoError(t, repository.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, route.ID, []uint{customB.ID, customA.ID}))
snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route})
require.NoError(t, err)
require.Len(t, snapshot.Bindings, 1)
assert.Equal(t, []uint{customB.ID, customA.ID}, snapshot.Bindings[0].RuleGroupIDs)
require.Len(t, snapshot.IPGroups, 1)
assert.Equal(t, referenced.ID, snapshot.IPGroups[0].ID)
var customBSnapshot *snapshotWAFRuleGroup
for index := range snapshot.RuleGroups {
if snapshot.RuleGroups[index].ID == customB.ID {
customBSnapshot = &snapshot.RuleGroups[index]
}
}
require.NotNil(t, customBSnapshot)
assert.Equal(t, "start", customBSnapshot.Graph.Entry)
assert.Equal(t, waf.RuleNodeIPMatch, customBSnapshot.Graph.Nodes["match"].Type)
raw, err := json.Marshal(customBSnapshot)
require.NoError(t, err)
assert.NotContains(t, string(raw), "position")
assert.NotContains(t, string(raw), "ip_whitelist")
}
func TestWAFGraphSnapshotEncodesEmptyBindingsAsArrays(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
route := &model.ProxyRoute{SiteName: "empty-binding.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true}
require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, route.SiteName)
snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route})
require.NoError(t, err)
require.Len(t, snapshot.Bindings, 1)
require.NotNil(t, snapshot.Bindings[0].RuleGroupIDs)
raw, err := json.Marshal(snapshot)
require.NoError(t, err)
assert.Contains(t, string(raw), `"rule_group_ids":[]`)
assert.NotContains(t, string(raw), `"rule_group_ids":null`)
}
func TestBuildSnapshotRejectsInvalidWAFGraph(t *testing.T) {
cleanup := setupConfigVersionTestDB(t)
defer cleanup()
ctx := context.Background()
invalid := &model.OpenFlareWAFRuleGroup{Name: "invalid", Enabled: true, Graph: `{"schema_version":1,"nodes":[],"edges":[]}`, Revision: 1}
require.NoError(t, db.DB(ctx).Create(invalid).Error)
_, err := buildSnapshotWAFDocument(ctx, nil)
require.ErrorContains(t, err, "invalid")
}
func createSnapshotRule(t *testing.T, ctx context.Context, name string, graph waf.RuleGraph) *model.OpenFlareWAFRuleGroup {
t.Helper()
raw, err := json.Marshal(graph)
require.NoError(t, err)
rule := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: true, Graph: string(raw), Revision: 1}
require.NoError(t, db.DB(ctx).Create(rule).Error)
return rule
}
func snapshotIPMatchGraph(ipGroupID uint) waf.RuleGraph {
return snapshotIPMatchGraphForGroups([]uint{ipGroupID})
}
func snapshotIPMatchGraphForGroups(ipGroupIDs []uint) waf.RuleGraph {
config, _ := json.Marshal(waf.IPMatchConfig{IPGroupIDs: ipGroupIDs})
return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
{ID: "start", Type: waf.RuleNodeStart, Position: waf.RulePosition{X: 1, Y: 2}, Config: json.RawMessage(`{}`)},
{ID: "match", Type: waf.RuleNodeIPMatch, Position: waf.RulePosition{X: 3, Y: 4}, Config: config},
{ID: "allow", Type: waf.RuleNodeAllow, Position: waf.RulePosition{X: 5, Y: 6}, Config: json.RawMessage(`{}`)},
}, Edges: []waf.RuleEdge{
{ID: "e1", Source: "start", SourceHandle: "next", Target: "match"},
{ID: "e2", Source: "match", SourceHandle: "true", Target: "allow"},
{ID: "e3", Source: "match", SourceHandle: "false", Target: "allow"},
}}
}
func snapshotPoWGraph() waf.RuleGraph {
config, _ := json.Marshal(waf.PoWNodeConfig{Algorithm: "fast", Difficulty: 4, SessionTTL: 600, ChallengeTTL: 300})
return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
{ID: "start", Type: waf.RuleNodeStart, Config: json.RawMessage(`{}`)},
{ID: "pow", Type: waf.RuleNodePoW, Config: config},
{ID: "allow", Type: waf.RuleNodeAllow, Config: json.RawMessage(`{}`)},
}, Edges: []waf.RuleEdge{
{ID: "e1", Source: "start", SourceHandle: "next", Target: "pow"},
{ID: "e2", Source: "pow", SourceHandle: "next", Target: "allow"},
}}
}
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package origin defines shared error messages for origin management.
package origin
const (
errOriginAddressRequired = "源站地址不能为空"
errOriginAddressInvalid = "源站地址格式不合法"
errOriginAddressExists = "源站地址已存在"
errOriginDeleteReferenced = "该源站仍被规则引用,无法删除"
errOriginMissingPort = "源站地址缺少端口"
errOriginNotFound = "源站不存在"
)
@@ -0,0 +1,89 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package origin
import (
"errors"
"fmt"
"net"
"net/url"
"strings"
"unicode"
)
const maxOriginHostnameLength = 253
func normalizeOriginAddress(raw string) string {
return strings.ToLower(strings.TrimSpace(raw))
}
func validateOriginAddress(address string) error {
if address == "" {
return errors.New(errOriginAddressRequired)
}
if strings.Contains(address, "://") || strings.ContainsAny(address, "/?#") {
return errors.New(errOriginAddressInvalid)
}
if strings.HasPrefix(address, "[") || strings.HasSuffix(address, "]") {
return errors.New(errOriginAddressInvalid)
}
if ip := net.ParseIP(address); ip != nil {
return nil
}
if len(address) > maxOriginHostnameLength {
return errors.New(errOriginAddressInvalid)
}
labels := strings.SplitSeq(address, ".")
for label := range labels {
if len(label) == 0 || len(label) > 63 {
return errors.New(errOriginAddressInvalid)
}
if label[0] == '-' || label[len(label)-1] == '-' {
return errors.New(errOriginAddressInvalid)
}
for _, r := range label {
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' {
continue
}
return errors.New(errOriginAddressInvalid)
}
}
return nil
}
func normalizeOriginName(name string, address string) string {
normalized := strings.TrimSpace(name)
if normalized != "" {
return normalized
}
return address
}
func formatOriginHost(address string, port string) string {
return net.JoinHostPort(address, port)
}
func rewriteOriginURLAddress(rawURL string, newAddress string) (string, error) {
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
if err != nil {
return "", fmt.Errorf("%s: %w", errOriginAddressInvalid, err)
}
address := normalizeOriginAddress(newAddress)
if err := validateOriginAddress(address); err != nil {
return "", err
}
port := parsed.Port()
if port == "" {
return "", errors.New(errOriginMissingPort)
}
parsed.Host = formatOriginHost(address, port)
return parsed.String(), nil
}
func isUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
@@ -0,0 +1,233 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package origin
import (
"context"
"encoding/json"
"errors"
"fmt"
"sort"
"strings"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"gorm.io/gorm"
)
// Input 源站创建/更新请求。
type Input struct {
Name string `json:"name"`
Address string `json:"address"`
Remark string `json:"remark"`
}
// RouteSummary 源站详情中的代理规则摘要。
type RouteSummary struct {
ID uint `json:"id"`
Domain string `json:"domain"`
OriginURL string `json:"origin_url"`
Enabled bool `json:"enabled"`
UpdatedAt string `json:"updated_at"`
}
// View 源站列表项。
type View struct {
ID uint `json:"id"`
Name string `json:"name"`
Address string `json:"address"`
Remark string `json:"remark"`
RouteCount int64 `json:"route_count"`
CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"`
}
// DetailView 源站详情。
type DetailView struct {
View
Routes []RouteSummary `json:"routes"`
}
// ListOrigins 列出全部源站。
func ListOrigins(ctx context.Context) ([]View, error) {
origins, err := repository.ListOrigins(ctx)
if err != nil {
return nil, err
}
return buildOriginViews(ctx, origins)
}
// GetOriginDetail 获取源站详情。
func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) {
origin, err := repository.GetOriginByID(ctx, id)
if err != nil {
return nil, err
}
views, err := buildOriginViews(ctx, []model.Origin{*origin})
if err != nil {
return nil, err
}
routes, err := repository.ListProxyRoutesByOriginID(ctx, id)
if err != nil {
return nil, err
}
items := make([]RouteSummary, 0, len(routes))
for _, route := range routes {
domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return nil, err
}
domain := ""
if len(domains) > 0 {
domain = domains[0].Domain
}
items = append(items, RouteSummary{
ID: route.ID,
Domain: domain,
OriginURL: route.OriginURL,
Enabled: route.Enabled,
UpdatedAt: route.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
})
}
sort.Slice(items, func(i, j int) bool {
return items[i].Domain < items[j].Domain
})
return &DetailView{
View: views[0],
Routes: items,
}, nil
}
// CreateOrigin 创建源站。
func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) {
origin, err := buildOrigin(nil, input)
if err != nil {
return nil, err
}
if err = repository.CreateOriginRecord(ctx, origin); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errOriginAddressExists)
}
return nil, err
}
return origin, nil
}
// UpdateOrigin 更新源站。
func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, error) {
origin, err := repository.GetOriginByID(ctx, id)
if err != nil {
return nil, err
}
previousAddress := origin.Address
nextOrigin, err := buildOrigin(origin, input)
if err != nil {
return nil, err
}
err = repository.WithOriginTx(ctx, func(tx *gorm.DB) error {
if err := repository.SaveOriginTx(tx, nextOrigin); err != nil {
if isUniqueConstraintError(err) {
return errors.New(errOriginAddressExists)
}
return err
}
if previousAddress == nextOrigin.Address {
return nil
}
return updateRoutesForOriginAddress(ctx, tx, nextOrigin.ID, nextOrigin.Address)
})
if err != nil {
return nil, err
}
return nextOrigin, nil
}
// DeleteOrigin 删除源站。
func DeleteOrigin(ctx context.Context, id uint) error {
count, err := repository.CountProxyRoutesByOriginID(ctx, id)
if err != nil {
return err
}
if count > 0 {
return errors.New(errOriginDeleteReferenced)
}
if _, err = repository.GetOriginByID(ctx, id); err != nil {
return err
}
return repository.DeleteOriginRecord(ctx, id)
}
func buildOrigin(existing *model.Origin, input Input) (*model.Origin, error) {
address := normalizeOriginAddress(input.Address)
if err := validateOriginAddress(address); err != nil {
return nil, err
}
if existing == nil {
existing = &model.Origin{}
}
existing.Address = address
existing.Name = normalizeOriginName(input.Name, address)
existing.Remark = strings.TrimSpace(input.Remark)
return existing, nil
}
func buildOriginViews(ctx context.Context, origins []model.Origin) ([]View, error) {
countRows, err := repository.ListOriginRouteCounts(ctx)
if err != nil {
return nil, err
}
countMap := make(map[uint]int64, len(countRows))
for _, row := range countRows {
countMap[row.OriginID] = row.RouteCount
}
views := make([]View, 0, len(origins))
for _, origin := range origins {
views = append(views, View{
ID: origin.ID,
Name: origin.Name,
Address: origin.Address,
Remark: origin.Remark,
RouteCount: countMap[origin.ID],
CreatedAt: origin.CreatedAt.Format("2006-01-02T15:04:05Z07:00"),
UpdatedAt: origin.UpdatedAt.Format("2006-01-02T15:04:05Z07:00"),
})
}
return views, nil
}
func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uint, address string) error {
if !repository.HasProxyRoutesTable(ctx) {
return nil
}
routes, err := repository.ListProxyRoutesByOriginIDAscTx(tx, originID)
if err != nil {
return fmt.Errorf("query routes for origin update failed: %w", err)
}
for _, route := range routes {
rewrittenOriginURL, err := rewriteOriginURLAddress(route.OriginURL, address)
if err != nil {
return fmt.Errorf("rewrite route %d origin failed: %w", route.ID, err)
}
upstreams := make([]string, 0)
if strings.TrimSpace(route.Upstreams) != "" {
if err := json.Unmarshal([]byte(route.Upstreams), &upstreams); err != nil {
return fmt.Errorf("decode route %d upstreams failed: %w", route.ID, err)
}
}
if len(upstreams) == 0 {
upstreams = append(upstreams, rewrittenOriginURL)
} else {
upstreams[0] = rewrittenOriginURL
}
upstreamsJSON, err := json.Marshal(upstreams)
if err != nil {
return fmt.Errorf("encode route %d upstreams failed: %w", route.ID, err)
}
if err := repository.UpdateProxyRouteOriginAddressTx(tx, route.ID, rewrittenOriginURL, string(upstreamsJSON)); err != nil {
return fmt.Errorf("update route %d origin address failed: %w", route.ID, err)
}
}
return nil
}
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package origin
import (
"context"
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupOriginTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.Origin{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func TestCreateOrigin(t *testing.T) {
cleanup := setupOriginTestDB(t)
defer cleanup()
ctx := context.Background()
origin, err := CreateOrigin(ctx, Input{
Name: "Primary Origin",
Address: "origin-a.internal",
Remark: "main upstream",
})
require.NoError(t, err)
assert.NotZero(t, origin.ID)
assert.Equal(t, "Primary Origin", origin.Name)
assert.Equal(t, "origin-a.internal", origin.Address)
assert.Equal(t, "main upstream", origin.Remark)
_, err = CreateOrigin(ctx, Input{
Address: "origin-a.internal",
})
require.Error(t, err)
assert.Equal(t, errOriginAddressExists, err.Error())
}
func TestListOrigins(t *testing.T) {
cleanup := setupOriginTestDB(t)
defer cleanup()
ctx := context.Background()
first, err := CreateOrigin(ctx, Input{
Name: "first-origin",
Address: "origin-a.internal",
})
require.NoError(t, err)
second, err := CreateOrigin(ctx, Input{
Name: "second-origin",
Address: "origin-b.internal",
})
require.NoError(t, err)
origins, err := ListOrigins(ctx)
require.NoError(t, err)
require.Len(t, origins, 2)
assert.Equal(t, second.ID, origins[0].ID)
assert.Equal(t, first.ID, origins[1].ID)
assert.Equal(t, int64(0), origins[0].RouteCount)
assert.Equal(t, int64(0), origins[1].RouteCount)
}
@@ -0,0 +1,141 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package origin
import (
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
func handleLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
return apiutil.AbortNotFoundIfMissing(c, err, errOriginNotFound)
}
// GetOrigins 列出全部源站。
// @Summary 获取源站列表
// @Description 返回所有源站及关联代理规则数量,需要管理员权限
// @Tags openflare-origin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]origin.View} "源站列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/origins [get]
func GetOrigins(c *gin.Context) {
origins, err := ListOrigins(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(origins))
}
// GetOrigin 获取源站详情。
// @Summary 获取源站详情
// @Description 返回指定源站信息及关联代理规则摘要,需要管理员权限
// @Tags openflare-origin
// @Produce json
// @Security SessionCookie
// @Param id path int true "源站 ID"
// @Success 200 {object} response.Any{data=origin.DetailView} "源站详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或源站不存在"
// @Router /api/v1/d/origins/{id} [get]
func GetOrigin(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
detail, err := GetOriginDetail(c.Request.Context(), id)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(detail))
}
// CreateOriginHandler 创建源站。
// @Summary 创建源站
// @Description 创建新的上游源站记录,需要管理员权限
// @Tags openflare-origin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body origin.Input true "源站参数"
// @Success 200 {object} response.Any{data=origin.View} "创建成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/origins [post]
func CreateOriginHandler(c *gin.Context) {
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
origin, err := CreateOrigin(c.Request.Context(), input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(origin))
}
// UpdateOriginHandler 更新源站。
// @Summary 更新源站
// @Description 更新指定源站的配置信息,需要管理员权限
// @Tags openflare-origin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "源站 ID"
// @Param body body origin.Input true "源站参数"
// @Success 200 {object} response.Any{data=origin.View} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或源站不存在"
// @Router /api/v1/d/origins/{id}/update [post]
func UpdateOriginHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
origin, err := UpdateOrigin(c.Request.Context(), id, input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(origin))
}
// DeleteOriginHandler 删除源站。
// @Summary 删除源站
// @Description 删除指定源站记录,需要管理员权限
// @Tags openflare-origin
// @Produce json
// @Security SessionCookie
// @Param id path int true "源站 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或源站不存在"
// @Router /api/v1/d/origins/{id}/delete [post]
func DeleteOriginHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if err := DeleteOrigin(c.Request.Context(), id); handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,142 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package proxy_route provides helpers for building proxy route configurations.
package proxy_route
import (
"context"
"encoding/json"
"errors"
"strings"
"Wavelet/openflare/plugins/server/kernel/model"
)
type proxyRouteJSONFields struct {
cacheRulesJSON string
upstreamsJSON string
customHeadersJSON string
}
func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input Input) (string, *uint, []string, error) {
switch upstreamType {
case proxyRouteUpstreamTypeTunnel, proxyRouteUpstreamTypePages:
if upstreamType == proxyRouteUpstreamTypePages {
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
return "", nil, nil, err
}
}
originURL := "http://127.0.0.1"
return originURL, nil, []string{originURL}, nil
default:
originURL, originID, err := resolveProxyRoutePrimaryOrigin(ctx, input)
if err != nil {
return "", nil, nil, err
}
upstreams, err := normalizeUpstreams(originURL, input.Upstreams)
if err != nil {
return "", nil, nil, err
}
return originURL, originID, upstreams, nil
}
}
func marshalProxyRouteJSONFields(
upstreams []string,
cacheRules []string,
customHeaders []CustomHeaderInput,
) (*proxyRouteJSONFields, error) {
cacheRulesJSON, err := json.Marshal(cacheRules)
if err != nil {
return nil, err
}
upstreamsJSON, err := json.Marshal(upstreams)
if err != nil {
return nil, err
}
customHeadersJSON, err := json.Marshal(customHeaders)
if err != nil {
return nil, err
}
return &proxyRouteJSONFields{
cacheRulesJSON: string(cacheRulesJSON),
upstreamsJSON: string(upstreamsJSON),
customHeadersJSON: string(customHeadersJSON),
}, nil
}
func normalizeProxyRouteBasicAuth(input *Input) error {
if !input.BasicAuthEnabled {
input.BasicAuthUsername = ""
input.BasicAuthPassword = ""
return nil
}
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
return errors.New(errProxyRouteBasicAuth)
}
return nil
}
func populateProxyRouteFields(
route *model.ProxyRoute,
input Input,
siteName string,
jsonFields *proxyRouteJSONFields,
originID *uint,
upstreams []string,
originHost, cachePolicy string,
limitConnPerServer, limitConnPerIP int,
limitRate, limitReqPerIP, upstreamType string,
) {
route.SiteName = siteName
route.OriginID = originID
route.OriginURL = upstreams[0]
route.OriginHost = originHost
route.Upstreams = jsonFields.upstreamsJSON
route.Enabled = input.Enabled
route.EnableHTTPS = input.EnableHTTPS
route.RedirectHTTP = input.RedirectHTTP
route.LimitConnPerServer = limitConnPerServer
route.LimitConnPerIP = limitConnPerIP
route.LimitRate = limitRate
route.LimitReqPerIP = limitReqPerIP
route.CacheEnabled = input.CacheEnabled
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
route.CacheRules = jsonFields.cacheRulesJSON
route.CustomHeaders = jsonFields.customHeadersJSON
route.BasicAuthEnabled = input.BasicAuthEnabled
route.BasicAuthUsername = input.BasicAuthUsername
route.BasicAuthPassword = input.BasicAuthPassword
route.UpstreamType = upstreamType
}
func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, upstreamType string, input Input) error {
switch upstreamType {
case proxyRouteUpstreamTypeTunnel:
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
if err != nil {
return err
}
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
return err
}
route.TunnelNodeID = tunnelNodeID
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
route.PagesProjectID = nil
case proxyRouteUpstreamTypePages:
route.TunnelNodeID = nil
route.TunnelTargetAddr = ""
route.TunnelTargetProtocol = ""
route.PagesProjectID = input.PagesProjectID
default:
route.TunnelNodeID = nil
route.TunnelTargetAddr = ""
route.TunnelTargetProtocol = ""
route.PagesProjectID = nil
}
return nil
}
@@ -0,0 +1,54 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
const (
errProxyRouteNotFound = "proxy route not found"
errProxyRouteIdentityExists = "proxy route identity already exists"
errProxyRouteSiteNameExists = "site_name already exists"
errProxyRouteDomainExists = "domain %s already exists"
errProxyRouteSiteNameEmpty = "site_name cannot be empty"
errProxyRouteZoneDomainsRequired = "at least one zone domain is required"
errProxyRouteZoneDomainNotFound = "selected zone domain does not exist"
errProxyRouteZoneDomainDuplicate = "zone_domain_ids must not contain duplicates"
errProxyRouteZoneDomainBound = "selected zone domain is already bound to another proxy route"
errProxyRouteOriginEmpty = "origin_url cannot be empty"
errProxyRouteOriginInvalid = "origin URL format is invalid"
errProxyRouteOriginScheme = "origin URL must start with http:// or https://"
errProxyRouteOriginHostInvalid = "origin_host format is invalid"
errProxyRouteUpstreamRequired = "at least one upstream is required"
errProxyRouteUpstreamScheme = "all upstreams must use the same scheme"
errProxyRouteUpstreamPath = "multi-upstream mode does not support origin paths"
errProxyRouteUpstreamQuery = "multi-upstream mode does not support origin query strings"
errProxyRouteOriginNotFound = "selected origin does not exist"
errProxyRouteCertNotFound = "selected certificate does not exist"
errProxyRouteCertRequired = "must select a certificate when HTTPS is enabled"
errProxyRouteCertDomainLength = "domain_cert_ids must match domains length"
errProxyRouteRedirectHTTP = "redirect_http requires enable_https"
errProxyRouteBasicAuth = "basic_auth_username and basic_auth_password cannot be empty when basic auth is enabled"
errProxyRouteLimitRate = "limit_rate must be a number or use the 512k / 1m format, or -1 to disable"
errProxyRouteCachePolicy = "cache policy is not supported"
errProxyRouteCacheSuffix = "cache suffix format is invalid"
errProxyRouteCachePath = "cache path rule format is invalid"
errProxyRouteCacheSuffixReq = "at least one suffix is required"
errProxyRouteCachePrefixReq = "at least one path prefix is required"
errProxyRouteCacheExactReq = "at least one exact path is required"
errProxyRouteHeaderKeyEmpty = "custom header key cannot be empty"
errProxyRouteHeaderKeyInvalid = "custom header key format is invalid"
errProxyRouteHeaderNewline = "custom headers cannot contain newlines"
errProxyRouteTunnelNodeReq = "tunnel_node_id is required for tunnel upstream"
errProxyRouteTunnelNodeMissing = "tunnel client node does not exist"
errProxyRouteTunnelNodeType = "tunnel_node_id must reference a tunnel_client node"
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
errProxyRoutePagesNotFound = "pages 项目不存在"
errProxyRoutePagesDisabled = "pages 项目未启用"
errProxyRoutePagesNoDeploy = "pages 项目没有激活部署"
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
errProxyRouteOriginPort = "端口格式不合法"
errProxyRouteOriginPortEmpty = "端口不能为空"
errProxyRouteOriginURI = "源站路径需以 / 或 ? 开头"
errProxyRouteOriginURIProto = "源站路径不能包含协议"
)
@@ -0,0 +1,699 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"crypto/x509"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"net"
"net/url"
"regexp"
"strconv"
"strings"
"unicode"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/model"
"gorm.io/gorm"
)
var proxyHeaderKeyPattern = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
var proxyRouteLimitRatePattern = regexp.MustCompile(`^\d+[kKmM]?$`)
var proxyRouteLimitReqPattern = regexp.MustCompile(`^\d+r/[sm]$`)
const (
proxyRouteCachePolicyStatic = "static"
proxyRouteCachePolicyAll = "all"
proxyRouteCachePolicyURL = "url" // legacy alias of all
proxyRouteCachePolicySuffix = "suffix"
proxyRouteCachePolicyPathPrefix = "path_prefix"
proxyRouteCachePolicyPathExact = "path_exact"
proxyRouteSchemeHTTP = "http"
proxyRouteSchemeHTTPS = "https"
proxyRouteUpstreamTypeTunnel = "tunnel"
proxyRouteUpstreamTypePages = "pages"
maxOriginHostnameLength = 253
originURIPathQueryParts = 2
)
func uniqueStrings(items []string) []string {
if len(items) == 0 {
return items
}
seen := make(map[string]struct{}, len(items))
result := make([]string, 0, len(items))
for _, item := range items {
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
result = append(result, item)
}
return result
}
func isUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
func normalizeOriginAddress(raw string) string {
return strings.ToLower(strings.TrimSpace(raw))
}
func validateOriginAddress(address string) error {
if address == "" {
return errors.New(errProxyRouteOriginEmpty)
}
if strings.Contains(address, "://") || strings.ContainsAny(address, "/?#") {
return errors.New(errProxyRouteOriginInvalid)
}
if strings.HasPrefix(address, "[") || strings.HasSuffix(address, "]") {
return errors.New(errProxyRouteOriginInvalid)
}
if ip := net.ParseIP(address); ip != nil {
return nil
}
if len(address) > maxOriginHostnameLength {
return errors.New(errProxyRouteOriginInvalid)
}
labels := strings.SplitSeq(address, ".")
for label := range labels {
if len(label) == 0 || len(label) > 63 {
return errors.New(errProxyRouteOriginInvalid)
}
if label[0] == '-' || label[len(label)-1] == '-' {
return errors.New(errProxyRouteOriginInvalid)
}
for _, r := range label {
if unicode.IsLetter(r) || unicode.IsDigit(r) || r == '-' {
continue
}
return errors.New(errProxyRouteOriginInvalid)
}
}
return nil
}
func normalizeOriginPort(raw string) (string, error) {
port := strings.TrimSpace(raw)
if port == "" {
return "", errors.New(errProxyRouteOriginPortEmpty)
}
value, err := strconv.Atoi(port)
if err != nil || value < 1 || value > 65535 {
return "", errors.New(errProxyRouteOriginPort)
}
return strconv.Itoa(value), nil
}
func normalizeOriginScheme(raw string) (string, error) {
scheme := strings.ToLower(strings.TrimSpace(raw))
switch scheme {
case proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
return scheme, nil
default:
return "", errors.New(errProxyRouteOriginSchemeOnly)
}
}
func normalizeOriginURI(raw string) (string, error) {
uri := strings.TrimSpace(raw)
if uri == "" {
return "", nil
}
if strings.Contains(uri, "://") {
return "", errors.New(errProxyRouteOriginURIProto)
}
if !strings.HasPrefix(uri, "/") && !strings.HasPrefix(uri, "?") {
return "", errors.New(errProxyRouteOriginURI)
}
return uri, nil
}
func formatOriginHost(address string, port string) string {
return net.JoinHostPort(address, port)
}
func buildOriginURLFromParts(scheme, address, port, uri string) (string, error) {
normalizedScheme, err := normalizeOriginScheme(scheme)
if err != nil {
return "", err
}
normalizedAddress := normalizeOriginAddress(address)
if err := validateOriginAddress(normalizedAddress); err != nil {
return "", err
}
normalizedPort, err := normalizeOriginPort(port)
if err != nil {
return "", err
}
normalizedURI, err := normalizeOriginURI(uri)
if err != nil {
return "", err
}
parsed := &url.URL{
Scheme: normalizedScheme,
Host: formatOriginHost(normalizedAddress, normalizedPort),
}
if normalizedURI != "" {
if after, ok := strings.CutPrefix(normalizedURI, "?"); ok {
parsed.RawQuery = after
} else {
pathQuery := strings.SplitN(normalizedURI, "?", originURIPathQueryParts)
parsed.Path = pathQuery[0]
if len(pathQuery) > 1 {
parsed.RawQuery = pathQuery[1]
}
}
}
return parsed.String(), nil
}
func extractOriginAddress(rawURL string) (string, error) {
parsed, err := url.ParseRequestURI(strings.TrimSpace(rawURL))
if err != nil {
return "", fmt.Errorf("%s: %w", errProxyRouteOriginInvalid, err)
}
address := normalizeOriginAddress(parsed.Hostname())
if err := validateOriginAddress(address); err != nil {
return "", err
}
return address, nil
}
func getOrCreateOriginByAddress(ctx context.Context, address string) (*model.Origin, error) {
normalizedAddress := normalizeOriginAddress(address)
if err := validateOriginAddress(normalizedAddress); err != nil {
return nil, err
}
existing, err := repository.GetOriginByAddress(ctx, normalizedAddress)
if err == nil {
return existing, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
origin := &model.Origin{
Name: normalizedAddress,
Address: normalizedAddress,
Remark: "",
}
if err := repository.CreateOriginRecord(ctx, origin); err != nil {
if isUniqueConstraintError(err) {
return repository.GetOriginByAddress(ctx, normalizedAddress)
}
return nil, err
}
return origin, nil
}
func lookupTLSCertificateByID(ctx context.Context, id uint) (*model.TLSCertificate, error) {
return repository.GetTLSCertificateByID(ctx, id)
}
func lookupTunnelNodeByID(ctx context.Context, id uint) (*model.OpenFlareNode, error) {
return repository.GetOpenFlareNodeByID(ctx, id)
}
func lookupPagesProjectByID(ctx context.Context, id uint) (*model.PagesProject, error) {
return repository.GetPagesProjectByID(ctx, id)
}
func parseLeafCertificate(certPEM string) (*x509.Certificate, error) {
certPEMBlock, _ := pem.Decode([]byte(certPEM))
if certPEMBlock == nil {
return nil, errors.New(errProxyRouteCertNotFound)
}
leaf, err := x509.ParseCertificate(certPEMBlock.Bytes)
if err != nil {
return nil, err
}
return leaf, nil
}
func validateCertificateCoverage(certificate *model.TLSCertificate, domains []string) error {
if certificate == nil {
return errors.New(errProxyRouteCertNotFound)
}
leaf, err := parseLeafCertificate(certificate.CertPEM)
if err != nil {
return err
}
for _, domain := range domains {
if err := leaf.VerifyHostname(domain); err != nil {
return fmt.Errorf("certificate does not cover domain %s", domain)
}
}
return nil
}
func loadProxyRouteZoneDomains(ctx context.Context, ids []uint) ([]model.ZoneDomain, error) {
if len(ids) == 0 {
return nil, errors.New(errProxyRouteZoneDomainsRequired)
}
seen := make(map[uint]struct{}, len(ids))
for _, id := range ids {
if id == 0 {
return nil, errors.New(errProxyRouteZoneDomainNotFound)
}
if _, ok := seen[id]; ok {
return nil, errors.New(errProxyRouteZoneDomainDuplicate)
}
seen[id] = struct{}{}
}
domains, err := repository.ListZoneDomainsByIDs(ctx, ids)
if err != nil {
return nil, errors.New(errProxyRouteZoneDomainNotFound)
}
return domains, nil
}
func validateProxyRouteSiteName(siteName string) error {
if strings.TrimSpace(siteName) == "" {
return errors.New(errProxyRouteSiteNameEmpty)
}
return nil
}
func validateProxyRouteSiteNameUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string) error {
routes, err := repository.ListProxyRoutes(ctx)
if err != nil {
return err
}
currentID := uint(0)
if route != nil {
currentID = route.ID
}
for _, item := range routes {
if item == nil || item.ID == currentID {
continue
}
if item.SiteName == siteName {
return errors.New(errProxyRouteSiteNameExists)
}
}
return nil
}
func validateProxyRouteZoneDomainCertificates(ctx context.Context, domains []model.ZoneDomain, enableHTTPS bool) error {
if !enableHTTPS {
return nil
}
for _, domain := range domains {
if domain.CertID == nil || *domain.CertID == 0 {
return errors.New(errProxyRouteCertRequired)
}
certificate, err := lookupTLSCertificateByID(ctx, *domain.CertID)
if err != nil {
return errors.New(errProxyRouteCertNotFound)
}
if err := validateCertificateCoverage(certificate, []string{domain.Domain}); err != nil {
return err
}
}
return nil
}
func normalizeProxyRouteLimitConnValue(value int, field string) (int, error) {
if value < -1 {
return 0, fmt.Errorf("%s must be greater than or equal to -1", field)
}
return value, nil
}
func normalizeProxyRouteLimitRate(raw string) (string, error) {
normalized := strings.ToLower(strings.TrimSpace(raw))
if normalized == "" || normalized == "0" {
return "", nil
}
if normalized == "-1" {
return "-1", nil
}
if !proxyRouteLimitRatePattern.MatchString(normalized) {
return "", errors.New(errProxyRouteLimitRate)
}
if strings.TrimRight(normalized, "km") == "" {
return "", nil
}
return normalized, nil
}
func normalizeProxyRouteLimitReqPerIP(raw string) (string, error) {
normalized := strings.ToLower(strings.TrimSpace(raw))
if normalized == "" || normalized == "0" {
return "", nil
}
if normalized == "-1" {
return "-1", nil
}
if !proxyRouteLimitReqPattern.MatchString(normalized) {
return "", errors.New("请求频率格式不合法,请使用类似 10r/s、100r/m,或 -1 禁用")
}
return normalized, nil
}
func hasStructuredOriginInput(input Input) bool {
return (input.OriginID != nil && *input.OriginID != 0) ||
strings.TrimSpace(input.OriginScheme) != "" ||
strings.TrimSpace(input.OriginAddress) != "" ||
strings.TrimSpace(input.OriginPort) != "" ||
strings.TrimSpace(input.OriginURI) != ""
}
func normalizeCustomHeaders(headers []CustomHeaderInput) ([]CustomHeaderInput, error) {
if len(headers) == 0 {
return []CustomHeaderInput{}, nil
}
normalized := make([]CustomHeaderInput, 0, len(headers))
for _, header := range headers {
key := strings.TrimSpace(header.Key)
value := strings.TrimSpace(header.Value)
if key == "" && value == "" {
continue
}
if key == "" {
return nil, errors.New(errProxyRouteHeaderKeyEmpty)
}
if !proxyHeaderKeyPattern.MatchString(key) {
return nil, errors.New(errProxyRouteHeaderKeyInvalid)
}
if strings.ContainsAny(key, "\r\n") || strings.ContainsAny(value, "\r\n") {
return nil, errors.New(errProxyRouteHeaderNewline)
}
normalized = append(normalized, CustomHeaderInput{Key: key, Value: value})
}
return normalized, nil
}
func normalizeUpstreams(originURL string, upstreams []string) ([]string, error) {
candidates := make([]string, 0, len(upstreams)+1)
if strings.TrimSpace(originURL) != "" {
candidates = append(candidates, originURL)
}
candidates = append(candidates, upstreams...)
trimmed := make([]string, 0, len(candidates))
for _, candidate := range candidates {
item := strings.TrimSpace(candidate)
if item == "" {
continue
}
trimmed = append(trimmed, item)
}
unique := uniqueStrings(trimmed)
normalized := make([]string, 0, len(unique))
var scheme string
multiUpstream := len(unique) > 1
for _, item := range unique {
if err := validateOriginURL(item); err != nil {
return nil, err
}
parsed, err := url.ParseRequestURI(item)
if err != nil {
return nil, errors.New(errProxyRouteOriginInvalid)
}
if multiUpstream && parsed.Path != "" && parsed.Path != "/" {
return nil, errors.New(errProxyRouteUpstreamPath)
}
if multiUpstream && parsed.RawQuery != "" {
return nil, errors.New(errProxyRouteUpstreamQuery)
}
if scheme == "" {
scheme = parsed.Scheme
} else if scheme != parsed.Scheme {
return nil, errors.New(errProxyRouteUpstreamScheme)
}
normalized = append(normalized, item)
}
if len(normalized) == 0 {
return nil, errors.New(errProxyRouteUpstreamRequired)
}
return normalized, nil
}
func decodeStoredCustomHeaders(raw string) ([]CustomHeaderInput, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []CustomHeaderInput{}, nil
}
var headers []CustomHeaderInput
if err := json.Unmarshal([]byte(text), &headers); err != nil {
return nil, errors.New("custom_headers payload is invalid")
}
return normalizeCustomHeaders(headers)
}
// normalizeCachePolicy stores API write values.
// When enabling with empty/url policy, keep legacy "all" semantics so old rows
// and clients that omit policy do not silently narrow cache to static extensions.
// New UI should send policy=static explicitly when choosing the recommended default.
func normalizeCachePolicy(enabled bool, raw string) string {
if !enabled {
return ""
}
policy := strings.TrimSpace(strings.ToLower(raw))
switch policy {
case "", proxyRouteCachePolicyURL, proxyRouteCachePolicyAll:
return proxyRouteCachePolicyAll
case proxyRouteCachePolicyStatic:
return proxyRouteCachePolicyStatic
case proxyRouteCachePolicySuffix, proxyRouteCachePolicyPathPrefix, proxyRouteCachePolicyPathExact:
return policy
default:
return policy
}
}
// displayCachePolicy normalizes values for API list/get (and UI).
func displayCachePolicy(enabled bool, raw string) string {
if !enabled {
return ""
}
return normalizeCachePolicy(true, raw)
}
func normalizeCacheRules(enabled bool, rawPolicy string, rules []string) ([]string, error) {
if !enabled {
return []string{}, nil
}
policy := normalizeCachePolicy(enabled, rawPolicy)
switch policy {
case proxyRouteCachePolicyStatic, proxyRouteCachePolicyAll, proxyRouteCachePolicyURL:
return []string{}, nil
case proxyRouteCachePolicySuffix:
return normalizeCacheSuffixRules(rules)
case proxyRouteCachePolicyPathPrefix:
return normalizeCachePathRules(rules, true)
case proxyRouteCachePolicyPathExact:
return normalizeCachePathRules(rules, false)
default:
return nil, errors.New(errProxyRouteCachePolicy)
}
}
func normalizeCacheSuffixRules(rules []string) ([]string, error) {
normalized := make([]string, 0, len(rules))
seen := make(map[string]struct{}, len(rules))
for _, rule := range rules {
item := strings.TrimSpace(strings.TrimPrefix(rule, "."))
if item == "" {
continue
}
if strings.ContainsAny(item, "/\\ \t\r\n") {
return nil, errors.New(errProxyRouteCacheSuffix)
}
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
normalized = append(normalized, item)
}
if len(normalized) == 0 {
return nil, errors.New(errProxyRouteCacheSuffixReq)
}
return normalized, nil
}
func normalizeCachePathRules(rules []string, allowPrefix bool) ([]string, error) {
normalized := make([]string, 0, len(rules))
seen := make(map[string]struct{}, len(rules))
for _, rule := range rules {
item := strings.TrimSpace(rule)
if item == "" {
continue
}
if !strings.HasPrefix(item, "/") || strings.Contains(item, "://") || strings.ContainsAny(item, " \t\r\n") {
return nil, errors.New(errProxyRouteCachePath)
}
if !allowPrefix && strings.HasSuffix(item, "/") && len(item) > 1 {
item = strings.TrimRight(item, "/")
}
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
normalized = append(normalized, item)
}
if len(normalized) == 0 {
if allowPrefix {
return nil, errors.New(errProxyRouteCachePrefixReq)
}
return nil, errors.New(errProxyRouteCacheExactReq)
}
return normalized, nil
}
func decodeStoredCacheRules(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []string{}, nil
}
var rules []string
if err := json.Unmarshal([]byte(text), &rules); err != nil {
return nil, errors.New("cache_rules payload is invalid")
}
normalized := make([]string, 0, len(rules))
for _, rule := range rules {
item := strings.TrimSpace(rule)
if item == "" {
continue
}
normalized = append(normalized, item)
}
return normalized, nil
}
func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return normalizeUpstreams(fallbackOriginURL, nil)
}
var upstreams []string
if err := json.Unmarshal([]byte(text), &upstreams); err != nil {
return nil, errors.New("upstreams payload is invalid")
}
return normalizeUpstreams(fallbackOriginURL, upstreams)
}
func validateOriginURL(raw string) error {
if raw == "" {
return errors.New(errProxyRouteOriginEmpty)
}
parsed, err := url.ParseRequestURI(raw)
if err != nil {
return errors.New(errProxyRouteOriginInvalid)
}
if parsed.Scheme != proxyRouteSchemeHTTP && parsed.Scheme != proxyRouteSchemeHTTPS {
return errors.New(errProxyRouteOriginScheme)
}
if parsed.Host == "" {
return errors.New(errProxyRouteOriginInvalid)
}
return nil
}
func validateOriginHost(raw string) error {
if raw == "" {
return nil
}
if strings.ContainsAny(raw, "/\\ \t\r\n") || strings.Contains(raw, "://") {
return errors.New(errProxyRouteOriginHostInvalid)
}
parsed, err := url.Parse("//" + raw)
if err != nil || parsed.Host == "" || parsed.Host != raw {
return errors.New(errProxyRouteOriginHostInvalid)
}
if parsed.Hostname() == "" {
return errors.New(errProxyRouteOriginHostInvalid)
}
return nil
}
func normalizeTunnelNodeID(tunnelNodeID, legacyTunnelID *uint) (*uint, error) {
if tunnelNodeID != nil && *tunnelNodeID != 0 {
return tunnelNodeID, nil
}
if legacyTunnelID != nil && *legacyTunnelID != 0 {
return legacyTunnelID, nil
}
return nil, errors.New(errProxyRouteTunnelNodeReq)
}
func validateTunnelRouteInput(ctx context.Context, tunnelNodeID *uint, targetAddr, targetProtocol string) error {
if tunnelNodeID == nil || *tunnelNodeID == 0 {
return errors.New(errProxyRouteTunnelNodeReq)
}
tunnelNode, err := lookupTunnelNodeByID(ctx, *tunnelNodeID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errProxyRouteTunnelNodeMissing)
}
return err
}
if tunnelNode.NodeType != "tunnel_client" {
return errors.New(errProxyRouteTunnelNodeType)
}
if strings.TrimSpace(targetAddr) == "" {
return errors.New(errProxyRouteTunnelAddrReq)
}
switch strings.ToLower(strings.TrimSpace(targetProtocol)) {
case "", proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
return nil
default:
return errors.New(errProxyRouteTunnelProtocol)
}
}
func validatePagesRouteInput(ctx context.Context, projectID *uint) error {
if projectID == nil || *projectID == 0 {
return errors.New(errProxyRoutePagesProjectReq)
}
project, err := lookupPagesProjectByID(ctx, *projectID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errProxyRoutePagesNotFound)
}
return err
}
if !project.Enabled {
return errors.New(errProxyRoutePagesDisabled)
}
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 {
return errors.New(errProxyRoutePagesNoDeploy)
}
return nil
}
func normalizeUpstreamType(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case proxyRouteUpstreamTypeTunnel:
return proxyRouteUpstreamTypeTunnel
case proxyRouteUpstreamTypePages:
return proxyRouteUpstreamTypePages
default:
return "direct"
}
}
func normalizeTunnelTargetProtocol(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case proxyRouteSchemeHTTPS:
return proxyRouteSchemeHTTPS
default:
return proxyRouteSchemeHTTP
}
}
@@ -0,0 +1,29 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import "testing"
func TestNormalizeProxyRouteLimitConnValue(t *testing.T) {
t.Parallel()
got, err := normalizeProxyRouteLimitConnValue(-1, "limit_conn_per_server")
if err != nil || got != -1 {
t.Fatalf("want -1, got %d err %v", got, err)
}
if _, err := normalizeProxyRouteLimitConnValue(-2, "limit_conn_per_server"); err == nil {
t.Fatal("expected error for -2")
}
}
func TestNormalizeProxyRouteLimitRate(t *testing.T) {
t.Parallel()
got, err := normalizeProxyRouteLimitRate("-1")
if err != nil || got != "-1" {
t.Fatalf("want -1, got %q err %v", got, err)
}
got, err = normalizeProxyRouteLimitRate("0")
if err != nil || got != "" {
t.Fatalf("want empty inherit, got %q err %v", got, err)
}
}
@@ -0,0 +1,408 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"errors"
"slices"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"gorm.io/gorm"
)
// CustomHeaderInput 自定义响应头。
type CustomHeaderInput struct {
Key string `json:"key"`
Value string `json:"value"`
}
// Input 代理规则创建/更新请求。
type Input struct {
SiteName string `json:"site_name"`
ZoneDomainIDs []uint `json:"zone_domain_ids"`
OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"`
OriginScheme string `json:"origin_scheme"`
OriginAddress string `json:"origin_address"`
OriginPort string `json:"origin_port"`
OriginURI string `json:"origin_uri"`
OriginHost string `json:"origin_host"`
Upstreams []string `json:"upstreams"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server"`
LimitConnPerIP int `json:"limit_conn_per_ip"`
LimitRate string `json:"limit_rate"`
LimitReqPerIP string `json:"limit_req_per_ip"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy"`
CacheRules []string `json:"cache_rules"`
CustomHeaders []CustomHeaderInput `json:"custom_headers"`
BasicAuthEnabled bool `json:"basic_auth_enabled"`
BasicAuthUsername string `json:"basic_auth_username"`
BasicAuthPassword string `json:"basic_auth_password"`
UpstreamType string `json:"upstream_type"`
TunnelNodeID *uint `json:"tunnel_node_id"`
TunnelID *uint `json:"tunnel_id"`
TunnelTargetAddr string `json:"tunnel_target_addr"`
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
PagesProjectID *uint `json:"pages_project_id"`
}
// View 代理规则视图。
type View struct {
ID uint `json:"id"`
SiteName string `json:"site_name"`
ZoneDomainIDs []uint `json:"zone_domain_ids"`
ZoneDomains []ZoneDomainView `json:"zone_domains"`
OriginID *uint `json:"origin_id"`
OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host"`
Upstreams string `json:"upstreams"`
UpstreamList []string `json:"upstream_list"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server"`
LimitConnPerIP int `json:"limit_conn_per_ip"`
LimitRate string `json:"limit_rate"`
LimitReqPerIP string `json:"limit_req_per_ip"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy"`
CacheRules string `json:"cache_rules"`
CacheRuleList []string `json:"cache_rule_list"`
CustomHeaders string `json:"custom_headers"`
CustomHeaderList []CustomHeaderInput `json:"custom_header_list"`
BasicAuthEnabled bool `json:"basic_auth_enabled"`
BasicAuthUsername string `json:"basic_auth_username"`
BasicAuthPassword string `json:"basic_auth_password"`
UpstreamType string `json:"upstream_type"`
TunnelNodeID *uint `json:"tunnel_node_id"`
TunnelID *uint `json:"tunnel_id"`
TunnelTargetAddr string `json:"tunnel_target_addr"`
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
PagesProjectID *uint `json:"pages_project_id"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// ZoneDomainView is the route-safe representation of a bound Zone domain.
type ZoneDomainView struct {
ID uint `json:"id"`
ZoneID uint `json:"zone_id"`
Domain string `json:"domain"`
CertID *uint `json:"cert_id"`
}
// ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*View, error) {
routes, err := repository.ListProxyRoutes(ctx)
if err != nil {
return nil, err
}
return buildProxyRouteViews(ctx, routes)
}
// GetProxyRoute 获取代理规则详情。
func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
route, err := repository.GetProxyRouteByID(ctx, id)
if err != nil {
return nil, err
}
return buildProxyRouteView(ctx, route)
}
// CreateProxyRoute 创建代理规则。
func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
route, err := buildProxyRoute(ctx, nil, input)
if err != nil {
return nil, err
}
if err = repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error {
if err := lockPagesProjectsForRouteMutation(tx, 0, route); err != nil {
return err
}
if err := repository.CreateProxyRouteRecordTx(tx, route); err != nil {
return err
}
return repository.ReplaceZoneDomainRouteBindingsTx(tx, route.ID, input.ZoneDomainIDs)
}); err != nil {
if mapped := mapProxyRoutePersistError(err); mapped != nil {
return nil, mapped
}
return nil, err
}
return buildProxyRouteView(ctx, route)
}
// UpdateProxyRoute 更新代理规则。
func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error) {
route, err := repository.GetProxyRouteByID(ctx, id)
if err != nil {
return nil, err
}
previousPagesProjectID := pagesProjectIDForRoute(route)
route, err = buildProxyRoute(ctx, route, input)
if err != nil {
return nil, err
}
if err = repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error {
if err := lockPagesProjectsForRouteMutation(tx, previousPagesProjectID, route); err != nil {
return err
}
if err := repository.UpdateProxyRouteRecordTx(tx, route); err != nil {
return err
}
return repository.ReplaceZoneDomainRouteBindingsTx(tx, route.ID, input.ZoneDomainIDs)
}); err != nil {
if mapped := mapProxyRoutePersistError(err); mapped != nil {
return nil, mapped
}
return nil, err
}
return buildProxyRouteView(ctx, route)
}
func mapProxyRoutePersistError(err error) error {
if err == nil {
return nil
}
if isUniqueConstraintError(err) {
return errors.New(errProxyRouteIdentityExists)
}
if errors.Is(err, repository.ErrZoneDomainBoundToAnotherRoute) {
return errors.New(errProxyRouteZoneDomainBound)
}
if errors.Is(err, repository.ErrZoneDomainNotFound) {
return errors.New(errProxyRouteZoneDomainNotFound)
}
return nil
}
func pagesProjectIDForRoute(route *model.ProxyRoute) uint {
if route == nil || route.UpstreamType != proxyRouteUpstreamTypePages || route.PagesProjectID == nil {
return 0
}
return *route.PagesProjectID
}
func lockPagesProjectsForRouteMutation(tx *gorm.DB, previousProjectID uint, route *model.ProxyRoute) error {
nextProjectID := pagesProjectIDForRoute(route)
var projectIDs []uint
if previousProjectID != 0 {
projectIDs = append(projectIDs, previousProjectID)
}
if nextProjectID != 0 && nextProjectID != previousProjectID {
projectIDs = append(projectIDs, nextProjectID)
}
slices.Sort(projectIDs)
for _, projectID := range projectIDs {
project, err := repository.LockPagesProjectByIDTx(tx, projectID)
if errors.Is(err, gorm.ErrRecordNotFound) && projectID != nextProjectID {
continue
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errProxyRoutePagesNotFound)
}
if err != nil {
return err
}
if projectID == nextProjectID {
if err := validateLockedPagesRouteProject(project); err != nil {
return err
}
}
}
return nil
}
func validateLockedPagesRouteProject(project *model.PagesProject) error {
if project == nil {
return errors.New(errProxyRoutePagesNotFound)
}
if !project.Enabled {
return errors.New(errProxyRoutePagesDisabled)
}
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 {
return errors.New(errProxyRoutePagesNoDeploy)
}
return nil
}
// DeleteProxyRoute 删除代理规则。
func DeleteProxyRoute(ctx context.Context, id uint) error {
if _, err := repository.GetProxyRouteByID(ctx, id); err != nil {
return err
}
return repository.DeleteProxyRouteAndUnbind(ctx, id)
}
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, error) {
domains, err := loadProxyRouteZoneDomains(ctx, input.ZoneDomainIDs)
if err != nil {
return nil, err
}
siteName := strings.TrimSpace(input.SiteName)
upstreamType := normalizeUpstreamType(input.UpstreamType)
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
if err != nil {
return nil, err
}
originHost := strings.TrimSpace(input.OriginHost)
cachePolicy := strings.TrimSpace(input.CachePolicy)
cacheRules, err := normalizeCacheRules(input.CacheEnabled, cachePolicy, input.CacheRules)
if err != nil {
return nil, err
}
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
if err != nil {
return nil, err
}
limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server")
if err != nil {
return nil, err
}
limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip")
if err != nil {
return nil, err
}
limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate)
if err != nil {
return nil, err
}
limitReqPerIP, err := normalizeProxyRouteLimitReqPerIP(input.LimitReqPerIP)
if err != nil {
return nil, err
}
if err := validateProxyRouteZoneDomainCertificates(ctx, domains, input.EnableHTTPS); err != nil {
return nil, err
}
jsonFields, err := marshalProxyRouteJSONFields(upstreams, cacheRules, customHeaders)
if err != nil {
return nil, err
}
if err := validateProxyRouteSiteName(siteName); err != nil {
return nil, err
}
if err := validateProxyRouteSiteNameUniqueness(ctx, route, siteName); err != nil {
return nil, err
}
if err := validateOriginHost(originHost); err != nil {
return nil, err
}
if input.RedirectHTTP && !input.EnableHTTPS {
return nil, errors.New(errProxyRouteRedirectHTTP)
}
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
return nil, err
}
if route == nil {
route = &model.ProxyRoute{}
}
populateProxyRouteFields(
route,
input,
siteName,
jsonFields,
originID,
upstreams,
originHost,
cachePolicy,
limitConnPerServer,
limitConnPerIP,
limitRate,
limitReqPerIP,
upstreamType,
)
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
return nil, err
}
return route, nil
}
func buildProxyRouteViews(ctx context.Context, routes []*model.ProxyRoute) ([]*View, error) {
views := make([]*View, 0, len(routes))
for _, route := range routes {
view, err := buildProxyRouteView(ctx, route)
if err != nil {
return nil, err
}
views = append(views, view)
}
return views, nil
}
func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, error) {
if route == nil {
return nil, errors.New("proxy route is nil")
}
domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil {
return nil, err
}
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
if err != nil {
return nil, err
}
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
if err != nil {
return nil, err
}
customHeaders, err := decodeStoredCustomHeaders(route.CustomHeaders)
if err != nil {
return nil, err
}
zoneDomainIDs := make([]uint, 0, len(domains))
zoneDomains := make([]ZoneDomainView, 0, len(domains))
for _, domain := range domains {
zoneDomainIDs = append(zoneDomainIDs, domain.ID)
zoneDomains = append(zoneDomains, ZoneDomainView{ID: domain.ID, ZoneID: domain.ZoneID, Domain: domain.Domain, CertID: domain.CertID})
}
return &View{
ID: route.ID,
SiteName: route.SiteName,
ZoneDomainIDs: zoneDomainIDs,
ZoneDomains: zoneDomains,
OriginID: route.OriginID,
OriginURL: route.OriginURL,
OriginHost: route.OriginHost,
Upstreams: route.Upstreams,
UpstreamList: upstreams,
Enabled: route.Enabled,
EnableHTTPS: route.EnableHTTPS,
RedirectHTTP: route.RedirectHTTP,
LimitConnPerServer: route.LimitConnPerServer,
LimitConnPerIP: route.LimitConnPerIP,
LimitRate: route.LimitRate,
LimitReqPerIP: route.LimitReqPerIP,
CacheEnabled: route.CacheEnabled,
CachePolicy: displayCachePolicy(route.CacheEnabled, route.CachePolicy),
CacheRules: route.CacheRules,
CacheRuleList: cacheRules,
CustomHeaders: route.CustomHeaders,
CustomHeaderList: customHeaders,
BasicAuthEnabled: route.BasicAuthEnabled,
BasicAuthUsername: route.BasicAuthUsername,
BasicAuthPassword: route.BasicAuthPassword,
UpstreamType: route.UpstreamType,
TunnelNodeID: route.TunnelNodeID,
TunnelID: route.TunnelNodeID,
TunnelTargetAddr: route.TunnelTargetAddr,
TunnelTargetProtocol: route.TunnelTargetProtocol,
PagesProjectID: route.PagesProjectID,
CreatedAt: route.CreatedAt,
UpdatedAt: route.UpdatedAt,
}, nil
}
@@ -0,0 +1,163 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupProxyRouteTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.ProxyRoute{},
&model.Origin{},
&model.Zone{},
&model.ZoneDomain{},
&model.TLSCertificate{},
&model.PagesProject{},
))
db.SetDB(sqliteDB)
return func() { db.SetDB(nil) }
}
func createZoneDomain(t *testing.T, ctx context.Context, domain string, certID *uint) *model.ZoneDomain {
t.Helper()
zone := &model.Zone{Domain: "example.com"}
var existing model.Zone
if err := db.DB(ctx).Where("domain = ?", zone.Domain).First(&existing).Error; err == nil {
zone = &existing
} else {
require.NoError(t, db.DB(ctx).Create(zone).Error)
}
item := &model.ZoneDomain{ZoneID: zone.ID, Domain: domain, CertID: certID}
require.NoError(t, db.DB(ctx).Create(item).Error)
return item
}
func TestCreateProxyRouteBindsZoneDomains(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domainA := createZoneDomain(t, ctx, "api.example.com", nil)
domainB := createZoneDomain(t, ctx, "www.example.com", nil)
view, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domainA.ID, domainB.ID}, OriginURL: "http://origin.example.com:8080", Enabled: true})
require.NoError(t, err)
assert.Equal(t, []uint{domainA.ID, domainB.ID}, view.ZoneDomainIDs)
require.Len(t, view.ZoneDomains, 2)
assert.Equal(t, "api.example.com", view.ZoneDomains[0].Domain)
}
func TestCreateProxyRouteRejectsInvalidZoneDomainBindings(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domain := createZoneDomain(t, ctx, "api.example.com", nil)
base := Input{SiteName: "api", OriginURL: "http://origin.example.com:8080"}
_, err := CreateProxyRoute(ctx, base)
require.EqualError(t, err, errProxyRouteZoneDomainsRequired)
base.ZoneDomainIDs = []uint{domain.ID, domain.ID}
_, err = CreateProxyRoute(ctx, base)
require.EqualError(t, err, errProxyRouteZoneDomainDuplicate)
first, err := CreateProxyRoute(ctx, Input{SiteName: "first", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://origin.example.com:8080"})
require.NoError(t, err)
_, err = CreateProxyRoute(ctx, Input{SiteName: "second", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://other.example.com:8080"})
require.Error(t, err)
require.NoError(t, DeleteProxyRoute(ctx, first.ID))
}
func TestCreateProxyRouteHTTPSRequiresCoveringCertificate(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domain := createZoneDomain(t, ctx, "api.example.com", nil)
_, err := CreateProxyRoute(ctx, Input{SiteName: "api", ZoneDomainIDs: []uint{domain.ID}, OriginURL: "http://origin.example.com:8080", EnableHTTPS: true})
require.EqualError(t, err, errProxyRouteCertRequired)
}
func TestPagesRouteLocksAndRevalidatesTargetProject(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domain := createZoneDomain(t, ctx, "pages.example.com", nil)
activeDeploymentID := uint(99)
project := &model.PagesProject{
Name: "Pages Site",
Slug: "pages-site",
Enabled: true,
ActiveDeploymentID: &activeDeploymentID,
}
require.NoError(t, db.DB(ctx).Create(project).Error)
view, err := CreateProxyRoute(ctx, Input{
SiteName: "pages",
ZoneDomainIDs: []uint{domain.ID},
UpstreamType: proxyRouteUpstreamTypePages,
PagesProjectID: &project.ID,
Enabled: true,
})
require.NoError(t, err)
require.NotNil(t, view.PagesProjectID)
assert.Equal(t, project.ID, *view.PagesProjectID)
require.NoError(t, db.DB(ctx).Delete(&model.PagesProject{}, project.ID).Error)
route := &model.ProxyRoute{UpstreamType: proxyRouteUpstreamTypePages, PagesProjectID: &project.ID}
err = repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error {
return lockPagesProjectsForRouteMutation(tx, 0, route)
})
require.EqualError(t, err, errProxyRoutePagesNotFound)
}
func TestRouteCanMoveAwayFromAlreadyMissingPagesProject(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
missingProjectID := uint(404)
route := &model.ProxyRoute{UpstreamType: "direct"}
err := repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error {
return lockPagesProjectsForRouteMutation(tx, missingProjectID, route)
})
require.NoError(t, err)
}
func TestNormalizeCachePolicyDefaultsAndLegacy(t *testing.T) {
assert.Empty(t, normalizeCachePolicy(false, "static"))
// Empty/url on write = legacy all (compat); UI sends static explicitly for new default.
assert.Equal(t, proxyRouteCachePolicyAll, normalizeCachePolicy(true, ""))
assert.Equal(t, proxyRouteCachePolicyStatic, normalizeCachePolicy(true, "static"))
assert.Equal(t, proxyRouteCachePolicyAll, normalizeCachePolicy(true, "url"))
assert.Equal(t, proxyRouteCachePolicyAll, normalizeCachePolicy(true, "all"))
assert.Equal(t, proxyRouteCachePolicySuffix, normalizeCachePolicy(true, "suffix"))
assert.Empty(t, displayCachePolicy(false, "all"))
assert.Equal(t, proxyRouteCachePolicyAll, displayCachePolicy(true, ""))
assert.Equal(t, proxyRouteCachePolicyAll, displayCachePolicy(true, "url"))
assert.Equal(t, proxyRouteCachePolicyStatic, displayCachePolicy(true, "static"))
rules, err := normalizeCacheRules(true, "url", []string{"css"})
require.NoError(t, err)
assert.Empty(t, rules)
rules, err = normalizeCacheRules(true, "static", nil)
require.NoError(t, err)
assert.Empty(t, rules)
_, err = normalizeCacheRules(true, "suffix", nil)
require.Error(t, err)
}
@@ -0,0 +1,86 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"errors"
"strings"
"Wavelet/openflare/plugins/server/kernel/repository"
"gorm.io/gorm"
)
func resolveStructuredOriginInput(ctx context.Context, input Input) (string, *uint, error) {
scheme, err := normalizeOriginScheme(input.OriginScheme)
if err != nil {
return "", nil, err
}
port, err := normalizeOriginPort(input.OriginPort)
if err != nil {
return "", nil, err
}
uri, err := normalizeOriginURI(input.OriginURI)
if err != nil {
return "", nil, err
}
if input.OriginID != nil && *input.OriginID != 0 {
return resolveOriginByID(ctx, scheme, port, uri, *input.OriginID)
}
return resolveOriginByAddress(ctx, scheme, port, uri, input.OriginAddress)
}
func resolveOriginByID(ctx context.Context, scheme, port, uri string, originID uint) (string, *uint, error) {
origin, err := repository.GetOriginByID(ctx, originID)
if err != nil {
return "", nil, errors.New(errProxyRouteOriginNotFound)
}
originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri)
if err != nil {
return "", nil, err
}
return originURL, &origin.ID, nil
}
func resolveOriginByAddress(ctx context.Context, scheme, port, uri, rawAddress string) (string, *uint, error) {
address := normalizeOriginAddress(rawAddress)
if err := validateOriginAddress(address); err != nil {
return "", nil, err
}
originURL, err := buildOriginURLFromParts(scheme, address, port, uri)
if err != nil {
return "", nil, err
}
origin, err := getOrCreateOriginByAddress(ctx, address)
if err != nil {
return "", nil, err
}
return originURL, &origin.ID, nil
}
func resolveLegacyOriginInput(ctx context.Context, originURL string) (string, *uint, error) {
if originURL == "" {
return "", nil, errors.New(errProxyRouteOriginEmpty)
}
address, err := extractOriginAddress(originURL)
if err != nil {
return "", nil, err
}
origin, findErr := repository.GetOriginByAddress(ctx, address)
if findErr == nil {
return originURL, &origin.ID, nil
}
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
return "", nil, findErr
}
return originURL, nil, nil
}
func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) {
if hasStructuredOriginInput(input) {
return resolveStructuredOriginInput(ctx, input)
}
return resolveLegacyOriginInput(ctx, strings.TrimSpace(input.OriginURL))
}
@@ -0,0 +1,141 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
func handleLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
return apiutil.AbortNotFoundIfMissing(c, err, errProxyRouteNotFound)
}
// GetProxyRoutes 列出全部代理规则。
// @Summary 获取代理规则列表
// @Description 返回所有代理规则配置,需要管理员权限
// @Tags openflare-proxy-route
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]proxy_route.View} "代理规则列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/proxy-routes [get]
func GetProxyRoutes(c *gin.Context) {
routes, err := ListProxyRoutes(c.Request.Context())
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(routes))
}
// GetProxyRouteHandler 获取代理规则详情。
// @Summary 获取代理规则详情
// @Description 返回指定代理规则的完整配置,需要管理员权限
// @Tags openflare-proxy-route
// @Produce json
// @Security SessionCookie
// @Param id path int true "代理规则 ID"
// @Success 200 {object} response.Any{data=proxy_route.View} "代理规则详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或规则不存在"
// @Router /api/v1/d/proxy-routes/{id} [get]
func GetProxyRouteHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
route, err := GetProxyRoute(c.Request.Context(), id)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(route))
}
// CreateProxyRouteHandler 创建代理规则。
// @Summary 创建代理规则
// @Description 创建新的反向代理规则,需要管理员权限
// @Tags openflare-proxy-route
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body proxy_route.Input true "代理规则参数"
// @Success 200 {object} response.Any{data=proxy_route.View} "创建成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/proxy-routes [post]
func CreateProxyRouteHandler(c *gin.Context) {
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
route, err := CreateProxyRoute(c.Request.Context(), input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(route))
}
// UpdateProxyRouteHandler 更新代理规则。
// @Summary 更新代理规则
// @Description 更新指定代理规则的配置,需要管理员权限
// @Tags openflare-proxy-route
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "代理规则 ID"
// @Param body body proxy_route.Input true "代理规则参数"
// @Success 200 {object} response.Any{data=proxy_route.View} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或规则不存在"
// @Router /api/v1/d/proxy-routes/{id}/update [post]
func UpdateProxyRouteHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
route, err := UpdateProxyRoute(c.Request.Context(), id, input)
if handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(route))
}
// DeleteProxyRouteHandler 删除代理规则。
// @Summary 删除代理规则
// @Description 删除指定代理规则,需要管理员权限
// @Tags openflare-proxy-route
// @Produce json
// @Security SessionCookie
// @Param id path int true "代理规则 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或规则不存在"
// @Router /api/v1/d/proxy-routes/{id}/delete [post]
func DeleteProxyRouteHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if err := DeleteProxyRoute(c.Request.Context(), id); handleLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,51 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package routeidentity resolves proxy route site names and normalized domains
// for OpenFlare control-plane and edge rendering.
package routeidentity
import (
"encoding/json"
"errors"
"fmt"
"strings"
)
// NormalizeDomains lowercases, deduplicates, and validates proxy route domains.
func NormalizeDomains(rawDomains []string) ([]string, error) {
normalized := make([]string, 0, len(rawDomains))
seen := make(map[string]struct{}, len(rawDomains))
for _, rawDomain := range rawDomains {
domain := strings.ToLower(strings.TrimSpace(rawDomain))
if domain == "" {
continue
}
if strings.Contains(domain, "://") || strings.Contains(domain, "/") {
return nil, fmt.Errorf("domain %q is invalid", rawDomain)
}
if _, ok := seen[domain]; ok {
continue
}
seen[domain] = struct{}{}
normalized = append(normalized, domain)
}
if len(normalized) == 0 {
return nil, errors.New("domain is required")
}
return normalized, nil
}
// DecodeDomains parses legacy route domain fields for the goose upgrade importer.
// Runtime consumers must read ZoneDomain bindings instead.
func DecodeDomains(raw string, fallbackDomain string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return NormalizeDomains([]string{fallbackDomain})
}
var domains []string
if err := json.Unmarshal([]byte(text), &domains); err != nil {
return nil, errors.New("domains payload is invalid")
}
return NormalizeDomains(domains)
}
@@ -0,0 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package routeidentity
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestDecodeDomainsNormalizesCaseAndOrder(t *testing.T) {
domains, err := DecodeDomains(`["WWW.Example.COM","example.com"]`, "fallback.example.com")
require.NoError(t, err)
assert.Equal(t, []string{"www.example.com", "example.com"}, domains)
}
@@ -0,0 +1,19 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package zone
const (
errZoneDomainRequired = "域名不能为空"
errZoneRootInvalid = "zone 必须是有效的注册根域"
errDomainInvalid = "域名格式不合法"
errDomainWildcardUnsupported = "不支持通配符域名"
errDomainOutsideZone = "域名不属于该 Zone"
errZoneNotFound = "Zone 不存在"
errDomainNotFound = "域名不存在"
errDomainExists = "域名已存在"
errCertificateNotFound = "所选证书不存在"
errDomainBoundToRoute = "域名已绑定反代路由,请先解除绑定"
errZoneHasDomains = "根域下仍有域名,请先删除全部域名"
errStatsRangeInvalid = "时间范围无效,请选择 24h、7d 或 30d"
)
@@ -0,0 +1,355 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package zone
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"Wavelet/openflare/plugins/server/domain/site/routeidentity"
)
// ImportReport describes the idempotent legacy migration result.
type ImportReport struct {
Zones int `json:"zones"`
Domains int `json:"domains"`
Conflicts []string `json:"conflicts,omitempty"`
}
// LogAndReturn decorates the import failure with its conflict count.
func (r ImportReport) LogAndReturn(err error) error {
if err != nil {
return fmt.Errorf("迁移 Zone 失败(%d 个冲突): %w", len(r.Conflicts), err)
}
return nil
}
type legacyDomain struct {
Domain string
CertID *uint
ProxyRouteID *uint
}
// ImportLegacyTx imports legacy proxy-route / managed-domain rows into Zone tables
// within an existing SQL transaction (goose runs this on Server upgrade).
// postgres selects $n placeholders; otherwise SQLite-style ? is used.
// Missing legacy columns or tables are skipped so re-runs after phase-2 cleanup are no-ops.
//
//nolint:cyclop,gocyclo // single-pass legacy importer validates every source before write.
func ImportLegacyTx(ctx context.Context, tx *sql.Tx, postgres bool) (report ImportReport, err error) {
if tx == nil {
return report, errors.New("transaction is required")
}
q := func(sqlText string) string { return rebindSQL(sqlText, postgres) }
items := make([]legacyDomain, 0)
hasRouteDomains := false
hasDomainCol, err := hasTableColumn(ctx, tx, q, postgres, "of_proxy_routes", "domain")
if err != nil {
return report, err
}
hasDomainsCol, err := hasTableColumn(ctx, tx, q, postgres, "of_proxy_routes", "domains")
if err != nil {
return report, err
}
if hasDomainCol && hasDomainsCol {
var collectErr error
items, hasRouteDomains, report.Conflicts, collectErr = collectLegacyRouteDomainsImpl(ctx, tx, q)
if collectErr != nil {
return report, collectErr
}
}
if !hasRouteDomains {
exists, tableErr := hasTable(ctx, tx, q, postgres, "of_managed_domains")
if tableErr != nil {
return report, tableErr
}
if exists {
managed, managedErr := collectLegacyManagedDomains(ctx, tx, q)
if managedErr != nil {
return report, managedErr
}
items = append(items, managed...)
}
}
for _, item := range items {
domain, normErr := normalizeDomain(item.Domain)
if normErr != nil {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %v", item.Domain, normErr))
continue
}
root, rootErr := zoneRoot(domain)
if rootErr != nil {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %v", domain, rootErr))
continue
}
var existingID uint
var existingZoneDomain string
scanErr := tx.QueryRowContext(ctx, q(`
SELECT zd.id, z.domain
FROM of_zone_domains zd
JOIN of_zones z ON z.id = zd.zone_id
WHERE zd.domain = ?
`), domain).Scan(&existingID, &existingZoneDomain)
if scanErr == nil {
if existingZoneDomain != root {
report.Conflicts = append(report.Conflicts, domain+": global domain conflict")
} else if item.ProxyRouteID != nil {
if _, bindErr := tx.ExecContext(ctx, q(`
UPDATE of_zone_domains
SET proxy_route_id = COALESCE(proxy_route_id, ?),
cert_id = COALESCE(cert_id, ?)
WHERE id = ?
`), *item.ProxyRouteID, nullableUint(item.CertID), existingID); bindErr != nil {
return report, bindErr
}
}
continue
}
if !errors.Is(scanErr, sql.ErrNoRows) {
return report, scanErr
}
zoneID, zoneErr := ensureZone(ctx, tx, q, root, &report)
if zoneErr != nil {
return report, zoneErr
}
if item.CertID != nil {
var certID uint
if certErr := tx.QueryRowContext(ctx, q(`SELECT id FROM of_tls_certificates WHERE id = ?`), *item.CertID).
Scan(&certID); certErr != nil {
if errors.Is(certErr, sql.ErrNoRows) {
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: %s", domain, errCertificateNotFound))
continue
}
return report, certErr
}
}
if _, insErr := tx.ExecContext(ctx, q(`
INSERT INTO of_zone_domains (zone_id, proxy_route_id, domain, cert_id, created_at, updated_at)
VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
`), zoneID, nullableUint(item.ProxyRouteID), domain, nullableUint(item.CertID)); insErr != nil {
return report, insErr
}
report.Domains++
}
if len(report.Conflicts) > 0 {
return report, errors.New("legacy data has conflicts")
}
return report, nil
}
func ensureZone(
ctx context.Context,
tx *sql.Tx,
q func(string) string,
root string,
report *ImportReport,
) (uint, error) {
var zoneID uint
err := tx.QueryRowContext(ctx, q(`SELECT id FROM of_zones WHERE domain = ?`), root).Scan(&zoneID)
if err == nil {
return zoneID, nil
}
if !errors.Is(err, sql.ErrNoRows) {
return 0, err
}
if _, execErr := tx.ExecContext(ctx, q(`
INSERT INTO of_zones (domain, created_at, updated_at)
VALUES (?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
`), root); execErr != nil {
return 0, execErr
}
if err := tx.QueryRowContext(ctx, q(`SELECT id FROM of_zones WHERE domain = ?`), root).Scan(&zoneID); err != nil {
return 0, err
}
report.Zones++
return zoneID, nil
}
func collectLegacyRouteDomainsImpl(
ctx context.Context,
tx *sql.Tx,
q func(string) string,
) (items []legacyDomain, hasRouteDomains bool, conflicts []string, err error) {
// Probe domain_cert_ids: if SELECT fails, fall back without it.
queryWithCert := q(`SELECT id, domain, domains, COALESCE(domain_cert_ids, '[]') FROM of_proxy_routes`)
rows, err := tx.QueryContext(ctx, queryWithCert)
useCert := true
if err != nil {
useCert = false
rows, err = tx.QueryContext(ctx, q(`SELECT id, domain, domains FROM of_proxy_routes`))
if err != nil {
return nil, false, nil, err
}
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var (
id uint
domain string
domains string
certIDs string
)
if useCert {
if err := rows.Scan(&id, &domain, &domains, &certIDs); err != nil {
return nil, false, nil, err
}
} else {
if err := rows.Scan(&id, &domain, &domains); err != nil {
return nil, false, nil, err
}
certIDs = "[]"
}
decoded, decodeErr := routeidentity.DecodeDomains(domains, domain)
if decodeErr != nil {
conflicts = append(conflicts, fmt.Sprintf("route %d: %v", id, decodeErr))
continue
}
if len(decoded) > 0 {
hasRouteDomains = true
}
ids := decodeLegacyCertIDs(certIDs, len(decoded))
routeID := id
for i, d := range decoded {
var certID *uint
if i < len(ids) && ids[i] > 0 {
v := ids[i]
certID = &v
}
items = append(items, legacyDomain{
Domain: d,
CertID: certID,
ProxyRouteID: &routeID,
})
}
}
return items, hasRouteDomains, conflicts, rows.Err()
}
func collectLegacyManagedDomains(ctx context.Context, tx *sql.Tx, q func(string) string) ([]legacyDomain, error) {
rows, err := tx.QueryContext(ctx, q(`SELECT domain, cert_id FROM of_managed_domains`))
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
items := make([]legacyDomain, 0)
for rows.Next() {
var (
domain string
certID sql.NullInt64
)
if err := rows.Scan(&domain, &certID); err != nil {
return nil, err
}
item := legacyDomain{Domain: domain}
if certID.Valid && certID.Int64 > 0 {
v := uint(certID.Int64)
item.CertID = &v
}
items = append(items, item)
}
return items, rows.Err()
}
func decodeLegacyCertIDs(raw string, count int) []uint {
var values []uint
if strings.TrimSpace(raw) == "" {
return make([]uint, count)
}
if json.Unmarshal([]byte(raw), &values) != nil {
return make([]uint, count)
}
return values
}
func nullableUint(v *uint) any {
if v == nil {
return nil
}
return *v
}
func rebindSQL(query string, postgres bool) string {
if !postgres {
return query
}
var b strings.Builder
b.Grow(len(query) + len(query)/4)
n := 0
for i := range len(query) {
if query[i] == '?' {
n++
b.WriteByte('$')
b.WriteString(strconv.Itoa(n))
continue
}
b.WriteByte(query[i])
}
return b.String()
}
func hasTable(
ctx context.Context,
tx *sql.Tx,
_ func(string) string,
postgres bool,
table string,
) (bool, error) {
var count int
var err error
if postgres {
err = tx.QueryRowContext(ctx, `
SELECT COUNT(*) FROM information_schema.tables
WHERE table_schema = 'public' AND table_name = $1
`, table).Scan(&count)
} else {
err = tx.QueryRowContext(ctx,
`SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?`, table,
).Scan(&count)
}
if err != nil {
return false, err
}
return count > 0, nil
}
func hasTableColumn(
ctx context.Context,
tx *sql.Tx,
_ func(string) string,
postgres bool,
table, column string,
) (bool, error) {
var count int
var err error
if postgres {
err = tx.QueryRowContext(ctx, `
SELECT COUNT(*) FROM information_schema.columns
WHERE table_schema = 'public' AND table_name = $1 AND column_name = $2
`, table, column).Scan(&count)
} else {
err = tx.QueryRowContext(ctx,
`SELECT COUNT(*) FROM pragma_table_info(?) WHERE name = ?`, table, column,
).Scan(&count)
}
if err != nil {
return false, err
}
return count > 0, nil
}
@@ -0,0 +1,151 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package zone
import (
"context"
"database/sql"
"testing"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupLegacyImportDB(t *testing.T) (*sql.DB, func()) {
t.Helper()
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
// Pre-phase-2 schema: legacy route columns + managed domains + zone tables.
stmts := []string{
`CREATE TABLE of_zones (
id INTEGER PRIMARY KEY AUTOINCREMENT,
domain TEXT NOT NULL UNIQUE,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`,
`CREATE TABLE of_zone_domains (
id INTEGER PRIMARY KEY AUTOINCREMENT,
zone_id INTEGER NOT NULL,
proxy_route_id INTEGER,
domain TEXT NOT NULL UNIQUE,
cert_id INTEGER,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
)`,
`CREATE TABLE of_proxy_routes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
site_name TEXT NOT NULL DEFAULT '',
domain TEXT NOT NULL DEFAULT '',
domains TEXT NOT NULL DEFAULT '[]',
domain_cert_ids TEXT NOT NULL DEFAULT '[]',
origin_url TEXT NOT NULL DEFAULT '',
remark TEXT NOT NULL DEFAULT ''
)`,
`CREATE TABLE of_tls_certificates (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL DEFAULT ''
)`,
`CREATE TABLE of_managed_domains (
id INTEGER PRIMARY KEY AUTOINCREMENT,
domain TEXT NOT NULL,
cert_id INTEGER,
remark TEXT NOT NULL DEFAULT ''
)`,
}
for _, stmt := range stmts {
_, err := sqlDB.Exec(stmt)
require.NoError(t, err)
}
previous := db.DB(context.Background())
db.SetDB(gormDB)
return sqlDB, func() {
db.SetDB(previous)
_ = sqlDB.Close()
}
}
func TestImportLegacyTxBindsRouteDomains(t *testing.T) {
sqlDB, cleanup := setupLegacyImportDB(t)
defer cleanup()
ctx := context.Background()
_, err := sqlDB.Exec(`INSERT INTO of_tls_certificates (id, name) VALUES (7, 'cert')`)
require.NoError(t, err)
_, err = sqlDB.Exec(`
INSERT INTO of_proxy_routes (id, site_name, domain, domains, domain_cert_ids, origin_url, remark)
VALUES (3, 'api', 'api.example.com', '["api.example.com","www.example.com"]', '[7,7]', 'http://origin', 'r')
`)
require.NoError(t, err)
tx, err := sqlDB.Begin()
require.NoError(t, err)
report, err := ImportLegacyTx(ctx, tx, false)
require.NoError(t, err)
require.NoError(t, tx.Commit())
assert.Equal(t, 1, report.Zones)
assert.Equal(t, 2, report.Domains)
var zoneDomain string
require.NoError(t, sqlDB.QueryRow(`SELECT domain FROM of_zones`).Scan(&zoneDomain))
assert.Equal(t, "example.com", zoneDomain)
var count int
require.NoError(t, sqlDB.QueryRow(`SELECT COUNT(*) FROM of_zone_domains WHERE proxy_route_id = 3`).Scan(&count))
assert.Equal(t, 2, count)
// Idempotent re-run
tx, err = sqlDB.Begin()
require.NoError(t, err)
report2, err := ImportLegacyTx(ctx, tx, false)
require.NoError(t, err)
require.NoError(t, tx.Commit())
assert.Equal(t, 0, report2.Domains)
}
func TestImportLegacyTxNoOpWithoutLegacyColumns(t *testing.T) {
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
defer sqlDB.Close()
_, err = sqlDB.Exec(`
CREATE TABLE of_zones (
id INTEGER PRIMARY KEY AUTOINCREMENT,
domain TEXT NOT NULL UNIQUE,
created_at DATETIME, updated_at DATETIME
);
CREATE TABLE of_zone_domains (
id INTEGER PRIMARY KEY AUTOINCREMENT,
zone_id INTEGER NOT NULL,
proxy_route_id INTEGER,
domain TEXT NOT NULL UNIQUE,
cert_id INTEGER,
created_at DATETIME, updated_at DATETIME
);
CREATE TABLE of_proxy_routes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
site_name TEXT NOT NULL DEFAULT '',
origin_url TEXT NOT NULL DEFAULT ''
);
`)
require.NoError(t, err)
tx, err := sqlDB.Begin()
require.NoError(t, err)
report, err := ImportLegacyTx(context.Background(), tx, false)
require.NoError(t, err)
require.NoError(t, tx.Commit())
assert.Equal(t, 0, report.Zones)
assert.Equal(t, 0, report.Domains)
}
@@ -0,0 +1,259 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package zone manages registered roots and their explicit hostnames.
package zone
import (
"context"
"errors"
"strings"
"time"
cf "Wavelet/openflare/plugins/server/domain/cloudflare"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/pkg/logger"
"golang.org/x/net/publicsuffix"
"gorm.io/gorm"
)
// Input is the mutable Zone payload.
type Input struct {
Domain string `json:"domain"`
}
// DomainInput is the mutable Zone-domain payload.
type DomainInput struct {
Domain string `json:"domain"`
CertID *uint `json:"cert_id"`
}
// Overview joins a Zone with its explicit domains.
type Overview struct {
Zone model.Zone `json:"zone"`
Domains []model.ZoneDomain `json:"domains"`
}
// ListItem is a Zone list row with denormalized domain count for the UI.
type ListItem struct {
ID uint `json:"id"`
Domain string `json:"domain"`
DomainCount int64 `json:"domain_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func zoneRoot(domain string) (string, error) {
return publicsuffix.EffectiveTLDPlusOne(strings.ToLower(strings.TrimSpace(domain)))
}
func normalizeDomain(raw string) (string, error) {
domain := strings.ToLower(strings.TrimSpace(raw))
if domain == "" {
return "", errors.New(errZoneDomainRequired)
}
if strings.Contains(domain, "*") {
return "", errors.New(errDomainWildcardUnsupported)
}
if strings.Contains(domain, "://") || strings.Contains(domain, "/") || strings.Contains(domain, "?") || strings.Contains(domain, "#") || strings.Contains(domain, "@") {
return "", errors.New(errDomainInvalid)
}
if _, err := zoneRoot(domain); err != nil {
return "", errors.New(errDomainInvalid)
}
return domain, nil
}
// Create persists a validated registered root.
func Create(ctx context.Context, input Input) (*model.Zone, error) {
domain, err := normalizeDomain(input.Domain)
if err != nil {
return nil, err
}
root, err := zoneRoot(domain)
if err != nil || root != domain {
return nil, errors.New(errZoneRootInvalid)
}
zone := &model.Zone{Domain: domain}
if err := repository.CreateZone(ctx, zone); err != nil {
if isUnique(err) {
return nil, errors.New(errDomainExists)
}
return nil, err
}
return zone, nil
}
// Update replaces a Zone's mutable fields.
func Update(ctx context.Context, id uint, input Input) (*model.Zone, error) {
zone, err := repository.GetZoneByID(ctx, id)
if err != nil {
return nil, err
}
domain, err := normalizeDomain(input.Domain)
if err != nil {
return nil, err
}
root, err := zoneRoot(domain)
if err != nil || root != domain {
return nil, errors.New(errZoneRootInvalid)
}
zone.Domain = domain
if err := repository.SaveZone(ctx, zone); err != nil {
if isUnique(err) {
return nil, errors.New(errDomainExists)
}
return nil, err
}
return zone, nil
}
// List returns all Zones in stable domain order, with domain counts for list cards.
func List(ctx context.Context) ([]ListItem, error) {
zones, err := repository.ListZones(ctx)
if err != nil {
return nil, err
}
rows, err := repository.ListZoneDomainCounts(ctx)
if err != nil {
return nil, err
}
counts := make(map[uint]int64, len(rows))
for _, row := range rows {
counts[row.ZoneID] = row.Count
}
items := make([]ListItem, 0, len(zones))
for _, zone := range zones {
items = append(items, ListItem{
ID: zone.ID,
Domain: zone.Domain,
DomainCount: counts[zone.ID],
CreatedAt: zone.CreatedAt,
UpdatedAt: zone.UpdatedAt,
})
}
return items, nil
}
// GetOverview returns a Zone and its domains.
func GetOverview(ctx context.Context, id uint) (*Overview, error) {
zone, err := repository.GetZoneByID(ctx, id)
if err != nil {
return nil, err
}
domains, err := repository.ListZoneDomainsByZoneID(ctx, id)
if err != nil {
return nil, err
}
return &Overview{Zone: *zone, Domains: domains}, nil
}
// CreateDomain adds a validated exact hostname to a Zone.
func CreateDomain(ctx context.Context, zoneID uint, input DomainInput) (*model.ZoneDomain, error) {
zone, err := repository.GetZoneByID(ctx, zoneID)
if err != nil {
return nil, err
}
domain, err := normalizeDomain(input.Domain)
if err != nil {
return nil, err
}
root, err := zoneRoot(domain)
if err != nil || root != zone.Domain {
return nil, errors.New(errDomainOutsideZone)
}
if input.CertID != nil {
if _, err := repository.GetTLSCertificateByID(ctx, *input.CertID); err != nil {
return nil, errors.New(errCertificateNotFound)
}
}
item := &model.ZoneDomain{ZoneID: zoneID, Domain: domain, CertID: input.CertID}
if err := repository.CreateZoneDomain(ctx, item); err != nil {
if isUnique(err) {
return nil, errors.New(errDomainExists)
}
return nil, err
}
return item, nil
}
// UpdateDomain replaces a Zone-domain's mutable fields.
func UpdateDomain(ctx context.Context, zoneID, id uint, input DomainInput) (*model.ZoneDomain, error) {
item, err := repository.GetZoneDomainByZoneAndID(ctx, zoneID, id)
if err != nil {
return nil, err
}
domain, err := normalizeDomain(input.Domain)
if err != nil {
return nil, err
}
zone, err := repository.GetZoneByID(ctx, zoneID)
if err != nil {
return nil, err
}
root, err := zoneRoot(domain)
if err != nil || root != zone.Domain {
return nil, errors.New(errDomainOutsideZone)
}
if input.CertID != nil {
if _, err = repository.GetTLSCertificateByID(ctx, *input.CertID); err != nil {
return nil, errors.New(errCertificateNotFound)
}
}
item.Domain, item.CertID = domain, input.CertID
if err = repository.SaveZoneDomain(ctx, item); err != nil {
if isUnique(err) {
return nil, errors.New(errDomainExists)
}
return nil, err
}
return item, nil
}
// DeleteDomain removes a Zone domain that is not bound to a proxy route.
func DeleteDomain(ctx context.Context, zoneID, id uint) error {
item, err := repository.GetZoneDomainByZoneAndID(ctx, zoneID, id)
if err != nil {
return err
}
if item.ProxyRouteID != nil {
return errors.New(errDomainBoundToRoute)
}
member, cfErr := repository.GetCFPointingMemberByZoneDomainID(ctx, item.ID)
if cfErr != nil && !errors.Is(cfErr, gorm.ErrRecordNotFound) {
return cfErr
}
if member != nil {
if delErr := cf.DeleteManagedRecord(ctx, member.ID); delErr != nil {
logger.WarnF(ctx, "[Zone] delete managed Cloudflare record failed for domain %s (member_id=%d): %v", item.Domain, member.ID, delErr)
}
if delMemberErr := repository.DeleteCFPointingMember(ctx, member); delMemberErr != nil {
logger.ErrorF(ctx, "[Zone] delete Cloudflare pointing member failed: member_id=%d error=%v", member.ID, delMemberErr)
return delMemberErr
}
}
return repository.DeleteZoneDomain(ctx, item)
}
// Delete removes a Zone that has no remaining domains.
func Delete(ctx context.Context, id uint) error {
if _, err := repository.GetZoneByID(ctx, id); err != nil {
return err
}
count, err := repository.CountZoneDomainsByZoneID(ctx, id)
if err != nil {
return err
}
if count > 0 {
return errors.New(errZoneHasDomains)
}
return repository.DeleteZone(ctx, id)
}
func isUnique(err error) bool {
return errors.Is(err, gorm.ErrDuplicatedKey) || strings.Contains(strings.ToLower(err.Error()), "unique constraint")
}
@@ -0,0 +1,128 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package zone
import (
"context"
"testing"
"time"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/testhelper"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupZoneDB(t *testing.T) context.Context {
t.Helper()
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true})
require.NoError(t, err)
require.NoError(t, conn.AutoMigrate(&model.Zone{}, &model.ZoneDomain{}, &model.TLSCertificate{}, &model.CFPointingGroup{}, &model.CFPointingMember{}))
db.SetDB(conn)
t.Cleanup(func() { db.SetDB(nil) })
return context.Background()
}
func TestCreateZoneDomainRejectsWildcard(t *testing.T) {
ctx := setupZoneDB(t)
zone, err := Create(ctx, Input{Domain: "example.com"})
require.NoError(t, err)
_, err = CreateDomain(ctx, zone.ID, DomainInput{Domain: "*.example.com"})
require.EqualError(t, err, errDomainWildcardUnsupported)
}
func TestDeleteDomainRejectsBoundRoute(t *testing.T) {
ctx := setupZoneDB(t)
zone, err := Create(ctx, Input{Domain: "example.com"})
require.NoError(t, err)
item, err := CreateDomain(ctx, zone.ID, DomainInput{Domain: "api.example.com"})
require.NoError(t, err)
routeID := uint(9)
item.ProxyRouteID = &routeID
require.NoError(t, repository.SaveZoneDomain(ctx, item))
err = DeleteDomain(ctx, zone.ID, item.ID)
require.EqualError(t, err, errDomainBoundToRoute)
item.ProxyRouteID = nil
require.NoError(t, repository.SaveZoneDomain(ctx, item))
require.NoError(t, DeleteDomain(ctx, zone.ID, item.ID))
}
func TestDeleteDomainCleansUpCloudflareMember(t *testing.T) {
ctx := setupZoneDB(t)
zone, err := Create(ctx, Input{Domain: "example.com"})
require.NoError(t, err)
domain, err := CreateDomain(ctx, zone.ID, DomainInput{Domain: "api.example.com"})
require.NoError(t, err)
member := model.CFPointingMember{GroupID: 1, ZoneDomainID: domain.ID}
require.NoError(t, repository.CreateCFPointingMember(ctx, &member))
require.NoError(t, DeleteDomain(ctx, zone.ID, domain.ID))
_, err = repository.GetCFPointingMemberByZoneDomainID(ctx, domain.ID)
require.Error(t, err)
}
func TestLegacyImportUsesEffectiveTLDPlusOne(t *testing.T) {
root, err := zoneRoot("api.example.co.uk")
require.NoError(t, err)
require.Equal(t, "example.co.uk", root)
}
func TestGetStatsAggregatesZoneHosts(t *testing.T) {
ctx := setupZoneDB(t)
testhelper.SetupLogStoresForTest(t)
zone, err := Create(ctx, Input{Domain: "example.com"})
require.NoError(t, err)
_, err = CreateDomain(ctx, zone.ID, DomainInput{Domain: "api.example.com"})
require.NoError(t, err)
_, err = CreateDomain(ctx, zone.ID, DomainInput{Domain: "www.example.com"})
require.NoError(t, err)
now := time.Now().UTC()
require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{
{NodeID: "n1", LoggedAt: now.Add(-1 * time.Hour), RemoteAddr: "1.1.1.1", Host: "api.example.com", Path: "/", StatusCode: 200, BytesSent: 1000},
{NodeID: "n1", LoggedAt: now.Add(-2 * time.Hour), RemoteAddr: "1.1.1.1", Host: "www.example.com", Path: "/", StatusCode: 200, BytesSent: 500},
{NodeID: "n1", LoggedAt: now.Add(-3 * time.Hour), RemoteAddr: "2.2.2.2", Host: "api.example.com", Path: "/x", StatusCode: 404, BytesSent: 200},
{NodeID: "n1", LoggedAt: now.Add(-3 * time.Hour), RemoteAddr: "3.3.3.3", Host: "other.com", Path: "/", StatusCode: 200, BytesSent: 100},
{NodeID: "n1", LoggedAt: now.Add(-48 * time.Hour), RemoteAddr: "4.4.4.4", Host: "api.example.com", Path: "/", StatusCode: 200, BytesSent: 800},
}))
stats, err := GetStats(ctx, zone.ID, "24h")
require.NoError(t, err)
require.Equal(t, StatsRange24h, stats.Range)
require.Equal(t, int64(3), stats.RequestCount)
require.Equal(t, int64(2), stats.UniqueVisitors)
require.Equal(t, int64(1700), stats.BytesSent)
require.Equal(t, 2, stats.DomainCount)
require.True(t, stats.Available)
require.NotEmpty(t, stats.Series)
require.Equal(t, 60, stats.BucketMinutes)
var seriesRequests int64
var seriesBytes int64
for _, point := range stats.Series {
seriesRequests += point.RequestCount
seriesBytes += point.BytesSent
}
require.Equal(t, int64(3), seriesRequests)
require.Equal(t, int64(1700), seriesBytes)
stats7d, err := GetStats(ctx, zone.ID, "7d")
require.NoError(t, err)
require.Equal(t, int64(4), stats7d.RequestCount)
require.Equal(t, int64(3), stats7d.UniqueVisitors)
require.Equal(t, int64(2500), stats7d.BytesSent)
require.NotEmpty(t, stats7d.Series)
_, err = GetStats(ctx, zone.ID, "1h")
require.EqualError(t, err, errStatsRangeInvalid)
}
@@ -0,0 +1,252 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package zone
import (
"errors"
"net/http"
"Wavelet/openflare/plugins/server/kernel/apiutil"
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
func abort(c *gin.Context, err error, missing string) bool {
if err == nil {
return false
}
switch {
case errors.Is(err, gorm.ErrRecordNotFound):
response.AbortNotFound(c, missing)
case err.Error() == errDomainExists:
response.AbortConflict(c, err.Error())
default:
response.AbortBadRequest(c, err.Error())
}
return true
}
// ListHandler lists registered Zones.
// @Summary 获取 Zone 列表
// @Tags openflare-zone
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]zone.ListItem}
// @Router /api/v1/d/zones [get]
func ListHandler(c *gin.Context) {
items, err := List(c.Request.Context())
if abort(c, err, errZoneNotFound) {
return
}
c.JSON(http.StatusOK, response.OK(items))
}
// CreateHandler creates a registered root domain.
// @Summary 创建 Zone
// @Tags openflare-zone
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body zone.Input true "Zone 参数"
// @Success 200 {object} response.Any{data=model.Zone}
// @Failure 400 {object} response.Any
// @Failure 409 {object} response.Any
// @Router /api/v1/d/zones [post]
func CreateHandler(c *gin.Context) {
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
item, err := Create(c.Request.Context(), input)
if abort(c, err, errZoneNotFound) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// GetOverviewHandler returns a Zone and its explicit domains.
// @Summary 获取 Zone 概览
// @Tags openflare-zone
// @Produce json
// @Security SessionCookie
// @Param id path int true "Zone ID"
// @Success 200 {object} response.Any{data=zone.Overview}
// @Failure 404 {object} response.Any
// @Router /api/v1/d/zones/{id}/overview [get]
func GetOverviewHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
item, err := GetOverview(c.Request.Context(), id)
if abort(c, err, errZoneNotFound) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// GetStatsHandler returns Zone traffic metrics for a time range.
// @Summary 获取 Zone 流量统计
// @Description 按 Zone 下全部域名聚合访问日志:唯一访问者、请求总数、已提供数据(字节)。range 支持 24h/7d/30d。
// @Tags openflare-zone
// @Produce json
// @Security SessionCookie
// @Param id path int true "Zone ID"
// @Param range query string false "时间范围:24h(默认)、7d、30d"
// @Success 200 {object} response.Any{data=zone.Stats}
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/d/zones/{id}/stats [get]
func GetStatsHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
item, err := GetStats(c.Request.Context(), id, c.Query("range"))
if abort(c, err, errZoneNotFound) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// UpdateHandler updates a Zone.
// @Summary 更新 Zone
// @Tags openflare-zone
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "Zone ID"
// @Param body body zone.Input true "Zone 参数"
// @Success 200 {object} response.Any{data=model.Zone}
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Failure 409 {object} response.Any
// @Router /api/v1/d/zones/{id}/update [post]
func UpdateHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input Input
if !apiutil.BindJSON(c, &input) {
return
}
item, err := Update(c.Request.Context(), id, input)
if abort(c, err, errZoneNotFound) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// DeleteHandler deletes a Zone with no remaining domains.
// @Summary 删除 Zone
// @Tags openflare-zone
// @Produce json
// @Security SessionCookie
// @Param id path int true "Zone ID"
// @Success 200 {object} response.Any
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/d/zones/{id}/delete [post]
func DeleteHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
if err := Delete(c.Request.Context(), id); abort(c, err, errZoneNotFound) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// CreateDomainHandler creates an explicit FQDN under a Zone.
// @Summary 创建 Zone 域名
// @Tags openflare-zone
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "Zone ID"
// @Param body body zone.DomainInput true "域名参数"
// @Success 200 {object} response.Any{data=model.ZoneDomain}
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Failure 409 {object} response.Any
// @Router /api/v1/d/zones/{id}/domains [post]
func CreateDomainHandler(c *gin.Context) {
id, ok := apiutil.IDParam(c)
if !ok {
return
}
var input DomainInput
if !apiutil.BindJSON(c, &input) {
return
}
item, err := CreateDomain(c.Request.Context(), id, input)
if abort(c, err, errZoneNotFound) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// UpdateDomainHandler updates a Zone domain.
// @Summary 更新 Zone 域名
// @Tags openflare-zone
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "Zone ID"
// @Param domainId path int true "域名 ID"
// @Param body body zone.DomainInput true "域名参数"
// @Success 200 {object} response.Any{data=model.ZoneDomain}
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Failure 409 {object} response.Any
// @Router /api/v1/d/zones/{id}/domains/{domainId}/update [post]
func UpdateDomainHandler(c *gin.Context) {
zoneID, ok := apiutil.IDParam(c)
if !ok {
return
}
domainID, ok := apiutil.NamedIDParam(c, "domainId")
if !ok {
return
}
var input DomainInput
if !apiutil.BindJSON(c, &input) {
return
}
item, err := UpdateDomain(c.Request.Context(), zoneID, domainID, input)
if abort(c, err, errDomainNotFound) {
return
}
c.JSON(http.StatusOK, response.OK(item))
}
// DeleteDomainHandler deletes a Zone domain not bound to a proxy route.
// @Summary 删除 Zone 域名
// @Tags openflare-zone
// @Produce json
// @Security SessionCookie
// @Param id path int true "Zone ID"
// @Param domainId path int true "域名 ID"
// @Success 200 {object} response.Any
// @Failure 400 {object} response.Any
// @Failure 404 {object} response.Any
// @Router /api/v1/d/zones/{id}/domains/{domainId}/delete [post]
func DeleteDomainHandler(c *gin.Context) {
zoneID, ok := apiutil.IDParam(c)
if !ok {
return
}
domainID, ok := apiutil.NamedIDParam(c, "domainId")
if !ok {
return
}
if err := DeleteDomain(c.Request.Context(), zoneID, domainID); abort(c, err, errDomainNotFound) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,209 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package zone
import (
"context"
"errors"
"strings"
"time"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
"gorm.io/gorm"
)
// StatsRange is a supported traffic window for Zone analytics.
type StatsRange string
// StatsRange constants representing supported analytics windows.
const (
// StatsRange24h represents a 24-hour time window.
StatsRange24h StatsRange = "24h"
// StatsRange7d represents a 7-day time window.
StatsRange7d StatsRange = "7d"
// StatsRange30d represents a 30-day time window.
StatsRange30d StatsRange = "30d"
)
const (
hoursPerDay = 24
daysPerWeek = 7
daysPerMonth = 30
minutesPerHour = 60
bucketMinutes24h = 60
bucketMinutes7d = 6 * minutesPerHour
bucketMinutes30d = 24 * minutesPerHour
)
// StatsPoint is one bucket on a Zone traffic chart.
type StatsPoint struct {
BucketStartedAt time.Time `json:"bucket_started_at"`
RequestCount int64 `json:"request_count"`
UniqueVisitors int64 `json:"unique_visitors"`
BytesSent int64 `json:"bytes_sent"`
}
// Stats summarizes edge traffic for all domains under a Zone.
type Stats struct {
Range StatsRange `json:"range"`
RangeHours int `json:"range_hours"`
WindowStartedAt time.Time `json:"window_started_at"`
WindowEndedAt time.Time `json:"window_ended_at"`
BucketMinutes int `json:"bucket_minutes"`
UniqueVisitors int64 `json:"unique_visitors"`
RequestCount int64 `json:"request_count"`
BytesSent int64 `json:"bytes_sent"`
DomainCount int `json:"domain_count"`
Available bool `json:"available"`
Series []StatsPoint `json:"series"`
}
func parseStatsRange(raw string) (StatsRange, time.Duration, int, error) {
switch StatsRange(strings.TrimSpace(raw)) {
case "", StatsRange24h:
return StatsRange24h, hoursPerDay * time.Hour, bucketMinutes24h, nil
case StatsRange7d:
return StatsRange7d, daysPerWeek * hoursPerDay * time.Hour, bucketMinutes7d, nil
case StatsRange30d:
return StatsRange30d, daysPerMonth * hoursPerDay * time.Hour, bucketMinutes30d, nil
default:
return "", 0, 0, errors.New(errStatsRangeInvalid)
}
}
// GetStats aggregates access-log traffic for a Zone over a time range.
func GetStats(ctx context.Context, id uint, rangeRaw string) (*Stats, error) {
statsRange, window, bucketMinutes, err := parseStatsRange(rangeRaw)
if err != nil {
return nil, err
}
if _, err := repository.GetZoneByID(ctx, id); err != nil {
return nil, err
}
domains, err := repository.ListZoneDomainsByZoneID(ctx, id)
if err != nil {
return nil, err
}
now := time.Now().UTC().Truncate(time.Minute)
since := now.Add(-window)
// Align chart window start to bucket boundary for cleaner x-axis labels.
bucket := time.Duration(bucketMinutes) * time.Minute
since = since.Truncate(bucket)
result := &Stats{
Range: statsRange,
RangeHours: int(window / time.Hour),
WindowStartedAt: since,
WindowEndedAt: now,
BucketMinutes: bucketMinutes,
DomainCount: len(domains),
Available: true,
Series: emptyStatsSeries(since, now, bucketMinutes),
}
if len(domains) == 0 {
return result, nil
}
hosts := make([]string, 0, len(domains))
for _, domain := range domains {
if host := strings.TrimSpace(domain.Domain); host != "" {
hosts = append(hosts, host)
}
}
if len(hosts) == 0 {
return result, nil
}
requestCount, uniqueVisitors, totalBytesSent, err := repository.CountOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{
Hosts: hosts,
Since: since,
Until: now,
})
if err != nil {
if isAnalyticsUnavailable(err) {
result.Available = false
return result, nil
}
return nil, err
}
result.RequestCount = requestCount
result.UniqueVisitors = uniqueVisitors
result.BytesSent = totalBytesSent
buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{
Hosts: hosts,
Since: since,
Until: now,
FoldMinutes: bucketMinutes,
SortBy: "logged_at",
SortOrder: "asc",
})
if err != nil {
if isAnalyticsUnavailable(err) {
result.Available = false
return result, nil
}
return nil, err
}
byEpoch := make(map[int64]model.OpenFlareAccessLogBucketRow, len(buckets))
for _, bucketRow := range buckets {
if bucketRow == nil {
continue
}
byEpoch[bucketRow.BucketEpoch] = *bucketRow
}
series := emptyStatsSeries(since, now, bucketMinutes)
for index := range series {
epoch := series[index].BucketStartedAt.Unix()
if row, ok := byEpoch[epoch]; ok {
series[index].RequestCount = row.RequestCount
series[index].UniqueVisitors = row.UniqueIPCount
series[index].BytesSent = row.BytesSent
}
}
result.Series = series
return result, nil
}
func emptyStatsSeries(since, until time.Time, bucketMinutes int) []StatsPoint {
if bucketMinutes <= 0 {
bucketMinutes = 60
}
bucket := time.Duration(bucketMinutes) * time.Minute
start := since.UTC().Truncate(bucket)
end := until.UTC()
if !end.After(start) {
return []StatsPoint{{BucketStartedAt: start}}
}
// Cap points to keep chart readable.
maxPoints := 120
capacity := min(int(end.Sub(start)/bucket)+1, maxPoints)
points := make([]StatsPoint, 0, capacity)
for cursor := start; !cursor.After(end) && len(points) < maxPoints; cursor = cursor.Add(bucket) {
points = append(points, StatsPoint{BucketStartedAt: cursor})
}
if len(points) == 0 {
points = append(points, StatsPoint{BucketStartedAt: start})
}
return points
}
func isAnalyticsUnavailable(err error) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrInvalidDB) {
return true
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "clickhouse connection is not initialized") ||
strings.Contains(msg, "clickhouse is not") ||
strings.Contains(msg, "database is not initialized")
}