mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-12 02:06:37 +08:00
refactor(user): decouple repository from direct database infra import
This commit is contained in:
@@ -12,8 +12,6 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
@@ -136,7 +134,7 @@ func Register(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
gormDB := database.DB(c.Request.Context())
|
gormDB := getDB(c.Request.Context())
|
||||||
if err := gormDB.Create(newUser).Error; err != nil {
|
if err := gormDB.Create(newUser).Error; err != nil {
|
||||||
response.AbortBadRequest(c, "创建用户失败: "+err.Error())
|
response.AbortBadRequest(c, "创建用户失败: "+err.Error())
|
||||||
return
|
return
|
||||||
@@ -183,7 +181,7 @@ func ChangePassword(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
gormDB := database.DB(c.Request.Context())
|
gormDB := getDB(c.Request.Context())
|
||||||
_ = gormDB.Save(&user)
|
_ = gormDB.Save(&user)
|
||||||
invalidateUserCache(c.Request.Context(), user.ID)
|
invalidateUserCache(c.Request.Context(), user.ID)
|
||||||
|
|
||||||
@@ -213,7 +211,7 @@ func UpdateProfile(c *gin.Context) {
|
|||||||
user.Website = req.Website
|
user.Website = req.Website
|
||||||
user.Location = req.Location
|
user.Location = req.Location
|
||||||
|
|
||||||
gormDB := database.DB(c.Request.Context())
|
gormDB := getDB(c.Request.Context())
|
||||||
_ = gormDB.Save(&user)
|
_ = gormDB.Save(&user)
|
||||||
invalidateUserCache(c.Request.Context(), user.ID)
|
invalidateUserCache(c.Request.Context(), user.ID)
|
||||||
|
|
||||||
@@ -224,7 +222,7 @@ func UpdateProfile(c *gin.Context) {
|
|||||||
func ListAccessTokens(c *gin.Context) {
|
func ListAccessTokens(c *gin.Context) {
|
||||||
userID := getUserIDFromSession(c)
|
userID := getUserIDFromSession(c)
|
||||||
var tokens []AccessToken
|
var tokens []AccessToken
|
||||||
gormDB := database.DB(c.Request.Context())
|
gormDB := getDB(c.Request.Context())
|
||||||
_ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error
|
_ = gormDB.Where("user_id = ?", userID).Find(&tokens).Error
|
||||||
c.JSON(http.StatusOK, response.OK(tokens))
|
c.JSON(http.StatusOK, response.OK(tokens))
|
||||||
}
|
}
|
||||||
@@ -262,7 +260,7 @@ func CreateAccessToken(c *gin.Context) {
|
|||||||
IsAdmin: req.IsAdmin,
|
IsAdmin: req.IsAdmin,
|
||||||
}
|
}
|
||||||
|
|
||||||
gormDB := database.DB(c.Request.Context())
|
gormDB := getDB(c.Request.Context())
|
||||||
if err := gormDB.Create(&token).Error; err != nil {
|
if err := gormDB.Create(&token).Error; err != nil {
|
||||||
response.AbortInternal(c, "创建令牌失败")
|
response.AbortInternal(c, "创建令牌失败")
|
||||||
return
|
return
|
||||||
@@ -285,7 +283,7 @@ func DeleteAccessToken(c *gin.Context) {
|
|||||||
|
|
||||||
userID := getUserIDFromSession(c)
|
userID := getUserIDFromSession(c)
|
||||||
var token AccessToken
|
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 {
|
if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil {
|
||||||
response.AbortNotFound(c, errTokenNotFound)
|
response.AbortNotFound(c, errTokenNotFound)
|
||||||
return
|
return
|
||||||
@@ -307,7 +305,7 @@ func RotateAccessToken(c *gin.Context) {
|
|||||||
|
|
||||||
userID := getUserIDFromSession(c)
|
userID := getUserIDFromSession(c)
|
||||||
var token AccessToken
|
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 {
|
if err := gormDB.Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil {
|
||||||
response.AbortNotFound(c, errTokenNotFound)
|
response.AbortNotFound(c, errTokenNotFound)
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ package user
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"embed"
|
"embed"
|
||||||
|
"reflect"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
@@ -52,6 +53,13 @@ func (p *Plugin) Name() string {
|
|||||||
return PluginName
|
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.
|
// Manifest returns the plugin metadata.
|
||||||
func (p *Plugin) Manifest() core.Manifest {
|
func (p *Plugin) Manifest() core.Manifest {
|
||||||
return 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.
|
// Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
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 loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||||
var noTokenMW 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 {
|
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||||
|
|||||||
@@ -36,7 +36,10 @@ func setupTestDB(t *testing.T) *gorm.DB {
|
|||||||
|
|
||||||
func TestUserPluginUnit(t *testing.T) {
|
func TestUserPluginUnit(t *testing.T) {
|
||||||
ctx := core.NewContext(context.Background())
|
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()
|
p := user.New()
|
||||||
assert.Equal(t, "user", p.Name())
|
assert.Equal(t, "user", p.Name())
|
||||||
|
|||||||
@@ -5,18 +5,47 @@ package user
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"Wavelet/core"
|
||||||
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
"gorm.io/gorm"
|
"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 获取用户
|
// GetUserByID 通过 ID 获取用户
|
||||||
func GetUserByID(ctx context.Context, id uint64) (*User, error) {
|
func GetUserByID(ctx context.Context, id uint64) (*User, error) {
|
||||||
var u User
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &u, nil
|
return &u, nil
|
||||||
@@ -25,7 +54,7 @@ func GetUserByID(ctx context.Context, id uint64) (*User, error) {
|
|||||||
// GetUserByUsername 通过用户名获取用户
|
// GetUserByUsername 通过用户名获取用户
|
||||||
func GetUserByUsername(ctx context.Context, username string) (*User, error) {
|
func GetUserByUsername(ctx context.Context, username string) (*User, error) {
|
||||||
var u User
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &u, nil
|
return &u, nil
|
||||||
@@ -34,7 +63,7 @@ func GetUserByUsername(ctx context.Context, username string) (*User, error) {
|
|||||||
// GetUserByEmail 通过邮箱获取用户
|
// GetUserByEmail 通过邮箱获取用户
|
||||||
func GetUserByEmail(ctx context.Context, email string) (*User, error) {
|
func GetUserByEmail(ctx context.Context, email string) (*User, error) {
|
||||||
var u User
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &u, nil
|
return &u, nil
|
||||||
@@ -42,17 +71,17 @@ func GetUserByEmail(ctx context.Context, email string) (*User, error) {
|
|||||||
|
|
||||||
// CreateUser 创建用户
|
// CreateUser 创建用户
|
||||||
func CreateUser(ctx context.Context, u *User) error {
|
func CreateUser(ctx context.Context, u *User) error {
|
||||||
return database.DB(ctx).Create(u).Error
|
return getDB(ctx).Create(u).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateUser 更新用户
|
// UpdateUser 更新用户
|
||||||
func UpdateUser(ctx context.Context, u *User) error {
|
func UpdateUser(ctx context.Context, u *User) error {
|
||||||
return database.DB(ctx).Save(u).Error
|
return getDB(ctx).Save(u).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListUsers 分页查询用户
|
// ListUsers 分页查询用户
|
||||||
func ListUsers(ctx context.Context, page, pageSize int, keyword string) ([]*User, int64, error) {
|
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 != "" {
|
if keyword != "" {
|
||||||
escaped := util.EscapeLike(keyword)
|
escaped := util.EscapeLike(keyword)
|
||||||
db = db.Where("username LIKE ? ESCAPE '\\' OR nickname LIKE ? ESCAPE '\\' OR email LIKE ? ESCAPE '\\'", "%"+escaped+"%", "%"+escaped+"%", "%"+escaped+"%")
|
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 查询访问令牌
|
// GetAccessTokenByHash 通过 Hash 查询访问令牌
|
||||||
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*AccessToken, error) {
|
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*AccessToken, error) {
|
||||||
var token AccessToken
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &token, nil
|
return &token, nil
|
||||||
@@ -90,7 +119,7 @@ type AdminUserListFilter struct {
|
|||||||
|
|
||||||
// ListAdminUsers 获取后台管理用户列表
|
// ListAdminUsers 获取后台管理用户列表
|
||||||
func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []User, error) {
|
func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []User, error) {
|
||||||
query := database.DB(ctx).Model(&User{})
|
query := getDB(ctx).Model(&User{})
|
||||||
if filter.Username != "" {
|
if filter.Username != "" {
|
||||||
escaped := util.EscapeLike(strings.ToLower(filter.Username))
|
escaped := util.EscapeLike(strings.ToLower(filter.Username))
|
||||||
query = query.Where("LOWER(username) LIKE ? ESCAPE '\\'", "%"+escaped+"%")
|
query = query.Where("LOWER(username) LIKE ? ESCAPE '\\'", "%"+escaped+"%")
|
||||||
@@ -116,13 +145,13 @@ func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []U
|
|||||||
|
|
||||||
// UpdateUserActive 更新用户激活状态
|
// UpdateUserActive 更新用户激活状态
|
||||||
func UpdateUserActive(ctx context.Context, id uint64, active bool) error {
|
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 获取处于激活状态的用户
|
// GetActiveUserByID 获取处于激活状态的用户
|
||||||
func GetActiveUserByID(ctx context.Context, id uint64) (*User, error) {
|
func GetActiveUserByID(ctx context.Context, id uint64) (*User, error) {
|
||||||
var u User
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &u, nil
|
return &u, nil
|
||||||
@@ -130,7 +159,7 @@ func GetActiveUserByID(ctx context.Context, id uint64) (*User, error) {
|
|||||||
|
|
||||||
// DeleteUserWithRelations 删除用户及其级联关系
|
// DeleteUserWithRelations 删除用户及其级联关系
|
||||||
func DeleteUserWithRelations(ctx context.Context, id uint64) error {
|
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 {
|
if err := tx.Where("user_id = ?", id).Delete(&AccessToken{}).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -141,7 +170,7 @@ func DeleteUserWithRelations(ctx context.Context, id uint64) error {
|
|||||||
// GetFirstAdminUser 获取第一个管理员用户
|
// GetFirstAdminUser 获取第一个管理员用户
|
||||||
func GetFirstAdminUser(ctx context.Context) (*User, error) {
|
func GetFirstAdminUser(ctx context.Context) (*User, error) {
|
||||||
var u User
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &u, nil
|
return &u, nil
|
||||||
@@ -151,7 +180,7 @@ func GetFirstAdminUser(ctx context.Context) (*User, error) {
|
|||||||
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
|
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
|
||||||
var usernames []string
|
var usernames []string
|
||||||
escaped := util.EscapeLike(strings.ToLower(base))
|
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+"%").
|
Where("LOWER(username) LIKE ? ESCAPE '\\'", escaped+"%").
|
||||||
Pluck("username", &usernames).Error; err != nil {
|
Pluck("username", &usernames).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ import (
|
|||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
pkgu "Wavelet/pkg/util"
|
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) {
|
func (s *userServiceImpl) GetUserByEmail(ctx context.Context, email string) (*contracts.UserDTO, error) {
|
||||||
var u User
|
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 nil, err
|
||||||
}
|
}
|
||||||
return toUserDTO(&u), nil
|
return toUserDTO(&u), nil
|
||||||
@@ -143,7 +142,7 @@ func (s *userServiceImpl) UpdateProfile(ctx context.Context, id uint64, req cont
|
|||||||
}
|
}
|
||||||
updates[columnUpdatedAt] = time.Now()
|
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
|
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 {
|
func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPassword, newPassword string) error {
|
||||||
var user User
|
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -164,7 +163,7 @@ func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPass
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return database.DB(ctx).Model(&User{}).Where("id = ?", id).
|
return getDB(ctx).Model(&User{}).Where("id = ?", id).
|
||||||
Updates(map[string]any{
|
Updates(map[string]any{
|
||||||
"password": user.Password,
|
"password": user.Password,
|
||||||
columnUpdatedAt: time.Now(),
|
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 {
|
func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, password string) bool {
|
||||||
var user User
|
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)
|
pkgu.DummyCheckPassword(password)
|
||||||
return false
|
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 {
|
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{
|
Updates(map[string]any{
|
||||||
"last_login_at": time.Now(),
|
"last_login_at": time.Now(),
|
||||||
columnUpdatedAt: 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 {
|
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) {
|
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) {
|
func (s *userServiceImpl) CountUsers(ctx context.Context) (int64, error) {
|
||||||
var count int64
|
var count int64
|
||||||
err := database.DB(ctx).Model(&User{}).Count(&count).Error
|
err := getDB(ctx).Model(&User{}).Count(&count).Error
|
||||||
return count, err
|
return count, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *userServiceImpl) CountActiveUsers(ctx context.Context) (int64, error) {
|
func (s *userServiceImpl) CountActiveUsers(ctx context.Context) (int64, error) {
|
||||||
var count int64
|
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
|
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) {
|
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 {
|
if filter.UserID != nil {
|
||||||
query = query.Where("id = ?", *filter.UserID)
|
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) {
|
func (s *userServiceImpl) AdminGetUser(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
|
||||||
var user contracts.UserDTO
|
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").
|
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).
|
Where("id = ?", id).
|
||||||
First(&user).Error; err != nil {
|
First(&user).Error; err != nil {
|
||||||
@@ -357,7 +356,7 @@ func (s *userServiceImpl) AdminCreateUser(ctx context.Context, req contracts.Adm
|
|||||||
}
|
}
|
||||||
|
|
||||||
var count int64
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
if count > 0 {
|
if count > 0 {
|
||||||
@@ -365,7 +364,7 @@ func (s *userServiceImpl) AdminCreateUser(ctx context.Context, req contracts.Adm
|
|||||||
}
|
}
|
||||||
|
|
||||||
var emailCount int64
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
if emailCount > 0 {
|
if emailCount > 0 {
|
||||||
@@ -404,7 +403,7 @@ func (s *userServiceImpl) AdminCreateUser(ctx context.Context, req contracts.Adm
|
|||||||
"created_at": now,
|
"created_at": now,
|
||||||
columnUpdatedAt: 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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -428,7 +427,7 @@ func (s *userServiceImpl) AdminUpdateUser(ctx context.Context, currentUserID uin
|
|||||||
}
|
}
|
||||||
|
|
||||||
var targetUser contracts.UserDTO
|
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -438,7 +437,7 @@ func (s *userServiceImpl) AdminUpdateUser(ctx context.Context, currentUserID uin
|
|||||||
|
|
||||||
if targetUser.Email != req.Email {
|
if targetUser.Email != req.Email {
|
||||||
var count int64
|
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
|
return err
|
||||||
}
|
}
|
||||||
if count > 0 {
|
if count > 0 {
|
||||||
@@ -469,7 +468,7 @@ func (s *userServiceImpl) AdminUpdateUser(ctx context.Context, currentUserID uin
|
|||||||
updates["password"] = hash
|
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 {
|
if err == nil && s.events != nil {
|
||||||
_ = s.events.Emit(ctx, contracts.EventTopicUserUpdated, &targetUser)
|
_ = s.events.Emit(ctx, contracts.EventTopicUserUpdated, &targetUser)
|
||||||
}
|
}
|
||||||
@@ -481,14 +480,14 @@ func (s *userServiceImpl) AdminUpdateUserStatus(ctx context.Context, id uint64,
|
|||||||
ID uint64
|
ID uint64
|
||||||
IsAdmin bool
|
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
|
return err
|
||||||
}
|
}
|
||||||
if !active && flags.IsAdmin {
|
if !active && flags.IsAdmin {
|
||||||
return errors.New("管理员账号无法被禁用")
|
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 {
|
if err == nil && s.events != nil {
|
||||||
_ = s.events.Emit(ctx, contracts.EventTopicUserStatusChanged, contracts.UserStatusChangedEvent{
|
_ = s.events.Emit(ctx, contracts.EventTopicUserStatusChanged, contracts.UserStatusChangedEvent{
|
||||||
UserID: id,
|
UserID: id,
|
||||||
@@ -506,14 +505,14 @@ func (s *userServiceImpl) AdminDeleteUser(ctx context.Context, currentUserID, ta
|
|||||||
ID uint64
|
ID uint64
|
||||||
IsAdmin bool
|
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
|
return err
|
||||||
}
|
}
|
||||||
if flags.IsAdmin {
|
if flags.IsAdmin {
|
||||||
return errors.New("管理员账号无法被删除")
|
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 {
|
if err := tx.Table("w_access_tokens").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user