diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index 08c88c7a..08287794 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -12,8 +12,6 @@ import ( "strconv" "time" - database "Wavelet/plugins/infra/database" - "Wavelet/core/contracts" "Wavelet/pkg/response" "github.com/gin-contrib/sessions" @@ -136,7 +134,7 @@ func Register(c *gin.Context) { return } - gormDB := database.DB(c.Request.Context()) + gormDB := getDB(c.Request.Context()) if err := gormDB.Create(newUser).Error; err != nil { response.AbortBadRequest(c, "创建用户失败: "+err.Error()) return @@ -183,7 +181,7 @@ func ChangePassword(c *gin.Context) { return } - gormDB := database.DB(c.Request.Context()) + gormDB := getDB(c.Request.Context()) _ = gormDB.Save(&user) invalidateUserCache(c.Request.Context(), user.ID) @@ -213,7 +211,7 @@ func UpdateProfile(c *gin.Context) { user.Website = req.Website user.Location = req.Location - gormDB := database.DB(c.Request.Context()) + gormDB := getDB(c.Request.Context()) _ = gormDB.Save(&user) invalidateUserCache(c.Request.Context(), user.ID) @@ -224,7 +222,7 @@ func UpdateProfile(c *gin.Context) { func ListAccessTokens(c *gin.Context) { userID := getUserIDFromSession(c) var tokens []AccessToken - gormDB := database.DB(c.Request.Context()) + gormDB := getDB(c.Request.Context()) _ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error c.JSON(http.StatusOK, response.OK(tokens)) } @@ -262,7 +260,7 @@ func CreateAccessToken(c *gin.Context) { IsAdmin: req.IsAdmin, } - gormDB := database.DB(c.Request.Context()) + gormDB := getDB(c.Request.Context()) if err := gormDB.Create(&token).Error; err != nil { response.AbortInternal(c, "创建令牌失败") return @@ -285,7 +283,7 @@ func DeleteAccessToken(c *gin.Context) { userID := getUserIDFromSession(c) var token AccessToken - gormDB := database.DB(c.Request.Context()) + gormDB := getDB(c.Request.Context()) if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { response.AbortNotFound(c, errTokenNotFound) return @@ -307,7 +305,7 @@ func RotateAccessToken(c *gin.Context) { userID := getUserIDFromSession(c) var token AccessToken - gormDB := database.DB(c.Request.Context()) + gormDB := getDB(c.Request.Context()) if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { response.AbortNotFound(c, errTokenNotFound) return diff --git a/backend/plugins/domain/user/plugin.go b/backend/plugins/domain/user/plugin.go index 11a49a10..d091e61f 100644 --- a/backend/plugins/domain/user/plugin.go +++ b/backend/plugins/domain/user/plugin.go @@ -7,6 +7,7 @@ package user import ( "context" "embed" + "reflect" "Wavelet/core" "Wavelet/core/contracts" @@ -52,6 +53,13 @@ func (p *Plugin) Name() string { return PluginName } +// Inject declares required dependencies for the user domain plugin. +func (p *Plugin) Inject() []reflect.Type { + return []reflect.Type{ + reflect.TypeFor[contracts.DBService](), + } +} + // Manifest returns the plugin metadata. func (p *Plugin) Manifest() core.Manifest { return core.Manifest{ @@ -64,7 +72,12 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { - // 0. Resolve auth service for middleware (via IoC, not direct import) + // 0. Bind DBService from Context + if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { + setDBService(db) + } + + // 0.1 Resolve auth service for middleware (via IoC, not direct import) var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() } var noTokenMW gin.HandlerFunc = func(c *gin.Context) { c.Next() } if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil { diff --git a/backend/plugins/domain/user/plugin_test.go b/backend/plugins/domain/user/plugin_test.go index 95365e37..fc30d164 100644 --- a/backend/plugins/domain/user/plugin_test.go +++ b/backend/plugins/domain/user/plugin_test.go @@ -36,7 +36,10 @@ func setupTestDB(t *testing.T) *gorm.DB { func TestUserPluginUnit(t *testing.T) { ctx := core.NewContext(context.Background()) - _ = setupTestDB(t) + testDB := setupTestDB(t) + + dbPlugin := database.New(database.WithDB(testDB)) + require.NoError(t, dbPlugin.Apply(ctx)) p := user.New() assert.Equal(t, "user", p.Name()) diff --git a/backend/plugins/domain/user/repository.go b/backend/plugins/domain/user/repository.go index 804b1bf7..04ae495c 100644 --- a/backend/plugins/domain/user/repository.go +++ b/backend/plugins/domain/user/repository.go @@ -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 diff --git a/backend/plugins/domain/user/service.go b/backend/plugins/domain/user/service.go index d37ea440..b764570f 100644 --- a/backend/plugins/domain/user/service.go +++ b/backend/plugins/domain/user/service.go @@ -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 }