package service 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.GetSysUserByAccount(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.GetActiveSysUserByID(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 }