Files
nl-im-service/internal/service/auth_service.go
2025-12-03 11:00:47 +08:00

236 lines
5.4 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 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"
}