Files
zsp-project/backend/internal/service/user_service.go
2026-06-03 20:59:39 +08:00

128 lines
3.2 KiB
Go
Executable File

// internal/service/user_service.go
package service
import (
"errors"
"fmt"
"time"
"github.com/golang-jwt/jwt/v4"
"github.com/rovina/zsp-backend/internal/dto"
"github.com/rovina/zsp-backend/internal/model"
"github.com/rovina/zsp-backend/internal/repository"
"golang.org/x/crypto/bcrypt"
)
type Claims struct {
UserID uint `json:"user_id"`
jwt.RegisteredClaims
}
type UserService interface {
Register(req *dto.CreateUserRequest) (*dto.UserResponse, error)
Login(req *dto.LoginRequest) (*dto.UserResponse, string, error)
GetUserByID(id uint) (*dto.UserResponse, error)
}
type userService struct {
userRepo repository.UserRepository
jwtSecret []byte // 改为字节切片
}
func NewUserService(userRepo repository.UserRepository, jwtSecretStr string) (*userService, error) {
// 直接使用字符串作为密钥,不再解析 PEM
return &userService{
userRepo: userRepo,
jwtSecret: []byte(jwtSecretStr),
}, nil
}
func (s *userService) Register(req *dto.CreateUserRequest) (*dto.UserResponse, error) {
// 检查邮箱是否已存在
existingUser, err := s.userRepo.FindByEmail(req.Email)
if err != nil {
return nil, fmt.Errorf("database error: %v", err)
}
if existingUser != nil {
return nil, errors.New("email already exists")
}
existingUser, err = s.userRepo.FindByUsername(req.Username)
if err != nil {
return nil, fmt.Errorf("database error: %v", err)
}
if existingUser != nil {
return nil, errors.New("username already exists")
}
// 加密密码
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(req.Password), bcrypt.DefaultCost)
if err != nil {
return nil, err
}
user := &model.User{
Username: req.Username,
Email: req.Email,
Password: string(hashedPassword),
}
if err := s.userRepo.Create(user); err != nil {
return nil, err
}
return &dto.UserResponse{
ID: user.ID,
Username: user.Username,
Email: user.Email,
}, nil
}
func (s *userService) Login(req *dto.LoginRequest) (*dto.UserResponse, string, error) {
user, err := s.userRepo.FindByEmail(req.Email)
if err != nil {
return nil, "", errors.New("invalid credentials")
}
if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(req.Password)); err != nil {
return nil, "", errors.New("invalid credentials")
}
expirationTime := time.Now().Add(24 * time.Hour)
claims := &Claims{
UserID: user.ID,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(expirationTime),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
Issuer: "zsp-backend",
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
tokenString, err := token.SignedString(s.jwtSecret)
if err != nil {
return nil, "", err
}
return &dto.UserResponse{
ID: user.ID,
Username: user.Username,
Email: user.Email,
}, tokenString, nil
}
func (s *userService) GetUserByID(id uint) (*dto.UserResponse, error) {
user, err := s.userRepo.FindByID(id)
if err != nil {
return nil, err
}
return &dto.UserResponse{
ID: user.ID,
Username: user.Username,
Email: user.Email,
}, nil
}