Files
backend_v2/internal/service/auth_service.go
toom1996 74ba700598 update
2026-09-07 00:01:48 +08:00

232 lines
8.3 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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
}