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
+44 -15
View File
@@ -5,18 +5,47 @@ package user
import (
"context"
"strings"
"sync"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/pkg/util"
database "Wavelet/plugins/infra/database"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
func setDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
}
func getDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s := c.DB(); s != nil {
return s.DB(ctx)
}
}
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
// GetUserByID 通过 ID 获取用户
func GetUserByID(ctx context.Context, id uint64) (*User, error) {
var u User
if err := database.DB(ctx).First(&u, id).Error; err != nil {
if err := getDB(ctx).First(&u, id).Error; err != nil {
return nil, err
}
return &u, nil
@@ -25,7 +54,7 @@ func GetUserByID(ctx context.Context, id uint64) (*User, error) {
// GetUserByUsername 通过用户名获取用户
func GetUserByUsername(ctx context.Context, username string) (*User, error) {
var u User
if err := database.DB(ctx).Where("username = ?", username).First(&u).Error; err != nil {
if err := getDB(ctx).Where("username = ?", username).First(&u).Error; err != nil {
return nil, err
}
return &u, nil
@@ -34,7 +63,7 @@ func GetUserByUsername(ctx context.Context, username string) (*User, error) {
// GetUserByEmail 通过邮箱获取用户
func GetUserByEmail(ctx context.Context, email string) (*User, 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 &u, nil
@@ -42,17 +71,17 @@ func GetUserByEmail(ctx context.Context, email string) (*User, error) {
// CreateUser 创建用户
func CreateUser(ctx context.Context, u *User) error {
return database.DB(ctx).Create(u).Error
return getDB(ctx).Create(u).Error
}
// UpdateUser 更新用户
func UpdateUser(ctx context.Context, u *User) error {
return database.DB(ctx).Save(u).Error
return getDB(ctx).Save(u).Error
}
// ListUsers 分页查询用户
func ListUsers(ctx context.Context, page, pageSize int, keyword string) ([]*User, int64, error) {
db := database.DB(ctx).Model(&User{})
db := getDB(ctx).Model(&User{})
if keyword != "" {
escaped := util.EscapeLike(keyword)
db = db.Where("username LIKE ? ESCAPE '\\' OR nickname LIKE ? ESCAPE '\\' OR email LIKE ? ESCAPE '\\'", "%"+escaped+"%", "%"+escaped+"%", "%"+escaped+"%")
@@ -74,7 +103,7 @@ func ListUsers(ctx context.Context, page, pageSize int, keyword string) ([]*User
// GetAccessTokenByHash 通过 Hash 查询访问令牌
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*AccessToken, error) {
var token AccessToken
if err := database.DB(ctx).Where("token_hash = ?", tokenHash).First(&token).Error; err != nil {
if err := getDB(ctx).Where("token_hash = ?", tokenHash).First(&token).Error; err != nil {
return nil, err
}
return &token, nil
@@ -90,7 +119,7 @@ type AdminUserListFilter struct {
// ListAdminUsers 获取后台管理用户列表
func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []User, error) {
query := database.DB(ctx).Model(&User{})
query := getDB(ctx).Model(&User{})
if filter.Username != "" {
escaped := util.EscapeLike(strings.ToLower(filter.Username))
query = query.Where("LOWER(username) LIKE ? ESCAPE '\\'", "%"+escaped+"%")
@@ -116,13 +145,13 @@ func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []U
// UpdateUserActive 更新用户激活状态
func UpdateUserActive(ctx context.Context, id uint64, active bool) error {
return database.DB(ctx).Model(&User{}).Where("id = ?", id).Update("is_active", active).Error
return getDB(ctx).Model(&User{}).Where("id = ?", id).Update("is_active", active).Error
}
// GetActiveUserByID 获取处于激活状态的用户
func GetActiveUserByID(ctx context.Context, id uint64) (*User, error) {
var u User
if err := database.DB(ctx).Where("id = ? AND is_active = ?", id, true).First(&u).Error; err != nil {
if err := getDB(ctx).Where("id = ? AND is_active = ?", id, true).First(&u).Error; err != nil {
return nil, err
}
return &u, nil
@@ -130,7 +159,7 @@ func GetActiveUserByID(ctx context.Context, id uint64) (*User, error) {
// DeleteUserWithRelations 删除用户及其级联关系
func DeleteUserWithRelations(ctx context.Context, id uint64) error {
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
return getDB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("user_id = ?", id).Delete(&AccessToken{}).Error; err != nil {
return err
}
@@ -141,7 +170,7 @@ func DeleteUserWithRelations(ctx context.Context, id uint64) error {
// GetFirstAdminUser 获取第一个管理员用户
func GetFirstAdminUser(ctx context.Context) (*User, error) {
var u User
if err := database.DB(ctx).Where("is_admin = ?", true).Order("id ASC").First(&u).Error; err != nil {
if err := getDB(ctx).Where("is_admin = ?", true).Order("id ASC").First(&u).Error; err != nil {
return nil, err
}
return &u, nil
@@ -151,7 +180,7 @@ func GetFirstAdminUser(ctx context.Context) (*User, error) {
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
var usernames []string
escaped := util.EscapeLike(strings.ToLower(base))
if err := database.DB(ctx).Model(&User{}).
if err := getDB(ctx).Model(&User{}).
Where("LOWER(username) LIKE ? ESCAPE '\\'", escaped+"%").
Pluck("username", &usernames).Error; err != nil {
return nil, err