Files
nl-blogs/server/repositories/user_repository.go
2026-01-15 13:51:44 +08:00

294 lines
6.1 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}