Files
CamTalk/backend/internal/auth/service.go

222 lines
5.5 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 auth
import (
"context"
"errors"
"time"
"github.com/hhs/camtalk/internal/store"
)
// 自定义业务错误。
var (
ErrUsernameTaken = errors.New("username already taken")
ErrInvalidCredentials = errors.New("invalid username or password")
ErrRefreshTokenUsed = errors.New("refresh token has been used or expired")
)
// RegisterRequest 注册请求。
type RegisterRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
// LoginRequest 登录请求。
type LoginRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
// RefreshRequest 刷新令牌请求。
type RefreshRequest struct {
RefreshToken string `json:"refresh_token"`
}
// AuthResponse 认证响应。
type AuthResponse struct {
User UserResponse `json:"user"`
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
}
// UserResponse 用户信息响应。
type UserResponse struct {
ID string `json:"id"`
Username string `json:"username"`
CreatedAt time.Time `json:"created_at"`
}
// Service 认证业务接口。
type Service interface {
Register(ctx context.Context, req RegisterRequest) (*AuthResponse, error)
Login(ctx context.Context, req LoginRequest) (*AuthResponse, error)
Refresh(ctx context.Context, req RefreshRequest) (*AuthResponse, error)
Logout(ctx context.Context, userID, refreshToken string) error
}
// authService 认证服务实现。
type authService struct {
tokenMgr *TokenManager
userRepo store.UserRepository
}
// NewAuthService 创建认证服务。
func NewAuthService(tokenMgr *TokenManager, userRepo store.UserRepository) Service {
return &authService{
tokenMgr: tokenMgr,
userRepo: userRepo,
}
}
// Register 用户注册。
func (s *authService) Register(ctx context.Context, req RegisterRequest) (*AuthResponse, error) {
// 检查用户名是否已存在
_, err := s.userRepo.FindByUsername(ctx, req.Username)
if err == nil {
return nil, ErrUsernameTaken
}
if !errors.Is(err, store.ErrUserNotFound) {
return nil, err
}
// 哈希密码
hash, err := HashPassword(req.Password)
if err != nil {
return nil, err
}
// 创建用户
userID, err := s.userRepo.Create(ctx, req.Username, hash)
if err != nil {
if errors.Is(err, store.ErrUsernameTaken) {
return nil, ErrUsernameTaken
}
return nil, err
}
// 生成令牌对
access, refresh, err := s.tokenMgr.GeneratePair(userID, req.Username)
if err != nil {
return nil, err
}
// 保存 refresh token hash 到 DB
if err := s.saveRefreshToken(ctx, userID, refresh); err != nil {
return nil, err
}
return &AuthResponse{
User: UserResponse{
ID: userID,
Username: req.Username,
},
AccessToken: access,
RefreshToken: refresh,
}, nil
}
// Login 用户登录。
func (s *authService) Login(ctx context.Context, req LoginRequest) (*AuthResponse, error) {
user, err := s.userRepo.FindByUsername(ctx, req.Username)
if err != nil {
if errors.Is(err, store.ErrUserNotFound) {
return nil, ErrInvalidCredentials
}
return nil, err
}
// 校验密码
if err := CheckPassword(user.PasswordHash, req.Password); err != nil {
return nil, ErrInvalidCredentials
}
// 生成令牌对
access, refresh, err := s.tokenMgr.GeneratePair(user.ID, user.Username)
if err != nil {
return nil, err
}
// 保存 refresh token hash
if err := s.saveRefreshToken(ctx, user.ID, refresh); err != nil {
return nil, err
}
return &AuthResponse{
User: UserResponse{
ID: user.ID,
Username: user.Username,
CreatedAt: user.CreatedAt,
},
AccessToken: access,
RefreshToken: refresh,
}, nil
}
// Refresh 刷新令牌Refresh Token Rotation
func (s *authService) Refresh(ctx context.Context, req RefreshRequest) (*AuthResponse, error) {
// 校验 refresh token
claims, err := s.tokenMgr.ValidateRefresh(req.RefreshToken)
if err != nil {
return nil, ErrRefreshTokenUsed
}
tokenHash := HashToken(req.RefreshToken)
// 查找 DB 中的 token hash确认未被使用
userID, err := s.userRepo.FindRefreshToken(ctx, tokenHash)
if err != nil {
if errors.Is(err, store.ErrRefreshTokenNotFound) {
return nil, ErrRefreshTokenUsed
}
return nil, err
}
// 确认 token 归属的用户与 claims 一致
if userID != claims.UserID {
return nil, ErrRefreshTokenUsed
}
// 删除旧 refresh tokenrotation
_ = s.userRepo.DeleteRefreshToken(ctx, tokenHash)
// 生成新的令牌对
access, refresh, err := s.tokenMgr.GeneratePair(claims.UserID, claims.Username)
if err != nil {
return nil, err
}
// 保存新 refresh token
if err := s.saveRefreshToken(ctx, claims.UserID, refresh); err != nil {
return nil, err
}
// 查用户信息
user, err := s.userRepo.FindByID(ctx, claims.UserID)
if err != nil {
return nil, err
}
return &AuthResponse{
User: UserResponse{
ID: user.ID,
Username: user.Username,
CreatedAt: user.CreatedAt,
},
AccessToken: access,
RefreshToken: refresh,
}, nil
}
// Logout 登出,删除 refresh token。
func (s *authService) Logout(ctx context.Context, userID, refreshToken string) error {
tokenHash := HashToken(refreshToken)
return s.userRepo.DeleteRefreshToken(ctx, tokenHash)
}
// saveRefreshToken 将 refresh token 的 hash 保存到 DB。
func (s *authService) saveRefreshToken(ctx context.Context, userID, refreshToken string) error {
tokenHash := HashToken(refreshToken)
expiresAt := time.Now().Add(s.tokenMgr.refreshTTL)
return s.userRepo.SaveRefreshToken(ctx, userID, tokenHash, expiresAt)
}