后端与全仓代码质量清理(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:
ryan
2026-08-16 21:23:37 +08:00
parent 5a8722ff07
commit 2f60329886
292 changed files with 1362 additions and 793 deletions
+3 -12
View File
@@ -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,
})
+3
View File
@@ -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
+1 -1
View File
@@ -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)
}
+2 -2
View File
@@ -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)
+18 -4
View File
@@ -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 (
+8 -7
View File
@@ -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) {
+13 -10
View File
@@ -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
}
+3 -3
View File
@@ -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"`
}
+2 -2
View File
@@ -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)
+1 -4
View File
@@ -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))
}
+6 -6
View File
@@ -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
}
+2 -2
View File
@@ -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)
}
+5 -10
View File
@@ -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"
}
+1 -1
View File
@@ -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
+9 -8
View File
@@ -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 {
+2 -2
View File
@@ -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]
+22 -22
View File
@@ -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"))
+2 -1
View File
@@ -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)
+1 -1
View File
@@ -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)
+10 -2
View File
@@ -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)
}
+15 -11
View File
@@ -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()
}
}
+1 -1
View File
@@ -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) {
+14 -5
View File
@@ -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)
+1 -1
View File
@@ -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)
+7 -7
View File
@@ -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
}
+11 -16
View File
@@ -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
}
+5 -3
View File
@@ -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
}
+4 -5
View File
@@ -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('$')
+1 -4
View File
@@ -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})