feat(framework): 回灌 OpenFlare 分层、安全与运行时改进

将平台域持久化收敛为 repository 唯一入口,model 去掉 IO。
邮件头写入前清除 CR/LF,防止 header 注入。
httppool 支持可配置 Transport;batchwriter 增加 MinBatchSize/Stats,flush 失败交回批次;任务 PermanentError 作为 SkipRetry 终态。
设置与推送页的确认改为 AlertDialog;axios 去尾斜杠并按 Gin 数组序列化查询参数。
升级共享 Go 依赖(Gin、Asynq、OTel、GORM、Redis 等)。
This commit is contained in:
ryan
2026-08-16 11:07:20 +08:00
parent b9b42e3174
commit 6a53619dd2
58 changed files with 2201 additions and 1175 deletions
+8 -8
View File
@@ -49,7 +49,7 @@ type ToggleAuthSourceRequest struct {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/auth-sources [get]
func ListAuthSources(c *gin.Context) {
sources, err := model.GetAuthSources(c.Request.Context())
sources, err := repository.GetAuthSources(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -88,7 +88,7 @@ func CreateAuthSource(c *gin.Context) {
Scopes: req.Scopes,
IconURL: req.IconURL,
}
if err := model.CreateAuthSource(c.Request.Context(), &source); err != nil {
if err := repository.CreateAuthSource(c.Request.Context(), &source); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -126,7 +126,7 @@ func UpdateAuthSource(c *gin.Context) {
}
// 记录更新前的 Discovery URL,以便更新成功后清除旧缓存条目。
existing, _ := model.GetAuthSourceByID(c.Request.Context(), id)
existing, _ := repository.GetAuthSourceByID(c.Request.Context(), id)
source := model.AuthSource{
ID: id,
@@ -141,7 +141,7 @@ func UpdateAuthSource(c *gin.Context) {
IconURL: req.IconURL,
}
keepSecret := source.ClientSecret == ""
if err := model.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil {
if err := repository.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -154,7 +154,7 @@ func UpdateAuthSource(c *gin.Context) {
oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL))
_ = repository.InvalidateAuthSourceCache(c.Request.Context())
updated, err := model.GetAuthSourceByID(c.Request.Context(), id)
updated, err := repository.GetAuthSourceByID(c.Request.Context(), id)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -190,7 +190,7 @@ func ToggleAuthSource(c *gin.Context) {
return
}
if err := model.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil {
if err := repository.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -216,7 +216,7 @@ func DeleteAuthSource(c *gin.Context) {
response.AbortBadRequest(c, err.Error())
return
}
if err := model.DeleteAuthSource(c.Request.Context(), id); err != nil {
if err := repository.DeleteAuthSource(c.Request.Context(), id); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -229,7 +229,7 @@ func parseSourceID(c *gin.Context) (uint64, error) {
if raw == "" {
return 0, errors.New(admin.InvalidAuthSourceID)
}
source, err := model.GetAuthSourceByName(c.Request.Context(), raw)
source, err := repository.GetAuthSourceByName(c.Request.Context(), raw)
if err == nil {
return source.ID, nil
}
+5 -9
View File
@@ -16,7 +16,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
@@ -163,10 +163,7 @@ func buildAccessLogFilter(ctx context.Context, c *gin.Context) (analyticsrepo.Ac
username := c.Query("username")
if username != "" {
var userIDs []uint64
err := db.DB(ctx).Model(&model.User{}).
Where("username LIKE ?", "%"+username+"%").
Pluck("id", &userIDs).Error
userIDs, err := repository.ListUserIDsByUsernameContains(ctx, username)
if err != nil {
return filter, fmt.Errorf("查询用户信息失败: %w", err)
}
@@ -215,8 +212,7 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
}
userMap := make(map[uint64]struct{ Username, Nickname string })
var users []model.User
if err := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; err == nil {
if users, err := repository.ListUsersByIDs(ctx, userIDs); err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
@@ -402,8 +398,8 @@ func GetLogsAnalytics(c *gin.Context) {
Username string
Nickname string
})
var users []model.User
if errProfile := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
users, errProfile := repository.ListUsersByIDs(ctx, userIDs)
if errProfile == nil {
for _, u := range users {
userProfileMap[u.ID] = struct {
Username string
+6 -7
View File
@@ -116,13 +116,12 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
return
}
if err := tx.Model(&model.TaskExecution{}).
Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed).
Updates(map[string]any{
"status": model.TaskExecutionStatusSucceeded,
"result": "存储配置直接更新,故障迁移任务自动标记为已解决",
"finished_at": time.Now(),
}).Error; err != nil {
if err := repository.MarkFailedTaskExecutionsSucceededTx(
tx,
"storage:migrate",
"存储配置直接更新,故障迁移任务自动标记为已解决",
time.Now(),
); err != nil {
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
}
}
+8 -7
View File
@@ -15,6 +15,7 @@ import (
"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/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
@@ -119,7 +120,7 @@ func ListTaskExecutions(c *gin.Context) {
}
}
executions, total, err := model.ListTaskExecutions(c.Request.Context(), req)
executions, total, err := repository.ListTaskExecutions(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -153,7 +154,7 @@ func GetTaskExecution(c *gin.Context) {
return
}
execution, err := model.GetTaskExecutionByID(c.Request.Context(), id)
execution, err := repository.GetTaskExecutionByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, TaskNotFound)
return
@@ -211,7 +212,7 @@ func RetryTask(c *gin.Context) {
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/schedules [get]
func ListSchedules(c *gin.Context) {
schedules, err := model.ListSchedules(c.Request.Context())
schedules, err := repository.ListSchedules(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -281,7 +282,7 @@ func CreateSchedule(c *gin.Context) {
IsActive: *req.IsActive,
}
if err := model.CreateSchedule(c.Request.Context(), schedule); err != nil {
if err := repository.CreateSchedule(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -333,7 +334,7 @@ func UpdateSchedule(c *gin.Context) {
}
// 检查定时任务是否存在
schedule, err := model.GetScheduleByID(c.Request.Context(), id)
schedule, err := repository.GetScheduleByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, ScheduleNotFound)
return
@@ -369,7 +370,7 @@ func UpdateSchedule(c *gin.Context) {
schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive
if err := model.UpdateSchedule(c.Request.Context(), schedule); err != nil {
if err := repository.UpdateSchedule(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -402,7 +403,7 @@ func DeleteSchedule(c *gin.Context) {
return
}
if err := model.DeleteSchedule(c.Request.Context(), id); err != nil {
if err := repository.DeleteSchedule(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return
}
+7 -6
View File
@@ -20,6 +20,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/platform/bootstrap"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
@@ -234,7 +235,7 @@ func TestListTaskExecutions(t *testing.T) {
{TaskID: "exec_003", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", Retryable: true, MaxRetry: 3},
}
for _, r := range records {
err := model.CreateTaskExecution(ctx, r)
err := repository.CreateTaskExecution(ctx, r)
require.NoError(t, err)
}
@@ -345,7 +346,7 @@ func TestGetTaskExecution(t *testing.T) {
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
t.Run("get existing execution", func(t *testing.T) {
@@ -410,7 +411,7 @@ func TestRetryTask(t *testing.T) {
StartedAt: &now,
FinishedAt: &now,
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
@@ -430,7 +431,7 @@ func TestRetryTask(t *testing.T) {
assert.True(t, ok)
assert.NotEmpty(t, newTaskID)
newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID)
newExecution, err := repository.GetTaskExecutionByTaskID(ctx, newTaskID)
require.NoError(t, err)
assert.Equal(t, 1, newExecution.RetryCount)
assert.Equal(t, "retry", newExecution.TriggeredBy)
@@ -446,7 +447,7 @@ func TestRetryTask(t *testing.T) {
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
@@ -466,7 +467,7 @@ func TestRetryTask(t *testing.T) {
Retryable: false,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
+3 -5
View File
@@ -10,7 +10,6 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"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"
@@ -41,7 +40,7 @@ func updateUserStatus(ctx context.Context, id uint64, active bool) error {
var tokens []model.AccessToken
if !active {
_ = db.DB(ctx).Where("user_id = ?", id).Find(&tokens).Error
tokens, _ = repository.ListAccessTokensByUserID(ctx, id)
}
err = repository.UpdateUserActive(ctx, id, active)
@@ -68,8 +67,7 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
return errors.New(cannotDelete)
}
var tokens []model.AccessToken
_ = db.DB(ctx).Where("user_id = ?", targetID).Find(&tokens).Error
tokens, _ := repository.ListAccessTokensByUserID(ctx, targetID)
err = repository.DeleteUserWithRelations(ctx, targetID)
if err == nil {
@@ -181,7 +179,7 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam
needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin)
var tokens []model.AccessToken
if needRevokeTokens {
_ = db.DB(ctx).Where("user_id = ?", param.ID).Find(&tokens).Error
tokens, _ = repository.ListAccessTokensByUserID(ctx, param.ID)
}
// 更新字段
+12 -10
View File
@@ -130,12 +130,12 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS
response.AbortUnauthorized(c, shared.UnAuthorized)
return
}
var user model.User
if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil {
user, err := repository.GetUserByID(ctx, userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
@@ -146,7 +146,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS
return
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
}
@@ -154,13 +154,15 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS
func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
var user model.User
account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub)
account, err := repository.FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch {
case err == nil:
if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil {
response.AbortInternal(c, err.Error())
loaded, loadErr := repository.GetUserByID(ctx, account.UserID)
if 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 {
@@ -173,7 +175,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
}
user.LastLoginAt = time.Now()
_ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
if err := setLoginSession(ctx, c, &user); err != nil {
response.AbortInternal(c, err.Error())
return
@@ -209,11 +211,11 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A
userInfo.Username = username
var user model.User
if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil {
if err := repository.CreateUserFromOAuth(ctx, &user, userInfo); err != nil {
response.AbortInternal(c, err.Error())
return model.User{}, false
}
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{
if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID,
UserID: user.ID,
ExternalID: userInfo.Sub,
@@ -8,7 +8,8 @@ import (
"strconv"
"strings"
"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/gin-gonic/gin"
@@ -26,7 +27,7 @@ import (
// @Router /api/v1/oauth/external-accounts [get]
func ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c)
accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID)
accounts, err := repository.ListExternalAccountsByUserID(c.Request.Context(), userID)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -57,7 +58,7 @@ func DeleteExternalAccount(c *gin.Context) {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return
}
if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
if err := repository.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
+8 -9
View File
@@ -8,8 +8,8 @@ import (
"context"
"errors"
"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"
"github.com/Rain-kl/Wavelet/internal/shared/response"
@@ -33,8 +33,8 @@ func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.A
tokenHash := model.HashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
var dbToken model.AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&dbToken).Error; err != nil {
dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err != nil {
return nil, nil, err
}
tokenRecord = &dbToken
@@ -43,8 +43,8 @@ func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.A
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive {
var dbUser model.User
if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err != nil {
return nil, nil, err
}
user = &dbUser
@@ -87,11 +87,10 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
user, err := GetCachedUser(ctx, userID)
if err != nil || !user.IsActive {
var dbUser model.User
// load user from db to make sure is active
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&dbUser)
if tx.Error != nil {
return nil, tx.Error
dbUser, loadErr := repository.GetActiveUserByID(ctx, userID)
if loadErr != nil {
return nil, loadErr
}
user = &dbUser
SetCachedUser(ctx, userID, user)
+3 -5
View File
@@ -9,8 +9,8 @@ import (
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
)
@@ -21,10 +21,8 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
base = "user"
}
var existingUsernames []string
if err := db.DB(ctx).Model(&model.User{}).
Where("username = ? OR username LIKE ?", base, base+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
existingUsernames, err := repository.ListUsernamesMatchingBase(ctx, base)
if err != nil {
return "", err
}
+2 -2
View File
@@ -50,8 +50,8 @@ func InitLogWriter(ctx context.Context) {
}
logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", path)
}),
batchwriter.WithFlushErrorHandler[*analytics.UserAccessLog](func(ctx context.Context, batchSize int, err error) {
logger.ErrorF(ctx, "[RiskControl] Send ClickHouse batch failed (batch=%d): %v", batchSize, err)
batchwriter.WithFlushErrorHandler[*analytics.UserAccessLog](func(ctx context.Context, items []*analytics.UserAccessLog, err error) {
logger.ErrorF(ctx, "[RiskControl] Send ClickHouse batch failed (batch=%d): %v", len(items), err)
}),
)
if err != nil {
+2 -14
View File
@@ -11,9 +11,8 @@ import (
"strings"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/repository"
)
// StorageMigrationTask is the Asynq task name for storage migration.
@@ -21,18 +20,7 @@ const StorageMigrationTask = "storage:migrate"
// LatestMigrationExecution returns the most recent storage migration task execution.
func LatestMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) {
var execution model.TaskExecution
err := db.DB(ctx).
Where("task_type = ?", StorageMigrationTask).
Order("id DESC").
First(&execution).Error
if err == nil {
return &execution, true, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
return nil, false, nil
return repository.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
}
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
+2 -1
View File
@@ -18,6 +18,7 @@ import (
"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/pkg/logger"
"gorm.io/gorm"
)
@@ -123,7 +124,7 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
}
task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
taskLogStats, err := model.CleanupTaskExecutionLogs(ctx, time.Now())
taskLogStats, err := repository.CleanupTaskExecutionLogs(ctx, time.Now())
if err != nil {
task.AppendLog(ctx, "清理任务执行日志失败: %v", err)
logger.ErrorF(ctx, "清理任务执行日志失败: %v", err)
+2 -1
View File
@@ -24,6 +24,7 @@ import (
"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/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -123,7 +124,7 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
UpdatedAt: now.AddDate(0, 0, -31),
TriggeredBy: "system",
}
err = model.CreateTaskExecution(ctx, oldTaskLog)
err = repository.CreateTaskExecution(ctx, oldTaskLog)
require.NoError(t, err)
// 执行 handler
+29 -30
View File
@@ -218,8 +218,8 @@ func sendRegisterEmailCode(ctx context.Context, email string) error {
return errors.New(errEmailRequired)
}
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil {
count, err := repository.CountUsersByEmail(ctx, email)
if err != nil {
return err
}
if count > 0 {
@@ -249,8 +249,8 @@ func validateRegisterEmailVerification(ctx context.Context, email, code string)
}
func updateUserProfile(ctx context.Context, userID uint64, input updateProfileInput) (*model.User, error) {
var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil {
dbUser, err := repository.GetUserByID(ctx, userID)
if err != nil {
return nil, errors.New(errUserNotFound)
}
@@ -260,8 +260,8 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn
return nil, errors.New(errEmailFormatInvalid)
}
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", input.Email, dbUser.ID).Count(&count).Error; err != nil {
count, err := repository.CountUsersByEmailExceptID(ctx, input.Email, dbUser.ID)
if err != nil {
return nil, err
}
if count > 0 {
@@ -281,26 +281,26 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn
dbUser.Website = strings.TrimSpace(input.Website)
dbUser.Location = strings.TrimSpace(input.Location)
if err := db.DB(ctx).Save(&dbUser).Error; err != nil {
if err := repository.UpdateUser(ctx, &dbUser); err != nil {
return nil, err
}
return &dbUser, nil
}
func getUserByUsernameOrEmail(ctx context.Context, input string) (*model.User, error) {
var user model.User
if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil {
user, err := repository.GetUserByUsernameOrEmail(ctx, input)
if err != nil {
return nil, err
}
return &user, nil
}
func updateLastLogin(ctx context.Context, user *model.User) error {
return db.DB(ctx).Model(user).Update("last_login_at", user.LastLoginAt).Error
return repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
}
func registerUserLogic(ctx context.Context, u *model.User) error {
if err := u.RegisterUser(ctx, db.DB(ctx)); err != nil {
if err := repository.RegisterUserWithChecks(ctx, u); err != nil {
if strings.Contains(err.Error(), "duplicate key") || strings.Contains(err.Error(), "UNIQUE") {
return errors.New("用户名或邮箱已被占用")
}
@@ -310,8 +310,8 @@ func registerUserLogic(ctx context.Context, u *model.User) error {
}
func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass string) error {
var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil {
dbUser, err := repository.GetUserByID(ctx, userID)
if err != nil {
return errors.New(errUserNotFound)
}
@@ -323,18 +323,17 @@ func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass st
return errors.New(errPasswordEncryptFailed)
}
if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil {
if err := repository.UpdateUserPassword(ctx, dbUser.ID, dbUser.Password); err != nil {
return errors.New("更新密码失败,请稍后再试")
}
// 吊销该用户所有的 Access Token
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Find(&tokens).Error; err == nil {
if tokens, err := repository.ListAccessTokensByUserID(ctx, dbUser.ID); err == nil {
for _, token := range tokens {
oauth.InvalidateCachedToken(ctx, token.TokenHash)
}
}
if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil {
if err := repository.DeleteAccessTokensByUserID(ctx, dbUser.ID); err != nil {
return errors.New("吊销 Access Token 失败,请稍后再试")
}
@@ -343,23 +342,23 @@ func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass st
}
func listAccessTokensLogic(ctx context.Context, userID uint64) ([]model.AccessToken, error) {
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil {
tokens, err := repository.ListAccessTokensByUserID(ctx, userID)
if err != nil {
return nil, errors.New("获取令牌列表失败,请稍后再试")
}
return tokens, nil
}
func countAccessTokensLogic(ctx context.Context, userID uint64) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil {
count, err := repository.CountAccessTokensByUserID(ctx, userID)
if err != nil {
return 0, errors.New("查询令牌数量失败,请稍后再试")
}
return count, nil
}
func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) error {
if err := db.DB(ctx).Create(record).Error; err != nil {
if err := repository.CreateAccessToken(ctx, record); err != nil {
return errors.New("创建令牌失败,请稍后再试")
}
oauth.SetCachedToken(ctx, record.TokenHash, record)
@@ -367,25 +366,25 @@ func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) erro
}
func deleteAccessTokenLogic(ctx context.Context, id, userID uint64) error {
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil {
tokenRecord, err := repository.GetAccessTokenByIDAndUserID(ctx, id, userID)
if err != nil {
return errors.New(errTokenNotFoundOrForbidden)
}
oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash)
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{})
if tx.Error != nil {
rows, err := repository.DeleteAccessTokenForUser(ctx, id, userID)
if err != nil {
return errors.New("删除令牌失败,请稍后再试")
}
if tx.RowsAffected == 0 {
if rows == 0 {
return errors.New(errTokenNotFoundOrForbidden)
}
return nil
}
func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *model.AccessToken, error) {
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil {
tokenRecord, err := repository.GetAccessTokenByIDAndUserID(ctx, id, userID)
if err != nil {
return "", nil, errors.New(errTokenNotFoundOrForbidden)
}
@@ -402,7 +401,7 @@ func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *mo
tokenRecord.TokenHash = newTokenHash
tokenRecord.MaskedToken = newMaskedToken
if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil {
if err := repository.SaveAccessToken(ctx, &tokenRecord); err != nil {
return "", nil, errors.New("轮换令牌失败,请稍后再试")
}
@@ -11,6 +11,7 @@ import (
const (
defaultQueueSize = 10_000
defaultMaxBatchSize = 1_000
defaultMinBatchSize = 50
defaultFlushEvery = time.Second
)
@@ -25,8 +26,17 @@ type Config struct {
// MaxBatchSize triggers a flush when the in-memory batch reaches this count.
MaxBatchSize int
// FlushInterval triggers a time-based flush even when the batch is smaller.
// MinBatchSize is the minimum in-memory batch size for time-based flushes.
// Zero disables the threshold and preserves legacy interval flush behavior.
// When set, interval flushes below this size are skipped unless MaxFlushWait elapses.
MinBatchSize int
// FlushInterval is how often the worker checks whether a time-based flush should run.
FlushInterval time.Duration
// MaxFlushWait forces a flush of any non-empty batch once the oldest item has waited
// this long, even if MinBatchSize has not been reached. Zero disables the force path.
MaxFlushWait time.Duration
}
// DefaultConfig returns production-friendly defaults aligned with audit log batching.
@@ -34,6 +44,7 @@ func DefaultConfig() Config {
return Config{
QueueSize: defaultQueueSize,
MaxBatchSize: defaultMaxBatchSize,
MinBatchSize: defaultMinBatchSize,
FlushInterval: defaultFlushEvery,
}
}
@@ -45,8 +56,14 @@ func (c Config) validate() error {
if c.MaxBatchSize <= 0 {
return fmt.Errorf("batchwriter: max batch size must be positive")
}
if c.MinBatchSize < 0 {
return fmt.Errorf("batchwriter: min batch size must be non-negative")
}
if c.FlushInterval <= 0 {
return fmt.Errorf("batchwriter: flush interval must be positive")
}
if c.MaxFlushWait < 0 {
return fmt.Errorf("batchwriter: max flush wait must be non-negative")
}
return nil
}
@@ -9,22 +9,34 @@ package batchwriter
import (
"context"
"sync"
"sync/atomic"
"time"
)
// FlushFunc persists a batch of queued items. It is invoked from the worker goroutine.
type FlushFunc[T any] func(ctx context.Context, items []T) error
// FlushErrorHandler is called when FlushFunc returns an error. The batch is discarded
// after the handler returns; the worker continues processing.
type FlushErrorHandler func(ctx context.Context, batchSize int, err error)
// FlushErrorHandler is called when FlushFunc returns an error after optional retries.
// The batch is discarded after the handler returns; the worker continues processing.
// Handlers receive the failed items so callers can release dedup keys or re-queue.
type FlushErrorHandler[T any] func(ctx context.Context, items []T, err error)
// Stats is a point-in-time snapshot of Writer queue and failure counters.
type Stats struct {
Name string
Depth int
Cap int
Drops int64
FlushErrors int64
Running bool
}
// Writer buffers items and flushes them by size or interval.
type Writer[T any] struct {
cfg Config
flush FlushFunc[T]
onFlushError FlushErrorHandler
onFlushError FlushErrorHandler[T]
onDrop func(T)
startOnce sync.Once
@@ -34,13 +46,16 @@ type Writer[T any] struct {
ch chan T
workerCtx context.Context
done chan struct{}
drops atomic.Int64
flushErrors atomic.Int64
}
// Option configures optional Writer callbacks.
type Option[T any] func(*Writer[T])
// WithFlushErrorHandler registers a callback for flush failures.
func WithFlushErrorHandler[T any](handler FlushErrorHandler) Option[T] {
func WithFlushErrorHandler[T any](handler FlushErrorHandler[T]) Option[T] {
return func(w *Writer[T]) {
w.onFlushError = handler
}
@@ -168,22 +183,37 @@ func (w *Writer[T]) Cap() int {
return w.cfg.QueueSize
}
// Stats returns a point-in-time snapshot of queue depth and failure counters.
func (w *Writer[T]) Stats() Stats {
return Stats{
Name: w.cfg.Name,
Depth: w.Len(),
Cap: w.Cap(),
Drops: w.drops.Load(),
FlushErrors: w.flushErrors.Load(),
Running: w.Running(),
}
}
func (w *Writer[T]) run() {
ticker := time.NewTicker(w.cfg.FlushInterval)
defer ticker.Stop()
batch := make([]T, 0, w.cfg.MaxBatchSize)
var batchStartedAt time.Time
flush := func() {
if len(batch) == 0 {
return
}
items := append([]T(nil), batch...)
if err := w.flush(w.workerCtx, items); err != nil {
w.flushErrors.Add(1)
if w.onFlushError != nil {
w.onFlushError(w.workerCtx, len(items), err)
w.onFlushError(w.workerCtx, items, err)
}
}
batch = batch[:0]
batchStartedAt = time.Time{}
}
defer func() {
@@ -197,17 +227,36 @@ func (w *Writer[T]) run() {
if !ok {
return
}
if len(batch) == 0 {
batchStartedAt = time.Now()
}
batch = append(batch, item)
if len(batch) >= w.cfg.MaxBatchSize {
flush()
}
case <-ticker.C:
flush()
if w.shouldFlushOnInterval(len(batch), batchStartedAt, time.Now()) {
flush()
}
}
}
}
func (w *Writer[T]) shouldFlushOnInterval(batchLen int, batchStartedAt time.Time, now time.Time) bool {
if batchLen == 0 {
return false
}
if w.cfg.MinBatchSize == 0 || batchLen >= w.cfg.MinBatchSize {
return true
}
if w.cfg.MaxFlushWait <= 0 || batchStartedAt.IsZero() {
return false
}
return !now.Before(batchStartedAt.Add(w.cfg.MaxFlushWait))
}
func (w *Writer[T]) notifyDrop(item T) {
w.drops.Add(1)
if w.onDrop == nil {
return
}
@@ -98,6 +98,7 @@ func TestWriterFlushesOnInterval(t *testing.T) {
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 0
cfg.FlushInterval = 20 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
@@ -228,18 +229,18 @@ func TestWriterInvokesFlushErrorHandler(t *testing.T) {
flushErr := errors.New("flush failed")
var (
mu sync.Mutex
errCount int
batchSize int
mu sync.Mutex
errCount int
gotItems []int
)
writer, err := New[int](cfg, func(context.Context, []int) error {
return flushErr
}, WithFlushErrorHandler[int](func(_ context.Context, size int, err error) {
}, WithFlushErrorHandler[int](func(_ context.Context, items []int, err error) {
mu.Lock()
defer mu.Unlock()
errCount++
batchSize = size
gotItems = append([]int(nil), items...)
if !errors.Is(err, flushErr) {
t.Errorf("flush error = %v, want %v", err, flushErr)
}
@@ -272,13 +273,219 @@ func TestWriterInvokesFlushErrorHandler(t *testing.T) {
mu.Lock()
gotCount := errCount
gotSize := batchSize
items := gotItems
mu.Unlock()
if gotCount != 1 {
t.Fatalf("flush error handler count = %d, want 1", gotCount)
}
if gotSize != 1 {
t.Fatalf("flush error handler batch size = %d, want 1", gotSize)
if diff := cmp.Diff([]int{7}, items); diff != "" {
t.Fatalf("flush error handler items mismatch (-want +got):\n%s", diff)
}
stats := writer.Stats()
if stats.FlushErrors != 1 {
t.Fatalf("Stats().FlushErrors = %d, want 1", stats.FlushErrors)
}
}
func TestWriterSkipsIntervalFlushBelowMinBatchSize(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batch []int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 5
cfg.FlushInterval = 20 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batch = append([]int(nil), items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
for i := range 3 {
if !writer.TryEnqueue(i + 1) {
t.Fatalf("TryEnqueue(%d) = false, want true", i+1)
}
}
time.Sleep(100 * time.Millisecond)
mu.Lock()
got := batch
mu.Unlock()
if len(got) != 0 {
t.Fatalf("interval flush with below-min batch = %v, want no flush", got)
}
}
func TestWriterFlushesOnIntervalWhenMinBatchSizeReached(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batch []int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 3
cfg.FlushInterval = 20 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batch = append([]int(nil), items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
for i := range 3 {
if !writer.TryEnqueue(i + 1) {
t.Fatalf("TryEnqueue(%d) = false, want true", i+1)
}
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := len(batch) == 3
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(5 * time.Millisecond)
}
mu.Lock()
got := batch
mu.Unlock()
want := []int{1, 2, 3}
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("interval flush at min batch size mismatch (-want +got):\n%s", diff)
}
}
func TestWriterForcesFlushAfterMaxFlushWait(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batch []int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 50
cfg.FlushInterval = 20 * time.Millisecond
cfg.MaxFlushWait = 80 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batch = append([]int(nil), items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
if !writer.TryEnqueue(1) {
t.Fatal("TryEnqueue(1) = false, want true")
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := len(batch) == 1
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(5 * time.Millisecond)
}
mu.Lock()
got := batch
mu.Unlock()
if diff := cmp.Diff([]int{1}, got); diff != "" {
t.Fatalf("max flush wait mismatch (-want +got):\n%s", diff)
}
}
func TestWriterStatsTracksDrops(t *testing.T) {
t.Parallel()
cfg := DefaultConfig()
cfg.Name = "test-drops"
cfg.QueueSize = 1
cfg.MaxBatchSize = 10
cfg.FlushInterval = time.Hour
writer, err := New[int](cfg, func(context.Context, []int) error { return nil })
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
if !writer.TryEnqueue(1) {
t.Fatal("TryEnqueue(1) = false, want true")
}
if writer.TryEnqueue(2) {
t.Fatal("TryEnqueue(2) = true, want false")
}
stats := writer.Stats()
if stats.Drops != 1 {
t.Fatalf("Stats().Drops = %d, want 1", stats.Drops)
}
if stats.Cap != 1 {
t.Fatalf("Stats().Cap = %d, want 1", stats.Cap)
}
if !stats.Running {
t.Fatal("Stats().Running = false, want true")
}
}
+19 -14
View File
@@ -13,6 +13,7 @@ import (
"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"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/hibiken/asynq"
@@ -110,7 +111,7 @@ func AppendLog(ctx context.Context, format string, args ...interface{}) {
}
logLine := fmt.Sprintf(format, args...)
if err := model.AppendTaskExecutionLog(ctx, taskID, logLine); err != nil {
if err := repository.AppendTaskExecutionLog(ctx, taskID, logLine); err != nil {
logger.ErrorF(ctx, "[TaskExecutor] 追加任务日志失败 taskID=%s: %v", taskID, err)
}
}
@@ -138,7 +139,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
TriggeredBy: triggeredBy,
}
if err := model.CreateTaskExecution(ctx, execution); err != nil {
if err := repository.CreateTaskExecution(ctx, execution); err != nil {
return "", fmt.Errorf(errCreateTaskExecutionFailed, err)
}
@@ -156,11 +157,11 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
now := time.Now()
execution.StartedAt = &now
execution.FinishedAt = &now
_ = model.UpdateTaskExecution(ctx, execution)
_ = repository.UpdateTaskExecution(ctx, execution)
return "", fmt.Errorf(errTaskEnqueueFailed, err)
}
if err := model.AppendTaskExecutionLog(ctx, taskID, fmt.Sprintf("[系统] 任务已成功入队,等待调度执行 (队列: %s, 最大重试次数: %d)", meta.Queue, meta.MaxRetry)); err != nil {
if err := repository.AppendTaskExecutionLog(ctx, taskID, fmt.Sprintf("[系统] 任务已成功入队,等待调度执行 (队列: %s, 最大重试次数: %d)", meta.Queue, meta.MaxRetry)); err != nil {
logger.ErrorF(ctx, "[TaskExecutor] 追加入队日志失败 taskID=%s: %v", taskID, err)
}
@@ -169,7 +170,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
// RetryTask 重试失败的任务
func RetryTask(ctx context.Context, id uint64) (string, error) {
execution, err := model.GetTaskExecutionByID(ctx, id)
execution, err := repository.GetTaskExecutionByID(ctx, id)
if err != nil {
return "", fmt.Errorf(errTaskExecutionNotFound, err)
}
@@ -198,7 +199,7 @@ func RetryTask(ctx context.Context, id uint64) (string, error) {
TriggeredBy: "retry",
}
if err := model.CreateTaskExecution(ctx, newExecution); err != nil {
if err := repository.CreateTaskExecution(ctx, newExecution); err != nil {
return "", fmt.Errorf(errCreateRetryExecutionFailed, err)
}
@@ -221,11 +222,11 @@ func RetryTask(ctx context.Context, id uint64) (string, error) {
now := time.Now()
newExecution.StartedAt = &now
newExecution.FinishedAt = &now
_ = model.UpdateTaskExecution(ctx, newExecution)
_ = repository.UpdateTaskExecution(ctx, newExecution)
return "", fmt.Errorf(errRetryTaskEnqueueFailed, err)
}
if err := model.AppendTaskExecutionLog(ctx, newTaskID, fmt.Sprintf("[系统] 手动触发重试,已重新创建任务并入队 (原任务ID: %s, 重试次数: %d/%d)", execution.TaskID, execution.RetryCount+1, execution.MaxRetry)); err != nil {
if err := repository.AppendTaskExecutionLog(ctx, newTaskID, fmt.Sprintf("[系统] 手动触发重试,已重新创建任务并入队 (原任务ID: %s, 重试次数: %d/%d)", execution.TaskID, execution.RetryCount+1, execution.MaxRetry)); err != nil {
logger.ErrorF(ctx, "[TaskExecutor] 追加重试日志失败 taskID=%s: %v", newTaskID, err)
}
@@ -350,7 +351,7 @@ func updateExecutionOnStart(ctx context.Context, execution *model.TaskExecution,
dirty = true
}
if dirty {
if updateErr := model.UpdateTaskExecution(ctx, execution); updateErr != nil {
if updateErr := repository.UpdateTaskExecution(ctx, execution); updateErr != nil {
logger.ErrorF(ctx, "[TaskExecutor] 更新执行状态失败 taskID=%s: %v", execution.TaskID, updateErr)
}
}
@@ -358,7 +359,7 @@ func updateExecutionOnStart(ctx context.Context, execution *model.TaskExecution,
// getOrCreateTaskExecution 获取已有的任务执行记录,如果不存在则针对已知任务类型动态创建记录
func getOrCreateTaskExecution(ctx context.Context, taskID string, t *asynq.Task, payload []byte, now time.Time) (*model.TaskExecution, error) {
execution, err := model.GetTaskExecutionByTaskID(ctx, taskID)
execution, err := repository.GetTaskExecutionByTaskID(ctx, taskID)
if err == nil {
return execution, nil
}
@@ -381,7 +382,7 @@ func getOrCreateTaskExecution(ctx context.Context, taskID string, t *asynq.Task,
StartedAt: &now,
}
if createErr := model.CreateTaskExecution(ctx, execution); createErr != nil {
if createErr := repository.CreateTaskExecution(ctx, execution); createErr != nil {
logger.ErrorF(ctx, "[TaskExecutor] 动态创建执行记录失败 taskID=%s: %v", taskID, createErr)
return nil, createErr
}
@@ -404,11 +405,11 @@ func completeTaskExecution(ctx context.Context, execution *model.TaskExecution,
handleSuccessfulTask(ctx, execution, t, duration, result)
}
if err := model.UpdateTaskExecution(ctx, execution); err != nil {
if err := repository.UpdateTaskExecution(ctx, execution); err != nil {
logger.ErrorF(ctx, "[TaskExecutor] 更新执行记录失败 taskID=%s: %v", execution.TaskID, err)
}
if shouldFlushTaskExecutionLog(ctx, execErr) {
if err := model.FlushTaskExecutionLog(ctx, execution.TaskID); err != nil {
if err := repository.FlushTaskExecutionLog(ctx, execution.TaskID); err != nil {
logger.ErrorF(ctx, "[TaskExecutor] 持久化任务日志失败 taskID=%s: %v", execution.TaskID, err)
}
}
@@ -428,7 +429,7 @@ func notifyTaskCompleted(ctx context.Context, execution *model.TaskExecution, re
}
func shouldFlushTaskExecutionLog(ctx context.Context, execErr error) bool {
if execErr == nil {
if isTerminalTaskExecutionError(execErr) {
return true
}
@@ -440,6 +441,10 @@ func shouldFlushTaskExecutionLog(ctx context.Context, execErr error) bool {
return retryCount >= maxRetry
}
func isTerminalTaskExecutionError(execErr error) bool {
return execErr == nil || errors.Is(execErr, asynq.SkipRetry)
}
func handleFailedTask(ctx context.Context, execution *model.TaskExecution, t *asynq.Task, duration time.Duration, execErr error, span trace.Span) {
execution.Status = model.TaskExecutionStatusFailed
execution.ErrorMessage = execErr.Error()
+20 -13
View File
@@ -6,11 +6,13 @@ package task
import (
"context"
"errors"
"fmt"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/hibiken/asynq"
"github.com/stretchr/testify/assert"
@@ -120,7 +122,7 @@ func TestAppendLogWithTaskID(t *testing.T) {
Status: model.TaskExecutionStatusRunning,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 注入 taskID 并追加日志
@@ -129,7 +131,7 @@ func TestAppendLogWithTaskID(t *testing.T) {
AppendLog(ctx, "处理了 %d 条数据", 50)
// 验证日志
found, err := model.GetTaskExecutionByTaskID(ctx, "log_test_001")
found, err := repository.GetTaskExecutionByTaskID(ctx, "log_test_001")
require.NoError(t, err)
assert.Contains(t, found.Log, "第一条日志")
assert.Contains(t, found.Log, "处理了 50 条数据")
@@ -188,7 +190,7 @@ func TestProcessTaskSuccess(t *testing.T) {
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 通过 asynq 的 Task 不能直接设置 taskID,ProcessTask 通过 t.ResultWriter().TaskID() 获取
@@ -206,7 +208,7 @@ func TestProcessTaskSuccess(t *testing.T) {
assert.Equal(t, "处理完成,共 100 条", result.Message)
// 验证日志被追加
found, err := model.GetTaskExecutionByTaskID(ctx, "process_success_001")
found, err := repository.GetTaskExecutionByTaskID(ctx, "process_success_001")
require.NoError(t, err)
assert.Contains(t, found.Log, "执行成功,处理了 100 条数据")
}
@@ -229,7 +231,7 @@ func TestProcessTaskFailure(t *testing.T) {
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 直接调用 handler
@@ -242,7 +244,7 @@ func TestProcessTaskFailure(t *testing.T) {
assert.Contains(t, err.Error(), "模拟执行失败")
// 验证日志
found, err := model.GetTaskExecutionByTaskID(ctx, "process_fail_001")
found, err := repository.GetTaskExecutionByTaskID(ctx, "process_fail_001")
require.NoError(t, err)
assert.Contains(t, found.Log, "开始执行任务")
}
@@ -259,7 +261,7 @@ func TestCompleteTaskExecutionFlushesLog(t *testing.T) {
Status: model.TaskExecutionStatusRunning,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
ctx = withTaskID(ctx, execution.TaskID)
@@ -277,13 +279,18 @@ func TestCompleteTaskExecutionFlushesLog(t *testing.T) {
trace.SpanFromContext(ctx),
)
found, err := model.GetTaskExecutionByTaskID(ctx, execution.TaskID)
found, err := repository.GetTaskExecutionByTaskID(ctx, execution.TaskID)
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusSucceeded, found.Status)
assert.Contains(t, found.Log, "任务执行中的日志")
assert.Contains(t, found.Log, "任务执行成功")
}
func TestPermanentErrorIsTerminalForLogFlush(t *testing.T) {
assert.True(t, isTerminalTaskExecutionError(PermanentError("配置无效")))
assert.False(t, isTerminalTaskExecutionError(errors.New("temporary failure")))
}
func TestRetryTask(t *testing.T) {
cleanup := setupTest(t)
defer cleanup()
@@ -305,7 +312,7 @@ func TestRetryTask(t *testing.T) {
Duration: 100,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 重试
@@ -315,7 +322,7 @@ func TestRetryTask(t *testing.T) {
assert.Contains(t, newTaskID, "retry_1_")
// 验证新记录
newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID)
newExecution, err := repository.GetTaskExecutionByTaskID(ctx, newTaskID)
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusPending, newExecution.Status)
assert.Equal(t, 1, newExecution.RetryCount)
@@ -324,7 +331,7 @@ func TestRetryTask(t *testing.T) {
assert.True(t, newExecution.Retryable)
// 原记录不变
original, err := model.GetTaskExecutionByID(ctx, execution.ID)
original, err := repository.GetTaskExecutionByID(ctx, execution.ID)
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusFailed, original.Status)
assert.Equal(t, 0, original.RetryCount)
@@ -345,7 +352,7 @@ func TestRetryTaskNotFailed(t *testing.T) {
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 尝试重试成功的任务
@@ -368,7 +375,7 @@ func TestRetryTaskNotRetryable(t *testing.T) {
MaxRetry: 0,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
_, err = RetryTask(ctx, execution.ID)
+35
View File
@@ -0,0 +1,35 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"strings"
"github.com/hibiken/asynq"
)
const defaultPermanentErrorMessage = "任务无法继续执行"
type permanentTaskError struct {
message string
}
// PermanentError marks a safe domain message as a non-retryable task failure.
// It intentionally accepts no underlying error so Error never exposes provider,
// URL, header, response-body, or other sensitive implementation details.
func PermanentError(message string) error {
message = strings.TrimSpace(message)
if message == "" {
message = defaultPermanentErrorMessage
}
return &permanentTaskError{message: message}
}
func (e *permanentTaskError) Error() string {
return e.message
}
func (e *permanentTaskError) Unwrap() error {
return asynq.SkipRetry
}
@@ -0,0 +1,27 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"errors"
"testing"
"github.com/hibiken/asynq"
"github.com/stretchr/testify/assert"
)
func TestPermanentErrorSkipsRetryWithoutExposingAsynqMessage(t *testing.T) {
err := PermanentError(" 来源配置无效 ")
assert.True(t, errors.Is(err, asynq.SkipRetry))
assert.Equal(t, "来源配置无效", err.Error())
assert.NotContains(t, err.Error(), asynq.SkipRetry.Error())
}
func TestPermanentErrorUsesSafeFallbackForBlankMessage(t *testing.T) {
err := PermanentError(" ")
assert.True(t, errors.Is(err, asynq.SkipRetry))
assert.Equal(t, defaultPermanentErrorMessage, err.Error())
}
+2 -2
View File
@@ -12,8 +12,8 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/platform/bootstrap"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/hibiken/asynq"
@@ -84,7 +84,7 @@ func ReloadScheduler() error {
}
// 2. 从数据库载入启用的定时任务配置
schedules, err := model.ListActiveSchedules(context.Background())
schedules, err := repository.ListActiveSchedules(context.Background())
if err != nil {
return fmt.Errorf("load schedules from db failed: %w", err)
}
-204
View File
@@ -4,14 +4,10 @@
package model
import (
"context"
"errors"
"regexp"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"gorm.io/gorm"
)
// 认证源类型
@@ -116,203 +112,3 @@ func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
// GetAuthSources 获取所有认证源(已脱敏)
func GetAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
// GetActiveAuthSources 获取所有已启用的认证源(已脱敏)
func GetActiveAuthSources(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
}
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
if id == 0 {
return nil, errors.New(errAuthSourceIDRequired)
}
var source AuthSource
if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写)
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
name = strings.TrimSpace(name)
if name == "" {
return nil, errors.New(errAuthSourceNameRequired)
}
var source AuthSource
if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
// CreateAuthSource 创建认证源
func CreateAuthSource(ctx context.Context, source *AuthSource) error {
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Create(source).Error
}
// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥
func UpdateAuthSource(ctx context.Context, source *AuthSource, keepSecret bool) error {
if source.ID == 0 {
return errors.New(errAuthSourceIDRequired)
}
var current AuthSource
if err := db.DB(ctx).First(&current, "id = ?", source.ID).Error; err != nil {
return err
}
if keepSecret {
source.ClientSecret = current.ClientSecret
}
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Model(&current).Updates(map[string]any{
"name": source.Name,
"type": source.Type,
"display_name": source.DisplayName,
"is_active": source.IsActive,
"client_id": source.ClientID,
"client_secret": source.ClientSecret,
"openid_discovery_url": source.OpenIDDiscoveryURL,
"scopes": source.Scopes,
"icon_url": source.IconURL,
}).Error
}
// ToggleAuthSource 切换认证源启用状态
func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error {
source, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
source.IsActive = isActive
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
}
// DeleteAuthSource 删除认证源及其关联的外部帐号绑定
func DeleteAuthSource(ctx context.Context, id uint64) error {
if id == 0 {
return errors.New(errAuthSourceIDRequired)
}
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
return err
}
return tx.Delete(&AuthSource{}, "id = ?", id).Error
})
}
// FindExternalAccount 查找外部帐号绑定记录
func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱)
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
return errors.New(errExternalAccountBindingIncomplete)
}
account.ExternalID = strings.TrimSpace(account.ExternalID)
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
account.Email = strings.TrimSpace(account.Email)
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var current ExternalAccount
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(&current).Error
if err == nil {
if current.UserID != account.UserID {
return errors.New(errExternalAccountAlreadyBoundToAnother)
}
return tx.Model(&current).Updates(map[string]any{
"external_username": account.ExternalUsername,
"email": account.Email,
}).Error
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
return tx.Create(account).Error
})
}
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccountView, error) {
if userID == 0 {
return nil, errors.New(errUserIDRequired)
}
var accounts []ExternalAccount
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
return nil, err
}
views := make([]ExternalAccountView, 0, len(accounts))
for _, account := range accounts {
var name, sourceType, label string
if account.AuthSourceID == 0 {
name = "default"
sourceType = "oidc"
label = "历史认证源"
} else {
source, err := GetAuthSourceByID(ctx, account.AuthSourceID)
if err != nil {
continue
}
name = source.Name
sourceType = source.Type
label = source.DisplayName
if label == "" {
label = source.Name
}
}
views = append(views, ExternalAccountView{
ID: account.ID,
AuthSourceID: account.AuthSourceID,
AuthSourceName: name,
AuthSourceType: sourceType,
AuthSourceLabel: label,
ExternalUsername: account.ExternalUsername,
Email: account.Email,
CreatedAt: account.CreatedAt,
})
}
return views, nil
}
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error {
if id == 0 || userID == 0 {
return errors.New(errExternalAccountBindingIDRequired)
}
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
+10 -23
View File
@@ -3,28 +3,15 @@
package model
// Domain validation messages used by model.Validate and other no-IO rules.
// Persistence / data-access messages belong in internal/repository (do not import repository).
const (
errRegistrationDisabled = "注册已关闭"
errDatabaseNotInitialized = "database not initialized"
errUsernameExists = "用户名已存在"
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
errTemplateKeyRequired = "模板标识符不能为空"
errTemplateNameRequired = "模板名称不能为空"
errTemplateContentRequired = "模板内容不能为空"
errTemplateUnavailable = "模板 %s 不存在或不可用: %w"
errTemplateRenderFailed = "模板 %s 渲染失败: %w"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头"
errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc"
errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errAuthSourceIDRequired = "认证源 ID 不能为空"
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
errTemplateKeyRequired = "模板标识符不能为空"
errTemplateNameRequired = "模板名称不能为空"
errTemplateContentRequired = "模板内容不能为空"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头"
errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc"
errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
)
-45
View File
@@ -4,10 +4,7 @@
package model
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// Schedule 定时任务配置表
@@ -26,45 +23,3 @@ type Schedule struct {
func (Schedule) TableName() string {
return "w_schedules"
}
// CreateSchedule 创建定时任务
func CreateSchedule(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Create(schedule).Error
}
// UpdateSchedule 更新定时任务
func UpdateSchedule(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Save(schedule).Error
}
// DeleteSchedule 删除定时任务
func DeleteSchedule(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
}
// ListSchedules 获取所有定时任务
func ListSchedules(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
}
+8 -239
View File
@@ -5,15 +5,7 @@
package model
import (
"context"
"errors"
"fmt"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/redis/go-redis/v9"
)
// TaskExecutionStatus 任务执行状态
@@ -25,10 +17,6 @@ const (
TaskExecutionStatusRunning TaskExecutionStatus = "running"
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// TaskExecution 任务执行记录
@@ -64,233 +52,14 @@ func (TaskExecution) TableName() string {
return "w_task_executions"
}
// CreateTaskExecution 创建任务执行记录
func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
}
// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecution(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
}
// 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
}
// ListTaskExecutionsRequest 查询任务执行记录列表请求
type ListTaskExecutionsRequest struct {
Status string `form:"status"`
TaskType string `form:"task_type"`
Page int `form:"page"`
PageSize int `form:"page_size"`
}
// ListTaskExecutions 分页查询任务执行记录
func ListTaskExecutions(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)
}
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
}
// 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
Status string `form:"status"`
TaskType string `form:"task_type"`
TaskTypePrefix string `form:"task_type_prefix"`
// TaskTypes is a comma-separated list of exact asynq task types (IN filter).
// Used when TaskType is empty; takes precedence over TaskTypePrefix.
TaskTypes string `form:"task_types"`
Page int `form:"page"`
PageSize int `form:"page_size"`
}
-75
View File
@@ -5,16 +5,13 @@
package model
import (
"context"
"errors"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/pkg/util"
"gorm.io/gorm"
)
// OAuthUserInfo 用户信息结构(同时支持 OIDC ID Token claims 和 UserEndpoint 响应)
@@ -104,14 +101,6 @@ func (u *User) CheckPassword(password string) bool {
return u.Password == password
}
// GetByID 根据 ID 查询用户
func (u *User) GetByID(tx *gorm.DB, id uint64) error {
if err := tx.Where("id = ?", id).First(u).Error; err != nil {
return err
}
return nil
}
// UpdateFromOAuthInfo 根据 OAuth 信息更新用户数据
func (u *User) UpdateFromOAuthInfo(oauthInfo *OAuthUserInfo) {
u.Username = oauthInfo.Username
@@ -129,67 +118,3 @@ func (u *User) CheckActive() error {
}
return nil
}
func (u *User) assignIDIfMissing() error {
if u.ID != 0 {
return nil
}
u.ID = idgen.NextUint64ID()
return nil
}
// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验)
func (u *User) CreateUser(_ context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
now := time.Now()
userID := oauthInfo.GetID()
newUser := User{
ID: userID,
Username: oauthInfo.Username,
Nickname: oauthInfo.Name,
Email: oauthInfo.Email,
AvatarURL: oauthInfo.AvatarURL,
IsActive: oauthInfo.Active,
LastLoginAt: now,
IsAdmin: false,
}
if err := newUser.assignIDIfMissing(); err != nil {
return err
}
if err := tx.Create(&newUser).Error; err != nil {
return err
}
*u = newUser
return nil
}
// RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验)
func (u *User) RegisterUser(_ context.Context, tx *gorm.DB) error {
// 检查用户名冲突
var count int64
if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil {
return err
}
if count > 0 {
return errors.New(errUsernameExists)
}
// 检查邮箱冲突
if u.Email != "" {
var emailCount int64
if err := tx.Model(&User{}).Where("email = ?", u.Email).Count(&emailCount).Error; err != nil {
return err
}
if emailCount > 0 {
return errors.New(errEmailAlreadyBound)
}
}
if err := u.assignIDIfMissing(); err != nil {
return err
}
if err := tx.Create(u).Error; err != nil {
return err
}
return nil
}
+69
View File
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListAccessTokensByUserID returns all access tokens for a user ordered by created_at desc.
func ListAccessTokensByUserID(ctx context.Context, userID uint64) ([]model.AccessToken, error) {
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil {
return nil, err
}
return tokens, nil
}
// CountAccessTokensByUserID returns how many access tokens a user owns.
func CountAccessTokensByUserID(ctx context.Context, userID uint64) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreateAccessToken inserts a new access token record.
func CreateAccessToken(ctx context.Context, record *model.AccessToken) error {
return db.DB(ctx).Create(record).Error
}
// GetAccessTokenByIDAndUserID loads a token owned by the given user.
func GetAccessTokenByIDAndUserID(ctx context.Context, id, userID uint64) (model.AccessToken, error) {
var token model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil {
return model.AccessToken{}, err
}
return token, nil
}
// DeleteAccessTokenForUser deletes a token if it belongs to the user.
// Returns the number of rows affected.
func DeleteAccessTokenForUser(ctx context.Context, id, userID uint64) (int64, error) {
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{})
return tx.RowsAffected, tx.Error
}
// GetAccessTokenByHash loads an access token by its token hash.
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (model.AccessToken, error) {
var token model.AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&token).Error; err != nil {
return model.AccessToken{}, err
}
return token, nil
}
// SaveAccessToken persists all fields of an existing access token.
func SaveAccessToken(ctx context.Context, record *model.AccessToken) error {
return db.DB(ctx).Save(record).Error
}
// DeleteAccessTokensByUserID deletes all access tokens for a user.
func DeleteAccessTokensByUserID(ctx context.Context, userID uint64) error {
return db.DB(ctx).Where("user_id = ?", userID).Delete(&model.AccessToken{}).Error
}
@@ -5,6 +5,7 @@ package analytics
import (
"context"
"io"
"testing"
"time"
@@ -169,6 +170,12 @@ func (m *mockConn) Exec(_ context.Context, _ string, _ ...any) error { return ni
func (m *mockConn) AsyncInsert(_ context.Context, _ string, _ bool, _ ...any) error { return nil }
func (m *mockConn) InsertFormat(_ context.Context, _ string, _ string, _ io.Reader) error { return nil }
func (m *mockConn) QueryFormat(_ context.Context, _ string, _ string, _ ...any) (io.ReadCloser, error) {
return nil, nil
}
func (m *mockConn) Ping(_ context.Context) error { return nil }
func (m *mockConn) Stats() driver.Stats { return driver.Stats{} }
+215
View File
@@ -0,0 +1,215 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"strings"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// GetAuthSources 获取所有认证源(已脱敏)
func GetAuthSources(ctx context.Context) ([]model.AuthSource, error) {
var sources []model.AuthSource
if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
// GetActiveAuthSources 获取所有已启用的认证源(已脱敏)
func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) {
var sources []model.AuthSource
if err := db.DB(ctx).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*model.AuthSource, error) {
if id == 0 {
return nil, errors.New(errAuthSourceIDRequired)
}
var source model.AuthSource
if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写)
func GetAuthSourceByName(ctx context.Context, name string) (*model.AuthSource, error) {
name = strings.TrimSpace(name)
if name == "" {
return nil, errors.New(errAuthSourceNameRequired)
}
var source model.AuthSource
if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
// CreateAuthSource 创建认证源
func CreateAuthSource(ctx context.Context, source *model.AuthSource) error {
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Create(source).Error
}
// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥
func UpdateAuthSource(ctx context.Context, source *model.AuthSource, keepSecret bool) error {
if source.ID == 0 {
return errors.New(errAuthSourceIDRequired)
}
var current model.AuthSource
if err := db.DB(ctx).First(&current, "id = ?", source.ID).Error; err != nil {
return err
}
if keepSecret {
source.ClientSecret = current.ClientSecret
}
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Model(&current).Updates(map[string]any{
colName: source.Name,
"type": source.Type,
"display_name": source.DisplayName,
"is_active": source.IsActive,
"client_id": source.ClientID,
"client_secret": source.ClientSecret,
"openid_discovery_url": source.OpenIDDiscoveryURL,
"scopes": source.Scopes,
"icon_url": source.IconURL,
}).Error
}
// ToggleAuthSource 切换认证源启用状态
func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error {
source, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
source.IsActive = isActive
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Model(&model.AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
}
// DeleteAuthSource 删除认证源及其关联的外部帐号绑定
func DeleteAuthSource(ctx context.Context, id uint64) error {
if id == 0 {
return errors.New(errAuthSourceIDRequired)
}
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("auth_source_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil {
return err
}
return tx.Delete(&model.AuthSource{}, "id = ?", id).Error
})
}
// FindExternalAccount 查找外部帐号绑定记录
func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*model.ExternalAccount, error) {
var account model.ExternalAccount
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱)
func BindExternalAccount(ctx context.Context, account *model.ExternalAccount) error {
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
return errors.New(errExternalAccountBindingIncomplete)
}
account.ExternalID = strings.TrimSpace(account.ExternalID)
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
account.Email = strings.TrimSpace(account.Email)
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var current model.ExternalAccount
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(&current).Error
if err == nil {
if current.UserID != account.UserID {
return errors.New(errExternalAccountAlreadyBoundToAnother)
}
return tx.Model(&current).Updates(map[string]any{
"external_username": account.ExternalUsername,
"email": account.Email,
}).Error
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
return tx.Create(account).Error
})
}
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]model.ExternalAccountView, error) {
if userID == 0 {
return nil, errors.New(errUserIDRequired)
}
var accounts []model.ExternalAccount
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
return nil, err
}
views := make([]model.ExternalAccountView, 0, len(accounts))
for _, account := range accounts {
var name, sourceType, label string
if account.AuthSourceID == 0 {
name = "default"
sourceType = "oidc"
label = "历史认证源"
} else {
source, err := GetAuthSourceByID(ctx, account.AuthSourceID)
if err != nil {
continue
}
name = source.Name
sourceType = source.Type
label = source.DisplayName
if label == "" {
label = source.Name
}
}
views = append(views, model.ExternalAccountView{
ID: account.ID,
AuthSourceID: account.AuthSourceID,
AuthSourceName: name,
AuthSourceType: sourceType,
AuthSourceLabel: label,
ExternalUsername: account.ExternalUsername,
Email: account.Email,
CreatedAt: account.CreatedAt,
})
}
return views, nil
}
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error {
if id == 0 || userID == 0 {
return errors.New(errExternalAccountBindingIDRequired)
}
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.ExternalAccount{}).Error
}
+3 -3
View File
@@ -178,7 +178,7 @@ func GetActiveAuthSourcesCached(ctx context.Context) ([]model.AuthSource, error)
}
}
sources, err := model.GetActiveAuthSources(ctx)
sources, err := GetActiveAuthSources(ctx)
if err != nil {
return nil, err
}
@@ -192,7 +192,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou
normalized := normalizeAuthSourceName(name)
if normalized == "" {
return model.GetAuthSourceByName(ctx, name)
return GetAuthSourceByName(ctx, name)
}
if source, ok := authSourceByNameRAM.GetIfPresent(normalized); ok {
@@ -210,7 +210,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou
}
}
source, err := model.GetAuthSourceByName(ctx, name)
source, err := GetAuthSourceByName(ctx, name)
if err != nil {
return nil, err
}
@@ -74,7 +74,7 @@ func TestGetActiveAuthSourcesCached_LoadsFromRedisBeforeDB(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
@@ -119,7 +119,7 @@ func TestGetAuthSourceByNameCached_LoadsFromRedisBeforeDB(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
@@ -164,7 +164,7 @@ func TestInvalidateAuthSourceCache_ClearsRedisKeys(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
if _, err := GetActiveAuthSourcesCached(ctx); err != nil {
@@ -209,7 +209,7 @@ func TestAuthSourceInvalidationPubSubClearsPeerRAM(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
+25
View File
@@ -0,0 +1,25 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
// Persistence and repository-layer parameter messages live here (unexported).
// Domain field validation used by model.Validate stays in internal/model/errs.go;
// repository may call model.Validate and return those errors as-is.
// Keep wording aligned with model where the same user-facing phrase applies,
// but do not import or re-export model unexported consts (would require exporting).
const (
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceIDRequired = "认证源 ID 不能为空"
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
)
const colName = "name"
+53
View File
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// CreateSchedule 创建定时任务
func CreateSchedule(ctx context.Context, schedule *model.Schedule) error {
return db.DB(ctx).Create(schedule).Error
}
// UpdateSchedule 更新定时任务
func UpdateSchedule(ctx context.Context, schedule *model.Schedule) error {
return db.DB(ctx).Save(schedule).Error
}
// DeleteSchedule 删除定时任务
func DeleteSchedule(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&model.Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*model.Schedule, error) {
var schedule model.Schedule
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedules 获取所有定时任务
func ListSchedules(ctx context.Context) ([]model.Schedule, error) {
var schedules []model.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) ([]model.Schedule, error) {
var schedules []model.Schedule
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
+1 -8
View File
@@ -18,14 +18,7 @@ import (
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
)
const (
configTypeSystem = "system"
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
)
const configTypeSystem = "system"
// PreheatSystemConfigs loads all system configs from database.
// This function strictly performs database read and does not perform any cache read or write operations.
+304
View File
@@ -0,0 +1,304 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"fmt"
"strings"
"time"
"github.com/redis/go-redis/v9"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// CreateTaskExecution 创建任务执行记录
func CreateTaskExecution(ctx context.Context, execution *model.TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
}
// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecution(ctx context.Context, execution *model.TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*model.TaskExecution, error) {
var execution model.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) (*model.TaskExecution, error) {
var execution model.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.
// ok is false when no row exists.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*model.TaskExecution, bool, error) {
var execution model.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(&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
}
// ListTaskExecutions 分页查询任务执行记录
func ListTaskExecutions(ctx context.Context, req model.ListTaskExecutionsRequest) ([]model.TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := db.DB(ctx).Model(&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 ?", req.TaskTypePrefix+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []model.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(&model.TaskExecution{}).
Where("task_type = ? AND status = ?", taskType, model.TaskExecutionStatusFailed).
Updates(map[string]any{
"status": model.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) (model.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 := []model.TaskExecutionStatus{model.TaskExecutionStatusSucceeded, model.TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
Model(&model.TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return model.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(&model.TaskExecution{})
if highFrequencyResult.Error != nil {
return model.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(&model.TaskExecution{})
if lowFrequencyResult.Error != nil {
return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return model.TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
func taskExecutionLogRedisKey(taskID string) string {
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
}
func loadTaskExecutionLog(ctx context.Context, execution *model.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 []model.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
}
@@ -2,7 +2,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
package repository
import (
"context"
@@ -10,7 +10,9 @@ import (
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
@@ -26,7 +28,7 @@ func setupTaskExecutionTestEnvironment(t *testing.T) func() {
})
require.NoError(t, err)
err = sqliteDB.AutoMigrate(&TaskExecution{})
err = sqliteDB.AutoMigrate(&model.TaskExecution{})
require.NoError(t, err)
miniRedis, err := miniredis.Run()
@@ -54,11 +56,11 @@ func TestCreateTaskExecution(t *testing.T) {
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "manual_cleanup_123",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Status: model.TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
RetryCount: 0,
@@ -79,11 +81,11 @@ func TestGetTaskExecutionByTaskID(t *testing.T) {
ctx := context.Background()
// 创建记录
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "test_task_id_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Status: model.TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
@@ -96,7 +98,7 @@ func TestGetTaskExecutionByTaskID(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, execution.ID, found.ID)
assert.Equal(t, "test_task_id_001", found.TaskID)
assert.Equal(t, TaskExecutionStatusPending, found.Status)
assert.Equal(t, model.TaskExecutionStatusPending, found.Status)
assert.True(t, found.Retryable)
assert.Equal(t, 3, found.MaxRetry)
@@ -110,11 +112,11 @@ func TestGetTaskExecutionByID(t *testing.T) {
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "test_by_id_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Status: model.TaskExecutionStatusPending,
TriggeredBy: "system",
}
err := CreateTaskExecution(ctx, execution)
@@ -132,11 +134,11 @@ func TestUpdateTaskExecution(t *testing.T) {
ctx := context.Background()
// 创建记录
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "test_update_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Status: model.TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
@@ -144,7 +146,7 @@ func TestUpdateTaskExecution(t *testing.T) {
// 更新状态为 running
now := time.Now()
execution.Status = TaskExecutionStatusRunning
execution.Status = model.TaskExecutionStatusRunning
execution.StartedAt = &now
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
@@ -152,12 +154,12 @@ func TestUpdateTaskExecution(t *testing.T) {
// 验证更新
found, err := GetTaskExecutionByTaskID(ctx, "test_update_001")
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusRunning, found.Status)
assert.Equal(t, model.TaskExecutionStatusRunning, found.Status)
assert.NotNil(t, found.StartedAt)
// 更新为 succeeded
finishTime := time.Now()
execution.Status = TaskExecutionStatusSucceeded
execution.Status = model.TaskExecutionStatusSucceeded
execution.FinishedAt = &finishTime
execution.Duration = 1500
execution.Result = "共清理 50 个文件"
@@ -166,7 +168,7 @@ func TestUpdateTaskExecution(t *testing.T) {
found, err = GetTaskExecutionByTaskID(ctx, "test_update_001")
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusSucceeded, found.Status)
assert.Equal(t, model.TaskExecutionStatusSucceeded, found.Status)
assert.Equal(t, int64(1500), found.Duration)
assert.Equal(t, "共清理 50 个文件", found.Result)
}
@@ -176,11 +178,11 @@ func TestUpdateTaskExecutionFailed(t *testing.T) {
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "test_fail_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Status: model.TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
@@ -190,7 +192,7 @@ func TestUpdateTaskExecutionFailed(t *testing.T) {
// 标记为失败
now := time.Now()
execution.Status = TaskExecutionStatusFailed
execution.Status = model.TaskExecutionStatusFailed
execution.StartedAt = &now
execution.FinishedAt = &now
execution.Duration = 200
@@ -200,7 +202,7 @@ func TestUpdateTaskExecutionFailed(t *testing.T) {
found, err := GetTaskExecutionByTaskID(ctx, "test_fail_001")
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusFailed, found.Status)
assert.Equal(t, model.TaskExecutionStatusFailed, found.Status)
assert.Equal(t, "S3 连接超时", found.ErrorMessage)
assert.Equal(t, int64(200), found.Duration)
}
@@ -210,11 +212,11 @@ func TestUpdateTaskExecutionDoesNotPersistBufferedLog(t *testing.T) {
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "test_omit_log_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Status: model.TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
@@ -226,15 +228,15 @@ func TestUpdateTaskExecutionDoesNotPersistBufferedLog(t *testing.T) {
assert.Empty(t, execution.Log)
execution.Status = TaskExecutionStatusSucceeded
execution.Status = model.TaskExecutionStatusSucceeded
execution.Duration = 100
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
var persisted TaskExecution
var persisted model.TaskExecution
err = db.DB(ctx).Where("task_id = ?", "test_omit_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusSucceeded, persisted.Status)
assert.Equal(t, model.TaskExecutionStatusSucceeded, persisted.Status)
assert.Empty(t, persisted.Log)
found, err := GetTaskExecutionByTaskID(ctx, "test_omit_log_001")
@@ -247,11 +249,11 @@ func TestAppendTaskExecutionLog(t *testing.T) {
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "test_log_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Status: model.TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
@@ -274,7 +276,7 @@ func TestAppendTaskExecutionLog(t *testing.T) {
assert.Contains(t, found.Log, "本批次找到 42 个待清理文件")
assert.Contains(t, found.Log, "清理完成,共删除 42 个文件")
var persisted TaskExecution
var persisted model.TaskExecution
err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Empty(t, persisted.Log)
@@ -332,11 +334,11 @@ func TestGetTaskExecutionLogPrefersRedis(t *testing.T) {
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: "redis_priority_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusRunning,
Status: model.TaskExecutionStatusRunning,
Log: "数据库旧日志",
TriggeredBy: "manual",
}
@@ -357,12 +359,12 @@ func TestListTaskExecutions(t *testing.T) {
ctx := context.Background()
// 创建多条记录,包含不同状态和类型
records := []*TaskExecution{
{TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusSucceeded, TriggeredBy: "manual"},
{TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusFailed, TriggeredBy: "system"},
{TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusPending, TriggeredBy: "manual"},
{TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusRunning, TriggeredBy: "manual"},
{TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusSucceeded, TriggeredBy: "system"},
records := []*model.TaskExecution{
{TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual"},
{TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system"},
{TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual"},
{TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusRunning, TriggeredBy: "manual"},
{TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "system"},
}
for _, r := range records {
err := CreateTaskExecution(ctx, r)
@@ -372,7 +374,7 @@ func TestListTaskExecutions(t *testing.T) {
require.NoError(t, err)
// 查询全部(分页)
items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 10})
items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(5), total)
assert.Len(t, items, 5)
@@ -383,24 +385,24 @@ func TestListTaskExecutions(t *testing.T) {
}
// 按状态筛选:failed
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10})
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Len(t, items, 1)
assert.Equal(t, "list_002", items[0].TaskID)
// 按类型筛选
_, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10})
_, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(2), total)
// 分页测试
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 2})
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 2})
require.NoError(t, err)
assert.Equal(t, int64(5), total)
assert.Len(t, items, 2)
items2, total2, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 2, PageSize: 2})
items2, total2, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 2, PageSize: 2})
require.NoError(t, err)
assert.Equal(t, int64(5), total2)
assert.Len(t, items2, 2)
@@ -409,10 +411,38 @@ func TestListTaskExecutions(t *testing.T) {
assert.NotEqual(t, items[0].ID, items2[0].ID)
// 状态 + 类型组合筛选
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "succeeded", TaskType: "system:cleanup", Page: 1, PageSize: 10})
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "succeeded", TaskType: "system:cleanup", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Equal(t, "list_001", items[0].TaskID)
// 按类型前缀筛选
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{TaskTypePrefix: "system:", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(3), total)
assert.Len(t, items, 3)
// 按多类型 IN 筛选
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{
TaskTypes: "system:cleanup,other:task",
Page: 1,
PageSize: 10,
})
require.NoError(t, err)
assert.Equal(t, int64(5), total)
assert.Len(t, items, 5)
// 精确类型优先于 task_types / 前缀
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{
TaskType: "other:task",
TaskTypes: "system:cleanup",
TaskTypePrefix: "system:",
Page: 1,
PageSize: 10,
})
require.NoError(t, err)
assert.Equal(t, int64(2), total)
assert.Len(t, items, 2)
}
func TestListTaskExecutionsDefaultPaging(t *testing.T) {
@@ -421,7 +451,7 @@ func TestListTaskExecutionsDefaultPaging(t *testing.T) {
ctx := context.Background()
// 不传分页参数,应使用默认值 page=1, pageSize=20
items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{})
items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{})
require.NoError(t, err)
assert.Equal(t, int64(0), total)
assert.Len(t, items, 0)
@@ -434,14 +464,14 @@ func TestCleanupTaskExecutionLogs(t *testing.T) {
now := time.Date(2026, 6, 17, 12, 0, 0, 0, time.UTC)
for i := 0; i < 31; i++ {
createTaskExecutionForCleanup(t, ctx, fmt.Sprintf("high_recent_%02d", i), "high:task", TaskExecutionStatusSucceeded, now.Add(-2*time.Hour))
createTaskExecutionForCleanup(t, ctx, fmt.Sprintf("high_recent_%02d", i), "high:task", model.TaskExecutionStatusSucceeded, now.Add(-2*time.Hour))
}
createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4))
createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", TaskExecutionStatusFailed, now.AddDate(0, 0, -40))
createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", TaskExecutionStatusRunning, now.AddDate(0, 0, -10))
createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31))
createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29))
createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", TaskExecutionStatusPending, now.AddDate(0, 0, -45))
createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4))
createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", model.TaskExecutionStatusFailed, now.AddDate(0, 0, -40))
createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", model.TaskExecutionStatusRunning, now.AddDate(0, 0, -10))
createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31))
createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29))
createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", model.TaskExecutionStatusPending, now.AddDate(0, 0, -45))
stats, err := CleanupTaskExecutionLogs(ctx, now)
require.NoError(t, err)
@@ -450,27 +480,27 @@ func TestCleanupTaskExecutionLogs(t *testing.T) {
for _, taskID := range []string{"high_old_4d", "high_old_40d", "low_old_31d"} {
var count int64
err := db.DB(ctx).Model(&TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
require.NoError(t, err)
assert.Equal(t, int64(0), count, "CleanupTaskExecutionLogs(%s) should delete expired log", taskID)
}
for _, taskID := range []string{"high_recent_00", "high_running_old", "low_recent_29d", "low_pending_old"} {
var count int64
err := db.DB(ctx).Model(&TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
require.NoError(t, err)
assert.Equal(t, int64(1), count, "CleanupTaskExecutionLogs(%s) should keep retained log", taskID)
}
}
func TestTaskExecutionTableName(t *testing.T) {
execution := TaskExecution{}
execution := model.TaskExecution{}
assert.Equal(t, "w_task_executions", execution.TableName())
}
func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status TaskExecutionStatus, createdAt time.Time) {
func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status model.TaskExecutionStatus, createdAt time.Time) {
t.Helper()
execution := &TaskExecution{
execution := &model.TaskExecution{
TaskID: taskID,
TaskType: taskType,
TaskName: taskType,
+114 -1
View File
@@ -5,8 +5,11 @@ package repository
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
@@ -175,7 +178,117 @@ func ListUsersByIDs(ctx context.Context, ids []uint64) ([]model.User, error) {
return users, nil
}
// ListUserIDsByUsernameContains returns user IDs whose username contains the given fragment.
func ListUserIDsByUsernameContains(ctx context.Context, username string) ([]uint64, error) {
if username == "" {
return []uint64{}, nil
}
var userIDs []uint64
if err := db.DB(ctx).Model(&model.User{}).
Where("username LIKE ?", "%"+username+"%").
Pluck("id", &userIDs).Error; err != nil {
return nil, err
}
return userIDs, nil
}
// UpdateUser updates all fields of an existing user.
func UpdateUser(ctx context.Context, user *model.User) error {
return db.DB(ctx).Save(user).Error
}
// CreateUserFromOAuth creates a user from OAuth profile data and fills userOut.
func CreateUserFromOAuth(ctx context.Context, userOut *model.User, oauthInfo *model.OAuthUserInfo) error {
now := time.Now()
userID := oauthInfo.GetID()
newUser := model.User{
ID: userID,
Username: oauthInfo.Username,
Nickname: oauthInfo.Name,
Email: oauthInfo.Email,
AvatarURL: oauthInfo.AvatarURL,
IsActive: oauthInfo.Active,
LastLoginAt: now,
IsAdmin: false,
}
if newUser.ID == 0 {
newUser.ID = idgen.NextUint64ID()
}
if err := db.DB(ctx).Create(&newUser).Error; err != nil {
return err
}
*userOut = newUser
return nil
}
// ListUsernamesMatchingBase returns usernames equal to base or prefixed with base+"-".
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
var names []string
if err := db.DB(ctx).Model(&model.User{}).
Where("username = ? OR username LIKE ?", base, base+"-%").
Pluck("username", &names).Error; err != nil {
return nil, err
}
return names, nil
}
// GetActiveUserByID loads a user by ID who is active.
func GetActiveUserByID(ctx context.Context, id uint64) (model.User, error) {
var user model.User
if err := db.DB(ctx).Where("id = ? AND is_active = ?", id, true).First(&user).Error; err != nil {
return model.User{}, err
}
return user, nil
}
// GetUserByUsernameOrEmail loads a user by username or email.
func GetUserByUsernameOrEmail(ctx context.Context, input string) (model.User, error) {
var user model.User
if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil {
return model.User{}, err
}
return user, nil
}
// CountUsersByEmailExceptID counts users with the email excluding a given user id.
func CountUsersByEmailExceptID(ctx context.Context, email string, exceptID uint64) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", email, exceptID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// UpdateUserLastLoginAt updates only last_login_at for a user.
func UpdateUserLastLoginAt(ctx context.Context, userID uint64, at time.Time) error {
return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("last_login_at", at).Error
}
// UpdateUserPassword updates only the password hash for a user.
func UpdateUserPassword(ctx context.Context, userID uint64, passwordHash string) error {
return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("password", passwordHash).Error
}
// RegisterUserWithChecks validates username/email uniqueness then creates the user.
func RegisterUserWithChecks(ctx context.Context, user *model.User) error {
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", user.Username).Count(&count).Error; err != nil {
return err
}
if count > 0 {
return errors.New("用户名已存在")
}
if user.Email != "" {
var emailCount int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", user.Email).Count(&emailCount).Error; err != nil {
return err
}
if emailCount > 0 {
return errors.New("该邮箱已被其他账号绑定")
}
}
if user.ID == 0 {
user.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(user).Error
}