181 lines
4.3 KiB
Go
181 lines
4.3 KiB
Go
package admin
|
|
|
|
import (
|
|
"context"
|
|
"server/internal/config"
|
|
"server/internal/db"
|
|
"server/internal/middleware"
|
|
"server/internal/model/auth"
|
|
"server/internal/model/request"
|
|
"server/internal/model/response"
|
|
"server/internal/pkg/cache"
|
|
"server/internal/pkg/cache/cachekey"
|
|
"server/internal/pkg/errs"
|
|
"time"
|
|
|
|
"golang.org/x/crypto/bcrypt"
|
|
)
|
|
|
|
type AuthService struct {
|
|
jwt *middleware.JWTMiddleware
|
|
cache *cache.Caches
|
|
store *db.Store
|
|
cfg *config.Config
|
|
}
|
|
|
|
func NewAuthService(jwt *middleware.JWTMiddleware, cache *cache.Caches, store *db.Store, cfg *config.Config) *AuthService {
|
|
return &AuthService{
|
|
jwt: jwt,
|
|
cache: cache,
|
|
store: store,
|
|
cfg: cfg,
|
|
}
|
|
}
|
|
|
|
// generatePasswordHash 生成密码哈希
|
|
func generatePasswordHash(password string) (string, error) {
|
|
hashed, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
return string(hashed), nil
|
|
}
|
|
|
|
// comparePasswordHash 比较密码哈希
|
|
func comparePasswordHash(passwordHash, inputPassword string) error {
|
|
return bcrypt.CompareHashAndPassword([]byte(passwordHash), []byte(inputPassword))
|
|
}
|
|
|
|
func (s *AuthService) Login(ctx context.Context, req request.LoginRequest) (*response.LoginResponse, error) {
|
|
user, err := s.store.GetUserByAccount(ctx, req.Account)
|
|
|
|
if err != nil {
|
|
return nil, errs.ErrInvalidCredentials
|
|
}
|
|
|
|
// 此处判断如果用户id不为1 且状态为0表示用户已被禁用
|
|
if user.ID != 1 && user.Status == 0 {
|
|
return nil, errs.ErrUserDisabled
|
|
}
|
|
|
|
if err = comparePasswordHash(user.PasswordHash, req.Password); err != nil {
|
|
return nil, errs.ErrInvalidCredentials
|
|
}
|
|
|
|
accessToken, accessTokenExp, err := s.jwt.GenerateAccessToken(user.ID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
refreshToken, err := s.jwt.GenerateRefreshToken()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// 哈希
|
|
hash := s.jwt.HashRefreshToken(refreshToken)
|
|
|
|
now := time.Now()
|
|
|
|
refreshTokenExp := now.Add(s.cfg.JWTConfig.RefreshExpire)
|
|
refreshTokenRecord := &auth.RefreshTokenRecord{
|
|
UserID: user.ID,
|
|
CreatedAt: now,
|
|
ExpiresAt: refreshTokenExp,
|
|
}
|
|
|
|
// 存入redis
|
|
err = s.cache.SetJSON(ctx, cachekey.AuthRefresh(hash), refreshTokenRecord, s.cfg.JWTConfig.RefreshExpire)
|
|
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// 反向索引
|
|
err = s.cache.SAdd(ctx, cachekey.AuthRefreshUser(user.ID), hash)
|
|
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// 给反向索引设置过期时间 这个过期时间需要覆盖最后一个token的过期时间 所以用最新的就行了
|
|
if err = s.cache.Expire(ctx, cachekey.AuthRefreshUser(user.ID), s.cfg.JWTConfig.RefreshExpire); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return &response.LoginResponse{
|
|
AccessToken: accessToken,
|
|
AccessTokenExp: accessTokenExp,
|
|
RefreshToken: refreshToken,
|
|
RefreshTokenExp: refreshTokenExp,
|
|
}, nil
|
|
}
|
|
|
|
func (s *AuthService) Logout(ctx context.Context, refreshToken string) error {
|
|
hash := s.jwt.HashRefreshToken(refreshToken)
|
|
|
|
// 获取用户信息
|
|
user, ok, err := cache.GetJSON[auth.RefreshTokenRecord](ctx, s.cache, cachekey.AuthRefresh(hash))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
if ok {
|
|
// 清理反向索引
|
|
if err = s.cache.SRem(ctx, cachekey.AuthRefreshUser(user.UserID), hash); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
|
|
// 清理当前登录的token
|
|
if err = s.cache.Del(ctx, cachekey.AuthRefresh(hash)); err != nil {
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *AuthService) GetActiveSysUser(ctx context.Context, id int32) error {
|
|
var err error
|
|
|
|
_, err = s.store.GetActiveUserByID(ctx, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *AuthService) RefreshToken(ctx context.Context, refreshToken string) (*response.LoginResponse, error) {
|
|
// 获取哈希
|
|
hash := s.jwt.HashRefreshToken(refreshToken)
|
|
// 从redis中获取数据
|
|
record, ok, err := cache.GetJSON[auth.RefreshTokenRecord](ctx, s.cache, cachekey.AuthRefresh(hash))
|
|
if err != nil {
|
|
return nil, errs.ErrInvalidRefreshToken
|
|
}
|
|
|
|
if !ok {
|
|
return nil, errs.ErrInvalidRefreshToken
|
|
}
|
|
|
|
// 如果不是超级用户 需要判断用户状态
|
|
if record.UserID != 1 {
|
|
err = s.GetActiveSysUser(ctx, record.UserID)
|
|
if err != nil {
|
|
return nil, errs.ErrInvalidRefreshToken
|
|
}
|
|
}
|
|
|
|
// 获取新的access token
|
|
accessToken, accessTokenExp, err := s.jwt.GenerateAccessToken(record.UserID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return &response.LoginResponse{
|
|
AccessToken: accessToken,
|
|
AccessTokenExp: accessTokenExp,
|
|
}, nil
|
|
}
|