mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 15:46:37 +08:00
后端与全仓代码质量清理(golangci 扩展集 · 测试质量 · 并发安全 · 文档同步)
代码质量全量清理,零行为变化:golangci 扩展集 13 类 linter(gosec/modernize/perfsprint/canonicalheader/usestdlibvars/wastedassign/intrange/errorlint/forcetypeassert/recvcheck/exhaustive/unparam)全量修复,测试代码质量(testifylint/thelper/usetesting)25→0,frpc 进程生命周期真 bug(进程组击杀)、全仓 go test -race 6 类数据竞争(含 1 个生产竞争)、SPDX license 头补齐 131 文件、前端测试套件 next-intl 迁移后 44 失败→全绿、过期 swagger 文档重新生成、pnpm-workspace 构建审批。 Experiments: #2-#17, #18, #20, #21, #23 Metric: total_issues 108 → 8 (-92.6%)
This commit is contained in:
@@ -179,18 +179,9 @@ func buildNodeAccessLogRecords(nodeID string, direct []NodeAccessLog, buffered [
|
||||
records := make([]*model.OpenFlareAccessLog, 0, total)
|
||||
appendLogs := func(logs []NodeAccessLog) {
|
||||
for _, item := range logs {
|
||||
bytesSent := item.BytesSent
|
||||
if bytesSent < 0 {
|
||||
bytesSent = 0
|
||||
}
|
||||
requestLength := item.RequestLength
|
||||
if requestLength < 0 {
|
||||
requestLength = 0
|
||||
}
|
||||
requestTimeMs := item.RequestTimeMs
|
||||
if requestTimeMs < 0 {
|
||||
requestTimeMs = 0
|
||||
}
|
||||
bytesSent := max(item.BytesSent, 0)
|
||||
requestLength := max(item.RequestLength, 0)
|
||||
requestTimeMs := max(item.RequestTimeMs, 0)
|
||||
record := &model.OpenFlareAccessLog{
|
||||
NodeID: nodeID,
|
||||
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -45,7 +46,7 @@ func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[s
|
||||
}
|
||||
changed := make([]WAFIPGroup, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum {
|
||||
if strings.TrimSpace(checksums[strconv.FormatUint(uint64(group.ID), 10)]) == group.Checksum {
|
||||
continue
|
||||
}
|
||||
changed = append(changed, group)
|
||||
@@ -97,7 +98,7 @@ func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error
|
||||
if len(ids) == 0 {
|
||||
return []WAFIPGroup{}, nil
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
slices.Sort(ids)
|
||||
groups, err := repository.ListOpenFlareWAFIPGroupsByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -204,7 +205,7 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
slices.Sort(ids)
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@ import (
|
||||
)
|
||||
|
||||
func setupApplyLogTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package openflare implements openflare configuration, service orchestration, and background tasks.
|
||||
package openflare
|
||||
|
||||
|
||||
@@ -192,7 +192,7 @@ func (client *HTTPClient) do(ctx context.Context, method, path string, query url
|
||||
return err
|
||||
}
|
||||
requestURL := buildRequestURL(client.baseURL, path, query)
|
||||
for attempt := 0; attempt < maxRequestAttempts; attempt++ {
|
||||
for attempt := range maxRequestAttempts {
|
||||
statusCode, retryHeader, responseBody, requestErr := client.send(ctx, method, requestURL, encodedBody)
|
||||
if requestErr != nil {
|
||||
return requestErr
|
||||
|
||||
@@ -94,6 +94,7 @@ func TestBuildSnapshotReadsZoneDomainCertificates(t *testing.T) {
|
||||
}
|
||||
|
||||
func generateTestCertKeyPairForSnapshot(t *testing.T) (certPEM string, keyPEM string) {
|
||||
t.Helper()
|
||||
return generateTestCertKeyPairForSnapshotForDomain(t, "test.example.com")
|
||||
}
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ package config_version
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"strconv"
|
||||
@@ -27,7 +28,7 @@ func normalizeSnapshotDomains(domains []string) ([]string, error) {
|
||||
for _, raw := range domains {
|
||||
domain := strings.ToLower(strings.TrimSpace(raw))
|
||||
if domain == "" || strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, fmt.Errorf("domains payload is invalid")
|
||||
return nil, errors.New("domains payload is invalid")
|
||||
}
|
||||
if _, ok := seen[domain]; ok {
|
||||
continue
|
||||
@@ -36,7 +37,7 @@ func normalizeSnapshotDomains(domains []string) ([]string, error) {
|
||||
normalized = append(normalized, domain)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, fmt.Errorf("domain is required")
|
||||
return nil, errors.New("domain is required")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
@@ -55,7 +56,7 @@ func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, erro
|
||||
}
|
||||
var upstreams []string
|
||||
if err := json.Unmarshal([]byte(text), &upstreams); err != nil {
|
||||
return nil, fmt.Errorf("upstreams payload is invalid")
|
||||
return nil, errors.New("upstreams payload is invalid")
|
||||
}
|
||||
return normalizeUpstreams(fallbackOriginURL, upstreams)
|
||||
}
|
||||
@@ -79,7 +80,7 @@ func normalizeUpstreams(originURL string, upstreams []string) ([]string, error)
|
||||
normalized = append(normalized, value)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, fmt.Errorf("upstream is required")
|
||||
return nil, errors.New("upstream is required")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
@@ -91,7 +92,7 @@ func decodeStoredCustomHeaders(raw string) ([]customHeaderInput, error) {
|
||||
}
|
||||
var headers []customHeaderInput
|
||||
if err := json.Unmarshal([]byte(text), &headers); err != nil {
|
||||
return nil, fmt.Errorf("custom_headers payload is invalid")
|
||||
return nil, errors.New("custom_headers payload is invalid")
|
||||
}
|
||||
return headers, nil
|
||||
}
|
||||
@@ -103,7 +104,7 @@ func decodeStoredCacheRules(raw string) ([]string, error) {
|
||||
}
|
||||
var rules []string
|
||||
if err := json.Unmarshal([]byte(text), &rules); err != nil {
|
||||
return nil, fmt.Errorf("cache_rules payload is invalid")
|
||||
return nil, errors.New("cache_rules payload is invalid")
|
||||
}
|
||||
normalized := make([]string, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
|
||||
@@ -494,51 +494,51 @@ func diffOpenRestyOptionDetails(left openRestyConfigSnapshot, right openRestyCon
|
||||
CurrentValue: current,
|
||||
})
|
||||
}
|
||||
appendIfChanged("OpenRestyDefaultServerReturnStatus", fmt.Sprintf("%d", left.DefaultServerReturnStatus), fmt.Sprintf("%d", right.DefaultServerReturnStatus))
|
||||
appendIfChanged("OpenRestyDefaultServerReturnStatus", strconv.Itoa(left.DefaultServerReturnStatus), strconv.Itoa(right.DefaultServerReturnStatus))
|
||||
appendIfChanged("OpenRestyWorkerProcesses", left.WorkerProcesses, right.WorkerProcesses)
|
||||
appendIfChanged("OpenRestyWorkerConnections", fmt.Sprintf("%d", left.WorkerConnections), fmt.Sprintf("%d", right.WorkerConnections))
|
||||
appendIfChanged("OpenRestyWorkerRlimitNofile", fmt.Sprintf("%d", left.WorkerRlimitNofile), fmt.Sprintf("%d", right.WorkerRlimitNofile))
|
||||
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", fmt.Sprintf("%t", left.EventsMultiAcceptEnabled), fmt.Sprintf("%t", right.EventsMultiAcceptEnabled))
|
||||
appendIfChanged("OpenRestyKeepaliveTimeout", fmt.Sprintf("%d", left.KeepaliveTimeout), fmt.Sprintf("%d", right.KeepaliveTimeout))
|
||||
appendIfChanged("OpenRestyKeepaliveRequests", fmt.Sprintf("%d", left.KeepaliveRequests), fmt.Sprintf("%d", right.KeepaliveRequests))
|
||||
appendIfChanged("OpenRestyClientHeaderTimeout", fmt.Sprintf("%d", left.ClientHeaderTimeout), fmt.Sprintf("%d", right.ClientHeaderTimeout))
|
||||
appendIfChanged("OpenRestyClientBodyTimeout", fmt.Sprintf("%d", left.ClientBodyTimeout), fmt.Sprintf("%d", right.ClientBodyTimeout))
|
||||
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", fmt.Sprintf("%d", left.SendTimeout), fmt.Sprintf("%d", right.SendTimeout))
|
||||
appendIfChanged("OpenRestyProxyConnectTimeout", fmt.Sprintf("%d", left.ProxyConnectTimeout), fmt.Sprintf("%d", right.ProxyConnectTimeout))
|
||||
appendIfChanged("OpenRestyProxySendTimeout", fmt.Sprintf("%d", left.ProxySendTimeout), fmt.Sprintf("%d", right.ProxySendTimeout))
|
||||
appendIfChanged("OpenRestyProxyReadTimeout", fmt.Sprintf("%d", left.ProxyReadTimeout), fmt.Sprintf("%d", right.ProxyReadTimeout))
|
||||
appendIfChanged("OpenRestyWebsocketEnabled", fmt.Sprintf("%t", left.WebsocketEnabled), fmt.Sprintf("%t", right.WebsocketEnabled))
|
||||
appendIfChanged("OpenRestyHTTP3Enabled", fmt.Sprintf("%t", left.HTTP3Enabled), fmt.Sprintf("%t", right.HTTP3Enabled))
|
||||
appendIfChanged("OpenRestyProxyRequestBufferingEnabled", fmt.Sprintf("%t", left.ProxyRequestBuffering), fmt.Sprintf("%t", right.ProxyRequestBuffering))
|
||||
appendIfChanged("OpenRestyProxyBufferingEnabled", fmt.Sprintf("%t", left.ProxyBufferingEnabled), fmt.Sprintf("%t", right.ProxyBufferingEnabled))
|
||||
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", fmt.Sprintf("%t", left.GzipEnabled), fmt.Sprintf("%t", right.GzipEnabled))
|
||||
appendIfChanged("OpenRestyGzipMinLength", fmt.Sprintf("%d", left.GzipMinLength), fmt.Sprintf("%d", right.GzipMinLength))
|
||||
appendIfChanged("OpenRestyGzipCompLevel", fmt.Sprintf("%d", left.GzipCompLevel), fmt.Sprintf("%d", right.GzipCompLevel))
|
||||
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", fmt.Sprintf("%t", left.CacheEnabled), fmt.Sprintf("%t", right.CacheEnabled))
|
||||
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", fmt.Sprintf("%t", left.CacheLockEnabled), fmt.Sprintf("%t", right.CacheLockEnabled))
|
||||
appendIfChanged("OpenRestyCacheLockEnabled", strconv.FormatBool(left.CacheLockEnabled), strconv.FormatBool(right.CacheLockEnabled))
|
||||
appendIfChanged("OpenRestyCacheLockTimeout", left.CacheLockTimeout, right.CacheLockTimeout)
|
||||
appendIfChanged("OpenRestyCacheUseStale", left.CacheUseStale, right.CacheUseStale)
|
||||
appendIfChanged("OpenRestyDefaultLimitConnPerServer", fmt.Sprintf("%d", left.DefaultLimitConnPerServer), fmt.Sprintf("%d", right.DefaultLimitConnPerServer))
|
||||
appendIfChanged("OpenRestyDefaultLimitConnPerIP", fmt.Sprintf("%d", left.DefaultLimitConnPerIP), fmt.Sprintf("%d", right.DefaultLimitConnPerIP))
|
||||
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", fmt.Sprintf("%t", left.OriginErrorPageEnabled), fmt.Sprintf("%t", right.OriginErrorPageEnabled))
|
||||
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", fmt.Sprintf("%t", left.OriginErrorPageGetOnly), fmt.Sprintf("%t", right.OriginErrorPageGetOnly))
|
||||
appendIfChanged("SWOfflineEnabled", fmt.Sprintf("%t", left.SWOfflineEnabled), fmt.Sprintf("%t", right.SWOfflineEnabled))
|
||||
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
|
||||
|
||||
@@ -98,7 +98,7 @@ func TestDiffOpenRestyOptionDetailsOriginErrorPage(t *testing.T) {
|
||||
assert.Equal(t, "false", keys["OriginErrorPageEnabled"].CurrentValue)
|
||||
assert.Equal(t, `["500-599"]`, keys["OriginErrorPageStatusCodes"].PreviousValue)
|
||||
assert.Equal(t, `["522"]`, keys["OriginErrorPageStatusCodes"].CurrentValue)
|
||||
assert.Equal(t, "", keys["OriginErrorPageHTML"].PreviousValue)
|
||||
assert.Empty(t, keys["OriginErrorPageHTML"].PreviousValue)
|
||||
assert.Equal(t, "<p>x</p>", keys["OriginErrorPageHTML"].CurrentValue)
|
||||
}
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -199,9 +200,7 @@ func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configB
|
||||
return nil, err
|
||||
}
|
||||
|
||||
mainConfig := ""
|
||||
routeConfig := ""
|
||||
checksum := ""
|
||||
var mainConfig, routeConfig, checksum string
|
||||
supportFiles := []SupportFile(nil)
|
||||
|
||||
rendered, renderErr := renderSnapshotConfig(string(snapshotJSON), certificateFiles)
|
||||
@@ -429,7 +428,7 @@ func buildSnapshotWAFIPGroups(ctx context.Context, idSet map[uint]struct{}) ([]s
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
slices.Sort(ids)
|
||||
groups, err := listWAFIPGroupsByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -477,7 +476,7 @@ func decodeIPList(raw string) ([]string, error) {
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(text), &items); err != nil {
|
||||
return nil, fmt.Errorf("ip_list payload is invalid")
|
||||
return nil, errors.New("ip_list payload is invalid")
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
@@ -629,7 +628,7 @@ func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) (
|
||||
for certID := range certIDSet {
|
||||
certIDs = append(certIDs, certID)
|
||||
}
|
||||
sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] })
|
||||
slices.Sort(certIDs)
|
||||
files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate)
|
||||
for _, certID := range certIDs {
|
||||
certificate, err := repository.GetTLSCertificateByID(ctx, certID)
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
)
|
||||
|
||||
func setupDashboardTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
@@ -124,8 +125,8 @@ func TestGetOverviewStructure(t *testing.T) {
|
||||
onlineNodeCheck := overview.Nodes
|
||||
require.NotEmpty(t, onlineNodeCheck)
|
||||
|
||||
assert.Equal(t, 55.0, overview.Capacity.AverageCPUUsagePercent)
|
||||
assert.Equal(t, 50.0, overview.Capacity.AverageMemoryUsagePercent)
|
||||
assert.InDelta(t, 55.0, overview.Capacity.AverageCPUUsagePercent, 1e-9)
|
||||
assert.InDelta(t, 50.0, overview.Capacity.AverageMemoryUsagePercent, 1e-9)
|
||||
assert.Equal(t, 0, overview.Capacity.HighCPUNodes)
|
||||
assert.Equal(t, 0, overview.Capacity.HighMemoryNodes)
|
||||
assert.Equal(t, 0, overview.Capacity.HighStorageNodes)
|
||||
@@ -171,11 +172,11 @@ func TestGetOverviewStructure(t *testing.T) {
|
||||
assert.Equal(t, "online", onlineNode[6])
|
||||
assert.Equal(t, "healthy", onlineNode[7])
|
||||
// Latest-per-node health fields (indexes match compressDashboardNodes).
|
||||
assert.Equal(t, 55.0, onlineNode[11]) // cpu_usage_percent from latest snapshot
|
||||
assert.Equal(t, 50.0, onlineNode[12]) // memory_usage_percent
|
||||
assert.Equal(t, int64(12), onlineNode[14]) // request_count from access logs
|
||||
assert.Equal(t, int64(1), onlineNode[15]) // error_count
|
||||
assert.Equal(t, int64(4), onlineNode[16]) // unique visitors
|
||||
assert.InDelta(t, 55.0, onlineNode[11], 1e-9) // cpu_usage_percent from latest snapshot
|
||||
assert.InDelta(t, 50.0, onlineNode[12], 1e-9) // memory_usage_percent
|
||||
assert.Equal(t, int64(12), onlineNode[14]) // request_count from access logs
|
||||
assert.Equal(t, int64(1), onlineNode[15]) // error_count
|
||||
assert.Equal(t, int64(4), onlineNode[16]) // unique visitors
|
||||
|
||||
pendingNode := nodeByID["node-dashboard-2"]
|
||||
require.NotNil(t, pendingNode)
|
||||
@@ -183,6 +184,6 @@ func TestGetOverviewStructure(t *testing.T) {
|
||||
assert.Equal(t, "pending", pendingNode[6])
|
||||
assert.Equal(t, "unknown", pendingNode[7])
|
||||
|
||||
assert.Equal(t, 55.0, overview.Capacity.AverageCPUUsagePercent)
|
||||
assert.InDelta(t, 55.0, overview.Capacity.AverageCPUUsagePercent, 1e-9)
|
||||
assert.Equal(t, 1, overview.Traffic.ReportedNodes)
|
||||
}
|
||||
|
||||
@@ -28,7 +28,7 @@ const (
|
||||
// Heartbeat processes an OpenFlared heartbeat and returns runtime settings.
|
||||
func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("tunnel client node is nil")
|
||||
return nil, errors.New("tunnel client node is nil")
|
||||
}
|
||||
if node.NodeType != "tunnel_client" {
|
||||
return nil, fmt.Errorf("node %s is not a tunnel_client", node.NodeID)
|
||||
@@ -95,7 +95,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
|
||||
// GetTunnelConfig builds the full tunnel routing config for an OpenFlared client.
|
||||
func GetTunnelConfig(ctx context.Context, node *model.OpenFlareNode) (*TunnelConfigResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("node is nil")
|
||||
return nil, errors.New("node is nil")
|
||||
}
|
||||
|
||||
activeVersion, err := getActiveConfigMeta(ctx)
|
||||
|
||||
@@ -37,7 +37,11 @@ func PostHeartbeat(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := Heartbeat(c.Request.Context(), node, payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
@@ -63,7 +67,11 @@ func GetActiveConfig(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
config, err := GetTunnelConfig(c.Request.Context(), node)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
@@ -91,7 +99,9 @@ func PostApplyLog(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
if authNode, ok := c.Get(ctxFlaredNodeKey); ok {
|
||||
payload.NodeID = authNode.(*model.OpenFlareNode).NodeID
|
||||
if node, ok := authNode.(*model.OpenFlareNode); ok {
|
||||
payload.NodeID = node.NodeID
|
||||
}
|
||||
}
|
||||
|
||||
log, err := ReportApplyLog(c.Request.Context(), payload)
|
||||
@@ -115,6 +125,10 @@ func GetWebSocket(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
ofws.ServeFlared(c, node.NodeID)
|
||||
}
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package data embeds the MaxMind GeoLite2 Country database for the control plane.
|
||||
//
|
||||
// Server keeps a Country-only embed so MaxMind provider can seed without network.
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package data
|
||||
|
||||
import (
|
||||
|
||||
@@ -26,7 +26,7 @@ const (
|
||||
|
||||
var (
|
||||
runtimeOnce sync.Once
|
||||
runtimeInitErr error
|
||||
errRuntimeInit error
|
||||
currentProviderMu sync.RWMutex
|
||||
currentProvider string
|
||||
)
|
||||
@@ -34,9 +34,9 @@ var (
|
||||
// EnsureRuntimeProvider loads GeoIP provider config from SystemConfig.
|
||||
func EnsureRuntimeProvider(ctx context.Context) error {
|
||||
runtimeOnce.Do(func() {
|
||||
runtimeInitErr = applyProviderFromSystemConfig(ctx)
|
||||
errRuntimeInit = applyProviderFromSystemConfig(ctx)
|
||||
})
|
||||
return runtimeInitErr
|
||||
return errRuntimeInit
|
||||
}
|
||||
|
||||
// RefreshRuntimeProvider reapplies GeoIPProvider after config updates.
|
||||
@@ -91,11 +91,12 @@ func ensureServerMMDB() (string, error) {
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
_, statErr := os.Stat(path)
|
||||
if statErr == nil {
|
||||
return path, nil
|
||||
}
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return "", err
|
||||
if !os.IsNotExist(statErr) {
|
||||
return "", statErr
|
||||
}
|
||||
|
||||
// Control plane: seed Country from embedded asset (no City; Agent uses disk/image).
|
||||
@@ -115,7 +116,7 @@ func ensureServerMMDB() (string, error) {
|
||||
// ResetRuntimeForTest clears lazy-init state for unit tests.
|
||||
func ResetRuntimeForTest() {
|
||||
runtimeOnce = sync.Once{}
|
||||
runtimeInitErr = nil
|
||||
errRuntimeInit = nil
|
||||
currentProviderMu.Lock()
|
||||
currentProvider = ""
|
||||
currentProviderMu.Unlock()
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package geoip
|
||||
|
||||
import (
|
||||
|
||||
@@ -171,7 +171,7 @@ func TestCoreChainMigrationFlow(t *testing.T) {
|
||||
zoneDomains := data["zone_domains"].([]any)
|
||||
assert.Len(t, zoneDomains, 1)
|
||||
assert.Equal(t, "core-chain.example.com", zoneDomains[0].(map[string]any)["domain"])
|
||||
assert.Equal(t, float64(originID), data["origin_id"])
|
||||
assert.InDelta(t, float64(originID), data["origin_id"], 1e-9)
|
||||
assert.Equal(t, "http://origin.core-chain.internal:8080", data["origin_url"])
|
||||
})
|
||||
|
||||
@@ -257,7 +257,7 @@ func TestCoreChainMigrationFlow(t *testing.T) {
|
||||
|
||||
listResp := requireAPIOK(t, listRec)
|
||||
listData := unmarshalAPIMap(t, listResp.Data)
|
||||
assert.Equal(t, float64(1), listData["total"])
|
||||
assert.InDelta(t, float64(1), listData["total"], 1e-9)
|
||||
|
||||
rows, ok := listData["rows"].([]any)
|
||||
require.True(t, ok)
|
||||
@@ -278,10 +278,10 @@ func TestCoreChainMigrationFlow(t *testing.T) {
|
||||
require.Len(t, nodes, 1)
|
||||
nodeView, ok := nodes[0].(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, float64(nodeID), nodeView["id"])
|
||||
assert.InDelta(t, float64(nodeID), nodeView["id"], 1e-9)
|
||||
assert.Equal(t, nodePublicID, nodeView["node_id"])
|
||||
assert.Equal(t, "success", nodeView["latest_apply_result"])
|
||||
assert.Equal(t, configChecksum, nodeView["latest_apply_checksum"])
|
||||
assert.Equal(t, float64(2), nodeView["latest_support_file_count"])
|
||||
assert.InDelta(t, float64(2), nodeView["latest_support_file_count"], 1e-9)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -119,7 +119,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
assert.NotZero(t, ruleGroupID)
|
||||
assert.Equal(t, "edge-security", data["name"])
|
||||
assert.Equal(t, false, data["is_global"])
|
||||
assert.Equal(t, float64(1), data["revision"])
|
||||
assert.InDelta(t, float64(1), data["revision"], 1e-9)
|
||||
assert.NotNil(t, data["graph"])
|
||||
})
|
||||
|
||||
@@ -161,7 +161,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.Equal(t, float64(ruleGroupID), data["id"])
|
||||
assert.InDelta(t, float64(ruleGroupID), data["id"], 1e-9)
|
||||
assert.Equal(t, "edge-security", data["name"])
|
||||
})
|
||||
|
||||
@@ -244,12 +244,12 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
|
||||
resp := requireAPIOK(t, rec)
|
||||
data := unmarshalAPIMap(t, resp.Data)
|
||||
assert.Equal(t, float64(proxyRouteID), data["route_id"])
|
||||
assert.InDelta(t, float64(proxyRouteID), data["route_id"], 1e-9)
|
||||
|
||||
appliedIDs, ok := data["applied_ids"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Len(t, appliedIDs, 1)
|
||||
assert.Equal(t, float64(ruleGroupID), appliedIDs[0])
|
||||
assert.InDelta(t, float64(ruleGroupID), appliedIDs[0], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("verify site rule groups binding", func(t *testing.T) {
|
||||
@@ -272,7 +272,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
require.Len(t, appliedGroups, 1)
|
||||
group, ok := appliedGroups[0].(map[string]any)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, float64(ruleGroupID), group["id"])
|
||||
assert.InDelta(t, float64(ruleGroupID), group["id"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("create TLS certificate with PEM", func(t *testing.T) {
|
||||
@@ -319,7 +319,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
|
||||
domainID = uint(data["id"].(float64))
|
||||
assert.NotZero(t, domainID)
|
||||
assert.Equal(t, "security.example.com", data["domain"])
|
||||
assert.Equal(t, float64(certID), data["cert_id"])
|
||||
assert.InDelta(t, float64(certID), data["cert_id"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("create DNS account", func(t *testing.T) {
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
@@ -308,6 +309,8 @@ func fetchLatestGitHubRelease(ctx context.Context, repo string, channel releaseC
|
||||
switch normalizeReleaseChannel(string(channel)) {
|
||||
case releaseChannelPreview:
|
||||
return fetchLatestPreviewGitHubRelease(ctx, repo)
|
||||
case releaseChannelStable:
|
||||
return fetchLatestStableGitHubRelease(ctx, repo)
|
||||
default:
|
||||
return fetchLatestStableGitHubRelease(ctx, repo)
|
||||
}
|
||||
@@ -317,11 +320,11 @@ func fetchLatestStableGitHubRelease(ctx context.Context, repo string) (*githubRe
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"/latest", strings.TrimSpace(repo))
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建更新请求失败")
|
||||
return nil, errors.New("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取最新版本失败: %v", err)
|
||||
return nil, fmt.Errorf("获取最新版本失败: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
@@ -334,11 +337,11 @@ func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubR
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"?per_page=20", strings.TrimSpace(repo))
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建更新请求失败")
|
||||
return nil, errors.New("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取 preview 版本失败: %v", err)
|
||||
return nil, fmt.Errorf("获取 preview 版本失败: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
@@ -346,7 +349,7 @@ func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubR
|
||||
}
|
||||
var releases []githubReleaseResponse
|
||||
if err = json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
return nil, fmt.Errorf("解析 preview 版本信息失败")
|
||||
return nil, errors.New("解析 preview 版本信息失败")
|
||||
}
|
||||
for _, release := range releases {
|
||||
if release.Draft || !release.Prerelease {
|
||||
@@ -355,22 +358,22 @@ func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubR
|
||||
releaseCopy := release
|
||||
return &releaseCopy, nil
|
||||
}
|
||||
return nil, fmt.Errorf("当前没有可用的 preview 发布")
|
||||
return nil, errors.New("当前没有可用的 preview 发布")
|
||||
}
|
||||
|
||||
func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*githubReleaseResponse, error) {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
return nil, fmt.Errorf("缺少发布版本号")
|
||||
return nil, errors.New("缺少发布版本号")
|
||||
}
|
||||
url := fmt.Sprintf(githubReleasesAPIBase+"/tags/%s", strings.TrimSpace(repo), tag)
|
||||
req, err := newGitHubReleaseRequest(ctx, url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建更新请求失败")
|
||||
return nil, errors.New("创建更新请求失败")
|
||||
}
|
||||
resp, err := releaseHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取指定版本失败: %v", err)
|
||||
return nil, fmt.Errorf("获取指定版本失败: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
@@ -395,7 +398,7 @@ func newGitHubReleaseRequest(ctx context.Context, url string) (*http.Request, er
|
||||
func decodeGitHubRelease(reader io.Reader) (*githubReleaseResponse, error) {
|
||||
var release githubReleaseResponse
|
||||
if err := json.NewDecoder(reader).Decode(&release); err != nil {
|
||||
return nil, fmt.Errorf("解析版本信息失败")
|
||||
return nil, errors.New("解析版本信息失败")
|
||||
}
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
@@ -361,7 +361,7 @@ func RequestForceSync(ctx context.Context, id uint) (*View, error) {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, fmt.Errorf("无法获取当前激活的配置版本:%s", errNoActiveConfigVersion)
|
||||
}
|
||||
return nil, fmt.Errorf("无法获取当前激活的配置版本:%v", err)
|
||||
return nil, fmt.Errorf("无法获取当前激活的配置版本:%w", err)
|
||||
}
|
||||
if !ofws.SendForceSyncConfig(node.NodeID, forceSyncConfigPayload{
|
||||
Version: activeConfig.Version,
|
||||
@@ -414,14 +414,14 @@ type forceSyncConfigPayload struct {
|
||||
func ValidateDiscoveryToken(ctx context.Context, token string) error {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return fmt.Errorf("缺少 Discovery Token")
|
||||
return errors.New("缺少 Discovery Token")
|
||||
}
|
||||
discoveryToken, err := ensureGlobalDiscoveryToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if token != discoveryToken {
|
||||
return fmt.Errorf("discovery Token 无效") // error 消息首字母小写
|
||||
return errors.New("discovery Token 无效") // error 消息首字母小写
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package observability provides monitoring, metrics, and access log analysis for OpenFlare.
|
||||
package observability
|
||||
|
||||
@@ -175,7 +178,7 @@ type AccessLogIPSummaryList struct {
|
||||
TotalIP int64 `json:"total_ip"`
|
||||
Hours int `json:"hours"`
|
||||
Since time.Time `json:"since"`
|
||||
Until time.Time `json:"until,omitempty"`
|
||||
Until time.Time `json:"until,omitzero"`
|
||||
SortBy string `json:"sort_by"`
|
||||
SortOrder string `json:"sort_order"`
|
||||
}
|
||||
|
||||
@@ -108,7 +108,7 @@ func updateOptions(ctx context.Context, options []model.OpenFlareOption) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) {
|
||||
func getStatus(ctx context.Context, baseAPIPath string) *statusView {
|
||||
authSources, err := publicAuthSources(ctx, baseAPIPath)
|
||||
if err != nil {
|
||||
authSources = []publicAuthSourceView{}
|
||||
@@ -128,7 +128,7 @@ func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) {
|
||||
PasswordRegisterEnabled: passwordRegisterEnabled,
|
||||
CapLoginEnabled: capLoginEnabled,
|
||||
AuthSources: authSources,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func publicAuthSources(ctx context.Context, baseAPIPath string) ([]publicAuthSourceView, error) {
|
||||
|
||||
@@ -189,7 +189,7 @@ func validateOpenRestyCacheUseStale(key, trimmed string) error {
|
||||
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
|
||||
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
|
||||
}
|
||||
for _, token := range strings.Fields(trimmed) {
|
||||
for token := range strings.FieldsSeq(trimmed) {
|
||||
if _, ok := allowedTokens[token]; !ok {
|
||||
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
|
||||
}
|
||||
@@ -239,7 +239,7 @@ func validateOriginErrorPageStatusCodes(key, trimmed string) error {
|
||||
}
|
||||
codes, err := openrestyrender.ExpandStatusCodeTags(tags)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %v", key, err)
|
||||
return fmt.Errorf("%s: %w", key, err)
|
||||
}
|
||||
if len(codes) == 0 {
|
||||
return fmt.Errorf("%s 展开后不能为空", key)
|
||||
|
||||
@@ -22,10 +22,7 @@ import (
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/d/status [get]
|
||||
func GetStatusHandler(c *gin.Context) {
|
||||
view, err := getStatus(c.Request.Context(), "/api/v1/d")
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
return
|
||||
}
|
||||
view := getStatus(c.Request.Context(), "/api/v1/d")
|
||||
c.JSON(http.StatusOK, response.OK(view))
|
||||
}
|
||||
|
||||
|
||||
@@ -176,16 +176,16 @@ func validateUptimeKumaEnabled(ctx context.Context, key, trimmed string, state m
|
||||
username := strings.TrimSpace(state[model.ConfigKeyUptimeKumaUsername])
|
||||
password := strings.TrimSpace(state[model.ConfigKeyUptimeKumaPassword])
|
||||
if url == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时地址不能为空")
|
||||
return errors.New("启用 Uptime Kuma 时地址不能为空")
|
||||
}
|
||||
if username == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
return errors.New("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
// 如果待验证的密码为空,且当前配置中也没有密码,则报错
|
||||
if password == "" {
|
||||
existingPwd, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaPassword)
|
||||
if strings.TrimSpace(existingPwd.Value) == "" {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时密码不能为空")
|
||||
return errors.New("启用 Uptime Kuma 时密码不能为空")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
@@ -193,21 +193,21 @@ func validateUptimeKumaEnabled(ctx context.Context, key, trimmed string, state m
|
||||
|
||||
func validateUptimeKumaUsername(trimmed string, state map[string]string) error {
|
||||
if trimmed == "" && state[model.ConfigKeyUptimeKumaEnabled] == optionValueTrue {
|
||||
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
|
||||
return errors.New("启用 Uptime Kuma 时用户名不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUptimeKumaURL(trimmed string) error {
|
||||
if trimmed != "" && !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") {
|
||||
return fmt.Errorf("uptime Kuma 地址必须以 http:// 或 https:// 开头")
|
||||
return errors.New("uptime Kuma 地址必须以 http:// 或 https:// 开头")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUptimeKumaMonitorScope(trimmed string) error {
|
||||
if trimmed != "all" && trimmed != "selected" {
|
||||
return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)")
|
||||
return errors.New("监控范围必须为全部站点 (all) 或选择站点 (selected)")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -34,8 +34,8 @@ func validateOriginAddress(address string) error {
|
||||
if len(address) > maxOriginHostnameLength {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
for _, label := range labels {
|
||||
labels := strings.SplitSeq(address, ".")
|
||||
for label := range labels {
|
||||
if len(label) == 0 || len(label) > 63 {
|
||||
return errors.New(errOriginAddressInvalid)
|
||||
}
|
||||
|
||||
@@ -80,18 +80,11 @@ func resolvePagesLimits(ctx context.Context) pagesLimits {
|
||||
|
||||
historyCount := defaultPagesMaxHistoryCount
|
||||
if value, err := repository.GetIntByKey(ctx, model.ConfigKeyPagesMaxHistoryCount); err == nil {
|
||||
if value < 0 {
|
||||
historyCount = 0
|
||||
} else {
|
||||
historyCount = value
|
||||
}
|
||||
historyCount = max(value, 0)
|
||||
}
|
||||
|
||||
packageBytes := int64(packageMB) * bytesPerMiB
|
||||
extractedBytes := packageBytes * pagesExtractedSizeMultiplier
|
||||
if extractedBytes < pagesMinExtractedSizeBytes {
|
||||
extractedBytes = pagesMinExtractedSizeBytes
|
||||
}
|
||||
extractedBytes := max(packageBytes*pagesExtractedSizeMultiplier, pagesMinExtractedSizeBytes)
|
||||
|
||||
return pagesLimits{
|
||||
PackageBytes: packageBytes,
|
||||
@@ -154,7 +147,7 @@ func normalizePagesFallbackPath(raw string) (string, error) {
|
||||
return "", errors.New("spa fallback 回退路径不能包含空白或控制字符")
|
||||
}
|
||||
}
|
||||
for _, segment := range strings.Split(value, "/") {
|
||||
for segment := range strings.SplitSeq(value, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "", errors.New("spa fallback 回退路径不能包含 . 或 .. 路径段")
|
||||
}
|
||||
@@ -233,6 +226,8 @@ func safeTempSuffix(format pagesarchive.Format) string {
|
||||
return "7z"
|
||||
case pagesarchive.FormatTar:
|
||||
return "tar"
|
||||
case pagesarchive.FormatZip:
|
||||
return "zip"
|
||||
default:
|
||||
return "zip"
|
||||
}
|
||||
|
||||
@@ -513,7 +513,7 @@ func pruneProjectDeploymentHistory(ctx context.Context, projectID uint, keepCoun
|
||||
// Two passes: first pass after upload, second pass heals a concurrent race
|
||||
// that inserted another deployment between our list and delete.
|
||||
var lastErr error
|
||||
for pass := 0; pass < 2; pass++ {
|
||||
for range 2 {
|
||||
deleted, err := pruneProjectDeploymentHistoryOnce(ctx, projectID, keepCount, preserveCandidateID)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
|
||||
@@ -6,6 +6,7 @@ package pages
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"path"
|
||||
"strings"
|
||||
@@ -89,11 +90,11 @@ func rebindPagesRouteMaps(ctx context.Context, routes []map[string]json.RawMessa
|
||||
if route == nil {
|
||||
continue
|
||||
}
|
||||
upstreamType, _ := rawJSONString(route["upstream_type"])
|
||||
upstreamType := rawJSONString(route["upstream_type"])
|
||||
if !strings.EqualFold(strings.TrimSpace(upstreamType), "pages") {
|
||||
continue
|
||||
}
|
||||
siteName, _ := rawJSONString(route["site_name"])
|
||||
siteName := rawJSONString(route["site_name"])
|
||||
projectID, err := resolveProjectIDFromRouteMap(route)
|
||||
if err != nil {
|
||||
if siteName == "" {
|
||||
@@ -151,7 +152,7 @@ func resolveProjectIDFromRouteMap(route map[string]json.RawMessage) (uint, error
|
||||
return deployment.ProjectID, nil
|
||||
}
|
||||
}
|
||||
return 0, fmt.Errorf("pages 配置无效: 缺少 pages_project_id")
|
||||
return 0, errors.New("pages 配置无效: 缺少 pages_project_id")
|
||||
}
|
||||
|
||||
func loadActivePagesProject(ctx context.Context, projectID uint, siteName string) (*model.PagesProject, *model.PagesDeployment, error) {
|
||||
@@ -224,15 +225,15 @@ func buildLivePagesDeployment(
|
||||
}, nil
|
||||
}
|
||||
|
||||
func rawJSONString(raw json.RawMessage) (string, bool) {
|
||||
func rawJSONString(raw json.RawMessage) string {
|
||||
if !isPresentJSON(raw) {
|
||||
return "", false
|
||||
return ""
|
||||
}
|
||||
var value string
|
||||
if err := json.Unmarshal(raw, &value); err != nil {
|
||||
return "", false
|
||||
return ""
|
||||
}
|
||||
return value, true
|
||||
return value
|
||||
}
|
||||
|
||||
func putJSON(route map[string]json.RawMessage, key string, value any) error {
|
||||
@@ -245,5 +246,5 @@ func putJSON(route map[string]json.RawMessage, key string, value any) error {
|
||||
}
|
||||
|
||||
func errorsIsNotFound(err error) bool {
|
||||
return err != nil && (err == gorm.ErrRecordNotFound || strings.Contains(strings.ToLower(err.Error()), "record not found"))
|
||||
return err != nil && (errors.Is(err, gorm.ErrRecordNotFound) || strings.Contains(strings.ToLower(err.Error()), "record not found"))
|
||||
}
|
||||
|
||||
@@ -111,6 +111,8 @@ func (summary *PagesOrphanCleanupSummary) add(outcome pagesOrphanCleanupOutcome)
|
||||
summary.LeaseBusy++
|
||||
case pagesOrphanCleanupInvalidMarker:
|
||||
summary.InvalidMarker++
|
||||
case pagesOrphanCleanupSkipped:
|
||||
summary.Skipped++
|
||||
default:
|
||||
summary.Skipped++
|
||||
}
|
||||
|
||||
@@ -266,6 +266,8 @@ func scanOneDueGitHubSource(
|
||||
case sourceLeaseStale:
|
||||
summary.StaleSources++
|
||||
return
|
||||
case sourceLeaseAcquired:
|
||||
// 获取执行权成功,继续执行扫描。
|
||||
}
|
||||
if snapshot == nil || snapshot.SourceType != PagesSourceTypeGitHubRelease ||
|
||||
snapshot.ReleaseSelector != githubReleaseSelectorLatest {
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path"
|
||||
"sort"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -611,7 +611,7 @@ func lockSourceDeploymentUploadsTx(
|
||||
if hasIngest && ingestResult.Upload.ID != 0 && ingestResult.Upload.ID != target.UploadID {
|
||||
uploadIDs = append(uploadIDs, ingestResult.Upload.ID)
|
||||
}
|
||||
sort.Slice(uploadIDs, func(i, j int) bool { return uploadIDs[i] < uploadIDs[j] })
|
||||
slices.Sort(uploadIDs)
|
||||
var records []model.Upload
|
||||
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
|
||||
Where("id IN ?", uploadIDs).
|
||||
|
||||
@@ -154,10 +154,8 @@ func (h *SourceActionHandler) Execute(ctx context.Context, payload []byte) (*tas
|
||||
if source.SourceType != PagesSourceTypeRemoteURL && source.SourceType != PagesSourceTypeGitHubRelease {
|
||||
return nil, task.PermanentError(errPagesSourceTypeUnsupported)
|
||||
}
|
||||
if source.SourceType == PagesSourceTypeRemoteURL && (input.TargetRevision != "" || input.ConfirmedRevision != "") {
|
||||
return nil, task.PermanentError(errPagesSourceActionInvalid)
|
||||
}
|
||||
if input.Action == sourceActionCheck && (input.TargetRevision != "" || input.ConfirmedRevision != "") {
|
||||
if (source.SourceType == PagesSourceTypeRemoteURL || input.Action == sourceActionCheck) &&
|
||||
(input.TargetRevision != "" || input.ConfirmedRevision != "") {
|
||||
return nil, task.PermanentError(errPagesSourceActionInvalid)
|
||||
}
|
||||
|
||||
@@ -174,6 +172,8 @@ func (h *SourceActionHandler) Execute(ctx context.Context, payload []byte) (*tas
|
||||
case sourceLeaseStale:
|
||||
task.AppendLog(ctx, "[resolve] 来源配置或执行权已变化,本次任务跳过")
|
||||
return &task.TaskResult{Message: errPagesSourceActionStale}, nil
|
||||
case sourceLeaseAcquired:
|
||||
// 已获取执行权,继续执行。
|
||||
}
|
||||
|
||||
if input.Action == sourceActionCheck {
|
||||
|
||||
@@ -86,8 +86,8 @@ func validateOriginAddress(address string) error {
|
||||
if len(address) > maxOriginHostnameLength {
|
||||
return errors.New(errProxyRouteOriginInvalid)
|
||||
}
|
||||
labels := strings.Split(address, ".")
|
||||
for _, label := range labels {
|
||||
labels := strings.SplitSeq(address, ".")
|
||||
for label := range labels {
|
||||
if len(label) == 0 || len(label) > 63 {
|
||||
return errors.New(errProxyRouteOriginInvalid)
|
||||
}
|
||||
@@ -167,8 +167,8 @@ func buildOriginURLFromParts(scheme, address, port, uri string) (string, error)
|
||||
Host: formatOriginHost(normalizedAddress, normalizedPort),
|
||||
}
|
||||
if normalizedURI != "" {
|
||||
if strings.HasPrefix(normalizedURI, "?") {
|
||||
parsed.RawQuery = strings.TrimPrefix(normalizedURI, "?")
|
||||
if after, ok := strings.CutPrefix(normalizedURI, "?"); ok {
|
||||
parsed.RawQuery = after
|
||||
} else {
|
||||
pathQuery := strings.SplitN(normalizedURI, "?", originURIPathQueryParts)
|
||||
parsed.Path = pathQuery[0]
|
||||
|
||||
@@ -6,7 +6,7 @@ package proxy_route
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sort"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -120,7 +120,7 @@ func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
|
||||
|
||||
// CreateProxyRoute 创建代理规则。
|
||||
func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
|
||||
route, _, err := buildProxyRoute(ctx, nil, input)
|
||||
route, err := buildProxyRoute(ctx, nil, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -148,7 +148,7 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
|
||||
return nil, err
|
||||
}
|
||||
previousPagesProjectID := pagesProjectIDForRoute(route)
|
||||
route, _, err = buildProxyRoute(ctx, route, input)
|
||||
route, err = buildProxyRoute(ctx, route, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -201,7 +201,7 @@ func lockPagesProjectsForRouteMutation(tx *gorm.DB, previousProjectID uint, rout
|
||||
if nextProjectID != 0 && nextProjectID != previousProjectID {
|
||||
projectIDs = append(projectIDs, nextProjectID)
|
||||
}
|
||||
sort.Slice(projectIDs, func(i int, j int) bool { return projectIDs[i] < projectIDs[j] })
|
||||
slices.Sort(projectIDs)
|
||||
|
||||
for _, projectID := range projectIDs {
|
||||
project, err := repository.LockPagesProjectByIDTx(tx, projectID)
|
||||
@@ -244,67 +244,67 @@ func DeleteProxyRoute(ctx context.Context, id uint) error {
|
||||
return repository.DeleteProxyRouteAndUnbind(ctx, id)
|
||||
}
|
||||
|
||||
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, []model.ZoneDomain, error) {
|
||||
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, error) {
|
||||
domains, err := loadProxyRouteZoneDomains(ctx, input.ZoneDomainIDs)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
siteName := strings.TrimSpace(input.SiteName)
|
||||
|
||||
upstreamType := normalizeUpstreamType(input.UpstreamType)
|
||||
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
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, nil, err
|
||||
return nil, err
|
||||
}
|
||||
customHeaders, err := normalizeCustomHeaders(input.CustomHeaders)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
limitConnPerServer, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerServer, "limit_conn_per_server")
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
limitConnPerIP, err := normalizeProxyRouteLimitConnValue(input.LimitConnPerIP, "limit_conn_per_ip")
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
limitRate, err := normalizeProxyRouteLimitRate(input.LimitRate)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
limitReqPerIP, err := normalizeProxyRouteLimitReqPerIP(input.LimitReqPerIP)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
if err := validateProxyRouteZoneDomainCertificates(ctx, domains, input.EnableHTTPS); err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
jsonFields, err := marshalProxyRouteJSONFields(upstreams, cacheRules, customHeaders)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := validateProxyRouteSiteName(siteName); err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
if err := validateProxyRouteSiteNameUniqueness(ctx, route, siteName); err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
if err := validateOriginHost(originHost); err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
if input.RedirectHTTP && !input.EnableHTTPS {
|
||||
return nil, nil, errors.New(errProxyRouteRedirectHTTP)
|
||||
return nil, errors.New(errProxyRouteRedirectHTTP)
|
||||
}
|
||||
|
||||
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if route == nil {
|
||||
@@ -326,9 +326,9 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
|
||||
upstreamType,
|
||||
)
|
||||
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
|
||||
return nil, nil, err
|
||||
return nil, err
|
||||
}
|
||||
return route, domains, nil
|
||||
return route, nil
|
||||
}
|
||||
|
||||
func buildProxyRouteViews(ctx context.Context, routes []*model.ProxyRoute) ([]*View, error) {
|
||||
|
||||
@@ -136,7 +136,7 @@ func TestRouteCanMoveAwayFromAlreadyMissingPagesProject(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestNormalizeCachePolicyDefaultsAndLegacy(t *testing.T) {
|
||||
assert.Equal(t, "", normalizeCachePolicy(false, "static"))
|
||||
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"))
|
||||
@@ -144,7 +144,7 @@ func TestNormalizeCachePolicyDefaultsAndLegacy(t *testing.T) {
|
||||
assert.Equal(t, proxyRouteCachePolicyAll, normalizeCachePolicy(true, "all"))
|
||||
assert.Equal(t, proxyRouteCachePolicySuffix, normalizeCachePolicy(true, "suffix"))
|
||||
|
||||
assert.Equal(t, "", displayCachePolicy(false, "all"))
|
||||
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"))
|
||||
|
||||
@@ -5,6 +5,7 @@ package relay
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -20,7 +21,7 @@ const nodeStatusOnline = "online"
|
||||
// Heartbeat processes a relay heartbeat, updates node status, and returns config.
|
||||
func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("relay node is nil")
|
||||
return nil, errors.New("relay node is nil")
|
||||
}
|
||||
|
||||
payload.Version = strings.TrimSpace(payload.Version)
|
||||
|
||||
@@ -117,7 +117,7 @@ func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) {
|
||||
snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, snapshots, 1)
|
||||
assert.Equal(t, 12.5, snapshots[0].CPUUsagePercent)
|
||||
assert.InDelta(t, 12.5, snapshots[0].CPUUsagePercent, 1e-9)
|
||||
|
||||
frpsObs, err := repository.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -38,7 +38,11 @@ func PostHeartbeat(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := Heartbeat(c.Request.Context(), node, payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
@@ -61,6 +65,10 @@ func GetWebSocket(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errAgentTokenInvalid)
|
||||
return
|
||||
}
|
||||
ofws.ServeRelay(c, node.NodeID)
|
||||
}
|
||||
|
||||
@@ -8,17 +8,16 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupSSLRenewTestDB(t *testing.T) func() {
|
||||
@@ -26,18 +25,23 @@ func setupSSLRenewTestDB(t *testing.T) func() {
|
||||
|
||||
task.RegisterTaskMeta(tls.SSLSingleRenewMeta)
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqliteDB.AutoMigrate(&model.TLSCertificate{}, &model.TaskExecution{}))
|
||||
_, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
require.NoError(t, db.DB(nil).AutoMigrate(&model.TLSCertificate{}, &model.TaskExecution{}))
|
||||
|
||||
// task 包 init() 会按配置创建指向真实 Redis 的客户端;测试显式改用
|
||||
// miniredis(与 executor_test 一致),避免依赖本地 redis 实例。
|
||||
oldClient := task.AsynqClient
|
||||
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{Addr: mr.Addr()})
|
||||
t.Cleanup(func() {
|
||||
_ = task.AsynqClient.Close()
|
||||
task.AsynqClient = oldClient
|
||||
})
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
oldSecret := config.Config.App.SessionSecret
|
||||
config.Config.App.SessionSecret = "test_session_secret_for_ssl_renew"
|
||||
return func() {
|
||||
db.SetDB(nil)
|
||||
config.Config.App.SessionSecret = oldSecret
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -173,7 +173,7 @@ func SetupDNSProvider(client *lego.Client, dnsType, dnsAuth string, dns1, dns2 s
|
||||
case "cloudflare":
|
||||
var creds map[string]string
|
||||
if err := json.Unmarshal([]byte(dnsAuth), &creds); err != nil {
|
||||
return fmt.Errorf("failed to parse cloudflare credentials: %v", err)
|
||||
return fmt.Errorf("failed to parse cloudflare credentials: %w", err)
|
||||
}
|
||||
|
||||
config := cloudflare.NewDefaultConfig()
|
||||
|
||||
@@ -5,7 +5,6 @@ package tls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -51,10 +50,19 @@ func TestApplyCertificateReturnsApplying(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
obtainDone := make(chan struct{})
|
||||
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, cert *model.TLSCertificate) error {
|
||||
defer close(obtainDone)
|
||||
return updateCertError(ctx, cert, "dns challenge failed")
|
||||
})
|
||||
defer restore()
|
||||
defer func() {
|
||||
select {
|
||||
case <-obtainDone:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("async certificate obtain did not finish")
|
||||
}
|
||||
restore()
|
||||
}()
|
||||
|
||||
cert, err := ApplyCertificate(ctx, ApplyInput{
|
||||
Name: "Test ACME Cert",
|
||||
@@ -126,7 +134,7 @@ func TestConvertCertificateToACMEPreservesUploadOnFailure(t *testing.T) {
|
||||
assert.Equal(t, "error", finalCert.ApplyStatus)
|
||||
assert.Equal(t, originalStoredCertPEM, finalCert.CertPEM)
|
||||
assert.Equal(t, originalStoredKeyPEM, finalCert.KeyPEM)
|
||||
assert.True(t, strings.Contains(finalCert.ApplyMessage, "dns challenge failed"))
|
||||
assert.Contains(t, finalCert.ApplyMessage, "dns challenge failed")
|
||||
}
|
||||
|
||||
func TestConvertCertificateToACMERejectsInvalidStates(t *testing.T) {
|
||||
|
||||
@@ -197,12 +197,17 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 先取响应快照再启动异步续签:sanitize 会整体拷贝 cert,若与异步 goroutine
|
||||
// 的字段写入并发会构成数据竞争(生产真实问题)。
|
||||
returned := sanitizeCertificateForResponse(cert)
|
||||
|
||||
obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换)
|
||||
go func(c *model.TLSCertificate) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
_ = obtainTLSCertificate(asyncCtx, c)
|
||||
_ = obtainFn(asyncCtx, c)
|
||||
}(cert)
|
||||
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
return returned, nil
|
||||
}
|
||||
|
||||
// UpdateACMECertificate 更新 ACME 证书配置。
|
||||
@@ -225,12 +230,15 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod
|
||||
return nil, err
|
||||
}
|
||||
|
||||
returned := sanitizeCertificateForResponse(cert)
|
||||
|
||||
obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换)
|
||||
go func(c *model.TLSCertificate) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
_ = obtainTLSCertificate(asyncCtx, c)
|
||||
_ = obtainFn(asyncCtx, c)
|
||||
}(cert)
|
||||
|
||||
return sanitizeCertificateForResponse(cert), nil
|
||||
return returned, nil
|
||||
}
|
||||
|
||||
// ConvertCertificateToACME 将上传证书转为 ACME 管理。
|
||||
@@ -257,9 +265,10 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*
|
||||
return nil, err
|
||||
}
|
||||
|
||||
obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换)
|
||||
go func(c *model.TLSCertificate) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
if err := obtainTLSCertificate(asyncCtx, c); err != nil {
|
||||
if err := obtainFn(asyncCtx, c); err != nil {
|
||||
return
|
||||
}
|
||||
latest, err := repository.GetTLSCertificateByID(asyncCtx, c.ID)
|
||||
|
||||
@@ -123,7 +123,7 @@ func splitAcmeDomains(primaryDomain, otherDomains string) []string {
|
||||
if !strings.Contains(otherDomains, "\n") && strings.Contains(otherDomains, ",") {
|
||||
separator = ","
|
||||
}
|
||||
for _, domain := range strings.Split(otherDomains, separator) {
|
||||
for domain := range strings.SplitSeq(otherDomains, separator) {
|
||||
domain = strings.TrimSpace(domain)
|
||||
if domain != "" {
|
||||
domains = append(domains, domain)
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"maps"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -88,7 +89,7 @@ func NewSocketIOClient(baseURL string) *SocketIOClient {
|
||||
// Connect performs the Engine.IO handshake and starts the polling loop.
|
||||
func (c *SocketIOClient) Connect() error {
|
||||
slog.Debug("Uptime Kuma client starting handshake", "baseURL", c.baseURL)
|
||||
u := fmt.Sprintf("%s/socket.io/?EIO=4&transport=polling", c.baseURL)
|
||||
u := c.baseURL + "/socket.io/?EIO=4&transport=polling"
|
||||
reqHandshake, err := http.NewRequestWithContext(c.ctx, http.MethodGet, u, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create handshake request failed: %w", err)
|
||||
@@ -180,8 +181,8 @@ func (c *SocketIOClient) pollLoop() {
|
||||
}
|
||||
|
||||
slog.Debug("Received polling payload from Uptime Kuma", "length", len(bodyStr))
|
||||
packets := strings.Split(bodyStr, "\x1e")
|
||||
for _, pkt := range packets {
|
||||
packets := strings.SplitSeq(bodyStr, "\x1e")
|
||||
for pkt := range packets {
|
||||
if len(pkt) == 0 {
|
||||
continue
|
||||
}
|
||||
@@ -280,7 +281,8 @@ func (c *SocketIOClient) Emit(event string, args ...any) (string, error) {
|
||||
c.ackChanMap[id] = ch
|
||||
c.ackMutex.Unlock()
|
||||
|
||||
payloadArr := []any{event}
|
||||
payloadArr := make([]any, 1, 1+len(args))
|
||||
payloadArr[0] = event
|
||||
payloadArr = append(payloadArr, args...)
|
||||
bs, err := json.Marshal(payloadArr)
|
||||
if err != nil {
|
||||
@@ -352,9 +354,7 @@ func (c *SocketIOClient) GetMonitorList() map[string]Monitor {
|
||||
defer c.monitorListMutex.RUnlock()
|
||||
|
||||
m := make(map[string]Monitor, len(c.monitorList))
|
||||
for k, v := range c.monitorList {
|
||||
m[k] = v
|
||||
}
|
||||
maps.Copy(m, c.monitorList)
|
||||
return m
|
||||
}
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ package uptimekuma
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
@@ -32,7 +33,7 @@ type kumaConfig struct {
|
||||
}
|
||||
|
||||
// loadKumaConfig 从 SystemConfig 加载 UptimeKuma 配置
|
||||
func loadKumaConfig(ctx context.Context) (*kumaConfig, error) {
|
||||
func loadKumaConfig(ctx context.Context) *kumaConfig {
|
||||
url, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaURL)
|
||||
username, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaUsername)
|
||||
password, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUptimeKumaPassword)
|
||||
@@ -67,7 +68,7 @@ func loadKumaConfig(ctx context.Context) (*kumaConfig, error) {
|
||||
Retry: retry,
|
||||
RetryInterval: retryInterval,
|
||||
Timeout: timeout,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
// SyncToUptimeKuma synchronizes enabled proxy routes to Uptime Kuma monitors.
|
||||
@@ -75,19 +76,16 @@ func SyncToUptimeKuma(ctx context.Context) error {
|
||||
// 检查是否启用
|
||||
enabled, _ := repository.GetBoolByKey(ctx, model.ConfigKeyUptimeKumaEnabled)
|
||||
if !enabled {
|
||||
return fmt.Errorf("uptime Kuma integration is disabled")
|
||||
return errors.New("uptime Kuma integration is disabled")
|
||||
}
|
||||
|
||||
if !isSyncing.CompareAndSwap(false, true) {
|
||||
return fmt.Errorf("sync task is already in progress, please try again later")
|
||||
return errors.New("sync task is already in progress, please try again later")
|
||||
}
|
||||
defer isSyncing.Store(false)
|
||||
|
||||
// 加载配置
|
||||
config, err := loadKumaConfig(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
config := loadKumaConfig(ctx)
|
||||
|
||||
// 验证配置
|
||||
if err := validateKumaConfig(config); err != nil {
|
||||
@@ -105,10 +103,7 @@ func SyncToUptimeKuma(ctx context.Context) error {
|
||||
return fmt.Errorf("failed to list local proxy routes: %w", err)
|
||||
}
|
||||
|
||||
expectedRoutes, err := filterExpectedRoutes(allRoutes, config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
expectedRoutes := filterExpectedRoutes(allRoutes, config)
|
||||
|
||||
client, err := connectAndLoginUptimeKuma(config.URL, config.Username, config.Password)
|
||||
if err != nil {
|
||||
@@ -128,7 +123,7 @@ func SyncToUptimeKuma(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func filterExpectedRoutes(allRoutes []*model.ProxyRoute, config *kumaConfig) ([]*model.ProxyRoute, error) {
|
||||
func filterExpectedRoutes(allRoutes []*model.ProxyRoute, config *kumaConfig) []*model.ProxyRoute {
|
||||
scope := config.MonitorScope
|
||||
if scope == "selected" {
|
||||
selectedList := strings.Split(config.SelectedSites, ",")
|
||||
@@ -145,7 +140,7 @@ func filterExpectedRoutes(allRoutes []*model.ProxyRoute, config *kumaConfig) ([]
|
||||
expectedRoutes = append(expectedRoutes, route)
|
||||
}
|
||||
}
|
||||
return expectedRoutes, nil
|
||||
return expectedRoutes
|
||||
}
|
||||
|
||||
var expectedRoutes []*model.ProxyRoute
|
||||
@@ -154,7 +149,7 @@ func filterExpectedRoutes(allRoutes []*model.ProxyRoute, config *kumaConfig) ([]
|
||||
expectedRoutes = append(expectedRoutes, route)
|
||||
}
|
||||
}
|
||||
return expectedRoutes, nil
|
||||
return expectedRoutes
|
||||
}
|
||||
|
||||
func ensureOpenFlareTag(client *SocketIOClient) (int, error) {
|
||||
@@ -225,7 +220,7 @@ func filterOpenFlareMonitors(monitors map[string]Monitor, openFlareTagID int) ma
|
||||
|
||||
func routeMonitorURL(ctx context.Context, route *model.ProxyRoute) (string, error) {
|
||||
if route == nil {
|
||||
return "", fmt.Errorf("proxy route is nil")
|
||||
return "", errors.New("proxy route is nil")
|
||||
}
|
||||
domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
|
||||
if err != nil {
|
||||
|
||||
@@ -5,6 +5,7 @@ package uptimekuma
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
@@ -18,13 +19,13 @@ const monitorListWaitTimeout = 5 * time.Second
|
||||
// validateKumaConfig 验证 kumaConfig 配置完整性
|
||||
func validateKumaConfig(config *kumaConfig) error {
|
||||
if strings.TrimSpace(config.URL) == "" {
|
||||
return fmt.Errorf("uptime Kuma URL is not configured")
|
||||
return errors.New("uptime Kuma URL is not configured")
|
||||
}
|
||||
if strings.TrimSpace(config.Username) == "" {
|
||||
return fmt.Errorf("uptime Kuma username is not configured")
|
||||
return errors.New("uptime Kuma username is not configured")
|
||||
}
|
||||
if strings.TrimSpace(config.Password) == "" {
|
||||
return fmt.Errorf("uptime Kuma password is not configured")
|
||||
return errors.New("uptime Kuma password is not configured")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -65,7 +66,7 @@ func connectAndLoginUptimeKuma(kumaURL, kumaUsername, kumaPassword string) (*Soc
|
||||
case <-time.After(monitorListWaitTimeout):
|
||||
client.Close()
|
||||
slog.Error("Timeout waiting for Uptime Kuma monitorList push event")
|
||||
return nil, fmt.Errorf("timeout waiting for monitorList event from Uptime Kuma")
|
||||
return nil, errors.New("timeout waiting for monitorList event from Uptime Kuma")
|
||||
}
|
||||
return client, nil
|
||||
}
|
||||
|
||||
@@ -4,7 +4,9 @@
|
||||
package waf
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
)
|
||||
|
||||
@@ -34,14 +36,14 @@ func CompileRuleGraph(graph RuleGraph) (RuntimeRuleGraph, error) {
|
||||
}
|
||||
if node.Type == RuleNodeStart {
|
||||
if runtime.Entry != "" {
|
||||
return RuntimeRuleGraph{}, fmt.Errorf("规则图包含多个开始节点")
|
||||
return RuntimeRuleGraph{}, errors.New("规则图包含多个开始节点")
|
||||
}
|
||||
runtime.Entry = node.ID
|
||||
}
|
||||
runtime.Nodes[node.ID] = RuntimeRuleNode{Type: node.Type, Config: config}
|
||||
}
|
||||
if runtime.Entry == "" {
|
||||
return RuntimeRuleGraph{}, fmt.Errorf("规则图缺少开始节点")
|
||||
return RuntimeRuleGraph{}, errors.New("规则图缺少开始节点")
|
||||
}
|
||||
for _, edge := range graph.Edges {
|
||||
node, ok := runtime.Nodes[edge.Source]
|
||||
@@ -142,7 +144,7 @@ func sortedUniqueStrings(values []string) []string {
|
||||
|
||||
func sortedUniqueUints(values []uint) []uint {
|
||||
result := append([]uint(nil), values...)
|
||||
sort.Slice(result, func(i, j int) bool { return result[i] < result[j] })
|
||||
slices.Sort(result)
|
||||
write := 0
|
||||
for _, value := range result {
|
||||
if write == 0 || result[write-1] != value {
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"io"
|
||||
"net/netip"
|
||||
"regexp"
|
||||
"slices"
|
||||
"strings"
|
||||
)
|
||||
|
||||
@@ -67,7 +68,7 @@ func validateRuleGraphLimits(graph RuleGraph) error {
|
||||
if raw, err := json.Marshal(graph); err != nil {
|
||||
return fmt.Errorf("规则图无法序列化: %w", err)
|
||||
} else if len(raw) > maxRuleGraphBytes {
|
||||
return fmt.Errorf("规则图大小不能超过 256 KiB")
|
||||
return errors.New("规则图大小不能超过 256 KiB")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -374,12 +375,7 @@ func decodeStrictConfig(raw json.RawMessage, dst any) error {
|
||||
}
|
||||
|
||||
func validSourceHandle(t RuleNodeType, handle string) bool {
|
||||
for _, expected := range requiredHandles(t) {
|
||||
if handle == expected {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
return slices.Contains(requiredHandles(t), handle)
|
||||
}
|
||||
func requiredHandles(t RuleNodeType) []string {
|
||||
switch t {
|
||||
@@ -387,6 +383,8 @@ func requiredHandles(t RuleNodeType) []string {
|
||||
return []string{"next"}
|
||||
case RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodeUACheck, RuleNodeSecurityCheck:
|
||||
return []string{"true", "false"}
|
||||
case RuleNodeAllow, RuleNodeBlock:
|
||||
return nil
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"maps"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
@@ -325,9 +326,7 @@ func evaluateParsedIPGroupAutoConfig(ctx context.Context, config ipGroupAutoConf
|
||||
lastSeen = time.Unix(item.LastSeenEpoch, 0).UTC()
|
||||
}
|
||||
statusCounts := make(map[int]int, len(item.StatusCounts))
|
||||
for code, count := range item.StatusCounts {
|
||||
statusCounts[code] = count
|
||||
}
|
||||
maps.Copy(statusCounts, item.StatusCounts)
|
||||
accumulators[ip] = &ipGroupAutoAccumulator{
|
||||
ip: ip,
|
||||
requestCount: item.RequestCount,
|
||||
@@ -404,7 +403,7 @@ func downloadIPGroupSubscription(ctx context.Context, rawURL string) ([]byte, er
|
||||
return nil, err
|
||||
}
|
||||
client := http.Client{Timeout: 15 * time.Second}
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", rawURL, nil)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("下载订阅失败: %w", err)
|
||||
}
|
||||
@@ -479,7 +478,7 @@ func selectJSONMappingNodes(payload any, mappingRule string) ([]any, error) {
|
||||
}
|
||||
rule = strings.TrimPrefix(rule, "$.")
|
||||
nodes := []any{payload}
|
||||
for _, rawSegment := range strings.Split(rule, ".") {
|
||||
for rawSegment := range strings.SplitSeq(rule, ".") {
|
||||
segment := strings.TrimSpace(rawSegment)
|
||||
if segment == "" {
|
||||
continue
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
@@ -63,12 +62,14 @@ func TestRuleHandlersMapFailures(t *testing.T) {
|
||||
{name: "invalid id", method: http.MethodGet, path: "/rules/nope", setup: setupWAFTestDB, want: http.StatusBadRequest},
|
||||
{name: "malformed json", method: http.MethodPost, path: "/rules", body: `{`, setup: setupWAFTestDB, want: http.StatusBadRequest},
|
||||
{name: "invalid graph", method: http.MethodPost, path: "/rules/1/graph", body: `{"revision":1,"graph":{"schema_version":1,"nodes":[],"edges":[]}}`, setup: func(t *testing.T) func() {
|
||||
t.Helper()
|
||||
cleanup := setupWAFTestDB(t)
|
||||
_, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
|
||||
require.NoError(t, err)
|
||||
return cleanup
|
||||
}, want: http.StatusBadRequest},
|
||||
{name: "manual IP group sync", method: http.MethodPost, path: "/ip-groups/1/sync", setup: func(t *testing.T) func() {
|
||||
t.Helper()
|
||||
cleanup := setupWAFTestDB(t)
|
||||
_, err := CreateIPGroup(context.Background(), IPGroupInput{Name: "manual", Type: wafIPGroupTypeManual, Enabled: true})
|
||||
require.NoError(t, err)
|
||||
@@ -76,12 +77,13 @@ func TestRuleHandlersMapFailures(t *testing.T) {
|
||||
}, want: http.StatusBadRequest},
|
||||
{name: "missing", method: http.MethodGet, path: "/rules/999", setup: setupWAFTestDB, want: http.StatusNotFound},
|
||||
{name: "conflict", method: http.MethodPost, path: "/rules/1/graph", body: mustGraphRequest(t, 0), setup: func(t *testing.T) func() {
|
||||
t.Helper()
|
||||
cleanup := setupWAFTestDB(t)
|
||||
_, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
|
||||
require.NoError(t, err)
|
||||
return cleanup
|
||||
}, want: http.StatusConflict},
|
||||
{name: "database failure", method: http.MethodGet, path: "/rules", setup: func(t *testing.T) func() { db.SetDB(nil); return func() {} }, want: http.StatusInternalServerError},
|
||||
{name: "database failure", method: http.MethodGet, path: "/rules", setup: func(t *testing.T) func() { t.Helper(); db.SetDB(nil); return func() {} }, want: http.StatusInternalServerError},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
@@ -171,7 +173,7 @@ func TestReplaceSiteRuleGroupsPreservesOrderAndRejectsGlobal(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
_, err = ReplaceSiteRuleGroups(ctx, 7, []uint{global.ID, second.ID})
|
||||
require.Error(t, err)
|
||||
assert.False(t, errors.Is(err, model.ErrWAFRuleRevisionConflict))
|
||||
require.NotErrorIs(t, err, model.ErrWAFRuleRevisionConflict)
|
||||
assert.Equal(t, []uint{third.ID, first.ID, second.ID}, mustListSiteRuleGroupIDs(t, ctx, 7))
|
||||
}
|
||||
|
||||
|
||||
@@ -103,7 +103,7 @@ func ImportLegacyTx(ctx context.Context, tx *sql.Tx, postgres bool) (report Impo
|
||||
`), domain).Scan(&existingID, &existingZoneDomain)
|
||||
if scanErr == nil {
|
||||
if existingZoneDomain != root {
|
||||
report.Conflicts = append(report.Conflicts, fmt.Sprintf("%s: global domain conflict", domain))
|
||||
report.Conflicts = append(report.Conflicts, domain+": global domain conflict")
|
||||
} else if item.ProxyRouteID != nil {
|
||||
if _, bindErr := tx.ExecContext(ctx, q(`
|
||||
UPDATE of_zone_domains
|
||||
@@ -147,7 +147,7 @@ func ImportLegacyTx(ctx context.Context, tx *sql.Tx, postgres bool) (report Impo
|
||||
}
|
||||
|
||||
if len(report.Conflicts) > 0 {
|
||||
return report, fmt.Errorf("legacy data has conflicts")
|
||||
return report, errors.New("legacy data has conflicts")
|
||||
}
|
||||
return report, nil
|
||||
}
|
||||
@@ -292,7 +292,7 @@ func rebindSQL(query string, postgres bool) string {
|
||||
var b strings.Builder
|
||||
b.Grow(len(query) + len(query)/4)
|
||||
n := 0
|
||||
for i := 0; i < len(query); i++ {
|
||||
for i := range len(query) {
|
||||
if query[i] == '?' {
|
||||
n++
|
||||
b.WriteByte('$')
|
||||
|
||||
@@ -183,10 +183,7 @@ func emptyStatsSeries(since, until time.Time, bucketMinutes int) []StatsPoint {
|
||||
}
|
||||
// Cap points to keep chart readable.
|
||||
maxPoints := 120
|
||||
capacity := int(end.Sub(start)/bucket) + 1
|
||||
if capacity > maxPoints {
|
||||
capacity = maxPoints
|
||||
}
|
||||
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})
|
||||
|
||||
Reference in New Issue
Block a user