package service import ( "context" "crypto/rand" "crypto/sha256" "encoding/hex" "errors" "time" "fashionapi/internal/dto" "fashionapi/internal/model" "fashionapi/internal/pkg/jwt" "fashionapi/internal/repository" "golang.org/x/crypto/bcrypt" ) // AuthService 账号体系业务接口。 // // 采用「access token(短命,无状态 JWT)+ refresh token(长命,落库可吊销)」双令牌模式: // - Login 签发两者;access 用于日常请求,refresh 仅在续期时使用。 // - Refresh 用 refresh token 换发新 access(refresh 复用、不轮换,足够内部工具)。 // - Logout 吊销当前 refresh;RevokeAllByRefresh 吊销该用户全部 refresh(踢下线)。 // // 由于 refresh 落库,服务端得以单方面作废会话——这是纯无状态 JWT 做不到的。 type AuthService interface { Me(ctx context.Context, userID uint32) (*dto.UserPayload, error) // Login 校验账号密码,签发 access + refresh;并吊销该用户既有会话(同账号互斥登录)。account 可为用户名或邮箱。 Login(ctx context.Context, account, password string) (access, refresh string, expiresIn int, user *dto.UserPayload, err error) // Refresh 用 refresh token 换取新的 access token。 Refresh(ctx context.Context, refreshToken string) (access string, expiresIn int, err error) // Logout 吊销当前 refresh token(单设备登出),幂等。 Logout(ctx context.Context, refreshToken string) error // RevokeAllByRefresh 根据 refresh token 找到所属用户,吊销其全部 refresh(踢下线)。 RevokeAllByRefresh(ctx context.Context, refreshToken string) (int64, error) // ListUsers 列出全量用户(后台管理提级 VIP 用)。 ListUsers(ctx context.Context) ([]dto.UserPayload, error) // SetTier 设置用户等级(free/vip),供后台把某账号提级为 VIP。 SetTier(ctx context.Context, id uint32, tier string) error } type authService struct { users repository.UserRepository refresh repository.RefreshTokenRepository jwt *jwt.Manager accessTTLSeconds int refreshExpireHours int } // NewAuthService 创建账号服务。 func NewAuthService( users repository.UserRepository, refresh repository.RefreshTokenRepository, jwtManager *jwt.Manager, accessExpireHours, refreshExpireHours int, ) AuthService { return &authService{ users: users, refresh: refresh, jwt: jwtManager, accessTTLSeconds: accessExpireHours * 3600, refreshExpireHours: refreshExpireHours, } } // Me 用令牌中的 uid 取回最新用户信息。 func (s *authService) Me(ctx context.Context, userID uint32) (*dto.UserPayload, error) { user, err := s.users.FindByID(ctx, userID) if err != nil { // 令牌有效但用户已不存在,同样视为未授权 return nil, unauthorized("unauthorized") } p := toUserPayload(*user) return &p, nil } // Login 校验账号密码并签发双令牌。account 可以是用户名或邮箱。 // // 登录即顶旧:签发前先吊销该用户全部既有 refresh 会话,保证同一账号同一时刻 // 只有本条新会话有效(互斥登录)。账号不存在与密码错误统一返回 unauthorized // (避免泄露账号是否存在)。 func (s *authService) Login(ctx context.Context, account, password string) (string, string, int, *dto.UserPayload, error) { user, err := s.users.FindByAccount(ctx, account) if err != nil { return "", "", 0, nil, unauthorized("invalid credentials") } if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil { return "", "", 0, nil, unauthorized("invalid credentials") } // 同账号互斥登录:本次登录先吊销该用户所有既有 refresh 会话, // 仅保留随后新签发的这一条,从而实现「同一账号同时只能有一个人登录」。 if _, err := s.refresh.RevokeAllByUser(ctx, user.ID); err != nil { return "", "", 0, nil, internalErr("revoke stale sessions: " + err.Error()) } access, err := s.jwt.Generate(user.ID, user.Username, user.Email, user.Tier) if err != nil { return "", "", 0, nil, err } refresh, err := s.issueRefresh(ctx, user.ID) if err != nil { return "", "", 0, nil, err } p := toUserPayload(*user) return access, refresh, s.accessTTLSeconds, &p, nil } // Refresh 用 refresh token 换取新的 access token(refresh 复用)。 func (s *authService) Refresh(ctx context.Context, refreshToken string) (string, int, error) { if refreshToken == "" { return "", 0, unauthorized("missing refresh token") } rt, err := s.lookupRefresh(ctx, refreshToken) if err != nil { return "", 0, err } user, err := s.users.FindByID(ctx, rt.UserID) if err != nil { return "", 0, unauthorized("invalid refresh token") } access, err := s.jwt.Generate(user.ID, user.Username, user.Email, user.Tier) if err != nil { return "", 0, err } return access, s.accessTTLSeconds, nil } // Logout 吊销当前 refresh token;令牌不存在/已失效时视为幂等成功。 func (s *authService) Logout(ctx context.Context, refreshToken string) error { if refreshToken == "" { return nil } rt, err := s.findRefresh(ctx, refreshToken) if err != nil { return nil } return s.refresh.Revoke(ctx, rt.ID) } // RevokeAllByRefresh 找到 refresh 所属用户,吊销其全部 refresh(踢下线)。 func (s *authService) RevokeAllByRefresh(ctx context.Context, refreshToken string) (int64, error) { if refreshToken == "" { return 0, unauthorized("missing refresh token") } rt, err := s.findRefresh(ctx, refreshToken) if err != nil { return 0, unauthorized("invalid refresh token") } return s.refresh.RevokeAllByUser(ctx, rt.UserID) } // issueRefresh 生成随机 refresh token,存哈希入表,返回原始串。 func (s *authService) issueRefresh(ctx context.Context, userID uint32) (string, error) { raw, hash, err := newRefreshToken() if err != nil { return "", internalErr("generate refresh token: " + err.Error()) } now := uint32(time.Now().Unix()) rt := &model.RefreshToken{ UserID: userID, TokenHash: hash, ExpiresAt: now + uint32(s.refreshExpireHours)*3600, Revoked: 0, CreatedAt: now, } if err := s.refresh.Create(ctx, rt); err != nil { return "", internalErr("store refresh token: " + err.Error()) } return raw, nil } // findRefresh 按原始串取哈希查库(用于登出/吊销),找不到返回 ErrNotFound。 func (s *authService) findRefresh(ctx context.Context, refreshToken string) (*model.RefreshToken, error) { hash := hashToken(refreshToken) return s.refresh.FindByHash(ctx, hash) } // lookupRefresh 同 findRefresh,但额外校验有效性(未吊销且未过期)。 func (s *authService) lookupRefresh(ctx context.Context, refreshToken string) (*model.RefreshToken, error) { rt, err := s.findRefresh(ctx, refreshToken) if err != nil { return nil, unauthorized("invalid refresh token") } if !rt.IsValid() { return nil, unauthorized("refresh token expired or revoked") } return rt, nil } // newRefreshToken 生成 32 字节随机串及其 SHA256 哈希(仅哈希落库)。 func newRefreshToken() (raw, hash string, err error) { b := make([]byte, 32) if _, err = rand.Read(b); err != nil { return "", "", err } raw = hex.EncodeToString(b) sum := sha256.Sum256([]byte(raw)) return raw, hex.EncodeToString(sum[:]), nil } // hashToken 计算 refresh token 的存储哈希。 func hashToken(raw string) string { sum := sha256.Sum256([]byte(raw)) return hex.EncodeToString(sum[:]) } func toUserPayload(u model.User) dto.UserPayload { return dto.UserPayload{ID: u.ID, Username: u.Username, Email: u.Email, Tier: u.Tier} } // ListUsers 列出全量用户(后台管理提级 VIP 用)。 func (s *authService) ListUsers(ctx context.Context) ([]dto.UserPayload, error) { users, err := s.users.List(ctx) if err != nil { return nil, internalErr(err.Error()) } out := make([]dto.UserPayload, 0, len(users)) for _, u := range users { out = append(out, toUserPayload(u)) } return out, nil } // SetTier 设置用户等级(free/vip),供后台把某账号提级为 VIP。 func (s *authService) SetTier(ctx context.Context, id uint32, tier string) error { if tier != "free" && tier != "vip" { return errors.New("tier must be free or vip") } if err := s.users.SetTier(ctx, id, tier); err != nil { return internalErr(err.Error()) } return nil }