mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 16:46:37 +08:00
fix lint
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 = "节点不存在"
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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 类型"
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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 (
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,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 (
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,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 (
|
||||
|
||||
@@ -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, ".")
|
||||
|
||||
@@ -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 = "配置快照格式无效"
|
||||
)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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 = "端口不能为空"
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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 类型"
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,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 (
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 ©
|
||||
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 ©
|
||||
certCopy := *account
|
||||
certCopy.Authorization = ""
|
||||
return &certCopy
|
||||
}
|
||||
|
||||
func sanitizeAcmeAccountForResponse(account *model.AcmeAccount) *model.AcmeAccount {
|
||||
if account == nil {
|
||||
return nil
|
||||
}
|
||||
copy := *account
|
||||
copy.PrivateKey = ""
|
||||
return ©
|
||||
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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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,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 组不存在"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,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()
|
||||
|
||||
Reference in New Issue
Block a user