344 lines
8.3 KiB
Go
344 lines
8.3 KiB
Go
package service
|
|
|
|
import (
|
|
"encoding/base64"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
"verification-platform-backend/internal/database"
|
|
"verification-platform-backend/internal/model"
|
|
"verification-platform-backend/pkg/jwt"
|
|
"verification-platform-backend/pkg/utils"
|
|
|
|
"golang.org/x/crypto/bcrypt"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
// AuthService 认证服务
|
|
type AuthService struct{}
|
|
|
|
// NewAuthService 创建认证服务实例
|
|
func NewAuthService() *AuthService {
|
|
return &AuthService{}
|
|
}
|
|
|
|
// Login 用户登录
|
|
func (s *AuthService) Login(username, password, agentPath string) (map[string]interface{}, error) {
|
|
// 查找用户
|
|
var user model.User
|
|
err := database.DB.Where("username = ?", username).First(&user).Error
|
|
if err != nil {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return nil, errors.New("用户名或密码错误")
|
|
}
|
|
return nil, err
|
|
}
|
|
|
|
// 验证密码
|
|
err = bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(password))
|
|
if err != nil {
|
|
return nil, errors.New("用户名或密码错误")
|
|
}
|
|
|
|
// 更新最后登录时间
|
|
user.LastLoginAt = &time.Time{}
|
|
*user.LastLoginAt = time.Now()
|
|
database.DB.Save(&user)
|
|
|
|
// 生成JWT令牌
|
|
token, err := jwt.GenerateToken(user.ID, user.Username, user.Role)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// 返回用户信息和令牌
|
|
return map[string]interface{}{
|
|
"user": map[string]interface{}{
|
|
"id": user.ID,
|
|
"username": user.Username,
|
|
"email": user.Email,
|
|
"avatar": user.Avatar,
|
|
"role": user.Role,
|
|
"status": user.Status,
|
|
"parent_agent_id": user.ParentAgentID,
|
|
},
|
|
"token": token,
|
|
}, nil
|
|
}
|
|
|
|
// Register 用户注册
|
|
func (s *AuthService) Register(username, email, password string) error {
|
|
return s.RegisterWithRole(username, email, "", password, "developer")
|
|
}
|
|
|
|
// RegisterWithRole 用户注册(带角色)
|
|
func (s *AuthService) RegisterWithRole(username, email, phone, password, role string) error {
|
|
if username == "" {
|
|
return errors.New("用户名不能为空")
|
|
}
|
|
if password == "" {
|
|
return errors.New("密码不能为空")
|
|
}
|
|
|
|
var existingUser model.User
|
|
err := database.DB.Where("username = ?", username).First(&existingUser).Error
|
|
if err == nil {
|
|
return errors.New("用户名已存在")
|
|
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return err
|
|
}
|
|
|
|
if email != "" {
|
|
var emailStr string = email
|
|
err = database.DB.Where("email = ?", emailStr).First(&existingUser).Error
|
|
if err == nil {
|
|
return errors.New("邮箱已被注册")
|
|
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return err
|
|
}
|
|
}
|
|
|
|
if phone != "" {
|
|
err = database.DB.Where("device_id = ?", phone).First(&existingUser).Error
|
|
if err == nil {
|
|
return errors.New("手机号已被注册")
|
|
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return err
|
|
}
|
|
}
|
|
|
|
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
if role == "" {
|
|
role = "developer"
|
|
}
|
|
|
|
user := model.User{
|
|
Username: username,
|
|
DeviceID: phone,
|
|
Password: string(hashedPassword),
|
|
Role: role,
|
|
Status: "active",
|
|
CreatedAt: time.Now(),
|
|
UpdatedAt: time.Now(),
|
|
}
|
|
|
|
if email != "" {
|
|
user.Email = &email
|
|
}
|
|
|
|
err = database.DB.Create(&user).Error
|
|
if err != nil {
|
|
// 处理数据库唯一约束错误
|
|
errMsg := err.Error()
|
|
if strings.Contains(errMsg, "UNIQUE constraint failed") {
|
|
if strings.Contains(errMsg, "users.username") {
|
|
return errors.New("用户名已存在")
|
|
}
|
|
if strings.Contains(errMsg, "users.email") {
|
|
return errors.New("邮箱已被注册")
|
|
}
|
|
if strings.Contains(errMsg, "users.device_id") {
|
|
return errors.New("手机号已被注册")
|
|
}
|
|
return errors.New("该账号已被注册")
|
|
}
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// GetProfile 获取用户资料
|
|
func (s *AuthService) GetProfile(userID uint) (*model.User, error) {
|
|
var user model.User
|
|
err := database.DB.First(&user, userID).Error
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// 不返回密码
|
|
user.Password = ""
|
|
return &user, nil
|
|
}
|
|
|
|
// UpdateProfile 更新用户资料
|
|
func (s *AuthService) UpdateProfile(userID uint, email, phone string) error {
|
|
// 检查邮箱是否已被其他用户使用
|
|
if email != "" {
|
|
var existingUser model.User
|
|
err := database.DB.Where("email = ? AND id != ?", email, userID).First(&existingUser).Error
|
|
if err == nil {
|
|
return errors.New("邮箱已被其他用户使用")
|
|
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return err
|
|
}
|
|
}
|
|
|
|
// 更新用户信息
|
|
updates := map[string]interface{}{
|
|
"email": email,
|
|
"phone": phone,
|
|
"updated_at": time.Now(),
|
|
}
|
|
|
|
err := database.DB.Model(&model.User{}).Where("id = ?", userID).Updates(updates).Error
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// ChangePassword 修改密码
|
|
func (s *AuthService) ChangePassword(userID uint, oldPassword, newPassword string) error {
|
|
// 获取用户
|
|
var user model.User
|
|
err := database.DB.First(&user, userID).Error
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// 验证旧密码
|
|
err = bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(oldPassword))
|
|
if err != nil {
|
|
return errors.New("原密码错误")
|
|
}
|
|
|
|
// 加密新密码
|
|
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// 更新密码
|
|
user.Password = string(hashedPassword)
|
|
user.UpdatedAt = time.Now()
|
|
|
|
err = database.DB.Save(&user).Error
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// ForgotPassword 忘记密码
|
|
func (s *AuthService) ForgotPassword(email string) error {
|
|
// 查找用户
|
|
var user model.User
|
|
err := database.DB.Where("email = ?", email).First(&user).Error
|
|
if err != nil {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return errors.New("邮箱未注册")
|
|
}
|
|
return err
|
|
}
|
|
|
|
// 生成重置token
|
|
resetToken := utils.GenerateRandomString(32)
|
|
user.ResetToken = resetToken
|
|
expiresAt := time.Now().Add(24 * time.Hour) // 24小时有效期
|
|
user.ResetTokenExpiresAt = &expiresAt
|
|
|
|
err = database.DB.Save(&user).Error
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// 这里应该发送邮件,简化处理
|
|
_ = resetToken
|
|
|
|
return nil
|
|
}
|
|
|
|
// ResetPassword 重置密码
|
|
func (s *AuthService) ResetPassword(token, newPassword string) error {
|
|
// 查找用户
|
|
var user model.User
|
|
err := database.DB.Where("reset_token = ? AND reset_token_expires_at > ?", token, time.Now()).First(&user).Error
|
|
if err != nil {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return errors.New("重置链接无效或已过期")
|
|
}
|
|
return err
|
|
}
|
|
|
|
// 加密新密码
|
|
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// 更新密码
|
|
user.Password = string(hashedPassword)
|
|
user.ResetToken = ""
|
|
var emptyTime time.Time
|
|
user.ResetTokenExpiresAt = &emptyTime
|
|
user.UpdatedAt = time.Now()
|
|
|
|
err = database.DB.Save(&user).Error
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// GetCaptcha 获取验证码
|
|
func (s *AuthService) GetCaptcha() (map[string]interface{}, error) {
|
|
// 生成随机验证码
|
|
captchaText := utils.GenerateRandomString(6)
|
|
|
|
// 生成验证码ID
|
|
captchaID := utils.GenerateRandomString(16)
|
|
|
|
// 存储验证码到数据库
|
|
captcha := model.Captcha{
|
|
CaptchaID: captchaID,
|
|
Code: captchaText,
|
|
ExpiresAt: time.Now().Add(5 * time.Minute),
|
|
CreatedAt: time.Now(),
|
|
}
|
|
|
|
if err := database.DB.Create(&captcha).Error; err != nil {
|
|
return nil, fmt.Errorf("存储验证码失败")
|
|
}
|
|
|
|
// 创建一个简单的验证码图片
|
|
// 这里使用SVG格式生成一个简单的验证码图片
|
|
svg := generateCaptchaSVG(captchaText)
|
|
|
|
// 将SVG转换为base64
|
|
captchaImage := "data:image/svg+xml;base64," + base64.StdEncoding.EncodeToString([]byte(svg))
|
|
|
|
return map[string]interface{}{
|
|
"captcha_id": captchaID,
|
|
"captcha_image": captchaImage,
|
|
}, nil
|
|
}
|
|
|
|
// generateCaptchaSVG 生成验证码SVG图片
|
|
func generateCaptchaSVG(text string) string {
|
|
width := 120
|
|
height := 40
|
|
|
|
// 生成随机颜色
|
|
colors := []string{"#FF5722", "#4CAF50", "#2196F3", "#FF9800", "#9C27B0"}
|
|
bgColor := "#F5F5F5"
|
|
|
|
// 创建SVG
|
|
svg := fmt.Sprintf(`<svg xmlns="http://www.w3.org/2000/svg" width="%d" height="%d" viewBox="0 0 %d %d">
|
|
<rect width="100%%" height="100%%" fill="%s"/>
|
|
<text x="50%%" y="50%%" text-anchor="middle" dominant-baseline="middle"
|
|
font-family="Arial, sans-serif" font-size="20" font-weight="bold" fill="%s">%s</text>
|
|
</svg>`, width, height, width, height, bgColor, colors[0], text)
|
|
|
|
return svg
|
|
}
|