Files

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)
}