mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 09:06:36 +08:00
refactor(backend): rename OpenFlare directory to lowercase openflare
This commit is contained in:
@@ -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
|
||||
}
|
||||
+119
@@ -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)
|
||||
}
|
||||
+128
@@ -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
|
||||
}
|
||||
+158
@@ -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")
|
||||
}
|
||||
Reference in New Issue
Block a user