mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
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:
@@ -13,7 +13,6 @@ import ("bytes"
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response")
|
||||||
@@ -26,7 +25,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
|
|||||||
// Mock authentication middleware
|
// Mock authentication middleware
|
||||||
adminGroup.Use(func(c *gin.Context) {
|
adminGroup.Use(func(c *gin.Context) {
|
||||||
if authUser != nil {
|
if authUser != nil {
|
||||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
|
|||||||
Vendored
+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)
|
||||||
|
}
|
||||||
Vendored
+4
-37
@@ -4,18 +4,16 @@
|
|||||||
// Package cache provides HTTP handlers for managing disk cache.
|
// Package cache provides HTTP handlers for managing disk cache.
|
||||||
package cache
|
package cache
|
||||||
|
|
||||||
import ("context"
|
import (
|
||||||
"errors"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/diskcache"
|
"github.com/Rain-kl/Wavelet/internal/diskcache"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/gin-gonic/gin"
|
"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 {
|
type updateCacheConfigRequest struct {
|
||||||
MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"`
|
MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"`
|
||||||
@@ -62,25 +60,21 @@ func UpdateCacheConfig(c *gin.Context) {
|
|||||||
|
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
// Update Max Size
|
|
||||||
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
|
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update Default TTL
|
|
||||||
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
|
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update LRU Enabled
|
|
||||||
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
|
if err := saveOrUpdateConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Trigger hot reloading in global cache
|
|
||||||
diskcache.GetGlobalCache().ReloadConfig(ctx)
|
diskcache.GetGlobalCache().ReloadConfig(ctx)
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
@@ -103,31 +97,4 @@ func ClearCache(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
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)
|
|
||||||
}
|
|
||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
)
|
)
|
||||||
|
|
||||||
// getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击
|
// getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击
|
||||||
@@ -32,8 +33,7 @@ func getUpgrader() *websocket.Upgrader {
|
|||||||
|
|
||||||
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
|
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
|
||||||
ctx := r.Context()
|
ctx := r.Context()
|
||||||
var sc model.SystemConfig
|
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
|
|
||||||
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
|
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
|
||||||
allowedOrigins := strings.Split(sc.Value, ",")
|
allowedOrigins := strings.Split(sc.Value, ",")
|
||||||
for _, allowed := range allowedOrigins {
|
for _, allowed := range allowedOrigins {
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ package admin
|
|||||||
import (
|
import (
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
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")
|
ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired")
|
||||||
defer span.End()
|
defer span.End()
|
||||||
|
|
||||||
user, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
user, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
|
|
||||||
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
|
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
|
||||||
if tokenAuth, _ := util.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
|
if tokenAuth, _ := oauth.GetFromContext[bool](c, oauth.TokenAuthKey); tokenAuth {
|
||||||
tokenAdmin, _ := util.GetFromContext[bool](c, oauth.TokenAdminKey)
|
tokenAdmin, _ := oauth.GetFromContext[bool](c, oauth.TokenAdminKey)
|
||||||
if !tokenAdmin {
|
if !tokenAdmin {
|
||||||
response.AbortNotFound(c, TokenAdminRequired)
|
response.AbortNotFound(c, TokenAdminRequired)
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -3,19 +3,20 @@
|
|||||||
|
|
||||||
package push
|
package push
|
||||||
|
|
||||||
import ("encoding/json"
|
import (
|
||||||
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
|
||||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"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 获取各种消息通道的表单配置定义列表
|
// ListChannelDefinitions 获取各种消息通道的表单配置定义列表
|
||||||
// @Summary 获取所有消息通道配置字段定义
|
// @Summary 获取所有消息通道配置字段定义
|
||||||
@@ -38,9 +39,8 @@ func ListChannelDefinitions(c *gin.Context) {
|
|||||||
// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表"
|
// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表"
|
||||||
// @Router /api/v1/admin/push/channels [get]
|
// @Router /api/v1/admin/push/channels [get]
|
||||||
func ListChannels(c *gin.Context) {
|
func ListChannels(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
channels, err := listPushChannels(c.Request.Context())
|
||||||
var channels []model.PushChannel
|
if err != nil {
|
||||||
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -75,41 +75,11 @@ func CreateChannel(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
channel, err := createPushChannel(c.Request.Context(), req)
|
||||||
|
if err != nil {
|
||||||
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 {
|
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
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))
|
c.JSON(http.StatusOK, response.OK(channel))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -135,8 +105,7 @@ type UpdateChannelRequest struct {
|
|||||||
// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功"
|
// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功"
|
||||||
// @Router /api/v1/admin/push/channels/{id} [put]
|
// @Router /api/v1/admin/push/channels/{id} [put]
|
||||||
func UpdateChannel(c *gin.Context) {
|
func UpdateChannel(c *gin.Context) {
|
||||||
idStr := c.Param("id")
|
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, "invalid channel id")
|
response.AbortBadRequest(c, "invalid channel id")
|
||||||
return
|
return
|
||||||
@@ -148,10 +117,8 @@ func UpdateChannel(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
channel, err := updatePushChannel(c.Request.Context(), id, req)
|
||||||
|
if err != nil {
|
||||||
var channel model.PushChannel
|
|
||||||
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
response.AbortNotFound(c, "channel not found")
|
response.AbortNotFound(c, "channel not found")
|
||||||
return
|
return
|
||||||
@@ -159,27 +126,6 @@ func UpdateChannel(c *gin.Context) {
|
|||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
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))
|
c.JSON(http.StatusOK, response.OK(channel))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -193,16 +139,13 @@ func UpdateChannel(c *gin.Context) {
|
|||||||
// @Success 200 {object} response.Any "删除成功"
|
// @Success 200 {object} response.Any "删除成功"
|
||||||
// @Router /api/v1/admin/push/channels/{id} [delete]
|
// @Router /api/v1/admin/push/channels/{id} [delete]
|
||||||
func DeleteChannel(c *gin.Context) {
|
func DeleteChannel(c *gin.Context) {
|
||||||
idStr := c.Param("id")
|
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, "invalid channel id")
|
response.AbortBadRequest(c, "invalid channel id")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
if err := deletePushChannel(c.Request.Context(), id); err != nil {
|
||||||
var channel model.PushChannel
|
|
||||||
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
response.AbortNotFound(c, "channel not found")
|
response.AbortNotFound(c, "channel not found")
|
||||||
return
|
return
|
||||||
@@ -210,15 +153,6 @@ func DeleteChannel(c *gin.Context) {
|
|||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
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())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -250,26 +184,12 @@ func TestChannel(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
var url, token, other, channelType string
|
url, token, other, channelType, err := loadChannelForTest(ctx, req)
|
||||||
|
if err != nil {
|
||||||
if req.Name != "" {
|
response.AbortBadRequest(c, err.Error())
|
||||||
var channel model.PushChannel
|
return
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 对邮件类型应用全局配置作为回退
|
|
||||||
if channelType == channelEmail {
|
if channelType == channelEmail {
|
||||||
url, token, other = resolveSMTPConfig(ctx, url, token, other)
|
url, token, other = resolveSMTPConfig(ctx, url, token, other)
|
||||||
}
|
}
|
||||||
@@ -282,7 +202,6 @@ func TestChannel(c *gin.Context) {
|
|||||||
Type: channelType,
|
Type: channelType,
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := tempChannel.Validate(); err != nil {
|
if err := tempChannel.Validate(); err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
@@ -291,34 +210,16 @@ func TestChannel(c *gin.Context) {
|
|||||||
|
|
||||||
var config pkgpush.Config
|
var config pkgpush.Config
|
||||||
var renderedJSON string
|
var renderedJSON string
|
||||||
|
|
||||||
switch channelType {
|
switch channelType {
|
||||||
case channelLark:
|
case channelLark:
|
||||||
config = pkgpush.Config{
|
config = pkgpush.Config{Channel: channelLark, URL: url, Secret: token}
|
||||||
Channel: channelLark,
|
|
||||||
URL: url,
|
|
||||||
Secret: token,
|
|
||||||
}
|
|
||||||
renderedJSON = other
|
renderedJSON = other
|
||||||
case channelEmail:
|
case channelEmail:
|
||||||
config = pkgpush.Config{
|
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
|
||||||
Channel: channelEmail,
|
|
||||||
URL: url,
|
|
||||||
Key: token,
|
|
||||||
Secret: other,
|
|
||||||
}
|
|
||||||
case channelTelegram:
|
case channelTelegram:
|
||||||
config = pkgpush.Config{
|
config = pkgpush.Config{Channel: channelTelegram, URL: url, Secret: token, Key: other}
|
||||||
Channel: channelTelegram,
|
|
||||||
URL: url,
|
|
||||||
Secret: token,
|
|
||||||
Key: other,
|
|
||||||
}
|
|
||||||
default:
|
default:
|
||||||
config = pkgpush.Config{
|
config = pkgpush.Config{Channel: channelCustom, URL: url}
|
||||||
Channel: channelCustom,
|
|
||||||
URL: url,
|
|
||||||
}
|
|
||||||
customPushReq := CustomPushRequest{
|
customPushReq := CustomPushRequest{
|
||||||
Title: "通道测试通知",
|
Title: "通道测试通知",
|
||||||
Content: "这是一条来自系统的消息通道连通性测试消息。",
|
Content: "这是一条来自系统的消息通道连通性测试消息。",
|
||||||
@@ -340,12 +241,10 @@ func TestChannel(c *gin.Context) {
|
|||||||
},
|
},
|
||||||
Template: renderedJSON,
|
Template: renderedJSON,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := enqueuePushTask(ctx, payload); err != nil {
|
if err := enqueuePushTask(ctx, payload); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
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, "$url", escapeJSONString(req.URL))
|
||||||
result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To))
|
result = strings.ReplaceAll(result, "$to", escapeJSONString(req.To))
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
@@ -9,12 +9,9 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"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/internal/task"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||||
@@ -73,30 +70,7 @@ type EventTrigger struct{}
|
|||||||
// DefaultTrigger is the singleton instance of EventTrigger.
|
// DefaultTrigger is the singleton instance of EventTrigger.
|
||||||
var DefaultTrigger = &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.
|
// 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
|
//nolint:contextcheck
|
||||||
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
|
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)
|
body["user"] = getSystemUser(asyncCtx)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 1. Check if the event is enabled (try Redis cache first)
|
eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key)
|
||||||
eventPtr, err := model.GetActivePushEventByKey(asyncCtx, meta.Key)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return
|
return
|
||||||
@@ -119,22 +92,16 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
event := *eventPtr
|
event := *eventPtr
|
||||||
|
|
||||||
if len(event.Channels) == 0 {
|
if len(event.Channels) == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 2. Build and render notification message
|
|
||||||
flatBody := getFlatBody(body)
|
flatBody := getFlatBody(body)
|
||||||
msg, _ := t.buildMessage(&event, meta, flatBody, body)
|
msg, _ := t.buildMessage(&event, meta, flatBody, body)
|
||||||
|
|
||||||
// 3. Enqueue tasks for each matching channel
|
|
||||||
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
|
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) {
|
func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
|
||||||
var msg NotificationMessage
|
var msg NotificationMessage
|
||||||
renderedTemplate := ""
|
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) {
|
func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, msg NotificationMessage, flatBody map[string]any) {
|
||||||
for _, channelName := range event.Channels {
|
for _, channelName := range event.Channels {
|
||||||
// 检查是不是自定义数据库渠道 (使用 Redis 缓存优先)
|
customChannel, err := repository.GetActivePushChannelByName(ctx, channelName)
|
||||||
customChannel, err := model.GetActivePushChannelByName(ctx, channelName)
|
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody)
|
t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err)
|
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 {
|
switch channel.Type {
|
||||||
case channelLark:
|
case channelLark:
|
||||||
config = pkgpush.Config{
|
config = pkgpush.Config{Channel: channelLark, URL: channel.URL, Secret: channel.Token}
|
||||||
Channel: channelLark,
|
renderedTemplate = channel.Other
|
||||||
URL: channel.URL,
|
|
||||||
Secret: channel.Token, // Feishu Bot Sign Secret
|
|
||||||
}
|
|
||||||
renderedTemplate = channel.Other // Optional custom template/card for lark
|
|
||||||
case channelEmail:
|
case channelEmail:
|
||||||
url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other)
|
url, token, other := resolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other)
|
||||||
config = pkgpush.Config{
|
config = pkgpush.Config{Channel: channelEmail, URL: url, Key: token, Secret: other}
|
||||||
Channel: channelEmail,
|
|
||||||
URL: url, // SMTP host:port
|
|
||||||
Key: token, // SMTP Username
|
|
||||||
Secret: other, // SMTP Password
|
|
||||||
}
|
|
||||||
case channelTelegram:
|
case channelTelegram:
|
||||||
config = pkgpush.Config{
|
config = pkgpush.Config{Channel: channelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other}
|
||||||
Channel: channelTelegram,
|
default:
|
||||||
URL: channel.URL,
|
config = pkgpush.Config{Channel: channelCustom, 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,
|
|
||||||
}
|
|
||||||
customPushReq := CustomPushRequest{
|
customPushReq := CustomPushRequest{
|
||||||
Title: msg.Title,
|
Title: msg.Title,
|
||||||
Content: msg.Content,
|
Content: msg.Content,
|
||||||
@@ -301,8 +249,6 @@ func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, m
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
func enqueuePushTask(ctx context.Context, payload SendPayload) error {
|
func enqueuePushTask(ctx context.Context, payload SendPayload) error {
|
||||||
payloadBytes, err := json.Marshal(payload)
|
payloadBytes, err := json.Marshal(payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -348,48 +294,23 @@ func resolveTarget(ctx context.Context, target string, flatBody map[string]any,
|
|||||||
}
|
}
|
||||||
|
|
||||||
resolved := resolveDynamicKeyword(target, flatBody)
|
resolved := resolveDynamicKeyword(target, flatBody)
|
||||||
|
|
||||||
// 2. 如果包含 @,说明已经是个邮箱,直接返回
|
|
||||||
if strings.Contains(resolved, "@") {
|
if strings.Contains(resolved, "@") {
|
||||||
return resolved
|
return resolved
|
||||||
}
|
}
|
||||||
|
|
||||||
// 2.5 如果为特殊的系统虚拟用户,自动映射为首位管理员
|
|
||||||
if val, matched := resolveSystemTarget(ctx, resolved, channel); matched {
|
if val, matched := resolveSystemTarget(ctx, resolved, channel); matched {
|
||||||
return val
|
return val
|
||||||
}
|
}
|
||||||
|
|
||||||
// 3. 不包含 @,说明可能是用户 ID 或用户名。我们需要从数据库中查询对应用户
|
user, found := resolveTargetUser(ctx, resolved, channel)
|
||||||
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 推送渠道进行转换
|
|
||||||
if !found {
|
if !found {
|
||||||
return resolved
|
return resolved
|
||||||
}
|
}
|
||||||
|
|
||||||
if channel == channelEmail && user.Email != "" {
|
if channel == channelEmail && user.Email != "" {
|
||||||
return user.Email
|
return user.Email
|
||||||
}
|
}
|
||||||
|
|
||||||
if channel != channelEmail && user.Username != "" {
|
if channel != channelEmail && user.Username != "" {
|
||||||
return user.Username
|
return user.Username
|
||||||
}
|
}
|
||||||
|
|
||||||
return resolved
|
return resolved
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -418,51 +339,4 @@ func resolveDynamicKeyword(target string, flatBody map[string]any) string {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
return target
|
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
|
|
||||||
}
|
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -16,7 +16,6 @@ import ("bytes"
|
|||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/task"
|
"github.com/Rain-kl/Wavelet/internal/task"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||||
"github.com/alicebob/miniredis/v2"
|
"github.com/alicebob/miniredis/v2"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -104,7 +103,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
|
|||||||
|
|
||||||
adminGroup.Use(func(c *gin.Context) {
|
adminGroup.Use(func(c *gin.Context) {
|
||||||
if authUser != nil {
|
if authUser != nil {
|
||||||
util.SetToContext(c, "user_obj", authUser)
|
oauth.SetToContext(c, "user_obj", authUser)
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -4,22 +4,21 @@
|
|||||||
// Package push defines push notification HTTP routes.
|
// Package push defines push notification HTTP routes.
|
||||||
package push
|
package push
|
||||||
|
|
||||||
import ("context"
|
import (
|
||||||
"encoding/json"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"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/Rain-kl/Wavelet/pkg/push"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
|
)
|
||||||
|
|
||||||
// UpdateEventRequest 更新事件请求参数
|
// UpdateEventRequest 更新事件请求参数
|
||||||
type UpdateEventRequest struct {
|
type UpdateEventRequest struct {
|
||||||
@@ -37,30 +36,7 @@ type TestPushRequest struct {
|
|||||||
|
|
||||||
// SyncEvents automatically registers/updates built-in events in the database.
|
// SyncEvents automatically registers/updates built-in events in the database.
|
||||||
func SyncEvents(ctx context.Context) error {
|
func SyncEvents(ctx context.Context) error {
|
||||||
for _, meta := range BuiltInEvents {
|
return syncBuiltInEvents(ctx)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListEvents 获取通知事件列表
|
// ListEvents 获取通知事件列表
|
||||||
@@ -73,9 +49,8 @@ func SyncEvents(ctx context.Context) error {
|
|||||||
// @Router /api/v1/admin/push/events [get]
|
// @Router /api/v1/admin/push/events [get]
|
||||||
func ListEvents(c *gin.Context) {
|
func ListEvents(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
events, err := listPushEvents(ctx)
|
||||||
var events []model.PushEvent
|
if err != nil {
|
||||||
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -113,46 +88,6 @@ func ListBuiltInEvents(c *gin.Context) {
|
|||||||
c.JSON(http.StatusOK, response.OK(BuiltInEvents))
|
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 创建通知事件
|
// CreateEvent 创建通知事件
|
||||||
// @Summary 创建通知事件
|
// @Summary 创建通知事件
|
||||||
// @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限
|
// @Description 绑定系统内置事件或异步任务、推送渠道、接收目标并创建通知事件配置,需要管理员权限
|
||||||
@@ -170,70 +105,11 @@ func CreateEvent(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
event, err := createPushEvent(c.Request.Context(), req)
|
||||||
|
|
||||||
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
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))
|
c.JSON(http.StatusOK, response.OK(event))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -247,32 +123,20 @@ func CreateEvent(c *gin.Context) {
|
|||||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||||
// @Router /api/v1/admin/push/events/{id} [delete]
|
// @Router /api/v1/admin/push/events/{id} [delete]
|
||||||
func DeleteEvent(c *gin.Context) {
|
func DeleteEvent(c *gin.Context) {
|
||||||
idStr := c.Param("id")
|
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, "invalid event id")
|
response.AbortBadRequest(c, "invalid event id")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
if err := deletePushEvent(c.Request.Context(), id); err != nil {
|
||||||
var event model.PushEvent
|
|
||||||
if err := db.DB(ctx).First(&event, id).Error; err != nil {
|
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
response.AbortNotFound(c, "notification event not found")
|
response.AbortNotFound(c, "notification event not found")
|
||||||
} else {
|
return
|
||||||
response.AbortInternal(c, err.Error())
|
|
||||||
}
|
}
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := db.DB(ctx).Delete(&event).Error; err != nil {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 缓存一致性:清除事件缓存
|
|
||||||
model.DeleteActivePushEventCache(ctx, event.EventKey)
|
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -288,8 +152,7 @@ func DeleteEvent(c *gin.Context) {
|
|||||||
// @Success 200 {object} response.Any{data=string} "修改成功"
|
// @Success 200 {object} response.Any{data=string} "修改成功"
|
||||||
// @Router /api/v1/admin/push/events/{id} [put]
|
// @Router /api/v1/admin/push/events/{id} [put]
|
||||||
func UpdateEvent(c *gin.Context) {
|
func UpdateEvent(c *gin.Context) {
|
||||||
idStr := c.Param("id")
|
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, "invalid event id")
|
response.AbortBadRequest(c, "invalid event id")
|
||||||
return
|
return
|
||||||
@@ -301,34 +164,14 @@ func UpdateEvent(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var event model.PushEvent
|
if err := updatePushEvent(c.Request.Context(), id, req); err != nil {
|
||||||
if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil {
|
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
response.AbortNotFound(c, "notification event not found")
|
response.AbortNotFound(c, "notification event not found")
|
||||||
} else {
|
return
|
||||||
response.AbortInternal(c, err.Error())
|
|
||||||
}
|
}
|
||||||
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())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
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())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -342,37 +185,22 @@ func UpdateEvent(c *gin.Context) {
|
|||||||
// @Success 200 {object} response.Any{data=string} "切换成功"
|
// @Success 200 {object} response.Any{data=string} "切换成功"
|
||||||
// @Router /api/v1/admin/push/events/{id}/toggle [post]
|
// @Router /api/v1/admin/push/events/{id}/toggle [post]
|
||||||
func ToggleEvent(c *gin.Context) {
|
func ToggleEvent(c *gin.Context) {
|
||||||
idStr := c.Param("id")
|
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, "invalid event id")
|
response.AbortBadRequest(c, "invalid event id")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var event model.PushEvent
|
enabled, err := togglePushEvent(c.Request.Context(), id)
|
||||||
if err := db.DB(c.Request.Context()).First(&event, id).Error; err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
response.AbortNotFound(c, "notification event not found")
|
response.AbortNotFound(c, "notification event not found")
|
||||||
} else {
|
return
|
||||||
response.AbortInternal(c, err.Error())
|
|
||||||
}
|
}
|
||||||
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
c.JSON(http.StatusOK, response.OK(enabled))
|
||||||
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))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// pushHistoriesResponse 推送历史分页响应
|
// pushHistoriesResponse 推送历史分页响应
|
||||||
@@ -396,37 +224,22 @@ type pushHistoriesResponse struct {
|
|||||||
// @Success 200 {object} response.Any{data=pushHistoriesResponse} "推送历史列表"
|
// @Success 200 {object} response.Any{data=pushHistoriesResponse} "推送历史列表"
|
||||||
// @Router /api/v1/admin/push/histories [get]
|
// @Router /api/v1/admin/push/histories [get]
|
||||||
func ListHistories(c *gin.Context) {
|
func ListHistories(c *gin.Context) {
|
||||||
pageStr := c.DefaultQuery("page", "1")
|
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||||
pageSizeStr := c.DefaultQuery("page_size", "20")
|
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
|
||||||
eventKey := c.Query("event_key")
|
if page < 1 {
|
||||||
status := c.Query("status")
|
|
||||||
|
|
||||||
page, err := strconv.Atoi(pageStr)
|
|
||||||
if err != nil || page < 1 {
|
|
||||||
page = 1
|
page = 1
|
||||||
}
|
}
|
||||||
pageSize, err := strconv.Atoi(pageSizeStr)
|
if pageSize < 1 {
|
||||||
if err != nil || pageSize < 1 {
|
|
||||||
pageSize = 20
|
pageSize = 20
|
||||||
}
|
}
|
||||||
|
|
||||||
query := db.DB(c.Request.Context()).Model(&model.PushHistory{}).Order("created_at DESC")
|
total, results, err := listPushHistories(c.Request.Context(), repository.PushHistoryListFilter{
|
||||||
if eventKey != "" {
|
EventKey: c.Query("event_key"),
|
||||||
query = query.Where("event_key = ?", eventKey)
|
Status: c.Query("status"),
|
||||||
}
|
Page: page,
|
||||||
if status != "" {
|
PageSize: pageSize,
|
||||||
query = query.Where("status = ?", status)
|
})
|
||||||
}
|
if err != nil {
|
||||||
|
|
||||||
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 {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -459,44 +272,21 @@ func TestPush(c *gin.Context) {
|
|||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 校验配置
|
|
||||||
if err := pusher.ValidateConfig(req.Config); err != nil {
|
if err := pusher.ValidateConfig(req.Config); err != nil {
|
||||||
response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err))
|
response.AbortBadRequest(c, fmt.Sprintf("validation failed: %v", err))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 邮件渠道需要从系统设置中拉取发件人 SMTP 信息做测试 (除非配了独立的)
|
applySMTPFallbackToPushConfig(c.Request.Context(), &req.Config)
|
||||||
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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
testBody := map[string]any{
|
testBody := map[string]any{
|
||||||
keyTitle: "测试通道推送",
|
keyTitle: "测试通道推送",
|
||||||
keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。",
|
keyContent: "当您收到这条消息,说明当前渠道连通性测试通过。",
|
||||||
keyLevel: defaultLevelInfo,
|
keyLevel: defaultLevelInfo,
|
||||||
}
|
}
|
||||||
|
if err := pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil); err != nil {
|
||||||
err = pusher.Send(c.Request.Context(), req.Config, req.Target, testBody, "", nil)
|
|
||||||
if err != nil {
|
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
@@ -9,7 +9,6 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/task"
|
"github.com/Rain-kl/Wavelet/internal/task"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
@@ -20,21 +19,16 @@ func RegisterTaskListeners() {
|
|||||||
task.OnTaskCompleted(handleTaskCompleted)
|
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) {
|
func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.TaskResult, execErr error) {
|
||||||
// Query all active push events configured for this task type
|
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
|
||||||
var events []model.PushEvent
|
|
||||||
err := db.DB(ctx).Where("task_type = ? AND enabled = ?", execution.TaskType, true).Find(&events).Error
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
|
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(events) == 0 {
|
if len(events) == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build the notification body context
|
|
||||||
body := map[string]any{
|
body := map[string]any{
|
||||||
"task_id": execution.TaskID,
|
"task_id": execution.TaskID,
|
||||||
"task_name": execution.TaskName,
|
"task_name": execution.TaskName,
|
||||||
@@ -43,20 +37,17 @@ func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, re
|
|||||||
"task_duration": execution.Duration,
|
"task_duration": execution.Duration,
|
||||||
"time": time.Now().Format("2006-01-02 15:04:05"),
|
"time": time.Now().Format("2006-01-02 15:04:05"),
|
||||||
}
|
}
|
||||||
|
|
||||||
if execErr != nil {
|
if execErr != nil {
|
||||||
body["task_error"] = execErr.Error()
|
body["task_error"] = execErr.Error()
|
||||||
} else {
|
} else {
|
||||||
body["task_error"] = ""
|
body["task_error"] = ""
|
||||||
}
|
}
|
||||||
|
|
||||||
if result != nil {
|
if result != nil {
|
||||||
body["task_result"] = result.Message
|
body["task_result"] = result.Message
|
||||||
} else {
|
} else {
|
||||||
body["task_result"] = ""
|
body["task_result"] = ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse payload parameters if it is valid JSON
|
|
||||||
var payloadMap map[string]any
|
var payloadMap map[string]any
|
||||||
if execution.Payload != "" {
|
if execution.Payload != "" {
|
||||||
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
|
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)
|
extractUserFromMap(ctx, payloadMap, body)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse result detail parameters if it is valid JSON
|
|
||||||
if result != nil && result.Detail != "" {
|
if result != nil && result.Detail != "" {
|
||||||
var detailMap map[string]any
|
var detailMap map[string]any
|
||||||
if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil {
|
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 {
|
for _, event := range events {
|
||||||
meta := EventMetadata{
|
meta := EventMetadata{
|
||||||
Key: event.EventKey,
|
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) {
|
func extractUserFromMap(ctx context.Context, data map[string]any, body map[string]any) {
|
||||||
if u, exists := body["user"]; exists && u != nil {
|
if u, exists := body["user"]; exists && u != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if user := loadUserFromPayload(ctx, data); user != nil {
|
||||||
if uVal, ok := data["user"]; ok && uVal != nil {
|
body["user"] = user
|
||||||
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
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// extractUserID extracts and validates a userID from map fields.
|
|
||||||
func extractUserID(data map[string]any) (uint64, bool) {
|
func extractUserID(data map[string]any) (uint64, bool) {
|
||||||
for _, k := range []string{"user_id", "userId", "uid"} {
|
for _, k := range []string{"user_id", "userId", "uid"} {
|
||||||
val, ok := data[k]
|
val, ok := data[k]
|
||||||
@@ -144,7 +112,6 @@ func extractUserID(data map[string]any) (uint64, bool) {
|
|||||||
return 0, false
|
return 0, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// extractUsername extracts a username string from map fields.
|
|
||||||
func extractUsername(data map[string]any) string {
|
func extractUsername(data map[string]any) string {
|
||||||
for _, k := range []string{"username", "user_name"} {
|
for _, k := range []string{"username", "user_name"} {
|
||||||
if val, ok := data[k]; ok && val != nil {
|
if val, ok := data[k]; ok && val != nil {
|
||||||
@@ -154,4 +121,4 @@ func extractUsername(data map[string]any) string {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
@@ -9,10 +9,7 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"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/internal/task"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/push"
|
"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) {
|
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
|
||||||
title := req.Body.Title
|
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
|
||||||
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 {
|
|
||||||
task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
|
task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -10,7 +10,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"time"
|
"strings"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
@@ -18,8 +18,8 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/cap"
|
"github.com/Rain-kl/Wavelet/internal/apps/cap"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"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/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/storage"
|
"github.com/Rain-kl/Wavelet/internal/storage"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
mail "github.com/Rain-kl/Wavelet/pkg/mail"
|
mail "github.com/Rain-kl/Wavelet/pkg/mail"
|
||||||
@@ -64,42 +64,15 @@ func CreateSystemConfig(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 检查配置键是否已存在
|
if err := createSystemConfig(c.Request.Context(), req); err != nil {
|
||||||
var existing model.SystemConfig
|
if err.Error() == ConfigKeyExists {
|
||||||
if err := db.DB(c.Request.Context()).Where("key = ?", req.Key).First(&existing).Error; err == nil {
|
response.AbortBadRequest(c, ConfigKeyExists)
|
||||||
response.AbortBadRequest(c, ConfigKeyExists)
|
return
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
|
||||||
}); err != nil {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
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())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -116,14 +89,8 @@ func CreateSystemConfig(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/system-configs [get]
|
// @Router /api/v1/admin/system-configs [get]
|
||||||
func ListSystemConfigs(c *gin.Context) {
|
func ListSystemConfigs(c *gin.Context) {
|
||||||
configType := c.Query("type")
|
configs, err := listSystemConfigs(c.Request.Context(), c.Query("type"))
|
||||||
query := db.DB(c.Request.Context()).Order("created_at DESC")
|
if err != nil {
|
||||||
if configType != "" {
|
|
||||||
query = query.Where("type = ?", configType)
|
|
||||||
}
|
|
||||||
|
|
||||||
var configs []model.SystemConfig
|
|
||||||
if err := query.Find(&configs).Error; err != nil {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -149,8 +116,8 @@ func ListSystemConfigs(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/system-configs/{key} [get]
|
// @Router /api/v1/admin/system-configs/{key} [get]
|
||||||
func GetSystemConfig(c *gin.Context) {
|
func GetSystemConfig(c *gin.Context) {
|
||||||
var config model.SystemConfig
|
config, err := getSystemConfig(c.Request.Context(), c.Param("key"))
|
||||||
if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&config).Error; err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
response.AbortNotFound(c, SystemConfigNotFound)
|
response.AbortNotFound(c, SystemConfigNotFound)
|
||||||
} else {
|
} else {
|
||||||
@@ -188,101 +155,24 @@ func UpdateSystemConfig(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
key := c.Param("key")
|
key := c.Param("key")
|
||||||
|
if err := updateSystemConfig(c.Request.Context(), key, req); err != nil {
|
||||||
// 检查配置是否存在
|
|
||||||
var config model.SystemConfig
|
|
||||||
if err := db.DB(c.Request.Context()).Where("key = ?", key).First(&config).Error; err != nil {
|
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
response.AbortNotFound(c, SystemConfigNotFound)
|
response.AbortNotFound(c, SystemConfigNotFound)
|
||||||
} else {
|
return
|
||||||
response.AbortInternal(c, err.Error())
|
|
||||||
}
|
}
|
||||||
return
|
if isStorageConfigValidationError(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(c.Request.Context(), req.Value, config.Value)
|
|
||||||
if err != nil {
|
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
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())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
invalidateCachesAfterConfigUpdate(c.Request.Context(), key)
|
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
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) {
|
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)
|
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
|
||||||
}
|
}
|
||||||
if cap.IsRuntimeConfigKey(key) {
|
if cap.IsRuntimeConfigKey(key) {
|
||||||
@@ -304,7 +194,7 @@ func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
|
|||||||
upload.PublishAccessCacheInvalidation(ctx)
|
upload.PublishAccessCacheInvalidation(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := model.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||||
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -345,8 +235,7 @@ func TestSMTP(c *gin.Context) {
|
|||||||
|
|
||||||
password := req.SMTPPassword
|
password := req.SMTPPassword
|
||||||
if password == maskedConfigValue {
|
if password == maskedConfigValue {
|
||||||
var sc model.SystemConfig
|
if sc, err := repository.GetSystemConfigByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil {
|
||||||
if err := sc.GetByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil {
|
|
||||||
password = sc.Value
|
password = sc.Value
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -375,6 +264,15 @@ func TestSMTP(c *gin.Context) {
|
|||||||
c.JSON(http.StatusOK, response.OK(resp))
|
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 {
|
func maskSensitiveConfig(key, value string) string {
|
||||||
if value == "" {
|
if value == "" {
|
||||||
return value
|
return value
|
||||||
|
|||||||
@@ -4,7 +4,8 @@
|
|||||||
|
|
||||||
package system_config
|
package system_config
|
||||||
|
|
||||||
import ("bufio"
|
import (
|
||||||
|
"bufio"
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
@@ -17,24 +18,24 @@ import ("bufio"
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"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/storage"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
|
)
|
||||||
|
|
||||||
const expectedDefaultConfigsCount = 30
|
const expectedDefaultConfigsCount = 30
|
||||||
|
|
||||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||||
gin.SetMode(gin.TestMode)
|
r := testhelper.NewTestGinEngine()
|
||||||
r := gin.New()
|
|
||||||
adminGroup := r.Group("/api/v1/admin")
|
adminGroup := r.Group("/api/v1/admin")
|
||||||
|
|
||||||
// Mock authentication middleware
|
// Mock authentication middleware
|
||||||
adminGroup.Use(func(c *gin.Context) {
|
adminGroup.Use(func(c *gin.Context) {
|
||||||
if authUser != nil {
|
if authUser != nil {
|
||||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
@@ -86,22 +87,22 @@ func TestCreateSystemConfig(t *testing.T) {
|
|||||||
// Verify caches are invalidated after create and repopulate on read
|
// Verify caches are invalidated after create and repopulate on read
|
||||||
_, err = db.Redis.HGet(
|
_, err = db.Redis.HGet(
|
||||||
context.Background(),
|
context.Background(),
|
||||||
db.PrefixedKey(model.SystemConfigRedisHashKey),
|
db.PrefixedKey(repository.SystemConfigRedisHashKey),
|
||||||
"custom_key",
|
"custom_key",
|
||||||
).Result()
|
).Result()
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatal("expected redis cache miss immediately after create")
|
t.Fatal("expected redis cache miss immediately after create")
|
||||||
}
|
}
|
||||||
|
|
||||||
var loaded model.SystemConfig
|
loaded, err := repository.GetSystemConfigByKey(context.Background(), "custom_key")
|
||||||
if err := loaded.GetByKey(context.Background(), "custom_key"); err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetByKey(custom_key) error = %v", err)
|
t.Fatalf("GetSystemConfigByKey(custom_key) error = %v", err)
|
||||||
}
|
}
|
||||||
if loaded.Value != "custom_value" {
|
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 {
|
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
|
// Verify caches are invalidated after update and repopulate on read
|
||||||
_, err := db.Redis.HGet(
|
_, err := db.Redis.HGet(
|
||||||
context.Background(),
|
context.Background(),
|
||||||
db.PrefixedKey(model.SystemConfigRedisHashKey),
|
db.PrefixedKey(repository.SystemConfigRedisHashKey),
|
||||||
model.ConfigKeySiteName,
|
model.ConfigKeySiteName,
|
||||||
).Result()
|
).Result()
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatal("expected redis cache miss immediately after update")
|
t.Fatal("expected redis cache miss immediately after update")
|
||||||
}
|
}
|
||||||
|
|
||||||
var loaded model.SystemConfig
|
loaded, err := repository.GetSystemConfigByKey(context.Background(), model.ConfigKeySiteName)
|
||||||
if err := loaded.GetByKey(context.Background(), model.ConfigKeySiteName); err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetByKey(site_name) error = %v", err)
|
t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
if loaded.Value != "Super Site Name" {
|
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 {
|
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)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ import ("bytes"
|
|||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/task"
|
"github.com/Rain-kl/Wavelet/internal/task"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/hibiken/asynq"
|
"github.com/hibiken/asynq"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
@@ -51,7 +50,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
|
|||||||
// Mock authentication middleware
|
// Mock authentication middleware
|
||||||
adminGroup.Use(func(c *gin.Context) {
|
adminGroup.Use(func(c *gin.Context) {
|
||||||
if authUser != nil {
|
if authUser != nil {
|
||||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -3,15 +3,15 @@
|
|||||||
|
|
||||||
package template
|
package template
|
||||||
|
|
||||||
import ("errors"
|
import (
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
|
)
|
||||||
|
|
||||||
// CreateTemplateRequest 创建模板请求
|
// CreateTemplateRequest 创建模板请求
|
||||||
type CreateTemplateRequest struct {
|
type CreateTemplateRequest struct {
|
||||||
@@ -32,6 +32,24 @@ type UpdateTemplateRequest struct {
|
|||||||
Description string `json:"description" binding:"max=255"`
|
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 创建模板
|
// CreateTemplate 创建模板
|
||||||
// @Summary 创建模板
|
// @Summary 创建模板
|
||||||
// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限
|
// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限
|
||||||
@@ -53,33 +71,8 @@ func CreateTemplate(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 检查模板 Key 是否已存在
|
tmpl, err := createTemplate(c.Request.Context(), req)
|
||||||
var existing model.Template
|
if abortTemplateLogicError(c, err) {
|
||||||
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())
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -98,8 +91,8 @@ func CreateTemplate(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/templates [get]
|
// @Router /api/v1/admin/templates [get]
|
||||||
func ListTemplates(c *gin.Context) {
|
func ListTemplates(c *gin.Context) {
|
||||||
var templates []model.Template
|
templates, err := listTemplates(c.Request.Context())
|
||||||
if err := db.DB(c.Request.Context()).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
|
if err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -121,13 +114,8 @@ func ListTemplates(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/templates/{key} [get]
|
// @Router /api/v1/admin/templates/{key} [get]
|
||||||
func GetTemplate(c *gin.Context) {
|
func GetTemplate(c *gin.Context) {
|
||||||
var tmpl model.Template
|
tmpl, err := getTemplate(c.Request.Context(), c.Param("key"))
|
||||||
if err := db.DB(c.Request.Context()).Where("key = ?", c.Param("key")).First(&tmpl).Error; err != nil {
|
if abortTemplateLogicError(c, err) {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
||||||
response.AbortNotFound(c, TemplateNotFound)
|
|
||||||
} else {
|
|
||||||
response.AbortInternal(c, err.Error())
|
|
||||||
}
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -157,32 +145,8 @@ func UpdateTemplate(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
key := c.Param("key")
|
tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req)
|
||||||
|
if abortTemplateLogicError(c, err) {
|
||||||
// 检查模板是否存在
|
|
||||||
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())
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -204,29 +168,9 @@ func UpdateTemplate(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/templates/{key} [delete]
|
// @Router /api/v1/admin/templates/{key} [delete]
|
||||||
func DeleteTemplate(c *gin.Context) {
|
func DeleteTemplate(c *gin.Context) {
|
||||||
key := c.Param("key")
|
if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
|
||||||
|
|
||||||
// 检查模板是否存在
|
|
||||||
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())
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
@@ -12,20 +12,18 @@ import ("bytes"
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response")
|
||||||
|
|
||||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||||
gin.SetMode(gin.TestMode)
|
r := testhelper.NewTestGinEngine()
|
||||||
r := gin.New()
|
|
||||||
adminGroup := r.Group("/api/v1/admin")
|
adminGroup := r.Group("/api/v1/admin")
|
||||||
|
|
||||||
// Mock authentication middleware
|
// Mock authentication middleware
|
||||||
adminGroup.Use(func(c *gin.Context) {
|
adminGroup.Use(func(c *gin.Context) {
|
||||||
if authUser != nil {
|
if authUser != nil {
|
||||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/buildinfo"
|
"github.com/Rain-kl/Wavelet/internal/buildinfo"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
"golang.org/x/mod/semver"
|
"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) {
|
func loadRepository(ctx context.Context) (string, error) {
|
||||||
var config model.SystemConfig
|
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository)
|
||||||
if err := config.GetByKey(ctx, model.ConfigKeyUpdateUpstreamRepository); err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
|
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
|
||||||
}
|
}
|
||||||
return parseRepository(config.Value)
|
return parseRepository(config.Value)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) {
|
func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) {
|
||||||
repository, err := loadRepository(ctx)
|
upstreamRepo, err := loadRepository(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return Status{}, releaseAsset{}, err
|
return Status{}, releaseAsset{}, err
|
||||||
}
|
}
|
||||||
release, asset, err := m.fetchRelease(ctx, repository)
|
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return Status{}, releaseAsset{}, err
|
return Status{}, releaseAsset{}, err
|
||||||
}
|
}
|
||||||
@@ -259,7 +260,7 @@ func (m *manager) status(ctx context.Context) (Status, releaseAsset, error) {
|
|||||||
ReleaseNotes: release.Body,
|
ReleaseNotes: release.Body,
|
||||||
ReleaseURL: release.HTMLURL,
|
ReleaseURL: release.HTMLURL,
|
||||||
PublishedAt: release.Published.Format(time.RFC3339),
|
PublishedAt: release.Published.Format(time.RFC3339),
|
||||||
UpstreamRepository: repository,
|
UpstreamRepository: upstreamRepo,
|
||||||
AssetName: asset.Name,
|
AssetName: asset.Name,
|
||||||
Platform: runtime.GOOS + "/" + runtime.GOARCH,
|
Platform: runtime.GOOS + "/" + runtime.GOARCH,
|
||||||
}, asset, nil
|
}, asset, nil
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -5,16 +5,13 @@
|
|||||||
package user
|
package user
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"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/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"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 获取用户列表
|
// ListUsers 获取用户列表
|
||||||
// @Summary 获取用户列表
|
// @Summary 获取用户列表
|
||||||
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
|
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
|
||||||
@@ -98,7 +120,6 @@ func toUser(u model.User) user {
|
|||||||
// @Failure 403 {object} response.Any "无管理员权限"
|
// @Failure 403 {object} response.Any "无管理员权限"
|
||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/users [get]
|
// @Router /api/v1/admin/users [get]
|
||||||
// ListUsers 获取用户列表
|
|
||||||
func ListUsers(c *gin.Context) {
|
func ListUsers(c *gin.Context) {
|
||||||
var req listUsersRequest
|
var req listUsersRequest
|
||||||
if err := c.ShouldBindQuery(&req); err != nil {
|
if err := c.ShouldBindQuery(&req); err != nil {
|
||||||
@@ -106,34 +127,8 @@ func ListUsers(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var modelUsers []model.User
|
total, modelUsers, err := listUsers(c.Request.Context(), req)
|
||||||
var total int64
|
if err != nil {
|
||||||
|
|
||||||
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 {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -169,17 +164,8 @@ func GetUser(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var targetUser model.User
|
targetUser, err := getUserDetail(c.Request.Context(), id)
|
||||||
if err := db.DB(c.Request.Context()).
|
if abortUserLogicError(c, err, userNotFound, nil, nil) {
|
||||||
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())
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -219,32 +205,10 @@ func UpdateUserStatus(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var targetUser struct {
|
if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
|
||||||
ID uint64 `gorm:"column:id"`
|
if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) {
|
||||||
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)
|
|
||||||
return
|
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)
|
response.AbortInternal(c, updateUserFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -272,43 +236,11 @@ func DeleteUser(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
if currUser != nil && currUser.ID == id {
|
if err := deleteUser(c.Request.Context(), currUser.ID, id); err != nil {
|
||||||
response.AbortForbidden(c, cannotDeleteSelf)
|
if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) {
|
||||||
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)
|
|
||||||
return
|
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)
|
response.AbortInternal(c, deleteUserFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -347,67 +279,10 @@ func CreateUser(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
req.Username = strings.TrimSpace(req.Username)
|
newUser, err := createUser(c.Request.Context(), req)
|
||||||
req.Nickname = strings.TrimSpace(req.Nickname)
|
if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) {
|
||||||
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())
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OK(toUser(newUser)))
|
c.JSON(http.StatusOK, response.OK(toUser(newUser)))
|
||||||
}
|
}
|
||||||
@@ -14,20 +14,18 @@ import ("bytes"
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response")
|
||||||
|
|
||||||
func setupTestRouter(authUser *model.User) *gin.Engine {
|
func setupTestRouter(authUser *model.User) *gin.Engine {
|
||||||
gin.SetMode(gin.TestMode)
|
r := testhelper.NewTestGinEngine()
|
||||||
r := gin.New()
|
|
||||||
adminGroup := r.Group("/api/v1/admin")
|
adminGroup := r.Group("/api/v1/admin")
|
||||||
|
|
||||||
// Mock authentication middleware
|
// Mock authentication middleware
|
||||||
adminGroup.Use(func(c *gin.Context) {
|
adminGroup.Use(func(c *gin.Context) {
|
||||||
if authUser != nil {
|
if authUser != nil {
|
||||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
|
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
|
||||||
)
|
)
|
||||||
@@ -68,7 +69,7 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to enable cap_login_enabled in DB: %v", err)
|
t.Fatalf("failed to enable cap_login_enabled in DB: %v", err)
|
||||||
}
|
}
|
||||||
if err := model.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled); err != nil {
|
if err := repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyCapLoginEnabled); err != nil {
|
||||||
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
|
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
|
||||||
}
|
}
|
||||||
InvalidateRuntimeSettings()
|
InvalidateRuntimeSettings()
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -130,7 +131,7 @@ func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, er
|
|||||||
}
|
}
|
||||||
|
|
||||||
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
|
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
|
||||||
configs, err := model.ListSystemConfigsByKeys(ctx, runtimeConfigKeys)
|
configs, err := repository.ListSystemConfigsByKeys(ctx, runtimeConfigKeys)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return RuntimeSettings{}, err
|
return RuntimeSettings{}, err
|
||||||
}
|
}
|
||||||
@@ -190,7 +191,7 @@ func startRuntimeSettingsInvalidationListener() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
pubsub := db.Redis.Subscribe(context.Background(), model.SystemConfigInvalidationChannel)
|
pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
}()
|
||||||
@@ -208,4 +209,4 @@ func startRuntimeSettingsInvalidationListener() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -19,7 +20,7 @@ func TestCurrentSettingsLoadsSnapshotOnce(t *testing.T) {
|
|||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
ResetRuntimeSettingsForTest()
|
ResetRuntimeSettingsForTest()
|
||||||
model.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
|
|
||||||
first, err := CurrentSettings(ctx)
|
first, err := CurrentSettings(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -34,7 +35,7 @@ func TestCurrentSettingsLoadsSnapshotOnce(t *testing.T) {
|
|||||||
Update("value", "4").Error; err != nil {
|
Update("value", "4").Error; err != nil {
|
||||||
t.Fatalf("Update(cap_challenge_count) error = %v", err)
|
t.Fatalf("Update(cap_challenge_count) error = %v", err)
|
||||||
}
|
}
|
||||||
if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapChallengeCount); err != nil {
|
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapChallengeCount); err != nil {
|
||||||
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
|
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
|
||||||
}
|
}
|
||||||
InvalidateRuntimeSettings()
|
InvalidateRuntimeSettings()
|
||||||
@@ -64,7 +65,7 @@ func TestProtectionEnabledReflectsLoginSwitch(t *testing.T) {
|
|||||||
Update("value", "true").Error; err != nil {
|
Update("value", "true").Error; err != nil {
|
||||||
t.Fatalf("Update(cap_login_enabled) error = %v", err)
|
t.Fatalf("Update(cap_login_enabled) error = %v", err)
|
||||||
}
|
}
|
||||||
if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapLoginEnabled); err != nil {
|
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyCapLoginEnabled); err != nil {
|
||||||
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
|
t.Fatalf("InvalidateSystemConfigCache() error = %v", err)
|
||||||
}
|
}
|
||||||
InvalidateRuntimeSettings()
|
InvalidateRuntimeSettings()
|
||||||
@@ -115,4 +116,4 @@ func TestInstallTestRuntimeSettings(t *testing.T) {
|
|||||||
if settings.ChallengeCount != 2 {
|
if settings.ChallengeCount != 2 {
|
||||||
t.Fatalf("ChallengeCount = %d, want %d", settings.ChallengeCount, 2)
|
t.Fatalf("ChallengeCount = %d, want %d", settings.ChallengeCount, 2)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -17,11 +18,11 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
|||||||
defer cleanup()
|
defer cleanup()
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
if err := model.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||||
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := model.ListVisibleSystemConfigs(ctx); err != nil {
|
if _, err := repository.ListVisibleSystemConfigs(ctx); err != nil {
|
||||||
t.Fatalf("ListVisibleSystemConfigs() warm cache error = %v", err)
|
t.Fatalf("ListVisibleSystemConfigs() warm cache error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -35,7 +36,7 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
|||||||
t.Fatalf("Create(cache_probe_public_key) error = %v", err)
|
t.Fatalf("Create(cache_probe_public_key) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
cached, err := model.ListVisibleSystemConfigs(ctx)
|
cached, err := repository.ListVisibleSystemConfigs(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("ListVisibleSystemConfigs() cached call error = %v", err)
|
t.Fatalf("ListVisibleSystemConfigs() cached call error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -45,7 +46,7 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
exists, err := db.Redis.Exists(ctx, db.PrefixedKey(model.SystemConfigVisibleListRedisKey)).Result()
|
exists, err := db.Redis.Exists(ctx, db.PrefixedKey(repository.SystemConfigVisibleListRedisKey)).Result()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Redis.Exists() error = %v", err)
|
t.Fatalf("Redis.Exists() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -53,11 +54,11 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
|||||||
t.Fatal("expected visible config list cache key to exist")
|
t.Fatal("expected visible config list cache key to exist")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := model.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||||
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
t.Fatalf("InvalidateVisibleSystemConfigsCache() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
refreshed, err := model.ListVisibleSystemConfigs(ctx)
|
refreshed, err := repository.ListVisibleSystemConfigs(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("ListVisibleSystemConfigs() refreshed call error = %v", err)
|
t.Fatalf("ListVisibleSystemConfigs() refreshed call error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -72,4 +73,4 @@ func TestListVisibleSystemConfigsUsesRedisCache(t *testing.T) {
|
|||||||
if !found {
|
if !found {
|
||||||
t.Fatal("refreshed visible config list should include newly created public config")
|
t.Fatal("refreshed visible config list should include newly created public config")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,12 +5,15 @@
|
|||||||
// Package config 提供公开配置查询接口
|
// Package config 提供公开配置查询接口
|
||||||
package config
|
package config
|
||||||
|
|
||||||
import ("net/http"
|
import (
|
||||||
|
"net/http"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
|
)
|
||||||
|
|
||||||
// GetPublicConfig 获取公共配置
|
// GetPublicConfig 获取公共配置
|
||||||
// @Summary 获取公共配置
|
// @Summary 获取公共配置
|
||||||
@@ -22,7 +25,7 @@ import ("net/http"
|
|||||||
// @Router /api/v1/config/public [get]
|
// @Router /api/v1/config/public [get]
|
||||||
func GetPublicConfig(c *gin.Context) {
|
func GetPublicConfig(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
configs, err := model.ListVisibleSystemConfigs(ctx)
|
configs, err := repository.ListVisibleSystemConfigs(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
@@ -45,7 +48,7 @@ func GetPublicConfig(c *gin.Context) {
|
|||||||
// @Router /robots.txt [get]
|
// @Router /robots.txt [get]
|
||||||
func GetRobotsTXT(c *gin.Context) {
|
func GetRobotsTXT(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
enabled, err := model.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled)
|
||||||
content := "User-Agent: *\nDisallow: /\n"
|
content := "User-Agent: *\nDisallow: /\n"
|
||||||
if err == nil && enabled {
|
if err == nil && enabled {
|
||||||
content = "User-Agent: *\nAllow: /\n"
|
content = "User-Agent: *\nAllow: /\n"
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -20,17 +21,17 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) {
|
|||||||
defer cleanup()
|
defer cleanup()
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
model.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
if err := model.InvalidateAllSystemConfigCaches(ctx); err != nil {
|
if err := repository.InvalidateAllSystemConfigCaches(ctx); err != nil {
|
||||||
t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err)
|
t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var warm model.SystemConfig
|
warm, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||||
if err := warm.GetByKey(ctx, model.ConfigKeySiteName); err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetByKey(site_name) warm error = %v", err)
|
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
|
||||||
}
|
}
|
||||||
if warm.Value != "Wavelet" {
|
if warm.Value != "Wavelet" {
|
||||||
t.Fatalf("GetByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
|
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := dbConn.Model(&model.SystemConfig{}).
|
if err := dbConn.Model(&model.SystemConfig{}).
|
||||||
@@ -38,31 +39,31 @@ func TestSystemConfigRAMCacheServesUntilInvalidated(t *testing.T) {
|
|||||||
Update("value", "ram_probe_value").Error; err != nil {
|
Update("value", "ram_probe_value").Error; err != nil {
|
||||||
t.Fatalf("Update(site_name) error = %v", err)
|
t.Fatalf("Update(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
if err := db.HDel(ctx, model.SystemConfigRedisHashKey, model.ConfigKeySiteName); err != nil {
|
if err := db.HDel(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeySiteName); err != nil {
|
||||||
t.Fatalf("HDel(site_name) error = %v", err)
|
t.Fatalf("HDel(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var cached model.SystemConfig
|
cached, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||||
if err := cached.GetByKey(ctx, model.ConfigKeySiteName); err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetByKey(site_name) cached error = %v", err)
|
t.Fatalf("GetSystemConfigByKey(site_name) cached error = %v", err)
|
||||||
}
|
}
|
||||||
if cached.Value != "Wavelet" {
|
if cached.Value != "Wavelet" {
|
||||||
t.Fatalf("GetByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "Wavelet")
|
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want stale RAM value %q", cached.Value, "Wavelet")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
||||||
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var refreshed model.SystemConfig
|
refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||||
if err := refreshed.GetByKey(ctx, model.ConfigKeySiteName); err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetByKey(site_name) refreshed error = %v", err)
|
t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err)
|
||||||
}
|
}
|
||||||
if refreshed.Value != "ram_probe_value" {
|
if refreshed.Value != "ram_probe_value" {
|
||||||
t.Fatalf("GetByKey(site_name).Value = %q, want %q", refreshed.Value, "ram_probe_value")
|
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "ram_probe_value")
|
||||||
}
|
}
|
||||||
|
|
||||||
exists, err := db.Redis.HExists(ctx, db.PrefixedKey(model.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
exists, err := db.Redis.HExists(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("HExists(site_name) error = %v", err)
|
t.Fatalf("HExists(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -76,16 +77,17 @@ func TestInvalidateSystemConfigCacheClearsRedisField(t *testing.T) {
|
|||||||
defer cleanup()
|
defer cleanup()
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
var sc model.SystemConfig
|
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeySiteName); err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetByKey(site_name) error = %v", err)
|
t.Fatalf("GetSystemConfigByKey(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
|
_ = sc
|
||||||
|
|
||||||
if err := model.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
||||||
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
t.Fatalf("InvalidateSystemConfigCache(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := db.Redis.HGet(ctx, db.PrefixedKey(model.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
_, err = db.Redis.HGet(ctx, db.PrefixedKey(repository.SystemConfigRedisHashKey), model.ConfigKeySiteName).Result()
|
||||||
if !errors.Is(err, redis.Nil) {
|
if !errors.Is(err, redis.Nil) {
|
||||||
t.Fatalf("HGet(site_name) error = %v, want redis.Nil", err)
|
t.Fatalf("HGet(site_name) error = %v, want redis.Nil", err)
|
||||||
}
|
}
|
||||||
@@ -96,11 +98,11 @@ func TestInvalidateSystemConfigCacheClearsRedisField(t *testing.T) {
|
|||||||
t.Fatalf("Update(site_name) error = %v", err)
|
t.Fatalf("Update(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var refreshed model.SystemConfig
|
refreshed, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||||
if err := refreshed.GetByKey(ctx, model.ConfigKeySiteName); err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetByKey(site_name) refreshed error = %v", err)
|
t.Fatalf("GetSystemConfigByKey(site_name) refreshed error = %v", err)
|
||||||
}
|
}
|
||||||
if refreshed.Value != "after_invalidate" {
|
if refreshed.Value != "after_invalidate" {
|
||||||
t.Fatalf("GetByKey(site_name).Value = %q, want %q", refreshed.Value, "after_invalidate")
|
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", refreshed.Value, "after_invalidate")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -9,12 +9,13 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/coreos/go-oidc/v3/oidc"
|
"github.com/coreos/go-oidc/v3/oidc"
|
||||||
"golang.org/x/oauth2"
|
"golang.org/x/oauth2"
|
||||||
)
|
)
|
||||||
|
|
||||||
func isOIDCLoginEnabled(ctx context.Context) bool {
|
func isOIDCLoginEnabled(ctx context.Context) bool {
|
||||||
enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
@@ -37,7 +38,7 @@ func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSourc
|
|||||||
}
|
}
|
||||||
|
|
||||||
func activeLoginSources(ctx context.Context) []AuthSourceView {
|
func activeLoginSources(ctx context.Context) []AuthSourceView {
|
||||||
enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
|
||||||
if err == nil && !enabled {
|
if err == nil && !enabled {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -62,8 +63,8 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
||||||
var sc model.SystemConfig
|
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || strings.TrimSpace(sc.Value) == "" {
|
if err != nil || strings.TrimSpace(sc.Value) == "" {
|
||||||
return "", errors.New(errServerAddressMissing)
|
return "", errors.New(errServerAddressMissing)
|
||||||
}
|
}
|
||||||
return strings.TrimRight(sc.Value, "/") + "/login", nil
|
return strings.TrimRight(sc.Value, "/") + "/login", nil
|
||||||
|
|||||||
@@ -1,13 +1,11 @@
|
|||||||
// Copyright 2025 linux.do
|
|
||||||
// Copyright 2026 Arctel.net
|
// Copyright 2026 Arctel.net
|
||||||
// SPDX-License-Identifier: Apache-2.0
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
// Package util 提供通用工具函数
|
package oauth
|
||||||
package util
|
|
||||||
|
|
||||||
import "github.com/gin-gonic/gin"
|
import "github.com/gin-gonic/gin"
|
||||||
|
|
||||||
// GetFromContext 从上下文获取指定类型的值
|
// GetFromContext 从 Gin 请求上下文获取指定类型的值。
|
||||||
func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
|
func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
|
||||||
value, exists := c.Get(key)
|
value, exists := c.Get(key)
|
||||||
if !exists {
|
if !exists {
|
||||||
@@ -18,7 +16,7 @@ func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
|
|||||||
return typed, ok
|
return typed, ok
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetToContext 设置值到上下文
|
// SetToContext 设置值到 Gin 请求上下文。
|
||||||
func SetToContext[T any](c *gin.Context, key string, value T) {
|
func SetToContext[T any](c *gin.Context, key string, value T) {
|
||||||
c.Set(key, value)
|
c.Set(key, value)
|
||||||
}
|
}
|
||||||
@@ -15,6 +15,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -188,7 +189,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
|
|||||||
// handleCallbackRegister 处理 OAuth 回调中的自动注册流程
|
// handleCallbackRegister 处理 OAuth 回调中的自动注册流程
|
||||||
// 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false
|
// 若注册被禁用则保存 pending 信息并返回 false;若注册成功则返回新用户;若出错则返回 false
|
||||||
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
|
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
|
||||||
registrationEnabled, regErr := model.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||||
if regErr != nil {
|
if regErr != nil {
|
||||||
registrationEnabled = true
|
registrationEnabled = true
|
||||||
}
|
}
|
||||||
@@ -223,4 +224,4 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A
|
|||||||
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
|
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
|
||||||
|
|
||||||
return user, true
|
return user, true
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -62,8 +62,8 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
|
|||||||
if user.Username == "system" {
|
if user.Username == "system" {
|
||||||
return nil, errors.New("system user is not allowed to login")
|
return nil, errors.New("system user is not allowed to login")
|
||||||
}
|
}
|
||||||
util.SetToContext(c, TokenAuthKey, true)
|
SetToContext(c, TokenAuthKey, true)
|
||||||
util.SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin)
|
SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin)
|
||||||
return user, nil
|
return user, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -91,8 +91,8 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// set keys in context for session auth
|
// set keys in context for session auth
|
||||||
util.SetToContext(c, TokenAuthKey, false)
|
SetToContext(c, TokenAuthKey, false)
|
||||||
util.SetToContext(c, TokenAdminKey, false)
|
SetToContext(c, TokenAdminKey, false)
|
||||||
|
|
||||||
// 强行阻止 system 用户任何会话/Token 鉴权通过
|
// 强行阻止 system 用户任何会话/Token 鉴权通过
|
||||||
if user.Username == "system" {
|
if user.Username == "system" {
|
||||||
@@ -119,7 +119,7 @@ func LoginRequired() gin.HandlerFunc {
|
|||||||
LogForAudit(ctx, user, c)
|
LogForAudit(ctx, user, c)
|
||||||
|
|
||||||
// set user info
|
// set user info
|
||||||
util.SetToContext(c, UserObjKey, user)
|
SetToContext(c, UserObjKey, user)
|
||||||
|
|
||||||
// next
|
// next
|
||||||
c.Next()
|
c.Next()
|
||||||
@@ -129,7 +129,7 @@ func LoginRequired() gin.HandlerFunc {
|
|||||||
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
|
// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点
|
||||||
func DisallowTokenAuth() gin.HandlerFunc {
|
func DisallowTokenAuth() gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
if tokenAuth, _ := util.GetFromContext[bool](c, TokenAuthKey); tokenAuth {
|
if tokenAuth, _ := GetFromContext[bool](c, TokenAuthKey); tokenAuth {
|
||||||
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
|
response.AbortForbidden(c, ErrTokenAuthNotAllowed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
"github.com/Rain-kl/Wavelet/internal/util"
|
||||||
)
|
)
|
||||||
@@ -1094,7 +1095,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
Key: model.ConfigKeyOIDCLoginEnabled,
|
Key: model.ConfigKeyOIDCLoginEnabled,
|
||||||
Value: "false",
|
Value: "false",
|
||||||
})
|
})
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||||
if wLoginDisabled.Code != http.StatusBadRequest {
|
if wLoginDisabled.Code != http.StatusBadRequest {
|
||||||
t.Errorf("expected 400 when OIDC globally disabled, got %d", wLoginDisabled.Code)
|
t.Errorf("expected 400 when OIDC globally disabled, got %d", wLoginDisabled.Code)
|
||||||
@@ -1102,7 +1103,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
|
|
||||||
// Re-enable globally, but deactivate source
|
// Re-enable globally, but deactivate source
|
||||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
|
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
|
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
|
||||||
|
|
||||||
wSourceInactive := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
wSourceInactive := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||||
@@ -1113,7 +1114,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
// --- 2. Test Authorize enforcement ---
|
// --- 2. Test Authorize enforcement ---
|
||||||
// Deactivate globally again
|
// Deactivate globally again
|
||||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false")
|
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false")
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
|
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
|
||||||
|
|
||||||
wAuthDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/"+testSourceName+"/authorize", nil, nil, nil)
|
wAuthDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/"+testSourceName+"/authorize", nil, nil, nil)
|
||||||
@@ -1124,7 +1125,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
// --- 3. Test Callback enforcement ---
|
// --- 3. Test Callback enforcement ---
|
||||||
// Set up a valid state beforehand (when enabled)
|
// Set up a valid state beforehand (when enabled)
|
||||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
|
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
|
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
|
||||||
|
|
||||||
wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||||
@@ -1149,7 +1150,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
|
|
||||||
// Now disable OIDC globally and attempt callback
|
// Now disable OIDC globally and attempt callback
|
||||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false")
|
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false")
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state)
|
reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state)
|
||||||
wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{
|
wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
@@ -1160,7 +1161,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
|
|
||||||
// Enable globally but deactivate source and attempt callback
|
// Enable globally but deactivate source and attempt callback
|
||||||
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
|
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
|
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
|
||||||
|
|
||||||
// Since callback deletes state, we need to generate state again
|
// Since callback deletes state, we need to generate state again
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ import ("net/http"
|
|||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response")
|
||||||
|
|
||||||
@@ -60,7 +60,7 @@ func BuildBasicUserInfo(user *model.User, needChange bool) BasicUserInfo {
|
|||||||
// @Router /api/v1/user-info [get]
|
// @Router /api/v1/user-info [get]
|
||||||
// @Router /api/v1/user/self [get]
|
// @Router /api/v1/user/self [get]
|
||||||
func UserInfo(c *gin.Context) {
|
func UserInfo(c *gin.Context) {
|
||||||
user, _ := util.GetFromContext[*model.User](c, UserObjKey)
|
user, _ := GetFromContext[*model.User](c, UserObjKey)
|
||||||
session := sessions.Default(c)
|
session := sessions.Default(c)
|
||||||
needChange := session.Get("need_change_password") == true
|
needChange := session.Get("need_change_password") == true
|
||||||
|
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
@@ -56,7 +57,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
|||||||
maxAge := config.Config.App.SessionAge
|
maxAge := config.Config.App.SessionAge
|
||||||
isSessionCookie := false
|
isSessionCookie := false
|
||||||
|
|
||||||
ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
switch {
|
switch {
|
||||||
case ttlHours == -1:
|
case ttlHours == -1:
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -39,7 +38,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
|
|||||||
c.Next()
|
c.Next()
|
||||||
|
|
||||||
// 3. 后置身份检查:仅记录通过认证的请求
|
// 3. 后置身份检查:仅记录通过认证的请求
|
||||||
userObj, exists := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
userObj, exists := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
if !exists || userObj == nil {
|
if !exists || userObj == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
@@ -51,7 +50,7 @@ func TestRiskControlMiddleware(t *testing.T) {
|
|||||||
r.Use(func(c *gin.Context) {
|
r.Use(func(c *gin.Context) {
|
||||||
// Mock authentication middleware placing user in context
|
// Mock authentication middleware placing user in context
|
||||||
user := &model.User{ID: 12345}
|
user := &model.User{ID: 12345}
|
||||||
util.SetToContext(c, oauth.UserObjKey, user)
|
oauth.SetToContext(c, oauth.UserObjKey, user)
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
r.Use(RiskControlMiddleware())
|
r.Use(RiskControlMiddleware())
|
||||||
|
|||||||
+3
-2
@@ -15,6 +15,7 @@ import (
|
|||||||
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"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/storage"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -112,8 +113,8 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func parseFileAccessWhitelist(ctx context.Context) []string {
|
func parseFileAccessWhitelist(ctx context.Context) []string {
|
||||||
var sc model.SystemConfig
|
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyFileAccessWhitelist)
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeyFileAccessWhitelist); err != nil || sc.Value == "" {
|
if err != nil || sc.Value == "" {
|
||||||
return []string{shared.DefaultPublicUploadType}
|
return []string{shared.DefaultPublicUploadType}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+5
-4
@@ -8,10 +8,11 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
|
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -67,10 +68,10 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
|||||||
if err := dbConn.Save(&sc).Error; err != nil {
|
if err := dbConn.Save(&sc).Error; err != nil {
|
||||||
t.Fatalf("save whitelist config: %v", err)
|
t.Fatalf("save whitelist config: %v", err)
|
||||||
}
|
}
|
||||||
if err := db.HSetJSON(ctx, model.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil {
|
if err := db.HSetJSON(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil {
|
||||||
t.Fatalf("refresh whitelist redis cache: %v", err)
|
t.Fatalf("refresh whitelist redis cache: %v", err)
|
||||||
}
|
}
|
||||||
model.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
|
|
||||||
ResetAccessCaches()
|
ResetAccessCaches()
|
||||||
if !IsFilePublic(ctx, "attachment") {
|
if !IsFilePublic(ctx, "attachment") {
|
||||||
@@ -97,4 +98,4 @@ func TestAccessCacheTTLExpires(t *testing.T) {
|
|||||||
if !IsFilePublic(ctx, "avatar") {
|
if !IsFilePublic(ctx, "avatar") {
|
||||||
t.Fatal("expected whitelist reload after TTL expiration")
|
t.Fatal("expected whitelist reload after TTL expiration")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/diskcache"
|
"github.com/Rain-kl/Wavelet/internal/diskcache"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
apputil "github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"golang.org/x/sync/singleflight"
|
"golang.org/x/sync/singleflight"
|
||||||
@@ -282,7 +282,7 @@ func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, er
|
|||||||
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
|
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
|
||||||
var currUser *model.User
|
var currUser *model.User
|
||||||
var err error
|
var err error
|
||||||
if u, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil {
|
if u, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil {
|
||||||
currUser = u
|
currUser = u
|
||||||
} else {
|
} else {
|
||||||
currUser, err = oauth.GetUserFromRequest(c)
|
currUser, err = oauth.GetUserFromRequest(c)
|
||||||
@@ -306,7 +306,7 @@ func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
||||||
if _, ok := apputil.GetFromContext[*model.User](c, oauth.UserObjKey); !ok {
|
if _, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); !ok {
|
||||||
if _, err := oauth.GetUserFromRequest(c); err != nil {
|
if _, err := oauth.GetUserFromRequest(c); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,22 +4,17 @@
|
|||||||
package handler
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"sort"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
|
||||||
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"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/model"
|
||||||
apputil "github.com/Rain-kl/Wavelet/internal/util"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type listFilesRequest struct {
|
type listFilesRequest struct {
|
||||||
@@ -69,31 +64,15 @@ func ListFiles(c *gin.Context) {
|
|||||||
req.PageSize = 20
|
req.PageSize = 20
|
||||||
}
|
}
|
||||||
|
|
||||||
query := db.DB(ctx).Model(&model.Upload{}).
|
total, items, err := listUploadFiles(ctx, repository.UploadListFilter{
|
||||||
Where("status != ?", model.UploadStatusDeleted)
|
UserID: req.UserID,
|
||||||
|
Keyword: req.Keyword,
|
||||||
if req.UserID != 0 {
|
Type: req.Type,
|
||||||
query = query.Where("user_id = ?", req.UserID)
|
Extension: req.Extension,
|
||||||
}
|
Page: req.Page,
|
||||||
if req.Keyword != "" {
|
PageSize: req.PageSize,
|
||||||
query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(req.Keyword)+"%")
|
})
|
||||||
}
|
if err != nil {
|
||||||
if req.Type != "" {
|
|
||||||
query = query.Where("type = ?", req.Type)
|
|
||||||
}
|
|
||||||
if req.Extension != "" {
|
|
||||||
query = query.Where("extension = ?", strings.ToLower(req.Extension))
|
|
||||||
}
|
|
||||||
|
|
||||||
var total int64
|
|
||||||
if err := query.Count(&total).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, shared.ErrQueryFileCountFailed)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var items []model.Upload
|
|
||||||
offset := (req.Page - 1) * req.PageSize
|
|
||||||
if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
|
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -130,20 +109,14 @@ func DeleteFile(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var upload model.Upload
|
if _, err := softDeleteUpload(ctx, uploadID); err != nil {
|
||||||
if err := db.DB(ctx).Where("id = ? AND status != ?", uploadID, model.UploadStatusDeleted).First(&upload).Error; err != nil {
|
if isRecordNotFound(err) {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
||||||
c.AbortWithStatus(http.StatusNotFound)
|
c.AbortWithStatus(http.StatusNotFound)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
|
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
uploadstats.RecordUploadStatsRemove(ctx, &upload)
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -159,16 +132,12 @@ func DeleteFile(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/uploads/types [get]
|
// @Router /api/v1/admin/uploads/types [get]
|
||||||
func GetDistinctUploadTypes(c *gin.Context) {
|
func GetDistinctUploadTypes(c *gin.Context) {
|
||||||
var dbTypes []string
|
types, err := listDistinctUploadTypes(c.Request.Context())
|
||||||
if err := db.DB(c.Request.Context()).Model(&model.Upload{}).
|
if err != nil {
|
||||||
Where("type IS NOT NULL AND type != ''").
|
|
||||||
Distinct().
|
|
||||||
Pluck("type", &dbTypes).Error; err != nil {
|
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
sort.Strings(dbTypes)
|
c.JSON(http.StatusOK, response.OK(types))
|
||||||
c.JSON(http.StatusOK, response.OK(dbTypes))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type listMyFilesRequest struct {
|
type listMyFilesRequest struct {
|
||||||
@@ -201,7 +170,7 @@ type listMyFilesResponse struct {
|
|||||||
// @Failure 401 {object} response.Any "未登录"
|
// @Failure 401 {object} response.Any "未登录"
|
||||||
// @Router /api/v1/upload/my [get]
|
// @Router /api/v1/upload/my [get]
|
||||||
func ListMyFiles(c *gin.Context) {
|
func ListMyFiles(c *gin.Context) {
|
||||||
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
var req listMyFilesRequest
|
var req listMyFilesRequest
|
||||||
@@ -216,28 +185,14 @@ func ListMyFiles(c *gin.Context) {
|
|||||||
req.PageSize = 20
|
req.PageSize = 20
|
||||||
}
|
}
|
||||||
|
|
||||||
query := db.DB(ctx).Model(&model.Upload{}).
|
total, items, err := listMyUploadFiles(ctx, currUser.ID, repository.UploadListFilter{
|
||||||
Where("user_id = ? AND status != ?", currUser.ID, model.UploadStatusDeleted)
|
Keyword: req.Keyword,
|
||||||
|
Type: req.Type,
|
||||||
if req.Keyword != "" {
|
Extension: req.Extension,
|
||||||
query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(req.Keyword)+"%")
|
Page: req.Page,
|
||||||
}
|
PageSize: req.PageSize,
|
||||||
if req.Type != "" {
|
})
|
||||||
query = query.Where("type = ?", req.Type)
|
if err != nil {
|
||||||
}
|
|
||||||
if req.Extension != "" {
|
|
||||||
query = query.Where("extension = ?", strings.ToLower(req.Extension))
|
|
||||||
}
|
|
||||||
|
|
||||||
var total int64
|
|
||||||
if err := query.Count(&total).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, shared.ErrQueryFileCountFailed)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var items []model.Upload
|
|
||||||
offset := (req.Page - 1) * req.PageSize
|
|
||||||
if err := query.Order("created_at DESC").Offset(offset).Limit(req.PageSize).Find(&items).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
|
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -262,7 +217,7 @@ func ListMyFiles(c *gin.Context) {
|
|||||||
// @Failure 404 {object} response.Any "文件不存在"
|
// @Failure 404 {object} response.Any "文件不存在"
|
||||||
// @Router /api/v1/upload/{id} [delete]
|
// @Router /api/v1/upload/{id} [delete]
|
||||||
func DeleteMyFile(c *gin.Context) {
|
func DeleteMyFile(c *gin.Context) {
|
||||||
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
if uploadstorage.ReadOnly(ctx) {
|
if uploadstorage.ReadOnly(ctx) {
|
||||||
response.AbortConflict(c, shared.ErrStorageReadOnly)
|
response.AbortConflict(c, shared.ErrStorageReadOnly)
|
||||||
@@ -275,26 +230,18 @@ func DeleteMyFile(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var upload model.Upload
|
if _, err := softDeleteOwnedUpload(ctx, currUser.ID, uploadID); err != nil {
|
||||||
if err := db.DB(ctx).Where("id = ? AND status != ?", uploadID, model.UploadStatusDeleted).First(&upload).Error; err != nil {
|
if isRecordNotFound(err) {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
||||||
c.AbortWithStatus(http.StatusNotFound)
|
c.AbortWithStatus(http.StatusNotFound)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
|
if err == errUploadForbidden {
|
||||||
return
|
c.AbortWithStatus(http.StatusForbidden)
|
||||||
}
|
return
|
||||||
|
}
|
||||||
if upload.UserID != currUser.ID {
|
|
||||||
c.AbortWithStatus(http.StatusForbidden)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := db.DB(ctx).Model(&upload).Update("status", model.UploadStatusDeleted).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
|
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
uploadstats.RecordUploadStatsRemove(ctx, &upload)
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -317,7 +264,7 @@ type updateMyFileRequest struct {
|
|||||||
// @Failure 404 {object} response.Any "文件不存在"
|
// @Failure 404 {object} response.Any "文件不存在"
|
||||||
// @Router /api/v1/upload/{id} [put]
|
// @Router /api/v1/upload/{id} [put]
|
||||||
func UpdateMyFile(c *gin.Context) {
|
func UpdateMyFile(c *gin.Context) {
|
||||||
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
if uploadstorage.ReadOnly(ctx) {
|
if uploadstorage.ReadOnly(ctx) {
|
||||||
response.AbortConflict(c, shared.ErrStorageReadOnly)
|
response.AbortConflict(c, shared.ErrStorageReadOnly)
|
||||||
@@ -336,34 +283,18 @@ func UpdateMyFile(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var upload model.Upload
|
upload, err := updateOwnedUpload(ctx, currUser.ID, uploadID, updateMyUploadInput(req))
|
||||||
if err := db.DB(ctx).Where("id = ? AND status != ?", uploadID, model.UploadStatusDeleted).First(&upload).Error; err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if isRecordNotFound(err) {
|
||||||
c.AbortWithStatus(http.StatusNotFound)
|
c.AbortWithStatus(http.StatusNotFound)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
|
if err == errUploadForbidden {
|
||||||
return
|
c.AbortWithStatus(http.StatusForbidden)
|
||||||
}
|
|
||||||
|
|
||||||
if upload.UserID != currUser.ID {
|
|
||||||
c.AbortWithStatus(http.StatusForbidden)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
updates := make(map[string]any)
|
|
||||||
if req.FileName != "" {
|
|
||||||
updates["file_name"] = req.FileName
|
|
||||||
}
|
|
||||||
if req.AccessMode != nil {
|
|
||||||
updates["access_mode"] = *req.AccessMode
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(updates) > 0 {
|
|
||||||
if err := db.DB(ctx).Model(&upload).Updates(updates).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, "更新文件记录失败")
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
response.AbortBadRequest(c, "更新文件记录失败")
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OK(upload))
|
c.JSON(http.StatusOK, response.OK(upload))
|
||||||
|
|||||||
@@ -0,0 +1,199 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
|
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
||||||
|
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||||
|
"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 listUploadFiles(ctx context.Context, filter repository.UploadListFilter) (int64, []model.Upload, error) {
|
||||||
|
return repository.ListUploads(ctx, filter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func listMyUploadFiles(ctx context.Context, userID uint64, filter repository.UploadListFilter) (int64, []model.Upload, error) {
|
||||||
|
filter.UserID = userID
|
||||||
|
return repository.ListUploads(ctx, filter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func softDeleteUpload(ctx context.Context, uploadID uint64) (model.Upload, error) {
|
||||||
|
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
|
||||||
|
if err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
if err := repository.SoftDeleteUpload(ctx, &upload); err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
uploadstats.RecordUploadStatsRemove(ctx, &upload)
|
||||||
|
return upload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func softDeleteOwnedUpload(ctx context.Context, userID, uploadID uint64) (model.Upload, error) {
|
||||||
|
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
|
||||||
|
if err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
if upload.UserID != userID {
|
||||||
|
return model.Upload{}, errUploadForbidden
|
||||||
|
}
|
||||||
|
if err := repository.SoftDeleteUpload(ctx, &upload); err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
uploadstats.RecordUploadStatsRemove(ctx, &upload)
|
||||||
|
return upload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func listDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||||
|
types, err := repository.ListDistinctUploadTypes(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
sort.Strings(types)
|
||||||
|
return types, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type updateMyUploadInput struct {
|
||||||
|
FileName string
|
||||||
|
AccessMode *int
|
||||||
|
}
|
||||||
|
|
||||||
|
func updateOwnedUpload(ctx context.Context, userID, uploadID uint64, input updateMyUploadInput) (model.Upload, error) {
|
||||||
|
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
|
||||||
|
if err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
if upload.UserID != userID {
|
||||||
|
return model.Upload{}, errUploadForbidden
|
||||||
|
}
|
||||||
|
|
||||||
|
updates := make(map[string]any)
|
||||||
|
if input.FileName != "" {
|
||||||
|
updates["file_name"] = input.FileName
|
||||||
|
}
|
||||||
|
if input.AccessMode != nil {
|
||||||
|
updates["access_mode"] = *input.AccessMode
|
||||||
|
}
|
||||||
|
if err := repository.UpdateUpload(ctx, &upload, updates); err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
if name, ok := updates["file_name"].(string); ok {
|
||||||
|
upload.FileName = name
|
||||||
|
}
|
||||||
|
if mode, ok := updates["access_mode"].(int); ok {
|
||||||
|
upload.AccessMode = mode
|
||||||
|
}
|
||||||
|
return upload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func listUploadsForBatchDownload(ctx context.Context, ids []uint64) ([]model.Upload, error) {
|
||||||
|
return repository.ListUploadsByIDs(ctx, ids)
|
||||||
|
}
|
||||||
|
|
||||||
|
type instantUploadInput struct {
|
||||||
|
UserID uint64
|
||||||
|
FileHash string
|
||||||
|
Size int64
|
||||||
|
MimeType string
|
||||||
|
Extension string
|
||||||
|
OrigName string
|
||||||
|
UploadType string
|
||||||
|
AccessMode int
|
||||||
|
}
|
||||||
|
|
||||||
|
func createInstantUpload(ctx context.Context, existing model.Upload, input instantUploadInput) (model.Upload, error) {
|
||||||
|
newUpload := model.Upload{
|
||||||
|
ID: idgen.NextUint64ID(),
|
||||||
|
UserID: input.UserID,
|
||||||
|
FileName: input.OrigName,
|
||||||
|
FilePath: existing.FilePath,
|
||||||
|
FileSize: input.Size,
|
||||||
|
MimeType: input.MimeType,
|
||||||
|
Extension: input.Extension,
|
||||||
|
Hash: input.FileHash,
|
||||||
|
StorageDriver: existing.StorageDriver,
|
||||||
|
Type: input.UploadType,
|
||||||
|
Status: model.UploadStatusUsed,
|
||||||
|
AccessMode: input.AccessMode,
|
||||||
|
Metadata: existing.Metadata,
|
||||||
|
}
|
||||||
|
if err := repository.CreateUpload(ctx, &newUpload); err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
uploadstats.RecordUploadStatsAdd(ctx, &newUpload)
|
||||||
|
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", newUpload.ID, existing.FilePath)
|
||||||
|
return newUpload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func findReusableUpload(ctx context.Context, hash string, size int64) (model.Upload, error) {
|
||||||
|
return repository.FindReusableUploadByHash(ctx, hash, size)
|
||||||
|
}
|
||||||
|
|
||||||
|
func saveNewUploadRecord(ctx context.Context, upload *model.Upload, storageDriver, filePath string) error {
|
||||||
|
if err := repository.CreateUpload(ctx, upload); err != nil {
|
||||||
|
backend, backendErr := storage.ForDriver(ctx, storage.Driver(storageDriver))
|
||||||
|
if backendErr == nil {
|
||||||
|
if deleteErr := backend.Delete(ctx, filePath); deleteErr != nil {
|
||||||
|
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
uploadstats.RecordUploadStatsAdd(ctx, upload)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadUploadStats(ctx context.Context) ([]model.UploadStat, error) {
|
||||||
|
return repository.ListUploadStats(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
var errUploadForbidden = errors.New("upload forbidden")
|
||||||
|
|
||||||
|
func storeUploadObject(ctx context.Context, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, error) {
|
||||||
|
if uploadstorage.ReadOnly(ctx) {
|
||||||
|
return "", "", errors.New(shared.ErrStorageReadOnly)
|
||||||
|
}
|
||||||
|
driver, backend, err := storage.Active(ctx)
|
||||||
|
if err != nil {
|
||||||
|
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
|
||||||
|
return "", "", errors.New(shared.ErrSaveFileFailed)
|
||||||
|
}
|
||||||
|
result, err := backend.Put(ctx, subPath, bytes.NewReader(buf.Bytes()), size, mimeType)
|
||||||
|
if err != nil {
|
||||||
|
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
|
||||||
|
return "", "", errors.New(shared.ErrSaveFileFailed)
|
||||||
|
}
|
||||||
|
meta.Bucket = result.Bucket
|
||||||
|
return string(driver), result.Key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateUploadAllowedExtension(ctx context.Context, ext string) string {
|
||||||
|
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUploadAllowedExtensions)
|
||||||
|
if err != nil || sc.Value == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
allowedExts := strings.Split(strings.ToLower(sc.Value), ",")
|
||||||
|
for _, allowedExt := range allowedExts {
|
||||||
|
if strings.TrimSpace(allowedExt) == ext {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return shared.ErrUnsupportedFormat
|
||||||
|
}
|
||||||
|
|
||||||
|
func isRecordNotFound(err error) bool {
|
||||||
|
return errors.Is(err, gorm.ErrRecordNotFound)
|
||||||
|
}
|
||||||
@@ -26,16 +26,12 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
|
||||||
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common"
|
"github.com/Rain-kl/Wavelet/internal/common"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/storage"
|
|
||||||
apputil "github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
@@ -68,7 +64,7 @@ func UploadFile(c *gin.Context) {
|
|||||||
|
|
||||||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
|
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
|
||||||
|
|
||||||
currUser, _ := apputil.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
header, err := c.FormFile("file")
|
header, err := c.FormFile("file")
|
||||||
@@ -95,7 +91,7 @@ func UploadFile(c *gin.Context) {
|
|||||||
ext = "bin"
|
ext = "bin"
|
||||||
}
|
}
|
||||||
|
|
||||||
if errMsg := validateUploadExtension(ctx, ext); errMsg != "" {
|
if errMsg := validateUploadAllowedExtension(ctx, ext); errMsg != "" {
|
||||||
response.AbortBadRequest(c, errMsg)
|
response.AbortBadRequest(c, errMsg)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -142,9 +138,9 @@ func UploadFile(c *gin.Context) {
|
|||||||
id := idgen.NextUint64ID()
|
id := idgen.NextUint64ID()
|
||||||
subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
|
subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
|
||||||
|
|
||||||
storageDriver, subPath, errMsg := storeUploadFile(ctx, subPath, size, mimeType, &buf, &meta)
|
storageDriver, subPath, err := storeUploadObject(ctx, subPath, size, mimeType, &buf, &meta)
|
||||||
if errMsg != "" {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, errMsg)
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -164,8 +160,8 @@ func UploadFile(c *gin.Context) {
|
|||||||
Metadata: meta,
|
Metadata: meta,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := saveUploadRecord(ctx, &newUpload, storageDriver, subPath); err != "" {
|
if err := saveNewUploadRecord(ctx, &newUpload, storageDriver, subPath); err != nil {
|
||||||
response.AbortBadRequest(c, err)
|
response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -253,8 +249,8 @@ func BatchDownloadFiles(c *gin.Context) {
|
|||||||
ids = append(ids, id)
|
ids = append(ids, id)
|
||||||
}
|
}
|
||||||
|
|
||||||
var uploads []model.Upload
|
uploads, err := listUploadsForBatchDownload(ctx, ids)
|
||||||
if err := db.DB(ctx).Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed).Find(&uploads).Error; err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, shared.ErrRetrieveUploadRecordsFailed)
|
response.AbortBadRequest(c, shared.ErrRetrieveUploadRecordsFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -325,27 +321,8 @@ func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) {
|
|||||||
return accessMode, ""
|
return accessMode, ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func validateUploadExtension(ctx context.Context, ext string) string {
|
|
||||||
var sc model.SystemConfig
|
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeyUploadAllowedExtensions); err == nil && sc.Value != "" {
|
|
||||||
allowedExts := strings.Split(strings.ToLower(sc.Value), ",")
|
|
||||||
allowed := false
|
|
||||||
for _, allowedExt := range allowedExts {
|
|
||||||
if strings.TrimSpace(allowedExt) == ext {
|
|
||||||
allowed = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !allowed {
|
|
||||||
return shared.ErrUnsupportedFormat
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, fileHash string, size int64, mimeType, ext, origName string, accessMode int) (bool, error) {
|
func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, fileHash string, size int64, mimeType, ext, origName string, accessMode int) (bool, error) {
|
||||||
var existing model.Upload
|
existing, err := findReusableUpload(ctx, fileHash, size)
|
||||||
err := db.DB(ctx).Where("hash = ? AND file_size = ? AND status IN (?, ?)", fileHash, size, model.UploadStatusPending, model.UploadStatusUsed).First(&existing).Error
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
@@ -354,52 +331,24 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User,
|
|||||||
return true, nil
|
return true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
id := idgen.NextUint64ID()
|
newUpload, err := createInstantUpload(ctx, existing, instantUploadInput{
|
||||||
newUpload := model.Upload{
|
UserID: currUser.ID,
|
||||||
ID: id,
|
FileHash: fileHash,
|
||||||
UserID: currUser.ID,
|
Size: size,
|
||||||
FileName: origName,
|
MimeType: mimeType,
|
||||||
FilePath: existing.FilePath,
|
Extension: ext,
|
||||||
FileSize: size,
|
OrigName: origName,
|
||||||
MimeType: mimeType,
|
UploadType: c.DefaultPostForm("type", "generic"),
|
||||||
Extension: ext,
|
AccessMode: accessMode,
|
||||||
Hash: fileHash,
|
})
|
||||||
StorageDriver: existing.StorageDriver,
|
if err != nil {
|
||||||
Type: c.DefaultPostForm("type", "generic"),
|
|
||||||
Status: model.UploadStatusUsed,
|
|
||||||
AccessMode: accessMode,
|
|
||||||
Metadata: existing.Metadata,
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := db.DB(ctx).Create(&newUpload).Error; err != nil {
|
|
||||||
response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed)
|
response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed)
|
||||||
return true, err
|
return true, err
|
||||||
}
|
}
|
||||||
uploadstats.RecordUploadStatsAdd(ctx, &newUpload)
|
|
||||||
|
|
||||||
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", id, existing.FilePath)
|
|
||||||
c.JSON(http.StatusOK, response.OK(newUpload))
|
c.JSON(http.StatusOK, response.OK(newUpload))
|
||||||
return true, nil
|
return true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func storeUploadFile(ctx context.Context, subPath string, size int64, mimeType string, buf *bytes.Buffer, meta *model.UploadMetadata) (string, string, string) {
|
|
||||||
if uploadstorage.ReadOnly(ctx) {
|
|
||||||
return "", "", shared.ErrStorageReadOnly
|
|
||||||
}
|
|
||||||
driver, backend, err := storage.Active(ctx)
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
|
|
||||||
return "", "", shared.ErrSaveFileFailed
|
|
||||||
}
|
|
||||||
result, err := backend.Put(ctx, subPath, bytes.NewReader(buf.Bytes()), size, mimeType)
|
|
||||||
if err != nil {
|
|
||||||
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
|
|
||||||
return "", "", shared.ErrSaveFileFailed
|
|
||||||
}
|
|
||||||
meta.Bucket = result.Bucket
|
|
||||||
return string(driver), result.Key, ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) {
|
func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) {
|
||||||
var meta model.UploadMetadata
|
var meta model.UploadMetadata
|
||||||
metadataStr := c.DefaultPostForm("metadata", "")
|
metadataStr := c.DefaultPostForm("metadata", "")
|
||||||
@@ -422,16 +371,3 @@ func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64)
|
|||||||
return mimeType
|
return mimeType
|
||||||
}
|
}
|
||||||
|
|
||||||
func saveUploadRecord(ctx context.Context, upload *model.Upload, storageDriver, filePath string) string {
|
|
||||||
if err := db.DB(ctx).Create(upload).Error; err != nil {
|
|
||||||
backend, backendErr := storage.ForDriver(ctx, storage.Driver(storageDriver))
|
|
||||||
if backendErr == nil {
|
|
||||||
if deleteErr := backend.Delete(ctx, filePath); deleteErr != nil {
|
|
||||||
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return shared.ErrSaveUploadRecordFailed
|
|
||||||
}
|
|
||||||
uploadstats.RecordUploadStatsAdd(ctx, upload)
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
@@ -21,13 +21,13 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
|
||||||
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"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/storage"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -43,7 +43,7 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
|
|||||||
|
|
||||||
authMiddleware := func(c *gin.Context) {
|
authMiddleware := func(c *gin.Context) {
|
||||||
if authUser != nil {
|
if authUser != nil {
|
||||||
util.SetToContext(c, oauth.UserObjKey, authUser)
|
oauth.SetToContext(c, oauth.UserObjKey, authUser)
|
||||||
}
|
}
|
||||||
c.Next()
|
c.Next()
|
||||||
}
|
}
|
||||||
@@ -297,8 +297,8 @@ func TestUploadFile(t *testing.T) {
|
|||||||
dbConn.Where("key = ?", model.ConfigKeyUploadAllowedExtensions).First(&sc)
|
dbConn.Where("key = ?", model.ConfigKeyUploadAllowedExtensions).First(&sc)
|
||||||
sc.Value = "jpg,png,webp,txt"
|
sc.Value = "jpg,png,webp,txt"
|
||||||
dbConn.Save(&sc)
|
dbConn.Save(&sc)
|
||||||
_ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, sc.Key, &sc)
|
_ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, sc.Key, &sc)
|
||||||
model.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
|
|
||||||
contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{
|
contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{
|
||||||
"type": "document",
|
"type": "document",
|
||||||
@@ -998,4 +998,3 @@ func TestUserUploadManagement(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"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/model"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
@@ -48,8 +47,8 @@ type fileStatsResponse struct {
|
|||||||
func GetFileStats(c *gin.Context) {
|
func GetFileStats(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
var stats []model.UploadStat
|
stats, err := loadUploadStats(ctx)
|
||||||
if err := db.DB(ctx).Find(&stats).Error; err != nil {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,17 +5,19 @@
|
|||||||
// Package user 提供用户认证与帐户管理功能
|
// Package user 提供用户认证与帐户管理功能
|
||||||
package user
|
package user
|
||||||
|
|
||||||
import ("net/http"
|
import (
|
||||||
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
|
)
|
||||||
|
|
||||||
type createTokenRequest struct {
|
type createTokenRequest struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
@@ -38,7 +40,7 @@ type tokenResponse struct {
|
|||||||
// @Router /api/v1/user/access-tokens [get]
|
// @Router /api/v1/user/access-tokens [get]
|
||||||
// ListAccessTokens 获取当前用户的 AccessToken 列表
|
// ListAccessTokens 获取当前用户的 AccessToken 列表
|
||||||
func ListAccessTokens(c *gin.Context) {
|
func ListAccessTokens(c *gin.Context) {
|
||||||
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
var tokens []model.AccessToken
|
var tokens []model.AccessToken
|
||||||
@@ -62,7 +64,7 @@ func ListAccessTokens(c *gin.Context) {
|
|||||||
// @Failure 400 {object} response.Any "参数错误或超限"
|
// @Failure 400 {object} response.Any "参数错误或超限"
|
||||||
// @Router /api/v1/user/access-tokens [post]
|
// @Router /api/v1/user/access-tokens [post]
|
||||||
func CreateAccessToken(c *gin.Context) {
|
func CreateAccessToken(c *gin.Context) {
|
||||||
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
var req createTokenRequest
|
var req createTokenRequest
|
||||||
@@ -85,7 +87,7 @@ func CreateAccessToken(c *gin.Context) {
|
|||||||
|
|
||||||
// 检查最大限制(基于 ConfigKeyMaxAPIKeysPerUser 配置,默认值为 5)
|
// 检查最大限制(基于 ConfigKeyMaxAPIKeysPerUser 配置,默认值为 5)
|
||||||
maxLimit := 5
|
maxLimit := 5
|
||||||
if val, err := model.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil {
|
if val, err := repository.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil {
|
||||||
maxLimit = val
|
maxLimit = val
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -140,7 +142,7 @@ func CreateAccessToken(c *gin.Context) {
|
|||||||
// @Failure 400 {object} response.Any "参数错误"
|
// @Failure 400 {object} response.Any "参数错误"
|
||||||
// @Router /api/v1/user/access-tokens/{id} [delete]
|
// @Router /api/v1/user/access-tokens/{id} [delete]
|
||||||
func DeleteAccessToken(c *gin.Context) {
|
func DeleteAccessToken(c *gin.Context) {
|
||||||
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
idStr := c.Param("id")
|
idStr := c.Param("id")
|
||||||
@@ -175,7 +177,7 @@ func DeleteAccessToken(c *gin.Context) {
|
|||||||
// @Failure 400 {object} response.Any "参数错误"
|
// @Failure 400 {object} response.Any "参数错误"
|
||||||
// @Router /api/v1/user/access-tokens/{id}/rotate [post]
|
// @Router /api/v1/user/access-tokens/{id}/rotate [post]
|
||||||
func RotateAccessToken(c *gin.Context) {
|
func RotateAccessToken(c *gin.Context) {
|
||||||
currUser, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
|
|
||||||
idStr := c.Param("id")
|
idStr := c.Param("id")
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"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/internal/task"
|
||||||
pkgu "github.com/Rain-kl/Wavelet/pkg/util"
|
pkgu "github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
@@ -47,8 +48,32 @@ type updateProfileInput struct {
|
|||||||
Location string
|
Location string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isPasswordLoginEnabled(ctx context.Context) bool {
|
||||||
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled)
|
||||||
|
if err != nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return enabled
|
||||||
|
}
|
||||||
|
|
||||||
|
func isPasswordRegisterEnabled(ctx context.Context) bool {
|
||||||
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled)
|
||||||
|
if err != nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return enabled
|
||||||
|
}
|
||||||
|
|
||||||
|
func isRegistrationEnabled(ctx context.Context) bool {
|
||||||
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||||
|
if err != nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return enabled
|
||||||
|
}
|
||||||
|
|
||||||
func isEmailLoginVerificationEnabled(ctx context.Context) bool {
|
func isEmailLoginVerificationEnabled(ctx context.Context) bool {
|
||||||
enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailLoginVerificationEnabled)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -56,7 +81,7 @@ func isEmailLoginVerificationEnabled(ctx context.Context) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func isEmailRegisterVerificationEnabled(ctx context.Context) bool {
|
func isEmailRegisterVerificationEnabled(ctx context.Context) bool {
|
||||||
enabled, err := model.GetBoolByKey(ctx, model.ConfigKeyEmailRegisterVerificationEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyEmailRegisterVerificationEnabled)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -64,26 +89,14 @@ func isEmailRegisterVerificationEnabled(ctx context.Context) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func isSMTPConfigured(ctx context.Context) bool {
|
func isSMTPConfigured(ctx context.Context) bool {
|
||||||
var host, port, username, password string
|
scHost, errHost := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost)
|
||||||
|
scPort, errPort := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort)
|
||||||
var scHost model.SystemConfig
|
scUser, errUser := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername)
|
||||||
if err := scHost.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil {
|
scPass, errPass := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword)
|
||||||
host = scHost.Value
|
if errHost != nil || errPort != nil || errUser != nil || errPass != nil {
|
||||||
|
return false
|
||||||
}
|
}
|
||||||
var scPort model.SystemConfig
|
return scHost.Value != "" && scPort.Value != "" && scUser.Value != "" && scPass.Value != ""
|
||||||
if err := scPort.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil {
|
|
||||||
port = scPort.Value
|
|
||||||
}
|
|
||||||
var scUser model.SystemConfig
|
|
||||||
if err := scUser.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil {
|
|
||||||
username = scUser.Value
|
|
||||||
}
|
|
||||||
var scPass model.SystemConfig
|
|
||||||
if err := scPass.GetByKey(ctx, model.ConfigKeySMTPPassword); err == nil {
|
|
||||||
password = scPass.Value
|
|
||||||
}
|
|
||||||
|
|
||||||
return host != "" && port != "" && username != "" && password != ""
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func generateVerificationCode() (string, error) {
|
func generateVerificationCode() (string, error) {
|
||||||
@@ -114,11 +127,11 @@ func sendEmailVerificationCode(ctx context.Context, email, scene, templateName s
|
|||||||
codeKey := getEmailCodeKey(scene, email)
|
codeKey := getEmailCodeKey(scene, email)
|
||||||
cooldownKey := getEmailCooldownKey(scene, email)
|
cooldownKey := getEmailCooldownKey(scene, email)
|
||||||
|
|
||||||
emailSubject, emailBody, err := model.RenderTemplate(
|
tmpl, err := repository.GetTemplateByKey(ctx, templateName)
|
||||||
ctx,
|
if err != nil {
|
||||||
templateName,
|
return fmt.Errorf("模板 %s 不存在或不可用: %w", templateName, err)
|
||||||
map[string]any{"Code": code},
|
}
|
||||||
)
|
emailSubject, emailBody, err := tmpl.Render(map[string]any{"Code": code})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf(errRenderEmailTemplateFailed, err)
|
return fmt.Errorf(errRenderEmailTemplateFailed, err)
|
||||||
}
|
}
|
||||||
@@ -271,4 +284,4 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &dbUser, nil
|
return &dbUser, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -39,7 +40,7 @@ func TestProcessLoginEmailVerificationSMTPFallback(t *testing.T) {
|
|||||||
Update("value", "").Error; err != nil {
|
Update("value", "").Error; err != nil {
|
||||||
t.Fatalf("clear SMTP host failed: %v", err)
|
t.Fatalf("clear SMTP host failed: %v", err)
|
||||||
}
|
}
|
||||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil {
|
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -104,7 +105,7 @@ func TestProcessLoginEmailVerificationEmptyEmailFallback(t *testing.T) {
|
|||||||
t.Fatalf("set %s failed: %v", cfg.key, err)
|
t.Fatalf("set %s failed: %v", cfg.key, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil {
|
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -148,4 +149,4 @@ func TestProcessLoginEmailVerificationInvalidCode(t *testing.T) {
|
|||||||
if result.Status != LoginEmailVerificationRejected || result.Message != errEmailCodeInvalidOrExpired {
|
if result.Status != LoginEmailVerificationRejected || result.Message != errEmailCodeInvalidOrExpired {
|
||||||
t.Fatalf("processLoginEmailVerification() = %+v, want rejected invalid code", result)
|
t.Fatalf("processLoginEmailVerification() = %+v, want rejected invalid code", result)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,23 +3,24 @@
|
|||||||
|
|
||||||
package user
|
package user
|
||||||
|
|
||||||
import ("context"
|
import (
|
||||||
|
"context"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common"
|
"github.com/Rain-kl/Wavelet/internal/common"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/util"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type loginRequest struct {
|
type loginRequest struct {
|
||||||
@@ -53,30 +54,6 @@ type updateProfileRequest struct {
|
|||||||
Location string `json:"location"`
|
Location string `json:"location"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func isPasswordLoginEnabled() bool {
|
|
||||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled)
|
|
||||||
if err != nil {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return enabled
|
|
||||||
}
|
|
||||||
|
|
||||||
func isPasswordRegisterEnabled() bool {
|
|
||||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordRegisterEnabled)
|
|
||||||
if err != nil {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return enabled
|
|
||||||
}
|
|
||||||
|
|
||||||
func isRegistrationEnabled() bool {
|
|
||||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyRegistrationEnabled)
|
|
||||||
if err != nil {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return enabled
|
|
||||||
}
|
|
||||||
|
|
||||||
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
||||||
session := sessions.Default(c)
|
session := sessions.Default(c)
|
||||||
session.Set(oauth.UserIDKey, user.ID)
|
session.Set(oauth.UserIDKey, user.ID)
|
||||||
@@ -87,7 +64,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
|||||||
maxAge := config.Config.App.SessionAge
|
maxAge := config.Config.App.SessionAge
|
||||||
isSessionCookie := false
|
isSessionCookie := false
|
||||||
|
|
||||||
ttlHours, err := model.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
switch {
|
switch {
|
||||||
case ttlHours == -1:
|
case ttlHours == -1:
|
||||||
@@ -124,7 +101,8 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
|||||||
// @Failure 500 {object} response.Any "服务内部错误"
|
// @Failure 500 {object} response.Any "服务内部错误"
|
||||||
// @Router /api/v1/user/login [post]
|
// @Router /api/v1/user/login [post]
|
||||||
func Login(c *gin.Context) {
|
func Login(c *gin.Context) {
|
||||||
if !isPasswordLoginEnabled() {
|
ctx := c.Request.Context()
|
||||||
|
if !isPasswordLoginEnabled(ctx) {
|
||||||
response.AbortBadRequest(c, errPasswordLoginDisabled)
|
response.AbortBadRequest(c, errPasswordLoginDisabled)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -140,7 +118,6 @@ func Login(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var user model.User
|
var user model.User
|
||||||
ctx := c.Request.Context()
|
|
||||||
if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil {
|
if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil {
|
||||||
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
|
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
|
||||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||||
@@ -211,7 +188,8 @@ func Login(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "服务内部错误"
|
// @Failure 500 {object} response.Any "服务内部错误"
|
||||||
// @Router /api/v1/user/register [post]
|
// @Router /api/v1/user/register [post]
|
||||||
func Register(c *gin.Context) {
|
func Register(c *gin.Context) {
|
||||||
if !isRegistrationEnabled() || !isPasswordRegisterEnabled() {
|
ctx := c.Request.Context()
|
||||||
|
if !isRegistrationEnabled(ctx) || !isPasswordRegisterEnabled(ctx) {
|
||||||
response.AbortBadRequest(c, errRegistrationDisabled)
|
response.AbortBadRequest(c, errRegistrationDisabled)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -242,8 +220,6 @@ func Register(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
|
||||||
|
|
||||||
// 邮箱注册验证校验
|
// 邮箱注册验证校验
|
||||||
if err := validateRegisterEmailVerification(ctx, req.Email, req.Code); err != nil {
|
if err := validateRegisterEmailVerification(ctx, req.Email, req.Code); err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
@@ -344,7 +320,7 @@ func ChangePassword(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
if userObj == nil {
|
if userObj == nil {
|
||||||
response.AbortUnauthorized(c, errLoginRequired)
|
response.AbortUnauthorized(c, errLoginRequired)
|
||||||
return
|
return
|
||||||
@@ -443,7 +419,7 @@ func UpdateProfile(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey)
|
userObj, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
|
||||||
if userObj == nil {
|
if userObj == nil {
|
||||||
response.AbortUnauthorized(c, errLoginRequired)
|
response.AbortUnauthorized(c, errLoginRequired)
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -3,7 +3,8 @@
|
|||||||
|
|
||||||
package user
|
package user
|
||||||
|
|
||||||
import ("bytes"
|
import (
|
||||||
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -16,12 +17,14 @@ import ("bytes"
|
|||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-contrib/sessions/cookie"
|
"github.com/gin-contrib/sessions/cookie"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
|
)
|
||||||
|
|
||||||
func setupUserTestRouter(t *testing.T) *gin.Engine {
|
func setupUserTestRouter(t *testing.T) *gin.Engine {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
@@ -281,7 +284,7 @@ func TestLoginEmailVerificationFallbackWhenSMTPUnconfigured(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 2.5 Invalidate the system config cache in Redis
|
// 2.5 Invalidate the system config cache in Redis
|
||||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil {
|
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -387,7 +390,7 @@ func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Invalidate the system config cache in Redis
|
// Invalidate the system config cache in Redis
|
||||||
if err := db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err(); err != nil {
|
if err := db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Err(); err != nil {
|
||||||
t.Fatalf("invalidate system config cache failed: %v", err)
|
t.Fatalf("invalidate system config cache failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"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/internal/task"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/mail"
|
"github.com/Rain-kl/Wavelet/pkg/mail"
|
||||||
)
|
)
|
||||||
@@ -111,17 +112,16 @@ func (h *SendEmailHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
|||||||
var smtpUsername string
|
var smtpUsername string
|
||||||
var smtpPassword string
|
var smtpPassword string
|
||||||
|
|
||||||
var sc model.SystemConfig
|
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPHost); err == nil {
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil {
|
|
||||||
smtpHost = sc.Value
|
smtpHost = sc.Value
|
||||||
}
|
}
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil {
|
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPort); err == nil {
|
||||||
smtpPortVal = sc.Value
|
smtpPortVal = sc.Value
|
||||||
}
|
}
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil {
|
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPUsername); err == nil {
|
||||||
smtpUsername = sc.Value
|
smtpUsername = sc.Value
|
||||||
}
|
}
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeySMTPPassword); err == nil {
|
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySMTPPassword); err == nil {
|
||||||
smtpPassword = sc.Value
|
smtpPassword = sc.Value
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/pressly/goose/v3"
|
"github.com/pressly/goose/v3"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -69,7 +69,7 @@ func Migrate() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func clearSystemConfigCache() {
|
func clearSystemConfigCache() {
|
||||||
if err := model.InvalidateAllSystemConfigCaches(context.Background()); err != nil {
|
if err := repository.InvalidateAllSystemConfigCaches(context.Background()); err != nil {
|
||||||
log.Printf("[%s] clear system config cache failed: %v\n", dbType(), err)
|
log.Printf("[%s] clear system config cache failed: %v\n", dbType(), err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/alicebob/miniredis/v2"
|
"github.com/alicebob/miniredis/v2"
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"github.com/redis/go-redis/v9"
|
"github.com/redis/go-redis/v9"
|
||||||
@@ -107,13 +108,13 @@ func TestMigrateClearsStaleSystemConfigCache(t *testing.T) {
|
|||||||
Value: "true",
|
Value: "true",
|
||||||
Type: "system",
|
Type: "system",
|
||||||
}
|
}
|
||||||
if err := db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, model.ConfigKeyCapLoginEnabled, &staleConfig); err != nil {
|
if err := db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, model.ConfigKeyCapLoginEnabled, &staleConfig); err != nil {
|
||||||
t.Fatalf("HSetJSON() error = %v", err)
|
t.Fatalf("HSetJSON() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
Migrate()
|
Migrate()
|
||||||
|
|
||||||
exists, err := db.Redis.Exists(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Result()
|
exists, err := db.Redis.Exists(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)).Result()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Redis.Exists() error = %v", err)
|
t.Fatalf("Redis.Exists() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -121,7 +122,7 @@ func TestMigrateClearsStaleSystemConfigCache(t *testing.T) {
|
|||||||
t.Fatalf("system config cache exists = %d, want 0", exists)
|
t.Fatalf("system config cache exists = %d, want 0", exists)
|
||||||
}
|
}
|
||||||
|
|
||||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyCapLoginEnabled)
|
enabled, err := repository.GetBoolByKey(context.Background(), model.ConfigKeyCapLoginEnabled)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetBoolByKey(%s) error = %v", model.ConfigKeyCapLoginEnabled, err)
|
t.Fatalf("GetBoolByKey(%s) error = %v", model.ConfigKeyCapLoginEnabled, err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
pkgcache "github.com/Rain-kl/Wavelet/pkg/cache/disk"
|
pkgcache "github.com/Rain-kl/Wavelet/pkg/cache/disk"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -69,27 +70,24 @@ func (c *DiskCache) ReloadConfig(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 1. Max Size
|
// 1. Max Size
|
||||||
var scMaxSize model.SystemConfig
|
|
||||||
maxSizeMB := int64(defaultMaxSizeMB)
|
maxSizeMB := int64(defaultMaxSizeMB)
|
||||||
if err := scMaxSize.GetByKey(ctx, model.ConfigKeyDiskCacheMaxSizeMB); err == nil && scMaxSize.Value != "" {
|
if scMaxSize, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyDiskCacheMaxSizeMB); err == nil && scMaxSize.Value != "" {
|
||||||
if val, err := strconv.ParseInt(scMaxSize.Value, 10, 64); err == nil && val > 0 {
|
if val, err := strconv.ParseInt(scMaxSize.Value, 10, 64); err == nil && val > 0 {
|
||||||
maxSizeMB = val
|
maxSizeMB = val
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 2. Default TTL
|
// 2. Default TTL
|
||||||
var scTTL model.SystemConfig
|
|
||||||
ttlMinutes := int64(defaultTTLMinutes)
|
ttlMinutes := int64(defaultTTLMinutes)
|
||||||
if err := scTTL.GetByKey(ctx, model.ConfigKeyDiskCacheTTLMinutes); err == nil && scTTL.Value != "" {
|
if scTTL, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyDiskCacheTTLMinutes); err == nil && scTTL.Value != "" {
|
||||||
if val, err := strconv.ParseInt(scTTL.Value, 10, 64); err == nil && val >= 0 {
|
if val, err := strconv.ParseInt(scTTL.Value, 10, 64); err == nil && val >= 0 {
|
||||||
ttlMinutes = val
|
ttlMinutes = val
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 3. LRU Enabled
|
// 3. LRU Enabled
|
||||||
var scLRU model.SystemConfig
|
|
||||||
lruEnabled := true
|
lruEnabled := true
|
||||||
if err := scLRU.GetByKey(ctx, model.ConfigKeyDiskCacheLRUEnabled); err == nil && scLRU.Value != "" {
|
if scLRU, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyDiskCacheLRUEnabled); err == nil && scLRU.Value != "" {
|
||||||
if val, err := strconv.ParseBool(scLRU.Value); err == nil {
|
if val, err := strconv.ParseBool(scLRU.Value); err == nil {
|
||||||
lruEnabled = val
|
lruEnabled = val
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -31,7 +32,7 @@ func TestDiskCacheReloadConfig(t *testing.T) {
|
|||||||
|
|
||||||
// Invalidate Redis config cache to force DB reload
|
// Invalidate Redis config cache to force DB reload
|
||||||
if db.Redis != nil {
|
if db.Redis != nil {
|
||||||
db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey))
|
db.Redis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reload config
|
// Reload config
|
||||||
|
|||||||
@@ -4,15 +4,11 @@
|
|||||||
package model
|
package model
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -70,8 +66,6 @@ func (pc *PushChannel) Validate() error {
|
|||||||
return errors.New("request URL/address is required")
|
return errors.New("request URL/address is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
// For custom and lark, we must enforce https:// URL prefix for security.
|
|
||||||
// For email, it is an SMTP host:port, so no need for https:// prefix.
|
|
||||||
if pc.Type != TypeEmail && !strings.HasPrefix(pc.URL, "https://") {
|
if pc.Type != TypeEmail && !strings.HasPrefix(pc.URL, "https://") {
|
||||||
return errors.New("request URL must use HTTPS protocol for security reasons")
|
return errors.New("request URL must use HTTPS protocol for security reasons")
|
||||||
}
|
}
|
||||||
@@ -102,58 +96,4 @@ func validateJSON(s string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return errors.New("payload schema must be a valid JSON format")
|
return errors.New("payload schema must be a valid JSON format")
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetPushChannelByName 根据名称获取消息通道
|
|
||||||
func GetPushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
|
|
||||||
var channel PushChannel
|
|
||||||
err := db.DB(ctx).Where("name = ?", name).First(&channel).Error
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &channel, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
const activePushChannelCacheTTL = 24 * time.Hour
|
|
||||||
|
|
||||||
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)
|
|
||||||
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
|
|
||||||
cacheKey := "push:channel:active:" + name
|
|
||||||
var channel PushChannel
|
|
||||||
if db.Redis != nil {
|
|
||||||
if err := db.GetJSON(ctx, cacheKey, &channel); err == nil {
|
|
||||||
return &channel, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if db.Redis != nil {
|
|
||||||
// 缓存有效时间设置为 24 小时
|
|
||||||
_ = db.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
|
||||||
}
|
|
||||||
|
|
||||||
return &channel, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeleteActivePushChannelCache 清理启用消息通道的缓存
|
|
||||||
func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
|
||||||
if db.Redis != nil {
|
|
||||||
_ = db.Redis.Del(ctx, db.PrefixedKey("push:channel:active:"+name)).Err()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// AfterSave GORM 保存后钩子,用于自动清理缓存
|
|
||||||
func (pc *PushChannel) AfterSave(tx *gorm.DB) error {
|
|
||||||
DeleteActivePushChannelCache(tx.Statement.Context, pc.Name)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AfterDelete GORM 删除后钩子,用于自动清理缓存
|
|
||||||
func (pc *PushChannel) AfterDelete(tx *gorm.DB) error {
|
|
||||||
DeleteActivePushChannelCache(tx.Statement.Context, pc.Name)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -4,13 +4,9 @@
|
|||||||
package model
|
package model
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// PushEvent 系统通知事件模型
|
// PushEvent 系统通知事件模型
|
||||||
@@ -51,48 +47,4 @@ func (pe *PushEvent) Validate() error {
|
|||||||
return errors.New("cannot enable event without any push channels configured")
|
return errors.New("cannot enable event without any push channels configured")
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
const activePushEventCacheTTL = 24 * time.Hour
|
|
||||||
|
|
||||||
// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)
|
|
||||||
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
|
|
||||||
cacheKey := "push:event:active:" + key
|
|
||||||
var event PushEvent
|
|
||||||
if db.Redis != nil {
|
|
||||||
if err := db.GetJSON(ctx, cacheKey, &event); err == nil {
|
|
||||||
return &event, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if db.Redis != nil {
|
|
||||||
// 缓存有效时间设置为 24 小时
|
|
||||||
_ = db.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
|
|
||||||
}
|
|
||||||
|
|
||||||
return &event, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeleteActivePushEventCache 清理启用通知事件的缓存
|
|
||||||
func DeleteActivePushEventCache(ctx context.Context, key string) {
|
|
||||||
if db.Redis != nil {
|
|
||||||
_ = db.Redis.Del(ctx, db.PrefixedKey("push:event:active:"+key)).Err()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// AfterSave GORM 保存后钩子,用于自动清理缓存
|
|
||||||
func (pe *PushEvent) AfterSave(tx *gorm.DB) error {
|
|
||||||
DeleteActivePushEventCache(tx.Statement.Context, pe.EventKey)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AfterDelete GORM 删除后钩子,用于自动清理缓存
|
|
||||||
func (pe *PushEvent) AfterDelete(tx *gorm.DB) error {
|
|
||||||
DeleteActivePushEventCache(tx.Statement.Context, pe.EventKey)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -3,19 +3,7 @@
|
|||||||
|
|
||||||
package model
|
package model
|
||||||
|
|
||||||
import (
|
import "time"
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"strconv"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/redis/go-redis/v9"
|
|
||||||
"github.com/shopspring/decimal"
|
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
)
|
|
||||||
|
|
||||||
// 配置键常量 - 所有系统配置的 key 定义
|
// 配置键常量 - 所有系统配置的 key 定义
|
||||||
const (
|
const (
|
||||||
@@ -51,13 +39,6 @@ const (
|
|||||||
ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON)
|
ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON)
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
|
||||||
// SystemConfigRedisHashKey Redis Hash key,存储所有系统配置
|
|
||||||
SystemConfigRedisHashKey = "system:system_configs"
|
|
||||||
// SystemConfigVisibleListRedisKey Redis key,缓存所有 visibility=1 的公共配置列表
|
|
||||||
SystemConfigVisibleListRedisKey = "system:visible_configs"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
|
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
|
||||||
ConfigVisibilityHidden = 0
|
ConfigVisibilityHidden = 0
|
||||||
@@ -79,181 +60,4 @@ type SystemConfig struct {
|
|||||||
// TableName 表名
|
// TableName 表名
|
||||||
func (SystemConfig) TableName() string {
|
func (SystemConfig) TableName() string {
|
||||||
return "w_system_configs"
|
return "w_system_configs"
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetByKey 通过 key 查询配置(带 RAM + Redis 缓存)
|
|
||||||
func (sc *SystemConfig) GetByKey(ctx context.Context, key string) error {
|
|
||||||
ensureSystemConfigCacheListener()
|
|
||||||
|
|
||||||
if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok {
|
|
||||||
*sc = cloneSystemConfig(cached)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if db.Redis != nil {
|
|
||||||
if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, sc); err == nil {
|
|
||||||
systemConfigRAMCache.Set(key, cloneSystemConfig(*sc))
|
|
||||||
return nil
|
|
||||||
} else if !errors.Is(err, redis.Nil) {
|
|
||||||
// Redis 服务错误,返回错误
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 查数据库
|
|
||||||
database := db.DB(ctx)
|
|
||||||
if database == nil {
|
|
||||||
return errors.New(errDatabaseNotInitialized)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := database.Where("key = ?", key).First(sc).Error; err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
populateSystemConfigCache(ctx, *sc)
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ListSystemConfigsByKeys loads multiple config keys in one database round trip.
|
|
||||||
// Keys already present in the process-local RAM cache are returned without querying PostgreSQL.
|
|
||||||
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]SystemConfig, error) {
|
|
||||||
if len(keys) == 0 {
|
|
||||||
return map[string]SystemConfig{}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
ensureSystemConfigCacheListener()
|
|
||||||
|
|
||||||
result := make(map[string]SystemConfig, len(keys))
|
|
||||||
missing := make([]string, 0, len(keys))
|
|
||||||
for _, key := range keys {
|
|
||||||
if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok {
|
|
||||||
result[key] = cloneSystemConfig(cached)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
missing = append(missing, key)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(missing) == 0 {
|
|
||||||
return result, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
database := db.DB(ctx)
|
|
||||||
if database == nil {
|
|
||||||
return nil, errors.New(errDatabaseNotInitialized)
|
|
||||||
}
|
|
||||||
|
|
||||||
var configs []SystemConfig
|
|
||||||
if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := range configs {
|
|
||||||
populateSystemConfigCache(ctx, configs[i])
|
|
||||||
result[configs[i].Key] = cloneSystemConfig(configs[i])
|
|
||||||
}
|
|
||||||
|
|
||||||
return result, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
|
|
||||||
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
|
|
||||||
if db.Redis == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return db.Redis.Del(ctx, db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
// ListVisibleSystemConfigs 查询所有可通过公共配置接口暴露的配置(带 Redis 列表缓存)
|
|
||||||
func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
|
||||||
if db.Redis != nil {
|
|
||||||
var cached []SystemConfig
|
|
||||||
if err := db.GetJSON(ctx, SystemConfigVisibleListRedisKey, &cached); err == nil {
|
|
||||||
return cached, nil
|
|
||||||
} else if !errors.Is(err, redis.Nil) {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
database := db.DB(ctx)
|
|
||||||
if database == nil {
|
|
||||||
return nil, errors.New(errDatabaseNotInitialized)
|
|
||||||
}
|
|
||||||
|
|
||||||
var configs []SystemConfig
|
|
||||||
if err := database.Where("visibility = ?", ConfigVisibilityVisible).Find(&configs).Error; err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if db.Redis != nil {
|
|
||||||
_ = db.SetJSON(ctx, SystemConfigVisibleListRedisKey, configs, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
return configs, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetIntByKey 通过 key 查询配置并转换为 int 类型
|
|
||||||
func GetIntByKey(ctx context.Context, key string) (int, error) {
|
|
||||||
var sc SystemConfig
|
|
||||||
if err := sc.GetByKey(ctx, key); err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
|
|
||||||
value, err := strconv.Atoi(sc.Value)
|
|
||||||
if err != nil {
|
|
||||||
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return value, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetDecimalByKey 通过 key 查询配置并转换为 decimal.Decimal 类型
|
|
||||||
// precision 指定保留的小数位数,多余的小数会被裁剪
|
|
||||||
func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) {
|
|
||||||
var sc SystemConfig
|
|
||||||
if err := sc.GetByKey(ctx, key); err != nil {
|
|
||||||
return decimal.Zero, err
|
|
||||||
}
|
|
||||||
|
|
||||||
value, err := decimal.NewFromString(sc.Value)
|
|
||||||
if err != nil {
|
|
||||||
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 裁剪到指定小数位数
|
|
||||||
return value.Truncate(precision), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetBoolByKey 通过 key 查询配置并转换为 bool 类型
|
|
||||||
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
|
|
||||||
var sc SystemConfig
|
|
||||||
if err := sc.GetByKey(ctx, key); err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
|
|
||||||
value, err := strconv.ParseBool(sc.Value)
|
|
||||||
if err != nil {
|
|
||||||
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return value, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetMenuDisplayConfig 获取目录显示配置,解析为 map[string]bool
|
|
||||||
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
|
|
||||||
var sc SystemConfig
|
|
||||||
if err := sc.GetByKey(ctx, ConfigKeyMenuDisplayConfig); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
config := make(map[string]bool)
|
|
||||||
if sc.Value == "" || sc.Value == "{}" {
|
|
||||||
return config, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
|
|
||||||
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return config, nil
|
|
||||||
}
|
|
||||||
@@ -5,14 +5,10 @@ package model
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
|
||||||
"strings"
|
"strings"
|
||||||
"text/template"
|
"text/template"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Template 邮件/消息模板实体
|
// Template 邮件/消息模板实体
|
||||||
@@ -90,17 +86,3 @@ func (t *Template) Render(data any) (string, string, error) {
|
|||||||
|
|
||||||
return subject, bodyBuf.String(), nil
|
return subject, bodyBuf.String(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// RenderTemplate 渲染指定模板。模板不存在或渲染失败时返回错误,由调用方决定如何处理。
|
|
||||||
func RenderTemplate(ctx context.Context, key string, data any) (string, string, error) {
|
|
||||||
var t Template
|
|
||||||
if err := db.DB(ctx).Where("key = ?", key).First(&t).Error; err != nil {
|
|
||||||
return "", "", fmt.Errorf(errTemplateUnavailable, key, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
subject, body, err := t.Render(data)
|
|
||||||
if err != nil {
|
|
||||||
return "", "", fmt.Errorf(errTemplateRenderFailed, key, err)
|
|
||||||
}
|
|
||||||
return subject, body, nil
|
|
||||||
}
|
|
||||||
|
|||||||
+2
-12
@@ -139,12 +139,7 @@ func (u *User) assignIDIfMissing() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验)
|
// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验)
|
||||||
func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
|
func (u *User) CreateUser(_ context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
|
||||||
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
|
|
||||||
if err == nil && !enabled {
|
|
||||||
return errors.New(errRegistrationDisabled)
|
|
||||||
}
|
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
userID := oauthInfo.GetID()
|
userID := oauthInfo.GetID()
|
||||||
newUser := User{
|
newUser := User{
|
||||||
@@ -169,12 +164,7 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser
|
|||||||
}
|
}
|
||||||
|
|
||||||
// RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验)
|
// RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验)
|
||||||
func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error {
|
func (u *User) RegisterUser(_ context.Context, tx *gorm.DB) error {
|
||||||
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
|
|
||||||
if err == nil && !enabled {
|
|
||||||
return errors.New(errRegistrationDisabled)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 检查用户名冲突
|
// 检查用户名冲突
|
||||||
var count int64
|
var count int64
|
||||||
if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil {
|
if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,105 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
const activePushChannelCacheTTL = 24 * time.Hour
|
||||||
|
|
||||||
|
// ListPushChannels returns all push channels ordered by creation time descending.
|
||||||
|
func ListPushChannels(ctx context.Context) ([]model.PushChannel, error) {
|
||||||
|
var channels []model.PushChannel
|
||||||
|
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return channels, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPushChannelByID loads a push channel by primary key.
|
||||||
|
func GetPushChannelByID(ctx context.Context, id uint64) (model.PushChannel, error) {
|
||||||
|
var channel model.PushChannel
|
||||||
|
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
||||||
|
return model.PushChannel{}, err
|
||||||
|
}
|
||||||
|
return channel, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPushChannelByName 根据名称获取消息通道。
|
||||||
|
func GetPushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) {
|
||||||
|
var channel model.PushChannel
|
||||||
|
if err := db.DB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &channel, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CountPushChannelsByName returns how many channels share the given name.
|
||||||
|
func CountPushChannelsByName(ctx context.Context, name string) (int64, error) {
|
||||||
|
var count int64
|
||||||
|
if err := db.DB(ctx).Model(&model.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return count, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreatePushChannel persists a new channel and invalidates cache.
|
||||||
|
func CreatePushChannel(ctx context.Context, channel *model.PushChannel) error {
|
||||||
|
if err := db.DB(ctx).Create(channel).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SavePushChannel updates a channel and invalidates cache.
|
||||||
|
func SavePushChannel(ctx context.Context, channel *model.PushChannel) error {
|
||||||
|
if err := db.DB(ctx).Save(channel).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeletePushChannel removes a channel and invalidates cache.
|
||||||
|
func DeletePushChannel(ctx context.Context, channel *model.PushChannel) error {
|
||||||
|
if err := db.DB(ctx).Delete(channel).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetActivePushChannelByName 根据名称获取启用的消息通道 (优先从 Redis 缓存获取)。
|
||||||
|
func GetActivePushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) {
|
||||||
|
cacheKey := "push:channel:active:" + name
|
||||||
|
var channel model.PushChannel
|
||||||
|
if db.Redis != nil {
|
||||||
|
if err := db.GetJSON(ctx, cacheKey, &channel); err == nil {
|
||||||
|
return &channel, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if db.Redis != nil {
|
||||||
|
_ = db.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &channel, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
|
||||||
|
func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
||||||
|
if db.Redis != nil {
|
||||||
|
_ = db.Redis.Del(ctx, db.PrefixedKey("push:channel:active:"+name)).Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
const activePushEventCacheTTL = 24 * time.Hour
|
||||||
|
|
||||||
|
// ListPushEvents returns all push events ordered by creation time descending.
|
||||||
|
func ListPushEvents(ctx context.Context) ([]model.PushEvent, error) {
|
||||||
|
var events []model.PushEvent
|
||||||
|
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return events, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPushEventByID loads a push event by primary key.
|
||||||
|
func GetPushEventByID(ctx context.Context, id uint64) (model.PushEvent, error) {
|
||||||
|
var event model.PushEvent
|
||||||
|
if err := db.DB(ctx).First(&event, id).Error; err != nil {
|
||||||
|
return model.PushEvent{}, err
|
||||||
|
}
|
||||||
|
return event, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPushEventByKey loads a push event by event key.
|
||||||
|
func GetPushEventByKey(ctx context.Context, key string) (model.PushEvent, error) {
|
||||||
|
var event model.PushEvent
|
||||||
|
if err := db.DB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
|
||||||
|
return model.PushEvent{}, err
|
||||||
|
}
|
||||||
|
return event, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CountPushEventsByKey returns how many events use the given event key.
|
||||||
|
func CountPushEventsByKey(ctx context.Context, key string) (int64, error) {
|
||||||
|
var count int64
|
||||||
|
if err := db.DB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return count, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreatePushEvent persists a new push event and invalidates cache.
|
||||||
|
func CreatePushEvent(ctx context.Context, event *model.PushEvent) error {
|
||||||
|
if err := db.DB(ctx).Create(event).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SavePushEvent updates a push event and invalidates cache.
|
||||||
|
func SavePushEvent(ctx context.Context, event *model.PushEvent) error {
|
||||||
|
if err := db.DB(ctx).Save(event).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdatePushEventEnabled toggles the enabled flag for a push event.
|
||||||
|
func UpdatePushEventEnabled(ctx context.Context, event *model.PushEvent, enabled bool) error {
|
||||||
|
event.Enabled = enabled
|
||||||
|
if err := db.DB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeletePushEvent removes a push event and invalidates cache.
|
||||||
|
func DeletePushEvent(ctx context.Context, event *model.PushEvent) error {
|
||||||
|
if err := db.DB(ctx).Delete(event).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListActivePushEventsByTaskType returns enabled events bound to a task type.
|
||||||
|
func ListActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) {
|
||||||
|
var events []model.PushEvent
|
||||||
|
if err := db.DB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return events, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetActivePushEventByKey 获取启用的通知事件 (优先从 Redis 缓存获取)。
|
||||||
|
func GetActivePushEventByKey(ctx context.Context, key string) (*model.PushEvent, error) {
|
||||||
|
cacheKey := "push:event:active:" + key
|
||||||
|
var event model.PushEvent
|
||||||
|
if db.Redis != nil {
|
||||||
|
if err := db.GetJSON(ctx, cacheKey, &event); err == nil {
|
||||||
|
return &event, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if db.Redis != nil {
|
||||||
|
_ = db.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &event, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteActivePushEventCache 清理启用通知事件的缓存。
|
||||||
|
func DeleteActivePushEventCache(ctx context.Context, key string) {
|
||||||
|
if db.Redis != nil {
|
||||||
|
_ = db.Redis.Del(ctx, db.PrefixedKey("push:event:active:"+key)).Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PushHistoryListFilter filters push history pagination queries.
|
||||||
|
type PushHistoryListFilter struct {
|
||||||
|
EventKey string
|
||||||
|
Status string
|
||||||
|
Page int
|
||||||
|
PageSize int
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListPushHistories returns paginated push history records.
|
||||||
|
func ListPushHistories(ctx context.Context, filter PushHistoryListFilter) (int64, []model.PushHistory, error) {
|
||||||
|
query := db.DB(ctx).Model(&model.PushHistory{}).Order("created_at DESC")
|
||||||
|
if filter.EventKey != "" {
|
||||||
|
query = query.Where("event_key = ?", filter.EventKey)
|
||||||
|
}
|
||||||
|
if filter.Status != "" {
|
||||||
|
query = query.Where("status = ?", filter.Status)
|
||||||
|
}
|
||||||
|
|
||||||
|
var total int64
|
||||||
|
if err := query.Count(&total).Error; err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var results []model.PushHistory
|
||||||
|
offset := (filter.Page - 1) * filter.PageSize
|
||||||
|
if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return total, results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreatePushHistory persists a push history audit record.
|
||||||
|
func CreatePushHistory(ctx context.Context, history *model.PushHistory) error {
|
||||||
|
return db.DB(ctx).Create(history).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// PushHistoryQuery returns a scoped query builder for push histories.
|
||||||
|
func PushHistoryQuery(ctx context.Context) *gorm.DB {
|
||||||
|
return db.DB(ctx).Model(&model.PushHistory{})
|
||||||
|
}
|
||||||
@@ -0,0 +1,198 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
// Package repository provides data access with caching and persistence boundaries.
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
|
"github.com/shopspring/decimal"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
errDatabaseNotInitialized = "database not initialized"
|
||||||
|
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
|
||||||
|
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
|
||||||
|
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
|
||||||
|
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
|
||||||
|
)
|
||||||
|
|
||||||
|
// GetSystemConfigByKey 通过 key 查询配置(带 RAM + Redis 缓存)。
|
||||||
|
func GetSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
|
||||||
|
ensureSystemConfigCacheListener()
|
||||||
|
|
||||||
|
if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok {
|
||||||
|
return cloneSystemConfig(cached), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var sc model.SystemConfig
|
||||||
|
if db.Redis != nil {
|
||||||
|
if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, &sc); err == nil {
|
||||||
|
systemConfigRAMCache.Set(key, cloneSystemConfig(sc))
|
||||||
|
return sc, nil
|
||||||
|
} else if !errors.Is(err, redis.Nil) {
|
||||||
|
return model.SystemConfig{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
database := db.DB(ctx)
|
||||||
|
if database == nil {
|
||||||
|
return model.SystemConfig{}, errors.New(errDatabaseNotInitialized)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
|
||||||
|
return model.SystemConfig{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
populateSystemConfigCache(ctx, sc)
|
||||||
|
return sc, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListSystemConfigsByKeys loads multiple config keys in one database round trip.
|
||||||
|
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]model.SystemConfig, error) {
|
||||||
|
if len(keys) == 0 {
|
||||||
|
return map[string]model.SystemConfig{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
ensureSystemConfigCacheListener()
|
||||||
|
|
||||||
|
result := make(map[string]model.SystemConfig, len(keys))
|
||||||
|
missing := make([]string, 0, len(keys))
|
||||||
|
for _, key := range keys {
|
||||||
|
if cached, ok := systemConfigRAMCache.GetIfPresent(key); ok {
|
||||||
|
result[key] = cloneSystemConfig(cached)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
missing = append(missing, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(missing) == 0 {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
database := db.DB(ctx)
|
||||||
|
if database == nil {
|
||||||
|
return nil, errors.New(errDatabaseNotInitialized)
|
||||||
|
}
|
||||||
|
|
||||||
|
var configs []model.SystemConfig
|
||||||
|
if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := range configs {
|
||||||
|
populateSystemConfigCache(ctx, configs[i])
|
||||||
|
result[configs[i].Key] = cloneSystemConfig(configs[i])
|
||||||
|
}
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
|
||||||
|
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
|
||||||
|
if db.Redis == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return db.Redis.Del(ctx, db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListVisibleSystemConfigs 查询所有可通过公共配置接口暴露的配置(带 Redis 列表缓存)。
|
||||||
|
func ListVisibleSystemConfigs(ctx context.Context) ([]model.SystemConfig, error) {
|
||||||
|
if db.Redis != nil {
|
||||||
|
var cached []model.SystemConfig
|
||||||
|
if err := db.GetJSON(ctx, SystemConfigVisibleListRedisKey, &cached); err == nil {
|
||||||
|
return cached, nil
|
||||||
|
} else if !errors.Is(err, redis.Nil) {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
database := db.DB(ctx)
|
||||||
|
if database == nil {
|
||||||
|
return nil, errors.New(errDatabaseNotInitialized)
|
||||||
|
}
|
||||||
|
|
||||||
|
var configs []model.SystemConfig
|
||||||
|
if err := database.Where("visibility = ?", model.ConfigVisibilityVisible).Find(&configs).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if db.Redis != nil {
|
||||||
|
_ = db.SetJSON(ctx, SystemConfigVisibleListRedisKey, configs, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
return configs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetIntByKey 通过 key 查询配置并转换为 int 类型。
|
||||||
|
func GetIntByKey(ctx context.Context, key string) (int, error) {
|
||||||
|
sc, err := GetSystemConfigByKey(ctx, key)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
value, err := strconv.Atoi(sc.Value)
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return value, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDecimalByKey 通过 key 查询配置并转换为 decimal.Decimal 类型。
|
||||||
|
func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) {
|
||||||
|
sc, err := GetSystemConfigByKey(ctx, key)
|
||||||
|
if err != nil {
|
||||||
|
return decimal.Zero, err
|
||||||
|
}
|
||||||
|
|
||||||
|
value, err := decimal.NewFromString(sc.Value)
|
||||||
|
if err != nil {
|
||||||
|
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return value.Truncate(precision), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetBoolByKey 通过 key 查询配置并转换为 bool 类型。
|
||||||
|
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
|
||||||
|
sc, err := GetSystemConfigByKey(ctx, key)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
|
||||||
|
value, err := strconv.ParseBool(sc.Value)
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return value, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetMenuDisplayConfig 获取目录显示配置,解析为 map[string]bool。
|
||||||
|
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
|
||||||
|
sc, err := GetSystemConfigByKey(ctx, model.ConfigKeyMenuDisplayConfig)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
config := make(map[string]bool)
|
||||||
|
if sc.Value == "" || sc.Value == "{}" {
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
|
||||||
|
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
|
||||||
|
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
|
||||||
|
query := db.DB(ctx).Order("created_at DESC")
|
||||||
|
if configType != "" {
|
||||||
|
query = query.Where("type = ?", configType)
|
||||||
|
}
|
||||||
|
var configs []model.SystemConfig
|
||||||
|
if err := query.Find(&configs).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return configs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetAdminSystemConfigByKey loads a config directly from PostgreSQL.
|
||||||
|
func GetAdminSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
|
||||||
|
var config model.SystemConfig
|
||||||
|
if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
|
||||||
|
return model.SystemConfig{}, err
|
||||||
|
}
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SystemConfigExists reports whether a config key already exists.
|
||||||
|
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
|
||||||
|
var existing model.SystemConfig
|
||||||
|
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
|
||||||
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateSystemConfig persists a new system config row.
|
||||||
|
func CreateSystemConfig(ctx context.Context, config *model.SystemConfig) error {
|
||||||
|
return db.DB(ctx).Create(config).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateSystemConfigFields applies partial updates to a system config row.
|
||||||
|
func UpdateSystemConfigFields(ctx context.Context, config *model.SystemConfig, updates map[string]any) error {
|
||||||
|
return db.DB(ctx).Model(config).Updates(updates).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
|
||||||
|
func SaveOrUpdateSystemConfig(ctx context.Context, key, 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: model.ConfigVisibilityHidden,
|
||||||
|
}
|
||||||
|
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 InvalidateSystemConfigCache(ctx, key)
|
||||||
|
}
|
||||||
@@ -1,7 +1,7 @@
|
|||||||
// Copyright 2026 Arctel.net
|
// Copyright 2026 Arctel.net
|
||||||
// SPDX-License-Identifier: Apache-2.0
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
package model
|
package repository
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
@@ -9,14 +9,20 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// SystemConfigInvalidationChannel broadcasts RAM cache eviction across nodes.
|
// SystemConfigInvalidationChannel broadcasts RAM cache eviction across nodes.
|
||||||
SystemConfigInvalidationChannel = "system:config_invalidation"
|
SystemConfigInvalidationChannel = "system:config_invalidation"
|
||||||
systemConfigInvalidateAllToken = "*"
|
// SystemConfigRedisHashKey Redis Hash key,存储所有系统配置。
|
||||||
systemConfigRAMMaximumSize = 512
|
SystemConfigRedisHashKey = "system:system_configs"
|
||||||
|
// SystemConfigVisibleListRedisKey Redis key,缓存所有 visibility=1 的公共配置列表。
|
||||||
|
SystemConfigVisibleListRedisKey = "system:visible_configs"
|
||||||
|
|
||||||
|
systemConfigInvalidateAllToken = "*"
|
||||||
|
systemConfigRAMMaximumSize = 512
|
||||||
)
|
)
|
||||||
|
|
||||||
type systemConfigInvalidationMessage struct {
|
type systemConfigInvalidationMessage struct {
|
||||||
@@ -24,7 +30,7 @@ type systemConfigInvalidationMessage struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
systemConfigRAMCache = ram.MustNew[string, SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize})
|
systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize})
|
||||||
systemConfigListenerOnce sync.Once
|
systemConfigListenerOnce sync.Once
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -58,11 +64,11 @@ func startSystemConfigCacheInvalidationListener() {
|
|||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
func cloneSystemConfig(sc SystemConfig) SystemConfig {
|
func cloneSystemConfig(sc model.SystemConfig) model.SystemConfig {
|
||||||
return sc
|
return sc
|
||||||
}
|
}
|
||||||
|
|
||||||
func populateSystemConfigCache(ctx context.Context, sc SystemConfig) {
|
func populateSystemConfigCache(ctx context.Context, sc model.SystemConfig) {
|
||||||
systemConfigRAMCache.Set(sc.Key, cloneSystemConfig(sc))
|
systemConfigRAMCache.Set(sc.Key, cloneSystemConfig(sc))
|
||||||
if db.Redis != nil {
|
if db.Redis != nil {
|
||||||
_ = db.HSetJSON(ctx, SystemConfigRedisHashKey, sc.Key, &sc)
|
_ = db.HSetJSON(ctx, SystemConfigRedisHashKey, sc.Key, &sc)
|
||||||
@@ -81,7 +87,6 @@ func publishSystemConfigRAMInvalidation(ctx context.Context, key string) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// InvalidateSystemConfigCache evicts one config key from local RAM and Redis.
|
// InvalidateSystemConfigCache evicts one config key from local RAM and Redis.
|
||||||
// It also publishes cluster-wide RAM invalidation when Redis is available.
|
|
||||||
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
||||||
ensureSystemConfigCacheListener()
|
ensureSystemConfigCacheListener()
|
||||||
|
|
||||||
@@ -96,7 +101,6 @@ func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// InvalidateAllSystemConfigCaches evicts all config entries from local RAM and Redis.
|
// InvalidateAllSystemConfigCaches evicts all config entries from local RAM and Redis.
|
||||||
// It also publishes cluster-wide RAM invalidation when Redis is available.
|
|
||||||
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
|
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
|
||||||
ensureSystemConfigCacheListener()
|
ensureSystemConfigCacheListener()
|
||||||
|
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ListTemplates returns all templates ordered by system flag and creation time.
|
||||||
|
func ListTemplates(ctx context.Context) ([]model.Template, error) {
|
||||||
|
var templates []model.Template
|
||||||
|
if err := db.DB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return templates, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetTemplateByKey loads a template by its key.
|
||||||
|
func GetTemplateByKey(ctx context.Context, key string) (model.Template, error) {
|
||||||
|
var tmpl model.Template
|
||||||
|
if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
|
||||||
|
return model.Template{}, err
|
||||||
|
}
|
||||||
|
return tmpl, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// TemplateExistsByKey reports whether a template key is already taken.
|
||||||
|
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
|
||||||
|
var existing model.Template
|
||||||
|
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
|
||||||
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateTemplate persists a new template.
|
||||||
|
func CreateTemplate(ctx context.Context, tmpl *model.Template) error {
|
||||||
|
return db.DB(ctx).Create(tmpl).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveTemplate updates an existing template.
|
||||||
|
func SaveTemplate(ctx context.Context, tmpl *model.Template) error {
|
||||||
|
return db.DB(ctx).Save(tmpl).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteTemplate removes a template record.
|
||||||
|
func DeleteTemplate(ctx context.Context, tmpl *model.Template) error {
|
||||||
|
return db.DB(ctx).Delete(tmpl).Error
|
||||||
|
}
|
||||||
@@ -0,0 +1,118 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// UploadListFilter filters paginated upload queries.
|
||||||
|
type UploadListFilter struct {
|
||||||
|
UserID uint64
|
||||||
|
Keyword string
|
||||||
|
Type string
|
||||||
|
Extension string
|
||||||
|
Page int
|
||||||
|
PageSize int
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListUploads returns paginated upload records matching the filter.
|
||||||
|
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []model.Upload, error) {
|
||||||
|
query := db.DB(ctx).Model(&model.Upload{}).
|
||||||
|
Where("status != ?", model.UploadStatusDeleted)
|
||||||
|
|
||||||
|
if filter.UserID != 0 {
|
||||||
|
query = query.Where("user_id = ?", filter.UserID)
|
||||||
|
}
|
||||||
|
if filter.Keyword != "" {
|
||||||
|
query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(filter.Keyword)+"%")
|
||||||
|
}
|
||||||
|
if filter.Type != "" {
|
||||||
|
query = query.Where("type = ?", filter.Type)
|
||||||
|
}
|
||||||
|
if filter.Extension != "" {
|
||||||
|
query = query.Where("extension = ?", strings.ToLower(filter.Extension))
|
||||||
|
}
|
||||||
|
|
||||||
|
var total int64
|
||||||
|
if err := query.Count(&total).Error; err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var items []model.Upload
|
||||||
|
offset := (filter.Page - 1) * filter.PageSize
|
||||||
|
if err := query.Order("created_at DESC").Offset(offset).Limit(filter.PageSize).Find(&items).Error; err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
return total, items, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||||
|
func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
|
||||||
|
var upload model.Upload
|
||||||
|
if err := db.DB(ctx).Where("id = ? AND status != ?", id, model.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||||
|
return model.Upload{}, err
|
||||||
|
}
|
||||||
|
return upload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SoftDeleteUpload marks an upload as deleted.
|
||||||
|
func SoftDeleteUpload(ctx context.Context, upload *model.Upload) error {
|
||||||
|
return db.DB(ctx).Model(upload).Update("status", model.UploadStatusDeleted).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateUpload applies partial field updates to an upload record.
|
||||||
|
func UpdateUpload(ctx context.Context, upload *model.Upload, updates map[string]any) error {
|
||||||
|
if len(updates) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return db.DB(ctx).Model(upload).Updates(updates).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||||
|
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||||
|
var types []string
|
||||||
|
if err := db.DB(ctx).Model(&model.Upload{}).
|
||||||
|
Where("type IS NOT NULL AND type != ''").
|
||||||
|
Distinct().
|
||||||
|
Pluck("type", &types).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return types, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||||
|
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (model.Upload, error) {
|
||||||
|
var existing model.Upload
|
||||||
|
err := db.DB(ctx).
|
||||||
|
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, model.UploadStatusPending, model.UploadStatusUsed).
|
||||||
|
First(&existing).Error
|
||||||
|
return existing, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateUpload persists a new upload record.
|
||||||
|
func CreateUpload(ctx context.Context, upload *model.Upload) error {
|
||||||
|
return db.DB(ctx).Create(upload).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||||
|
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]model.Upload, error) {
|
||||||
|
var uploads []model.Upload
|
||||||
|
if err := db.DB(ctx).
|
||||||
|
Where("id IN ? AND status IN (?, ?)", ids, model.UploadStatusPending, model.UploadStatusUsed).
|
||||||
|
Find(&uploads).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return uploads, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UploadQuery returns a scoped GORM query for uploads.
|
||||||
|
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||||
|
return db.DB(ctx).Model(&model.Upload{})
|
||||||
|
}
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ListUploadStats returns all upload statistics rows.
|
||||||
|
func ListUploadStats(ctx context.Context) ([]model.UploadStat, error) {
|
||||||
|
var stats []model.UploadStat
|
||||||
|
if err := db.DB(ctx).Find(&stats).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return stats, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,172 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package repository
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// GetUserByID loads an active user by ID.
|
||||||
|
func GetUserByID(ctx context.Context, id uint64) (model.User, error) {
|
||||||
|
var user model.User
|
||||||
|
if err := db.DB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
|
||||||
|
return model.User{}, err
|
||||||
|
}
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserByUsername loads a user by username.
|
||||||
|
func GetUserByUsername(ctx context.Context, username string) (model.User, error) {
|
||||||
|
var user model.User
|
||||||
|
if err := db.DB(ctx).Where("username = ?", username).First(&user).Error; err != nil {
|
||||||
|
return model.User{}, err
|
||||||
|
}
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSystemUser loads the built-in system user, or returns a synthetic fallback.
|
||||||
|
func GetSystemUser(ctx context.Context) model.User {
|
||||||
|
var user model.User
|
||||||
|
if err := db.DB(ctx).Where("username = ?", "system").First(&user).Error; err == nil {
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
return model.User{
|
||||||
|
ID: 999,
|
||||||
|
Username: "system",
|
||||||
|
Nickname: "系统",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetFirstAdminUser loads the earliest admin user.
|
||||||
|
func GetFirstAdminUser(ctx context.Context) (model.User, error) {
|
||||||
|
var user model.User
|
||||||
|
if err := db.DB(ctx).Where("is_admin = ?", true).Order("id asc").First(&user).Error; err != nil {
|
||||||
|
return model.User{}, err
|
||||||
|
}
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AdminUserListFilter filters admin user list queries.
|
||||||
|
type AdminUserListFilter struct {
|
||||||
|
UserID *uint64
|
||||||
|
Username string
|
||||||
|
Page int
|
||||||
|
PageSize int
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListAdminUsers returns paginated users for the admin console.
|
||||||
|
func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []model.User, error) {
|
||||||
|
query := db.DB(ctx).Model(&model.User{})
|
||||||
|
if filter.UserID != nil {
|
||||||
|
query = query.Where("id = ?", *filter.UserID)
|
||||||
|
}
|
||||||
|
if filter.Username != "" {
|
||||||
|
query = query.Where("username LIKE ?", filter.Username+"%")
|
||||||
|
}
|
||||||
|
|
||||||
|
var total int64
|
||||||
|
if err := query.Count(&total).Error; err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var users []model.User
|
||||||
|
offset := (filter.Page - 1) * filter.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(filter.PageSize).
|
||||||
|
Find(&users).Error; err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
return total, users, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetAdminUserDetail loads full user profile fields for admin detail view.
|
||||||
|
func GetAdminUserDetail(ctx context.Context, id uint64) (model.User, error) {
|
||||||
|
var user model.User
|
||||||
|
if err := db.DB(ctx).
|
||||||
|
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(&user).Error; err != nil {
|
||||||
|
return model.User{}, err
|
||||||
|
}
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UserAdminFlags stores minimal user authorization flags.
|
||||||
|
type UserAdminFlags struct {
|
||||||
|
ID uint64
|
||||||
|
IsAdmin bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserAdminFlags loads id and is_admin for authorization checks.
|
||||||
|
func GetUserAdminFlags(ctx context.Context, id uint64) (UserAdminFlags, error) {
|
||||||
|
var flags UserAdminFlags
|
||||||
|
if err := db.DB(ctx).
|
||||||
|
Model(&model.User{}).
|
||||||
|
Select("id, is_admin").
|
||||||
|
Where("id = ?", id).
|
||||||
|
First(&flags).Error; err != nil {
|
||||||
|
return UserAdminFlags{}, err
|
||||||
|
}
|
||||||
|
return flags, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateUserActive updates the is_active flag for a user.
|
||||||
|
func UpdateUserActive(ctx context.Context, id uint64, active bool) error {
|
||||||
|
return db.DB(ctx).Model(&model.User{}).Where("id = ?", id).Update("is_active", active).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteUserWithRelations removes a user and related access tokens / external accounts.
|
||||||
|
func DeleteUserWithRelations(ctx context.Context, id uint64) error {
|
||||||
|
return db.DB(ctx).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
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// CountUsersByUsername returns how many users share the username.
|
||||||
|
func CountUsersByUsername(ctx context.Context, username string) (int64, error) {
|
||||||
|
var count int64
|
||||||
|
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", username).Count(&count).Error; err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return count, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CountUsersByEmail returns how many users share the email.
|
||||||
|
func CountUsersByEmail(ctx context.Context, email string) (int64, error) {
|
||||||
|
var count int64
|
||||||
|
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return count, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateUser persists a new user record.
|
||||||
|
func CreateUser(ctx context.Context, user *model.User) error {
|
||||||
|
return db.DB(ctx).Create(user).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListUsersByIDs loads users matching the given IDs.
|
||||||
|
func ListUsersByIDs(ctx context.Context, ids []uint64) ([]model.User, error) {
|
||||||
|
if len(ids) == 0 {
|
||||||
|
return []model.User{}, nil
|
||||||
|
}
|
||||||
|
var users []model.User
|
||||||
|
if err := db.DB(ctx).Where("id IN ?", ids).Find(&users).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return users, nil
|
||||||
|
}
|
||||||
@@ -15,6 +15,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -72,8 +73,8 @@ func loggerMiddleware() gin.HandlerFunc {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func isOriginAllowed(ctx context.Context, origin string) bool {
|
func isOriginAllowed(ctx context.Context, origin string) bool {
|
||||||
var sc model.SystemConfig
|
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
|
||||||
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || sc.Value == "" {
|
if err != nil || sc.Value == "" {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
allowedOrigins := strings.Split(sc.Value, ",")
|
allowedOrigins := strings.Split(sc.Value, ",")
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
@@ -21,7 +22,7 @@ func TestCORSMiddleware(t *testing.T) {
|
|||||||
gin.SetMode(gin.TestMode)
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
clearConfigCache := func() {
|
clearConfigCache := func() {
|
||||||
if err := model.InvalidateAllSystemConfigCaches(context.Background()); err != nil {
|
if err := repository.InvalidateAllSystemConfigCaches(context.Background()); err != nil {
|
||||||
t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err)
|
t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -109,8 +110,8 @@ func LoadConfig(ctx context.Context) (Config, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func loadConfigByKey(ctx context.Context, key string, fallback Config) (Config, error) {
|
func loadConfigByKey(ctx context.Context, key string, fallback Config) (Config, error) {
|
||||||
var sc model.SystemConfig
|
sc, err := repository.GetSystemConfigByKey(ctx, key)
|
||||||
if err := sc.GetByKey(ctx, key); err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return fallback, nil
|
return fallback, nil
|
||||||
}
|
}
|
||||||
@@ -208,7 +209,7 @@ func upsertSystemConfig(ctx context.Context, tx *gorm.DB, key string, value any,
|
|||||||
FirstOrCreate(&sc).Error; err != nil {
|
FirstOrCreate(&sc).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return model.InvalidateSystemConfigCache(ctx, key)
|
return repository.InvalidateSystemConfigCache(ctx, key)
|
||||||
}
|
}
|
||||||
|
|
||||||
// MergeMaskedSecrets restores unchanged secrets from the current configuration.
|
// MergeMaskedSecrets restores unchanged secrets from the current configuration.
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -123,8 +124,7 @@ func Active(ctx context.Context) (Driver, Backend, error) {
|
|||||||
return activeDriver, activeBackend, nil
|
return activeDriver, activeBackend, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var sc model.SystemConfig
|
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyStorageConfig)
|
||||||
err := sc.GetByKey(ctx, model.ConfigKeyStorageConfig)
|
|
||||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return "", nil, err
|
return "", nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/alicebob/miniredis/v2"
|
"github.com/alicebob/miniredis/v2"
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"github.com/redis/go-redis/v9"
|
"github.com/redis/go-redis/v9"
|
||||||
@@ -76,7 +77,7 @@ func SetupTestEnvironment(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func())
|
|||||||
// Cleanup function
|
// Cleanup function
|
||||||
cleanup := func() {
|
cleanup := func() {
|
||||||
runExtraCleanups()
|
runExtraCleanups()
|
||||||
model.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
_ = redisClient.Close()
|
_ = redisClient.Close()
|
||||||
mr.Close()
|
mr.Close()
|
||||||
// Reset database and Redis references
|
// Reset database and Redis references
|
||||||
@@ -316,6 +317,6 @@ func seedDefaultConfigs(t *testing.T, tx *gorm.DB) {
|
|||||||
if _, ok := publicKeys[config.Key]; ok {
|
if _, ok := publicKeys[config.Key]; ok {
|
||||||
config.Visibility = model.ConfigVisibilityVisible
|
config.Visibility = model.ConfigVisibilityVisible
|
||||||
}
|
}
|
||||||
_ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, config.Key, &config)
|
_ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, config.Key, &config)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
// Copyright 2026 Arctel.net
|
// Copyright 2026 Arctel.net
|
||||||
// SPDX-License-Identifier: Apache-2.0
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
// Package util provides framework-agnostic helper types and HTTP utilities.
|
||||||
package util
|
package util
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
|||||||
Reference in New Issue
Block a user