merge: replace legacy openflare-server with Wavelet rename

This commit is contained in:
ryan
2026-06-19 11:30:05 +08:00
1303 changed files with 40520 additions and 157754 deletions
@@ -0,0 +1,245 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package auth_source 提供认证源管理功能
package auth_source
import (
"errors"
"fmt"
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// AuthSourceRequest 创建或更新认证源的请求参数
type AuthSourceRequest struct {
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
ClientID string `json:"client_id"`
ClientSecret string `json:"client_secret"`
OpenIDDiscoveryURL string `json:"openid_discovery_url"`
Scopes string `json:"scopes"`
IconURL string `json:"icon_url"`
}
// ToggleAuthSourceRequest 切换认证源启用状态的请求参数
type ToggleAuthSourceRequest struct {
IsActive bool `json:"is_active"`
}
// ListAuthSources 获取认证源列表
// @Summary 获取认证源列表
// @Description 返回所有已配置的 OAuth/OIDC 认证源列表,包括已启用和未启用的,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.AuthSource} "认证源列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/auth-sources [get]
func ListAuthSources(c *gin.Context) {
sources, err := model.GetAuthSources(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(sources))
}
// CreateAuthSource 创建认证源
// @Summary 创建认证源
// @Description 创建一个新的 OAuth/OIDC 认证源配置,认证源名称必须唯一且符合命名规范,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body auth_source.AuthSourceRequest true "创建认证源参数"
// @Success 200 {object} response.Any{data=model.AuthSource} "创建成功,返回认证源信息"
// @Failure 400 {object} response.Any "参数错误或验证失败"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/auth-sources [post]
func CreateAuthSource(c *gin.Context) {
var req AuthSourceRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
source := model.AuthSource{
Name: req.Name,
Type: req.Type,
DisplayName: req.DisplayName,
IsActive: req.IsActive,
ClientID: req.ClientID,
ClientSecret: req.ClientSecret,
OpenIDDiscoveryURL: req.OpenIDDiscoveryURL,
Scopes: req.Scopes,
IconURL: req.IconURL,
}
if err := model.CreateAuthSource(c.Request.Context(), &source); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
source.Sanitize()
c.JSON(http.StatusOK, response.OK(source))
}
// UpdateAuthSource 更新认证源
// @Summary 更新认证源
// @Description 更新指定 ID 的认证源配置。若 client_secret 字段为空,则保留原有密钥不变,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "认证源 ID 或名称"
// @Param request body auth_source.AuthSourceRequest true "更新认证源参数"
// @Success 200 {object} response.Any{data=model.AuthSource} "更新成功,返回更新后的认证源信息"
// @Failure 400 {object} response.Any "参数错误或验证失败"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/auth-sources/{id} [put]
func UpdateAuthSource(c *gin.Context) {
id, err := parseSourceID(c)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
var req AuthSourceRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
// 记录更新前的 Discovery URL,以便更新成功后清除旧缓存条目。
existing, _ := model.GetAuthSourceByID(c.Request.Context(), id)
source := model.AuthSource{
ID: id,
Name: req.Name,
Type: req.Type,
DisplayName: req.DisplayName,
IsActive: req.IsActive,
ClientID: req.ClientID,
ClientSecret: req.ClientSecret,
OpenIDDiscoveryURL: req.OpenIDDiscoveryURL,
Scopes: req.Scopes,
IconURL: req.IconURL,
}
keepSecret := source.ClientSecret == ""
if err := model.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
// Discovery URL 可能已变更,清除旧、新 issuer 的 provider 缓存,
// 确保下次登录时重新拉取最新 OIDC 元数据。
if existing != nil {
oauth.InvalidateOIDCProviderCache(normalizeIssuer(existing.OpenIDDiscoveryURL))
}
oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL))
updated, err := model.GetAuthSourceByID(c.Request.Context(), id)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
updated.Sanitize()
c.JSON(http.StatusOK, response.OK(updated))
}
// ToggleAuthSource 切换认证源启用状态
// @Summary 切换认证源启用状态
// @Description 启用或禁用指定认证源。尝试启用时将验证 Client ID 和 Client Secret 是否已配置,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "认证源 ID 或名称"
// @Param request body auth_source.ToggleAuthSourceRequest true "启用状态"
// @Success 200 {object} response.Any{data=string} "切换成功"
// @Failure 400 {object} response.Any "验证失败或认证源不存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/auth-sources/{id}/toggle [put]
func ToggleAuthSource(c *gin.Context) {
id, err := parseSourceID(c)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
var req ToggleAuthSourceRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := model.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// DeleteAuthSource 删除认证源
// @Summary 删除认证源
// @Description 删除指定认证源及其关联的所有外部帐号绑定记录,警告:删除后相关用户将无法通过该源登录,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "认证源 ID 或名称"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "ID 无效或删除失败"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/auth-sources/{id} [delete]
func DeleteAuthSource(c *gin.Context) {
id, err := parseSourceID(c)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := model.DeleteAuthSource(c.Request.Context(), id); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func parseSourceID(c *gin.Context) (uint64, error) {
raw := c.Param("id")
if raw == "" {
return 0, errors.New(admin.InvalidAuthSourceID)
}
source, err := model.GetAuthSourceByName(c.Request.Context(), raw)
if err == nil {
return source.ID, nil
}
var id uint64
if _, scanErr := fmt.Sscanf(raw, "%d", &id); scanErr != nil || id == 0 {
return 0, errors.New(admin.InvalidAuthSourceID)
}
return id, nil
}
// normalizeIssuer 将 Discovery URL 规范化为 issuer 基础 URL,
// 与 oauth.buildOAuthConfig 中的规范化逻辑保持一致。
func normalizeIssuer(discoveryURL string) string {
issuer := strings.TrimSuffix(strings.TrimSpace(discoveryURL), "/")
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
return issuer
}
@@ -0,0 +1,333 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth_source
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
func setupTestRouter(authUser *model.User) *gin.Engine {
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
adminGroup.GET("/auth-sources", ListAuthSources)
adminGroup.POST("/auth-sources", CreateAuthSource)
adminGroup.PUT("/auth-sources/:id", UpdateAuthSource)
adminGroup.PUT("/auth-sources/:id/toggle", ToggleAuthSource)
adminGroup.DELETE("/auth-sources/:id", DeleteAuthSource)
return r
}
func TestListAuthSources(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Seed source
source := model.AuthSource{
ID: 1,
Name: "google",
Type: "oidc",
DisplayName: "Google Auth",
IsActive: true,
ClientID: "client_id_123",
ClientSecret: "client_secret_456",
OpenIDDiscoveryURL: "https://accounts.google.com",
}
dbConn.Create(&source)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
req, _ := http.NewRequest("GET", "/api/v1/admin/auth-sources", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d", w.Code)
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var sources []model.AuthSource
_ = json.Unmarshal(dataBytes, &sources)
if len(sources) != 1 {
t.Errorf("expected 1 auth source, got %d", len(sources))
}
if sources[0].Name != "google" {
t.Errorf("expected name 'google', got '%s'", sources[0].Name)
}
// Verify sanitize removed the secret
if sources[0].ClientSecret != "" {
t.Error("client secret should be sanitized")
}
if !sources[0].ClientSecretConfigured {
t.Error("client secret configured flag should be true")
}
}
func TestCreateAuthSource(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("create successfully", func(t *testing.T) {
reqPayload := AuthSourceRequest{
Name: "github",
Type: "oidc",
DisplayName: "GitHub OIDC",
IsActive: true,
ClientID: "client_id_gh",
ClientSecret: "client_secret_gh",
OpenIDDiscoveryURL: "https://github.com",
}
body, _ := json.Marshal(reqPayload)
req, _ := http.NewRequest("POST", "/api/v1/admin/auth-sources", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
// Verify database
var src model.AuthSource
dbConn.Where("name = ?", "github").First(&src)
if src.ClientID != "client_id_gh" {
t.Errorf("expected client_id_gh, got '%s'", src.ClientID)
}
})
t.Run("create invalid validation failure", func(t *testing.T) {
reqPayload := AuthSourceRequest{
Name: "invalid name!",
Type: "oidc",
DisplayName: "Invalid",
IsActive: true,
ClientID: "client_id_val",
ClientSecret: "client_secret_val",
OpenIDDiscoveryURL: "https://discovery.url",
}
body, _ := json.Marshal(reqPayload)
req, _ := http.NewRequest("POST", "/api/v1/admin/auth-sources", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request, got %d", w.Code)
}
})
}
func TestUpdateAuthSource(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Seed source
source := model.AuthSource{
ID: 1,
Name: "microsoft",
Type: "oidc",
DisplayName: "Microsoft",
IsActive: true,
ClientID: "old_client_id",
ClientSecret: "old_secret",
OpenIDDiscoveryURL: "https://login.microsoftonline.com",
}
dbConn.Create(&source)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("update keep client secret", func(t *testing.T) {
reqPayload := AuthSourceRequest{
Name: "microsoft",
Type: "oidc",
DisplayName: "Microsoft Updated",
IsActive: true,
ClientID: "new_client_id",
ClientSecret: "", // empty implies keeping existing secret
OpenIDDiscoveryURL: "https://login.microsoftonline.com",
}
body, _ := json.Marshal(reqPayload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var src model.AuthSource
dbConn.First(&src, 1)
if src.DisplayName != "Microsoft Updated" {
t.Errorf("expected display name update, got '%s'", src.DisplayName)
}
if src.ClientSecret != "old_secret" {
t.Errorf("expected old secret to be preserved, got '%s'", src.ClientSecret)
}
})
t.Run("update new client secret", func(t *testing.T) {
reqPayload := AuthSourceRequest{
Name: "microsoft",
Type: "oidc",
DisplayName: "Microsoft Updated Again",
IsActive: true,
ClientID: "new_client_id",
ClientSecret: "brand_new_secret",
OpenIDDiscoveryURL: "https://login.microsoftonline.com",
}
body, _ := json.Marshal(reqPayload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/microsoft", bytes.NewBuffer(body)) // Using Name instead of ID
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var src model.AuthSource
dbConn.First(&src, 1)
if src.ClientSecret != "brand_new_secret" {
t.Errorf("expected secret update, got '%s'", src.ClientSecret)
}
})
}
func TestToggleAuthSource(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
source := model.AuthSource{
ID: 1,
Name: "test_source",
Type: "oidc",
DisplayName: "Test Source",
IsActive: false,
ClientID: "",
ClientSecret: "",
OpenIDDiscoveryURL: "https://test.discovery.url",
}
dbConn.Create(&source)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("cannot activate without credentials", func(t *testing.T) {
payload := ToggleAuthSourceRequest{IsActive: true}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1/toggle", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request when activating without client_id/secret, got %d", w.Code)
}
})
t.Run("toggle success after setting credentials", func(t *testing.T) {
// Set credentials first
dbConn.Model(&model.AuthSource{}).Where("id = ?", 1).Updates(map[string]interface{}{
"client_id": "id",
"client_secret": "secret",
})
payload := ToggleAuthSourceRequest{IsActive: true}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/auth-sources/1/toggle", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var src model.AuthSource
dbConn.First(&src, 1)
if !src.IsActive {
t.Error("auth source should be activated")
}
})
}
func TestDeleteAuthSource(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
source := model.AuthSource{
ID: 1,
Name: "delete_me",
Type: "oidc",
DisplayName: "Delete Me",
IsActive: true,
ClientID: "id",
ClientSecret: "secret",
OpenIDDiscoveryURL: "https://delete.me",
}
dbConn.Create(&source)
externalAccount := model.ExternalAccount{
ID: 10,
AuthSourceID: 1,
UserID: 50,
ExternalID: "ext_50",
}
dbConn.Create(&externalAccount)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
req, _ := http.NewRequest("DELETE", "/api/v1/admin/auth-sources/1", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d", w.Code)
}
// Verify AuthSource is deleted
var srcCount int64
dbConn.Model(&model.AuthSource{}).Where("id = ?", 1).Count(&srcCount)
if srcCount != 0 {
t.Error("AuthSource should be deleted from the database")
}
// Verify ExternalAccount bindings are also deleted
var extCount int64
dbConn.Model(&model.ExternalAccount{}).Where("auth_source_id = ?", 1).Count(&extCount)
if extCount != 0 {
t.Error("related ExternalAccount bindings should be deleted")
}
}
+14
View File
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cache
import (
"context"
"github.com/Rain-kl/Wavelet/internal/repository"
)
func saveOrUpdateConfig(ctx context.Context, key, value string) error {
return repository.SaveOrUpdateSystemConfig(ctx, key, value)
}
+100
View File
@@ -0,0 +1,100 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cache provides HTTP handlers for managing disk cache.
package cache
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
type updateCacheConfigRequest struct {
MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"`
TTLMinutes int64 `json:"ttl_minutes" binding:"required,min=0"`
LRUEnabled bool `json:"lru_enabled"`
}
// GetCacheStatus 获取磁盘缓存状态与当前统计数据
// @Summary 获取缓存状态
// @Description 获取当前系统磁盘缓存的使用情况(已占用字节、Key 数量等)与策略配置
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/cache/status [get]
func GetCacheStatus(c *gin.Context) {
status := diskcache.GetGlobalCache().Status()
c.JSON(http.StatusOK, response.OK(status))
}
// UpdateCacheConfig 更新磁盘缓存策略配置
// @Summary 更新缓存配置
// @Description 更改磁盘缓存最大容量限制、文件生存时间(TTL)以及是否启用 LRU 淘汰淘汰算法,并进行热更新
// @Tags admin
// @Accept json
// @Produce json
// @Param request body cache.updateCacheConfigRequest true "缓存配置请求体"
// @Security SessionCookie
// @Success 200 {object} response.Any "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/config [post]
func UpdateCacheConfig(c *gin.Context) {
var req updateCacheConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
response.AbortInternal(c, err.Error())
return
}
diskcache.GetGlobalCache().ReloadConfig(ctx)
c.JSON(http.StatusOK, response.OKNil())
}
// ClearCache 一键清空所有磁盘缓存数据
// @Summary 清空缓存
// @Description 清除系统磁盘缓存目录中的所有临时文件,并重置缓存容量和 Key 追踪数据
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "清理成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/clear [post]
func ClearCache(c *gin.Context) {
if err := diskcache.GetGlobalCache().Clear(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,498 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package db_manage provides router handlers for managing database tables,
// overview information, and executing custom SQL queries.
package db_manage
import (
"database/sql"
"fmt"
"math"
"net/http"
"os"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
const (
binaryKB = 0
binaryMB = 1
binaryGB = 2
valueThreshold = 10
maxStringLength = 200
)
// DBOverviewResponse 数据库运行概览响应结构体
type DBOverviewResponse struct {
Type string `json:"type"`
Version string `json:"version"`
Name string `json:"name"`
Size string `json:"size"`
TableCount int64 `json:"table_count"`
Connections int64 `json:"connections"`
}
// GetTableDataRequest 分页拉取表数据请求结构体
type GetTableDataRequest struct {
Table string `form:"table" binding:"required"`
Page int `form:"page,default=1"`
PageSize int `form:"pageSize,default=10"`
}
// TableDataResponse 动态数据表响应结构体
type TableDataResponse struct {
Columns []string `json:"columns"`
Total int64 `json:"total"`
Results []map[string]interface{} `json:"results"`
}
// ExecuteSQLRequest 执行自定义 SQL 请求结构体
type ExecuteSQLRequest struct {
SQL string `json:"sql" binding:"required"`
}
// 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"`
}
// formatBytes 格式化字节大小为可读字符串
func formatBytes(bytes uint64) string {
const unit = 1024
if bytes < unit {
return fmt.Sprintf("%d B", bytes)
}
div, exp := int64(unit), 0
for n := bytes / unit; n >= unit; n /= unit {
div *= unit
exp++
}
value := float64(bytes) / float64(div)
var suffix string
switch exp {
case binaryKB:
suffix = "KiB"
case binaryMB:
suffix = "MiB"
case binaryGB:
suffix = "GiB"
default:
suffix = "TiB"
}
if value == math.Trunc(value) {
if value >= valueThreshold {
return fmt.Sprintf("%.0f %s", value, suffix)
}
return fmt.Sprintf("%.1f %s", value, suffix)
}
return fmt.Sprintf("%.1f %s", value, suffix)
}
// getSQLiteOverview 获取 SQLite 数据库概览信息
func getSQLiteOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
name := config.Config.Database.SQLitePath
if name == "" {
name = "./data/openflare.db"
}
var version string
var ver string
if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil {
version = "SQLite " + ver
} else {
version = "SQLite"
}
var sizeStr string
if fi, err := os.Stat(name); err == nil {
size := fi.Size()
if size < 0 {
size = 0
}
sizeStr = formatBytes(uint64(size))
} else {
sizeStr = "0 B"
}
var tableCount int64
if err := gormDB.Raw("SELECT count(*) FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'").Scan(&tableCount).Error; err != nil {
tableCount = 0
}
var connCount int64
if sqlDB, err := gormDB.DB(); err == nil {
connCount = int64(sqlDB.Stats().OpenConnections)
} else {
connCount = 1
}
return DBOverviewResponse{
Type: "sqlite",
Version: version,
Name: name,
Size: sizeStr,
TableCount: tableCount,
Connections: connCount,
}, nil
}
// getPostgresOverview 获取 PostgreSQL 数据库概览信息
func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
name := config.Config.Database.Database
var version string
var ver string
if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil {
version = ver
} else {
version = "PostgreSQL"
}
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
}
sizeStr = formatBytes(uint64(size))
} else {
sizeStr = "0 B"
}
var tableCount int64
if err := gormDB.Raw("SELECT count(*) FROM information_schema.tables WHERE table_schema = current_schema()").Scan(&tableCount).Error; err != nil {
tableCount = 0
}
var connCount int64
var pgc sql.NullInt64
if err := gormDB.Raw("SELECT count(*) FROM pg_stat_activity WHERE datname = current_database()").Scan(&pgc).Error; err == nil && pgc.Valid {
connCount = pgc.Int64
} else {
if sqlDB, err := gormDB.DB(); err == nil {
connCount = int64(sqlDB.Stats().OpenConnections)
} else {
connCount = 1
}
}
return DBOverviewResponse{
Type: "postgres",
Version: version,
Name: name,
Size: sizeStr,
TableCount: tableCount,
Connections: connCount,
}, nil
}
// GetDBOverview 获取数据库运行概览
// @Summary 获取数据库运行概览
// @Description 获取数据库类型、版本、名称、文件大小、表数量及当前连接数,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=db_manage.DBOverviewResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/overview [get]
func GetDBOverview(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
var overview DBOverviewResponse
var err error
if !config.Config.Database.Enabled {
overview, err = getSQLiteOverview(gormDB)
} else {
overview, err = getPostgresOverview(gormDB)
}
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(overview))
}
// ListDBTables 获取数据库所有表名
// @Summary 获取数据库所有表名
// @Description 返回当前数据库的所有用户自定义表名称列表,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]string} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/tables [get]
func ListDBTables(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
var tables []string
var err error
if !config.Config.Database.Enabled {
err = gormDB.Raw("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name").Scan(&tables).Error
} else {
err = gormDB.Raw("SELECT table_name FROM information_schema.tables WHERE table_schema = current_schema() ORDER BY table_name").Scan(&tables).Error
}
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(tables))
}
// GetDBTableData 获取数据表 data
func GetDBTableData(c *gin.Context) {
var req GetTableDataRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
// 安全转义表名并拼接
quotedTable := `"` + strings.ReplaceAll(req.Table, `"`, `""`) + `"`
var total int64
if err := gormDB.Raw("SELECT count(*) FROM " + quotedTable).Scan(&total).Error; err != nil {
response.AbortBadRequest(c, err.Error())
return
}
offset := (req.Page - 1) * req.PageSize
if offset < 0 {
offset = 0
}
limit := req.PageSize
if limit <= 0 {
limit = 10
}
rows, err := gormDB.Raw("SELECT * FROM "+quotedTable+" LIMIT ? OFFSET ?", limit, offset).Rows()
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
defer func() {
_ = rows.Close()
}()
cols, err := rows.Columns()
if err != nil {
response.AbortInternal(c, err.Error())
return
}
results, err := scanTableRows(rows, cols)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(TableDataResponse{
Columns: cols,
Total: total,
Results: results,
}))
}
// scanTableRows 扫描并提取数据表行数据,做截断处理
func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]interface{}, error) {
results := make([]map[string]interface{}, 0)
for rows.Next() {
columns := make([]interface{}, len(cols))
columnPointers := make([]interface{}, len(cols))
for i := range columns {
columnPointers[i] = &columns[i]
}
if err := rows.Scan(columnPointers...); err != nil {
return nil, err
}
rowMap := make(map[string]interface{})
for i, colName := range cols {
val := columns[i]
if b, ok := val.([]byte); ok {
strVal := string(b)
runes := []rune(strVal)
if len(runes) > maxStringLength {
strVal = string(runes[:maxStringLength]) + "..."
}
rowMap[colName] = strVal
} else if str, ok := val.(string); ok {
runes := []rune(str)
if len(runes) > maxStringLength {
str = string(runes[:maxStringLength]) + "..."
}
rowMap[colName] = str
} else {
rowMap[colName] = val
}
}
results = append(results, rowMap)
}
return results, nil
}
// executeSQLQuery 执行并解析查询类 SQL 语句
func executeSQLQuery(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) {
rows, err := gormDB.Raw(sqlStr).Rows()
if err != nil {
return ExecuteSQLResponse{}, err
}
defer func() {
_ = rows.Close()
}()
cols, err := rows.Columns()
if err != nil {
return ExecuteSQLResponse{}, err
}
results := make([]map[string]interface{}, 0)
for rows.Next() {
columns := make([]interface{}, len(cols))
columnPointers := make([]interface{}, len(cols))
for i := range columns {
columnPointers[i] = &columns[i]
}
if err := rows.Scan(columnPointers...); err != nil {
return ExecuteSQLResponse{}, err
}
rowMap := make(map[string]interface{})
for i, colName := range cols {
val := columns[i]
if b, ok := val.([]byte); ok {
rowMap[colName] = string(b)
} else {
rowMap[colName] = val
}
}
results = append(results, rowMap)
}
executionTime := time.Since(startTime).Milliseconds()
return ExecuteSQLResponse{
Type: "select",
Columns: cols,
Results: results,
AffectedRows: int64(len(results)),
ExecutionTimeMs: executionTime,
}, nil
}
// executeSQLMutation 执行修改/更新类 SQL 语句
func executeSQLMutation(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) {
tx := gormDB.Exec(sqlStr)
if tx.Error != nil {
return ExecuteSQLResponse{}, tx.Error
}
executionTime := time.Since(startTime).Milliseconds()
return ExecuteSQLResponse{
Type: "exec",
AffectedRows: tx.RowsAffected,
ExecutionTimeMs: executionTime,
}, nil
}
// ExecuteSQL 执行 SQL 查询
// @Summary 执行 SQL 查询
// @Description 在当前数据库中执行任意自定义 SQL,如果是查询语句将返回格式化后的列与数据集,否则返回受影响行数,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body db_manage.ExecuteSQLRequest true "SQL 请求参数"
// @Success 200 {object} response.Any{data=db_manage.ExecuteSQLResponse} "执行完毕"
// @Failure 400 {object} response.Any "SQL 语句错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/query [post]
func ExecuteSQL(c *gin.Context) {
var req ExecuteSQLRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
gormDB := db.DB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
}
trimmedSQL := strings.TrimSpace(req.SQL)
if trimmedSQL == "" {
response.AbortBadRequest(c, "SQL 语句不能为空")
return
}
startTime := time.Now()
// 识别是否是查询语句(SELECT, SHOW, EXPLAIN 等)
isQuery := false
lowerSQL := strings.ToLower(trimmedSQL)
queryKeywords := []string{"select", "show", "explain", "describe", "pragma"}
for _, kw := range queryKeywords {
if strings.HasPrefix(lowerSQL, kw) {
isQuery = true
break
}
}
var resp ExecuteSQLResponse
var err error
if isQuery {
resp, err = executeSQLQuery(gormDB, trimmedSQL, startTime)
} else {
resp, err = executeSQLMutation(gormDB, trimmedSQL, startTime)
}
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
@@ -0,0 +1,15 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package admin 提供管理后台功能
package admin
// 管理后台错误消息常量
const (
AdminRequired = "未经授权访问"
TokenAdminRequired = "该访问令牌没有管理员权限,无法访问管理端点" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
InvalidAuthSourceID = "认证源 ID 无效"
InvalidCursorParam = "无效的 cursor 参数"
InvalidTaskExecutionID = "无效的任务执行记录 ID"
)
@@ -0,0 +1,46 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/gin-gonic/gin"
)
// LoginAdminRequired 返回管理员权限校验中间件
func LoginAdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
// init trace
ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey)
if !tokenAdmin {
response.AbortNotFound(c, TokenAdminRequired)
return
}
}
if !user.IsAdmin {
response.AbortNotFound(c, AdminRequired)
return
}
// log
logger.InfoF(ctx, "[LoginAdminRequired] %d %s", user.ID, user.Username)
// next
c.Next()
}
}
@@ -0,0 +1,278 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"encoding/json"
"errors"
"net/http"
"strconv"
"strings"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListChannelDefinitions 获取各种消息通道的表单配置定义列表
// @Summary 获取所有消息通道配置字段定义
// @Description 返回系统支持的所有消息通道类型(如飞书、邮件、自定义、Telegram)的动态表单定义,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]Definition} "通道配置定义列表"
// @Router /api/v1/admin/push/channels/definitions [get]
func ListChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(ListDefinitions()))
}
// ListChannels 获取消息通道列表
// @Summary 获取所有消息通道
// @Description 返回系统配置的所有消息通道列表,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表"
// @Router /api/v1/admin/push/channels [get]
func ListChannels(c *gin.Context) {
channels, err := listPushChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channels))
}
// CreateChannelRequest 创建通道参数
type CreateChannelRequest struct {
Name string `json:"name" binding:"required"`
Description string `json:"description"`
Type string `json:"type" binding:"required"`
Token string `json:"token"`
URL string `json:"url"`
Other string `json:"other"`
Enabled bool `json:"enabled"`
}
// CreateChannel 创建消息通道
// @Summary 创建消息通道
// @Description 新建一个消息通道配置,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateChannelRequest true "创建参数"
// @Success 200 {object} response.Any{data=model.PushChannel} "创建成功"
// @Router /api/v1/admin/push/channels [post]
func CreateChannel(c *gin.Context) {
var req CreateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
channel, err := createPushChannel(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channel))
}
// UpdateChannelRequest 修改通道参数
type UpdateChannelRequest struct {
Description string `json:"description"`
Type string `json:"type" binding:"required"`
Token string `json:"token"`
URL string `json:"url"`
Other string `json:"other"`
Enabled bool `json:"enabled"`
}
// UpdateChannel 更新消息通道
// @Summary 更新消息通道
// @Description 修改消息通道配置,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "通道ID"
// @Param request body UpdateChannelRequest true "更新参数"
// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功"
// @Router /api/v1/admin/push/channels/{id} [put]
func UpdateChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
var req UpdateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
channel, err := updatePushChannel(c.Request.Context(), id, req)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(channel))
}
// DeleteChannel 删除消息通道
// @Summary 删除消息通道
// @Description 根据ID删除消息通道,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "通道ID"
// @Success 200 {object} response.Any "删除成功"
// @Router /api/v1/admin/push/channels/{id} [delete]
func DeleteChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
if err := deletePushChannel(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// TestChannelRequest 测试通道连通性参数
type TestChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Token string `json:"token"`
URL string `json:"url"`
Other string `json:"other"`
Target string `json:"target"`
}
// TestChannel 测试通道连通性
// @Summary 测试通道连通性
// @Description 触发一次临时的或现有的通道连通性推送测试,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body TestChannelRequest true "测试参数"
// @Success 200 {object} response.Any "测试触发成功"
// @Router /api/v1/admin/push/channels/test [post]
func TestChannel(c *gin.Context) {
var req TestChannelRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
url, token, other, channelType, err := loadChannelForTest(ctx, req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if channelType == channelEmail {
url, token, other = resolveSMTPConfig(ctx, url, token, other)
}
tempChannel := model.PushChannel{
Name: "test_temp",
URL: url,
Token: token,
Other: other,
Type: channelType,
Enabled: true,
}
if err := tempChannel.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
url = tempChannel.URL
var config pkgpush.Config
var renderedJSON string
switch channelType {
case channelLark:
config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token}
renderedJSON = other
case channelEmail:
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other}
default:
config = pkgpush.Config{Channel: channelCustom, URL: url}
customPushReq := CustomPushRequest{
Title: "通道测试通知",
Content: "这是一条来自系统的消息通道连通性测试消息。",
Description: "系统通道测试",
URL: "https://example.com",
To: req.Target,
}
renderedJSON = renderCustomPayload(other, customPushReq)
}
payload := SendPayload{
EventKey: "test_channel",
Config: config,
Target: req.Target,
Body: NotificationMessage{
Title: "通道测试通知",
Content: "这是一条来自系统的消息通道连通性测试消息。",
Level: defaultLevelInfo,
},
Template: renderedJSON,
}
if err := enqueuePushTask(ctx, payload); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// CustomPushRequest 外部公开推送请求参数
type CustomPushRequest struct {
Title string `json:"title" form:"title"`
Description string `json:"description" form:"description"`
Content string `json:"content" form:"content"`
URL string `json:"url" form:"url"`
To string `json:"to" form:"to"`
Token string `json:"token" form:"token"`
}
func escapeJSONString(s string) string {
b, _ := json.Marshal(s)
const minJSONLen = 2
if len(b) >= minJSONLen {
return string(b[1 : len(b)-1])
}
return s
}
func renderCustomPayload(template string, req CustomPushRequest) string {
result := template
result = strings.ReplaceAll(result, "$title", escapeJSONString(req.Title))
result = strings.ReplaceAll(result, "$description", escapeJSONString(req.Description))
result = strings.ReplaceAll(result, "$content", escapeJSONString(req.Content))
result = strings.ReplaceAll(result, "$url", escapeJSONString(req.URL))
result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To))
return result
}
@@ -0,0 +1,183 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import "sync"
const (
// KeyURL represents the URL field key
KeyURL = "url"
// KeyToken represents the Token field key
KeyToken = "token"
// KeyOther represents the Other field key
KeyOther = "other"
// TypeText represents standard text input type
TypeText = "text"
// TypePassword represents password input type
TypePassword = "password"
// TypeTextarea represents textarea input type
TypeTextarea = "textarea"
)
// Field represents a form field configuration for a channel.
type Field struct {
Key string `json:"key"` // unique key for the field (e.g. url, token, other)
Label string `json:"label"` // human readable label (e.g. "Webhook 地址")
Type string `json:"type"` // input type: "text" | "password" | "textarea"
Required bool `json:"required"` // whether this field is required
Placeholder string `json:"placeholder"` // input placeholder
Description string `json:"description"` // field explanation/help text
}
// Definition represents the metadata and form schema for a notification channel.
type Definition struct {
Type string `json:"type"` // channel type (e.g., custom, lark, email)
Name string `json:"name"` // display name
Description string `json:"description"` // short description
Fields []Field `json:"fields"` // form fields
}
var (
defMu sync.RWMutex
definitions = make(map[string]Definition)
)
// RegisterChannelDefinition registers a channel definition.
func RegisterChannelDefinition(def Definition) {
defMu.Lock()
defer defMu.Unlock()
definitions[def.Type] = def
}
// ListDefinitions returns all registered channel definitions.
func ListDefinitions() []Definition {
defMu.RLock()
defer defMu.RUnlock()
// We want a stable order: custom, lark, telegram, email
order := []string{channelCustom, channelLark, channelTelegram, channelEmail}
res := make([]Definition, 0, len(definitions))
for _, t := range order {
if d, ok := definitions[t]; ok {
res = append(res, d)
}
}
// Add any others
for t, d := range definitions {
found := false
for _, o := range order {
if o == t {
found = true
break
}
}
if !found {
res = append(res, d)
}
}
return res
}
func init() {
// Register custom webhook channel
RegisterChannelDefinition(Definition{
Type: channelCustom,
Name: "自定义消息通道",
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
Fields: []Field{
{
Key: KeyURL,
Label: "请求地址",
Type: TypeText,
Required: true,
Placeholder: "在此填写完整的请求地址,必须使用 HTTPS 协议",
Description: "接口请求的完整 HTTPS URL,例如 https://api.example.com/webhook",
},
{
Key: KeyOther,
Label: "请求体 (JSON)",
Type: TypeTextarea,
Required: true,
Placeholder: "在此输入请求体,支持模板变量,必须为合法的 JSON 格式",
Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}",
},
},
})
// Register Lark robot channel
RegisterChannelDefinition(Definition{
Type: channelLark,
Name: "飞书群机器人",
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
Fields: []Field{
{
Key: KeyURL,
Label: "Webhook 地址",
Type: TypeText,
Required: true,
Placeholder: "https://open.feishu.cn/open-apis/bot/v2/hook/YOUR_TOKEN",
Description: "从飞书群机器人设置中复制 of Webhook URL",
// Note: using 'of' was in feishu.go, let's keep original wording or fix it
},
{
Key: KeyToken,
Label: "签名校验密钥 (Secret) (可选)",
Type: TypeText,
Required: false,
Placeholder: "可选,若机器人启用了安全设置中的签名校验,请在此输入",
Description: "飞书群机器人安全设置中的签名校验 Key",
},
{
Key: KeyOther,
Label: "自定义卡片 JSON 模版 (可选)",
Type: TypeTextarea,
Required: false,
Placeholder: "可选,留空则默认使用系统内置的精美互动卡片",
Description: "若填写,必须是合法的飞书卡片 JSON 格式",
},
},
})
// Register Telegram channel
RegisterChannelDefinition(Definition{
Type: channelTelegram,
Name: "Telegram 机器人",
Description: "配置 Telegram 机器人推送消息。",
Fields: []Field{
{
Key: KeyURL,
Label: "API 基础地址 (可选)",
Type: TypeText,
Required: false,
Placeholder: "https://api.telegram.org",
Description: "接口请求的 HTTPS 基础地址,留空默认为 https://api.telegram.org",
},
{
Key: KeyToken,
Label: "机器人 Token (Bot Token)",
Type: TypePassword,
Required: true,
Placeholder: "在此输入 Telegram 机器人的 Bot Token",
Description: "通过 BotFather 申请到的机器人 Access Token",
},
{
Key: KeyOther,
Label: "默认会话 ID (Chat ID) (可选)",
Type: TypeText,
Required: false,
Placeholder: "例如 -100123456789 或 @channel_name",
Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID",
},
},
})
// Register Email channel
RegisterChannelDefinition(Definition{
Type: channelEmail,
Name: "邮件推送通道",
Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。",
Fields: []Field{},
})
}
@@ -0,0 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
const (
channelCustom = "custom"
channelEmail = "email"
channelLark = "lark"
channelTelegram = "telegram"
defaultLevelInfo = "INFO"
keyTitle = "title"
keyContent = "content"
keyLevel = "level"
)
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package custom_events defines custom push notification events.
package custom_events
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
"github.com/Rain-kl/Wavelet/internal/listener"
)
// AdminLogin is the metadata definition for the admin login event.
var AdminLogin = push.EventMetadata{
Key: "admin_login",
Name: "管理员登录",
DefaultTemplate: push.NotificationMessage{
Title: "管理员登录提醒",
Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。",
Level: "INFO",
},
Description: "当管理员成功登录系统时触发此通知",
}
func handleAdminLogin(ctx context.Context, event listener.AdminLoggedIn) {
if event.User == nil {
return
}
body := map[string]any{
"user": event.User,
"ip": event.IP,
"time": time.Now().Format("2006-01-02 15:04:05"),
}
push.DefaultTrigger.Trigger(ctx, AdminLogin, body)
}
@@ -0,0 +1,175 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package custom_events
import (
"context"
"encoding/json"
"sync"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/hibiken/asynq"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
var registerOnce sync.Once
func ensureRegistered() {
registerOnce.Do(Register)
}
func setupAdminLoginIntegrationTest(t *testing.T) (*gorm.DB, func()) {
t.Helper()
dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t)
err := dbConn.AutoMigrate(
&model.PushEvent{},
&model.PushHistory{},
&model.PushChannel{},
)
require.NoError(t, err)
sysUser := &model.User{
ID: 999,
Username: "system",
Nickname: "系统",
Password: "*",
IsActive: true,
}
require.NoError(t, dbConn.Create(sysUser).Error)
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{Addr: mr.Addr()})
task.RegisterHandler(push.SendNotificationTask, &push.PushHandler{})
task.RegisterTaskMeta(push.SendNotificationMeta)
ensureRegistered()
require.NoError(t, push.SyncEvents(context.Background()))
return dbConn, func() {
cleanup()
if task.AsynqClient != nil {
task.AsynqClient.Close()
task.AsynqClient = nil
}
}
}
func seedMockPushChannel(t *testing.T, dbConn *gorm.DB) *model.PushChannel {
t.Helper()
channel := &model.PushChannel{
Name: "mock_channel",
Type: "custom",
URL: "https://webhook.site/admin-login",
Other: `{"text": "$content"}`,
Enabled: true,
}
require.NoError(t, dbConn.Create(channel).Error)
return channel
}
func enableAdminLoginEvent(t *testing.T, dbConn *gorm.DB, channelName string, targets []string) {
t.Helper()
var event model.PushEvent
require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error)
event.Enabled = true
event.Channels = []string{channelName}
event.Targets = targets
require.NoError(t, dbConn.Save(&event).Error)
}
func waitForAsyncTrigger(t *testing.T) {
t.Helper()
time.Sleep(100 * time.Millisecond)
}
func countPushTasks(t *testing.T, dbConn *gorm.DB) int64 {
t.Helper()
var count int64
require.NoError(t, dbConn.Model(&model.TaskExecution{}).
Where("task_type = ?", push.SendNotificationTask).
Count(&count).Error)
return count
}
func TestAdminLoginPushIntegration(t *testing.T) {
dbConn, cleanup := setupAdminLoginIntegrationTest(t)
defer cleanup()
channel := seedMockPushChannel(t, dbConn)
defer dbConn.Delete(channel)
enableAdminLoginEvent(t, dbConn, channel.Name, []string{"ops_team"})
adminUser := &model.User{
ID: 1001,
Username: "super_admin",
IsAdmin: true,
IsActive: true,
}
require.NoError(t, dbConn.Create(adminUser).Error)
t.Run("admin login emits push task with user and ip", func(t *testing.T) {
dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
listener.EmitAdminLoggedIn(context.Background(), adminUser, "203.0.113.42")
waitForAsyncTrigger(t)
var execution model.TaskExecution
require.NoError(t, dbConn.Where("task_type = ?", push.SendNotificationTask).First(&execution).Error)
var payload push.SendPayload
require.NoError(t, json.Unmarshal([]byte(execution.Payload), &payload))
assert.Equal(t, AdminLogin.Key, payload.EventKey)
assert.Equal(t, "ops_team", payload.Target)
assert.Equal(t, "管理员登录提醒", payload.Body.Title)
assert.Contains(t, payload.Body.Content, "super_admin")
assert.Contains(t, payload.Body.Content, "203.0.113.42")
})
t.Run("non-admin login does not trigger push", func(t *testing.T) {
dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
nonAdmin := &model.User{
ID: 2002,
Username: "regular_user",
IsAdmin: false,
IsActive: true,
}
require.NoError(t, dbConn.Create(nonAdmin).Error)
listener.EmitAdminLoggedIn(context.Background(), nonAdmin, "198.51.100.1")
waitForAsyncTrigger(t)
assert.Equal(t, int64(0), countPushTasks(t, dbConn))
})
t.Run("disabled admin login event does not enqueue push", func(t *testing.T) {
dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
var event model.PushEvent
require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error)
event.Enabled = false
require.NoError(t, dbConn.Save(&event).Error)
listener.EmitAdminLoggedIn(context.Background(), adminUser, "10.0.0.1")
waitForAsyncTrigger(t)
assert.Equal(t, int64(0), countPushTasks(t, dbConn))
})
}
@@ -0,0 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package custom_events
import (
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
"github.com/Rain-kl/Wavelet/internal/listener"
)
// Register wires push notification handlers for domain events and registers
// built-in event metadata. Must be called once during application bootstrap
// before push.SyncEvents.
func Register() {
push.RegisterBuiltInEvent(AdminLogin)
listener.OnAdminLoggedIn(handleAdminLogin)
}
@@ -0,0 +1,343 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package push defines push notification HTTP routes, background tasks, and events.
package push
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"gorm.io/gorm"
)
// NotificationMessage represents the structured notification message payload.
type NotificationMessage struct {
Title string `json:"title"`
Content string `json:"content"`
Level string `json:"level"`
Ext map[string]any `json:"ext,omitempty"`
}
// Flatten converts the structured NotificationMessage back to a flat map (original json structure).
func (m NotificationMessage) Flatten() map[string]any {
res := map[string]any{
keyTitle: m.Title,
keyContent: m.Content,
keyLevel: m.Level,
}
for k, v := range m.Ext {
res[k] = v
}
return res
}
// EventMetadata represents the metadata of a push notification event.
type EventMetadata struct {
Key string `json:"key"`
Name string `json:"name"`
DefaultTemplate NotificationMessage `json:"default_template"`
Description string `json:"description"`
}
// SendPayload 异步投递推送载荷 (供 task/Worker 使用)
type SendPayload struct {
EventKey string `json:"event_key"`
Config pkgpush.Config `json:"config"`
Target string `json:"target"`
Body NotificationMessage `json:"body"`
Template string `json:"template"`
}
// BuiltInEvents lists all built-in events defined in custom_events.
var BuiltInEvents []EventMetadata
// RegisterBuiltInEvent registers a built-in event definition.
func RegisterBuiltInEvent(meta EventMetadata) {
BuiltInEvents = append(BuiltInEvents, meta)
}
// EventTrigger represents the unified event trigger class.
type EventTrigger struct{}
// DefaultTrigger is the singleton instance of EventTrigger.
var DefaultTrigger = &EventTrigger{}
// Trigger receives event metadata and processes the event notification dispatch asynchronously.
//
//nolint:contextcheck
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
asyncCtx := context.WithoutCancel(ctx)
go func() {
if body == nil {
body = make(map[string]any)
}
if _, hasUser := body["user"]; !hasUser || body["user"] == nil {
body["user"] = getSystemUser(asyncCtx)
}
eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return
}
logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err)
return
}
event := *eventPtr
if len(event.Channels) == 0 {
return
}
flatBody := getFlatBody(body)
msg, _ := t.buildMessage(&event, meta, flatBody, body)
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
}()
}
func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
var msg NotificationMessage
renderedTemplate := ""
templateSource := event.Template
if templateSource != "" {
var err error
msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody)
if err != nil {
msg.Title = event.Name
msg.Content = renderedTemplate
msg.Level = defaultLevelInfo
}
} else {
msg = t.parseDefaultTemplate(meta, flatBody)
}
if msg.Ext == nil {
msg.Ext = make(map[string]any)
}
for k, v := range body {
if k == keyTitle || k == keyContent || k == keyLevel {
continue
}
if _, exists := msg.Ext[k]; !exists {
msg.Ext[k] = v
}
}
return msg, renderedTemplate
}
func (t *EventTrigger) parseCustomTemplate(event *model.PushEvent, templateSource string, flatBody map[string]any) (NotificationMessage, string, error) {
var msg NotificationMessage
renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody)
var tMap map[string]any
if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil {
return msg, renderedTemplate, err
}
if title, ok := tMap[keyTitle].(string); ok && title != "" {
msg.Title = title
} else {
msg.Title = event.Name
}
delete(tMap, keyTitle)
if content, ok := tMap[keyContent].(string); ok && content != "" {
msg.Content = content
} else {
msg.Content = renderedTemplate
}
delete(tMap, keyContent)
if level, ok := tMap[keyLevel].(string); ok && level != "" {
msg.Level = level
} else {
msg.Level = defaultLevelInfo
}
delete(tMap, keyLevel)
msg.Ext = tMap
return msg, renderedTemplate, nil
}
func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[string]any) NotificationMessage {
var msg NotificationMessage
msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody)
msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody)
msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody)
if meta.DefaultTemplate.Ext != nil {
msg.Ext = make(map[string]any)
for k, v := range meta.DefaultTemplate.Ext {
if strVal, ok := v.(string); ok {
msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody)
} else {
msg.Ext[k] = v
}
}
}
return msg
}
func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, msg NotificationMessage, flatBody map[string]any) {
for _, channelName := range event.Channels {
customChannel, err := repository.GetActivePushChannelByName(ctx, channelName)
if err == nil {
t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody)
continue
}
logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err)
}
}
func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, channel *model.PushChannel, msg NotificationMessage, flatBody map[string]any) {
if len(event.Targets) == 0 {
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg)
return
}
for _, target := range event.Targets {
resolvedTarget := resolveTarget(ctx, target, flatBody, channel.Name)
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg)
}
}
func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta EventMetadata, channel *model.PushChannel, target string, msg NotificationMessage) {
var config pkgpush.Config
var renderedTemplate string
switch channel.Type {
case channelLark:
config = pkgpush.Config{Channel: channelLark, URL: channel.URL, Secret: channel.Token}
renderedTemplate = channel.Other
case channelEmail:
url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other)
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{Channel: channelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other}
default:
config = pkgpush.Config{Channel: channelCustom, URL: channel.URL}
customPushReq := CustomPushRequest{
Title: msg.Title,
Content: msg.Content,
Description: meta.Description,
To: target,
}
if urlVal, ok := msg.Ext["url"].(string); ok {
customPushReq.URL = urlVal
}
renderedTemplate = renderCustomPayload(channel.Other, customPushReq)
}
payload := SendPayload{
EventKey: meta.Key,
Config: config,
Target: target,
Body: msg,
Template: renderedTemplate,
}
if err := enqueuePushTask(ctx, payload); err != nil {
logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err)
}
}
func enqueuePushTask(ctx context.Context, payload SendPayload) error {
payloadBytes, err := json.Marshal(payload)
if err != nil {
return err
}
_, err = task.DispatchTask(ctx, "send_notification", payloadBytes, "system")
return err
}
func getFlatBody(body map[string]any) map[string]any {
jsonBytes, err := json.Marshal(body)
if err != nil {
return body
}
var jsonMap map[string]any
if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil {
return body
}
flatResult := make(map[string]any)
flattenMap("", jsonMap, flatResult)
return flatResult
}
func flattenMap(prefix string, m map[string]any, result map[string]any) {
for k, v := range m {
key := k
if prefix != "" {
key = prefix + "." + k
}
if nestedMap, ok := v.(map[string]any); ok {
flattenMap(key, nestedMap, result)
} else {
result[key] = v
}
}
}
func resolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string {
target = strings.TrimSpace(target)
if target == "" {
return ""
}
resolved := resolveDynamicKeyword(target, flatBody)
if strings.Contains(resolved, "@") {
return resolved
}
if val, matched := resolveSystemTarget(ctx, resolved, channel); matched {
return val
}
user, found := resolveTargetUser(ctx, resolved, channel)
if !found {
return resolved
}
if channel == channelEmail && user.Email != "" {
return user.Email
}
if channel != channelEmail && user.Username != "" {
return user.Username
}
return resolved
}
func resolveDynamicKeyword(target string, flatBody map[string]any) string {
switch target {
case "user.id", "id":
if val, ok := flatBody["user.id"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["id"]; ok {
return fmt.Sprintf("%v", val)
}
case "user.username", "username":
if val, ok := flatBody["user.username"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["username"]; ok {
return fmt.Sprintf("%v", val)
}
case "user.email", channelEmail:
if val, ok := flatBody["user.email"]; ok {
return fmt.Sprintf("%v", val)
}
if val, ok := flatBody["email"]; ok {
return fmt.Sprintf("%v", val)
}
}
return target
}
@@ -0,0 +1,413 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"encoding/json"
"errors"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/task"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"gorm.io/gorm"
)
type smtpConfig struct {
Host string
Port string
Username string
Password string
}
func loadSMTPConfig(ctx context.Context) smtpConfig {
host, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost)
port, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort)
user, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername)
pass, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword)
return smtpConfig{
Host: host.Value,
Port: port.Value,
Username: user.Value,
Password: pass.Value,
}
}
func syncBuiltInEvents(ctx context.Context) error {
for _, meta := range BuiltInEvents {
_, err := repository.GetPushEventByKey(ctx, meta.Key)
if errors.Is(err, gorm.ErrRecordNotFound) {
var defaultTemplateStr string
if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil {
defaultTemplateStr = string(defaultTemplateBytes)
}
event := model.PushEvent{
EventKey: meta.Key,
Name: meta.Name,
Channels: []string{},
Targets: []string{},
Template: defaultTemplateStr,
Enabled: false,
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
return err
}
} else if err != nil {
return err
}
}
return nil
}
func listPushEvents(ctx context.Context) ([]model.PushEvent, error) {
return repository.ListPushEvents(ctx)
}
func createPushEvent(ctx context.Context, req CreateEventRequest) (model.PushEvent, error) {
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
if err != nil {
return model.PushEvent{}, err
}
count, err := repository.CountPushEventsByKey(ctx, eventKey)
if err != nil {
return model.PushEvent{}, err
}
if count > 0 {
return model.PushEvent{}, errors.New("this notification event is already configured")
}
templateStr := strings.TrimSpace(req.Template)
if templateStr == "" {
templateStr = string(defaultTemplateBytes)
} else {
var tempMap map[string]any
if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil {
return model.PushEvent{}, errors.New("custom template is not a valid JSON format")
}
}
channels := req.Channels
if channels == nil {
channels = []string{}
}
targets := req.Targets
if targets == nil {
targets = []string{}
}
event := model.PushEvent{
EventKey: eventKey,
Name: eventName,
TaskType: req.TaskType,
Channels: channels,
Targets: targets,
Template: templateStr,
Enabled: req.Enabled,
}
if err := event.Validate(); err != nil {
return model.PushEvent{}, err
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
return model.PushEvent{}, err
}
return event, nil
}
func deletePushEvent(ctx context.Context, id uint64) error {
event, err := repository.GetPushEventByID(ctx, id)
if err != nil {
return err
}
return repository.DeletePushEvent(ctx, &event)
}
func updatePushEvent(ctx context.Context, id uint64, req UpdateEventRequest) error {
event, err := repository.GetPushEventByID(ctx, id)
if err != nil {
return err
}
event.Channels = req.Channels
event.Targets = req.Targets
event.Template = req.Template
event.Enabled = req.Enabled
if err := event.Validate(); err != nil {
return err
}
return repository.SavePushEvent(ctx, &event)
}
func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
event, err := repository.GetPushEventByID(ctx, id)
if err != nil {
return false, err
}
enabled := !event.Enabled
if enabled && len(event.Channels) == 0 {
return false, errors.New("cannot enable event without any push channels configured")
}
if err := repository.UpdatePushEventEnabled(ctx, &event, enabled); err != nil {
return false, err
}
return enabled, nil
}
func listPushHistories(ctx context.Context, filter repository.PushHistoryListFilter) (int64, []model.PushHistory, error) {
return repository.ListPushHistories(ctx, filter)
}
func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) {
if cfg.Channel != channelEmail || (cfg.URL != "" && cfg.Key != "") {
return
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return
}
port := smtp.Port
if port == "" {
port = "587"
}
cfg.URL = smtp.Host + ":" + port
cfg.Key = smtp.Username
cfg.Secret = smtp.Password
}
func listPushChannels(ctx context.Context) ([]model.PushChannel, error) {
return repository.ListPushChannels(ctx)
}
func createPushChannel(ctx context.Context, req CreateChannelRequest) (model.PushChannel, error) {
count, err := repository.CountPushChannelsByName(ctx, req.Name)
if err != nil {
return model.PushChannel{}, err
}
if count > 0 {
return model.PushChannel{}, errors.New("channel name already exists")
}
channel := model.PushChannel{
Name: req.Name,
Description: req.Description,
Type: req.Type,
Token: req.Token,
URL: req.URL,
Other: req.Other,
Enabled: req.Enabled,
}
if err := channel.Validate(); err != nil {
return model.PushChannel{}, err
}
if err := repository.CreatePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
}
return channel, nil
}
func updatePushChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (model.PushChannel, error) {
channel, err := repository.GetPushChannelByID(ctx, id)
if err != nil {
return model.PushChannel{}, err
}
channel.Description = req.Description
channel.Type = req.Type
channel.Token = req.Token
channel.URL = req.URL
channel.Other = req.Other
channel.Enabled = req.Enabled
if err := channel.Validate(); err != nil {
return model.PushChannel{}, err
}
if err := repository.SavePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
}
return channel, nil
}
func deletePushChannel(ctx context.Context, id uint64) error {
channel, err := repository.GetPushChannelByID(ctx, id)
if err != nil {
return err
}
return repository.DeletePushChannel(ctx, &channel)
}
func loadChannelForTest(ctx context.Context, req TestChannelRequest) (string, string, string, string, error) {
if req.Name != "" {
channel, err := repository.GetPushChannelByName(ctx, req.Name)
if err != nil {
return "", "", "", "", errors.New("channel not found")
}
return channel.URL, channel.Token, channel.Other, channel.Type, nil
}
return req.URL, req.Token, req.Other, req.Type, nil
}
func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) {
return repository.ListActivePushEventsByTaskType(ctx, taskType)
}
func loadUserFromPayload(ctx context.Context, data map[string]any) any {
if u, exists := data["user"]; exists && u != nil {
return u
}
if userID, ok := extractUserID(data); ok && userID > 0 {
if user, err := repository.GetUserByID(ctx, userID); err == nil {
return &user
}
}
if username := extractUsername(data); username != "" {
if user, err := repository.GetUserByUsername(ctx, username); err == nil {
return &user
}
}
return nil
}
func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg string) error {
title := req.Body.Title
content := req.Body.Content
level := req.Body.Level
if title == "" {
title = "系统通知"
}
if level == "" {
level = defaultLevelInfo
}
target := req.Target
if target == "" {
if req.Config.URL != "" {
target = req.Config.URL
const maxTargetLen = 50
const truncatedLen = 47
if len(target) > maxTargetLen {
target = target[:truncatedLen] + "..."
}
} else {
target = "default"
}
}
history := model.PushHistory{
EventKey: req.EventKey,
Channel: req.Config.Channel,
Target: target,
Title: title,
Content: content,
Level: level,
Status: status,
ErrorMsg: errMsg,
}
return repository.CreatePushHistory(ctx, &history)
}
func resolveTargetUser(ctx context.Context, resolved string, _ string) (model.User, bool) {
found := false
var user model.User
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if u, err := repository.GetUserByID(ctx, id); err == nil {
user = u
found = true
}
}
if !found {
if u, err := repository.GetUserByUsername(ctx, resolved); err == nil {
user = u
found = true
}
}
return user, found
}
func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) {
if resolved != "系统" && resolved != "system" && resolved != "0" {
return "", false
}
adminUser, err := repository.GetFirstAdminUser(ctx)
if err != nil {
return resolved, true
}
if channel == channelEmail && adminUser.Email != "" {
return adminUser.Email, true
}
if channel != channelEmail && adminUser.Username != "" {
return adminUser.Username, true
}
return resolved, true
}
func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) {
if url != "" && token != "" {
return url, token, other
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return url, token, other
}
port := smtp.Port
if port == "" {
port = "587"
}
if url == "" {
url = smtp.Host + ":" + port
}
if token == "" {
token = smtp.Username
}
if other == "" {
other = smtp.Password
}
return url, token, other
}
func getSystemUser(ctx context.Context) *model.User {
user := repository.GetSystemUser(ctx)
return &user
}
func getEventInfo(req CreateEventRequest) (string, string, []byte, error) {
if req.TaskType != "" {
meta := task.GetTaskMetaByAsynqTask(req.TaskType)
if meta == nil {
return "", "", nil, errors.New("unsupported task type")
}
eventKey := "task_completed:" + req.TaskType
eventName := "任务完成: " + meta.Name
defaultTemplate := NotificationMessage{
Title: "任务完成: " + meta.Name,
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
Level: defaultLevelInfo,
}
defaultTemplateBytes, err := json.Marshal(defaultTemplate)
if err != nil {
return "", "", nil, err
}
return eventKey, eventName, defaultTemplateBytes, nil
}
if req.EventKey == "" {
return "", "", nil, errors.New("either event_key or task_type must be provided")
}
meta, found := findBuiltInEvent(req.EventKey)
if !found {
return "", "", nil, errors.New("unsupported built-in event key")
}
defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate)
if err != nil {
return "", "", nil, err
}
return req.EventKey, meta.Name, defaultTemplateBytes, nil
}
@@ -0,0 +1,759 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"sync"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/alicebob/miniredis/v2"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
var adminLoginEvent = EventMetadata{
Key: "admin_login",
Name: "管理员登录",
DefaultTemplate: NotificationMessage{
Title: "管理员登录提醒",
Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。",
Level: "INFO",
},
Description: "当管理员成功登录系统时触发此通知",
}
func init() {
RegisterBuiltInEvent(adminLoginEvent)
}
// mockPusher mock implementation of pkgpush.Pusher
type mockPusher struct {
mu sync.Mutex
sentBody map[string]any
sentTgt string
}
func (m *mockPusher) Send(ctx context.Context, cfg pkgpush.Config, target string, body map[string]any, template string, ext map[string]any) error {
m.mu.Lock()
defer m.mu.Unlock()
m.sentBody = body
m.sentTgt = target
return nil
}
func (m *mockPusher) ValidateConfig(cfg pkgpush.Config) error {
return nil
}
func setupPushTest(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) {
dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t)
// AutoMigrate push tables in SQLite test environment
err := dbConn.AutoMigrate(&model.PushEvent{}, &model.PushHistory{}, &model.User{}, &model.PushChannel{}, &model.SystemConfig{})
require.NoError(t, err)
// 写入数据库系统默认用户 Seed 记录
sysUser := &model.User{
ID: 999,
Username: "system",
Nickname: "系统",
Password: "*",
IsActive: true,
}
err = dbConn.Create(sysUser).Error
require.NoError(t, err)
// Initialize AsynqClient pointing to miniredis
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
Addr: mr.Addr(),
})
// Register the task handler and metadata
task.RegisterHandler(SendNotificationTask, &PushHandler{})
task.RegisterTaskMeta(SendNotificationMeta)
return dbConn, mr, func() {
cleanup()
if task.AsynqClient != nil {
task.AsynqClient.Close()
task.AsynqClient = nil
}
}
}
func setupTestRouter(authUser *model.User) *gin.Engine {
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin/push")
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
oauth.SetToContext(c, "user_obj", authUser)
}
c.Next()
})
adminGroup.GET("/events", ListEvents)
adminGroup.GET("/events/builtin", ListBuiltInEvents)
adminGroup.POST("/events", CreateEvent)
adminGroup.PUT("/events/:id", UpdateEvent)
adminGroup.DELETE("/events/:id", DeleteEvent)
adminGroup.POST("/events/:id/toggle", ToggleEvent)
adminGroup.GET("/histories", ListHistories)
adminGroup.POST("/test", TestPush)
return r
}
func TestSyncEvents(t *testing.T) {
dbConn, _, cleanup := setupPushTest(t)
defer cleanup()
// 1. SyncEvents first time
err := SyncEvents(context.Background())
require.NoError(t, err)
// Verify event exists in DB
var event model.PushEvent
err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error
require.NoError(t, err)
assert.Equal(t, "管理员登录", event.Name)
assert.False(t, event.Enabled)
// Verify DefaultTemplate matches GORM template field
var defaultMsg NotificationMessage
err = json.Unmarshal([]byte(event.Template), &defaultMsg)
require.NoError(t, err)
assert.Equal(t, adminLoginEvent.DefaultTemplate.Title, defaultMsg.Title)
assert.Equal(t, adminLoginEvent.DefaultTemplate.Content, defaultMsg.Content)
}
func TestEventTrigger(t *testing.T) {
dbConn, _, cleanup := setupPushTest(t)
defer cleanup()
// SyncEvents
err := SyncEvents(context.Background())
require.NoError(t, err)
t.Run("trigger disabled event silently ignored", func(t *testing.T) {
body := map[string]any{
"user": map[string]any{"username": "test_admin"},
"ip": "127.0.0.1",
}
DefaultTrigger.Trigger(context.Background(), adminLoginEvent, body)
// Sleep briefly since Trigger runs in goroutine
time.Sleep(50 * time.Millisecond)
// Verify no tasks enqueued in TaskExecution GORM table
var count int64
dbConn.Model(&model.TaskExecution{}).Count(&count)
assert.Equal(t, int64(0), count)
})
t.Run("trigger enabled event enqueues task", func(t *testing.T) {
// Create an enabled custom channel in GORM
customChan := &model.PushChannel{
Name: "mock_channel",
Type: "custom",
URL: "https://webhook.site/trigger",
Other: `{"text": "$content"}`,
Enabled: true,
}
err = dbConn.Create(customChan).Error
require.NoError(t, err)
defer dbConn.Delete(customChan)
// Enable the push event in DB using struct to trigger JSON serializer
var event model.PushEvent
err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error
require.NoError(t, err)
event.Enabled = true
event.Channels = []string{"mock_channel"}
event.Targets = []string{"admin_user"}
err = dbConn.Save(&event).Error
require.NoError(t, err)
// Trigger
body := map[string]any{
"user": map[string]any{
"username": "super_admin",
},
"ip": "1.1.1.1",
"time": "2026-06-14 18:00:00",
}
DefaultTrigger.Trigger(context.Background(), adminLoginEvent, body)
// Wait for goroutine execution
time.Sleep(50 * time.Millisecond)
// Verify TaskExecution enqueued record
var execution model.TaskExecution
err = dbConn.Where("task_type = ?", SendNotificationTask).First(&execution).Error
require.NoError(t, err)
// Verify enqueued payload structure
var payload SendPayload
err = json.Unmarshal([]byte(execution.Payload), &payload)
require.NoError(t, err)
assert.Equal(t, "admin_login", payload.EventKey)
assert.Equal(t, "custom", payload.Config.Channel)
assert.Equal(t, "https://webhook.site/trigger", payload.Config.URL)
assert.Equal(t, "admin_user", payload.Target)
assert.Equal(t, "管理员登录提醒", payload.Body.Title)
assert.Contains(t, payload.Body.Content, "super_admin")
assert.Contains(t, payload.Body.Content, "1.1.1.1")
})
t.Run("trigger without user injects virtual system user", func(t *testing.T) {
// Create an enabled custom channel in GORM
customChan := &model.PushChannel{
Name: "mock_channel",
Type: "custom",
URL: "https://webhook.site/trigger",
Other: `{"text": "$content"}`,
Enabled: true,
}
err = dbConn.Create(customChan).Error
require.NoError(t, err)
defer dbConn.Delete(customChan)
// Enable the push event in DB
var event model.PushEvent
err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error
require.NoError(t, err)
// 清理旧任务执行记录
dbConn.Where("task_type = ?", SendNotificationTask).Delete(&model.TaskExecution{})
event.Enabled = true
event.Channels = []string{"mock_channel"}
event.Targets = []string{"user.username"} // 动态目标
err = dbConn.Save(&event).Error
require.NoError(t, err)
// Trigger with empty body (simulates cron scheduler triggering)
DefaultTrigger.Trigger(context.Background(), adminLoginEvent, nil)
// Wait for goroutine execution
time.Sleep(50 * time.Millisecond)
// Verify TaskExecution enqueued record
var execution model.TaskExecution
err = dbConn.Where("task_type = ?", SendNotificationTask).First(&execution).Error
require.NoError(t, err)
var payload SendPayload
err = json.Unmarshal([]byte(execution.Payload), &payload)
require.NoError(t, err)
// 检查 payload 是否将 target (user.username) 成功替换为 "system"
assert.Equal(t, "system", payload.Target)
// 检查 payload 中的 Content,应当被替换为 "system" 变量
assert.Contains(t, payload.Body.Content, "system")
})
}
func TestPushHandler(t *testing.T) {
dbConn, _, cleanup := setupPushTest(t)
defer cleanup()
mPusher := &mockPusher{}
pkgpush.Register("mock_channel", mPusher)
handler := &PushHandler{}
payload := SendPayload{
EventKey: "admin_login",
Config: pkgpush.Config{
Channel: "mock_channel",
URL: "http://mock-url",
},
Target: "admin_user",
Body: NotificationMessage{
Title: "Structured Alert",
Content: "Hello World",
Level: "WARNING",
Ext: map[string]any{"extra_val": 42},
},
}
payloadBytes, err := json.Marshal(payload)
require.NoError(t, err)
t.Run("validate payload", func(t *testing.T) {
validated, valErr := handler.ValidatePayload(payloadBytes)
require.NoError(t, valErr)
assert.NotEmpty(t, validated)
})
t.Run("execute task successfully", func(t *testing.T) {
res, execErr := handler.Execute(context.Background(), payloadBytes)
require.NoError(t, execErr)
assert.Contains(t, res.Message, "推送成功")
// Verify mock pusher received flattened variables
mPusher.mu.Lock()
assert.Equal(t, "admin_user", mPusher.sentTgt)
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
mPusher.mu.Unlock()
// Verify PushHistory recorded
var history model.PushHistory
err = dbConn.First(&history).Error
require.NoError(t, err)
assert.Equal(t, "admin_login", history.EventKey)
assert.Equal(t, "mock_channel", history.Channel)
assert.Equal(t, "success", history.Status)
assert.Equal(t, "Structured Alert", history.Title)
})
}
func TestPushRouters(t *testing.T) {
dbConn, _, cleanup := setupPushTest(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
r := setupTestRouter(adminUser)
// Sync events to populate db
err := SyncEvents(context.Background())
require.NoError(t, err)
t.Run("list events", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/push/events", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
err = json.Unmarshal(w.Body.Bytes(), &resp)
require.NoError(t, err)
dataBytes, _ := json.Marshal(resp.Data)
var events []model.PushEvent
err = json.Unmarshal(dataBytes, &events)
require.NoError(t, err)
assert.Len(t, events, 1)
assert.Equal(t, "admin_login", events[0].EventKey)
})
t.Run("toggle event status", func(t *testing.T) {
var event model.PushEvent
dbConn.First(&event)
// 1. 未配置任何渠道时开启,应该被拒绝
req, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
// 2. 为该事件关联渠道后,再切换开启,应当成功
event.Channels = []string{"email"}
dbConn.Save(&event)
req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil)
w2 := httptest.NewRecorder()
r.ServeHTTP(w2, req2)
assert.Equal(t, http.StatusOK, w2.Code)
var updated model.PushEvent
dbConn.First(&updated)
assert.True(t, updated.Enabled)
})
t.Run("update event", func(t *testing.T) {
var event model.PushEvent
dbConn.First(&event)
updateReq := UpdateEventRequest{
Channels: []string{"email"},
Targets: []string{"user@test.com"},
Template: `{"title": "Custom Login Alert", "content": "Alert", "level": "WARNING"}`,
Enabled: true,
}
bodyBytes, _ := json.Marshal(updateReq)
req, _ := http.NewRequest("PUT", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10), bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var updated model.PushEvent
dbConn.First(&updated)
assert.Equal(t, []string{"email"}, updated.Channels)
assert.Equal(t, []string{"user@test.com"}, updated.Targets)
assert.Contains(t, updated.Template, "Custom Login Alert")
})
t.Run("list push histories", func(t *testing.T) {
// Populate history record
hist := model.PushHistory{
EventKey: "admin_login",
Channel: "email",
Target: "user@test.com",
Title: "Custom Login Alert",
Content: "Alert",
Level: "WARNING",
Status: "success",
}
dbConn.Create(&hist)
req, _ := http.NewRequest("GET", "/api/v1/admin/push/histories?page=1&page_size=10", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataMap, ok := resp.Data.(map[string]any)
assert.True(t, ok)
assert.Equal(t, float64(1), dataMap["total"])
})
t.Run("test push endpoint", func(t *testing.T) {
mPusher := &mockPusher{}
pkgpush.Register("test_channel", mPusher)
testReq := TestPushRequest{
Config: pkgpush.Config{
Channel: "test_channel",
URL: "http://test-url",
},
Target: "test_target",
}
bodyBytes, _ := json.Marshal(testReq)
req, _ := http.NewRequest("POST", "/api/v1/admin/push/test", bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
})
t.Run("list built-in events", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/push/events/builtin", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
err := json.Unmarshal(w.Body.Bytes(), &resp)
require.NoError(t, err)
builtins, ok := resp.Data.([]any)
assert.True(t, ok)
assert.NotEmpty(t, builtins)
})
t.Run("create and delete push event", func(t *testing.T) {
// Clean up any existing admin_login event first
dbConn.Where("event_key = ?", "admin_login").Delete(&model.PushEvent{})
// 1. Create event
createReq := CreateEventRequest{
EventKey: "admin_login",
Channels: []string{"email"},
Targets: []string{"admin@test.com"},
Enabled: true,
}
bodyBytes, _ := json.Marshal(createReq)
req, _ := http.NewRequest("POST", "/api/v1/admin/push/events", bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
// Verify created in DB
var event model.PushEvent
err := dbConn.Where("event_key = ?", "admin_login").First(&event).Error
require.NoError(t, err)
assert.Equal(t, "admin_login", event.EventKey)
assert.Equal(t, "管理员登录", event.Name)
assert.True(t, event.Enabled)
// 2. Try creating again (should fail)
w2 := httptest.NewRecorder()
req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events", bytes.NewBuffer(bodyBytes))
req2.Header.Set("Content-Type", "application/json")
r.ServeHTTP(w2, req2)
assert.Equal(t, http.StatusBadRequest, w2.Code)
// 3. Delete event
w3 := httptest.NewRecorder()
req3, _ := http.NewRequest("DELETE", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10), nil)
r.ServeHTTP(w3, req3)
assert.Equal(t, http.StatusOK, w3.Code)
// Verify deleted from DB
var count int64
dbConn.Model(&model.PushEvent{}).Where("event_key = ?", "admin_login").Count(&count)
assert.Equal(t, int64(0), count)
})
}
func TestResolveTarget(t *testing.T) {
dbConn, _, cleanup := setupPushTest(t)
defer cleanup()
// 1. 创建测试用户与管理员用户
testUser := &model.User{
ID: 9999,
Username: "target_user",
Email: "target@test.com",
IsAdmin: false,
}
err := dbConn.Create(testUser).Error
require.NoError(t, err)
adminUser := &model.User{
ID: 8888,
Username: "admin_user",
Email: "admin@test.com",
IsAdmin: true,
}
err = dbConn.Create(adminUser).Error
require.NoError(t, err)
flatBody := map[string]any{
"user.id": float64(9999), // JSON 反序列化后一般是 float64
"user.username": "target_user",
"user.email": "target@test.com",
}
ctx := context.Background()
t.Run("dynamic user.id resolved and converted for email channel", func(t *testing.T) {
res := resolveTarget(ctx, "user.id", flatBody, "email")
assert.Equal(t, "target@test.com", res)
})
t.Run("dynamic user.username resolved and converted for email channel", func(t *testing.T) {
res := resolveTarget(ctx, "user.username", flatBody, "email")
assert.Equal(t, "target@test.com", res)
})
t.Run("dynamic user.email resolved directly for email channel", func(t *testing.T) {
res := resolveTarget(ctx, "user.email", flatBody, "email")
assert.Equal(t, "target@test.com", res)
})
t.Run("fixed user id resolved and converted for email channel", func(t *testing.T) {
res := resolveTarget(ctx, "9999", flatBody, "email")
assert.Equal(t, "target@test.com", res)
})
t.Run("fixed username resolved and converted for email channel", func(t *testing.T) {
res := resolveTarget(ctx, "target_user", flatBody, "email")
assert.Equal(t, "target@test.com", res)
})
t.Run("fixed email address resolved directly for email channel", func(t *testing.T) {
res := resolveTarget(ctx, "fixed@example.com", flatBody, "email")
assert.Equal(t, "fixed@example.com", res)
})
t.Run("fixed username resolved for non-email channel", func(t *testing.T) {
res := resolveTarget(ctx, "target_user", flatBody, "lark")
assert.Equal(t, "target_user", res)
})
t.Run("non-exist user resolved as fallback", func(t *testing.T) {
res := resolveTarget(ctx, "non_exist_user", flatBody, "email")
assert.Equal(t, "non_exist_user", res)
})
t.Run("system target resolves to admin email for email channel", func(t *testing.T) {
res := resolveTarget(ctx, "系统", flatBody, "email")
assert.Equal(t, "admin@test.com", res)
res2 := resolveTarget(ctx, "system", flatBody, "email")
assert.Equal(t, "admin@test.com", res2)
res3 := resolveTarget(ctx, "0", flatBody, "email")
assert.Equal(t, "admin@test.com", res3)
})
t.Run("system target resolves to admin username for lark channel", func(t *testing.T) {
res := resolveTarget(ctx, "系统", flatBody, "lark")
assert.Equal(t, "admin_user", res)
})
}
func TestPushChannelAPI(t *testing.T) {
// 1. 模型校验测试
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())
// 校验 URL 安全前缀 HTTPS
c2 := &model.PushChannel{Name: "custom_channel", URL: "http://insecure-hook.com", Other: "{}"}
assert.Error(t, c2.Validate())
// 校验 JSON 格式
c3 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{invalid-json}"}
assert.Error(t, c3.Validate())
// 正确配置
c4 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{\"content\":\"$content\"}"}
assert.NoError(t, c4.Validate())
// 飞书渠道校验:非 HTTPS 地址报错
c5 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "http://open.feishu.cn", Other: ""}
assert.Error(t, c5.Validate())
// 飞书正确配置
c6 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "https://open.feishu.cn", Other: ""}
assert.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())
cTelegramErr2 := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "http://api.telegram.org", Token: "123:abc", Other: ""}
assert.Error(t, cTelegramErr2.Validate())
cTelegramOk := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "", Token: "123:abc", Other: "-100123"}
assert.NoError(t, cTelegramOk.Validate())
assert.Equal(t, "https://api.telegram.org", cTelegramOk.URL)
// 邮件配置校验:允许空配置以复用系统全局设置
c7 := &model.PushChannel{Name: "email_channel", Type: "email", URL: "", Token: "", Other: ""}
assert.NoError(t, c7.Validate())
// 邮件正确配置 (非 HTTPS 协议 URL 允许)
c8 := &model.PushChannel{Name: "email_channel", Type: "email", URL: "smtp.exmail.qq.com:465", Token: "user@example.com", Other: "authcode"}
assert.NoError(t, c8.Validate())
})
// 2. HTTP CRUD & 触发鉴权测试
dbConn, _, cleanup := setupPushTest(t)
defer cleanup()
// 构建路由以进行 HTTP 模拟请求
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
{
adminGroup.GET("/push/channels", ListChannels)
adminGroup.POST("/push/channels", CreateChannel)
adminGroup.PUT("/push/channels/:id", UpdateChannel)
adminGroup.DELETE("/push/channels/:id", DeleteChannel)
adminGroup.POST("/push/channels/test", TestChannel)
}
var createdID uint64
t.Run("admin create channel", func(t *testing.T) {
reqBody := CreateChannelRequest{
Name: "my_custom_channel",
Description: "My custom channel webhook",
Type: "custom",
Token: "my_chan_token",
URL: "https://webhook.site/test",
Other: `{"title": "$title", "body": "$content"}`,
Enabled: true,
}
bodyBytes, _ := json.Marshal(reqBody)
req, _ := http.NewRequest("POST", "/api/v1/admin/push/channels", bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataMap, ok := resp.Data.(map[string]any)
assert.True(t, ok)
assert.Equal(t, "my_custom_channel", dataMap["name"])
createdID = uint64(dataMap["id"].(float64))
})
t.Run("admin list channels", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/push/channels", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
list, ok := resp.Data.([]any)
assert.True(t, ok)
assert.Len(t, list, 1)
})
t.Run("admin update channel", func(t *testing.T) {
updateReq := UpdateChannelRequest{
Description: "Updated remark",
Type: "custom",
Token: "new_chan_token",
URL: "https://webhook.site/updated",
Other: `{"text": "$content"}`,
Enabled: true,
}
bodyBytes, _ := json.Marshal(updateReq)
req, _ := http.NewRequest("PUT", "/api/v1/admin/push/channels/"+strconv.FormatUint(createdID, 10), bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var updated model.PushChannel
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)
})
t.Run("admin test channel endpoint", func(t *testing.T) {
testReq := TestChannelRequest{
Name: "my_custom_channel",
Target: "test_target",
}
bodyBytes, _ := json.Marshal(testReq)
req, _ := http.NewRequest("POST", "/api/v1/admin/push/channels/test", bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
})
t.Run("admin delete channel", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/push/channels/"+strconv.FormatUint(createdID, 10), nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var count int64
dbConn.Model(&model.PushChannel{}).Where("id = ?", createdID).Count(&count)
assert.Equal(t, int64(0), count)
})
}
@@ -0,0 +1,292 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package push defines push notification HTTP routes.
package push
import (
"context"
"errors"
"fmt"
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/push"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// UpdateEventRequest 更新事件请求参数
type UpdateEventRequest struct {
Channels []string `json:"channels"`
Targets []string `json:"targets"`
Template string `json:"template" binding:"required"`
Enabled bool `json:"enabled"`
}
// TestPushRequest 测试推送通道请求参数
type TestPushRequest struct {
Config push.Config `json:"config" binding:"required"`
Target string `json:"target"`
}
// SyncEvents automatically registers/updates built-in events in the database.
func SyncEvents(ctx context.Context) error {
return syncBuiltInEvents(ctx)
}
// ListEvents 获取通知事件列表
// @Summary 获取所有通知事件
// @Description 返回系统配置的通知事件列表,包括预置和自定义事件,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.PushEvent} "通知事件列表"
// @Router /api/v1/admin/push/events [get]
func ListEvents(c *gin.Context) {
ctx := c.Request.Context()
events, err := listPushEvents(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(events))
}
// CreateEventRequest 创建事件请求参数
type CreateEventRequest struct {
EventKey string `json:"event_key"`
TaskType string `json:"task_type"` // 关联的异步任务类型
Channels []string `json:"channels"`
Targets []string `json:"targets"`
Template string `json:"template"`
Enabled bool `json:"enabled"`
}
func findBuiltInEvent(key string) (EventMetadata, bool) {
for _, meta := range BuiltInEvents {
if meta.Key == key {
return meta, true
}
}
return EventMetadata{}, false
}
// ListBuiltInEvents 获取内置通知事件列表
// @Summary 获取所有内置通知事件
// @Description 返回系统定义的所有内置通知事件元数据,供前端下拉框选择,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]EventMetadata} "内置通知事件列表"
// @Router /api/v1/admin/push/events/builtin [get]
func ListBuiltInEvents(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(BuiltInEvents))
}
// CreateEvent 创建通知事件
// @Summary 创建通知事件
// @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateEventRequest true "创建参数"
// @Success 200 {object} response.Any{data=model.PushEvent} "创建成功"
// @Router /api/v1/admin/push/events [post]
func CreateEvent(c *gin.Context) {
var req CreateEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
event, err := createPushEvent(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(event))
}
// DeleteEvent 删除通知事件配置
// @Summary 删除通知事件配置
// @Description 删除数据库中的特定通知事件配置,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param id path int true "事件 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Router /api/v1/admin/push/events/{id} [delete]
func DeleteEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
if err := deletePushEvent(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// UpdateEvent 更新通知事件
// @Summary 更新通知事件
// @Description 更新已有通知事件的推送渠道、接收目标和内容模板,需要管理员权限
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "事件 ID"
// @Param request body push.UpdateEventRequest true "更新参数"
// @Success 200 {object} response.Any{data=string} "修改成功"
// @Router /api/v1/admin/push/events/{id} [put]
func UpdateEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
var req UpdateEventRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := updatePushEvent(c.Request.Context(), id, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ToggleEvent 快捷切换通知事件启用状态
// @Summary 快捷切换通知事件启用状态
// @Description 启用或禁用指定的通知事件
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param id path int true "事件 ID"
// @Success 200 {object} response.Any{data=string} "切换成功"
// @Router /api/v1/admin/push/events/{id}/toggle [post]
func ToggleEvent(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
enabled, err := togglePushEvent(c.Request.Context(), id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
return
}
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(enabled))
}
// pushHistoriesResponse 推送历史分页响应
//
//nolint:unused
type pushHistoriesResponse struct {
Total int64 `json:"total"`
Results []model.PushHistory `json:"results"`
}
// ListHistories 分页获取通知推送历史
// @Summary 分页获取通知推送历史
// @Description 返回分页的通知历史日志数据,需要管理员权限
// @Tags admin-push
// @Produce json
// @Security SessionCookie
// @Param page query int false "当前页码"
// @Param page_size query int false "分页大小"
// @Param event_key query string false "过滤事件名称"
// @Param status query string false "过滤发送状态"
// @Success 200 {object} response.Any{data=pushHistoriesResponse} "推送历史列表"
// @Router /api/v1/admin/push/histories [get]
func ListHistories(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
total, results, err := listPushHistories(c.Request.Context(), repository.PushHistoryListFilter{
EventKey: c.Query("event_key"),
Status: c.Query("status"),
Page: page,
PageSize: pageSize,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(map[string]any{
"total": total,
"results": results,
}))
}
// TestPush 测试推送通道发送
// @Summary 测试推送通道发送
// @Description 接收临时通知渠道配置并在本地同步调用 Pusher.Send 发送测试消息
// @Tags admin-push
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body push.TestPushRequest true "测试请求体"
// @Success 200 {object} response.Any{data=string} "测试成功"
// @Router /api/v1/admin/push/test [post]
func TestPush(c *gin.Context) {
var req TestPushRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
pusher, err := push.GetPusher(req.Config.Channel)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := pusher.ValidateConfig(req.Config); err != nil {
response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err))
return
}
applySMTPFallbackToPushConfig(c.Request.Context(), &req.Config)
testBody := map[string]any{
keyTitle: "测试通道推送",
keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。",
keyLevel: defaultLevelInfo,
}
if err := pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,124 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"encoding/json"
"strconv"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
// RegisterTaskListeners subscribes push notification handlers to task completion events.
func RegisterTaskListeners() {
task.OnTaskCompleted(handleTaskCompleted)
}
func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.TaskResult, execErr error) {
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
if err != nil {
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
return
}
if len(events) == 0 {
return
}
body := map[string]any{
"task_id": execution.TaskID,
"task_name": execution.TaskName,
"task_type": execution.TaskType,
"task_status": string(execution.Status),
"task_duration": execution.Duration,
"time": time.Now().Format("2006-01-02 15:04:05"),
}
if execErr != nil {
body["task_error"] = execErr.Error()
} else {
body["task_error"] = ""
}
if result != nil {
body["task_result"] = result.Message
} else {
body["task_result"] = ""
}
var payloadMap map[string]any
if execution.Payload != "" {
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
body["payload"] = payloadMap
extractUserFromMap(ctx, payloadMap, body)
}
}
if result != nil && result.Detail != "" {
var detailMap map[string]any
if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil {
body["detail"] = detailMap
extractUserFromMap(ctx, detailMap, body)
}
}
for _, event := range events {
meta := EventMetadata{
Key: event.EventKey,
Name: event.Name,
Description: "异步任务执行完毕触发的自动通知",
}
DefaultTrigger.Trigger(ctx, meta, body)
}
}
func extractUserFromMap(ctx context.Context, data map[string]any, body map[string]any) {
if u, exists := body["user"]; exists && u != nil {
return
}
if user := loadUserFromPayload(ctx, data); user != nil {
body["user"] = user
}
}
func extractUserID(data map[string]any) (uint64, bool) {
for _, k := range []string{"user_id", "userId", "uid"} {
val, ok := data[k]
if !ok || val == nil {
continue
}
switch v := val.(type) {
case float64:
if v >= 0 {
return uint64(v), true
}
case int:
if v >= 0 {
return uint64(v), true
}
case int64:
if v >= 0 {
return uint64(v), true
}
case uint64:
return v, true
case string:
if id, err := strconv.ParseUint(v, 10, 64); err == nil {
return id, true
}
}
}
return 0, false
}
func extractUsername(data map[string]any) string {
for _, k := range []string{"username", "user_name"} {
if val, ok := data[k]; ok && val != nil {
if s, ok := val.(string); ok && s != "" {
return s
}
}
}
return ""
}
@@ -0,0 +1,121 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package push defines push notification HTTP routes and background tasks.
package push
import (
"context"
"encoding/json"
"errors"
"fmt"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/push"
)
const (
// SendNotificationTask 发送推送通知任务标识
SendNotificationTask = "push:send"
// TaskTypeSendNotification 推送通知管理类型
TaskTypeSendNotification = "send_notification"
)
// SendNotificationMeta represents the task metadata.
var SendNotificationMeta = task.TaskMeta{
Type: TaskTypeSendNotification,
AsynqTask: SendNotificationTask,
Name: "推送通知",
Description: "异步执行系统通知的多渠道派发与推送",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.TaskParam{
{
Name: "event_key",
Label: "事件标识",
Type: "string",
Required: true,
Placeholder: "admin_login",
},
{
Name: "target",
Label: "目标接收者",
Type: "string",
Required: false,
},
},
}
// PushHandler 通知推送异步任务处理器
//
//nolint:revive
type PushHandler struct{}
// ValidatePayload 校验并标准化推送参数
func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
return nil, errors.New("payload is required")
}
var req SendPayload
if err := json.Unmarshal(payload, &req); err != nil {
return nil, fmt.Errorf("invalid json format: %w", err)
}
if req.Config.Channel == "" {
return nil, errors.New("channel type is required")
}
return json.Marshal(req)
}
// Execute 异步执行推送操作并记录推送历史审计
func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
var req SendPayload
if err := json.Unmarshal(payload, &req); err != nil {
task.AppendLog(ctx, "解析推送参数失败: %v", err)
return nil, fmt.Errorf("parse payload failed: %w", err)
}
task.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
pusher, err := push.GetPusher(req.Config.Channel)
if err != nil {
errWrap := fmt.Errorf("get pusher failed: %w", err)
task.AppendLog(ctx, "推送失败: %v", errWrap)
if task.IsFinalAttempt(ctx) {
h.recordHistory(ctx, req, "failed", errWrap.Error())
}
return nil, errWrap
}
// 执行真正的消息推送,扁平化为原始 json 格式
flatBody := req.Body.Flatten()
err = pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil)
title := req.Body.Title
content := req.Body.Content
if err != nil {
task.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
if task.IsFinalAttempt(ctx) {
h.recordHistory(ctx, req, "failed", err.Error())
}
return nil, fmt.Errorf("pusher.Send failed: %w", err)
}
task.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
h.recordHistory(ctx, req, "success", "")
return &task.TaskResult{
Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target),
}, nil
}
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
}
}
@@ -0,0 +1,355 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package status 提供系统状态查询接口
package status
import (
"context"
"fmt"
"log"
"math"
"net/http"
"os"
"os/exec"
"runtime"
"time"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// startTime 记录服务启动时间
var startTime = time.Now()
const (
hoursInDay = 24
minutesInHour = 60
secondsInMinute = 60
nanosPerSecond = 1e9
binaryKB = 0
binaryMB = 1
binaryGB = 2
valueThreshold = 10 // 格式化时区分整数显示的阈值
)
// SystemStatusResponse 系统状态响应结构体
type SystemStatusResponse struct {
Uptime string `json:"uptime"`
NumGoroutine int `json:"num_goroutine"`
Alloc string `json:"alloc"`
TotalAlloc string `json:"total_alloc"`
Sys string `json:"sys"`
Lookups uint64 `json:"lookups"`
Mallocs uint64 `json:"mallocs"`
Frees uint64 `json:"frees"`
HeapAlloc string `json:"heap_alloc"`
HeapSys string `json:"heap_sys"`
HeapIdle string `json:"heap_idle"`
HeapInuse string `json:"heap_inuse"`
HeapReleased string `json:"heap_released"`
HeapObjects uint64 `json:"heap_objects"`
StackInuse string `json:"stack_inuse"`
StackSys string `json:"stack_sys"`
MSpanInuse string `json:"mspan_inuse"`
MSpanSys string `json:"mspan_sys"`
MCacheInuse string `json:"mcache_inuse"`
MCacheSys string `json:"mcache_sys"`
BuckHashSys string `json:"buck_hash_sys"`
GCSys string `json:"gc_sys"`
OtherSys string `json:"other_sys"`
NextGC string `json:"next_gc"`
LastGCTime string `json:"last_gc_time"`
PauseTotalNs string `json:"pause_total_ns"`
LastPause string `json:"last_pause"`
NumGC uint32 `json:"num_gc"`
}
// formatBytes 格式化字节大小
func formatBytes(bytes uint64) string {
const unit = 1024
if bytes < unit {
return fmt.Sprintf("%d B", bytes)
}
div, exp := int64(unit), 0
for n := bytes / unit; n >= unit; n /= unit {
div *= unit
exp++
}
value := float64(bytes) / float64(div)
var suffix string
switch exp {
case binaryKB:
suffix = "KiB"
case binaryMB:
suffix = "MiB"
case binaryGB:
suffix = "GiB"
default:
suffix = "TiB"
}
// 格式化规则:
// - 如果是整数(如 16, 73, 105, 986, 112):
// - 如果 >= 10,则格式化为 "%.0f" (e.g. "16 KiB")
// - 如果 < 10,则格式化为 "%.1f" (e.g. "9.0 KiB")
// - 如果不是整数(如 5.8, 9.1, 7.6, 4.8):格式化为 "%.1f"
if value == math.Trunc(value) {
if value >= valueThreshold {
return fmt.Sprintf("%.0f %s", value, suffix)
}
return fmt.Sprintf("%.1f %s", value, suffix)
}
return fmt.Sprintf("%.1f %s", value, suffix)
}
// formatDuration 格式化时间持续时间
func formatDuration(d time.Duration) string {
days := int(d.Hours()) / hoursInDay
hours := int(d.Hours()) % hoursInDay
minutes := int(d.Minutes()) % minutesInHour
seconds := int(d.Seconds()) % secondsInMinute
var res string
if days > 0 {
res += fmt.Sprintf("%d天", days)
}
if hours > 0 {
res += fmt.Sprintf("%d小时", hours)
}
if minutes > 0 {
res += fmt.Sprintf("%d分钟", minutes)
}
if seconds > 0 || res == "" {
res += fmt.Sprintf("%d秒钟", seconds)
}
return res
}
// GetSystemStatus 获取系统状态信息
// @Summary 获取系统状态信息
// @Description 获取后端服务运行状态、Goroutine、内存指标等详细统计数据,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=status.SystemStatusResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/status [get]
func GetSystemStatus(c *gin.Context) {
var m runtime.MemStats
runtime.ReadMemStats(&m)
uptime := formatDuration(time.Since(startTime))
numGoroutine := runtime.NumGoroutine()
var lastGCTime string
switch {
case m.LastGC > 0 && m.LastGC <= math.MaxInt64:
lastGCTime = formatDuration(time.Since(time.Unix(0, int64(m.LastGC))))
case m.LastGC > 0:
lastGCTime = "未知"
default:
lastGCTime = "无"
}
var lastPause string
if m.NumGC > 0 {
lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond)
} else {
lastPause = "0.000s"
}
res := SystemStatusResponse{
Uptime: uptime,
NumGoroutine: numGoroutine,
Alloc: formatBytes(m.Alloc),
TotalAlloc: formatBytes(m.TotalAlloc),
Sys: formatBytes(m.Sys),
Lookups: m.Lookups,
Mallocs: m.Mallocs,
Frees: m.Frees,
HeapAlloc: formatBytes(m.HeapAlloc),
HeapSys: formatBytes(m.HeapSys),
HeapIdle: formatBytes(m.HeapIdle),
HeapInuse: formatBytes(m.HeapInuse),
HeapReleased: formatBytes(m.HeapReleased),
HeapObjects: m.HeapObjects,
StackInuse: formatBytes(m.StackInuse),
StackSys: formatBytes(m.StackSys),
MSpanInuse: formatBytes(m.MSpanInuse),
MSpanSys: formatBytes(m.MSpanSys),
MCacheInuse: formatBytes(m.MCacheInuse),
MCacheSys: formatBytes(m.MCacheSys),
BuckHashSys: formatBytes(m.BuckHashSys),
GCSys: formatBytes(m.GCSys),
OtherSys: formatBytes(m.OtherSys),
NextGC: formatBytes(m.NextGC),
LastGCTime: lastGCTime,
PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond),
LastPause: lastPause,
NumGC: m.NumGC,
}
c.JSON(http.StatusOK, response.OK(res))
}
// DatabaseInfoResponse 数据库信息响应结构体
type DatabaseInfoResponse struct {
Type string `json:"type"`
Name string `json:"name"`
Version string `json:"version"`
}
// getSQLiteInfo 返回 SQLite 数据库信息
func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
info := DatabaseInfoResponse{
Type: "sqlite",
Name: config.Config.Database.SQLitePath,
Version: "SQLite",
}
if info.Name == "" {
info.Name = "./data/openflare.db"
}
gormDB := db.DB(ctx)
if gormDB == nil {
return info
}
var ver string
if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil && ver != "" {
info.Version = "SQLite " + ver
}
return info
}
// getPostgresInfo 返回 PostgreSQL 数据库信息
func getPostgresInfo(ctx context.Context) DatabaseInfoResponse {
info := DatabaseInfoResponse{
Type: "postgres",
Name: config.Config.Database.Database,
Version: "PostgreSQL",
}
gormDB := db.DB(ctx)
if gormDB == nil {
return info
}
var ver string
if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil && ver != "" {
info.Version = ver
}
return info
}
// GetDatabaseInfo 获取当前数据库类型及版本信息
// @Summary 获取数据库信息
// @Description 返回当前使用的数据库类型(sqlite/postgres)、名称/路径及版本字符串,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=status.DatabaseInfoResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/db-info [get]
func GetDatabaseInfo(c *gin.Context) {
var info DatabaseInfoResponse
if !config.Config.Database.Enabled {
info = getSQLiteInfo(c.Request.Context())
} else {
info = getPostgresInfo(c.Request.Context())
}
c.JSON(http.StatusOK, response.OK(info))
}
// ExportDatabase 导出数据库
// @Summary 导出数据库
// @Description SQLite 时直接下载 .db 文件;PostgreSQL 时执行 pg_dump 并流式下载 .sql 文件,需要管理员权限
// @Tags admin
// @Produce application/octet-stream
// @Security SessionCookie
// @Success 200 {file} binary "数据库文件"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "导出失败"
// @Router /api/v1/admin/db-export [get]
func ExportDatabase(c *gin.Context) {
if !config.Config.Database.Enabled {
exportSQLite(c)
} else {
exportPostgres(c)
}
}
// exportSQLite 以 HTTP 附件方式下载 SQLite .db 文件
func exportSQLite(c *gin.Context) {
path := config.Config.Database.SQLitePath
if path == "" {
path = "./data/openflare.db"
}
f, err := os.Open(path) //nolint:gosec // path is loaded from server startup configuration, not user input
if err != nil {
response.AbortInternal(c, "无法打开数据库文件: "+err.Error())
return
}
defer func() {
if closeErr := f.Close(); closeErr != nil {
_ = closeErr
}
}()
fi, err := f.Stat()
if err != nil {
response.AbortInternal(c, "无法读取数据库文件信息: "+err.Error())
return
}
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.Status(http.StatusOK)
http.ServeContent(c.Writer, c.Request, "openflare.db", fi.ModTime(), f)
}
// exportPostgres 执行 pg_dump 并将输出流式传输给客户端
func exportPostgres(c *gin.Context) {
dbCfg := config.Config.Database
// 检查 pg_dump 是否可用
pgDumpPath, err := exec.LookPath("pg_dump")
if err != nil {
response.AbortInternal(c, "pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具")
return
}
args := []string{
"--no-password",
"-h", dbCfg.Host,
"-p", fmt.Sprintf("%d", dbCfg.Port),
"-U", dbCfg.Username,
dbCfg.Database,
}
cmd := exec.CommandContext(c.Request.Context(), pgDumpPath, args...) //nolint:gosec // pgDumpPath is a looked up command path, args are from database configuration
if dbCfg.Password != "" {
cmd.Env = append(os.Environ(), "PGPASSWORD="+dbCfg.Password)
} else {
cmd.Env = os.Environ()
}
fileName := fmt.Sprintf("openflare_%s.sql", time.Now().Format("20060102_150405"))
c.Header("Content-Disposition", `attachment; filename="`+fileName+`"`)
c.Header("Content-Type", "application/octet-stream")
c.Status(http.StatusOK)
cmd.Stdout = c.Writer
cmd.Stderr = nil // 忽略 stderr 以避免污染输出流
if err := cmd.Run(); err != nil {
// 响应头已发出,无法再写 JSON 错误,记录到服务器日志
log.Printf("[db-export] pg_dump failed: %v\n", err)
}
}
@@ -0,0 +1,15 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package system_config 提供系统配置管理功能
package system_config
// 系统配置错误消息常量
const (
SystemConfigNotFound = "系统配置不存在"
ConfigKeyRequired = "配置键不能为空"
ConfigValueRequired = "配置值不能为空"
ConfigKeyExists = "配置键已存在"
StorageDriverSwitchRequiresMigration = "存在存量文件,请通过存储迁移任务切换存储引擎"
)
@@ -0,0 +1,128 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package system_config
import (
"context"
"encoding/json"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
)
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
exists, err := repository.SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
if exists {
return errors.New(ConfigKeyExists)
}
config := model.SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := repository.CreateSystemConfig(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
return nil
}
func listSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
return repository.ListAdminSystemConfigs(ctx, configType)
}
func getSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) {
return repository.GetAdminSystemConfigByKey(ctx, key)
}
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return err
}
var originalDriver storage.Driver
if key == model.ConfigKeyStorageConfig {
var currentCfg storage.Config
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value)
if err != nil {
return err
}
req.Value = validatedVal
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
updates := map[string]any{
"description": req.Description,
}
if req.Visibility != nil {
updates["visibility"] = *req.Visibility
config.Visibility = *req.Visibility
}
if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue {
updates["value"] = req.Value
config.Value = req.Value
}
if err := tx.Model(&config).Updates(updates).Error; err != nil {
return err
}
resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value)
return nil
}); err != nil {
return err
}
invalidateCachesAfterConfigUpdate(ctx, key)
return nil
}
func resolveStorageMigrationTasksOnDirectDriverUpdate(
ctx context.Context,
tx *gorm.DB,
key string,
originalDriver storage.Driver,
newValue string,
) {
if key != model.ConfigKeyStorageConfig || originalDriver == "" {
return
}
var newCfg storage.Config
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
return
}
if newCfg.Driver != originalDriver {
return
}
if err := tx.Model(&model.TaskExecution{}).
Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed).
Updates(map[string]any{
"status": model.TaskExecutionStatusSucceeded,
"result": "存储配置直接更新,故障迁移任务自动标记为已解决",
"finished_at": time.Now(),
}).Error; err != nil {
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
}
}
@@ -0,0 +1,365 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package system_config
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/apps/cap"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/pkg/logger"
mail "github.com/Rain-kl/Wavelet/pkg/mail"
)
const maskedConfigValue = "******"
// CreateSystemConfigRequest 创建系统配置请求
type CreateSystemConfigRequest struct {
Key string `json:"key" binding:"required,max=64"`
Value string `json:"value" binding:"required"`
Type string `json:"type" binding:"required,oneof=system business"`
Visibility int `json:"visibility" binding:"oneof=0 1"`
Description string `json:"description" binding:"max=255"`
}
// UpdateSystemConfigRequest 更新系统配置请求
type UpdateSystemConfigRequest struct {
Value string `json:"value" binding:"required"`
Visibility *int `json:"visibility" binding:"omitempty,oneof=0 1"`
Description string `json:"description" binding:"max=255"`
}
// CreateSystemConfig 创建系统配置
// @Summary 创建系统配置
// @Description 创建一条新的系统配置项,配置键不可重复,同时将新配置同步到 Redis,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body system_config.CreateSystemConfigRequest true "创建请求参数"
// @Success 200 {object} response.Any{data=string} "创建成功"
// @Failure 400 {object} response.Any "参数错误或配置键已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs [post]
func CreateSystemConfig(c *gin.Context) {
var req CreateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := createSystemConfig(c.Request.Context(), req); err != nil {
if err.Error() == ConfigKeyExists {
response.AbortBadRequest(c, ConfigKeyExists)
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// ListSystemConfigs 获取系统配置列表
// @Summary 获取系统配置列表
// @Description 返回所有系统配置列表,支持按配置类型(system/business)过滤,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param type query string false "配置类型(system/business)"
// @Success 200 {object} response.Any{data=[]model.SystemConfig} "系统配置列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs [get]
func ListSystemConfigs(c *gin.Context) {
configs, err := listSystemConfigs(c.Request.Context(), c.Query("type"))
if err != nil {
response.AbortInternal(c, err.Error())
return
}
for i := range configs {
configs[i].Value = maskSensitiveConfig(configs[i].Key, configs[i].Value)
}
c.JSON(http.StatusOK, response.OK(configs))
}
// GetSystemConfig 获取单个系统配置
// @Summary 获取单个系统配置
// @Description 根据配置键获取对应的系统配置详情,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Success 200 {object} response.Any{data=model.SystemConfig} "系统配置详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs/{key} [get]
func GetSystemConfig(c *gin.Context) {
config, err := getSystemConfig(c.Request.Context(), c.Param("key"))
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
} else {
response.AbortInternal(c, err.Error())
}
return
}
config.Value = maskSensitiveConfig(config.Key, config.Value)
c.JSON(http.StatusOK, response.OK(config))
}
// UpdateSystemConfig 更新系统配置
// @Summary 更新系统配置
// @Description 根据配置键更新对应的配置内容,同时将更新同步到 Redis,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Param request body system_config.UpdateSystemConfigRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs/{key} [put]
func UpdateSystemConfig(c *gin.Context) {
var req UpdateSystemConfigRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
key := c.Param("key")
if err := updateSystemConfig(c.Request.Context(), key, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
return
}
if isStorageConfigValidationError(err) {
response.AbortBadRequest(c, err.Error())
return
}
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
if cap.IsRuntimeConfigKey(key) {
cap.InvalidateRuntimeSettings()
}
}
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if key == model.ConfigKeyStorageConfig {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
storage.ResetCache()
storage.PublishCacheInvalidation(ctx)
}
if key == model.ConfigKeyFileAccessWhitelist {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
}
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
}
// TestSMTPRequest 测试 SMTP 配置请求
type TestSMTPRequest struct {
SMTPHost string `json:"smtp_host" binding:"required,max=255"`
SMTPPort int `json:"smtp_port" binding:"required"`
SMTPUsername string `json:"smtp_username" binding:"required,max=255"`
SMTPPassword string `json:"smtp_password" binding:"required,max=255"`
To string `json:"to" binding:"required,email"`
}
// TestSMTPResponse 测试 SMTP 配置响应
type TestSMTPResponse struct {
Success bool `json:"success"`
Log string `json:"log"`
Error string `json:"error"`
}
// TestSMTP 测试 SMTP 邮件发送
// @Summary 测试 SMTP 邮件发送
// @Description 使用传入的配置进行 SMTP 邮件发送测试,支持使用 ****** 占位符使用保存的数据库密码
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body system_config.TestSMTPRequest true "测试请求参数"
// @Success 200 {object} response.Any{data=system_config.TestSMTPResponse} "测试执行完毕"
// @Failure 400 {object} response.Any "参数错误"
// @Router /api/v1/admin/system-configs/smtp/test [post]
func TestSMTP(c *gin.Context) {
var req TestSMTPRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
password := req.SMTPPassword
if password == maskedConfigValue {
if sc, err := repository.GetSystemConfigByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil {
password = sc.Value
}
}
cfg := mail.Config{
Host: req.SMTPHost,
Port: req.SMTPPort,
Username: req.SMTPUsername,
Password: password,
}
subject := "OpenFlare SMTP Test Mail"
body := `<h3>SMTP Mail Connection Test</h3>
<p>If you received this message, your SMTP configuration is correct and mail sending is working properly.</p>
<p>Sent from OpenFlare.</p>`
logs, err := mail.SendMailWithLog(c.Request.Context(), cfg, req.To, subject, body)
resp := TestSMTPResponse{
Success: err == nil,
Log: logs,
}
if err != nil {
resp.Error = err.Error()
}
c.JSON(http.StatusOK, response.OK(resp))
}
func isStorageConfigValidationError(err error) bool {
msg := err.Error()
return msg == StorageDriverSwitchRequiresMigration ||
strings.HasPrefix(msg, "解析") ||
strings.HasPrefix(msg, "验证") ||
strings.HasPrefix(msg, "初始化测试") ||
strings.HasPrefix(msg, "存储连通性") ||
strings.HasPrefix(msg, "序列化") ||
strings.HasPrefix(msg, "检查存量文件")
}
func maskSensitiveConfig(key, value string) string {
if value == "" {
return value
}
switch key {
case model.ConfigKeySMTPPassword:
return maskedConfigValue
case model.ConfigKeyStorageConfig:
var cfg storage.Config
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
masked := storage.MaskSecrets(cfg)
if val, err := json.Marshal(masked); err == nil {
return string(val)
}
}
}
return value
}
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
// and tests connectivity of the new storage configuration.
func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) {
var currentCfg storage.Config
if err := json.Unmarshal([]byte(currentConfig), &currentCfg); err != nil {
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
}
var newCfg storage.Config
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
}
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
targetCfg := storage.MergeMaskedSecrets(newCfg, currentCfg)
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
return "", err
}
// 序列化为最终保存的真实明文配置,防止保存屏蔽的 ****** 字符
unmaskedVal, err := json.Marshal(targetCfg)
if err != nil {
return "", fmt.Errorf("序列化存储配置失败: %w", err)
}
return string(unmaskedVal), nil
}
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg storage.Config) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64
if err := db.DB(ctx).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
}
if uploadCount > 0 {
return errors.New(StorageDriverSwitchRequiresMigration)
}
if err := validateDriverConfig(targetCfg, newCfg.Driver); err != nil {
return fmt.Errorf("验证目标存储配置参数失败: %w", err)
}
pendingCfg := targetCfg
pendingCfg.Driver = newCfg.Driver
return testStorageBackend(ctx, pendingCfg, newCfg.Driver)
}
if err := storage.ValidateConfig(targetCfg); err != nil {
return fmt.Errorf("验证存储配置参数失败: %w", err)
}
return testStorageBackend(ctx, targetCfg, targetCfg.Driver)
}
func validateDriverConfig(cfg storage.Config, driver storage.Driver) error {
cfg.Driver = driver
return storage.ValidateConfig(cfg)
}
func testStorageBackend(ctx context.Context, cfg storage.Config, driver storage.Driver) error {
testBackend, err := storage.NewBackend(ctx, cfg, driver)
if err != nil {
return fmt.Errorf("初始化测试存储实例失败: %w", err)
}
if err := testBackend.Test(ctx); err != nil {
return fmt.Errorf("存储连通性测试失败: %w", err)
}
return nil
}
@@ -0,0 +1,543 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package system_config
import (
"bufio"
"bytes"
"context"
"encoding/json"
"net"
"net/http"
"net/http/httptest"
"net/textproto"
"strings"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
const expectedDefaultConfigsCount = 30
func setupTestRouter(authUser *model.User) *gin.Engine {
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
adminGroup.POST("/system-configs", CreateSystemConfig)
adminGroup.GET("/system-configs", ListSystemConfigs)
systemConfigRouter := adminGroup.Group("/system-configs/:key")
{
systemConfigRouter.GET("", GetSystemConfig)
systemConfigRouter.PUT("", UpdateSystemConfig)
}
return r
}
func TestCreateSystemConfig(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("create successfully", func(t *testing.T) {
payload := CreateSystemConfigRequest{
Key: "custom_key",
Value: "custom_value",
Type: "system",
Visibility: model.ConfigVisibilityVisible,
Description: "desc",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
// Verify database
var cfg model.SystemConfig
err := dbConn.Where("key = ?", "custom_key").First(&cfg).Error
if err != nil {
t.Fatalf("failed to find system config in DB: %v", err)
}
// Verify caches are invalidated after create and repopulate on read
_, err = db.Redis.HGet(
context.Background(),
db.PrefixedKey(repository.SystemConfigRedisHashKey),
"custom_key",
).Result()
if err == nil {
t.Fatal("expected redis cache miss immediately after create")
}
loaded, err := repository.GetSystemConfigByKey(context.Background(), "custom_key")
if err != nil {
t.Fatalf("GetSystemConfigByKey(custom_key) error = %v", err)
}
if loaded.Value != "custom_value" {
t.Errorf("GetSystemConfigByKey(custom_key).Value = %q, want %q", loaded.Value, "custom_value")
}
if loaded.Visibility != model.ConfigVisibilityVisible {
t.Errorf("GetSystemConfigByKey(custom_key).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityVisible)
}
})
t.Run("create duplicate key error", func(t *testing.T) {
// Key "custom_key" already exists from previous test
payload := CreateSystemConfigRequest{
Key: "custom_key",
Value: "another_value",
Type: "system",
Description: "desc",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request on duplicate key, got %d", w.Code)
}
})
}
func TestListSystemConfigs(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("list all seeded configurations", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d", w.Code)
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var configs []model.SystemConfig
_ = json.Unmarshal(dataBytes, &configs)
// Defaults seed configurations
if len(configs) != expectedDefaultConfigsCount {
t.Errorf("expected %d default configs, got %d", expectedDefaultConfigsCount, len(configs))
}
})
t.Run("filter by type business", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs?type=business", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var configs []model.SystemConfig
_ = json.Unmarshal(dataBytes, &configs)
if len(configs) != 1 || configs[0].Key != model.ConfigKeyMaxAPIKeysPerUser {
t.Errorf("expected 1 business config (max_api_keys_per_user), got %d: %v", len(configs), configs)
}
})
}
func TestGetSystemConfig(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("get existing configuration", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d", w.Code)
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var cfg model.SystemConfig
_ = json.Unmarshal(dataBytes, &cfg)
if cfg.Value != "OpenFlare" {
t.Errorf("expected 'OpenFlare', got '%s'", cfg.Value)
}
})
t.Run("get non-existent config", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/system-configs/non_existent_key", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404 Not Found, got %d", w.Code)
}
})
}
func TestUpdateSystemConfig(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("update successfully", func(t *testing.T) {
hidden := model.ConfigVisibilityHidden
payload := UpdateSystemConfigRequest{
Value: "Super Site Name",
Visibility: &hidden,
Description: "Updated Description",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/"+model.ConfigKeySiteName, bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
// Verify database
var cfg model.SystemConfig
dbConn.Where("key = ?", model.ConfigKeySiteName).First(&cfg)
if cfg.Value != "Super Site Name" || cfg.Description != "Updated Description" || cfg.Visibility != model.ConfigVisibilityHidden {
t.Errorf("database values not updated: %+v", cfg)
}
// Verify caches are invalidated after update and repopulate on read
_, err := db.Redis.HGet(
context.Background(),
db.PrefixedKey(repository.SystemConfigRedisHashKey),
model.ConfigKeySiteName,
).Result()
if err == nil {
t.Fatal("expected redis cache miss immediately after update")
}
loaded, err := repository.GetSystemConfigByKey(context.Background(), model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err)
}
if loaded.Value != "Super Site Name" {
t.Errorf("GetSystemConfigByKey(site_name).Value = %q, want %q", loaded.Value, "Super Site Name")
}
if loaded.Visibility != model.ConfigVisibilityHidden {
t.Errorf("GetSystemConfigByKey(site_name).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityHidden)
}
})
t.Run("update non-existent config", func(t *testing.T) {
payload := UpdateSystemConfigRequest{
Value: "New Value",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/invalid_key", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404 Not Found, got %d", w.Code)
}
})
}
func TestTestSMTP(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
r := setupTestRouter(adminUser)
r.POST("/api/v1/admin/system-configs/smtp/test", TestSMTP)
// Start a mock SMTP server
l, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("failed to start mock smtp server: %v", err)
}
defer func() { _ = l.Close() }()
port := l.Addr().(*net.TCPAddr).Port
go func() {
conn, err := l.Accept()
if err != nil {
return
}
defer func() { _ = conn.Close() }()
writer := bufio.NewWriter(conn)
reader := bufio.NewReader(conn)
tp := textproto.NewReader(reader)
// 220 Ready
_, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n")
_ = writer.Flush()
// Read HELO/EHLO
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n")
_ = writer.Flush()
// Read AUTH PLAIN
_, _ = tp.ReadLine()
_, _ = writer.WriteString("235 Authentication successful\r\n")
_ = writer.Flush()
// Read MAIL FROM
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read RCPT TO
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read DATA
_, _ = tp.ReadLine()
_, _ = writer.WriteString("354 Start mail input\r\n")
_ = writer.Flush()
// Read body lines until dot
for {
line, err := tp.ReadLine()
if err != nil || line == "." {
break
}
}
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read QUIT
_, _ = tp.ReadLine()
_, _ = writer.WriteString("221 Bye\r\n")
_ = writer.Flush()
}()
payload := TestSMTPRequest{
SMTPHost: "127.0.0.1",
SMTPPort: port,
SMTPUsername: "sender@example.com",
SMTPPassword: "password",
To: "recipient@example.com",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/system-configs/smtp/test", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var testResp TestSMTPResponse
json.Unmarshal(dataBytes, &testResp)
if !testResp.Success {
t.Errorf("expected test success, got failed: %s. Log: %s", testResp.Error, testResp.Log)
}
}
func TestUpdateStorageConfigValidation(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("update storage config successfully", func(t *testing.T) {
tempDir := t.TempDir()
cfg := storage.DefaultConfig()
cfg.Local.Root = tempDir
cfgBytes, _ := json.Marshal(cfg)
payload := UpdateSystemConfigRequest{
Value: string(cfgBytes),
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
// Verify database
var dbCfg model.SystemConfig
dbConn.Where("key = ?", "storage_config").First(&dbCfg)
var savedCfg storage.Config
_ = json.Unmarshal([]byte(dbCfg.Value), &savedCfg)
if savedCfg.Local.Root != tempDir {
t.Errorf("expected local root to be updated to %s, got %s", tempDir, savedCfg.Local.Root)
}
})
t.Run("update storage config failed connectivity check", func(t *testing.T) {
cfg := storage.DefaultConfig()
cfg.Driver = storage.DriverS3
cfg.S3.Bucket = "non-existent-bucket"
cfg.S3.Endpoint = "http://127.0.0.1:9999" // Will fail connectivity check
cfgBytes, _ := json.Marshal(cfg)
payload := UpdateSystemConfigRequest{
Value: string(cfgBytes),
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String())
}
})
t.Run("reject driver switch when uploads exist", func(t *testing.T) {
upload := model.Upload{
ID: 88001,
UserID: 1,
FileName: "keep.txt",
FilePath: "uploads/keep.txt",
FileSize: 4,
MimeType: "text/plain",
Extension: "txt",
Type: "attachment",
Status: model.UploadStatusUsed,
}
if err := dbConn.Create(&upload).Error; err != nil {
t.Fatalf("seed upload failed: %v", err)
}
tempDir := t.TempDir()
cfg := storage.DefaultConfig()
cfg.Driver = storage.DriverS3
cfg.S3.Endpoint = "http://127.0.0.1:19998"
cfg.S3.Region = "us-east-1"
cfg.S3.Bucket = "wavelet"
cfg.S3.AccessKeyID = "test"
cfg.S3.SecretAccessKey = "test"
cfg.Local.Root = tempDir
cfgBytes, _ := json.Marshal(cfg)
payload := UpdateSystemConfigRequest{Value: string(cfgBytes)}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Fatalf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String())
}
if !strings.Contains(w.Body.String(), StorageDriverSwitchRequiresMigration) {
t.Fatalf("expected migration-required error, got: %s", w.Body.String())
}
})
t.Run("switch to local while active s3 is unreachable", func(t *testing.T) {
if err := dbConn.Where("1 = 1").Delete(&model.Upload{}).Error; err != nil {
t.Fatalf("clear uploads failed: %v", err)
}
activeCfg := storage.DefaultConfig()
activeCfg.Driver = storage.DriverS3
activeCfg.S3.Endpoint = "http://127.0.0.1:9999"
activeCfg.S3.Region = "us-east-1"
activeCfg.S3.Bucket = "wavelet"
activeCfg.S3.AccessKeyID = "test"
activeCfg.S3.SecretAccessKey = "test"
activeBytes, _ := json.Marshal(activeCfg)
seedCfg := model.SystemConfig{
Key: "storage_config",
Value: string(activeBytes),
Type: "system",
}
if err := dbConn.Where("key = ?", "storage_config").
Assign(map[string]any{"value": seedCfg.Value, "type": seedCfg.Type}).
FirstOrCreate(&seedCfg).Error; err != nil {
t.Fatalf("seed active storage config failed: %v", err)
}
tempDir := t.TempDir()
stagedCfg := activeCfg
stagedCfg.Driver = storage.DriverLocal
stagedCfg.Local.Root = tempDir
cfgBytes, _ := json.Marshal(stagedCfg)
payload := UpdateSystemConfigRequest{Value: string(cfgBytes)}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/system-configs/storage_config", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var dbCfg model.SystemConfig
if err := dbConn.Where("key = ?", "storage_config").First(&dbCfg).Error; err != nil {
t.Fatalf("load saved storage config failed: %v", err)
}
var savedCfg storage.Config
if err := json.Unmarshal([]byte(dbCfg.Value), &savedCfg); err != nil {
t.Fatalf("parse saved storage config failed: %v", err)
}
if savedCfg.Driver != storage.DriverLocal {
t.Fatalf("active driver = %q, want %q after save", savedCfg.Driver, storage.DriverLocal)
}
if savedCfg.Local.Root != tempDir {
t.Fatalf("staged local root = %q, want %q", savedCfg.Local.Root, tempDir)
}
})
}
@@ -0,0 +1,23 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package task 提供任务管理接口
package task
// 任务管理相关错误消息
const (
InvalidTaskType = "无效的任务类型"
InvalidTimeRange = "无效的时间范围"
TaskDispatchFailed = "任务下发失败"
UserIDRequired = "用户ID必填"
TaskNotFound = "任务执行记录不存在"
TaskNotRetryable = "该任务不支持重试"
TaskNotFailed = "只有失败的任务才能重试"
TaskMaxRetryExceeded = "已达到最大重试次数"
TaskRetryFailed = "任务重试失败"
InvalidCronExpression = "无效的 Cron 表达式"
ScheduleNotFound = "定时任务不存在"
ScheduleSaveFailed = "保存定时任务失败"
ScheduleDeleteFailed = "删除定时任务失败"
)
@@ -0,0 +1,416 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/task/scheduler"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// ListTaskTypes 获取支持的任务类型列表
// @Summary 获取支持的任务类型
// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]task.TaskMeta} "任务类型列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/types [get]
func ListTaskTypes(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(task.GetDispatchableTasks()))
}
// DispatchTaskRequest 下发任务请求
type DispatchTaskRequest struct {
TaskType string `json:"task_type" binding:"required"`
StartTime *time.Time `json:"start_time"`
EndTime *time.Time `json:"end_time"`
UserID *uint64 `json:"user_id"`
Payload string `json:"payload"`
}
// DispatchTask 下发任务
// @Summary 下发异步任务
// @Description 手动触发指定类型的异步任务,支持指定时间范围和用户,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body DispatchTaskRequest true "任务请求参数"
// @Success 200 {object} response.Any{data=string} "任务已入队"
// @Failure 400 {object} response.Any "任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "任务入队失败"
// @Router /api/v1/admin/tasks/dispatch [post]
func DispatchTask(c *gin.Context) {
var req DispatchTaskRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
meta := task.GetTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
}
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskID, err := task.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
if err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
return
}
c.JSON(http.StatusOK, response.OK(taskID))
}
// ListTaskExecutions 查询任务执行记录列表
// @Summary 查询任务执行记录
// @Description 分页查询任务执行记录,支持按状态和任务类型筛选,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param status query string false "状态筛选 (pending/running/succeeded/failed)"
// @Param task_type query string false "任务类型筛选"
// @Param page query int false "页码" default(1)
// @Param page_size query int false "每页条数" default(20)
// @Success 200 {object} response.Any{data=object} "任务执行记录列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/executions [get]
func ListTaskExecutions(c *gin.Context) {
var req model.ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if req.TaskType != "" {
if meta := task.GetTaskMeta(req.TaskType); meta != nil {
req.TaskType = meta.AsynqTask
}
}
executions, total, err := model.ListTaskExecutions(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(gin.H{
"items": executions,
"total": total,
"page": req.Page,
"page_size": req.PageSize,
}))
}
// GetTaskExecution 查询单条任务执行详情
// @Summary 查询任务执行详情
// @Description 根据 ID 查询任务执行记录详情,包含完整执行日志,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=model.TaskExecution} "任务执行详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "记录不存在"
// @Router /api/v1/admin/tasks/executions/{id} [get]
func GetTaskExecution(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, admin.InvalidTaskExecutionID)
return
}
execution, err := model.GetTaskExecutionByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, TaskNotFound)
return
}
c.JSON(http.StatusOK, response.OK(execution))
}
// RetryTask 重试失败的任务
// @Summary 重试失败任务
// @Description 重新下发一条失败的任务,创建新的执行记录,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=string} "新任务的 TaskID"
// @Failure 400 {object} response.Any "任务不支持重试或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "记录不存在"
// @Failure 500 {object} response.Any "重试失败"
// @Router /api/v1/admin/tasks/executions/{id}/retry [post]
func RetryTask(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, admin.InvalidTaskExecutionID)
return
}
newTaskID, err := task.RetryTask(c.Request.Context(), id)
if err != nil {
errMsg := err.Error()
switch {
case strings.Contains(errMsg, "不存在"):
response.AbortNotFound(c, errMsg)
case strings.Contains(errMsg, "只有失败的任务") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"):
response.AbortBadRequest(c, errMsg)
default:
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err))
}
return
}
c.JSON(http.StatusOK, response.OK(newTaskID))
}
// ListSchedules 获取定时任务列表
// @Summary 获取定时任务列表
// @Description 返回系统所有的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Schedule} "定时任务列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/schedules [get]
func ListSchedules(c *gin.Context) {
schedules, err := model.ListSchedules(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(schedules))
}
// CreateScheduleRequest 创建定时任务请求
type CreateScheduleRequest struct {
Name string `json:"name" binding:"required"`
TaskType string `json:"task_type" binding:"required"`
Cron string `json:"cron" binding:"required"`
Payload string `json:"payload"`
IsActive *bool `json:"is_active" binding:"required"`
}
// CreateSchedule 创建定时任务
// @Summary 创建定时任务
// @Description 新增一个动态定时任务配置,关联已有的异步任务,配置 Cron 表达式和执行参数,并触发调度器热加载,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body CreateScheduleRequest true "创建定时任务请求参数"
// @Success 200 {object} response.Any{data=model.Schedule} "创建成功的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、异步任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "保存定时任务失败"
// @Router /api/v1/admin/tasks/schedules [post]
func CreateSchedule(c *gin.Context) {
var req CreateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
response.AbortBadRequest(c, InvalidCronExpression)
return
}
// 校验关联的异步任务类型
meta := task.GetTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
}
// 校验并规范化 Payload
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule := &model.Schedule{
Name: req.Name,
TaskType: req.TaskType,
Cron: req.Cron,
Payload: string(validated),
IsActive: *req.IsActive,
}
if err := model.CreateSchedule(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := scheduler.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OK(schedule))
}
// UpdateScheduleRequest 修改定时任务请求
type UpdateScheduleRequest struct {
Name string `json:"name" binding:"required"`
TaskType string `json:"task_type" binding:"required"`
Cron string `json:"cron" binding:"required"`
Payload string `json:"payload"`
IsActive *bool `json:"is_active" binding:"required"`
}
// UpdateSchedule 修改定时任务
// @Summary 修改定时任务
// @Description 修改一个定时任务的配置(名称、Cron 表达式、异步任务参数和是否启用等),并触发调度器热加载,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Param request body UpdateScheduleRequest true "修改定时任务请求参数"
// @Success 200 {object} response.Any{data=model.Schedule} "修改后的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "定时任务不存在"
// @Failure 500 {object} response.Any "修改定时任务失败"
// @Router /api/v1/admin/tasks/schedules/{id} [put]
func UpdateSchedule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的定时任务ID")
return
}
var req UpdateScheduleRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
// 检查定时任务是否存在
schedule, err := model.GetScheduleByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, ScheduleNotFound)
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
response.AbortBadRequest(c, InvalidCronExpression)
return
}
// 校验关联的异步任务类型
meta := task.GetTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
}
// 校验并规范化 Payload
var payloadBytes []byte
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
schedule.Name = req.Name
schedule.TaskType = req.TaskType
schedule.Cron = req.Cron
schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive
if err := model.UpdateSchedule(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
// 触发调度服务重载
if err := scheduler.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OK(schedule))
}
// DeleteSchedule 删除定时任务
// @Summary 删除定时任务
// @Description 删除指定的定时任务配置,并触发调度器热加载,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Success 200 {object} response.Any{data=string} "删除结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "删除定时任务失败"
// @Router /api/v1/admin/tasks/schedules/{id} [delete]
func DeleteSchedule(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "无效的定时任务ID")
return
}
if err := model.DeleteSchedule(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return
}
// 触发调度服务重载
if err := scheduler.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,495 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
"github.com/Rain-kl/Wavelet/internal/apps/user"
"github.com/Rain-kl/Wavelet/internal/bootstrap"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
func setupTaskTestEnvironment(t *testing.T) func() {
_, mr, cleanup := testhelper.SetupTestEnvironment(t)
bootstrap.RegisterTasks()
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
Addr: mr.Addr(),
})
return func() {
if task.AsynqClient != nil {
_ = task.AsynqClient.Close()
task.AsynqClient = nil
}
cleanup()
}
}
func setupTestRouter(authUser *model.User) *gin.Engine {
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
adminGroup.GET("/tasks/types", ListTaskTypes)
adminGroup.POST("/tasks/dispatch", DispatchTask)
adminGroup.GET("/tasks/executions", ListTaskExecutions)
adminGroup.GET("/tasks/executions/:id", GetTaskExecution)
adminGroup.POST("/tasks/executions/:id/retry", RetryTask)
return r
}
func TestListTaskTypes(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/types", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d", w.Code)
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var taskMetas []task.TaskMeta
_ = json.Unmarshal(dataBytes, &taskMetas)
if len(taskMetas) == 0 {
t.Error("expected at least one dispatchable task type")
}
foundCleanup := false
foundWarmImageCache := false
for _, m := range taskMetas {
if m.Type == uploadtask.TaskTypeSystemCleanup {
foundCleanup = true
}
if m.Type == uploadtask.TaskTypeWarmImageCache {
foundWarmImageCache = true
}
}
if !foundCleanup {
t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeSystemCleanup)
}
if !foundWarmImageCache {
t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeWarmImageCache)
}
}
func TestDispatchTask(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("dispatch valid task successfully", func(t *testing.T) {
payload := DispatchTaskRequest{
TaskType: uploadtask.TaskTypeSystemCleanup,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code, "Body: %s", w.Body.String())
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
assert.Empty(t, resp.ErrorMsg)
assert.NotNil(t, resp.Data)
// 返回的 data 应该是 taskID
taskID, ok := resp.Data.(string)
assert.True(t, ok)
assert.NotEmpty(t, taskID)
})
t.Run("dispatch send_email task successfully with valid payload", func(t *testing.T) {
payload := DispatchTaskRequest{
TaskType: user.TaskTypeSendEmail,
Payload: `{"to":"receiver@example.com","subject":"Test Subject","body":"Test Body"}`,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code, "Body: %s", w.Body.String())
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
assert.Empty(t, resp.ErrorMsg)
assert.NotNil(t, resp.Data)
})
t.Run("dispatch send_email task failure with invalid payload json", func(t *testing.T) {
payload := DispatchTaskRequest{
TaskType: user.TaskTypeSendEmail,
Payload: `{"to":`,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
assert.Contains(t, resp.ErrorMsg, "无效的 JSON 格式")
})
t.Run("dispatch send_email task failure with missing fields", func(t *testing.T) {
payload := DispatchTaskRequest{
TaskType: user.TaskTypeSendEmail,
Payload: `{"to":"","subject":"Test","body":"Test"}`,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
assert.Contains(t, resp.ErrorMsg, "不能为空")
})
t.Run("dispatch invalid task type failure", func(t *testing.T) {
payload := DispatchTaskRequest{
TaskType: "invalid_task_type",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
assert.Equal(t, InvalidTaskType, resp.ErrorMsg)
})
t.Run("dispatch with empty body failure", func(t *testing.T) {
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer([]byte("{}")))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
})
}
func TestListTaskExecutions(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
ctx := context.Background()
// 准备测试数据
now := time.Now()
records := []*model.TaskExecution{
{TaskID: "exec_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual", Duration: 1500, Result: "清理完成", StartedAt: &now, FinishedAt: &now},
{TaskID: "exec_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system", ErrorMessage: "连接超时", StartedAt: &now, FinishedAt: &now},
{TaskID: "exec_003", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", Retryable: true, MaxRetry: 3},
}
for _, r := range records {
err := model.CreateTaskExecution(ctx, r)
require.NoError(t, err)
}
t.Run("list all executions", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var data map[string]interface{}
json.Unmarshal(dataBytes, &data)
assert.Equal(t, float64(3), data["total"])
})
t.Run("filter by status", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?status=failed", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var data map[string]interface{}
json.Unmarshal(dataBytes, &data)
assert.Equal(t, float64(1), data["total"])
})
t.Run("filter by task_type (asynq task name)", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?task_type=system:cleanup", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var data map[string]interface{}
json.Unmarshal(dataBytes, &data)
assert.Equal(t, float64(3), data["total"])
})
t.Run("filter by task_type (management task type)", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?task_type=system_cleanup", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var data map[string]interface{}
json.Unmarshal(dataBytes, &data)
assert.Equal(t, float64(3), data["total"])
})
t.Run("pagination", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?page=1&page_size=2", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var data map[string]interface{}
json.Unmarshal(dataBytes, &data)
assert.Equal(t, float64(3), data["total"])
})
}
func TestGetTaskExecution(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
ctx := context.Background()
// 创建测试记录
execution := &model.TaskExecution{
TaskID: "detail_001",
TaskType: "system:cleanup",
TaskName: "系统垃圾清理",
Status: model.TaskExecutionStatusSucceeded,
Log: "[10:00:01] 开始扫描\n[10:00:02] 找到 50 个文件\n[10:00:03] 清理完成",
Result: "共清理 50 个文件",
Duration: 2000,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
t.Run("get existing execution", func(t *testing.T) {
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d", execution.ID)
req, _ := http.NewRequest("GET", url, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var detail model.TaskExecution
json.Unmarshal(dataBytes, &detail)
assert.Equal(t, "detail_001", detail.TaskID)
assert.Equal(t, model.TaskExecutionStatusSucceeded, detail.Status)
assert.Contains(t, detail.Log, "开始扫描")
assert.Contains(t, detail.Log, "清理完成")
assert.Equal(t, int64(2000), detail.Duration)
})
t.Run("get non-existent execution", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/99999999", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
})
t.Run("invalid ID format", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/invalid", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
})
}
func TestRetryTask(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
ctx := context.Background()
t.Run("retry failed task successfully", func(t *testing.T) {
now := time.Now()
execution := &model.TaskExecution{
TaskID: "retry_api_001",
TaskType: "system:cleanup",
TaskName: "系统垃圾清理",
Status: model.TaskExecutionStatusFailed,
ErrorMessage: "S3 连接超时",
Retryable: true,
MaxRetry: 3,
RetryCount: 0,
TriggeredBy: "manual",
StartedAt: &now,
FinishedAt: &now,
}
err := model.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
req, _ := http.NewRequest("POST", url, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
json.Unmarshal(w.Body.Bytes(), &resp)
assert.Empty(t, resp.ErrorMsg)
assert.NotNil(t, resp.Data)
// 验证新记录
newTaskID, ok := resp.Data.(string)
assert.True(t, ok)
assert.NotEmpty(t, newTaskID)
newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID)
require.NoError(t, err)
assert.Equal(t, 1, newExecution.RetryCount)
assert.Equal(t, "retry", newExecution.TriggeredBy)
})
t.Run("retry succeeded task fails", func(t *testing.T) {
execution := &model.TaskExecution{
TaskID: "retry_succeeded_001",
TaskType: "system:cleanup",
TaskName: "系统垃圾清理",
Status: model.TaskExecutionStatusSucceeded,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
req, _ := http.NewRequest("POST", url, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
})
t.Run("retry non-retryable task fails", func(t *testing.T) {
execution := &model.TaskExecution{
TaskID: "retry_not_allowed_001",
TaskType: "system:cleanup",
TaskName: "系统垃圾清理",
Status: model.TaskExecutionStatusFailed,
Retryable: false,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
req, _ := http.NewRequest("POST", url, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
})
t.Run("retry non-existent task", func(t *testing.T) {
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/99999999/retry", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
})
t.Run("retry with invalid ID", func(t *testing.T) {
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/invalid/retry", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
})
}
@@ -0,0 +1,16 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package template 提供模板管理功能
package template
// 模板管理相关错误消息
const (
TemplateNotFound = "模板不存在"
TemplateKeyRequired = "模板标识符不能为空"
TemplateNameRequired = "模板名称不能为空"
TemplateContentRequired = "模板内容不能为空"
TemplateKeyExists = "模板标识符已存在"
SystemTemplateCannotDelete = "系统预置模板不可删除"
SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符"
)
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package template
import (
"context"
"errors"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
)
func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Template, error) {
exists, err := repository.TemplateExistsByKey(ctx, req.Key)
if err != nil {
return model.Template{}, err
}
if exists {
return model.Template{}, errors.New(TemplateKeyExists)
}
tmpl := model.Template{
Key: req.Key,
Name: req.Name,
Type: req.Type,
Subject: req.Subject,
Content: req.Content,
Description: req.Description,
IsSystem: false,
}
if err := tmpl.Validate(); err != nil {
return model.Template{}, err
}
if err := repository.CreateTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
}
return tmpl, nil
}
func listTemplates(ctx context.Context) ([]model.Template, error) {
return repository.ListTemplates(ctx)
}
func getTemplate(ctx context.Context, key string) (model.Template, error) {
return repository.GetTemplateByKey(ctx, key)
}
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (model.Template, error) {
tmpl, err := repository.GetTemplateByKey(ctx, key)
if err != nil {
return model.Template{}, err
}
tmpl.Name = req.Name
tmpl.Type = req.Type
tmpl.Subject = req.Subject
tmpl.Content = req.Content
tmpl.Description = req.Description
if err := tmpl.Validate(); err != nil {
return model.Template{}, err
}
if err := repository.SaveTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
}
return tmpl, nil
}
func deleteTemplate(ctx context.Context, key string) error {
tmpl, err := repository.GetTemplateByKey(ctx, key)
if err != nil {
return err
}
if tmpl.IsSystem {
return errors.New(SystemTemplateCannotDelete)
}
return repository.DeleteTemplate(ctx, &tmpl)
}
@@ -0,0 +1,176 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package template
import (
"errors"
"net/http"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// CreateTemplateRequest 创建模板请求
type CreateTemplateRequest struct {
Key string `json:"key" binding:"required,max=80"`
Name string `json:"name" binding:"required,max=100"`
Type string `json:"type" binding:"required,max=20"`
Subject string `json:"subject" binding:"max=255"`
Content string `json:"content" binding:"required"`
Description string `json:"description" binding:"max=255"`
}
// UpdateTemplateRequest 更新模板请求
type UpdateTemplateRequest struct {
Name string `json:"name" binding:"required,max=100"`
Type string `json:"type" binding:"required,max=20"`
Subject string `json:"subject" binding:"max=255"`
Content string `json:"content" binding:"required"`
Description string `json:"description" binding:"max=255"`
}
func abortTemplateLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, TemplateNotFound)
return true
}
msg := err.Error()
switch msg {
case TemplateKeyExists, SystemTemplateCannotDelete:
response.AbortBadRequest(c, msg)
return true
}
response.AbortInternal(c, msg)
return true
}
// CreateTemplate 创建模板
// @Summary 创建模板
// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body template.CreateTemplateRequest true "创建请求参数"
// @Success 200 {object} response.Any{data=string} "创建成功"
// @Failure 400 {object} response.Any "参数错误或模板标识符已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [post]
func CreateTemplate(c *gin.Context) {
var req CreateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := createTemplate(c.Request.Context(), req)
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// ListTemplates 获取模板列表
// @Summary 获取模板列表
// @Description 返回所有通知模板列表,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Template} "模板列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [get]
func ListTemplates(c *gin.Context) {
templates, err := listTemplates(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(templates))
}
// GetTemplate 获取单个模板
// @Summary 获取单个模板
// @Description 根据模板标识符获取对应的模板详情,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=model.Template} "模板详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [get]
func GetTemplate(c *gin.Context) {
tmpl, err := getTemplate(c.Request.Context(), c.Param("key"))
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// UpdateTemplate 更新模板
// @Summary 更新模板
// @Description 根据模板标识符更新对应的模板内容,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Param request body template.UpdateTemplateRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=model.Template} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [put]
func UpdateTemplate(c *gin.Context) {
var req UpdateTemplateRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req)
if abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(tmpl))
}
// DeleteTemplate 删除模板
// @Summary 删除模板
// @Description 根据模板标识符删除对应模板,系统预置模板不可删除,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "不可删除系统模板"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [delete]
func DeleteTemplate(c *gin.Context) {
if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,242 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package template
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
func setupTestRouter(authUser *model.User) *gin.Engine {
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
adminGroup.GET("/templates", ListTemplates)
adminGroup.POST("/templates", CreateTemplate)
templateRouter := adminGroup.Group("/templates/:key")
{
templateRouter.GET("", GetTemplate)
templateRouter.PUT("", UpdateTemplate)
templateRouter.DELETE("", DeleteTemplate)
}
return r
}
func TestCreateTemplate(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("create successfully", func(t *testing.T) {
payload := CreateTemplateRequest{
Key: "test_template",
Name: "Test Template",
Type: "email",
Subject: "Test Subject",
Content: "Hello {{.Name}}",
Description: "Test Desc",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/templates", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var tmpl model.Template
err := dbConn.Where("key = ?", "test_template").First(&tmpl).Error
if err != nil {
t.Fatalf("failed to find template in DB: %v", err)
}
if tmpl.Name != "Test Template" {
t.Errorf("expected Name 'Test Template', got '%s'", tmpl.Name)
}
})
t.Run("create duplicate key error", func(t *testing.T) {
payload := CreateTemplateRequest{
Key: "test_template",
Name: "Another Name",
Type: "email",
Subject: "Another Subject",
Content: "Hello",
Description: "desc",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/templates", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request on duplicate key, got %d", w.Code)
}
})
}
func TestListTemplates(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Seed system templates manually for testing
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
t2 := model.Template{Key: "register_email", Name: "Register Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
dbConn.Create(&t1)
dbConn.Create(&t2)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("list templates", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/templates", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d", w.Code)
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var templates []model.Template
_ = json.Unmarshal(dataBytes, &templates)
if len(templates) != 2 {
t.Errorf("expected 2 templates, got %d", len(templates))
}
})
}
func TestGetTemplate(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
dbConn.Create(&t1)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("get existing", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/templates/login_email", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d", w.Code)
}
})
t.Run("get non-existent", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/templates/non_existent", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404 Not Found, got %d", w.Code)
}
})
}
func TestUpdateTemplate(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
dbConn.Create(&t1)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("update successfully", func(t *testing.T) {
payload := UpdateTemplateRequest{
Name: "Updated Login Code",
Type: "email",
Subject: "New Subject",
Content: "new code {{.Code}}",
Description: "new desc",
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/templates/login_email", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var tmpl model.Template
dbConn.Where("key = ?", "login_email").First(&tmpl)
if tmpl.Name != "Updated Login Code" || tmpl.Subject != "New Subject" {
t.Errorf("database values not updated: %+v", tmpl)
}
})
}
func TestDeleteTemplate(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
t2 := model.Template{Key: "custom_tmpl", Name: "Custom", Type: "email", Content: "hi", IsSystem: false}
dbConn.Create(&t1)
dbConn.Create(&t2)
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("delete system template should fail", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/templates/login_email", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request when deleting system template, got %d", w.Code)
}
})
t.Run("delete custom template should succeed", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/templates/custom_tmpl", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d", w.Code)
}
var count int64
dbConn.Model(&model.Template{}).Where("key = ?", "custom_tmpl").Count(&count)
if count != 0 {
t.Error("custom template was not deleted from DB")
}
})
}
@@ -0,0 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package updater manages GitHub Release checks and in-place application upgrades.
package updater
const (
errInvalidRepository = "上游仓库地址无效"
errReleaseRequestFailed = "获取上游版本失败"
errReleaseResponseInvalid = "上游版本响应无效"
errNoCompatibleRelease = "未找到兼容的 Release"
errNoCompatibleAsset = "未找到当前系统对应的 Release 资产"
errDevelopmentBuild = "开发版本无法执行自动升级"
errAlreadyUpToDate = "当前已是最新版本"
errUpgradeAlreadyRunning = "已有升级任务正在执行"
errAutomaticUpgradeBlocked = "当前平台暂不支持自动替换二进制"
)
@@ -0,0 +1,40 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import "context"
// GetStatus returns the current build and the newest compatible upstream release.
func GetStatus(ctx context.Context) (Status, error) {
status, _, err := defaultManager.status(ctx)
return status, err
}
// PrepareUpgrade downloads and stages the upgrade binary for the current platform.
func PrepareUpgrade(ctx context.Context) (executable string, stagedBinary string, status Status, err error) {
status, _, err = defaultManager.status(ctx)
if err != nil {
return "", "", Status{}, err
}
executable, stagedBinary, err = defaultManager.prepareUpgrade(ctx)
return executable, stagedBinary, status, err
}
// ApplyPreparedUpgrade replaces the running binary and restarts the process.
func ApplyPreparedUpgrade(executable, stagedBinary string) error {
return replaceAndRestart(executable, stagedBinary)
}
// FinishUpgrade clears the in-progress upgrade flag after a failed restart.
func FinishUpgrade() {
defaultManager.finishUpgrade()
}
// IsUpgrading reports whether an upgrade task is currently running.
func IsUpgrading() bool {
defaultManager.mu.Lock()
defer defaultManager.mu.Unlock()
return defaultManager.upgrading
}
@@ -0,0 +1,647 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import (
"archive/tar"
"archive/zip"
"compress/gzip"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/buildinfo"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger"
"golang.org/x/mod/semver"
)
const (
githubAPIBaseURL = "https://api.github.com"
maxArchiveSize = int64(1024 * 1024 * 1024)
maxReleaseSize = int64(4 * 1024 * 1024)
repositoryParts = 2
windowsOS = "windows"
archiveFileMode = 0o600
stagedBinaryMode = 0o700
)
type releaseAsset struct {
Name string `json:"name"`
BrowserDownloadURL string `json:"browser_download_url"`
Size int64 `json:"size"`
State string `json:"state"`
}
type githubRelease struct {
TagName string `json:"tag_name"`
Name string `json:"name"`
Body string `json:"body"`
HTMLURL string `json:"html_url"`
Draft bool `json:"draft"`
Prerelease bool `json:"prerelease"`
Published time.Time `json:"published_at"`
Assets []releaseAsset `json:"assets"`
}
// Status describes the current build and the newest compatible upstream release.
type Status struct {
CurrentVersion string `json:"current_version"`
BuildTime string `json:"build_time"`
LatestVersion string `json:"latest_version"`
UpdateAvailable bool `json:"update_available"`
CanUpgrade bool `json:"can_upgrade"`
Prerelease bool `json:"prerelease"`
ReleaseName string `json:"release_name"`
ReleaseNotes string `json:"release_notes"`
ReleaseURL string `json:"release_url"`
PublishedAt string `json:"published_at"`
UpstreamRepository string `json:"upstream_repository"`
AssetName string `json:"asset_name"`
Platform string `json:"platform"`
}
type releaseClient interface {
Do(req *http.Request) (*http.Response, error)
}
type manager struct {
client releaseClient
mu sync.Mutex
upgrading bool
}
var defaultManager = &manager{
client: &http.Client{Timeout: 10 * time.Minute},
}
func normalizeVersion(version string) string {
version = strings.TrimSpace(version)
if version == "" || version == "dev" {
return ""
}
if !strings.HasPrefix(version, "v") {
version = "v" + version
}
if !semver.IsValid(version) {
return ""
}
return version
}
func parseRepository(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", errors.New(errInvalidRepository)
}
if !strings.Contains(raw, "://") {
repo := strings.TrimSuffix(strings.Trim(raw, "/"), ".git")
if len(strings.Split(repo, "/")) == repositoryParts {
return repo, nil
}
return "", errors.New(errInvalidRepository)
}
parsed, err := url.Parse(raw)
if err != nil || !strings.EqualFold(parsed.Hostname(), "github.com") {
return "", errors.New(errInvalidRepository)
}
repo := strings.TrimSuffix(strings.Trim(parsed.Path, "/"), ".git")
if len(strings.Split(repo, "/")) != repositoryParts {
return "", errors.New(errInvalidRepository)
}
return repo, nil
}
func expectedAssetName(tag string) string {
extension := "tar.gz"
if runtime.GOOS == windowsOS {
extension = "zip"
}
return fmt.Sprintf("wavelet_%s_%s_%s.%s", tag, runtime.GOOS, runtime.GOARCH, extension)
}
func expectedAssetNames(repository, tag string) []string {
names := []string{expectedAssetName(tag)}
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
repoName := parts[1]
if repoName != "wavelet" {
extension := "tar.gz"
if runtime.GOOS == windowsOS {
extension = "zip"
}
names = append(names, fmt.Sprintf("%s_%s_%s_%s.%s", repoName, tag, runtime.GOOS, runtime.GOARCH, extension))
}
}
return names
}
func selectLatestRelease(repository string, releases []githubRelease) (githubRelease, releaseAsset, error) {
var selected githubRelease
var selectedAsset releaseAsset
selectedVersion := ""
for _, release := range releases {
version := normalizeVersion(release.TagName)
if release.Draft || version == "" {
continue
}
expectedNames := expectedAssetNames(repository, release.TagName)
for _, asset := range release.Assets {
matched := false
for _, name := range expectedNames {
if asset.Name == name {
matched = true
break
}
}
if !matched || asset.BrowserDownloadURL == "" || asset.State != "uploaded" {
continue
}
if selectedVersion == "" || semver.Compare(version, selectedVersion) > 0 {
selected = release
selectedAsset = asset
selectedVersion = version
}
}
}
if selectedVersion == "" {
return githubRelease{}, releaseAsset{}, errors.New(errNoCompatibleRelease)
}
return selected, selectedAsset, nil
}
func (m *manager) fetchRelease(ctx context.Context, repository string) (githubRelease, releaseAsset, error) {
req, err := http.NewRequestWithContext(
ctx,
http.MethodGet,
fmt.Sprintf("%s/repos/%s/releases?per_page=30", githubAPIBaseURL, repository),
nil,
)
if err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
}
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")
resp, err := m.client.Do(req)
if err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
}
defer func() {
// The response body is read-only; close errors cannot affect the parsed result.
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: HTTP %d", errReleaseRequestFailed, resp.StatusCode)
}
var releases []githubRelease
decoder := json.NewDecoder(io.LimitReader(resp.Body, maxReleaseSize))
if err := decoder.Decode(&releases); err != nil {
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseResponseInvalid, err)
}
release, asset, err := selectLatestRelease(repository, releases)
if err != nil {
return githubRelease{}, releaseAsset{}, err
}
logger.InfoF(ctx, "[Updater] Selected latest compatible release: %s (Asset: %s)", release.TagName, asset.Name)
return release, asset, nil
}
func loadRepository(ctx context.Context) (string, error) {
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository)
if err != nil {
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
}
return parseRepository(config.Value)
}
func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) {
upstreamRepo, err := loadRepository(ctx)
if err != nil {
return Status{}, releaseAsset{}, err
}
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
if err != nil {
return Status{}, releaseAsset{}, err
}
currentVersion := normalizeVersion(buildinfo.Version)
latestVersion := normalizeVersion(release.TagName)
updateAvailable := currentVersion != "" && semver.Compare(latestVersion, currentVersion) > 0
logger.InfoF(ctx, "[Updater] Check update complete. current: %s, latest: %s, update_available: %t", buildinfo.Version, release.TagName, updateAvailable)
return Status{
CurrentVersion: buildinfo.Version,
BuildTime: buildinfo.BuildTime,
LatestVersion: release.TagName,
UpdateAvailable: updateAvailable,
CanUpgrade: updateAvailable && runtime.GOOS != windowsOS,
Prerelease: release.Prerelease,
ReleaseName: release.Name,
ReleaseNotes: release.Body,
ReleaseURL: release.HTMLURL,
PublishedAt: release.Published.Format(time.RFC3339),
UpstreamRepository: upstreamRepo,
AssetName: asset.Name,
Platform: runtime.GOOS + "/" + runtime.GOARCH,
}, asset, nil
}
func downloadArchive(ctx context.Context, client releaseClient, asset releaseAsset, destination string) error {
if asset.Size <= 0 || asset.Size > maxArchiveSize {
return fmt.Errorf("release 资产大小无效: %d", asset.Size)
}
logger.InfoF(ctx, "[Updater] Downloading release asset: %s", asset.Name)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, asset.BrowserDownloadURL, nil)
if err != nil {
return fmt.Errorf("创建升级下载请求失败: %w", err)
}
req.Header.Set("User-Agent", "OpenFlare-Updater")
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf("下载升级资产失败: %w", err)
}
defer func() {
// The downloaded body has already been validated by size before use.
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("下载升级资产失败: HTTP %d", resp.StatusCode)
}
file, err := os.OpenFile(destination, os.O_CREATE|os.O_EXCL|os.O_WRONLY, archiveFileMode) //nolint:gosec // destination is created inside the verified executable directory.
if err != nil {
return fmt.Errorf("创建升级归档失败: %w", err)
}
written, err := io.Copy(file, io.LimitReader(resp.Body, maxArchiveSize+1))
if err != nil {
_ = file.Close()
return fmt.Errorf("写入升级归档失败: %w", err)
}
if err := file.Close(); err != nil {
return fmt.Errorf("关闭升级归档失败: %w", err)
}
if written > maxArchiveSize || written != asset.Size {
return fmt.Errorf("升级归档大小不匹配: got %d, want %d", written, asset.Size)
}
logger.InfoF(ctx, "[Updater] Successfully downloaded release asset to %s", destination)
return nil
}
func safeArchivePath(destination, name string) (string, error) {
cleanName := filepath.Clean(name)
if filepath.IsAbs(cleanName) || cleanName == "." || strings.HasPrefix(cleanName, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("归档包含非法路径: %s", name)
}
target := filepath.Join(destination, cleanName)
relative, err := filepath.Rel(destination, target)
if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("归档路径越界: %s", name)
}
return target, nil
}
func matchBinaryName(name string, candidates []string) bool {
for _, candidate := range candidates {
if runtime.GOOS == windowsOS {
if strings.EqualFold(name, candidate) {
return true
}
} else {
if name == candidate {
return true
}
}
}
return false
}
func getCandidateBinaryNames(executable string, repository string) []string {
execName := filepath.Base(executable)
names := []string{execName}
addName := func(base string) {
name := base
if runtime.GOOS == windowsOS && !strings.HasSuffix(strings.ToLower(name), ".exe") {
name += ".exe"
}
for _, existing := range names {
if existing == name {
return
}
}
names = append(names, name)
}
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
addName(parts[1])
}
addName("wavelet")
return names
}
func isLikelyBinary(name string, isDir bool, mode os.FileMode) bool {
if isDir {
return false
}
base := strings.ToLower(filepath.Base(name))
// Exclude typical non-binary metadata files
exclusions := []string{
"license", "licence", "copying", "notice", "readme", "changelog",
}
for _, excl := range exclusions {
if strings.HasPrefix(base, excl) {
return false
}
}
if runtime.GOOS == windowsOS {
return filepath.Ext(base) == ".exe"
}
// On Unix, it should either have the executable permission bit set, OR have no extension
return (mode.Perm()&0111 != 0) || (filepath.Ext(base) == "")
}
func findBinaryInTarGz(archivePath string, candidates []string) (string, error) {
file, err := os.Open(archivePath) //nolint:gosec // archivePath is created by prepareUpgrade in the executable directory.
if err != nil {
return "", err
}
defer func() {
_ = file.Close()
}()
gzipReader, err := gzip.NewReader(file)
if err != nil {
return "", err
}
defer func() {
_ = gzipReader.Close()
}()
reader := tar.NewReader(gzipReader)
var binaries []string
for {
header, err := reader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return "", err
}
if header.Typeflag == tar.TypeReg && isLikelyBinary(header.Name, false, header.FileInfo().Mode()) {
binaries = append(binaries, header.Name)
}
}
if len(binaries) == 1 {
return binaries[0], nil
}
// Fallback to candidate match if multiple or zero likely binaries found
for _, name := range binaries {
if matchBinaryName(filepath.Base(name), candidates) {
return name, nil
}
}
return "", errors.New(errNoCompatibleAsset)
}
func findBinaryInZip(archivePath string, candidates []string) (string, error) {
reader, err := zip.OpenReader(archivePath)
if err != nil {
return "", err
}
defer func() {
_ = reader.Close()
}()
var binaries []string
for _, file := range reader.File {
if !file.FileInfo().IsDir() && isLikelyBinary(file.Name, false, file.FileInfo().Mode()) {
binaries = append(binaries, file.Name)
}
}
if len(binaries) == 1 {
return binaries[0], nil
}
// Fallback to candidate match if multiple or zero likely binaries found
for _, name := range binaries {
if matchBinaryName(filepath.Base(name), candidates) {
return name, nil
}
}
return "", errors.New(errNoCompatibleAsset)
}
func extractTarGz(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
binaryPathInArchive, err := findBinaryInTarGz(archivePath, candidates)
if err != nil {
return "", err
}
logger.InfoF(ctx, "[Updater] Extracting tar.gz archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
file, err := os.Open(archivePath) //nolint:gosec // archivePath is created by prepareUpgrade in the executable directory.
if err != nil {
return "", err
}
defer func() {
// Read-only archive close errors do not change extraction validity.
_ = file.Close()
}()
gzipReader, err := gzip.NewReader(file)
if err != nil {
return "", err
}
defer func() {
// The gzip checksum is verified while reading the selected file.
_ = gzipReader.Close()
}()
reader := tar.NewReader(gzipReader)
for {
header, err := reader.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return "", err
}
if header.Name != binaryPathInArchive {
continue
}
target, err := safeArchivePath(destination, targetName)
if err != nil {
return "", err
}
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode) //nolint:gosec // target is constrained by safeArchivePath.
if err != nil {
return "", err
}
written, copyErr := io.Copy(output, io.LimitReader(reader, maxArchiveSize+1))
closeErr := output.Close()
if copyErr != nil {
return "", copyErr
}
if closeErr != nil {
return "", closeErr
}
if written > maxArchiveSize {
return "", errors.New("解压后的程序文件超过大小限制")
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errNoCompatibleAsset)
}
func extractZip(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
binaryPathInArchive, err := findBinaryInZip(archivePath, candidates)
if err != nil {
return "", err
}
logger.InfoF(ctx, "[Updater] Extracting zip archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
reader, err := zip.OpenReader(archivePath)
if err != nil {
return "", err
}
defer func() {
// Read-only archive close errors do not change extraction validity.
_ = reader.Close()
}()
for _, file := range reader.File {
if file.Name != binaryPathInArchive {
continue
}
target, err := safeArchivePath(destination, targetName)
if err != nil {
return "", err
}
input, err := file.Open()
if err != nil {
return "", err
}
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode) //nolint:gosec // target is constrained by safeArchivePath.
if err != nil {
_ = input.Close()
return "", err
}
written, copyErr := io.Copy(output, io.LimitReader(input, maxArchiveSize+1))
inputCloseErr := input.Close()
outputCloseErr := output.Close()
if copyErr != nil {
return "", copyErr
}
if inputCloseErr != nil {
return "", inputCloseErr
}
if outputCloseErr != nil {
return "", outputCloseErr
}
if written > maxArchiveSize {
return "", errors.New("解压后的程序文件超过大小限制")
}
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
return target, nil
}
return "", errors.New(errNoCompatibleAsset)
}
func (m *manager) prepareUpgrade(ctx context.Context) (string, string, error) {
if runtime.GOOS == windowsOS {
return "", "", errors.New(errAutomaticUpgradeBlocked)
}
if normalizeVersion(buildinfo.Version) == "" {
return "", "", errors.New(errDevelopmentBuild)
}
m.mu.Lock()
defer m.mu.Unlock()
if m.upgrading {
return "", "", errors.New(errUpgradeAlreadyRunning)
}
status, asset, err := m.status(ctx)
if err != nil {
return "", "", err
}
if !status.UpdateAvailable {
return "", "", errors.New(errAlreadyUpToDate)
}
logger.InfoF(ctx, "[Updater] Preparing upgrade. current: %s, latest: %s", status.CurrentVersion, status.LatestVersion)
executable, err := os.Executable()
if err != nil {
return "", "", fmt.Errorf("定位当前程序失败: %w", err)
}
executable, err = filepath.EvalSymlinks(executable)
if err != nil {
return "", "", fmt.Errorf("解析当前程序路径失败: %w", err)
}
tempDir, err := os.MkdirTemp(filepath.Dir(executable), ".wavelet-update-*")
if err != nil {
return "", "", fmt.Errorf("创建升级目录失败: %w", err)
}
archivePath := filepath.Join(tempDir, asset.Name)
if err := downloadArchive(ctx, m.client, asset, archivePath); err != nil {
// Cleanup is best effort because the download error is the actionable failure.
_ = os.RemoveAll(tempDir)
return "", "", err
}
targetName := filepath.Base(executable)
candidates := getCandidateBinaryNames(executable, status.UpstreamRepository)
var stagedBinary string
if strings.HasSuffix(asset.Name, ".zip") {
stagedBinary, err = extractZip(ctx, archivePath, tempDir, targetName, candidates)
} else {
stagedBinary, err = extractTarGz(ctx, archivePath, tempDir, targetName, candidates)
}
if err != nil {
// Cleanup is best effort because the extraction error is the actionable failure.
_ = os.RemoveAll(tempDir)
return "", "", fmt.Errorf("解压升级资产失败: %w", err)
}
logger.InfoF(ctx, "[Updater] Staged binary successfully prepared: %s", stagedBinary)
m.upgrading = true
return executable, stagedBinary, nil
}
func (m *manager) finishUpgrade() {
m.mu.Lock()
defer m.mu.Unlock()
m.upgrading = false
}
@@ -0,0 +1,130 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import (
"runtime"
"testing"
"time"
)
func TestParseRepository(t *testing.T) {
tests := []struct {
name string
input string
want string
wantErr bool
}{
{name: "short form", input: "Rain-kl/OpenFlare", want: "Rain-kl/OpenFlare"},
{name: "GitHub URL", input: "https://github.com/Rain-kl/OpenFlare.git", want: "Rain-kl/OpenFlare"},
{name: "unsupported host", input: "https://example.com/Rain-kl/OpenFlare", wantErr: true},
{name: "missing owner", input: "OpenFlare", wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := parseRepository(tt.input)
if gotErr := err != nil; gotErr != tt.wantErr {
t.Errorf("parseRepository(%q) error = %v, want error presence = %t", tt.input, err, tt.wantErr)
}
if got != tt.want {
t.Errorf("parseRepository(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
func TestSelectLatestRelease(t *testing.T) {
assetNameV1 := expectedAssetName("v1.0.0")
assetNameV2 := expectedAssetName("v2.0.0")
releases := []githubRelease{
{
TagName: "v1.0.0",
Published: time.Date(2026, time.June, 1, 0, 0, 0, 0, time.UTC),
Assets: []releaseAsset{{
Name: assetNameV1,
BrowserDownloadURL: "https://example.com/v1",
State: "uploaded",
}},
},
{
TagName: "v2.0.0",
Published: time.Date(2026, time.June, 2, 0, 0, 0, 0, time.UTC),
Assets: []releaseAsset{{
Name: assetNameV2,
BrowserDownloadURL: "https://example.com/v2",
State: "uploaded",
}},
},
{
TagName: "v3.0.0",
Assets: []releaseAsset{{
Name: "wavelet_v3.0.0_other_platform.tar.gz",
BrowserDownloadURL: "https://example.com/v3",
State: "uploaded",
}},
},
}
release, asset, err := selectLatestRelease("Rain-kl/OpenFlare", releases)
if err != nil {
t.Fatalf("selectLatestRelease() error = %v", err)
}
if release.TagName != "v2.0.0" {
t.Errorf("selectLatestRelease() tag = %q, want %q", release.TagName, "v2.0.0")
}
if asset.Name != assetNameV2 {
t.Errorf("selectLatestRelease() asset = %q, want %q", asset.Name, assetNameV2)
}
}
func TestSelectLatestReleaseWithCustomRepo(t *testing.T) {
extension := "tar.gz"
if runtime.GOOS == "windows" {
extension = "zip"
}
releases := []githubRelease{
{
TagName: "v1.0.0",
Published: time.Date(2026, time.June, 1, 0, 0, 0, 0, time.UTC),
Assets: []releaseAsset{{
Name: "wavelet_v1.0.0_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension,
BrowserDownloadURL: "https://example.com/v1",
State: "uploaded",
}},
},
{
TagName: "v2.0.0",
Published: time.Date(2026, time.June, 2, 0, 0, 0, 0, time.UTC),
Assets: []releaseAsset{{
Name: "PixezSync_v2.0.0_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension,
BrowserDownloadURL: "https://example.com/v2",
State: "uploaded",
}},
},
}
release, asset, err := selectLatestRelease("Rain-kl/PixezSync", releases)
if err != nil {
t.Fatalf("selectLatestRelease() error = %v", err)
}
if release.TagName != "v2.0.0" {
t.Errorf("selectLatestRelease() tag = %q, want %q", release.TagName, "v2.0.0")
}
expectedName := "PixezSync_v2.0.0_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension
if asset.Name != expectedName {
t.Errorf("selectLatestRelease() asset = %q, want %q", asset.Name, expectedName)
}
}
func TestExpectedAssetName(t *testing.T) {
extension := "tar.gz"
if runtime.GOOS == "windows" {
extension = "zip"
}
want := "wavelet_v1.2.3_" + runtime.GOOS + "_" + runtime.GOARCH + "." + extension
if got := expectedAssetName("v1.2.3"); got != want {
t.Errorf("expectedAssetName(%q) = %q, want %q", "v1.2.3", got, want)
}
}
@@ -0,0 +1,50 @@
//go:build !windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import (
"context"
"fmt"
"os"
"path/filepath"
"syscall"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
const installedBinaryMode = 0o755
func replaceAndRestart(executable, stagedBinary string) error {
ctx := context.Background()
logger.InfoF(ctx, "[Updater] Swapping executable: %s -> %s", executable, stagedBinary)
backup := executable + ".old"
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("删除旧备份失败: %w", err)
}
if err := os.Rename(executable, backup); err != nil {
return fmt.Errorf("备份当前程序失败: %w", err)
}
if err := os.Rename(stagedBinary, executable); err != nil {
_ = os.Rename(backup, executable)
return fmt.Errorf("替换当前程序失败: %w", err)
}
if err := os.Chmod(executable, installedBinaryMode); err != nil { //nolint:gosec // the installed application binary must be executable.
_ = os.Remove(executable)
_ = os.Rename(backup, executable)
return fmt.Errorf("设置程序执行权限失败: %w", err)
}
stagingDir := filepath.Dir(stagedBinary)
// Cleanup is best effort; a leftover staging directory must not block restart.
_ = os.RemoveAll(stagingDir)
logger.InfoF(ctx, "[Updater] Executing syscall.Exec to restart service: %s %v", executable, os.Args)
return syscall.Exec(executable, os.Args, os.Environ()) //nolint:gosec // executable is resolved from os.Executable and never supplied by a request.
}
@@ -0,0 +1,12 @@
//go:build windows
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import "errors"
func replaceAndRestart(_, _ string) error {
return errors.New(errAutomaticUpgradeBlocked)
}
@@ -0,0 +1,68 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package updater
import (
"context"
"net/http"
"time"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// GetUpdateStatus 获取应用更新状态
// @Summary 获取应用更新状态
// @Description 从系统配置指定的 GitHub 上游仓库查询最新兼容 Release,并与当前服务版本比较
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=updater.Status} "更新状态"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "查询失败"
// @Router /api/v1/admin/update [get]
func GetUpdateStatus(c *gin.Context) {
status, _, err := defaultManager.status(c.Request.Context())
if err != nil {
logger.ErrorF(c.Request.Context(), "[Updater] check release failed: %v", err)
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(status))
}
// ApplyUpdate 下载并应用应用更新
// @Summary 下载并应用应用更新
// @Description 下载当前平台对应的 GitHub Actions Release 资产,替换当前二进制并重启进程
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any "升级已准备并即将重启"
// @Failure 400 {object} response.Any "当前版本不可升级"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "升级准备失败"
// @Router /api/v1/admin/update/apply [post]
func ApplyUpdate(c *gin.Context) {
executable, stagedBinary, err := defaultManager.prepareUpgrade(c.Request.Context())
if err != nil {
logger.ErrorF(c.Request.Context(), "[Updater] prepare upgrade failed: %v", err)
response.AbortBadRequest(c, err.Error())
return
}
logger.InfoF(c.Request.Context(), "[Updater] upgrade prepared; restarting with %s", stagedBinary)
c.JSON(http.StatusOK, response.OKNil())
go func() {
time.Sleep(time.Second)
if err := replaceAndRestart(executable, stagedBinary); err != nil {
defaultManager.finishUpgrade()
logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err)
}
}()
}
@@ -0,0 +1,21 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package user 提供用户管理功能
package user
const (
userNotFound = "用户不存在"
cannotDisable = "不能禁用管理员用户"
cannotDelete = "不能删除管理员用户"
cannotDeleteSelf = "不能删除当前登录用户"
updateUserFailed = "更新用户状态失败"
deleteUserFailed = "删除用户失败"
usernameExists = "用户名已存在"
usernameRequired = "用户名不能为空"
passwordTooShort = "密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
createUserFailed = "创建用户失败"
emailRequired = "邮箱不能为空"
emailExists = "邮箱已被注册"
)
@@ -0,0 +1,106 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package user
import (
"context"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
)
func listUsers(ctx context.Context, req listUsersRequest) (int64, []model.User, error) {
return repository.ListAdminUsers(ctx, repository.AdminUserListFilter{
UserID: req.UserID,
Username: strings.TrimSpace(req.Username),
Page: req.Page,
PageSize: req.PageSize,
})
}
func getUserDetail(ctx context.Context, id uint64) (model.User, error) {
return repository.GetAdminUserDetail(ctx, id)
}
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
flags, err := repository.GetUserAdminFlags(ctx, id)
if err != nil {
return err
}
if !active && flags.IsAdmin {
return errors.New(cannotDisable)
}
return repository.UpdateUserActive(ctx, id, active)
}
func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
if currentUserID == targetID {
return errors.New(cannotDeleteSelf)
}
flags, err := repository.GetUserAdminFlags(ctx, targetID)
if err != nil {
return err
}
if flags.IsAdmin {
return errors.New(cannotDelete)
}
return repository.DeleteUserWithRelations(ctx, targetID)
}
func createUser(ctx context.Context, req createUserRequest) (model.User, error) {
req.Username = strings.TrimSpace(req.Username)
req.Nickname = strings.TrimSpace(req.Nickname)
req.Password = strings.TrimSpace(req.Password)
req.Email = strings.TrimSpace(req.Email)
if req.Username == "" {
return model.User{}, errors.New(usernameRequired)
}
if req.Email == "" {
return model.User{}, errors.New(emailRequired)
}
if len(req.Password) < minPasswordLength {
return model.User{}, errors.New(passwordTooShort)
}
count, err := repository.CountUsersByUsername(ctx, req.Username)
if err != nil {
return model.User{}, err
}
if count > 0 {
return model.User{}, errors.New(usernameExists)
}
emailCount, err := repository.CountUsersByEmail(ctx, req.Email)
if err != nil {
return model.User{}, err
}
if emailCount > 0 {
return model.User{}, errors.New(emailExists)
}
newUser := model.User{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
LastLoginAt: time.Time{},
}
if newUser.Nickname == "" {
newUser.Nickname = req.Username
}
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
return model.User{}, err
}
if err := repository.CreateUser(ctx, &newUser); err != nil {
return model.User{}, err
}
return newUser, nil
}
@@ -0,0 +1,288 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package user
import (
"errors"
"net/http"
"strconv"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// minPasswordLength 密码最小长度
const minPasswordLength = 8
// listUsersRequest 用户列表查询请求
type listUsersRequest struct {
Page int `form:"page" binding:"min=1"`
PageSize int `form:"page_size" binding:"min=1,max=100"`
UserID *uint64 `form:"user_id" binding:"omitempty,gt=0"`
Username string `form:"username"`
}
type user struct {
ID uint64 `json:"id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsActive bool `json:"is_active"`
IsAdmin bool `json:"is_admin"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
LastLoginAt time.Time `json:"last_login_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// listUsersResponse 用户列表响应
type listUsersResponse struct {
Users []user `json:"users"`
Total int64 `json:"total"`
}
func parseUserID(c *gin.Context) (uint64, bool) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, userNotFound)
return 0, false
}
return id, true
}
func toUser(u model.User) user {
return user{
ID: u.ID,
Username: u.Username,
Nickname: u.Nickname,
Email: u.Email,
AvatarURL: u.AvatarURL,
IsActive: u.IsActive,
IsAdmin: u.IsAdmin,
Bio: u.Bio,
Phone: u.Phone,
Gender: u.Gender,
Website: u.Website,
Location: u.Location,
LastLoginAt: u.LastLoginAt,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
}
}
func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbiddenMsgs, badRequestMsgs []string) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, notFoundMsg)
return true
}
msg := err.Error()
for _, m := range badRequestMsgs {
if msg == m {
response.AbortBadRequest(c, msg)
return true
}
}
for _, m := range forbiddenMsgs {
if msg == m {
response.AbortForbidden(c, msg)
return true
}
}
response.AbortInternal(c, msg)
return true
}
// ListUsers 获取用户列表
// @Summary 获取用户列表
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param request query listUsersRequest true "查询参数"
// @Success 200 {object} response.Any{data=user.listUsersResponse} "用户列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users [get]
func ListUsers(c *gin.Context) {
var req listUsersRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
total, modelUsers, err := listUsers(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
users := make([]user, 0, len(modelUsers))
for _, modelUser := range modelUsers {
users = append(users, toUser(modelUser))
}
c.JSON(http.StatusOK, response.OK(listUsersResponse{
Users: users,
Total: total,
}))
}
// GetUser 获取用户详情
// @Summary 获取用户详情
// @Description 返回指定用户的完整个人资料和系统状态,需要管理员权限,不返回密码等敏感字段
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Success 200 {object} response.Any{data=user.user} "用户详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [get]
func GetUser(c *gin.Context) {
id, ok := parseUserID(c)
if !ok {
return
}
targetUser, err := getUserDetail(c.Request.Context(), id)
if abortUserLogicError(c, err, userNotFound, nil, nil) {
return
}
c.JSON(http.StatusOK, response.OK(toUser(targetUser)))
}
// updateUserStatusRequest 更新用户状态请求
type updateUserStatusRequest struct {
IsActive bool `json:"is_active"`
}
// UpdateUserStatus 更新用户状态(启用/禁用)
// @Summary 更新用户状态
// @Description 启用或禁用指定用户,管理员账号无法被禁用,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Param request body updateUserStatusRequest true "状态参数"
// @Success 200 {object} response.Any{data=string} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限或尝试禁用管理员"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id}/status [put]
func UpdateUserStatus(c *gin.Context) {
var req updateUserStatusRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
id, ok := parseUserID(c)
if !ok {
return
}
if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) {
return
}
response.AbortInternal(c, updateUserFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// DeleteUser 删除用户
// @Summary 删除用户
// @Description 删除指定非管理员用户,需要管理员权限,不能删除当前登录用户
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param id path int true "用户 ID"
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限、尝试删除管理员或当前用户"
// @Failure 404 {object} response.Any "用户不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users/{id} [delete]
func DeleteUser(c *gin.Context) {
id, ok := parseUserID(c)
if !ok {
return
}
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
if err := deleteUser(c.Request.Context(), currUser.ID, id); err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) {
return
}
response.AbortInternal(c, deleteUserFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// createUserRequest 创建用户请求
type createUserRequest struct {
Username string `json:"username" binding:"required,min=3,max=64"`
Password string `json:"password" binding:"required,min=8,max=64"`
Nickname string `json:"nickname" binding:"omitempty,max=64"`
Email string `json:"email" binding:"required,email,max=255"`
IsActive bool `json:"is_active"`
IsAdmin bool `json:"is_admin"`
}
// CreateUser 创建用户
// @Summary 创建用户
// @Description 创建一个本地密码登录的新用户,需要管理员权限
// @Tags admin
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param request body user.createUserRequest true "创建用户参数"
// @Success 200 {object} response.Any{data=user.user} "创建成功"
// @Failure 400 {object} response.Any "参数错误或用户名已存在"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users [post]
func CreateUser(c *gin.Context) {
var req createUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
newUser, err := createUser(c.Request.Context(), req)
if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) {
return
}
c.JSON(http.StatusOK, response.OK(toUser(newUser)))
}
@@ -0,0 +1,582 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package user
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
func setupTestRouter(authUser *model.User) *gin.Engine {
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
adminGroup.GET("/users", ListUsers)
adminGroup.POST("/users", CreateUser)
adminGroup.GET("/users/:id", GetUser)
adminGroup.PUT("/users/:id/status", UpdateUserStatus)
adminGroup.DELETE("/users/:id", DeleteUser)
return r
}
func TestListUsers(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Seed users
users := []model.User{
{
ID: 1001,
Username: "alice",
Nickname: "Alice Nickname",
IsActive: true,
IsAdmin: false,
LastLoginAt: time.Now(),
},
{
ID: 1002,
Username: "bob",
Nickname: "Bob Nickname",
IsActive: true,
IsAdmin: false,
LastLoginAt: time.Now(),
},
{
ID: 1003,
Username: "charlie",
Nickname: "Charlie Nickname",
IsActive: false,
IsAdmin: true,
LastLoginAt: time.Now(),
},
}
for _, u := range users {
if err := dbConn.Create(&u).Error; err != nil {
t.Fatalf("failed to seed user: %v", err)
}
}
adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("basic pagination list", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=2", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var resp response.Any
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse response: %v", err)
}
// Parse data map to our structure
dataBytes, _ := json.Marshal(resp.Data)
var listResp listUsersResponse
if err := json.Unmarshal(dataBytes, &listResp); err != nil {
t.Fatalf("failed to parse list response: %v", err)
}
if len(listResp.Users) != 2 {
t.Errorf("expected 2 users, got %d", len(listResp.Users))
}
if listResp.Total != 3 {
t.Errorf("expected total 3, got %d", listResp.Total)
}
// Ordered by ID ASC
if listResp.Users[0].ID != 1001 || listResp.Users[1].ID != 1002 {
t.Errorf("expected ordered ASC, got first ID %d, second ID %d", listResp.Users[0].ID, listResp.Users[1].ID)
}
})
t.Run("filter by user_id", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&user_id=1001", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var listResp listUsersResponse
_ = json.Unmarshal(dataBytes, &listResp)
if len(listResp.Users) != 1 || listResp.Users[0].ID != 1001 {
t.Errorf("expected 1 user with ID 1001, got total %d", len(listResp.Users))
}
})
t.Run("filter by username prefix", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&username=bo", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
dataBytes, _ := json.Marshal(resp.Data)
var listResp listUsersResponse
_ = json.Unmarshal(dataBytes, &listResp)
if len(listResp.Users) != 1 || listResp.Users[0].Username != "bob" {
t.Errorf("expected bob, got %v", listResp.Users)
}
})
t.Run("invalid pagination parameter", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=0&page_size=10", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request, got %d", w.Code)
}
})
}
func TestGetUser(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
targetUser := model.User{
ID: 1001,
Username: "alice",
Password: "secret-hash",
Nickname: "Alice Nickname",
Email: "alice@example.com",
AvatarURL: "https://example.com/avatar.png",
IsActive: true,
IsAdmin: false,
Bio: "hello",
Phone: "123456",
Gender: "female",
Website: "https://example.com",
Location: "Shanghai",
}
if err := dbConn.Create(&targetUser).Error; err != nil {
t.Fatalf("failed to seed user: %v", err)
}
adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("get full user profile", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/users/1001", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var resp response.Any
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse response: %v", err)
}
dataBytes, _ := json.Marshal(resp.Data)
var resUser user
if err := json.Unmarshal(dataBytes, &resUser); err != nil {
t.Fatalf("failed to parse response data: %v", err)
}
if resUser.Email != targetUser.Email || resUser.Bio != targetUser.Bio || resUser.Phone != targetUser.Phone ||
resUser.Gender != targetUser.Gender || resUser.Website != targetUser.Website || resUser.Location != targetUser.Location {
t.Errorf("profile fields were not returned correctly: %+v", resUser)
}
if bytes.Contains(dataBytes, []byte("secret-hash")) {
t.Error("response should not include password")
}
})
t.Run("get non-existent user", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/users/9999", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String())
}
})
}
func TestUpdateUserStatus(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Seed users
regularUser := model.User{
ID: 1001,
Username: "alice",
IsActive: true,
IsAdmin: false,
}
adminUser := model.User{
ID: 1002,
Username: "bob",
IsActive: true,
IsAdmin: true,
}
dbConn.Create(&regularUser)
dbConn.Create(&adminUser)
router := setupTestRouter(&adminUser)
t.Run("disable regular user successfully", func(t *testing.T) {
payload := updateUserStatusRequest{IsActive: false}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/users/1001/status", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
// Verify DB status
var u model.User
dbConn.First(&u, 1001)
if u.IsActive {
t.Error("user should be deactivated in the database")
}
})
t.Run("cannot disable admin user", func(t *testing.T) {
payload := updateUserStatusRequest{IsActive: false}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/users/1002/status", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Errorf("expected 403 Forbidden, got %d. Body: %s", w.Code, w.Body.String())
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
if resp.ErrorMsg != cannotDisable {
t.Errorf("expected error message '%s', got '%s'", cannotDisable, resp.ErrorMsg)
}
})
t.Run("cannot enable/disable non-existent user", func(t *testing.T) {
payload := updateUserStatusRequest{IsActive: false}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("PUT", "/api/v1/admin/users/9999/status", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String())
}
})
}
func TestCreateUser(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true}
router := setupTestRouter(adminUser)
t.Run("create user successfully", func(t *testing.T) {
payload := createUserRequest{
Username: "newuser",
Password: "newpassword123",
Nickname: "New Nickname",
Email: "newuser@example.com",
IsActive: true,
IsAdmin: false,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var resp response.Any
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse response: %v", err)
}
if resp.ErrorMsg != "" {
t.Errorf("expected empty error message, got '%s'", resp.ErrorMsg)
}
dataBytes, _ := json.Marshal(resp.Data)
var resUser user
if err := json.Unmarshal(dataBytes, &resUser); err != nil {
t.Fatalf("failed to parse response data: %v", err)
}
if resUser.Username != "newuser" || resUser.Nickname != "New Nickname" || !resUser.IsActive || resUser.IsAdmin {
t.Errorf("unexpected user values: %+v", resUser)
}
// Verify in DB
var dbUser model.User
if err := dbConn.Where("username = ?", "newuser").First(&dbUser).Error; err != nil {
t.Fatalf("failed to find user in db: %v", err)
}
if dbUser.Email != "newuser@example.com" {
t.Errorf("expected email 'newuser@example.com', got '%s'", dbUser.Email)
}
if !dbUser.CheckPassword("newpassword123") {
t.Error("password was not hashed correctly")
}
})
t.Run("create user with duplicate username", func(t *testing.T) {
// Create the first user
existing := model.User{
ID: 2001,
Username: "dupuser",
Nickname: "Dup User",
Email: "dupuser@example.com",
}
dbConn.Create(&existing)
payload := createUserRequest{
Username: "dupuser",
Password: "password123",
Nickname: "Another Nick",
Email: "another@example.com",
IsActive: true,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String())
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
if resp.ErrorMsg != usernameExists {
t.Errorf("expected error '%s', got '%s'", usernameExists, resp.ErrorMsg)
}
})
t.Run("create user with duplicate email", func(t *testing.T) {
existing := model.User{
ID: 2002,
Username: "existingemail",
Nickname: "Existing Email",
Email: "dupemail@example.com",
}
dbConn.Create(&existing)
payload := createUserRequest{
Username: "newuser2",
Password: "password123",
Nickname: "New User 2",
Email: "dupemail@example.com",
IsActive: true,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String())
}
var resp response.Any
_ = json.Unmarshal(w.Body.Bytes(), &resp)
if resp.ErrorMsg != emailExists {
t.Errorf("expected error '%s', got '%s'", emailExists, resp.ErrorMsg)
}
})
t.Run("validation error - password too short", func(t *testing.T) {
payload := createUserRequest{
Username: "shortpass",
Password: "123",
Email: "shortpass@example.com",
IsActive: true,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String())
}
})
t.Run("validation error - invalid email format", func(t *testing.T) {
payload := map[string]interface{}{
"username": "bademail",
"password": "password123",
"email": "not-an-email",
"is_active": true,
}
body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/users", bytes.NewBuffer(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Errorf("expected 400 Bad Request, got %d. Body: %s", w.Code, w.Body.String())
}
})
}
func TestDeleteUser(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
if err := dbConn.AutoMigrate(&model.AccessToken{}, &model.ExternalAccount{}); err != nil {
t.Fatalf("failed to migrate delete-related tables: %v", err)
}
regularUser := model.User{
ID: 1001,
Username: "alice",
IsActive: true,
IsAdmin: false,
}
adminUser := model.User{
ID: 1002,
Username: "bob",
IsActive: true,
IsAdmin: true,
}
selfUser := model.User{
ID: 1003,
Username: "charlie",
IsActive: true,
IsAdmin: false,
}
if err := dbConn.Create(&regularUser).Error; err != nil {
t.Fatalf("failed to seed regular user: %v", err)
}
if err := dbConn.Create(&adminUser).Error; err != nil {
t.Fatalf("failed to seed admin user: %v", err)
}
if err := dbConn.Create(&selfUser).Error; err != nil {
t.Fatalf("failed to seed self user: %v", err)
}
if err := dbConn.Create(&model.AccessToken{
UserID: regularUser.ID,
Name: "api",
TokenHash: "hash-for-delete-user-test",
MaskedToken: "at_****test",
}).Error; err != nil {
t.Fatalf("failed to seed access token: %v", err)
}
if err := dbConn.Create(&model.ExternalAccount{
ID: 5001,
AuthSourceID: 1,
UserID: regularUser.ID,
ExternalID: "external-alice",
}).Error; err != nil {
t.Fatalf("failed to seed external account: %v", err)
}
router := setupTestRouter(&selfUser)
t.Run("delete regular user successfully", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/1001", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var count int64
if err := dbConn.Model(&model.User{}).Where("id = ?", 1001).Count(&count).Error; err != nil {
t.Fatalf("failed to count deleted user: %v", err)
}
if count != 0 {
t.Errorf("expected deleted user count 0, got %d", count)
}
if err := dbConn.Model(&model.AccessToken{}).Where("user_id = ?", 1001).Count(&count).Error; err != nil {
t.Fatalf("failed to count deleted access tokens: %v", err)
}
if count != 0 {
t.Errorf("expected deleted access token count 0, got %d", count)
}
if err := dbConn.Model(&model.ExternalAccount{}).Where("user_id = ?", 1001).Count(&count).Error; err != nil {
t.Fatalf("failed to count deleted external accounts: %v", err)
}
if count != 0 {
t.Errorf("expected deleted external account count 0, got %d", count)
}
})
t.Run("cannot delete admin user", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/1002", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Fatalf("expected 403 Forbidden, got %d. Body: %s", w.Code, w.Body.String())
}
})
t.Run("cannot delete current user", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/1003", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Fatalf("expected 403 Forbidden, got %d. Body: %s", w.Code, w.Body.String())
}
})
t.Run("delete non-existent user", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/admin/users/9999", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String())
}
})
}
@@ -0,0 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap 提供人机验证中间件
package cap
const (
errCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
)
@@ -0,0 +1,196 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap provides CAPTCHA and proof-of-work (PoW) verification services.
package cap
import (
"context"
"crypto/sha256"
"encoding/hex"
"strconv"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
)
const (
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
tokenPartsCount = 2 // 兑换 Token 由两部分组成
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
)
// Manager orchestrates challenge generation and solution validation.
type Manager struct {
secret []byte
store pkgcap.Store
}
// NewManager creates a new CAPTCHA Manager.
func NewManager(secret []byte, store pkgcap.Store) *Manager {
return &Manager{
secret: secret,
store: store,
}
}
// Generate creates a challenge response.
func (m *Manager) Generate(ctx context.Context, scope string) (*pkgcap.ChallengeResponse, error) {
settings, err := CurrentSettings(ctx)
if err != nil {
return nil, err
}
challengeConfig := pkgcap.ChallengeConfig{
Count: settings.ChallengeCount,
Size: settings.ChallengeSize,
Difficulty: settings.ChallengeDifficulty,
Expires: settings.ChallengeTTL,
}
return pkgcap.GenerateChallenge(m.secret, challengeConfig, scope)
}
// RedeemResponse is returned to the client on redeem.
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
// Redeem verifies PoW solutions and returns a one-time redeem token.
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
sigHex := pkgcap.JwtSigHex(token)
if sigHex == "" {
return &RedeemResponse{Success: false, Error: "invalid_token"}, nil
}
nonceKey := "cap:nonce:" + sigHex
payload, err := pkgcap.VerifyChallengeSolutions(token, solutions, m.secret, scope)
if err != nil {
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors
}
now := time.Now().UnixNano() / int64(time.Millisecond)
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
if nonceTTL < time.Second {
nonceTTL = time.Second
}
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
if err != nil {
return &RedeemResponse{Success: false, Error: "nonce_store_error"}, err
}
if !set {
return &RedeemResponse{Success: false, Error: "already_redeemed"}, nil
}
settings, err := CurrentSettings(ctx)
if err != nil {
return &RedeemResponse{Success: false, Error: "settings_load_error"}, err
}
id := pkgcap.RandomHex(redeemTokenIDLength)
verToken := pkgcap.RandomHex(redeemVerTokenLength)
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
tokenKey := "cap:token:" + id + ":" + verHashHex
tokenExpires := time.Now().Add(settings.TokenTTL)
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
return &RedeemResponse{Success: false, Error: "token_store_error"}, err
}
return &RedeemResponse{
Success: true,
Token: id + ":" + verToken,
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
}, nil
}
// VerifyToken validates and consumes the redeem token (single-use).
func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope string) (bool, error) {
if token == "" {
return false, nil
}
parts := strings.Split(token, ":")
if len(parts) != tokenPartsCount {
return false, nil
}
id := parts[0]
verToken := parts[1]
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
tokenKey := "cap:token:" + id + ":" + verHashHex
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
if err != nil {
return false, err
}
if !exists {
return false, nil
}
valParts := strings.Split(val, "|")
if len(valParts) != valuePartsCount {
return false, nil
}
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
if err != nil {
return false, nil //nolint:nilerr // invalid format is treated as validation failure
}
tokenScope := valParts[1]
if expectedScope != "" && tokenScope != expectedScope {
return false, nil
}
if time.Now().UnixNano() > expNano {
return false, nil
}
return true, nil
}
func sGetAndDelete(ctx context.Context, store pkgcap.Store, key string) (string, bool, error) {
if store == nil {
return "", false, nil
}
return store.GetAndDelete(ctx, key)
}
var (
defaultManager *Manager
once sync.Once
)
// GetDefaultManager yields the global singleton CAPTCHA manager.
func GetDefaultManager() *Manager {
once.Do(func() {
secret := []byte("default-captcha-secret-key-at-least-16-bytes")
if config.Config != nil && config.Config.App.SessionSecret != "" {
secret = []byte(config.Config.App.SessionSecret)
}
var store pkgcap.Store
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
store = pkgcap.NewRedisStore(db.Redis)
} else {
store = pkgcap.NewMemoryStore(1 * time.Minute)
}
defaultManager = NewManager(secret, store)
})
return defaultManager
}
@@ -0,0 +1,162 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"context"
"sync"
"sync/atomic"
"testing"
"time"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
)
func installTestManagerSettings(t *testing.T) func() {
t.Helper()
return InstallTestRuntimeSettings(RuntimeSettings{
ChallengeCount: 3,
ChallengeSize: 32,
ChallengeDifficulty: 3,
ChallengeTTL: 5 * time.Second,
TokenTTL: 10 * time.Second,
})
}
func TestCapFullFlow(t *testing.T) {
cleanup := installTestManagerSettings(t)
defer cleanup()
secret := []byte("a-very-long-secret-key-at-least-16-bytes")
store := pkgcap.NewMemoryStore(1 * time.Minute)
manager := NewManager(secret, store)
scope := "test-scope"
ctx := context.Background()
resp, err := manager.Generate(ctx, scope)
if err != nil {
t.Fatalf("Generate() error = %v", err)
}
if resp.Challenge.C != 3 {
t.Fatalf("Generate().Challenge.C = %d, want %d", resp.Challenge.C, 3)
}
solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, scope)
if err != nil {
t.Fatalf("Redeem() error = %v", err)
}
if !redeemResp.Success {
t.Fatalf("Redeem().Success = false, error = %s", redeemResp.Error)
}
if redeemResp.Token == "" {
t.Fatal("Redeem().Token is empty")
}
valid, err := manager.VerifyToken(ctx, redeemResp.Token, scope)
if err != nil {
t.Fatalf("VerifyToken() error = %v", err)
}
if !valid {
t.Fatal("VerifyToken() = false, want true")
}
validAgain, err := manager.VerifyToken(ctx, redeemResp.Token, scope)
if err != nil {
t.Fatalf("VerifyToken() second call error = %v", err)
}
if validAgain {
t.Fatal("VerifyToken() second call = true, want false")
}
}
func TestRedeemConcurrentRace(t *testing.T) {
const goroutines = 50
cleanup := installTestManagerSettings(t)
defer cleanup()
secret := []byte("race-test-secret-key-at-least-16-bytes")
store := pkgcap.NewMemoryStore(1 * time.Minute)
manager := NewManager(secret, store)
ctx := context.Background()
resp, err := manager.Generate(ctx, "login")
if err != nil {
t.Fatalf("Generate() error = %v", err)
}
solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
var (
wg sync.WaitGroup
success atomic.Int32
barrier = make(chan struct{})
)
for range goroutines {
wg.Add(1)
go func() {
defer wg.Done()
<-barrier
r, _ := manager.Redeem(ctx, resp.Token, solutions, "login")
if r != nil && r.Success {
success.Add(1)
}
}()
}
close(barrier)
wg.Wait()
if got := success.Load(); got != 1 {
t.Fatalf("successful Redeem count = %d, want %d", got, 1)
}
}
func TestVerifyTokenConcurrentRace(t *testing.T) {
const goroutines = 50
cleanup := installTestManagerSettings(t)
defer cleanup()
secret := []byte("race-test-secret-key-at-least-16-bytes")
store := pkgcap.NewMemoryStore(1 * time.Minute)
manager := NewManager(secret, store)
ctx := context.Background()
resp, err := manager.Generate(ctx, "login")
if err != nil {
t.Fatalf("Generate() error = %v", err)
}
solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, "login")
if err != nil || !redeemResp.Success {
t.Fatalf("Redeem() error = %v, resp = %+v", err, redeemResp)
}
var (
wg sync.WaitGroup
success atomic.Int32
barrier = make(chan struct{})
)
for range goroutines {
wg.Add(1)
go func() {
defer wg.Done()
<-barrier
ok, _ := manager.VerifyToken(ctx, redeemResp.Token, "login")
if ok {
success.Add(1)
}
}()
}
close(barrier)
wg.Wait()
if got := success.Load(); got != 1 {
t.Fatalf("successful VerifyToken count = %d, want %d", got, 1)
}
}
@@ -0,0 +1,34 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
return func(c *gin.Context) {
if !ProtectionEnabled(c.Request.Context()) {
c.Next()
return
}
token := c.GetHeader("X-Cap-Token")
if token == "" {
response.AbortUnauthorized(c, errCapTokenMissing)
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
return
}
c.Next()
}
}
@@ -0,0 +1,94 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"net/http"
"github.com/gin-gonic/gin"
)
type challengeRequest struct {
Scope string `json:"scope" form:"scope"`
}
type redeemRequest struct {
Token string `json:"token" binding:"required"`
Solutions []int `json:"solutions" binding:"required"`
Scope string `json:"scope" form:"scope"`
}
// Challenge 生成 PoW 人机验证难题
// @Summary 生成人机验证难题
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
// @Tags cap
// @Accept json
// @Produce json
// @Param request body challengeRequest false "可选范围限制参数"
// @Success 200 {object} cap.ChallengeResponse "成功返回 PoW 难题"
// @Failure 500 {object} RedeemResponse "内部服务错误"
// @Router /api/cap/challenge [post]
func Challenge(c *gin.Context) {
var req challengeRequest
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultManager()
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, RedeemResponse{
Success: false,
Error: err.Error(),
})
return
}
c.JSON(http.StatusOK, resp)
}
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
// @Summary 校验人机验证解答
// @Description 提交 PoW 解答进行核销,成功后返回一次性 X-Cap-Token 凭证
// @Tags cap
// @Accept json
// @Produce json
// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组"
// @Success 200 {object} RedeemResponse "核销成功,返回 X-Cap-Token"
// @Failure 400 {object} RedeemResponse "参数错误或核销失败"
// @Failure 500 {object} RedeemResponse "内部服务错误"
// @Router /api/cap/redeem [post]
func Redeem(c *gin.Context) {
var req redeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, RedeemResponse{
Success: false,
Error: "无效的参数",
})
return
}
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultManager()
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, RedeemResponse{
Success: false,
Error: err.Error(),
})
return
}
if !resp.Success {
c.JSON(http.StatusBadRequest, resp)
return
}
c.JSON(http.StatusOK, resp)
}
@@ -0,0 +1,128 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
)
func TestCapEndpointsAndMiddleware(t *testing.T) {
sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
r := testhelper.NewTestGinEngine()
// Mount CAPTCHA API endpoints
capGroup := r.Group("/api/cap")
{
capGroup.POST("/challenge", Challenge)
capGroup.POST("/redeem", Redeem)
}
r.POST("/api/v1/user/login", VerifyMiddleware(GetDefaultManager(), "login"), func(c *gin.Context) {
c.JSON(http.StatusOK, response.OK("login success"))
})
// 1. Test challenge generation
w := httptest.NewRecorder()
req, _ := http.NewRequest("POST", "/api/cap/challenge", nil)
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var challengeResp pkgcap.ChallengeResponse
if err := json.Unmarshal(w.Body.Bytes(), &challengeResp); err != nil {
t.Fatalf("failed to unmarshal challenge response: %v", err)
}
if challengeResp.Token == "" {
t.Fatalf("expected token in challenge response")
}
// 2. Test login with CAPTCHA disabled (should pass)
w = httptest.NewRecorder()
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK when CAPTCHA is disabled, got %d. Body: %s", w.Code, w.Body.String())
}
// 3. Enable CAPTCHA in DB and invalidate runtime snapshot
err := sqliteDB.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyCapLoginEnabled).Update("value", "true").Error
if err != nil {
t.Fatalf("failed to enable cap_login_enabled in DB: %v", err)
}
if err := repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled); err != nil {
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
}
InvalidateRuntimeSettings()
// 4. Test login with CAPTCHA enabled but no header (should be blocked)
w = httptest.NewRecorder()
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Fatalf("expected 401 Unauthorized, got %d. Body: %s", w.Code, w.Body.String())
}
// 5. Solve the challenge
solutions := pkgcap.Solve(challengeResp.Token, challengeResp.Challenge.C, challengeResp.Challenge.S, challengeResp.Challenge.D)
// 6. Redeem solutions
redeemReqPayload := redeemRequest{
Token: challengeResp.Token,
Solutions: solutions,
}
bodyBytes, _ := json.Marshal(redeemReqPayload)
w = httptest.NewRecorder()
req, _ = http.NewRequest("POST", "/api/cap/redeem", bytes.NewBuffer(bodyBytes))
req.Header.Set("Content-Type", "application/json")
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK for redeem, got %d. Body: %s", w.Code, w.Body.String())
}
var redeemResp RedeemResponse
if err := json.Unmarshal(w.Body.Bytes(), &redeemResp); err != nil {
t.Fatalf("failed to unmarshal redeem response: %v", err)
}
if !redeemResp.Success || redeemResp.Token == "" {
t.Fatalf("redeem failed or returned empty token: %+v", redeemResp)
}
// 7. Login with valid redeem token (should pass)
w = httptest.NewRecorder()
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
req.Header.Set("X-Cap-Token", redeemResp.Token)
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK with valid cap token, got %d. Body: %s", w.Code, w.Body.String())
}
// 8. Replay attack: Login with the same redeem token again (should be blocked as it is single-use)
w = httptest.NewRecorder()
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
req.Header.Set("X-Cap-Token", redeemResp.Token)
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Fatalf("expected 401 Unauthorized on replayed token, got %d. Body: %s", w.Code, w.Body.String())
}
}
@@ -0,0 +1,212 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"context"
"encoding/json"
"errors"
"strconv"
"sync"
"sync/atomic"
"time"
"golang.org/x/sync/singleflight"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
)
const (
defaultChallengeCount = 1
defaultChallengeSize = 32
defaultChallengeDifficulty = 4
defaultChallengeTTL = 10 * time.Minute
defaultTokenTTL = 20 * time.Minute
)
// RuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
type RuntimeSettings struct {
LoginEnabled bool
ChallengeCount int
ChallengeSize int
ChallengeDifficulty int
ChallengeTTL time.Duration
TokenTTL time.Duration
}
var runtimeConfigKeys = []string{
model.ConfigKeyCapLoginEnabled,
model.ConfigKeyCapChallengeCount,
model.ConfigKeyCapChallengeSize,
model.ConfigKeyCapChallengeDifficulty,
model.ConfigKeyCapChallengeTTL,
model.ConfigKeyCapTokenTTL,
}
var runtimeConfigKeySet = func() map[string]struct{} {
set := make(map[string]struct{}, len(runtimeConfigKeys))
for _, key := range runtimeConfigKeys {
set[key] = struct{}{}
}
return set
}()
type runtimeSettingsStore struct {
snapshot atomic.Pointer[RuntimeSettings]
loadGroup singleflight.Group
listenerOnce sync.Once
}
var settingsStore = &runtimeSettingsStore{}
// IsRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
func IsRuntimeConfigKey(key string) bool {
_, ok := runtimeConfigKeySet[key]
return ok
}
// CurrentSettings returns the cached CAPTCHA runtime settings snapshot.
func CurrentSettings(ctx context.Context) (RuntimeSettings, error) {
return settingsStore.current(ctx)
}
// ProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
func ProtectionEnabled(ctx context.Context) bool {
settings, err := CurrentSettings(ctx)
if err != nil {
return false
}
return settings.LoginEnabled
}
// InvalidateRuntimeSettings drops the in-process CAPTCHA settings snapshot.
func InvalidateRuntimeSettings() {
settingsStore.snapshot.Store(nil)
}
// ResetRuntimeSettingsForTest clears the CAPTCHA runtime snapshot.
func ResetRuntimeSettingsForTest() {
InvalidateRuntimeSettings()
}
// InstallTestRuntimeSettings installs a fixed snapshot for unit tests.
func InstallTestRuntimeSettings(settings RuntimeSettings) func() {
snapshot := settings
settingsStore.snapshot.Store(&snapshot)
return InvalidateRuntimeSettings
}
func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, error) {
s.ensureInvalidationListener()
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
loaded, err, _ := s.loadGroup.Do("cap-runtime-settings", func() (any, error) {
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
settings, loadErr := loadRuntimeSettings(ctx)
if loadErr != nil {
return RuntimeSettings{}, loadErr
}
s.snapshot.Store(&settings)
return settings, nil
})
if err != nil {
return RuntimeSettings{}, err
}
settings, ok := loaded.(RuntimeSettings)
if !ok {
return RuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
}
return settings, nil
}
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
configs, err := repository.ListSystemConfigsByKeys(ctx, runtimeConfigKeys)
if err != nil {
return RuntimeSettings{}, err
}
return parseRuntimeSettings(configs), nil
}
func parseRuntimeSettings(configs map[string]model.SystemConfig) RuntimeSettings {
settings := RuntimeSettings{
ChallengeCount: defaultChallengeCount,
ChallengeSize: defaultChallengeSize,
ChallengeDifficulty: defaultChallengeDifficulty,
ChallengeTTL: defaultChallengeTTL,
TokenTTL: defaultTokenTTL,
}
if sc, ok := configs[model.ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(sc.Value); err == nil {
settings.LoginEnabled = enabled
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(sc.Value); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(sc.Value); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeDifficulty]; ok {
if difficulty, err := strconv.Atoi(sc.Value); err == nil && difficulty > 0 {
settings.ChallengeDifficulty = difficulty
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(sc.Value); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if sc, ok := configs[model.ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(sc.Value); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
return settings
}
func (s *runtimeSettingsStore) ensureInvalidationListener() {
s.listenerOnce.Do(startRuntimeSettingsInvalidationListener)
}
func startRuntimeSettingsInvalidationListener() {
if db.Redis == nil {
return
}
go func() {
pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
for msg := range pubsub.Channel() {
var payload struct {
Key string `json:"key"`
}
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
InvalidateRuntimeSettings()
continue
}
if payload.Key == "" || payload.Key == "*" || IsRuntimeConfigKey(payload.Key) {
InvalidateRuntimeSettings()
}
}
}()
}
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
func TestCurrentSettingsLoadsSnapshotOnce(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ResetRuntimeSettingsForTest()
repository.ResetSystemConfigRAMCacheForTest()
first, err := CurrentSettings(ctx)
if err != nil {
t.Fatalf("CurrentSettings() first error = %v", err)
}
if first.ChallengeCount != 1 {
t.Fatalf("CurrentSettings().ChallengeCount = %d, want %d", first.ChallengeCount, 1)
}
if err := db.DB(ctx).Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeyCapChallengeCount).
Update("value", "4").Error; err != nil {
t.Fatalf("Update(cap_challenge_count) error = %v", err)
}
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapChallengeCount); err != nil {
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
}
InvalidateRuntimeSettings()
second, err := CurrentSettings(ctx)
if err != nil {
t.Fatalf("CurrentSettings() second error = %v", err)
}
if second.ChallengeCount != 4 {
t.Fatalf("CurrentSettings().ChallengeCount = %d, want %d", second.ChallengeCount, 4)
}
}
func TestProtectionEnabledReflectsLoginSwitch(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ResetRuntimeSettingsForTest()
if !ProtectionEnabled(ctx) {
t.Fatal("ProtectionEnabled() = false, want true from seed defaults")
}
if err := db.DB(ctx).Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeyCapLoginEnabled).
Update("value", "false").Error; err != nil {
t.Fatalf("Update(cap_login_enabled) error = %v", err)
}
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapLoginEnabled); err != nil {
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
}
InvalidateRuntimeSettings()
if ProtectionEnabled(ctx) {
t.Fatal("ProtectionEnabled() = true, want false after config update")
}
}
func TestParseRuntimeSettingsUsesDefaultsForMissingKeys(t *testing.T) {
settings := parseRuntimeSettings(map[string]model.SystemConfig{})
if settings.ChallengeCount != defaultChallengeCount {
t.Fatalf("ChallengeCount = %d, want %d", settings.ChallengeCount, defaultChallengeCount)
}
if settings.ChallengeTTL != defaultChallengeTTL {
t.Fatalf("ChallengeTTL = %s, want %s", settings.ChallengeTTL, defaultChallengeTTL)
}
if settings.TokenTTL != defaultTokenTTL {
t.Fatalf("TokenTTL = %s, want %s", settings.TokenTTL, defaultTokenTTL)
}
}
func TestIsRuntimeConfigKey(t *testing.T) {
if !IsRuntimeConfigKey(model.ConfigKeyCapChallengeCount) {
t.Fatalf("IsRuntimeConfigKey(%s) = false, want true", model.ConfigKeyCapChallengeCount)
}
if IsRuntimeConfigKey(model.ConfigKeySiteName) {
t.Fatalf("IsRuntimeConfigKey(%s) = true, want false", model.ConfigKeySiteName)
}
}
func TestInstallTestRuntimeSettings(t *testing.T) {
cleanup := InstallTestRuntimeSettings(RuntimeSettings{
LoginEnabled: true,
ChallengeCount: 2,
TokenTTL: 30 * time.Minute,
})
defer cleanup()
settings, err := CurrentSettings(context.Background())
if err != nil {
t.Fatalf("CurrentSettings() error = %v", err)
}
if !settings.LoginEnabled {
t.Fatal("LoginEnabled = false, want true")
}
if settings.ChallengeCount != 2 {
t.Fatalf("ChallengeCount = %d, want %d", settings.ChallengeCount, 2)
}
}
@@ -0,0 +1,10 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import "github.com/Rain-kl/Wavelet/internal/testhelper"
func init() {
testhelper.RegisterCleanup(ResetRuntimeSettingsForTest)
}
@@ -0,0 +1,76 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
}
if _, err := repository.ListVisibleSystemConfigs(ctx); err != nil {
t.Fatalf("ListVisibleSystemConfigs() warm cache error = %v", err)
}
if err := dbConn.Create(&model.SystemConfig{
Key: "cache_probe_public_key",
Value: "cache_probe_public_value",
Type: "system",
Visibility: model.ConfigVisibilityVisible,
Description: "cache probe",
}).Error; err != nil {
t.Fatalf("Create(cache_probe_public_key) error = %v", err)
}
cached, err := repository.ListVisibleSystemConfigs(ctx)
if err != nil {
t.Fatalf("ListVisibleSystemConfigs() cached call error = %v", err)
}
for _, item := range cached {
if item.Key == "cache_probe_public_key" {
t.Fatal("cached visible config list should be stale before invalidation")
}
}
exists, err := db.Redis.Exists(ctx, db.PrefixedKey(repository.SystemConfigVisibleListRedisKey)).Result()
if err != nil {
t.Fatalf("Redis.Exists() error = %v", err)
}
if exists == 0 {
t.Fatal("expected visible config list cache key to exist")
}
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
}
refreshed, err := repository.ListVisibleSystemConfigs(ctx)
if err != nil {
t.Fatalf("ListVisibleSystemConfigs() refreshed call error = %v", err)
}
var found bool
for _, item := range refreshed {
if item.Key == "cache_probe_public_key" {
found = true
break
}
}
if !found {
t.Fatal("refreshed visible config list should include newly created public config")
}
}
@@ -0,0 +1,57 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package config 提供公开配置查询接口
package config
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// GetPublicConfig 获取公共配置
// @Summary 获取公共配置
// @Description 返回系统配置表中 visibility 为 1 的配置键值集合
// @Tags config
// @Accept json
// @Produce json
// @Success 200 {object} response.Any
// @Router /api/v1/config/public [get]
func GetPublicConfig(c *gin.Context) {
ctx := c.Request.Context()
configs, err := repository.ListVisibleSystemConfigs(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
resp := make(map[string]string, len(configs))
for _, config := range configs {
resp[config.Key] = config.Value
}
c.JSON(http.StatusOK, response.OK(resp))
}
// GetRobotsTXT 动态生成 robots.txt
// @Summary 获取 robots.txt
// @Description 根据系统配置决定是否允许搜索引擎检索,并返回相应的 robots.txt 文件内容
// @Tags config
// @Produce text/plain
// @Success 200 {string} string "robots.txt 内容"
// @Router /robots.txt [get]
func GetRobotsTXT(c *gin.Context) {
ctx := c.Request.Context()
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled)
content := "User-Agent: *\nDisallow: /\n"
if err == nil && enabled {
content = "User-Agent: *\nAllow: /\n"
}
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(content))
}
@@ -0,0 +1,70 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config
import (
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
func TestGetPublicConfigUsesVisibility(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
if err := dbConn.Create(&model.SystemConfig{
Key: "custom_public_key",
Value: "custom_public_value",
Type: "system",
Visibility: model.ConfigVisibilityVisible,
Description: "custom public config",
}).Error; err != nil {
t.Fatalf("Create(custom_public_key) error = %v", err)
}
if err := dbConn.Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeySiteName).
Update("visibility", model.ConfigVisibilityHidden).Error; err != nil {
t.Fatalf("Update(%s.visibility) error = %v", model.ConfigKeySiteName, err)
}
gin.SetMode(gin.TestMode)
router := gin.New()
router.GET("/api/v1/config/public", GetPublicConfig)
req := httptest.NewRequest(http.MethodGet, "/api/v1/config/public", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("GetPublicConfig() status = %d, want %d; body = %s", w.Code, http.StatusOK, w.Body.String())
}
var resp response.Any
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("json.Unmarshal(GetPublicConfig()) error = %v", err)
}
dataBytes, err := json.Marshal(resp.Data)
if err != nil {
t.Fatalf("json.Marshal(GetPublicConfig().data) error = %v", err)
}
var configs map[string]string
if err := json.Unmarshal(dataBytes, &configs); err != nil {
t.Fatalf("json.Unmarshal(GetPublicConfig().data) error = %v", err)
}
if got := configs["custom_public_key"]; got != "custom_public_value" {
t.Errorf("GetPublicConfig()[custom_public_key] = %q, want %q", got, "custom_public_value")
}
if _, ok := configs[model.ConfigKeySiteName]; ok {
t.Errorf("GetPublicConfig()[%s] is present, want hidden", model.ConfigKeySiteName)
}
}
@@ -0,0 +1,108 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config
import (
"context"
"errors"
"testing"
"github.com/redis/go-redis/v9"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
repository.ResetSystemConfigRAMCacheForTest()
if err := repository.InvalidateAllSystemConfigCaches(ctx); err != nil {
t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err)
}
warm, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
}
if warm.Value != "OpenFlare" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "OpenFlare")
}
if err := dbConn.Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeySiteName).
Update("value", "ram_probe_value").Error; err != nil {
t.Fatalf("Update(site_name) error = %v", err)
}
if err := db.HDel(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeySiteName); err != nil {
t.Fatalf("HDel(site_name) error = %v", err)
}
cached, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) cached error = %v", err)
}
if cached.Value != "OpenFlare" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "OpenFlare")
}
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
}
refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err)
}
if refreshed.Value != "ram_probe_value" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "ram_probe_value")
}
exists, err := db.Redis.HExists(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
if err != nil {
t.Fatalf("HExists(site_name) error = %v", err)
}
if !exists {
t.Fatal("expected redis hash field to be repopulated after refresh")
}
}
func TestInvalidateSystemConfigCacheClearsRedisField(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err)
}
_ = sc
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
}
_, err = db.Redis.HGet(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
if !errors.Is(err, redis.Nil) {
t.Fatalf("HGet(site_name) error = %v, want redis.Nil", err)
}
if err := dbConn.Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeySiteName).
Update("value", "after_invalidate").Error; err != nil {
t.Fatalf("Update(site_name) error = %v", err)
}
refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err)
}
if refreshed.Value != "after_invalidate" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "after_invalidate")
}
}
@@ -0,0 +1,24 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package custom provides custom business handlers
package custom
import (
"net/http"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// Hello is a sample handler for custom business logic
// @Summary Sample Hello API
// @Description A sample business API for customization
// @Tags custom
// @Produce json
// @Success 200 {object} response.Any{data=string} "成功"
// @Router /api/v1/custom/hello [get]
func Hello(c *gin.Context) {
c.JSON(http.StatusOK, response.OK("Hello from custom business module!"))
}
@@ -0,0 +1,25 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package health 提供健康检查端点
package health
import (
"net/http"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// Health 健康检查
// @Summary 健康检查
// @Description 检查服务是否正常运行,可用于负载均衡存活探测
// @Tags health
// @Produce json
// @Success 200 {object} response.Any{data=string} "服务正常"
// @Router /api/health [get]
func Health(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,36 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package oauth 提供 OAuth/OIDC 认证与会话管理
package oauth
import (
"context"
"encoding/json"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) {
auditLog := loginRequiredAuditLog{
UserID: user.ID,
Username: user.Username,
ClientIP: c.ClientIP(),
Method: c.Request.Method,
Path: c.Request.URL.Path,
RequestURI: c.Request.RequestURI,
UserAgent: c.Request.UserAgent(),
Referer: c.Request.Referer(),
}
auditJSON, err := json.Marshal(auditLog)
if err != nil {
logger.ErrorF(ctx, "[LoginRequiredAudit] marshal failed: %v", err)
logger.DebugF(ctx, "[LoginRequiredAudit] %s %d %s", c.ClientIP(), user.ID, user.Username)
} else {
logger.DebugF(ctx, "[LoginRequiredAudit] %s", auditJSON)
}
}
@@ -0,0 +1,118 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"errors"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func isOIDCLoginEnabled(ctx context.Context) bool {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
if err != nil {
return true
}
return enabled
}
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := model.GetActiveAuthSources(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(errNoActiveAuthSource)
}
return &sources[0], nil
}
return model.GetAuthSourceByName(ctx, name)
}
func activeLoginSources(ctx context.Context) []AuthSourceView {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
if err == nil && !enabled {
return nil
}
dbSources, err := model.GetActiveAuthSources(ctx)
if err != nil {
return nil
}
sources := make([]AuthSourceView, 0, len(dbSources))
for _, source := range dbSources {
sources = append(sources, AuthSourceView{
ID: source.ID,
Name: source.Name,
Type: source.Type,
DisplayName: source.DisplayName,
IsActive: source.IsActive,
IconURL: source.IconURL,
ClientSecretConfigured: source.ClientSecretConfigured,
})
}
return sources
}
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
if err != nil || strings.TrimSpace(sc.Value) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(sc.Value, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New(errAuthSourceRequired)
}
if source.OpenIDDiscoveryURL == "" {
return nil, nil, errors.New(errDiscoveryURLRequired)
}
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/")
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
// 使用进程级缓存获取 provider,避免每次调用都向 issuer 发起
// /.well-known/openid-configuration HTTP 请求。
provider, err := globalOIDCProviderCache.get(ctx, issuer)
if err != nil {
return nil, nil, err
}
verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID})
scopes := strings.Fields(source.Scopes)
if len(scopes) == 0 {
scopes = []string{oidc.ScopeOpenID, "profile", "email"}
}
if !containsScope(scopes, oidc.ScopeOpenID) {
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
}
return &oauth2.Config{
ClientID: source.ClientID,
ClientSecret: source.ClientSecret,
RedirectURL: redirectURL,
Scopes: scopes,
Endpoint: provider.Endpoint(),
}, verifier, nil
}
func containsScope(scopes []string, scope string) bool {
for _, item := range scopes {
if item == scope {
return true
}
}
return false
}
@@ -0,0 +1,56 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"encoding/json"
"time"
)
// Session 用户信息字段 Key
const (
UserNameKey = "username"
UserIDKey = "user_id"
UserObjKey = "user_obj"
TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权
TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限
SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials
PasswordHashKey = "password_hash"
)
// OAuth State 缓存 Key 格式与过期时间
const (
OAuthStateCacheKeyFormat = "oauth:state:%s"
OAuthStateCacheKeyExpiration = 10 * time.Minute
)
// OAuth 授权用途常量
const (
OAuthPurposeLogin = "login"
OAuthPurposeBind = "bind"
)
type oauthStatePayload struct {
SourceName string `json:"source_name"`
Purpose string `json:"purpose"`
UserID uint64 `json:"user_id,omitempty"`
SessionHash string `json:"session_hash"`
}
func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) {
data, err := json.Marshal(payload)
if err != nil {
return "", err
}
return string(data), nil
}
func decodeOAuthStatePayload(value string) (oauthStatePayload, error) {
var payload oauthStatePayload
if err := json.Unmarshal([]byte(value), &payload); err != nil {
return oauthStatePayload{}, err
}
return payload, nil
}
@@ -0,0 +1,22 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
// OAuth 认证相关错误消息
const (
errInvalidState = "非法登录请求"
errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errIDTokenVerifyFailedFormat = "%s: %w"
errNonceMismatch = "nonce 不匹配,可能存在重放攻击"
errNoActiveAuthSource = "未配置可用认证源"
errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试"
errAuthSourceRequired = "认证源不能为空"
errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errUsernameGenerateFailed = "无法生成可用用户名"
errUsernameFromSourceFailed = "无法从认证源获取用户名"
errAuthSourceDisabled = "认证源未启用"
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
)
@@ -0,0 +1,22 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import "github.com/gin-gonic/gin"
// GetFromContext 从 Gin 请求上下文获取指定类型的值。
func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
value, exists := c.Get(key)
if !exists {
var zero T
return zero, false
}
typed, ok := value.(T)
return typed, ok
}
// SetToContext 设置值到 Gin 请求上下文。
func SetToContext[T any](c *gin.Context, key string, value T) {
c.Set(key, value)
}
@@ -0,0 +1,173 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"fmt"
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
)
// GetLoginURL 获取登录授权地址
// @Summary 获取登录授权地址
// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。
// @Tags oauth
// @Produce json
// @Param source query string false "认证源名称,为空使用第一个启用的认证源"
// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
// @Failure 400 {object} response.Any "认证源不存在或未配置"
// @Failure 500 {object} response.Any "Redis 异常 or 构造 URL 失败"
// @Router /api/v1/oauth/login [get]
func GetLoginURL(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Query("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
session := sessions.Default(c)
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
userID := GetUserIDFromSession(session)
sessionHash := hashSessionToken(token)
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: OAuthPurposeLogin,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
response.AbortInternal(c, err.Error())
return
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) {
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
return "", err
}
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return "", err
}
if verifier != nil {
return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil
}
return authConfig.AuthCodeURL(state), nil
}
// Authorize 发起指定认证源授权
// @Summary 发起指定认证源授权
// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
// @Tags oauth
// @Produce json
// @Param source path string true "认证源名称"
// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login"
// @Success 200 {object} response.Any{data=oauth.OAuthAuthorizeResponse} "授权 URL"
// @Failure 400 {object} response.Any "认证源不存在或未启用"
// @Failure 500 {object} response.Any "Redis 异常或构造 URL 失败"
// @Router /api/v1/oauth/{source}/authorize [get]
func Authorize(c *gin.Context) {
ctx := c.Request.Context()
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, c.Param("source"))
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose")))
if purpose != OAuthPurposeBind {
purpose = OAuthPurposeLogin
}
session := sessions.Default(c)
userID := GetUserIDFromSession(session)
if purpose == OAuthPurposeBind && userID == 0 {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
token, isNew := ensureSessionToken(session)
if isNew {
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
}
sessionHash := hashSessionToken(token)
state := uuid.NewString()
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
SourceName: source.Name,
Purpose: purpose,
UserID: userID,
SessionHash: sessionHash,
})
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
response.AbortInternal(c, err.Error())
return
}
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
@@ -0,0 +1,227 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"errors"
"fmt"
"net/http"
"time"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// Callback OAuth 回调处理
// @Summary OAuth 回调处理
// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。
// @Tags oauth
// @Accept json
// @Produce json
// @Param request body oauth.CallbackRequest true "回调请求参数"
// @Success 200 {object} response.Any{data=oauth.OAuthCallbackResult} "登录或绑定成功"
// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误"
// @Failure 401 {object} response.Any "绑定场景未登录"
// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误"
// @Router /api/v1/oauth/callback [post]
func Callback(c *gin.Context) {
var req CallbackRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
ctx := c.Request.Context()
stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
if err != nil {
response.AbortBadRequest(c, errInvalidState)
return
}
_ = db.Redis.Del(ctx, stateKey)
payload, err := decodeOAuthStatePayload(payloadRaw)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
session := sessions.Default(c)
currentUserID := GetUserIDFromSession(session)
if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
token, ok := session.Get(SessionTokenKey).(string)
if !ok || token == "" {
response.AbortBadRequest(c, "invalid session context")
return
}
if hashSessionToken(token) != payload.SessionHash {
response.AbortBadRequest(c, "session mismatch for oauth state")
return
}
if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID {
response.AbortBadRequest(c, "user context mismatch for oauth binding")
return
}
if !isOIDCLoginEnabled(ctx) {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
source, err := resolveAuthSource(ctx, payload.SourceName)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if !source.IsActive {
response.AbortBadRequest(c, errAuthSourceDisabled)
return
}
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := normalizeOAuthUserInfo(userInfo); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if userInfo.Sub == "" {
userInfo.Sub = userInfo.Username
}
if payload.Purpose == OAuthPurposeBind {
handleCallbackBind(ctx, c, source, userInfo)
return
}
handleCallbackLogin(ctx, c, source, userInfo)
}
// handleCallbackBind 处理 OAuth 回调中的帐号绑定流程
func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
var user model.User
if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
}
// handleCallbackLogin 处理 OAuth 回调中的登录流程(查找已有帐号或自动注册)
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
var user model.User
account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
case errors.Is(err, gorm.ErrRecordNotFound):
newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
if !ok {
return
}
user = newUser
default:
response.AbortInternal(c, err.Error())
return
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
if err := setLoginSession(ctx, c, &user); err != nil {
response.AbortInternal(c, err.Error())
return
}
logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
}
// handleCallbackRegister 处理 OAuth 回调中的自动注册流程
// 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
if regErr != nil {
registrationEnabled = true
}
if !registrationEnabled {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return model.User{}, false
}
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
response.AbortInternal(c, uniqueErr.Error())
return model.User{}, false
}
userInfo.Username = username
var user model.User
if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil {
response.AbortInternal(c, err.Error())
return model.User{}, false
}
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
ExternalUsername: userInfo.Username,
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return model.User{}, false
}
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
return user, true
}
@@ -0,0 +1,65 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"net/http"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
)
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
// @Summary 获取外部帐号列表
// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.ExternalAccountView} "外部帐号列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/oauth/external-accounts [get]
func ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(accounts))
}
// DeleteExternalAccount 解除外部帐号绑定
// @Summary 解除外部帐号绑定
// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Param id path uint64 true "外部帐号绑定记录 ID"
// @Success 200 {object} response.Any{data=string} "解除绑定成功"
// @Failure 400 {object} response.Any "ID 无效或解除失败"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/oauth/external-accounts/{id}/delete [post]
func DeleteExternalAccount(c *gin.Context) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
rawID := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseUint(rawID, 10, 64)
if err != nil || id == 0 {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return
}
if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,22 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/gin-gonic/gin"
)
// GetLoginSources 获取可用登录源列表
// @Summary 获取可用登录源
// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用
// @Tags oauth
// @Produce json
// @Success 200 {object} response.Any{data=[]oauth.AuthSourceView} "登录源列表"
// @Router /api/v1/oauth/sources [get]
func GetLoginSources(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context())))
}
@@ -0,0 +1,138 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"errors"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
type loginRequiredAuditLog struct {
UserID uint64 `json:"user_id"`
Username string `json:"username"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Path string `json:"path"`
RequestURI string `json:"request_uri"`
UserAgent string `json:"user_agent"`
Referer string `json:"referer"`
}
func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) {
tokenHash := model.HashToken(tokenStr)
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err != nil {
return nil, nil, err
}
var user model.User
if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&user).Error; err != nil {
return nil, nil, err
}
return &user, &tokenRecord, nil
}
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
func GetUserFromRequest(c *gin.Context) (*model.User, error) {
ctx := c.Request.Context()
// check token in headers
tokenStr := c.GetHeader("X-Access-Token")
if tokenStr == "" {
authHeader := c.GetHeader("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
tokenStr = authHeader[7:]
}
}
// 优先使用 Access Token 鉴权
if tokenStr != "" {
if user, tokenRecord, err := getUserByToken(ctx, tokenStr); err == nil {
// 强行阻止 system 用户任何会话/Token 鉴权通过
if user.Username == "system" {
return nil, errors.New("system user is not allowed to login")
}
SetToContext(c, TokenAuthKey, true)
SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin)
return user, nil
}
}
// 降级使用 Session 鉴权
userID := GetUserIDFromContext(c)
if userID <= 0 {
return nil, errors.New("unauthorized")
}
var user model.User
// load user from db to make sure is active
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&user)
if tx.Error != nil {
return nil, tx.Error
}
// 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致
if user.Password != "" {
session := sessions.Default(c)
sessionHash, _ := session.Get(PasswordHashKey).(string)
if sessionHash != user.Password {
return nil, errors.New("session expired due to password change")
}
}
// set keys in context for session auth
SetToContext(c, TokenAuthKey, false)
SetToContext(c, TokenAdminKey, false)
// 强行阻止 system 用户任何会话/Token 鉴权通过
if user.Username == "system" {
return nil, errors.New("system user is not allowed to login")
}
return &user, nil
}
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
func LoginRequired() gin.HandlerFunc {
return func(c *gin.Context) {
// init trace
ctx, span := otel_trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
user, err := GetUserFromRequest(c)
if err != nil {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
// log
LogForAudit(ctx, user, c)
// set user info
SetToContext(c, UserObjKey, user)
// next
c.Next()
}
}
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
func DisallowTokenAuth() gin.HandlerFunc {
return func(c *gin.Context) {
if tokenAuth, _ := GetFromContext[bool](c, TokenAuthKey); tokenAuth {
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
return
}
c.Next()
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,36 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
// AuthSourceView 登录源展示信息
type AuthSourceView struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IsActive bool `json:"is_active"`
IconURL string `json:"icon_url"`
ClientSecretConfigured bool `json:"client_secret_configured"`
}
// OAuthAuthorizeResponse 授权 URL 响应
//
//nolint:revive // OAuth 前缀保持包内语义清晰
type OAuthAuthorizeResponse struct {
AuthorizeURL string `json:"authorize_url"`
}
// OAuthCallbackResult 回调处理结果
//
//nolint:revive // OAuth 前缀保持包内语义清晰
type OAuthCallbackResult struct {
Status string `json:"status"`
User *BasicUserInfo `json:"user,omitempty"`
}
// CallbackRequest OAuth 回调请求参数
type CallbackRequest struct {
State string `json:"state" binding:"required"`
Code string `json:"code" binding:"required"`
}
@@ -0,0 +1,141 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"errors"
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
var existingUsernames []string
if err := db.DB(ctx).Model(&model.User{}).
Where("username = ? OR username LIKE ?", base, base+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return "", err
}
// 将现有的用户名放入 map 中,以便 O(1) 查找
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
// 检查 base 是否被占用
if !exists[strings.ToLower(base)] {
return base, nil
}
// 顺序查找第一个可用的带后缀用户名
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New(errUsernameGenerateFailed)
}
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return nil, err
}
token, err := authConfig.Exchange(ctx, code)
if err != nil {
return nil, err
}
userInfo := &model.OAuthUserInfo{Active: true}
if verifier != nil {
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
return nil, verifyErr
}
}
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
return userInfo, nil
}
// verifyIDToken 验证 OIDC ID Token 并将 Claims 解析到 userInfo
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
return nil
}
idToken, verifyErr := verifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr)
}
if nonce != "" && idToken.Nonce != nonce {
return errors.New(errNonceMismatch)
}
if claimsErr := idToken.Claims(userInfo); claimsErr != nil {
return claimsErr
}
return nil
}
func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
userInfo.Username = strings.TrimSpace(userInfo.Username)
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
userInfo.Email = strings.TrimSpace(userInfo.Email)
userInfo.Name = strings.TrimSpace(userInfo.Name)
userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL)
if userInfo.Username == "" && userInfo.PreferredUsername != "" {
userInfo.Username = userInfo.PreferredUsername
}
if userInfo.Username == "" && userInfo.Email != "" {
userInfo.Username = strings.Split(userInfo.Email, "@")[0]
}
if userInfo.Username == "" && userInfo.Sub != "" {
userInfo.Username = userInfo.Sub
}
if userInfo.Username == "" {
return errors.New(errUsernameFromSourceFailed)
}
if userInfo.Name == "" {
userInfo.Name = userInfo.Username
}
if !userInfo.Active {
userInfo.Active = true
}
return nil
}
func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
result := OAuthCallbackResult{Status: status}
if user != nil {
info := BuildBasicUserInfo(user, false)
result.User = &info
}
return result
}
@@ -0,0 +1,99 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"net/http"
"sync"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
"golang.org/x/sync/singleflight"
)
// oidcProviderCache 进程级 OIDC provider 缓存。
//
// oidc.NewProvider 每次调用都会向远端 issuer 的
// /.well-known/openid-configuration 发起 HTTP 请求拉取元数据。
// 由于 provider 元数据极少变动,将其缓存后可消除登录发起与回调时的
// 重复外部 HTTP 往返。
//
// 并发安全性:
// - mu + entries 防止并发读写 map。
// - sfGroup 保证同一 issuer 同时只有一次在途的 NewProvider 调用
// (singleflight),后续等待者复用同一结果,彻底消除 thundering herd。
type oidcProviderCache struct {
mu sync.RWMutex
entries map[string]*oidc.Provider // key: normalized issuer URL
sfGroup singleflight.Group
}
// globalOIDCProviderCache 是包级单例缓存,与进程同生命周期。
var globalOIDCProviderCache = &oidcProviderCache{
entries: make(map[string]*oidc.Provider),
}
// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。
// 这样既能在测试中注入 mock 客户端,又避免请求取消导致 provider 拉取失败。
func discoveryContext(ctx context.Context) context.Context {
bg := context.Background()
if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil {
bg = oidc.ClientContext(bg, client)
}
return bg
}
// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
// 同一 issuer 并发调用时,singleflight 保证只有一次实际 HTTP 请求。
func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) {
// 快路径:已有缓存则直接返回。
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
return p, nil
}
c.mu.RUnlock()
// 慢路径:通过 singleflight 合并并发的首次请求。
discCtx := discoveryContext(ctx)
v, err, _ := c.sfGroup.Do(issuer, func() (any, error) {
// 双检:singleflight 内再次检查,前一个并发组可能已写入缓存。
c.mu.RLock()
if p, ok := c.entries[issuer]; ok {
c.mu.RUnlock()
return p, nil
}
c.mu.RUnlock()
p, err := oidc.NewProvider(discCtx, issuer)
if err != nil {
return nil, err
}
c.mu.Lock()
c.entries[issuer] = p
c.mu.Unlock()
return p, nil
})
if err != nil {
return nil, err
}
return v.(*oidc.Provider), nil //nolint:forcetypeassert // singleflight value 由同函数写入,类型确定
}
// invalidate 从缓存中移除指定 issuer 对应的 provider。
// 在认证源的 Discovery URL 被修改时调用,强制下次请求重新拉取元数据。
func (c *oidcProviderCache) invalidate(issuer string) {
c.mu.Lock()
delete(c.entries, issuer)
c.mu.Unlock()
}
// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。
// 当管理员更新认证源的 Discovery URL 后调用,以确保下次登录时重新拉取最新元数据。
// issuer 值应为去掉 /.well-known/openid-configuration 后缀的规范化 URL。
func InvalidateOIDCProviderCache(issuer string) {
globalOIDCProviderCache.invalidate(issuer)
}
@@ -0,0 +1,106 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// BasicUserInfo 用户基本信息结构体
type BasicUserInfo struct {
ID uint64 `json:"id"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Email string `json:"email"`
AvatarURL string `json:"avatar_url"`
IsAdmin bool `json:"is_admin"`
NeedChangePassword bool `json:"need_change_password"`
Bio string `json:"bio"`
Phone string `json:"phone"`
Gender string `json:"gender"`
Website string `json:"website"`
Location string `json:"location"`
}
// BuildBasicUserInfo 将 User 模型转换为 BasicUserInfo
func BuildBasicUserInfo(user *model.User, needChange bool) BasicUserInfo {
return BasicUserInfo{
ID: user.ID,
Username: user.Username,
Nickname: user.Nickname,
Email: user.Email,
AvatarURL: user.AvatarURL,
IsAdmin: user.IsAdmin,
NeedChangePassword: needChange,
Bio: user.Bio,
Phone: user.Phone,
Gender: user.Gender,
Website: user.Website,
Location: user.Location,
}
}
// UserInfo 获取当前登录用户信息
// @Summary 获取当前登录用户信息
// @Description 返回当前登录用户的基本信息及余额数据,需要登录。包括用户 ID、用户名、信任等级、各类余额信息等。
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "用户信息"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/oauth/user-info [get]
// @Router /api/v1/user-info [get]
// @Router /api/v1/user/self [get]
func UserInfo(c *gin.Context) {
user, _ := GetFromContext[*model.User](c, UserObjKey)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true
c.JSON(
http.StatusOK,
response.OK(BuildBasicUserInfo(user, needChange)),
)
}
// GetLoginURL 获取登录地址
// @Summary 获取登录地址
// @Description 生成 OAuth 登录 URL,前端跳转至该地址完成授权。返回的 URL 中包含 state 参数用于 CSRF 防护。
// @Tags oauth
// @Produce json
// @Success 200 {object} response.Any{data=string} "OAuth 登录 URL"
// @Failure 500 {object} response.Any "Redis 异常或内部错误"
// @Router /api/v1/oauth/login [get]
// Logout 退出登录
// @Summary 退出登录
// @Description 清除当前用户的登录会话,完成退出。清除 Cookie 中的 Session 数据。
// @Tags oauth
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=string} "退出成功"
// @Failure 500 {object} response.Any "Session 清除失败"
// @Router /api/v1/oauth/logout [get]
func Logout(c *gin.Context) {
session := sessions.Default(c)
userID := session.Get(UserIDKey)
username := session.Get(UserNameKey)
if userID != nil {
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
}
session.Options(GetSessionOptions(-1))
session.Clear()
if err := session.Save(); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package oauth provides authentication and OAuth integration.
package oauth
import (
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/gin-contrib/sessions"
)
// GetSessionOptions 根据配置构建 Session 选项
func GetSessionOptions(maxAge int) sessions.Options {
return sessions.Options{
Path: "/",
Domain: config.Config.App.SessionDomain,
MaxAge: maxAge,
HttpOnly: config.Config.App.SessionHTTPOnly,
Secure: config.Config.App.SessionSecure,
SameSite: http.SameSiteLaxMode,
}
}
// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie
func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) {
headers := header["Set-Cookie"]
if len(headers) == 0 {
return
}
newHeaders := make([]string, 0, len(headers))
for _, h := range headers {
if strings.HasPrefix(h, cookieName+"=") {
parts := strings.Split(h, ";")
newParts := make([]string, 0, len(parts))
for _, p := range parts {
trimmed := strings.TrimSpace(p)
lower := strings.ToLower(trimmed)
if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") {
continue
}
newParts = append(newParts, p)
}
newHeaders = append(newHeaders, strings.Join(newParts, ";"))
} else {
newHeaders = append(newHeaders, h)
}
}
header["Set-Cookie"] = newHeaders
}
@@ -0,0 +1,83 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"crypto/sha256"
"encoding/hex"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
)
// GetUserIDFromSession 从 Session 中提取用户 ID
func GetUserIDFromSession(s sessions.Session) uint64 {
userID, ok := s.Get(UserIDKey).(uint64)
if !ok {
return 0
}
return userID
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) uint64 {
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
func ensureSessionToken(s sessions.Session) (string, bool) {
token, ok := s.Get(SessionTokenKey).(string)
if !ok || token == "" {
token = uuid.NewString()
s.Set(SessionTokenKey, token)
return token, true
}
return token, false
}
func hashSessionToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
session := sessions.Default(c)
session.Set(UserIDKey, user.ID)
session.Set(UserNameKey, user.Username)
session.Set(PasswordHashKey, user.Password)
// 根据系统配置动态设置 Session 过期时间
maxAge := config.Config.App.SessionAge
isSessionCookie := false
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
if err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
session.Options(GetSessionOptions(maxAge))
if err := session.Save(); err != nil {
return err
}
if isSessionCookie {
StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
}
return nil
}
@@ -1,19 +1,22 @@
package service
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"log/slog"
"net"
"strings"
"github.com/rain-kl/openflare/pkg/geoip"
pkggeoip "github.com/rain-kl/openflare/pkg/geoip"
)
var accessLogGeoProviderFactory = func() (geoip.GeoIPService, error) {
return geoip.NewMaxMindGeoIPService()
var accessLogGeoProviderFactory = func() (pkggeoip.GeoIPService, error) {
return pkggeoip.NewMaxMindGeoIPService()
}
type accessLogRegionResolver struct {
provider geoip.GeoIPService
provider pkggeoip.GeoIPService
cache map[string]string
}
@@ -0,0 +1,143 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
const (
agentTokenPositiveCacheTTL = 2 * time.Minute
agentTokenNegativeCacheTTL = 10 * time.Minute
)
type cachedAgentNode struct {
node *model.OpenFlareNode
expiresAt time.Time
}
type accessTokenAuthCache struct {
mu sync.RWMutex
positive map[string]cachedAgentNode
negative map[string]time.Time
now func() time.Time
loadNodeByToken func(context.Context, string) (*model.OpenFlareNode, error)
}
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)
},
}
}
func (c *accessTokenAuthCache) authenticate(ctx context.Context, token string) (*model.OpenFlareNode, error) {
now := c.now()
if node, ok := c.getNode(token, now); ok {
return node, nil
}
if c.isMissing(token, now) {
return nil, gorm.ErrRecordNotFound
}
node, err := c.loadNodeByToken(ctx, token)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL))
}
return nil, err
}
c.storeNode(token, node)
return cloneNode(node), nil
}
func (c *accessTokenAuthCache) getNode(token string, now time.Time) (*model.OpenFlareNode, bool) {
c.mu.RLock()
entry, ok := c.positive[token]
c.mu.RUnlock()
if !ok {
return nil, false
}
if now.After(entry.expiresAt) {
c.mu.Lock()
delete(c.positive, token)
c.mu.Unlock()
return nil, false
}
return cloneNode(entry.node), true
}
func (c *accessTokenAuthCache) isMissing(token string, now time.Time) bool {
c.mu.RLock()
expiresAt, ok := c.negative[token]
c.mu.RUnlock()
if !ok {
return false
}
if now.After(expiresAt) {
c.mu.Lock()
delete(c.negative, token)
c.mu.Unlock()
return false
}
return true
}
func (c *accessTokenAuthCache) storeNode(token string, node *model.OpenFlareNode) {
if token == "" || node == nil {
return
}
c.mu.Lock()
defer c.mu.Unlock()
delete(c.negative, token)
c.positive[token] = cachedAgentNode{
node: cloneNode(node),
expiresAt: c.now().Add(agentTokenPositiveCacheTTL),
}
}
func (c *accessTokenAuthCache) storeMissing(token string, expiresAt time.Time) {
if token == "" {
return
}
c.mu.Lock()
defer c.mu.Unlock()
delete(c.positive, token)
c.negative[token] = expiresAt
}
func (c *accessTokenAuthCache) reset() {
c.mu.Lock()
defer c.mu.Unlock()
c.positive = make(map[string]cachedAgentNode)
c.negative = make(map[string]time.Time)
}
// ResetAuthCacheForTest clears the in-memory access token cache for integration tests.
func ResetAuthCacheForTest() {
tokenCache.reset()
}
// AuthenticateAccessToken validates X-Agent-Token against of_nodes.access_token.
func AuthenticateAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
token = strings.TrimSpace(token)
if token == "" {
return nil, errors.New(errMissingAgentToken)
}
return tokenCache.authenticate(ctx, token)
}
@@ -0,0 +1,103 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty"
"gorm.io/gorm"
)
type configVersionRecord struct {
ID uint `gorm:"primaryKey"`
Version string `gorm:"column:version"`
SnapshotJSON string `gorm:"column:snapshot_json"`
SupportFilesJSON string `gorm:"column:support_files_json"`
Checksum string `gorm:"column:checksum"`
IsActive bool `gorm:"column:is_active"`
CreatedAt time.Time `gorm:"column:created_at"`
}
func (configVersionRecord) TableName() string {
return "of_config_versions"
}
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
version, err := loadActiveConfigVersion(ctx)
if err != nil {
return nil, err
}
return &ActiveConfigMeta{
Version: version.Version,
Checksum: version.Checksum,
}, nil
}
func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) {
version, err := loadActiveConfigVersion(ctx)
if err != nil {
return nil, err
}
var supportFiles []SupportFile
if strings.TrimSpace(version.SupportFilesJSON) != "" {
if err = json.Unmarshal([]byte(version.SupportFilesJSON), &supportFiles); err != nil {
return nil, err
}
}
return &ConfigResponse{
Version: version.Version,
Checksum: version.Checksum,
SourceConfigJSON: version.SnapshotJSON,
SupportFiles: sourceSupportFiles(supportFiles),
CreatedAt: version.CreatedAt,
}, nil
}
func loadActiveConfigVersion(ctx context.Context) (*configVersionRecord, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New("database not initialized")
}
version := &configVersionRecord{}
err := conn.Where("is_active = ?", true).Order("id desc").First(version).Error
if err != nil {
return nil, err
}
return version, nil
}
func sourceSupportFiles(files []SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
result := make([]SupportFile, 0, len(files))
for _, file := range files {
if isRuntimeGeneratedSupportFile(file.Path) {
continue
}
result = append(result, file)
}
return result
}
func isRuntimeGeneratedSupportFile(path string) bool {
switch strings.TrimSpace(path) {
case "pow_config.json", "waf_config.json", openrestyrender.SourceConfigFileName:
return true
default:
return false
}
}
func isActiveConfigNotFound(err error) bool {
return errors.Is(err, gorm.ErrRecordNotFound)
}
@@ -0,0 +1,46 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"testing"
openrestyrender "github.com/rain-kl/openflare/pkg/render/openresty"
)
func TestIsRuntimeGeneratedSupportFile(t *testing.T) {
tests := []struct {
path string
want bool
}{
{path: "pow_config.json", want: true},
{path: "waf_config.json", want: true},
{path: openrestyrender.SourceConfigFileName, want: true},
{path: "runtime/custom.json", want: false},
{path: "certs/example.pem", want: false},
}
for _, tc := range tests {
if got := isRuntimeGeneratedSupportFile(tc.path); got != tc.want {
t.Fatalf("isRuntimeGeneratedSupportFile(%q) = %v, want %v", tc.path, got, tc.want)
}
}
}
func TestSourceSupportFilesFiltersRuntimeGeneratedFiles(t *testing.T) {
files := []SupportFile{
{Path: "certs/example.pem", Content: "pem"},
{Path: "pow_config.json", Content: "{}"},
{Path: "waf_config.json", Content: "{}"},
{Path: openrestyrender.SourceConfigFileName, Content: "{}"},
{Path: "routes/extra.json", Content: "{}"},
}
filtered := sourceSupportFiles(files)
if len(filtered) != 2 {
t.Fatalf("expected 2 support files, got %d: %+v", len(filtered), filtered)
}
if filtered[0].Path != "certs/example.pem" || filtered[1].Path != "routes/extra.json" {
t.Fatalf("unexpected filtered files: %+v", filtered)
}
}
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
const (
errMissingAgentToken = "缺少 Agent Token"
errInvalidAgentToken = "无权进行此操作,Agent Token 无效"
errInvalidDiscoveryToken = "无权进行此操作,注册 Token 无效"
errNodeMissingFromContext = "Node object missing from context"
errNoActiveConfig = "当前没有激活版本"
errNodeNotFound = "节点不存在"
errNodeIDRequired = "node_id 不能为空"
errVersionRequired = "version 不能为空"
errInvalidApplyResult = "result 仅支持 success、warning 或 failed"
errIPRequired = "ip 不能为空"
errIPInvalid = "ip 格式无效"
errAgentVersionRequired = "version 不能为空"
errNodeIDConflict = "节点标识生成冲突,请重试"
)
@@ -0,0 +1,308 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"crypto/rand"
"encoding/hex"
"net"
"strings"
"time"
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
openrestyStatusHealthy = "healthy"
openrestyStatusUnhealthy = "unhealthy"
openrestyStatusUnknown = "unknown"
releaseChannelStable = "stable"
)
func newRandomToken() (string, error) {
buf := make([]byte, 16)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return hex.EncodeToString(buf), nil
}
func newServerNodeID() (string, error) {
token, err := newRandomToken()
if err != nil {
return "", err
}
return "node-" + token, nil
}
func normalizeOpenrestyStatus(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case openrestyStatusHealthy:
return openrestyStatusHealthy
case openrestyStatusUnhealthy:
return openrestyStatusUnhealthy
default:
return openrestyStatusUnknown
}
}
func normalizeNodePayload(payload NodePayload) NodePayload {
payload.Name = strings.TrimSpace(payload.Name)
payload.IP = strings.TrimSpace(payload.IP)
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.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
payload.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
return payload
}
func validateNodePayload(payload NodePayload) error {
if payload.IP == "" {
return errPayload(errIPRequired)
}
if net.ParseIP(payload.IP) == nil {
return errPayload(errIPInvalid)
}
if payload.Version == "" {
return errPayload(errAgentVersionRequired)
}
return nil
}
type payloadError string
func (e payloadError) Error() string { return string(e) }
func errPayload(message string) error { return payloadError(message) }
func applyNodeRuntime(node *model.OpenFlareNode, payload NodePayload, preserveName bool) {
if !preserveName || strings.TrimSpace(node.Name) == "" {
if strings.TrimSpace(payload.Name) != "" {
node.Name = strings.TrimSpace(payload.Name)
}
}
if !node.IPManualOverride {
node.IP = strings.TrimSpace(payload.IP)
}
node.Version = strings.TrimSpace(payload.Version)
node.ExtVersion = strings.TrimSpace(payload.ExtVersion)
node.OpenrestyStatus = normalizeOpenrestyStatus(payload.OpenrestyStatus)
node.OpenrestyMessage = truncateForDatabase(payload.OpenrestyMessage, 16000)
node.Status = nodeStatusOnline
node.CurrentVersion = strings.TrimSpace(payload.CurrentVersion)
now := time.Now()
node.LastSeenAt = &now
node.LastError = truncateForDatabase(payload.LastError, 16000)
if !node.GeoManualOverride {
applyGeoInfoFromIP(node, node.IP)
}
}
func applyGeoInfoFromIP(node *model.OpenFlareNode, rawIP string) {
if node == nil {
return
}
node.GeoName = ""
node.GeoLatitude = nil
node.GeoLongitude = nil
ip := net.ParseIP(strings.TrimSpace(rawIP))
if ip == nil {
return
}
info, err := ofgeoip.GeoInfoFromIP(ip)
if err != nil || info == nil {
return
}
if strings.TrimSpace(info.Name) != "" {
node.GeoName = strings.TrimSpace(info.Name)
}
if info.Latitude != nil && info.Longitude != nil {
node.GeoLatitude = cloneCoordinate(info.Latitude)
node.GeoLongitude = cloneCoordinate(info.Longitude)
}
}
func cloneCoordinate(value *float64) *float64 {
if value == nil {
return nil
}
cloned := *value
return &cloned
}
func truncateForDatabase(value string, max int) string {
if max <= 0 {
return ""
}
runes := []rune(strings.TrimSpace(value))
if len(runes) <= max {
return string(runes)
}
return string(runes[:max])
}
func resolveReportedNodeIP(reportedIP string, remoteAddr string) string {
reported := normalizeIP(reportedIP)
remote := normalizeRemoteAddr(remoteAddr)
if reported == "" {
return remote
}
if isPublicNodeIP(reported) {
return reported
}
if isPublicNodeIP(remote) {
return remote
}
return reported
}
func normalizeIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
host := raw
if strings.Contains(raw, ":") {
if h, _, err := net.SplitHostPort(raw); err == nil {
host = h
}
}
host = strings.TrimPrefix(host, "[")
host = strings.TrimSuffix(host, "]")
if ip := net.ParseIP(host); ip != nil {
return ip.String()
}
return ""
}
func normalizeRemoteAddr(remoteAddr string) string {
remoteAddr = strings.TrimSpace(remoteAddr)
if remoteAddr == "" {
return ""
}
host, _, err := net.SplitHostPort(remoteAddr)
if err != nil {
return normalizeIP(remoteAddr)
}
return normalizeIP(host)
}
func isPublicNodeIP(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
if ip == nil {
return false
}
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
return false
}
return true
}
func buildAgentSettings(node *model.OpenFlareNode, updateNow bool, updateChannel string, updateTag string, restartOpenrestyNow bool) *Settings {
autoUpdate := false
if node != nil {
autoUpdate = node.AutoUpdateEnabled
}
if strings.TrimSpace(updateChannel) == "" {
updateChannel = releaseChannelStable
}
return &Settings{
HeartbeatInterval: model.AgentHeartbeatInterval,
WebsocketUpgradeEnabled: model.AgentWebsocketUpgradeEnabled,
AutoUpdate: autoUpdate,
UpdateRepo: model.AgentUpdateRepo,
UpdateNow: updateNow,
UpdateChannel: updateChannel,
UpdateTag: strings.TrimSpace(updateTag),
RestartOpenrestyNow: restartOpenrestyNow,
}
}
func collectHeartbeatChanges(previous *model.OpenFlareNode, current *model.OpenFlareNode) map[string]any {
if previous == nil || current == nil {
return map[string]any{}
}
changes := make(map[string]any)
appendIfChanged := func(key string, before any, after any) {
if before != after {
changes[key] = after
}
}
appendIfChanged("name", previous.Name, current.Name)
appendIfChanged("ip", previous.IP, current.IP)
appendIfChanged("geo_name", previous.GeoName, current.GeoName)
appendIfChanged("version", previous.Version, current.Version)
appendIfChanged("ext_version", previous.ExtVersion, current.ExtVersion)
appendIfChanged("openresty_status", previous.OpenrestyStatus, current.OpenrestyStatus)
appendIfChanged("openresty_message", previous.OpenrestyMessage, current.OpenrestyMessage)
appendIfChanged("status", previous.Status, current.Status)
appendIfChanged("current_version", previous.CurrentVersion, current.CurrentVersion)
appendIfChanged("last_error", previous.LastError, current.LastError)
appendIfChanged("update_requested", previous.UpdateRequested, current.UpdateRequested)
appendIfChanged("update_channel", previous.UpdateChannel, current.UpdateChannel)
appendIfChanged("update_tag", previous.UpdateTag, current.UpdateTag)
appendIfChanged("restart_openresty_requested", previous.RestartOpenrestyRequested, current.RestartOpenrestyRequested)
if !coordinatesEqual(previous.GeoLatitude, current.GeoLatitude) {
changes["geo_latitude"] = current.GeoLatitude
}
if !coordinatesEqual(previous.GeoLongitude, current.GeoLongitude) {
changes["geo_longitude"] = current.GeoLongitude
}
if !lastSeenAtEqual(previous.LastSeenAt, current.LastSeenAt) {
changes["last_seen_at"] = current.LastSeenAt
}
return changes
}
func coordinatesEqual(before *float64, after *float64) bool {
if before == nil || after == nil {
return before == after
}
return *before == *after
}
func lastSeenAtEqual(before *time.Time, after *time.Time) bool {
if before == nil || after == nil {
return before == after
}
return before.Equal(*after)
}
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.Checksum = strings.TrimSpace(payload.Checksum)
payload.MainConfigChecksum = strings.TrimSpace(payload.MainConfigChecksum)
payload.RouteConfigChecksum = strings.TrimSpace(payload.RouteConfigChecksum)
return payload
}
func isUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
// RefreshAccessTokenCache updates the in-memory node cache after heartbeat mutations.
func RefreshAccessTokenCache(ctx context.Context, node *model.OpenFlareNode) {
if node == nil {
return
}
tokenCache.storeNode(node.AccessToken, cloneNode(node))
}
func cloneNode(node *model.OpenFlareNode) *model.OpenFlareNode {
if node == nil {
return nil
}
cloned := *node
return &cloned
}
@@ -0,0 +1,131 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"net"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
pkggeoip "github.com/rain-kl/openflare/pkg/geoip"
)
type fakeGeoIPProvider struct {
info *pkggeoip.GeoInfo
}
func (f *fakeGeoIPProvider) Name() string { return "fake-geoip" }
func (f *fakeGeoIPProvider) GetGeoInfo(ip net.IP) (*pkggeoip.GeoInfo, error) {
return f.info, nil
}
func (f *fakeGeoIPProvider) UpdateDatabase() error { return nil }
func (f *fakeGeoIPProvider) Close() error { return nil }
func withFakeGeoIPProvider(t *testing.T, info *pkggeoip.GeoInfo) {
t.Helper()
previous := pkggeoip.CurrentProvider
pkggeoip.CurrentProvider = &fakeGeoIPProvider{info: info}
t.Cleanup(func() {
pkggeoip.CurrentProvider = previous
})
}
func geoipFloat(value float64) *float64 {
return &value
}
func TestApplyGeoInfoFromIP(t *testing.T) {
latitude := 31.2304
longitude := 121.4737
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
Name: "Shanghai",
Latitude: geoipFloat(latitude),
Longitude: geoipFloat(longitude),
})
node := &model.OpenFlareNode{IP: "203.0.113.10"}
applyGeoInfoFromIP(node, node.IP)
if node.GeoName != "Shanghai" {
t.Fatalf("expected geo_name Shanghai, got %q", node.GeoName)
}
if node.GeoLatitude == nil || *node.GeoLatitude != latitude {
t.Fatalf("unexpected geo_latitude: %+v", node.GeoLatitude)
}
if node.GeoLongitude == nil || *node.GeoLongitude != longitude {
t.Fatalf("unexpected geo_longitude: %+v", node.GeoLongitude)
}
}
func TestApplyGeoInfoFromIPSkipsInvalidIP(t *testing.T) {
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{Name: "Should Not Apply"})
node := &model.OpenFlareNode{
IP: "203.0.113.10",
GeoName: "Existing",
GeoLatitude: geoipFloat(1),
GeoLongitude: geoipFloat(2),
}
applyGeoInfoFromIP(node, "not-an-ip")
if node.GeoName != "" || node.GeoLatitude != nil || node.GeoLongitude != nil {
t.Fatalf("expected geo fields to be cleared on invalid IP, got %+v", node)
}
}
func TestApplyNodeRuntimeRespectsGeoManualOverride(t *testing.T) {
withFakeGeoIPProvider(t, &pkggeoip.GeoInfo{
Name: "Shanghai",
Latitude: geoipFloat(31.2304),
Longitude: geoipFloat(121.4737),
})
node := &model.OpenFlareNode{
GeoManualOverride: true,
GeoName: "Manual",
GeoLatitude: geoipFloat(10),
GeoLongitude: geoipFloat(20),
}
applyNodeRuntime(node, NodePayload{
IP: "203.0.113.10",
Version: "1.0.0",
}, true)
if node.GeoName != "Manual" {
t.Fatalf("expected manual geo_name to be preserved, got %q", node.GeoName)
}
if node.GeoLatitude == nil || *node.GeoLatitude != 10 {
t.Fatalf("expected manual geo_latitude to be preserved, got %+v", node.GeoLatitude)
}
}
func TestCollectHeartbeatChangesTracksGeoFields(t *testing.T) {
before := &model.OpenFlareNode{
IP: "10.0.0.1",
GeoName: "Old Region",
}
after := &model.OpenFlareNode{
IP: "203.0.113.10",
GeoName: "New Region",
GeoLatitude: geoipFloat(31.2304),
GeoLongitude: geoipFloat(121.4737),
}
changes := collectHeartbeatChanges(before, after)
if changes["ip"] != after.IP {
t.Fatalf("expected ip change, got %+v", changes)
}
if changes["geo_name"] != after.GeoName {
t.Fatalf("expected geo_name change, got %+v", changes)
}
if changes["geo_latitude"] != after.GeoLatitude {
t.Fatalf("expected geo_latitude change, got %+v", changes)
}
if changes["geo_longitude"] != after.GeoLongitude {
t.Fatalf("expected geo_longitude change, got %+v", changes)
}
}
@@ -0,0 +1,224 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/node"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
// RegisterWithAccessToken registers an agent on a reserved node token.
func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*RegistrationResponse, error) {
payload = normalizeNodePayload(payload)
if authNode == nil {
return nil, errors.New(errNodeNotFound)
}
if err := validateNodePayload(payload); err != nil {
return nil, err
}
applyNodeRuntime(authNode, payload, true)
if err := model.SaveOpenFlareNode(ctx, authNode); err != nil {
return nil, err
}
RefreshAccessTokenCache(ctx, authNode)
return &RegistrationResponse{
NodeID: authNode.NodeID,
AccessToken: authNode.AccessToken,
Name: authNode.Name,
}, nil
}
// RegisterWithDiscovery registers a new node using the global discovery token.
func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*RegistrationResponse, error) {
payload = normalizeNodePayload(payload)
if err := validateNodePayload(payload); err != nil {
return nil, err
}
nodeID, err := newServerNodeID()
if err != nil {
return nil, err
}
accessToken, err := newRandomToken()
if err != nil {
return nil, err
}
nodeName := payload.Name
if nodeName == "" {
nodeName = nodeID
}
record := &model.OpenFlareNode{
NodeID: nodeID,
Name: nodeName,
AccessToken: accessToken,
Status: nodeStatusOnline,
NodeType: "edge_node",
CapabilitiesJSON: "[]",
UpdateChannel: releaseChannelStable,
}
applyNodeRuntime(record, payload, false)
if err = model.CreateOpenFlareNode(ctx, record); err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errNodeIDConflict)
}
return nil, err
}
RefreshAccessTokenCache(ctx, record)
return &RegistrationResponse{
NodeID: record.NodeID,
AccessToken: record.AccessToken,
Name: record.Name,
}, nil
}
// HeartbeatNode updates runtime state and returns agent settings.
func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload NodePayload) (*HeartbeatResponse, error) {
if authNode == nil {
return nil, errors.New(errNodeNotFound)
}
payload.NodeID = authNode.NodeID
payload = normalizeNodePayload(payload)
if err := validateNodePayload(payload); err != nil {
return nil, err
}
previous := *authNode
updateNow := authNode.UpdateRequested
restartOpenrestyNow := authNode.RestartOpenrestyRequested
updateChannel := strings.TrimSpace(authNode.UpdateChannel)
updateTag := strings.TrimSpace(authNode.UpdateTag)
applyNodeRuntime(authNode, payload, true)
authNode.UpdateRequested = false
authNode.UpdateChannel = releaseChannelStable
authNode.UpdateTag = ""
authNode.RestartOpenrestyRequested = false
changes := collectHeartbeatChanges(&previous, authNode)
if len(changes) > 0 {
fields := make([]string, 0, len(changes))
for field := range changes {
fields = append(fields, field)
}
if err := model.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil {
return nil, err
}
}
RefreshAccessTokenCache(ctx, authNode)
reportedAt := time.Now()
if authNode.LastSeenAt != nil {
reportedAt = *authNode.LastSeenAt
}
PersistHeartbeatObservability(ctx, authNode.NodeID, payload, reportedAt)
activeConfig, err := getActiveConfigMeta(ctx)
if err != nil && !isActiveConfigNotFound(err) {
return nil, err
}
wafIPGroups, err := ChangedWAFIPGroupsForAgent(ctx, nil, payload.WAFIPGroupChecksums)
if err != nil {
return nil, err
}
return &HeartbeatResponse{
Node: authNode,
AgentSettings: buildAgentSettings(authNode, updateNow, updateChannel, updateTag, restartOpenrestyNow),
ActiveConfig: activeConfig,
WAFIPGroups: wafIPGroups,
}, nil
}
// GetActiveConfig returns the active configuration for an agent.
func GetActiveConfig(ctx context.Context) (*ConfigResponse, error) {
config, err := getActiveConfigForAgent(ctx)
if err != nil {
if isActiveConfigNotFound(err) {
return nil, errors.New(errNoActiveConfig)
}
return nil, err
}
return config, nil
}
// SyncWAFIPGroups returns WAF IP groups whose checksums differ from the agent state.
func SyncWAFIPGroups(ctx context.Context, input WAFIPGroupSyncInput) (*WAFIPGroupSyncResult, error) {
groups, err := ChangedWAFIPGroupsForAgent(ctx, input.IDs, input.Checksums)
if err != nil {
return nil, err
}
return &WAFIPGroupSyncResult{Groups: groups}, nil
}
// ReportApplyLog records an agent apply result.
func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFlareApplyLog, error) {
now := time.Now()
payload = normalizeApplyLogPayload(payload)
if payload.NodeID == "" {
return nil, errors.New(errNodeIDRequired)
}
if payload.Version == "" {
return nil, errors.New(errVersionRequired)
}
if payload.Result != applyResultOK && payload.Result != applyResultWarn && payload.Result != applyResultFailed {
return nil, errors.New(errInvalidApplyResult)
}
log := &model.OpenFlareApplyLog{
NodeID: payload.NodeID,
Version: payload.Version,
Result: payload.Result,
Message: payload.Message,
Checksum: payload.Checksum,
MainConfigChecksum: payload.MainConfigChecksum,
RouteConfigChecksum: payload.RouteConfigChecksum,
SupportFileCount: payload.SupportFileCount,
CreatedAt: now,
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New("database not initialized")
}
err := conn.Transaction(func(tx *gorm.DB) error {
record := &model.OpenFlareNode{}
if err := tx.Where("node_id = ?", payload.NodeID).First(record).Error; err != nil {
return err
}
record.Status = nodeStatusOnline
record.LastSeenAt = &now
if payload.Result == applyResultOK {
record.CurrentVersion = payload.Version
record.LastError = ""
} else {
record.LastError = payload.Message
}
if err := tx.Create(log).Error; err != nil {
return err
}
return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).Error
})
if err != nil {
return nil, err
}
return log, nil
}
// ValidateDiscoveryToken delegates to the node package discovery token helper.
func ValidateDiscoveryToken(ctx context.Context, token string) error {
return node.ValidateDiscoveryToken(ctx, token)
}
@@ -0,0 +1,59 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"strings"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
)
const (
agentTokenHeader = "X-Agent-Token"
agentNodeContextKey = "agent_node"
)
// AgentAuth validates X-Agent-Token against of_nodes.access_token.
func AgentAuth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
node, err := AuthenticateAccessToken(c.Request.Context(), token)
if err != nil {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
c.Set(agentNodeContextKey, node)
c.Next()
}
}
// AgentRegisterAuth accepts either a node access token or the global discovery token.
func AgentRegisterAuth() gin.HandlerFunc {
return func(c *gin.Context) {
token := strings.TrimSpace(c.GetHeader(agentTokenHeader))
if node, err := AuthenticateAccessToken(c.Request.Context(), token); err == nil {
c.Set(agentNodeContextKey, node)
c.Next()
return
}
if err := ValidateDiscoveryToken(c.Request.Context(), token); err != nil {
response.AbortUnauthorized(c, errInvalidDiscoveryToken)
return
}
c.Set("discovery_enabled", true)
c.Next()
}
}
// AgentNodeFromContext returns the authenticated agent node.
func AgentNodeFromContext(c *gin.Context) (*model.OpenFlareNode, bool) {
value, ok := c.Get(agentNodeContextKey)
if !ok {
return nil, false
}
node, ok := value.(*model.OpenFlareNode)
return node, ok
}
@@ -0,0 +1,199 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/option"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupAgentAuthTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareNode{},
&model.OpenFlareOption{},
))
db.SetDB(sqliteDB)
option.ResetInitializationForTest()
tokenCache.reset()
return func() {
db.SetDB(nil)
option.ResetInitializationForTest()
tokenCache.reset()
}
}
func TestAuthenticateAccessToken(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-auth-1",
Name: "edge",
AccessToken: "valid-agent-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
t.Run("valid token", func(t *testing.T) {
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
require.NoError(t, err)
assert.Equal(t, "node-auth-1", node.NodeID)
})
t.Run("cached token", func(t *testing.T) {
originalLoader := tokenCache.loadNodeByToken
t.Cleanup(func() {
tokenCache.loadNodeByToken = originalLoader
})
tokenCache.loadNodeByToken = func(context.Context, string) (*model.OpenFlareNode, error) {
t.Fatal("db should not be queried for cached token")
return nil, nil
}
node, err := AuthenticateAccessToken(ctx, "valid-agent-token")
require.NoError(t, err)
assert.Equal(t, "node-auth-1", node.NodeID)
})
t.Run("missing token", func(t *testing.T) {
_, err := AuthenticateAccessToken(ctx, "")
require.Error(t, err)
assert.Contains(t, err.Error(), errMissingAgentToken)
})
t.Run("invalid token", func(t *testing.T) {
_, err := AuthenticateAccessToken(ctx, "invalid-token")
require.Error(t, err)
})
}
func TestAgentAuthMiddleware(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-mw-1",
Name: "edge",
AccessToken: "middleware-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
router := testhelper.NewTestGinEngine()
router.GET("/protected", AgentAuth(), func(c *gin.Context) {
node, ok := AgentNodeFromContext(c)
if !ok {
c.Status(http.StatusInternalServerError)
return
}
c.JSON(http.StatusOK, response.OK(gin.H{"node_id": node.NodeID}))
})
t.Run("authorized request", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set(agentTokenHeader, "middleware-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
assert.Empty(t, apiResp.ErrorMsg)
})
t.Run("unauthorized request", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set(agentTokenHeader, "bad-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusUnauthorized, resp.Code)
})
}
func TestAgentRegisterAuthMiddleware(t *testing.T) {
cleanup := setupAgentAuthTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{
NodeID: "node-register-1",
Name: "edge",
AccessToken: "existing-node-token",
Status: nodeStatusOnline,
LastSeenAt: &now,
NodeType: "edge_node",
}).Error)
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 {
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "node", "node_id": node.NodeID}))
return
}
if _, ok := c.Get("discovery_enabled"); ok {
c.JSON(http.StatusOK, response.OK(gin.H{"mode": "discovery"}))
return
}
c.Status(http.StatusInternalServerError)
})
t.Run("existing node token", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/register", nil)
req.Header.Set(agentTokenHeader, "existing-node-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
data, ok := apiResp.Data.(map[string]any)
require.True(t, ok)
assert.Equal(t, "node", data["mode"])
})
t.Run("discovery token", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/register", nil)
req.Header.Set(agentTokenHeader, "discovery-token")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assert.Equal(t, http.StatusOK, resp.Code)
var apiResp response.Any
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &apiResp))
data, ok := apiResp.Data.(map[string]any)
require.True(t, ok)
assert.Equal(t, "discovery", data["mode"])
})
}
@@ -1,28 +1,36 @@
package service
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"errors"
"log/slog"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/model"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"go.uber.org/zap"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
NodeHealthEventStatusActive = "active"
NodeHealthEventStatusResolved = "resolved"
NodeHealthSeverityInfo = "info"
NodeHealthSeverityWarning = "warning"
NodeHealthSeverityCritical = "critical"
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
nodeAccessLogPathMaxLength = 100
healthEventStatusActive = "active"
healthEventStatusResolved = "resolved"
healthSeverityInfo = "info"
healthSeverityWarning = "warning"
healthSeverityCritical = "critical"
nodeAccessLogRetentionDays = 90
nodeAccessLogRetentionWindow = nodeAccessLogRetentionDays * 24 * time.Hour
accessLogPathMaxLength = 100
)
type AgentNodeSystemProfile struct {
// NodeSystemProfile is the agent-reported system profile.
type NodeSystemProfile struct {
Hostname string `json:"hostname"`
OSName string `json:"os_name"`
OSVersion string `json:"os_version"`
@@ -36,7 +44,8 @@ type AgentNodeSystemProfile struct {
ReportedAtUnix int64 `json:"reported_at_unix"`
}
type AgentNodeMetricSnapshot struct {
// NodeMetricSnapshot is the agent-reported capacity snapshot.
type NodeMetricSnapshot struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
@@ -49,14 +58,16 @@ type AgentNodeMetricSnapshot struct {
NetworkTxBytes int64 `json:"network_tx_bytes"`
}
type AgentNodeOpenrestyObservation struct {
// NodeOpenrestyObservation is the agent-reported openresty network observation.
type NodeOpenrestyObservation struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
OpenrestyConnections int64 `json:"openresty_connections"`
}
type AgentNodeTrafficReport struct {
// NodeTrafficReport is the agent-reported traffic window.
type NodeTrafficReport struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
RequestCount int64 `json:"request_count"`
@@ -67,7 +78,8 @@ type AgentNodeTrafficReport struct {
SourceCountries map[string]int64 `json:"source_countries"`
}
type AgentNodeAccessLog struct {
// NodeAccessLog is a single access log row from the agent.
type NodeAccessLog struct {
LoggedAtUnix int64 `json:"logged_at_unix"`
RemoteAddr string `json:"remote_addr"`
Host string `json:"host"`
@@ -75,15 +87,17 @@ type AgentNodeAccessLog struct {
StatusCode int `json:"status_code"`
}
type AgentBufferedObservabilityRecord struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
Snapshot *AgentNodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *AgentNodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *AgentNodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []AgentNodeAccessLog `json:"access_logs,omitempty"`
// BufferedObservabilityRecord is a buffered observability window from the agent.
type BufferedObservabilityRecord struct {
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
}
type AgentNodeHealthEvent struct {
// NodeHealthEvent is an agent-reported health event.
type NodeHealthEvent struct {
EventType string `json:"event_type"`
Severity string `json:"severity"`
Message string `json:"message"`
@@ -91,15 +105,26 @@ type AgentNodeHealthEvent struct {
Metadata map[string]string `json:"metadata"`
}
func persistHeartbeatObservability(nodeID string, payload AgentNodePayload, reportedAt time.Time) {
// PersistHeartbeatObservability stores profile, snapshots, traffic, access logs, and health events.
func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload NodePayload, reportedAt time.Time) {
if strings.TrimSpace(nodeID) == "" {
return
}
if payload.Profile == nil && payload.Snapshot == nil && payload.TrafficReport == nil && len(payload.AccessLogs) == 0 && len(payload.BufferedObservability) == 0 && payload.HealthEvents == nil {
if payload.Profile == nil &&
payload.Snapshot == nil &&
payload.TrafficReport == nil &&
len(payload.AccessLogs) == 0 &&
len(payload.BufferedObservability) == 0 &&
payload.HealthEvents == nil {
return
}
if err := model.DB.Transaction(func(tx *gorm.DB) error {
conn := db.DB(ctx)
if conn == nil {
return
}
if err := conn.Transaction(func(tx *gorm.DB) error {
if err := persistNodeSystemProfile(tx, nodeID, payload.Profile, reportedAt); err != nil {
return err
}
@@ -125,11 +150,11 @@ func persistHeartbeatObservability(nodeID string, payload AgentNodePayload, repo
}
return nil
}); err != nil {
slog.Error("persist heartbeat observability failed", "node_id", nodeID, "error", err)
zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err))
}
}
func persistBufferedObservability(tx *gorm.DB, nodeID string, records []AgentBufferedObservabilityRecord, reportedAt time.Time) error {
func persistBufferedObservability(tx *gorm.DB, nodeID string, records []BufferedObservabilityRecord, reportedAt time.Time) error {
for _, record := range records {
if err := persistNodeMetricSnapshot(tx, nodeID, record.Snapshot, reportedAt); err != nil {
return err
@@ -147,11 +172,11 @@ func persistBufferedObservability(tx *gorm.DB, nodeID string, records []AgentBuf
return nil
}
func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *AgentNodeSystemProfile, reportedAt time.Time) error {
func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *NodeSystemProfile, reportedAt time.Time) error {
if profile == nil {
return nil
}
record := &model.NodeSystemProfile{
record := &model.OpenFlareNodeSystemProfile{
NodeID: nodeID,
Hostname: strings.TrimSpace(profile.Hostname),
OSName: strings.TrimSpace(profile.OSName),
@@ -165,14 +190,30 @@ func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *AgentNodeSyst
UptimeSeconds: profile.UptimeSeconds,
ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt),
}
return tx.Model(&model.NodeSystemProfile{}).Where("node_id = ?", nodeID).Assign(record).FirstOrCreate(record).Error
return tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "node_id"}},
DoUpdates: clause.AssignmentColumns([]string{
"hostname",
"os_name",
"os_version",
"kernel_version",
"architecture",
"cpu_model",
"cpu_cores",
"total_memory_bytes",
"total_disk_bytes",
"uptime_seconds",
"reported_at",
"updated_at",
}),
}).Create(record).Error
}
func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *AgentNodeMetricSnapshot, reportedAt time.Time) error {
func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *NodeMetricSnapshot, reportedAt time.Time) error {
if snapshot == nil {
return nil
}
record := &model.NodeMetricSnapshot{
record := &model.OpenFlareMetricSnapshot{
NodeID: nodeID,
CapturedAt: timeFromUnix(snapshot.CapturedAtUnix, reportedAt),
CPUUsagePercent: snapshot.CPUUsagePercent,
@@ -185,7 +226,7 @@ func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *AgentNodeMe
NetworkRxBytes: snapshot.NetworkRxBytes,
NetworkTxBytes: snapshot.NetworkTxBytes,
}
exists, err := model.NodeMetricSnapshotExists(tx, nodeID, record.CapturedAt)
exists, err := metricSnapshotExists(tx, nodeID, record.CapturedAt)
if err != nil {
return err
}
@@ -195,11 +236,11 @@ func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *AgentNodeMe
return tx.Create(record).Error
}
func persistNodeOpenrestyObservation(tx *gorm.DB, nodeID string, obs *AgentNodeOpenrestyObservation, reportedAt time.Time) error {
func persistNodeOpenrestyObservation(tx *gorm.DB, nodeID string, obs *NodeOpenrestyObservation, reportedAt time.Time) error {
if obs == nil {
return nil
}
record := &model.NodeObservationOpenresty{
record := &model.OpenFlareNodeObservationOpenresty{
NodeID: nodeID,
CapturedAt: timeFromUnix(obs.CapturedAtUnix, reportedAt),
OpenrestyRxBytes: obs.OpenrestyRxBytes,
@@ -209,14 +250,14 @@ func persistNodeOpenrestyObservation(tx *gorm.DB, nodeID string, obs *AgentNodeO
return tx.Create(record).Error
}
func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *AgentNodeTrafficReport, reportedAt time.Time) error {
func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *NodeTrafficReport, reportedAt time.Time) error {
if report == nil {
return nil
}
if report.WindowEndedAtUnix > 0 && report.WindowStartedAtUnix > report.WindowEndedAtUnix {
return errors.New("traffic report window_started_at_unix 不能大于 window_ended_at_unix")
}
record := &model.NodeRequestReport{
record := &model.OpenFlareRequestReport{
NodeID: nodeID,
WindowStartedAt: timeFromUnix(report.WindowStartedAtUnix, reportedAt),
WindowEndedAt: timeFromUnix(report.WindowEndedAtUnix, reportedAt),
@@ -227,7 +268,7 @@ func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *AgentNodeTraff
TopDomainsJSON: marshalJSON(report.TopDomains),
SourceCountriesJSON: marshalJSON(report.SourceCountries),
}
exists, err := model.NodeRequestReportExists(tx, nodeID, record.WindowStartedAt, record.WindowEndedAt)
exists, err := requestReportExists(tx, nodeID, record.WindowStartedAt, record.WindowEndedAt)
if err != nil {
return err
}
@@ -237,7 +278,7 @@ func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *AgentNodeTraff
return tx.Create(record).Error
}
func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []AgentNodeAccessLog, reportedAt time.Time) error {
func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []NodeAccessLog, reportedAt time.Time) error {
if len(logs) == 0 {
return nil
}
@@ -249,19 +290,19 @@ func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []AgentNodeAccessLog
defer resolver.Close()
}
for _, item := range logs {
record := &model.NodeAccessLog{
record := &model.OpenFlareAccessLog{
NodeID: nodeID,
LoggedAt: timeFromUnix(item.LoggedAtUnix, reportedAt),
RemoteAddr: strings.TrimSpace(item.RemoteAddr),
Region: "",
Host: strings.TrimSpace(item.Host),
Path: truncateForDatabase(strings.TrimSpace(item.Path), nodeAccessLogPathMaxLength),
Path: truncateForDatabase(strings.TrimSpace(item.Path), accessLogPathMaxLength),
StatusCode: item.StatusCode,
}
if resolver != nil {
record.Region = resolver.Resolve(record.RemoteAddr)
}
exists, err := model.NodeAccessLogExists(tx, record)
exists, err := accessLogExists(tx, record)
if err != nil {
return err
}
@@ -272,16 +313,17 @@ func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []AgentNodeAccessLog
return err
}
}
_, err = model.DeleteNodeAccessLogsByNodeBefore(tx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow))
_, err = deleteAccessLogsByNodeBefore(tx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow))
return err
}
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentNodeHealthEvent, reportedAt time.Time) error {
return reconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error {
return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
}
func reconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentNodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
activeTypes := make(map[string]AgentNodeHealthEvent, len(events))
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types.
func ReconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
activeTypes := make(map[string]NodeHealthEvent, len(events))
for _, event := range events {
eventType := normalizeHealthEventType(event.EventType)
if eventType == "" {
@@ -300,8 +342,8 @@ func reconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentN
activeTypes[eventType] = event
}
var activeEvents []*model.NodeHealthEvent
query := tx.Where("node_id = ? AND status = ?", nodeID, NodeHealthEventStatusActive)
var activeEvents []*model.OpenFlareHealthEvent
query := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive)
if len(managedEventTypes) > 0 {
scopedTypes := make([]string, 0, len(managedEventTypes))
for eventType := range managedEventTypes {
@@ -319,7 +361,7 @@ func reconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentN
return err
}
activeByType := make(map[string]*model.NodeHealthEvent, len(activeEvents))
activeByType := make(map[string]*model.OpenFlareHealthEvent, len(activeEvents))
for _, event := range activeEvents {
activeByType[event.EventType] = event
}
@@ -338,11 +380,11 @@ func reconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentN
}
continue
}
record := &model.NodeHealthEvent{
record := &model.OpenFlareHealthEvent{
NodeID: nodeID,
EventType: eventType,
Severity: event.Severity,
Status: NodeHealthEventStatusActive,
Status: healthEventStatusActive,
Message: normalizeHealthEventMessage(event.Message),
FirstTriggeredAt: triggeredAt,
LastTriggeredAt: triggeredAt,
@@ -359,7 +401,7 @@ func reconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentN
continue
}
resolvedAt := reportedAt
existing.Status = NodeHealthEventStatusResolved
existing.Status = healthEventStatusResolved
existing.ReportedAt = reportedAt
existing.ResolvedAt = &resolvedAt
if err := tx.Save(existing).Error; err != nil {
@@ -370,6 +412,52 @@ func reconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentN
return nil
}
func metricSnapshotExists(tx *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) {
var count int64
if err := tx.Model(&model.OpenFlareMetricSnapshot{}).
Where("node_id = ? AND captured_at = ?", nodeID, capturedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
return count > 0, nil
}
func requestReportExists(tx *gorm.DB, nodeID string, windowStartedAt, windowEndedAt time.Time) (bool, error) {
var count int64
if err := tx.Model(&model.OpenFlareRequestReport{}).
Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, windowStartedAt, windowEndedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
return count > 0, nil
}
func deleteAccessLogsByNodeBefore(tx *gorm.DB, nodeID string, before time.Time) (int64, error) {
result := tx.Where("node_id = ? AND logged_at < ?", nodeID, before).Delete(&model.OpenFlareAccessLog{})
return result.RowsAffected, result.Error
}
func accessLogExists(tx *gorm.DB, record *model.OpenFlareAccessLog) (bool, error) {
var count int64
if err := tx.Model(&model.OpenFlareAccessLog{}).
Where(
"node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?",
record.NodeID,
record.LoggedAt,
record.RemoteAddr,
record.Host,
record.Path,
record.StatusCode,
).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
return count > 0, nil
}
func normalizeHealthEventType(eventType string) string {
eventType = strings.TrimSpace(strings.ToLower(eventType))
eventType = strings.ReplaceAll(eventType, " ", "_")
@@ -378,12 +466,12 @@ func normalizeHealthEventType(eventType string) string {
func normalizeHealthSeverity(severity string) string {
switch strings.ToLower(strings.TrimSpace(severity)) {
case NodeHealthSeverityCritical:
return NodeHealthSeverityCritical
case NodeHealthSeverityInfo:
return NodeHealthSeverityInfo
case healthSeverityCritical:
return healthSeverityCritical
case healthSeverityInfo:
return healthSeverityInfo
default:
return NodeHealthSeverityWarning
return healthSeverityWarning
}
}
@@ -398,6 +486,11 @@ func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
return time.Unix(unixSeconds, 0).UTC()
}
// MarshalJSON serializes a value for database JSON columns.
func MarshalJSON(value any) string {
return marshalJSON(value)
}
func marshalJSON(value any) string {
if value == nil {
return ""
@@ -0,0 +1,212 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/gin-gonic/gin"
)
// RegisterHandler registers or discovers an agent node.
// @Summary 注册或发现 Agent 节点
// @Description 使用节点 access token 重新注册,或使用全局 discovery token 发现新节点;请求头需携带 X-Agent-Token
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.NodePayload true "节点上报数据"
// @Success 200 {object} response.Any{data=agent.RegistrationResponse} "注册成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/nodes/register [post]
func RegisterHandler(c *gin.Context) {
var payload NodePayload
if !apiutil.BindJSON(c, &payload) {
return
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
var (
result *RegistrationResponse
err error
)
if authNode, ok := AgentNodeFromContext(c); ok {
result, err = RegisterWithAccessToken(c.Request.Context(), authNode, payload)
} else {
result, err = RegisterWithDiscovery(c.Request.Context(), payload)
}
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// HeartbeatHandler records agent heartbeat state.
// @Summary Agent 心跳上报
// @Description 上报节点状态、指标与健康事件,返回远程控制配置与活跃配置元信息
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.NodePayload true "心跳数据"
// @Success 200 {object} response.Any{data=agent.HeartbeatResponse} "心跳成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/nodes/heartbeat [post]
func HeartbeatHandler(c *gin.Context) {
var payload NodePayload
if !apiutil.BindJSON(c, &payload) {
return
}
payload.IP = resolveReportedNodeIP(payload.IP, c.Request.RemoteAddr)
authNode, ok := AgentNodeFromContext(c)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
heartbeat, err := HeartbeatNode(c.Request.Context(), authNode, payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(heartbeat))
}
// GetActiveConfigHandler returns the active configuration version.
// @Summary 获取活跃配置版本
// @Description 返回当前生效的完整配置包,供 Agent 拉取并应用
// @Tags openflare-agent
// @Produce json
// @Security AgentTokenAuth
// @Success 200 {object} response.Any{data=agent.ConfigResponse} "活跃配置"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/config-versions/active [get]
func GetActiveConfigHandler(c *gin.Context) {
if _, ok := AgentNodeFromContext(c); !ok {
response.AbortUnauthorized(c, errNodeMissingFromContext)
return
}
config, err := GetActiveConfig(c.Request.Context())
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(config))
}
// SyncWAFIPGroupsHandler syncs WAF IP groups for an agent.
// @Summary 同步 WAF IP 组
// @Description 按 ID 与校验和增量同步 WAF IP 组定义
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.WAFIPGroupSyncInput true "同步请求"
// @Success 200 {object} response.Any{data=agent.WAFIPGroupSyncResult} "同步结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/waf/ip-groups/sync [post]
func SyncWAFIPGroupsHandler(c *gin.Context) {
var input WAFIPGroupSyncInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := SyncWAFIPGroups(c.Request.Context(), input)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// ReportApplyLogHandler records an agent apply log entry.
// @Summary 上报配置应用日志
// @Description 记录 Agent 配置下发与应用结果
// @Tags openflare-agent
// @Accept json
// @Produce json
// @Security AgentTokenAuth
// @Param body body agent.ApplyLogPayload true "应用日志"
// @Success 200 {object} response.Any{data=model.OpenFlareApplyLog} "日志记录"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/apply-logs [post]
func ReportApplyLogHandler(c *gin.Context) {
var payload ApplyLogPayload
if !apiutil.BindJSON(c, &payload) {
return
}
if authNode, ok := AgentNodeFromContext(c); ok {
payload.NodeID = authNode.NodeID
}
log, err := ReportApplyLog(c.Request.Context(), payload)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(log))
}
// DownloadPagesPackageHandler streams the Pages deployment artifact to an authenticated agent.
// @Summary 下载 Pages 部署包
// @Description 流式下载指定部署的静态资源压缩包,供 Agent 边缘分发
// @Tags openflare-agent
// @Produce application/octet-stream
// @Security AgentTokenAuth
// @Param deployment_id path int true "部署 ID"
// @Success 200 {file} binary "部署包文件"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "Token 无效"
// @Router /api/v1/agent/pages/deployments/{deployment_id}/package [get]
func DownloadPagesPackageHandler(c *gin.Context) {
deploymentID, ok := pagesDeploymentIDParam(c)
if !ok {
return
}
packageObj, fileName, err := pages.OpenDeploymentPackage(c.Request.Context(), deploymentID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
defer packageObj.Body.Close()
c.Header("Content-Disposition", "attachment; filename="+fileName)
if packageObj.ContentType != "" {
c.Header("Content-Type", packageObj.ContentType)
}
c.DataFromReader(http.StatusOK, packageObj.ContentLength, packageObj.ContentType, packageObj.Body, nil)
}
func pagesDeploymentIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("deployment_id")
if raw == "" {
response.AbortBadRequest(c, "无效的 ID")
return 0, false
}
id64, err := strconv.ParseUint(raw, 10, 64)
if err != nil || id64 == 0 {
response.AbortBadRequest(c, "无效的 ID")
return 0, false
}
return uint(id64), true
}
// AgentWebSocketHandler 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)
if !ok {
response.AbortUnauthorized(c, errInvalidAgentToken)
return
}
websocket.ServeAgent(c, authNode.NodeID, HandleWSStatus)
}
@@ -0,0 +1,119 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"time"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
nodeStatusOnline = "online"
applyResultOK = "success"
applyResultWarn = "warning"
applyResultFailed = "failed"
)
// NodePayload is the agent register/heartbeat payload.
type NodePayload struct {
NodeID string `json:"node_id"`
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
ExtVersion string `json:"ext_version"`
CurrentVersion string `json:"current_version"`
LastError string `json:"last_error"`
OpenrestyStatus string `json:"openresty_status"`
OpenrestyMessage string `json:"openresty_message"`
Profile *NodeSystemProfile `json:"profile,omitempty"`
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
OpenrestyObservation *NodeOpenrestyObservation `json:"openresty_observation,omitempty"`
TrafficReport *NodeTrafficReport `json:"traffic_report,omitempty"`
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
BufferedObservability []BufferedObservabilityRecord `json:"buffered_observability,omitempty"`
HealthEvents []NodeHealthEvent `json:"health_events"`
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
}
// ApplyLogPayload is the agent apply log report payload.
type ApplyLogPayload struct {
NodeID string `json:"node_id"`
Version string `json:"version"`
Result string `json:"result"`
Message string `json:"message"`
Checksum string `json:"checksum"`
MainConfigChecksum string `json:"main_config_checksum"`
RouteConfigChecksum string `json:"route_config_checksum"`
SupportFileCount int `json:"support_file_count"`
}
// RegistrationResponse is returned after agent registration.
type RegistrationResponse struct {
NodeID string `json:"node_id"`
AccessToken string `json:"access_token"`
Name string `json:"name"`
}
// Settings carries remote agent control flags.
type Settings struct {
HeartbeatInterval int `json:"heartbeat_interval"`
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
AutoUpdate bool `json:"auto_update"`
UpdateRepo string `json:"update_repo"`
UpdateNow bool `json:"update_now"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
RestartOpenrestyNow bool `json:"restart_openresty_now"`
}
// ActiveConfigMeta summarizes the active configuration version.
type ActiveConfigMeta struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
// SupportFile is a configuration support artifact shipped to agents.
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
// ConfigResponse is the full active config payload for agents.
type ConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
SourceConfigJSON string `json:"source_config_json"`
SupportFiles []SupportFile `json:"support_files"`
CreatedAt time.Time `json:"created_at"`
}
// WAFIPGroup is a WAF IP group snapshot for agents.
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
Checksum string `json:"checksum"`
}
// WAFIPGroupSyncInput requests changed WAF IP groups.
type WAFIPGroupSyncInput struct {
IDs []uint `json:"ids"`
Checksums map[string]string `json:"checksums"`
}
// WAFIPGroupSyncResult returns synced WAF IP groups.
type WAFIPGroupSyncResult struct {
Groups []WAFIPGroup `json:"groups"`
}
// HeartbeatResponse is the heartbeat handler result.
type HeartbeatResponse struct {
Node *model.OpenFlareNode `json:"node"`
AgentSettings *Settings `json:"agent_settings"`
ActiveConfig *ActiveConfigMeta `json:"active_config"`
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
}
@@ -0,0 +1,205 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"sort"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
)
type snapshotWAFRuleGroupRef struct {
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
}
type snapshotWAFSection struct {
RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"`
}
type activeConfigSnapshot struct {
WAF snapshotWAFSection `json:"waf"`
}
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
return buildAgentWAFIPGroups(ctx, ids)
}
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids)
if len(targetIDs) == 0 {
activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
if err != nil {
return nil, err
}
targetIDs = activeIDs
}
if len(targetIDs) == 0 {
return []WAFIPGroup{}, nil
}
groups, err := buildAgentWAFIPGroups(ctx, targetIDs)
if err != nil {
return nil, err
}
changed := make([]WAFIPGroup, 0, len(groups))
for _, group := range groups {
if strings.TrimSpace(checksums[fmt.Sprintf("%d", group.ID)]) == group.Checksum {
continue
}
changed = append(changed, group)
}
return changed, nil
}
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
return []WAFIPGroup{}, nil
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
groups, err := model.ListOpenFlareWAFIPGroupsByIDs(ctx, ids)
if err != nil {
return nil, err
}
groupByID := make(map[uint]*model.OpenFlareWAFIPGroup, len(groups))
for _, group := range groups {
groupByID[group.ID] = group
}
result := make([]WAFIPGroup, 0, len(ids))
for _, id := range ids {
group := groupByID[id]
if group == nil {
continue
}
agentGroup, err := buildAgentWAFIPGroup(group)
if err != nil {
return nil, err
}
result = append(result, agentGroup)
}
return result, nil
}
func buildAgentWAFIPGroup(group *model.OpenFlareWAFIPGroup) (WAFIPGroup, error) {
if group == nil {
return WAFIPGroup{}, errors.New("IP 组不存在")
}
ips, err := decodeWAFIPGroupStringList(group.IPList)
if err != nil {
return WAFIPGroup{}, err
}
if !group.Enabled {
ips = []string{}
}
agentGroup := WAFIPGroup{
ID: group.ID,
Name: group.Name,
Type: group.Type,
Enabled: group.Enabled,
IPList: ips,
}
agentGroup.Checksum = checksumAgentWAFIPGroup(agentGroup)
return agentGroup, nil
}
func checksumAgentWAFIPGroup(group WAFIPGroup) string {
payload := struct {
ID uint `json:"id"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list"`
}{
ID: group.ID,
Enabled: group.Enabled,
IPList: append([]string{}, group.IPList...),
}
sort.Strings(payload.IPList)
data, _ := json.Marshal(payload)
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
version, err := loadActiveConfigVersion(ctx)
if err != nil {
if isActiveConfigNotFound(err) {
return []uint{}, nil
}
return nil, err
}
snapshot, err := parseActiveConfigSnapshot(version.SnapshotJSON)
if err != nil {
return nil, err
}
idSet := make(map[uint]struct{})
for _, group := range snapshot.WAF.RuleGroups {
for _, id := range group.IPWhitelistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
for _, id := range group.IPBlacklistGroups {
if id > 0 {
idSet[id] = struct{}{}
}
}
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids, nil
}
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
text := strings.TrimSpace(snapshotJSON)
if text == "" {
return &activeConfigSnapshot{}, nil
}
var snapshot activeConfigSnapshot
if err := json.Unmarshal([]byte(text), &snapshot); err != nil {
return nil, err
}
if snapshot.WAF.RuleGroups == nil {
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
}
return &snapshot, nil
}
func decodeWAFIPGroupStringList(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
return []string{}, nil
}
var items []string
if err := json.Unmarshal([]byte(text), &items); err != nil {
return nil, err
}
return items, nil
}
func uniqueUintIDs(ids []uint) []uint {
normalized := make([]uint, 0, len(ids))
seen := make(map[uint]struct{}, len(ids))
for _, id := range ids {
if id == 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
}
return normalized
}
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"strconv"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupWAFIPGroupTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(
&model.OpenFlareWAFIPGroup{},
&configVersionRecord{},
))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
t.Helper()
snapshot := map[string]any{
"routes": []any{},
"waf": map[string]any{
"rule_groups": []map[string]any{
{
"id": 1,
"name": "agent refs",
"enabled": true,
"ip_blacklist_group_ids": []uint{ipGroupID},
},
},
"bindings": []any{},
},
}
snapshotJSON, err := json.Marshal(snapshot)
require.NoError(t, err)
require.NoError(t, db.DB(ctx).Create(&configVersionRecord{
Version: "20260618-001",
SnapshotJSON: string(snapshotJSON),
Checksum: "test-checksum",
IsActive: true,
}).Error)
}
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "agent runtime group",
Type: "manual",
Enabled: true,
IPList: `["203.0.113.44"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.NoError(t, err)
require.Len(t, groups, 1)
assert.Equal(t, ipGroup.ID, groups[0].ID)
assert.Equal(t, "203.0.113.44", groups[0].IPList[0])
assert.NotEmpty(t, groups[0].Checksum)
groupKey := strconv.FormatUint(uint64(ipGroup.ID), 10)
same, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
require.NoError(t, err)
assert.Empty(t, same)
ipGroup.IPList = `["203.0.113.45"]`
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
require.NoError(t, err)
require.Len(t, delta, 1)
assert.Equal(t, ipGroup.ID, delta[0].ID)
assert.Equal(t, "203.0.113.45", delta[0].IPList[0])
assert.NotEqual(t, groups[0].Checksum, delta[0].Checksum)
}
func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "sync group",
Type: "manual",
Enabled: true,
IPList: `["198.51.100.10"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
IDs: []uint{ipGroup.ID},
Checksums: map[string]string{},
})
require.NoError(t, err)
require.Len(t, result.Groups, 1)
assert.Equal(t, ipGroup.ID, result.Groups[0].ID)
assert.Equal(t, "198.51.100.10", result.Groups[0].IPList[0])
result, err = SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
IDs: []uint{ipGroup.ID},
Checksums: map[string]string{
strconv.FormatUint(uint64(ipGroup.ID), 10): result.Groups[0].Checksum,
},
})
require.NoError(t, err)
assert.Empty(t, result.Groups)
}
func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
ctx := context.Background()
ipGroup := &model.OpenFlareWAFIPGroup{
Name: "disabled group",
Type: "manual",
Enabled: true,
IPList: `["203.0.113.10"]`,
}
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
ipGroup.Enabled = false
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
require.NoError(t, err)
require.Len(t, groups, 1)
assert.False(t, groups[0].Enabled)
assert.Empty(t, groups[0].IPList)
assert.NotEmpty(t, groups[0].Checksum)
}
@@ -0,0 +1,57 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package agent
import (
"context"
"encoding/json"
"log/slog"
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/Rain-kl/Wavelet/internal/model"
)
// HandleWSStatus processes an agent websocket status payload (replaces HTTP heartbeat in WS mode).
func HandleWSStatus(ctx context.Context, nodeID, remoteAddr string, rawPayload json.RawMessage) {
var payload NodePayload
if err := json.Unmarshal(rawPayload, &payload); err != nil {
slog.Debug("agent ws status payload decode failed", "node_id", nodeID, "error", err)
return
}
authNode, err := model.GetOpenFlareNodeByNodeID(ctx, nodeID)
if err != nil {
slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err)
return
}
payload.IP = resolveReportedNodeIP(payload.IP, remoteAddr)
response, err := HeartbeatNode(ctx, authNode, payload)
if err != nil {
slog.Debug("agent ws status handling failed", "node_id", nodeID, "error", err)
return
}
settingsSent := false
if response.AgentSettings != nil {
settingsSent = ofws.SendAgentSettings(nodeID, response.AgentSettings)
}
activeConfigSent := false
if response.ActiveConfig != nil {
activeConfigSent = ofws.SendAgentActiveConfig(nodeID, response.ActiveConfig)
}
wafIPGroupsSent := false
if len(response.WAFIPGroups) > 0 {
wafIPGroupsSent = ofws.SendAgentWAFIPGroups(nodeID, response.WAFIPGroups)
}
slog.Debug("agent ws status processed",
"node_id", nodeID,
"current_version", payload.CurrentVersion,
"openresty_status", payload.OpenrestyStatus,
"settings_sent", settingsSent,
"active_config_sent", activeConfigSent,
"waf_ip_groups_sent", wafIPGroupsSent,
)
}
@@ -0,0 +1,39 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package apiutil provides HTTP helpers for OpenFlare v1 custom API handlers.
package apiutil
import (
"strconv"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/gin-gonic/gin"
)
const errInvalidParams = "参数错误"
const errInvalidID = "无效的 ID"
// BindJSON binds JSON body; returns false after aborting with 400.
func BindJSON(c *gin.Context, dst any) bool {
if err := c.ShouldBindJSON(dst); err != nil {
response.AbortBadRequest(c, errInvalidParams)
return false
}
return true
}
// IDParam parses :id from the URL path.
func IDParam(c *gin.Context) (uint, bool) {
raw := c.Param("id")
if raw == "" {
response.AbortBadRequest(c, errInvalidID)
return 0, false
}
id64, err := strconv.ParseUint(raw, 10, 64)
if err != nil || id64 == 0 {
response.AbortBadRequest(c, errInvalidID)
return 0, false
}
return uint(id64), true
}
@@ -0,0 +1,34 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apiutil
import (
"errors"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// AbortNotFoundIfMissing maps gorm.ErrRecordNotFound to 404; other errors to 400.
func AbortNotFoundIfMissing(c *gin.Context, err error, notFoundMsg string) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, notFoundMsg)
return true
}
response.AbortBadRequest(c, err.Error())
return true
}
// AbortBadRequestOnError writes a 400 for any non-nil error.
func AbortBadRequestOnError(c *gin.Context, err error) bool {
if err == nil {
return false
}
response.AbortBadRequest(c, err.Error())
return true
}
@@ -0,0 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apiutil
import (
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/gin-gonic/gin"
)
// AdminMiddlewares returns Wavelet-standard middlewares for OpenFlare console routes.
// OpenFlare no longer distinguishes Admin vs Root tiers; all management endpoints share
// the same gate: user.IsAdmin for session users, token_admin for Access Token callers.
func AdminMiddlewares() []gin.HandlerFunc {
return []gin.HandlerFunc{oauth.LoginRequired(), admin.LoginAdminRequired()}
}
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apiutil
import (
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupAdminMiddlewareTest(t *testing.T) (*gin.Engine, *gorm.DB, func()) {
t.Helper()
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, dbConn.AutoMigrate(&model.User{}, &model.AccessToken{}))
db.SetDB(dbConn)
sessionCookieName := "test_admin_middleware_session"
if config.Config.App.SessionCookieName != "" {
sessionCookieName = config.Config.App.SessionCookieName
}
store := cookie.NewStore([]byte("test_admin_middleware_session_secret"))
store.Options(oauth.GetSessionOptions(3600))
engine := testhelper.NewTestGinEngine(sessions.Sessions(sessionCookieName, store))
protected := engine.Group("/protected", AdminMiddlewares()...)
protected.GET("", func(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(gin.H{"ok": true}))
})
cleanup := func() {
db.SetDB(nil)
}
return engine, dbConn, cleanup
}
func seedUser(t *testing.T, dbConn *gorm.DB, username string, isAdmin bool) *model.User {
t.Helper()
user := &model.User{
ID: idgen.NextUint64ID(),
Username: username,
Nickname: username,
Email: username + "@openflare.test",
IsActive: true,
IsAdmin: isAdmin,
}
require.NoError(t, dbConn.Create(user).Error)
return user
}
func seedAccessToken(t *testing.T, dbConn *gorm.DB, user *model.User, isAdmin bool) string {
t.Helper()
token, err := model.GenerateTokenString()
require.NoError(t, err)
require.NoError(t, dbConn.Create(&model.AccessToken{
UserID: user.ID,
Name: user.Username + "-token",
TokenHash: model.HashToken(token),
MaskedToken: model.MaskTokenString(token),
IsAdmin: isAdmin,
}).Error)
return token
}
func decodeResponse(t *testing.T, rec *httptest.ResponseRecorder) response.Any {
t.Helper()
var resp response.Any
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
return resp
}
func TestAdminRequiredUnauthenticated(t *testing.T) {
engine, _, cleanup := setupAdminMiddlewareTest(t)
defer cleanup()
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusUnauthorized, rec.Code)
resp := decodeResponse(t, rec)
assert.NotEmpty(t, resp.ErrorMsg)
}
func TestAdminRequiredNonAdminToken(t *testing.T) {
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
defer cleanup()
user := seedUser(t, dbConn, "regular", false)
token := seedAccessToken(t, dbConn, user, false)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set("X-Access-Token", token)
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusNotFound, rec.Code)
resp := decodeResponse(t, rec)
assert.Equal(t, admin.TokenAdminRequired, resp.ErrorMsg)
}
func TestAdminRequiredAdminWithoutTokenAdmin(t *testing.T) {
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
defer cleanup()
user := seedUser(t, dbConn, "admin-no-token-admin", true)
token := seedAccessToken(t, dbConn, user, false)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set("X-Access-Token", token)
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusNotFound, rec.Code)
resp := decodeResponse(t, rec)
assert.Equal(t, admin.TokenAdminRequired, resp.ErrorMsg)
}
func TestAdminRequiredAdminWithTokenAdmin(t *testing.T) {
engine, dbConn, cleanup := setupAdminMiddlewareTest(t)
defer cleanup()
user := seedUser(t, dbConn, "admin", true)
token := seedAccessToken(t, dbConn, user, true)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/protected", nil)
req.Header.Set("X-Access-Token", token)
engine.ServeHTTP(rec, req)
assert.Equal(t, http.StatusOK, rec.Code)
resp := decodeResponse(t, rec)
assert.Empty(t, resp.ErrorMsg)
}
@@ -0,0 +1,19 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apiutil
import (
"strings"
"github.com/gin-gonic/gin"
)
// RegisterCollection registers a collection endpoint on both "" and "/" so requests
// work with or without a trailing slash.
func RegisterCollection(route *gin.RouterGroup, method string, handlers ...gin.HandlerFunc) {
route.Handle(method, "/", handlers...)
if !strings.HasSuffix(route.BasePath(), "/") {
route.Handle(method, "", handlers...)
}
}
@@ -0,0 +1,8 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apply_log
const (
errRetentionDaysOutOfRange = "retention_days 必须在 1 到 3650 之间"
)
@@ -0,0 +1,128 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apply_log
import (
"context"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
defaultApplyLogPageSize = 20
maxApplyLogPageSize = 200
maxApplyLogRetentionDays = 3650
)
// ListQuery filters apply logs for paginated listing.
type ListQuery struct {
NodeID string `json:"node_id"`
PageNo int `json:"pageNo"`
PageSize int `json:"pageSize"`
}
// ListResult is the paginated apply log list response.
type ListResult struct {
Rows []*model.OpenFlareApplyLog `json:"rows"`
Current int `json:"current"`
Total int `json:"total"`
TotalPage int `json:"totalPage"`
}
// CleanupInput controls apply log cleanup behavior.
type CleanupInput struct {
DeleteAll bool `json:"delete_all"`
RetentionDays int `json:"retention_days"`
}
// CleanupResult reports apply log cleanup outcome.
type CleanupResult struct {
DeleteAll bool `json:"delete_all"`
RetentionDays int `json:"retention_days"`
DeletedCount int64 `json:"deleted_count"`
Cutoff *time.Time `json:"cutoff,omitempty"`
}
// ListPage returns paginated apply logs with optional node_id filter.
func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) {
pageNo := normalizePageNo(input.PageNo)
pageSize := normalizePageSize(input.PageSize)
nodeID := strings.TrimSpace(input.NodeID)
rows, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: nodeID,
PageNo: pageNo,
PageSize: pageSize,
})
if err != nil {
return nil, err
}
total, err := model.CountOpenFlareApplyLogs(ctx, nodeID)
if err != nil {
return nil, err
}
totalPage := 0
if total > 0 {
totalPage = int((total + int64(pageSize) - 1) / int64(pageSize))
}
return &ListResult{
Rows: rows,
Current: pageNo,
Total: int(total),
TotalPage: totalPage,
}, nil
}
// Cleanup removes old apply logs or deletes all records.
func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) {
if input.DeleteAll {
deleted, err := model.DeleteAllOpenFlareApplyLogs(ctx)
if err != nil {
return nil, err
}
return &CleanupResult{
DeleteAll: true,
DeletedCount: deleted,
}, nil
}
if input.RetentionDays <= 0 || input.RetentionDays > maxApplyLogRetentionDays {
return nil, errors.New(errRetentionDaysOutOfRange)
}
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour)
deleted, err := model.DeleteOpenFlareApplyLogsBefore(ctx, cutoff)
if err != nil {
return nil, err
}
return &CleanupResult{
RetentionDays: input.RetentionDays,
DeletedCount: deleted,
Cutoff: &cutoff,
}, nil
}
func normalizePageNo(pageNo int) int {
if pageNo <= 0 {
return 1
}
return pageNo
}
func normalizePageSize(pageSize int) int {
if pageSize <= 0 {
return defaultApplyLogPageSize
}
if pageSize > maxApplyLogPageSize {
return maxApplyLogPageSize
}
return pageSize
}
@@ -0,0 +1,105 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apply_log
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupApplyLogTestDB(t *testing.T) func() {
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
err = sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{})
require.NoError(t, err)
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func TestListPageAndCleanup(t *testing.T) {
cleanup := setupApplyLogTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
logs := []model.OpenFlareApplyLog{
{NodeID: "node-logs", Version: "v1", Result: "success", Message: "1", CreatedAt: now.Add(-10 * 24 * time.Hour)},
{NodeID: "node-logs", Version: "v2", Result: "success", Message: "2", CreatedAt: now.Add(-5 * 24 * time.Hour)},
{NodeID: "node-logs", Version: "v3", Result: "success", Message: "3", CreatedAt: now},
}
for i := range logs {
require.NoError(t, db.DB(ctx).Create(&logs[i]).Error)
}
pageResult, err := ListPage(ctx, ListQuery{
NodeID: "node-logs",
PageNo: 1,
PageSize: 2,
})
require.NoError(t, err)
assert.Equal(t, 3, pageResult.Total)
assert.Len(t, pageResult.Rows, 2)
assert.Equal(t, 2, pageResult.TotalPage)
assert.Equal(t, 1, pageResult.Current)
cleanupResult, err := Cleanup(ctx, CleanupInput{
DeleteAll: false,
RetentionDays: 7,
})
require.NoError(t, err)
assert.Equal(t, int64(1), cleanupResult.DeletedCount)
assert.NotNil(t, cleanupResult.Cutoff)
remaining, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: "node-logs",
PageNo: 1,
PageSize: 10,
})
require.NoError(t, err)
assert.Len(t, remaining, 2)
cleanupAll, err := Cleanup(ctx, CleanupInput{DeleteAll: true})
require.NoError(t, err)
assert.Equal(t, int64(2), cleanupAll.DeletedCount)
assert.True(t, cleanupAll.DeleteAll)
finalLogs, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: "node-logs",
PageNo: 1,
PageSize: 10,
})
require.NoError(t, err)
assert.Empty(t, finalLogs)
}
func TestCleanupInvalidRetentionDays(t *testing.T) {
cleanup := setupApplyLogTestDB(t)
defer cleanup()
ctx := context.Background()
_, err := Cleanup(ctx, CleanupInput{RetentionDays: 0})
require.Error(t, err)
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
_, err = Cleanup(ctx, CleanupInput{RetentionDays: 4000})
require.Error(t, err)
assert.Equal(t, errRetentionDaysOutOfRange, err.Error())
}
@@ -0,0 +1,77 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package apply_log
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/gin-gonic/gin"
)
// GetApplyLogs lists apply logs with pagination and optional node_id filter.
// @Summary 获取配置下发日志
// @Description 分页返回节点配置下发记录,支持按节点 ID 筛选,需要管理员权限
// @Tags openflare-apply-log
// @Produce json
// @Security SessionCookie
// @Param node_id query string false "节点 ID 筛选"
// @Param pageNo query int false "页码"
// @Param page_no query int false "页码(别名)"
// @Param pageSize query int false "每页数量"
// @Param page_size query int false "每页数量(别名)"
// @Success 200 {object} response.Any{data=apply_log.ListResult} "下发日志列表"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/apply-logs [get]
func GetApplyLogs(c *gin.Context) {
result, err := ListPage(c.Request.Context(), ListQuery{
NodeID: c.Query("node_id"),
PageNo: readIntQuery(c, "pageNo", "page_no"),
PageSize: readIntQuery(c, "pageSize", "page_size"),
})
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// CleanupApplyLogs removes old apply logs or deletes all records.
// @Summary 清理配置下发日志
// @Description 按保留天数清理历史下发记录,或删除全部记录,需要管理员权限
// @Tags openflare-apply-log
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param body body apply_log.CleanupInput true "清理参数"
// @Success 200 {object} response.Any{data=apply_log.CleanupResult} "清理结果"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "无权限或不存在"
// @Router /api/v1/d/apply-logs/cleanup [post]
func CleanupApplyLogs(c *gin.Context) {
var input CleanupInput
if !apiutil.BindJSON(c, &input) {
return
}
result, err := Cleanup(c.Request.Context(), input)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
func readIntQuery(c *gin.Context, primary, secondary string) int {
value := c.Query(primary)
if value == "" {
value = c.Query(secondary)
}
parsed, _ := strconv.Atoi(value)
return parsed
}
@@ -0,0 +1,200 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package openflare
import (
"context"
"fmt"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tasks"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/uptimekuma"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
)
const (
// SSLRenewTask renews due ACME TLS certificates.
SSLRenewTask = "openflare:ssl_renew"
// TaskTypeSSLRenew is the admin task type for SSL renewal.
TaskTypeSSLRenew = "of_ssl_renew"
// DatabaseAutoCleanupTask prunes observability tables by retention policy.
DatabaseAutoCleanupTask = "openflare:database_auto_cleanup"
// TaskTypeDatabaseAutoCleanup is the admin task type for observability cleanup.
TaskTypeDatabaseAutoCleanup = "of_database_auto_cleanup"
// WAFIPGroupSyncTask syncs due automatic/subscription WAF IP groups.
WAFIPGroupSyncTask = "openflare:waf_ip_group_sync"
// TaskTypeWAFIPGroupSync is the admin task type for WAF IP group sync.
TaskTypeWAFIPGroupSync = "of_waf_ip_group_sync"
// UptimeKumaSyncTask synchronizes proxy routes to Uptime Kuma monitors.
UptimeKumaSyncTask = "openflare:uptime_kuma_sync"
// TaskTypeUptimeKumaSync is the admin task type for Uptime Kuma sync.
TaskTypeUptimeKumaSync = "of_uptime_kuma_sync"
)
var (
lastUptimeKumaSyncTime time.Time
uptimeKumaSyncMutex sync.Mutex
)
// SSLRenewMeta describes the SSL renewal task.
var SSLRenewMeta = task.TaskMeta{
Type: TaskTypeSSLRenew,
AsynqTask: SSLRenewTask,
Name: "OpenFlare SSL 自动续期",
Description: "扫描即将到期的 ACME 证书并触发自动续期",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// DatabaseAutoCleanupMeta describes the observability auto-cleanup task.
var DatabaseAutoCleanupMeta = task.TaskMeta{
Type: TaskTypeDatabaseAutoCleanup,
AsynqTask: DatabaseAutoCleanupTask,
Name: "OpenFlare 可观测数据自动清理",
Description: "按保留天数清理访问日志、性能快照与请求聚合数据",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// WAFIPGroupSyncMeta describes the WAF IP group sync task.
var WAFIPGroupSyncMeta = task.TaskMeta{
Type: TaskTypeWAFIPGroupSync,
AsynqTask: WAFIPGroupSyncTask,
Name: "OpenFlare WAF IP 组同步",
Description: "同步到期的自动规则与订阅类型 WAF IP 组",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// UptimeKumaSyncMeta describes the Uptime Kuma sync task.
var UptimeKumaSyncMeta = task.TaskMeta{
Type: TaskTypeUptimeKumaSync,
AsynqTask: UptimeKumaSyncTask,
Name: "OpenFlare Uptime Kuma 同步",
Description: "将启用的代理规则同步到 Uptime Kuma 监控",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// SSLRenewHandler renews due TLS certificates.
type SSLRenewHandler struct{}
// Execute runs SSL certificate renewal for all due certificates.
func (h *SSLRenewHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
task.AppendLog(ctx, "开始扫描待续期证书")
if err := tasks.RunSSLRenewJob(ctx); err != nil {
task.AppendLog(ctx, "SSL 自动续期失败: %v", err)
return nil, err
}
msg := "SSL 自动续期任务完成"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
// DatabaseAutoCleanupHandler prunes observability data when auto-cleanup is enabled.
type DatabaseAutoCleanupHandler struct{}
// Execute runs retention-based cleanup for all observability targets.
func (h *DatabaseAutoCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
if !model.DatabaseAutoCleanupEnabled {
msg := "自动清理未启用,跳过执行"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
task.AppendLog(ctx, "开始执行可观测数据自动清理,保留天数=%d", model.DatabaseAutoCleanupRetentionDays)
summary, err := tasks.RunDatabaseAutoCleanupOnce(time.Now())
if err != nil {
task.AppendLog(ctx, "可观测数据自动清理失败: %v", err)
return nil, err
}
if summary == nil {
msg := "自动清理未启用,跳过执行"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
var totalDeleted int64
for _, item := range summary.Results {
totalDeleted += item.DeletedCount
task.AppendLog(ctx, "清理 %s:删除 %d 条", item.TargetLabel, item.DeletedCount)
}
msg := fmt.Sprintf(
"可观测数据自动清理完成,保留 %d 天,共删除 %d 条",
summary.RetentionDays,
totalDeleted,
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
// WAFIPGroupSyncHandler syncs due WAF IP groups to agents.
type WAFIPGroupSyncHandler struct{}
// Execute syncs all due automatic/subscription WAF IP groups.
func (h *WAFIPGroupSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
task.AppendLog(ctx, "开始同步到期的 WAF IP 组")
if err := waf.SyncDueWAFIPGroups(ctx); err != nil {
task.AppendLog(ctx, "WAF IP 组同步失败: %v", err)
return nil, err
}
msg := "WAF IP 组同步完成"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
// UptimeKumaSyncHandler synchronizes proxy routes to Uptime Kuma.
type UptimeKumaSyncHandler struct{}
// Execute runs Uptime Kuma sync when integration is enabled and the interval has elapsed.
func (h *UptimeKumaSyncHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
if !model.UptimeKumaEnabled {
msg := "Uptime Kuma 集成未启用,跳过执行"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
interval := model.UptimeKumaSyncInterval
if interval <= 0 {
interval = 5
}
if time.Since(lastUptimeKumaSyncTime) < time.Duration(interval)*time.Minute {
msg := fmt.Sprintf("距上次同步不足 %d 分钟,跳过执行", interval)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
if !uptimeKumaSyncMutex.TryLock() {
msg := "Uptime Kuma 同步任务正在执行,跳过本次调度"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
defer uptimeKumaSyncMutex.Unlock()
task.AppendLog(ctx, "开始同步代理规则到 Uptime Kuma")
if err := uptimekuma.SyncToUptimeKuma(ctx); err != nil {
task.AppendLog(ctx, "Uptime Kuma 同步失败: %v", err)
return nil, err
}
lastUptimeKumaSyncTime = time.Now()
msg := "Uptime Kuma 同步完成"
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}

Some files were not shown because too many files have changed in this diff Show More