294 lines
6.1 KiB
Go
294 lines
6.1 KiB
Go
package repositories
|
||
|
||
import (
|
||
"database/sql"
|
||
"log"
|
||
|
||
"github.com/niangaodev/art-code/config"
|
||
"github.com/niangaodev/art-code/models"
|
||
)
|
||
|
||
// GetUserByUsername 根据用户名获取用户
|
||
func GetUserByUsername(username string) (*models.User, error) {
|
||
query := `
|
||
SELECT u.id, u.username, u.email, u.password_hash, u.role_id, COALESCE(r.name, u.role), u.is_active, u.created_at, u.updated_at
|
||
FROM users u
|
||
LEFT JOIN roles r ON u.role_id = r.id
|
||
WHERE u.username = ?
|
||
`
|
||
row := config.DB.QueryRow(query, username)
|
||
|
||
var user models.User
|
||
var roleID sql.NullInt64 // Handle nullable role_id
|
||
var roleName sql.NullString // Handle nullable role name
|
||
|
||
if err := row.Scan(
|
||
&user.ID,
|
||
&user.Username,
|
||
&user.Email,
|
||
&user.PasswordHash,
|
||
&roleID,
|
||
&roleName,
|
||
&user.IsActive,
|
||
&user.CreatedAt,
|
||
&user.UpdatedAt,
|
||
); err != nil {
|
||
if err == sql.ErrNoRows {
|
||
return nil, nil
|
||
}
|
||
log.Printf("Error scanning user by username: %v", err)
|
||
return nil, err
|
||
}
|
||
|
||
if roleID.Valid {
|
||
user.RoleID = uint(roleID.Int64)
|
||
}
|
||
if roleName.Valid {
|
||
user.Role = roleName.String
|
||
}
|
||
|
||
return &user, nil
|
||
}
|
||
|
||
// GetUserByID 根据ID获取用户
|
||
func GetUserByID(id uint) (*models.User, error) {
|
||
query := `
|
||
SELECT u.id, u.username, u.email, u.password_hash, u.role_id, COALESCE(r.name, u.role), u.is_active, u.created_at, u.updated_at
|
||
FROM users u
|
||
LEFT JOIN roles r ON u.role_id = r.id
|
||
WHERE u.id = ?
|
||
`
|
||
row := config.DB.QueryRow(query, id)
|
||
|
||
var user models.User
|
||
var roleID sql.NullInt64
|
||
var roleName sql.NullString
|
||
|
||
if err := row.Scan(
|
||
&user.ID,
|
||
&user.Username,
|
||
&user.Email,
|
||
&user.PasswordHash,
|
||
&roleID,
|
||
&roleName,
|
||
&user.IsActive,
|
||
&user.CreatedAt,
|
||
&user.UpdatedAt,
|
||
); err != nil {
|
||
if err == sql.ErrNoRows {
|
||
return nil, nil
|
||
}
|
||
log.Printf("Error scanning user by ID: %v", err)
|
||
return nil, err
|
||
}
|
||
|
||
if roleID.Valid {
|
||
user.RoleID = uint(roleID.Int64)
|
||
}
|
||
if roleName.Valid {
|
||
user.Role = roleName.String
|
||
}
|
||
|
||
return &user, nil
|
||
}
|
||
|
||
// GetUsers 获取所有用户
|
||
func GetUsers() ([]models.User, error) {
|
||
query := `
|
||
SELECT u.id, u.username, u.email, u.password_hash, u.role_id, COALESCE(r.name, u.role), u.is_active, u.created_at, u.updated_at
|
||
FROM users u
|
||
LEFT JOIN roles r ON u.role_id = r.id
|
||
ORDER BY u.created_at DESC
|
||
`
|
||
rows, err := config.DB.Query(query)
|
||
if err != nil {
|
||
log.Printf("Error querying users: %v", err)
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
|
||
var users []models.User
|
||
for rows.Next() {
|
||
var user models.User
|
||
var roleID sql.NullInt64
|
||
var roleName sql.NullString
|
||
|
||
if err := rows.Scan(
|
||
&user.ID,
|
||
&user.Username,
|
||
&user.Email,
|
||
&user.PasswordHash,
|
||
&roleID,
|
||
&roleName,
|
||
&user.IsActive,
|
||
&user.CreatedAt,
|
||
&user.UpdatedAt,
|
||
); err != nil {
|
||
log.Printf("Error scanning user: %v", err)
|
||
continue
|
||
}
|
||
|
||
if roleID.Valid {
|
||
user.RoleID = uint(roleID.Int64)
|
||
}
|
||
if roleName.Valid {
|
||
user.Role = roleName.String
|
||
}
|
||
|
||
users = append(users, user)
|
||
}
|
||
|
||
return users, 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
|
||
}
|
||
}
|
||
|
||
query := `
|
||
INSERT INTO users (username, email, password_hash, role_id, role, is_active, created_at, updated_at)
|
||
VALUES (?, ?, ?, ?, ?, ?, NOW(), NOW())
|
||
`
|
||
|
||
var roleID interface{}
|
||
if user.RoleID != 0 {
|
||
roleID = user.RoleID
|
||
} else {
|
||
roleID = nil
|
||
}
|
||
|
||
result, err := config.DB.Exec(
|
||
query,
|
||
user.Username,
|
||
user.Email,
|
||
user.PasswordHash,
|
||
roleID,
|
||
user.Role, // Fallback legacy column
|
||
user.IsActive,
|
||
)
|
||
if err != nil {
|
||
log.Printf("Error creating user: %v", err)
|
||
return err
|
||
}
|
||
|
||
// 获取自增ID
|
||
id, err := result.LastInsertId()
|
||
if err != nil {
|
||
log.Printf("Error getting last insert ID: %v", err)
|
||
return err
|
||
}
|
||
user.ID = uint(id)
|
||
|
||
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
|
||
}
|
||
}
|
||
|
||
query := `
|
||
UPDATE users SET username = ?, email = ?, role_id = ?, role = ?, is_active = ?, updated_at = NOW()
|
||
WHERE id = ?
|
||
`
|
||
|
||
var roleID interface{}
|
||
if user.RoleID != 0 {
|
||
roleID = user.RoleID
|
||
} else {
|
||
roleID = nil
|
||
}
|
||
|
||
_, err := config.DB.Exec(
|
||
query,
|
||
user.Username,
|
||
user.Email,
|
||
roleID,
|
||
user.Role,
|
||
user.IsActive,
|
||
user.ID,
|
||
)
|
||
if err != nil {
|
||
log.Printf("Error updating user: %v", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// UpdateUserPassword 更新用户密码
|
||
func UpdateUserPassword(id uint, passwordHash string) error {
|
||
query := `
|
||
UPDATE users SET password_hash = ?, updated_at = NOW()
|
||
WHERE id = ?
|
||
`
|
||
_, err := config.DB.Exec(query, passwordHash, id)
|
||
if err != nil {
|
||
log.Printf("Error updating user password: %v", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// DeleteUser 删除用户
|
||
func DeleteUser(id uint) error {
|
||
query := "DELETE FROM users WHERE id = ?"
|
||
_, err := config.DB.Exec(query, id)
|
||
if err != nil {
|
||
log.Printf("Error deleting user: %v", err)
|
||
return err
|
||
}
|
||
|
||
return nil
|
||
}
|
||
|
||
// GetUserCount 获取用户总数
|
||
func GetUserCount() (int, error) {
|
||
var count int
|
||
query := "SELECT COUNT(*) FROM users"
|
||
row := config.DB.QueryRow(query)
|
||
|
||
err := row.Scan(&count)
|
||
if err != nil {
|
||
log.Printf("Error getting user count: %v", err)
|
||
return 0, err
|
||
}
|
||
|
||
return 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: user.CreatedAt.Format("2006-01-02 15:04:05"),
|
||
UpdatedAt: user.UpdatedAt.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
|
||
}
|