mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 23:06:36 +08:00
后端与全仓代码质量清理(golangci 扩展集 · 测试质量 · 并发安全 · 文档同步)
代码质量全量清理,零行为变化:golangci 扩展集 13 类 linter(gosec/modernize/perfsprint/canonicalheader/usestdlibvars/wastedassign/intrange/errorlint/forcetypeassert/recvcheck/exhaustive/unparam)全量修复,测试代码质量(testifylint/thelper/usetesting)25→0,frpc 进程生命周期真 bug(进程组击杀)、全仓 go test -race 6 类数据竞争(含 1 个生产竞争)、SPDX license 头补齐 131 文件、前端测试套件 next-intl 迁移后 44 失败→全绿、过期 swagger 文档重新生成、pnpm-workspace 构建审批。 Experiments: #2-#17, #18, #20, #21, #23 Metric: total_issues 108 → 8 (-92.6%)
This commit is contained in:
@@ -50,9 +50,9 @@ type GetTableDataRequest struct {
|
||||
|
||||
// TableDataResponse 动态数据表响应结构体
|
||||
type TableDataResponse struct {
|
||||
Columns []string `json:"columns"`
|
||||
Total int64 `json:"total"`
|
||||
Results []map[string]interface{} `json:"results"`
|
||||
Columns []string `json:"columns"`
|
||||
Total int64 `json:"total"`
|
||||
Results []map[string]any `json:"results"`
|
||||
}
|
||||
|
||||
// ExecuteSQLRequest 执行自定义 SQL 请求结构体
|
||||
@@ -62,11 +62,11 @@ type ExecuteSQLRequest struct {
|
||||
|
||||
// ExecuteSQLResponse 执行自定义 SQL 响应结构体
|
||||
type ExecuteSQLResponse struct {
|
||||
Type string `json:"type"` // "select" 或 "exec"
|
||||
Columns []string `json:"columns,omitempty"`
|
||||
Results []map[string]interface{} `json:"results,omitempty"`
|
||||
AffectedRows int64 `json:"affected_rows"`
|
||||
ExecutionTimeMs int64 `json:"execution_time_ms"`
|
||||
Type string `json:"type"` // "select" 或 "exec"
|
||||
Columns []string `json:"columns,omitempty"`
|
||||
Results []map[string]any `json:"results,omitempty"`
|
||||
AffectedRows int64 `json:"affected_rows"`
|
||||
ExecutionTimeMs int64 `json:"execution_time_ms"`
|
||||
}
|
||||
|
||||
// formatBytes 格式化字节大小为可读字符串
|
||||
@@ -103,7 +103,7 @@ func formatBytes(bytes uint64) string {
|
||||
}
|
||||
|
||||
// getSQLiteOverview 获取 SQLite 数据库概览信息
|
||||
func getSQLiteOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
func getSQLiteOverview(gormDB *gorm.DB) DBOverviewResponse {
|
||||
name := config.Config.Database.SQLitePath
|
||||
if name == "" {
|
||||
name = "./data/openflare.db"
|
||||
@@ -119,10 +119,7 @@ func getSQLiteOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
|
||||
var sizeStr string
|
||||
if fi, err := os.Stat(name); err == nil {
|
||||
size := fi.Size()
|
||||
if size < 0 {
|
||||
size = 0
|
||||
}
|
||||
size := max(fi.Size(), 0)
|
||||
sizeStr = formatBytes(uint64(size))
|
||||
} else {
|
||||
sizeStr = "0 B"
|
||||
@@ -147,11 +144,11 @@ func getSQLiteOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
Size: sizeStr,
|
||||
TableCount: tableCount,
|
||||
Connections: connCount,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
// getPostgresOverview 获取 PostgreSQL 数据库概览信息
|
||||
func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
func getPostgresOverview(gormDB *gorm.DB) DBOverviewResponse {
|
||||
name := config.Config.Database.Database
|
||||
|
||||
var version string
|
||||
@@ -165,10 +162,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
var sizeStr string
|
||||
var sizeBytes sql.NullInt64
|
||||
if err := gormDB.Raw("SELECT pg_database_size(current_database())").Scan(&sizeBytes).Error; err == nil && sizeBytes.Valid {
|
||||
size := sizeBytes.Int64
|
||||
if size < 0 {
|
||||
size = 0
|
||||
}
|
||||
size := max(sizeBytes.Int64, 0)
|
||||
sizeStr = formatBytes(uint64(size))
|
||||
} else {
|
||||
sizeStr = "0 B"
|
||||
@@ -198,7 +192,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
Size: sizeStr,
|
||||
TableCount: tableCount,
|
||||
Connections: connCount,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
// GetDBOverview 获取数据库运行概览
|
||||
@@ -220,17 +214,11 @@ func GetDBOverview(c *gin.Context) {
|
||||
}
|
||||
|
||||
var overview DBOverviewResponse
|
||||
var err error
|
||||
|
||||
if !config.Config.Database.Enabled {
|
||||
overview, err = getSQLiteOverview(gormDB)
|
||||
overview = getSQLiteOverview(gormDB)
|
||||
} else {
|
||||
overview, err = getPostgresOverview(gormDB)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
overview = getPostgresOverview(gormDB)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(overview))
|
||||
@@ -294,10 +282,7 @@ func GetDBTableData(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
offset := (req.Page - 1) * req.PageSize
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
offset := max((req.Page-1)*req.PageSize, 0)
|
||||
limit := req.PageSize
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
@@ -332,11 +317,11 @@ func GetDBTableData(c *gin.Context) {
|
||||
}
|
||||
|
||||
// scanTableRows 扫描并提取数据表行数据,做截断处理
|
||||
func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]interface{}, error) {
|
||||
results := make([]map[string]interface{}, 0)
|
||||
func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]any, error) {
|
||||
results := make([]map[string]any, 0)
|
||||
for rows.Next() {
|
||||
columns := make([]interface{}, len(cols))
|
||||
columnPointers := make([]interface{}, len(cols))
|
||||
columns := make([]any, len(cols))
|
||||
columnPointers := make([]any, len(cols))
|
||||
for i := range columns {
|
||||
columnPointers[i] = &columns[i]
|
||||
}
|
||||
@@ -345,7 +330,7 @@ func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]interface{}, err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
rowMap := make(map[string]interface{})
|
||||
rowMap := make(map[string]any)
|
||||
for i, colName := range cols {
|
||||
val := columns[i]
|
||||
if b, ok := val.([]byte); ok {
|
||||
@@ -385,10 +370,10 @@ func executeSQLQuery(gormDB *gorm.DB, sqlStr string, startTime time.Time) (Execu
|
||||
return ExecuteSQLResponse{}, err
|
||||
}
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
results := make([]map[string]any, 0)
|
||||
for rows.Next() {
|
||||
columns := make([]interface{}, len(cols))
|
||||
columnPointers := make([]interface{}, len(cols))
|
||||
columns := make([]any, len(cols))
|
||||
columnPointers := make([]any, len(cols))
|
||||
for i := range columns {
|
||||
columnPointers[i] = &columns[i]
|
||||
}
|
||||
@@ -397,7 +382,7 @@ func executeSQLQuery(gormDB *gorm.DB, sqlStr string, startTime time.Time) (Execu
|
||||
return ExecuteSQLResponse{}, err
|
||||
}
|
||||
|
||||
rowMap := make(map[string]interface{})
|
||||
rowMap := make(map[string]any)
|
||||
for i, colName := range cols {
|
||||
val := columns[i]
|
||||
if b, ok := val.([]byte); ok {
|
||||
|
||||
@@ -56,11 +56,11 @@ func GetLogs(c *gin.Context) {
|
||||
limitStr := c.DefaultQuery("limit", "200")
|
||||
|
||||
var cursor, limit int
|
||||
if _, err := parsePositiveInt(cursorStr, &cursor); err != nil {
|
||||
if err := parsePositiveInt(cursorStr, &cursor); err != nil {
|
||||
response.AbortWithError(c, http.StatusBadRequest, admin.InvalidCursorParam)
|
||||
return
|
||||
}
|
||||
if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
|
||||
if err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
|
||||
limit = defaultLimit
|
||||
}
|
||||
if limit > maxLimit {
|
||||
|
||||
@@ -35,8 +35,8 @@ func getUpgrader() *websocket.Upgrader {
|
||||
ctx := r.Context()
|
||||
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
|
||||
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
|
||||
allowedOrigins := strings.Split(sc.Value, ",")
|
||||
for _, allowed := range allowedOrigins {
|
||||
allowedOrigins := strings.SplitSeq(sc.Value, ",")
|
||||
for allowed := range allowedOrigins {
|
||||
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
|
||||
if allowed != "" && strings.EqualFold(allowed, originToCheck) {
|
||||
return true
|
||||
@@ -49,15 +49,15 @@ func getUpgrader() *websocket.Upgrader {
|
||||
}
|
||||
|
||||
// parsePositiveInt 解析非负整数字符串
|
||||
func parsePositiveInt(s string, result *int) (bool, error) {
|
||||
func parsePositiveInt(s string, result *int) error {
|
||||
if s == "" {
|
||||
*result = 0
|
||||
return true, nil
|
||||
return nil
|
||||
}
|
||||
n, err := strconv.Atoi(s)
|
||||
if err != nil || n < 0 {
|
||||
return false, err
|
||||
return err
|
||||
}
|
||||
*result = n
|
||||
return true, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
)
|
||||
|
||||
func setupTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to open sqlite in memory: %v", err)
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
|
||||
package push
|
||||
|
||||
import "slices"
|
||||
|
||||
import "sync"
|
||||
|
||||
const (
|
||||
@@ -66,13 +68,7 @@ func ListDefinitions() []Definition {
|
||||
}
|
||||
// Add any others
|
||||
for t, d := range definitions {
|
||||
found := false
|
||||
for _, o := range order {
|
||||
if o == t {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
found := slices.Contains(order, t)
|
||||
if !found {
|
||||
res = append(res, d)
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"maps"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||
@@ -34,9 +35,7 @@ func (m NotificationMessage) Flatten() map[string]any {
|
||||
keyContent: m.Content,
|
||||
keyLevel: m.Level,
|
||||
}
|
||||
for k, v := range m.Ext {
|
||||
res[k] = v
|
||||
}
|
||||
maps.Copy(res, m.Ext)
|
||||
return res
|
||||
}
|
||||
|
||||
|
||||
@@ -65,6 +65,7 @@ func (m *mockPusher) ValidateConfig(cfg pkgpush.Config) error {
|
||||
}
|
||||
|
||||
func setupPushTest(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) {
|
||||
t.Helper()
|
||||
dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
|
||||
// AutoMigrate push tables in SQLite test environment
|
||||
@@ -319,7 +320,7 @@ func TestPushHandler(t *testing.T) {
|
||||
assert.Equal(t, "Structured Alert", mPusher.sentBody["title"])
|
||||
assert.Equal(t, "Hello World", mPusher.sentBody["content"])
|
||||
assert.Equal(t, "WARNING", mPusher.sentBody["level"])
|
||||
assert.Equal(t, float64(42), mPusher.sentBody["extra_val"]) // unmarshaled json numbers are float64 by default
|
||||
assert.InDelta(t, float64(42), mPusher.sentBody["extra_val"], 1e-9) // unmarshaled json numbers are float64 by default
|
||||
mPusher.mu.Unlock()
|
||||
|
||||
// Verify PushHistory recorded
|
||||
@@ -436,7 +437,7 @@ func TestPushRouters(t *testing.T) {
|
||||
|
||||
dataMap, ok := resp.Data.(map[string]any)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, float64(1), dataMap["total"])
|
||||
assert.InDelta(t, float64(1), dataMap["total"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("test push endpoint", func(t *testing.T) {
|
||||
@@ -614,37 +615,37 @@ func TestPushChannelAPI(t *testing.T) {
|
||||
t.Run("validate push channel model constraints", func(t *testing.T) {
|
||||
// 校验名称合法性
|
||||
c1 := &model.PushChannel{Name: "invalid-name!", URL: "https://hook.com", Other: "{}"}
|
||||
assert.Error(t, c1.Validate())
|
||||
require.Error(t, c1.Validate())
|
||||
|
||||
// 校验 URL 安全前缀 HTTPS
|
||||
c2 := &model.PushChannel{Name: "custom_channel", URL: "http://insecure-hook.com", Other: "{}"}
|
||||
assert.Error(t, c2.Validate())
|
||||
require.Error(t, c2.Validate())
|
||||
|
||||
// 校验 JSON 格式
|
||||
c3 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{invalid-json}"}
|
||||
assert.Error(t, c3.Validate())
|
||||
require.Error(t, c3.Validate())
|
||||
|
||||
// 正确配置
|
||||
c4 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{\"content\":\"$content\"}"}
|
||||
assert.NoError(t, c4.Validate())
|
||||
require.NoError(t, c4.Validate())
|
||||
|
||||
// 飞书渠道校验:非 HTTPS 地址报错
|
||||
c5 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "http://open.feishu.cn", Other: ""}
|
||||
assert.Error(t, c5.Validate())
|
||||
require.Error(t, c5.Validate())
|
||||
|
||||
// 飞书正确配置
|
||||
c6 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "https://open.feishu.cn", Other: ""}
|
||||
assert.NoError(t, c6.Validate())
|
||||
require.NoError(t, c6.Validate())
|
||||
|
||||
// Telegram 渠道校验
|
||||
cTelegramErr := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "https://api.telegram.org", Token: "", Other: ""}
|
||||
assert.Error(t, cTelegramErr.Validate())
|
||||
require.Error(t, cTelegramErr.Validate())
|
||||
|
||||
cTelegramErr2 := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "http://api.telegram.org", Token: "123:abc", Other: ""}
|
||||
assert.Error(t, cTelegramErr2.Validate())
|
||||
require.Error(t, cTelegramErr2.Validate())
|
||||
|
||||
cTelegramOk := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "", Token: "123:abc", Other: "-100123"}
|
||||
assert.NoError(t, cTelegramOk.Validate())
|
||||
require.NoError(t, cTelegramOk.Validate())
|
||||
assert.Equal(t, "https://api.telegram.org", cTelegramOk.URL)
|
||||
|
||||
// 邮件配置校验:允许空配置以复用系统全局设置
|
||||
@@ -733,7 +734,7 @@ func TestPushChannelAPI(t *testing.T) {
|
||||
dbConn.First(&updated, createdID)
|
||||
assert.Equal(t, "Updated remark", updated.Description)
|
||||
assert.Equal(t, "new_chan_token", updated.Token)
|
||||
assert.Equal(t, `{"text": "$content"}`, updated.Other)
|
||||
assert.JSONEq(t, `{"text": "$content"}`, updated.Other)
|
||||
})
|
||||
|
||||
t.Run("admin test channel endpoint", func(t *testing.T) {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||
@@ -309,7 +310,7 @@ func exportSQLite(c *gin.Context) {
|
||||
|
||||
c.Header("Content-Disposition", `attachment; filename="openflare.db"`)
|
||||
c.Header("Content-Type", "application/octet-stream")
|
||||
c.Header("Content-Length", fmt.Sprintf("%d", fi.Size()))
|
||||
c.Header("Content-Length", strconv.FormatInt(fi.Size(), 10))
|
||||
c.Status(http.StatusOK)
|
||||
http.ServeContent(c.Writer, c.Request, "openflare.db", fi.ModTime(), f)
|
||||
}
|
||||
@@ -328,7 +329,7 @@ func exportPostgres(c *gin.Context) {
|
||||
args := []string{
|
||||
"--no-password",
|
||||
"-h", dbCfg.Host,
|
||||
"-p", fmt.Sprintf("%d", dbCfg.Port),
|
||||
"-p", strconv.Itoa(dbCfg.Port),
|
||||
"-U", dbCfg.Username,
|
||||
dbCfg.Database,
|
||||
}
|
||||
|
||||
@@ -46,6 +46,7 @@ func registerInternalOnlyTaskMeta() {
|
||||
}
|
||||
|
||||
func setupTaskTestEnvironment(t *testing.T) func() {
|
||||
t.Helper()
|
||||
_, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
bootstrap.RegisterTasks()
|
||||
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
|
||||
@@ -490,7 +491,7 @@ func TestListTaskExecutions(t *testing.T) {
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
assert.InDelta(t, float64(3), data["total"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("filter by status", func(t *testing.T) {
|
||||
@@ -507,7 +508,7 @@ func TestListTaskExecutions(t *testing.T) {
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(1), data["total"])
|
||||
assert.InDelta(t, float64(1), data["total"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("filter by task_type (asynq task name)", func(t *testing.T) {
|
||||
@@ -524,7 +525,7 @@ func TestListTaskExecutions(t *testing.T) {
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
assert.InDelta(t, float64(3), data["total"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("filter by task_type (management task type)", func(t *testing.T) {
|
||||
@@ -541,7 +542,7 @@ func TestListTaskExecutions(t *testing.T) {
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
assert.InDelta(t, float64(3), data["total"], 1e-9)
|
||||
})
|
||||
|
||||
t.Run("pagination", func(t *testing.T) {
|
||||
@@ -558,7 +559,7 @@ func TestListTaskExecutions(t *testing.T) {
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
assert.InDelta(t, float64(3), data["total"], 1e-9)
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -163,13 +164,7 @@ func selectLatestRelease(repository string, releases []githubRelease) (githubRel
|
||||
}
|
||||
expectedNames := expectedAssetNames(repository, release.TagName)
|
||||
for _, asset := range release.Assets {
|
||||
matched := false
|
||||
for _, name := range expectedNames {
|
||||
if asset.Name == name {
|
||||
matched = true
|
||||
break
|
||||
}
|
||||
}
|
||||
matched := slices.Contains(expectedNames, asset.Name)
|
||||
if !matched || asset.BrowserDownloadURL == "" || asset.State != "uploaded" {
|
||||
continue
|
||||
}
|
||||
@@ -199,7 +194,7 @@ func (m *manager) fetchRelease(ctx context.Context, repository string) (githubRe
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
req.Header.Set("User-Agent", "OpenFlare-Updater")
|
||||
req.Header.Set("X-GitHub-Api-Version", "2022-11-28")
|
||||
req.Header.Set("X-Github-Api-Version", "2022-11-28")
|
||||
|
||||
resp, err := m.client.Do(req)
|
||||
if err != nil {
|
||||
@@ -348,10 +343,8 @@ func getCandidateBinaryNames(executable string, repository string) []string {
|
||||
if runtime.GOOS == windowsOS && !strings.HasSuffix(strings.ToLower(name), ".exe") {
|
||||
name += ".exe"
|
||||
}
|
||||
for _, existing := range names {
|
||||
if existing == name {
|
||||
return
|
||||
}
|
||||
if slices.Contains(names, name) {
|
||||
return
|
||||
}
|
||||
names = append(names, name)
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ package user
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"slices"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
@@ -93,17 +94,13 @@ func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbidde
|
||||
return true
|
||||
}
|
||||
msg := err.Error()
|
||||
for _, m := range badRequestMsgs {
|
||||
if msg == m {
|
||||
response.AbortBadRequest(c, msg)
|
||||
return true
|
||||
}
|
||||
if slices.Contains(badRequestMsgs, msg) {
|
||||
response.AbortBadRequest(c, msg)
|
||||
return true
|
||||
}
|
||||
for _, m := range forbiddenMsgs {
|
||||
if msg == m {
|
||||
response.AbortForbidden(c, msg)
|
||||
return true
|
||||
}
|
||||
if slices.Contains(forbiddenMsgs, msg) {
|
||||
response.AbortForbidden(c, msg)
|
||||
return true
|
||||
}
|
||||
logger.ErrorF(c.Request.Context(), "Admin user error: %v", err)
|
||||
response.AbortInternal(c, "内部服务器错误")
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package agent implements the local OpenFlare agent runtime loop.
|
||||
package agent
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package agent
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package config loads and persists agent daemon configuration.
|
||||
package config
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import edgeconfig "github.com/Rain-kl/Wavelet/internal/apps/edge/config"
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
// Version is the current agent version string, overridden at build time.
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package geoipdata holds shared GeoIP database filename constants.
|
||||
//
|
||||
// MaxMind MMDB files are NOT embedded into the agent binary. Docker images
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package geoipupdate schedules local MaxMind GeoIP database updates for the agent.
|
||||
package geoipupdate
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package geoipupdate
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package heartbeat implements the periodic heartbeat cycle executed by the agent,
|
||||
// including payload preparation, config sync, WAF IP group application, and observability buffering.
|
||||
package heartbeat
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logging configures structured logging for the agent process.
|
||||
package logging
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package nginx manages OpenResty configuration, runtime, and supporting assets.
|
||||
package nginx
|
||||
|
||||
@@ -134,7 +137,7 @@ func (e *PathExecutor) Reload(ctx context.Context) error {
|
||||
slog.Warn("openresty reload reported runtime is not running, starting binary", "path", e.Path)
|
||||
startOutput, startErr := e.Runner.Run(ctx, e.Path, "-c", e.ConfigPath)
|
||||
if startErr != nil {
|
||||
return fmt.Errorf("openresty reload failed: %w: %s; start failed: %v: %s", err, string(output), startErr, string(startOutput))
|
||||
return fmt.Errorf("openresty reload failed: %w: %s; start failed: %w: %s", err, string(output), startErr, string(startOutput))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -355,14 +358,14 @@ func (m *Manager) activateConfig(ctx context.Context) error {
|
||||
func (m *Manager) rollbackAfterFailedApply(ctx context.Context, backup *backupState, applyErr error) ApplyOutcome {
|
||||
slog.Warn("openresty apply failed, restoring previous config", "error", applyErr)
|
||||
if err := m.restore(backup); err != nil {
|
||||
return fatalApplyOutcome(fmt.Errorf("restore openresty backup failed after apply error %v: %w", applyErr, err))
|
||||
return fatalApplyOutcome(fmt.Errorf("restore openresty backup failed after apply error %w: %w", applyErr, err))
|
||||
}
|
||||
if err := m.activateConfig(ctx); err != nil {
|
||||
if backup != nil && backup.MainExisted {
|
||||
return fatalApplyOutcome(fmt.Errorf("apply failed: %v; rollback recovery failed: %w", applyErr, err))
|
||||
return fatalApplyOutcome(fmt.Errorf("apply failed: %w; rollback recovery failed: %w", applyErr, err))
|
||||
}
|
||||
if fallbackErr := m.EnsureSafeFallbackRuntime(ctx, fmt.Sprintf("apply failed: %v; rollback recovery failed: %v", applyErr, err)); fallbackErr != nil {
|
||||
return fatalApplyOutcome(fmt.Errorf("apply failed: %v; rollback recovery failed: %w; fallback recovery failed: %v", applyErr, err, fallbackErr))
|
||||
return fatalApplyOutcome(fmt.Errorf("apply failed: %w; rollback recovery failed: %w; fallback recovery failed: %w", applyErr, err, fallbackErr))
|
||||
}
|
||||
message := fmt.Sprintf("apply failed, but fallback runtime started: %v; rollback recovery failed: %v", applyErr, err)
|
||||
slog.Warn("openresty apply recovered with safe default fallback", "message", message)
|
||||
@@ -516,7 +519,7 @@ func (m *Manager) CurrentChecksum() (string, error) {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, listen, openrestyrender.ObservabilityListenPlaceholder)
|
||||
}
|
||||
if m.OpenrestyObservabilityPort > 0 {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, fmt.Sprintf("%d", m.OpenrestyObservabilityPort), openrestyrender.ObservabilityPortPlaceholder)
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, strconv.Itoa(m.OpenrestyObservabilityPort), openrestyrender.ObservabilityPortPlaceholder)
|
||||
}
|
||||
if resolverDirective := strings.TrimSpace(m.OpenrestyResolverDirective); resolverDirective != "" {
|
||||
normalizedMain = strings.ReplaceAll(normalizedMain, resolverDirective, ResolverDirectivePlaceholder)
|
||||
@@ -1463,7 +1466,7 @@ func (m *Manager) renderMainConfig(content string) string {
|
||||
rendered = strings.ReplaceAll(rendered, openrestyrender.ObservabilityListenPlaceholder, listen)
|
||||
}
|
||||
if m.OpenrestyObservabilityPort > 0 {
|
||||
rendered = strings.ReplaceAll(rendered, openrestyrender.ObservabilityPortPlaceholder, fmt.Sprintf("%d", m.OpenrestyObservabilityPort))
|
||||
rendered = strings.ReplaceAll(rendered, openrestyrender.ObservabilityPortPlaceholder, strconv.Itoa(m.OpenrestyObservabilityPort))
|
||||
}
|
||||
if resolverDirective := strings.TrimSpace(m.OpenrestyResolverDirective); resolverDirective != "" {
|
||||
rendered = strings.ReplaceAll(rendered, ResolverDirectivePlaceholder, resolverDirective)
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
// DefaultMimeTypes is the embedded nginx mime.types map used by generated configs.
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package nginx
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package observability provides system and service level observability data collection for the agent.
|
||||
package observability
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
@@ -62,7 +65,7 @@ func CollectEdgeHealth(ctx context.Context, cfg *config.Config) *EdgeHealthSnaps
|
||||
}
|
||||
|
||||
func fetchLocalJSON(ctx context.Context, client *http.Client, url string, target any) error {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", url, nil)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package protocol defines type aliases and constants for the agent protocol.
|
||||
package protocol
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package runtimeuser defines the shared OS account used by the agent process
|
||||
// and OpenResty worker processes so file ownership stays aligned.
|
||||
package runtimeuser
|
||||
@@ -120,7 +123,7 @@ func ensureWorldTraversablePath(targetDir string) error {
|
||||
if current == "" || current == "." {
|
||||
return nil
|
||||
}
|
||||
for depth := 0; depth < maxDepth; depth++ {
|
||||
for range maxDepth {
|
||||
if err := os.Chmod(current, DefaultDirPerm); err != nil { //nolint:gosec // parent dirs must be traversable by the runtime user
|
||||
if os.IsNotExist(err) || os.IsPermission(err) {
|
||||
break
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package runtimeuser
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
//go:build unix
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package runtimeuser
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package state persists agent runtime state and observability snapshots.
|
||||
package state
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package state
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package state
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package state
|
||||
|
||||
import (
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
|
||||
@@ -198,7 +199,7 @@ func (s *Service) ensurePagesProject(ctx context.Context, snapshot *state.Snapsh
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
for attempt := 0; attempt < pagesLatestPullAttempts; attempt++ {
|
||||
for attempt := range pagesLatestPullAttempts {
|
||||
latest, err := s.client.GetPagesProjectLatestHash(ctx, projectID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("fetch Pages project %d latest hash: %w", projectID, err)
|
||||
@@ -359,10 +360,7 @@ func validatePagesPackageMetadata(
|
||||
if extractedBytes == 0 {
|
||||
extractedBytes = 1
|
||||
}
|
||||
maxFileBytes := extractedBytes
|
||||
if maxFileBytes > agentPagesMaxFileBytes {
|
||||
maxFileBytes = agentPagesMaxFileBytes
|
||||
}
|
||||
maxFileBytes := min(extractedBytes, agentPagesMaxFileBytes)
|
||||
|
||||
return pagesPackageLimits{
|
||||
PackageBytes: metadata.PackageSize,
|
||||
@@ -395,7 +393,7 @@ func (s *Service) downloadPagesProjectPackage(
|
||||
metadata *protocol.PagesProjectLatestHashResponse,
|
||||
maxBytes int64,
|
||||
) (packagePath string, hash string, err error) {
|
||||
releasesRoot := filepath.Join(s.pagesDir, "projects", fmt.Sprintf("%d", projectID), "releases")
|
||||
releasesRoot := filepath.Join(s.pagesDir, "projects", strconv.FormatUint(uint64(projectID), 10), "releases")
|
||||
if err := os.MkdirAll(releasesRoot, pagesDirPerm); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
@@ -444,7 +442,7 @@ func cleanupPagesProjectStaleReleases(baseDir string, projectID uint, keepHash s
|
||||
if projectID == 0 || keepHash == "" {
|
||||
return nil
|
||||
}
|
||||
releasesRoot := filepath.Join(baseDir, "projects", fmt.Sprintf("%d", projectID), "releases")
|
||||
releasesRoot := filepath.Join(baseDir, "projects", strconv.FormatUint(uint64(projectID), 10), "releases")
|
||||
entries, err := os.ReadDir(releasesRoot) //nolint:gosec // managed PagesDir
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
@@ -985,9 +983,9 @@ func writePagesMarker(dir string, project pagesProjectRef) error {
|
||||
}
|
||||
|
||||
func pagesProjectCurrentDir(baseDir string, projectID uint) string {
|
||||
return filepath.Join(baseDir, "projects", fmt.Sprintf("%d", projectID), "current")
|
||||
return filepath.Join(baseDir, "projects", strconv.FormatUint(uint64(projectID), 10), "current")
|
||||
}
|
||||
|
||||
func pagesProjectReleaseDir(baseDir string, projectID uint, checksum string) string {
|
||||
return filepath.Join(baseDir, "projects", fmt.Sprintf("%d", projectID), "releases", checksum)
|
||||
return filepath.Join(baseDir, "projects", strconv.FormatUint(uint64(projectID), 10), "releases", checksum)
|
||||
}
|
||||
|
||||
@@ -238,10 +238,7 @@ func TestPromoteSameHashReleaseFailureRestoresPreviousCurrent(t *testing.T) {
|
||||
if err := switchPagesProjectCurrentDir(pagesDir, projectID, releaseDir); err != nil {
|
||||
t.Fatalf("seed same-hash current error = %v", err)
|
||||
}
|
||||
stagingDir, err := os.MkdirTemp(filepath.Dir(releaseDir), ".same-hash-*.tmp")
|
||||
if err != nil {
|
||||
t.Fatalf("create same-hash staging error = %v", err)
|
||||
}
|
||||
stagingDir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(stagingDir, "index.html"), []byte("new"), pagesFilePerm); err != nil {
|
||||
t.Fatalf("write repaired same-hash release error = %v", err)
|
||||
}
|
||||
@@ -249,7 +246,7 @@ func TestPromoteSameHashReleaseFailureRestoresPreviousCurrent(t *testing.T) {
|
||||
t.Fatalf("write repaired same-hash marker error = %v", err)
|
||||
}
|
||||
copyErr := errors.New("injected same-hash copy failure")
|
||||
err = promotePagesReleaseWithCopy(
|
||||
err := promotePagesReleaseWithCopy(
|
||||
stagingDir,
|
||||
releaseDir,
|
||||
project,
|
||||
@@ -298,10 +295,7 @@ func TestPromotePagesReleaseRepairsDanglingCurrent(t *testing.T) {
|
||||
}
|
||||
|
||||
requireTestMkdirAll(t, filepath.Dir(releaseDir))
|
||||
stagingDir, err := os.MkdirTemp(filepath.Dir(releaseDir), ".dangling-*.tmp")
|
||||
if err != nil {
|
||||
t.Fatalf("create dangling repair staging error = %v", err)
|
||||
}
|
||||
stagingDir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(stagingDir, "index.html"), []byte("repaired"), pagesFilePerm); err != nil {
|
||||
t.Fatalf("write dangling repair staging error = %v", err)
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"sort"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -363,7 +363,7 @@ func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) ([]uint, error
|
||||
for id := range seen {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
slices.Sort(ids)
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
@@ -426,7 +426,7 @@ func (s *Service) ensureRuntimeForCurrentConfig(ctx context.Context, mode string
|
||||
snapshot.OpenrestyMessage = "safe default fallback runtime started"
|
||||
return nil
|
||||
}
|
||||
err = fmt.Errorf("%v; fallback recovery failed: %w", err, fallbackErr)
|
||||
err = fmt.Errorf("%w; fallback recovery failed: %w", err, fallbackErr)
|
||||
}
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusUnhealthy
|
||||
snapshot.OpenrestyMessage = err.Error()
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package sync
|
||||
|
||||
import (
|
||||
@@ -220,6 +223,9 @@ func updateSnapshotFromApplyOutcome(mode string, snapshot *state.Snapshot, confi
|
||||
snapshot.OpenrestyStatus = protocol.OpenrestyStatusHealthy
|
||||
snapshot.OpenrestyMessage = result.message
|
||||
result.reportResult = ApplyResultWarning
|
||||
case nginx.ApplyStatusFatal:
|
||||
// 致命错误与普通失败同走失败路径:标记阻塞并上报 Unhealthy。
|
||||
fallthrough
|
||||
default:
|
||||
if result.message == "" {
|
||||
result.message = "openresty apply failed"
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package updater provides agent self-update integration with the edge updater.
|
||||
package updater
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package wsclient provides the agent-side WebSocket client for connecting to the OpenFlare server.
|
||||
package wsclient
|
||||
|
||||
|
||||
@@ -78,10 +78,7 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
|
||||
if nonceTTL < time.Second {
|
||||
nonceTTL = time.Second
|
||||
}
|
||||
nonceTTL := max(time.Duration(payload.Expires-now)*time.Millisecond, time.Second)
|
||||
|
||||
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
|
||||
if err != nil {
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package config provides shared configuration types for edge applications.
|
||||
package config
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package heartbeat handles periodic heartbeat and update checks.
|
||||
package heartbeat
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package httpclient provides an authenticated HTTP client for edge services.
|
||||
package httpclient
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logging configures structured logging for edge applications.
|
||||
package logging
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package nodeip detects the preferred public IP address for edge nodes.
|
||||
package nodeip
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package observability provides helpers that read Linux /proc and /sys metrics for system monitoring.
|
||||
package observability
|
||||
|
||||
@@ -96,10 +99,7 @@ func ReadMemInfo() (int64, int64) {
|
||||
if total == 0 {
|
||||
return 0, 0
|
||||
}
|
||||
used := total - (memAvailableKB * 1024)
|
||||
if used < 0 {
|
||||
used = 0
|
||||
}
|
||||
used := max(total-(memAvailableKB*1024), 0)
|
||||
return total, used
|
||||
}
|
||||
|
||||
@@ -138,8 +138,8 @@ func ReadLinuxCPUStat() (uint64, uint64) {
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
lines := strings.Split(string(content), "\n")
|
||||
for _, line := range lines {
|
||||
lines := strings.SplitSeq(string(content), "\n")
|
||||
for line := range lines {
|
||||
if !strings.HasPrefix(line, "cpu ") {
|
||||
continue
|
||||
}
|
||||
@@ -258,23 +258,23 @@ func StatFilesystem(path string) (int64, int64) {
|
||||
if err := syscall.Statfs(absPath, &stat); err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
total := multiplyUint64ToInt64(stat.Blocks, uint64(stat.Bsize))
|
||||
free := multiplyUint64ToInt64(stat.Bavail, uint64(stat.Bsize))
|
||||
used := total - free
|
||||
if used < 0 {
|
||||
used = 0
|
||||
}
|
||||
total := multiplyUint64Int64(stat.Blocks, stat.Bsize)
|
||||
free := multiplyUint64Int64(stat.Bavail, stat.Bsize)
|
||||
used := max(total-free, 0)
|
||||
return total, used
|
||||
}
|
||||
|
||||
func multiplyUint64ToInt64(a uint64, b uint64) int64 {
|
||||
if a == 0 || b == 0 {
|
||||
// multiplyUint64Int64 multiplies a uint64 by a positive int64, saturating at
|
||||
// math.MaxInt64 to avoid int64 overflow.
|
||||
func multiplyUint64Int64(a uint64, b int64) int64 {
|
||||
if a == 0 || b <= 0 {
|
||||
return 0
|
||||
}
|
||||
if a > math.MaxInt64/b {
|
||||
v := a * uint64(b)
|
||||
if v > math.MaxInt64 {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return int64(a * b) //nolint:gosec // product is bounded to math.MaxInt64 above
|
||||
return int64(v)
|
||||
}
|
||||
|
||||
// ReadFirstLine reads and returns the trimmed first line of a file.
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package observability
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package runner provides shared WebSocket reconnect helpers for edge daemons.
|
||||
package runner
|
||||
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
//go:build !windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package updater provides capabilities to check for, download, and apply updates.
|
||||
package updater
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
@@ -19,7 +23,7 @@ func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
renameErr := err
|
||||
if err := os.Remove(tmpPath); err != nil && !os.IsNotExist(err) {
|
||||
slog.Error("remove tmp binary failed", "path", tmpPath, "error", err)
|
||||
return fmt.Errorf("backup current binary: %w; remove tmp binary: %v", renameErr, err)
|
||||
return fmt.Errorf("backup current binary: %w; remove tmp binary: %w", renameErr, err)
|
||||
}
|
||||
return fmt.Errorf("backup current binary: %w", renameErr)
|
||||
}
|
||||
@@ -27,7 +31,7 @@ func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
replaceErr := err
|
||||
if err := os.Rename(backupPath, execPath); err != nil {
|
||||
slog.Error("restore backup binary failed", "path", backupPath, "error", err)
|
||||
return fmt.Errorf("replace binary: %w; restore backup binary: %v", replaceErr, err)
|
||||
return fmt.Errorf("replace binary: %w; restore backup binary: %w", replaceErr, err)
|
||||
}
|
||||
return fmt.Errorf("replace binary: %w", replaceErr)
|
||||
}
|
||||
@@ -37,7 +41,7 @@ func replaceAndRestart(execPath string, tmpPath string) error {
|
||||
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil { //nolint:gosec // execPath is the validated edge updater binary path
|
||||
return fmt.Errorf("exec restart: %w", err)
|
||||
}
|
||||
return fmt.Errorf("unreachable after exec")
|
||||
return errors.New("unreachable after exec")
|
||||
}
|
||||
|
||||
func removeBackupBinary(path string) error {
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
//go:build !windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
//go:build windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package updater provides capabilities to check for, download, and apply updates.
|
||||
package updater
|
||||
|
||||
@@ -6,6 +9,7 @@ import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
@@ -285,7 +289,7 @@ func (s *Service) downloadChecksum(ctx context.Context, url string, assetName st
|
||||
|
||||
func parseSHA256Checksum(content string, assetName string) (string, error) {
|
||||
assetName = strings.TrimSpace(assetName)
|
||||
for _, line := range strings.Split(content, "\n") {
|
||||
for line := range strings.SplitSeq(content, "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "#") {
|
||||
continue
|
||||
@@ -295,7 +299,7 @@ func parseSHA256Checksum(content string, assetName string) (string, error) {
|
||||
}
|
||||
}
|
||||
if assetName == "" {
|
||||
return "", fmt.Errorf("checksum asset does not contain a valid sha256 digest")
|
||||
return "", errors.New("checksum asset does not contain a valid sha256 digest")
|
||||
}
|
||||
return "", fmt.Errorf("checksum asset does not contain a sha256 digest for %q", assetName)
|
||||
}
|
||||
@@ -340,7 +344,7 @@ func isSHA256Hex(value string) bool {
|
||||
func (s *Service) downloadAndRestart(ctx context.Context, url string, expectedChecksum string, targetPath string) error {
|
||||
expectedChecksum = strings.ToLower(strings.TrimSpace(expectedChecksum))
|
||||
if !isSHA256Hex(expectedChecksum) {
|
||||
return fmt.Errorf("invalid expected sha256 checksum")
|
||||
return errors.New("invalid expected sha256 checksum")
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package updater
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package wsclient provides WebSocket client abstractions for edge node communication.
|
||||
package wsclient
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package config loads and persists flared daemon configuration.
|
||||
package config
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
// Version is the flared daemon build version string.
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package flared implements the tunnel client daemon runtime loop.
|
||||
package flared
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package frpc manages frpc child processes for tunnel relay connections.
|
||||
package frpc
|
||||
|
||||
@@ -195,6 +198,17 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
var stderrBuf bytes.Buffer
|
||||
cmd.Stderr = &stderrBuf
|
||||
|
||||
// frpc 及其中间子进程必须整体随上下文终止:CommandContext 默认只杀
|
||||
// 直接子进程,孤儿孙进程会继续持有 stderr 管道导致 cmd.Wait 阻塞到其
|
||||
// 自然退出。这里为 frpc 单独建进程组并整组 SIGKILL。
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
cmd.Cancel = func() error {
|
||||
if cmd.Process == nil {
|
||||
return nil
|
||||
}
|
||||
return syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
proc.Cmd = cmd
|
||||
proc.Status = "running"
|
||||
@@ -203,7 +217,7 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
startedAt := time.Now()
|
||||
err := cmd.Start()
|
||||
if err == nil {
|
||||
_ = os.WriteFile(pidPath, []byte(fmt.Sprintf("%d", cmd.Process.Pid)), frpcConfigFilePerm)
|
||||
_ = os.WriteFile(pidPath, fmt.Appendf(nil, "%d", cmd.Process.Pid), frpcConfigFilePerm)
|
||||
err = cmd.Wait()
|
||||
}
|
||||
_ = os.Remove(pidPath)
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package frpc
|
||||
|
||||
import (
|
||||
@@ -7,6 +10,7 @@ import (
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -16,6 +20,7 @@ import (
|
||||
|
||||
// Helper to write control file for the dummy script
|
||||
func writeControl(t *testing.T, dir string, exitCode int, delaySeconds int) {
|
||||
t.Helper()
|
||||
controlPath := filepath.Join(dir, "control.txt")
|
||||
content := fmt.Sprintf("%d %d\n", exitCode, delaySeconds)
|
||||
err := os.WriteFile(controlPath, []byte(content), 0644)
|
||||
@@ -26,6 +31,7 @@ func writeControl(t *testing.T, dir string, exitCode int, delaySeconds int) {
|
||||
|
||||
// Setup a dummy executable script that reads control.txt to decide exit code and sleep duration
|
||||
func setupDummyScript(t *testing.T) (string, string) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
scriptPath := filepath.Join(dir, "dummy_frpc")
|
||||
|
||||
@@ -53,12 +59,17 @@ exit "${EXIT_CODE:-0}"
|
||||
|
||||
// Helper to poll for status to eliminate timing flakiness in tests
|
||||
func assertStatusEventually(t *testing.T, m *Manager, relayID string, expectedStatus string, timeout time.Duration) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
m.mu.RLock()
|
||||
proc, ok := m.processes[relayID]
|
||||
var status string
|
||||
if ok {
|
||||
status = proc.Status
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
if ok && proc.Status == expectedStatus {
|
||||
if ok && status == expectedStatus {
|
||||
return
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
@@ -77,6 +88,7 @@ func assertStatusEventually(t *testing.T, m *Manager, relayID string, expectedSt
|
||||
t.Fatalf("expected status eventually %s, got %s (err: %s)", expectedStatus, got, errStr)
|
||||
}
|
||||
|
||||
// assertCommandExitedEventually 等待测试自建进程退出(本测试持有其 Wait 权)。
|
||||
func assertCommandExitedEventually(t *testing.T, cmd *exec.Cmd, timeout time.Duration) {
|
||||
t.Helper()
|
||||
|
||||
@@ -92,6 +104,22 @@ func assertCommandExitedEventually(t *testing.T, cmd *exec.Cmd, timeout time.Dur
|
||||
}
|
||||
}
|
||||
|
||||
// assertManagedCommandExitedEventually 探测受管进程是否已退出。不能对其调用
|
||||
// Wait —— Wait 由 Manager 拥有,测试并发 Wait 会与 os/exec 的 ctxResult
|
||||
// 通道竞争而永久挂起;Signal(0) 在进程被 Manager 收割后即报错。
|
||||
func assertManagedCommandExitedEventually(t *testing.T, cmd *exec.Cmd, timeout time.Duration) {
|
||||
t.Helper()
|
||||
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
if err := cmd.Process.Signal(syscall.Signal(0)); err != nil {
|
||||
return
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("expected managed process pid=%d to exit within %s", cmd.Process.Pid, timeout)
|
||||
}
|
||||
|
||||
func TestStartProcessSuccess(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 5) // exit code 0, sleep 5s
|
||||
@@ -377,7 +405,7 @@ func TestStopCancelsRunningProcesses(t *testing.T) {
|
||||
m.mu.RUnlock()
|
||||
|
||||
m.Stop()
|
||||
assertCommandExitedEventually(t, cmd, 2*time.Second)
|
||||
assertManagedCommandExitedEventually(t, cmd, 2*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package heartbeat runs the periodic flared heartbeat loop against the control plane.
|
||||
package heartbeat
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package httpclient provides the HTTP client used by the flared agent to communicate with the Wavelet server.
|
||||
package httpclient
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package sync periodically fetches and applies the active tunnel configuration.
|
||||
package sync
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package updater provides update service capabilities for flared.
|
||||
package updater
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package wsclient provides a WebSocket client for flared control-plane communication.
|
||||
package wsclient
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ package oauth
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
@@ -109,10 +110,5 @@ func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL
|
||||
}
|
||||
|
||||
func containsScope(scopes []string, scope string) bool {
|
||||
for _, item := range scopes {
|
||||
if item == scope {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
return slices.Contains(scopes, scope)
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"sync"
|
||||
@@ -31,10 +32,12 @@ var (
|
||||
tokenListenerOnce sync.Once
|
||||
tokenListenerCtx context.Context
|
||||
tokenListenerCancel context.CancelFunc
|
||||
tokenListenerDone chan struct{}
|
||||
|
||||
userListenerOnce sync.Once
|
||||
userListenerCtx context.Context
|
||||
userListenerCancel context.CancelFunc
|
||||
userListenerDone chan struct{}
|
||||
)
|
||||
|
||||
func tokenCacheKey(tokenHash string) string {
|
||||
@@ -54,15 +57,19 @@ func ensureTokenCacheListener() {
|
||||
|
||||
func startTokenCacheInvalidationListener() {
|
||||
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
|
||||
tokenListenerDone = make(chan struct{})
|
||||
|
||||
go func() {
|
||||
pubsub := db.Redis.Subscribe(tokenListenerCtx, oauthTokenInvalidationChannel)
|
||||
listenerCtx := tokenListenerCtx
|
||||
defer close(tokenListenerDone)
|
||||
|
||||
pubsub := db.Redis.Subscribe(listenerCtx, oauthTokenInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
<-tokenListenerCtx.Done()
|
||||
<-listenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
@@ -93,15 +100,19 @@ func ensureUserCacheListener() {
|
||||
|
||||
func startUserCacheInvalidationListener() {
|
||||
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
|
||||
userListenerDone = make(chan struct{})
|
||||
|
||||
go func() {
|
||||
pubsub := db.Redis.Subscribe(userListenerCtx, oauthUserInvalidationChannel)
|
||||
listenerCtx := userListenerCtx
|
||||
defer close(userListenerDone)
|
||||
|
||||
pubsub := db.Redis.Subscribe(listenerCtx, oauthUserInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
<-userListenerCtx.Done()
|
||||
<-listenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
@@ -140,7 +151,7 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken,
|
||||
return &token, nil
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("cache miss")
|
||||
return nil, errors.New("cache miss")
|
||||
}
|
||||
|
||||
// SetCachedToken 设置 AccessToken 缓存
|
||||
@@ -183,7 +194,7 @@ func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
|
||||
return &u, nil
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("cache miss")
|
||||
return nil, errors.New("cache miss")
|
||||
}
|
||||
|
||||
// SetCachedUser 设置 User 缓存
|
||||
@@ -213,13 +224,21 @@ func InvalidateCachedUser(ctx context.Context, userID uint64) {
|
||||
func StopOauthCacheListener() {
|
||||
if tokenListenerCancel != nil {
|
||||
tokenListenerCancel()
|
||||
if tokenListenerDone != nil {
|
||||
<-tokenListenerDone
|
||||
}
|
||||
tokenListenerCancel = nil
|
||||
tokenListenerDone = nil
|
||||
}
|
||||
tokenListenerOnce = sync.Once{}
|
||||
|
||||
if userListenerCancel != nil {
|
||||
userListenerCancel()
|
||||
if userListenerDone != nil {
|
||||
<-userListenerDone
|
||||
}
|
||||
userListenerCancel = nil
|
||||
userListenerDone = nil
|
||||
}
|
||||
userListenerOnce = sync.Once{}
|
||||
}
|
||||
|
||||
@@ -336,6 +336,7 @@ func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, user
|
||||
}
|
||||
|
||||
func setupTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
repository.ResetAuthSourceRAMCacheForTest()
|
||||
|
||||
@@ -380,6 +381,11 @@ func resetOIDCProviderCacheForTest() {
|
||||
}
|
||||
|
||||
func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *http.Client) *gin.Engine {
|
||||
// 停止各层 Pub/Sub 监听 goroutine(可能在先前测试的 API 调用中随 sync.Once
|
||||
// 启动),否则它们在 db.Redis 被替换时仍读取旧值,产生数据竞争。
|
||||
StopOauthCacheListener()
|
||||
repository.StopAuthSourceCacheListener()
|
||||
repository.StopSystemConfigCacheListener()
|
||||
resetOIDCProviderCacheForTest()
|
||||
|
||||
r := testhelper.NewTestGinEngine(gin.Recovery())
|
||||
|
||||
@@ -179,18 +179,9 @@ func buildNodeAccessLogRecords(nodeID string, direct []NodeAccessLog, buffered [
|
||||
records := make([]*model.OpenFlareAccessLog, 0, total)
|
||||
appendLogs := func(logs []NodeAccessLog) {
|
||||
for _, item := range logs {
|
||||
bytesSent := item.BytesSent
|
||||
if bytesSent < 0 {
|
||||
bytesSent = 0
|
||||
}
|
||||
requestLength := item.RequestLength
|
||||
if requestLength < 0 {
|
||||
requestLength = 0
|
||||
}
|
||||
requestTimeMs := item.RequestTimeMs
|
||||
if requestTimeMs < 0 {
|
||||
requestTimeMs = 0
|
||||
}
|
||||
bytesSent := max(item.BytesSent, 0)
|
||||
requestLength := max(item.RequestLength, 0)
|
||||
requestTimeMs := max(item.RequestTimeMs, 0)
|
||||
record := &model.OpenFlareAccessLog{
|
||||
NodeID: nodeID,
|
||||
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -45,7 +46,7 @@ func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[s
|
||||
}
|
||||
changed := make([]WAFIPGroup, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum {
|
||||
if strings.TrimSpace(checksums[strconv.FormatUint(uint64(group.ID), 10)]) == group.Checksum {
|
||||
continue
|
||||
}
|
||||
changed = append(changed, group)
|
||||
@@ -97,7 +98,7 @@ func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error
|
||||
if len(ids) == 0 {
|
||||
return []WAFIPGroup{}, nil
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
slices.Sort(ids)
|
||||
groups, err := repository.ListOpenFlareWAFIPGroupsByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -204,7 +205,7 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
slices.Sort(ids)
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@ import (
|
||||
)
|
||||
|
||||
func setupApplyLogTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package openflare implements openflare configuration, service orchestration, and background tasks.
|
||||
package openflare
|
||||
|
||||
|
||||
@@ -192,7 +192,7 @@ func (client *HTTPClient) do(ctx context.Context, method, path string, query url
|
||||
return err
|
||||
}
|
||||
requestURL := buildRequestURL(client.baseURL, path, query)
|
||||
for attempt := 0; attempt < maxRequestAttempts; attempt++ {
|
||||
for attempt := range maxRequestAttempts {
|
||||
statusCode, retryHeader, responseBody, requestErr := client.send(ctx, method, requestURL, encodedBody)
|
||||
if requestErr != nil {
|
||||
return requestErr
|
||||
|
||||
@@ -94,6 +94,7 @@ func TestBuildSnapshotReadsZoneDomainCertificates(t *testing.T) {
|
||||
}
|
||||
|
||||
func generateTestCertKeyPairForSnapshot(t *testing.T) (certPEM string, keyPEM string) {
|
||||
t.Helper()
|
||||
return generateTestCertKeyPairForSnapshotForDomain(t, "test.example.com")
|
||||
}
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ package config_version
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"strconv"
|
||||
@@ -27,7 +28,7 @@ func normalizeSnapshotDomains(domains []string) ([]string, error) {
|
||||
for _, raw := range domains {
|
||||
domain := strings.ToLower(strings.TrimSpace(raw))
|
||||
if domain == "" || strings.Contains(domain, "://") || strings.Contains(domain, "/") {
|
||||
return nil, fmt.Errorf("domains payload is invalid")
|
||||
return nil, errors.New("domains payload is invalid")
|
||||
}
|
||||
if _, ok := seen[domain]; ok {
|
||||
continue
|
||||
@@ -36,7 +37,7 @@ func normalizeSnapshotDomains(domains []string) ([]string, error) {
|
||||
normalized = append(normalized, domain)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, fmt.Errorf("domain is required")
|
||||
return nil, errors.New("domain is required")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
@@ -55,7 +56,7 @@ func decodeStoredUpstreams(raw string, fallbackOriginURL string) ([]string, erro
|
||||
}
|
||||
var upstreams []string
|
||||
if err := json.Unmarshal([]byte(text), &upstreams); err != nil {
|
||||
return nil, fmt.Errorf("upstreams payload is invalid")
|
||||
return nil, errors.New("upstreams payload is invalid")
|
||||
}
|
||||
return normalizeUpstreams(fallbackOriginURL, upstreams)
|
||||
}
|
||||
@@ -79,7 +80,7 @@ func normalizeUpstreams(originURL string, upstreams []string) ([]string, error)
|
||||
normalized = append(normalized, value)
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil, fmt.Errorf("upstream is required")
|
||||
return nil, errors.New("upstream is required")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
@@ -91,7 +92,7 @@ func decodeStoredCustomHeaders(raw string) ([]customHeaderInput, error) {
|
||||
}
|
||||
var headers []customHeaderInput
|
||||
if err := json.Unmarshal([]byte(text), &headers); err != nil {
|
||||
return nil, fmt.Errorf("custom_headers payload is invalid")
|
||||
return nil, errors.New("custom_headers payload is invalid")
|
||||
}
|
||||
return headers, nil
|
||||
}
|
||||
@@ -103,7 +104,7 @@ func decodeStoredCacheRules(raw string) ([]string, error) {
|
||||
}
|
||||
var rules []string
|
||||
if err := json.Unmarshal([]byte(text), &rules); err != nil {
|
||||
return nil, fmt.Errorf("cache_rules payload is invalid")
|
||||
return nil, errors.New("cache_rules payload is invalid")
|
||||
}
|
||||
normalized := make([]string, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
|
||||
@@ -494,51 +494,51 @@ func diffOpenRestyOptionDetails(left openRestyConfigSnapshot, right openRestyCon
|
||||
CurrentValue: current,
|
||||
})
|
||||
}
|
||||
appendIfChanged("OpenRestyDefaultServerReturnStatus", fmt.Sprintf("%d", left.DefaultServerReturnStatus), fmt.Sprintf("%d", right.DefaultServerReturnStatus))
|
||||
appendIfChanged("OpenRestyDefaultServerReturnStatus", strconv.Itoa(left.DefaultServerReturnStatus), strconv.Itoa(right.DefaultServerReturnStatus))
|
||||
appendIfChanged("OpenRestyWorkerProcesses", left.WorkerProcesses, right.WorkerProcesses)
|
||||
appendIfChanged("OpenRestyWorkerConnections", fmt.Sprintf("%d", left.WorkerConnections), fmt.Sprintf("%d", right.WorkerConnections))
|
||||
appendIfChanged("OpenRestyWorkerRlimitNofile", fmt.Sprintf("%d", left.WorkerRlimitNofile), fmt.Sprintf("%d", right.WorkerRlimitNofile))
|
||||
appendIfChanged("OpenRestyWorkerConnections", strconv.Itoa(left.WorkerConnections), strconv.Itoa(right.WorkerConnections))
|
||||
appendIfChanged("OpenRestyWorkerRlimitNofile", strconv.Itoa(left.WorkerRlimitNofile), strconv.Itoa(right.WorkerRlimitNofile))
|
||||
appendIfChanged("OpenRestyEventsUse", left.EventsUse, right.EventsUse)
|
||||
appendIfChanged("OpenRestyEventsMultiAcceptEnabled", fmt.Sprintf("%t", left.EventsMultiAcceptEnabled), fmt.Sprintf("%t", right.EventsMultiAcceptEnabled))
|
||||
appendIfChanged("OpenRestyKeepaliveTimeout", fmt.Sprintf("%d", left.KeepaliveTimeout), fmt.Sprintf("%d", right.KeepaliveTimeout))
|
||||
appendIfChanged("OpenRestyKeepaliveRequests", fmt.Sprintf("%d", left.KeepaliveRequests), fmt.Sprintf("%d", right.KeepaliveRequests))
|
||||
appendIfChanged("OpenRestyClientHeaderTimeout", fmt.Sprintf("%d", left.ClientHeaderTimeout), fmt.Sprintf("%d", right.ClientHeaderTimeout))
|
||||
appendIfChanged("OpenRestyClientBodyTimeout", fmt.Sprintf("%d", left.ClientBodyTimeout), fmt.Sprintf("%d", right.ClientBodyTimeout))
|
||||
appendIfChanged("OpenRestyEventsMultiAcceptEnabled", strconv.FormatBool(left.EventsMultiAcceptEnabled), strconv.FormatBool(right.EventsMultiAcceptEnabled))
|
||||
appendIfChanged("OpenRestyKeepaliveTimeout", strconv.Itoa(left.KeepaliveTimeout), strconv.Itoa(right.KeepaliveTimeout))
|
||||
appendIfChanged("OpenRestyKeepaliveRequests", strconv.Itoa(left.KeepaliveRequests), strconv.Itoa(right.KeepaliveRequests))
|
||||
appendIfChanged("OpenRestyClientHeaderTimeout", strconv.Itoa(left.ClientHeaderTimeout), strconv.Itoa(right.ClientHeaderTimeout))
|
||||
appendIfChanged("OpenRestyClientBodyTimeout", strconv.Itoa(left.ClientBodyTimeout), strconv.Itoa(right.ClientBodyTimeout))
|
||||
appendIfChanged("OpenRestyClientMaxBodySize", left.ClientMaxBodySize, right.ClientMaxBodySize)
|
||||
appendIfChanged("OpenRestyLargeClientHeaderBuffers", left.LargeClientHeaderBuffers, right.LargeClientHeaderBuffers)
|
||||
appendIfChanged("OpenRestySendTimeout", fmt.Sprintf("%d", left.SendTimeout), fmt.Sprintf("%d", right.SendTimeout))
|
||||
appendIfChanged("OpenRestyProxyConnectTimeout", fmt.Sprintf("%d", left.ProxyConnectTimeout), fmt.Sprintf("%d", right.ProxyConnectTimeout))
|
||||
appendIfChanged("OpenRestyProxySendTimeout", fmt.Sprintf("%d", left.ProxySendTimeout), fmt.Sprintf("%d", right.ProxySendTimeout))
|
||||
appendIfChanged("OpenRestyProxyReadTimeout", fmt.Sprintf("%d", left.ProxyReadTimeout), fmt.Sprintf("%d", right.ProxyReadTimeout))
|
||||
appendIfChanged("OpenRestyWebsocketEnabled", fmt.Sprintf("%t", left.WebsocketEnabled), fmt.Sprintf("%t", right.WebsocketEnabled))
|
||||
appendIfChanged("OpenRestyHTTP3Enabled", fmt.Sprintf("%t", left.HTTP3Enabled), fmt.Sprintf("%t", right.HTTP3Enabled))
|
||||
appendIfChanged("OpenRestyProxyRequestBufferingEnabled", fmt.Sprintf("%t", left.ProxyRequestBuffering), fmt.Sprintf("%t", right.ProxyRequestBuffering))
|
||||
appendIfChanged("OpenRestyProxyBufferingEnabled", fmt.Sprintf("%t", left.ProxyBufferingEnabled), fmt.Sprintf("%t", right.ProxyBufferingEnabled))
|
||||
appendIfChanged("OpenRestySendTimeout", strconv.Itoa(left.SendTimeout), strconv.Itoa(right.SendTimeout))
|
||||
appendIfChanged("OpenRestyProxyConnectTimeout", strconv.Itoa(left.ProxyConnectTimeout), strconv.Itoa(right.ProxyConnectTimeout))
|
||||
appendIfChanged("OpenRestyProxySendTimeout", strconv.Itoa(left.ProxySendTimeout), strconv.Itoa(right.ProxySendTimeout))
|
||||
appendIfChanged("OpenRestyProxyReadTimeout", strconv.Itoa(left.ProxyReadTimeout), strconv.Itoa(right.ProxyReadTimeout))
|
||||
appendIfChanged("OpenRestyWebsocketEnabled", strconv.FormatBool(left.WebsocketEnabled), strconv.FormatBool(right.WebsocketEnabled))
|
||||
appendIfChanged("OpenRestyHTTP3Enabled", strconv.FormatBool(left.HTTP3Enabled), strconv.FormatBool(right.HTTP3Enabled))
|
||||
appendIfChanged("OpenRestyProxyRequestBufferingEnabled", strconv.FormatBool(left.ProxyRequestBuffering), strconv.FormatBool(right.ProxyRequestBuffering))
|
||||
appendIfChanged("OpenRestyProxyBufferingEnabled", strconv.FormatBool(left.ProxyBufferingEnabled), strconv.FormatBool(right.ProxyBufferingEnabled))
|
||||
appendIfChanged("OpenRestyProxyBuffers", left.ProxyBuffers, right.ProxyBuffers)
|
||||
appendIfChanged("OpenRestyProxyBufferSize", left.ProxyBufferSize, right.ProxyBufferSize)
|
||||
appendIfChanged("OpenRestyProxyBusyBuffersSize", left.ProxyBusyBuffersSize, right.ProxyBusyBuffersSize)
|
||||
appendIfChanged("OpenRestyGzipEnabled", fmt.Sprintf("%t", left.GzipEnabled), fmt.Sprintf("%t", right.GzipEnabled))
|
||||
appendIfChanged("OpenRestyGzipMinLength", fmt.Sprintf("%d", left.GzipMinLength), fmt.Sprintf("%d", right.GzipMinLength))
|
||||
appendIfChanged("OpenRestyGzipCompLevel", fmt.Sprintf("%d", left.GzipCompLevel), fmt.Sprintf("%d", right.GzipCompLevel))
|
||||
appendIfChanged("OpenRestyGzipEnabled", strconv.FormatBool(left.GzipEnabled), strconv.FormatBool(right.GzipEnabled))
|
||||
appendIfChanged("OpenRestyGzipMinLength", strconv.Itoa(left.GzipMinLength), strconv.Itoa(right.GzipMinLength))
|
||||
appendIfChanged("OpenRestyGzipCompLevel", strconv.Itoa(left.GzipCompLevel), strconv.Itoa(right.GzipCompLevel))
|
||||
appendIfChanged("OpenRestyResolvers", left.Resolvers, right.Resolvers)
|
||||
appendIfChanged("OpenRestyCacheEnabled", fmt.Sprintf("%t", left.CacheEnabled), fmt.Sprintf("%t", right.CacheEnabled))
|
||||
appendIfChanged("OpenRestyCacheEnabled", strconv.FormatBool(left.CacheEnabled), strconv.FormatBool(right.CacheEnabled))
|
||||
appendIfChanged("OpenRestyCachePath", left.CachePath, right.CachePath)
|
||||
appendIfChanged("OpenRestyCacheLevels", left.CacheLevels, right.CacheLevels)
|
||||
appendIfChanged("OpenRestyCacheInactive", left.CacheInactive, right.CacheInactive)
|
||||
appendIfChanged("OpenRestyCacheMaxSize", left.CacheMaxSize, right.CacheMaxSize)
|
||||
appendIfChanged("OpenRestyCacheKeyTemplate", left.CacheKeyTemplate, right.CacheKeyTemplate)
|
||||
appendIfChanged("OpenRestyCacheLockEnabled", fmt.Sprintf("%t", left.CacheLockEnabled), fmt.Sprintf("%t", right.CacheLockEnabled))
|
||||
appendIfChanged("OpenRestyCacheLockEnabled", strconv.FormatBool(left.CacheLockEnabled), strconv.FormatBool(right.CacheLockEnabled))
|
||||
appendIfChanged("OpenRestyCacheLockTimeout", left.CacheLockTimeout, right.CacheLockTimeout)
|
||||
appendIfChanged("OpenRestyCacheUseStale", left.CacheUseStale, right.CacheUseStale)
|
||||
appendIfChanged("OpenRestyDefaultLimitConnPerServer", fmt.Sprintf("%d", left.DefaultLimitConnPerServer), fmt.Sprintf("%d", right.DefaultLimitConnPerServer))
|
||||
appendIfChanged("OpenRestyDefaultLimitConnPerIP", fmt.Sprintf("%d", left.DefaultLimitConnPerIP), fmt.Sprintf("%d", right.DefaultLimitConnPerIP))
|
||||
appendIfChanged("OpenRestyDefaultLimitConnPerServer", strconv.Itoa(left.DefaultLimitConnPerServer), strconv.Itoa(right.DefaultLimitConnPerServer))
|
||||
appendIfChanged("OpenRestyDefaultLimitConnPerIP", strconv.Itoa(left.DefaultLimitConnPerIP), strconv.Itoa(right.DefaultLimitConnPerIP))
|
||||
appendIfChanged("OpenRestyDefaultLimitRate", left.DefaultLimitRate, right.DefaultLimitRate)
|
||||
appendIfChanged("OpenRestyDefaultLimitReqPerIP", left.DefaultLimitReqPerIP, right.DefaultLimitReqPerIP)
|
||||
appendIfChanged("OriginErrorPageEnabled", fmt.Sprintf("%t", left.OriginErrorPageEnabled), fmt.Sprintf("%t", right.OriginErrorPageEnabled))
|
||||
appendIfChanged("OriginErrorPageEnabled", strconv.FormatBool(left.OriginErrorPageEnabled), strconv.FormatBool(right.OriginErrorPageEnabled))
|
||||
appendIfChanged("OriginErrorPageStatusCodes", encodeOriginErrorPageStatusCodes(left.OriginErrorPageStatusCodes), encodeOriginErrorPageStatusCodes(right.OriginErrorPageStatusCodes))
|
||||
appendIfChanged("OriginErrorPageHTML", left.OriginErrorPageHTML, right.OriginErrorPageHTML)
|
||||
appendIfChanged("OriginErrorPageGetOnly", fmt.Sprintf("%t", left.OriginErrorPageGetOnly), fmt.Sprintf("%t", right.OriginErrorPageGetOnly))
|
||||
appendIfChanged("SWOfflineEnabled", fmt.Sprintf("%t", left.SWOfflineEnabled), fmt.Sprintf("%t", right.SWOfflineEnabled))
|
||||
appendIfChanged("OriginErrorPageGetOnly", strconv.FormatBool(left.OriginErrorPageGetOnly), strconv.FormatBool(right.OriginErrorPageGetOnly))
|
||||
appendIfChanged("SWOfflineEnabled", strconv.FormatBool(left.SWOfflineEnabled), strconv.FormatBool(right.SWOfflineEnabled))
|
||||
appendIfChanged("SWOfflineHTML", left.SWOfflineHTML, right.SWOfflineHTML)
|
||||
appendIfChanged("SWOfflineDomains", encodeSWOfflineDomains(left.SWOfflineDomains), encodeSWOfflineDomains(right.SWOfflineDomains))
|
||||
return changes
|
||||
|
||||
@@ -98,7 +98,7 @@ func TestDiffOpenRestyOptionDetailsOriginErrorPage(t *testing.T) {
|
||||
assert.Equal(t, "false", keys["OriginErrorPageEnabled"].CurrentValue)
|
||||
assert.Equal(t, `["500-599"]`, keys["OriginErrorPageStatusCodes"].PreviousValue)
|
||||
assert.Equal(t, `["522"]`, keys["OriginErrorPageStatusCodes"].CurrentValue)
|
||||
assert.Equal(t, "", keys["OriginErrorPageHTML"].PreviousValue)
|
||||
assert.Empty(t, keys["OriginErrorPageHTML"].PreviousValue)
|
||||
assert.Equal(t, "<p>x</p>", keys["OriginErrorPageHTML"].CurrentValue)
|
||||
}
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -199,9 +200,7 @@ func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configB
|
||||
return nil, err
|
||||
}
|
||||
|
||||
mainConfig := ""
|
||||
routeConfig := ""
|
||||
checksum := ""
|
||||
var mainConfig, routeConfig, checksum string
|
||||
supportFiles := []SupportFile(nil)
|
||||
|
||||
rendered, renderErr := renderSnapshotConfig(string(snapshotJSON), certificateFiles)
|
||||
@@ -429,7 +428,7 @@ func buildSnapshotWAFIPGroups(ctx context.Context, idSet map[uint]struct{}) ([]s
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
|
||||
slices.Sort(ids)
|
||||
groups, err := listWAFIPGroupsByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -477,7 +476,7 @@ func decodeIPList(raw string) ([]string, error) {
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(text), &items); err != nil {
|
||||
return nil, fmt.Errorf("ip_list payload is invalid")
|
||||
return nil, errors.New("ip_list payload is invalid")
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
@@ -629,7 +628,7 @@ func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) (
|
||||
for certID := range certIDSet {
|
||||
certIDs = append(certIDs, certID)
|
||||
}
|
||||
sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] })
|
||||
slices.Sort(certIDs)
|
||||
files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate)
|
||||
for _, certID := range certIDs {
|
||||
certificate, err := repository.GetTLSCertificateByID(ctx, certID)
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
)
|
||||
|
||||
func setupDashboardTestDB(t *testing.T) func() {
|
||||
t.Helper()
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
@@ -124,8 +125,8 @@ func TestGetOverviewStructure(t *testing.T) {
|
||||
onlineNodeCheck := overview.Nodes
|
||||
require.NotEmpty(t, onlineNodeCheck)
|
||||
|
||||
assert.Equal(t, 55.0, overview.Capacity.AverageCPUUsagePercent)
|
||||
assert.Equal(t, 50.0, overview.Capacity.AverageMemoryUsagePercent)
|
||||
assert.InDelta(t, 55.0, overview.Capacity.AverageCPUUsagePercent, 1e-9)
|
||||
assert.InDelta(t, 50.0, overview.Capacity.AverageMemoryUsagePercent, 1e-9)
|
||||
assert.Equal(t, 0, overview.Capacity.HighCPUNodes)
|
||||
assert.Equal(t, 0, overview.Capacity.HighMemoryNodes)
|
||||
assert.Equal(t, 0, overview.Capacity.HighStorageNodes)
|
||||
@@ -171,11 +172,11 @@ func TestGetOverviewStructure(t *testing.T) {
|
||||
assert.Equal(t, "online", onlineNode[6])
|
||||
assert.Equal(t, "healthy", onlineNode[7])
|
||||
// Latest-per-node health fields (indexes match compressDashboardNodes).
|
||||
assert.Equal(t, 55.0, onlineNode[11]) // cpu_usage_percent from latest snapshot
|
||||
assert.Equal(t, 50.0, onlineNode[12]) // memory_usage_percent
|
||||
assert.Equal(t, int64(12), onlineNode[14]) // request_count from access logs
|
||||
assert.Equal(t, int64(1), onlineNode[15]) // error_count
|
||||
assert.Equal(t, int64(4), onlineNode[16]) // unique visitors
|
||||
assert.InDelta(t, 55.0, onlineNode[11], 1e-9) // cpu_usage_percent from latest snapshot
|
||||
assert.InDelta(t, 50.0, onlineNode[12], 1e-9) // memory_usage_percent
|
||||
assert.Equal(t, int64(12), onlineNode[14]) // request_count from access logs
|
||||
assert.Equal(t, int64(1), onlineNode[15]) // error_count
|
||||
assert.Equal(t, int64(4), onlineNode[16]) // unique visitors
|
||||
|
||||
pendingNode := nodeByID["node-dashboard-2"]
|
||||
require.NotNil(t, pendingNode)
|
||||
@@ -183,6 +184,6 @@ func TestGetOverviewStructure(t *testing.T) {
|
||||
assert.Equal(t, "pending", pendingNode[6])
|
||||
assert.Equal(t, "unknown", pendingNode[7])
|
||||
|
||||
assert.Equal(t, 55.0, overview.Capacity.AverageCPUUsagePercent)
|
||||
assert.InDelta(t, 55.0, overview.Capacity.AverageCPUUsagePercent, 1e-9)
|
||||
assert.Equal(t, 1, overview.Traffic.ReportedNodes)
|
||||
}
|
||||
|
||||
@@ -28,7 +28,7 @@ const (
|
||||
// Heartbeat processes an OpenFlared heartbeat and returns runtime settings.
|
||||
func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload HeartbeatPayload) (*HeartbeatResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("tunnel client node is nil")
|
||||
return nil, errors.New("tunnel client node is nil")
|
||||
}
|
||||
if node.NodeType != "tunnel_client" {
|
||||
return nil, fmt.Errorf("node %s is not a tunnel_client", node.NodeID)
|
||||
@@ -95,7 +95,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
|
||||
// GetTunnelConfig builds the full tunnel routing config for an OpenFlared client.
|
||||
func GetTunnelConfig(ctx context.Context, node *model.OpenFlareNode) (*TunnelConfigResponse, error) {
|
||||
if node == nil {
|
||||
return nil, fmt.Errorf("node is nil")
|
||||
return nil, errors.New("node is nil")
|
||||
}
|
||||
|
||||
activeVersion, err := getActiveConfigMeta(ctx)
|
||||
|
||||
@@ -37,7 +37,11 @@ func PostHeartbeat(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := Heartbeat(c.Request.Context(), node, payload)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
@@ -63,7 +67,11 @@ func GetActiveConfig(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
config, err := GetTunnelConfig(c.Request.Context(), node)
|
||||
if apiutil.AbortBadRequestOnError(c, err) {
|
||||
@@ -91,7 +99,9 @@ func PostApplyLog(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
if authNode, ok := c.Get(ctxFlaredNodeKey); ok {
|
||||
payload.NodeID = authNode.(*model.OpenFlareNode).NodeID
|
||||
if node, ok := authNode.(*model.OpenFlareNode); ok {
|
||||
payload.NodeID = node.NodeID
|
||||
}
|
||||
}
|
||||
|
||||
log, err := ReportApplyLog(c.Request.Context(), payload)
|
||||
@@ -115,6 +125,10 @@ func GetWebSocket(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
node := authNode.(*model.OpenFlareNode)
|
||||
node, ok := authNode.(*model.OpenFlareNode)
|
||||
if !ok {
|
||||
response.AbortUnauthorized(c, errTunnelTokenInvalid)
|
||||
return
|
||||
}
|
||||
ofws.ServeFlared(c, node.NodeID)
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user