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) +}