From 5211d14d50eb91416df6653003210f746ad81f48 Mon Sep 17 00:00:00 2001 From: hhs <386998068@qq.com> Date: Sun, 14 Jun 2026 17:14:52 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=9A=E4=B9=89=E5=B9=B6=E5=AE=9E?= =?UTF-8?q?=E7=8E=B0=20AuthService=EF=BC=88=E6=B3=A8=E5=86=8C/=E7=99=BB?= =?UTF-8?q?=E5=BD=95/=E5=88=B7=E6=96=B0/=E7=99=BB=E5=87=BA=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/internal/auth/service.go | 221 +++++++++++++++++++++++++++++++ 1 file changed, 221 insertions(+) create mode 100644 backend/internal/auth/service.go diff --git a/backend/internal/auth/service.go b/backend/internal/auth/service.go new file mode 100644 index 0000000..4f18447 --- /dev/null +++ b/backend/internal/auth/service.go @@ -0,0 +1,221 @@ +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 token(rotation) + _ = 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) +}