This commit is contained in:
ryan
2026-06-19 15:13:24 +08:00
parent 0b34792709
commit 32861c5db9
376 changed files with 3648 additions and 19957 deletions
@@ -1,6 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package agent implements the OpenFlare agent protocol: node registration,
// heartbeat processing, access-log ingestion, and related middleware.
package agent
import (
@@ -11,12 +13,12 @@ import (
pkggeoip "github.com/Rain-kl/Wavelet/pkg/geoip"
)
var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) {
var accessLogGeoProviderFactory = func() (pkggeoip.Service, error) {
return pkggeoip.NewMaxMindGeoIPService()
}
type accessLogRegionResolver struct {
provider pkggeoip.GeoIPService
provider pkggeoip.Service
cache map[string]string
}
+4 -6
View File
@@ -36,12 +36,10 @@ var tokenCache = newAccessTokenAuthCache()
func newAccessTokenAuthCache() *accessTokenAuthCache {
return &accessTokenAuthCache{
positive: make(map[string]cachedAgentNode),
negative: make(map[string]time.Time),
now: time.Now,
loadNodeByToken: func(ctx context.Context, token string) (*model.OpenFlareNode, error) {
return model.GetOpenFlareNodeByAccessToken(ctx, token)
},
positive: make(map[string]cachedAgentNode),
negative: make(map[string]time.Time),
now: time.Now,
loadNodeByToken: model.GetOpenFlareNodeByAccessToken,
}
}
+3 -3
View File
@@ -4,9 +4,9 @@
package agent
const (
errMissingAgentToken = "缺少 Agent Token"
errInvalidAgentToken = "无权进行此操作,Agent Token 无效"
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效"
errMissingAgentToken = "缺少 Agent Token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errInvalidAgentToken = "无权进行此操作,Agent Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errNodeMissingFromContext = "Node object missing from context"
errNoActiveConfig = "当前没有激活版本"
errNodeNotFound = "节点不存在"
+13 -11
View File
@@ -20,10 +20,12 @@ const (
openrestyStatusUnhealthy = "unhealthy"
openrestyStatusUnknown = "unknown"
releaseChannelStable = "stable"
randomTokenBytes = 16
maxDatabaseTextLength = 16000
)
func newRandomToken() (string, error) {
buf := make([]byte, 16)
buf := make([]byte, randomTokenBytes)
if _, err := rand.Read(buf); err != nil {
return "", err
}
@@ -55,9 +57,9 @@ func normalizeNodePayload(payload NodePayload) NodePayload {
payload.Version = strings.TrimSpace(payload.Version)
payload.ExtVersion = strings.TrimSpace(payload.ExtVersion)
payload.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
payload.LastError = truncateForDatabase(payload.LastError, 16000)
payload.LastError = truncateForDatabase(payload.LastError, maxDatabaseTextLength)
payload.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, maxDatabaseTextLength)
return payload
}
@@ -92,12 +94,12 @@ func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveNa
node.Version = strings.TrimSpace(payload.Version)
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, maxDatabaseTextLength)
node.Status = nodeStatusOnline
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
now := time.Now()
node.LastSeenAt = &now
node.LastError = truncateForDatabase(payload.LastError, 16000)
node.LastError = truncateForDatabase(payload.LastError, maxDatabaseTextLength)
if !node.GeoManualOverride {
applyGeoInfoFromIP(node, node.IP)
}
@@ -135,15 +137,15 @@ func cloneCoordinate(value *float64) *float64 {
return &cloned
}
func truncateForDatabase(value string, max int) string {
if max <= 0 {
func truncateForDatabase(value string, maxVal int) string {
if maxVal <= 0 {
return ""
}
runes := []rune(strings.TrimSpace(value))
if len(runes) <= max {
if len(runes) <= maxVal {
return string(runes)
}
return string(runes[:max])
return string(runes[:maxVal])
}
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
@@ -277,7 +279,7 @@ func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
payload.NodeID = strings.TrimSpace(payload.NodeID)
payload.Version = strings.TrimSpace(payload.Version)
payload.Result = strings.ToLower(strings.TrimSpace(payload.Result))
payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), 16000)
payload.Message = truncateForDatabase(strings.TrimSpace(payload.Message), maxDatabaseTextLength)
payload.Checksum = strings.TrimSpace(payload.Checksum)
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
@@ -292,7 +294,7 @@ func isUniqueConstraintError(err error) bool {
}
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
func RefreshAccessTokenCache(_ context.Context, node *model.OpenFlareNode) {
if node == nil {
return
}
+8 -8
View File
@@ -12,12 +12,12 @@ import (
)
const (
agentTokenHeader = "X-Agent-Token"
agentTokenHeader = "X-Agent-Token" //nolint:gosec // HTTP header name, not a credential value
agentNodeContextKey = "agent_node"
)
// AgentAuth validates X-Agent-Token against of_nodes.access_token.
func AgentAuth() gin.HandlerFunc {
// Auth validates X-Agent-Token against of_nodes.access_token.
func Auth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
node, err := AuthenticateAccessToken(c.Request.Context(), token)
@@ -30,8 +30,8 @@ func AgentAuth() gin.HandlerFunc {
}
}
// AgentRegisterAuth accepts either a node access token or the global discovery token.
func AgentRegisterAuth() gin.HandlerFunc {
// RegisterAuth accepts either a node access token or the global discovery token.
func RegisterAuth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil {
@@ -48,12 +48,12 @@ func AgentRegisterAuth() gin.HandlerFunc {
}
}
// AgentNodeFromContext returns the authenticated agent node.
func AgentNodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
// NodeFromContext returns the authenticated agent node.
func NodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
value, ok := c.Get(agentNodeContextKey)
if !ok {
return nil, false
}
node, ok := value.(*model.OpenFlareNode)
return node, ok
}
}
@@ -109,8 +109,8 @@ func TestAgentAuthMiddleware(t *testing.T) {
}).Error)
router := testhelper.NewTestGinEngine()
router.GET("/protected", AgentAuth(), func(c *gin.Context) {
node, ok := AgentNodeFromContext(c)
router.GET("/protected", Auth(), func(c *gin.Context) {
node, ok := NodeFromContext(c)
if !ok {
c.Status(http.StatusInternalServerError)
return
@@ -157,8 +157,8 @@ func TestAgentRegisterAuthMiddleware(t *testing.T) {
require.NoError(t, model.UpdateOpenFlareOption(ctx, "AgentDiscoveryToken", "discovery-token"))
router := testhelper.NewTestGinEngine()
router.POST("/register", AgentRegisterAuth(), func(c *gin.Context) {
if node, ok := AgentNodeFromContext(c); ok {
router.POST("/register", RegisterAuth(), func(c *gin.Context) {
if node, ok := NodeFromContext(c); ok {
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "node", "node_id": node.NodeID}))
return
}
@@ -196,4 +196,4 @@ func TestAgentRegisterAuthMiddleware(t *testing.T) {
require.True(t, ok)
assert.Equal(t, "discovery", data["mode"])
})
}
}
@@ -27,6 +27,7 @@ const (
nodeAccessLogRetentionDays = 90
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
accessLogPathMaxLength = 100
healthEventMessageMaxLength = 4096
)
// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events.
@@ -396,7 +397,7 @@ func normalizeHealthSeverity(severity string) string {
}
func normalizeHealthEventMessage(message string) string {
return truncateForDatabase(message, 4096)
return truncateForDatabase(message, healthEventMessageMaxLength)
}
func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
@@ -5,22 +5,55 @@ package agent
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
// NodePayload is the data sent by an agent on registration or heartbeat.
type NodePayload = pkgprotocol.NodePayload
// NodeSystemProfile carries static host information reported by an agent.
type NodeSystemProfile = pkgprotocol.NodeSystemProfile
// NodeMetricSnapshot holds a point-in-time resource-usage sample from an agent.
type NodeMetricSnapshot = pkgprotocol.NodeMetricSnapshot
// NodeOpenrestyObservation reports the OpenResty process health observed by an agent.
type NodeOpenrestyObservation = pkgprotocol.NodeOpenrestyObservation
// NodeTrafficReport aggregates traffic counters collected by an agent.
type NodeTrafficReport = pkgprotocol.NodeTrafficReport
// NodeAccessLog is a single access-log record forwarded by an agent.
type NodeAccessLog = pkgprotocol.NodeAccessLog
// BufferedObservabilityRecord bundles multiple observability payloads into one upload.
type BufferedObservabilityRecord = pkgprotocol.BufferedObservabilityRecord
// NodeHealthEvent represents a discrete health-state change on an agent node.
type NodeHealthEvent = pkgprotocol.NodeHealthEvent
// ApplyLogPayload carries the result of a configuration-apply attempt reported by an agent.
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
// Settings contains remote-control directives sent from the server to an agent.
type Settings = pkgprotocol.AgentSettings
// ActiveConfigMeta describes the currently active configuration version on the server.
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
// SupportFile represents a supplementary file bundled with an agent configuration package.
type SupportFile = pkgprotocol.SupportFile
// WAFIPGroup is a named IP-address group used in WAF allow/block rules.
type WAFIPGroup = pkgprotocol.WAFIPGroup
// WAFIPGroupSyncRequest is sent by an agent to request an incremental WAF IP-group sync.
type WAFIPGroupSyncRequest = pkgprotocol.WAFIPGroupSyncRequest
// WAFIPGroupSyncResponse carries the server's reply to a WAF IP-group sync request.
type WAFIPGroupSyncResponse = pkgprotocol.WAFIPGroupSyncResponse
// Backward-compatible names used by server routers and handlers.
// WAFIPGroupSyncInput is an alias for WAFIPGroupSyncRequest kept for backward compatibility.
type WAFIPGroupSyncInput = WAFIPGroupSyncRequest
type WAFIPGroupSyncResult = WAFIPGroupSyncResponse
// WAFIPGroupSyncResult is an alias for WAFIPGroupSyncResponse kept for backward compatibility.
type WAFIPGroupSyncResult = WAFIPGroupSyncResponse
+9 -9
View File
@@ -37,7 +37,7 @@ func RegisterHandler(c *gin.Context) {
result *RegistrationResponse
err error
)
if authNode, ok := AgentNodeFromContext(c); ok {
if authNode, ok := NodeFromContext(c); ok {
result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload)
} else {
result, err = RegisterWithDiscovery(c.Request.Context(), payload)
@@ -67,7 +67,7 @@ func HeartbeatHandler(c *gin.Context) {
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := AgentNodeFromContext(c)
authNode, ok := NodeFromContext(c)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
@@ -91,7 +91,7 @@ func HeartbeatHandler(c *gin.Context) {
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/config-versions/active [get]
func GetActiveConfigHandler(c *gin.Context) {
if _, ok := AgentNodeFromContext(c); !ok {
if _, ok := NodeFromContext(c); !ok {
response.AbortUnauthorized(c, errNodeMissingFromContext)
return
}
@@ -143,7 +143,7 @@ func ReportApplyLogHandler(c *gin.Context) {
if !apiutil.BindJSON(c, &payload) {
return
}
if authNode, ok := AgentNodeFromContext(c); ok {
if authNode, ok := NodeFromContext(c); ok {
payload.NodeID = authNode.NodeID
}
log, err := ReportApplyLog(c.Request.Context(), payload)
@@ -173,7 +173,7 @@ func DownloadPagesPackageHandler(c *gin.Context) {
if apiutil.AbortBadRequestOnError(c, err) {
return
}
defer packageObj.Body.Close()
defer func() { _ = packageObj.Body.Close() }()
c.Header("Content-Disposition", "attachment; filename="+fileName)
if packageObj.ContentType != "" {
c.Header("Content-Type", packageObj.ContentType)
@@ -195,18 +195,18 @@ func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
return uint(id64), true
}
// AgentWebSocketHandler upgrades an authenticated agent websocket connection.
// WebSocketHandler upgrades an authenticated agent websocket connection.
// @Summary Agent WebSocket 连接
// @Description 升级为 WebSocket 长连接,用于实时推送配置同步、WAF IP 组等指令;需携带 X-Agent-Token
// @Tags openflare-agent
// @Security AgentTokenAuth
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/ws [get]
func AgentWebSocketHandler(c *gin.Context) {
authNode, ok := AgentNodeFromContext(c)
func WebSocketHandler(c *gin.Context) {
authNode, ok := NodeFromContext(c)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
}
}
@@ -1,6 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package apply_log manages the application of configuration change logs,
// including validation and retention policy enforcement.
package apply_log
const (
+3 -5
View File
@@ -1,6 +1,4 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package openflare implements openflare configuration, service orchestration, and background tasks.
package openflare
import (
@@ -118,7 +116,7 @@ func (h *DatabaseAutoCleanupHandler) Execute(ctx context.Context, _ []byte) (*ta
}
task.AppendLog(ctx, "开始执行可观测数据自动清理,保留天数=%d", model.DatabaseAutoCleanupRetentionDays)
summary, err := tasks.RunDatabaseAutoCleanupOnce(time.Now())
summary, err := tasks.RunDatabaseAutoCleanupOnce(ctx, time.Now())
if err != nil {
task.AppendLog(ctx, "可观测数据自动清理失败: %v", err)
return nil, err
@@ -197,4 +195,4 @@ func (h *UptimeKumaSyncHandler) Execute(ctx context.Context, _ []byte) (*task.Ta
msg := "Uptime Kuma 同步完成"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
}
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package config_version defines shared error messages for configuration versions.
package config_version
const (
@@ -8,6 +8,7 @@ import (
"encoding/json"
"errors"
"fmt"
"slices"
"sort"
"strconv"
"strings"
@@ -17,6 +18,11 @@ import (
"gorm.io/gorm"
)
const (
cleanupSuccessMessage = "清理成功"
minConfigVersionKeepCount = 3
)
// ConfigPreviewResult is the preview response payload.
type ConfigPreviewResult struct {
SnapshotJSON string `json:"snapshot_json"`
@@ -243,15 +249,15 @@ func ActivateConfigVersion(ctx context.Context, id uint) (*model.ConfigVersion,
// CleanupConfigVersions removes old inactive config versions.
func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, error) {
if keepCount < 3 {
keepCount = 3
if keepCount < minConfigVersionKeepCount {
keepCount = minConfigVersionKeepCount
}
versions, err := model.ListConfigVersionSummaries(ctx)
if err != nil {
return nil, err
}
if len(versions) <= keepCount {
return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil
return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil
}
var deleteIDs []uint
for index, version := range versions {
@@ -264,13 +270,13 @@ func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult,
deleteIDs = append(deleteIDs, version.ID)
}
if len(deleteIDs) == 0 {
return &CleanupResult{DeletedCount: 0, Message: "清理成功"}, nil
return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil
}
deletedCount, err := model.DeleteConfigVersionsByIDs(ctx, deleteIDs)
if err != nil {
return nil, err
}
return &CleanupResult{DeletedCount: deletedCount, Message: "清理成功"}, nil
return &CleanupResult{DeletedCount: deletedCount, Message: cleanupSuccessMessage}, nil
}
func nextVersionNumber(ctx context.Context, now time.Time) (string, error) {
@@ -368,50 +374,50 @@ func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRo
}
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL ||
left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP ||
left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP ||
left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy ||
left.BasicAuthEnabled != right.BasicAuthEnabled || left.BasicAuthUsername != right.BasicAuthUsername ||
left.BasicAuthPassword != right.BasicAuthPassword || left.UpstreamType != right.UpstreamType ||
!uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) || left.TunnelTargetAddr != right.TunnelTargetAddr ||
left.TunnelTargetProto != right.TunnelTargetProto || !uintPtrEqual(left.PagesProjectID, right.PagesProjectID) ||
!uintSliceEqual(left.CertIDs, right.CertIDs) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) {
return false
}
if len(left.Domains) != len(right.Domains) {
return false
}
for index := range left.Domains {
if left.Domains[index] != right.Domains[index] {
return false
}
}
if len(left.Upstreams) != len(right.Upstreams) {
return false
}
for index := range left.Upstreams {
if left.Upstreams[index] != right.Upstreams[index] {
return false
}
}
if len(left.CacheRules) != len(right.CacheRules) {
return false
}
for index := range left.CacheRules {
if left.CacheRules[index] != right.CacheRules[index] {
return false
}
}
if len(left.CustomHeaders) != len(right.CustomHeaders) {
return false
}
for index := range left.CustomHeaders {
if left.CustomHeaders[index] != right.CustomHeaders[index] {
return false
}
}
return true
return snapshotRouteScalarsEqual(left, right) &&
slices.Equal(left.Domains, right.Domains) &&
slices.Equal(left.Upstreams, right.Upstreams) &&
slices.Equal(left.CacheRules, right.CacheRules) &&
slices.Equal(left.CustomHeaders, right.CustomHeaders)
}
func snapshotRouteScalarsEqual(left, right snapshotRoute) bool {
return snapshotRouteIdentityEqual(left, right) &&
snapshotRouteOriginEqual(left, right) &&
snapshotRoutePolicyEqual(left, right) &&
snapshotRouteTunnelEqual(left, right) &&
uintSliceEqual(left.CertIDs, right.CertIDs) &&
uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs)
}
func snapshotRouteIdentityEqual(left, right snapshotRoute) bool {
return left.SiteName == right.SiteName && left.Domain == right.Domain
}
func snapshotRouteOriginEqual(left, right snapshotRoute) bool {
return left.OriginURL == right.OriginURL &&
left.OriginHost == right.OriginHost &&
left.UpstreamType == right.UpstreamType
}
func snapshotRoutePolicyEqual(left, right snapshotRoute) bool {
return left.EnableHTTPS == right.EnableHTTPS &&
left.RedirectHTTP == right.RedirectHTTP &&
left.LimitConnPerServer == right.LimitConnPerServer &&
left.LimitConnPerIP == right.LimitConnPerIP &&
left.LimitRate == right.LimitRate &&
left.CacheEnabled == right.CacheEnabled &&
left.CachePolicy == right.CachePolicy &&
left.BasicAuthEnabled == right.BasicAuthEnabled &&
left.BasicAuthUsername == right.BasicAuthUsername &&
left.BasicAuthPassword == right.BasicAuthPassword
}
func snapshotRouteTunnelEqual(left, right snapshotRoute) bool {
return left.TunnelTargetAddr == right.TunnelTargetAddr &&
left.TunnelTargetProto == right.TunnelTargetProto &&
uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) &&
uintPtrEqual(left.PagesProjectID, right.PagesProjectID)
}
func snapshotWAFConfigEqual(left snapshotWAFDocument, right snapshotWAFDocument) bool {
@@ -16,6 +16,8 @@ import (
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
)
const supportFilesPerCertificate = 2
type snapshotRoute struct {
ID uint `json:"id,omitempty"`
SiteName string `json:"site_name,omitempty"`
@@ -227,13 +229,14 @@ func buildSnapshotRoutes(ctx context.Context, routes []*model.ProxyRoute) ([]sna
var tunnelTargetAddr string
var tunnelTargetProtocol string
var pagesProjectID *uint
if upstreamType == "tunnel" {
switch upstreamType {
case "tunnel":
originURL = resolveTunnelOpenRestyUpstreamURL(ctx)
upstreams = []string{originURL}
tunnelNodeID = route.TunnelNodeID
tunnelTargetAddr = strings.TrimSpace(route.TunnelTargetAddr)
tunnelTargetProtocol = normalizeTunnelTargetProtocol(route.TunnelTargetProtocol)
} else if upstreamType == "pages" {
case "pages":
return nil, fmt.Errorf("路由 %s Pages 配置无效: pages module is not available", route.Domain)
}
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
@@ -495,7 +498,7 @@ func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) (
certIDs = append(certIDs, certID)
}
sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] })
files := make([]SupportFile, 0, len(certIDs)*2)
files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate)
for _, certID := range certIDs {
certificate, err := model.GetTLSCertificateByID(ctx, certID)
if err != nil {
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package dashboard provides helper utilities for dashboard API handlers.
package dashboard
import (
@@ -15,6 +16,11 @@ const (
nodeStatusOnline = "online"
nodeStatusOffline = "offline"
nodeStatusPending = "pending"
dashboardDistributionLimit = 8
highCPUUsagePercentThreshold = 80
highMemoryUsagePercentThreshold = 85
highStorageUsagePercentThreshold = 85
)
func computeNodeStatus(node *model.OpenFlareNode) string {
+47 -37
View File
@@ -120,7 +120,7 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
if err != nil {
return nil, err
}
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, 8)
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, dashboardDistributionLimit)
if err != nil {
return nil, err
}
@@ -136,7 +136,7 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
view := &OverviewView{
GeneratedAt: now,
Nodes: make([]NodeHealth, 0, len(nodes)),
Distributions: observability.BuildTrafficDistributions(reports, accessLogRegions, 8),
Distributions: observability.BuildTrafficDistributions(reports, accessLogRegions, dashboardDistributionLimit),
Trends: observability.NodeTrends{
Traffic24h: observability.BuildTrafficTrendPoints(now, reports),
Capacity24h: observability.BuildCapacityTrendPoints(now, snapshots),
@@ -183,41 +183,8 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
ActiveEventCount: len(nodeActiveEvents),
}
if latestSnapshot != nil {
nodeHealth.CPUUsagePercent = latestSnapshot.CPUUsagePercent
nodeHealth.MemoryUsagePercent = observability.Percentage(latestSnapshot.MemoryUsedBytes, latestSnapshot.MemoryTotalBytes)
nodeHealth.StorageUsagePercent = observability.Percentage(latestSnapshot.StorageUsedBytes, latestSnapshot.StorageTotalBytes)
if latestSnapshot.CPUUsagePercent > 0 {
view.Capacity.AverageCPUUsagePercent += latestSnapshot.CPUUsagePercent
cpuNodeCount++
}
if nodeHealth.MemoryUsagePercent > 0 {
view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent
memoryNodeCount++
}
if latestSnapshot.CPUUsagePercent >= 80 {
view.Capacity.HighCPUNodes++
}
if nodeHealth.MemoryUsagePercent >= 85 {
view.Capacity.HighMemoryNodes++
}
if nodeHealth.StorageUsagePercent >= 85 {
view.Capacity.HighStorageNodes++
}
}
if latestTraffic != nil {
nodeHealth.RequestCount = latestTraffic.RequestCount
nodeHealth.ErrorCount = latestTraffic.ErrorCount
nodeHealth.UniqueVisitorCount = latestTraffic.UniqueVisitorCount
view.Traffic.RequestCount += latestTraffic.RequestCount
view.Traffic.UniqueVisitors += latestTraffic.UniqueVisitorCount
view.Traffic.ErrorCount += latestTraffic.ErrorCount
if duration := latestTraffic.WindowEndedAt.Sub(latestTraffic.WindowStartedAt).Seconds(); duration > 0 {
view.Traffic.EstimatedQPS += float64(latestTraffic.RequestCount) / duration
}
view.Traffic.ReportedNodes++
}
cpuNodeCount, memoryNodeCount = applyNodeSnapshotMetrics(&nodeHealth, latestSnapshot, view, cpuNodeCount, memoryNodeCount)
applyNodeTrafficMetrics(&nodeHealth, latestTraffic, view)
view.Nodes = append(view.Nodes, nodeHealth)
}
@@ -240,6 +207,49 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
return view, nil
}
func applyNodeSnapshotMetrics(nodeHealth *NodeHealth, snapshot *model.OpenFlareMetricSnapshot, view *OverviewView, cpuNodeCount, memoryNodeCount int) (int, int) {
if snapshot == nil {
return cpuNodeCount, memoryNodeCount
}
nodeHealth.CPUUsagePercent = snapshot.CPUUsagePercent
nodeHealth.MemoryUsagePercent = observability.Percentage(snapshot.MemoryUsedBytes, snapshot.MemoryTotalBytes)
nodeHealth.StorageUsagePercent = observability.Percentage(snapshot.StorageUsedBytes, snapshot.StorageTotalBytes)
if snapshot.CPUUsagePercent > 0 {
view.Capacity.AverageCPUUsagePercent += snapshot.CPUUsagePercent
cpuNodeCount++
}
if nodeHealth.MemoryUsagePercent > 0 {
view.Capacity.AverageMemoryUsagePercent += nodeHealth.MemoryUsagePercent
memoryNodeCount++
}
if snapshot.CPUUsagePercent >= highCPUUsagePercentThreshold {
view.Capacity.HighCPUNodes++
}
if nodeHealth.MemoryUsagePercent >= highMemoryUsagePercentThreshold {
view.Capacity.HighMemoryNodes++
}
if nodeHealth.StorageUsagePercent >= highStorageUsagePercentThreshold {
view.Capacity.HighStorageNodes++
}
return cpuNodeCount, memoryNodeCount
}
func applyNodeTrafficMetrics(nodeHealth *NodeHealth, traffic *model.OpenFlareRequestReport, view *OverviewView) {
if traffic == nil {
return
}
nodeHealth.RequestCount = traffic.RequestCount
nodeHealth.ErrorCount = traffic.ErrorCount
nodeHealth.UniqueVisitorCount = traffic.UniqueVisitorCount
view.Traffic.RequestCount += traffic.RequestCount
view.Traffic.UniqueVisitors += traffic.UniqueVisitorCount
view.Traffic.ErrorCount += traffic.ErrorCount
if duration := traffic.WindowEndedAt.Sub(traffic.WindowStartedAt).Seconds(); duration > 0 {
view.Traffic.EstimatedQPS += float64(traffic.RequestCount) / duration
}
view.Traffic.ReportedNodes++
}
func compressOverview(view *OverviewView) *OverviewPayload {
if view == nil {
return &OverviewPayload{
+2 -1
View File
@@ -1,9 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package flared defines shared error messages for tunnel client operations.
package flared
const (
errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效"
errTunnelTokenInvalid = "无权进行此操作,Tunnel Token 无效" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errTunnelNodeTypeMismatch = "此节点不是 TunnelClient 类型"
)
+10 -5
View File
@@ -18,6 +18,11 @@ import (
"gorm.io/gorm"
)
const (
updateChannelStable = "stable"
defaultTunnelTargetPort = 80
)
type configVersionRow struct {
Version string `gorm:"column:version"`
Checksum string `gorm:"column:checksum"`
@@ -31,7 +36,7 @@ func normalizeReleaseChannel(channel string) string {
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
return "preview"
}
return "stable"
return updateChannelStable
}
func normalizeFlaredHeartbeatPayload(payload HeartbeatPayload) HeartbeatPayload {
@@ -146,20 +151,20 @@ func decodeStoredDomains(raw string, fallbackDomain string) ([]string, error) {
func parseTunnelTargetAddr(addr string) (string, int) {
addr = strings.TrimSpace(addr)
if addr == "" {
return "127.0.0.1", 80
return "127.0.0.1", defaultTunnelTargetPort
}
host, portStr, err := net.SplitHostPort(addr)
if err != nil {
lastColon := strings.LastIndex(addr, ":")
if lastColon < 0 {
return addr, 80
return addr, defaultTunnelTargetPort
}
host = addr[:lastColon]
portStr = addr[lastColon+1:]
}
port := 80
port := defaultTunnelTargetPort
if _, scanErr := fmt.Sscanf(portStr, "%d", &port); scanErr != nil {
port = 80
port = defaultTunnelTargetPort
}
if host == "" {
host = "127.0.0.1"
+10 -9
View File
@@ -17,10 +17,11 @@ import (
)
const (
nodeStatusOnline = "online"
applyResultOK = "success"
applyResultWarn = "warning"
applyResultFail = "failed"
nodeStatusOnline = "online"
applyResultOK = "success"
applyResultWarn = "warning"
applyResultFail = "failed"
maxApplyLogMessageLength = 16000
)
// Heartbeat processes an OpenFlared heartbeat and returns runtime settings.
@@ -46,13 +47,13 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
"last_seen_at": now,
"status": nodeStatusOnline,
"update_requested": false,
"update_channel": "stable",
"update_channel": updateChannelStable,
"update_tag": "",
}
if !previous.UpdateRequested {
delete(changes, "update_requested")
}
if previous.UpdateChannel == "stable" {
if previous.UpdateChannel == updateChannelStable {
delete(changes, "update_channel")
}
if previous.UpdateTag == "" {
@@ -67,7 +68,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
node.ExtVersion = payload.FrpVersion
node.CurrentVersion = payload.CurrentVersion
node.UpdateRequested = false
node.UpdateChannel = "stable"
node.UpdateChannel = updateChannelStable
node.UpdateTag = ""
lastSeen := now
node.LastSeenAt = &lastSeen
@@ -228,8 +229,8 @@ func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
payload.Checksum = strings.TrimSpace(payload.Checksum)
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
if len(payload.Message) > 16000 {
payload.Message = payload.Message[:16000]
if len(payload.Message) > maxApplyLogMessageLength {
payload.Message = payload.Message[:maxApplyLogMessageLength]
}
return payload
}
@@ -5,11 +5,26 @@ package flared
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
// HeartbeatPayload is an alias for FlaredHeartbeatPayload.
type HeartbeatPayload = pkgprotocol.FlaredHeartbeatPayload
// ConnectedRelay is an alias for FlaredConnectedRelay.
type ConnectedRelay = pkgprotocol.FlaredConnectedRelay
// ActiveConfigMeta is an alias for ActiveConfigMeta.
type ActiveConfigMeta = pkgprotocol.ActiveConfigMeta
// HeartbeatResponse is an alias for FlaredHeartbeatResponse.
type HeartbeatResponse = pkgprotocol.FlaredHeartbeatResponse
// TunnelConfigResponse is an alias for FlaredTunnelConfigResponse.
type TunnelConfigResponse = pkgprotocol.FlaredTunnelConfigResponse
// RelayInfo is an alias for FlaredRelayInfo.
type RelayInfo = pkgprotocol.FlaredRelayInfo
// ProxyEntry is an alias for FlaredProxyEntry.
type ProxyEntry = pkgprotocol.FlaredProxyEntry
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
// ApplyLogPayload is an alias for ApplyLogPayload.
type ApplyLogPayload = pkgprotocol.ApplyLogPayload
+1 -1
View File
@@ -31,7 +31,7 @@ func (f *fakeLookupProvider) Close() error { return nil }
func TestLookupWithProvider(t *testing.T) {
previousFactory := pkggeoip.ProviderFactoryForTest()
pkggeoip.SetProviderFactoryForTest(func(provider string) (pkggeoip.GeoIPService, error) {
pkggeoip.SetProviderFactoryForTest(func(provider string) (pkggeoip.Service, error) {
return &fakeLookupProvider{}, nil
})
t.Cleanup(func() {
+1
View File
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package node defines node validation and management error messages.
package node
const (
+52 -45
View File
@@ -28,6 +28,13 @@ const (
openrestyStatusUnhealthy = "unhealthy"
openrestyStatusUnknown = "unknown"
githubReleasesAPIBase = "https://api.github.com/repos/%s/releases"
nodeTypeTunnelRelay = "tunnel_relay"
nodeTypeTunnelClient = "tunnel_client"
nodeTypeEdgeNode = "edge_node"
nodeTokenByteLength = 16
maxNodeIPLength = 64
maxNodeGeoNameLength = 128
)
type releaseChannel string
@@ -49,7 +56,7 @@ type githubReleaseResponse struct {
}
func newRandomToken() (string, error) {
buf := make([]byte, 16)
buf := make([]byte, nodeTokenByteLength)
if _, err := rand.Read(buf); err != nil {
return "", err
}
@@ -66,12 +73,12 @@ func newServerNodeID() (string, error) {
func normalizeNodeType(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "tunnel_relay":
return "tunnel_relay"
case "tunnel_client":
return "tunnel_client"
case nodeTypeTunnelRelay:
return nodeTypeTunnelRelay
case nodeTypeTunnelClient:
return nodeTypeTunnelClient
default:
return "edge_node"
return nodeTypeEdgeNode
}
}
@@ -134,40 +141,50 @@ func normalizeNodeInput(input Input) (string, string, string, *float64, *float64
name := strings.TrimSpace(input.Name)
ip := strings.TrimSpace(input.IP)
geoName := strings.TrimSpace(input.GeoName)
manualOverride := input.GeoManualOverride || geoName != "" || input.GeoLatitude != nil || input.GeoLongitude != nil
if len(ip) > 64 {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPTooLong)
if err := validateNodeIPInput(input, ip); err != nil {
return "", "", "", nil, nil, false, err
}
if ip != "" && net.ParseIP(ip) == nil {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPInvalid)
}
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeIPManualRequired)
}
if len(geoName) > 128 {
if len(geoName) > maxNodeGeoNameLength {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoNameTooLong)
}
geoLatitude := cloneCoordinate(input.GeoLatitude)
geoLongitude := cloneCoordinate(input.GeoLongitude)
if err := validateNodeGeoCoordinates(geoLatitude, geoLongitude); err != nil {
return "", "", "", nil, nil, false, err
}
manualOverride := input.GeoManualOverride || geoName != "" || geoLatitude != nil || geoLongitude != nil
if !manualOverride || (geoLatitude == nil && geoLongitude == nil && geoName == "") {
return name, ip, "", nil, nil, false, nil
}
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
}
func validateNodeIPInput(input Input, ip string) error {
if len(ip) > maxNodeIPLength {
return fmt.Errorf("%s", errNodeIPTooLong)
}
if ip != "" && net.ParseIP(ip) == nil {
return fmt.Errorf("%s", errNodeIPInvalid)
}
if input.IPManualOverride != nil && *input.IPManualOverride && ip == "" {
return fmt.Errorf("%s", errNodeIPManualRequired)
}
return nil
}
func validateNodeGeoCoordinates(geoLatitude, geoLongitude *float64) error {
if (geoLatitude == nil) != (geoLongitude == nil) {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoCoordinateMismatch)
return fmt.Errorf("%s", errNodeGeoCoordinateMismatch)
}
if geoLatitude != nil && (*geoLatitude < -90 || *geoLatitude > 90) {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLatitudeInvalid)
return fmt.Errorf("%s", errNodeGeoLatitudeInvalid)
}
if geoLongitude != nil && (*geoLongitude < -180 || *geoLongitude > 180) {
return "", "", "", nil, nil, false, fmt.Errorf("%s", errNodeGeoLongitudeInvalid)
return fmt.Errorf("%s", errNodeGeoLongitudeInvalid)
}
if !manualOverride {
return name, ip, "", nil, nil, false, nil
}
if geoLatitude == nil && geoLongitude == nil && geoName == "" {
return name, ip, "", nil, nil, false, nil
}
return name, ip, geoName, geoLatitude, geoLongitude, true, nil
return nil
}
func computeNodeStatus(node *model.OpenFlareNode) string {
@@ -189,12 +206,12 @@ func nodeViewLastSeenAt(node *model.OpenFlareNode) any {
}
nodeType := strings.TrimSpace(node.NodeType)
if nodeType == "" {
nodeType = "edge_node"
nodeType = nodeTypeEdgeNode
}
if nodeType == "tunnel_relay" && ofws.IsRelayConnected(node.NodeID) {
if nodeType == nodeTypeTunnelRelay && ofws.IsRelayConnected(node.NodeID) {
return ofws.RelayWSConnectedLastSeenValue
}
if nodeType == "tunnel_client" && ofws.IsFlaredConnected(node.NodeID) {
if nodeType == nodeTypeTunnelClient && ofws.IsFlaredConnected(node.NodeID) {
return ofws.FlaredWSConnectedLastSeenValue
}
if ofws.IsAgentConnected(node.NodeID) {
@@ -250,7 +267,7 @@ func buildNodeView(node *model.OpenFlareNode) *View {
view.UpdateChannel = releaseChannelStable.String()
}
if view.NodeType == "" {
view.NodeType = "edge_node"
view.NodeType = nodeTypeEdgeNode
}
return view
}
@@ -378,7 +395,7 @@ func fetchLatestStableGitHubRelease(ctx context.Context, repo string) (*githubRe
if err != nil {
return nil, fmt.Errorf("获取最新版本失败: %v", err)
}
defer resp.Body.Close()
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
}
@@ -395,7 +412,7 @@ func fetchLatestPreviewGitHubRelease(ctx context.Context, repo string) (*githubR
if err != nil {
return nil, fmt.Errorf("获取 preview 版本失败: %v", err)
}
defer resp.Body.Close()
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("GitHub 返回异常状态: %s", resp.Status)
}
@@ -427,7 +444,7 @@ func fetchGitHubReleaseByTag(ctx context.Context, repo string, tag string) (*git
if err != nil {
return nil, fmt.Errorf("获取指定版本失败: %v", err)
}
defer resp.Body.Close()
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode == http.StatusNotFound {
return nil, fmt.Errorf("未找到指定版本: %s", tag)
}
@@ -461,13 +478,3 @@ func isUniqueConstraintError(err error) bool {
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
func setReleaseHTTPClientForTest(client *http.Client) *http.Client {
previous := releaseHTTPClient
if client == nil {
releaseHTTPClient = &http.Client{Timeout: 30 * time.Second}
} else {
releaseHTTPClient = client
}
return previous
}
+8 -3
View File
@@ -17,6 +17,11 @@ import (
"gorm.io/gorm"
)
const (
defaultRelayBindPort = 7000
defaultRelayVhostHTTPPort = 8080
)
// Input is the create/update node payload.
type Input struct {
Name string `json:"name"`
@@ -184,8 +189,8 @@ func CreateNode(ctx context.Context, input Input) (*View, error) {
return nil, err
}
if node.NodeType == "tunnel_relay" {
node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, 7000)
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, 8080)
node.RelayBindPort = normalizeRelayPort(input.RelayBindPort, defaultRelayBindPort)
node.RelayVhostHTTPPort = normalizeRelayPort(input.RelayVhostHTTPPort, defaultRelayVhostHTTPPort)
node.RelayAuthToken, err = newRandomToken()
if err != nil {
return nil, err
@@ -399,7 +404,7 @@ func ValidateDiscoveryToken(ctx context.Context, token string) error {
return err
}
if token != discoveryToken {
return fmt.Errorf("Discovery Token 无效")
return fmt.Errorf("discovery Token 无效") // error 消息首字母小写
}
return nil
}
@@ -1,6 +1,4 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package observability provides monitoring, metrics, and access log analysis for OpenFlare.
package observability
import (
@@ -22,6 +20,8 @@ const (
defaultIPTrendBucketMinute = 30
maxIPTrendHours = 168
nodeAccessLogRetentionDays = 90
accessLogFieldRemoteAddr = "remote_addr"
accessLogFieldRequestCount = "request_count"
)
var nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
@@ -524,9 +524,9 @@ func normalizeFoldedAccessLogIPQuery(input FoldedAccessLogIPQuery) (FoldedAccess
}
normalizedSortBy := strings.TrimSpace(input.SortBy)
switch normalizedSortBy {
case "last_seen_at", "remote_addr":
case "last_seen_at", accessLogFieldRemoteAddr:
default:
normalizedSortBy = "request_count"
normalizedSortBy = accessLogFieldRequestCount
}
return FoldedAccessLogIPQuery{
NodeID: strings.TrimSpace(input.NodeID),
@@ -591,7 +591,7 @@ func normalizeAccessLogPageSize(pageSize int) int {
func normalizeAccessLogSortBy(sortBy string) string {
switch strings.TrimSpace(sortBy) {
case "status_code", "remote_addr", "host", "path":
case "status_code", accessLogFieldRemoteAddr, "host", "path":
return strings.TrimSpace(sortBy)
default:
return defaultAccessLogSortBy
@@ -607,8 +607,8 @@ func normalizeAccessLogSortOrder(sortOrder string) string {
func normalizeFoldSortBy(sortBy string) string {
switch strings.TrimSpace(sortBy) {
case "request_count":
return "request_count"
case accessLogFieldRequestCount:
return accessLogFieldRequestCount
default:
return "bucket_started_at"
}
@@ -616,7 +616,7 @@ func normalizeFoldSortBy(sortBy string) string {
func normalizeIPSummarySortBy(sortBy string) string {
switch strings.TrimSpace(sortBy) {
case "recent_requests", "last_seen_at", "remote_addr":
case "recent_requests", "last_seen_at", accessLogFieldRemoteAddr:
return strings.TrimSpace(sortBy)
default:
return "total_requests"
@@ -19,6 +19,7 @@ const (
healthEventStatusResolved = "resolved"
healthSeverityCritical = "critical"
healthSeverityWarning = "warning"
percentageMultiplier = 100
)
// DistributionItem is a key/value distribution entry.
@@ -415,7 +416,7 @@ func Percentage(used int64, total int64) float64 {
if used <= 0 || total <= 0 {
return 0
}
return (float64(used) / float64(total)) * 100
return (float64(used) / float64(total)) * percentageMultiplier
}
func mergeJSONCounts(target distributionAccumulator, raw string) {
@@ -14,9 +14,10 @@ import (
)
const (
defaultObservabilityWindow = 24 * time.Hour
defaultObservabilityLimit = 120
maxObservabilityLimit = 500
defaultObservabilityWindow = 24 * time.Hour
defaultObservabilityLimit = 120
maxObservabilityLimit = 500
defaultTrafficDistributionLimit = 8
)
// NodeQuery filters node observability data.
@@ -106,7 +107,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV
if err != nil {
return nil, err
}
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, 8)
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, defaultTrafficDistributionLimit)
if err != nil {
return nil, err
}
@@ -135,7 +136,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV
HealthEvents: events,
Analytics: NodeAnalytics{
Traffic: buildTrafficWindowSummary(latestTrafficReport(reports)),
Distributions: BuildTrafficDistributions(reports, accessLogRegions, 8),
Distributions: BuildTrafficDistributions(reports, accessLogRegions, defaultTrafficDistributionLimit),
Health: buildHealthSummary(latestMetricSnapshot(snapshots), latestTrafficReport(reports), events),
},
Trends: NodeTrends{
@@ -40,7 +40,7 @@ func GetAccessLogsHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(logs))
}
// getFoldedAccessLogsHandler 分页列出折叠访问日志。
// GetFoldedAccessLogsHandler 分页列出折叠访问日志。
// @Summary 列出折叠访问日志
// @Description 按时间桶聚合访问日志并分页返回,需要管理员权限
// @Tags openflare-observability
@@ -71,7 +71,7 @@ func GetFoldedAccessLogsHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(logs))
}
// getFoldedAccessLogIPsHandler 列出折叠桶内的 IP 汇总。
// GetFoldedAccessLogIPsHandler 列出折叠桶内的 IP 汇总。
// @Summary 列出折叠访问日志 IP 汇总
// @Description 在指定时间桶内按 IP 聚合访问统计,需要管理员权限
// @Tags openflare-observability
@@ -112,7 +112,7 @@ func GetFoldedAccessLogIPsHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(result))
}
// getAccessLogIPSummariesHandler 列出访问日志 IP 汇总。
// GetAccessLogIPSummariesHandler 列出访问日志 IP 汇总。
// @Summary 列出访问日志 IP 汇总
// @Description 按 IP 聚合访问日志统计并分页返回,需要管理员权限
// @Tags openflare-observability
@@ -147,7 +147,7 @@ func GetAccessLogIPSummariesHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(result))
}
// getAccessLogIPTrendHandler 获取 IP 访问趋势。
// GetAccessLogIPTrendHandler 获取 IP 访问趋势。
// @Summary 获取访问日志 IP 趋势
// @Description 返回指定 IP 在时间范围内的访问趋势数据,需要管理员权限
// @Tags openflare-observability
@@ -178,7 +178,7 @@ func GetAccessLogIPTrendHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(result))
}
// cleanupAccessLogsHandler 清理过期访问日志。
// CleanupAccessLogsHandler 清理过期访问日志。
// @Summary 清理访问日志
// @Description 按保留天数清理过期访问日志记录,需要管理员权限
// @Tags openflare-observability
@@ -220,4 +220,4 @@ func readAccessLogQuery(c *gin.Context) AccessLogQuery {
func readQueryInt(c *gin.Context, key string) int {
value, _ := strconv.Atoi(c.DefaultQuery(key, "0"))
return value
}
}
+1
View File
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package option provides handler-level error message constants for the openflare option module.
package option
const (
+1 -1
View File
@@ -161,7 +161,7 @@ func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) {
StartTime: model.StartTime,
EmailVerification: model.EmailVerificationEnabled,
GitHubOAuth: model.GitHubOAuthEnabled,
GitHubClientID: model.GitHubClientId,
GitHubClientID: model.GitHubClientID,
SystemName: model.SystemName,
HomePageLink: model.HomePageLink,
FooterHTML: model.Footer,
@@ -0,0 +1,179 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package option
import (
"fmt"
"regexp"
"strconv"
"strings"
)
var openRestyOptionValidators = map[string]func(key, value string) error{
"OpenRestyDefaultServerReturnStatus": validateOpenRestyDefaultServerReturnStatus,
"OpenRestyWorkerProcesses": validateOpenRestyWorkerProcesses,
"OpenRestyWorkerConnections": validatePositiveIntegerOption,
"OpenRestyWorkerRlimitNofile": validatePositiveIntegerOption,
"OpenRestyKeepaliveTimeout": validatePositiveIntegerOption,
"OpenRestyKeepaliveRequests": validatePositiveIntegerOption,
"OpenRestyClientHeaderTimeout": validatePositiveIntegerOption,
"OpenRestyClientBodyTimeout": validatePositiveIntegerOption,
"OpenRestySendTimeout": validatePositiveIntegerOption,
"OpenRestyProxyConnectTimeout": validatePositiveIntegerOption,
"OpenRestyProxySendTimeout": validatePositiveIntegerOption,
"OpenRestyProxyReadTimeout": validatePositiveIntegerOption,
"OpenRestyGzipMinLength": validatePositiveIntegerOption,
"OpenRestyGzipCompLevel": validateOpenRestyGzipCompLevel,
"OpenRestyEventsUse": validateOpenRestyEventsUse,
"OpenRestyResolvers": validateOpenRestyResolvers,
"OpenRestyEventsMultiAcceptEnabled": validateBooleanOption,
"OpenRestyWebsocketEnabled": validateBooleanOption,
"OpenRestyHTTP3Enabled": validateBooleanOption,
"OpenRestyProxyRequestBufferingEnabled": validateBooleanOption,
"OpenRestyProxyBufferingEnabled": validateBooleanOption,
"OpenRestyGzipEnabled": validateBooleanOption,
"OpenRestyCacheEnabled": validateBooleanOption,
"OpenRestyCacheLockEnabled": validateBooleanOption,
"OpenRestyProxyBuffers": validateOpenRestyProxyBuffers,
"OpenRestyLargeClientHeaderBuffers": validateOpenRestyProxyBuffers,
"OpenRestyProxyBufferSize": validateOpenRestySizeValue,
"OpenRestyProxyBusyBuffersSize": validateOpenRestySizeValue,
"OpenRestyCacheMaxSize": validateOpenRestySizeValue,
"OpenRestyClientMaxBodySize": validateOpenRestySizeValue,
"OpenRestyCachePath": validateOpenRestyCachePath,
"OpenRestyCacheLevels": validateOpenRestyCacheLevels,
"OpenRestyCacheInactive": validateOpenRestyDurationToken,
"OpenRestyCacheLockTimeout": validateOpenRestyDurationToken,
"OpenRestyCacheKeyTemplate": validateOpenRestyCacheKeyTemplate,
"OpenRestyCacheUseStale": validateOpenRestyCacheUseStale,
"OpenRestyMainConfigTemplate": validateOpenRestyMainConfigTemplate,
}
func validateOpenRestyOption(key, value string) error {
trimmed := strings.TrimSpace(value)
if validator, ok := openRestyOptionValidators[key]; ok {
return validator(key, trimmed)
}
return nil
}
func validateOpenRestyDefaultServerReturnStatus(key, trimmed string) error {
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
statusCode, _ := strconv.Atoi(trimmed)
if statusCode < 100 || statusCode > 999 {
return fmt.Errorf("%s 必须在 100 到 999 之间", key)
}
return nil
}
func validateOpenRestyWorkerProcesses(key, trimmed string) error {
if trimmed == "auto" {
return nil
}
return validatePositiveIntegerOption(key, trimmed)
}
func validateOpenRestyGzipCompLevel(key, trimmed string) error {
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
level, _ := strconv.Atoi(trimmed)
if level > maxOpenRestyGzipCompLevel {
return fmt.Errorf("%s 不能大于 %d", key, maxOpenRestyGzipCompLevel)
}
return nil
}
func validateOpenRestyEventsUse(key, trimmed string) error {
if trimmed == "" {
return nil
}
switch trimmed {
case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport":
return nil
default:
return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key)
}
}
func validateOpenRestyResolvers(key, trimmed string) error {
if trimmed == "" {
return nil
}
if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) {
return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key)
}
return nil
}
func validateOpenRestyProxyBuffers(key, trimmed string) error {
if openRestyProxyBuffersPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key)
}
func validateOpenRestySizeValue(key, trimmed string) error {
if openRestySizePattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key)
}
func validateOpenRestyCachePath(key, trimmed string) error {
if strings.ContainsAny(trimmed, "\r\n\t") {
return fmt.Errorf("%s 不能包含换行或制表符", key)
}
return nil
}
func validateOpenRestyCacheLevels(key, trimmed string) error {
if openRestyCacheLevelsPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key)
}
func validateOpenRestyDurationToken(key, trimmed string) error {
if openRestyDurationTokenPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key)
}
func validateOpenRestyCacheKeyTemplate(key, trimmed string) error {
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
if strings.ContainsAny(trimmed, "\r\n") {
return fmt.Errorf("%s 不能包含换行", key)
}
return nil
}
func validateOpenRestyCacheUseStale(key, trimmed string) error {
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
allowedTokens := map[string]struct{}{
"error": {}, "timeout": {}, "invalid_header": {}, "updating": {},
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
}
for _, token := range strings.Fields(trimmed) {
if _, ok := allowedTokens[token]; !ok {
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
}
}
return nil
}
func validateOpenRestyMainConfigTemplate(key, value string) error {
if strings.TrimSpace(value) == "" {
return fmt.Errorf("%s 不能为空", key)
}
return nil
}
+8 -15
View File
@@ -32,7 +32,7 @@ func GetStatusHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(view))
}
// getNoticeHandler 获取系统公告。
// GetNoticeHandler 获取系统公告。
// @Summary 获取系统公告
// @Description 返回 OpenFlare 控制台公告文本,无需登录
// @Tags openflare-option
@@ -41,7 +41,6 @@ func GetStatusHandler(c *gin.Context) {
// @Failure 400 {object} response.Any "参数错误"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/notice [get]
// GetNoticeHandler returns the notice content.
func GetNoticeHandler(c *gin.Context) {
notice, err := getNotice(c.Request.Context())
if apiutil.AbortBadRequestOnError(c, err) {
@@ -50,7 +49,7 @@ func GetNoticeHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(notice))
}
// listOptionsHandler 列出全部配置项。
// ListOptionsHandler 列出全部配置项。
// @Summary 列出 OpenFlare 配置项
// @Description 返回全部非敏感 OpenFlare 配置项,需要管理员权限
// @Tags openflare-option
@@ -62,7 +61,6 @@ func GetNoticeHandler(c *gin.Context) {
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option [get]
// ListOptionsHandler lists OpenFlare options.
func ListOptionsHandler(c *gin.Context) {
options, err := listOptions(c.Request.Context())
if apiutil.AbortBadRequestOnError(c, err) {
@@ -71,7 +69,7 @@ func ListOptionsHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(options))
}
// updateOptionHandler 更新单个配置项。
// UpdateOptionHandler 更新单个配置项。
// @Summary 更新 OpenFlare 配置项
// @Description 更新单个 OpenFlare 配置项,需要管理员权限
// @Tags openflare-option
@@ -85,7 +83,6 @@ func ListOptionsHandler(c *gin.Context) {
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option/update [post]
// UpdateOptionHandler updates a single option.
func UpdateOptionHandler(c *gin.Context) {
var option model.OpenFlareOption
if !apiutil.BindJSON(c, &option) {
@@ -97,7 +94,7 @@ func UpdateOptionHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
// updateOptionsBatchHandler 批量更新配置项。
// UpdateOptionsBatchHandler 批量更新配置项。
// @Summary 批量更新 OpenFlare 配置项
// @Description 批量更新多个 OpenFlare 配置项,需要管理员权限
// @Tags openflare-option
@@ -111,7 +108,6 @@ func UpdateOptionHandler(c *gin.Context) {
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option/update-batch [post]
// UpdateOptionsBatchHandler updates options in batch.
func UpdateOptionsBatchHandler(c *gin.Context) {
var payload optionBatchPayload
if !apiutil.BindJSON(c, &payload) {
@@ -123,7 +119,7 @@ func UpdateOptionsBatchHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
// lookupGeoIPHandler 查询 GeoIP 信息。
// LookupGeoIPHandler 查询 GeoIP 信息。
// @Summary GeoIP 地址查询
// @Description 按提供商与 IP 查询地理位置信息,需要管理员权限
// @Tags openflare-option
@@ -137,7 +133,6 @@ func UpdateOptionsBatchHandler(c *gin.Context) {
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option/geoip/lookup [post]
// LookupGeoIPHandler performs a GeoIP lookup.
func LookupGeoIPHandler(c *gin.Context) {
var request geoIPLookupRequest
if !apiutil.BindJSON(c, &request) {
@@ -150,7 +145,7 @@ func LookupGeoIPHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(view))
}
// cleanupDatabaseHandler 清理可观测性数据库数据。
// CleanupDatabaseHandler 清理可观测性数据库数据。
// @Summary 清理可观测性数据库
// @Description 按目标与保留天数清理可观测性相关数据表,需要管理员权限
// @Tags openflare-option
@@ -164,7 +159,6 @@ func LookupGeoIPHandler(c *gin.Context) {
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/option/database/cleanup [post]
// CleanupDatabaseHandler cleans up observability data.
func CleanupDatabaseHandler(c *gin.Context) {
var input databaseCleanupInput
if err := bindOptionalJSON(c.Request.Body, &input); err != nil {
@@ -178,7 +172,7 @@ func CleanupDatabaseHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(result))
}
// syncUptimeKumaHandler 同步 Uptime Kuma 监控。
// SyncUptimeKumaHandler 同步 Uptime Kuma 监控。
// @Summary 同步 Uptime Kuma
// @Description 将 OpenFlare 节点同步到 Uptime Kuma,需要管理员权限
// @Tags openflare-option
@@ -191,7 +185,6 @@ func CleanupDatabaseHandler(c *gin.Context) {
// @Failure 404 {object} response.Any "无权限或不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/uptimekuma/sync [post]
// SyncUptimeKumaHandler triggers UptimeKuma sync.
func SyncUptimeKumaHandler(c *gin.Context) {
if apiutil.AbortBadRequestOnError(c, syncUptimeKuma(c.Request.Context())) {
return
@@ -204,4 +197,4 @@ func bindOptionalJSON(body io.Reader, target any) error {
return err
}
return nil
}
}
+61 -146
View File
@@ -4,6 +4,7 @@
package option
import (
"errors"
"fmt"
"regexp"
"strconv"
@@ -13,6 +14,8 @@ import (
"github.com/Rain-kl/Wavelet/internal/model"
)
const maxOpenRestyGzipCompLevel = 9
var (
openRestySizePattern = regexp.MustCompile(`^\d+[kKmMgG]?$`)
openRestyProxyBuffersPattern = regexp.MustCompile(`^\d+\s+\d+[kKmMgG]?$`)
@@ -20,6 +23,8 @@ var (
openRestyDurationTokenPattern = regexp.MustCompile(`^\d+[smhdwSMHDW]$`)
)
const optionValueTrue = "true"
func buildOptionValidationState(options []model.OpenFlareOption) map[string]string {
model.OptionMapRWMutex.RLock()
state := make(map[string]string, len(model.OptionMap)+len(options))
@@ -37,11 +42,11 @@ func buildOptionValidationState(options []model.OpenFlareOption) map[string]stri
func validateOptionWithState(option model.OpenFlareOption, state map[string]string) error {
switch option.Key {
case "GitHubOAuthEnabled":
if option.Value == "true" && strings.TrimSpace(state["GitHubClientId"]) == "" {
if option.Value == optionValueTrue && strings.TrimSpace(state["GitHubClientId"]) == "" {
return fmt.Errorf("无法启用 GitHub OAuth,请先填入 GitHub Client ID 以及 GitHub Client Secret!")
}
case "WeChatAuthEnabled":
if option.Value == "true" && strings.TrimSpace(state["WeChatServerAddress"]) == "" {
if option.Value == optionValueTrue && strings.TrimSpace(state["WeChatServerAddress"]) == "" {
return fmt.Errorf("无法启用微信登录,请先填入微信登录相关配置信息!")
}
}
@@ -71,7 +76,7 @@ func validatePositiveIntegerOption(key, value string) error {
func validateBooleanOption(key, value string) error {
switch value {
case "true", "false":
case optionValueTrue, "false":
return nil
default:
return fmt.Errorf("%s 必须为 true 或 false", key)
@@ -112,171 +117,81 @@ func validateUptimeKumaOption(key, value string, state map[string]string) error
trimmed := strings.TrimSpace(value)
switch key {
case "UptimeKumaEnabled":
if err := validateBooleanOption(key, trimmed); err != nil {
return err
}
if trimmed == "true" {
url := strings.TrimSpace(state["UptimeKumaUrl"])
username := strings.TrimSpace(state["UptimeKumaUsername"])
password := strings.TrimSpace(state["UptimeKumaPassword"])
if url == "" {
return fmt.Errorf("启用 Uptime Kuma 时地址不能为空")
}
if username == "" {
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
}
if password == "" && model.UptimeKumaPassword == "" {
return fmt.Errorf("启用 Uptime Kuma 时密码不能为空")
}
}
return validateUptimeKumaEnabled(key, trimmed, state)
case "UptimeKumaUsername":
if trimmed == "" && state["UptimeKumaEnabled"] == "true" {
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
}
return validateUptimeKumaUsername(trimmed, state)
case "UptimeKumaUrl":
if trimmed != "" && !strings.HasPrefix(trimmed, "http://") && !strings.HasPrefix(trimmed, "https://") {
return fmt.Errorf("Uptime Kuma 地址必须以 http:// 或 https:// 开头")
}
return validateUptimeKumaURL(trimmed)
case "UptimeKumaMonitorScope":
if trimmed != "all" && trimmed != "selected" {
return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)")
}
return validateUptimeKumaMonitorScope(trimmed)
case "UptimeKumaSyncInterval", "UptimeKumaInterval", "UptimeKumaRetryInterval", "UptimeKumaTimeout":
return validatePositiveIntegerOption(key, trimmed)
case "UptimeKumaRetry":
intValue, err := strconv.Atoi(trimmed)
if err != nil || intValue < 0 {
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
}
return validateUptimeKumaRetry(key, trimmed)
}
return nil
}
func validateOpenRestyOption(key, value string) error {
trimmed := strings.TrimSpace(value)
func validateUptimeKumaEnabled(key, trimmed string, state map[string]string) error {
if err := validateBooleanOption(key, trimmed); err != nil {
return err
}
if trimmed != optionValueTrue {
return nil
}
url := strings.TrimSpace(state["UptimeKumaUrl"])
username := strings.TrimSpace(state["UptimeKumaUsername"])
password := strings.TrimSpace(state["UptimeKumaPassword"])
if url == "" {
return fmt.Errorf("启用 Uptime Kuma 时地址不能为空")
}
if username == "" {
return fmt.Errorf("启用 Uptime Kuma 时用户名不能为空")
}
if password == "" && model.UptimeKumaPassword == "" {
return fmt.Errorf("启用 Uptime Kuma 时密码不能为空")
}
return nil
}
switch key {
case "OpenRestyDefaultServerReturnStatus":
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
statusCode, _ := strconv.Atoi(trimmed)
if statusCode < 100 || statusCode > 999 {
return fmt.Errorf("%s 必须在 100 到 999 之间", key)
}
case "OpenRestyWorkerProcesses":
if trimmed == "auto" {
return nil
}
return validatePositiveIntegerOption(key, trimmed)
case "OpenRestyWorkerConnections",
"OpenRestyWorkerRlimitNofile",
"OpenRestyKeepaliveTimeout",
"OpenRestyKeepaliveRequests",
"OpenRestyClientHeaderTimeout",
"OpenRestyClientBodyTimeout",
"OpenRestySendTimeout",
"OpenRestyProxyConnectTimeout",
"OpenRestyProxySendTimeout",
"OpenRestyProxyReadTimeout",
"OpenRestyGzipMinLength":
return validatePositiveIntegerOption(key, trimmed)
case "OpenRestyGzipCompLevel":
if err := validatePositiveIntegerOption(key, trimmed); err != nil {
return err
}
level, _ := strconv.Atoi(trimmed)
if level > 9 {
return fmt.Errorf("%s 不能大于 9", key)
}
case "OpenRestyEventsUse":
if trimmed == "" {
return nil
}
switch trimmed {
case "epoll", "kqueue", "poll", "select", "rtsig", "/dev/poll", "eventport":
return nil
default:
return fmt.Errorf("%s 仅支持 epoll、kqueue、poll、select、rtsig、/dev/poll、eventport 或留空", key)
}
case "OpenRestyResolvers":
if trimmed == "" {
return nil
}
if !regexp.MustCompile(`^[a-zA-Z0-9.:\-\s]+$`).MatchString(trimmed) {
return fmt.Errorf("%s 包含非法字符,请填入有效的 IP 地址或域名,以空格分隔", key)
}
case "OpenRestyEventsMultiAcceptEnabled",
"OpenRestyWebsocketEnabled",
"OpenRestyHTTP3Enabled",
"OpenRestyProxyRequestBufferingEnabled",
"OpenRestyProxyBufferingEnabled",
"OpenRestyGzipEnabled",
"OpenRestyCacheEnabled",
"OpenRestyCacheLockEnabled":
return validateBooleanOption(key, trimmed)
case "OpenRestyProxyBuffers", "OpenRestyLargeClientHeaderBuffers":
if openRestyProxyBuffersPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"16 16k\"", key)
case "OpenRestyProxyBufferSize", "OpenRestyProxyBusyBuffersSize", "OpenRestyCacheMaxSize", "OpenRestyClientMaxBodySize":
if openRestySizePattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为整数或带 k/m/g 单位的大小值", key)
case "OpenRestyCachePath":
if strings.ContainsAny(trimmed, "\r\n\t") {
return fmt.Errorf("%s 不能包含换行或制表符", key)
}
case "OpenRestyCacheLevels":
if openRestyCacheLevelsPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须类似 \"1:2\" 或 \"1:2:2\"", key)
case "OpenRestyCacheInactive", "OpenRestyCacheLockTimeout":
if openRestyDurationTokenPattern.MatchString(trimmed) {
return nil
}
return fmt.Errorf("%s 格式必须为带单位的时长,例如 30m 或 5s", key)
case "OpenRestyCacheKeyTemplate":
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
if strings.ContainsAny(trimmed, "\r\n") {
return fmt.Errorf("%s 不能包含换行", key)
}
case "OpenRestyCacheUseStale":
if trimmed == "" {
return fmt.Errorf("%s 不能为空", key)
}
allowedTokens := map[string]struct{}{
"error": {}, "timeout": {}, "invalid_header": {}, "updating": {},
"http_500": {}, "http_502": {}, "http_503": {}, "http_504": {},
"http_403": {}, "http_404": {}, "http_429": {}, "off": {},
}
for _, token := range strings.Fields(trimmed) {
if _, ok := allowedTokens[token]; !ok {
return fmt.Errorf("%s 包含不支持的值 %q", key, token)
}
}
case "OpenRestyMainConfigTemplate":
if strings.TrimSpace(value) == "" {
return fmt.Errorf("%s 不能为空", key)
}
func validateUptimeKumaUsername(trimmed string, state map[string]string) error {
if trimmed == "" && state["UptimeKumaEnabled"] == optionValueTrue {
return fmt.Errorf("启用 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 nil
}
func validateUptimeKumaMonitorScope(trimmed string) error {
if trimmed != "all" && trimmed != "selected" {
return fmt.Errorf("监控范围必须为全部站点 (all) 或选择站点 (selected)")
}
return nil
}
func validateUptimeKumaRetry(key, trimmed string) error {
intValue, err := strconv.Atoi(trimmed)
if err != nil || intValue < 0 {
return fmt.Errorf("%s 必须为大于等于 0 的整数", key)
}
return nil
}
func validateOptions(options []model.OpenFlareOption) error {
if len(options) == 0 {
return fmt.Errorf(errInvalidParams)
return errors.New(errInvalidParams)
}
state := buildOptionValidationState(options)
for _, option := range options {
if strings.TrimSpace(option.Key) == "" {
return fmt.Errorf(errInvalidParams)
return errors.New(errInvalidParams)
}
if err := validateOptionWithState(option, state); err != nil {
return err
+1
View File
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package origin defines shared error messages for origin management.
package origin
const (
+3 -1
View File
@@ -12,6 +12,8 @@ import (
"unicode"
)
const maxOriginHostnameLength = 253
func normalizeOriginAddress(raw string) string {
return strings.ToLower(strings.TrimSpace(raw))
}
@@ -29,7 +31,7 @@ func validateOriginAddress(address string) error {
if ip := net.ParseIP(address); ip != nil {
return nil
}
if len(address) > 253 {
if len(address) > maxOriginHostnameLength {
return errors.New(errOriginAddressInvalid)
}
labels := strings.Split(address, ".")
+16 -15
View File
@@ -1,27 +1,28 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package pages provides logics and management for OpenFlare static page deployments.
package pages
const (
errPagesProjectNotFound = "Pages 项目不存在"
errPagesSlugExists = "Pages 项目标识已存在"
errPagesNameRequired = "Pages 项目名称不能为空"
errPagesSlugInvalid = "Pages 项目标识只能包含小写字母、数字和连字符"
errPagesDeleteReferenced = "Pages 项目已被规则引用,不能删除"
errPagesDeploymentNotFound = "Pages 部署不存在"
errPagesDeploymentMismatch = "Pages 部署不属于该项目"
errPagesProjectNotFound = "pages 项目不存在"
errPagesSlugExists = "pages 项目标识已存在"
errPagesNameRequired = "pages 项目名称不能为空"
errPagesSlugInvalid = "pages 项目标识只能包含小写字母、数字和连字符"
errPagesDeleteReferenced = "pages 项目已被规则引用,不能删除"
errPagesDeploymentNotFound = "pages 部署不存在"
errPagesDeploymentMismatch = "pages 部署不属于该项目"
errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署"
errPagesPackageMissing = "缺少 Pages 部署包"
errPagesPackageNotZip = "Pages 部署包必须是 .zip 文件"
errPagesPackageInvalidZip = "Pages 部署包不是有效 zip 文件"
errPagesPackageEmpty = "Pages 部署包不能为空"
errPagesPackageNotZip = "pages 部署包必须是 .zip 文件"
errPagesPackageInvalidZip = "pages 部署包不是有效 zip 文件"
errPagesPackageEmpty = "pages 部署包不能为空"
errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空"
errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头"
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空"
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL"
errPagesPackagePathEmpty = "Pages 部署包路径为空"
errPagesPackageUploadMissing = "Pages 部署包上传记录不存在"
errPagesPackageNotInActiveConfig = "Pages 部署尚未进入激活配置"
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errPagesPackagePathEmpty = "pages 部署包路径为空"
errPagesPackageUploadMissing = "pages 部署包上传记录不存在"
errPagesPackageNotInActiveConfig = "pages 部署尚未进入激活配置"
errPagesInvalidSnapshotFormat = "配置快照格式无效"
)
+49 -36
View File
@@ -11,6 +11,7 @@ import (
"errors"
"fmt"
"io"
"math"
"mime/multipart"
"os"
"path"
@@ -24,11 +25,14 @@ import (
)
const (
pagesMaxDeploymentFiles = 1000
pagesMaxDeploymentBytes = 100 * 1024 * 1024
defaultPagesEntryFile = "index.html"
defaultPagesFallbackPath = "/index.html"
pagesMaxDeploymentFiles = 1000
pagesMaxDeploymentBytes = 100 * 1024 * 1024
defaultPagesEntryFile = "index.html"
defaultPagesFallbackPath = "/index.html"
pagesDeploymentUploadType = "openflare_pages_deployment"
mimeTypeApplicationZip = "application/zip"
pagesMaxPathLength = 512
bytesPerKiB = 1024
)
var pagesSlugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,126}[a-z0-9]$|^[a-z0-9]$`)
@@ -71,15 +75,15 @@ func validateAndNormalizePagesRootDir(raw string) (string, error) {
if value == "" {
return "", nil
}
if len(value) > 512 {
return "", errors.New("Pages 根目录长度不能超过 512")
if len(value) > pagesMaxPathLength {
return "", errors.New("pages 根目录长度不能超过 512") // error 消息首字母小写
}
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
return "", errors.New("Pages 根目录包含不支持的字符")
return "", errors.New("pages 根目录包含不支持的字符")
}
for _, r := range value {
if r <= 0x20 || r == 0x7f {
return "", errors.New("Pages 根目录不能包含空白或控制字符")
return "", errors.New("pages 根目录不能包含空白或控制字符")
}
}
cleaned := path.Clean(filepath.ToSlash(value))
@@ -88,7 +92,7 @@ func validateAndNormalizePagesRootDir(raw string) (string, error) {
}
for _, segment := range strings.Split(cleaned, "/") {
if segment == "." || segment == ".." {
return "", errors.New("Pages 根目录不能包含 . 或 .. 路径段")
return "", errors.New("pages 根目录不能包含 . 或 .. 路径段")
}
}
return strings.TrimPrefix(cleaned, "/"), nil
@@ -99,34 +103,34 @@ func normalizePagesFallbackPath(raw string) (string, error) {
if value == "" {
value = defaultPagesFallbackPath
}
if len(value) > 512 {
return "", errors.New("SPA fallback 回退路径长度不能超过 512")
if len(value) > pagesMaxPathLength {
return "", errors.New("spa fallback 回退路径长度不能超过 512")
}
if !strings.HasPrefix(value, "/") {
return "", errors.New("SPA fallback 回退路径必须以 / 开头")
return "", errors.New("spa fallback 回退路径必须以 / 开头")
}
if value == "/" || strings.HasSuffix(value, "/") {
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
return "", errors.New("spa fallback 回退路径必须指向具体文件")
}
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
return "", errors.New("SPA fallback 回退路径包含不支持的字符")
return "", errors.New("spa fallback 回退路径包含不支持的字符")
}
for _, r := range value {
if r <= 0x20 || r == 0x7f {
return "", errors.New("SPA fallback 回退路径不能包含空白或控制字符")
return "", errors.New("spa fallback 回退路径不能包含空白或控制字符")
}
}
for _, segment := range strings.Split(value, "/") {
if segment == "." || segment == ".." {
return "", errors.New("SPA fallback 回退路径不能包含 . 或 .. 路径段")
return "", errors.New("spa fallback 回退路径不能包含 . 或 .. 路径段")
}
}
cleaned := path.Clean(value)
if cleaned == "." || !strings.HasPrefix(cleaned, "/") {
return "", errors.New("SPA fallback 回退路径不合法")
return "", errors.New("spa fallback 回退路径不合法")
}
if cleaned == "/" || strings.HasSuffix(cleaned, "/") {
return "", errors.New("SPA fallback 回退路径必须指向具体文件")
return "", errors.New("spa fallback 回退路径必须指向具体文件")
}
return cleaned, nil
}
@@ -152,12 +156,12 @@ func persistPagesUploadTemp(fileHeader *multipart.FileHeader) (string, string, i
if err != nil {
return "", "", 0, err
}
defer file.Close()
defer func() { _ = file.Close() }()
temp, err := os.CreateTemp("", "openflare-pages-*.zip")
if err != nil {
return "", "", 0, err
}
defer temp.Close()
defer func() { _ = temp.Close() }()
hash := sha256.New()
limited := io.LimitReader(file, pagesMaxDeploymentBytes+1)
written, err := io.Copy(io.MultiWriter(temp, hash), limited)
@@ -167,7 +171,7 @@ func persistPagesUploadTemp(fileHeader *multipart.FileHeader) (string, string, i
}
if written > pagesMaxDeploymentBytes {
_ = os.Remove(temp.Name())
return "", "", 0, fmt.Errorf("Pages 部署包不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
return "", "", 0, fmt.Errorf("pages 部署包不能超过 %d MiB", pagesMaxDeploymentBytes/bytesPerKiB/bytesPerKiB)
}
return temp.Name(), hex.EncodeToString(hash.Sum(nil)), written, nil
}
@@ -180,11 +184,11 @@ func ingestPagesDeploymentPackage(
projectSlug string,
fileName string,
) (upload.IngestResult, error) {
file, err := os.Open(tempPath)
file, err := os.Open(tempPath) //nolint:gosec // tempPath is a validated pages deployment staging file
if err != nil {
return upload.IngestResult{}, err
}
defer file.Close()
defer func() { _ = file.Close() }()
systemUser := repository.GetSystemUser(ctx)
accessMode := 0
@@ -193,7 +197,7 @@ func ingestPagesDeploymentPackage(
Reader: file,
Size: size,
FileName: fileName,
MimeType: "application/zip",
MimeType: mimeTypeApplicationZip,
Extension: "zip",
Hash: checksum,
Type: pagesDeploymentUploadType,
@@ -270,7 +274,7 @@ func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploym
if err != nil {
return nil, errors.New(errPagesPackageInvalidZip)
}
defer reader.Close()
defer func() { _ = reader.Close() }()
commonPrefix, err := findCommonRootPrefix(reader.File)
if err != nil {
@@ -298,18 +302,18 @@ func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploym
normalizedPath = strings.TrimPrefix(normalizedPath, commonPrefix)
}
if item.FileInfo().Mode()&os.ModeSymlink != 0 {
return nil, fmt.Errorf("Pages 部署包不支持符号链接: %s", normalizedPath)
return nil, fmt.Errorf("pages 部署包不支持符号链接: %s", normalizedPath)
}
if item.UncompressedSize64 > pagesMaxDeploymentBytes {
return nil, fmt.Errorf("Pages 文件过大: %s", normalizedPath)
return nil, fmt.Errorf("pages 文件过大: %s", normalizedPath)
}
manifest.FileCount++
if manifest.FileCount > pagesMaxDeploymentFiles {
return nil, fmt.Errorf("Pages 部署文件数不能超过 %d", pagesMaxDeploymentFiles)
return nil, fmt.Errorf("pages 部署文件数不能超过 %d", pagesMaxDeploymentFiles)
}
manifest.TotalSize += int64(item.UncompressedSize64)
if manifest.TotalSize > pagesMaxDeploymentBytes {
return nil, fmt.Errorf("Pages 部署展开后不能超过 %d MiB", pagesMaxDeploymentBytes/1024/1024)
return nil, fmt.Errorf("pages 部署展开后不能超过 %d MiB", pagesMaxDeploymentBytes/bytesPerKiB/bytesPerKiB)
}
checksum, err := checksumZipFile(item)
if err != nil {
@@ -328,7 +332,7 @@ func inspectPagesZip(zipPath string, rootDir string, entryFile string) (*deploym
return nil, errors.New(errPagesPackageEmpty)
}
if !entrySeen {
return nil, fmt.Errorf("Pages 部署包缺少入口文件 %s", targetEntryPath)
return nil, fmt.Errorf("pages 部署包缺少入口文件 %s", targetEntryPath)
}
return manifest, nil
}
@@ -342,29 +346,38 @@ func normalizePagesZipPath(raw string) (string, bool, error) {
return "", true, nil
}
if strings.HasPrefix(name, "/") || path.IsAbs(name) {
return "", false, fmt.Errorf("Pages 部署包不能包含绝对路径: %s", raw)
return "", false, fmt.Errorf("pages 部署包不能包含绝对路径: %s", raw)
}
cleaned := path.Clean(name)
if cleaned == "." {
return "", true, nil
}
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || strings.Contains(cleaned, "/../") {
return "", false, fmt.Errorf("Pages 部署包路径不能逃逸目录: %s", raw)
return "", false, fmt.Errorf("pages 部署包路径不能逃逸目录: %s", raw)
}
return cleaned, false, nil
}
func pagesZipEntryCopyLimit(size uint64) (int64, error) {
if size == 0 || size > pagesMaxDeploymentBytes || size > uint64(math.MaxInt64) {
return 0, errors.New("pages file size out of bounds")
}
return int64(size), nil //nolint:gosec // size is bounded to math.MaxInt64 above
}
func checksumZipFile(item *zip.File) (string, error) {
file, err := item.Open()
if err != nil {
return "", err
}
defer file.Close()
defer func() { _ = file.Close() }()
hash := sha256.New()
if _, err = io.Copy(hash, file); err != nil {
limit, err := pagesZipEntryCopyLimit(item.UncompressedSize64)
if err != nil {
return "", err
}
if _, err = io.CopyN(hash, file, limit); err != nil {
return "", err
}
return hex.EncodeToString(hash.Sum(nil)), nil
}
+6 -6
View File
@@ -256,7 +256,7 @@ func UploadDeployment(ctx context.Context, projectID uint, fileHeader *multipart
if err != nil {
return nil, err
}
defer os.Remove(tempPath)
defer func() { _ = os.Remove(tempPath) }()
manifest, err := inspectPagesZip(tempPath, rootDir, entryFile)
if err != nil {
return nil, err
@@ -370,10 +370,10 @@ func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (*storage.Obj
}
obj, err := uploadstorage.OpenStoredObject(ctx, &uploadRecord)
if err != nil {
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
return nil, "", fmt.Errorf("pages 部署包不存在: %w", err)
}
if obj.ContentType == "" {
obj.ContentType = "application/zip"
obj.ContentType = mimeTypeApplicationZip
}
return obj, fileName, nil
}
@@ -382,17 +382,17 @@ func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (*storage.Obj
}
file, err := os.Open(deployment.ArtifactPath)
if err != nil {
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
return nil, "", fmt.Errorf("pages 部署包不存在: %w", err)
}
info, err := file.Stat()
if err != nil {
_ = file.Close()
return nil, "", fmt.Errorf("Pages 部署包不存在: %w", err)
return nil, "", fmt.Errorf("pages 部署包不存在: %w", err)
}
return &storage.Object{
Body: file,
ContentLength: info.Size(),
ContentType: "application/zip",
ContentType: mimeTypeApplicationZip,
}, fileName, nil
}
@@ -0,0 +1,178 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package proxy_route provides helpers for building proxy route configurations.
package proxy_route
import (
"context"
"encoding/json"
"errors"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
)
type proxyRouteJSONFields struct {
cacheRulesJSON string
upstreamsJSON string
customHeadersJSON string
certIDsJSON string
domainCertIDsJSON string
domainsJSON string
}
func resolveProxyRouteUpstreams(ctx context.Context, upstreamType string, input Input) (string, *uint, []string, error) {
switch upstreamType {
case proxyRouteUpstreamTypeTunnel, proxyRouteUpstreamTypePages:
if upstreamType == proxyRouteUpstreamTypePages {
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
return "", nil, nil, err
}
}
originURL := "http://127.0.0.1"
return originURL, nil, []string{originURL}, nil
default:
originURL, originID, err := resolveProxyRoutePrimaryOrigin(ctx, input)
if err != nil {
return "", nil, nil, err
}
upstreams, err := normalizeUpstreams(originURL, input.Upstreams)
if err != nil {
return "", nil, nil, err
}
return originURL, originID, upstreams, nil
}
}
func marshalProxyRouteJSONFields(
domains []string,
upstreams []string,
cacheRules []string,
customHeaders []CustomHeaderInput,
certIDs []uint,
domainCertIDs []uint,
) (*proxyRouteJSONFields, error) {
cacheRulesJSON, err := json.Marshal(cacheRules)
if err != nil {
return nil, err
}
upstreamsJSON, err := json.Marshal(upstreams)
if err != nil {
return nil, err
}
customHeadersJSON, err := json.Marshal(customHeaders)
if err != nil {
return nil, err
}
certIDsJSON, err := json.Marshal(certIDs)
if err != nil {
return nil, err
}
domainCertIDsJSON, err := json.Marshal(domainCertIDs)
if err != nil {
return nil, err
}
domainsJSON, err := json.Marshal(domains)
if err != nil {
return nil, err
}
return &proxyRouteJSONFields{
cacheRulesJSON: string(cacheRulesJSON),
upstreamsJSON: string(upstreamsJSON),
customHeadersJSON: string(customHeadersJSON),
certIDsJSON: string(certIDsJSON),
domainCertIDsJSON: string(domainCertIDsJSON),
domainsJSON: string(domainsJSON),
}, nil
}
func normalizeProxyRouteHTTPSInput(input *Input) {
if input.EnableHTTPS {
return
}
input.RedirectHTTP = false
input.CertID = nil
input.CertIDs = nil
input.DomainCertIDs = nil
}
func normalizeProxyRouteBasicAuth(input *Input) error {
if !input.BasicAuthEnabled {
input.BasicAuthUsername = ""
input.BasicAuthPassword = ""
return nil
}
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
return errors.New(errProxyRouteBasicAuth)
}
return nil
}
func populateProxyRouteFields(
route *model.ProxyRoute,
input Input,
siteName, domain string,
jsonFields *proxyRouteJSONFields,
originID *uint,
upstreams []string,
originHost, remark, cachePolicy string,
limitConnPerServer, limitConnPerIP int,
limitRate, upstreamType string,
) {
route.SiteName = siteName
route.Domain = domain
route.Domains = jsonFields.domainsJSON
route.OriginID = originID
route.OriginURL = upstreams[0]
route.OriginHost = originHost
route.Upstreams = jsonFields.upstreamsJSON
route.Enabled = input.Enabled
route.EnableHTTPS = input.EnableHTTPS
route.CertID = input.CertID
route.CertIDs = jsonFields.certIDsJSON
route.DomainCertIDs = jsonFields.domainCertIDsJSON
route.RedirectHTTP = input.RedirectHTTP
route.LimitConnPerServer = limitConnPerServer
route.LimitConnPerIP = limitConnPerIP
route.LimitRate = limitRate
route.CacheEnabled = input.CacheEnabled
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
route.CacheRules = jsonFields.cacheRulesJSON
route.CustomHeaders = jsonFields.customHeadersJSON
route.BasicAuthEnabled = input.BasicAuthEnabled
route.BasicAuthUsername = input.BasicAuthUsername
route.BasicAuthPassword = input.BasicAuthPassword
route.Remark = remark
route.UpstreamType = upstreamType
}
func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, upstreamType string, input Input) error {
switch upstreamType {
case proxyRouteUpstreamTypeTunnel:
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
if err != nil {
return err
}
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
return err
}
route.TunnelNodeID = tunnelNodeID
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
route.PagesProjectID = nil
case proxyRouteUpstreamTypePages:
route.TunnelNodeID = nil
route.TunnelTargetAddr = ""
route.TunnelTargetProtocol = ""
route.PagesProjectID = input.PagesProjectID
default:
route.TunnelNodeID = nil
route.TunnelTargetAddr = ""
route.TunnelTargetProtocol = ""
route.PagesProjectID = nil
}
return nil
}
@@ -0,0 +1,71 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"errors"
)
func normalizeExplicitDomainCertIDs(ctx context.Context, domains []string, rawDomainCertIDs []uint) ([]uint, []uint, *uint, error) {
if len(rawDomainCertIDs) != len(domains) {
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
}
normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs))
uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs))
seen := make(map[uint]struct{}, len(rawDomainCertIDs))
hasAssignedCertificate := false
for index, item := range rawDomainCertIDs {
if item == 0 {
continue
}
if _, err := lookupTLSCertificateByID(ctx, item); err != nil {
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
}
normalizedDomainCertIDs[index] = item
hasAssignedCertificate = true
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
uniqueCertIDs = append(uniqueCertIDs, item)
}
if !hasAssignedCertificate {
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
}
primaryCertID := &uniqueCertIDs[0]
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
}
func normalizeDerivedDomainCertIDs(
ctx context.Context,
domains []string,
normalizedCertIDs []uint,
) ([]uint, []uint, *uint, error) {
switch {
case len(normalizedCertIDs) == 0:
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
case len(normalizedCertIDs) == 1:
domainCertIDs := make([]uint, len(domains))
for index := range domainCertIDs {
domainCertIDs[index] = normalizedCertIDs[0]
}
primaryCertID := &normalizedCertIDs[0]
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
case len(normalizedCertIDs) == len(domains):
domainCertIDs := make([]uint, len(normalizedCertIDs))
copy(domainCertIDs, normalizedCertIDs)
primaryCertID := &normalizedCertIDs[0]
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
default:
domainCertIDs, err := deriveDomainCertIDsFromCertificateSet(ctx, domains, normalizedCertIDs)
if err != nil {
return nil, nil, nil, err
}
primaryCertID := &normalizedCertIDs[0]
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
}
}
+3 -3
View File
@@ -42,9 +42,9 @@ const (
errProxyRouteTunnelAddrReq = "tunnel_target_addr is required for tunnel upstream"
errProxyRouteTunnelProtocol = "tunnel_target_protocol must be http or https"
errProxyRoutePagesProjectReq = "pages_project_id is required for Pages upstream"
errProxyRoutePagesNotFound = "Pages 项目不存在"
errProxyRoutePagesDisabled = "Pages 项目未启用"
errProxyRoutePagesNoDeploy = "Pages 项目没有激活部署"
errProxyRoutePagesNotFound = "pages 项目不存在"
errProxyRoutePagesDisabled = "pages 项目未启用"
errProxyRoutePagesNoDeploy = "pages 项目没有激活部署"
errProxyRouteOriginSchemeOnly = "源站协议仅支持 http 或 https"
errProxyRouteOriginPort = "端口格式不合法"
errProxyRouteOriginPortEmpty = "端口不能为空"
+21 -124
View File
@@ -30,6 +30,13 @@ const (
proxyRouteCachePolicySuffix = "suffix"
proxyRouteCachePolicyPathPrefix = "path_prefix"
proxyRouteCachePolicyPathExact = "path_exact"
proxyRouteSchemeHTTP = "http"
proxyRouteSchemeHTTPS = "https"
proxyRouteUpstreamTypeTunnel = "tunnel"
proxyRouteUpstreamTypePages = "pages"
maxOriginHostnameLength = 253
originURIPathQueryParts = 2
)
type tlsCertificateRow struct {
@@ -100,7 +107,7 @@ func validateOriginAddress(address string) error {
if ip := net.ParseIP(address); ip != nil {
return nil
}
if len(address) > 253 {
if len(address) > maxOriginHostnameLength {
return errors.New(errProxyRouteOriginInvalid)
}
labels := strings.Split(address, ".")
@@ -136,7 +143,7 @@ func normalizeOriginPort(raw string) (string, error) {
func normalizeOriginScheme(raw string) (string, error) {
scheme := strings.ToLower(strings.TrimSpace(raw))
switch scheme {
case "http", "https":
case proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
return scheme, nil
default:
return "", errors.New(errProxyRouteOriginSchemeOnly)
@@ -187,7 +194,7 @@ func buildOriginURLFromParts(scheme, address, port, uri string) (string, error)
if strings.HasPrefix(normalizedURI, "?") {
parsed.RawQuery = strings.TrimPrefix(normalizedURI, "?")
} else {
pathQuery := strings.SplitN(normalizedURI, "?", 2)
pathQuery := strings.SplitN(normalizedURI, "?", originURIPathQueryParts)
parsed.Path = pathQuery[0]
if len(pathQuery) > 1 {
parsed.RawQuery = pathQuery[1]
@@ -465,65 +472,14 @@ func normalizeProxyRouteDomainCertificateIDs(
}
if len(rawDomainCertIDs) > 0 {
if len(rawDomainCertIDs) != len(domains) {
return nil, nil, nil, errors.New(errProxyRouteCertDomainLength)
}
normalizedDomainCertIDs := make([]uint, len(rawDomainCertIDs))
uniqueCertIDs := make([]uint, 0, len(rawDomainCertIDs))
seen := make(map[uint]struct{}, len(rawDomainCertIDs))
hasAssignedCertificate := false
for index, item := range rawDomainCertIDs {
if item == 0 {
continue
}
if _, err := lookupTLSCertificateByID(ctx, item); err != nil {
return nil, nil, nil, errors.New(errProxyRouteCertNotFound)
}
normalizedDomainCertIDs[index] = item
hasAssignedCertificate = true
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
uniqueCertIDs = append(uniqueCertIDs, item)
}
if !hasAssignedCertificate {
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
}
primaryCertID := &uniqueCertIDs[0]
return normalizedDomainCertIDs, uniqueCertIDs, primaryCertID, nil
return normalizeExplicitDomainCertIDs(ctx, domains, rawDomainCertIDs)
}
normalizedCertIDs, err := normalizeProxyRouteCertificateIDs(ctx, enableHTTPS, certID, certIDs)
if err != nil {
return nil, nil, nil, err
}
switch {
case len(normalizedCertIDs) == 0:
return nil, nil, nil, errors.New(errProxyRouteCertRequired)
case len(normalizedCertIDs) == 1:
domainCertIDs := make([]uint, len(domains))
for index := range domainCertIDs {
domainCertIDs[index] = normalizedCertIDs[0]
}
primaryCertID := &normalizedCertIDs[0]
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
case len(normalizedCertIDs) == len(domains):
domainCertIDs := make([]uint, len(normalizedCertIDs))
copy(domainCertIDs, normalizedCertIDs)
primaryCertID := &normalizedCertIDs[0]
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
default:
domainCertIDs, err := deriveDomainCertIDsFromCertificateSet(ctx, domains, normalizedCertIDs)
if err != nil {
return nil, nil, nil, err
}
primaryCertID := &normalizedCertIDs[0]
return domainCertIDs, normalizedCertIDs, primaryCertID, nil
}
return normalizeDerivedDomainCertIDs(ctx, domains, normalizedCertIDs)
}
func validateProxyRouteDomainCertificateCoverage(ctx context.Context, domains []string, domainCertIDs []uint) error {
@@ -637,65 +593,6 @@ func hasStructuredOriginInput(input Input) bool {
strings.TrimSpace(input.OriginURI) != ""
}
func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) {
if hasStructuredOriginInput(input) {
scheme, err := normalizeOriginScheme(input.OriginScheme)
if err != nil {
return "", nil, err
}
port, err := normalizeOriginPort(input.OriginPort)
if err != nil {
return "", nil, err
}
uri, err := normalizeOriginURI(input.OriginURI)
if err != nil {
return "", nil, err
}
if input.OriginID != nil && *input.OriginID != 0 {
origin, err := model.GetOriginByID(ctx, *input.OriginID)
if err != nil {
return "", nil, errors.New(errProxyRouteOriginNotFound)
}
originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri)
if err != nil {
return "", nil, err
}
return originURL, &origin.ID, nil
}
address := normalizeOriginAddress(input.OriginAddress)
if err := validateOriginAddress(address); err != nil {
return "", nil, err
}
originURL, err := buildOriginURLFromParts(scheme, address, port, uri)
if err != nil {
return "", nil, err
}
origin, err := getOrCreateOriginByAddress(ctx, address)
if err != nil {
return "", nil, err
}
return originURL, &origin.ID, nil
}
originURL := strings.TrimSpace(input.OriginURL)
if originURL == "" {
return "", nil, errors.New(errProxyRouteOriginEmpty)
}
address, err := extractOriginAddress(originURL)
if err != nil {
return "", nil, err
}
origin, findErr := model.GetOriginByAddress(ctx, address)
if findErr == nil {
return originURL, &origin.ID, nil
}
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
return "", nil, findErr
}
return originURL, nil, nil
}
func normalizeCustomHeaders(headers []CustomHeaderInput) ([]CustomHeaderInput, error) {
if len(headers) == 0 {
return []CustomHeaderInput{}, nil
@@ -942,7 +839,7 @@ func validateOriginURL(raw string) error {
if err != nil {
return errors.New(errProxyRouteOriginInvalid)
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
if parsed.Scheme != proxyRouteSchemeHTTP && parsed.Scheme != proxyRouteSchemeHTTPS {
return errors.New(errProxyRouteOriginScheme)
}
if parsed.Host == "" {
@@ -996,7 +893,7 @@ func validateTunnelRouteInput(ctx context.Context, tunnelNodeID *uint, targetAdd
return errors.New(errProxyRouteTunnelAddrReq)
}
switch strings.ToLower(strings.TrimSpace(targetProtocol)) {
case "", "http", "https":
case "", proxyRouteSchemeHTTP, proxyRouteSchemeHTTPS:
return nil
default:
return errors.New(errProxyRouteTunnelProtocol)
@@ -1025,10 +922,10 @@ func validatePagesRouteInput(ctx context.Context, projectID *uint) error {
func normalizeUpstreamType(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "tunnel":
return "tunnel"
case "pages":
return "pages"
case proxyRouteUpstreamTypeTunnel:
return proxyRouteUpstreamTypeTunnel
case proxyRouteUpstreamTypePages:
return proxyRouteUpstreamTypePages
default:
return "direct"
}
@@ -1036,9 +933,9 @@ func normalizeUpstreamType(raw string) string {
func normalizeTunnelTargetProtocol(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "https":
return "https"
case proxyRouteSchemeHTTPS:
return proxyRouteSchemeHTTPS
default:
return "http"
return proxyRouteSchemeHTTP
}
}
+25 -107
View File
@@ -5,7 +5,6 @@ package proxy_route
import (
"context"
"encoding/json"
"errors"
"strings"
"time"
@@ -168,28 +167,9 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
siteName := normalizeProxyRouteSiteNameInput(route, input.SiteName, domain)
upstreamType := normalizeUpstreamType(input.UpstreamType)
var originURL string
var originID *uint
var upstreams []string
if upstreamType == "tunnel" {
originURL = "http://127.0.0.1"
upstreams = []string{originURL}
} else if upstreamType == "pages" {
if err := validatePagesRouteInput(ctx, input.PagesProjectID); err != nil {
return nil, err
}
originURL = "http://127.0.0.1"
upstreams = []string{originURL}
} else {
originURL, originID, err = resolveProxyRoutePrimaryOrigin(ctx, input)
if err != nil {
return nil, err
}
upstreams, err = normalizeUpstreams(originURL, input.Upstreams)
if err != nil {
return nil, err
}
_, originID, upstreams, err := resolveProxyRouteUpstreams(ctx, upstreamType, input)
if err != nil {
return nil, err
}
originHost := strings.TrimSpace(input.OriginHost)
remark := strings.TrimSpace(input.Remark)
@@ -215,25 +195,7 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
return nil, err
}
cacheRulesJSON, err := json.Marshal(cacheRules)
if err != nil {
return nil, err
}
upstreamsJSON, err := json.Marshal(upstreams)
if err != nil {
return nil, err
}
customHeadersJSON, err := json.Marshal(customHeaders)
if err != nil {
return nil, err
}
if !input.EnableHTTPS {
input.RedirectHTTP = false
input.CertID = nil
input.CertIDs = nil
input.DomainCertIDs = nil
}
normalizeProxyRouteHTTPSInput(&input)
domainCertIDs, certIDs, primaryCertID, err := normalizeProxyRouteDomainCertificateIDs(
ctx,
domains,
@@ -248,15 +210,7 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
if err := validateProxyRouteDomainCertificateCoverage(ctx, domains, domainCertIDs); err != nil {
return nil, err
}
certIDsJSON, err := json.Marshal(certIDs)
if err != nil {
return nil, err
}
domainCertIDsJSON, err := json.Marshal(domainCertIDs)
if err != nil {
return nil, err
}
domainsJSON, err := json.Marshal(domains)
jsonFields, err := marshalProxyRouteJSONFields(domains, upstreams, cacheRules, customHeaders, certIDs, domainCertIDs)
if err != nil {
return nil, err
}
@@ -277,67 +231,31 @@ func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input)
return nil, errors.New(errProxyRouteRedirectHTTP)
}
if input.BasicAuthEnabled {
input.BasicAuthUsername = strings.TrimSpace(input.BasicAuthUsername)
input.BasicAuthPassword = strings.TrimSpace(input.BasicAuthPassword)
if input.BasicAuthUsername == "" || input.BasicAuthPassword == "" {
return nil, errors.New(errProxyRouteBasicAuth)
}
} else {
input.BasicAuthUsername = ""
input.BasicAuthPassword = ""
if err := normalizeProxyRouteBasicAuth(&input); err != nil {
return nil, err
}
if route == nil {
route = &model.ProxyRoute{}
}
route.SiteName = siteName
route.Domain = domain
route.Domains = string(domainsJSON)
route.OriginID = originID
route.OriginURL = upstreams[0]
route.OriginHost = originHost
route.Upstreams = string(upstreamsJSON)
route.Enabled = input.Enabled
route.EnableHTTPS = input.EnableHTTPS
route.CertID = input.CertID
route.CertIDs = string(certIDsJSON)
route.DomainCertIDs = string(domainCertIDsJSON)
route.RedirectHTTP = input.RedirectHTTP
route.LimitConnPerServer = limitConnPerServer
route.LimitConnPerIP = limitConnPerIP
route.LimitRate = limitRate
route.CacheEnabled = input.CacheEnabled
route.CachePolicy = normalizeCachePolicy(input.CacheEnabled, cachePolicy)
route.CacheRules = string(cacheRulesJSON)
route.CustomHeaders = string(customHeadersJSON)
route.BasicAuthEnabled = input.BasicAuthEnabled
route.BasicAuthUsername = input.BasicAuthUsername
route.BasicAuthPassword = input.BasicAuthPassword
route.Remark = remark
route.UpstreamType = upstreamType
if upstreamType == "tunnel" {
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
if err != nil {
return nil, err
}
if err := validateTunnelRouteInput(ctx, tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
return nil, err
}
route.TunnelNodeID = tunnelNodeID
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
route.PagesProjectID = nil
} else if upstreamType == "pages" {
route.TunnelNodeID = nil
route.TunnelTargetAddr = ""
route.TunnelTargetProtocol = ""
route.PagesProjectID = input.PagesProjectID
} else {
route.TunnelNodeID = nil
route.TunnelTargetAddr = ""
route.TunnelTargetProtocol = ""
route.PagesProjectID = nil
populateProxyRouteFields(
route,
input,
siteName,
domain,
jsonFields,
originID,
upstreams,
originHost,
remark,
cachePolicy,
limitConnPerServer,
limitConnPerIP,
limitRate,
upstreamType,
)
if err := applyProxyRouteUpstreamType(ctx, route, upstreamType, input); err != nil {
return nil, err
}
return route, nil
}
@@ -0,0 +1,85 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package proxy_route
import (
"context"
"errors"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
func resolveStructuredOriginInput(ctx context.Context, input Input) (string, *uint, error) {
scheme, err := normalizeOriginScheme(input.OriginScheme)
if err != nil {
return "", nil, err
}
port, err := normalizeOriginPort(input.OriginPort)
if err != nil {
return "", nil, err
}
uri, err := normalizeOriginURI(input.OriginURI)
if err != nil {
return "", nil, err
}
if input.OriginID != nil && *input.OriginID != 0 {
return resolveOriginByID(ctx, scheme, port, uri, *input.OriginID)
}
return resolveOriginByAddress(ctx, scheme, port, uri, input.OriginAddress)
}
func resolveOriginByID(ctx context.Context, scheme, port, uri string, originID uint) (string, *uint, error) {
origin, err := model.GetOriginByID(ctx, originID)
if err != nil {
return "", nil, errors.New(errProxyRouteOriginNotFound)
}
originURL, err := buildOriginURLFromParts(scheme, origin.Address, port, uri)
if err != nil {
return "", nil, err
}
return originURL, &origin.ID, nil
}
func resolveOriginByAddress(ctx context.Context, scheme, port, uri, rawAddress string) (string, *uint, error) {
address := normalizeOriginAddress(rawAddress)
if err := validateOriginAddress(address); err != nil {
return "", nil, err
}
originURL, err := buildOriginURLFromParts(scheme, address, port, uri)
if err != nil {
return "", nil, err
}
origin, err := getOrCreateOriginByAddress(ctx, address)
if err != nil {
return "", nil, err
}
return originURL, &origin.ID, nil
}
func resolveLegacyOriginInput(ctx context.Context, originURL string) (string, *uint, error) {
if originURL == "" {
return "", nil, errors.New(errProxyRouteOriginEmpty)
}
address, err := extractOriginAddress(originURL)
if err != nil {
return "", nil, err
}
origin, findErr := model.GetOriginByAddress(ctx, address)
if findErr == nil {
return originURL, &origin.ID, nil
}
if !errors.Is(findErr, gorm.ErrRecordNotFound) {
return "", nil, findErr
}
return originURL, nil, nil
}
func resolveProxyRoutePrimaryOrigin(ctx context.Context, input Input) (string, *uint, error) {
if hasStructuredOriginInput(input) {
return resolveStructuredOriginInput(ctx, input)
}
return resolveLegacyOriginInput(ctx, strings.TrimSpace(input.OriginURL))
}
+2
View File
@@ -1,9 +1,11 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package relay provides relay node management and authentication for the OpenFlare platform.
package relay
const (
//nolint:gosec // error message text, not a credential
errAgentTokenInvalid = "无权进行此操作,Agent Token 无效"
errRelayNodeTypeMismatch = "此节点不是 TunnelRelay 类型"
)
+9 -4
View File
@@ -10,12 +10,17 @@ import (
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
relayStatusUnhealthy = "unhealthy"
releaseChannelStable = "stable"
)
func normalizeRelayStatus(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case "healthy":
return "healthy"
case "unhealthy":
return "unhealthy"
case relayStatusUnhealthy:
return relayStatusUnhealthy
default:
return "unknown"
}
@@ -25,7 +30,7 @@ func normalizeReleaseChannel(channel string) string {
if strings.ToLower(strings.TrimSpace(channel)) == "preview" {
return "preview"
}
return "stable"
return releaseChannelStable
}
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
@@ -98,7 +103,7 @@ func BuildSettings(node *model.OpenFlareNode, updateNow bool, updateChannel, upd
autoUpdate = node.AutoUpdateEnabled
}
if strings.TrimSpace(updateChannel) == "" {
updateChannel = "stable"
updateChannel = releaseChannelStable
}
return &Settings{
HeartbeatInterval: model.AgentHeartbeatInterval,
+3 -3
View File
@@ -41,7 +41,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
"last_seen_at": now,
"status": nodeStatusOnline,
"update_requested": false,
"update_channel": "stable",
"update_channel": releaseChannelStable,
"update_tag": "",
}
if payload.Name != "" && strings.TrimSpace(node.Name) == "" {
@@ -55,7 +55,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
if !previous.UpdateRequested {
delete(changes, "update_requested")
}
if previous.UpdateChannel == "stable" {
if previous.UpdateChannel == releaseChannelStable {
delete(changes, "update_channel")
}
if previous.UpdateTag == "" {
@@ -66,7 +66,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
node.ExtVersion = payload.ExtVersion
node.RelayStatus = payload.RelayStatus
node.UpdateRequested = false
node.UpdateChannel = "stable"
node.UpdateChannel = releaseChannelStable
node.UpdateTag = ""
lastSeen := now
node.LastSeenAt = &lastSeen
+3 -3
View File
@@ -16,8 +16,8 @@ import (
const ctxRelayNodeKey = "relay_node"
// RelayAuth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type.
func RelayAuth() gin.HandlerFunc {
// Auth authenticates relay requests using X-Agent-Token and verifies tunnel_relay type.
func Auth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader("X-Agent-Token"))
node, err := authenticateAccessToken(c.Request.Context(), token)
@@ -46,4 +46,4 @@ func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlar
return nil, err
}
return node, nil
}
}
@@ -55,7 +55,7 @@ func TestRelayAuthMissingToken(t *testing.T) {
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) {
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
c.Status(http.StatusOK)
})
@@ -74,7 +74,7 @@ func TestRelayAuthRejectsWrongNodeType(t *testing.T) {
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) {
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
c.Status(http.StatusOK)
})
@@ -93,7 +93,7 @@ func TestRelayAuthAcceptsTunnelRelay(t *testing.T) {
gin.SetMode(gin.TestMode)
engine := gin.New()
engine.GET("/relay/test", RelayAuth(), func(c *gin.Context) {
engine.GET("/relay/test", Auth(), func(c *gin.Context) {
authNode, ok := c.Get(ctxRelayNodeKey)
require.True(t, ok)
assert.Equal(t, node.NodeID, authNode.(*model.OpenFlareNode).NodeID)
@@ -24,7 +24,7 @@ func reconcileRelayHealthEvents(ctx context.Context, nodeID string, relayStatus
relayFrpsUnhealthyEventType: {},
}
events := []agent.NodeHealthEvent{}
if relayStatus == "unhealthy" {
if relayStatus == relayStatusUnhealthy {
events = append(events, agent.NodeHealthEvent{
EventType: relayFrpsUnhealthyEventType,
Severity: "critical",
@@ -5,8 +5,17 @@ package relay
import pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
// ProxyStat is an alias for protocol.RelayProxyStat.
type ProxyStat = pkgprotocol.RelayProxyStat
// HeartbeatPayload is an alias for protocol.RelayHeartbeatPayload.
type HeartbeatPayload = pkgprotocol.RelayHeartbeatPayload
// Config is an alias for protocol.RelayConfig.
type Config = pkgprotocol.RelayConfig
// Settings is an alias for protocol.RelaySettings.
type Settings = pkgprotocol.RelaySettings
type HeartbeatResponse = pkgprotocol.RelayHeartbeatResponse
// HeartbeatResponse is an alias for protocol.RelayHeartbeatResponse.
type HeartbeatResponse = pkgprotocol.RelayHeartbeatResponse
@@ -90,7 +90,7 @@ func CleanupDatabaseObservability(ctx context.Context, input DatabaseCleanupInpu
}
// RunDatabaseAutoCleanupOnce runs retention-based cleanup for all observability targets.
func RunDatabaseAutoCleanupOnce(now time.Time) (*DatabaseAutoCleanupSummary, error) {
func RunDatabaseAutoCleanupOnce(ctx context.Context, now time.Time) (*DatabaseAutoCleanupSummary, error) {
if !model.DatabaseAutoCleanupEnabled {
return nil, nil
}
@@ -99,7 +99,6 @@ func RunDatabaseAutoCleanupOnce(now time.Time) (*DatabaseAutoCleanupSummary, err
}
retentionDays := model.DatabaseAutoCleanupRetentionDays
ctx := context.Background()
results := make([]DatabaseCleanupResult, 0, len(databaseCleanupTargets))
for _, target := range []string{
DatabaseCleanupTargetAccessLogs,
@@ -134,7 +134,7 @@ func TestRunDatabaseAutoCleanupOnceDeletesAllObservabilityTargets(t *testing.T)
model.DatabaseAutoCleanupRetentionDays = previousRetentionDays
})
summary, err := RunDatabaseAutoCleanupOnce(now)
summary, err := RunDatabaseAutoCleanupOnce(ctx, now)
require.NoError(t, err)
require.NotNil(t, summary)
require.Len(t, summary.Results, 3)
+16 -10
View File
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package acme implements ACME certificate issuance and renewal.
package acme
import (
@@ -25,22 +26,27 @@ import (
"github.com/go-acme/lego/v4/registration"
)
// AcmeUser implements lego's user interface.
type AcmeUser struct {
const dnsChallengePrecheckDelay = 20 * time.Second
// User implements lego's registration.User interface for ACME account management.
type User struct {
Email string
Registration *registration.Resource
key crypto.PrivateKey
}
func (u *AcmeUser) GetEmail() string {
// GetEmail returns the email address associated with this ACME account.
func (u *User) GetEmail() string {
return u.Email
}
func (u *AcmeUser) GetRegistration() *registration.Resource {
// GetRegistration returns the ACME account registration resource.
func (u *User) GetRegistration() *registration.Resource {
return u.Registration
}
func (u *AcmeUser) GetPrivateKey() crypto.PrivateKey {
// GetPrivateKey returns the private key used to authenticate with the ACME server.
func (u *User) GetPrivateKey() crypto.PrivateKey {
return u.key
}
@@ -88,7 +94,7 @@ func encodePrivateKey(key crypto.PrivateKey) (string, error) {
}
// GetOrCreateLegoClient returns a configured lego client and optional new account credentials.
func GetOrCreateLegoClient(acmeEmail, privateKeyPEM, accountURL string, keyAlgorithm string) (*lego.Client, *AcmeUser, string, string, error) {
func GetOrCreateLegoClient(acmeEmail, privateKeyPEM, accountURL string, keyAlgorithm string) (*lego.Client, *User, string, string, error) {
var privateKey crypto.PrivateKey
var err error
var newPrivateKeyPEM string
@@ -111,7 +117,7 @@ func GetOrCreateLegoClient(acmeEmail, privateKeyPEM, accountURL string, keyAlgor
}
}
user := &AcmeUser{
user := &User{
Email: acmeEmail,
key: privateKey,
}
@@ -197,12 +203,12 @@ func SetupDNSProvider(client *lego.Client, dnsType, dnsAuth string, dns1, dns2 s
}
if disableCNAME {
opts = append(opts, dns01.DisableCompletePropagationRequirement())
opts = append(opts, dns01.DisableAuthoritativeNssPropagationRequirement())
}
if skipDNS {
opts = append(opts, dns01.WrapPreCheck(func(domain, fqdn, value string, check dns01.PreCheckFunc) (bool, error) {
time.Sleep(20 * time.Second)
opts = append(opts, dns01.WrapPreCheck(func(_, _, _ string, _ dns01.PreCheckFunc) (bool, error) {
time.Sleep(dnsChallengePrecheckDelay)
return true, nil
}))
}
@@ -58,7 +58,7 @@ func TestApplyCertificateReturnsApplying(t *testing.T) {
Name: "Test ACME Cert",
PrimaryDomain: "example.com",
OtherDomains: "*.example.com",
DnsAccountID: dnsAccount.ID,
DNSAccountID: dnsAccount.ID,
KeyAlgorithm: "RSA2048",
AutoRenew: true,
})
+1
View File
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package tls defines shared error messages for certificate management.
package tls
const (
+1 -1
View File
@@ -31,7 +31,7 @@ func readMultipartFile(fileHeader *multipart.FileHeader) (string, error) {
if err != nil {
return "", err
}
defer file.Close()
defer func() { _ = file.Close() }()
data, err := io.ReadAll(file)
if err != nil {
return "", err
+34 -30
View File
@@ -32,7 +32,7 @@ type CertificateContent struct {
Remark string `json:"remark"`
Provider string `json:"provider"`
AcmeAccountID uint `json:"acme_account_id"`
DnsAccountID uint `json:"dns_account_id"`
DNSAccountID uint `json:"dns_account_id"`
KeyAlgorithm string `json:"key_algorithm"`
AutoRenew bool `json:"auto_renew"`
PrimaryDomain string `json:"primary_domain"`
@@ -50,7 +50,7 @@ type ApplyInput struct {
Name string `json:"name"`
Remark string `json:"remark"`
AcmeAccountID uint `json:"acme_account_id"`
DnsAccountID uint `json:"dns_account_id"`
DNSAccountID uint `json:"dns_account_id"`
KeyAlgorithm string `json:"key_algorithm"`
AutoRenew bool `json:"auto_renew"`
PrimaryDomain string `json:"primary_domain"`
@@ -96,7 +96,7 @@ func GetCertificateContent(ctx context.Context, id uint) (*CertificateContent, e
Remark: certificate.Remark,
Provider: certificate.Provider,
AcmeAccountID: certificate.AcmeAccountID,
DnsAccountID: certificate.DnsAccountID,
DNSAccountID: certificate.DNSAccountID,
KeyAlgorithm: certificate.KeyAlgorithm,
AutoRenew: certificate.AutoRenew,
PrimaryDomain: certificate.PrimaryDomain,
@@ -179,7 +179,7 @@ func DeleteCertificate(ctx context.Context, id uint) error {
// ApplyCertificate 申请 ACME 证书。
func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertificate, error) {
cert := &model.TLSCertificate{
Provider: "acme",
Provider: tlsProviderACME,
CertPEM: " ",
KeyPEM: " ",
}
@@ -195,7 +195,8 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific
}
go func(c *model.TLSCertificate) {
_ = obtainTLSCertificate(context.Background(), c)
asyncCtx := context.WithoutCancel(ctx)
_ = obtainTLSCertificate(asyncCtx, c)
}(cert)
return sanitizeCertificateForResponse(cert), nil
@@ -207,7 +208,7 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod
if err != nil {
return nil, err
}
if cert.Provider != "acme" {
if cert.Provider != tlsProviderACME {
return nil, errors.New(errCertificateOnlyACME)
}
fillAcmeCertificateFields(cert, input)
@@ -222,7 +223,8 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod
}
go func(c *model.TLSCertificate) {
_ = obtainTLSCertificate(context.Background(), c)
asyncCtx := context.WithoutCancel(ctx)
_ = obtainTLSCertificate(asyncCtx, c)
}(cert)
return sanitizeCertificateForResponse(cert), nil
@@ -237,7 +239,7 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*
if cert.Provider != "upload" {
return nil, errors.New(errCertificateOnlyUploadConvert)
}
if cert.ApplyStatus == "applying" {
if cert.ApplyStatus == tlsApplyStatusApplying {
return nil, errors.New(errCertificateAlreadyApplying)
}
fillAcmeCertificateFields(cert, input)
@@ -253,17 +255,18 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*
}
go func(c *model.TLSCertificate) {
if err := obtainTLSCertificate(context.Background(), c); err != nil {
asyncCtx := context.WithoutCancel(ctx)
if err := obtainTLSCertificate(asyncCtx, c); err != nil {
return
}
latest, err := model.GetTLSCertificateByID(context.Background(), c.ID)
latest, err := model.GetTLSCertificateByID(asyncCtx, c.ID)
if err != nil {
return
}
latest.Provider = "acme"
latest.ApplyStatus = "ready"
latest.Provider = tlsProviderACME
latest.ApplyStatus = tlsApplyStatusReady
latest.ApplyMessage = ""
_ = model.SaveTLSCertificate(context.Background(), latest)
_ = model.SaveTLSCertificate(asyncCtx, latest)
}(cert)
return sanitizeCertificateForResponse(cert), nil
@@ -275,15 +278,16 @@ func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, erro
if err != nil {
return nil, err
}
if cert.Provider != "acme" {
if cert.Provider != tlsProviderACME {
return nil, errors.New(errCertificateOnlyACMERenew)
}
go func(c *model.TLSCertificate) {
_ = obtainTLSCertificate(context.Background(), c)
asyncCtx := context.WithoutCancel(ctx)
_ = obtainTLSCertificate(asyncCtx, c)
}(cert)
cert.ApplyStatus = "applying"
cert.ApplyStatus = tlsApplyStatusApplying
cert.ApplyMessage = ""
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
return nil, err
@@ -362,7 +366,7 @@ func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) {
return sanitizeAcmeAccountForResponse(account), nil
}
func buildCertificate(ctx context.Context, existing *model.TLSCertificate, input CertificateInput) (*model.TLSCertificate, error) {
func buildCertificate(_ context.Context, existing *model.TLSCertificate, input CertificateInput) (*model.TLSCertificate, error) {
name := strings.TrimSpace(input.Name)
certPEM := strings.TrimSpace(input.CertPEM)
keyPEM := strings.TrimSpace(input.KeyPEM)
@@ -391,7 +395,7 @@ func buildCertificate(ctx context.Context, existing *model.TLSCertificate, input
if existing == nil {
existing = &model.TLSCertificate{
Provider: "upload",
ApplyStatus: "ready",
ApplyStatus: tlsApplyStatusReady,
}
}
existing.Name = name
@@ -407,7 +411,7 @@ func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) {
cert.Name = strings.TrimSpace(input.Name)
cert.Remark = strings.TrimSpace(input.Remark)
cert.AcmeAccountID = input.AcmeAccountID
cert.DnsAccountID = input.DnsAccountID
cert.DNSAccountID = input.DNSAccountID
cert.KeyAlgorithm = input.KeyAlgorithm
cert.AutoRenew = input.AutoRenew
cert.PrimaryDomain = strings.TrimSpace(input.PrimaryDomain)
@@ -416,7 +420,7 @@ func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) {
cert.SkipDNS = input.SkipDNS
cert.DNS1 = strings.TrimSpace(input.DNS1)
cert.DNS2 = strings.TrimSpace(input.DNS2)
cert.ApplyStatus = "applying"
cert.ApplyStatus = tlsApplyStatusApplying
}
func ensureCertificateNotReferenced(ctx context.Context, id uint) error {
@@ -457,26 +461,26 @@ func sanitizeCertificateForResponse(certificate *model.TLSCertificate) *model.TL
if certificate == nil {
return nil
}
copy := *certificate
copy.CertPEM = ""
copy.KeyPEM = ""
return &copy
certCopy := *certificate
certCopy.CertPEM = ""
certCopy.KeyPEM = ""
return &certCopy
}
func sanitizeDNSAccountForResponse(account *model.DNSAccount) *model.DNSAccount {
if account == nil {
return nil
}
copy := *account
copy.Authorization = ""
return &copy
certCopy := *account
certCopy.Authorization = ""
return &certCopy
}
func sanitizeAcmeAccountForResponse(account *model.AcmeAccount) *model.AcmeAccount {
if account == nil {
return nil
}
copy := *account
copy.PrivateKey = ""
return &copy
certCopy := *account
certCopy.PrivateKey = ""
return &certCopy
}
@@ -17,6 +17,9 @@ import (
const (
managedDomainMatchTypeExact = "exact"
managedDomainMatchTypeWildcard = "wildcard"
maxManagedDomainLength = 253
minManagedDomainLabelCount = 2
)
// ManagedDomainInput 托管域名创建/更新请求。
@@ -182,11 +185,11 @@ func validateHostname(domain string) error {
if domain == "" {
return errors.New(errManagedDomainRequired)
}
if len(domain) > 253 {
if len(domain) > maxManagedDomainLength {
return errors.New(errManagedDomainInvalid)
}
labels := strings.Split(domain, ".")
if len(labels) < 2 {
if len(labels) < minManagedDomainLabelCount {
return errors.New(errManagedDomainInvalid)
}
for _, label := range labels {
+16 -47
View File
@@ -13,7 +13,12 @@ import (
"github.com/Rain-kl/Wavelet/internal/model"
)
const acmeRenewLeadTime = 7 * 24 * time.Hour
const (
acmeRenewLeadTime = 7 * 24 * time.Hour
tlsProviderACME = "acme"
tlsApplyStatusApplying = "applying"
tlsApplyStatusReady = "ready"
)
var obtainTLSCertificate = obtainCertificate
@@ -27,24 +32,17 @@ func SetObtainCertificateFuncForTest(fn func(context.Context, *model.TLSCertific
}
func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error {
cert.ApplyStatus = "applying"
cert.ApplyStatus = tlsApplyStatusApplying
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
return err
}
acmeAccount, err := model.GetAcmeAccountByID(ctx, cert.AcmeAccountID)
acmeAccount, err := resolveAcmeAccount(ctx, cert)
if err != nil {
acmeAccount, err = model.GetDefaultAcmeAccount(ctx)
if err != nil {
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get ACME account: %v", err))
}
cert.AcmeAccountID = acmeAccount.ID
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
return err
}
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get ACME account: %v", err))
}
dnsAccount, err := model.GetDNSAccountByID(ctx, cert.DnsAccountID)
dnsAccount, err := model.GetDNSAccountByID(ctx, cert.DNSAccountID)
if err != nil {
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get DNS account: %v", err))
}
@@ -75,47 +73,18 @@ func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error {
domains,
)
if (newPrivateKeyPEM != "" && acmePrivateKey != newPrivateKeyPEM) || (newAccountURL != "" && acmeAccount.URL != newAccountURL) {
if newPrivateKeyPEM != "" {
sealedKey, sealErr := sealSensitive(newPrivateKeyPEM)
if sealErr != nil {
return updateCertError(ctx, cert, fmt.Sprintf("Failed to seal ACME account key: %v", sealErr))
}
acmeAccount.PrivateKey = sealedKey
}
if newAccountURL != "" {
acmeAccount.URL = newAccountURL
}
if acmeAccount.ID == 0 {
if dbErr := model.CreateAcmeAccountRecord(ctx, acmeAccount); dbErr != nil {
return updateCertError(ctx, cert, fmt.Sprintf("Failed to create ACME account: %v", dbErr))
}
} else if dbErr := model.SaveAcmeAccount(ctx, acmeAccount); dbErr != nil {
return updateCertError(ctx, cert, fmt.Sprintf("Failed to save ACME account: %v", dbErr))
}
cert.AcmeAccountID = acmeAccount.ID
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
return err
}
if err := persistAcmeAccountUpdates(ctx, cert, acmeAccount, newAccountURL, newPrivateKeyPEM, acmePrivateKey); err != nil {
return updateCertError(ctx, cert, err.Error())
}
if err != nil {
return updateCertError(ctx, cert, err.Error())
}
sealedKey, err := sealSensitive(result.KeyPEM)
if err != nil {
return updateCertError(ctx, cert, fmt.Sprintf("Failed to seal certificate key: %v", err))
if err := saveObtainedCertificate(ctx, cert, result); err != nil {
return updateCertError(ctx, cert, err.Error())
}
cert.CertPEM = result.CertPEM
cert.KeyPEM = sealedKey
cert.NotBefore = result.NotBefore
cert.NotAfter = result.NotAfter
cert.ApplyStatus = "ready"
cert.ApplyMessage = ""
return model.SaveTLSCertificate(ctx, cert)
return nil
}
func updateCertError(ctx context.Context, cert *model.TLSCertificate, message string) error {
@@ -155,7 +124,7 @@ func splitAcmeDomains(primaryDomain, otherDomains string) []string {
func CertificatesDueForRenewal(certificates []model.TLSCertificate, now time.Time) []model.TLSCertificate {
due := make([]model.TLSCertificate, 0)
for _, cert := range certificates {
if !cert.AutoRenew || cert.Provider != "acme" || cert.ApplyStatus == "applying" {
if !cert.AutoRenew || cert.Provider != tlsProviderACME || cert.ApplyStatus == tlsApplyStatusApplying {
continue
}
if cert.NotAfter.IsZero() {
@@ -0,0 +1,74 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package tls
import (
"context"
"fmt"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme"
"github.com/Rain-kl/Wavelet/internal/model"
)
func resolveAcmeAccount(ctx context.Context, cert *model.TLSCertificate) (*model.AcmeAccount, error) {
acmeAccount, err := model.GetAcmeAccountByID(ctx, cert.AcmeAccountID)
if err == nil {
return acmeAccount, nil
}
acmeAccount, err = model.GetDefaultAcmeAccount(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get ACME account: %w", err)
}
cert.AcmeAccountID = acmeAccount.ID
if err := model.SaveTLSCertificate(ctx, cert); err != nil {
return nil, err
}
return acmeAccount, nil
}
func persistAcmeAccountUpdates(
ctx context.Context,
cert *model.TLSCertificate,
acmeAccount *model.AcmeAccount,
newAccountURL, newPrivateKeyPEM, acmePrivateKey string,
) error {
accountChanged := (newPrivateKeyPEM != "" && acmePrivateKey != newPrivateKeyPEM) ||
(newAccountURL != "" && acmeAccount.URL != newAccountURL)
if !accountChanged {
return nil
}
if newPrivateKeyPEM != "" && acmePrivateKey != newPrivateKeyPEM {
sealedKey, sealErr := sealSensitive(newPrivateKeyPEM)
if sealErr != nil {
return fmt.Errorf("failed to seal ACME account key: %w", sealErr)
}
acmeAccount.PrivateKey = sealedKey
}
if newAccountURL != "" {
acmeAccount.URL = newAccountURL
}
if acmeAccount.ID == 0 {
if dbErr := model.CreateAcmeAccountRecord(ctx, acmeAccount); dbErr != nil {
return fmt.Errorf("failed to create ACME account: %w", dbErr)
}
} else if dbErr := model.SaveAcmeAccount(ctx, acmeAccount); dbErr != nil {
return fmt.Errorf("failed to save ACME account: %w", dbErr)
}
cert.AcmeAccountID = acmeAccount.ID
return model.SaveTLSCertificate(ctx, cert)
}
func saveObtainedCertificate(ctx context.Context, cert *model.TLSCertificate, result *acme.CertificateResult) error {
sealedKey, err := sealSensitive(result.KeyPEM)
if err != nil {
return fmt.Errorf("failed to seal certificate key: %w", err)
}
cert.CertPEM = result.CertPEM
cert.KeyPEM = sealedKey
cert.NotBefore = result.NotBefore
cert.NotAfter = result.NotAfter
cert.ApplyStatus = tlsApplyStatusReady
cert.ApplyMessage = ""
return model.SaveTLSCertificate(ctx, cert)
}
+30 -27
View File
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package uptimekuma provides a Socket.IO client and sync implementation for Uptime Kuma.
package uptimekuma
import (
@@ -16,28 +17,30 @@ import (
"time"
)
// UptimeKumaMonitor represents a monitor entry from Uptime Kuma.
type UptimeKumaMonitor struct {
ID int `json:"id"`
Name string `json:"name"`
Url string `json:"url"`
Type string `json:"type"`
Interval int `json:"interval"`
MaxRetries int `json:"maxretries"`
RetryInterval int `json:"retryInterval"`
Timeout int `json:"timeout"`
Tags []UptimeKumaTag `json:"tags"`
const emitAckTimeout = 10 * time.Second
// Monitor represents a monitor entry from Uptime Kuma.
type Monitor struct {
ID int `json:"id"`
Name string `json:"name"`
URL string `json:"url"`
Type string `json:"type"`
Interval int `json:"interval"`
MaxRetries int `json:"maxretries"`
RetryInterval int `json:"retryInterval"`
Timeout int `json:"timeout"`
Tags []Tag `json:"tags"`
}
// UptimeKumaTag represents a tag attached to a monitor.
type UptimeKumaTag struct {
// Tag represents a tag attached to a monitor.
type Tag struct {
ID int `json:"tag_id"`
Name string `json:"name"`
Color string `json:"color"`
}
// UptimeKumaTagItem represents a tag returned by getTags.
type UptimeKumaTagItem struct {
// TagItem represents a tag returned by getTags.
type TagItem struct {
ID int `json:"id"`
Name string `json:"name"`
Color string `json:"color"`
@@ -55,7 +58,7 @@ type SocketIOClient struct {
closeOnce sync.Once
monitorListMutex sync.RWMutex
monitorList map[string]UptimeKumaMonitor
monitorList map[string]Monitor
monitorListChan chan struct{}
monitorListOnce sync.Once
@@ -76,7 +79,7 @@ func NewSocketIOClient(baseURL string) *SocketIOClient {
ackChanMap: make(map[int]chan string),
doneChan: make(chan struct{}),
monitorListChan: make(chan struct{}),
monitorList: make(map[string]UptimeKumaMonitor),
monitorList: make(map[string]Monitor),
ctx: ctx,
cancel: cancel,
}
@@ -95,7 +98,7 @@ func (c *SocketIOClient) Connect() error {
slog.Error("Uptime Kuma handshake connection failed", "url", u, "error", err)
return fmt.Errorf("handshake request failed: %w", err)
}
defer resp.Body.Close()
defer func() { _ = resp.Body.Close() }()
bs, err := io.ReadAll(resp.Body)
if err != nil {
@@ -130,7 +133,7 @@ func (c *SocketIOClient) Connect() error {
slog.Error("Uptime Kuma namespace connect request failed", "sid", c.sid, "error", err)
return fmt.Errorf("namespace connect failed: %w", err)
}
respConnect.Body.Close()
_ = respConnect.Body.Close()
slog.Debug("Namespace connected successfully to Uptime Kuma", "sid", c.sid)
go c.pollLoop()
@@ -164,7 +167,7 @@ func (c *SocketIOClient) pollLoop() {
}
bs, err := io.ReadAll(resp.Body)
resp.Body.Close()
_ = resp.Body.Close()
if err != nil {
slog.Error("Failed to read Uptime Kuma polling body", "sid", c.sid, "error", err)
c.err = err
@@ -218,7 +221,7 @@ func (c *SocketIOClient) sendPong() {
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
resp, err := c.httpClient.Do(req)
if err == nil {
resp.Body.Close()
_ = resp.Body.Close()
}
}
@@ -232,7 +235,7 @@ func (c *SocketIOClient) handleEvent(payload string) {
return
}
if eventName == "monitorList" {
var list map[string]UptimeKumaMonitor
var list map[string]Monitor
if err := json.Unmarshal(arr[1], &list); err == nil {
c.monitorListMutex.Lock()
c.monitorList = list
@@ -309,13 +312,13 @@ func (c *SocketIOClient) Emit(event string, args ...any) (string, error) {
slog.Error("Failed to send Emit request", "event", event, "ackID", id, "error", err)
return "", err
}
resp.Body.Close()
_ = resp.Body.Close()
select {
case result := <-ch:
slog.Debug("Received Ack for event", "event", event, "ackID", id, "response", result)
return result, nil
case <-time.After(10 * time.Second):
case <-time.After(emitAckTimeout):
c.ackMutex.Lock()
delete(c.ackChanMap, id)
c.ackMutex.Unlock()
@@ -344,11 +347,11 @@ func (c *SocketIOClient) GetMonitorListChan() <-chan struct{} {
}
// GetMonitorList returns a copy of the current monitor list.
func (c *SocketIOClient) GetMonitorList() map[string]UptimeKumaMonitor {
func (c *SocketIOClient) GetMonitorList() map[string]Monitor {
c.monitorListMutex.RLock()
defer c.monitorListMutex.RUnlock()
m := make(map[string]UptimeKumaMonitor, len(c.monitorList))
m := make(map[string]Monitor, len(c.monitorList))
for k, v := range c.monitorList {
m[k] = v
}
@@ -372,7 +375,7 @@ func ParseAckResponse(response string, target any) error {
if errMsg == "" {
errMsg = "unknown error from Uptime Kuma"
}
return fmt.Errorf("Uptime Kuma error response: %s", errMsg)
return fmt.Errorf("uptime Kuma error response: %s", errMsg)
}
}
+20 -92
View File
@@ -10,17 +10,18 @@ import (
"log/slog"
"strings"
"sync/atomic"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
)
const uptimeKumaTagOpenFlare = "OpenFlare"
var isSyncing atomic.Bool
// SyncToUptimeKuma synchronizes enabled proxy routes to Uptime Kuma monitors.
func SyncToUptimeKuma(ctx context.Context) error {
if !model.UptimeKumaEnabled {
return fmt.Errorf("Uptime Kuma integration is disabled")
return fmt.Errorf("uptime Kuma integration is disabled")
}
if !isSyncing.CompareAndSwap(false, true) {
@@ -28,14 +29,9 @@ func SyncToUptimeKuma(ctx context.Context) error {
}
defer isSyncing.Store(false)
kumaURL := strings.TrimSpace(model.UptimeKumaUrl)
kumaUsername := strings.TrimSpace(model.UptimeKumaUsername)
kumaPassword := strings.TrimSpace(model.UptimeKumaPassword)
if kumaURL == "" || kumaUsername == "" || kumaPassword == "" {
return fmt.Errorf(
"Uptime Kuma URL, username, or password is not configured (URL: %q, Username: %q, PasswordLength: %d)",
kumaURL, kumaUsername, len(kumaPassword),
)
kumaURL, kumaUsername, kumaPassword, err := validateUptimeKumaConfig()
if err != nil {
return err
}
slog.Info("Starting Uptime Kuma sync process",
@@ -54,88 +50,20 @@ func SyncToUptimeKuma(ctx context.Context) error {
return err
}
slog.Debug("Connecting to Uptime Kuma socket endpoint", "url", kumaURL)
client := NewSocketIOClient(kumaURL)
if err := client.Connect(); err != nil {
slog.Error("Failed to connect to Uptime Kuma endpoint", "url", kumaURL, "error", err)
return fmt.Errorf("failed to connect to Uptime Kuma: %w", err)
client, err := connectAndLoginUptimeKuma(kumaURL, kumaUsername, kumaPassword)
if err != nil {
return err
}
defer client.Close()
slog.Debug("Sending login request to Uptime Kuma", "username", kumaUsername)
loginPayload := map[string]string{
"username": kumaUsername,
"password": kumaPassword,
}
loginAck, err := client.Emit("login", loginPayload)
if err != nil {
slog.Error("Failed to send login request to Uptime Kuma", "username", kumaUsername, "error", err)
return fmt.Errorf("login request failed: %w", err)
}
var loginResult struct {
Ok bool `json:"ok"`
}
if err := ParseAckResponse(loginAck, &loginResult); err != nil || !loginResult.Ok {
slog.Error("Uptime Kuma login verification failed", "username", kumaUsername, "error", err)
return fmt.Errorf("login failed: %w", err)
}
slog.Debug("Successfully logged into Uptime Kuma", "username", kumaUsername)
slog.Debug("Waiting for monitor list push from Uptime Kuma")
select {
case <-client.GetMonitorListChan():
slog.Debug("Received monitor list from Uptime Kuma")
case <-time.After(5 * time.Second):
slog.Error("Timeout waiting for Uptime Kuma monitorList push event")
return fmt.Errorf("timeout waiting for monitorList event from Uptime Kuma")
}
openFlareTagID, err := ensureOpenFlareTag(client)
if err != nil {
return err
}
existingOpenFlareMonitors := filterOpenFlareMonitors(client.GetMonitorList(), openFlareTagID)
expectedSitesMap := make(map[string]bool)
for _, route := range expectedRoutes {
expectedSitesMap[route.SiteName] = true
targetURL, urlErr := routeMonitorURL(route)
if urlErr != nil {
slog.Error("Failed to resolve monitor URL", "name", route.SiteName, "error", urlErr)
continue
}
existing, exists := existingOpenFlareMonitors[route.SiteName]
if !exists {
if err := createMonitor(client, route.SiteName, targetURL, openFlareTagID); err != nil {
slog.Error("Failed to add monitor to Uptime Kuma", "name", route.SiteName, "error", err)
}
continue
}
if monitorNeedsUpdate(existing, targetURL) {
if err := updateMonitor(client, existing.ID, route.SiteName, targetURL); err != nil {
slog.Error("Failed to edit monitor in Uptime Kuma", "name", route.SiteName, "error", err)
}
}
}
for name, monitor := range existingOpenFlareMonitors {
if expectedSitesMap[name] {
continue
}
slog.Info("Deleting monitor in Uptime Kuma", "name", name, "monitorID", monitor.ID)
deleteAck, err := client.Emit("deleteMonitor", monitor.ID)
if err != nil {
slog.Error("Failed to delete monitor in Uptime Kuma", "name", name, "monitorID", monitor.ID, "error", err)
continue
}
if err := ParseAckResponse(deleteAck, nil); err != nil {
slog.Error("Failed to parse delete monitor result", "name", name, "monitorID", monitor.ID, "error", err)
}
}
expectedSitesMap := syncRouteMonitors(client, expectedRoutes, existingOpenFlareMonitors, openFlareTagID)
removeStaleMonitors(client, existingOpenFlareMonitors, expectedSitesMap)
return nil
}
@@ -178,8 +106,8 @@ func ensureOpenFlareTag(client *SocketIOClient) (int, error) {
}
var tagsResult struct {
Ok bool `json:"ok"`
Tags []UptimeKumaTagItem `json:"tags"`
Ok bool `json:"ok"`
Tags []TagItem `json:"tags"`
}
if err := ParseAckResponse(tagsAck, &tagsResult); err != nil {
slog.Error("Failed to parse tags response from Uptime Kuma", "error", err)
@@ -187,7 +115,7 @@ func ensureOpenFlareTag(client *SocketIOClient) (int, error) {
}
for _, tag := range tagsResult.Tags {
if tag.Name == "OpenFlare" {
if tag.Name == uptimeKumaTagOpenFlare {
slog.Debug("Found existing OpenFlare tag", "tag_id", tag.ID)
return tag.ID, nil
}
@@ -195,7 +123,7 @@ func ensureOpenFlareTag(client *SocketIOClient) (int, error) {
slog.Debug("OpenFlare tag not found, creating new tag")
addTagAck, err := client.Emit("addTag", map[string]string{
"name": "OpenFlare",
"name": uptimeKumaTagOpenFlare,
"color": "#4f46e5",
})
if err != nil {
@@ -218,12 +146,12 @@ func ensureOpenFlareTag(client *SocketIOClient) (int, error) {
return tagResult.Tag.ID, nil
}
func filterOpenFlareMonitors(monitors map[string]UptimeKumaMonitor, openFlareTagID int) map[string]UptimeKumaMonitor {
existingOpenFlareMonitors := make(map[string]UptimeKumaMonitor)
func filterOpenFlareMonitors(monitors map[string]Monitor, openFlareTagID int) map[string]Monitor {
existingOpenFlareMonitors := make(map[string]Monitor)
for _, monitor := range monitors {
hasOpenFlareTag := false
for _, tag := range monitor.Tags {
if tag.Name == "OpenFlare" || tag.ID == openFlareTagID {
if tag.Name == uptimeKumaTagOpenFlare || tag.ID == openFlareTagID {
hasOpenFlareTag = true
break
}
@@ -295,8 +223,8 @@ func monitorPayload(id int, name, targetURL string) map[string]any {
return payload
}
func monitorNeedsUpdate(existing UptimeKumaMonitor, targetURL string) bool {
return existing.Url != targetURL ||
func monitorNeedsUpdate(existing Monitor, targetURL string) bool {
return existing.URL != targetURL ||
existing.Interval != model.UptimeKumaInterval ||
existing.MaxRetries != model.UptimeKumaRetry ||
existing.RetryInterval != model.UptimeKumaRetryInterval ||
@@ -0,0 +1,112 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package uptimekuma
import (
"fmt"
"log/slog"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
)
const monitorListWaitTimeout = 5 * time.Second
func validateUptimeKumaConfig() (string, string, string, error) {
kumaURL := strings.TrimSpace(model.UptimeKumaURL)
kumaUsername := strings.TrimSpace(model.UptimeKumaUsername)
kumaPassword := strings.TrimSpace(model.UptimeKumaPassword)
if kumaURL == "" || kumaUsername == "" || kumaPassword == "" {
return kumaURL, kumaUsername, kumaPassword, fmt.Errorf(
"uptime Kuma URL, username, or password is not configured (URL: %q, Username: %q, PasswordLength: %d)",
kumaURL, kumaUsername, len(kumaPassword),
)
}
return kumaURL, kumaUsername, kumaPassword, nil
}
func connectAndLoginUptimeKuma(kumaURL, kumaUsername, kumaPassword string) (*SocketIOClient, error) {
slog.Debug("Connecting to Uptime Kuma socket endpoint", "url", kumaURL)
client := NewSocketIOClient(kumaURL)
if err := client.Connect(); err != nil {
slog.Error("Failed to connect to Uptime Kuma endpoint", "url", kumaURL, "error", err)
return nil, fmt.Errorf("failed to connect to Uptime Kuma: %w", err)
}
slog.Debug("Sending login request to Uptime Kuma", "username", kumaUsername)
loginAck, err := client.Emit("login", map[string]string{
"username": kumaUsername,
"password": kumaPassword,
})
if err != nil {
client.Close()
slog.Error("Failed to send login request to Uptime Kuma", "username", kumaUsername, "error", err)
return nil, fmt.Errorf("login request failed: %w", err)
}
var loginResult struct {
Ok bool `json:"ok"`
}
if err := ParseAckResponse(loginAck, &loginResult); err != nil || !loginResult.Ok {
client.Close()
slog.Error("Uptime Kuma login verification failed", "username", kumaUsername, "error", err)
return nil, fmt.Errorf("login failed: %w", err)
}
slog.Debug("Successfully logged into Uptime Kuma", "username", kumaUsername)
slog.Debug("Waiting for monitor list push from Uptime Kuma")
select {
case <-client.GetMonitorListChan():
slog.Debug("Received monitor list from Uptime Kuma")
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 client, nil
}
func syncRouteMonitors(client *SocketIOClient, expectedRoutes []*model.ProxyRoute, existingMonitors map[string]Monitor, openFlareTagID int) map[string]bool {
expectedSitesMap := make(map[string]bool, len(expectedRoutes))
for _, route := range expectedRoutes {
expectedSitesMap[route.SiteName] = true
targetURL, urlErr := routeMonitorURL(route)
if urlErr != nil {
slog.Error("Failed to resolve monitor URL", "name", route.SiteName, "error", urlErr)
continue
}
existing, exists := existingMonitors[route.SiteName]
if !exists {
if err := createMonitor(client, route.SiteName, targetURL, openFlareTagID); err != nil {
slog.Error("Failed to add monitor to Uptime Kuma", "name", route.SiteName, "error", err)
}
continue
}
if monitorNeedsUpdate(existing, targetURL) {
if err := updateMonitor(client, existing.ID, route.SiteName, targetURL); err != nil {
slog.Error("Failed to edit monitor in Uptime Kuma", "name", route.SiteName, "error", err)
}
}
}
return expectedSitesMap
}
func removeStaleMonitors(client *SocketIOClient, existingMonitors map[string]Monitor, expectedSitesMap map[string]bool) {
for name, monitor := range existingMonitors {
if expectedSitesMap[name] {
continue
}
slog.Info("Deleting monitor in Uptime Kuma", "name", name, "monitorID", monitor.ID)
deleteAck, err := client.Emit("deleteMonitor", monitor.ID)
if err != nil {
slog.Error("Failed to delete monitor in Uptime Kuma", "name", name, "monitorID", monitor.ID, "error", err)
continue
}
if err := ParseAckResponse(deleteAck, nil); err != nil {
slog.Error("Failed to parse delete monitor result", "name", name, "monitorID", monitor.ID, "error", err)
}
}
}
@@ -131,7 +131,7 @@ func setupSyncTestDB(t *testing.T) func() {
func backupUptimeKumaConfig() func() {
oldEnabled := model.UptimeKumaEnabled
oldURL := model.UptimeKumaUrl
oldURL := model.UptimeKumaURL
oldUsername := model.UptimeKumaUsername
oldPassword := model.UptimeKumaPassword
oldScope := model.UptimeKumaMonitorScope
@@ -143,7 +143,7 @@ func backupUptimeKumaConfig() func() {
return func() {
model.UptimeKumaEnabled = oldEnabled
model.UptimeKumaUrl = oldURL
model.UptimeKumaURL = oldURL
model.UptimeKumaUsername = oldUsername
model.UptimeKumaPassword = oldPassword
model.UptimeKumaMonitorScope = oldScope
@@ -228,7 +228,7 @@ func TestSyncToUptimeKumaSuccess(t *testing.T) {
defer server.Close()
model.UptimeKumaEnabled = true
model.UptimeKumaUrl = server.URL
model.UptimeKumaURL = server.URL
model.UptimeKumaUsername = "admin"
model.UptimeKumaPassword = "password"
model.UptimeKumaMonitorScope = "all"
@@ -313,7 +313,7 @@ func TestSyncToUptimeKumaSelectedScope(t *testing.T) {
defer server.Close()
model.UptimeKumaEnabled = true
model.UptimeKumaUrl = server.URL
model.UptimeKumaURL = server.URL
model.UptimeKumaUsername = "admin"
model.UptimeKumaPassword = "password"
model.UptimeKumaMonitorScope = "selected"
+1 -5
View File
@@ -1,9 +1,5 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package waf defines shared error messages for WAF management.
package waf
const (
errWAFRuleGroupNotFound = "WAF 规则组不存在"
errWAFIPGroupNotFound = "IP 组不存在"
)
+9 -5
View File
@@ -95,7 +95,7 @@ func syncOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGro
}
func syncIPGroupSubscription(ctx context.Context, group *model.OpenFlareWAFIPGroup, now time.Time) (*IPGroupSyncResult, error) {
content, err := downloadIPGroupSubscription(group.SubscriptionURL)
content, err := downloadIPGroupSubscription(ctx, group.SubscriptionURL)
if err != nil {
recordIPGroupSyncFailure(ctx, group, now, err)
return nil, err
@@ -266,7 +266,7 @@ func evaluateParsedIPGroupAutoConfig(ctx context.Context, config ipGroupAutoConf
if item.StatusCode >= 400 && item.StatusCode < 500 {
acc.clientErrorCount++
}
if item.StatusCode >= 500 {
if item.StatusCode >= http.StatusInternalServerError {
acc.serverErrorCount++
}
if hostIsIPLiteral(item.Host) {
@@ -341,16 +341,20 @@ func hostIsIPLiteral(value string) bool {
return ok
}
func downloadIPGroupSubscription(rawURL string) ([]byte, error) {
func downloadIPGroupSubscription(ctx context.Context, rawURL string) ([]byte, error) {
if err := validateSubscriptionURL(rawURL); err != nil {
return nil, err
}
client := http.Client{Timeout: 15 * time.Second}
resp, err := client.Get(rawURL)
req, err := http.NewRequestWithContext(ctx, "GET", rawURL, nil)
if err != nil {
return nil, fmt.Errorf("下载订阅失败: %w", err)
}
defer resp.Body.Close()
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("下载订阅失败: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("订阅返回状态码 %d", resp.StatusCode)
}
+26 -73
View File
@@ -8,10 +8,8 @@ import (
"encoding/json"
"errors"
"fmt"
"net"
"net/netip"
"net/url"
"regexp"
"sort"
"strings"
"time"
@@ -38,6 +36,9 @@ const (
defaultWAFIPGroupAutoLookbackMinutes = 60
minWAFIPGroupSyncIntervalMinutes = 5
maxWAFIPGroupSyncIntervalMinutes = 43200
minPoWSessionTTLSeconds = 60
minPoWChallengeTTLSeconds = 30
)
// RuleGroupInput is the create/update payload for WAF rule groups.
@@ -986,82 +987,34 @@ func defaultPoWConfig() PoWConfig {
}
func normalizePoWConfig(enabled bool, raw string) (PoWConfig, error) {
if !enabled {
return defaultPoWConfig(), nil
cfg, err := parsePoWConfigRaw(enabled, raw)
if err != nil {
return cfg, err
}
cfg := defaultPoWConfig()
text := strings.TrimSpace(raw)
if text != "" && text != "{}" {
if err := json.Unmarshal([]byte(text), &cfg); err != nil {
return cfg, errors.New("pow_config 格式无效")
}
if err := validatePoWCoreSettings(cfg); err != nil {
return cfg, err
}
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
return cfg, errors.New("pow_config.difficulty 必须在 1-16 之间")
if err := validatePoWCIDRs(cfg.Whitelist.IPCidrs, "白名单"); err != nil {
return cfg, err
}
if !powAlgorithmValues[cfg.Algorithm] {
return cfg, errors.New("pow_config.algorithm 必须为 fast 或 slow")
if err := validatePoWCIDRs(cfg.Blacklist.IPCidrs, "黑名单"); err != nil {
return cfg, err
}
if cfg.SessionTTL < 60 {
return cfg, errors.New("pow_config.session_ttl 不能小于 60 秒")
if err := validatePoWPathRegexes(cfg.Whitelist.PathRegexes, "白名单"); err != nil {
return cfg, err
}
if cfg.ChallengeTTL < 30 {
return cfg, errors.New("pow_config.challenge_ttl 不能小于 30 秒")
if err := validatePoWPathRegexes(cfg.Blacklist.PathRegexes, "黑名单"); err != nil {
return cfg, err
}
for _, cidr := range cfg.Whitelist.IPCidrs {
if _, _, err := net.ParseCIDR(cidr); err != nil {
return cfg, fmt.Errorf("pow_config 白名单 IP CIDR 格式无效: %s", cidr)
}
if err := validatePoWIPs(cfg.Whitelist.IPs, "白名单"); err != nil {
return cfg, err
}
for _, cidr := range cfg.Blacklist.IPCidrs {
if _, _, err := net.ParseCIDR(cidr); err != nil {
return cfg, fmt.Errorf("pow_config 黑名单 IP CIDR 格式无效: %s", cidr)
}
if err := validatePoWIPs(cfg.Blacklist.IPs, "黑名单"); err != nil {
return cfg, err
}
for _, re := range cfg.Whitelist.PathRegexes {
if _, err := regexp.Compile(re); err != nil {
return cfg, fmt.Errorf("pow_config 白名单路径正则格式无效: %s", re)
}
if err := validatePoWListMutualExclusion(cfg); err != nil {
return cfg, err
}
for _, re := range cfg.Blacklist.PathRegexes {
if _, err := regexp.Compile(re); err != nil {
return cfg, fmt.Errorf("pow_config 黑名单路径正则格式无效: %s", re)
}
}
for _, ip := range cfg.Whitelist.IPs {
if net.ParseIP(ip) == nil {
return cfg, fmt.Errorf("pow_config 白名单 IP 格式无效: %s", ip)
}
}
for _, ip := range cfg.Blacklist.IPs {
if net.ParseIP(ip) == nil {
return cfg, fmt.Errorf("pow_config 黑名单 IP 格式无效: %s", ip)
}
}
type dimension struct {
name string
wl []string
bl []string
}
dimensions := []dimension{
{"IP", cfg.Whitelist.IPs, cfg.Blacklist.IPs},
{"IP CIDR", cfg.Whitelist.IPCidrs, cfg.Blacklist.IPCidrs},
{"路径", cfg.Whitelist.Paths, cfg.Blacklist.Paths},
{"路径正则", cfg.Whitelist.PathRegexes, cfg.Blacklist.PathRegexes},
{"User-Agent", cfg.Whitelist.UserAgents, cfg.Blacklist.UserAgents},
}
for _, dim := range dimensions {
if len(dim.wl) > 0 && len(dim.bl) > 0 {
return cfg, fmt.Errorf("pow_config %s 不能同时配置白名单和黑名单", dim.name)
}
}
return cfg, nil
}
@@ -1111,11 +1064,11 @@ func parseIPGroupAutoConfig(raw json.RawMessage) (ipGroupAutoConfig, error) {
if config.LookbackMinutes <= 0 {
config.LookbackMinutes = defaultWAFIPGroupAutoLookbackMinutes
}
if config.LookbackMinutes < 5 {
config.LookbackMinutes = 5
if config.LookbackMinutes < minWAFIPGroupSyncIntervalMinutes {
config.LookbackMinutes = minWAFIPGroupSyncIntervalMinutes
}
if config.LookbackMinutes > 43200 {
config.LookbackMinutes = 43200
if config.LookbackMinutes > maxWAFIPGroupSyncIntervalMinutes {
config.LookbackMinutes = maxWAFIPGroupSyncIntervalMinutes
}
if config.TTL == 0 {
config.TTL = -1
@@ -0,0 +1,92 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package waf
import (
"encoding/json"
"errors"
"fmt"
"net"
"regexp"
"strings"
)
func parsePoWConfigRaw(enabled bool, raw string) (PoWConfig, error) {
if !enabled {
return defaultPoWConfig(), nil
}
cfg := defaultPoWConfig()
text := strings.TrimSpace(raw)
if text == "" || text == "{}" {
return cfg, nil
}
if err := json.Unmarshal([]byte(text), &cfg); err != nil {
return cfg, errors.New("pow_config 格式无效")
}
return cfg, nil
}
func validatePoWCoreSettings(cfg PoWConfig) error {
if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
return errors.New("pow_config.difficulty 必须在 1-16 之间")
}
if !powAlgorithmValues[cfg.Algorithm] {
return errors.New("pow_config.algorithm 必须为 fast 或 slow")
}
if cfg.SessionTTL < minPoWSessionTTLSeconds {
return errors.New("pow_config.session_ttl 不能小于 60 秒")
}
if cfg.ChallengeTTL < minPoWChallengeTTLSeconds {
return errors.New("pow_config.challenge_ttl 不能小于 30 秒")
}
return nil
}
func validatePoWCIDRs(cidrs []string, listName string) error {
for _, cidr := range cidrs {
if _, _, err := net.ParseCIDR(cidr); err != nil {
return fmt.Errorf("pow_config %s IP CIDR 格式无效: %s", listName, cidr)
}
}
return nil
}
func validatePoWPathRegexes(regexes []string, listName string) error {
for _, re := range regexes {
if _, err := regexp.Compile(re); err != nil {
return fmt.Errorf("pow_config %s路径正则格式无效: %s", listName, re)
}
}
return nil
}
func validatePoWIPs(ips []string, listName string) error {
for _, ip := range ips {
if net.ParseIP(ip) == nil {
return fmt.Errorf("pow_config %s IP 格式无效: %s", listName, ip)
}
}
return nil
}
func validatePoWListMutualExclusion(cfg PoWConfig) error {
type dimension struct {
name string
wl []string
bl []string
}
dimensions := []dimension{
{"IP", cfg.Whitelist.IPs, cfg.Blacklist.IPs},
{"IP CIDR", cfg.Whitelist.IPCidrs, cfg.Blacklist.IPCidrs},
{"路径", cfg.Whitelist.Paths, cfg.Blacklist.Paths},
{"路径正则", cfg.Whitelist.PathRegexes, cfg.Blacklist.PathRegexes},
{"User-Agent", cfg.Whitelist.UserAgents, cfg.Blacklist.UserAgents},
}
for _, dim := range dimensions {
if len(dim.wl) > 0 && len(dim.bl) > 0 {
return fmt.Errorf("pow_config %s 不能同时配置白名单和黑名单", dim.name)
}
}
return nil
}
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package websocket manages persistent WebSocket connections between the OpenFlare server and its agents.
package websocket
import (
@@ -67,7 +68,7 @@ func ServeAgent(c *gin.Context, nodeID string, onStatus AgentStatusHandler) {
nodeID: nodeID,
remoteAddr: c.Request.RemoteAddr,
conn: conn,
send: make(chan Message, 16),
send: make(chan Message, wsChannelBuf),
done: make(chan struct{}),
onStatus: onStatus,
}
@@ -202,15 +203,15 @@ func (c *agentClient) readPump() {
}
func agentWSReadTimeout() time.Duration {
timeout := 90 * time.Second
if timeout < 30*time.Second {
return 30 * time.Second
timeout := wsReadDeadline
if timeout < minAgentWSReadTimeout {
return minAgentWSReadTimeout
}
return timeout
}
func (c *agentClient) writePump() {
ticker := time.NewTicker(30 * time.Second)
ticker := time.NewTicker(wsPingInterval)
defer ticker.Stop()
for {
@@ -218,7 +219,7 @@ func (c *agentClient) writePump() {
case <-c.done:
return
case message := <-c.send:
_ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteDeadline))
if err := c.conn.WriteJSON(message); err != nil {
slog.Debug("agent ws write failed", "node_id", c.nodeID, "error", err)
c.close()
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import "time"
const (
wsChannelBuf = 16
wsPingInterval = 30 * time.Second
wsReadDeadline = 90 * time.Second
wsWriteDeadline = 10 * time.Second
minAgentWSReadTimeout = 30 * time.Second
)
@@ -4,7 +4,6 @@
package websocket
import (
"encoding/json"
"log/slog"
"sync"
"time"
@@ -58,7 +57,7 @@ func ServeFlared(c *gin.Context, nodeID string) {
client := &flaredClient{
nodeID: nodeID,
conn: conn,
send: make(chan Message, 16),
send: make(chan Message, wsChannelBuf),
done: make(chan struct{}),
}
defaultFlaredHub.register(client)
@@ -136,37 +135,11 @@ func SendFlaredPong(nodeID string) bool {
}
func (c *flaredClient) readPump() {
defer c.close()
_ = c.conn.SetReadDeadline(time.Now().Add(90 * time.Second))
c.conn.SetPongHandler(func(string) error {
return c.conn.SetReadDeadline(time.Now().Add(90 * time.Second))
})
for {
_, data, err := c.conn.ReadMessage()
if err != nil {
slog.Debug("flared ws read closed", "node_id", c.nodeID, "error", err)
return
}
var message Message
if err = json.Unmarshal(data, &message); err != nil {
slog.Debug("flared ws invalid message", "node_id", c.nodeID, "error", err)
continue
}
switch message.Type {
case messageTypePing:
_ = SendFlaredPong(c.nodeID)
case flaredMessageTypePong:
default:
slog.Debug("flared ws unsupported message", "node_id", c.nodeID, "type", message.Type)
}
}
runReadPump(c.nodeID, c.conn, c.close, "flared ws", SendFlaredPong, flaredMessageTypePong)
}
func (c *flaredClient) writePump() {
ticker := time.NewTicker(30 * time.Second)
ticker := time.NewTicker(wsPingInterval)
defer ticker.Stop()
for {
@@ -174,7 +147,7 @@ func (c *flaredClient) writePump() {
case <-c.done:
return
case message := <-c.send:
_ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteDeadline))
if err := c.conn.WriteJSON(message); err != nil {
slog.Debug("flared ws write failed", "node_id", c.nodeID, "error", err)
c.close()
@@ -0,0 +1,49 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package websocket
import (
"encoding/json"
"log/slog"
"time"
"github.com/gorilla/websocket"
)
func runReadPump(
nodeID string,
conn *websocket.Conn,
closeFn func(),
logLabel string,
sendPong func(string) bool,
clientPongType string,
) {
defer closeFn()
_ = conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
conn.SetPongHandler(func(string) error {
return conn.SetReadDeadline(time.Now().Add(wsReadDeadline))
})
for {
_, data, err := conn.ReadMessage()
if err != nil {
slog.Debug(logLabel+" read closed", "node_id", nodeID, "error", err)
return
}
var message Message
if err = json.Unmarshal(data, &message); err != nil {
slog.Debug(logLabel+" invalid message", "node_id", nodeID, "error", err)
continue
}
switch message.Type {
case messageTypePing:
_ = sendPong(nodeID)
case clientPongType:
default:
slog.Debug(logLabel+" unsupported message", "node_id", nodeID, "type", message.Type)
}
}
}
+4 -31
View File
@@ -4,7 +4,6 @@
package websocket
import (
"encoding/json"
"log/slog"
"sync"
"time"
@@ -52,7 +51,7 @@ func ServeRelay(c *gin.Context, nodeID string) {
client := &relayClient{
nodeID: nodeID,
conn: conn,
send: make(chan Message, 16),
send: make(chan Message, wsChannelBuf),
done: make(chan struct{}),
}
defaultRelayHub.register(client)
@@ -117,37 +116,11 @@ func SendRelayPong(nodeID string) bool {
}
func (c *relayClient) readPump() {
defer c.close()
_ = c.conn.SetReadDeadline(time.Now().Add(90 * time.Second))
c.conn.SetPongHandler(func(string) error {
return c.conn.SetReadDeadline(time.Now().Add(90 * time.Second))
})
for {
_, data, err := c.conn.ReadMessage()
if err != nil {
slog.Debug("relay ws read closed", "node_id", c.nodeID, "error", err)
return
}
var message Message
if err = json.Unmarshal(data, &message); err != nil {
slog.Debug("relay ws invalid message", "node_id", c.nodeID, "error", err)
continue
}
switch message.Type {
case messageTypePing:
_ = SendRelayPong(c.nodeID)
case messageTypePong:
default:
slog.Debug("relay ws unsupported message", "node_id", c.nodeID, "type", message.Type)
}
}
runReadPump(c.nodeID, c.conn, c.close, "relay ws", SendRelayPong, messageTypePong)
}
func (c *relayClient) writePump() {
ticker := time.NewTicker(30 * time.Second)
ticker := time.NewTicker(wsPingInterval)
defer ticker.Stop()
for {
@@ -155,7 +128,7 @@ func (c *relayClient) writePump() {
case <-c.done:
return
case message := <-c.send:
_ = c.conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteDeadline))
if err := c.conn.WriteJSON(message); err != nil {
slog.Debug("relay ws write failed", "node_id", c.nodeID, "error", err)
c.close()