refactor(user): decouple repository from direct database infra import

This commit is contained in:
ryan
2026-08-28 13:37:13 +08:00
parent 3570ccd5c3
commit a296818922
5 changed files with 90 additions and 48 deletions
+21 -22
View File
@@ -14,7 +14,6 @@ import (
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
database "Wavelet/plugins/infra/database"
"gorm.io/gorm"
pkgu "Wavelet/pkg/util"
@@ -75,7 +74,7 @@ func (s *userServiceImpl) GetUserByUsername(ctx context.Context, username string
func (s *userServiceImpl) GetUserByEmail(ctx context.Context, email string) (*contracts.UserDTO, error) {
var u User
if err := database.DB(ctx).Where("email = ?", email).First(&u).Error; err != nil {
if err := getDB(ctx).Where("email = ?", email).First(&u).Error; err != nil {
return nil, err
}
return toUserDTO(&u), nil
@@ -143,7 +142,7 @@ func (s *userServiceImpl) UpdateProfile(ctx context.Context, id uint64, req cont
}
updates[columnUpdatedAt] = time.Now()
if err := database.DB(ctx).Model(&User{}).Where("id = ?", id).Updates(updates).Error; err != nil {
if err := getDB(ctx).Model(&User{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err
}
@@ -152,7 +151,7 @@ func (s *userServiceImpl) UpdateProfile(ctx context.Context, id uint64, req cont
func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPassword, newPassword string) error {
var user User
if err := database.DB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
if err := getDB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
return err
}
@@ -164,7 +163,7 @@ func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPass
return err
}
return database.DB(ctx).Model(&User{}).Where("id = ?", id).
return getDB(ctx).Model(&User{}).Where("id = ?", id).
Updates(map[string]any{
"password": user.Password,
columnUpdatedAt: time.Now(),
@@ -173,7 +172,7 @@ func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPass
func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, password string) bool {
var user User
if err := database.DB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
if err := getDB(ctx).Where("id = ?", id).First(&user).Error; err != nil {
pkgu.DummyCheckPassword(password)
return false
}
@@ -181,7 +180,7 @@ func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, passwor
}
func (s *userServiceImpl) UpdateLastLogin(ctx context.Context, id uint64, _ string) error {
return database.DB(ctx).Model(&User{}).Where("id = ?", id).
return getDB(ctx).Model(&User{}).Where("id = ?", id).
Updates(map[string]any{
"last_login_at": time.Now(),
columnUpdatedAt: time.Now(),
@@ -220,7 +219,7 @@ func (s *userServiceImpl) SetUserActive(ctx context.Context, id uint64, active b
}
func (s *userServiceImpl) SetUserAdmin(ctx context.Context, id uint64, admin bool) error {
return database.DB(ctx).Model(&User{}).Where("id = ?", id).Update("is_admin", admin).Error
return getDB(ctx).Model(&User{}).Where("id = ?", id).Update("is_admin", admin).Error
}
func (s *userServiceImpl) VerifyAccessToken(ctx context.Context, tokenHash string) (*contracts.UserDTO, bool, error) {
@@ -243,13 +242,13 @@ func (s *userServiceImpl) DeleteUser(ctx context.Context, id uint64) error {
func (s *userServiceImpl) CountUsers(ctx context.Context) (int64, error) {
var count int64
err := database.DB(ctx).Model(&User{}).Count(&count).Error
err := getDB(ctx).Model(&User{}).Count(&count).Error
return count, err
}
func (s *userServiceImpl) CountActiveUsers(ctx context.Context) (int64, error) {
var count int64
err := database.DB(ctx).Model(&User{}).Where("is_active = ?", true).Count(&count).Error
err := getDB(ctx).Model(&User{}).Where("is_active = ?", true).Count(&count).Error
return count, err
}
@@ -292,7 +291,7 @@ func (s *userServiceImpl) UniqueUsername(ctx context.Context, base string) (stri
}
func (s *userServiceImpl) AdminListUsers(ctx context.Context, filter contracts.AdminListUsersFilter) (int64, []*contracts.UserDTO, error) {
query := database.DB(ctx).Table("w_users")
query := getDB(ctx).Table("w_users")
if filter.UserID != nil {
query = query.Where("id = ?", *filter.UserID)
}
@@ -330,7 +329,7 @@ func (s *userServiceImpl) AdminListUsers(ctx context.Context, filter contracts.A
func (s *userServiceImpl) AdminGetUser(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := database.DB(ctx).Table("w_users").
if err := getDB(ctx).Table("w_users").
Select("id, username, nickname, email, avatar_url, is_active, is_admin, bio, phone, gender, website, location, last_login_at, created_at, updated_at").
Where("id = ?", id).
First(&user).Error; err != nil {
@@ -357,7 +356,7 @@ func (s *userServiceImpl) AdminCreateUser(ctx context.Context, req contracts.Adm
}
var count int64
if err := database.DB(ctx).Table("w_users").Where("username = ?", req.Username).Count(&count).Error; err != nil {
if err := getDB(ctx).Table("w_users").Where("username = ?", req.Username).Count(&count).Error; err != nil {
return nil, err
}
if count > 0 {
@@ -365,7 +364,7 @@ func (s *userServiceImpl) AdminCreateUser(ctx context.Context, req contracts.Adm
}
var emailCount int64
if err := database.DB(ctx).Table("w_users").Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
if err := getDB(ctx).Table("w_users").Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
return nil, err
}
if emailCount > 0 {
@@ -404,7 +403,7 @@ func (s *userServiceImpl) AdminCreateUser(ctx context.Context, req contracts.Adm
"created_at": now,
columnUpdatedAt: now,
}
if err := database.DB(ctx).Table("w_users").Create(row).Error; err != nil {
if err := getDB(ctx).Table("w_users").Create(row).Error; err != nil {
return nil, err
}
@@ -428,7 +427,7 @@ func (s *userServiceImpl) AdminUpdateUser(ctx context.Context, currentUserID uin
}
var targetUser contracts.UserDTO
if err := database.DB(ctx).Table("w_users").Where("id = ?", req.ID).First(&targetUser).Error; err != nil {
if err := getDB(ctx).Table("w_users").Where("id = ?", req.ID).First(&targetUser).Error; err != nil {
return err
}
@@ -438,7 +437,7 @@ func (s *userServiceImpl) AdminUpdateUser(ctx context.Context, currentUserID uin
if targetUser.Email != req.Email {
var count int64
if err := database.DB(ctx).Table("w_users").Where("email = ? AND id != ?", req.Email, req.ID).Count(&count).Error; err != nil {
if err := getDB(ctx).Table("w_users").Where("email = ? AND id != ?", req.Email, req.ID).Count(&count).Error; err != nil {
return err
}
if count > 0 {
@@ -469,7 +468,7 @@ func (s *userServiceImpl) AdminUpdateUser(ctx context.Context, currentUserID uin
updates["password"] = hash
}
err := database.DB(ctx).Table("w_users").Where("id = ?", req.ID).Updates(updates).Error
err := getDB(ctx).Table("w_users").Where("id = ?", req.ID).Updates(updates).Error
if err == nil && s.events != nil {
_ = s.events.Emit(ctx, contracts.EventTopicUserUpdated, &targetUser)
}
@@ -481,14 +480,14 @@ func (s *userServiceImpl) AdminUpdateUserStatus(ctx context.Context, id uint64,
ID uint64
IsAdmin bool
}
if err := database.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", id).First(&flags).Error; err != nil {
if err := getDB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", id).First(&flags).Error; err != nil {
return err
}
if !active && flags.IsAdmin {
return errors.New("管理员账号无法被禁用")
}
err := database.DB(ctx).Table("w_users").Where("id = ?", id).Update("is_active", active).Error
err := getDB(ctx).Table("w_users").Where("id = ?", id).Update("is_active", active).Error
if err == nil && s.events != nil {
_ = s.events.Emit(ctx, contracts.EventTopicUserStatusChanged, contracts.UserStatusChangedEvent{
UserID: id,
@@ -506,14 +505,14 @@ func (s *userServiceImpl) AdminDeleteUser(ctx context.Context, currentUserID, ta
ID uint64
IsAdmin bool
}
if err := database.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", targetID).First(&flags).Error; err != nil {
if err := getDB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", targetID).First(&flags).Error; err != nil {
return err
}
if flags.IsAdmin {
return errors.New("管理员账号无法被删除")
}
err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
err := getDB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Table("w_access_tokens").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
return err
}