236 lines
5.4 KiB
Go
236 lines
5.4 KiB
Go
/**
|
||
* package service
|
||
* 作用:认证服务,处理登录、注册、验证码等功能
|
||
*/
|
||
package service
|
||
|
||
import (
|
||
"crypto/rand"
|
||
"encoding/hex"
|
||
"errors"
|
||
"fmt"
|
||
"math/big"
|
||
"strings"
|
||
"time"
|
||
"xk-websocket-v2/internal/model"
|
||
"xk-websocket-v2/internal/utils"
|
||
|
||
"github.com/go-redis/redis/v8"
|
||
"gorm.io/gorm"
|
||
)
|
||
|
||
// AuthService 认证服务结构体
|
||
type AuthService struct {
|
||
DB *gorm.DB
|
||
Redis *redis.Client
|
||
}
|
||
|
||
// AuthSvc 全局单例
|
||
var AuthSvc *AuthService
|
||
|
||
/**
|
||
* InitAuthService
|
||
* 功能:初始化认证服务
|
||
*/
|
||
func InitAuthService(db *gorm.DB, rdb *redis.Client) {
|
||
AuthSvc = &AuthService{DB: db, Redis: rdb}
|
||
}
|
||
|
||
/**
|
||
* Login
|
||
* 功能:用户登录验证
|
||
* @param account 账号(邮箱或手机号)
|
||
* @param password 密码
|
||
* @returns 用户信息和错误
|
||
*/
|
||
func (s *AuthService) Login(account, password string) (*model.User, error) {
|
||
var user model.User
|
||
|
||
// 根据邮箱或手机号查询用户
|
||
result := s.DB.Where("email = ? OR phone = ?", account, account).First(&user)
|
||
if result.Error != nil {
|
||
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
||
return nil, errors.New("用户不存在")
|
||
}
|
||
return nil, result.Error
|
||
}
|
||
|
||
// 验证密码
|
||
if !utils.CheckPassword(password, user.Password) {
|
||
return nil, errors.New("密码错误")
|
||
}
|
||
|
||
return &user, nil
|
||
}
|
||
|
||
/**
|
||
* Register
|
||
* 功能:用户注册
|
||
* @param req 注册请求
|
||
* @returns 用户信息和错误
|
||
*/
|
||
func (s *AuthService) Register(req *model.RegisterReq) (*model.User, error) {
|
||
// 验证密码一致性
|
||
if req.Password != req.ConfirmPassword {
|
||
return nil, errors.New("两次密码输入不一致")
|
||
}
|
||
|
||
// 检查邮箱是否已存在
|
||
var existingUser model.User
|
||
if err := s.DB.Where("email = ?", req.Email).First(&existingUser).Error; err == nil {
|
||
return nil, errors.New("邮箱已被注册")
|
||
}
|
||
|
||
// 检查手机号是否已存在
|
||
if err := s.DB.Where("phone = ?", req.Phone).First(&existingUser).Error; err == nil {
|
||
return nil, errors.New("手机号已被注册")
|
||
}
|
||
|
||
// 加密密码
|
||
hashedPassword, err := utils.HashPassword(req.Password)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("密码加密失败: %v", err)
|
||
}
|
||
|
||
// 生成用户ID(使用时间戳+随机数)
|
||
userID := fmt.Sprintf("%d%s", time.Now().UnixNano(), generateRandomString(6))
|
||
|
||
// 生成默认名称(使用邮箱前缀)
|
||
defaultName := req.Email
|
||
if atIndex := strings.Index(req.Email, "@"); atIndex > 0 {
|
||
defaultName = req.Email[:atIndex]
|
||
}
|
||
if len(defaultName) > 20 {
|
||
defaultName = defaultName[:20]
|
||
}
|
||
|
||
// 创建用户
|
||
user := model.User{
|
||
ID: userID,
|
||
Email: req.Email,
|
||
Phone: req.Phone,
|
||
Password: hashedPassword,
|
||
Name: defaultName,
|
||
Avatar: generateAvatar(userID),
|
||
Desc: "",
|
||
Region: "",
|
||
}
|
||
|
||
if err := s.DB.Create(&user).Error; err != nil {
|
||
return nil, fmt.Errorf("创建用户失败: %v", err)
|
||
}
|
||
|
||
// 清除密码字段
|
||
user.Password = ""
|
||
|
||
return &user, nil
|
||
}
|
||
|
||
/**
|
||
* SendEmailCode
|
||
* 功能:发送邮箱验证码(模拟)
|
||
* @param email 邮箱
|
||
* @returns 验证码和错误
|
||
*/
|
||
func (s *AuthService) SendEmailCode(email string) (string, error) {
|
||
// 生成6位验证码
|
||
code := generateCode(6)
|
||
|
||
// 保存验证码到数据库(5分钟过期)
|
||
vc := model.VerificationCode{
|
||
Target: email,
|
||
Code: code,
|
||
Type: "email",
|
||
ExpiresAt: time.Now().Add(5 * time.Minute),
|
||
}
|
||
|
||
if err := s.DB.Create(&vc).Error; err != nil {
|
||
return "", fmt.Errorf("保存验证码失败: %v", err)
|
||
}
|
||
|
||
// 模拟发送(实际应调用邮件服务)
|
||
// 这里直接返回验证码,生产环境应通过邮件发送
|
||
return code, nil
|
||
}
|
||
|
||
/**
|
||
* SendSmsCode
|
||
* 功能:发送短信验证码(模拟)
|
||
* @param phone 手机号
|
||
* @returns 验证码和错误
|
||
*/
|
||
func (s *AuthService) SendSmsCode(phone string) (string, error) {
|
||
// 生成6位验证码
|
||
code := generateCode(6)
|
||
|
||
// 保存验证码到数据库(5分钟过期)
|
||
vc := model.VerificationCode{
|
||
Target: phone,
|
||
Code: code,
|
||
Type: "sms",
|
||
ExpiresAt: time.Now().Add(5 * time.Minute),
|
||
}
|
||
|
||
if err := s.DB.Create(&vc).Error; err != nil {
|
||
return "", fmt.Errorf("保存验证码失败: %v", err)
|
||
}
|
||
|
||
// 模拟发送(实际应调用短信服务)
|
||
// 这里直接返回验证码,生产环境应通过短信发送
|
||
return code, nil
|
||
}
|
||
|
||
/**
|
||
* VerifyCode
|
||
* 功能:验证验证码
|
||
* @param target 邮箱或手机号
|
||
* @param code 验证码
|
||
* @param codeType 类型(email/sms)
|
||
* @returns 是否有效
|
||
*/
|
||
func (s *AuthService) VerifyCode(target, code, codeType string) (bool, error) {
|
||
var vc model.VerificationCode
|
||
|
||
// 查询验证码
|
||
result := s.DB.Where("target = ? AND code = ? AND type = ? AND expires_at > ?",
|
||
target, code, codeType, time.Now()).
|
||
Order("created_at DESC").
|
||
First(&vc)
|
||
|
||
if result.Error != nil {
|
||
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
||
return false, errors.New("验证码无效或已过期")
|
||
}
|
||
return false, result.Error
|
||
}
|
||
|
||
return true, nil
|
||
}
|
||
|
||
// 辅助函数:生成随机字符串
|
||
func generateRandomString(length int) string {
|
||
b := make([]byte, length)
|
||
rand.Read(b)
|
||
return hex.EncodeToString(b)[:length]
|
||
}
|
||
|
||
// 辅助函数:生成验证码
|
||
func generateCode(length int) string {
|
||
code := ""
|
||
for i := 0; i < length; i++ {
|
||
n, _ := rand.Int(rand.Reader, big.NewInt(10))
|
||
code += n.String()
|
||
}
|
||
return code
|
||
}
|
||
|
||
// 辅助函数:生成默认头像
|
||
func generateAvatar(userID string) string {
|
||
// 简单实现:使用用户ID的第一个字符
|
||
if len(userID) > 0 {
|
||
return string(userID[0])
|
||
}
|
||
return "U"
|
||
}
|
||
|