refactor(api): extract repository layer and thin HTTP handlers

Introduce internal/repository for data access and cache-backed system
config reads. Move business logic into logics.go across admin push,
user, template, cache, system_config, and upload/handler packages.

Remove Gin from internal/util by relocating request-scoped helpers to
oauth/gin_context.go. Propagate request context for config lookups in
user flows. Slim model entities and delete model-level DB/cache helpers.

Wire handlers to logics/repository so targeted packages no longer call
db.DB directly. Update admin router tests to use ErrorHandlerMiddleware.
This commit is contained in:
ryan
2026-06-18 12:12:49 +08:00
parent e5b3a60f73
commit 1b2e083aec
77 changed files with 2370 additions and 1783 deletions
@@ -13,7 +13,6 @@ import ("bytes"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response")
@@ -26,7 +25,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
util.SetToContext(c, oauth.UserObjKey, authUser)
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
+14
View File
@@ -0,0 +1,14 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cache
import (
"context"
"github.com/Rain-kl/Wavelet/internal/repository"
)
func saveOrUpdateConfig(ctx context.Context, key, value string) error {
return repository.SaveOrUpdateSystemConfig(ctx, key, value)
}
+4 -37
View File
@@ -4,18 +4,16 @@
// Package cache provides HTTP handlers for managing disk cache.
package cache
import ("context"
"errors"
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response")
"github.com/Rain-kl/Wavelet/internal/common/response"
)
type updateCacheConfigRequest struct {
MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"`
@@ -62,25 +60,21 @@ func UpdateCacheConfig(c *gin.Context) {
ctx := c.Request.Context()
// Update Max Size
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
// Update Default TTL
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
// Update LRU Enabled
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
response.AbortInternal(c, err.Error())
return
}
// Trigger hot reloading in global cache
diskcache.GetGlobalCache().ReloadConfig(ctx)
c.JSON(http.StatusOK, response.OKNil())
@@ -103,31 +97,4 @@ func ClearCache(c *gin.Context) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
func saveOrUpdateConfig(ctx context.Context, key string, value string) error {
var sc model.SystemConfig
err := db.DB(ctx).Where("key = ?", key).First(&sc).Error
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
if errors.Is(err, gorm.ErrRecordNotFound) {
sc = model.SystemConfig{
Key: key,
Value: value,
Type: "system",
Visibility: 0,
}
if err := db.DB(ctx).Create(&sc).Error; err != nil {
return err
}
} else {
sc.Value = value
if err := db.DB(ctx).Save(&sc).Error; err != nil {
return err
}
}
return model.InvalidateSystemConfigCache(ctx, key)
}
}
+2 -2
View File
@@ -13,6 +13,7 @@ import (
"github.com/gorilla/websocket"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
)
// getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击
@@ -32,8 +33,7 @@ func getUpgrader() *websocket.Upgrader {
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
ctx := r.Context()
var sc model.SystemConfig
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
allowedOrigins := strings.Split(sc.Value, ",")
for _, allowed := range allowedOrigins {
+3 -4
View File
@@ -7,7 +7,6 @@ package admin
import (
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
@@ -22,11 +21,11 @@ func LoginAdminRequired() gin.HandlerFunc {
ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
if tokenAuth, _ := util.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
tokenAdmin, _ := util.GetFromContext[bool](c, oauth.TokenAdminKey)
if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey)
if !tokenAdmin {
response.AbortNotFound(c, TokenAdminRequired)
return
+23 -124
View File
@@ -3,19 +3,20 @@
package push
import ("encoding/json"
import (
"encoding/json"
"errors"
"net/http"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
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/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListChannelDefinitions 获取各种消息通道的表单配置定义列表
// @Summary 获取所有消息通道配置字段定义
@@ -38,9 +39,8 @@ func ListChannelDefinitions(c *gin.Context) {
// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表"
// @Router /api/v1/admin/push/channels [get]
func ListChannels(c *gin.Context) {
ctx := c.Request.Context()
var channels []model.PushChannel
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
channels, err := listPushChannels(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -75,41 +75,11 @@ func CreateChannel(c *gin.Context) {
return
}
ctx := c.Request.Context()
var count int64
if err := db.DB(ctx).Model(&model.PushChannel{}).Where("name = ?", req.Name).Count(&count).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if count > 0 {
response.AbortBadRequest(c, "channel name already exists")
return
}
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 {
channel, err := createPushChannel(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(ctx).Create(&channel).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
// 缓存一致性:清除渠道缓存
model.DeleteActivePushChannelCache(ctx, channel.Name)
c.JSON(http.StatusOK, response.OK(channel))
}
@@ -135,8 +105,7 @@ type UpdateChannelRequest struct {
// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功"
// @Router /api/v1/admin/push/channels/{id} [put]
func UpdateChannel(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
@@ -148,10 +117,8 @@ func UpdateChannel(c *gin.Context) {
return
}
ctx := c.Request.Context()
var channel model.PushChannel
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
channel, err := updatePushChannel(c.Request.Context(), id, req)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
@@ -159,27 +126,6 @@ func UpdateChannel(c *gin.Context) {
response.AbortInternal(c, err.Error())
return
}
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 {
response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(ctx).Save(&channel).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
// 缓存一致性:清除渠道缓存
model.DeleteActivePushChannelCache(ctx, channel.Name)
c.JSON(http.StatusOK, response.OK(channel))
}
@@ -193,16 +139,13 @@ func UpdateChannel(c *gin.Context) {
// @Success 200 {object} response.Any "删除成功"
// @Router /api/v1/admin/push/channels/{id} [delete]
func DeleteChannel(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid channel id")
return
}
ctx := c.Request.Context()
var channel model.PushChannel
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
if err := deletePushChannel(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "channel not found")
return
@@ -210,15 +153,6 @@ func DeleteChannel(c *gin.Context) {
response.AbortInternal(c, err.Error())
return
}
if err := db.DB(ctx).Delete(&channel).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
// 缓存一致性:清除渠道缓存
model.DeleteActivePushChannelCache(ctx, channel.Name)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -250,26 +184,12 @@ func TestChannel(c *gin.Context) {
}
ctx := c.Request.Context()
var url, token, other, channelType string
if req.Name != "" {
var channel model.PushChannel
if err := db.DB(ctx).Where("name = ?", req.Name).First(&channel).Error; err != nil {
response.AbortBadRequest(c, "channel not found")
return
}
url = channel.URL
token = channel.Token
other = channel.Other
channelType = channel.Type
} else {
url = req.URL
token = req.Token
other = req.Other
channelType = req.Type
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)
}
@@ -282,7 +202,6 @@ func TestChannel(c *gin.Context) {
Type: channelType,
Enabled: true,
}
if err := tempChannel.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -291,34 +210,16 @@ func TestChannel(c *gin.Context) {
var config pkgpush.Config
var renderedJSON string
switch channelType {
case channelLark:
config = pkgpush.Config{
Channel: channelLark,
URL: url,
Secret: token,
}
config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token}
renderedJSON = other
case channelEmail:
config = pkgpush.Config{
Channel: channelEmail,
URL: url,
Key: token,
Secret: other,
}
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{
Channel: channelTelegram,
URL: url,
Secret: token,
Key: other,
}
config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other}
default:
config = pkgpush.Config{
Channel: channelCustom,
URL: url,
}
config = pkgpush.Config{Channel: channelCustom, URL: url}
customPushReq := CustomPushRequest{
Title: "通道测试通知",
Content: "这是一条来自系统的消息通道连通性测试消息。",
@@ -340,12 +241,10 @@ func TestChannel(c *gin.Context) {
},
Template: renderedJSON,
}
if err := enqueuePushTask(ctx, payload); err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -376,4 +275,4 @@ func renderCustomPayload(template string, req CustomPushRequest) string {
result = strings.ReplaceAll(result, "$url", escapeJSONString(req.URL))
result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To))
return result
}
}
+11 -137
View File
@@ -9,12 +9,9 @@ import (
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"sync"
"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/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
@@ -73,30 +70,7 @@ type EventTrigger struct{}
// DefaultTrigger is the singleton instance of EventTrigger.
var DefaultTrigger = &EventTrigger{}
var (
systemUser *model.User
systemOnce sync.Once
)
func getSystemUser(ctx context.Context) *model.User {
systemOnce.Do(func() {
var u model.User
if err := db.DB(ctx).Where("username = ?", "system").First(&u).Error; err == nil {
systemUser = &u
} else {
systemUser = &model.User{
ID: 999,
Username: "system",
Nickname: "系统",
Email: "",
}
}
})
return systemUser
}
// Trigger receives event metadata and processes the event notification dispatch asynchronously.
// It automatically enqueues tasks using a background goroutine and avoids blocking the calling thread.
//
//nolint:contextcheck
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
@@ -109,8 +83,7 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map
body["user"] = getSystemUser(asyncCtx)
}
// 1. Check if the event is enabled (try Redis cache first)
eventPtr, err := model.GetActivePushEventByKey(asyncCtx, meta.Key)
eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return
@@ -119,22 +92,16 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map
return
}
event := *eventPtr
if len(event.Channels) == 0 {
return
}
// 2. Build and render notification message
flatBody := getFlatBody(body)
msg, _ := t.buildMessage(&event, meta, flatBody, body)
// 3. Enqueue tasks for each matching channel
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 := ""
@@ -222,13 +189,11 @@ func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[str
func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, msg NotificationMessage, flatBody map[string]any) {
for _, channelName := range event.Channels {
// 检查是不是自定义数据库渠道 (使用 Redis 缓存优先)
customChannel, err := model.GetActivePushChannelByName(ctx, channelName)
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)
}
}
@@ -251,32 +216,15 @@ func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, m
switch channel.Type {
case channelLark:
config = pkgpush.Config{
Channel: channelLark,
URL: channel.URL,
Secret: channel.Token, // Feishu Bot Sign Secret
}
renderedTemplate = channel.Other // Optional custom template/card for lark
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, // SMTP host:port
Key: token, // SMTP Username
Secret: other, // SMTP Password
}
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
case channelTelegram:
config = pkgpush.Config{
Channel: channelTelegram,
URL: channel.URL,
Secret: channel.Token, // Telegram Bot Token
Key: channel.Other, // Default Chat ID
}
default: // custom
config = pkgpush.Config{
Channel: channelCustom,
URL: channel.URL,
}
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,
@@ -301,8 +249,6 @@ func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, m
}
}
func enqueuePushTask(ctx context.Context, payload SendPayload) error {
payloadBytes, err := json.Marshal(payload)
if err != nil {
@@ -348,48 +294,23 @@ func resolveTarget(ctx context.Context, target string, flatBody map[string]any,
}
resolved := resolveDynamicKeyword(target, flatBody)
// 2. 如果包含 @,说明已经是个邮箱,直接返回
if strings.Contains(resolved, "@") {
return resolved
}
// 2.5 如果为特殊的系统虚拟用户,自动映射为首位管理员
if val, matched := resolveSystemTarget(ctx, resolved, channel); matched {
return val
}
// 3. 不包含 @,说明可能是用户 ID 或用户名。我们需要从数据库中查询对应用户
var user model.User
found := false
// 尝试作为用户 ID 查询(纯数字)
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if err := db.DB(ctx).Where("id = ?", id).First(&user).Error; err == nil {
found = true
}
}
// 如果没有按 ID 查到,尝试作为用户名查询
if !found {
if err := db.DB(ctx).Where("username = ?", resolved).First(&user).Error; err == nil {
found = true
}
}
// 4. 根据查询结果 and 推送渠道进行转换
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
}
@@ -418,51 +339,4 @@ func resolveDynamicKeyword(target string, flatBody map[string]any) string {
}
}
return target
}
func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) {
if resolved != "系统" && resolved != "system" && resolved != "0" {
return "", false
}
var adminUser model.User
if err := db.DB(ctx).Where("is_admin = ?", true).Order("id asc").First(&adminUser).Error; 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
}
// resolveSMTPConfig resolves SMTP configuration by falling back to system-wide global configuration if any inputs are blank.
func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) {
if url != "" && token != "" {
return url, token, other
}
var smtpHost, smtpPort, smtpUser, smtpPass model.SystemConfig
_ = smtpHost.GetByKey(ctx, model.ConfigKeySMTPHost)
_ = smtpPort.GetByKey(ctx, model.ConfigKeySMTPPort)
_ = smtpUser.GetByKey(ctx, model.ConfigKeySMTPUsername)
_ = smtpPass.GetByKey(ctx, model.ConfigKeySMTPPassword)
if smtpHost.Value == "" || smtpUser.Value == "" {
return url, token, other
}
port := smtpPort.Value
if port == "" {
port = "587"
}
if url == "" {
url = smtpHost.Value + ":" + port
}
if token == "" {
token = smtpUser.Value
}
if other == "" {
other = smtpPass.Value
}
return url, token, other
}
}
+413
View File
@@ -0,0 +1,413 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"encoding/json"
"errors"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/task"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"gorm.io/gorm"
)
type smtpConfig struct {
Host string
Port string
Username string
Password string
}
func loadSMTPConfig(ctx context.Context) smtpConfig {
host, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost)
port, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort)
user, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername)
pass, _ := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword)
return smtpConfig{
Host: host.Value,
Port: port.Value,
Username: user.Value,
Password: pass.Value,
}
}
func syncBuiltInEvents(ctx context.Context) error {
for _, meta := range BuiltInEvents {
_, err := repository.GetPushEventByKey(ctx, meta.Key)
if errors.Is(err, gorm.ErrRecordNotFound) {
var defaultTemplateStr string
if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil {
defaultTemplateStr = string(defaultTemplateBytes)
}
event := model.PushEvent{
EventKey: meta.Key,
Name: meta.Name,
Channels: []string{},
Targets: []string{},
Template: defaultTemplateStr,
Enabled: false,
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
return err
}
} else if err != nil {
return err
}
}
return nil
}
func listPushEvents(ctx context.Context) ([]model.PushEvent, error) {
return repository.ListPushEvents(ctx)
}
func createPushEvent(ctx context.Context, req CreateEventRequest) (model.PushEvent, error) {
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
if err != nil {
return model.PushEvent{}, err
}
count, err := repository.CountPushEventsByKey(ctx, eventKey)
if err != nil {
return model.PushEvent{}, err
}
if count > 0 {
return model.PushEvent{}, errors.New("this notification event is already configured")
}
templateStr := strings.TrimSpace(req.Template)
if templateStr == "" {
templateStr = string(defaultTemplateBytes)
} else {
var tempMap map[string]any
if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil {
return model.PushEvent{}, errors.New("custom template is not a valid JSON format")
}
}
channels := req.Channels
if channels == nil {
channels = []string{}
}
targets := req.Targets
if targets == nil {
targets = []string{}
}
event := model.PushEvent{
EventKey: eventKey,
Name: eventName,
TaskType: req.TaskType,
Channels: channels,
Targets: targets,
Template: templateStr,
Enabled: req.Enabled,
}
if err := event.Validate(); err != nil {
return model.PushEvent{}, err
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
return model.PushEvent{}, err
}
return event, nil
}
func deletePushEvent(ctx context.Context, id uint64) error {
event, err := repository.GetPushEventByID(ctx, id)
if err != nil {
return err
}
return repository.DeletePushEvent(ctx, &event)
}
func updatePushEvent(ctx context.Context, id uint64, req UpdateEventRequest) error {
event, err := repository.GetPushEventByID(ctx, id)
if err != nil {
return err
}
event.Channels = req.Channels
event.Targets = req.Targets
event.Template = req.Template
event.Enabled = req.Enabled
if err := event.Validate(); err != nil {
return err
}
return repository.SavePushEvent(ctx, &event)
}
func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
event, err := repository.GetPushEventByID(ctx, id)
if err != nil {
return false, err
}
enabled := !event.Enabled
if enabled && len(event.Channels) == 0 {
return false, errors.New("cannot enable event without any push channels configured")
}
if err := repository.UpdatePushEventEnabled(ctx, &event, enabled); err != nil {
return false, err
}
return enabled, nil
}
func listPushHistories(ctx context.Context, filter repository.PushHistoryListFilter) (int64, []model.PushHistory, error) {
return repository.ListPushHistories(ctx, filter)
}
func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) {
if cfg.Channel != channelEmail || (cfg.URL != "" && cfg.Key != "") {
return
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return
}
port := smtp.Port
if port == "" {
port = "587"
}
cfg.URL = smtp.Host + ":" + port
cfg.Key = smtp.Username
cfg.Secret = smtp.Password
}
func listPushChannels(ctx context.Context) ([]model.PushChannel, error) {
return repository.ListPushChannels(ctx)
}
func createPushChannel(ctx context.Context, req CreateChannelRequest) (model.PushChannel, error) {
count, err := repository.CountPushChannelsByName(ctx, req.Name)
if err != nil {
return model.PushChannel{}, err
}
if count > 0 {
return model.PushChannel{}, errors.New("channel name already exists")
}
channel := model.PushChannel{
Name: req.Name,
Description: req.Description,
Type: req.Type,
Token: req.Token,
URL: req.URL,
Other: req.Other,
Enabled: req.Enabled,
}
if err := channel.Validate(); err != nil {
return model.PushChannel{}, err
}
if err := repository.CreatePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
}
return channel, nil
}
func updatePushChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (model.PushChannel, error) {
channel, err := repository.GetPushChannelByID(ctx, id)
if err != nil {
return model.PushChannel{}, err
}
channel.Description = req.Description
channel.Type = req.Type
channel.Token = req.Token
channel.URL = req.URL
channel.Other = req.Other
channel.Enabled = req.Enabled
if err := channel.Validate(); err != nil {
return model.PushChannel{}, err
}
if err := repository.SavePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
}
return channel, nil
}
func deletePushChannel(ctx context.Context, id uint64) error {
channel, err := repository.GetPushChannelByID(ctx, id)
if err != nil {
return err
}
return repository.DeletePushChannel(ctx, &channel)
}
func loadChannelForTest(ctx context.Context, req TestChannelRequest) (string, string, string, string, error) {
if req.Name != "" {
channel, err := repository.GetPushChannelByName(ctx, req.Name)
if err != nil {
return "", "", "", "", errors.New("channel not found")
}
return channel.URL, channel.Token, channel.Other, channel.Type, nil
}
return req.URL, req.Token, req.Other, req.Type, nil
}
func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) {
return repository.ListActivePushEventsByTaskType(ctx, taskType)
}
func loadUserFromPayload(ctx context.Context, data map[string]any) any {
if u, exists := data["user"]; exists && u != nil {
return u
}
if userID, ok := extractUserID(data); ok && userID > 0 {
if user, err := repository.GetUserByID(ctx, userID); err == nil {
return &user
}
}
if username := extractUsername(data); username != "" {
if user, err := repository.GetUserByUsername(ctx, username); err == nil {
return &user
}
}
return nil
}
func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg string) error {
title := req.Body.Title
content := req.Body.Content
level := req.Body.Level
if title == "" {
title = "系统通知"
}
if level == "" {
level = defaultLevelInfo
}
target := req.Target
if target == "" {
if req.Config.URL != "" {
target = req.Config.URL
const maxTargetLen = 50
const truncatedLen = 47
if len(target) > maxTargetLen {
target = target[:truncatedLen] + "..."
}
} else {
target = "default"
}
}
history := model.PushHistory{
EventKey: req.EventKey,
Channel: req.Config.Channel,
Target: target,
Title: title,
Content: content,
Level: level,
Status: status,
ErrorMsg: errMsg,
}
return repository.CreatePushHistory(ctx, &history)
}
func resolveTargetUser(ctx context.Context, resolved string, _ string) (model.User, bool) {
found := false
var user model.User
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if u, err := repository.GetUserByID(ctx, id); err == nil {
user = u
found = true
}
}
if !found {
if u, err := repository.GetUserByUsername(ctx, resolved); err == nil {
user = u
found = true
}
}
return user, found
}
func resolveSystemTarget(ctx context.Context, resolved string, channel string) (string, bool) {
if resolved != "系统" && resolved != "system" && resolved != "0" {
return "", false
}
adminUser, err := repository.GetFirstAdminUser(ctx)
if err != nil {
return resolved, true
}
if channel == channelEmail && adminUser.Email != "" {
return adminUser.Email, true
}
if channel != channelEmail && adminUser.Username != "" {
return adminUser.Username, true
}
return resolved, true
}
func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) {
if url != "" && token != "" {
return url, token, other
}
smtp := loadSMTPConfig(ctx)
if smtp.Host == "" || smtp.Username == "" {
return url, token, other
}
port := smtp.Port
if port == "" {
port = "587"
}
if url == "" {
url = smtp.Host + ":" + port
}
if token == "" {
token = smtp.Username
}
if other == "" {
other = smtp.Password
}
return url, token, other
}
func getSystemUser(ctx context.Context) *model.User {
user := repository.GetSystemUser(ctx)
return &user
}
func getEventInfo(req CreateEventRequest) (string, string, []byte, error) {
if req.TaskType != "" {
meta := task.GetTaskMetaByAsynqTask(req.TaskType)
if meta == nil {
return "", "", nil, errors.New("unsupported task type")
}
eventKey := "task_completed:" + req.TaskType
eventName := "任务完成: " + meta.Name
defaultTemplate := NotificationMessage{
Title: "任务完成: " + meta.Name,
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
Level: defaultLevelInfo,
}
defaultTemplateBytes, err := json.Marshal(defaultTemplate)
if err != nil {
return "", "", nil, err
}
return eventKey, eventName, defaultTemplateBytes, nil
}
if req.EventKey == "" {
return "", "", nil, errors.New("either event_key or task_type must be provided")
}
meta, found := findBuiltInEvent(req.EventKey)
if !found {
return "", "", nil, errors.New("unsupported built-in event key")
}
defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate)
if err != nil {
return "", "", nil, err
}
return req.EventKey, meta.Name, defaultTemplateBytes, nil
}
+1 -2
View File
@@ -16,7 +16,6 @@ import ("bytes"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/alicebob/miniredis/v2"
"github.com/gin-gonic/gin"
@@ -104,7 +103,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
util.SetToContext(c, "user_obj", authUser)
oauth.SetToContext(c, "user_obj", authUser)
}
c.Next()
})
+35 -245
View File
@@ -4,22 +4,21 @@
// Package push defines push notification HTTP routes.
package push
import ("context"
"encoding/json"
import (
"context"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"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")
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// UpdateEventRequest 更新事件请求参数
type UpdateEventRequest struct {
@@ -37,30 +36,7 @@ type TestPushRequest struct {
// SyncEvents automatically registers/updates built-in events in the database.
func SyncEvents(ctx context.Context) error {
for _, meta := range BuiltInEvents {
var event model.PushEvent
err := db.DB(ctx).Where("event_key = ?", meta.Key).First(&event).Error
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 := db.DB(ctx).Create(&event).Error; err != nil {
return err
}
} else if err != nil {
return err
}
}
return nil
return syncBuiltInEvents(ctx)
}
// ListEvents 获取通知事件列表
@@ -73,9 +49,8 @@ func SyncEvents(ctx context.Context) error {
// @Router /api/v1/admin/push/events [get]
func ListEvents(c *gin.Context) {
ctx := c.Request.Context()
var events []model.PushEvent
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
events, err := listPushEvents(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -113,46 +88,6 @@ func ListBuiltInEvents(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(BuiltInEvents))
}
func getEventInfo(req CreateEventRequest) (string, string, []byte, error) {
if req.TaskType != "" {
// 1. 检查关联任务是否存在
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")
}
// 1. 检查内置事件是否存在
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
}
// CreateEvent 创建通知事件
// @Summary 创建通知事件
// @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限
@@ -170,70 +105,11 @@ func CreateEvent(c *gin.Context) {
return
}
ctx := c.Request.Context()
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
event, err := createPushEvent(c.Request.Context(), req)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
// 2. 检查是否已经创建过该事件的配置
var count int64
if err := db.DB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", eventKey).Count(&count).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if count > 0 {
response.AbortBadRequest(c, "this notification event is already configured")
return
}
// 3. 模板处理
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 {
response.AbortBadRequest(c, "custom template is not a valid JSON format")
return
}
}
// 4. 创建事件记录
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 {
response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(ctx).Create(&event).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
// 缓存一致性:清除旧事件缓存
model.DeleteActivePushEventCache(ctx, event.EventKey)
c.JSON(http.StatusOK, response.OK(event))
}
@@ -247,32 +123,20 @@ func CreateEvent(c *gin.Context) {
// @Success 200 {object} response.Any{data=string} "删除成功"
// @Router /api/v1/admin/push/events/{id} [delete]
func DeleteEvent(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
ctx := c.Request.Context()
var event model.PushEvent
if err := db.DB(ctx).First(&event, id).Error; err != nil {
if err := deletePushEvent(c.Request.Context(), id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
} else {
response.AbortInternal(c, err.Error())
return
}
return
}
if err := db.DB(ctx).Delete(&event).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
// 缓存一致性:清除事件缓存
model.DeleteActivePushEventCache(ctx, event.EventKey)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -288,8 +152,7 @@ func DeleteEvent(c *gin.Context) {
// @Success 200 {object} response.Any{data=string} "修改成功"
// @Router /api/v1/admin/push/events/{id} [put]
func UpdateEvent(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
@@ -301,34 +164,14 @@ func UpdateEvent(c *gin.Context) {
return
}
var event model.PushEvent
if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil {
if err := updatePushEvent(c.Request.Context(), id, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
} else {
response.AbortInternal(c, err.Error())
return
}
return
}
event.Channels = req.Channels
event.Targets = req.Targets
event.Template = req.Template
event.Enabled = req.Enabled
if err := event.Validate(); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(c.Request.Context()).Save(&event).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
// 缓存一致性:清除事件缓存
model.DeleteActivePushEventCache(c.Request.Context(), event.EventKey)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -342,37 +185,22 @@ func UpdateEvent(c *gin.Context) {
// @Success 200 {object} response.Any{data=string} "切换成功"
// @Router /api/v1/admin/push/events/{id}/toggle [post]
func ToggleEvent(c *gin.Context) {
idStr := c.Param("id")
id, err := strconv.ParseUint(idStr, 10, 64)
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, "invalid event id")
return
}
var event model.PushEvent
if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil {
enabled, err := togglePushEvent(c.Request.Context(), id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, "notification event not found")
} else {
response.AbortInternal(c, err.Error())
return
}
response.AbortBadRequest(c, err.Error())
return
}
event.Enabled = !event.Enabled
if event.Enabled && len(event.Channels) == 0 {
response.AbortBadRequest(c, "cannot enable event without any push channels configured")
return
}
if err := db.DB(c.Request.Context()).Model(&event).Update("enabled", event.Enabled).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
// 缓存一致性:清除事件缓存
model.DeleteActivePushEventCache(c.Request.Context(), event.EventKey)
c.JSON(http.StatusOK, response.OK(event.Enabled))
c.JSON(http.StatusOK, response.OK(enabled))
}
// pushHistoriesResponse 推送历史分页响应
@@ -396,37 +224,22 @@ type pushHistoriesResponse struct {
// @Success 200 {object} response.Any{data=pushHistoriesResponse} "推送历史列表"
// @Router /api/v1/admin/push/histories [get]
func ListHistories(c *gin.Context) {
pageStr := c.DefaultQuery("page", "1")
pageSizeStr := c.DefaultQuery("page_size", "20")
eventKey := c.Query("event_key")
status := c.Query("status")
page, err := strconv.Atoi(pageStr)
if err != nil || page < 1 {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
if page < 1 {
page = 1
}
pageSize, err := strconv.Atoi(pageSizeStr)
if err != nil || pageSize < 1 {
if pageSize < 1 {
pageSize = 20
}
query := db.DB(c.Request.Context()).Model(&model.PushHistory{}).Order("created_at DESC")
if eventKey != "" {
query = query.Where("event_key = ?", eventKey)
}
if status != "" {
query = query.Where("status = ?", status)
}
var total int64
if err := query.Count(&total).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
var results []model.PushHistory
offset := (page - 1) * pageSize
if err := query.Offset(offset).Limit(pageSize).Find(&results).Error; err != nil {
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
}
@@ -459,44 +272,21 @@ func TestPush(c *gin.Context) {
response.AbortBadRequest(c, err.Error())
return
}
// 校验配置
if err := pusher.ValidateConfig(req.Config); err != nil {
response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err))
return
}
// 邮件渠道需要从系统设置中拉取发件人 SMTP 信息做测试 (除非配了独立的)
if req.Config.Channel == channelEmail && (req.Config.URL == "" || req.Config.Key == "") {
var smtpHost, smtpPort, smtpUser, smtpPass model.SystemConfig
ctx := c.Request.Context()
_ = smtpHost.GetByKey(ctx, model.ConfigKeySMTPHost)
_ = smtpPort.GetByKey(ctx, model.ConfigKeySMTPPort)
_ = smtpUser.GetByKey(ctx, model.ConfigKeySMTPUsername)
_ = smtpPass.GetByKey(ctx, model.ConfigKeySMTPPassword)
if smtpHost.Value != "" && smtpUser.Value != "" {
port := smtpPort.Value
if port == "" {
port = "587"
}
req.Config.URL = smtpHost.Value + ":" + port
req.Config.Key = smtpUser.Value
req.Config.Secret = smtpPass.Value
}
}
applySMTPFallbackToPushConfig(c.Request.Context(), &req.Config)
testBody := map[string]any{
keyTitle: "测试通道推送",
keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。",
keyLevel: defaultLevelInfo,
}
err = pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil)
if err != nil {
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())
}
}
+4 -37
View File
@@ -9,7 +9,6 @@ import (
"strconv"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
@@ -20,21 +19,16 @@ func RegisterTaskListeners() {
task.OnTaskCompleted(handleTaskCompleted)
}
// handleTaskCompleted handles task completions and triggers appropriate push events.
func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.TaskResult, execErr error) {
// Query all active push events configured for this task type
var events []model.PushEvent
err := db.DB(ctx).Where("task_type = ? AND enabled = ?", execution.TaskType, true).Find(&events).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
}
// Build the notification body context
body := map[string]any{
"task_id": execution.TaskID,
"task_name": execution.TaskName,
@@ -43,20 +37,17 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re
"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"] = ""
}
// Parse payload parameters if it is valid JSON
var payloadMap map[string]any
if execution.Payload != "" {
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
@@ -64,8 +55,6 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re
extractUserFromMap(ctx, payloadMap, body)
}
}
// Parse result detail parameters if it is valid JSON
if result != nil && result.Detail != "" {
var detailMap map[string]any
if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil {
@@ -74,7 +63,6 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re
}
}
// Trigger notifications for all configured events
for _, event := range events {
meta := EventMetadata{
Key: event.EventKey,
@@ -85,35 +73,15 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re
}
}
// extractUserFromMap tries to find user information from a map and load the full User model.
func extractUserFromMap(ctx context.Context, data map[string]any, body map[string]any) {
if u, exists := body["user"]; exists && u != nil {
return
}
if uVal, ok := data["user"]; ok && uVal != nil {
body["user"] = uVal
return
}
if userID, ok := extractUserID(data); ok && userID > 0 {
var u model.User
if err := db.DB(ctx).Where("id = ?", userID).First(&u).Error; err == nil {
body["user"] = &u
return
}
}
if username := extractUsername(data); username != "" {
var u model.User
if err := db.DB(ctx).Where("username = ?", username).First(&u).Error; err == nil {
body["user"] = &u
return
}
if user := loadUserFromPayload(ctx, data); user != nil {
body["user"] = user
}
}
// extractUserID extracts and validates a userID from map fields.
func extractUserID(data map[string]any) (uint64, bool) {
for _, k := range []string{"user_id", "userId", "uid"} {
val, ok := data[k]
@@ -144,7 +112,6 @@ func extractUserID(data map[string]any) (uint64, bool) {
return 0, false
}
// extractUsername extracts a username string from map fields.
func extractUsername(data map[string]any) string {
for _, k := range []string{"username", "user_name"} {
if val, ok := data[k]; ok && val != nil {
@@ -154,4 +121,4 @@ func extractUsername(data map[string]any) string {
}
}
return ""
}
}
+1 -43
View File
@@ -9,10 +9,7 @@ import (
"encoding/json"
"errors"
"fmt"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/push"
)
@@ -118,46 +115,7 @@ func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*task.TaskRe
}
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
title := req.Body.Title
content := req.Body.Content
level := req.Body.Level
if title == "" {
title = "系统通知"
}
if level == "" {
level = defaultLevelInfo
}
target := req.Target
if target == "" {
// 如果目标人为空 (例如 webhook bot),用其地址填充前缀或默认词作为归档
if req.Config.URL != "" {
target = req.Config.URL
// 隐藏敏感 URL 细节
//nolint:mnd
if len(target) > 50 {
target = target[:47] + "..."
}
} 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,
CreatedAt: time.Now(),
}
// 记录到数据库
if dbErr := db.DB(ctx).Create(&history).Error; dbErr != nil {
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
}
}
+128
View File
@@ -0,0 +1,128 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package system_config
import (
"context"
"encoding/json"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
)
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
exists, err := repository.SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
if exists {
return errors.New(ConfigKeyExists)
}
config := model.SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := repository.CreateSystemConfig(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
return nil
}
func listSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
return repository.ListAdminSystemConfigs(ctx, configType)
}
func getSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) {
return repository.GetAdminSystemConfigByKey(ctx, key)
}
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return err
}
var originalDriver storage.Driver
if key == model.ConfigKeyStorageConfig {
var currentCfg storage.Config
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value)
if err != nil {
return err
}
req.Value = validatedVal
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
updates := map[string]any{
"description": req.Description,
}
if req.Visibility != nil {
updates["visibility"] = *req.Visibility
config.Visibility = *req.Visibility
}
if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue {
updates["value"] = req.Value
config.Value = req.Value
}
if err := tx.Model(&config).Updates(updates).Error; err != nil {
return err
}
resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value)
return nil
}); err != nil {
return err
}
invalidateCachesAfterConfigUpdate(ctx, key)
return nil
}
func resolveStorageMigrationTasksOnDirectDriverUpdate(
ctx context.Context,
tx *gorm.DB,
key string,
originalDriver storage.Driver,
newValue string,
) {
if key != model.ConfigKeyStorageConfig || originalDriver == "" {
return
}
var newCfg storage.Config
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
return
}
if newCfg.Driver != originalDriver {
return
}
if err := tx.Model(&model.TaskExecution{}).
Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed).
Updates(map[string]any{
"status": model.TaskExecutionStatusSucceeded,
"result": "存储配置直接更新,故障迁移任务自动标记为已解决",
"finished_at": time.Now(),
}).Error; err != nil {
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
}
}
+25 -127
View File
@@ -10,7 +10,7 @@ import (
"errors"
"fmt"
"net/http"
"time"
"strings"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
@@ -18,8 +18,8 @@ import (
"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"
@@ -64,42 +64,15 @@ func CreateSystemConfig(c *gin.Context) {
return
}
// 检查配置键是否已存在
var existing model.SystemConfig
if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil {
response.AbortBadRequest(c, ConfigKeyExists)
return
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortInternal(c, err.Error())
return
}
config := model.SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error {
// 创建配置
if err := tx.Create(&config).Error; err != nil {
return err
if err := createSystemConfig(c.Request.Context(), req); err != nil {
if err.Error() == ConfigKeyExists {
response.AbortBadRequest(c, ConfigKeyExists)
return
}
return nil
}); err != nil {
response.AbortInternal(c, err.Error())
return
}
invalidateSystemConfigCaches(c.Request.Context(), req.Key)
if err := model.InvalidateVisibleSystemConfigsCache(c.Request.Context()); err != nil {
logger.WarnF(c.Request.Context(), "清理公共配置列表缓存失败: %v", err)
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -116,14 +89,8 @@ func CreateSystemConfig(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs [get]
func ListSystemConfigs(c *gin.Context) {
configType := c.Query("type")
query := db.DB(c.Request.Context()).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
var configs []model.SystemConfig
if err := query.Find(&configs).Error; err != nil {
configs, err := listSystemConfigs(c.Request.Context(), c.Query("type"))
if err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -149,8 +116,8 @@ func ListSystemConfigs(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/system-configs/{key} [get]
func GetSystemConfig(c *gin.Context) {
var config model.SystemConfig
if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&config).Error; err != nil {
config, err := getSystemConfig(c.Request.Context(), c.Param("key"))
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
} else {
@@ -188,101 +155,24 @@ func UpdateSystemConfig(c *gin.Context) {
}
key := c.Param("key")
// 检查配置是否存在
var config model.SystemConfig
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&config).Error; err != nil {
if err := updateSystemConfig(c.Request.Context(), key, req); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, SystemConfigNotFound)
} else {
response.AbortInternal(c, err.Error())
return
}
return
}
var originalDriver storage.Driver
if key == model.ConfigKeyStorageConfig {
var currentCfg storage.Config
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
validatedVal, err := validateAndMergeStorageConfig(c.Request.Context(), req.Value, config.Value)
if err != nil {
if isStorageConfigValidationError(err) {
response.AbortBadRequest(c, err.Error())
return
}
req.Value = validatedVal
}
if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error {
// 更新配置
updates := map[string]interface{}{
"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(
c.Request.Context(),
tx,
key,
originalDriver,
req.Value,
)
return nil
}); err != nil {
response.AbortInternal(c, err.Error())
return
}
invalidateCachesAfterConfigUpdate(c.Request.Context(), key)
c.JSON(http.StatusOK, response.OKNil())
}
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)
}
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := model.InvalidateSystemConfigCache(ctx, key); err != nil {
if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
if cap.IsRuntimeConfigKey(key) {
@@ -304,7 +194,7 @@ func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
upload.PublishAccessCacheInvalidation(ctx)
}
if err := model.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
}
@@ -345,8 +235,7 @@ func TestSMTP(c *gin.Context) {
password := req.SMTPPassword
if password == maskedConfigValue {
var sc model.SystemConfig
if err := sc.GetByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil {
if sc, err := repository.GetSystemConfigByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil {
password = sc.Value
}
}
@@ -375,6 +264,15 @@ func TestSMTP(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(resp))
}
func isStorageConfigValidationError(err error) bool {
msg := err.Error()
return 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
@@ -4,7 +4,8 @@
package system_config
import ("bufio"
import (
"bufio"
"bytes"
"context"
"encoding/json"
@@ -17,24 +18,24 @@ import ("bufio"
"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/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response")
"github.com/Rain-kl/Wavelet/internal/common/response"
)
const expectedDefaultConfigsCount = 30
func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
util.SetToContext(c, oauth.UserObjKey, authUser)
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
@@ -86,22 +87,22 @@ func TestCreateSystemConfig(t *testing.T) {
// Verify caches are invalidated after create and repopulate on read
_, err = db.Redis.HGet(
context.Background(),
db.PrefixedKey(model.SystemConfigRedisHashKey),
db.PrefixedKey(repository.SystemConfigRedisHashKey),
"custom_key",
).Result()
if err == nil {
t.Fatal("expected redis cache miss immediately after create")
}
var loaded model.SystemConfig
if err := loaded.GetByKey(context.Background(), "custom_key"); err != nil {
t.Fatalf("GetByKey(custom_key) error = %v", err)
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("GetByKey(custom_key).Value = %q, want %q", loaded.Value, "custom_value")
t.Errorf("GetSystemConfigByKey(custom_key).Value = %q, want %q", loaded.Value, "custom_value")
}
if loaded.Visibility != model.ConfigVisibilityVisible {
t.Errorf("GetByKey(custom_key).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityVisible)
t.Errorf("GetSystemConfigByKey(custom_key).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityVisible)
}
})
@@ -245,22 +246,22 @@ func TestUpdateSystemConfig(t *testing.T) {
// Verify caches are invalidated after update and repopulate on read
_, err := db.Redis.HGet(
context.Background(),
db.PrefixedKey(model.SystemConfigRedisHashKey),
db.PrefixedKey(repository.SystemConfigRedisHashKey),
model.ConfigKeySiteName,
).Result()
if err == nil {
t.Fatal("expected redis cache miss immediately after update")
}
var loaded model.SystemConfig
if err := loaded.GetByKey(context.Background(), model.ConfigKeySiteName); err != nil {
t.Fatalf("GetByKey(site_name) error = %v", err)
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("GetByKey(site_name).Value = %q, want %q", 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("GetByKey(site_name).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityHidden)
t.Errorf("GetSystemConfigByKey(site_name).Visibility = %d, want %d", loaded.Visibility, model.ConfigVisibilityHidden)
}
})
+1 -2
View File
@@ -20,7 +20,6 @@ import ("bytes"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
"github.com/stretchr/testify/assert"
@@ -51,7 +50,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
util.SetToContext(c, oauth.UserObjKey, authUser)
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
+78
View File
@@ -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)
}
+32 -88
View File
@@ -3,15 +3,15 @@
package template
import ("errors"
import (
"errors"
"net/http"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response")
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// CreateTemplateRequest 创建模板请求
type CreateTemplateRequest struct {
@@ -32,6 +32,24 @@ type UpdateTemplateRequest struct {
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)不可重复,需要管理员权限
@@ -53,33 +71,8 @@ func CreateTemplate(c *gin.Context) {
return
}
// 检查模板 Key 是否已存在
var existing model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil {
response.AbortBadRequest(c, TemplateKeyExists)
return
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortInternal(c, err.Error())
return
}
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 {
response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(c.Request.Context()).Create(&tmpl).Error; err != nil {
response.AbortInternal(c, err.Error())
tmpl, err := createTemplate(c.Request.Context(), req)
if abortTemplateLogicError(c, err) {
return
}
@@ -98,8 +91,8 @@ func CreateTemplate(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates [get]
func ListTemplates(c *gin.Context) {
var templates []model.Template
if err := db.DB(c.Request.Context()).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
templates, err := listTemplates(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -121,13 +114,8 @@ func ListTemplates(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [get]
func GetTemplate(c *gin.Context) {
var tmpl model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&tmpl).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, TemplateNotFound)
} else {
response.AbortInternal(c, err.Error())
}
tmpl, err := getTemplate(c.Request.Context(), c.Param("key"))
if abortTemplateLogicError(c, err) {
return
}
@@ -157,32 +145,8 @@ func UpdateTemplate(c *gin.Context) {
return
}
key := c.Param("key")
// 检查模板是否存在
var tmpl model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&tmpl).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, TemplateNotFound)
} else {
response.AbortInternal(c, err.Error())
}
return
}
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 {
response.AbortBadRequest(c, err.Error())
return
}
if err := db.DB(c.Request.Context()).Save(&tmpl).Error; err != nil {
response.AbortInternal(c, err.Error())
tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req)
if abortTemplateLogicError(c, err) {
return
}
@@ -204,29 +168,9 @@ func UpdateTemplate(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/templates/{key} [delete]
func DeleteTemplate(c *gin.Context) {
key := c.Param("key")
// 检查模板是否存在
var tmpl model.Template
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&tmpl).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.AbortNotFound(c, TemplateNotFound)
} else {
response.AbortInternal(c, err.Error())
}
return
}
// 限制系统模板删除
if tmpl.IsSystem {
response.AbortBadRequest(c, SystemTemplateCannotDelete)
return
}
if err := db.DB(c.Request.Context()).Delete(&tmpl).Error; err != nil {
response.AbortInternal(c, err.Error())
if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
}
}
+2 -4
View File
@@ -12,20 +12,18 @@ import ("bytes"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response")
func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
util.SetToContext(c, oauth.UserObjKey, authUser)
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})
+6 -5
View File
@@ -23,6 +23,7 @@ import (
"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"
)
@@ -225,19 +226,19 @@ func (m *manager) fetchRelease(ctx context.Context, repository string) (githubRe
}
func loadRepository(ctx context.Context) (string, error) {
var config model.SystemConfig
if err := config.GetByKey(ctx, model.ConfigKeyUpdateUpstreamRepository); err != nil {
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) {
repository, err := loadRepository(ctx)
upstreamRepo, err := loadRepository(ctx)
if err != nil {
return Status{}, releaseAsset{}, err
}
release, asset, err := m.fetchRelease(ctx, repository)
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
if err != nil {
return Status{}, releaseAsset{}, err
}
@@ -259,7 +260,7 @@ func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) {
ReleaseNotes: release.Body,
ReleaseURL: release.HTMLURL,
PublishedAt: release.Published.Format(time.RFC3339),
UpstreamRepository: repository,
UpstreamRepository: upstreamRepo,
AssetName: asset.Name,
Platform: runtime.GOOS + "/" + runtime.GOARCH,
}, asset, nil
+106
View File
@@ -0,0 +1,106 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package user
import (
"context"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
)
func listUsers(ctx context.Context, req listUsersRequest) (int64, []model.User, error) {
return repository.ListAdminUsers(ctx, repository.AdminUserListFilter{
UserID: req.UserID,
Username: strings.TrimSpace(req.Username),
Page: req.Page,
PageSize: req.PageSize,
})
}
func getUserDetail(ctx context.Context, id uint64) (model.User, error) {
return repository.GetAdminUserDetail(ctx, id)
}
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
flags, err := repository.GetUserAdminFlags(ctx, id)
if err != nil {
return err
}
if !active && flags.IsAdmin {
return errors.New(cannotDisable)
}
return repository.UpdateUserActive(ctx, id, active)
}
func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
if currentUserID == targetID {
return errors.New(cannotDeleteSelf)
}
flags, err := repository.GetUserAdminFlags(ctx, targetID)
if err != nil {
return err
}
if flags.IsAdmin {
return errors.New(cannotDelete)
}
return repository.DeleteUserWithRelations(ctx, targetID)
}
func createUser(ctx context.Context, req createUserRequest) (model.User, error) {
req.Username = strings.TrimSpace(req.Username)
req.Nickname = strings.TrimSpace(req.Nickname)
req.Password = strings.TrimSpace(req.Password)
req.Email = strings.TrimSpace(req.Email)
if req.Username == "" {
return model.User{}, errors.New(usernameRequired)
}
if req.Email == "" {
return model.User{}, errors.New(emailRequired)
}
if len(req.Password) < minPasswordLength {
return model.User{}, errors.New(passwordTooShort)
}
count, err := repository.CountUsersByUsername(ctx, req.Username)
if err != nil {
return model.User{}, err
}
if count > 0 {
return model.User{}, errors.New(usernameExists)
}
emailCount, err := repository.CountUsersByEmail(ctx, req.Email)
if err != nil {
return model.User{}, err
}
if emailCount > 0 {
return model.User{}, errors.New(emailExists)
}
newUser := model.User{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
LastLoginAt: time.Time{},
}
if newUser.Nickname == "" {
newUser.Nickname = req.Username
}
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
return model.User{}, err
}
if err := repository.CreateUser(ctx, &newUser); err != nil {
return model.User{}, err
}
return newUser, nil
}
+38 -163
View File
@@ -5,16 +5,13 @@
package user
import (
"errors"
"net/http"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
@@ -85,6 +82,31 @@ func toUser(u model.User) user {
}
}
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 和用户名筛选,需要管理员权限
@@ -98,7 +120,6 @@ func toUser(u model.User) user {
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/users [get]
// ListUsers 获取用户列表
func ListUsers(c *gin.Context) {
var req listUsersRequest
if err := c.ShouldBindQuery(&req); err != nil {
@@ -106,34 +127,8 @@ func ListUsers(c *gin.Context) {
return
}
var modelUsers []model.User
var total int64
query := db.DB(c.Request.Context()).Model(&model.User{})
username := strings.TrimSpace(req.Username)
if req.UserID != nil {
query = query.Where("id = ?", *req.UserID)
}
if username != "" {
query = query.Where("username LIKE ?", username+"%")
}
if err := query.Count(&total).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
offset := (req.Page - 1) * req.PageSize
if err := query.
Select("id, username, nickname, avatar_url, is_active, is_admin, " +
"last_login_at, created_at, updated_at").
Order("id DESC").
Offset(offset).
Limit(req.PageSize).
Find(&modelUsers).Error; err != nil {
total, modelUsers, err := listUsers(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -169,17 +164,8 @@ func GetUser(c *gin.Context) {
return
}
var targetUser model.User
if err := db.DB(c.Request.Context()).
Select("id, username, nickname, email, avatar_url, is_active, is_admin, "+
"bio, phone, gender, website, location, last_login_at, created_at, updated_at").
Where("id = ?", id).
First(&targetUser).Error; err != nil {
if err == gorm.ErrRecordNotFound {
response.AbortNotFound(c, userNotFound)
return
}
response.AbortInternal(c, err.Error())
targetUser, err := getUserDetail(c.Request.Context(), id)
if abortUserLogicError(c, err, userNotFound, nil, nil) {
return
}
@@ -219,32 +205,10 @@ func UpdateUserStatus(c *gin.Context) {
return
}
var targetUser struct {
ID uint64 `gorm:"column:id"`
IsAdmin bool `gorm:"column:is_admin"`
}
if err := db.DB(c.Request.Context()).
Model(&model.User{}).
Select("id, is_admin").
Where("id = ?", id).
First(&targetUser).Error; err != nil {
if err == gorm.ErrRecordNotFound {
response.AbortNotFound(c, userNotFound)
if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) {
return
}
response.AbortInternal(c, err.Error())
return
}
if !req.IsActive && targetUser.IsAdmin {
response.AbortForbidden(c, cannotDisable)
return
}
if err := db.DB(c.Request.Context()).
Model(&model.User{}).
Where("id = ?", id).
Update("is_active", req.IsActive).Error; err != nil {
response.AbortInternal(c, updateUserFailed)
return
}
@@ -272,43 +236,11 @@ func DeleteUser(c *gin.Context) {
return
}
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
if currUser != nil && currUser.ID == id {
response.AbortForbidden(c, cannotDeleteSelf)
return
}
var targetUser struct {
ID uint64 `gorm:"column:id"`
IsAdmin bool `gorm:"column:is_admin"`
}
if err := db.DB(c.Request.Context()).
Model(&model.User{}).
Select("id, is_admin").
Where("id = ?", id).
First(&targetUser).Error; err != nil {
if err == gorm.ErrRecordNotFound {
response.AbortNotFound(c, userNotFound)
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, err.Error())
return
}
if targetUser.IsAdmin {
response.AbortForbidden(c, cannotDelete)
return
}
if err := db.DB(c.Request.Context()).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("user_id = ?", id).Delete(&model.AccessToken{}).Error; err != nil {
return err
}
if err := tx.Where("user_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil {
return err
}
return tx.Where("id = ?", id).Delete(&model.User{}).Error
}); err != nil {
response.AbortInternal(c, deleteUserFailed)
return
}
@@ -347,67 +279,10 @@ func CreateUser(c *gin.Context) {
return
}
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 == "" {
response.AbortBadRequest(c, usernameRequired)
return
}
if req.Email == "" {
response.AbortBadRequest(c, emailRequired)
return
}
if len(req.Password) < minPasswordLength {
response.AbortBadRequest(c, passwordTooShort)
return
}
ctx := c.Request.Context()
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", req.Username).Count(&count).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if count > 0 {
response.AbortBadRequest(c, usernameExists)
return
}
var emailCount int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if emailCount > 0 {
response.AbortBadRequest(c, emailExists)
return
}
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 {
response.AbortInternal(c, err.Error())
return
}
if err := db.DB(ctx).Create(&newUser).Error; err != nil {
response.AbortInternal(c, err.Error())
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)))
}
}
+2 -4
View File
@@ -14,20 +14,18 @@ import ("bytes"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response")
func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware
adminGroup.Use(func(c *gin.Context) {
if authUser != nil {
util.SetToContext(c, oauth.UserObjKey, authUser)
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
})