package admin import ( "context" "errors" "server/internal/config" "server/internal/db" "server/internal/db/sqlc" "server/internal/middleware" "server/internal/model/auth" "server/internal/model/enum" "server/internal/model/request" "server/internal/model/response" "server/internal/pkg/cache" "server/internal/pkg/cache/cachekey" "server/internal/pkg/errs" "time" "github.com/jackc/pgx/v5" "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) GetAuthState(ctx context.Context, userID int32) (*sqlc.GetUserAuthStateRow, error) { // 过期时间使用最短窗口的那个 也就是access_token的过期时间 return cache.GetOrSetJSON[*sqlc.GetUserAuthStateRow](ctx, s.cache, cachekey.UserAuthState(userID), s.cfg.JWTConfig.Expire, func() (*sqlc.GetUserAuthStateRow, error) { state, err := s.store.GetUserAuthState(ctx, userID) if err != nil { return nil, err } return &state, nil }) } func (s *AuthService) discardRefreshToken(ctx context.Context, userID int32, hash string) { _ = s.cache.Del(ctx, cachekey.AuthRefresh(hash)) _ = s.cache.SRem(ctx, cachekey.AuthRefreshUser(userID), hash) } func (s *AuthService) Login(ctx context.Context, req request.LoginRequest) (*response.LoginResponse, error) { user, err := s.store.GetUser(ctx, sqlc.GetUserParams{Account: req.Account}) if err != nil { return nil, errs.ErrInvalidCredentials } // 判断用户不为超管 且状态为0表示用户已被禁用 if !auth.IsAdmin(user.ID) && 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, user.TokenVersion) if err != nil { return nil, errs.ErrInternalServer } refreshToken, err := s.jwt.GenerateRefreshToken() if err != nil { return nil, errs.ErrInternalServer } // 哈希 hash := s.jwt.HashRefreshToken(refreshToken) now := time.Now() refreshTokenExp := now.Add(s.cfg.JWTConfig.RefreshExpire) refreshTokenRecord := &auth.RefreshTokenRecord{ UserID: user.ID, TokenVersion: user.TokenVersion, CreatedAt: now, ExpiresAt: refreshTokenExp, } // 存入redis err = s.cache.SetJSON(ctx, cachekey.AuthRefresh(hash), refreshTokenRecord, s.cfg.JWTConfig.RefreshExpire) if err != nil { return nil, errs.ErrInternalServer } // 反向索引 err = s.cache.SAdd(ctx, cachekey.AuthRefreshUser(user.ID), hash) if err != nil { return nil, errs.ErrInternalServer } // 给反向索引设置过期时间 这个过期时间需要覆盖最后一个token的过期时间 所以用最新的就行了 if err = s.cache.Expire(ctx, cachekey.AuthRefreshUser(user.ID), s.cfg.JWTConfig.RefreshExpire); err != nil { return nil, errs.ErrInternalServer } 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 { s.discardRefreshToken(ctx, user.UserID, hash) } 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 { // redis 错误返回500 return nil, errs.ErrInternalServer } if !ok { return nil, errs.ErrInvalidRefreshToken } // 拿到用户信息 authState, err := s.GetAuthState(ctx, record.UserID) if err != nil { if errors.Is(err, pgx.ErrNoRows) { // 用户已被删除 惰性清理 s.discardRefreshToken(ctx, record.UserID, hash) return nil, errs.ErrInvalidRefreshToken } return nil, errs.ErrInternalServer } // 如果不是超级用户 需要判断用户状态 if !auth.IsAdmin(record.UserID) { if enum.Status(authState.Status) == enum.StatusDisabled { // 用户被禁用 清理redis缓存 s.discardRefreshToken(ctx, record.UserID, hash) return nil, errs.ErrInvalidRefreshToken } } // 比对token version 如果不相等 此时 惰性清理掉redis中的缓存 if authState.TokenVersion != record.TokenVersion { s.discardRefreshToken(ctx, record.UserID, hash) return nil, errs.ErrInvalidRefreshToken } // 获取新的access token accessToken, accessTokenExp, err := s.jwt.GenerateAccessToken(record.UserID, authState.TokenVersion) if err != nil { return nil, errs.ErrInternalServer } return &response.LoginResponse{ AccessToken: accessToken, AccessTokenExp: accessTokenExp, }, nil }