203 lines
7.2 KiB
Go
203 lines
7.2 KiB
Go
package service
|
||
|
||
import (
|
||
"context"
|
||
"crypto/rand"
|
||
"crypto/sha256"
|
||
"encoding/hex"
|
||
"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)
|
||
}
|
||
|
||
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)
|
||
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)
|
||
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}
|
||
}
|