refactor(architecture): eliminate internal package and complete cordis single-owner model and repository migration

- Physically purged all legacy internal/ packages, centralized pkg/model/ and pkg/repository/
- Migrated domain models and database repositories into self-contained owner plugins (user, auth, message_gateway, admin, upload, risk_control)
- Decoupled cross-plugin interactions via pure core/contracts and typed EventBus
- Ensured 100% test coverage pass, zero data races (-race clean), and 0 lint issues in make code-check
This commit is contained in:
ryan
2026-08-28 08:40:43 +08:00
parent 1f348fd425
commit fb6a3edb89
323 changed files with 8222 additions and 17693 deletions
+2 -2
View File
@@ -7,8 +7,8 @@ import (
"net/http"
"strconv"
persistence "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/shared/response"
persistence "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-gonic/gin"
)
+6 -8
View File
@@ -10,10 +10,8 @@ import (
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/infra/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/diskcache"
)
type updateCacheConfigRequest struct {
@@ -61,17 +59,17 @@ func UpdateCacheConfig(c *gin.Context) {
ctx := c.Request.Context()
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -101,5 +99,5 @@ func ClearCache(c *gin.Context) {
}
func saveOrUpdateCacheConfig(ctx context.Context, key, value string) error {
return repository.SaveOrUpdateSystemConfig(ctx, key, value)
return SaveOrUpdateSystemConfig(ctx, key, value)
}
+36 -37
View File
@@ -13,15 +13,12 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
mail "github.com/Rain-kl/Wavelet/pkg/mail"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/cap"
"github.com/Rain-kl/Wavelet/plugins/domain/upload"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -54,7 +51,7 @@ type UpdateSystemConfigRequest struct {
// @Router /api/v1/config/public [get]
func GetPublicConfig(c *gin.Context) {
ctx := c.Request.Context()
configs, err := repository.ListVisibleSystemConfigs(ctx)
configs, err := ListVisibleSystemConfigs(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -77,7 +74,7 @@ func GetPublicConfig(c *gin.Context) {
// @Router /robots.txt [get]
func GetRobotsTXT(c *gin.Context) {
ctx := c.Request.Context()
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled)
enabled, err := GetBoolByKey(ctx, ConfigKeySearchEngineIndexingEnabled)
content := "User-Agent: *\nDisallow: /\n"
if err == nil && enabled {
content = "User-Agent: *\nAllow: /\n"
@@ -129,7 +126,7 @@ func CreateSystemConfig(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param type query string false "配置类型(system/business)"
// @Success 200 {object} response.Any{data=[]model.SystemConfig} "系统配置列表"
// @Success 200 {object} response.Any{data=[]SystemConfig} "系统配置列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
@@ -155,7 +152,7 @@ func ListSystemConfigs(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Success 200 {object} response.Any{data=model.SystemConfig} "系统配置详情"
// @Success 200 {object} response.Any{data=SystemConfig} "系统配置详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
@@ -222,14 +219,14 @@ func UpdateSystemConfig(c *gin.Context) {
}
func isProtectedConfigKey(key string) bool {
return key == model.ConfigKeyLogDatabase || key == model.ConfigKeyLogDBMigration
return key == ConfigKeyLogDatabase || key == ConfigKeyLogDBMigration
}
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
if isProtectedConfigKey(req.Key) {
return errors.New(protectedConfigKeyMessage)
}
exists, err := repository.SystemConfigExists(ctx, req.Key)
exists, err := SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
@@ -237,43 +234,43 @@ func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) erro
return errors.New(ConfigKeyExists)
}
config := model.SystemConfig{
config := SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := repository.CreateSystemConfig(ctx, &config); err != nil {
if err := CreateSystemConfigRecord(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
if err := 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 listSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
return ListAdminSystemConfigs(ctx, configType)
}
func getSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) {
return repository.GetAdminSystemConfigByKey(ctx, key)
func getSystemConfig(ctx context.Context, key string) (SystemConfig, error) {
return GetAdminSystemConfigByKey(ctx, key)
}
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
if isProtectedConfigKey(key) {
return errors.New(protectedConfigKeyMessage)
}
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
config, err := GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return err
}
var originalDriver objectstore.Driver
if key == model.ConfigKeyStorageConfig {
if key == ConfigKeyStorageConfig {
var currentCfg objectstore.Config
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
@@ -294,7 +291,7 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
updates["visibility"] = *req.Visibility
config.Visibility = *req.Visibility
}
if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue {
if key != ConfigKeySMTPPassword || req.Value != maskedConfigValue {
updates["value"] = req.Value
config.Value = req.Value
}
@@ -318,7 +315,7 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
originalDriver objectstore.Driver,
newValue string,
) {
if key != model.ConfigKeyStorageConfig || originalDriver == "" {
if key != ConfigKeyStorageConfig || originalDriver == "" {
return
}
@@ -330,7 +327,7 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
return
}
if err := repository.MarkFailedTaskExecutionsSucceededTx(
if err := MarkFailedTaskExecutionsSucceededTx(
tx,
"storage:migrate",
"存储配置直接更新,故障迁移任务自动标记为已解决",
@@ -341,7 +338,7 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil {
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
if cap.IsRuntimeConfigKey(key) {
@@ -352,18 +349,20 @@ func invalidateSystemConfigCaches(ctx context.Context, key string) {
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if key == model.ConfigKeyStorageConfig {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
if key == ConfigKeyStorageConfig {
if db.Redis != nil {
_ = db.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
}
objectstore.ResetCache()
objectstore.PublishCacheInvalidation(ctx)
}
if key == model.ConfigKeyFileAccessWhitelist {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
if key == ConfigKeyFileAccessWhitelist {
if db.Redis != nil {
_ = db.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
}
}
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
}
@@ -404,7 +403,7 @@ func TestSMTP(c *gin.Context) {
password := req.SMTPPassword
if password == maskedConfigValue {
if sc, err := repository.GetSystemConfigByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil {
if sc, err := GetSystemConfigByKey(c.Request.Context(), ConfigKeySMTPPassword); err == nil {
password = sc.Value
}
}
@@ -449,9 +448,9 @@ func maskSensitiveConfig(key, value string) string {
return value
}
switch key {
case model.ConfigKeySMTPPassword:
case ConfigKeySMTPPassword:
return maskedConfigValue
case model.ConfigKeyStorageConfig:
case ConfigKeyStorageConfig:
var cfg objectstore.Config
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
masked := objectstore.MaskSecrets(cfg)
@@ -494,8 +493,8 @@ func validateAndMergeStorageConfig(ctx context.Context, value string, currentCon
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64
if err := db.DB(ctx).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
if err := db.DB(ctx).Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
}
+3 -3
View File
@@ -19,9 +19,9 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/infra/config"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/config"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
)
const (
+25 -15
View File
@@ -15,13 +15,12 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
"github.com/gin-gonic/gin"
@@ -164,8 +163,10 @@ func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessL
username := c.Query("username")
if username != "" {
userIDs, err := repository.ListUserIDsByUsernameContains(ctx, username)
if err != nil {
var userIDs []uint64
if err := db.DB(ctx).Table("w_users").
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
Pluck("id", &userIDs).Error; err != nil {
return filter, fmt.Errorf("查询用户信息失败: %w", err)
}
filter.UserIDs = userIDs
@@ -213,7 +214,12 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
}
userMap := make(map[uint64]struct{ Username, Nickname string })
if users, err := repository.ListUsersByIDs(ctx, userIDs); err == nil {
var users []struct {
ID uint64
Username string
Nickname string
}
if err := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
@@ -402,8 +408,12 @@ func GetLogsAnalytics(c *gin.Context) {
Username string
Nickname string
})
users, errProfile := repository.ListUsersByIDs(ctx, userIDs)
if errProfile == nil {
var users []struct {
ID uint64
Username string
Nickname string
}
if errProfile := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
for _, u := range users {
userProfileMap[u.ID] = struct {
Username string
@@ -445,7 +455,7 @@ func getUpgrader() *websocket.Upgrader {
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
ctx := r.Context()
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
if sc, err := GetSystemConfigByKey(ctx, ConfigKeyServerAddress); err == nil && sc.Value != "" {
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
allowedOrigins := strings.Split(sc.Value, ",")
for _, allowed := range allowedOrigins {
@@ -632,7 +642,7 @@ func validateSwitch(ctx context.Context, target string) error {
}
func currentLogDatabase(ctx context.Context) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase)
cfg, err := GetSystemConfigByKey(ctx, ConfigKeyLogDatabase)
if err != nil {
return "", fmt.Errorf("读取日志主库失败: %w", err)
}
@@ -643,11 +653,11 @@ func currentLogDatabase(ctx context.Context) (string, error) {
}
func setMigrationFlag(ctx context.Context, v string) error {
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDBMigration, v)
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDBMigration, v)
}
func flipLogDatabase(ctx context.Context, target string) error {
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, target)
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
}
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
+7 -9
View File
@@ -16,12 +16,10 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
"github.com/Rain-kl/Wavelet/pkg/response"
)
var startTime = time.Now()
@@ -200,16 +198,16 @@ func GetLogDatabaseStatus(c *gin.Context) {
ActiveDatabase: activeDB,
Migration: migration,
RetentionDays: map[string]int{
logDBNamePostgres: retentionOr(ctx, model.ConfigKeyLogRetentionDaysPostgres),
logDBNameSQLite: retentionOr(ctx, model.ConfigKeyLogRetentionDaysSQLite),
logDBNameClickHouse: retentionOr(ctx, model.ConfigKeyLogRetentionDaysClickHouse),
logDBNamePostgres: retentionOr(ctx, ConfigKeyLogRetentionDaysPostgres),
logDBNameSQLite: retentionOr(ctx, ConfigKeyLogRetentionDaysSQLite),
logDBNameClickHouse: retentionOr(ctx, ConfigKeyLogRetentionDaysClickHouse),
},
AvailableTargets: availableLogTargets(activeDB),
}))
}
func retentionOr(ctx context.Context, key string) int {
v, err := repository.GetIntByKey(ctx, key)
v, err := GetIntByKey(ctx, key)
if err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) {
logger.ErrorF(ctx, "读取日志保留天数配置失败 key=%s: %v", key, err)
+16 -18
View File
@@ -14,12 +14,10 @@ import (
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/infra/task/scheduler"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/pkg/task/scheduler"
)
// ListTaskTypes 获取支持的任务类型列表
@@ -107,7 +105,7 @@ func DispatchTask(c *gin.Context) {
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/executions [get]
func ListTaskExecutions(c *gin.Context) {
var req model.ListTaskExecutionsRequest
var req ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -119,7 +117,7 @@ func ListTaskExecutions(c *gin.Context) {
}
}
executions, total, err := repository.ListTaskExecutions(c.Request.Context(), req)
executions, total, err := ListTaskExecutionRecords(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -140,7 +138,7 @@ func ListTaskExecutions(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=model.TaskExecution} "任务执行详情"
// @Success 200 {object} response.Any{data=TaskExecution} "任务执行详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -153,7 +151,7 @@ func GetTaskExecution(c *gin.Context) {
return
}
execution, err := repository.GetTaskExecutionByID(c.Request.Context(), id)
execution, err := GetTaskExecutionByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, TaskNotFound)
return
@@ -206,12 +204,12 @@ func RetryTask(c *gin.Context) {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Schedule} "定时任务列表"
// @Success 200 {object} response.Any{data=[]Schedule} "定时任务列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/schedules [get]
func ListSchedules(c *gin.Context) {
schedules, err := repository.ListSchedules(c.Request.Context())
schedules, err := ListSchedulesRecord(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -236,7 +234,7 @@ type CreateScheduleRequest struct {
// @Produce json
// @Security SessionCookie
// @Param request body CreateScheduleRequest true "创建定时任务请求参数"
// @Success 200 {object} response.Any{data=model.Schedule} "创建成功的定时任务信息"
// @Success 200 {object} response.Any{data=Schedule} "创建成功的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、异步任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -273,7 +271,7 @@ func CreateSchedule(c *gin.Context) {
return
}
schedule := &model.Schedule{
schedule := &Schedule{
Name: req.Name,
TaskType: req.TaskType,
Cron: req.Cron,
@@ -281,7 +279,7 @@ func CreateSchedule(c *gin.Context) {
IsActive: *req.IsActive,
}
if err := repository.CreateSchedule(c.Request.Context(), schedule); err != nil {
if err := CreateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -312,7 +310,7 @@ type UpdateScheduleRequest struct {
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Param request body UpdateScheduleRequest true "修改定时任务请求参数"
// @Success 200 {object} response.Any{data=model.Schedule} "修改后的定时任务信息"
// @Success 200 {object} response.Any{data=Schedule} "修改后的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -332,7 +330,7 @@ func UpdateSchedule(c *gin.Context) {
return
}
schedule, err := repository.GetScheduleByID(c.Request.Context(), id)
schedule, err := GetScheduleByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, ScheduleNotFound)
return
@@ -368,7 +366,7 @@ func UpdateSchedule(c *gin.Context) {
schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive
if err := repository.UpdateSchedule(c.Request.Context(), schedule); err != nil {
if err := UpdateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -401,7 +399,7 @@ func DeleteSchedule(c *gin.Context) {
return
}
if err := repository.DeleteSchedule(c.Request.Context(), id); err != nil {
if err := DeleteScheduleRecord(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return
}
+24 -26
View File
@@ -11,9 +11,7 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/response"
)
// CreateTemplateRequest 创建模板请求
@@ -88,7 +86,7 @@ func CreateTemplate(c *gin.Context) {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Template} "模板列表"
// @Success 200 {object} response.Any{data=[]Template} "模板列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
@@ -110,7 +108,7 @@ func ListTemplates(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=model.Template} "模板详情"
// @Success 200 {object} response.Any{data=Template} "模板详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
@@ -134,7 +132,7 @@ func GetTemplate(c *gin.Context) {
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Param request body UpdateTemplateRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=model.Template} "更新成功"
// @Success 200 {object} response.Any{data=Template} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -178,16 +176,16 @@ func DeleteTemplate(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Template, error) {
exists, err := repository.TemplateExistsByKey(ctx, req.Key)
func createTemplate(ctx context.Context, req CreateTemplateRequest) (Template, error) {
exists, err := TemplateExistsByKey(ctx, req.Key)
if err != nil {
return model.Template{}, err
return Template{}, err
}
if exists {
return model.Template{}, errors.New(TemplateKeyExists)
return Template{}, errors.New(TemplateKeyExists)
}
tmpl := model.Template{
tmpl := Template{
Key: req.Key,
Name: req.Name,
Type: req.Type,
@@ -197,26 +195,26 @@ func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Templ
IsSystem: false,
}
if err := tmpl.Validate(); err != nil {
return model.Template{}, err
return Template{}, err
}
if err := repository.CreateTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
if err := CreateTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func listTemplates(ctx context.Context) ([]model.Template, error) {
return repository.ListTemplates(ctx)
func listTemplates(ctx context.Context) ([]Template, error) {
return ListTemplatesRecord(ctx)
}
func getTemplate(ctx context.Context, key string) (model.Template, error) {
return repository.GetTemplateByKey(ctx, key)
func getTemplate(ctx context.Context, key string) (Template, error) {
return GetTemplateByKey(ctx, key)
}
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (model.Template, error) {
tmpl, err := repository.GetTemplateByKey(ctx, key)
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (Template, error) {
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return model.Template{}, err
return Template{}, err
}
tmpl.Name = req.Name
@@ -225,21 +223,21 @@ func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest)
tmpl.Content = req.Content
tmpl.Description = req.Description
if err := tmpl.Validate(); err != nil {
return model.Template{}, err
return Template{}, err
}
if err := repository.SaveTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
if err := SaveTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func deleteTemplate(ctx context.Context, key string) error {
tmpl, err := repository.GetTemplateByKey(ctx, key)
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return err
}
if tmpl.IsSystem {
return errors.New(SystemTemplateCannotDelete)
}
return repository.DeleteTemplate(ctx, &tmpl)
return DeleteTemplateRecord(ctx, &tmpl)
}
+3 -5
View File
@@ -24,11 +24,9 @@ import (
"github.com/gin-gonic/gin"
"golang.org/x/mod/semver"
"github.com/Rain-kl/Wavelet/internal/buildinfo"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/buildinfo"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/util"
)
@@ -282,7 +280,7 @@ func (m *updaterManager) fetchRelease(ctx context.Context, repository string) (g
}
func loadRepository(ctx context.Context) (string, error) {
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository)
config, err := GetSystemConfigByKey(ctx, ConfigKeyUpdateUpstreamRepository)
if err != nil {
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
}
+144 -79
View File
@@ -15,11 +15,12 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
)
@@ -67,7 +68,10 @@ func parseUserID(c *gin.Context) (uint64, bool) {
return id, true
}
func toUserResponse(u model.User) userResponse {
func toUserResponse(u *contracts.UserDTO) userResponse {
if u == nil {
return userResponse{}
}
return userResponse{
ID: u.ID,
Username: u.Username,
@@ -133,16 +137,16 @@ func ListUsers(c *gin.Context) {
return
}
total, modelUsers, err := listUsers(c.Request.Context(), req)
total, dtos, err := listUsers(c.Request.Context(), req)
if err != nil {
logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err)
response.AbortInternal(c, "获取用户列表失败")
return
}
users := make([]userResponse, 0, len(modelUsers))
for _, modelUser := range modelUsers {
users = append(users, toUserResponse(modelUser))
users := make([]userResponse, 0, len(dtos))
for _, dto := range dtos {
users = append(users, toUserResponse(dto))
}
c.JSON(http.StatusOK, response.OK(listUsersResponse{
@@ -243,7 +247,7 @@ func DeleteUser(c *gin.Context) {
return
}
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
@@ -334,7 +338,7 @@ func UpdateUser(c *gin.Context) {
return
}
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
@@ -358,40 +362,70 @@ func UpdateUser(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
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),
Email: strings.TrimSpace(req.Email),
Page: req.Page,
PageSize: req.PageSize,
})
func listUsers(ctx context.Context, req listUsersRequest) (int64, []*contracts.UserDTO, error) {
query := db.DB(ctx).Table("w_users")
if req.UserID != nil {
query = query.Where("id = ?", *req.UserID)
}
if req.Username != "" {
query = query.Where("username LIKE ? ESCAPE '\\'", util.EscapeLike(req.Username)+"%")
}
if req.Email != "" {
query = query.Where("email LIKE ? ESCAPE '\\'", util.EscapeLike(req.Email)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var users []*contracts.UserDTO
offset := (req.Page - 1) * req.PageSize
if err := query.
Select("id, username, nickname, email, avatar_url, is_active, is_admin, last_login_at, created_at, updated_at").
Order("id ASC").
Offset(offset).
Limit(req.PageSize).
Find(&users).Error; err != nil {
return 0, nil, err
}
return total, users, nil
}
func getUserDetail(ctx context.Context, id uint64) (model.User, error) {
return repository.GetAdminUserDetail(ctx, id)
func getUserDetail(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").
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 nil, err
}
return &user, nil
}
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
flags, err := repository.GetUserAdminFlags(ctx, id)
if err != nil {
var flags struct {
ID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", id).First(&flags).Error; err != nil {
return err
}
if !active && flags.IsAdmin {
return errors.New(cannotDisable)
}
var tokens []model.AccessToken
var tokenHashes []string
if !active {
tokens, _ = repository.ListAccessTokensByUserID(ctx, id)
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", id).Pluck("token_hash", &tokenHashes).Error
}
err = repository.UpdateUserActive(ctx, id, active)
err := db.DB(ctx).Table("w_users").Where("id = ?", id).Update("is_active", active).Error
if err == nil {
auth.InvalidateCachedUser(ctx, id)
if !active {
for _, token := range tokens {
auth.InvalidateCachedToken(ctx, token.TokenHash)
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
}
@@ -402,77 +436,106 @@ 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 {
var flags struct {
ID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", targetID).First(&flags).Error; err != nil {
return err
}
if flags.IsAdmin {
return errors.New(cannotDelete)
}
tokens, _ := repository.ListAccessTokensByUserID(ctx, targetID)
var tokenHashes []string
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", targetID).Pluck("token_hash", &tokenHashes).Error
err = repository.DeleteUserWithRelations(ctx, targetID)
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Table("w_access_tokens").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
return err
}
if err := tx.Table("w_external_accounts").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
return err
}
return tx.Table("w_users").Where("id = ?", targetID).Delete(map[string]any{}).Error
})
if err == nil {
auth.InvalidateCachedUser(ctx, targetID)
for _, token := range tokens {
auth.InvalidateCachedToken(ctx, token.TokenHash)
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
return err
}
func createUser(ctx context.Context, req createUserRequest) (model.User, error) {
func createUser(ctx context.Context, req createUserRequest) (*contracts.UserDTO, 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)
return nil, errors.New(usernameRequired)
}
if req.Email == "" {
return model.User{}, errors.New(emailRequired)
return nil, errors.New(emailRequired)
}
if len(req.Password) < minPasswordLength {
return model.User{}, errors.New(passwordTooShort)
return nil, errors.New(passwordTooShort)
}
count, err := repository.CountUsersByUsername(ctx, req.Username)
if err != nil {
return model.User{}, err
var count int64
if err := db.DB(ctx).Table("w_users").Where("username = ?", req.Username).Count(&count).Error; err != nil {
return nil, err
}
if count > 0 {
return model.User{}, errors.New(usernameExists)
return nil, errors.New(usernameExists)
}
emailCount, err := repository.CountUsersByEmail(ctx, req.Email)
if err != nil {
return model.User{}, err
var emailCount int64
if err := db.DB(ctx).Table("w_users").Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
return nil, err
}
if emailCount > 0 {
return model.User{}, errors.New(emailExists)
return nil, 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{},
hash, err := util.HashPassword(req.Password)
if err != nil {
return nil, err
}
if newUser.Nickname == "" {
newUser.Nickname = req.Username
if req.Nickname == "" {
req.Nickname = req.Username
}
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
return model.User{}, err
now := time.Now()
newUser := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
CreatedAt: now,
UpdatedAt: now,
}
if err := repository.CreateUser(ctx, &newUser); err != nil {
return model.User{}, err
row := map[string]any{
"id": newUser.ID,
"username": newUser.Username,
"password": hash,
"nickname": newUser.Nickname,
"email": newUser.Email,
"is_active": newUser.IsActive,
"is_admin": newUser.IsAdmin,
"created_at": now,
"updated_at": now,
}
return newUser, nil
if err := db.DB(ctx).Table("w_users").Create(row).Error; err != nil {
return nil, err
}
return &newUser, nil
}
type updateUserParam struct {
@@ -492,20 +555,18 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam
return errors.New(emailRequired)
}
targetUser, err := repository.GetAdminUserDetail(ctx, param.ID)
if err != nil {
var targetUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).First(&targetUser).Error; err != nil {
return err
}
// 不能撤销当前登录用户的管理员权限
if currentUserID == param.ID && !param.IsAdmin && targetUser.IsAdmin {
return errors.New(cannotRevokeSelfAdmin)
}
// 如果修改了邮箱,检查邮箱是否被其他用户占用
if targetUser.Email != param.Email {
count, err := repository.CountUsersByEmail(ctx, param.Email)
if err != nil {
var count int64
if err := db.DB(ctx).Table("w_users").Where("email = ? AND id != ?", param.Email, param.ID).Count(&count).Error; err != nil {
return err
}
if count > 0 {
@@ -513,36 +574,40 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam
}
}
// 密码强度校验(如果输入了新密码)
if param.Password != "" && len(param.Password) < minPasswordLength {
return errors.New(passwordTooShort)
}
needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin)
var tokens []model.AccessToken
var tokenHashes []string
if needRevokeTokens {
tokens, _ = repository.ListAccessTokensByUserID(ctx, param.ID)
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", param.ID).Pluck("token_hash", &tokenHashes).Error
}
targetUser.Nickname = param.Nickname
if targetUser.Nickname == "" {
targetUser.Nickname = targetUser.Username
if param.Nickname == "" {
param.Nickname = targetUser.Username
}
targetUser.Email = param.Email
targetUser.IsAdmin = param.IsAdmin
updates := map[string]any{
"nickname": param.Nickname,
"email": param.Email,
"is_admin": param.IsAdmin,
"updated_at": time.Now(),
}
if param.Password != "" {
if err := targetUser.SetEncryptedPassword(param.Password); err != nil {
hash, err := util.HashPassword(param.Password)
if err != nil {
return err
}
updates["password"] = hash
}
err = repository.UpdateUser(ctx, &targetUser)
err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).Updates(updates).Error
if err == nil {
auth.InvalidateCachedUser(ctx, param.ID)
if needRevokeTokens {
for _, token := range tokens {
auth.InvalidateCachedToken(ctx, token.TokenHash)
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
}
+3 -3
View File
@@ -4,9 +4,9 @@
package admin
import (
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-gonic/gin"
@@ -18,7 +18,7 @@ func LoginAdminRequired() gin.HandlerFunc {
ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
user, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
if user == nil {
response.AbortNotFound(c, AdminRequired)
return
+20 -1
View File
@@ -92,6 +92,9 @@ func (Template) TableName() string {
return "w_templates"
}
// TemplateTypeEmail 邮件模板类型
const TemplateTypeEmail = "email"
// Normalize 规范化模板字段
func (t *Template) Normalize() {
t.Key = strings.TrimSpace(t.Key)
@@ -101,7 +104,7 @@ func (t *Template) Normalize() {
t.Content = strings.TrimSpace(t.Content)
t.Description = strings.TrimSpace(t.Description)
if t.Type == "" {
t.Type = "email"
t.Type = TemplateTypeEmail
}
}
@@ -201,3 +204,19 @@ type TaskExecution struct {
func (TaskExecution) TableName() string {
return "w_task_executions"
}
// ListTaskExecutionsRequest 分页查询任务执行记录请求参数
type ListTaskExecutionsRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Status string `form:"status"`
TaskType string `form:"task_type"`
TaskTypes string `form:"task_types"`
TaskTypePrefix string `form:"task_type_prefix"`
}
// TaskExecutionCleanupStats 任务日志清理结果统计
type TaskExecutionCleanupStats struct {
HighFrequencyDeleted int64 `json:"high_frequency_deleted"`
LowFrequencyDeleted int64 `json:"low_frequency_deleted"`
}
+704
View File
@@ -0,0 +1,704 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/redis/go-redis/v9"
"github.com/shopspring/decimal"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/util"
)
const (
configTypeSystem = "system"
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// PreheatSystemConfigs loads all system configs from database.
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
database := db.DB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// PreheatSystemConfigByKey loads a single config key from database.
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
database := db.DB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
return sc, nil
}
// GetSystemConfigByGroup queries a configuration by Type and Key.
func GetSystemConfigByGroup(ctx context.Context, configType string, key string) (SystemConfig, error) {
ensureSystemConfigCacheListener()
if item, ok := ram.Get(configType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
return sc, nil
}
}
database := db.DB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
valBytes, err := json.Marshal(sc)
if err == nil {
ram.Set(ram.CacheItem{
Key: sc.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(sc.Key),
})
}
return sc, nil
}
// GetSystemConfigByKey queries config by key.
func GetSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
return GetSystemConfigByGroup(ctx, ConfigCacheType, key)
}
// ListSystemConfigsByKeys loads multiple config keys.
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 item, ok := ram.Get(ConfigCacheType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
result[key] = sc
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 {
valBytes, err := json.Marshal(configs[i])
if err == nil {
ram.Set(ram.CacheItem{
Key: configs[i].Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(configs[i].Key),
})
}
result[configs[i].Key] = configs[i]
}
return result, nil
}
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
return InvalidateAllSystemConfigCaches(ctx)
}
// ListVisibleSystemConfigs queries visible configs using local cache store.
func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
ensureSystemConfigCacheListener()
items := ram.GetTypeItems(ConfigCacheType)
if len(items) > 0 {
var list []SystemConfig
for _, item := range items {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
if sc.Visibility == ConfigVisibilityVisible {
list = append(list, sc)
}
}
}
return list, nil
}
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
}
for _, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(cfg.Key),
})
}
}
return configs, nil
}
// GetIntByKey queries config and converts to 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 queries config and converts to 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 queries config and converts to 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 queries and parses menu config.
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
sc, err := GetSystemConfigByKey(ctx, 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
}
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
query := db.DB(ctx).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
var configs []SystemConfig
if err := query.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// GetAdminSystemConfigByKey loads a config directly from DB.
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
var config SystemConfig
if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
return SystemConfig{}, err
}
return config, nil
}
// SystemConfigExists reports whether a config key already exists.
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
var existing 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
}
// CreateSystemConfigRecord persists a new system config row.
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
return db.DB(ctx).Create(config).Error
}
// UpdateSystemConfigFields applies partial updates to a system config row.
func UpdateSystemConfigFields(ctx context.Context, config *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 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 = SystemConfig{
Key: key,
Value: value,
Type: configTypeSystem,
Visibility: 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)
}
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
var templates []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) (Template, error) {
var tmpl Template
if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
return 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 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
}
// CreateTemplateRecord persists a new template.
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Create(tmpl).Error
}
// SaveTemplateRecord updates an existing template.
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Save(tmpl).Error
}
// DeleteTemplateRecord removes a template record.
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Delete(tmpl).Error
}
// CreateScheduleRecord 创建定时任务
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Create(schedule).Error
}
// UpdateScheduleRecord 更新定时任务
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Save(schedule).Error
}
// DeleteScheduleRecord 删除定时任务
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
var schedule Schedule
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedulesRecord 获取所有定时任务
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// CreateTaskExecutionRecord 创建任务执行记录
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
}
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
var execution TaskExecution
err := db.DB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
if err == nil {
if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil {
return nil, false, loadErr
}
return &execution, true, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
_, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
pipe.RPush(ctx, key, line)
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
pipe.Expire(ctx, key, taskExecutionLogExpiration)
return nil
})
if err != nil {
return fmt.Errorf("append task execution log to redis: %w", err)
}
return nil
}
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
logText := strings.Join(logLines, "")
result := db.DB(ctx).Model(&TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
return fmt.Errorf("persist task execution log: %w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
if err := db.Redis.Del(ctx, key).Err(); err != nil {
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
}
return nil
}
// ListTaskExecutionRecords 分页查询任务执行记录
func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := db.DB(ctx).Model(&TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
}
if req.TaskType != "" {
query = query.Where("task_type = ?", req.TaskType)
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
query = query.Where("task_type IN ?", types)
} else if req.TaskTypePrefix != "" {
query = query.Where("task_type LIKE ? ESCAPE '\\'", util.EscapeLike(req.TaskTypePrefix)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []TaskExecution
offset := (req.Page - 1) * req.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
return nil, 0, err
}
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
return nil, 0, err
}
return executions, total, nil
}
func parseTaskTypesFilter(raw string) []string {
if strings.TrimSpace(raw) == "" {
return nil
}
parts := strings.Split(raw, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction.
func MarkFailedTaskExecutionsSucceededTx(
tx *gorm.DB,
taskType string,
result string,
finishedAt time.Time,
) error {
return tx.Model(&TaskExecution{}).
Where("task_type = ? AND status = ?", taskType, TaskExecutionStatusFailed).
Updates(map[string]any{
"status": TaskExecutionStatusSucceeded,
"result": result,
"finished_at": finishedAt,
}).Error
}
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecutionCleanupStats, error) {
const (
frequencyWindowDays = 30
highFrequencyThreshold = frequencyWindowDays
)
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
highFrequencyCutoff := now.AddDate(0, 0, -3)
lowFrequencyCutoff := now.AddDate(0, 0, -30)
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
Model(&TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
}
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
Delete(&TaskExecution{})
if highFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
}
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
}
lowFrequencyResult := lowFrequencyQuery.Delete(&TaskExecution{})
if lowFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
func taskExecutionLogRedisKey(taskID string) string {
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
}
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
if db.Redis == nil {
return nil
}
logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
execution.Log = strings.Join(logLines, "")
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
if db.Redis == nil || len(executions) == 0 {
return nil
}
commands := make([]*redis.StringSliceCmd, len(executions))
_, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
for i := range executions {
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
}
return nil
})
if err != nil {
return fmt.Errorf("get task execution logs from redis: %w", err)
}
for i := range executions {
logLines := commands[i].Val()
if len(logLines) > 0 {
executions[i].Log = strings.Join(logLines, "")
}
}
return nil
}
+208
View File
@@ -0,0 +1,208 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"sync"
"time"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
)
const (
// SystemConfigBroadcastChannel broadcasts system config cache updates across nodes.
SystemConfigBroadcastChannel = "system:config_broadcast"
// SystemConfigInvalidationChannel is kept as an alias for backward compatibility.
SystemConfigInvalidationChannel = SystemConfigBroadcastChannel
// SystemConfigRedisHashKey is kept for backward compatibility in tests.
SystemConfigRedisHashKey = "system:system_configs"
// SystemConfigVisibleListRedisKey is kept for backward compatibility in tests.
SystemConfigVisibleListRedisKey = "system:visible_configs"
// ConfigCacheType is the cache type for all system configs.
ConfigCacheType = "config"
)
type systemConfigBroadcastMessage struct {
Type string `json:"type"`
Key string `json:"key"`
}
// ConfigLoader loads configuration data from the database.
type ConfigLoader struct{}
// LoadAll loads all system configs from database as CacheItems.
func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.CacheItem, error) {
configs, err := PreheatSystemConfigs(ctx)
if err != nil {
return nil, err
}
items := make([]ram.CacheItem, len(configs))
for i, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err != nil {
return nil, err
}
items[i] = ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}
}
return items, nil
}
// LoadOne loads a single system config from database as a CacheItem.
func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string) (ram.CacheItem, error) {
cfg, err := PreheatSystemConfigByKey(ctx, key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ram.CacheItem{}, ram.ErrNotFound
}
return ram.CacheItem{}, err
}
valBytes, err := json.Marshal(cfg)
if err != nil {
return ram.CacheItem{}, err
}
return ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}, nil
}
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
func PreloadSystemConfigs(ctx context.Context) error {
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
}
var (
systemConfigListenerOnce sync.Once
systemConfigListenerCtx context.Context
systemConfigListenerCancel context.CancelFunc
systemConfigListenerDone chan struct{}
)
func ensureSystemConfigCacheListener() {
systemConfigListenerOnce.Do(startSystemConfigCacheInvalidationListener)
}
func startSystemConfigCacheInvalidationListener() {
if db.Redis == nil {
return
}
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
systemConfigListenerDone = make(chan struct{})
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
util.Go(func() {
listenerCtx := systemConfigListenerCtx
defer close(systemConfigListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, SystemConfigBroadcastChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
var payload systemConfigBroadcastMessage
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
ram.UpdateTypeItems(ConfigCacheType, nil)
continue
}
key := payload.Key
if key == "*" || key == "" {
ram.UpdateTypeItems(payload.Type, nil)
} else {
ram.Delete(payload.Type, key)
}
}
})
}
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
func StopSystemConfigCacheListener() {
if systemConfigListenerCancel != nil {
systemConfigListenerCancel()
if systemConfigListenerDone != nil {
<-systemConfigListenerDone
}
systemConfigListenerCancel = nil
systemConfigListenerDone = nil
}
systemConfigListenerOnce = sync.Once{}
}
func determineTTL(_ string) time.Duration {
// Program-determined TTL: -1 means never expire for all configs by default
return -1
}
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
ensureSystemConfigCacheListener()
// Invalidate local cache synchronously first
ram.Delete(ConfigCacheType, key)
// Broadcast to other nodes and clean legacy Redis cache key
if db.Redis != nil {
_ = db.HDel(ctx, SystemConfigRedisHashKey, key)
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
}
return nil
}
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
ensureSystemConfigCacheListener()
// Invalidate all items of type ConfigCacheType synchronously first
ram.UpdateTypeItems(ConfigCacheType, nil)
// Broadcast to other nodes and clean legacy Redis cache keys
if db.Redis != nil {
_ = db.Redis.Del(ctx, db.PrefixedKey(SystemConfigRedisHashKey), db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
}
return nil
}
func publishSystemConfigBroadcast(ctx context.Context, configType string, key string) {
if db.Redis == nil {
return
}
payload, err := json.Marshal(systemConfigBroadcastMessage{Type: configType, Key: key})
if err != nil {
return
}
_ = db.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
}
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
func ResetSystemConfigRAMCacheForTest() {
ram.ResetForTest()
}
+157
View File
@@ -0,0 +1,157 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
)
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("gorm.Open(sqlite) error = %v", err)
}
if err := sqliteDB.AutoMigrate(&SystemConfig{}); err != nil {
t.Fatalf("AutoMigrate(SystemConfig) error = %v", err)
}
siteConfig := SystemConfig{
Key: ConfigKeySiteName,
Value: "Wavelet",
Type: "system",
Description: "系统平台的展示名称",
}
if err := sqliteDB.Create(&siteConfig).Error; err != nil {
t.Fatalf("Create(site_name) error = %v", err)
}
mr, err := miniredis.Run()
if err != nil {
t.Fatalf("miniredis.Run() error = %v", err)
}
redisClient := redis.NewClient(&redis.Options{
Addr: mr.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
previousRedis := db.Redis
db.SetDB(sqliteDB)
db.Redis = redisClient
cleanup := func() {
StopSystemConfigCacheListener()
ResetSystemConfigRAMCacheForTest()
db.SetDB(nil)
db.Redis = previousRedis
_ = redisClient.Close()
mr.Close()
}
return sqliteDB, cleanup
}
func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) {
result, err := ListSystemConfigsByKeys(context.Background(), nil)
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(nil) error = %v", err)
}
if len(result) != 0 {
t.Fatalf("ListSystemConfigsByKeys(nil) = %#v, want empty map", result)
}
}
func TestListSystemConfigsByKeys_LoadsFromRAMCache(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
// Initial load
warm, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
}
if warm.Value != "Wavelet" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
}
// Update DB directly
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
Update("value", "db_only_value").Error; err != nil {
t.Fatalf("Update(site_name) error = %v", err)
}
// Fetch via ListSystemConfigsByKeys should serve from local store (meaning the old value "Wavelet")
configs, err := ListSystemConfigsByKeys(ctx, []string{ConfigKeySiteName})
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err)
}
sc, ok := configs[ConfigKeySiteName]
if !ok {
t.Fatal("ListSystemConfigsByKeys(site_name) missing site_name entry")
}
if sc.Value != "Wavelet" {
t.Fatalf("ListSystemConfigsByKeys(site_name).Value = %q, want cached value %q", sc.Value, "Wavelet")
}
}
func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
// Get via specific group/type
cfg, err := GetSystemConfigByGroup(ctx, ConfigCacheType, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByGroup error = %v", err)
}
if cfg.Value != "Wavelet" {
t.Fatalf("value = %q, want %q", cfg.Value, "Wavelet")
}
// Direct DB update
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
Update("value", "new_site_name").Error; err != nil {
t.Fatalf("DB Update error = %v", err)
}
// Invalidate
if err := InvalidateSystemConfigCache(ctx, ConfigKeySiteName); err != nil {
t.Fatalf("InvalidateSystemConfigCache error = %v", err)
}
// Wait for broadcast execution
time.Sleep(100 * time.Millisecond)
// Fetch again
updated, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey error = %v", err)
}
if updated.Value != "new_site_name" {
t.Fatalf("value = %q, want %q", updated.Value, "new_site_name")
}
}
+2 -2
View File
@@ -8,13 +8,13 @@ import (
"context"
"encoding/json"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
// LogForAudit 将登录鉴权审计日志写入 Logger
func LogForAudit(ctx context.Context, user *model.User, c *gin.Context) {
func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) {
if user == nil || c == nil {
return
}
+34 -50
View File
@@ -7,44 +7,58 @@ import (
"context"
"errors"
"fmt"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
func isOIDCLoginEnabled(ctx context.Context) bool {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
return true
}
b, err := strconv.ParseBool(val)
if err != nil {
return true
}
return enabled
return b
}
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
func resolveAuthSource(ctx context.Context, sourceName string) (*AuthSource, error) {
name := strings.TrimSpace(strings.ToLower(sourceName))
if name == "" {
sources, err := repository.GetActiveAuthSourcesCached(ctx)
sources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil, err
}
if len(sources) == 0 {
return nil, errors.New(errNoActiveAuthSource)
}
return repository.GetAuthSourceByNameCached(ctx, sources[0].Name)
src, err := GetAuthSourceByNameCached(ctx, sources[0].Name)
if err != nil {
return nil, err
}
return src, nil
}
return repository.GetAuthSourceByNameCached(ctx, name)
src, err := GetAuthSourceByNameCached(ctx, name)
if err != nil {
return nil, err
}
return src, nil
}
func activeLoginSources(ctx context.Context) []AuthSourceView {
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyOIDCLoginEnabled)
if err == nil && !enabled {
if !isOIDCLoginEnabled(ctx) {
return nil
}
dbSources, err := repository.GetActiveAuthSourcesCached(ctx)
dbSources, err := GetActiveAuthSourcesCached(ctx)
if err != nil {
return nil
}
@@ -64,14 +78,14 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
}
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
if err != nil || strings.TrimSpace(sc.Value) == "" {
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || strings.TrimSpace(val) == "" {
return "", errors.New(errServerAddressMissing)
}
return strings.TrimRight(sc.Value, "/") + "/login", nil
return strings.TrimRight(val, "/") + "/login", nil
}
func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
func buildOAuthConfig(ctx context.Context, source *AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) {
if source == nil {
return nil, nil, errors.New(errAuthSourceRequired)
}
@@ -116,37 +130,7 @@ func containsScope(scopes []string, scope string) bool {
return false
}
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
existingUsernames, err := repository.ListUsernamesMatchingBase(ctx, base)
if err != nil {
return "", err
}
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
if !exists[strings.ToLower(base)] {
return base, nil
}
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New(errUsernameGenerateFailed)
}
func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code string, nonce string, redirectURL string) (*model.OAuthUserInfo, error) {
func buildOAuthUserInfo(ctx context.Context, source *AuthSource, code string, nonce string, redirectURL string) (*contracts.OAuthUserInfoDTO, error) {
authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL)
if err != nil {
return nil, err
@@ -157,7 +141,7 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri
return nil, err
}
userInfo := &model.OAuthUserInfo{Active: true}
userInfo := &contracts.OAuthUserInfoDTO{Active: true}
if verifier != nil {
if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil {
return nil, verifyErr
@@ -180,7 +164,7 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri
return userInfo, nil
}
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *model.OAuthUserInfo) error {
func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
return nil
@@ -198,7 +182,7 @@ func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *o
return nil
}
func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
func normalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error {
userInfo.Username = strings.TrimSpace(userInfo.Username)
userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername)
userInfo.Email = strings.TrimSpace(userInfo.Email)
@@ -226,7 +210,7 @@ func normalizeOAuthUserInfo(userInfo *model.OAuthUserInfo) error {
return nil
}
func buildCallbackResult(user *model.User, status string) OAuthCallbackResult {
func buildCallbackResult(user *contracts.UserDTO, status string) OAuthCallbackResult {
result := OAuthCallbackResult{Status: status}
if user != nil {
info := BuildBasicUserInfo(user, false)
+22 -15
View File
@@ -10,9 +10,9 @@ import (
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
)
@@ -25,9 +25,16 @@ const (
oauthUserInvalidationChannel = "oauth:user_invalidation"
)
// CachedToken represents the minimal cached representation of an access token.
type CachedToken struct {
ID uint64 `json:"id"`
UserID uint64 `json:"user_id"`
IsAdmin bool `json:"is_admin"`
}
var (
tokenRAM = ram.MustNew[string, *model.AccessToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *model.User](ram.Options{MaximumSize: 2048})
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
userRAM = ram.MustNew[uint64, *contracts.UserDTO](ram.Options{MaximumSize: 2048})
tokenListenerOnce sync.Once
tokenListenerCtx context.Context
@@ -136,8 +143,8 @@ func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
_ = db.Redis.Publish(ctx, oauthUserInvalidationChannel, strconv.FormatUint(userID, 10)).Err()
}
// GetCachedToken 获取缓存的 AccessToken
func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) {
// GetCachedToken 获取缓存的 Token
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
ensureTokenCacheListener()
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
@@ -145,7 +152,7 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken,
}
if db.Redis != nil {
var token model.AccessToken
var token CachedToken
key := tokenCacheKey(tokenHash)
if err := db.GetJSON(ctx, key, &token); err == nil {
// Write back to local cache
@@ -156,8 +163,8 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken,
return nil, fmt.Errorf("cache miss")
}
// SetCachedToken 设置 AccessToken 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) {
// SetCachedToken 设置 Token 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
ensureTokenCacheListener()
tokenRAM.Set(tokenHash, token)
@@ -179,8 +186,8 @@ func InvalidateCachedToken(ctx context.Context, tokenHash string) {
}
}
// GetCachedUser 获取缓存的 User
func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
// GetCachedUser 获取缓存的 UserDTO
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
ensureUserCacheListener()
if val, ok := userRAM.GetIfPresent(userID); ok {
@@ -188,7 +195,7 @@ func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
}
if db.Redis != nil {
var u model.User
var u contracts.UserDTO
key := userCacheKey(userID)
if err := db.GetJSON(ctx, key, &u); err == nil {
// Write back to local cache
@@ -199,8 +206,8 @@ func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
return nil, fmt.Errorf("cache miss")
}
// SetCachedUser 设置 User 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *model.User) {
// SetCachedUser 设置 UserDTO 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
ensureUserCacheListener()
userRAM.Set(userID, u)
@@ -210,7 +217,7 @@ func SetCachedUser(ctx context.Context, userID uint64, u *model.User) {
}
}
// InvalidateCachedUser 吊销/失效 User 缓存
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
ensureUserCacheListener()
+8 -9
View File
@@ -11,8 +11,8 @@ import (
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
)
@@ -49,11 +49,10 @@ func TestTokenCache_GetSetInvalidate(t *testing.T) {
ctx := context.Background()
tokenHash := "test-token-hash"
token := &model.AccessToken{
ID: 123,
UserID: 456,
TokenHash: tokenHash,
Name: "test-token",
token := &auth.CachedToken{
ID: 123,
UserID: 456,
IsAdmin: true,
}
// 1. Get from empty cache -> miss
@@ -70,7 +69,7 @@ func TestTokenCache_GetSetInvalidate(t *testing.T) {
if err != nil {
t.Fatalf("GetCachedToken() failed: %v", err)
}
if cached.ID != token.ID || cached.UserID != token.UserID {
if cached.ID != token.ID || cached.UserID != token.UserID || cached.IsAdmin != token.IsAdmin {
t.Fatalf("expected cached token %+v, got %+v", token, cached)
}
@@ -90,7 +89,7 @@ func TestUserCache_GetSetInvalidate(t *testing.T) {
ctx := context.Background()
userID := uint64(789)
user := &model.User{
user := &contracts.UserDTO{
ID: userID,
Username: "testuser",
Email: "test@example.com",
+83 -38
View File
@@ -13,13 +13,16 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/shared"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
@@ -91,7 +94,7 @@ func GetLoginURL(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL}))
}
func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state string) (string, error) {
func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (string, error) {
redirectURL, err := getFrontendLoginRedirectURL(ctx)
if err != nil {
return "", err
@@ -282,18 +285,18 @@ func Callback(c *gin.Context) {
handleCallbackLogin(ctx, c, source, userInfo)
}
func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
userID := GetUserIDFromContext(c)
if userID == 0 {
response.AbortUnauthorized(c, shared.UnAuthorized)
return
}
user, err := repository.GetUserByID(ctx, userID)
if err != nil {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
@@ -304,22 +307,20 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS
return
}
user.LastLoginAt = time.Now()
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
}
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
var user model.User
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
var user contracts.UserDTO
account, err := repository.FindExternalAccount(ctx, source.ID, userInfo.Sub)
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
loaded, loadErr := repository.GetUserByID(ctx, account.UserID)
if loadErr != nil {
if loadErr := db.DB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
response.AbortInternal(c, loadErr.Error())
return
}
user = loaded
case errors.Is(err, gorm.ErrRecordNotFound):
newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
if !ok {
@@ -332,7 +333,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
}
user.LastLoginAt = time.Now()
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
if err := SetLoginSession(ctx, c, &user); err != nil {
response.AbortInternal(c, err.Error())
return
@@ -340,37 +341,81 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
SetCachedUser(ctx, user.ID, &user)
logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "logged_in")))
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
if regErr != nil {
registrationEnabled = false
func uniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = "user"
}
var existingUsernames []string
if err := db.DB(ctx).Table("w_users").
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return "", err
}
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
if !exists[strings.ToLower(base)] {
return base, nil
}
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New(errUsernameGenerateFailed)
}
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
registrationEnabled := true
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
if b, err := strconv.ParseBool(val); err == nil {
registrationEnabled = b
}
}
if !registrationEnabled {
c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind")))
return model.User{}, false
return contracts.UserDTO{}, false
}
username, uniqueErr := uniqueUsername(ctx, userInfo.Username)
if uniqueErr != nil {
response.AbortInternal(c, uniqueErr.Error())
return model.User{}, false
return contracts.UserDTO{}, false
}
userInfo.Username = username
var user model.User
if err := repository.CreateUserFromOAuth(ctx, &user, userInfo); err != nil {
response.AbortInternal(c, err.Error())
return model.User{}, false
now := time.Now()
user := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: userInfo.Username,
Nickname: userInfo.Name,
Email: userInfo.Email,
AvatarURL: userInfo.AvatarURL,
IsActive: userInfo.Active,
LastLoginAt: now,
CreatedAt: now,
UpdatedAt: now,
}
if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
if err := db.DB(ctx).Table("w_users").Create(&user).Error; err != nil {
response.AbortInternal(c, err.Error())
return contracts.UserDTO{}, false
}
if err := BindExternalAccount(ctx, &ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
@@ -378,7 +423,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A
Email: userInfo.Email,
}); err != nil {
response.AbortBadRequest(c, err.Error())
return model.User{}, false
return contracts.UserDTO{}, false
}
logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP())
@@ -387,7 +432,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A
// UserInfo 获取当前登录用户信息
func UserInfo(c *gin.Context) {
user, _ := GetFromContext[*model.User](c, UserObjKey)
user, _ := GetFromContext[*contracts.UserDTO](c, UserObjKey)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true
@@ -420,7 +465,7 @@ func Logout(c *gin.Context) {
// ListExternalAccounts 获取当前用户的外部帐号绑定列表
func ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := repository.ListExternalAccountsByUserID(c.Request.Context(), userID)
accounts, err := ListExternalAccountsByUserID(c.Request.Context(), userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -441,7 +486,7 @@ func DeleteExternalAccount(c *gin.Context) {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return
}
if err := repository.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
if err := UnbindExternalAccount(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
+43 -36
View File
@@ -6,14 +6,15 @@ package auth
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/shared"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
@@ -33,32 +34,47 @@ func SetToContext[T any](c *gin.Context, key string, value T) {
c.Set(key, value)
}
func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) {
tokenHash := model.HashToken(tokenStr)
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) {
tokenHash := hashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
if err == nil {
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err == nil && user != nil && user.IsActive {
return user, tokenRecord, nil
}
tokenRecord = &dbToken
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive {
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
user = &dbUser
SetCachedUser(ctx, tokenRecord.UserID, user)
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
return user, tokenRecord, nil
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, nil, err
}
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
var userRow contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRow.UserID, true).First(&userRow).Error; err != nil {
return nil, nil, err
}
SetCachedUser(ctx, userRow.ID, &userRow)
return &userRow, tokenRecord, nil
}
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
func GetUserFromRequest(c *gin.Context) (*model.User, error) {
func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
ctx := c.Request.Context()
// Check token in headers
@@ -89,24 +105,15 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
}
user, err := GetCachedUser(ctx, userID)
if err != nil || !user.IsActive {
dbUser, loadErr := repository.GetActiveUserByID(ctx, userID)
if loadErr != nil {
return nil, loadErr
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
SetCachedUser(ctx, userID, user)
}
// 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致
if user.Password != "" {
session := sessions.Default(c)
sessionHash, _ := session.Get(PasswordHashKey).(string)
if sessionHash != user.Password {
return nil, errors.New("session expired due to password change")
}
}
SetToContext(c, TokenAuthKey, false)
SetToContext(c, TokenAdminKey, false)
+3 -3
View File
@@ -12,7 +12,7 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/core/contracts"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
@@ -162,8 +162,8 @@ type BasicUserInfo struct {
Location string `json:"location"`
}
// BuildBasicUserInfo 将 User 模型转换为 BasicUserInfo
func BuildBasicUserInfo(user *model.User, needChange bool) BasicUserInfo {
// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo
func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo {
if user == nil {
return BasicUserInfo{}
}
+37 -10
View File
@@ -5,8 +5,11 @@ package auth_test
import (
"context"
"crypto/sha256"
"encoding/hex"
"path/filepath"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
@@ -15,11 +18,35 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
)
type testUser struct {
ID uint64 `gorm:"primaryKey"`
Username string
IsActive bool
LastLoginAt time.Time
}
func (testUser) TableName() string { return "w_users" }
type testAccessToken struct {
ID uint64 `gorm:"primaryKey"`
UserID uint64
TokenHash string
Name string
IsAdmin bool
}
func (testAccessToken) TableName() string { return "w_access_tokens" }
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
func setupTestDB(t *testing.T) *gorm.DB {
t.Helper()
dbPath := filepath.Join(t.TempDir(), "auth_test.db")
@@ -27,10 +54,10 @@ func setupTestDB(t *testing.T) *gorm.DB {
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&model.User{},
&model.AccessToken{},
&model.AuthSource{},
&model.ExternalAccount{},
&testUser{},
&testAccessToken{},
&auth.AuthSource{},
&auth.ExternalAccount{},
))
db.SetDB(testDB)
@@ -75,7 +102,7 @@ func TestAuthPluginUnit(t *testing.T) {
assert.Equal(t, "custom", prov.Name())
// Test User Token Verification with dummy token
user := model.User{
user := testUser{
ID: 101,
Username: "token_user",
IsActive: true,
@@ -83,8 +110,8 @@ func TestAuthPluginUnit(t *testing.T) {
require.NoError(t, testDB.Create(&user).Error)
tokenStr := "test-secret-token-123456"
tokenHash := model.HashToken(tokenStr)
tokenRecord := model.AccessToken{
tokenHash := hashToken(tokenStr)
tokenRecord := testAccessToken{
ID: 201,
UserID: user.ID,
TokenHash: tokenHash,
@@ -106,7 +133,7 @@ func TestAuthPluginUnit(t *testing.T) {
require.NoError(t, authSvc.RevokeUserSessions(context.Background(), user.ID))
// GetCurrentUser from context
userCtx := context.WithValue(context.Background(), "user_obj", userDTO)
userCtx := context.WithValue(context.Background(), auth.UserObjKey, userDTO)
current, err := authSvc.GetCurrentUser(userCtx)
require.NoError(t, err)
assert.Equal(t, user.ID, current.ID)
+76
View File
@@ -0,0 +1,76 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
)
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
var src AuthSource
if err := db.DB(ctx).First(&src, id).Error; err != nil {
return nil, err
}
return &src, nil
}
// GetAuthSourceByName 根据名称获取认证源
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
var src AuthSource
if err := db.DB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
return nil, err
}
return &src, nil
}
// ListActiveAuthSources 获取所有启用的认证源
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := db.DB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
return nil, err
}
return sources, nil
}
// GetActiveAuthSourcesCached 获取所有启用的认证源(带缓存或直接查询)
func GetActiveAuthSourcesCached(ctx context.Context) ([]AuthSource, error) {
return ListActiveAuthSources(ctx)
}
// GetAuthSourceByNameCached 根据名称获取认证源(带缓存或直接查询)
func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, error) {
return GetAuthSourceByName(ctx, name)
}
// FindExternalAccount 查询指定认证源的外部账号绑定
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部账号
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
return db.DB(ctx).Create(account).Error
}
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
var accounts []ExternalAccount
if err := db.DB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// UnbindExternalAccount 解绑外部账号
func UnbindExternalAccount(ctx context.Context, id uint64, userID uint64) error {
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
+19 -38
View File
@@ -9,34 +9,10 @@ import (
"sync"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/gin-gonic/gin"
)
func toUserDTO(u *model.User) *contracts.UserDTO {
if u == nil {
return nil
}
return &contracts.UserDTO{
ID: u.ID,
Username: u.Username,
Nickname: u.Nickname,
Email: u.Email,
AvatarURL: u.AvatarURL,
IsActive: u.IsActive,
IsAdmin: u.IsAdmin,
Bio: u.Bio,
Phone: u.Phone,
Gender: u.Gender,
Website: u.Website,
Location: u.Location,
LastLoginAt: u.LastLoginAt,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
}
}
type authServiceImpl struct{}
func newAuthService() contracts.AuthService {
@@ -53,15 +29,12 @@ func (s *authServiceImpl) RequireAdminMiddleware() any {
func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
if ginCtx, ok := ctx.(*gin.Context); ok {
if u, ok := GetFromContext[*model.User](ginCtx, UserObjKey); ok && u != nil {
return toUserDTO(u), nil
if u, ok := GetFromContext[*contracts.UserDTO](ginCtx, UserObjKey); ok && u != nil {
return u, nil
}
}
if v := ctx.Value(UserObjKey); v != nil {
if u, ok := v.(*model.User); ok && u != nil {
return toUserDTO(u), nil
}
if u, ok := v.(*contracts.UserDTO); ok && u != nil {
return u, nil
}
@@ -75,21 +48,29 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
return nil, errors.New("auth: empty token")
}
tokenHash := model.HashToken(token)
tokenHash := hashToken(token)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
var tokenRow struct {
ID uint64
UserID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
return nil, err
}
tokenRecord = &dbToken
tokenRecord = &CachedToken{
ID: tokenRow.ID,
UserID: tokenRow.UserID,
IsAdmin: tokenRow.IsAdmin,
}
SetCachedToken(ctx, tokenHash, tokenRecord)
}
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive {
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
if err != nil || user == nil || !user.IsActive {
var dbUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
return nil, err
}
user = &dbUser
@@ -100,7 +81,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
return nil, errors.New("auth: system user token not allowed")
}
return toUserDTO(user), nil
return user, nil
}
func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
+21 -16
View File
@@ -8,11 +8,12 @@ import (
"crypto/sha256"
"encoding/hex"
"net/http"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/config"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
@@ -66,7 +67,10 @@ func GetUserIDFromSession(s sessions.Session) uint64 {
}
// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID
func GetUserIDFromContext(c *gin.Context) uint64 {
func GetUserIDFromContext(c *gin.Context) (uid uint64) {
defer func() {
_ = recover()
}()
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
@@ -96,14 +100,13 @@ func rotateSessionID(s sessions.Session) {
}
// SetLoginSession writes the authenticated user into a freshly rotated session.
func SetLoginSession(ctx context.Context, c *gin.Context, user *model.User, extras ...map[string]any) error {
func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error {
session := sessions.Default(c)
session.Clear()
rotateSessionID(session)
session.Set(UserIDKey, user.ID)
session.Set(UserNameKey, user.Username)
session.Set(PasswordHashKey, user.Password)
if len(extras) > 0 {
for key, value := range extras[0] {
session.Set(key, value)
@@ -114,16 +117,18 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *model.User, extr
maxAge := config.Config.App.SessionAge
isSessionCookie := false
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
if err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "login_session_ttl_hours").Pluck("value", &val).Error; err == nil && val != "" {
if ttlHours, err := strconv.Atoi(val); err == nil {
switch {
case ttlHours == -1:
// 永不过期,设置为 10 年
maxAge = 10 * 365 * 24 * 3600
case ttlHours > 0:
maxAge = ttlHours * 3600
case ttlHours == 0:
isSessionCookie = true
}
}
}
session.Options(GetSessionOptions(maxAge))
+1 -1
View File
@@ -6,9 +6,9 @@ package cap
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/shared/response"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
+2 -2
View File
@@ -13,9 +13,9 @@ import (
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/persistence"
)
const (
+1 -1
View File
@@ -6,7 +6,7 @@ package cap
import (
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/response"
)
// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
+50 -27
View File
@@ -14,9 +14,8 @@ import (
"golang.org/x/sync/singleflight"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
)
@@ -38,13 +37,25 @@ type RuntimeSettings struct {
TokenTTL time.Duration
}
// CAP 动态配置键常量
const (
ConfigKeyCapLoginEnabled = "cap_login_enabled"
ConfigKeyCapChallengeCount = "cap_challenge_count"
ConfigKeyCapChallengeSize = "cap_challenge_size"
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty"
ConfigKeyCapChallengeTTL = "cap_challenge_ttl"
// ConfigKeyCapTokenTTL 验证码 Token 过期时间键
// #nosec G101
ConfigKeyCapTokenTTL = "cap_token_ttl"
)
var runtimeConfigKeys = []string{
model.ConfigKeyCapLoginEnabled,
model.ConfigKeyCapChallengeCount,
model.ConfigKeyCapChallengeSize,
model.ConfigKeyCapChallengeDifficulty,
model.ConfigKeyCapChallengeTTL,
model.ConfigKeyCapTokenTTL,
ConfigKeyCapLoginEnabled,
ConfigKeyCapChallengeCount,
ConfigKeyCapChallengeSize,
ConfigKeyCapChallengeDifficulty,
ConfigKeyCapChallengeTTL,
ConfigKeyCapTokenTTL,
}
var runtimeConfigKeySet = func() map[string]struct{} {
@@ -132,14 +143,22 @@ func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, er
}
func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
configs, err := repository.ListSystemConfigsByKeys(ctx, runtimeConfigKeys)
if err != nil {
type configRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
var records []configRecord
if err := db.DB(ctx).Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
return RuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
for _, r := range records {
configs[r.Key] = r.Value
}
return parseRuntimeSettings(configs), nil
}
func parseRuntimeSettings(configs map[string]model.SystemConfig) RuntimeSettings {
func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
settings := RuntimeSettings{
ChallengeCount: defaultChallengeCount,
ChallengeSize: defaultChallengeSize,
@@ -148,33 +167,33 @@ func parseRuntimeSettings(configs map[string]model.SystemConfig) RuntimeSettings
TokenTTL: defaultTokenTTL,
}
if sc, ok := configs[model.ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(sc.Value); err == nil {
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(sc.Value); err == nil && count > 0 {
if val, ok := configs[ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(val); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(sc.Value); err == nil && size > 0 {
if val, ok := configs[ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(val); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeDifficulty]; ok {
if difficulty, err := strconv.Atoi(sc.Value); err == nil && difficulty > 0 {
settings.ChallengeDifficulty = difficulty
if val, ok := configs[ConfigKeyCapChallengeDifficulty]; ok {
if diff, err := strconv.Atoi(val); err == nil && diff > 0 {
settings.ChallengeDifficulty = diff
}
}
if sc, ok := configs[model.ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(sc.Value); err == nil && ttlSeconds > 0 {
if val, ok := configs[ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if sc, ok := configs[model.ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(sc.Value); err == nil && ttlSeconds > 0 {
if val, ok := configs[ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
@@ -186,13 +205,17 @@ func (s *runtimeSettingsStore) ensureInvalidationListener() {
s.listenerOnce.Do(startRuntimeSettingsInvalidationListener)
}
// SystemConfigInvalidationChannel 系统配置失效广播通道
const SystemConfigInvalidationChannel = "system_config:invalidation"
func startRuntimeSettingsInvalidationListener() {
if db.Redis == nil {
rdb := db.Redis
if rdb == nil {
return
}
util.Go(func() {
pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel)
pubsub := rdb.Subscribe(context.Background(), SystemConfigInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
+13 -13
View File
@@ -18,8 +18,8 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/admin"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
@@ -38,17 +38,17 @@ func setupTestDB(t *testing.T) *gorm.DB {
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&model.User{},
&model.AccessToken{},
&model.AuthSource{},
&model.ExternalAccount{},
&model.MessageChannel{},
&model.MessageBinding{},
&model.MessagePairingCode{},
&model.SystemConfig{},
&model.PushChannel{},
&model.PushEvent{},
&model.PushHistory{},
&user.User{},
&user.AccessToken{},
&auth.AuthSource{},
&auth.ExternalAccount{},
&message_gateway.MessageChannel{},
&message_gateway.MessageBinding{},
&message_gateway.MessagePairingCode{},
&admin.SystemConfig{},
&message_gateway.PushChannel{},
&message_gateway.PushEvent{},
&message_gateway.PushHistory{},
))
db.SetDB(testDB)
@@ -7,7 +7,7 @@ import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
@@ -20,7 +20,7 @@ import (
// @Success 200 {object} response.Any{data=[]Definition}
// @Router /api/v1/admin/message-gateway/channels/definitions [get]
func ListAdminChannelDefinitions(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(channelDefinitions()))
c.JSON(http.StatusOK, response.OK(listDefinitions()))
}
// ListAdminChannels lists configured messaging channels with secrets masked.
+182 -190
View File
@@ -13,8 +13,6 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/tencent-connect/botgo/token"
"gorm.io/gorm"
)
@@ -31,68 +29,52 @@ type Field struct {
// Definition describes a channel type form.
type Definition struct {
Type string `json:"type"`
Name string `json:"name"`
Fields []Field `json:"fields"`
}
// CreateChannelRequest is the admin create body.
type CreateChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Enabled *bool `json:"enabled"`
BotToken string `json:"bot_token"`
AppID string `json:"app_id"`
AppSecret string `json:"app_secret"`
BaseURL string `json:"base_url"`
PortalHost string `json:"portal_host"`
Sandbox string `json:"sandbox"`
}
// UpdateChannelRequest is the admin patch body.
type UpdateChannelRequest struct {
Name *string `json:"name"`
Enabled *bool `json:"enabled"`
BotToken string `json:"bot_token"`
AppID string `json:"app_id"`
AppSecret string `json:"app_secret"`
BaseURL *string `json:"base_url"`
PortalHost *string `json:"portal_host"`
Sandbox *string `json:"sandbox"`
}
// ChannelDTO is a list/detail view with secrets masked.
// ChannelDTO represents a channel for admin consumption.
type ChannelDTO struct {
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
OwnerScope string `json:"owner_scope"`
Enabled bool `json:"enabled"`
BotToken string `json:"bot_token,omitempty"`
AppID string `json:"app_id,omitempty"`
AppSecret string `json:"app_secret,omitempty"`
BaseURL string `json:"base_url,omitempty"`
PortalHost string `json:"portal_host,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
ID uint64 `json:"id,string"`
Name string `json:"name"`
Type string `json:"type"`
OwnerScope string `json:"owner_scope"`
OwnerID *uint64 `json:"owner_id,string,omitempty"`
Enabled bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
func channelDefinitions() []Definition {
// CreateChannelRequest is admin create payload.
type CreateChannelRequest struct {
Name string `json:"name"`
Type string `json:"type"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
// UpdateChannelRequest is admin update payload.
type UpdateChannelRequest struct {
Name string `json:"name"`
Enabled *bool `json:"enabled"`
Credentials map[string]string `json:"credentials"`
Extra map[string]string `json:"extra"`
}
func listDefinitions() []Definition {
return []Definition{
{
Type: model.MessageChannelTypeTelegram,
Name: "Telegram",
Type: MessageChannelTypeTelegram,
Fields: []Field{
{Key: "bot_token", Type: "password", Required: true},
{Key: "base_url", Type: "text"},
{Key: "token", Type: "password", Required: true},
{Key: "api_base", Type: "text", Required: false},
},
},
{
Type: model.MessageChannelTypeQQ,
Name: "QQ",
Type: MessageChannelTypeQQ,
Fields: []Field{
{Key: "app_id", Required: true},
{Key: "app_secret", Type: "password", Required: true},
{Key: "portal_host", Type: "text"},
{Key: "app_id", Type: "text", Required: true},
{Key: "client_secret", Type: "password", Required: true},
},
},
}
@@ -103,35 +85,45 @@ func createChannel(ctx context.Context, req CreateChannelRequest) (ChannelDTO, e
if name == "" {
return ChannelDTO{}, errors.New(errNameRequired)
}
typ := strings.TrimSpace(req.Type)
creds, extra, err := credentialsFromCreate(req)
if err != nil {
channelType := strings.TrimSpace(req.Type)
if channelType != MessageChannelTypeTelegram && channelType != MessageChannelTypeQQ {
return ChannelDTO{}, errors.New(errTypeInvalid)
}
creds := req.Credentials
if creds == nil {
creds = map[string]string{}
}
if err := validateCredentials(channelType, creds, false); err != nil {
return ChannelDTO{}, err
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
extra := req.Extra
if extra == nil {
extra = map[string]string{}
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
row := &model.MessageChannel{
row := &MessageChannel{
Name: name,
Type: typ,
OwnerScope: model.MessageOwnerScopeSystem,
Type: channelType,
OwnerScope: MessageOwnerScopeSystem,
Enabled: enabled,
Credentials: cipher,
Extra: EncodeExtra(extra),
}
if err := repository.CreateMessageChannel(ctx, row); err != nil {
if err := CreateMessageChannel(ctx, row); err != nil {
return ChannelDTO{}, err
}
return toDTO(row, creds, extra), nil
}
func updateChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (ChannelDTO, error) {
row, err := repository.GetMessageChannel(ctx, id)
row, err := GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ChannelDTO{}, errors.New(errChannelNotFound)
@@ -140,80 +132,74 @@ func updateChannel(ctx context.Context, id uint64, req UpdateChannelRequest) (Ch
}
creds, err := DecryptCredentials(row.Credentials)
if err != nil {
creds = map[string]string{}
return ChannelDTO{}, err
}
extra := ParseExtra(row.Extra)
if req.Name != nil {
name := strings.TrimSpace(*req.Name)
if name == "" {
return ChannelDTO{}, errors.New(errNameRequired)
}
if name := strings.TrimSpace(req.Name); name != "" {
row.Name = name
}
if req.Enabled != nil {
row.Enabled = *req.Enabled
}
if token := strings.TrimSpace(req.BotToken); token != "" {
creds["bot_token"] = token
if req.Extra != nil {
extra = req.Extra
}
if appID := strings.TrimSpace(req.AppID); appID != "" {
creds["app_id"] = appID
}
if secret := strings.TrimSpace(req.AppSecret); secret != "" {
creds["app_secret"] = secret
}
if req.BaseURL != nil {
extra["base_url"] = strings.TrimSpace(*req.BaseURL)
}
if req.PortalHost != nil {
extra["portal_host"] = strings.TrimSpace(*req.PortalHost)
}
if req.Sandbox != nil {
extra["sandbox"] = strings.TrimSpace(*req.Sandbox)
}
if err := validateCredentials(row.Type, creds); err != nil {
return ChannelDTO{}, err
if len(req.Credentials) > 0 {
merged := make(map[string]string, len(creds))
for k, v := range creds {
merged[k] = v
}
for k, v := range req.Credentials {
if strings.TrimSpace(v) == "" {
continue
}
merged[k] = v
}
if err := validateCredentials(row.Type, merged, true); err != nil {
return ChannelDTO{}, err
}
creds = merged
}
cipher, err := EncryptCredentials(creds)
if err != nil {
return ChannelDTO{}, err
}
row.Credentials = cipher
row.Extra = EncodeExtra(extra)
if err := repository.UpdateMessageChannel(ctx, row); err != nil {
if err := UpdateMessageChannel(ctx, row); err != nil {
return ChannelDTO{}, err
}
return toDTO(row, creds, extra), nil
}
func listChannels(ctx context.Context) ([]ChannelDTO, error) {
rows, err := repository.ListMessageChannels(ctx)
rows, err := ListMessageChannels(ctx)
if err != nil {
return nil, err
}
out := make([]ChannelDTO, 0, len(rows))
for i := range rows {
creds, err := DecryptCredentials(rows[i].Credentials)
if err != nil {
creds = map[string]string{}
}
out = append(out, toDTO(&rows[i], creds, ParseExtra(rows[i].Extra)))
creds, _ := DecryptCredentials(rows[i].Credentials)
extra := ParseExtra(rows[i].Extra)
out = append(out, toDTO(&rows[i], creds, extra))
}
return out, nil
}
func deleteChannel(ctx context.Context, id uint64) error {
if _, err := repository.GetMessageChannel(ctx, id); err != nil {
if _, err := GetMessageChannel(ctx, id); err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errChannelNotFound)
}
return err
}
return repository.DeleteMessageChannel(ctx, id)
return DeleteMessageChannel(ctx, id)
}
func probeChannel(ctx context.Context, id uint64) error {
row, err := repository.GetMessageChannel(ctx, id)
row, err := GetMessageChannel(ctx, id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errChannelNotFound)
@@ -224,51 +210,90 @@ func probeChannel(ctx context.Context, id uint64) error {
if err != nil {
return err
}
extra := ParseExtra(row.Extra)
if err := probeCredentials(ctx, row.Type, creds, extra); err != nil {
return fmt.Errorf("%s: %w", errChannelProbeFailed, err)
switch row.Type {
case MessageChannelTypeTelegram:
return probeTelegram(ctx, creds)
case MessageChannelTypeQQ:
return probeQQ(ctx, creds)
default:
return errors.New(errTypeInvalid)
}
}
func probeTelegram(ctx context.Context, creds map[string]string) error {
tok := creds["token"]
if strings.TrimSpace(tok) == "" {
return errors.New("missing telegram bot token")
}
base := creds["api_base"]
base = strings.TrimRight(strings.TrimSpace(base), "/")
if base == "" {
base = defaultTelegramAPI
}
url := fmt.Sprintf("%s/bot%s/getMe", base, tok)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("telegram getMe failed (%d): %s", resp.StatusCode, string(body))
}
var res struct {
OK bool `json:"ok"`
}
if err := json.Unmarshal(body, &res); err != nil {
return err
}
if !res.OK {
return fmt.Errorf("telegram returned ok=false: %s", string(body))
}
return nil
}
func credentialsFromCreate(req CreateChannelRequest) (map[string]string, map[string]string, error) {
typ := strings.TrimSpace(req.Type)
creds := map[string]string{}
extra := map[string]string{}
switch typ {
case model.MessageChannelTypeTelegram:
creds["bot_token"] = strings.TrimSpace(req.BotToken)
if base := strings.TrimSpace(req.BaseURL); base != "" {
extra["base_url"] = base
}
case model.MessageChannelTypeQQ:
creds["app_id"] = strings.TrimSpace(req.AppID)
creds["app_secret"] = strings.TrimSpace(req.AppSecret)
if host := strings.TrimSpace(req.PortalHost); host != "" {
extra["portal_host"] = host
} else {
extra["portal_host"] = "q.qq.com"
}
if sandbox := strings.TrimSpace(req.Sandbox); sandbox != "" {
extra["sandbox"] = sandbox
}
default:
return nil, nil, errors.New(errTypeInvalid)
func probeQQ(_ context.Context, creds map[string]string) error {
appID := strings.TrimSpace(creds["app_id"])
secret := strings.TrimSpace(creds["app_secret"])
if appID == "" || secret == "" {
return errors.New("missing qq app_id or app_secret")
}
if err := validateCredentials(typ, creds); err != nil {
return nil, nil, err
credentials := &token.QQBotCredentials{
AppID: appID,
AppSecret: secret,
}
return creds, extra, nil
tokSrc := token.NewQQBotTokenSource(credentials)
tok, err := tokSrc.Token()
if err != nil {
return fmt.Errorf("qq token fetch failed: %w", err)
}
if tok == nil || tok.AccessToken == "" {
return errors.New("qq returned empty access token")
}
return nil
}
func validateCredentials(typ string, creds map[string]string) error {
switch typ {
case model.MessageChannelTypeTelegram:
if strings.TrimSpace(creds["bot_token"]) == "" {
func validateCredentials(t string, creds map[string]string, isUpdate bool) error {
switch t {
case MessageChannelTypeTelegram:
tok := creds["token"]
if strings.TrimSpace(tok) == "" && !isUpdate {
return errors.New(errTelegramTokenRequired)
}
case model.MessageChannelTypeQQ:
if strings.TrimSpace(creds["app_id"]) == "" || strings.TrimSpace(creds["app_secret"]) == "" {
if base, ok := creds["api_base"]; ok && strings.TrimSpace(base) != "" {
if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") {
return errors.New("api_base must start with http:// or https://")
}
}
case MessageChannelTypeQQ:
appID := creds["app_id"]
secret := creds["client_secret"]
if (strings.TrimSpace(appID) == "" || strings.TrimSpace(secret) == "") && !isUpdate {
return errors.New(errQQCredentialsRequired)
}
default:
@@ -277,70 +302,37 @@ func validateCredentials(typ string, creds map[string]string) error {
return nil
}
func toDTO(row *model.MessageChannel, creds, extra map[string]string) ChannelDTO {
dto := ChannelDTO{
ID: row.ID,
Name: row.Name,
Type: row.Type,
OwnerScope: row.OwnerScope,
Enabled: row.Enabled,
CreatedAt: row.CreatedAt,
UpdatedAt: row.UpdatedAt,
func toDTO(row *MessageChannel, creds, extra map[string]string) ChannelDTO {
return ChannelDTO{
ID: row.ID,
Name: row.Name,
Type: row.Type,
OwnerScope: row.OwnerScope,
OwnerID: row.OwnerID,
Enabled: row.Enabled,
Credentials: maskCredentials(row.Type, creds),
Extra: extra,
}
if strings.TrimSpace(creds["bot_token"]) != "" {
dto.BotToken = maskedSecret
}
if id := strings.TrimSpace(creds["app_id"]); id != "" {
dto.AppID = id
}
if strings.TrimSpace(creds["app_secret"]) != "" {
dto.AppSecret = maskedSecret
}
dto.BaseURL = extra["base_url"]
dto.PortalHost = extra["portal_host"]
return dto
}
func probeCredentials(ctx context.Context, typ string, creds, extra map[string]string) error {
switch typ {
case model.MessageChannelTypeTelegram:
base := strings.TrimSpace(extra["base_url"])
if base == "" {
base = defaultTelegramAPI
func maskCredentials(_ string, in map[string]string) map[string]string {
out := make(map[string]string, len(in))
for k, v := range in {
if k == "token" || k == "client_secret" {
out[k] = maskSecret(v)
} else {
out[k] = v
}
url := strings.TrimRight(base, "/") + "/bot" + creds["bot_token"] + "/getMe"
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
const probeBodyLimit = 4096
body, _ := io.ReadAll(io.LimitReader(resp.Body, probeBodyLimit))
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("telegram getMe status %d", resp.StatusCode)
}
var parsed struct {
OK bool `json:"ok"`
}
if err := json.Unmarshal(body, &parsed); err != nil {
return err
}
if !parsed.OK {
return errors.New("telegram getMe returned ok=false")
}
return nil
case model.MessageChannelTypeQQ:
src := token.NewQQBotTokenSource(&token.QQBotCredentials{
AppID: creds["app_id"],
AppSecret: creds["app_secret"],
})
_, err := src.Token()
return err
default:
return errors.New(errTypeInvalid)
}
return out
}
const minMaskSecretLength = 8
func maskSecret(s string) string {
s = strings.TrimSpace(s)
if len(s) <= minMaskSecretLength {
return "******"
}
return s[:4] + "..." + s[len(s)-4:]
}
@@ -7,7 +7,7 @@ import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/core/contracts"
)
// AdminLogin is the metadata definition for the admin login event.
@@ -22,7 +22,8 @@ var AdminLogin = EventMetadata{
Description: "当管理员成功登录系统时触发此通知",
}
func handleAdminLogin(ctx context.Context, event listener.AdminLoggedIn) {
// HandleAdminLoggedIn 处理管理员登录事件并触发通知
func HandleAdminLoggedIn(ctx context.Context, event contracts.AdminLoggedIn) {
if event.User == nil {
return
}
@@ -38,5 +39,4 @@ func handleAdminLogin(ctx context.Context, event listener.AdminLoggedIn) {
// RegisterCustomEvents registers default domain push notification events.
func RegisterCustomEvents() {
RegisterBuiltInEvent(AdminLogin)
listener.OnAdminLoggedIn(handleAdminLogin)
}
+4 -4
View File
@@ -8,14 +8,14 @@ import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-gonic/gin"
)
func currentUser(c *gin.Context) (*model.User, bool) {
return auth.GetFromContext[*model.User](c, auth.UserObjKey)
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
return auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
}
// ListChannels lists enabled channels a user can bind.
+14 -15
View File
@@ -10,9 +10,8 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
pkgmg "github.com/Rain-kl/Wavelet/pkg/message_gateway"
"gorm.io/gorm"
)
@@ -42,7 +41,7 @@ func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDT
if code == "" {
return BindingDTO{}, errCodeInvalid
}
pairing, err := repository.GetPairingCode(ctx, code)
pairing, err := GetPairingCode(ctx, code)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
@@ -55,7 +54,7 @@ func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDT
if pairing.ChannelID != channelID {
return BindingDTO{}, errChannelMismatch
}
ch, err := repository.GetMessageChannel(ctx, channelID)
ch, err := GetMessageChannel(ctx, channelID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, errCodeInvalid
@@ -66,7 +65,7 @@ func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDT
return BindingDTO{}, errChannelDisabled
}
existing, err := repository.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
existing, err := GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID)
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return BindingDTO{}, err
}
@@ -74,19 +73,19 @@ func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDT
if existing.UserID != userID {
return BindingDTO{}, errPlatformAlreadyBound
}
_ = repository.DeletePairingCode(ctx, pairing.Code)
_ = DeletePairingCode(ctx, pairing.Code)
return toBindingDTO(existing, ch), nil
}
row := &model.MessageBinding{
row := &MessageBinding{
UserID: userID,
ChannelID: channelID,
PlatformUserID: pairing.PlatformUserID,
}
if err := repository.CreateMessageBinding(ctx, row); err != nil {
if err := CreateMessageBinding(ctx, row); err != nil {
return BindingDTO{}, err
}
if err := repository.DeletePairingCode(ctx, pairing.Code); err != nil {
if err := DeletePairingCode(ctx, pairing.Code); err != nil {
return BindingDTO{}, err
}
return toBindingDTO(row, ch), nil
@@ -100,7 +99,7 @@ type PublicChannelDTO struct {
}
func listEnabledPublicChannels(ctx context.Context) ([]PublicChannelDTO, error) {
rows, err := repository.ListEnabledMessageChannels(ctx)
rows, err := ListEnabledMessageChannels(ctx)
if err != nil {
return nil, err
}
@@ -112,13 +111,13 @@ func listEnabledPublicChannels(ctx context.Context) ([]PublicChannelDTO, error)
}
func listUserBindings(ctx context.Context, userID uint64) ([]BindingDTO, error) {
rows, err := repository.ListBindingsByUser(ctx, userID)
rows, err := ListBindingsByUser(ctx, userID)
if err != nil {
return nil, err
}
out := make([]BindingDTO, 0, len(rows))
for i := range rows {
ch, err := repository.GetMessageChannel(ctx, rows[i].ChannelID)
ch, err := GetMessageChannel(ctx, rows[i].ChannelID)
if err != nil {
out = append(out, toBindingDTO(&rows[i], nil))
continue
@@ -129,7 +128,7 @@ func listUserBindings(ctx context.Context, userID uint64) ([]BindingDTO, error)
}
func unbindChannel(ctx context.Context, userID, bindingID uint64) error {
row, err := repository.GetMessageBinding(ctx, bindingID)
row, err := GetMessageBinding(ctx, bindingID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errBindingNotFound
@@ -139,10 +138,10 @@ func unbindChannel(ctx context.Context, userID, bindingID uint64) error {
if row.UserID != userID {
return errBindingForbidden
}
return repository.DeleteMessageBinding(ctx, bindingID)
return DeleteMessageBinding(ctx, bindingID)
}
func toBindingDTO(row *model.MessageBinding, ch *model.MessageChannel) BindingDTO {
func toBindingDTO(row *MessageBinding, ch *MessageChannel) BindingDTO {
dto := BindingDTO{
ID: row.ID,
UserID: row.UserID,
@@ -0,0 +1,30 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway_test
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
)
func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
if err != nil {
t.Fatal(err)
}
second, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ZZZZ9999", time.Now().Add(15*time.Minute))
if err != nil {
t.Fatal(err)
}
if first.Code != second.Code || first.Code != "ABCD1234" {
t.Fatalf("reuse failed: %+v %+v", first, second)
}
}
+165
View File
@@ -0,0 +1,165 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"errors"
"strings"
"time"
)
// Message channel and push channel constants.
const (
MessageChannelTypeTelegram = "telegram"
MessageChannelTypeQQ = "qq"
MessageOwnerScopeSystem = "system"
TypeCustom = "custom"
TypeEmail = "email"
TypeTelegram = "telegram"
)
// MessageChannel is an admin-configured messaging adapter.
type MessageChannel struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Type string `json:"type" gorm:"size:32;not null"`
Name string `json:"name" gorm:"size:64;not null"`
OwnerScope string `json:"owner_scope" gorm:"size:32;not null;default:'system'"`
OwnerID *uint64 `json:"owner_id,omitempty"`
Credentials string `json:"credentials" gorm:"type:text;not null"`
Extra string `json:"extra" gorm:"type:text"`
Enabled bool `json:"enabled" gorm:"default:false;not null"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (MessageChannel) TableName() string {
return "w_message_channels"
}
// MessageBinding maps a platform user to a Wavelet user on one channel.
type MessageBinding struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
UserID uint64 `json:"user_id" gorm:"not null;index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (MessageBinding) TableName() string {
return "w_message_bindings"
}
// MessagePairingCode is a one-time bind code.
type MessagePairingCode struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Code string `json:"code" gorm:"size:32;uniqueIndex;not null"`
ChannelID uint64 `json:"channel_id" gorm:"not null;index"`
PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"`
UserID uint64 `json:"user_id" gorm:"not null;index"`
ExpiresAt time.Time `json:"expires_at" gorm:"not null;index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (MessagePairingCode) TableName() string {
return "w_message_pairing_codes"
}
// PushChannel 消息通道模型
type PushChannel struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:100;not null"`
Description string `json:"description" gorm:"size:255"`
Type string `json:"type" gorm:"size:50;not null;index"`
URL string `json:"url" gorm:"type:text"`
Token string `json:"token" gorm:"type:text"`
Other string `json:"other" gorm:"type:text"`
Enabled bool `json:"enabled" gorm:"index;not null;default:true"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushChannel) TableName() string {
return "w_push_channels"
}
// Validate 验证与标准化字段
func (c *PushChannel) Validate() error {
c.Name = strings.TrimSpace(c.Name)
if c.Name == "" {
return errors.New("channel name is required")
}
c.Type = strings.TrimSpace(c.Type)
if c.Type == "" {
return errors.New("channel type is required")
}
return nil
}
// PushEvent 系统通知事件模型
type PushEvent struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"`
Name string `json:"name" gorm:"size:100;not null"`
TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"`
Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"`
Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"`
Template string `json:"template" gorm:"type:text;not null"`
Enabled bool `json:"enabled" gorm:"index;not null;default:false"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushEvent) TableName() string {
return "w_push_events"
}
// Validate 验证 PushEvent 实体字段
func (e *PushEvent) Validate() error {
e.EventKey = strings.TrimSpace(e.EventKey)
if e.EventKey == "" {
return errors.New("event_key is required")
}
e.Name = strings.TrimSpace(e.Name)
if e.Name == "" {
return errors.New("name is required")
}
return nil
}
// PushHistory 推送日志/历史实体
type PushHistory struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"size:80;not null;index"`
Channel string `json:"channel" gorm:"size:50;not null;index"`
Target string `json:"target" gorm:"size:255;not null"`
Title string `json:"title" gorm:"size:255;not null"`
Content string `json:"content" gorm:"type:text;not null"`
Level string `json:"level" gorm:"size:20;not null;default:'INFO'"`
Status string `json:"status" gorm:"size:20;not null;index"`
ErrorMsg string `json:"error_msg" gorm:"type:text"`
Payload string `json:"payload" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
}
// TableName 指定 GORM 表名
func (PushHistory) TableName() string {
return "w_push_histories"
}
// PushHistoryListFilter filters push history pagination queries.
type PushHistoryListFilter struct {
EventKey string
Channel string
Status string
StartTime *time.Time
EndTime *time.Time
Page int
PageSize int
}
+2 -3
View File
@@ -10,7 +10,6 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/extpoints"
"github.com/Rain-kl/Wavelet/plugins/domain/admin"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/hibiken/asynq"
)
@@ -84,7 +83,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
}
// 3. Register Admin Message Gateway HTTP Routes
adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", auth.LoginRequired(), admin.LoginAdminRequired())
adminMgGroup := ctx.Router().Group("/api/v1/admin/message-gateway", auth.LoginRequired(), auth.LoginAdminRequired())
{
adminMgGroup.GET("/channels/definitions", ListAdminChannelDefinitions)
adminMgGroup.GET("/channels", ListAdminChannels)
@@ -95,7 +94,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
}
// 4. Register Admin Push HTTP Routes
adminPushGroup := ctx.Router().Group("/api/v1/admin/push", auth.LoginRequired(), admin.LoginAdminRequired())
adminPushGroup := ctx.Router().Group("/api/v1/admin/push", auth.LoginRequired(), auth.LoginAdminRequired())
{
events := adminPushGroup.Group("/events")
{
@@ -11,9 +11,8 @@ import (
"strings"
"sync"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -314,7 +313,7 @@ func TestPushChannel(c *gin.Context) {
url, token, other = resolveSMTPConfig(ctx, url, token, other)
}
tempChannel := model.PushChannel{
tempChannel := PushChannel{
Name: "test_temp",
URL: url,
Token: token,
@@ -9,10 +9,10 @@ import (
"errors"
"sync"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/Rain-kl/Wavelet/pkg/util"
"gorm.io/gorm"
)
@@ -102,7 +102,7 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map
body["user"] = getSystemUser(asyncCtx)
}
eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key)
eventPtr, err := GetActivePushEventByKey(asyncCtx, meta.Key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return
@@ -121,7 +121,7 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map
})
}
func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
func (t *EventTrigger) buildMessage(event *PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
var msg NotificationMessage
renderedTemplate := ""
@@ -153,7 +153,7 @@ func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata,
return msg, renderedTemplate
}
func (t *EventTrigger) parseCustomTemplate(event *model.PushEvent, templateSource string, flatBody map[string]any) (NotificationMessage, string, error) {
func (t *EventTrigger) parseCustomTemplate(event *PushEvent, templateSource string, flatBody map[string]any) (NotificationMessage, string, error) {
var msg NotificationMessage
renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody)
@@ -206,9 +206,9 @@ func (t *EventTrigger) parseDefaultTemplate(meta EventMetadata, flatBody map[str
return msg
}
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 *PushEvent, msg NotificationMessage, flatBody map[string]any) {
for _, channelName := range event.Channels {
customChannel, err := repository.GetActivePushChannelByName(ctx, channelName)
customChannel, err := GetActivePushChannelByName(ctx, channelName)
if err == nil {
t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody)
continue
@@ -217,7 +217,7 @@ func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta EventMetadata,
}
}
func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta EventMetadata, event *model.PushEvent, channel *model.PushChannel, msg NotificationMessage, flatBody map[string]any) {
func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta EventMetadata, event *PushEvent, channel *PushChannel, msg NotificationMessage, flatBody map[string]any) {
if len(event.Targets) == 0 {
t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg)
return
@@ -229,7 +229,7 @@ func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta E
}
}
func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta EventMetadata, channel *model.PushChannel, target string, msg NotificationMessage) {
func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta EventMetadata, channel *PushChannel, target string, msg NotificationMessage) {
var config pkgpush.Config
var renderedTemplate string
@@ -9,9 +9,9 @@ import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -147,7 +147,7 @@ func ListPushHistories(c *gin.Context) {
pageSize = 20
}
total, results, err := listPushHistories(c.Request.Context(), repository.PushHistoryListFilter{
total, results, err := listPushHistories(c.Request.Context(), PushHistoryListFilter{
EventKey: c.Query("event_key"),
Status: c.Query("status"),
Page: page,
+82 -79
View File
@@ -11,10 +11,10 @@ import (
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"github.com/Rain-kl/Wavelet/pkg/task"
"gorm.io/gorm"
)
@@ -26,27 +26,28 @@ type smtpConfig struct {
}
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,
}
var cfg smtpConfig
var host, port, user, pass string
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
cfg.Host = host
cfg.Port = port
cfg.Username = user
cfg.Password = pass
return cfg
}
func syncBuiltInEvents(ctx context.Context) error {
for _, meta := range GetBuiltInEvents() {
_, err := repository.GetPushEventByKey(ctx, meta.Key)
_, err := GetPushEventByKeyRecord(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{
event := PushEvent{
EventKey: meta.Key,
Name: meta.Name,
Channels: []string{},
@@ -54,7 +55,7 @@ func syncBuiltInEvents(ctx context.Context) error {
Template: defaultTemplateStr,
Enabled: false,
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
if err := CreatePushEventRecord(ctx, &event); err != nil {
return err
}
} else if err != nil {
@@ -64,22 +65,22 @@ func syncBuiltInEvents(ctx context.Context) error {
return nil
}
func listPushEvents(ctx context.Context) ([]model.PushEvent, error) {
return repository.ListPushEvents(ctx)
func listPushEvents(ctx context.Context) ([]PushEvent, error) {
return ListPushEventsRecord(ctx)
}
func createPushEvent(ctx context.Context, req CreatePushEventRequest) (model.PushEvent, error) {
func createPushEvent(ctx context.Context, req CreatePushEventRequest) (PushEvent, error) {
eventKey, eventName, defaultTemplateBytes, err := getEventInfo(req)
if err != nil {
return model.PushEvent{}, err
return PushEvent{}, err
}
count, err := repository.CountPushEventsByKey(ctx, eventKey)
count, err := CountPushEventsByKeyRecord(ctx, eventKey)
if err != nil {
return model.PushEvent{}, err
return PushEvent{}, err
}
if count > 0 {
return model.PushEvent{}, errors.New("this notification event is already configured")
return PushEvent{}, errors.New("this notification event is already configured")
}
templateStr := strings.TrimSpace(req.Template)
@@ -88,7 +89,7 @@ func createPushEvent(ctx context.Context, req CreatePushEventRequest) (model.Pus
} 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")
return PushEvent{}, errors.New("custom template is not a valid JSON format")
}
}
@@ -101,7 +102,7 @@ func createPushEvent(ctx context.Context, req CreatePushEventRequest) (model.Pus
targets = []string{}
}
event := model.PushEvent{
event := PushEvent{
EventKey: eventKey,
Name: eventName,
TaskType: req.TaskType,
@@ -111,24 +112,24 @@ func createPushEvent(ctx context.Context, req CreatePushEventRequest) (model.Pus
Enabled: req.Enabled,
}
if err := event.Validate(); err != nil {
return model.PushEvent{}, err
return PushEvent{}, err
}
if err := repository.CreatePushEvent(ctx, &event); err != nil {
return model.PushEvent{}, err
if err := CreatePushEventRecord(ctx, &event); err != nil {
return PushEvent{}, err
}
return event, nil
}
func deletePushEvent(ctx context.Context, id uint64) error {
event, err := repository.GetPushEventByID(ctx, id)
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return err
}
return repository.DeletePushEvent(ctx, &event)
return DeletePushEventRecord(ctx, &event)
}
func updatePushEvent(ctx context.Context, id uint64, req UpdatePushEventRequest) error {
event, err := repository.GetPushEventByID(ctx, id)
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return err
}
@@ -140,11 +141,11 @@ func updatePushEvent(ctx context.Context, id uint64, req UpdatePushEventRequest)
if err := event.Validate(); err != nil {
return err
}
return repository.SavePushEvent(ctx, &event)
return SavePushEventRecord(ctx, &event)
}
func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
event, err := repository.GetPushEventByID(ctx, id)
event, err := GetPushEventByIDRecord(ctx, id)
if err != nil {
return false, err
}
@@ -153,14 +154,14 @@ func togglePushEvent(ctx context.Context, id uint64) (bool, error) {
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 {
if err := UpdatePushEventEnabledRecord(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 listPushHistories(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
return ListPushHistoriesRecord(ctx, filter)
}
func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) {
@@ -180,20 +181,20 @@ func applySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) {
cfg.Secret = smtp.Password
}
func listPushChannels(ctx context.Context) ([]model.PushChannel, error) {
return repository.ListPushChannels(ctx)
func listPushChannels(ctx context.Context) ([]PushChannel, error) {
return ListPushChannelsRecord(ctx)
}
func createPushChannel(ctx context.Context, req CreatePushChannelRequest) (model.PushChannel, error) {
count, err := repository.CountPushChannelsByName(ctx, req.Name)
func createPushChannel(ctx context.Context, req CreatePushChannelRequest) (PushChannel, error) {
count, err := CountPushChannelsByNameRecord(ctx, req.Name)
if err != nil {
return model.PushChannel{}, err
return PushChannel{}, err
}
if count > 0 {
return model.PushChannel{}, errors.New("channel name already exists")
return PushChannel{}, errors.New("channel name already exists")
}
channel := model.PushChannel{
channel := PushChannel{
Name: req.Name,
Description: req.Description,
Type: req.Type,
@@ -203,18 +204,18 @@ func createPushChannel(ctx context.Context, req CreatePushChannelRequest) (model
Enabled: req.Enabled,
}
if err := channel.Validate(); err != nil {
return model.PushChannel{}, err
return PushChannel{}, err
}
if err := repository.CreatePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
if err := CreatePushChannelRecord(ctx, &channel); err != nil {
return PushChannel{}, err
}
return channel, nil
}
func updatePushChannel(ctx context.Context, id uint64, req UpdatePushChannelRequest) (model.PushChannel, error) {
channel, err := repository.GetPushChannelByID(ctx, id)
func updatePushChannel(ctx context.Context, id uint64, req UpdatePushChannelRequest) (PushChannel, error) {
channel, err := GetPushChannelByIDRecord(ctx, id)
if err != nil {
return model.PushChannel{}, err
return PushChannel{}, err
}
channel.Description = req.Description
@@ -224,25 +225,25 @@ func updatePushChannel(ctx context.Context, id uint64, req UpdatePushChannelRequ
channel.Other = req.Other
channel.Enabled = req.Enabled
if err := channel.Validate(); err != nil {
return model.PushChannel{}, err
return PushChannel{}, err
}
if err := repository.SavePushChannel(ctx, &channel); err != nil {
return model.PushChannel{}, err
if err := SavePushChannelRecord(ctx, &channel); err != nil {
return PushChannel{}, err
}
return channel, nil
}
func deletePushChannel(ctx context.Context, id uint64) error {
channel, err := repository.GetPushChannelByID(ctx, id)
channel, err := GetPushChannelByIDRecord(ctx, id)
if err != nil {
return err
}
return repository.DeletePushChannel(ctx, &channel)
return DeletePushChannelRecord(ctx, &channel)
}
func loadChannelForTest(ctx context.Context, req TestPushChannelRequest) (string, string, string, string, error) {
if req.Name != "" {
channel, err := repository.GetPushChannelByName(ctx, req.Name)
channel, err := GetPushChannelByNameRecord(ctx, req.Name)
if err != nil {
return "", "", "", "", errors.New("channel not found")
}
@@ -251,8 +252,8 @@ func loadChannelForTest(ctx context.Context, req TestPushChannelRequest) (string
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 listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]PushEvent, error) {
return ListActivePushEventsByTaskTypeRecord(ctx, taskType)
}
func loadUserFromPayload(ctx context.Context, data map[string]any) any {
@@ -261,13 +262,15 @@ func loadUserFromPayload(ctx context.Context, data map[string]any) any {
}
if userID, ok := extractUserID(data); ok && userID > 0 {
if user, err := repository.GetUserByID(ctx, userID); err == nil {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err == nil {
return &user
}
}
if username := extractUsername(data); username != "" {
if user, err := repository.GetUserByUsername(ctx, username); err == nil {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
return &user
}
}
@@ -299,7 +302,7 @@ func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg stri
}
}
history := model.PushHistory{
history := PushHistory{
EventKey: req.EventKey,
Channel: req.Config.Channel,
Target: target,
@@ -309,7 +312,7 @@ func recordPushHistory(ctx context.Context, req SendPayload, status, errMsg stri
Status: status,
ErrorMsg: errMsg,
}
return repository.CreatePushHistory(ctx, &history)
return CreatePushHistoryRecord(ctx, &history)
}
func resolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string {
@@ -366,31 +369,25 @@ func resolveDynamicKeyword(target string, flatBody map[string]any) string {
return target
}
func resolveTargetUser(ctx context.Context, resolved string, _ string) (model.User, bool) {
found := false
var user model.User
func resolveTargetUser(ctx context.Context, resolved string, _ string) (contracts.UserDTO, bool) {
var user contracts.UserDTO
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if u, err := repository.GetUserByID(ctx, id); err == nil {
user = u
found = true
if err := db.DB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
return user, true
}
}
if !found {
if u, err := repository.GetUserByUsername(ctx, resolved); err == nil {
user = u
found = true
}
if err := db.DB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
return user, true
}
return user, found
return user, false
}
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 {
var adminUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
return resolved, true
}
if channel == channelEmail && adminUser.Email != "" {
@@ -426,9 +423,15 @@ func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, s
return url, token, other
}
func getSystemUser(ctx context.Context) *model.User {
user := repository.GetSystemUser(ctx)
return &user
func getSystemUser(ctx context.Context) *contracts.UserDTO {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
return &user
}
return &contracts.UserDTO{
Username: "system",
Nickname: "系统管理员",
}
}
func findBuiltInEvent(key string) (EventMetadata, bool) {
@@ -9,9 +9,8 @@ import (
"strconv"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/task"
)
// RegisterTaskListeners subscribes push notification handlers to task completion events.
@@ -19,7 +18,7 @@ func RegisterTaskListeners() {
task.OnTaskCompleted(handleTaskCompleted)
}
func handleTaskCompleted(ctx context.Context, execution *model.TaskExecution, result *task.TaskResult, execErr error) {
func handleTaskCompleted(ctx context.Context, execution *task.TaskExecution, result *task.TaskResult, execErr error) {
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
if err != nil {
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
+1 -1
View File
@@ -9,8 +9,8 @@ import (
"errors"
"fmt"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/pkg/push"
"github.com/Rain-kl/Wavelet/pkg/task"
)
const (
@@ -0,0 +1,392 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"errors"
"time"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
)
const (
activePushChannelCacheTTL = 24 * time.Hour
activePushEventCacheTTL = 24 * time.Hour
)
// CreateMessageChannel inserts a channel row.
func CreateMessageChannel(ctx context.Context, ch *MessageChannel) error {
if ch.ID == 0 {
ch.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(ch).Error
}
// UpdateMessageChannel saves a channel row.
func UpdateMessageChannel(ctx context.Context, ch *MessageChannel) error {
return db.DB(ctx).Save(ch).Error
}
// GetMessageChannel loads a channel by id.
func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error) {
var ch MessageChannel
if err := db.DB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
return nil, err
}
return &ch, nil
}
// ListMessageChannels returns all channels newest first.
func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := db.DB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// DeleteMessageChannel removes pairings, bindings, then the channel.
func DeleteMessageChannel(ctx context.Context, id uint64) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("channel_id = ?", id).Delete(&MessagePairingCode{}).Error; err != nil {
return err
}
if err := tx.Where("channel_id = ?", id).Delete(&MessageBinding{}).Error; err != nil {
return err
}
return tx.Delete(&MessageChannel{}, id).Error
})
}
// CreateMessageBinding inserts a binding.
func CreateMessageBinding(ctx context.Context, b *MessageBinding) error {
if b.ID == 0 {
b.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(b).Error
}
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*MessageBinding, error) {
var b MessageBinding
err := db.DB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
if err != nil {
return nil, err
}
return &b, nil
}
// ListBindingsByUser lists bindings for a Wavelet user.
func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, error) {
var rows []MessageBinding
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// GetMessageBinding loads a binding by id.
func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error) {
var b MessageBinding
if err := db.DB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
return nil, err
}
return &b, nil
}
// DeleteMessageBinding deletes a binding by id.
func DeleteMessageBinding(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&MessageBinding{}, id).Error
}
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*MessagePairingCode, error) {
var existing MessagePairingCode
err := db.DB(ctx).
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
First(&existing).Error
if err == nil {
return &existing, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
row := &MessagePairingCode{
Code: code,
ChannelID: channelID,
PlatformUserID: platformUserID,
ExpiresAt: expiresAt,
}
if err := db.DB(ctx).Create(row).Error; err != nil {
return nil, err
}
return row, nil
}
// GetPairingCode loads a pairing code by normalized code string.
func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, error) {
var row MessagePairingCode
if err := db.DB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
return nil, err
}
return &row, nil
}
// DeletePairingCode removes a pairing code.
func DeletePairingCode(ctx context.Context, code string) error {
return db.DB(ctx).Where("code = ?", code).Delete(&MessagePairingCode{}).Error
}
// DeleteExpiredPairingCodes removes expired pairing rows.
func DeleteExpiredPairingCodes(ctx context.Context) error {
return db.DB(ctx).Where("expires_at <= ?", time.Now()).Delete(&MessagePairingCode{}).Error
}
// ListEnabledMessageChannels returns enabled channels.
func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
var rows []MessageChannel
if err := db.DB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
var channels []PushChannel
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
return nil, err
}
return channels, nil
}
// GetPushChannelByIDRecord loads a push channel by primary key.
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, error) {
var channel PushChannel
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
return PushChannel{}, err
}
return channel, nil
}
// GetPushChannelByNameRecord 根据名称获取消息通道。
func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel, error) {
var channel PushChannel
if err := db.DB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
return nil, err
}
return &channel, nil
}
// CountPushChannelsByNameRecord returns how many channels share the given name.
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushChannelRecord persists a new channel and invalidates cache.
func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := db.DB(ctx).Create(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// SavePushChannelRecord updates a channel and invalidates cache.
func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
if err := db.DB(ctx).Save(channel).Error; err != nil {
return err
}
DeleteActivePushChannelCache(ctx, channel.Name)
return nil
}
// DeletePushChannelRecord removes a channel and invalidates cache.
func DeletePushChannelRecord(ctx context.Context, channel *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) (*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
}
}
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()
}
}
// ListPushEventsRecord returns all push events ordered by creation time descending.
func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
var events []PushEvent
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
return nil, err
}
return events, nil
}
// GetPushEventByIDRecord loads a push event by primary key.
func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
var event PushEvent
if err := db.DB(ctx).First(&event, id).Error; err != nil {
return PushEvent{}, err
}
return event, nil
}
// GetPushEventByKeyRecord loads a push event by event key.
func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error) {
var event PushEvent
if err := db.DB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
return PushEvent{}, err
}
return event, nil
}
// CountPushEventsByKeyRecord returns how many events use the given event key.
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreatePushEventRecord persists a new push event and invalidates cache.
func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := db.DB(ctx).Create(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// SavePushEventRecord updates a push event and invalidates cache.
func SavePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := db.DB(ctx).Save(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
func UpdatePushEventEnabledRecord(ctx context.Context, event *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
}
// DeletePushEventRecord removes a push event and invalidates cache.
func DeletePushEventRecord(ctx context.Context, event *PushEvent) error {
if err := db.DB(ctx).Delete(event).Error; err != nil {
return err
}
DeleteActivePushEventCache(ctx, event.EventKey)
return nil
}
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]PushEvent, error) {
var events []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) (*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
}
}
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()
}
}
// ListPushHistoriesRecord returns paginated push history records.
func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
query := db.DB(ctx).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 []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
}
// CreatePushHistoryRecord persists a push history audit record.
func CreatePushHistoryRecord(ctx context.Context, history *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(&PushHistory{})
}
+1 -1
View File
@@ -8,7 +8,7 @@ import (
"encoding/hex"
"encoding/json"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/util"
)
+11 -13
View File
@@ -8,16 +8,15 @@ import (
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter"
"github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/platform/lifecycle"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/persistence/batchwriter"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
)
var (
logWriterMu sync.RWMutex
logWriter *batchwriter.Writer[*analytics.UserAccessLog]
logWriter *batchwriter.Writer[*logstore.UserAccessLog]
)
// InitLogWriter initializes the access-log batch writer for the active log database.
@@ -29,8 +28,8 @@ func InitLogWriter(ctx context.Context) {
}
cfg := batchwriter.DefaultConfig()
writer, err := batchwriter.New[*analytics.UserAccessLog](cfg, func(ctx context.Context, items []*analytics.UserAccessLog) error {
rows := make([]analytics.UserAccessLog, 0, len(items))
writer, err := batchwriter.New[*logstore.UserAccessLog](cfg, func(ctx context.Context, items []*logstore.UserAccessLog) error {
rows := make([]logstore.UserAccessLog, 0, len(items))
for _, item := range items {
if item == nil {
continue
@@ -43,14 +42,14 @@ func InitLogWriter(ctx context.Context) {
}
return store.UserAccessLogs.BatchInsert(ctx, rows)
},
batchwriter.WithDropHandler[*analytics.UserAccessLog](func(item *analytics.UserAccessLog) {
batchwriter.WithDropHandler[*logstore.UserAccessLog](func(item *logstore.UserAccessLog) {
path := ""
if item != nil {
path = item.Path
}
logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", path)
}),
batchwriter.WithFlushErrorHandler[*analytics.UserAccessLog](func(ctx context.Context, items []*analytics.UserAccessLog, err error) {
batchwriter.WithFlushErrorHandler[*logstore.UserAccessLog](func(ctx context.Context, items []*logstore.UserAccessLog, err error) {
logger.ErrorF(ctx, "[RiskControl] flush access-log batch failed (batch=%d): %v", len(items), err)
}),
)
@@ -61,7 +60,6 @@ func InitLogWriter(ctx context.Context) {
writer.Start(ctx)
logWriter = writer
lifecycle.OnShutdown("risk_control_log_writer", StopLogWriter)
}
// StopLogWriter stops the ClickHouse access-log batch writer and drains pending logs.
@@ -83,7 +81,7 @@ func IsBufferFull() bool {
}
// QueueAccessLog enqueues an access log without blocking.
func QueueAccessLog(logItem *analytics.UserAccessLog) {
func QueueAccessLog(logItem *logstore.UserAccessLog) {
writer := currentLogWriter()
if writer == nil || logItem == nil {
return
@@ -92,7 +90,7 @@ func QueueAccessLog(logItem *analytics.UserAccessLog) {
}
// SetLogWriterForTest swaps the access-log writer for unit tests.
func SetLogWriterForTest(writer *batchwriter.Writer[*analytics.UserAccessLog]) func() {
func SetLogWriterForTest(writer *batchwriter.Writer[*logstore.UserAccessLog]) func() {
logWriterMu.Lock()
previous := logWriter
logWriter = writer
@@ -104,7 +102,7 @@ func SetLogWriterForTest(writer *batchwriter.Writer[*analytics.UserAccessLog]) f
}
}
func currentLogWriter() *batchwriter.Writer[*analytics.UserAccessLog] {
func currentLogWriter() *batchwriter.Writer[*logstore.UserAccessLog] {
logWriterMu.RLock()
defer logWriterMu.RUnlock()
return logWriter
+7 -7
View File
@@ -9,11 +9,11 @@ import (
"net/http"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-gonic/gin"
)
@@ -42,7 +42,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
c.Next()
// 3. 后置身份检查:仅记录通过认证的请求
userObj, exists := auth.GetFromContext[*model.User](c, auth.UserObjKey)
userObj, exists := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
if !exists || userObj == nil {
return
}
@@ -72,7 +72,7 @@ func RiskControlMiddleware() gin.HandlerFunc {
status = maxHTTPStatus
}
logItem := &analytics.UserAccessLog{
logItem := &logstore.UserAccessLog{
ID: idgen.NextUint64ID(),
UserID: userObj.ID, // 直接从 Context 获取已登录用户ID,避免数据库查询
Path: c.Request.URL.Path,
+41 -16
View File
@@ -12,25 +12,25 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/persistence/batchwriter"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
)
func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.Writer[*analytics.UserAccessLog], func() []*analytics.UserAccessLog) {
func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.Writer[*logstore.UserAccessLog], func() []*logstore.UserAccessLog) {
t.Helper()
var (
mu sync.Mutex
captured []*analytics.UserAccessLog
captured []*logstore.UserAccessLog
)
writer, err := batchwriter.New(cfg, func(_ context.Context, items []*analytics.UserAccessLog) error {
writer, err := batchwriter.New(cfg, func(_ context.Context, items []*logstore.UserAccessLog) error {
mu.Lock()
captured = append(captured, items...)
mu.Unlock()
@@ -49,14 +49,14 @@ func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.
_ = writer.Stop(stopCtx)
})
return writer, func() []*analytics.UserAccessLog {
return writer, func() []*logstore.UserAccessLog {
mu.Lock()
defer mu.Unlock()
return append([]*analytics.UserAccessLog(nil), captured...)
return append([]*logstore.UserAccessLog(nil), captured...)
}
}
func drainAccessLogWriter(t *testing.T, writer *batchwriter.Writer[*analytics.UserAccessLog]) {
func drainAccessLogWriter(t *testing.T, writer *batchwriter.Writer[*logstore.UserAccessLog]) {
t.Helper()
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
@@ -98,7 +98,7 @@ func TestRiskControlMiddleware(t *testing.T) {
r := gin.New()
r.Use(func(c *gin.Context) {
user := &model.User{ID: 12345}
user := &contracts.UserDTO{ID: 12345}
auth.SetToContext(c, auth.UserObjKey, user)
c.Next()
})
@@ -167,13 +167,38 @@ func TestRiskControlMiddleware(t *testing.T) {
cfg := batchwriter.DefaultConfig()
cfg.QueueSize = 2
cfg.MaxBatchSize = 100
cfg.MaxBatchSize = 1
cfg.FlushInterval = time.Hour
writer, _ := newTestAccessLogWriter(t, cfg)
blockCh := make(chan struct{})
enteredCh := make(chan struct{})
writer, err := batchwriter.New(cfg, func(_ context.Context, items []*logstore.UserAccessLog) error {
select {
case enteredCh <- struct{}{}:
default:
}
<-blockCh
return nil
})
assert.NoError(t, err)
writer.Start(context.Background())
restore := risk_control.SetLogWriterForTest(writer)
defer func() {
close(blockCh)
restore()
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
}()
// 1. 推入 1 个 item,worker 立即取走并触发 flush(),阻塞在 <-blockCh
writer.TryEnqueue(&logstore.UserAccessLog{})
<-enteredCh
// 2. 此时 worker 卡在 flush(),无法从 channel 取数据,推入 2 个 item 填满 channel
for range cfg.QueueSize {
writer.TryEnqueue(&analytics.UserAccessLog{})
writer.TryEnqueue(&logstore.UserAccessLog{})
}
if !risk_control.IsBufferFull() {
t.Fatal("IsBufferFull() = false, want true")
@@ -191,7 +216,7 @@ func TestRiskControlMiddleware(t *testing.T) {
assert.Equal(t, http.StatusTooManyRequests, w.Code)
var resp map[string]interface{}
err := json.Unmarshal(w.Body.Bytes(), &resp)
err = json.Unmarshal(w.Body.Bytes(), &resp)
assert.NoError(t, err)
assert.Contains(t, resp["error_msg"], "系统繁忙")
})
+43
View File
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package risk_control
import (
"fmt"
"time"
)
const (
userAccessLogTableName = "w_user_access_logs"
userAccessLogInsertColumns = "id, user_id, path, method, ip, user_agent, headers, status, latency, created_at"
)
// UserAccessLog stores HTTP access records in ClickHouse/database.
type UserAccessLog struct {
ID uint64 `gorm:"column:id"`
UserID uint64 `gorm:"column:user_id"`
Path string `gorm:"column:path"`
Method string `gorm:"column:method"`
IP string `gorm:"column:ip"`
UserAgent string `gorm:"column:user_agent"`
Headers string `gorm:"column:headers"`
Status int32 `gorm:"column:status"`
Latency int64 `gorm:"column:latency"`
CreatedAt time.Time `gorm:"column:created_at"`
}
// TableName returns the table name.
func (UserAccessLog) TableName() string {
return userAccessLogTableName
}
// InsertColumns returns comma-separated column names for batch insert.
func (UserAccessLog) InsertColumns() string {
return userAccessLogInsertColumns
}
// BatchInsertSQL returns the INSERT prefix used by native batch writers.
func (UserAccessLog) BatchInsertSQL() string {
return fmt.Sprintf("INSERT INTO %s (%s)", userAccessLogTableName, userAccessLogInsertColumns)
}
+8 -7
View File
@@ -8,9 +8,9 @@ import (
"net/http"
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/config"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
@@ -49,11 +49,12 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 2. Public config
ctx.Router().GET("/api/v1/config/public", func(c *gin.Context) {
configs, err := repository.ListVisibleSystemConfigs(c.Request.Context())
if err != nil {
response.AbortInternal(c, "获取公开配置失败")
return
type configItem struct {
Key string `json:"key"`
Value string `json:"value"`
}
var configs []configItem
_ = db.DB(c.Request.Context()).Table("w_system_configs").Where("visibility = ?", "visible").Find(&configs).Error
c.JSON(http.StatusOK, response.OK(gin.H{
"configs": configs,
"app": gin.H{
+7 -7
View File
@@ -11,13 +11,11 @@ import (
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
)
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
@@ -53,12 +51,13 @@ func ensureAccessCacheListener() {
}
func startAccessCacheInvalidationListener() {
if db.Redis == nil {
rdb := db.Redis
if rdb == nil {
return
}
util.Go(func() {
pubsub := db.Redis.Subscribe(
pubsub := rdb.Subscribe(
context.Background(),
objectstore.ConfigInvalidationChannel,
fileAccessInvalidationChannel,
@@ -114,7 +113,8 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
}
func parseFileAccessWhitelist(ctx context.Context) []string {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyFileAccessWhitelist)
var sc struct{ Value string }
err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
if err != nil || sc.Value == "" {
return []string{shared.DefaultPublicUploadType}
}
+3 -15
View File
@@ -8,10 +8,7 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"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/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
)
@@ -60,18 +57,9 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
t.Fatal("expected seeded avatar whitelist before reset")
}
var sc model.SystemConfig
if err := dbConn.Where("key = ?", model.ConfigKeyFileAccessWhitelist).First(&sc).Error; err != nil {
t.Fatalf("load whitelist config: %v", err)
if err := dbConn.Table("w_system_configs").Where("key = ?", "file_access_whitelist").Update("value", `["attachment"]`).Error; err != nil {
t.Fatalf("update whitelist config: %v", err)
}
sc.Value = `["attachment"]`
if err := dbConn.Save(&sc).Error; err != nil {
t.Fatalf("save whitelist config: %v", err)
}
if err := db.HSetJSON(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil {
t.Fatalf("refresh whitelist redis cache: %v", err)
}
repository.ResetSystemConfigRAMCacheForTest()
ResetAccessCaches()
if !IsFilePublic(ctx, "attachment") {
+23 -23
View File
@@ -9,10 +9,10 @@ import (
"fmt"
"sync"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
)
const (
@@ -26,7 +26,7 @@ type uploadMetaInvalidationMessage struct {
}
var (
uploadMetaRAM = ram.MustNew[uint64, model.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
uploadMetaListenerOnce sync.Once
uploadMetaListenerCtx context.Context
uploadMetaListenerCancel context.CancelFunc
@@ -37,8 +37,8 @@ func uploadMetaRedisKey(id uint64) string {
return fmt.Sprintf("upload:meta:%d", id)
}
func cloneUpload(upload model.Upload) model.Upload {
return upload
func cloneUpload(u models.Upload) models.Upload {
return u
}
func ensureUploadMetaCacheListener() {
@@ -88,45 +88,45 @@ func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
}
// GetUploadByID loads upload metadata from RAM, Redis, or the database.
func GetUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
func GetUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
ensureUploadMetaCacheListener()
if upload, ok := uploadMetaRAM.GetIfPresent(id); ok {
return cloneUpload(upload), nil
if u, ok := uploadMetaRAM.GetIfPresent(id); ok {
return cloneUpload(u), nil
}
key := uploadMetaRedisKey(id)
if db.Redis != nil {
var upload model.Upload
if err := db.GetJSON(ctx, key, &upload); err == nil {
uploadMetaRAM.Set(id, cloneUpload(upload))
return upload, nil
var u models.Upload
if err := db.GetJSON(ctx, key, &u); err == nil {
uploadMetaRAM.Set(id, cloneUpload(u))
return u, nil
}
}
var upload model.Upload
var u models.Upload
if err := db.DB(ctx).
Where("id = ? AND status IN (?, ?)", id, model.UploadStatusPending, model.UploadStatusUsed).
First(&upload).Error; err != nil {
return model.Upload{}, err
Where("id = ? AND status IN (?, ?)", id, models.UploadStatusPending, models.UploadStatusUsed).
First(&u).Error; err != nil {
return models.Upload{}, err
}
SetUploadMetaCache(ctx, &upload)
return upload, nil
SetUploadMetaCache(ctx, &u)
return u, nil
}
// SetUploadMetaCache populates RAM and Redis upload metadata caches.
func SetUploadMetaCache(ctx context.Context, upload *model.Upload) {
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
ensureUploadMetaCacheListener()
if upload == nil {
if u == nil {
return
}
cloned := cloneUpload(*upload)
uploadMetaRAM.Set(upload.ID, cloned)
cloned := cloneUpload(*u)
uploadMetaRAM.Set(u.ID, cloned)
if db.Redis != nil {
_ = db.SetJSON(ctx, uploadMetaRedisKey(upload.ID), cloned, uploadMetaRedisCacheTTL)
_ = db.SetJSON(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL)
}
}
+22 -22
View File
@@ -9,9 +9,9 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"gorm.io/gorm"
)
@@ -22,7 +22,7 @@ func init() {
})
}
func seedUpload(t *testing.T, dbConn *gorm.DB, upload model.Upload) {
func seedUpload(t *testing.T, dbConn *gorm.DB, upload models.Upload) {
t.Helper()
if err := dbConn.Create(&upload).Error; err != nil {
t.Fatalf("create upload: %v", err)
@@ -35,7 +35,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
ResetUploadMetaCacheForTest()
ctx := context.Background()
upload := model.Upload{
upload := models.Upload{
ID: 91001,
UserID: 1,
FileName: "cached.png",
@@ -44,7 +44,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
AccessMode: 1,
}
seedUpload(t, dbConn, upload)
@@ -57,7 +57,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
t.Fatalf("unexpected upload: %+v", got)
}
var redisUpload model.Upload
var redisUpload models.Upload
if err := db.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err != nil {
t.Fatalf("redis cache miss after DB load: %v", err)
}
@@ -65,7 +65,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
t.Fatalf("redis upload id mismatch: got=%d want=%d", redisUpload.ID, upload.ID)
}
if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil {
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
t.Fatalf("delete upload from db: %v", err)
}
@@ -84,7 +84,7 @@ func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
ResetUploadMetaCacheForTest()
ctx := context.Background()
upload := model.Upload{
upload := models.Upload{
ID: 91002,
UserID: 1,
FileName: "redis.png",
@@ -93,14 +93,14 @@ func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusPending,
Status: models.UploadStatusPending,
AccessMode: 0,
}
seedUpload(t, dbConn, upload)
SetUploadMetaCache(ctx, &upload)
ResetUploadMetaCacheForTest()
if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil {
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
t.Fatalf("delete upload from db: %v", err)
}
@@ -119,7 +119,7 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
ResetUploadMetaCacheForTest()
ctx := context.Background()
upload := model.Upload{
upload := models.Upload{
ID: 91003,
UserID: 1,
FileName: "invalidate.png",
@@ -128,7 +128,7 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
AccessMode: 1,
}
seedUpload(t, dbConn, upload)
@@ -136,7 +136,7 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
InvalidateUploadMetaCache(ctx, upload.ID)
var redisUpload model.Upload
var redisUpload models.Upload
if err := db.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err == nil {
t.Fatal("expected redis cache to be invalidated")
}
@@ -159,7 +159,7 @@ func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
ResetUploadMetaCacheForTest()
ctx := context.Background()
upload := model.Upload{
upload := models.Upload{
ID: 91006,
UserID: 1,
FileName: "pubsub.png",
@@ -168,7 +168,7 @@ func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
AccessMode: 1,
}
seedUpload(t, dbConn, upload)
@@ -177,7 +177,7 @@ func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
t.Fatalf("GetUploadByID: %v", err)
}
time.Sleep(50 * time.Millisecond) // allow pub/sub listener to subscribe
if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil {
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
t.Fatalf("delete upload from db: %v", err)
}
if _, err := GetUploadByID(ctx, upload.ID); err != nil {
@@ -219,7 +219,7 @@ func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
ResetUploadMetaCacheForTest()
ctx := context.Background()
upload := model.Upload{
upload := models.Upload{
ID: 91004,
UserID: 1,
FileName: "deleted.png",
@@ -228,7 +228,7 @@ func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusDeleted,
Status: models.UploadStatusDeleted,
AccessMode: 1,
}
seedUpload(t, dbConn, upload)
@@ -251,7 +251,7 @@ func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) {
})
ctx := context.Background()
upload := model.Upload{
upload := models.Upload{
ID: 91005,
UserID: 1,
FileName: "ram-only.png",
@@ -260,7 +260,7 @@ func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
AccessMode: 1,
}
seedUpload(t, dbConn, upload)
@@ -273,7 +273,7 @@ func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) {
t.Fatalf("unexpected upload: %+v", got)
}
if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil {
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
t.Fatalf("delete upload from db: %v", err)
}
+1 -1
View File
@@ -4,7 +4,7 @@
package upload
import (
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/handler"
+34 -25
View File
@@ -15,15 +15,16 @@ import (
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
appshared "github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response"
appshared "github.com/Rain-kl/Wavelet/pkg/shared"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/util"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/diskcache"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
@@ -85,7 +86,7 @@ func ServeFileByID(c *gin.Context) {
}
// GetUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。
func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) {
func GetUploadRecordByID(c *gin.Context) (*models.Upload, error) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("Content-Security-Policy", "sandbox")
@@ -103,7 +104,7 @@ func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) {
return &upload, nil
}
func getFileTypeCategory(upload *model.Upload) fileTypeCategory {
func getFileTypeCategory(upload *models.Upload) fileTypeCategory {
mime := strings.ToLower(upload.MimeType)
ext := strings.ToLower(upload.Extension)
@@ -120,7 +121,7 @@ func getFileTypeCategory(upload *model.Upload) fileTypeCategory {
}
// ServeUpload 将已存在的文件内容读取并流式响应给客户端。
func ServeUpload(c *gin.Context, upload *model.Upload) {
func ServeUpload(c *gin.Context, upload *models.Upload) {
setCacheHeaders(c, upload)
category := getFileTypeCategory(upload)
@@ -138,7 +139,7 @@ func ServeUpload(c *gin.Context, upload *model.Upload) {
}
}
func setCacheHeaders(c *gin.Context, upload *model.Upload) {
func setCacheHeaders(c *gin.Context, upload *models.Upload) {
if cache.IsFilePublic(c.Request.Context(), upload.Type) {
c.Header("Cache-Control", "public, max-age=31536000")
} else {
@@ -146,7 +147,7 @@ func setCacheHeaders(c *gin.Context, upload *model.Upload) {
}
}
func serveOriginalWithConditionalCheck(c *gin.Context, upload *model.Upload) {
func serveOriginalWithConditionalCheck(c *gin.Context, upload *models.Upload) {
etag := fmt.Sprintf(`W/"%s"`, upload.Hash)
c.Header("ETag", etag)
@@ -158,7 +159,7 @@ func serveOriginalWithConditionalCheck(c *gin.Context, upload *model.Upload) {
serveOriginal(c, upload)
}
func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string) {
func serveCompressedImage(c *gin.Context, upload *models.Upload, quality string) {
etag := fmt.Sprintf(`W/"%s-%s"`, upload.Hash, quality)
c.Header("ETag", etag)
@@ -167,7 +168,12 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string)
return
}
webpBytes, _, err := EnsureCompressedImageCache(c.Request.Context(), upload, quality)
webpBytes, hit, err := EnsureCompressedImageCache(c.Request.Context(), upload, quality)
if hit {
c.Header("X-Cache", "HIT")
} else {
c.Header("X-Cache", "MISS")
}
if err != nil {
if len(webpBytes) > 0 {
logger.WarnF(c.Request.Context(), "failed to cache compressed image: %v", err)
@@ -185,7 +191,7 @@ func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string)
// EnsureCompressedImageCache returns cached or freshly generated WebP bytes for an upload.
func EnsureCompressedImageCache(
ctx context.Context,
upload *model.Upload,
upload *models.Upload,
quality string,
) ([]byte, bool, error) {
cacheStore := diskcache.GetGlobalCache()
@@ -211,7 +217,7 @@ func EnsureCompressedImageCache(
func generateCompressedImageCache(
ctx context.Context,
upload *model.Upload,
upload *models.Upload,
quality string,
cacheKey string,
) (compressedImageCacheResult, error) {
@@ -246,7 +252,7 @@ func generateCompressedImageCache(
}
// ImageCompressionCacheKey returns the disk cache key for a compressed upload image.
func ImageCompressionCacheKey(upload *model.Upload, quality string) string {
func ImageCompressionCacheKey(upload *models.Upload, quality string) string {
return fmt.Sprintf(
"upload_webp_v1_%d_%d_%d_%s_%s",
upload.ID,
@@ -257,7 +263,7 @@ func ImageCompressionCacheKey(upload *model.Upload, quality string) string {
)
}
func serveOriginal(c *gin.Context, upload *model.Upload) {
func serveOriginal(c *gin.Context, upload *models.Upload) {
obj, err := uploadstorage.OpenStoredObject(c.Request.Context(), upload)
if err != nil {
response.AbortNotFound(c, "文件未找到")
@@ -267,7 +273,7 @@ func serveOriginal(c *gin.Context, upload *model.Upload) {
c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil)
}
func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, error) {
func getOriginalFileBytes(ctx context.Context, upload *models.Upload) ([]byte, error) {
obj, err := uploadstorage.OpenStoredObject(ctx, upload)
if err != nil {
return nil, err
@@ -277,33 +283,36 @@ func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, er
}
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
var currUser *model.User
var err error
if u, ok := auth.GetFromContext[*model.User](c, auth.UserObjKey); ok && u != nil {
currUser = u
var currUserID uint64
var isAdmin bool
if u, ok := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey); ok && u != nil {
currUserID = u.ID
isAdmin = u.IsAdmin
} else {
currUser, err = auth.GetUserFromRequest(c)
u, err := auth.GetUserFromRequest(c)
if err != nil {
return err
}
currUserID = u.ID
isAdmin = u.IsAdmin
}
if currUser.IsAdmin {
if isAdmin {
return nil
}
if currUser.ID != ownerID {
if currUserID != ownerID {
return errors.New("forbidden: cross-user access denied")
}
return nil
}
// CheckFileAccessPermission 校验文件是否可以被当前请求访问
func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error {
func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
if upload.AccessMode == 0 {
return checkPrivateFileOwner(c, upload.UserID)
}
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
if _, ok := auth.GetFromContext[*model.User](c, auth.UserObjKey); !ok {
if _, ok := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey); !ok {
if _, err := auth.GetUserFromRequest(c); err != nil {
return err
}
+128 -128
View File
@@ -6,8 +6,9 @@ package filesrv
import (
"bytes"
"context"
"crypto/sha256"
"encoding/json"
"fmt"
"image"
"image/color"
"image/png"
@@ -17,21 +18,20 @@ import (
"path/filepath"
"testing"
"github.com/Rain-kl/Wavelet/internal/infra/diskcache"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
appshared "github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/util"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadutil "github.com/Rain-kl/Wavelet/plugins/domain/upload/util"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/diskcache"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
)
func init() {
@@ -47,29 +47,30 @@ func TestServeFileByIDAccessControl(t *testing.T) {
configureLocalStorageRoot(t, dbConn, tempDir)
// Create a user in DB
user := model.User{
user := contracts.UserDTO{
ID: 12345,
Username: "file_test_user",
IsActive: true,
}
if err := dbConn.Create(&user).Error; err != nil {
if err := dbConn.Table("w_users").Create(&user).Error; err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create an access token for this user
tokenStr := "test-secret-token-123"
tokenHash := model.HashToken(tokenStr)
tokenRecord := model.AccessToken{
UserID: user.ID,
Name: "test_token",
TokenHash: tokenHash,
tokenHash := fmt.Sprintf("%x", sha256.Sum256([]byte(tokenStr)))
tokenRecord := map[string]any{
"user_id": user.ID,
"name": "test_token",
"token_hash": tokenHash,
"masked_token": "test-***",
}
if err := dbConn.Create(&tokenRecord).Error; err != nil {
if err := dbConn.Table("w_access_tokens").Create(&tokenRecord).Error; err != nil {
t.Fatalf("failed to create token: %v", err)
}
// Create two files: one in whitelist (avatar), one not in whitelist (attachment)
avatarFile := model.Upload{
avatarFile := models.Upload{
ID: 8001,
UserID: user.ID,
FileName: "avatar.png",
@@ -78,10 +79,10 @@ func TestServeFileByIDAccessControl(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
AccessMode: 1,
}
attachmentFile := model.Upload{
attachmentFile := models.Upload{
ID: 8002,
UserID: user.ID,
FileName: "doc.pdf",
@@ -90,7 +91,7 @@ func TestServeFileByIDAccessControl(t *testing.T) {
MimeType: "application/pdf",
Extension: "pdf",
Type: "attachment",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
AccessMode: 1,
}
@@ -104,73 +105,75 @@ func TestServeFileByIDAccessControl(t *testing.T) {
dbConn.Create(&avatarFile)
dbConn.Create(&attachmentFile)
// Set up router
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(response.ErrorHandlerMiddleware())
store := cookie.NewStore([]byte("secret"))
r.Use(sessions.Sessions("test_session", store))
r.Use(sessions.Sessions("wavelet_session_id", store))
r.GET("/f/:id", ServeFileByID)
t.Run("whitelisted file type (avatar) accessed without authentication", func(t *testing.T) {
t.Run("public access allowed for whitelist type (avatar)", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/8001", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200, got %d. Body: %s", w.Code, w.Body.String())
t.Fatalf("expected status 200 for public file, got %d", w.Code)
}
if w.Body.String() != "image" {
t.Errorf("expected 'image', got %q", w.Body.String())
t.Fatalf("expected body 'image', got '%s'", w.Body.String())
}
})
t.Run("non-whitelisted file type (attachment) accessed without authentication returns 401", func(t *testing.T) {
t.Run("public access rejected for non-whitelist type (attachment)", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/8002", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected 401, got %d. Body: %s", w.Code, w.Body.String())
}
var body map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("failed to parse JSON: %v", err)
}
if body["error_msg"] != appshared.UnAuthorized {
t.Errorf("expected error_msg %q, got %v", appshared.UnAuthorized, body["error_msg"])
t.Fatalf("expected status 401 for private file without auth, got %d", w.Code)
}
})
t.Run("non-whitelisted file type (attachment) accessed with valid token succeeds", func(t *testing.T) {
t.Run("authenticated access allowed for non-whitelist type (attachment)", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/8002", nil)
req.Header.Set("X-Access-Token", tokenStr)
req.Header.Set("Authorization", "Bearer "+tokenStr)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200, got %d. Body: %s", w.Code, w.Body.String())
t.Fatalf("expected status 200 for authenticated request, got %d", w.Code)
}
if w.Body.String() != "bytes" {
t.Errorf("expected 'bytes', got %q", w.Body.String())
t.Fatalf("expected body 'bytes', got '%s'", w.Body.String())
}
})
t.Run("accessing non-existent file returns 404", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/9999", nil)
t.Run("non-existent file returns 404", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/99999", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404, got %d", w.Code)
t.Fatalf("expected status 404 for non-existent file, got %d", w.Code)
}
})
t.Run("invalid id format returns 400", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/invalid-id", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Fatalf("expected status 400 for invalid ID, got %d", w.Code)
}
})
}
func TestImageCompression(t *testing.T) {
func TestServeFileByIDImageCompression(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
cache.ResetAccessCaches()
tempDir := t.TempDir()
configureLocalStorageRoot(t, dbConn, tempDir)
@@ -187,12 +190,12 @@ func TestImageCompression(t *testing.T) {
}()
// Create test user
user := model.User{
user := contracts.UserDTO{
ID: 555,
Username: "compress_tester",
IsActive: true,
}
dbConn.Create(&user)
dbConn.Table("w_users").Create(&user)
// Create a 1x1 pixel PNG image
img := image.NewRGBA(image.Rect(0, 0, 1, 1))
@@ -208,7 +211,7 @@ func TestImageCompression(t *testing.T) {
}
// Save upload record to DB
uploadRecord := model.Upload{
uploadRecord := models.Upload{
ID: 3001,
UserID: user.ID,
FileName: "test_image.png",
@@ -217,7 +220,7 @@ func TestImageCompression(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "avatar", // Whitelisted by default
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
AccessMode: 1,
}
dbConn.Create(&uploadRecord)
@@ -239,53 +242,15 @@ func TestImageCompression(t *testing.T) {
if w.Header().Get("Content-Type") != "image/png" {
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
}
if len(w.Body.Bytes()) != pngBuf.Len() {
t.Errorf("expected body size %d, got %d", pngBuf.Len(), len(w.Body.Bytes()))
if w.Header().Get("X-Cache") != "" {
t.Errorf("expected no X-Cache header for original file, got %s", w.Header().Get("X-Cache"))
}
if w.Header().Get("ETag") == "" {
t.Errorf("expected ETag header for original file")
}
})
t.Run("serve compressed WebP file with medium quality", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
// Content-Type should be image/webp
if w.Header().Get("Content-Type") != "image/webp" {
t.Errorf("expected Content-Type image/webp, got %s", w.Header().Get("Content-Type"))
}
cacheKey := ImageCompressionCacheKey(&uploadRecord, shared.ImageQualityMedium)
cachedBytes, err := cache.Get(cacheKey)
if err != nil {
t.Fatalf("disk cache Get(%q) returned error: %v", cacheKey, err)
}
if !bytes.Equal(cachedBytes, w.Body.Bytes()) {
t.Errorf("cached compressed image differs from response")
}
if err := os.Remove(filePath); err != nil {
t.Fatalf("failed to remove source image before cache-hit request: %v", err)
}
t.Cleanup(func() {
if err := os.WriteFile(filePath, pngBuf.Bytes(), 0644); err != nil {
t.Errorf("failed to restore source image: %v", err)
}
})
w2 := httptest.NewRecorder()
r.ServeHTTP(w2, req)
if w2.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w2.Code)
}
if !bytes.Equal(w2.Body.Bytes(), cachedBytes) {
t.Errorf("cache-hit response differs from cached compressed image")
}
})
t.Run("serve compressed WebP file and check cache headers and 304 Not Modified", func(t *testing.T) {
t.Run("first request with quality=medium produces cache MISS and converts to WebP", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
@@ -293,29 +258,37 @@ func TestImageCompression(t *testing.T) {
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w.Code)
}
etag := w.Header().Get("ETag")
if etag == "" {
t.Error("expected ETag header, got empty")
if w.Header().Get("Content-Type") != "image/webp" {
t.Errorf("expected Content-Type image/webp, got %s", w.Header().Get("Content-Type"))
}
cacheControl := w.Header().Get("Cache-Control")
if cacheControl != "public, max-age=31536000" {
t.Errorf("expected Cache-Control 'public, max-age=31536000', got %q", cacheControl)
if w.Header().Get("X-Cache") != "MISS" {
t.Errorf("expected X-Cache MISS on first compress request, got %s", w.Header().Get("X-Cache"))
}
// Perform conditional GET request
reqCond, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
reqCond.Header.Set("If-None-Match", etag)
wCond := httptest.NewRecorder()
r.ServeHTTP(wCond, reqCond)
if wCond.Code != http.StatusNotModified {
t.Errorf("expected status 304, got %d", wCond.Code)
if w.Header().Get("ETag") == "" {
t.Errorf("expected ETag header")
}
if len(w.Body.Bytes()) == 0 {
t.Errorf("expected non-empty body")
}
})
t.Run("serve original file with origin quality", func(t *testing.T) {
t.Run("second request with quality=medium produces cache HIT", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w.Code)
}
if w.Header().Get("Content-Type") != "image/webp" {
t.Errorf("expected Content-Type image/webp, got %s", w.Header().Get("Content-Type"))
}
if w.Header().Get("X-Cache") != "HIT" {
t.Errorf("expected X-Cache HIT on second compress request, got %s", w.Header().Get("X-Cache"))
}
})
t.Run("request with quality=origin behaves like original request", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=origin", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
@@ -326,8 +299,32 @@ func TestImageCompression(t *testing.T) {
if w.Header().Get("Content-Type") != "image/png" {
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
}
if !bytes.Equal(w.Body.Bytes(), pngBuf.Bytes()) {
t.Errorf("origin-quality response differs from original image")
if w.Header().Get("X-Cache") != "" {
t.Errorf("expected no X-Cache header for origin quality, got %s", w.Header().Get("X-Cache"))
}
})
t.Run("conditional GET with matching If-None-Match returns 304", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
etag := w.Header().Get("ETag")
if etag == "" {
t.Fatalf("expected ETag header from initial request")
}
// Second request with If-None-Match
req2, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
req2.Header.Set("If-None-Match", etag)
w2 := httptest.NewRecorder()
r.ServeHTTP(w2, req2)
if w2.Code != http.StatusNotModified {
t.Fatalf("expected status 304 Not Modified, got %d", w2.Code)
}
if w2.Body.Len() != 0 {
t.Errorf("expected empty body on 304 response, got %d bytes", w2.Body.Len())
}
})
}
@@ -338,18 +335,20 @@ func TestNormalizeImageQuality(t *testing.T) {
quality string
want string
}{
{name: shared.ImageQualityLow, quality: shared.ImageQualityLow, want: shared.ImageQualityLow},
{name: shared.ImageQualityMedium, quality: shared.ImageQualityMedium, want: shared.ImageQualityMedium},
{name: shared.ImageQualityHigh, quality: shared.ImageQualityHigh, want: shared.ImageQualityHigh},
{name: "origin", quality: "origin", want: "origin"},
{name: "uppercase", quality: "LOW", want: shared.ImageQualityLow},
{name: "empty", quality: "", want: "origin"},
{name: "invalid", quality: "maximum", want: "origin"},
{name: "empty quality returns origin", quality: "", want: shared.ImageQualityOrigin},
{name: "origin returns origin", quality: "origin", want: shared.ImageQualityOrigin},
{name: "ORIGIN case-insensitive returns origin", quality: "ORIGIN", want: shared.ImageQualityOrigin},
{name: "low returns low", quality: "low", want: shared.ImageQualityLow},
{name: "LOW returns low", quality: "LOW", want: shared.ImageQualityLow},
{name: "medium returns medium", quality: "medium", want: shared.ImageQualityMedium},
{name: "high returns high", quality: "high", want: shared.ImageQualityHigh},
{name: "unknown quality defaults to origin", quality: "ultra_hd", want: shared.ImageQualityOrigin},
{name: "whitespace padded quality is trimmed", quality: " medium ", want: shared.ImageQualityMedium},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := util.NormalizeImageQuality(tt.quality); got != tt.want {
if got := uploadutil.NormalizeImageQuality(tt.quality); got != tt.want {
t.Errorf("NormalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want)
}
})
@@ -357,8 +356,11 @@ func TestNormalizeImageQuality(t *testing.T) {
}
func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) {
var sc model.SystemConfig
if err := dbConn.Where("key = ?", model.ConfigKeyStorageConfig).First(&sc).Error; err != nil {
var sc struct {
Key string
Value string
}
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").First(&sc).Error; err != nil {
t.Fatalf("failed to find storage config: %v", err)
}
var cfg objectstore.Config
@@ -371,10 +373,8 @@ func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) {
t.Fatalf("failed to marshal storage config: %v", err)
}
sc.Value = string(newVal)
if err := dbConn.Save(&sc).Error; err != nil {
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").Update("value", sc.Value).Error; err != nil {
t.Fatalf("failed to save storage config: %v", err)
}
_ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, sc.Key, &sc)
repository.ResetSystemConfigRAMCacheForTest()
objectstore.ResetCache()
}
@@ -7,15 +7,16 @@ import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
"github.com/gin-gonic/gin"
)
type listFilesRequest struct {
@@ -28,10 +29,10 @@ type listFilesRequest struct {
}
type listFilesResponse struct {
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
Items []model.Upload `json:"items"`
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
Items []models.Upload `json:"items"`
}
// ListFiles 获取系统上传的文件列表
@@ -44,12 +45,11 @@ type listFilesResponse struct {
// @Param keyword query string false "文件名关键词(模糊匹配)"
// @Param type query string false "业务分类过滤"
// @Param extension query string false "扩展名过滤"
// @Param user_id query uint64 false "上传用户 ID"
// @Param user_id query int false "上传用户 ID 过滤"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=listFilesResponse} "查询成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/uploads [get]
// @Failure 400 {object} response.Any "参数错误"
// @Router /api/v1/admin/uploads/files [get]
func ListFiles(c *gin.Context) {
ctx := c.Request.Context()
@@ -86,17 +86,16 @@ func ListFiles(c *gin.Context) {
}))
}
// DeleteFile 软删除文件记录
// DeleteFile 软删除指定的文件记录
// @Summary 删除文件
// @Description 将文件状态置为 deleted(软删除),不会立即清理底层存储对象
// @Description 将指定 ID 的文件状态置为 deleted(软删除)
// @Tags admin
// @Produce json
// @Param id path string true "文件 ID"
// @Security SessionCookie
// @Success 200 {object} response.Any "删除成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/admin/uploads/{id} [delete]
// @Router /api/v1/admin/uploads/files/{id} [delete]
func DeleteFile(c *gin.Context) {
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
@@ -121,21 +120,19 @@ func DeleteFile(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
// GetDistinctUploadTypes 获取数据库中所有已存在的文件业务类型
// @Summary 获取文件业务类型列表
// @Description 返回数据库中所有已上传文件实际拥有的业务类型列表
// GetDistinctUploadTypes 获取所有已存在的文件业务分类列表
// @Summary 获取业务分类列表
// @Description 查询系统内所有不重复的上传业务分类标识(如 avatar, doc 等)
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]string} "业务类型列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Success 200 {object} response.Any{data=[]string} "查询成功"
// @Router /api/v1/admin/uploads/types [get]
func GetDistinctUploadTypes(c *gin.Context) {
types, err := listDistinctUploadTypes(c.Request.Context())
ctx := c.Request.Context()
types, err := listDistinctUploadTypes(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
response.AbortBadRequest(c, shared.ErrQueryTypeListFailed)
return
}
c.JSON(http.StatusOK, response.OK(types))
@@ -150,10 +147,10 @@ type listMyFilesRequest struct {
}
type listMyFilesResponse struct {
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
Items []model.Upload `json:"items"`
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
Items []models.Upload `json:"items"`
}
// ListMyFiles 获取当前用户上传的文件列表
@@ -171,7 +168,7 @@ type listMyFilesResponse struct {
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/upload/my [get]
func ListMyFiles(c *gin.Context) {
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
ctx := c.Request.Context()
var req listMyFilesRequest
@@ -218,7 +215,7 @@ func ListMyFiles(c *gin.Context) {
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) {
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
@@ -260,12 +257,12 @@ type updateMyFileRequest struct {
// @Param id path string true "文件 ID"
// @Param request body updateMyFileRequest true "更新字段"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.Upload} "更新成功"
// @Success 200 {object} response.Any{data=models.Upload} "更新成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [put]
func UpdateMyFile(c *gin.Context) {
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
@@ -284,7 +281,7 @@ func UpdateMyFile(c *gin.Context) {
return
}
upload, err := updateOwnedUpload(ctx, currUser.ID, uploadID, updateMyUploadInput(req))
updated, err := updateOwnedUpload(ctx, currUser.ID, uploadID, updateMyUploadInput(req))
if err != nil {
if isRecordNotFound(err) {
response.AbortNotFound(c, "文件记录未找到")
@@ -294,9 +291,9 @@ func UpdateMyFile(c *gin.Context) {
response.AbortForbidden(c, "无权操作")
return
}
response.AbortBadRequest(c, "更新文件记录失败")
response.AbortBadRequest(c, shared.ErrUpdateFileFailed)
return
}
c.JSON(http.StatusOK, response.OK(upload))
c.JSON(http.StatusOK, response.OK(updated))
}
@@ -9,19 +9,21 @@ import (
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
)
func TestGetDistinctUploadTypes(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
user := model.User{ID: 2222, Username: "test_user_2"}
dbConn.Create(&user)
user := contracts.UserDTO{ID: 2222, Username: "test_user_2"}
dbConn.Table("w_users").Create(&user)
customUpload := model.Upload{
customUpload := models.Upload{
ID: 9001,
UserID: user.ID,
FileName: "custom.txt",
@@ -30,7 +32,7 @@ func TestGetDistinctUploadTypes(t *testing.T) {
MimeType: "text/plain",
Extension: "txt",
Type: "custom_type_xyz",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
}
dbConn.Create(&customUpload)
+20 -19
View File
@@ -8,26 +8,27 @@ import (
"errors"
"sort"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
)
func listUploadFiles(ctx context.Context, filter repository.UploadListFilter) (int64, []model.Upload, error) {
func listUploadFiles(ctx context.Context, filter repository.UploadListFilter) (int64, []models.Upload, error) {
return repository.ListUploads(ctx, filter)
}
func listMyUploadFiles(ctx context.Context, userID uint64, filter repository.UploadListFilter) (int64, []model.Upload, error) {
func listMyUploadFiles(ctx context.Context, userID uint64, filter repository.UploadListFilter) (int64, []models.Upload, error) {
filter.UserID = userID
return repository.ListUploads(ctx, filter)
}
func softDeleteUpload(ctx context.Context, uploadID uint64) (model.Upload, error) {
func softDeleteUpload(ctx context.Context, uploadID uint64) (models.Upload, error) {
return ingest.Remove(ctx, uploadID)
}
func softDeleteOwnedUpload(ctx context.Context, userID, uploadID uint64) (model.Upload, error) {
func softDeleteOwnedUpload(ctx context.Context, userID, uploadID uint64) (models.Upload, error) {
return ingest.RemoveOwned(ctx, userID, uploadID)
}
@@ -45,13 +46,13 @@ type updateMyUploadInput struct {
AccessMode *int
}
func updateOwnedUpload(ctx context.Context, userID, uploadID uint64, input updateMyUploadInput) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
func updateOwnedUpload(ctx context.Context, userID, uploadID uint64, input updateMyUploadInput) (models.Upload, error) {
u, err := repository.GetActiveUploadByID(ctx, uploadID)
if err != nil {
return model.Upload{}, err
return models.Upload{}, err
}
if upload.UserID != userID {
return model.Upload{}, ingest.ErrForbidden
if u.UserID != userID {
return models.Upload{}, ingest.ErrForbidden
}
updates := make(map[string]any)
@@ -61,23 +62,23 @@ func updateOwnedUpload(ctx context.Context, userID, uploadID uint64, input updat
if input.AccessMode != nil {
updates["access_mode"] = *input.AccessMode
}
if err := repository.UpdateUpload(ctx, &upload, updates); err != nil {
return model.Upload{}, err
if err := repository.UpdateUpload(ctx, &u, updates); err != nil {
return models.Upload{}, err
}
if name, ok := updates["file_name"].(string); ok {
upload.FileName = name
u.FileName = name
}
if mode, ok := updates["access_mode"].(int); ok {
upload.AccessMode = mode
u.AccessMode = mode
}
return upload, nil
return u, nil
}
func listUploadsForBatchDownload(ctx context.Context, ids []uint64) ([]model.Upload, error) {
func listUploadsForBatchDownload(ctx context.Context, ids []uint64) ([]models.Upload, error) {
return repository.ListUploadsByIDs(ctx, ids)
}
func loadUploadStats(ctx context.Context) ([]model.UploadStat, error) {
func loadUploadStats(ctx context.Context) ([]models.UploadStat, error) {
return repository.ListUploadStats(ctx)
}
+8 -7
View File
@@ -22,13 +22,14 @@ import (
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
appshared "github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
appshared "github.com/Rain-kl/Wavelet/pkg/shared"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/util"
@@ -50,7 +51,7 @@ type batchDownloadRequest struct {
// @Param type formData string false "业务分类 (例如: avatar, attachment, doc,默认为 generic)"
// @Param metadata formData string false "额外的 JSON 格式元数据"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.Upload} "上传成功"
// @Success 200 {object} response.Any{data=models.Upload} "上传成功"
// @Failure 400 {object} response.Any "请求参数错误或文件受限"
// @Failure 401 {object} response.Any "未登录"
// @Failure 500 {object} response.Any "内部错误"
@@ -63,7 +64,7 @@ func UploadFile(c *gin.Context) {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
ctx := c.Request.Context()
header, err := c.FormFile("file")
@@ -310,8 +311,8 @@ func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) {
return accessMode, ""
}
func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) {
var meta model.UploadMetadata
func parseUploadMetadata(c *gin.Context, mimeType string) (models.UploadMetadata, string) {
var meta models.UploadMetadata
metadataStr := c.DefaultPostForm("metadata", "")
if metadataStr != "" {
if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil {
+45 -51
View File
@@ -19,15 +19,14 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"github.com/gin-gonic/gin"
)
@@ -36,7 +35,7 @@ type testResponse struct {
Data json.RawMessage `json:"data"`
}
func setupTestRouter(authUser *model.User) *gin.Engine {
func setupTestRouter(authUser *contracts.UserDTO) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(response.ErrorHandlerMiddleware())
@@ -106,7 +105,7 @@ func TestUploadFile(t *testing.T) {
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests
authUser := &model.User{ID: 1001, Username: "test_user"}
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Mock Storage Client
@@ -175,12 +174,12 @@ func TestUploadFile(t *testing.T) {
}
// Verify database record
var uploadRecord model.Upload
var uploadRecord models.Upload
if err := json.Unmarshal(resp.Data, &uploadRecord); err != nil {
t.Fatalf("failed to unmarshal upload record: %v", err)
}
var dbRecord model.Upload
var dbRecord models.Upload
if err := dbConn.First(&dbRecord, uploadRecord.ID).Error; err != nil {
t.Fatalf("failed to retrieve database record: %v", err)
}
@@ -258,7 +257,7 @@ func TestUploadFile(t *testing.T) {
t.Fatalf("second upload was unsuccessful: %s", resp2.ErrorMsg)
}
var uploadRecord2 model.Upload
var uploadRecord2 models.Upload
if err := json.Unmarshal(resp2.Data, &uploadRecord2); err != nil {
t.Fatalf("failed to unmarshal second upload record: %v", err)
}
@@ -269,7 +268,7 @@ func TestUploadFile(t *testing.T) {
}
// Check if database contains both records sharing the same FilePath
var records []model.Upload
var records []models.Upload
dbConn.Where("hash = ?", uploadRecord2.Hash).Find(&records)
if len(records) != 2 {
t.Errorf("expected 2 database records sharing the same hash, got %d", len(records))
@@ -289,12 +288,7 @@ func TestUploadFile(t *testing.T) {
objectstore.IsEnabledFunc = func() bool { return false }
// Seed allowed extensions configuration to allow txt files
var sc model.SystemConfig
dbConn.Where("key = ?", model.ConfigKeyUploadAllowedExtensions).First(&sc)
sc.Value = "jpg,png,webp,txt"
dbConn.Save(&sc)
_ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, sc.Key, &sc)
repository.ResetSystemConfigRAMCacheForTest()
dbConn.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Update("value", "jpg,png,webp,txt")
contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{
"type": "document",
@@ -316,7 +310,7 @@ func TestUploadFile(t *testing.T) {
t.Fatalf("local upload failed: %s", resp.ErrorMsg)
}
var localRecord model.Upload
var localRecord models.Upload
if err := json.Unmarshal(resp.Data, &localRecord); err != nil {
t.Fatalf("failed to unmarshal local upload record: %v", err)
}
@@ -338,11 +332,11 @@ func TestDownloadFile(t *testing.T) {
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
authUser := &model.User{ID: 1001, Username: "test_user"}
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Seed upload records in DB
localUpload := model.Upload{
localUpload := models.Upload{
ID: 2001,
UserID: 1001,
FileName: "中文文件名.txt",
@@ -350,7 +344,7 @@ func TestDownloadFile(t *testing.T) {
FileSize: 12,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
}
// Create local file
@@ -405,10 +399,10 @@ func TestListFiles(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
authUser := &model.User{ID: 1001, Username: "test_user"}
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
uploads := []model.Upload{
uploads := []models.Upload{
{
ID: 2101,
UserID: authUser.ID,
@@ -417,7 +411,7 @@ func TestListFiles(t *testing.T) {
FileSize: 10,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
{
ID: 2102,
@@ -427,7 +421,7 @@ func TestListFiles(t *testing.T) {
FileSize: 20,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
{
ID: 2103,
@@ -437,7 +431,7 @@ func TestListFiles(t *testing.T) {
FileSize: 30,
MimeType: "text/markdown",
Extension: "md",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
{
ID: 2104,
@@ -447,7 +441,7 @@ func TestListFiles(t *testing.T) {
FileSize: 40,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
}
for i := range uploads {
@@ -546,7 +540,7 @@ func TestBatchDownloadFiles(t *testing.T) {
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
authUser := &model.User{ID: 1001, Username: "test_user"}
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Create and write files locally
@@ -560,7 +554,7 @@ func TestBatchDownloadFiles(t *testing.T) {
_ = os.WriteFile("uploads/f3.txt", []byte("duplicate name file content"), 0644)
// Seed upload records. Note f2 and f3 have the same FileName "file_a.txt" to trigger name collision resolution.
uploads := []model.Upload{
uploads := []models.Upload{
{
ID: 3001,
UserID: 1001,
@@ -569,7 +563,7 @@ func TestBatchDownloadFiles(t *testing.T) {
FileSize: 13,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
{
ID: 3002,
@@ -579,7 +573,7 @@ func TestBatchDownloadFiles(t *testing.T) {
FileSize: 13,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
{
ID: 3003,
@@ -589,7 +583,7 @@ func TestBatchDownloadFiles(t *testing.T) {
FileSize: 28,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
}
@@ -658,8 +652,8 @@ func TestUploadAccessModeAccessControl(t *testing.T) {
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
user1 := &model.User{ID: 1001, Username: "user1"}
user2 := &model.User{ID: 1002, Username: "user2"}
user1 := &contracts.UserDTO{ID: 1001, Username: "user1"}
user2 := &contracts.UserDTO{ID: 1002, Username: "user2"}
// Seed user1
if err := dbConn.Create(user1).Error; err != nil {
@@ -689,7 +683,7 @@ func TestUploadAccessModeAccessControl(t *testing.T) {
t.Logf("Raw upload response: %s", w.Body.String())
var resp1 testResponse
_ = json.Unmarshal(w.Body.Bytes(), &resp1)
var upload1 model.Upload
var upload1 models.Upload
_ = json.Unmarshal(resp1.Data, &upload1)
if upload1.AccessMode != 0 {
@@ -707,7 +701,7 @@ func TestUploadAccessModeAccessControl(t *testing.T) {
router.ServeHTTP(w2, req2)
var resp2 testResponse
_ = json.Unmarshal(w2.Body.Bytes(), &resp2)
var upload2 model.Upload
var upload2 models.Upload
_ = json.Unmarshal(resp2.Data, &upload2)
if upload2.AccessMode != 1 {
@@ -744,11 +738,11 @@ func TestGetFileStats(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
authUser := &model.User{ID: 1001, Username: "test_user"}
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Insert some dummy uploads
uploads := []model.Upload{
uploads := []models.Upload{
{
ID: 3101,
UserID: authUser.ID,
@@ -758,7 +752,7 @@ func TestGetFileStats(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "generic",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
},
{
@@ -770,7 +764,7 @@ func TestGetFileStats(t *testing.T) {
MimeType: "video/mp4",
Extension: "mp4",
Type: "generic",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now().AddDate(0, 0, -2), // 2 days ago
},
{
@@ -782,7 +776,7 @@ func TestGetFileStats(t *testing.T) {
MimeType: "application/pdf",
Extension: "pdf",
Type: "avatar", // different type
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now().AddDate(0, 0, -10), // older than 7 days
},
}
@@ -854,8 +848,8 @@ func TestUserUploadManagement(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
user1 := &model.User{ID: 1001, Username: "user1"}
user2 := &model.User{ID: 1002, Username: "user2"}
user1 := &contracts.UserDTO{ID: 1001, Username: "user1"}
user2 := &contracts.UserDTO{ID: 1002, Username: "user2"}
_ = dbConn.Create(user1)
_ = dbConn.Create(user2)
@@ -864,7 +858,7 @@ func TestUserUploadManagement(t *testing.T) {
router2 := setupTestRouter(user2)
// Seed upload records
upload1 := model.Upload{
upload1 := models.Upload{
ID: 4001,
UserID: 1001,
FileName: "user1-file.txt",
@@ -872,10 +866,10 @@ func TestUserUploadManagement(t *testing.T) {
FileSize: 100,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
upload2 := model.Upload{
upload2 := models.Upload{
ID: 4002,
UserID: 1002,
FileName: "user2-file.png",
@@ -883,7 +877,7 @@ func TestUserUploadManagement(t *testing.T) {
FileSize: 200,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
@@ -927,7 +921,7 @@ func TestUserUploadManagement(t *testing.T) {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
var updated model.Upload
var updated models.Upload
dbConn.First(&updated, 4001)
if updated.FileName != "renamed.txt" {
t.Errorf("expected file name renamed.txt, got %s", updated.FileName)
@@ -970,9 +964,9 @@ func TestUserUploadManagement(t *testing.T) {
t.Fatalf("expected status 200, got %d", w.Code)
}
var deleted model.Upload
var deleted models.Upload
dbConn.First(&deleted, 4001)
if deleted.Status != model.UploadStatusDeleted {
if deleted.Status != models.UploadStatusDeleted {
t.Errorf("expected status deleted, got %s", deleted.Status)
}
})
+7 -7
View File
@@ -7,10 +7,10 @@ import (
"net/http"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
)
type trendItem struct {
@@ -79,22 +79,22 @@ func GetFileStats(c *gin.Context) {
for _, stat := range stats {
switch stat.Dimension {
case model.UploadStatDimensionTotal:
case shared.UploadStatDimensionTotal:
totalCount = stat.FileCount
totalSize = stat.FileSize
case model.UploadStatDimensionType:
case shared.UploadStatDimensionType:
types = append(types, distributionItem{
Name: stat.StatKey,
Count: stat.FileCount,
Size: stat.FileSize,
})
case model.UploadStatDimensionCategory:
case shared.UploadStatDimensionCategory:
if item, ok := categoryMap[stat.StatKey]; ok {
item.Count = stat.FileCount
item.Size = stat.FileSize
categoryMap[stat.StatKey] = item
}
case model.UploadStatDimensionTrend:
case shared.UploadStatDimensionTrend:
if _, ok := trendCountMap[stat.StatKey]; ok {
trendCountMap[stat.StatKey] = stat.FileCount
trendSizeMap[stat.StatKey] = stat.FileSize
+29 -16
View File
@@ -5,22 +5,23 @@ package ingest
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"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/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"gorm.io/gorm"
)
@@ -33,7 +34,7 @@ func normalizeRequest(req *Request) {
req.Type = "generic"
}
if req.Status == "" {
req.Status = model.UploadStatusUsed
req.Status = models.UploadStatusUsed
}
}
@@ -48,18 +49,30 @@ func resolveAccessMode(uploadType string, explicit *int) int {
}
func validateAllowedExtension(ctx context.Context, ext string) error {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUploadAllowedExtensions)
var val string
err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return err
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
return nil
}
if sc.Value == "" {
if val == "" {
return nil
}
allowedExts := strings.Split(strings.ToLower(sc.Value), ",")
var list []string
if err := json.Unmarshal([]byte(val), &list); err == nil {
for _, allowedExt := range list {
if strings.EqualFold(strings.TrimSpace(allowedExt), ext) {
return nil
}
}
return errors.New(shared.ErrUnsupportedFormat)
}
allowedExts := strings.Split(strings.ToLower(val), ",")
for _, allowedExt := range allowedExts {
if strings.TrimSpace(allowedExt) == ext {
return nil
@@ -79,7 +92,7 @@ func buildObjectKey(req Request, id uint64) string {
return defaultObjectKey(id, req.Extension)
}
func storeObject(ctx context.Context, objectKey string, reader io.Reader, size int64, mimeType string, meta *model.UploadMetadata) (string, error) {
func storeObject(ctx context.Context, objectKey string, reader io.Reader, size int64, mimeType string, meta *models.UploadMetadata) (string, error) {
if uploadstorage.ReadOnly(ctx) {
return "", ErrStorageReadOnly
}
@@ -100,7 +113,7 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
return result.Key, nil
}
func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey string) error {
func persistUploadRecord(ctx context.Context, upload *models.Upload, objectKey string) error {
if err := createUploadWithStats(ctx, upload); err != nil {
_, backend, backendErr := objectstore.Active(ctx)
if backendErr == nil {
@@ -114,7 +127,7 @@ func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey st
return nil
}
func createUploadWithStats(ctx context.Context, upload *model.Upload) error {
func createUploadWithStats(ctx context.Context, upload *models.Upload) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := repository.CreateUploadTx(tx, upload); err != nil {
return err
@@ -123,9 +136,9 @@ func createUploadWithStats(ctx context.Context, upload *model.Upload) error {
})
}
func createDedupRecord(ctx context.Context, existing model.Upload, req Request) (Result, error) {
func createDedupRecord(ctx context.Context, existing models.Upload, req Request) (Result, error) {
accessMode := resolveAccessMode(req.Type, req.AccessMode)
newUpload := model.Upload{
newUpload := models.Upload{
ID: idgen.NextUint64ID(),
UserID: req.UserID,
FileName: req.FileName,
@@ -172,7 +185,7 @@ func createNewUpload(ctx context.Context, req Request) (Result, error) {
}
accessMode := resolveAccessMode(req.Type, req.AccessMode)
upload := model.Upload{
upload := models.Upload{
ID: id,
UserID: req.UserID,
FileName: req.FileName,
+3 -3
View File
@@ -7,8 +7,8 @@ import (
"context"
"errors"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
"gorm.io/gorm"
)
@@ -36,7 +36,7 @@ func Ingest(ctx context.Context, req Request) (Result, error) {
}
// FindByHash returns a reusable active upload with the same hash and size.
func FindByHash(ctx context.Context, hash string, size int64) (model.Upload, error) {
func FindByHash(ctx context.Context, hash string, size int64) (models.Upload, error) {
return repository.FindReusableUploadByHash(ctx, hash, size)
}
+13 -13
View File
@@ -13,10 +13,10 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
)
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
@@ -67,7 +67,7 @@ func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
existing := model.Upload{
existing := models.Upload{
ID: 88001,
UserID: 42,
FileName: "existing.png",
@@ -77,7 +77,7 @@ func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
Extension: "png",
Hash: hashStr,
Type: "pixez_mirror",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&existing).Error; err != nil {
@@ -175,7 +175,7 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
}
var count int64
if err := dbConn.Model(&model.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
if err := dbConn.Model(&models.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
t.Fatalf("count uploads failed: %v", err)
}
if count != 2 {
@@ -188,7 +188,7 @@ func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
defer cleanup()
ctx := context.Background()
existing := model.Upload{
existing := models.Upload{
ID: 99001,
UserID: 1001,
FileName: "existing.png",
@@ -197,14 +197,14 @@ func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "generic",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&existing).Error; err != nil {
t.Fatalf("seed upload failed: %v", err)
}
duplicate := &model.Upload{
duplicate := &models.Upload{
ID: existing.ID,
UserID: 1002,
FileName: "duplicate.png",
@@ -213,7 +213,7 @@ func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
MimeType: "image/png",
Extension: "png",
Type: "generic",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := createUploadWithStats(ctx, duplicate); err == nil {
@@ -275,8 +275,8 @@ type totalStatsSnapshot struct {
}
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
var rows []model.UploadStat
if err := db.DB(ctx).Where("dimension = ?", model.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
var rows []models.UploadStat
if err := db.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
return totalStatsSnapshot{}, err
}
if len(rows) == 0 {
+13 -13
View File
@@ -6,44 +6,44 @@ package ingest
import (
"context"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/persistence"
uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
"gorm.io/gorm"
)
// Remove soft-deletes an upload and decrements incremental stats.
func Remove(ctx context.Context, uploadID uint64) (model.Upload, error) {
func Remove(ctx context.Context, uploadID uint64) (models.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
if err != nil {
return model.Upload{}, err
return models.Upload{}, err
}
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
return model.Upload{}, err
return models.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
upload.Status = models.UploadStatusDeleted
return upload, nil
}
// RemoveOwned soft-deletes an upload owned by userID and decrements incremental stats.
func RemoveOwned(ctx context.Context, userID, uploadID uint64) (model.Upload, error) {
func RemoveOwned(ctx context.Context, userID, uploadID uint64) (models.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
if err != nil {
return model.Upload{}, err
return models.Upload{}, err
}
if upload.UserID != userID {
return model.Upload{}, ErrForbidden
return models.Upload{}, ErrForbidden
}
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
return model.Upload{}, err
return models.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
upload.Status = models.UploadStatusDeleted
return upload, nil
}
func softDeleteUploadWithStats(ctx context.Context, upload *model.Upload) error {
func softDeleteUploadWithStats(ctx context.Context, upload *models.Upload) error {
statsSnapshot := *upload
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
+4 -4
View File
@@ -7,7 +7,7 @@ package ingest
import (
"io"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
)
// Policy controls how ingest handles hash collisions and record creation.
@@ -33,7 +33,7 @@ type Request struct {
Type string
AccessMode *int
Status model.UploadStatus
Status models.UploadStatus
Reader io.Reader
Size int64
@@ -42,7 +42,7 @@ type Request struct {
Extension string
Hash string
Metadata model.UploadMetadata
Metadata models.UploadMetadata
Policy Policy
ObjectKeyFn ObjectKeyFn
@@ -53,7 +53,7 @@ type Request struct {
// Result reports the outcome of an ingest operation.
type Result struct {
Upload model.Upload
Upload models.Upload
Created bool
Stored bool
Resolved bool
+34
View File
@@ -0,0 +1,34 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package upload 提供上传域的门面与类型重导出。
package upload
import (
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
)
// UploadStatus 上传状态类型别名
//
//nolint:revive
type UploadStatus = models.UploadStatus
// UploadMetadata 上传元数据类型别名
//
//nolint:revive
type UploadMetadata = models.UploadMetadata
// Upload 上传实体类型别名
type Upload = models.Upload
// UploadStat 上传统计实体类型别名
//
//nolint:revive
type UploadStat = models.UploadStat
// 上传状态常量别名
const (
UploadStatusPending = models.UploadStatusPending
UploadStatusUsed = models.UploadStatusUsed
UploadStatusDeleted = models.UploadStatusDeleted
)
+77
View File
@@ -0,0 +1,77 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package models 提供上传域核心数据模型。
package models
import (
"time"
)
// UploadStatus 上传状态
type UploadStatus string
// 上传状态
const (
UploadStatusPending UploadStatus = "pending" // 待使用
UploadStatusUsed UploadStatus = "used" // 已使用
UploadStatusDeleted UploadStatus = "deleted" // 已删除
)
// UploadMetadata 自定义可扩展的 JSON 字段存储非核心或可选的文件元数据
type UploadMetadata struct {
Width int `json:"width,omitempty"`
Height int `json:"height,omitempty"`
Duration float64 `json:"duration,omitempty"`
OriginalMime string `json:"original_mime,omitempty"`
UserAgent string `json:"user_agent,omitempty"`
ClientIP string `json:"client_ip,omitempty"`
Bucket string `json:"bucket,omitempty"`
Extra map[string]any `json:"extra,omitempty"`
}
// Upload 上传文件记录
type Upload struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
UserID uint64 `json:"user_id,string" gorm:"index;not null"`
FileName string `json:"file_name" gorm:"size:255;not null"`
FilePath string `json:"file_path" gorm:"size:500;not null;index"`
FileSize int64 `json:"file_size" gorm:"not null"`
MimeType string `json:"mime_type" gorm:"size:100;not null"`
Extension string `json:"extension" gorm:"size:50;not null"`
Hash string `json:"hash" gorm:"size:64;index"`
Type string `json:"type" gorm:"column:type;size:50;not null;index"`
Status UploadStatus `json:"status" gorm:"type:varchar(20);not null"`
AccessMode int `json:"access_mode" gorm:"column:access_mode;not null;default:0"`
Metadata UploadMetadata `json:"metadata" gorm:"serializer:json;type:jsonb"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (Upload) TableName() string {
return "w_uploads"
}
// Upload stats dimension keys stored in w_upload_stats.dimension.
const (
UploadStatDimensionTotal = "total"
UploadStatDimensionType = "type"
UploadStatDimensionCategory = "category"
UploadStatDimensionTrend = "trend"
)
// UploadStat 聚合统计记录
type UploadStat struct {
Dimension string `json:"dimension" gorm:"primaryKey;size:32;not null"`
StatKey string `json:"stat_key" gorm:"primaryKey;size:64;not null;default:''"`
FileCount int64 `json:"file_count" gorm:"not null;default:0"`
FileSize int64 `json:"file_size" gorm:"not null;default:0"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (UploadStat) TableName() string {
return "w_upload_stats"
}
+147
View File
@@ -0,0 +1,147 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
import (
"context"
"strings"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/util"
)
// UploadListFilter filters paginated upload queries.
//
//nolint:revive
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, []Upload, error) {
query := db.DB(ctx).Model(&Upload{}).
Where("status != ?", UploadStatusDeleted)
if filter.UserID != 0 {
query = query.Where("user_id = ?", filter.UserID)
}
if filter.Keyword != "" {
query = query.Where("LOWER(file_name) LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(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 []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) (Upload, error) {
var upload Upload
if err := db.DB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
return Upload{}, err
}
return upload, nil
}
// SoftDeleteUpload marks an upload as deleted.
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
func SoftDeleteUpload(ctx context.Context, upload *Upload) error {
return SoftDeleteUploadTx(db.DB(ctx), upload)
}
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
func SoftDeleteUploadTx(tx *gorm.DB, upload *Upload) error {
return tx.Model(upload).Update("status", UploadStatusDeleted).Error
}
// UpdateUpload applies partial field updates to an upload record.
func UpdateUpload(ctx context.Context, upload *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(&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) (Upload, error) {
var existing Upload
err := db.DB(ctx).
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, UploadStatusPending, UploadStatusUsed).
First(&existing).Error
return existing, err
}
// CreateUpload persists a new upload record.
func CreateUpload(ctx context.Context, upload *Upload) error {
return CreateUploadTx(db.DB(ctx), upload)
}
// CreateUploadTx persists a new upload record within an existing transaction.
func CreateUploadTx(tx *gorm.DB, upload *Upload) error {
if upload.ID == 0 {
upload.ID = idgen.NextUint64ID()
}
return tx.Create(upload).Error
}
// ListUploadsByIDs returns active uploads matching the given IDs.
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
var uploads []Upload
if err := db.DB(ctx).
Where("id IN ? AND status IN (?, ?)", ids, UploadStatusPending, UploadStatusUsed).
Find(&uploads).Error; err != nil {
return nil, err
}
return uploads, nil
}
// UploadQuery returns a scoped GORM query for uploads.
//
//nolint:revive
func UploadQuery(ctx context.Context) *gorm.DB {
return db.DB(ctx).Model(&Upload{})
}
// ListUploadStats returns all upload statistics rows.
func ListUploadStats(ctx context.Context) ([]UploadStat, error) {
var stats []UploadStat
if err := db.DB(ctx).Find(&stats).Error; err != nil {
return nil, err
}
return stats, nil
}
@@ -0,0 +1,140 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package repository 提供上传域数据库仓储层操作。
package repository
import (
"context"
"strings"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"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, []models.Upload, error) {
query := db.DB(ctx).Model(&models.Upload{}).
Where("status != ?", models.UploadStatusDeleted)
if filter.UserID != 0 {
query = query.Where("user_id = ?", filter.UserID)
}
if filter.Keyword != "" {
query = query.Where("LOWER(file_name) LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(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 []models.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) (models.Upload, error) {
var upload models.Upload
if err := db.DB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
return models.Upload{}, err
}
return upload, nil
}
// SoftDeleteUpload marks an upload as deleted.
func SoftDeleteUpload(ctx context.Context, upload *models.Upload) error {
return SoftDeleteUploadTx(db.DB(ctx), upload)
}
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
func SoftDeleteUploadTx(tx *gorm.DB, upload *models.Upload) error {
return tx.Model(upload).Update("status", models.UploadStatusDeleted).Error
}
// UpdateUpload applies partial field updates to an upload record.
func UpdateUpload(ctx context.Context, upload *models.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(&models.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) (models.Upload, error) {
var existing models.Upload
err := db.DB(ctx).
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, models.UploadStatusPending, models.UploadStatusUsed).
First(&existing).Error
return existing, err
}
// CreateUpload persists a new upload record.
func CreateUpload(ctx context.Context, upload *models.Upload) error {
return CreateUploadTx(db.DB(ctx), upload)
}
// CreateUploadTx persists a new upload record within an existing transaction.
func CreateUploadTx(tx *gorm.DB, upload *models.Upload) error {
return tx.Create(upload).Error
}
// ListUploadsByIDs returns active uploads matching the given IDs.
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error) {
var uploads []models.Upload
if err := db.DB(ctx).
Where("id IN ? AND status IN (?, ?)", ids, models.UploadStatusPending, models.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(&models.Upload{})
}
// ListUploadStats returns all upload statistics rows.
func ListUploadStats(ctx context.Context) ([]models.UploadStat, error) {
var stats []models.UploadStat
if err := db.DB(ctx).Find(&stats).Error; err != nil {
return nil, err
}
return stats, nil
}
@@ -17,4 +17,9 @@ const (
FileStatsTrendDays = 7
MaxS3KeyLength = 1024
AccessCacheTTL = 5 // seconds; multiplied by time.Second at use site
UploadStatDimensionTotal = "total"
UploadStatDimensionType = "type"
UploadStatDimensionCategory = "category"
UploadStatDimensionTrend = "trend"
)
+2
View File
@@ -38,4 +38,6 @@ const (
ErrInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high"
ErrParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w"
ErrQueryImagesForCacheWarmup = "查询待预热图片失败: %w"
ErrQueryTypeListFailed = "查询文件类型列表失败"
ErrUpdateFileFailed = "更新文件失败"
)
+18 -18
View File
@@ -7,32 +7,32 @@ import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// ApplyUploadStatsAdd increments incremental stats for a newly active upload record.
func ApplyUploadStatsAdd(ctx context.Context, upload *model.Upload) error {
func ApplyUploadStatsAdd(ctx context.Context, upload *models.Upload) error {
return applyUploadStatsDelta(ctx, upload, 1)
}
// ApplyUploadStatsRemove decrements incremental stats for a removed active upload record.
func ApplyUploadStatsRemove(ctx context.Context, upload *model.Upload) error {
func ApplyUploadStatsRemove(ctx context.Context, upload *models.Upload) error {
return applyUploadStatsDelta(ctx, upload, -1)
}
// RebuildUploadStats rebuilds all incremental stats from current upload records.
func RebuildUploadStats(ctx context.Context) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("1 = 1").Delete(&model.UploadStat{}).Error; err != nil {
if err := tx.Where("1 = 1").Delete(&models.UploadStat{}).Error; err != nil {
return err
}
var uploads []model.Upload
if err := tx.Where("status != ?", model.UploadStatusDeleted).Find(&uploads).Error; err != nil {
var uploads []models.Upload
if err := tx.Where("status != ?", models.UploadStatusDeleted).Find(&uploads).Error; err != nil {
return err
}
@@ -45,7 +45,7 @@ func RebuildUploadStats(ctx context.Context) error {
})
}
func applyUploadStatsDelta(ctx context.Context, upload *model.Upload, sign int64) error {
func applyUploadStatsDelta(ctx context.Context, upload *models.Upload, sign int64) error {
if upload == nil || !isActiveUploadStatus(upload.Status) {
return nil
}
@@ -55,7 +55,7 @@ func applyUploadStatsDelta(ctx context.Context, upload *model.Upload, sign int64
}
// ApplyUploadStatsDeltaTx applies incremental upload stats within an existing transaction.
func ApplyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) error {
func ApplyUploadStatsDeltaTx(tx *gorm.DB, upload *models.Upload, sign int64) error {
if upload == nil || !isActiveUploadStatus(upload.Status) || sign == 0 {
return nil
}
@@ -71,10 +71,10 @@ func ApplyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) erro
dimension string
key string
}{
{model.UploadStatDimensionTotal, ""},
{model.UploadStatDimensionType, typeKey},
{model.UploadStatDimensionCategory, GetFileCategory(upload.MimeType, upload.Extension)},
{model.UploadStatDimensionTrend, upload.CreatedAt.Format("2006-01-02")},
{models.UploadStatDimensionTotal, ""},
{models.UploadStatDimensionType, typeKey},
{models.UploadStatDimensionCategory, GetFileCategory(upload.MimeType, upload.Extension)},
{models.UploadStatDimensionTrend, upload.CreatedAt.Format("2006-01-02")},
}
for _, entry := range entries {
@@ -104,7 +104,7 @@ func upsertUploadStatDelta(tx *gorm.DB, dimension, key string, countDelta, sizeD
),
"updated_at": time.Now(),
}),
}).Create(&model.UploadStat{
}).Create(&models.UploadStat{
Dimension: dimension,
StatKey: key,
FileCount: countDelta,
@@ -113,19 +113,19 @@ func upsertUploadStatDelta(tx *gorm.DB, dimension, key string, countDelta, sizeD
}
// RecordUploadStatsAdd logs and applies upload stats increment.
func RecordUploadStatsAdd(ctx context.Context, upload *model.Upload) {
func RecordUploadStatsAdd(ctx context.Context, upload *models.Upload) {
if err := ApplyUploadStatsAdd(ctx, upload); err != nil {
logger.WarnF(ctx, "increment upload stats failed: %v", err)
}
}
// RecordUploadStatsRemove logs and applies upload stats decrement.
func RecordUploadStatsRemove(ctx context.Context, upload *model.Upload) {
func RecordUploadStatsRemove(ctx context.Context, upload *models.Upload) {
if err := ApplyUploadStatsRemove(ctx, upload); err != nil {
logger.WarnF(ctx, "decrement upload stats failed: %v", err)
}
}
func isActiveUploadStatus(status model.UploadStatus) bool {
return status == model.UploadStatusPending || status == model.UploadStatusUsed
func isActiveUploadStatus(status models.UploadStatus) bool {
return status == models.UploadStatusPending || status == models.UploadStatusUsed
}
@@ -8,9 +8,9 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"gorm.io/gorm"
)
@@ -19,13 +19,13 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
defer cleanup()
ctx := context.Background()
upload := &model.Upload{
upload := &models.Upload{
ID: 42002,
FileSize: 256,
MimeType: "image/jpeg",
Extension: "jpg",
Type: "avatar",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
@@ -49,13 +49,13 @@ func TestApplyUploadStatsAddAndRemove(t *testing.T) {
defer cleanup()
ctx := context.Background()
upload := &model.Upload{
upload := &models.Upload{
ID: 42001,
FileSize: 128,
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := ApplyUploadStatsAdd(ctx, upload); err != nil {
@@ -89,8 +89,8 @@ type uploadStatsSnapshot struct {
}
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
var rows []model.UploadStat
if err := db.DB(ctx).Where("dimension = ?", model.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
var rows []models.UploadStat
if err := db.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
return uploadStatsSnapshot{}, err
}
if len(rows) == 0 {
@@ -9,9 +9,10 @@ import (
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
)
// MigrationAccessState captures cached migration maintenance state.
@@ -70,9 +71,9 @@ func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
}
state := MigrationAccessState{
ReadOnly: execution.Status != model.TaskExecutionStatusSucceeded,
ReadOnly: execution.Status != task.TaskExecutionStatusSucceeded,
}
if execution.Status == model.TaskExecutionStatusSucceeded {
if execution.Status == task.TaskExecutionStatusSucceeded {
return state
}
+5 -5
View File
@@ -10,17 +10,17 @@ import (
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
)
// StorageMigrationTask is the Asynq task name for storage migration.
const StorageMigrationTask = "storage:migrate"
// LatestMigrationExecution returns the most recent storage migration task execution.
func LatestMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) {
return repository.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
func LatestMigrationExecution(ctx context.Context) (*task.TaskExecution, bool, error) {
return task.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
}
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
+3 -3
View File
@@ -6,9 +6,9 @@ package storage
import (
"context"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
)
// ReadOnly checks if the storage system is in read-only maintenance mode.
@@ -22,7 +22,7 @@ func ReadOnly(ctx context.Context) bool {
}
// OpenStoredObject opens a stored upload object from the active storage backend.
func OpenStoredObject(ctx context.Context, upload *model.Upload) (*objectstore.Object, error) {
func OpenStoredObject(ctx context.Context, upload *models.Upload) (*objectstore.Object, error) {
_, backend, err := objectstore.Active(ctx)
if err != nil {
return nil, err
+14 -14
View File
@@ -10,17 +10,17 @@ import (
"fmt"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
"github.com/Rain-kl/Wavelet/pkg/task"
uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"gorm.io/gorm"
)
@@ -61,9 +61,9 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
task.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339))
for {
var unusedUploads []model.Upload
var unusedUploads []models.Upload
if err := db.DB(ctx).
Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, oneHourAgo).
Where("id > ? AND status = ? AND created_at < ?", lastID, models.UploadStatusPending, oneHourAgo).
Order("id ASC").
Limit(batchSize).
Find(&unusedUploads).Error; err != nil {
@@ -81,9 +81,9 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
totalProcessed++
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&model.Upload{}).
Where("id = ? AND status = ?", u.ID, model.UploadStatusPending).
Update("status", model.UploadStatusDeleted).Error; err != nil {
if err := tx.Model(&models.Upload{}).
Where("id = ? AND status = ?", u.ID, models.UploadStatusPending).
Update("status", models.UploadStatusDeleted).Error; err != nil {
return err
}
@@ -112,10 +112,10 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
task.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...")
cutoff := time.Now().AddDate(0, 0, -7)
var pushHistoryCount int64
if err := db.DB(ctx).Model(&model.PushHistory{}).Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil {
if err := db.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil {
task.AppendLog(ctx, "统计待清理的历史推送记录失败: %v", err)
} else if pushHistoryCount > 0 {
if err := db.DB(ctx).Where("created_at < ?", cutoff).Delete(&model.PushHistory{}).Error; err != nil {
if err := db.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Delete(map[string]any{}).Error; err != nil {
task.AppendLog(ctx, "删除历史推送记录失败: %v", err)
} else {
task.AppendLog(ctx, "成功删除 %d 条历史推送记录 (截止时间: %s)", pushHistoryCount, cutoff.Format("2006-01-02 15:04:05"))
@@ -125,7 +125,7 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
}
task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
taskLogStats, err := repository.CleanupTaskExecutionLogs(ctx, time.Now())
taskLogStats, err := task.CleanupTaskExecutionLogs(ctx, time.Now())
if err != nil {
task.AppendLog(ctx, "清理任务执行日志失败: %v", err)
logger.ErrorF(ctx, "清理任务执行日志失败: %v", err)
+7 -7
View File
@@ -7,9 +7,9 @@ import (
"context"
"fmt"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
)
@@ -39,8 +39,8 @@ type RebuildUploadStatsHandler struct{}
func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
var activeCount int64
if err := db.DB(ctx).
Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
Model(&models.Upload{}).
Where("status != ?", models.UploadStatusDeleted).
Count(&activeCount).Error; err != nil {
task.AppendLog(ctx, "统计活跃上传记录失败: %v", err)
return nil, fmt.Errorf("count active uploads: %w", err)
@@ -53,9 +53,9 @@ func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*tas
return nil, fmt.Errorf("rebuild upload stats: %w", err)
}
var totalStat model.UploadStat
var totalStat models.UploadStat
if err := db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").
First(&totalStat).Error; err != nil {
task.AppendLog(ctx, "读取总量统计失败: %v", err)
return nil, fmt.Errorf("load total upload stats: %w", err)
@@ -8,9 +8,9 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
)
func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
@@ -20,16 +20,16 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
ctx := context.Background()
now := time.Now()
uploads := []model.Upload{
uploads := []models.Upload{
{
UserID: 1001, FileName: "a.jpg", FilePath: "uploads/a.jpg",
FileSize: 100, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash-a",
Type: "pixez_mirror", Status: model.UploadStatusUsed, CreatedAt: now,
Type: "pixez_mirror", Status: models.UploadStatusUsed, CreatedAt: now,
},
{
UserID: 1001, FileName: "b.png", FilePath: "uploads/b.png",
FileSize: 200, MimeType: "image/png", Extension: "png", Hash: "hash-b",
Type: "attachment", Status: model.UploadStatusUsed, CreatedAt: now,
Type: "attachment", Status: models.UploadStatusUsed, CreatedAt: now,
},
}
for i := range uploads {
@@ -39,8 +39,8 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
}
// Corrupt stats to ensure rebuild recalculates from uploads.
if err := db.DB(ctx).Create(&model.UploadStat{
Dimension: model.UploadStatDimensionTotal,
if err := db.DB(ctx).Create(&models.UploadStat{
Dimension: models.UploadStatDimensionTotal,
StatKey: "",
FileCount: 0,
FileSize: 0,
@@ -57,9 +57,9 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
t.Fatalf("Execute() returned empty result: %+v", result)
}
var totalStat model.UploadStat
var totalStat models.UploadStat
if err := db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").
First(&totalStat).Error; err != nil {
t.Fatalf("load total stat failed: %v", err)
}
+15 -15
View File
@@ -15,13 +15,13 @@ import (
"sync/atomic"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"golang.org/x/sync/errgroup"
)
@@ -173,8 +173,8 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
func countStorageObjects(ctx context.Context) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
err := db.DB(ctx).Model(&models.Upload{}).
Where("status != ?", models.UploadStatusDeleted).
Distinct("file_path").
Count(&count).Error
return count, err
@@ -185,7 +185,7 @@ func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) {
if err != nil || !ok {
return false, err
}
return execution.Status == model.TaskExecutionStatusPending || execution.Status == model.TaskExecutionStatusRunning, nil
return execution.Status == task.TaskExecutionStatusPending || execution.Status == task.TaskExecutionStatusRunning, nil
}
type migrationObject struct {
@@ -214,9 +214,9 @@ func migrateObjects(
task.AppendLog(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total)
var objects []migrationObject
query := db.DB(ctx).Model(&model.Upload{}).
query := db.DB(ctx).Model(&models.Upload{}).
Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash").
Where("status != ?", model.UploadStatusDeleted)
Where("status != ?", models.UploadStatusDeleted)
if lastFilePath != "" {
query = query.Where("file_path > ?", lastFilePath)
}
@@ -311,8 +311,8 @@ func migrateSingleObject(
if targetResult.Key != obj.FilePath {
task.AppendLog(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key)
if err := db.DB(ctx).Model(&model.Upload{}).
Where("file_path = ? AND status != ?", obj.FilePath, model.UploadStatusDeleted).
if err := db.DB(ctx).Model(&models.Upload{}).
Where("file_path = ? AND status != ?", obj.FilePath, models.UploadStatusDeleted).
Update("file_path", targetResult.Key).Error; err != nil {
return fmt.Errorf("update migrated object %q: %w", obj.FilePath, err)
}
@@ -344,15 +344,15 @@ func markMissingMigrationObjectDeleted(
) error {
task.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr)
var affectedUploads []model.Upload
var affectedUploads []models.Upload
if err := db.DB(ctx).
Where("file_path = ? AND status != ?", filePath, model.UploadStatusDeleted).
Where("file_path = ? AND status != ?", filePath, models.UploadStatusDeleted).
Find(&affectedUploads).Error; err != nil {
return fmt.Errorf("load missing object uploads %q: %w", filePath, err)
}
if err := db.DB(ctx).Model(&model.Upload{}).
if err := db.DB(ctx).Model(&models.Upload{}).
Where("file_path = ?", filePath).
Update("status", model.UploadStatusDeleted).Error; err != nil {
Update("status", models.UploadStatusDeleted).Error; err != nil {
return fmt.Errorf("update missing object %q: %w", filePath, err)
}
for i := range affectedUploads {
@@ -16,10 +16,10 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
)
@@ -59,7 +59,7 @@ func TestMigrationHandlerExecute(t *testing.T) {
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
}
upload := model.Upload{
upload := models.Upload{
ID: 99101,
UserID: 1,
FileName: "test.txt",
@@ -69,7 +69,7 @@ func TestMigrationHandlerExecute(t *testing.T) {
Extension: "txt",
Hash: "hash",
Type: "attachment",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
}
if err := dbConn.Create(&upload).Error; err != nil {
t.Fatalf("Create(upload) returned error: %v", err)
@@ -101,7 +101,7 @@ func TestMigrationHandlerExecute(t *testing.T) {
t.Errorf("migrated content = %q, want %q", copied.String(), content)
}
var migrated model.Upload
var migrated models.Upload
if err := dbConn.First(&migrated, upload.ID).Error; err != nil {
t.Fatalf("First(upload) returned error: %v", err)
}
@@ -156,7 +156,7 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
}
// Case 1: Incorrect Hash (should fail validation)
uploadIncorrect := model.Upload{
uploadIncorrect := models.Upload{
ID: 99102,
UserID: 1,
FileName: "test-hash.txt",
@@ -166,7 +166,7 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
Extension: "txt",
Hash: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", // Invalid hash
Type: "attachment",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
}
if err := dbConn.Create(&uploadIncorrect).Error; err != nil {
t.Fatalf("Create(uploadIncorrect) returned error: %v", err)
@@ -202,7 +202,7 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
}
// Case 2: Correct Hash (should succeed)
if err := dbConn.Model(&model.Upload{}).Where("id = ?", uploadIncorrect.ID).Update("hash", correctHash).Error; err != nil {
if err := dbConn.Model(&models.Upload{}).Where("id = ?", uploadIncorrect.ID).Update("hash", correctHash).Error; err != nil {
t.Fatalf("Update hash to correct value returned error: %v", err)
}
@@ -215,7 +215,7 @@ func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
t.Fatal("Execute() result = nil, want non-nil")
}
var migrated model.Upload
var migrated models.Upload
if err := dbConn.First(&migrated, uploadIncorrect.ID).Error; err != nil {
t.Fatalf("First(upload) returned error: %v", err)
}
+5 -5
View File
@@ -12,10 +12,10 @@ import (
"strings"
"sync"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
)
@@ -113,11 +113,11 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
return nil, fmt.Errorf("image cache warmup canceled: %w", err)
}
var uploads []model.Upload
var uploads []models.Upload
if err := db.DB(ctx).
Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)",
lastID,
model.UploadStatusDeleted,
models.UploadStatusDeleted,
"image/%",
[]string{"jpg", "jpeg", "png", "webp", "gif"},
).
+30 -30
View File
@@ -17,15 +17,15 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/diskcache"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"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/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
msg "github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/diskcache"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
@@ -48,39 +48,39 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
objectstore.ResetCache()
ctx := context.Background()
err := db.DB(ctx).AutoMigrate(&model.PushHistory{})
err := db.DB(ctx).AutoMigrate(&msg.PushHistory{})
require.NoError(t, err)
// 准备测试数据:创建一些上传记录
now := time.Now()
twoHoursAgo := now.Add(-2 * time.Hour)
records := []*model.Upload{
records := []*models.Upload{
// 超过1小时且状态为 pending 的记录 —— 应被清理
{
UserID: 1001, FileName: "old_file_1.jpg", FilePath: "uploads/old_1.jpg",
FileSize: 1024, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash1",
Type: "attachment", Status: model.UploadStatusPending,
Type: "attachment", Status: models.UploadStatusPending,
CreatedAt: twoHoursAgo,
},
{
UserID: 1001, FileName: "old_file_2.png", FilePath: "uploads/old_2.png",
FileSize: 2048, MimeType: "image/png", Extension: "png", Hash: "hash2",
Type: "attachment", Status: model.UploadStatusPending,
Type: "attachment", Status: models.UploadStatusPending,
CreatedAt: twoHoursAgo,
},
// 状态为 used 的记录 —— 不应被清理
{
UserID: 1001, FileName: "used_file.jpg", FilePath: "uploads/used.jpg",
FileSize: 512, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash3",
Type: "attachment", Status: model.UploadStatusUsed,
Type: "attachment", Status: models.UploadStatusUsed,
CreatedAt: twoHoursAgo,
},
// 不到1小时的 pending 记录 —— 不应被清理
{
UserID: 1001, FileName: "recent_file.jpg", FilePath: "uploads/recent.jpg",
FileSize: 256, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash4",
Type: "attachment", Status: model.UploadStatusPending,
Type: "attachment", Status: models.UploadStatusPending,
CreatedAt: now.Add(-10 * time.Minute),
},
}
@@ -90,7 +90,7 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
}
// 准备推送历史测试数据:1个旧的(应删除),1个新的(应保留)
oldPush := &model.PushHistory{
oldPush := &msg.PushHistory{
EventKey: "admin_login",
Channel: "email",
Target: "admin@test.com",
@@ -100,7 +100,7 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
Status: "success",
CreatedAt: now.AddDate(0, 0, -10),
}
newPush := &model.PushHistory{
newPush := &msg.PushHistory{
EventKey: "admin_login",
Channel: "lark",
Target: "http://webhook.com",
@@ -115,16 +115,16 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
err = db.DB(ctx).Create(newPush).Error
require.NoError(t, err)
oldTaskLog := &model.TaskExecution{
oldTaskLog := &task.TaskExecution{
TaskID: "old_low_frequency_task_log",
TaskType: "low:frequency",
TaskName: "低频任务",
Status: model.TaskExecutionStatusSucceeded,
Status: task.TaskExecutionStatusSucceeded,
CreatedAt: now.AddDate(0, 0, -31),
UpdatedAt: now.AddDate(0, 0, -31),
TriggeredBy: "system",
}
err = repository.CreateTaskExecution(ctx, oldTaskLog)
err = task.CreateTaskExecution(ctx, oldTaskLog)
require.NoError(t, err)
// 执行 handler
@@ -138,29 +138,29 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
// 验证数据库状态:pending 且超过1小时的应被标记为 deleted
var pendingCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusPending).Count(&pendingCount)
db.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusPending).Count(&pendingCount)
assert.Equal(t, int64(1), pendingCount, "应只剩1条 pending 记录(最近的文件)")
var deletedCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusDeleted).Count(&deletedCount)
db.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusDeleted).Count(&deletedCount)
assert.Equal(t, int64(2), deletedCount, "应有2条被标记为 deleted")
var usedCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusUsed).Count(&usedCount)
db.DB(ctx).Model(&models.Upload{}).Where("status = ?", models.UploadStatusUsed).Count(&usedCount)
assert.Equal(t, int64(1), usedCount, "used 状态的文件不应受影响")
// 验证推送历史数据状态:10天前的应被删除,今天的应保留
var pushCount int64
db.DB(ctx).Model(&model.PushHistory{}).Count(&pushCount)
db.DB(ctx).Model(&msg.PushHistory{}).Count(&pushCount)
assert.Equal(t, int64(1), pushCount, "应只剩1条推送历史记录")
var remainingPush model.PushHistory
var remainingPush msg.PushHistory
err = db.DB(ctx).First(&remainingPush).Error
require.NoError(t, err)
assert.Equal(t, "New Login", remainingPush.Title)
var taskLogCount int64
err = db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", "old_low_frequency_task_log").Count(&taskLogCount).Error
err = db.DB(ctx).Model(&task.TaskExecution{}).Where("task_id = ?", "old_low_frequency_task_log").Count(&taskLogCount).Error
require.NoError(t, err)
assert.Equal(t, int64(0), taskLogCount, "过期低频任务日志应被清理")
}
@@ -180,7 +180,7 @@ func TestSystemCleanupHandler_ExecuteNoFiles(t *testing.T) {
defer storageMock()
ctx := context.Background()
err := db.DB(ctx).AutoMigrate(&model.PushHistory{})
err := db.DB(ctx).AutoMigrate(&msg.PushHistory{})
require.NoError(t, err)
// 没有任何上传记录
@@ -279,7 +279,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
writeTaskTestPNG(t, firstPath, color.RGBA{R: 255, A: 255})
writeTaskTestPNG(t, secondPath, color.RGBA{G: 255, A: 255})
records := []model.Upload{
records := []models.Upload{
{
ID: 4101,
UserID: 1001,
@@ -287,7 +287,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
{
ID: 4102,
@@ -296,7 +296,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: secondPath,
MimeType: "application/octet-stream",
Extension: "jpg",
Status: model.UploadStatusPending,
Status: models.UploadStatusPending,
},
{
ID: 4103,
@@ -305,7 +305,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: filepath.Join(testDir, "notes.txt"),
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
Status: models.UploadStatusUsed,
},
{
ID: 4104,
@@ -314,7 +314,7 @@ func TestWarmImageCacheHandlerExecute(t *testing.T) {
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusDeleted,
Status: models.UploadStatusDeleted,
},
}
for i := range records {
+3 -2
View File
@@ -41,9 +41,10 @@ func IsDocumentExtension(ext string) bool {
// NormalizeImageQuality normalizes the requested image quality query parameter.
func NormalizeImageQuality(quality string) string {
switch strings.ToLower(quality) {
q := strings.TrimSpace(strings.ToLower(quality))
switch q {
case shared.ImageQualityLow, shared.ImageQualityMedium, shared.ImageQualityHigh:
return strings.ToLower(quality)
return q
default:
return shared.ImageQualityOrigin
}
+11 -12
View File
@@ -11,10 +11,9 @@ import (
"strconv"
"time"
persistence "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
persistence "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
@@ -60,7 +59,7 @@ func Login(c *gin.Context) {
return
}
user, err := repository.GetUserByUsername(c.Request.Context(), req.Username)
user, err := GetUserByUsername(c.Request.Context(), req.Username)
if err != nil {
response.AbortUnauthorized(c, errPasswordMismatch)
return
@@ -87,7 +86,7 @@ func Register(c *gin.Context) {
return
}
newUser := &model.User{
newUser := &User{
Username: req.Username,
Email: req.Email,
IsActive: true,
@@ -128,7 +127,7 @@ func ChangePassword(c *gin.Context) {
}
userID := auth.GetUserIDFromContext(c)
user, err := repository.GetUserByID(c.Request.Context(), userID)
user, err := GetUserByID(c.Request.Context(), userID)
if err != nil {
response.AbortNotFound(c, errUserNotFound)
return
@@ -160,7 +159,7 @@ func UpdateProfile(c *gin.Context) {
}
userID := auth.GetUserIDFromContext(c)
user, err := repository.GetUserByID(c.Request.Context(), userID)
user, err := GetUserByID(c.Request.Context(), userID)
if err != nil {
response.AbortNotFound(c, errUserNotFound)
return
@@ -184,7 +183,7 @@ func UpdateProfile(c *gin.Context) {
// ListAccessTokens lists access tokens for the current user.
func ListAccessTokens(c *gin.Context) {
userID := auth.GetUserIDFromContext(c)
var tokens []model.AccessToken
var tokens []AccessToken
gormDB := persistence.DB(c.Request.Context())
_ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error
c.JSON(http.StatusOK, response.OK(tokens))
@@ -215,7 +214,7 @@ func CreateAccessToken(c *gin.Context) {
masked = rawToken[:4] + "..." + rawToken[len(rawToken)-4:]
}
token := model.AccessToken{
token := AccessToken{
UserID: userID,
Name: req.Name,
TokenHash: tokenHash,
@@ -245,7 +244,7 @@ func DeleteAccessToken(c *gin.Context) {
}
userID := auth.GetUserIDFromContext(c)
var token model.AccessToken
var token AccessToken
gormDB := persistence.DB(c.Request.Context())
if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil {
response.AbortNotFound(c, errTokenNotFound)
@@ -267,7 +266,7 @@ func RotateAccessToken(c *gin.Context) {
}
userID := auth.GetUserIDFromContext(c)
var token model.AccessToken
var token AccessToken
gormDB := persistence.DB(c.Request.Context())
if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil {
response.AbortNotFound(c, errTokenNotFound)
+79
View File
@@ -0,0 +1,79 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package user
import (
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/util"
)
// AccessToken 个人访问令牌实体
type AccessToken struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
Name string `json:"name" gorm:"size:128;not null"`
TokenHash string `json:"-" gorm:"size:64;uniqueIndex;not null"`
MaskedToken string `json:"masked_token" gorm:"size:64;not null"`
IsAdmin bool `json:"is_admin" gorm:"default:false"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (AccessToken) TableName() string {
return "w_access_tokens"
}
// User 用户表实体
type User struct {
ID uint64 `json:"id,string" gorm:"primaryKey;not null"`
Username string `json:"username" gorm:"size:64;uniqueIndex"`
Password string `json:"password,omitempty" gorm:"size:255"`
Nickname string `json:"nickname" gorm:"size:255"`
Email string `json:"email" gorm:"size:255;index"`
AvatarURL string `json:"avatar_url" gorm:"size:255"`
IsActive bool `json:"is_active" gorm:"default:true;index"`
IsAdmin bool `json:"is_admin" gorm:"default:false"`
Bio string `json:"bio" gorm:"size:500"`
Phone string `json:"phone" gorm:"size:32"`
Gender string `json:"gender" gorm:"size:16"`
Website string `json:"website" gorm:"size:255"`
Location string `json:"location" gorm:"size:255"`
LastLoginAt time.Time `json:"last_login_at" gorm:"index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 表名
func (User) TableName() string {
return "w_users"
}
// SetEncryptedPassword 设置加密密码
func (u *User) SetEncryptedPassword(password string) error {
trimmed := strings.TrimSpace(password)
if trimmed == "" {
return errors.New("password cannot be empty")
}
hash, err := util.HashPassword(trimmed)
if err != nil {
return err
}
u.Password = hash
return nil
}
// CheckPassword 校验密码
func (u *User) CheckPassword(password string) bool {
if u.Password == "" {
util.DummyCheckPassword(password)
return false
}
return util.CheckPasswordHash(u.Password, password)
}
+5 -2
View File
@@ -44,15 +44,18 @@ func New(opts ...Option) *Plugin {
return p
}
// PluginName 用户插件唯一名称标识
const PluginName = "user"
// Name returns the unique identifier for the user domain plugin.
func (p *Plugin) Name() string {
return "user"
return PluginName
}
// Manifest returns the plugin metadata.
func (p *Plugin) Manifest() core.Manifest {
return core.Manifest{
Name: "user",
Name: PluginName,
Version: "1.0.0",
Description: "User profiles, credentials, role management, and access token domain plugin",
Author: "Wavelet Team",
+3 -4
View File
@@ -15,8 +15,7 @@ import (
"github.com/Rain-kl/Wavelet/core"
"github.com/Rain-kl/Wavelet/core/contracts"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/domain/user"
)
@@ -27,8 +26,8 @@ func setupTestDB(t *testing.T) *gorm.DB {
require.NoError(t, err)
require.NoError(t, testDB.AutoMigrate(
&model.User{},
&model.AccessToken{},
&user.User{},
&user.AccessToken{},
))
db.SetDB(testDB)
+161
View File
@@ -0,0 +1,161 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package user
import (
"context"
"strings"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
"gorm.io/gorm"
)
// GetUserByID 通过 ID 获取用户
func GetUserByID(ctx context.Context, id uint64) (*User, error) {
var u User
if err := db.DB(ctx).First(&u, id).Error; err != nil {
return nil, err
}
return &u, nil
}
// GetUserByUsername 通过用户名获取用户
func GetUserByUsername(ctx context.Context, username string) (*User, error) {
var u User
if err := db.DB(ctx).Where("username = ?", username).First(&u).Error; err != nil {
return nil, err
}
return &u, nil
}
// GetUserByEmail 通过邮箱获取用户
func GetUserByEmail(ctx context.Context, email string) (*User, error) {
var u User
if err := db.DB(ctx).Where("email = ?", email).First(&u).Error; err != nil {
return nil, err
}
return &u, nil
}
// CreateUser 创建用户
func CreateUser(ctx context.Context, u *User) error {
return db.DB(ctx).Create(u).Error
}
// UpdateUser 更新用户
func UpdateUser(ctx context.Context, u *User) error {
return db.DB(ctx).Save(u).Error
}
// ListUsers 分页查询用户
func ListUsers(ctx context.Context, page, pageSize int, keyword string) ([]*User, int64, error) {
db := db.DB(ctx).Model(&User{})
if keyword != "" {
escaped := util.EscapeLike(keyword)
db = db.Where("username LIKE ? ESCAPE '\\' OR nickname LIKE ? ESCAPE '\\' OR email LIKE ? ESCAPE '\\'", "%"+escaped+"%", "%"+escaped+"%", "%"+escaped+"%")
}
var total int64
if err := db.Count(&total).Error; err != nil {
return nil, 0, err
}
var users []*User
offset := (page - 1) * pageSize
if err := db.Offset(offset).Limit(pageSize).Order("id DESC").Find(&users).Error; err != nil {
return nil, 0, err
}
return users, total, nil
}
// GetAccessTokenByHash 通过 Hash 查询访问令牌
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*AccessToken, error) {
var token AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&token).Error; err != nil {
return nil, err
}
return &token, nil
}
// AdminUserListFilter 包含后台用户列表过滤条件
type AdminUserListFilter struct {
Username string
Keyword string
Page int
PageSize int
}
// ListAdminUsers 获取后台管理用户列表
func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []User, error) {
query := db.DB(ctx).Model(&User{})
if filter.Username != "" {
escaped := util.EscapeLike(strings.ToLower(filter.Username))
query = query.Where("LOWER(username) LIKE ? ESCAPE '\\'", "%"+escaped+"%")
}
if filter.Keyword != "" {
escaped := util.EscapeLike(strings.ToLower(filter.Keyword))
query = query.Where("LOWER(username) LIKE ? ESCAPE '\\' OR LOWER(nickname) LIKE ? ESCAPE '\\' OR LOWER(email) LIKE ? ESCAPE '\\'",
"%"+escaped+"%", "%"+escaped+"%", "%"+escaped+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var users []User
offset := (filter.Page - 1) * filter.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(filter.PageSize).Find(&users).Error; err != nil {
return 0, nil, err
}
return total, users, nil
}
// UpdateUserActive 更新用户激活状态
func UpdateUserActive(ctx context.Context, id uint64, active bool) error {
return db.DB(ctx).Model(&User{}).Where("id = ?", id).Update("is_active", active).Error
}
// GetActiveUserByID 获取处于激活状态的用户
func GetActiveUserByID(ctx context.Context, id uint64) (*User, error) {
var u User
if err := db.DB(ctx).Where("id = ? AND is_active = ?", id, true).First(&u).Error; err != nil {
return nil, err
}
return &u, nil
}
// DeleteUserWithRelations 删除用户及其级联关系
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(&AccessToken{}).Error; err != nil {
return err
}
return tx.Where("id = ?", id).Delete(&User{}).Error
})
}
// GetFirstAdminUser 获取第一个管理员用户
func GetFirstAdminUser(ctx context.Context) (*User, error) {
var u User
if err := db.DB(ctx).Where("is_admin = ?", true).Order("id ASC").First(&u).Error; err != nil {
return nil, err
}
return &u, nil
}
// ListUsernamesMatchingBase 列出匹配基础用户名的所有用户名
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
var usernames []string
escaped := util.EscapeLike(strings.ToLower(base))
if err := db.DB(ctx).Model(&User{}).
Where("LOWER(username) LIKE ? ESCAPE '\\'", escaped+"%").
Pluck("username", &usernames).Error; err != nil {
return nil, err
}
return usernames, nil
}
+93 -21
View File
@@ -1,20 +1,24 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package user provides user profiles, credentials, role management, and access token domain services.
package user
import (
"context"
"errors"
"fmt"
"strings"
"time"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
pkgu "github.com/Rain-kl/Wavelet/pkg/util"
)
func toUserDTO(u *model.User) *contracts.UserDTO {
func toUserDTO(u *User) *contracts.UserDTO {
if u == nil {
return nil
}
@@ -44,23 +48,23 @@ func newUserService() contracts.UserService {
}
func (s *userServiceImpl) GetUserByID(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
u, err := repository.GetUserByID(ctx, id)
u, err := GetUserByID(ctx, id)
if err != nil {
return nil, err
}
return toUserDTO(&u), nil
return toUserDTO(u), nil
}
func (s *userServiceImpl) GetUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) {
u, err := repository.GetUserByUsername(ctx, username)
u, err := GetUserByUsername(ctx, username)
if err != nil {
return nil, err
}
return toUserDTO(&u), nil
return toUserDTO(u), nil
}
func (s *userServiceImpl) GetUserByEmail(ctx context.Context, email string) (*contracts.UserDTO, error) {
var u model.User
var u User
if err := db.DB(ctx).Where("email = ?", email).First(&u).Error; err != nil {
return nil, err
}
@@ -72,7 +76,7 @@ func (s *userServiceImpl) CreateUser(ctx context.Context, req contracts.CreateUs
return nil, errors.New("user: username cannot be empty")
}
user := model.User{
user := User{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
@@ -94,7 +98,7 @@ func (s *userServiceImpl) CreateUser(ctx context.Context, req contracts.CreateUs
}
}
if err := repository.CreateUser(ctx, &user); err != nil {
if err := CreateUser(ctx, &user); err != nil {
return nil, err
}
@@ -129,7 +133,7 @@ func (s *userServiceImpl) UpdateProfile(ctx context.Context, id uint64, req cont
}
updates["updated_at"] = time.Now()
if err := db.DB(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error; err != nil {
if err := db.DB(ctx).Model(&User{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err
}
@@ -137,7 +141,7 @@ func (s *userServiceImpl) UpdateProfile(ctx context.Context, id uint64, req cont
}
func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPassword, newPassword string) error {
var user model.User
var user User
if err := db.DB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
return err
}
@@ -150,7 +154,7 @@ func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPass
return err
}
return db.DB(ctx).Model(&model.User{}).Where("id = ?", id).
return db.DB(ctx).Model(&User{}).Where("id = ?", id).
Updates(map[string]any{
"password": user.Password,
"updated_at": time.Now(),
@@ -158,7 +162,7 @@ func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPass
}
func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, password string) bool {
var user model.User
var user User
if err := db.DB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
pkgu.DummyCheckPassword(password)
return false
@@ -167,7 +171,7 @@ func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, passwor
}
func (s *userServiceImpl) UpdateLastLogin(ctx context.Context, id uint64, _ string) error {
return db.DB(ctx).Model(&model.User{}).Where("id = ?", id).
return db.DB(ctx).Model(&User{}).Where("id = ?", id).
Updates(map[string]any{
"last_login_at": time.Now(),
"updated_at": time.Now(),
@@ -182,13 +186,13 @@ func (s *userServiceImpl) ListUsers(ctx context.Context, page, pageSize int, key
pageSize = 20
}
filter := repository.AdminUserListFilter{
filter := AdminUserListFilter{
Username: keyword,
Page: page,
PageSize: pageSize,
}
total, users, err := repository.ListAdminUsers(ctx, filter)
total, users, err := ListAdminUsers(ctx, filter)
if err != nil {
return nil, 0, err
}
@@ -202,9 +206,77 @@ func (s *userServiceImpl) ListUsers(ctx context.Context, page, pageSize int, key
}
func (s *userServiceImpl) SetUserActive(ctx context.Context, id uint64, active bool) error {
return repository.UpdateUserActive(ctx, id, active)
return UpdateUserActive(ctx, id, active)
}
func (s *userServiceImpl) SetUserAdmin(ctx context.Context, id uint64, admin bool) error {
return db.DB(ctx).Model(&model.User{}).Where("id = ?", id).Update("is_admin", admin).Error
return db.DB(ctx).Model(&User{}).Where("id = ?", id).Update("is_admin", admin).Error
}
func (s *userServiceImpl) VerifyAccessToken(ctx context.Context, tokenHash string) (*contracts.UserDTO, bool, error) {
tokenRecord, err := GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, false, err
}
user, err := GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, false, err
}
return toUserDTO(user), tokenRecord.IsAdmin, nil
}
func (s *userServiceImpl) DeleteUser(ctx context.Context, id uint64) error {
return DeleteUserWithRelations(ctx, id)
}
func (s *userServiceImpl) CountUsers(ctx context.Context) (int64, error) {
var count int64
err := db.DB(ctx).Model(&User{}).Count(&count).Error
return count, err
}
func (s *userServiceImpl) CountActiveUsers(ctx context.Context) (int64, error) {
var count int64
err := db.DB(ctx).Model(&User{}).Where("is_active = ?", true).Count(&count).Error
return count, err
}
func (s *userServiceImpl) GetFirstAdminUser(ctx context.Context) (*contracts.UserDTO, error) {
u, err := GetFirstAdminUser(ctx)
if err != nil {
return nil, err
}
return toUserDTO(u), nil
}
func (s *userServiceImpl) UniqueUsername(ctx context.Context, base string) (string, error) {
base = strings.TrimSpace(base)
if base == "" {
base = PluginName
}
existingUsernames, err := ListUsernamesMatchingBase(ctx, base)
if err != nil {
return "", err
}
exists := make(map[string]bool, len(existingUsernames))
for _, u := range existingUsernames {
exists[strings.ToLower(u)] = true
}
if !exists[strings.ToLower(base)] {
return base, nil
}
for i := 1; i <= 1000; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if !exists[strings.ToLower(candidate)] {
return candidate, nil
}
}
return "", errors.New("failed to generate unique username")
}
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package driver_asynq_cron provides the Asynq cron schedule driver plugin for Cordis.
package driver_asynq_cron
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package driver_asynq_worker provides the Asynq worker driver plugin for Cordis.
package driver_asynq_worker
+123
View File
@@ -0,0 +1,123 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package driver_http
import (
"context"
"errors"
"log"
"net"
"net/http"
"os"
"os/signal"
"strconv"
"syscall"
"time"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/redis"
"github.com/gin-gonic/gin"
"go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin"
)
// BuildEngine 构建并初始化 Gin 路由引擎及全部中间件和路由
func BuildEngine() (*gin.Engine, error) {
// 运行模式
if config.Config.App.IsProduction() {
gin.SetMode(gin.ReleaseMode)
}
// 初始化路由
r := gin.New()
r.Use(gin.Recovery())
r.Use(corsMiddleware())
cfg := config.Config.Redis
addrs := cfg.Addrs
sessionAddr := "localhost:6379"
if len(addrs) > 0 {
sessionAddr = addrs[0]
}
sessionStore, err := redis.NewStoreWithDB(
cfg.MinIdleConn,
"tcp",
sessionAddr,
cfg.Username,
cfg.Password,
strconv.Itoa(cfg.DB),
[]byte(config.Config.App.SessionSecret),
)
if err != nil {
return nil, err
}
// 设置 Session Redis Key 前缀
if cfg.KeyPrefix != "" {
if err := redis.SetKeyPrefix(sessionStore, cfg.KeyPrefix+"session:"); err != nil {
log.Printf("[API] set session key prefix failed: %v\n", err)
}
}
sessionStore.Options(auth.GetSessionOptions(config.Config.App.SessionAge))
r.Use(sessions.Sessions(config.Config.App.SessionCookieName, sessionStore))
// 补充中间件
r.Use(otelgin.Middleware(config.Config.App.AppName), errorHandlerMiddleware(), loggerMiddleware(), risk_control.Middleware())
return r, nil
}
// Serve 启动 HTTP API 服务。onStarted 仅会在 HTTP 地址成功绑定后调用。
func Serve(onStarted func()) {
r, err := BuildEngine()
if err != nil {
log.Fatalf("[API] init session store failed: %v\n", err)
}
srv := &http.Server{
Addr: config.Config.App.Addr,
Handler: r,
ReadHeaderTimeout: 10 * time.Second,
}
listener, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", config.Config.App.Addr)
if err != nil {
log.Fatalf("[API] server failed to listen on %s: %v\n", config.Config.App.Addr, err)
}
if onStarted != nil {
onStarted()
}
util.Go(func() {
log.Printf("[API] server listening on %s\n", config.Config.App.Addr)
if err := srv.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Fatalf("[API] server failed: %v\n", err)
}
})
quit := make(chan os.Signal, 1)
signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM)
<-quit
shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Duration(config.Config.App.GracefulShutdownTimeout)*time.Second)
trace.Shutdown(shutdownCtx)
if err := srv.Shutdown(shutdownCtx); err != nil {
log.Printf("[API] server forced to shutdown: %v\n", err)
cancel()
os.Exit(1)
}
cancel()
log.Println("[API] server exited")
}
+111
View File
@@ -0,0 +1,111 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package driver_http 提供 HTTP 路由中间件与服务启动
package driver_http
import (
"context"
"net/http"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-gonic/gin"
"go.opentelemetry.io/otel/codes"
"go.opentelemetry.io/otel/trace"
)
func loggerMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
// 初始化 Trace
ctx, span := otel_trace.Start(c.Request.Context(), "LoggerMiddleware")
defer span.End()
// 开始计时
start := time.Now()
// 记录请求路径和 Query
path := c.Request.URL.Path
raw := c.Request.URL.RawQuery
if raw != "" {
path = path + "?" + raw
}
// 执行请求
c.Next()
// 停止计时
end := time.Now()
latency := end.Sub(start)
// 打印日志
// 排除健康检查接口
healthPath := config.Config.App.APIPrefix + "/health"
if c.Request.URL.Path != healthPath {
logger.InfoF(
ctx,
"[LoggerMiddleware] %s %s\nStartTime: %s\nEndTime: %s\nLatency: %d\nClientIP: %s\nResponse: %d %d",
c.Request.Method,
path,
start.Format(time.RFC3339),
end.Format(time.RFC3339),
latency.Milliseconds(),
c.ClientIP(),
c.Writer.Status(),
c.Writer.Size(),
)
}
// 设置 Span 状态
if c.Writer.Status() >= http.StatusBadRequest {
span := trace.SpanFromContext(ctx)
span.SetStatus(codes.Error, strconv.Itoa(c.Writer.Status()))
}
}
}
func isOriginAllowed(ctx context.Context, origin string) bool {
var val string
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || val == "" {
return false
}
allowedOrigins := strings.Split(val, ",")
for _, allowed := range allowedOrigins {
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
if allowed != "" && strings.EqualFold(allowed, origin) {
return true
}
}
return false
}
func corsMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
origin := c.Request.Header.Get("Origin")
if origin != "" && isOriginAllowed(c.Request.Context(), origin) {
c.Writer.Header().Set("Access-Control-Allow-Origin", origin)
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With, X-Access-Token, X-Cap-Token")
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE, PATCH")
}
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(http.StatusNoContent)
return
}
c.Next()
}
}
// errorHandlerMiddleware 委托给 response.ErrorHandlerMiddleware,保持路由层单一入口。
func errorHandlerMiddleware() gin.HandlerFunc {
return response.ErrorHandlerMiddleware()
}
@@ -0,0 +1,100 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package driver_http
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/pkg/testhelper"
"github.com/gin-gonic/gin"
)
func TestCORSMiddleware(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
gin.SetMode(gin.TestMode)
clearConfigCache := func() {}
t.Run("missing server_address configuration returns no CORS headers", func(t *testing.T) {
clearConfigCache()
if err := dbConn.Table("w_system_configs").Where("key = ?", "server_address").Update("value", "").Error; err != nil {
t.Fatalf("failed to update config: %v", err)
}
r := gin.New()
r.Use(corsMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
req.Header.Set("Origin", "http://attacker.com")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Header().Get("Access-Control-Allow-Origin") != "" {
t.Errorf("expected no Access-Control-Allow-Origin, got %s", w.Header().Get("Access-Control-Allow-Origin"))
}
})
t.Run("configured server_address allows exact origin match", func(t *testing.T) {
if err := dbConn.Table("w_system_configs").Where("key = ?", "server_address").Update("value", "http://trusted.com").Error; err != nil {
t.Fatalf("failed to update config: %v", err)
}
r := gin.New()
r.Use(corsMiddleware())
r.GET("/test", func(c *gin.Context) {
c.String(http.StatusOK, "ok")
})
// Trusted origin
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
req.Header.Set("Origin", "http://trusted.com")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Header().Get("Access-Control-Allow-Origin") != "http://trusted.com" {
t.Errorf("expected Access-Control-Allow-Origin http://trusted.com, got %s", w.Header().Get("Access-Control-Allow-Origin"))
}
if w.Header().Get("Access-Control-Allow-Credentials") != "true" {
t.Errorf("expected Access-Control-Allow-Credentials true, got %s", w.Header().Get("Access-Control-Allow-Credentials"))
}
// Untrusted origin
req, _ = http.NewRequest(http.MethodGet, "/test", nil)
req.Header.Set("Origin", "http://attacker.com")
w = httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Header().Get("Access-Control-Allow-Origin") != "" {
t.Errorf("expected no Access-Control-Allow-Origin for attacker, got %s", w.Header().Get("Access-Control-Allow-Origin"))
}
})
t.Run("preflight OPTIONS request responds with 204", func(t *testing.T) {
if err := dbConn.Table("w_system_configs").Where("key = ?", "server_address").Update("value", "http://trusted.com").Error; err != nil {
t.Fatalf("failed to update config: %v", err)
}
r := gin.New()
r.Use(corsMiddleware())
req, _ := http.NewRequest(http.MethodOptions, "/test", nil)
req.Header.Set("Origin", "http://trusted.com")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusNoContent {
t.Errorf("expected status 204, got %d", w.Code)
}
if w.Header().Get("Access-Control-Allow-Methods") == "" {
t.Error("expected Access-Control-Allow-Methods header")
}
})
}

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