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("轮换令牌失败,请稍后再试")
}