mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 23:06:36 +08:00
wavelet init
This commit is contained in:
@@ -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
@@ -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
@@ -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))
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package admin 提供管理后台功能
|
||||
package admin
|
||||
|
||||
// 管理后台错误消息常量
|
||||
const (
|
||||
AdminRequired = "未经授权访问"
|
||||
TokenAdminRequired = "该访问令牌没有管理员权限,无法访问管理端点" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
InvalidAuthSourceID = "认证源 ID 无效"
|
||||
InvalidCursorParam = "无效的 cursor 参数"
|
||||
InvalidTaskExecutionID = "无效的任务执行记录 ID"
|
||||
)
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// LoginAdminRequired 返回管理员权限校验中间件
|
||||
func LoginAdminRequired() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// init trace
|
||||
ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired")
|
||||
defer span.End()
|
||||
|
||||
user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
|
||||
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
|
||||
if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
|
||||
tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey)
|
||||
if !tokenAdmin {
|
||||
response.AbortNotFound(c, TokenAdminRequired)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if !user.IsAdmin {
|
||||
response.AbortNotFound(c, AdminRequired)
|
||||
return
|
||||
}
|
||||
|
||||
// log
|
||||
logger.InfoF(ctx, "[LoginAdminRequired] %d %s", user.ID, user.Username)
|
||||
|
||||
// next
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,278 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// ListChannelDefinitions 获取各种消息通道的表单配置定义列表
|
||||
// @Summary 获取所有消息通道配置字段定义
|
||||
// @Description 返回系统支持的所有消息通道类型(如飞书、邮件、自定义、Telegram)的动态表单定义,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]Definition} "通道配置定义列表"
|
||||
// @Router /api/v1/admin/push/channels/definitions [get]
|
||||
func ListChannelDefinitions(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(ListDefinitions()))
|
||||
}
|
||||
|
||||
// ListChannels 获取消息通道列表
|
||||
// @Summary 获取所有消息通道
|
||||
// @Description 返回系统配置的所有消息通道列表,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表"
|
||||
// @Router /api/v1/admin/push/channels [get]
|
||||
func ListChannels(c *gin.Context) {
|
||||
channels, err := listPushChannels(c.Request.Context())
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(channels))
|
||||
}
|
||||
|
||||
// CreateChannelRequest 创建通道参数
|
||||
type CreateChannelRequest struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Description string `json:"description"`
|
||||
Type string `json:"type" binding:"required"`
|
||||
Token string `json:"token"`
|
||||
URL string `json:"url"`
|
||||
Other string `json:"other"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// CreateChannel 创建消息通道
|
||||
// @Summary 创建消息通道
|
||||
// @Description 新建一个消息通道配置,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body CreateChannelRequest true "创建参数"
|
||||
// @Success 200 {object} response.Any{data=model.PushChannel} "创建成功"
|
||||
// @Router /api/v1/admin/push/channels [post]
|
||||
func CreateChannel(c *gin.Context) {
|
||||
var req CreateChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
channel, err := createPushChannel(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(channel))
|
||||
}
|
||||
|
||||
// UpdateChannelRequest 修改通道参数
|
||||
type UpdateChannelRequest struct {
|
||||
Description string `json:"description"`
|
||||
Type string `json:"type" binding:"required"`
|
||||
Token string `json:"token"`
|
||||
URL string `json:"url"`
|
||||
Other string `json:"other"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// UpdateChannel 更新消息通道
|
||||
// @Summary 更新消息通道
|
||||
// @Description 修改消息通道配置,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "通道ID"
|
||||
// @Param request body UpdateChannelRequest true "更新参数"
|
||||
// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功"
|
||||
// @Router /api/v1/admin/push/channels/{id} [put]
|
||||
func UpdateChannel(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
return
|
||||
}
|
||||
|
||||
var req UpdateChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
channel, err := updatePushChannel(c.Request.Context(), id, req)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "channel not found")
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(channel))
|
||||
}
|
||||
|
||||
// DeleteChannel 删除消息通道
|
||||
// @Summary 删除消息通道
|
||||
// @Description 根据ID删除消息通道,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path uint64 true "通道ID"
|
||||
// @Success 200 {object} response.Any "删除成功"
|
||||
// @Router /api/v1/admin/push/channels/{id} [delete]
|
||||
func DeleteChannel(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid channel id")
|
||||
return
|
||||
}
|
||||
|
||||
if err := deletePushChannel(c.Request.Context(), id); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "channel not found")
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// TestChannelRequest 测试通道连通性参数
|
||||
type TestChannelRequest struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Token string `json:"token"`
|
||||
URL string `json:"url"`
|
||||
Other string `json:"other"`
|
||||
Target string `json:"target"`
|
||||
}
|
||||
|
||||
// TestChannel 测试通道连通性
|
||||
// @Summary 测试通道连通性
|
||||
// @Description 触发一次临时的或现有的通道连通性推送测试,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body TestChannelRequest true "测试参数"
|
||||
// @Success 200 {object} response.Any "测试触发成功"
|
||||
// @Router /api/v1/admin/push/channels/test [post]
|
||||
func TestChannel(c *gin.Context) {
|
||||
var req TestChannelRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
url, token, other, channelType, err := loadChannelForTest(ctx, req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if channelType == channelEmail {
|
||||
url, token, other = resolveSMTPConfig(ctx, url, token, other)
|
||||
}
|
||||
|
||||
tempChannel := model.PushChannel{
|
||||
Name: "test_temp",
|
||||
URL: url,
|
||||
Token: token,
|
||||
Other: other,
|
||||
Type: channelType,
|
||||
Enabled: true,
|
||||
}
|
||||
if err := tempChannel.Validate(); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
url = tempChannel.URL
|
||||
|
||||
var config pkgpush.Config
|
||||
var renderedJSON string
|
||||
switch channelType {
|
||||
case channelLark:
|
||||
config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token}
|
||||
renderedJSON = other
|
||||
case channelEmail:
|
||||
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
|
||||
case channelTelegram:
|
||||
config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other}
|
||||
default:
|
||||
config = pkgpush.Config{Channel: channelCustom, URL: url}
|
||||
customPushReq := CustomPushRequest{
|
||||
Title: "通道测试通知",
|
||||
Content: "这是一条来自系统的消息通道连通性测试消息。",
|
||||
Description: "系统通道测试",
|
||||
URL: "https://example.com",
|
||||
To: req.Target,
|
||||
}
|
||||
renderedJSON = renderCustomPayload(other, customPushReq)
|
||||
}
|
||||
|
||||
payload := SendPayload{
|
||||
EventKey: "test_channel",
|
||||
Config: config,
|
||||
Target: req.Target,
|
||||
Body: NotificationMessage{
|
||||
Title: "通道测试通知",
|
||||
Content: "这是一条来自系统的消息通道连通性测试消息。",
|
||||
Level: defaultLevelInfo,
|
||||
},
|
||||
Template: renderedJSON,
|
||||
}
|
||||
if err := enqueuePushTask(ctx, payload); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// CustomPushRequest 外部公开推送请求参数
|
||||
type CustomPushRequest struct {
|
||||
Title string `json:"title" form:"title"`
|
||||
Description string `json:"description" form:"description"`
|
||||
Content string `json:"content" form:"content"`
|
||||
URL string `json:"url" form:"url"`
|
||||
To string `json:"to" form:"to"`
|
||||
Token string `json:"token" form:"token"`
|
||||
}
|
||||
|
||||
func escapeJSONString(s string) string {
|
||||
b, _ := json.Marshal(s)
|
||||
const minJSONLen = 2
|
||||
if len(b) >= minJSONLen {
|
||||
return string(b[1 : len(b)-1])
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func renderCustomPayload(template string, req CustomPushRequest) string {
|
||||
result := template
|
||||
result = strings.ReplaceAll(result, "$title", escapeJSONString(req.Title))
|
||||
result = strings.ReplaceAll(result, "$description", escapeJSONString(req.Description))
|
||||
result = strings.ReplaceAll(result, "$content", escapeJSONString(req.Content))
|
||||
result = strings.ReplaceAll(result, "$url", escapeJSONString(req.URL))
|
||||
result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To))
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import "sync"
|
||||
|
||||
const (
|
||||
// KeyURL represents the URL field key
|
||||
KeyURL = "url"
|
||||
// KeyToken represents the Token field key
|
||||
KeyToken = "token"
|
||||
// KeyOther represents the Other field key
|
||||
KeyOther = "other"
|
||||
|
||||
// TypeText represents standard text input type
|
||||
TypeText = "text"
|
||||
// TypePassword represents password input type
|
||||
TypePassword = "password"
|
||||
// TypeTextarea represents textarea input type
|
||||
TypeTextarea = "textarea"
|
||||
)
|
||||
|
||||
// Field represents a form field configuration for a channel.
|
||||
type Field struct {
|
||||
Key string `json:"key"` // unique key for the field (e.g. url, token, other)
|
||||
Label string `json:"label"` // human readable label (e.g. "Webhook 地址")
|
||||
Type string `json:"type"` // input type: "text" | "password" | "textarea"
|
||||
Required bool `json:"required"` // whether this field is required
|
||||
Placeholder string `json:"placeholder"` // input placeholder
|
||||
Description string `json:"description"` // field explanation/help text
|
||||
}
|
||||
|
||||
// Definition represents the metadata and form schema for a notification channel.
|
||||
type Definition struct {
|
||||
Type string `json:"type"` // channel type (e.g., custom, lark, email)
|
||||
Name string `json:"name"` // display name
|
||||
Description string `json:"description"` // short description
|
||||
Fields []Field `json:"fields"` // form fields
|
||||
}
|
||||
|
||||
var (
|
||||
defMu sync.RWMutex
|
||||
definitions = make(map[string]Definition)
|
||||
)
|
||||
|
||||
// RegisterChannelDefinition registers a channel definition.
|
||||
func RegisterChannelDefinition(def Definition) {
|
||||
defMu.Lock()
|
||||
defer defMu.Unlock()
|
||||
definitions[def.Type] = def
|
||||
}
|
||||
|
||||
// ListDefinitions returns all registered channel definitions.
|
||||
func ListDefinitions() []Definition {
|
||||
defMu.RLock()
|
||||
defer defMu.RUnlock()
|
||||
|
||||
// We want a stable order: custom, lark, telegram, email
|
||||
order := []string{channelCustom, channelLark, channelTelegram, channelEmail}
|
||||
res := make([]Definition, 0, len(definitions))
|
||||
for _, t := range order {
|
||||
if d, ok := definitions[t]; ok {
|
||||
res = append(res, d)
|
||||
}
|
||||
}
|
||||
// Add any others
|
||||
for t, d := range definitions {
|
||||
found := false
|
||||
for _, o := range order {
|
||||
if o == t {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
res = append(res, d)
|
||||
}
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func init() {
|
||||
// Register custom webhook channel
|
||||
RegisterChannelDefinition(Definition{
|
||||
Type: channelCustom,
|
||||
Name: "自定义消息通道",
|
||||
Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。",
|
||||
Fields: []Field{
|
||||
{
|
||||
Key: KeyURL,
|
||||
Label: "请求地址",
|
||||
Type: TypeText,
|
||||
Required: true,
|
||||
Placeholder: "在此填写完整的请求地址,必须使用 HTTPS 协议",
|
||||
Description: "接口请求的完整 HTTPS URL,例如 https://api.example.com/webhook",
|
||||
},
|
||||
{
|
||||
Key: KeyOther,
|
||||
Label: "请求体 (JSON)",
|
||||
Type: TypeTextarea,
|
||||
Required: true,
|
||||
Placeholder: "在此输入请求体,支持模板变量,必须为合法的 JSON 格式",
|
||||
Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
// Register Lark robot channel
|
||||
RegisterChannelDefinition(Definition{
|
||||
Type: channelLark,
|
||||
Name: "飞书群机器人",
|
||||
Description: "配置飞书群自定义机器人的 Webhook 接口投递。",
|
||||
Fields: []Field{
|
||||
{
|
||||
Key: KeyURL,
|
||||
Label: "Webhook 地址",
|
||||
Type: TypeText,
|
||||
Required: true,
|
||||
Placeholder: "https://open.feishu.cn/open-apis/bot/v2/hook/YOUR_TOKEN",
|
||||
Description: "从飞书群机器人设置中复制 of Webhook URL",
|
||||
// Note: using 'of' was in feishu.go, let's keep original wording or fix it
|
||||
},
|
||||
{
|
||||
Key: KeyToken,
|
||||
Label: "签名校验密钥 (Secret) (可选)",
|
||||
Type: TypeText,
|
||||
Required: false,
|
||||
Placeholder: "可选,若机器人启用了安全设置中的签名校验,请在此输入",
|
||||
Description: "飞书群机器人安全设置中的签名校验 Key",
|
||||
},
|
||||
{
|
||||
Key: KeyOther,
|
||||
Label: "自定义卡片 JSON 模版 (可选)",
|
||||
Type: TypeTextarea,
|
||||
Required: false,
|
||||
Placeholder: "可选,留空则默认使用系统内置的精美互动卡片",
|
||||
Description: "若填写,必须是合法的飞书卡片 JSON 格式",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
// Register Telegram channel
|
||||
RegisterChannelDefinition(Definition{
|
||||
Type: channelTelegram,
|
||||
Name: "Telegram 机器人",
|
||||
Description: "配置 Telegram 机器人推送消息。",
|
||||
Fields: []Field{
|
||||
{
|
||||
Key: KeyURL,
|
||||
Label: "API 基础地址 (可选)",
|
||||
Type: TypeText,
|
||||
Required: false,
|
||||
Placeholder: "https://api.telegram.org",
|
||||
Description: "接口请求的 HTTPS 基础地址,留空默认为 https://api.telegram.org",
|
||||
},
|
||||
{
|
||||
Key: KeyToken,
|
||||
Label: "机器人 Token (Bot Token)",
|
||||
Type: TypePassword,
|
||||
Required: true,
|
||||
Placeholder: "在此输入 Telegram 机器人的 Bot Token",
|
||||
Description: "通过 BotFather 申请到的机器人 Access Token",
|
||||
},
|
||||
{
|
||||
Key: KeyOther,
|
||||
Label: "默认会话 ID (Chat ID) (可选)",
|
||||
Type: TypeText,
|
||||
Required: false,
|
||||
Placeholder: "例如 -100123456789 或 @channel_name",
|
||||
Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID",
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
// Register Email channel
|
||||
RegisterChannelDefinition(Definition{
|
||||
Type: channelEmail,
|
||||
Name: "邮件推送通道",
|
||||
Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。",
|
||||
Fields: []Field{},
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
const (
|
||||
channelCustom = "custom"
|
||||
channelEmail = "email"
|
||||
channelLark = "lark"
|
||||
channelTelegram = "telegram"
|
||||
defaultLevelInfo = "INFO"
|
||||
keyTitle = "title"
|
||||
keyContent = "content"
|
||||
keyLevel = "level"
|
||||
)
|
||||
@@ -0,0 +1,38 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package custom_events defines custom push notification events.
|
||||
package custom_events
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
|
||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||
)
|
||||
|
||||
// AdminLogin is the metadata definition for the admin login event.
|
||||
var AdminLogin = push.EventMetadata{
|
||||
Key: "admin_login",
|
||||
Name: "管理员登录",
|
||||
DefaultTemplate: push.NotificationMessage{
|
||||
Title: "管理员登录提醒",
|
||||
Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。",
|
||||
Level: "INFO",
|
||||
},
|
||||
Description: "当管理员成功登录系统时触发此通知",
|
||||
}
|
||||
|
||||
func handleAdminLogin(ctx context.Context, event listener.AdminLoggedIn) {
|
||||
if event.User == nil {
|
||||
return
|
||||
}
|
||||
|
||||
body := map[string]any{
|
||||
"user": event.User,
|
||||
"ip": event.IP,
|
||||
"time": time.Now().Format("2006-01-02 15:04:05"),
|
||||
}
|
||||
push.DefaultTrigger.Trigger(ctx, AdminLogin, body)
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package custom_events
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
|
||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var registerOnce sync.Once
|
||||
|
||||
func ensureRegistered() {
|
||||
registerOnce.Do(Register)
|
||||
}
|
||||
|
||||
func setupAdminLoginIntegrationTest(t *testing.T) (*gorm.DB, func()) {
|
||||
t.Helper()
|
||||
|
||||
dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
|
||||
err := dbConn.AutoMigrate(
|
||||
&model.PushEvent{},
|
||||
&model.PushHistory{},
|
||||
&model.PushChannel{},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
sysUser := &model.User{
|
||||
ID: 999,
|
||||
Username: "system",
|
||||
Nickname: "系统",
|
||||
Password: "*",
|
||||
IsActive: true,
|
||||
}
|
||||
require.NoError(t, dbConn.Create(sysUser).Error)
|
||||
|
||||
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{Addr: mr.Addr()})
|
||||
task.RegisterHandler(push.SendNotificationTask, &push.PushHandler{})
|
||||
task.RegisterTaskMeta(push.SendNotificationMeta)
|
||||
|
||||
ensureRegistered()
|
||||
|
||||
require.NoError(t, push.SyncEvents(context.Background()))
|
||||
|
||||
return dbConn, func() {
|
||||
cleanup()
|
||||
if task.AsynqClient != nil {
|
||||
task.AsynqClient.Close()
|
||||
task.AsynqClient = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func seedMockPushChannel(t *testing.T, dbConn *gorm.DB) *model.PushChannel {
|
||||
t.Helper()
|
||||
|
||||
channel := &model.PushChannel{
|
||||
Name: "mock_channel",
|
||||
Type: "custom",
|
||||
URL: "https://webhook.site/admin-login",
|
||||
Other: `{"text": "$content"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
require.NoError(t, dbConn.Create(channel).Error)
|
||||
return channel
|
||||
}
|
||||
|
||||
func enableAdminLoginEvent(t *testing.T, dbConn *gorm.DB, channelName string, targets []string) {
|
||||
t.Helper()
|
||||
|
||||
var event model.PushEvent
|
||||
require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error)
|
||||
|
||||
event.Enabled = true
|
||||
event.Channels = []string{channelName}
|
||||
event.Targets = targets
|
||||
require.NoError(t, dbConn.Save(&event).Error)
|
||||
}
|
||||
|
||||
func waitForAsyncTrigger(t *testing.T) {
|
||||
t.Helper()
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
|
||||
func countPushTasks(t *testing.T, dbConn *gorm.DB) int64 {
|
||||
t.Helper()
|
||||
|
||||
var count int64
|
||||
require.NoError(t, dbConn.Model(&model.TaskExecution{}).
|
||||
Where("task_type = ?", push.SendNotificationTask).
|
||||
Count(&count).Error)
|
||||
return count
|
||||
}
|
||||
|
||||
func TestAdminLoginPushIntegration(t *testing.T) {
|
||||
dbConn, cleanup := setupAdminLoginIntegrationTest(t)
|
||||
defer cleanup()
|
||||
|
||||
channel := seedMockPushChannel(t, dbConn)
|
||||
defer dbConn.Delete(channel)
|
||||
|
||||
enableAdminLoginEvent(t, dbConn, channel.Name, []string{"ops_team"})
|
||||
|
||||
adminUser := &model.User{
|
||||
ID: 1001,
|
||||
Username: "super_admin",
|
||||
IsAdmin: true,
|
||||
IsActive: true,
|
||||
}
|
||||
require.NoError(t, dbConn.Create(adminUser).Error)
|
||||
|
||||
t.Run("admin login emits push task with user and ip", func(t *testing.T) {
|
||||
dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
|
||||
|
||||
listener.EmitAdminLoggedIn(context.Background(), adminUser, "203.0.113.42")
|
||||
waitForAsyncTrigger(t)
|
||||
|
||||
var execution model.TaskExecution
|
||||
require.NoError(t, dbConn.Where("task_type = ?", push.SendNotificationTask).First(&execution).Error)
|
||||
|
||||
var payload push.SendPayload
|
||||
require.NoError(t, json.Unmarshal([]byte(execution.Payload), &payload))
|
||||
|
||||
assert.Equal(t, AdminLogin.Key, payload.EventKey)
|
||||
assert.Equal(t, "ops_team", payload.Target)
|
||||
assert.Equal(t, "管理员登录提醒", payload.Body.Title)
|
||||
assert.Contains(t, payload.Body.Content, "super_admin")
|
||||
assert.Contains(t, payload.Body.Content, "203.0.113.42")
|
||||
})
|
||||
|
||||
t.Run("non-admin login does not trigger push", func(t *testing.T) {
|
||||
dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
|
||||
|
||||
nonAdmin := &model.User{
|
||||
ID: 2002,
|
||||
Username: "regular_user",
|
||||
IsAdmin: false,
|
||||
IsActive: true,
|
||||
}
|
||||
require.NoError(t, dbConn.Create(nonAdmin).Error)
|
||||
|
||||
listener.EmitAdminLoggedIn(context.Background(), nonAdmin, "198.51.100.1")
|
||||
waitForAsyncTrigger(t)
|
||||
|
||||
assert.Equal(t, int64(0), countPushTasks(t, dbConn))
|
||||
})
|
||||
|
||||
t.Run("disabled admin login event does not enqueue push", func(t *testing.T) {
|
||||
dbConn.Where("task_type = ?", push.SendNotificationTask).Delete(&model.TaskExecution{})
|
||||
|
||||
var event model.PushEvent
|
||||
require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error)
|
||||
event.Enabled = false
|
||||
require.NoError(t, dbConn.Save(&event).Error)
|
||||
|
||||
listener.EmitAdminLoggedIn(context.Background(), adminUser, "10.0.0.1")
|
||||
waitForAsyncTrigger(t)
|
||||
|
||||
assert.Equal(t, int64(0), countPushTasks(t, dbConn))
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package custom_events
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
|
||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||
)
|
||||
|
||||
// Register wires push notification handlers for domain events and registers
|
||||
// built-in event metadata. Must be called once during application bootstrap
|
||||
// before push.SyncEvents.
|
||||
func Register() {
|
||||
push.RegisterBuiltInEvent(AdminLogin)
|
||||
listener.OnAdminLoggedIn(handleAdminLogin)
|
||||
}
|
||||
@@ -0,0 +1,343 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package push defines push notification HTTP routes, background tasks, and events.
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// NotificationMessage represents the structured notification message payload.
|
||||
type NotificationMessage struct {
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Level string `json:"level"`
|
||||
Ext map[string]any `json:"ext,omitempty"`
|
||||
}
|
||||
|
||||
// Flatten converts the structured NotificationMessage back to a flat map (original json structure).
|
||||
func (m NotificationMessage) Flatten() map[string]any {
|
||||
res := map[string]any{
|
||||
keyTitle: m.Title,
|
||||
keyContent: m.Content,
|
||||
keyLevel: m.Level,
|
||||
}
|
||||
for k, v := range m.Ext {
|
||||
res[k] = v
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
// EventMetadata represents the metadata of a push notification event.
|
||||
type EventMetadata struct {
|
||||
Key string `json:"key"`
|
||||
Name string `json:"name"`
|
||||
DefaultTemplate NotificationMessage `json:"default_template"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// SendPayload 异步投递推送载荷 (供 task/Worker 使用)
|
||||
type SendPayload struct {
|
||||
EventKey string `json:"event_key"`
|
||||
Config pkgpush.Config `json:"config"`
|
||||
Target string `json:"target"`
|
||||
Body NotificationMessage `json:"body"`
|
||||
Template string `json:"template"`
|
||||
}
|
||||
|
||||
// BuiltInEvents lists all built-in events defined in custom_events.
|
||||
var BuiltInEvents []EventMetadata
|
||||
|
||||
// RegisterBuiltInEvent registers a built-in event definition.
|
||||
func RegisterBuiltInEvent(meta EventMetadata) {
|
||||
BuiltInEvents = append(BuiltInEvents, meta)
|
||||
}
|
||||
|
||||
// EventTrigger represents the unified event trigger class.
|
||||
type EventTrigger struct{}
|
||||
|
||||
// DefaultTrigger is the singleton instance of EventTrigger.
|
||||
var DefaultTrigger = &EventTrigger{}
|
||||
|
||||
// Trigger receives event metadata and processes the event notification dispatch asynchronously.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
|
||||
asyncCtx := context.WithoutCancel(ctx)
|
||||
go func() {
|
||||
if body == nil {
|
||||
body = make(map[string]any)
|
||||
}
|
||||
if _, hasUser := body["user"]; !hasUser || body["user"] == nil {
|
||||
body["user"] = getSystemUser(asyncCtx)
|
||||
}
|
||||
|
||||
eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return
|
||||
}
|
||||
logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err)
|
||||
return
|
||||
}
|
||||
event := *eventPtr
|
||||
if len(event.Channels) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
flatBody := getFlatBody(body)
|
||||
msg, _ := t.buildMessage(&event, meta, flatBody, body)
|
||||
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
|
||||
}()
|
||||
}
|
||||
|
||||
func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
|
||||
var msg NotificationMessage
|
||||
renderedTemplate := ""
|
||||
|
||||
templateSource := event.Template
|
||||
if templateSource != "" {
|
||||
var err error
|
||||
msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody)
|
||||
if err != nil {
|
||||
msg.Title = event.Name
|
||||
msg.Content = renderedTemplate
|
||||
msg.Level = defaultLevelInfo
|
||||
}
|
||||
} else {
|
||||
msg = t.parseDefaultTemplate(meta, flatBody)
|
||||
}
|
||||
|
||||
if msg.Ext == nil {
|
||||
msg.Ext = make(map[string]any)
|
||||
}
|
||||
for k, v := range body {
|
||||
if k == keyTitle || k == keyContent || k == keyLevel {
|
||||
continue
|
||||
}
|
||||
if _, exists := msg.Ext[k]; !exists {
|
||||
msg.Ext[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
return msg, renderedTemplate
|
||||
}
|
||||
|
||||
func (t *EventTrigger) parseCustomTemplate(event *model.PushEvent, templateSource string, flatBody map[string]any) (NotificationMessage, string, error) {
|
||||
var msg NotificationMessage
|
||||
renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody)
|
||||
|
||||
var tMap map[string]any
|
||||
if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil {
|
||||
return msg, renderedTemplate, err
|
||||
}
|
||||
|
||||
if title, ok := tMap[keyTitle].(string); ok && title != "" {
|
||||
msg.Title = title
|
||||
} else {
|
||||
msg.Title = event.Name
|
||||
}
|
||||
delete(tMap, keyTitle)
|
||||
|
||||
if content, ok := tMap[keyContent].(string); ok && content != "" {
|
||||
msg.Content = content
|
||||
} else {
|
||||
msg.Content = renderedTemplate
|
||||
}
|
||||
delete(tMap, keyContent)
|
||||
|
||||
if level, ok := tMap[keyLevel].(string); ok && level != "" {
|
||||
msg.Level = level
|
||||
} else {
|
||||
msg.Level = defaultLevelInfo
|
||||
}
|
||||
delete(tMap, keyLevel)
|
||||
|
||||
msg.Ext = tMap
|
||||
return msg, renderedTemplate, nil
|
||||
}
|
||||
|
||||
func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[string]any) NotificationMessage {
|
||||
var msg NotificationMessage
|
||||
msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody)
|
||||
msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody)
|
||||
msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody)
|
||||
|
||||
if meta.DefaultTemplate.Ext != nil {
|
||||
msg.Ext = make(map[string]any)
|
||||
for k, v := range meta.DefaultTemplate.Ext {
|
||||
if strVal, ok := v.(string); ok {
|
||||
msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody)
|
||||
} else {
|
||||
msg.Ext[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
return msg
|
||||
}
|
||||
|
||||
func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, msg NotificationMessage, flatBody map[string]any) {
|
||||
for _, channelName := range event.Channels {
|
||||
customChannel, err := repository.GetActivePushChannelByName(ctx, channelName)
|
||||
if err == nil {
|
||||
t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody)
|
||||
continue
|
||||
}
|
||||
logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, channel *model.PushChannel, msg NotificationMessage, flatBody map[string]any) {
|
||||
if len(event.Targets) == 0 {
|
||||
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg)
|
||||
return
|
||||
}
|
||||
|
||||
for _, target := range event.Targets {
|
||||
resolvedTarget := resolveTarget(ctx, target, flatBody, channel.Name)
|
||||
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg)
|
||||
}
|
||||
}
|
||||
|
||||
func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta EventMetadata, channel *model.PushChannel, target string, msg NotificationMessage) {
|
||||
var config pkgpush.Config
|
||||
var renderedTemplate string
|
||||
|
||||
switch channel.Type {
|
||||
case channelLark:
|
||||
config = pkgpush.Config{Channel: channelLark, URL: channel.URL, Secret: channel.Token}
|
||||
renderedTemplate = channel.Other
|
||||
case channelEmail:
|
||||
url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other)
|
||||
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
|
||||
case channelTelegram:
|
||||
config = pkgpush.Config{Channel: channelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other}
|
||||
default:
|
||||
config = pkgpush.Config{Channel: channelCustom, URL: channel.URL}
|
||||
customPushReq := CustomPushRequest{
|
||||
Title: msg.Title,
|
||||
Content: msg.Content,
|
||||
Description: meta.Description,
|
||||
To: target,
|
||||
}
|
||||
if urlVal, ok := msg.Ext["url"].(string); ok {
|
||||
customPushReq.URL = urlVal
|
||||
}
|
||||
renderedTemplate = renderCustomPayload(channel.Other, customPushReq)
|
||||
}
|
||||
|
||||
payload := SendPayload{
|
||||
EventKey: meta.Key,
|
||||
Config: config,
|
||||
Target: target,
|
||||
Body: msg,
|
||||
Template: renderedTemplate,
|
||||
}
|
||||
if err := enqueuePushTask(ctx, payload); err != nil {
|
||||
logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err)
|
||||
}
|
||||
}
|
||||
|
||||
func enqueuePushTask(ctx context.Context, payload SendPayload) error {
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = task.DispatchTask(ctx, "send_notification", payloadBytes, "system")
|
||||
return err
|
||||
}
|
||||
|
||||
func getFlatBody(body map[string]any) map[string]any {
|
||||
jsonBytes, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return body
|
||||
}
|
||||
var jsonMap map[string]any
|
||||
if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil {
|
||||
return body
|
||||
}
|
||||
|
||||
flatResult := make(map[string]any)
|
||||
flattenMap("", jsonMap, flatResult)
|
||||
return flatResult
|
||||
}
|
||||
|
||||
func flattenMap(prefix string, m map[string]any, result map[string]any) {
|
||||
for k, v := range m {
|
||||
key := k
|
||||
if prefix != "" {
|
||||
key = prefix + "." + k
|
||||
}
|
||||
if nestedMap, ok := v.(map[string]any); ok {
|
||||
flattenMap(key, nestedMap, result)
|
||||
} else {
|
||||
result[key] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func resolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string {
|
||||
target = strings.TrimSpace(target)
|
||||
if target == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
resolved := resolveDynamicKeyword(target, flatBody)
|
||||
if strings.Contains(resolved, "@") {
|
||||
return resolved
|
||||
}
|
||||
if val, matched := resolveSystemTarget(ctx, resolved, channel); matched {
|
||||
return val
|
||||
}
|
||||
|
||||
user, found := resolveTargetUser(ctx, resolved, channel)
|
||||
if !found {
|
||||
return resolved
|
||||
}
|
||||
if channel == channelEmail && user.Email != "" {
|
||||
return user.Email
|
||||
}
|
||||
if channel != channelEmail && user.Username != "" {
|
||||
return user.Username
|
||||
}
|
||||
return resolved
|
||||
}
|
||||
|
||||
func resolveDynamicKeyword(target string, flatBody map[string]any) string {
|
||||
switch target {
|
||||
case "user.id", "id":
|
||||
if val, ok := flatBody["user.id"]; ok {
|
||||
return fmt.Sprintf("%v", val)
|
||||
}
|
||||
if val, ok := flatBody["id"]; ok {
|
||||
return fmt.Sprintf("%v", val)
|
||||
}
|
||||
case "user.username", "username":
|
||||
if val, ok := flatBody["user.username"]; ok {
|
||||
return fmt.Sprintf("%v", val)
|
||||
}
|
||||
if val, ok := flatBody["username"]; ok {
|
||||
return fmt.Sprintf("%v", val)
|
||||
}
|
||||
case "user.email", channelEmail:
|
||||
if val, ok := flatBody["user.email"]; ok {
|
||||
return fmt.Sprintf("%v", val)
|
||||
}
|
||||
if val, ok := flatBody["email"]; ok {
|
||||
return fmt.Sprintf("%v", val)
|
||||
}
|
||||
}
|
||||
return target
|
||||
}
|
||||
@@ -0,0 +1,413 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type smtpConfig struct {
|
||||
Host string
|
||||
Port string
|
||||
Username string
|
||||
Password string
|
||||
}
|
||||
|
||||
func loadSMTPConfig(ctx context.Context) smtpConfig {
|
||||
host, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost)
|
||||
port, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort)
|
||||
user, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername)
|
||||
pass, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword)
|
||||
return smtpConfig{
|
||||
Host: host.Value,
|
||||
Port: port.Value,
|
||||
Username: user.Value,
|
||||
Password: pass.Value,
|
||||
}
|
||||
}
|
||||
|
||||
func syncBuiltInEvents(ctx context.Context) error {
|
||||
for _, meta := range BuiltInEvents {
|
||||
_, err := repository.GetPushEventByKey(ctx, meta.Key)
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
var defaultTemplateStr string
|
||||
if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil {
|
||||
defaultTemplateStr = string(defaultTemplateBytes)
|
||||
}
|
||||
event := model.PushEvent{
|
||||
EventKey: meta.Key,
|
||||
Name: meta.Name,
|
||||
Channels: []string{},
|
||||
Targets: []string{},
|
||||
Template: defaultTemplateStr,
|
||||
Enabled: false,
|
||||
}
|
||||
if err := repository.CreatePushEvent(ctx, &event); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func listPushEvents(ctx context.Context) ([]model.PushEvent, error) {
|
||||
return repository.ListPushEvents(ctx)
|
||||
}
|
||||
|
||||
func createPushEvent(ctx context.Context, req CreateEventRequest) (model.PushEvent, error) {
|
||||
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
|
||||
if err != nil {
|
||||
return model.PushEvent{}, err
|
||||
}
|
||||
|
||||
count, err := repository.CountPushEventsByKey(ctx, eventKey)
|
||||
if err != nil {
|
||||
return model.PushEvent{}, err
|
||||
}
|
||||
if count > 0 {
|
||||
return model.PushEvent{}, errors.New("this notification event is already configured")
|
||||
}
|
||||
|
||||
templateStr := strings.TrimSpace(req.Template)
|
||||
if templateStr == "" {
|
||||
templateStr = string(defaultTemplateBytes)
|
||||
} else {
|
||||
var tempMap map[string]any
|
||||
if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil {
|
||||
return model.PushEvent{}, errors.New("custom template is not a valid JSON format")
|
||||
}
|
||||
}
|
||||
|
||||
channels := req.Channels
|
||||
if channels == nil {
|
||||
channels = []string{}
|
||||
}
|
||||
targets := req.Targets
|
||||
if targets == nil {
|
||||
targets = []string{}
|
||||
}
|
||||
|
||||
event := model.PushEvent{
|
||||
EventKey: eventKey,
|
||||
Name: eventName,
|
||||
TaskType: req.TaskType,
|
||||
Channels: channels,
|
||||
Targets: targets,
|
||||
Template: templateStr,
|
||||
Enabled: req.Enabled,
|
||||
}
|
||||
if err := event.Validate(); err != nil {
|
||||
return model.PushEvent{}, err
|
||||
}
|
||||
if err := repository.CreatePushEvent(ctx, &event); err != nil {
|
||||
return model.PushEvent{}, err
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
|
||||
func deletePushEvent(ctx context.Context, id uint64) error {
|
||||
event, err := repository.GetPushEventByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return repository.DeletePushEvent(ctx, &event)
|
||||
}
|
||||
|
||||
func updatePushEvent(ctx context.Context, id uint64, req UpdateEventRequest) error {
|
||||
event, err := repository.GetPushEventByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
event.Channels = req.Channels
|
||||
event.Targets = req.Targets
|
||||
event.Template = req.Template
|
||||
event.Enabled = req.Enabled
|
||||
if err := event.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return repository.SavePushEvent(ctx, &event)
|
||||
}
|
||||
|
||||
func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
|
||||
event, err := repository.GetPushEventByID(ctx, id)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
enabled := !event.Enabled
|
||||
if enabled && len(event.Channels) == 0 {
|
||||
return false, errors.New("cannot enable event without any push channels configured")
|
||||
}
|
||||
if err := repository.UpdatePushEventEnabled(ctx, &event, enabled); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return enabled, nil
|
||||
}
|
||||
|
||||
func listPushHistories(ctx context.Context, filter repository.PushHistoryListFilter) (int64, []model.PushHistory, error) {
|
||||
return repository.ListPushHistories(ctx, filter)
|
||||
}
|
||||
|
||||
func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) {
|
||||
if cfg.Channel != channelEmail || (cfg.URL != "" && cfg.Key != "") {
|
||||
return
|
||||
}
|
||||
smtp := loadSMTPConfig(ctx)
|
||||
if smtp.Host == "" || smtp.Username == "" {
|
||||
return
|
||||
}
|
||||
port := smtp.Port
|
||||
if port == "" {
|
||||
port = "587"
|
||||
}
|
||||
cfg.URL = smtp.Host + ":" + port
|
||||
cfg.Key = smtp.Username
|
||||
cfg.Secret = smtp.Password
|
||||
}
|
||||
|
||||
func listPushChannels(ctx context.Context) ([]model.PushChannel, error) {
|
||||
return repository.ListPushChannels(ctx)
|
||||
}
|
||||
|
||||
func createPushChannel(ctx context.Context, req CreateChannelRequest) (model.PushChannel, error) {
|
||||
count, err := repository.CountPushChannelsByName(ctx, req.Name)
|
||||
if err != nil {
|
||||
return model.PushChannel{}, err
|
||||
}
|
||||
if count > 0 {
|
||||
return model.PushChannel{}, errors.New("channel name already exists")
|
||||
}
|
||||
|
||||
channel := model.PushChannel{
|
||||
Name: req.Name,
|
||||
Description: req.Description,
|
||||
Type: req.Type,
|
||||
Token: req.Token,
|
||||
URL: req.URL,
|
||||
Other: req.Other,
|
||||
Enabled: req.Enabled,
|
||||
}
|
||||
if err := channel.Validate(); err != nil {
|
||||
return model.PushChannel{}, err
|
||||
}
|
||||
if err := repository.CreatePushChannel(ctx, &channel); err != nil {
|
||||
return model.PushChannel{}, err
|
||||
}
|
||||
return channel, nil
|
||||
}
|
||||
|
||||
func updatePushChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (model.PushChannel, error) {
|
||||
channel, err := repository.GetPushChannelByID(ctx, id)
|
||||
if err != nil {
|
||||
return model.PushChannel{}, err
|
||||
}
|
||||
|
||||
channel.Description = req.Description
|
||||
channel.Type = req.Type
|
||||
channel.Token = req.Token
|
||||
channel.URL = req.URL
|
||||
channel.Other = req.Other
|
||||
channel.Enabled = req.Enabled
|
||||
if err := channel.Validate(); err != nil {
|
||||
return model.PushChannel{}, err
|
||||
}
|
||||
if err := repository.SavePushChannel(ctx, &channel); err != nil {
|
||||
return model.PushChannel{}, err
|
||||
}
|
||||
return channel, nil
|
||||
}
|
||||
|
||||
func deletePushChannel(ctx context.Context, id uint64) error {
|
||||
channel, err := repository.GetPushChannelByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return repository.DeletePushChannel(ctx, &channel)
|
||||
}
|
||||
|
||||
func loadChannelForTest(ctx context.Context, req TestChannelRequest) (string, string, string, string, error) {
|
||||
if req.Name != "" {
|
||||
channel, err := repository.GetPushChannelByName(ctx, req.Name)
|
||||
if err != nil {
|
||||
return "", "", "", "", errors.New("channel not found")
|
||||
}
|
||||
return channel.URL, channel.Token, channel.Other, channel.Type, nil
|
||||
}
|
||||
return req.URL, req.Token, req.Other, req.Type, nil
|
||||
}
|
||||
|
||||
func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) {
|
||||
return repository.ListActivePushEventsByTaskType(ctx, taskType)
|
||||
}
|
||||
|
||||
func loadUserFromPayload(ctx context.Context, data map[string]any) any {
|
||||
if u, exists := data["user"]; exists && u != nil {
|
||||
return u
|
||||
}
|
||||
|
||||
if userID, ok := extractUserID(data); ok && userID > 0 {
|
||||
if user, err := repository.GetUserByID(ctx, userID); err == nil {
|
||||
return &user
|
||||
}
|
||||
}
|
||||
|
||||
if username := extractUsername(data); username != "" {
|
||||
if user, err := repository.GetUserByUsername(ctx, username); err == nil {
|
||||
return &user
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg string) error {
|
||||
title := req.Body.Title
|
||||
content := req.Body.Content
|
||||
level := req.Body.Level
|
||||
if title == "" {
|
||||
title = "系统通知"
|
||||
}
|
||||
if level == "" {
|
||||
level = defaultLevelInfo
|
||||
}
|
||||
|
||||
target := req.Target
|
||||
if target == "" {
|
||||
if req.Config.URL != "" {
|
||||
target = req.Config.URL
|
||||
const maxTargetLen = 50
|
||||
const truncatedLen = 47
|
||||
if len(target) > maxTargetLen {
|
||||
target = target[:truncatedLen] + "..."
|
||||
}
|
||||
} else {
|
||||
target = "default"
|
||||
}
|
||||
}
|
||||
|
||||
history := model.PushHistory{
|
||||
EventKey: req.EventKey,
|
||||
Channel: req.Config.Channel,
|
||||
Target: target,
|
||||
Title: title,
|
||||
Content: content,
|
||||
Level: level,
|
||||
Status: status,
|
||||
ErrorMsg: errMsg,
|
||||
}
|
||||
return repository.CreatePushHistory(ctx, &history)
|
||||
}
|
||||
|
||||
func resolveTargetUser(ctx context.Context, resolved string, _ string) (model.User, bool) {
|
||||
found := false
|
||||
var user model.User
|
||||
|
||||
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
|
||||
if u, err := repository.GetUserByID(ctx, id); err == nil {
|
||||
user = u
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
if u, err := repository.GetUserByUsername(ctx, resolved); err == nil {
|
||||
user = u
|
||||
found = true
|
||||
}
|
||||
}
|
||||
return user, found
|
||||
}
|
||||
|
||||
func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) {
|
||||
if resolved != "系统" && resolved != "system" && resolved != "0" {
|
||||
return "", false
|
||||
}
|
||||
adminUser, err := repository.GetFirstAdminUser(ctx)
|
||||
if err != nil {
|
||||
return resolved, true
|
||||
}
|
||||
if channel == channelEmail && adminUser.Email != "" {
|
||||
return adminUser.Email, true
|
||||
}
|
||||
if channel != channelEmail && adminUser.Username != "" {
|
||||
return adminUser.Username, true
|
||||
}
|
||||
return resolved, true
|
||||
}
|
||||
|
||||
func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) {
|
||||
if url != "" && token != "" {
|
||||
return url, token, other
|
||||
}
|
||||
smtp := loadSMTPConfig(ctx)
|
||||
if smtp.Host == "" || smtp.Username == "" {
|
||||
return url, token, other
|
||||
}
|
||||
port := smtp.Port
|
||||
if port == "" {
|
||||
port = "587"
|
||||
}
|
||||
if url == "" {
|
||||
url = smtp.Host + ":" + port
|
||||
}
|
||||
if token == "" {
|
||||
token = smtp.Username
|
||||
}
|
||||
if other == "" {
|
||||
other = smtp.Password
|
||||
}
|
||||
return url, token, other
|
||||
}
|
||||
|
||||
func getSystemUser(ctx context.Context) *model.User {
|
||||
user := repository.GetSystemUser(ctx)
|
||||
return &user
|
||||
}
|
||||
|
||||
func getEventInfo(req CreateEventRequest) (string, string, []byte, error) {
|
||||
if req.TaskType != "" {
|
||||
meta := task.GetTaskMetaByAsynqTask(req.TaskType)
|
||||
if meta == nil {
|
||||
return "", "", nil, errors.New("unsupported task type")
|
||||
}
|
||||
eventKey := "task_completed:" + req.TaskType
|
||||
eventName := "任务完成: " + meta.Name
|
||||
defaultTemplate := NotificationMessage{
|
||||
Title: "任务完成: " + meta.Name,
|
||||
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
|
||||
Level: defaultLevelInfo,
|
||||
}
|
||||
defaultTemplateBytes, err := json.Marshal(defaultTemplate)
|
||||
if err != nil {
|
||||
return "", "", nil, err
|
||||
}
|
||||
return eventKey, eventName, defaultTemplateBytes, nil
|
||||
}
|
||||
|
||||
if req.EventKey == "" {
|
||||
return "", "", nil, errors.New("either event_key or task_type must be provided")
|
||||
}
|
||||
|
||||
meta, found := findBuiltInEvent(req.EventKey)
|
||||
if !found {
|
||||
return "", "", nil, errors.New("unsupported built-in event key")
|
||||
}
|
||||
|
||||
defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate)
|
||||
if err != nil {
|
||||
return "", "", nil, err
|
||||
}
|
||||
return req.EventKey, meta.Name, defaultTemplateBytes, nil
|
||||
}
|
||||
@@ -0,0 +1,759 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
var adminLoginEvent = EventMetadata{
|
||||
Key: "admin_login",
|
||||
Name: "管理员登录",
|
||||
DefaultTemplate: NotificationMessage{
|
||||
Title: "管理员登录提醒",
|
||||
Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。",
|
||||
Level: "INFO",
|
||||
},
|
||||
Description: "当管理员成功登录系统时触发此通知",
|
||||
}
|
||||
|
||||
func init() {
|
||||
RegisterBuiltInEvent(adminLoginEvent)
|
||||
}
|
||||
|
||||
// mockPusher mock implementation of pkgpush.Pusher
|
||||
type mockPusher struct {
|
||||
mu sync.Mutex
|
||||
sentBody map[string]any
|
||||
sentTgt string
|
||||
}
|
||||
|
||||
func (m *mockPusher) Send(ctx context.Context, cfg pkgpush.Config, target string, body map[string]any, template string, ext map[string]any) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.sentBody = body
|
||||
m.sentTgt = target
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockPusher) ValidateConfig(cfg pkgpush.Config) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func setupPushTest(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) {
|
||||
dbConn, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
|
||||
// AutoMigrate push tables in SQLite test environment
|
||||
err := dbConn.AutoMigrate(&model.PushEvent{}, &model.PushHistory{}, &model.User{}, &model.PushChannel{}, &model.SystemConfig{})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 写入数据库系统默认用户 Seed 记录
|
||||
sysUser := &model.User{
|
||||
ID: 999,
|
||||
Username: "system",
|
||||
Nickname: "系统",
|
||||
Password: "*",
|
||||
IsActive: true,
|
||||
}
|
||||
err = dbConn.Create(sysUser).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
// Initialize AsynqClient pointing to miniredis
|
||||
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
|
||||
Addr: mr.Addr(),
|
||||
})
|
||||
|
||||
// Register the task handler and metadata
|
||||
task.RegisterHandler(SendNotificationTask, &PushHandler{})
|
||||
task.RegisterTaskMeta(SendNotificationMeta)
|
||||
|
||||
return dbConn, mr, func() {
|
||||
cleanup()
|
||||
if task.AsynqClient != nil {
|
||||
task.AsynqClient.Close()
|
||||
task.AsynqClient = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
r := testhelper.NewTestGinEngine()
|
||||
adminGroup := r.Group("/api/v1/admin/push")
|
||||
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
oauth.SetToContext(c, "user_obj", authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.GET("/events", ListEvents)
|
||||
adminGroup.GET("/events/builtin", ListBuiltInEvents)
|
||||
adminGroup.POST("/events", CreateEvent)
|
||||
adminGroup.PUT("/events/:id", UpdateEvent)
|
||||
adminGroup.DELETE("/events/:id", DeleteEvent)
|
||||
adminGroup.POST("/events/:id/toggle", ToggleEvent)
|
||||
adminGroup.GET("/histories", ListHistories)
|
||||
adminGroup.POST("/test", TestPush)
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
func TestSyncEvents(t *testing.T) {
|
||||
dbConn, _, cleanup := setupPushTest(t)
|
||||
defer cleanup()
|
||||
|
||||
// 1. SyncEvents first time
|
||||
err := SyncEvents(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify event exists in DB
|
||||
var event model.PushEvent
|
||||
err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "管理员登录", event.Name)
|
||||
assert.False(t, event.Enabled)
|
||||
|
||||
// Verify DefaultTemplate matches GORM template field
|
||||
var defaultMsg NotificationMessage
|
||||
err = json.Unmarshal([]byte(event.Template), &defaultMsg)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, adminLoginEvent.DefaultTemplate.Title, defaultMsg.Title)
|
||||
assert.Equal(t, adminLoginEvent.DefaultTemplate.Content, defaultMsg.Content)
|
||||
}
|
||||
|
||||
func TestEventTrigger(t *testing.T) {
|
||||
dbConn, _, cleanup := setupPushTest(t)
|
||||
defer cleanup()
|
||||
|
||||
// SyncEvents
|
||||
err := SyncEvents(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("trigger disabled event silently ignored", func(t *testing.T) {
|
||||
body := map[string]any{
|
||||
"user": map[string]any{"username": "test_admin"},
|
||||
"ip": "127.0.0.1",
|
||||
}
|
||||
DefaultTrigger.Trigger(context.Background(), adminLoginEvent, body)
|
||||
|
||||
// Sleep briefly since Trigger runs in goroutine
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Verify no tasks enqueued in TaskExecution GORM table
|
||||
var count int64
|
||||
dbConn.Model(&model.TaskExecution{}).Count(&count)
|
||||
assert.Equal(t, int64(0), count)
|
||||
})
|
||||
|
||||
t.Run("trigger enabled event enqueues task", func(t *testing.T) {
|
||||
// Create an enabled custom channel in GORM
|
||||
customChan := &model.PushChannel{
|
||||
Name: "mock_channel",
|
||||
Type: "custom",
|
||||
URL: "https://webhook.site/trigger",
|
||||
Other: `{"text": "$content"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
err = dbConn.Create(customChan).Error
|
||||
require.NoError(t, err)
|
||||
defer dbConn.Delete(customChan)
|
||||
|
||||
// Enable the push event in DB using struct to trigger JSON serializer
|
||||
var event model.PushEvent
|
||||
err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
event.Enabled = true
|
||||
event.Channels = []string{"mock_channel"}
|
||||
event.Targets = []string{"admin_user"}
|
||||
err = dbConn.Save(&event).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
// Trigger
|
||||
body := map[string]any{
|
||||
"user": map[string]any{
|
||||
"username": "super_admin",
|
||||
},
|
||||
"ip": "1.1.1.1",
|
||||
"time": "2026-06-14 18:00:00",
|
||||
}
|
||||
DefaultTrigger.Trigger(context.Background(), adminLoginEvent, body)
|
||||
|
||||
// Wait for goroutine execution
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Verify TaskExecution enqueued record
|
||||
var execution model.TaskExecution
|
||||
err = dbConn.Where("task_type = ?", SendNotificationTask).First(&execution).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify enqueued payload structure
|
||||
var payload SendPayload
|
||||
err = json.Unmarshal([]byte(execution.Payload), &payload)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "admin_login", payload.EventKey)
|
||||
assert.Equal(t, "custom", payload.Config.Channel)
|
||||
assert.Equal(t, "https://webhook.site/trigger", payload.Config.URL)
|
||||
assert.Equal(t, "admin_user", payload.Target)
|
||||
assert.Equal(t, "管理员登录提醒", payload.Body.Title)
|
||||
assert.Contains(t, payload.Body.Content, "super_admin")
|
||||
assert.Contains(t, payload.Body.Content, "1.1.1.1")
|
||||
})
|
||||
|
||||
t.Run("trigger without user injects virtual system user", func(t *testing.T) {
|
||||
// Create an enabled custom channel in GORM
|
||||
customChan := &model.PushChannel{
|
||||
Name: "mock_channel",
|
||||
Type: "custom",
|
||||
URL: "https://webhook.site/trigger",
|
||||
Other: `{"text": "$content"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
err = dbConn.Create(customChan).Error
|
||||
require.NoError(t, err)
|
||||
defer dbConn.Delete(customChan)
|
||||
|
||||
// Enable the push event in DB
|
||||
var event model.PushEvent
|
||||
err = dbConn.Where("event_key = ?", "admin_login").First(&event).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
// 清理旧任务执行记录
|
||||
dbConn.Where("task_type = ?", SendNotificationTask).Delete(&model.TaskExecution{})
|
||||
|
||||
event.Enabled = true
|
||||
event.Channels = []string{"mock_channel"}
|
||||
event.Targets = []string{"user.username"} // 动态目标
|
||||
err = dbConn.Save(&event).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
// Trigger with empty body (simulates cron scheduler triggering)
|
||||
DefaultTrigger.Trigger(context.Background(), adminLoginEvent, nil)
|
||||
|
||||
// Wait for goroutine execution
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Verify TaskExecution enqueued record
|
||||
var execution model.TaskExecution
|
||||
err = dbConn.Where("task_type = ?", SendNotificationTask).First(&execution).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
var payload SendPayload
|
||||
err = json.Unmarshal([]byte(execution.Payload), &payload)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 检查 payload 是否将 target (user.username) 成功替换为 "system"
|
||||
assert.Equal(t, "system", payload.Target)
|
||||
// 检查 payload 中的 Content,应当被替换为 "system" 变量
|
||||
assert.Contains(t, payload.Body.Content, "system")
|
||||
})
|
||||
}
|
||||
|
||||
func TestPushHandler(t *testing.T) {
|
||||
dbConn, _, cleanup := setupPushTest(t)
|
||||
defer cleanup()
|
||||
|
||||
mPusher := &mockPusher{}
|
||||
pkgpush.Register("mock_channel", mPusher)
|
||||
|
||||
handler := &PushHandler{}
|
||||
|
||||
payload := SendPayload{
|
||||
EventKey: "admin_login",
|
||||
Config: pkgpush.Config{
|
||||
Channel: "mock_channel",
|
||||
URL: "http://mock-url",
|
||||
},
|
||||
Target: "admin_user",
|
||||
Body: NotificationMessage{
|
||||
Title: "Structured Alert",
|
||||
Content: "Hello World",
|
||||
Level: "WARNING",
|
||||
Ext: map[string]any{"extra_val": 42},
|
||||
},
|
||||
}
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("validate payload", func(t *testing.T) {
|
||||
validated, valErr := handler.ValidatePayload(payloadBytes)
|
||||
require.NoError(t, valErr)
|
||||
assert.NotEmpty(t, validated)
|
||||
})
|
||||
|
||||
t.Run("execute task successfully", func(t *testing.T) {
|
||||
res, execErr := handler.Execute(context.Background(), payloadBytes)
|
||||
require.NoError(t, execErr)
|
||||
assert.Contains(t, res.Message, "推送成功")
|
||||
|
||||
// Verify mock pusher received flattened variables
|
||||
mPusher.mu.Lock()
|
||||
assert.Equal(t, "admin_user", mPusher.sentTgt)
|
||||
assert.Equal(t, "Structured Alert", mPusher.sentBody["title"])
|
||||
assert.Equal(t, "Hello World", mPusher.sentBody["content"])
|
||||
assert.Equal(t, "WARNING", mPusher.sentBody["level"])
|
||||
assert.Equal(t, float64(42), mPusher.sentBody["extra_val"]) // unmarshaled json numbers are float64 by default
|
||||
mPusher.mu.Unlock()
|
||||
|
||||
// Verify PushHistory recorded
|
||||
var history model.PushHistory
|
||||
err = dbConn.First(&history).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "admin_login", history.EventKey)
|
||||
assert.Equal(t, "mock_channel", history.Channel)
|
||||
assert.Equal(t, "success", history.Status)
|
||||
assert.Equal(t, "Structured Alert", history.Title)
|
||||
})
|
||||
}
|
||||
|
||||
func TestPushRouters(t *testing.T) {
|
||||
dbConn, _, cleanup := setupPushTest(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
r := setupTestRouter(adminUser)
|
||||
|
||||
// Sync events to populate db
|
||||
err := SyncEvents(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("list events", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/push/events", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
var resp response.Any
|
||||
err = json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
require.NoError(t, err)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var events []model.PushEvent
|
||||
err = json.Unmarshal(dataBytes, &events)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Len(t, events, 1)
|
||||
assert.Equal(t, "admin_login", events[0].EventKey)
|
||||
})
|
||||
|
||||
t.Run("toggle event status", func(t *testing.T) {
|
||||
var event model.PushEvent
|
||||
dbConn.First(&event)
|
||||
|
||||
// 1. 未配置任何渠道时开启,应该被拒绝
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
|
||||
// 2. 为该事件关联渠道后,再切换开启,应当成功
|
||||
event.Channels = []string{"email"}
|
||||
dbConn.Save(&event)
|
||||
|
||||
req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil)
|
||||
w2 := httptest.NewRecorder()
|
||||
r.ServeHTTP(w2, req2)
|
||||
assert.Equal(t, http.StatusOK, w2.Code)
|
||||
|
||||
var updated model.PushEvent
|
||||
dbConn.First(&updated)
|
||||
assert.True(t, updated.Enabled)
|
||||
})
|
||||
|
||||
t.Run("update event", func(t *testing.T) {
|
||||
var event model.PushEvent
|
||||
dbConn.First(&event)
|
||||
|
||||
updateReq := UpdateEventRequest{
|
||||
Channels: []string{"email"},
|
||||
Targets: []string{"user@test.com"},
|
||||
Template: `{"title": "Custom Login Alert", "content": "Alert", "level": "WARNING"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
bodyBytes, _ := json.Marshal(updateReq)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10), bytes.NewBuffer(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var updated model.PushEvent
|
||||
dbConn.First(&updated)
|
||||
assert.Equal(t, []string{"email"}, updated.Channels)
|
||||
assert.Equal(t, []string{"user@test.com"}, updated.Targets)
|
||||
assert.Contains(t, updated.Template, "Custom Login Alert")
|
||||
})
|
||||
|
||||
t.Run("list push histories", func(t *testing.T) {
|
||||
// Populate history record
|
||||
hist := model.PushHistory{
|
||||
EventKey: "admin_login",
|
||||
Channel: "email",
|
||||
Target: "user@test.com",
|
||||
Title: "Custom Login Alert",
|
||||
Content: "Alert",
|
||||
Level: "WARNING",
|
||||
Status: "success",
|
||||
}
|
||||
dbConn.Create(&hist)
|
||||
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/push/histories?page=1&page_size=10", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataMap, ok := resp.Data.(map[string]any)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, float64(1), dataMap["total"])
|
||||
})
|
||||
|
||||
t.Run("test push endpoint", func(t *testing.T) {
|
||||
mPusher := &mockPusher{}
|
||||
pkgpush.Register("test_channel", mPusher)
|
||||
|
||||
testReq := TestPushRequest{
|
||||
Config: pkgpush.Config{
|
||||
Channel: "test_channel",
|
||||
URL: "http://test-url",
|
||||
},
|
||||
Target: "test_target",
|
||||
}
|
||||
bodyBytes, _ := json.Marshal(testReq)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/push/test", bytes.NewBuffer(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
})
|
||||
|
||||
t.Run("list built-in events", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/push/events/builtin", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
var resp response.Any
|
||||
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
require.NoError(t, err)
|
||||
|
||||
builtins, ok := resp.Data.([]any)
|
||||
assert.True(t, ok)
|
||||
assert.NotEmpty(t, builtins)
|
||||
})
|
||||
|
||||
t.Run("create and delete push event", func(t *testing.T) {
|
||||
// Clean up any existing admin_login event first
|
||||
dbConn.Where("event_key = ?", "admin_login").Delete(&model.PushEvent{})
|
||||
|
||||
// 1. Create event
|
||||
createReq := CreateEventRequest{
|
||||
EventKey: "admin_login",
|
||||
Channels: []string{"email"},
|
||||
Targets: []string{"admin@test.com"},
|
||||
Enabled: true,
|
||||
}
|
||||
bodyBytes, _ := json.Marshal(createReq)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/push/events", bytes.NewBuffer(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
// Verify created in DB
|
||||
var event model.PushEvent
|
||||
err := dbConn.Where("event_key = ?", "admin_login").First(&event).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "admin_login", event.EventKey)
|
||||
assert.Equal(t, "管理员登录", event.Name)
|
||||
assert.True(t, event.Enabled)
|
||||
|
||||
// 2. Try creating again (should fail)
|
||||
w2 := httptest.NewRecorder()
|
||||
req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events", bytes.NewBuffer(bodyBytes))
|
||||
req2.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w2, req2)
|
||||
assert.Equal(t, http.StatusBadRequest, w2.Code)
|
||||
|
||||
// 3. Delete event
|
||||
w3 := httptest.NewRecorder()
|
||||
req3, _ := http.NewRequest("DELETE", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10), nil)
|
||||
r.ServeHTTP(w3, req3)
|
||||
assert.Equal(t, http.StatusOK, w3.Code)
|
||||
|
||||
// Verify deleted from DB
|
||||
var count int64
|
||||
dbConn.Model(&model.PushEvent{}).Where("event_key = ?", "admin_login").Count(&count)
|
||||
assert.Equal(t, int64(0), count)
|
||||
})
|
||||
}
|
||||
|
||||
func TestResolveTarget(t *testing.T) {
|
||||
dbConn, _, cleanup := setupPushTest(t)
|
||||
defer cleanup()
|
||||
|
||||
// 1. 创建测试用户与管理员用户
|
||||
testUser := &model.User{
|
||||
ID: 9999,
|
||||
Username: "target_user",
|
||||
Email: "target@test.com",
|
||||
IsAdmin: false,
|
||||
}
|
||||
err := dbConn.Create(testUser).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
adminUser := &model.User{
|
||||
ID: 8888,
|
||||
Username: "admin_user",
|
||||
Email: "admin@test.com",
|
||||
IsAdmin: true,
|
||||
}
|
||||
err = dbConn.Create(adminUser).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
flatBody := map[string]any{
|
||||
"user.id": float64(9999), // JSON 反序列化后一般是 float64
|
||||
"user.username": "target_user",
|
||||
"user.email": "target@test.com",
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("dynamic user.id resolved and converted for email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "user.id", flatBody, "email")
|
||||
assert.Equal(t, "target@test.com", res)
|
||||
})
|
||||
|
||||
t.Run("dynamic user.username resolved and converted for email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "user.username", flatBody, "email")
|
||||
assert.Equal(t, "target@test.com", res)
|
||||
})
|
||||
|
||||
t.Run("dynamic user.email resolved directly for email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "user.email", flatBody, "email")
|
||||
assert.Equal(t, "target@test.com", res)
|
||||
})
|
||||
|
||||
t.Run("fixed user id resolved and converted for email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "9999", flatBody, "email")
|
||||
assert.Equal(t, "target@test.com", res)
|
||||
})
|
||||
|
||||
t.Run("fixed username resolved and converted for email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "target_user", flatBody, "email")
|
||||
assert.Equal(t, "target@test.com", res)
|
||||
})
|
||||
|
||||
t.Run("fixed email address resolved directly for email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "fixed@example.com", flatBody, "email")
|
||||
assert.Equal(t, "fixed@example.com", res)
|
||||
})
|
||||
|
||||
t.Run("fixed username resolved for non-email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "target_user", flatBody, "lark")
|
||||
assert.Equal(t, "target_user", res)
|
||||
})
|
||||
|
||||
t.Run("non-exist user resolved as fallback", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "non_exist_user", flatBody, "email")
|
||||
assert.Equal(t, "non_exist_user", res)
|
||||
})
|
||||
|
||||
t.Run("system target resolves to admin email for email channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "系统", flatBody, "email")
|
||||
assert.Equal(t, "admin@test.com", res)
|
||||
|
||||
res2 := resolveTarget(ctx, "system", flatBody, "email")
|
||||
assert.Equal(t, "admin@test.com", res2)
|
||||
|
||||
res3 := resolveTarget(ctx, "0", flatBody, "email")
|
||||
assert.Equal(t, "admin@test.com", res3)
|
||||
})
|
||||
|
||||
t.Run("system target resolves to admin username for lark channel", func(t *testing.T) {
|
||||
res := resolveTarget(ctx, "系统", flatBody, "lark")
|
||||
assert.Equal(t, "admin_user", res)
|
||||
})
|
||||
}
|
||||
|
||||
func TestPushChannelAPI(t *testing.T) {
|
||||
// 1. 模型校验测试
|
||||
t.Run("validate push channel model constraints", func(t *testing.T) {
|
||||
// 校验名称合法性
|
||||
c1 := &model.PushChannel{Name: "invalid-name!", URL: "https://hook.com", Other: "{}"}
|
||||
assert.Error(t, c1.Validate())
|
||||
|
||||
// 校验 URL 安全前缀 HTTPS
|
||||
c2 := &model.PushChannel{Name: "custom_channel", URL: "http://insecure-hook.com", Other: "{}"}
|
||||
assert.Error(t, c2.Validate())
|
||||
|
||||
// 校验 JSON 格式
|
||||
c3 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{invalid-json}"}
|
||||
assert.Error(t, c3.Validate())
|
||||
|
||||
// 正确配置
|
||||
c4 := &model.PushChannel{Name: "custom_channel", URL: "https://hook.com", Other: "{\"content\":\"$content\"}"}
|
||||
assert.NoError(t, c4.Validate())
|
||||
|
||||
// 飞书渠道校验:非 HTTPS 地址报错
|
||||
c5 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "http://open.feishu.cn", Other: ""}
|
||||
assert.Error(t, c5.Validate())
|
||||
|
||||
// 飞书正确配置
|
||||
c6 := &model.PushChannel{Name: "lark_channel", Type: "lark", URL: "https://open.feishu.cn", Other: ""}
|
||||
assert.NoError(t, c6.Validate())
|
||||
|
||||
// Telegram 渠道校验
|
||||
cTelegramErr := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "https://api.telegram.org", Token: "", Other: ""}
|
||||
assert.Error(t, cTelegramErr.Validate())
|
||||
|
||||
cTelegramErr2 := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "http://api.telegram.org", Token: "123:abc", Other: ""}
|
||||
assert.Error(t, cTelegramErr2.Validate())
|
||||
|
||||
cTelegramOk := &model.PushChannel{Name: "tg_channel", Type: "telegram", URL: "", Token: "123:abc", Other: "-100123"}
|
||||
assert.NoError(t, cTelegramOk.Validate())
|
||||
assert.Equal(t, "https://api.telegram.org", cTelegramOk.URL)
|
||||
|
||||
// 邮件配置校验:允许空配置以复用系统全局设置
|
||||
c7 := &model.PushChannel{Name: "email_channel", Type: "email", URL: "", Token: "", Other: ""}
|
||||
assert.NoError(t, c7.Validate())
|
||||
|
||||
// 邮件正确配置 (非 HTTPS 协议 URL 允许)
|
||||
c8 := &model.PushChannel{Name: "email_channel", Type: "email", URL: "smtp.exmail.qq.com:465", Token: "user@example.com", Other: "authcode"}
|
||||
assert.NoError(t, c8.Validate())
|
||||
})
|
||||
|
||||
// 2. HTTP CRUD & 触发鉴权测试
|
||||
dbConn, _, cleanup := setupPushTest(t)
|
||||
defer cleanup()
|
||||
|
||||
// 构建路由以进行 HTTP 模拟请求
|
||||
r := testhelper.NewTestGinEngine()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
{
|
||||
adminGroup.GET("/push/channels", ListChannels)
|
||||
adminGroup.POST("/push/channels", CreateChannel)
|
||||
adminGroup.PUT("/push/channels/:id", UpdateChannel)
|
||||
adminGroup.DELETE("/push/channels/:id", DeleteChannel)
|
||||
adminGroup.POST("/push/channels/test", TestChannel)
|
||||
}
|
||||
|
||||
var createdID uint64
|
||||
|
||||
t.Run("admin create channel", func(t *testing.T) {
|
||||
reqBody := CreateChannelRequest{
|
||||
Name: "my_custom_channel",
|
||||
Description: "My custom channel webhook",
|
||||
Type: "custom",
|
||||
Token: "my_chan_token",
|
||||
URL: "https://webhook.site/test",
|
||||
Other: `{"title": "$title", "body": "$content"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
bodyBytes, _ := json.Marshal(reqBody)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/push/channels", bytes.NewBuffer(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataMap, ok := resp.Data.(map[string]any)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, "my_custom_channel", dataMap["name"])
|
||||
createdID = uint64(dataMap["id"].(float64))
|
||||
})
|
||||
|
||||
t.Run("admin list channels", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/push/channels", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
list, ok := resp.Data.([]any)
|
||||
assert.True(t, ok)
|
||||
assert.Len(t, list, 1)
|
||||
})
|
||||
|
||||
t.Run("admin update channel", func(t *testing.T) {
|
||||
updateReq := UpdateChannelRequest{
|
||||
Description: "Updated remark",
|
||||
Type: "custom",
|
||||
Token: "new_chan_token",
|
||||
URL: "https://webhook.site/updated",
|
||||
Other: `{"text": "$content"}`,
|
||||
Enabled: true,
|
||||
}
|
||||
bodyBytes, _ := json.Marshal(updateReq)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/push/channels/"+strconv.FormatUint(createdID, 10), bytes.NewBuffer(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var updated model.PushChannel
|
||||
dbConn.First(&updated, createdID)
|
||||
assert.Equal(t, "Updated remark", updated.Description)
|
||||
assert.Equal(t, "new_chan_token", updated.Token)
|
||||
assert.Equal(t, `{"text": "$content"}`, updated.Other)
|
||||
})
|
||||
|
||||
t.Run("admin test channel endpoint", func(t *testing.T) {
|
||||
testReq := TestChannelRequest{
|
||||
Name: "my_custom_channel",
|
||||
Target: "test_target",
|
||||
}
|
||||
bodyBytes, _ := json.Marshal(testReq)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/push/channels/test", bytes.NewBuffer(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
})
|
||||
|
||||
t.Run("admin delete channel", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("DELETE", "/api/v1/admin/push/channels/"+strconv.FormatUint(createdID, 10), nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var count int64
|
||||
dbConn.Model(&model.PushChannel{}).Where("id = ?", createdID).Count(&count)
|
||||
assert.Equal(t, int64(0), count)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package push defines push notification HTTP routes.
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
"github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// UpdateEventRequest 更新事件请求参数
|
||||
type UpdateEventRequest struct {
|
||||
Channels []string `json:"channels"`
|
||||
Targets []string `json:"targets"`
|
||||
Template string `json:"template" binding:"required"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// TestPushRequest 测试推送通道请求参数
|
||||
type TestPushRequest struct {
|
||||
Config push.Config `json:"config" binding:"required"`
|
||||
Target string `json:"target"`
|
||||
}
|
||||
|
||||
// SyncEvents automatically registers/updates built-in events in the database.
|
||||
func SyncEvents(ctx context.Context) error {
|
||||
return syncBuiltInEvents(ctx)
|
||||
}
|
||||
|
||||
// ListEvents 获取通知事件列表
|
||||
// @Summary 获取所有通知事件
|
||||
// @Description 返回系统配置的通知事件列表,包括预置和自定义事件,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.PushEvent} "通知事件列表"
|
||||
// @Router /api/v1/admin/push/events [get]
|
||||
func ListEvents(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
events, err := listPushEvents(ctx)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(events))
|
||||
}
|
||||
|
||||
// CreateEventRequest 创建事件请求参数
|
||||
type CreateEventRequest struct {
|
||||
EventKey string `json:"event_key"`
|
||||
TaskType string `json:"task_type"` // 关联的异步任务类型
|
||||
Channels []string `json:"channels"`
|
||||
Targets []string `json:"targets"`
|
||||
Template string `json:"template"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
func findBuiltInEvent(key string) (EventMetadata, bool) {
|
||||
for _, meta := range BuiltInEvents {
|
||||
if meta.Key == key {
|
||||
return meta, true
|
||||
}
|
||||
}
|
||||
return EventMetadata{}, false
|
||||
}
|
||||
|
||||
// ListBuiltInEvents 获取内置通知事件列表
|
||||
// @Summary 获取所有内置通知事件
|
||||
// @Description 返回系统定义的所有内置通知事件元数据,供前端下拉框选择,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]EventMetadata} "内置通知事件列表"
|
||||
// @Router /api/v1/admin/push/events/builtin [get]
|
||||
func ListBuiltInEvents(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(BuiltInEvents))
|
||||
}
|
||||
|
||||
// CreateEvent 创建通知事件
|
||||
// @Summary 创建通知事件
|
||||
// @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body CreateEventRequest true "创建参数"
|
||||
// @Success 200 {object} response.Any{data=model.PushEvent} "创建成功"
|
||||
// @Router /api/v1/admin/push/events [post]
|
||||
func CreateEvent(c *gin.Context) {
|
||||
var req CreateEventRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
event, err := createPushEvent(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(event))
|
||||
}
|
||||
|
||||
// DeleteEvent 删除通知事件配置
|
||||
// @Summary 删除通知事件配置
|
||||
// @Description 删除数据库中的特定通知事件配置,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "事件 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||
// @Router /api/v1/admin/push/events/{id} [delete]
|
||||
func DeleteEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid event id")
|
||||
return
|
||||
}
|
||||
|
||||
if err := deletePushEvent(c.Request.Context(), id); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "notification event not found")
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// UpdateEvent 更新通知事件
|
||||
// @Summary 更新通知事件
|
||||
// @Description 更新已有通知事件的推送渠道、接收目标和内容模板,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "事件 ID"
|
||||
// @Param request body push.UpdateEventRequest true "更新参数"
|
||||
// @Success 200 {object} response.Any{data=string} "修改成功"
|
||||
// @Router /api/v1/admin/push/events/{id} [put]
|
||||
func UpdateEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid event id")
|
||||
return
|
||||
}
|
||||
|
||||
var req UpdateEventRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := updatePushEvent(c.Request.Context(), id, req); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "notification event not found")
|
||||
return
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// ToggleEvent 快捷切换通知事件启用状态
|
||||
// @Summary 快捷切换通知事件启用状态
|
||||
// @Description 启用或禁用指定的通知事件
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "事件 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "切换成功"
|
||||
// @Router /api/v1/admin/push/events/{id}/toggle [post]
|
||||
func ToggleEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "invalid event id")
|
||||
return
|
||||
}
|
||||
|
||||
enabled, err := togglePushEvent(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, "notification event not found")
|
||||
return
|
||||
}
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(enabled))
|
||||
}
|
||||
|
||||
// pushHistoriesResponse 推送历史分页响应
|
||||
//
|
||||
//nolint:unused
|
||||
type pushHistoriesResponse struct {
|
||||
Total int64 `json:"total"`
|
||||
Results []model.PushHistory `json:"results"`
|
||||
}
|
||||
|
||||
// ListHistories 分页获取通知推送历史
|
||||
// @Summary 分页获取通知推送历史
|
||||
// @Description 返回分页的通知历史日志数据,需要管理员权限
|
||||
// @Tags admin-push
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param page query int false "当前页码"
|
||||
// @Param page_size query int false "分页大小"
|
||||
// @Param event_key query string false "过滤事件名称"
|
||||
// @Param status query string false "过滤发送状态"
|
||||
// @Success 200 {object} response.Any{data=pushHistoriesResponse} "推送历史列表"
|
||||
// @Router /api/v1/admin/push/histories [get]
|
||||
func ListHistories(c *gin.Context) {
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 {
|
||||
pageSize = 20
|
||||
}
|
||||
|
||||
total, results, err := listPushHistories(c.Request.Context(), repository.PushHistoryListFilter{
|
||||
EventKey: c.Query("event_key"),
|
||||
Status: c.Query("status"),
|
||||
Page: page,
|
||||
PageSize: pageSize,
|
||||
})
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(map[string]any{
|
||||
"total": total,
|
||||
"results": results,
|
||||
}))
|
||||
}
|
||||
|
||||
// TestPush 测试推送通道发送
|
||||
// @Summary 测试推送通道发送
|
||||
// @Description 接收临时通知渠道配置并在本地同步调用 Pusher.Send 发送测试消息
|
||||
// @Tags admin-push
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body push.TestPushRequest true "测试请求体"
|
||||
// @Success 200 {object} response.Any{data=string} "测试成功"
|
||||
// @Router /api/v1/admin/push/test [post]
|
||||
func TestPush(c *gin.Context) {
|
||||
var req TestPushRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
pusher, err := push.GetPusher(req.Config.Channel)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := pusher.ValidateConfig(req.Config); err != nil {
|
||||
response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
applySMTPFallbackToPushConfig(c.Request.Context(), &req.Config)
|
||||
|
||||
testBody := map[string]any{
|
||||
keyTitle: "测试通道推送",
|
||||
keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。",
|
||||
keyLevel: defaultLevelInfo,
|
||||
}
|
||||
if err := pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
)
|
||||
|
||||
// RegisterTaskListeners subscribes push notification handlers to task completion events.
|
||||
func RegisterTaskListeners() {
|
||||
task.OnTaskCompleted(handleTaskCompleted)
|
||||
}
|
||||
|
||||
func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.TaskResult, execErr error) {
|
||||
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
|
||||
return
|
||||
}
|
||||
if len(events) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
body := map[string]any{
|
||||
"task_id": execution.TaskID,
|
||||
"task_name": execution.TaskName,
|
||||
"task_type": execution.TaskType,
|
||||
"task_status": string(execution.Status),
|
||||
"task_duration": execution.Duration,
|
||||
"time": time.Now().Format("2006-01-02 15:04:05"),
|
||||
}
|
||||
if execErr != nil {
|
||||
body["task_error"] = execErr.Error()
|
||||
} else {
|
||||
body["task_error"] = ""
|
||||
}
|
||||
if result != nil {
|
||||
body["task_result"] = result.Message
|
||||
} else {
|
||||
body["task_result"] = ""
|
||||
}
|
||||
|
||||
var payloadMap map[string]any
|
||||
if execution.Payload != "" {
|
||||
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
|
||||
body["payload"] = payloadMap
|
||||
extractUserFromMap(ctx, payloadMap, body)
|
||||
}
|
||||
}
|
||||
if result != nil && result.Detail != "" {
|
||||
var detailMap map[string]any
|
||||
if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil {
|
||||
body["detail"] = detailMap
|
||||
extractUserFromMap(ctx, detailMap, body)
|
||||
}
|
||||
}
|
||||
|
||||
for _, event := range events {
|
||||
meta := EventMetadata{
|
||||
Key: event.EventKey,
|
||||
Name: event.Name,
|
||||
Description: "异步任务执行完毕触发的自动通知",
|
||||
}
|
||||
DefaultTrigger.Trigger(ctx, meta, body)
|
||||
}
|
||||
}
|
||||
|
||||
func extractUserFromMap(ctx context.Context, data map[string]any, body map[string]any) {
|
||||
if u, exists := body["user"]; exists && u != nil {
|
||||
return
|
||||
}
|
||||
if user := loadUserFromPayload(ctx, data); user != nil {
|
||||
body["user"] = user
|
||||
}
|
||||
}
|
||||
|
||||
func extractUserID(data map[string]any) (uint64, bool) {
|
||||
for _, k := range []string{"user_id", "userId", "uid"} {
|
||||
val, ok := data[k]
|
||||
if !ok || val == nil {
|
||||
continue
|
||||
}
|
||||
switch v := val.(type) {
|
||||
case float64:
|
||||
if v >= 0 {
|
||||
return uint64(v), true
|
||||
}
|
||||
case int:
|
||||
if v >= 0 {
|
||||
return uint64(v), true
|
||||
}
|
||||
case int64:
|
||||
if v >= 0 {
|
||||
return uint64(v), true
|
||||
}
|
||||
case uint64:
|
||||
return v, true
|
||||
case string:
|
||||
if id, err := strconv.ParseUint(v, 10, 64); err == nil {
|
||||
return id, true
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func extractUsername(data map[string]any) string {
|
||||
for _, k := range []string{"username", "user_name"} {
|
||||
if val, ok := data[k]; ok && val != nil {
|
||||
if s, ok := val.(string); ok && s != "" {
|
||||
return s
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package push defines push notification HTTP routes and background tasks.
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/push"
|
||||
)
|
||||
|
||||
const (
|
||||
// SendNotificationTask 发送推送通知任务标识
|
||||
SendNotificationTask = "push:send"
|
||||
// TaskTypeSendNotification 推送通知管理类型
|
||||
TaskTypeSendNotification = "send_notification"
|
||||
)
|
||||
|
||||
// SendNotificationMeta represents the task metadata.
|
||||
var SendNotificationMeta = task.TaskMeta{
|
||||
Type: TaskTypeSendNotification,
|
||||
AsynqTask: SendNotificationTask,
|
||||
Name: "推送通知",
|
||||
Description: "异步执行系统通知的多渠道派发与推送",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
{
|
||||
Name: "event_key",
|
||||
Label: "事件标识",
|
||||
Type: "string",
|
||||
Required: true,
|
||||
Placeholder: "admin_login",
|
||||
},
|
||||
{
|
||||
Name: "target",
|
||||
Label: "目标接收者",
|
||||
Type: "string",
|
||||
Required: false,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// PushHandler 通知推送异步任务处理器
|
||||
//
|
||||
//nolint:revive
|
||||
type PushHandler struct{}
|
||||
|
||||
// ValidatePayload 校验并标准化推送参数
|
||||
func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
if len(payload) == 0 {
|
||||
return nil, errors.New("payload is required")
|
||||
}
|
||||
|
||||
var req SendPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
return nil, fmt.Errorf("invalid json format: %w", err)
|
||||
}
|
||||
|
||||
if req.Config.Channel == "" {
|
||||
return nil, errors.New("channel type is required")
|
||||
}
|
||||
|
||||
return json.Marshal(req)
|
||||
}
|
||||
|
||||
// Execute 异步执行推送操作并记录推送历史审计
|
||||
func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
var req SendPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
task.AppendLog(ctx, "解析推送参数失败: %v", err)
|
||||
return nil, fmt.Errorf("parse payload failed: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
|
||||
|
||||
pusher, err := push.GetPusher(req.Config.Channel)
|
||||
if err != nil {
|
||||
errWrap := fmt.Errorf("get pusher failed: %w", err)
|
||||
task.AppendLog(ctx, "推送失败: %v", errWrap)
|
||||
if task.IsFinalAttempt(ctx) {
|
||||
h.recordHistory(ctx, req, "failed", errWrap.Error())
|
||||
}
|
||||
return nil, errWrap
|
||||
}
|
||||
|
||||
// 执行真正的消息推送,扁平化为原始 json 格式
|
||||
flatBody := req.Body.Flatten()
|
||||
err = pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil)
|
||||
|
||||
title := req.Body.Title
|
||||
content := req.Body.Content
|
||||
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
|
||||
if task.IsFinalAttempt(ctx) {
|
||||
h.recordHistory(ctx, req, "failed", err.Error())
|
||||
}
|
||||
return nil, fmt.Errorf("pusher.Send failed: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
|
||||
h.recordHistory(ctx, req, "success", "")
|
||||
|
||||
return &task.TaskResult{
|
||||
Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
|
||||
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
|
||||
task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,355 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package status 提供系统状态查询接口
|
||||
package status
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"math"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// startTime 记录服务启动时间
|
||||
var startTime = time.Now()
|
||||
|
||||
const (
|
||||
hoursInDay = 24
|
||||
minutesInHour = 60
|
||||
secondsInMinute = 60
|
||||
nanosPerSecond = 1e9
|
||||
binaryKB = 0
|
||||
binaryMB = 1
|
||||
binaryGB = 2
|
||||
valueThreshold = 10 // 格式化时区分整数显示的阈值
|
||||
)
|
||||
|
||||
// SystemStatusResponse 系统状态响应结构体
|
||||
type SystemStatusResponse struct {
|
||||
Uptime string `json:"uptime"`
|
||||
NumGoroutine int `json:"num_goroutine"`
|
||||
Alloc string `json:"alloc"`
|
||||
TotalAlloc string `json:"total_alloc"`
|
||||
Sys string `json:"sys"`
|
||||
Lookups uint64 `json:"lookups"`
|
||||
Mallocs uint64 `json:"mallocs"`
|
||||
Frees uint64 `json:"frees"`
|
||||
HeapAlloc string `json:"heap_alloc"`
|
||||
HeapSys string `json:"heap_sys"`
|
||||
HeapIdle string `json:"heap_idle"`
|
||||
HeapInuse string `json:"heap_inuse"`
|
||||
HeapReleased string `json:"heap_released"`
|
||||
HeapObjects uint64 `json:"heap_objects"`
|
||||
StackInuse string `json:"stack_inuse"`
|
||||
StackSys string `json:"stack_sys"`
|
||||
MSpanInuse string `json:"mspan_inuse"`
|
||||
MSpanSys string `json:"mspan_sys"`
|
||||
MCacheInuse string `json:"mcache_inuse"`
|
||||
MCacheSys string `json:"mcache_sys"`
|
||||
BuckHashSys string `json:"buck_hash_sys"`
|
||||
GCSys string `json:"gc_sys"`
|
||||
OtherSys string `json:"other_sys"`
|
||||
NextGC string `json:"next_gc"`
|
||||
LastGCTime string `json:"last_gc_time"`
|
||||
PauseTotalNs string `json:"pause_total_ns"`
|
||||
LastPause string `json:"last_pause"`
|
||||
NumGC uint32 `json:"num_gc"`
|
||||
}
|
||||
|
||||
// formatBytes 格式化字节大小
|
||||
func formatBytes(bytes uint64) string {
|
||||
const unit = 1024
|
||||
if bytes < unit {
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
div, exp := int64(unit), 0
|
||||
for n := bytes / unit; n >= unit; n /= unit {
|
||||
div *= unit
|
||||
exp++
|
||||
}
|
||||
value := float64(bytes) / float64(div)
|
||||
var suffix string
|
||||
switch exp {
|
||||
case binaryKB:
|
||||
suffix = "KiB"
|
||||
case binaryMB:
|
||||
suffix = "MiB"
|
||||
case binaryGB:
|
||||
suffix = "GiB"
|
||||
default:
|
||||
suffix = "TiB"
|
||||
}
|
||||
|
||||
// 格式化规则:
|
||||
// - 如果是整数(如 16, 73, 105, 986, 112):
|
||||
// - 如果 >= 10,则格式化为 "%.0f" (e.g. "16 KiB")
|
||||
// - 如果 < 10,则格式化为 "%.1f" (e.g. "9.0 KiB")
|
||||
// - 如果不是整数(如 5.8, 9.1, 7.6, 4.8):格式化为 "%.1f"
|
||||
if value == math.Trunc(value) {
|
||||
if value >= valueThreshold {
|
||||
return fmt.Sprintf("%.0f %s", value, suffix)
|
||||
}
|
||||
return fmt.Sprintf("%.1f %s", value, suffix)
|
||||
}
|
||||
return fmt.Sprintf("%.1f %s", value, suffix)
|
||||
}
|
||||
|
||||
// formatDuration 格式化时间持续时间
|
||||
func formatDuration(d time.Duration) string {
|
||||
days := int(d.Hours()) / hoursInDay
|
||||
hours := int(d.Hours()) % hoursInDay
|
||||
minutes := int(d.Minutes()) % minutesInHour
|
||||
seconds := int(d.Seconds()) % secondsInMinute
|
||||
|
||||
var res string
|
||||
if days > 0 {
|
||||
res += fmt.Sprintf("%d天", days)
|
||||
}
|
||||
if hours > 0 {
|
||||
res += fmt.Sprintf("%d小时", hours)
|
||||
}
|
||||
if minutes > 0 {
|
||||
res += fmt.Sprintf("%d分钟", minutes)
|
||||
}
|
||||
if seconds > 0 || res == "" {
|
||||
res += fmt.Sprintf("%d秒钟", seconds)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
// GetSystemStatus 获取系统状态信息
|
||||
// @Summary 获取系统状态信息
|
||||
// @Description 获取后端服务运行状态、Goroutine、内存指标等详细统计数据,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=status.SystemStatusResponse} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/status [get]
|
||||
func GetSystemStatus(c *gin.Context) {
|
||||
var m runtime.MemStats
|
||||
runtime.ReadMemStats(&m)
|
||||
|
||||
uptime := formatDuration(time.Since(startTime))
|
||||
numGoroutine := runtime.NumGoroutine()
|
||||
|
||||
var lastGCTime string
|
||||
switch {
|
||||
case m.LastGC > 0 && m.LastGC <= math.MaxInt64:
|
||||
lastGCTime = formatDuration(time.Since(time.Unix(0, int64(m.LastGC))))
|
||||
case m.LastGC > 0:
|
||||
lastGCTime = "未知"
|
||||
default:
|
||||
lastGCTime = "无"
|
||||
}
|
||||
|
||||
var lastPause string
|
||||
if m.NumGC > 0 {
|
||||
lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond)
|
||||
} else {
|
||||
lastPause = "0.000s"
|
||||
}
|
||||
|
||||
res := SystemStatusResponse{
|
||||
Uptime: uptime,
|
||||
NumGoroutine: numGoroutine,
|
||||
Alloc: formatBytes(m.Alloc),
|
||||
TotalAlloc: formatBytes(m.TotalAlloc),
|
||||
Sys: formatBytes(m.Sys),
|
||||
Lookups: m.Lookups,
|
||||
Mallocs: m.Mallocs,
|
||||
Frees: m.Frees,
|
||||
HeapAlloc: formatBytes(m.HeapAlloc),
|
||||
HeapSys: formatBytes(m.HeapSys),
|
||||
HeapIdle: formatBytes(m.HeapIdle),
|
||||
HeapInuse: formatBytes(m.HeapInuse),
|
||||
HeapReleased: formatBytes(m.HeapReleased),
|
||||
HeapObjects: m.HeapObjects,
|
||||
StackInuse: formatBytes(m.StackInuse),
|
||||
StackSys: formatBytes(m.StackSys),
|
||||
MSpanInuse: formatBytes(m.MSpanInuse),
|
||||
MSpanSys: formatBytes(m.MSpanSys),
|
||||
MCacheInuse: formatBytes(m.MCacheInuse),
|
||||
MCacheSys: formatBytes(m.MCacheSys),
|
||||
BuckHashSys: formatBytes(m.BuckHashSys),
|
||||
GCSys: formatBytes(m.GCSys),
|
||||
OtherSys: formatBytes(m.OtherSys),
|
||||
NextGC: formatBytes(m.NextGC),
|
||||
LastGCTime: lastGCTime,
|
||||
PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond),
|
||||
LastPause: lastPause,
|
||||
NumGC: m.NumGC,
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(res))
|
||||
}
|
||||
|
||||
// DatabaseInfoResponse 数据库信息响应结构体
|
||||
type DatabaseInfoResponse struct {
|
||||
Type string `json:"type"`
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
// getSQLiteInfo 返回 SQLite 数据库信息
|
||||
func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
|
||||
info := DatabaseInfoResponse{
|
||||
Type: "sqlite",
|
||||
Name: config.Config.Database.SQLitePath,
|
||||
Version: "SQLite",
|
||||
}
|
||||
if info.Name == "" {
|
||||
info.Name = "./data/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), ¤tCfg); 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), ¤tCfg); 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)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package task 提供任务管理接口
|
||||
package task
|
||||
|
||||
// 任务管理相关错误消息
|
||||
const (
|
||||
InvalidTaskType = "无效的任务类型"
|
||||
InvalidTimeRange = "无效的时间范围"
|
||||
TaskDispatchFailed = "任务下发失败"
|
||||
UserIDRequired = "用户ID必填"
|
||||
TaskNotFound = "任务执行记录不存在"
|
||||
TaskNotRetryable = "该任务不支持重试"
|
||||
TaskNotFailed = "只有失败的任务才能重试"
|
||||
TaskMaxRetryExceeded = "已达到最大重试次数"
|
||||
TaskRetryFailed = "任务重试失败"
|
||||
InvalidCronExpression = "无效的 Cron 表达式"
|
||||
ScheduleNotFound = "定时任务不存在"
|
||||
ScheduleSaveFailed = "保存定时任务失败"
|
||||
ScheduleDeleteFailed = "删除定时任务失败"
|
||||
)
|
||||
@@ -0,0 +1,416 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package task
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/task/scheduler"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/robfig/cron/v3"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// ListTaskTypes 获取支持的任务类型列表
|
||||
// @Summary 获取支持的任务类型
|
||||
// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]task.TaskMeta} "任务类型列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/types [get]
|
||||
func ListTaskTypes(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(task.GetDispatchableTasks()))
|
||||
}
|
||||
|
||||
// DispatchTaskRequest 下发任务请求
|
||||
type DispatchTaskRequest struct {
|
||||
TaskType string `json:"task_type" binding:"required"`
|
||||
StartTime *time.Time `json:"start_time"`
|
||||
EndTime *time.Time `json:"end_time"`
|
||||
UserID *uint64 `json:"user_id"`
|
||||
Payload string `json:"payload"`
|
||||
}
|
||||
|
||||
// DispatchTask 下发任务
|
||||
// @Summary 下发异步任务
|
||||
// @Description 手动触发指定类型的异步任务,支持指定时间范围和用户,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body DispatchTaskRequest true "任务请求参数"
|
||||
// @Success 200 {object} response.Any{data=string} "任务已入队"
|
||||
// @Failure 400 {object} response.Any "任务类型不存在或参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "任务入队失败"
|
||||
// @Router /api/v1/admin/tasks/dispatch [post]
|
||||
func DispatchTask(c *gin.Context) {
|
||||
var req DispatchTaskRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
meta := task.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
|
||||
var payloadBytes []byte
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
|
||||
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
taskID, err := task.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
|
||||
if err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(taskID))
|
||||
}
|
||||
|
||||
// ListTaskExecutions 查询任务执行记录列表
|
||||
// @Summary 查询任务执行记录
|
||||
// @Description 分页查询任务执行记录,支持按状态和任务类型筛选,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param status query string false "状态筛选 (pending/running/succeeded/failed)"
|
||||
// @Param task_type query string false "任务类型筛选"
|
||||
// @Param page query int false "页码" default(1)
|
||||
// @Param page_size query int false "每页条数" default(20)
|
||||
// @Success 200 {object} response.Any{data=object} "任务执行记录列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/executions [get]
|
||||
func ListTaskExecutions(c *gin.Context) {
|
||||
var req model.ListTaskExecutionsRequest
|
||||
if err := c.ShouldBindQuery(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.TaskType != "" {
|
||||
if meta := task.GetTaskMeta(req.TaskType); meta != nil {
|
||||
req.TaskType = meta.AsynqTask
|
||||
}
|
||||
}
|
||||
|
||||
executions, total, err := model.ListTaskExecutions(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{
|
||||
"items": executions,
|
||||
"total": total,
|
||||
"page": req.Page,
|
||||
"page_size": req.PageSize,
|
||||
}))
|
||||
}
|
||||
|
||||
// GetTaskExecution 查询单条任务执行详情
|
||||
// @Summary 查询任务执行详情
|
||||
// @Description 根据 ID 查询任务执行记录详情,包含完整执行日志,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "任务执行记录 ID"
|
||||
// @Success 200 {object} response.Any{data=model.TaskExecution} "任务执行详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Router /api/v1/admin/tasks/executions/{id} [get]
|
||||
func GetTaskExecution(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, admin.InvalidTaskExecutionID)
|
||||
return
|
||||
}
|
||||
|
||||
execution, err := model.GetTaskExecutionByID(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
response.AbortNotFound(c, TaskNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(execution))
|
||||
}
|
||||
|
||||
// RetryTask 重试失败的任务
|
||||
// @Summary 重试失败任务
|
||||
// @Description 重新下发一条失败的任务,创建新的执行记录,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "任务执行记录 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "新任务的 TaskID"
|
||||
// @Failure 400 {object} response.Any "任务不支持重试或参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "重试失败"
|
||||
// @Router /api/v1/admin/tasks/executions/{id}/retry [post]
|
||||
func RetryTask(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, admin.InvalidTaskExecutionID)
|
||||
return
|
||||
}
|
||||
|
||||
newTaskID, err := task.RetryTask(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
errMsg := err.Error()
|
||||
switch {
|
||||
case strings.Contains(errMsg, "不存在"):
|
||||
response.AbortNotFound(c, errMsg)
|
||||
case strings.Contains(errMsg, "只有失败的任务") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"):
|
||||
response.AbortBadRequest(c, errMsg)
|
||||
default:
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(newTaskID))
|
||||
}
|
||||
|
||||
// ListSchedules 获取定时任务列表
|
||||
// @Summary 获取定时任务列表
|
||||
// @Description 返回系统所有的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.Schedule} "定时任务列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/schedules [get]
|
||||
func ListSchedules(c *gin.Context) {
|
||||
schedules, err := model.ListSchedules(c.Request.Context())
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(schedules))
|
||||
}
|
||||
|
||||
// CreateScheduleRequest 创建定时任务请求
|
||||
type CreateScheduleRequest struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
TaskType string `json:"task_type" binding:"required"`
|
||||
Cron string `json:"cron" binding:"required"`
|
||||
Payload string `json:"payload"`
|
||||
IsActive *bool `json:"is_active" binding:"required"`
|
||||
}
|
||||
|
||||
// CreateSchedule 创建定时任务
|
||||
// @Summary 创建定时任务
|
||||
// @Description 新增一个动态定时任务配置,关联已有的异步任务,配置 Cron 表达式和执行参数,并触发调度器热加载,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body CreateScheduleRequest true "创建定时任务请求参数"
|
||||
// @Success 200 {object} response.Any{data=model.Schedule} "创建成功的定时任务信息"
|
||||
// @Failure 400 {object} response.Any "Cron 表达式无效、异步任务类型不存在或参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "保存定时任务失败"
|
||||
// @Router /api/v1/admin/tasks/schedules [post]
|
||||
func CreateSchedule(c *gin.Context) {
|
||||
var req CreateScheduleRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 校验 Cron 表达式
|
||||
if _, err := cron.ParseStandard(req.Cron); err != nil {
|
||||
response.AbortBadRequest(c, InvalidCronExpression)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := task.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验并规范化 Payload
|
||||
var payloadBytes []byte
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
schedule := &model.Schedule{
|
||||
Name: req.Name,
|
||||
TaskType: req.TaskType,
|
||||
Cron: req.Cron,
|
||||
Payload: string(validated),
|
||||
IsActive: *req.IsActive,
|
||||
}
|
||||
|
||||
if err := model.CreateSchedule(c.Request.Context(), schedule); err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := scheduler.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(schedule))
|
||||
}
|
||||
|
||||
// UpdateScheduleRequest 修改定时任务请求
|
||||
type UpdateScheduleRequest struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
TaskType string `json:"task_type" binding:"required"`
|
||||
Cron string `json:"cron" binding:"required"`
|
||||
Payload string `json:"payload"`
|
||||
IsActive *bool `json:"is_active" binding:"required"`
|
||||
}
|
||||
|
||||
// UpdateSchedule 修改定时任务
|
||||
// @Summary 修改定时任务
|
||||
// @Description 修改一个定时任务的配置(名称、Cron 表达式、异步任务参数和是否启用等),并触发调度器热加载,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "定时任务 ID"
|
||||
// @Param request body UpdateScheduleRequest true "修改定时任务请求参数"
|
||||
// @Success 200 {object} response.Any{data=model.Schedule} "修改后的定时任务信息"
|
||||
// @Failure 400 {object} response.Any "Cron 表达式无效、参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "定时任务不存在"
|
||||
// @Failure 500 {object} response.Any "修改定时任务失败"
|
||||
// @Router /api/v1/admin/tasks/schedules/{id} [put]
|
||||
func UpdateSchedule(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "无效的定时任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req UpdateScheduleRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 检查定时任务是否存在
|
||||
schedule, err := model.GetScheduleByID(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
response.AbortNotFound(c, ScheduleNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验 Cron 表达式
|
||||
if _, err := cron.ParseStandard(req.Cron); err != nil {
|
||||
response.AbortBadRequest(c, InvalidCronExpression)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := task.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验并规范化 Payload
|
||||
var payloadBytes []byte
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
schedule.Name = req.Name
|
||||
schedule.TaskType = req.TaskType
|
||||
schedule.Cron = req.Cron
|
||||
schedule.Payload = string(validated)
|
||||
schedule.IsActive = *req.IsActive
|
||||
|
||||
if err := model.UpdateSchedule(c.Request.Context(), schedule); err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := scheduler.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(schedule))
|
||||
}
|
||||
|
||||
// DeleteSchedule 删除定时任务
|
||||
// @Summary 删除定时任务
|
||||
// @Description 删除指定的定时任务配置,并触发调度器热加载,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "定时任务 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "删除结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "删除定时任务失败"
|
||||
// @Router /api/v1/admin/tasks/schedules/{id} [delete]
|
||||
func DeleteSchedule(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "无效的定时任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := model.DeleteSchedule(c.Request.Context(), id); err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := scheduler.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,495 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package task
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/user"
|
||||
"github.com/Rain-kl/Wavelet/internal/bootstrap"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/task"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
func setupTaskTestEnvironment(t *testing.T) func() {
|
||||
_, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
bootstrap.RegisterTasks()
|
||||
task.AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
|
||||
Addr: mr.Addr(),
|
||||
})
|
||||
return func() {
|
||||
if task.AsynqClient != nil {
|
||||
_ = task.AsynqClient.Close()
|
||||
task.AsynqClient = nil
|
||||
}
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
r := testhelper.NewTestGinEngine()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
|
||||
// Mock authentication middleware
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.GET("/tasks/types", ListTaskTypes)
|
||||
adminGroup.POST("/tasks/dispatch", DispatchTask)
|
||||
adminGroup.GET("/tasks/executions", ListTaskExecutions)
|
||||
adminGroup.GET("/tasks/executions/:id", GetTaskExecution)
|
||||
adminGroup.POST("/tasks/executions/:id/retry", RetryTask)
|
||||
return r
|
||||
}
|
||||
|
||||
func TestListTaskTypes(t *testing.T) {
|
||||
cleanup := setupTaskTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/types", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
var resp response.Any
|
||||
_ = json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var taskMetas []task.TaskMeta
|
||||
_ = json.Unmarshal(dataBytes, &taskMetas)
|
||||
|
||||
if len(taskMetas) == 0 {
|
||||
t.Error("expected at least one dispatchable task type")
|
||||
}
|
||||
|
||||
foundCleanup := false
|
||||
foundWarmImageCache := false
|
||||
for _, m := range taskMetas {
|
||||
if m.Type == uploadtask.TaskTypeSystemCleanup {
|
||||
foundCleanup = true
|
||||
}
|
||||
if m.Type == uploadtask.TaskTypeWarmImageCache {
|
||||
foundWarmImageCache = true
|
||||
}
|
||||
}
|
||||
if !foundCleanup {
|
||||
t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeSystemCleanup)
|
||||
}
|
||||
if !foundWarmImageCache {
|
||||
t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeWarmImageCache)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDispatchTask(t *testing.T) {
|
||||
cleanup := setupTaskTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("dispatch valid task successfully", func(t *testing.T) {
|
||||
payload := DispatchTaskRequest{
|
||||
TaskType: uploadtask.TaskTypeSystemCleanup,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code, "Body: %s", w.Body.String())
|
||||
|
||||
var resp response.Any
|
||||
_ = json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Empty(t, resp.ErrorMsg)
|
||||
assert.NotNil(t, resp.Data)
|
||||
|
||||
// 返回的 data 应该是 taskID
|
||||
taskID, ok := resp.Data.(string)
|
||||
assert.True(t, ok)
|
||||
assert.NotEmpty(t, taskID)
|
||||
})
|
||||
|
||||
t.Run("dispatch send_email task successfully with valid payload", func(t *testing.T) {
|
||||
payload := DispatchTaskRequest{
|
||||
TaskType: user.TaskTypeSendEmail,
|
||||
Payload: `{"to":"receiver@example.com","subject":"Test Subject","body":"Test Body"}`,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code, "Body: %s", w.Body.String())
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Empty(t, resp.ErrorMsg)
|
||||
assert.NotNil(t, resp.Data)
|
||||
})
|
||||
|
||||
t.Run("dispatch send_email task failure with invalid payload json", func(t *testing.T) {
|
||||
payload := DispatchTaskRequest{
|
||||
TaskType: user.TaskTypeSendEmail,
|
||||
Payload: `{"to":`,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Contains(t, resp.ErrorMsg, "无效的 JSON 格式")
|
||||
})
|
||||
|
||||
t.Run("dispatch send_email task failure with missing fields", func(t *testing.T) {
|
||||
payload := DispatchTaskRequest{
|
||||
TaskType: user.TaskTypeSendEmail,
|
||||
Payload: `{"to":"","subject":"Test","body":"Test"}`,
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Contains(t, resp.ErrorMsg, "不能为空")
|
||||
})
|
||||
|
||||
t.Run("dispatch invalid task type failure", func(t *testing.T) {
|
||||
payload := DispatchTaskRequest{
|
||||
TaskType: "invalid_task_type",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Equal(t, InvalidTaskType, resp.ErrorMsg)
|
||||
})
|
||||
|
||||
t.Run("dispatch with empty body failure", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer([]byte("{}")))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func TestListTaskExecutions(t *testing.T) {
|
||||
cleanup := setupTaskTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
ctx := context.Background()
|
||||
|
||||
// 准备测试数据
|
||||
now := time.Now()
|
||||
records := []*model.TaskExecution{
|
||||
{TaskID: "exec_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual", Duration: 1500, Result: "清理完成", StartedAt: &now, FinishedAt: &now},
|
||||
{TaskID: "exec_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system", ErrorMessage: "连接超时", StartedAt: &now, FinishedAt: &now},
|
||||
{TaskID: "exec_003", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", Retryable: true, MaxRetry: 3},
|
||||
}
|
||||
for _, r := range records {
|
||||
err := model.CreateTaskExecution(ctx, r)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
t.Run("list all executions", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
})
|
||||
|
||||
t.Run("filter by status", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?status=failed", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(1), data["total"])
|
||||
})
|
||||
|
||||
t.Run("filter by task_type (asynq task name)", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?task_type=system:cleanup", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
})
|
||||
|
||||
t.Run("filter by task_type (management task type)", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?task_type=system_cleanup", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
})
|
||||
|
||||
t.Run("pagination", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions?page=1&page_size=2", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var data map[string]interface{}
|
||||
json.Unmarshal(dataBytes, &data)
|
||||
|
||||
assert.Equal(t, float64(3), data["total"])
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetTaskExecution(t *testing.T) {
|
||||
cleanup := setupTaskTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建测试记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "detail_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "系统垃圾清理",
|
||||
Status: model.TaskExecutionStatusSucceeded,
|
||||
Log: "[10:00:01] 开始扫描\n[10:00:02] 找到 50 个文件\n[10:00:03] 清理完成",
|
||||
Result: "共清理 50 个文件",
|
||||
Duration: 2000,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("get existing execution", func(t *testing.T) {
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d", execution.ID)
|
||||
req, _ := http.NewRequest("GET", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var detail model.TaskExecution
|
||||
json.Unmarshal(dataBytes, &detail)
|
||||
|
||||
assert.Equal(t, "detail_001", detail.TaskID)
|
||||
assert.Equal(t, model.TaskExecutionStatusSucceeded, detail.Status)
|
||||
assert.Contains(t, detail.Log, "开始扫描")
|
||||
assert.Contains(t, detail.Log, "清理完成")
|
||||
assert.Equal(t, int64(2000), detail.Duration)
|
||||
})
|
||||
|
||||
t.Run("get non-existent execution", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/99999999", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, w.Code)
|
||||
})
|
||||
|
||||
t.Run("invalid ID format", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/tasks/executions/invalid", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRetryTask(t *testing.T) {
|
||||
cleanup := setupTaskTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("retry failed task successfully", func(t *testing.T) {
|
||||
now := time.Now()
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_api_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "系统垃圾清理",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
ErrorMessage: "S3 连接超时",
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
RetryCount: 0,
|
||||
TriggeredBy: "manual",
|
||||
StartedAt: &now,
|
||||
FinishedAt: &now,
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
|
||||
req, _ := http.NewRequest("POST", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp response.Any
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.Empty(t, resp.ErrorMsg)
|
||||
assert.NotNil(t, resp.Data)
|
||||
|
||||
// 验证新记录
|
||||
newTaskID, ok := resp.Data.(string)
|
||||
assert.True(t, ok)
|
||||
assert.NotEmpty(t, newTaskID)
|
||||
|
||||
newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 1, newExecution.RetryCount)
|
||||
assert.Equal(t, "retry", newExecution.TriggeredBy)
|
||||
})
|
||||
|
||||
t.Run("retry succeeded task fails", func(t *testing.T) {
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_succeeded_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "系统垃圾清理",
|
||||
Status: model.TaskExecutionStatusSucceeded,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
|
||||
req, _ := http.NewRequest("POST", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
t.Run("retry non-retryable task fails", func(t *testing.T) {
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "retry_not_allowed_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "系统垃圾清理",
|
||||
Status: model.TaskExecutionStatusFailed,
|
||||
Retryable: false,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := model.CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
|
||||
req, _ := http.NewRequest("POST", url, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
t.Run("retry non-existent task", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/99999999/retry", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusNotFound, w.Code)
|
||||
})
|
||||
|
||||
t.Run("retry with invalid ID", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/executions/invalid/retry", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package template 提供模板管理功能
|
||||
package template
|
||||
|
||||
// 模板管理相关错误消息
|
||||
const (
|
||||
TemplateNotFound = "模板不存在"
|
||||
TemplateKeyRequired = "模板标识符不能为空"
|
||||
TemplateNameRequired = "模板名称不能为空"
|
||||
TemplateContentRequired = "模板内容不能为空"
|
||||
TemplateKeyExists = "模板标识符已存在"
|
||||
SystemTemplateCannotDelete = "系统预置模板不可删除"
|
||||
SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符"
|
||||
)
|
||||
@@ -0,0 +1,78 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package template
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
)
|
||||
|
||||
func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Template, error) {
|
||||
exists, err := repository.TemplateExistsByKey(ctx, req.Key)
|
||||
if err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
if exists {
|
||||
return model.Template{}, errors.New(TemplateKeyExists)
|
||||
}
|
||||
|
||||
tmpl := model.Template{
|
||||
Key: req.Key,
|
||||
Name: req.Name,
|
||||
Type: req.Type,
|
||||
Subject: req.Subject,
|
||||
Content: req.Content,
|
||||
Description: req.Description,
|
||||
IsSystem: false,
|
||||
}
|
||||
if err := tmpl.Validate(); err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
if err := repository.CreateTemplate(ctx, &tmpl); err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
func listTemplates(ctx context.Context) ([]model.Template, error) {
|
||||
return repository.ListTemplates(ctx)
|
||||
}
|
||||
|
||||
func getTemplate(ctx context.Context, key string) (model.Template, error) {
|
||||
return repository.GetTemplateByKey(ctx, key)
|
||||
}
|
||||
|
||||
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (model.Template, error) {
|
||||
tmpl, err := repository.GetTemplateByKey(ctx, key)
|
||||
if err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
|
||||
tmpl.Name = req.Name
|
||||
tmpl.Type = req.Type
|
||||
tmpl.Subject = req.Subject
|
||||
tmpl.Content = req.Content
|
||||
tmpl.Description = req.Description
|
||||
if err := tmpl.Validate(); err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
if err := repository.SaveTemplate(ctx, &tmpl); err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
func deleteTemplate(ctx context.Context, key string) error {
|
||||
tmpl, err := repository.GetTemplateByKey(ctx, key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if tmpl.IsSystem {
|
||||
return errors.New(SystemTemplateCannotDelete)
|
||||
}
|
||||
return repository.DeleteTemplate(ctx, &tmpl)
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package template
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// CreateTemplateRequest 创建模板请求
|
||||
type CreateTemplateRequest struct {
|
||||
Key string `json:"key" binding:"required,max=80"`
|
||||
Name string `json:"name" binding:"required,max=100"`
|
||||
Type string `json:"type" binding:"required,max=20"`
|
||||
Subject string `json:"subject" binding:"max=255"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
Description string `json:"description" binding:"max=255"`
|
||||
}
|
||||
|
||||
// UpdateTemplateRequest 更新模板请求
|
||||
type UpdateTemplateRequest struct {
|
||||
Name string `json:"name" binding:"required,max=100"`
|
||||
Type string `json:"type" binding:"required,max=20"`
|
||||
Subject string `json:"subject" binding:"max=255"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
Description string `json:"description" binding:"max=255"`
|
||||
}
|
||||
|
||||
func abortTemplateLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, TemplateNotFound)
|
||||
return true
|
||||
}
|
||||
msg := err.Error()
|
||||
switch msg {
|
||||
case TemplateKeyExists, SystemTemplateCannotDelete:
|
||||
response.AbortBadRequest(c, msg)
|
||||
return true
|
||||
}
|
||||
response.AbortInternal(c, msg)
|
||||
return true
|
||||
}
|
||||
|
||||
// CreateTemplate 创建模板
|
||||
// @Summary 创建模板
|
||||
// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body template.CreateTemplateRequest true "创建请求参数"
|
||||
// @Success 200 {object} response.Any{data=string} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误或模板标识符已存在"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates [post]
|
||||
func CreateTemplate(c *gin.Context) {
|
||||
var req CreateTemplateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
tmpl, err := createTemplate(c.Request.Context(), req)
|
||||
if abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tmpl))
|
||||
}
|
||||
|
||||
// ListTemplates 获取模板列表
|
||||
// @Summary 获取模板列表
|
||||
// @Description 返回所有通知模板列表,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]model.Template} "模板列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates [get]
|
||||
func ListTemplates(c *gin.Context) {
|
||||
templates, err := listTemplates(c.Request.Context())
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(templates))
|
||||
}
|
||||
|
||||
// GetTemplate 获取单个模板
|
||||
// @Summary 获取单个模板
|
||||
// @Description 根据模板标识符获取对应的模板详情,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "模板标识符"
|
||||
// @Success 200 {object} response.Any{data=model.Template} "模板详情"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "模板不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates/{key} [get]
|
||||
func GetTemplate(c *gin.Context) {
|
||||
tmpl, err := getTemplate(c.Request.Context(), c.Param("key"))
|
||||
if abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tmpl))
|
||||
}
|
||||
|
||||
// UpdateTemplate 更新模板
|
||||
// @Summary 更新模板
|
||||
// @Description 根据模板标识符更新对应的模板内容,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "模板标识符"
|
||||
// @Param request body template.UpdateTemplateRequest true "更新请求参数"
|
||||
// @Success 200 {object} response.Any{data=model.Template} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "模板不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates/{key} [put]
|
||||
func UpdateTemplate(c *gin.Context) {
|
||||
var req UpdateTemplateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req)
|
||||
if abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tmpl))
|
||||
}
|
||||
|
||||
// DeleteTemplate 删除模板
|
||||
// @Summary 删除模板
|
||||
// @Description 根据模板标识符删除对应模板,系统预置模板不可删除,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "模板标识符"
|
||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||
// @Failure 400 {object} response.Any "不可删除系统模板"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "模板不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates/{key} [delete]
|
||||
func DeleteTemplate(c *gin.Context) {
|
||||
if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package template
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
r := testhelper.NewTestGinEngine()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
|
||||
// Mock authentication middleware
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.GET("/templates", ListTemplates)
|
||||
adminGroup.POST("/templates", CreateTemplate)
|
||||
|
||||
templateRouter := adminGroup.Group("/templates/:key")
|
||||
{
|
||||
templateRouter.GET("", GetTemplate)
|
||||
templateRouter.PUT("", UpdateTemplate)
|
||||
templateRouter.DELETE("", DeleteTemplate)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
func TestCreateTemplate(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("create successfully", func(t *testing.T) {
|
||||
payload := CreateTemplateRequest{
|
||||
Key: "test_template",
|
||||
Name: "Test Template",
|
||||
Type: "email",
|
||||
Subject: "Test Subject",
|
||||
Content: "Hello {{.Name}}",
|
||||
Description: "Test Desc",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/templates", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var tmpl model.Template
|
||||
err := dbConn.Where("key = ?", "test_template").First(&tmpl).Error
|
||||
if err != nil {
|
||||
t.Fatalf("failed to find template in DB: %v", err)
|
||||
}
|
||||
if tmpl.Name != "Test Template" {
|
||||
t.Errorf("expected Name 'Test Template', got '%s'", tmpl.Name)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("create duplicate key error", func(t *testing.T) {
|
||||
payload := CreateTemplateRequest{
|
||||
Key: "test_template",
|
||||
Name: "Another Name",
|
||||
Type: "email",
|
||||
Subject: "Another Subject",
|
||||
Content: "Hello",
|
||||
Description: "desc",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("POST", "/api/v1/admin/templates", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request on duplicate key, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestListTemplates(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed system templates manually for testing
|
||||
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
|
||||
t2 := model.Template{Key: "register_email", Name: "Register Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
|
||||
dbConn.Create(&t1)
|
||||
dbConn.Create(&t2)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("list templates", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/templates", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
var resp response.Any
|
||||
_ = json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var templates []model.Template
|
||||
_ = json.Unmarshal(dataBytes, &templates)
|
||||
|
||||
if len(templates) != 2 {
|
||||
t.Errorf("expected 2 templates, got %d", len(templates))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetTemplate(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
|
||||
dbConn.Create(&t1)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("get existing", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/templates/login_email", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("get non-existent", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/templates/non_existent", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("expected 404 Not Found, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateTemplate(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
|
||||
dbConn.Create(&t1)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("update successfully", func(t *testing.T) {
|
||||
payload := UpdateTemplateRequest{
|
||||
Name: "Updated Login Code",
|
||||
Type: "email",
|
||||
Subject: "New Subject",
|
||||
Content: "new code {{.Code}}",
|
||||
Description: "new desc",
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
req, _ := http.NewRequest("PUT", "/api/v1/admin/templates/login_email", bytes.NewBuffer(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var tmpl model.Template
|
||||
dbConn.Where("key = ?", "login_email").First(&tmpl)
|
||||
if tmpl.Name != "Updated Login Code" || tmpl.Subject != "New Subject" {
|
||||
t.Errorf("database values not updated: %+v", tmpl)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestDeleteTemplate(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
t1 := model.Template{Key: "login_email", Name: "Login Code", Type: "email", Content: "code {{.Code}}", IsSystem: true}
|
||||
t2 := model.Template{Key: "custom_tmpl", Name: "Custom", Type: "email", Content: "hi", IsSystem: false}
|
||||
dbConn.Create(&t1)
|
||||
dbConn.Create(&t2)
|
||||
|
||||
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("delete system template should fail", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("DELETE", "/api/v1/admin/templates/login_email", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request when deleting system template, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("delete custom template should succeed", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("DELETE", "/api/v1/admin/templates/custom_tmpl", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
|
||||
var count int64
|
||||
dbConn.Model(&model.Template{}).Where("key = ?", "custom_tmpl").Count(&count)
|
||||
if count != 0 {
|
||||
t.Error("custom template was not deleted from DB")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package updater manages GitHub Release checks and in-place application upgrades.
|
||||
package updater
|
||||
|
||||
const (
|
||||
errInvalidRepository = "上游仓库地址无效"
|
||||
errReleaseRequestFailed = "获取上游版本失败"
|
||||
errReleaseResponseInvalid = "上游版本响应无效"
|
||||
errNoCompatibleRelease = "未找到兼容的 Release"
|
||||
errNoCompatibleAsset = "未找到当前系统对应的 Release 资产"
|
||||
errDevelopmentBuild = "开发版本无法执行自动升级"
|
||||
errAlreadyUpToDate = "当前已是最新版本"
|
||||
errUpgradeAlreadyRunning = "已有升级任务正在执行"
|
||||
errAutomaticUpgradeBlocked = "当前平台暂不支持自动替换二进制"
|
||||
)
|
||||
@@ -0,0 +1,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)
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package user 提供用户管理功能
|
||||
package user
|
||||
|
||||
const (
|
||||
userNotFound = "用户不存在"
|
||||
cannotDisable = "不能禁用管理员用户"
|
||||
cannotDelete = "不能删除管理员用户"
|
||||
cannotDeleteSelf = "不能删除当前登录用户"
|
||||
updateUserFailed = "更新用户状态失败"
|
||||
deleteUserFailed = "删除用户失败"
|
||||
usernameExists = "用户名已存在"
|
||||
usernameRequired = "用户名不能为空"
|
||||
passwordTooShort = "密码长度不能少于 8 位" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
createUserFailed = "创建用户失败"
|
||||
emailRequired = "邮箱不能为空"
|
||||
emailExists = "邮箱已被注册"
|
||||
)
|
||||
@@ -0,0 +1,106 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||
)
|
||||
|
||||
func listUsers(ctx context.Context, req listUsersRequest) (int64, []model.User, error) {
|
||||
return repository.ListAdminUsers(ctx, repository.AdminUserListFilter{
|
||||
UserID: req.UserID,
|
||||
Username: strings.TrimSpace(req.Username),
|
||||
Page: req.Page,
|
||||
PageSize: req.PageSize,
|
||||
})
|
||||
}
|
||||
|
||||
func getUserDetail(ctx context.Context, id uint64) (model.User, error) {
|
||||
return repository.GetAdminUserDetail(ctx, id)
|
||||
}
|
||||
|
||||
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
|
||||
flags, err := repository.GetUserAdminFlags(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !active && flags.IsAdmin {
|
||||
return errors.New(cannotDisable)
|
||||
}
|
||||
return repository.UpdateUserActive(ctx, id, active)
|
||||
}
|
||||
|
||||
func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
|
||||
if currentUserID == targetID {
|
||||
return errors.New(cannotDeleteSelf)
|
||||
}
|
||||
flags, err := repository.GetUserAdminFlags(ctx, targetID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if flags.IsAdmin {
|
||||
return errors.New(cannotDelete)
|
||||
}
|
||||
return repository.DeleteUserWithRelations(ctx, targetID)
|
||||
}
|
||||
|
||||
func createUser(ctx context.Context, req createUserRequest) (model.User, error) {
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
req.Nickname = strings.TrimSpace(req.Nickname)
|
||||
req.Password = strings.TrimSpace(req.Password)
|
||||
req.Email = strings.TrimSpace(req.Email)
|
||||
|
||||
if req.Username == "" {
|
||||
return model.User{}, errors.New(usernameRequired)
|
||||
}
|
||||
if req.Email == "" {
|
||||
return model.User{}, errors.New(emailRequired)
|
||||
}
|
||||
if len(req.Password) < minPasswordLength {
|
||||
return model.User{}, errors.New(passwordTooShort)
|
||||
}
|
||||
|
||||
count, err := repository.CountUsersByUsername(ctx, req.Username)
|
||||
if err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
if count > 0 {
|
||||
return model.User{}, errors.New(usernameExists)
|
||||
}
|
||||
|
||||
emailCount, err := repository.CountUsersByEmail(ctx, req.Email)
|
||||
if err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
if emailCount > 0 {
|
||||
return model.User{}, errors.New(emailExists)
|
||||
}
|
||||
|
||||
newUser := model.User{
|
||||
ID: idgen.NextUint64ID(),
|
||||
Username: req.Username,
|
||||
Nickname: req.Nickname,
|
||||
Email: req.Email,
|
||||
IsActive: req.IsActive,
|
||||
IsAdmin: req.IsAdmin,
|
||||
LastLoginAt: time.Time{},
|
||||
}
|
||||
if newUser.Nickname == "" {
|
||||
newUser.Nickname = req.Username
|
||||
}
|
||||
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
if err := repository.CreateUser(ctx, &newUser); err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
return newUser, nil
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
// minPasswordLength 密码最小长度
|
||||
const minPasswordLength = 8
|
||||
|
||||
// listUsersRequest 用户列表查询请求
|
||||
type listUsersRequest struct {
|
||||
Page int `form:"page" binding:"min=1"`
|
||||
PageSize int `form:"page_size" binding:"min=1,max=100"`
|
||||
UserID *uint64 `form:"user_id" binding:"omitempty,gt=0"`
|
||||
Username string `form:"username"`
|
||||
}
|
||||
|
||||
type user struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
Email string `json:"email"`
|
||||
AvatarURL string `json:"avatar_url"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
Bio string `json:"bio"`
|
||||
Phone string `json:"phone"`
|
||||
Gender string `json:"gender"`
|
||||
Website string `json:"website"`
|
||||
Location string `json:"location"`
|
||||
LastLoginAt time.Time `json:"last_login_at"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// listUsersResponse 用户列表响应
|
||||
type listUsersResponse struct {
|
||||
Users []user `json:"users"`
|
||||
Total int64 `json:"total"`
|
||||
}
|
||||
|
||||
func parseUserID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil || id == 0 {
|
||||
response.AbortBadRequest(c, userNotFound)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
func toUser(u model.User) user {
|
||||
return user{
|
||||
ID: u.ID,
|
||||
Username: u.Username,
|
||||
Nickname: u.Nickname,
|
||||
Email: u.Email,
|
||||
AvatarURL: u.AvatarURL,
|
||||
IsActive: u.IsActive,
|
||||
IsAdmin: u.IsAdmin,
|
||||
Bio: u.Bio,
|
||||
Phone: u.Phone,
|
||||
Gender: u.Gender,
|
||||
Website: u.Website,
|
||||
Location: u.Location,
|
||||
LastLoginAt: u.LastLoginAt,
|
||||
CreatedAt: u.CreatedAt,
|
||||
UpdatedAt: u.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbiddenMsgs, badRequestMsgs []string) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, notFoundMsg)
|
||||
return true
|
||||
}
|
||||
msg := err.Error()
|
||||
for _, m := range badRequestMsgs {
|
||||
if msg == m {
|
||||
response.AbortBadRequest(c, msg)
|
||||
return true
|
||||
}
|
||||
}
|
||||
for _, m := range forbiddenMsgs {
|
||||
if msg == m {
|
||||
response.AbortForbidden(c, msg)
|
||||
return true
|
||||
}
|
||||
}
|
||||
response.AbortInternal(c, msg)
|
||||
return true
|
||||
}
|
||||
|
||||
// ListUsers 获取用户列表
|
||||
// @Summary 获取用户列表
|
||||
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request query listUsersRequest true "查询参数"
|
||||
// @Success 200 {object} response.Any{data=user.listUsersResponse} "用户列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users [get]
|
||||
func ListUsers(c *gin.Context) {
|
||||
var req listUsersRequest
|
||||
if err := c.ShouldBindQuery(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
total, modelUsers, err := listUsers(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
users := make([]user, 0, len(modelUsers))
|
||||
for _, modelUser := range modelUsers {
|
||||
users = append(users, toUser(modelUser))
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(listUsersResponse{
|
||||
Users: users,
|
||||
Total: total,
|
||||
}))
|
||||
}
|
||||
|
||||
// GetUser 获取用户详情
|
||||
// @Summary 获取用户详情
|
||||
// @Description 返回指定用户的完整个人资料和系统状态,需要管理员权限,不返回密码等敏感字段
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Success 200 {object} response.Any{data=user.user} "用户详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "用户不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users/{id} [get]
|
||||
func GetUser(c *gin.Context) {
|
||||
id, ok := parseUserID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
targetUser, err := getUserDetail(c.Request.Context(), id)
|
||||
if abortUserLogicError(c, err, userNotFound, nil, nil) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(toUser(targetUser)))
|
||||
}
|
||||
|
||||
// updateUserStatusRequest 更新用户状态请求
|
||||
type updateUserStatusRequest struct {
|
||||
IsActive bool `json:"is_active"`
|
||||
}
|
||||
|
||||
// UpdateUserStatus 更新用户状态(启用/禁用)
|
||||
// @Summary 更新用户状态
|
||||
// @Description 启用或禁用指定用户,管理员账号无法被禁用,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Param request body updateUserStatusRequest true "状态参数"
|
||||
// @Success 200 {object} response.Any{data=string} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限或尝试禁用管理员"
|
||||
// @Failure 404 {object} response.Any "用户不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users/{id}/status [put]
|
||||
func UpdateUserStatus(c *gin.Context) {
|
||||
var req updateUserStatusRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
id, ok := parseUserID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
|
||||
if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) {
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, updateUserFailed)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// DeleteUser 删除用户
|
||||
// @Summary 删除用户
|
||||
// @Description 删除指定非管理员用户,需要管理员权限,不能删除当前登录用户
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限、尝试删除管理员或当前用户"
|
||||
// @Failure 404 {object} response.Any "用户不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users/{id} [delete]
|
||||
func DeleteUser(c *gin.Context) {
|
||||
id, ok := parseUserID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||
if err := deleteUser(c.Request.Context(), currUser.ID, id); err != nil {
|
||||
if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) {
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, deleteUserFailed)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// createUserRequest 创建用户请求
|
||||
type createUserRequest struct {
|
||||
Username string `json:"username" binding:"required,min=3,max=64"`
|
||||
Password string `json:"password" binding:"required,min=8,max=64"`
|
||||
Nickname string `json:"nickname" binding:"omitempty,max=64"`
|
||||
Email string `json:"email" binding:"required,email,max=255"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
}
|
||||
|
||||
// CreateUser 创建用户
|
||||
// @Summary 创建用户
|
||||
// @Description 创建一个本地密码登录的新用户,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body user.createUserRequest true "创建用户参数"
|
||||
// @Success 200 {object} response.Any{data=user.user} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误或用户名已存在"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users [post]
|
||||
func CreateUser(c *gin.Context) {
|
||||
var req createUserRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
newUser, err := createUser(c.Request.Context(), req)
|
||||
if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(toUser(newUser)))
|
||||
}
|
||||
@@ -0,0 +1,582 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package user
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
)
|
||||
|
||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||
r := testhelper.NewTestGinEngine()
|
||||
adminGroup := r.Group("/api/v1/admin")
|
||||
|
||||
// Mock authentication middleware
|
||||
adminGroup.Use(func(c *gin.Context) {
|
||||
if authUser != nil {
|
||||
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
adminGroup.GET("/users", ListUsers)
|
||||
adminGroup.POST("/users", CreateUser)
|
||||
adminGroup.GET("/users/:id", GetUser)
|
||||
adminGroup.PUT("/users/:id/status", UpdateUserStatus)
|
||||
adminGroup.DELETE("/users/:id", DeleteUser)
|
||||
return r
|
||||
}
|
||||
|
||||
func TestListUsers(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed users
|
||||
users := []model.User{
|
||||
{
|
||||
ID: 1001,
|
||||
Username: "alice",
|
||||
Nickname: "Alice Nickname",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
LastLoginAt: time.Now(),
|
||||
},
|
||||
{
|
||||
ID: 1002,
|
||||
Username: "bob",
|
||||
Nickname: "Bob Nickname",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
LastLoginAt: time.Now(),
|
||||
},
|
||||
{
|
||||
ID: 1003,
|
||||
Username: "charlie",
|
||||
Nickname: "Charlie Nickname",
|
||||
IsActive: false,
|
||||
IsAdmin: true,
|
||||
LastLoginAt: time.Now(),
|
||||
},
|
||||
}
|
||||
|
||||
for _, u := range users {
|
||||
if err := dbConn.Create(&u).Error; err != nil {
|
||||
t.Fatalf("failed to seed user: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("basic pagination list", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=2", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var resp response.Any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("failed to parse response: %v", err)
|
||||
}
|
||||
|
||||
// Parse data map to our structure
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var listResp listUsersResponse
|
||||
if err := json.Unmarshal(dataBytes, &listResp); err != nil {
|
||||
t.Fatalf("failed to parse list response: %v", err)
|
||||
}
|
||||
|
||||
if len(listResp.Users) != 2 {
|
||||
t.Errorf("expected 2 users, got %d", len(listResp.Users))
|
||||
}
|
||||
if listResp.Total != 3 {
|
||||
t.Errorf("expected total 3, got %d", listResp.Total)
|
||||
}
|
||||
// Ordered by ID ASC
|
||||
if listResp.Users[0].ID != 1001 || listResp.Users[1].ID != 1002 {
|
||||
t.Errorf("expected ordered ASC, got first ID %d, second ID %d", listResp.Users[0].ID, listResp.Users[1].ID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filter by user_id", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&user_id=1001", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
var resp response.Any
|
||||
_ = json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var listResp listUsersResponse
|
||||
_ = json.Unmarshal(dataBytes, &listResp)
|
||||
|
||||
if len(listResp.Users) != 1 || listResp.Users[0].ID != 1001 {
|
||||
t.Errorf("expected 1 user with ID 1001, got total %d", len(listResp.Users))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("filter by username prefix", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=1&page_size=10&username=bo", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
var resp response.Any
|
||||
_ = json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var listResp listUsersResponse
|
||||
_ = json.Unmarshal(dataBytes, &listResp)
|
||||
|
||||
if len(listResp.Users) != 1 || listResp.Users[0].Username != "bob" {
|
||||
t.Errorf("expected bob, got %v", listResp.Users)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid pagination parameter", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users?page=0&page_size=10", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("expected 400 Bad Request, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestGetUser(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
targetUser := model.User{
|
||||
ID: 1001,
|
||||
Username: "alice",
|
||||
Password: "secret-hash",
|
||||
Nickname: "Alice Nickname",
|
||||
Email: "alice@example.com",
|
||||
AvatarURL: "https://example.com/avatar.png",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
Bio: "hello",
|
||||
Phone: "123456",
|
||||
Gender: "female",
|
||||
Website: "https://example.com",
|
||||
Location: "Shanghai",
|
||||
}
|
||||
if err := dbConn.Create(&targetUser).Error; err != nil {
|
||||
t.Fatalf("failed to seed user: %v", err)
|
||||
}
|
||||
|
||||
adminUser := &model.User{ID: 1003, Username: "charlie", IsAdmin: true}
|
||||
router := setupTestRouter(adminUser)
|
||||
|
||||
t.Run("get full user profile", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users/1001", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var resp response.Any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("failed to parse response: %v", err)
|
||||
}
|
||||
|
||||
dataBytes, _ := json.Marshal(resp.Data)
|
||||
var resUser user
|
||||
if err := json.Unmarshal(dataBytes, &resUser); err != nil {
|
||||
t.Fatalf("failed to parse response data: %v", err)
|
||||
}
|
||||
|
||||
if resUser.Email != targetUser.Email || resUser.Bio != targetUser.Bio || resUser.Phone != targetUser.Phone ||
|
||||
resUser.Gender != targetUser.Gender || resUser.Website != targetUser.Website || resUser.Location != targetUser.Location {
|
||||
t.Errorf("profile fields were not returned correctly: %+v", resUser)
|
||||
}
|
||||
if bytes.Contains(dataBytes, []byte("secret-hash")) {
|
||||
t.Error("response should not include password")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("get non-existent user", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/api/v1/admin/users/9999", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("expected 404 Not Found, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestUpdateUserStatus(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed users
|
||||
regularUser := model.User{
|
||||
ID: 1001,
|
||||
Username: "alice",
|
||||
IsActive: true,
|
||||
IsAdmin: false,
|
||||
}
|
||||
adminUser := model.User{
|
||||
ID: 1002,
|
||||
Username: "bob",
|
||||
IsActive: true,
|
||||
IsAdmin: true,
|
||||
}
|
||||
|
||||
dbConn.Create(®ularUser)
|
||||
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(®ularUser).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())
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user