205 lines
4.8 KiB
Go
205 lines
4.8 KiB
Go
package repositories
|
||
|
||
import (
|
||
"log"
|
||
"time"
|
||
|
||
"github.com/niangaodev/art-code/config"
|
||
"github.com/niangaodev/art-code/models"
|
||
"gorm.io/gorm"
|
||
)
|
||
|
||
// GetUserByUsername 根据用户名获取用户
|
||
func GetUserByUsername(username string) (*models.User, error) {
|
||
var user models.User
|
||
err := config.DB.Model(&models.User{}).
|
||
Select("users.*, COALESCE(roles.name, users.role) as role").
|
||
Joins("LEFT JOIN roles ON users.role_id = roles.id").
|
||
Where("users.username = ? AND users.deleted_at = ?", username, 0).
|
||
First(&user).Error
|
||
|
||
if err != nil {
|
||
if err == gorm.ErrRecordNotFound {
|
||
return nil, nil
|
||
}
|
||
log.Printf("Error getting user by username: %v", err)
|
||
return nil, err
|
||
}
|
||
|
||
return &user, nil
|
||
}
|
||
|
||
// GetUserByID 根据ID获取用户
|
||
func GetUserByID(id uint) (*models.User, error) {
|
||
var user models.User
|
||
err := config.DB.Model(&models.User{}).
|
||
Select("users.*, COALESCE(roles.name, users.role) as role").
|
||
Joins("LEFT JOIN roles ON users.role_id = roles.id").
|
||
Where("users.id = ? AND users.deleted_at = ?", id, 0).
|
||
First(&user).Error
|
||
|
||
if err != nil {
|
||
if err == gorm.ErrRecordNotFound {
|
||
return nil, nil
|
||
}
|
||
log.Printf("Error getting user by ID: %v", err)
|
||
return nil, err
|
||
}
|
||
|
||
return &user, nil
|
||
}
|
||
|
||
// GetUsers 获取所有用户 (分页)
|
||
func GetUsers(page, pageSize int) ([]models.User, int, error) {
|
||
offset := (page - 1) * pageSize
|
||
|
||
var users []models.User
|
||
var total int64
|
||
|
||
// 获取总数
|
||
err := config.DB.Model(&models.User{}).
|
||
Where("deleted_at = ?", 0).
|
||
Count(&total).Error
|
||
if err != nil {
|
||
log.Printf("Error getting user count: %v", err)
|
||
return nil, 0, err
|
||
}
|
||
|
||
// 获取用户列表
|
||
err = config.DB.Model(&models.User{}).
|
||
Select("users.*, COALESCE(roles.name, users.role) as role").
|
||
Joins("LEFT JOIN roles ON users.role_id = roles.id").
|
||
Where("users.deleted_at = ?", 0).
|
||
Order("users.created_at DESC").
|
||
Limit(pageSize).
|
||
Offset(offset).
|
||
Find(&users).Error
|
||
|
||
if err != nil {
|
||
log.Printf("Error querying users: %v", err)
|
||
return nil, 0, err
|
||
}
|
||
|
||
return users, int(total), nil
|
||
}
|
||
|
||
// CreateUser 创建用户
|
||
func CreateUser(user *models.User) error {
|
||
// 如果提供了RoleID,确保它有效。如果没有RoleID但有Role name,尝试查找RoleID
|
||
if user.RoleID == 0 && user.Role != "" {
|
||
role, err := GetRoleByName(user.Role)
|
||
if err == nil && role != nil {
|
||
user.RoleID = role.ID
|
||
}
|
||
}
|
||
|
||
err := config.DB.Create(user).Error
|
||
if err != nil {
|
||
log.Printf("Error creating user: %v", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// UpdateUser 更新用户
|
||
func UpdateUser(user *models.User) error {
|
||
// 同样尝试解析RoleID
|
||
if user.RoleID == 0 && user.Role != "" {
|
||
role, err := GetRoleByName(user.Role)
|
||
if err == nil && role != nil {
|
||
user.RoleID = role.ID
|
||
}
|
||
}
|
||
|
||
updates := map[string]interface{}{
|
||
"username": user.Username,
|
||
"email": user.Email,
|
||
"role": user.Role,
|
||
"is_active": user.IsActive,
|
||
"updated_at": time.Now().Unix(),
|
||
}
|
||
|
||
if user.RoleID != 0 {
|
||
updates["role_id"] = user.RoleID
|
||
} else {
|
||
updates["role_id"] = nil
|
||
}
|
||
|
||
err := config.DB.Model(&models.User{}).
|
||
Where("id = ? AND deleted_at = ?", user.ID, 0).
|
||
Updates(updates).Error
|
||
if err != nil {
|
||
log.Printf("Error updating user: %v", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// UpdateUserPassword 更新用户密码
|
||
func UpdateUserPassword(id uint, passwordHash string) error {
|
||
err := config.DB.Model(&models.User{}).
|
||
Where("id = ? AND deleted_at = ?", id, 0).
|
||
Updates(map[string]interface{}{
|
||
"password_hash": passwordHash,
|
||
"updated_at": time.Now().Unix(),
|
||
}).Error
|
||
if err != nil {
|
||
log.Printf("Error updating user password: %v", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// DeleteUser 删除用户 (Soft Delete)
|
||
func DeleteUser(id uint) error {
|
||
err := config.DB.Model(&models.User{}).
|
||
Where("id = ?", id).
|
||
Update("deleted_at", time.Now().Unix()).Error
|
||
if err != nil {
|
||
log.Printf("Error deleting user: %v", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// GetUserCount 获取用户总数
|
||
func GetUserCount() (int, error) {
|
||
var count int64
|
||
err := config.DB.Model(&models.User{}).
|
||
Where("deleted_at = ?", 0).
|
||
Count(&count).Error
|
||
if err != nil {
|
||
log.Printf("Error getting user count: %v", err)
|
||
return 0, err
|
||
}
|
||
|
||
return int(count), nil
|
||
}
|
||
|
||
// BuildUserResponse 构建用户响应
|
||
func BuildUserResponse(user *models.User) *models.UserResponse {
|
||
return &models.UserResponse{
|
||
ID: user.ID,
|
||
Username: user.Username,
|
||
Email: user.Email,
|
||
RoleID: user.RoleID,
|
||
Role: user.Role,
|
||
IsActive: user.IsActive,
|
||
CreatedAt: time.Unix(user.CreatedAt, 0).Format("2006-01-02 15:04:05"),
|
||
UpdatedAt: time.Unix(user.UpdatedAt, 0).Format("2006-01-02 15:04:05"),
|
||
}
|
||
}
|
||
|
||
// BuildUsersResponse 构建用户列表响应
|
||
func BuildUsersResponse(users []models.User) []models.UserResponse {
|
||
var responses []models.UserResponse
|
||
for _, user := range users {
|
||
responses = append(responses, *BuildUserResponse(&user))
|
||
}
|
||
return responses
|
||
}
|