wavelet init

This commit is contained in:
ryan
2026-06-18 15:24:48 +08:00
parent d6a7011885
commit 99738bbc17
714 changed files with 139987 additions and 0 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/wavelet.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))
}
+15
View File
@@ -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)
}
+343
View File
@@ -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
}
+413
View File
@@ -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)
})
}
+292
View File
@@ -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 ""
}
+121
View File
@@ -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/wavelet.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/wavelet.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="wavelet.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, "wavelet.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("wavelet_%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 := "Wavelet 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 Wavelet.</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 != "Wavelet" {
t.Errorf("expected 'Wavelet', 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)
}
})
}
+23
View File
@@ -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 = "删除定时任务失败"
)
+416
View File
@@ -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,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", "Wavelet-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", "Wavelet-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/Wavelet", want: "Rain-kl/Wavelet"},
{name: "GitHub URL", input: "https://github.com/Rain-kl/Wavelet.git", want: "Rain-kl/Wavelet"},
{name: "unsupported host", input: "https://example.com/Rain-kl/Wavelet", wantErr: true},
{name: "missing owner", input: "Wavelet", 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/Wavelet", 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)
}
}()
}
+21
View File
@@ -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 = "邮箱已被注册"
)
+106
View File
@@ -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
}
+288
View File
@@ -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())
}
})
}