286 lines
8.1 KiB
Go
286 lines
8.1 KiB
Go
package service
|
|
|
|
import (
|
|
"errors"
|
|
"fmt"
|
|
"math/rand"
|
|
"photowall/internal/model"
|
|
"photowall/pkg/captcha"
|
|
"photowall/pkg/hash"
|
|
"photowall/pkg/jwt"
|
|
"time"
|
|
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
type AuthService struct {
|
|
db *gorm.DB
|
|
jm *jwt.Manager
|
|
configSvc *ConfigService
|
|
}
|
|
|
|
func NewAuthService(db *gorm.DB, jm *jwt.Manager, configSvc *ConfigService) *AuthService {
|
|
return &AuthService{db: db, jm: jm, configSvc: configSvc}
|
|
}
|
|
|
|
// ============ 请求结构 ============
|
|
|
|
type RegisterReq struct {
|
|
Username string `json:"username" binding:"required,min=3,max=32"`
|
|
Password string `json:"password" binding:"required,min=6,max=64"`
|
|
Nickname string `json:"nickname"`
|
|
Email string `json:"email"`
|
|
CaptchaID string `json:"captcha_id" binding:"required"`
|
|
CaptchaCode string `json:"captcha_code" binding:"required"`
|
|
}
|
|
|
|
type LoginReq struct {
|
|
Username string `json:"username"`
|
|
Password string `json:"password"`
|
|
LoginType string `json:"login_type"` // password / sms
|
|
Phone string `json:"phone"`
|
|
SmsCode string `json:"sms_code"`
|
|
CaptchaID string `json:"captcha_id"`
|
|
CaptchaCode string `json:"captcha_code"`
|
|
}
|
|
|
|
type LoginResp struct {
|
|
Token string `json:"token"`
|
|
User model.User `json:"user"`
|
|
ExpireIn int `json:"expire_in"`
|
|
}
|
|
|
|
type ForgotPasswordReq struct {
|
|
Username string `json:"username" binding:"required"`
|
|
Email string `json:"email" binding:"required"`
|
|
}
|
|
|
|
type ResetPasswordReq struct {
|
|
Username string `json:"username" binding:"required"`
|
|
ResetCode string `json:"reset_code" binding:"required"`
|
|
NewPassword string `json:"new_password" binding:"required,min=6,max=64"`
|
|
}
|
|
|
|
type SendSmsReq struct {
|
|
Phone string `json:"phone" binding:"required"`
|
|
Scene string `json:"scene" binding:"required"` // login / register / reset_password
|
|
}
|
|
|
|
// ============ 注册 ============
|
|
|
|
func (s *AuthService) Register(req *RegisterReq) (*model.User, error) {
|
|
// 校验图形验证码
|
|
if !captcha.Verify(req.CaptchaID, req.CaptchaCode) {
|
|
return nil, errors.New("验证码错误或已过期")
|
|
}
|
|
var count int64
|
|
if err := s.db.Model(&model.User{}).Where("username = ?", req.Username).Count(&count).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
if count > 0 {
|
|
return nil, errors.New("用户名已被注册")
|
|
}
|
|
salt, pwdHash := hash.Password(req.Password)
|
|
nickname := req.Nickname
|
|
if nickname == "" {
|
|
nickname = req.Username
|
|
}
|
|
user := &model.User{
|
|
Username: req.Username,
|
|
PasswordHash: pwdHash,
|
|
Salt: salt,
|
|
Email: req.Email,
|
|
Nickname: nickname,
|
|
Role: "user",
|
|
}
|
|
if err := s.db.Create(user).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
return user, nil
|
|
}
|
|
|
|
// ============ 登录 ============
|
|
|
|
func (s *AuthService) Login(req *LoginReq) (*LoginResp, error) {
|
|
// 图形验证码校验(密码登录必须,短信登录可选)
|
|
if req.LoginType != "sms" {
|
|
if !captcha.Verify(req.CaptchaID, req.CaptchaCode) {
|
|
return nil, errors.New("验证码错误或已过期")
|
|
}
|
|
}
|
|
|
|
var user model.User
|
|
if req.LoginType == "sms" {
|
|
// 短信验证码登录
|
|
if req.Phone == "" || req.SmsCode == "" {
|
|
return nil, errors.New("手机号和验证码不能为空")
|
|
}
|
|
// 校验短信验证码
|
|
var sms model.SmsCode
|
|
if err := s.db.Where("phone = ? AND scene = ? AND used = ?", req.Phone, "login", false).
|
|
Order("id DESC").First(&sms).Error; err != nil {
|
|
return nil, errors.New("验证码错误或已过期")
|
|
}
|
|
if sms.Code != req.SmsCode || time.Now().After(sms.ExpireAt) {
|
|
return nil, errors.New("验证码错误或已过期")
|
|
}
|
|
sms.Used = true
|
|
s.db.Save(&sms)
|
|
// 查找或创建用户
|
|
if err := s.db.Where("username = ?", req.Phone).First(&user).Error; err != nil {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
// 自动注册
|
|
salt, pwdHash := hash.Password(randomCode(8))
|
|
user = model.User{
|
|
Username: req.Phone,
|
|
PasswordHash: pwdHash,
|
|
Salt: salt,
|
|
Nickname: "手机用户" + req.Phone[len(req.Phone)-4:],
|
|
Role: "user",
|
|
}
|
|
s.db.Create(&user)
|
|
} else {
|
|
return nil, err
|
|
}
|
|
}
|
|
} else {
|
|
// 账号密码登录
|
|
if req.Username == "" || req.Password == "" {
|
|
return nil, errors.New("用户名和密码不能为空")
|
|
}
|
|
if err := s.db.Where("username = ?", req.Username).First(&user).Error; err != nil {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return nil, errors.New("用户名或密码错误")
|
|
}
|
|
return nil, err
|
|
}
|
|
if !hash.Verify(user.PasswordHash, user.Salt, req.Password) {
|
|
return nil, errors.New("用户名或密码错误")
|
|
}
|
|
}
|
|
|
|
// 检查账号状态
|
|
if user.Status == "banned" {
|
|
return nil, fmt.Errorf("账号已被封禁:%s", user.BanReason)
|
|
}
|
|
|
|
token, err := s.jm.Generate(user.ID, user.Username, user.Role)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &LoginResp{Token: token, User: user, ExpireIn: s.jm.ExpireHrs}, nil
|
|
}
|
|
|
|
// ============ 忘记密码 ============
|
|
|
|
func (s *AuthService) ForgotPassword(req *ForgotPasswordReq) (string, error) {
|
|
var user model.User
|
|
if err := s.db.Where("username = ? AND email = ?", req.Username, req.Email).First(&user).Error; err != nil {
|
|
return "", errors.New("用户名与邮箱不匹配")
|
|
}
|
|
// 生成6位重置码
|
|
code := fmt.Sprintf("%06d", rand.Intn(1000000))
|
|
reset := &model.PasswordReset{
|
|
UserID: user.ID,
|
|
Code: code,
|
|
ExpireAt: time.Now().Add(30 * time.Minute),
|
|
}
|
|
if err := s.db.Create(reset).Error; err != nil {
|
|
return "", err
|
|
}
|
|
// 开发阶段:重置码直接返回(生产环境应发送邮件)
|
|
// TODO: 接入邮件服务后,通过邮箱发送重置码
|
|
return code, nil
|
|
}
|
|
|
|
// ============ 重置密码 ============
|
|
|
|
func (s *AuthService) ResetPassword(req *ResetPasswordReq) error {
|
|
var user model.User
|
|
if err := s.db.Where("username = ?", req.Username).First(&user).Error; err != nil {
|
|
return errors.New("用户不存在")
|
|
}
|
|
var reset model.PasswordReset
|
|
if err := s.db.Where("user_id = ? AND code = ? AND used = ?", user.ID, req.ResetCode, false).
|
|
Order("id DESC").First(&reset).Error; err != nil {
|
|
return errors.New("重置码错误或已过期")
|
|
}
|
|
if time.Now().After(reset.ExpireAt) {
|
|
return errors.New("重置码已过期")
|
|
}
|
|
salt, pwdHash := hash.Password(req.NewPassword)
|
|
user.PasswordHash = pwdHash
|
|
user.Salt = salt
|
|
reset.Used = true
|
|
return s.db.Transaction(func(tx *gorm.DB) error {
|
|
if err := tx.Save(&user).Error; err != nil {
|
|
return err
|
|
}
|
|
return tx.Save(&reset).Error
|
|
})
|
|
}
|
|
|
|
// ============ 发送短信验证码(预留) ============
|
|
|
|
func (s *AuthService) SendSmsCode(req *SendSmsReq) (string, error) {
|
|
var accessKey, secretKey string
|
|
if s.configSvc != nil {
|
|
_, accessKey, secretKey, _, _ = s.configSvc.GetSmsConfig()
|
|
}
|
|
// 参数未配置时返回提示(开发阶段直接返回验证码)
|
|
if accessKey == "" || secretKey == "" {
|
|
code := fmt.Sprintf("%06d", rand.Intn(1000000))
|
|
sms := &model.SmsCode{
|
|
Phone: req.Phone,
|
|
Code: code,
|
|
Scene: req.Scene,
|
|
ExpireAt: time.Now().Add(5 * time.Minute),
|
|
}
|
|
s.db.Create(sms)
|
|
// 开发阶段直接返回验证码(生产环境应通过短信服务商发送)
|
|
return code, nil
|
|
}
|
|
// TODO: 接入阿里云/腾讯云短信SDK后发送真实短信
|
|
code := fmt.Sprintf("%06d", rand.Intn(1000000))
|
|
sms := &model.SmsCode{
|
|
Phone: req.Phone,
|
|
Code: code,
|
|
Scene: req.Scene,
|
|
ExpireAt: time.Now().Add(5 * time.Minute),
|
|
}
|
|
s.db.Create(sms)
|
|
return code, nil
|
|
}
|
|
|
|
// ============ 微信扫码登录(预留) ============
|
|
|
|
type WechatLoginReq struct {
|
|
Code string `json:"code"` // 微信授权code
|
|
}
|
|
|
|
func (s *AuthService) WechatLogin(req *WechatLoginReq) (*LoginResp, error) {
|
|
var appID, appSecret string
|
|
if s.configSvc != nil {
|
|
appID, appSecret, _ = s.configSvc.GetWechatConfig()
|
|
}
|
|
if appID == "" || appSecret == "" {
|
|
return nil, errors.New("微信登录未配置,请联系管理员")
|
|
}
|
|
// TODO: 接入微信开放平台API
|
|
// 1. 用 code 换取 access_token 和 openid
|
|
// 2. 用 openid 查找或创建用户
|
|
// 3. 签发 JWT
|
|
return nil, errors.New("微信登录功能开发中")
|
|
}
|
|
|
|
// ============ 工具 ============
|
|
|
|
func randomCode(length int) string {
|
|
const chars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
|
|
b := make([]byte, length)
|
|
for i := range b {
|
|
b[i] = chars[rand.Intn(len(chars))]
|
|
}
|
|
return string(b)
|
|
}
|