Files
cloudnest/internal/application/auth/service.go
T

159 lines
4.3 KiB
Go

package auth
import (
"fmt"
"math/rand"
"strings"
"time"
"cloudnest/internal/domain/auth/entity"
"cloudnest/internal/domain/auth/repository"
"cloudnest/internal/pkg/crypto"
"cloudnest/internal/pkg/errors"
"cloudnest/internal/pkg/jwt"
)
type EmailVerifier interface {
VerifyCode(email, code string) bool
}
type AuthService struct {
userRepo repository.UserRepository
emailService EmailVerifier
jwtSecret string
jwtExpires int
}
func NewAuthService(userRepo repository.UserRepository, emailService EmailVerifier, jwtSecret string, jwtExpires int) *AuthService {
return &AuthService{
userRepo: userRepo,
emailService: emailService,
jwtSecret: jwtSecret,
jwtExpires: jwtExpires,
}
}
func generateUserCode() string {
const charset = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
r := rand.New(rand.NewSource(time.Now().UnixNano()))
b := make([]byte, 8)
for i := range b {
b[i] = charset[r.Intn(len(charset))]
}
return string(b)
}
func (s *AuthService) Register(username, password, email, nickname, emailCode string) error {
existing, err := s.userRepo.FindByUsername(username)
if err == nil && existing != nil {
return errors.Conflict("用户名已存在")
}
if email != "" {
existing, err := s.userRepo.FindByEmail(email)
if err == nil && existing != nil {
return errors.Conflict("邮箱已被注册")
}
if emailCode == "" {
return errors.BadRequest("邮箱验证码不能为空")
}
if !s.emailService.VerifyCode(email, emailCode) {
return errors.BadRequest("邮箱验证码错误")
}
}
hashed, err := crypto.HashPassword(password)
if err != nil {
return errors.InternalError("密码加密失败", err)
}
userCode := generateUserCode()
for {
existing, err := s.userRepo.FindByUserCode(userCode)
if err != nil || existing == nil {
break
}
userCode = generateUserCode()
}
user := entity.NewUser(username, hashed, userCode, nickname)
if email != "" {
user.Email = email
}
return s.userRepo.Create(user)
}
func (s *AuthService) Login(loginName, password string) (string, *entity.User, error) {
var user *entity.User
var err error
if strings.Contains(loginName, "@") {
user, err = s.userRepo.FindByEmail(loginName)
} else {
user, err = s.userRepo.FindByUsername(loginName)
}
if err != nil {
return "", nil, err
}
if !crypto.CheckPasswordHash(password, user.Password) {
return "", nil, errors.Unauthorized("用户名或密码错误")
}
expiresIn := time.Duration(s.jwtExpires) * time.Hour
token, err := jwt.GenerateToken(user.ID, user.Username, user.UserCode, s.jwtSecret, expiresIn)
return token, user, err
}
func (s *AuthService) CheckStorageLimit(userCode string, fileSize int64) error {
user, err := s.userRepo.FindByUserCode(userCode)
if err != nil {
return err
}
if user.StorageUsed+fileSize > user.StorageLimit {
return errors.BadRequest("存储空间不足,当前已使用 " + formatSize(user.StorageUsed) + ",限制 " + formatSize(user.StorageLimit))
}
return nil
}
func (s *AuthService) AddStorageUsed(userCode string, fileSize int64) error {
return s.userRepo.UpdateStorageUsed(userCode, fileSize)
}
func (s *AuthService) SubtractStorageUsed(userCode string, fileSize int64) error {
return s.userRepo.UpdateStorageUsed(userCode, -fileSize)
}
func (s *AuthService) UpdateProfile(userCode, email, phone, nickname string) error {
return s.userRepo.UpdateProfile(userCode, email, phone, nickname)
}
func (s *AuthService) ChangePassword(userCode, currentPassword, newPassword string) error {
user, err := s.userRepo.FindByUserCode(userCode)
if err != nil {
return err
}
if !crypto.CheckPasswordHash(currentPassword, user.Password) {
return errors.Unauthorized("当前密码不正确")
}
hashed, err := crypto.HashPassword(newPassword)
if err != nil {
return errors.InternalError("密码加密失败", err)
}
return s.userRepo.UpdatePassword(userCode, hashed)
}
func (s *AuthService) GetUserInfo(userCode string) (*entity.User, error) {
return s.userRepo.FindByUserCode(userCode)
}
func formatSize(bytes int64) string {
if bytes < 1024 {
return fmt.Sprintf("%d B", bytes)
} else if bytes < 1024*1024 {
return fmt.Sprintf("%.2f KB", float64(bytes)/1024)
} else if bytes < 1024*1024*1024 {
return fmt.Sprintf("%.2f MB", float64(bytes)/(1024*1024))
}
return fmt.Sprintf("%.2f GB", float64(bytes)/(1024*1024*1024))
}