diff --git a/backend/go.mod b/backend/go.mod index 069496b..d14bd93 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -28,6 +28,7 @@ require ( github.com/go-playground/validator/v10 v10.20.0 // indirect github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/goccy/go-json v0.10.2 // indirect + github.com/golang-jwt/jwt/v5 v5.3.1 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect diff --git a/backend/go.sum b/backend/go.sum index d4c2b90..38b2a4b 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -37,6 +37,8 @@ github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9L github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU= github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I= +github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= +github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= diff --git a/backend/internal/auth/jwt.go b/backend/internal/auth/jwt.go new file mode 100644 index 0000000..1da1ac0 --- /dev/null +++ b/backend/internal/auth/jwt.go @@ -0,0 +1,111 @@ +package auth + +import ( + "crypto/sha256" + "encoding/hex" + "errors" + "time" + + "github.com/golang-jwt/jwt/v5" + "github.com/google/uuid" +) + +// 自定义错误。 +var ( + ErrInvalidToken = errors.New("invalid or expired token") +) + +// Claims JWT 声明。 +type Claims struct { + UserID string `json:"user_id"` + Username string `json:"username"` + jwt.RegisteredClaims +} + +// TokenManager JWT 令牌管理器。 +type TokenManager struct { + secret []byte + accessTTL time.Duration + refreshTTL time.Duration +} + +// NewTokenManager 创建 TokenManager。 +// secret: JWT 签名密钥;accessTTL/refreshTTL: 令牌有效期。 +func NewTokenManager(secret string, accessTTL, refreshTTL time.Duration) *TokenManager { + return &TokenManager{ + secret: []byte(secret), + accessTTL: accessTTL, + refreshTTL: refreshTTL, + } +} + +// GeneratePair 生成 access + refresh 令牌对。 +func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh string, err error) { + now := time.Now() + + // access token + accessClaims := &Claims{ + UserID: userID, + Username: username, + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)), + IssuedAt: jwt.NewNumericDate(now), + Issuer: "camtalk", + }, + } + accessTkn := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims) + access, err = accessTkn.SignedString(tm.secret) + if err != nil { + return "", "", err + } + + // refresh token(含唯一 token_id 用于 DB 关联) + tokenID := uuid.New().String() + refreshClaims := &Claims{ + UserID: userID, + Username: username, + RegisteredClaims: jwt.RegisteredClaims{ + ID: tokenID, + ExpiresAt: jwt.NewNumericDate(now.Add(tm.refreshTTL)), + IssuedAt: jwt.NewNumericDate(now), + Issuer: "camtalk", + }, + } + refreshTkn := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims) + refresh, err = refreshTkn.SignedString(tm.secret) + return +} + +// ValidateAccess 校验 access token 并返回 Claims。 +func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error) { + return tm.validate(tokenStr) +} + +// ValidateRefresh 校验 refresh token 并返回 Claims。 +func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error) { + return tm.validate(tokenStr) +} + +// validate 解析并校验 JWT。 +func (tm *TokenManager) validate(tokenStr string) (*Claims, error) { + token, err := jwt.ParseWithClaims(tokenStr, &Claims{}, func(t *jwt.Token) (interface{}, error) { + if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, ErrInvalidToken + } + return tm.secret, nil + }) + if err != nil { + return nil, ErrInvalidToken + } + claims, ok := token.Claims.(*Claims) + if !ok || !token.Valid { + return nil, ErrInvalidToken + } + return claims, nil +} + +// HashToken 对 token 做 SHA256 哈希,用于 DB 存储。 +func HashToken(token string) string { + h := sha256.Sum256([]byte(token)) + return hex.EncodeToString(h[:]) +} diff --git a/backend/internal/auth/jwt_test.go b/backend/internal/auth/jwt_test.go new file mode 100644 index 0000000..35b12fa --- /dev/null +++ b/backend/internal/auth/jwt_test.go @@ -0,0 +1,127 @@ +package auth + +import ( + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGeneratePair_ReturnsNonEmptyTokens(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + access, refresh, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + assert.NotEmpty(t, access) + assert.NotEmpty(t, refresh) + assert.NotEqual(t, access, refresh) +} + +func TestValidateAccess_ValidToken(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + claims, err := tm.ValidateAccess(access) + require.NoError(t, err) + assert.Equal(t, "user-123", claims.UserID) + assert.Equal(t, "alice", claims.Username) + assert.Equal(t, "camtalk", claims.Issuer) +} + +func TestValidateRefresh_ValidToken(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + _, refresh, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + claims, err := tm.ValidateRefresh(refresh) + require.NoError(t, err) + assert.Equal(t, "user-123", claims.UserID) + assert.Equal(t, "alice", claims.Username) + assert.NotEmpty(t, claims.ID) // refresh token 应含唯一 ID +} + +func TestValidateAccess_ExpiredToken(t *testing.T) { + // 使用极短的 TTL + tm := NewTokenManager("test-secret-key", -1*time.Second, -1*time.Second) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + _, err = tm.ValidateAccess(access) + assert.ErrorIs(t, err, ErrInvalidToken) +} + +func TestValidateAccess_WrongSecret(t *testing.T) { + tm1 := NewTokenManager("secret-1", 15*time.Minute, 7*24*time.Hour) + tm2 := NewTokenManager("secret-2", 15*time.Minute, 7*24*time.Hour) + + access, _, err := tm1.GeneratePair("user-123", "alice") + require.NoError(t, err) + + _, err = tm2.ValidateAccess(access) + assert.ErrorIs(t, err, ErrInvalidToken) +} + +func TestValidateAccess_InvalidFormat(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + _, err := tm.ValidateAccess("not-a-valid-token") + assert.ErrorIs(t, err, ErrInvalidToken) +} + +func TestValidateAccess_EmptyString(t *testing.T) { + tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour) + + _, err := tm.ValidateAccess("") + assert.ErrorIs(t, err, ErrInvalidToken) +} + +func TestHashToken_Deterministic(t *testing.T) { + hash1 := HashToken("some-token-value") + hash2 := HashToken("some-token-value") + assert.Equal(t, hash1, hash2) + assert.Len(t, hash1, 64) // SHA256 hex = 64 chars +} + +func TestHashToken_DifferentInputsDifferentHashes(t *testing.T) { + hash1 := HashToken("token-a") + hash2 := HashToken("token-b") + assert.NotEqual(t, hash1, hash2) +} + +func TestValidateRefresh_ExpiredToken(t *testing.T) { + tm := NewTokenManager("test-secret-key", -1*time.Minute, -1*time.Minute) + + _, refresh, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + _, err = tm.ValidateRefresh(refresh) + assert.ErrorIs(t, err, ErrInvalidToken) +} + +func TestGeneratePair_TokenClaimsContainCorrectExpiry(t *testing.T) { + accessTTL := 15 * time.Minute + refreshTTL := 7 * 24 * time.Hour + tm := NewTokenManager("test-secret-key", accessTTL, refreshTTL) + + before := time.Now() + access, refresh, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + after := time.Now() + + // 校验 access token 有效期范围 + accessClaims, err := tm.ValidateAccess(access) + require.NoError(t, err) + assert.True(t, accessClaims.ExpiresAt.Time.After(before.Add(accessTTL).Add(-1*time.Second))) + assert.True(t, accessClaims.ExpiresAt.Time.Before(after.Add(accessTTL).Add(1*time.Second))) + + // 校验 refresh token 有效期范围 + refreshClaims, err := tm.ValidateRefresh(refresh) + require.NoError(t, err) + assert.True(t, refreshClaims.ExpiresAt.Time.After(before.Add(refreshTTL).Add(-1*time.Second))) + assert.True(t, refreshClaims.ExpiresAt.Time.Before(after.Add(refreshTTL).Add(1*time.Second))) +} diff --git a/backend/internal/auth/middleware.go b/backend/internal/auth/middleware.go new file mode 100644 index 0000000..a932240 --- /dev/null +++ b/backend/internal/auth/middleware.go @@ -0,0 +1,54 @@ +package auth + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" +) + +// contextKey 用于在 Gin context 中存储 Claims 的 key。 +const ( + ContextKeyUserID = "user_id" + ContextKeyUsername = "username" +) + +// AuthMiddleware 返回 Gin 中间件,从 Authorization: Bearer 提取并校验 JWT。 +// 校验成功后将 user_id 和 username 写入 Gin Context。 +func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc { + return func(c *gin.Context) { + authHeader := c.GetHeader("Authorization") + if authHeader == "" { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{ + "code": "INVALID_TOKEN", + "message": "missing authorization header", + }) + return + } + + // 提取 Bearer token + parts := strings.SplitN(authHeader, " ", 2) + if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{ + "code": "INVALID_TOKEN", + "message": "invalid authorization format", + }) + return + } + + claims, err := tokenMgr.ValidateAccess(parts[1]) + if err != nil { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{ + "code": "INVALID_TOKEN", + "message": "invalid or expired token", + }) + return + } + + // 将用户信息写入 context + c.Set(ContextKeyUserID, claims.UserID) + c.Set(ContextKeyUsername, claims.Username) + + c.Next() + } +} diff --git a/backend/internal/auth/password.go b/backend/internal/auth/password.go new file mode 100644 index 0000000..0664268 --- /dev/null +++ b/backend/internal/auth/password.go @@ -0,0 +1,19 @@ +package auth + +import "golang.org/x/crypto/bcrypt" + +const bcryptCost = 10 + +// HashPassword 使用 bcrypt 对密码进行哈希。 +func HashPassword(password string) (string, error) { + hash, err := bcrypt.GenerateFromPassword([]byte(password), bcryptCost) + if err != nil { + return "", err + } + return string(hash), nil +} + +// CheckPassword 校验密码与哈希是否匹配。 +func CheckPassword(hashedPassword, password string) error { + return bcrypt.CompareHashAndPassword([]byte(hashedPassword), []byte(password)) +} 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) +} diff --git a/backend/internal/auth/service_test.go b/backend/internal/auth/service_test.go new file mode 100644 index 0000000..9fc91a5 --- /dev/null +++ b/backend/internal/auth/service_test.go @@ -0,0 +1,190 @@ +package auth_test + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/hhs/camtalk/internal/auth" + "github.com/hhs/camtalk/internal/store" +) + +// newTestService 创建测试用的 AuthService + MemUserRepository。 +func newTestService(t *testing.T) (auth.Service, *store.MemUserRepository) { + t.Helper() + tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour) + repo := store.NewMemUserRepository() + svc := auth.NewAuthService(tm, repo) + return svc, repo +} + +// --- Register --- + +func TestRegister_Success(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + resp, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "alice", + Password: "password123", + }) + require.NoError(t, err) + assert.NotEmpty(t, resp.User.ID) + assert.Equal(t, "alice", resp.User.Username) + assert.NotEmpty(t, resp.AccessToken) + assert.NotEmpty(t, resp.RefreshToken) +} + +func TestRegister_DuplicateUsername(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + _, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "alice", + Password: "password123", + }) + require.NoError(t, err) + + // 同名再次注册 + _, err = svc.Register(ctx, auth.RegisterRequest{ + Username: "alice", + Password: "another-password", + }) + assert.ErrorIs(t, err, auth.ErrUsernameTaken) +} + +// --- Login --- + +func TestLogin_Success(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + // 先注册 + _, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "bob", + Password: "password123", + }) + require.NoError(t, err) + + // 登录 + resp, err := svc.Login(ctx, auth.LoginRequest{ + Username: "bob", + Password: "password123", + }) + require.NoError(t, err) + assert.Equal(t, "bob", resp.User.Username) + assert.NotEmpty(t, resp.AccessToken) + assert.NotEmpty(t, resp.RefreshToken) +} + +func TestLogin_WrongPassword(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + _, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "bob", + Password: "password123", + }) + require.NoError(t, err) + + _, err = svc.Login(ctx, auth.LoginRequest{ + Username: "bob", + Password: "wrong-password", + }) + assert.ErrorIs(t, err, auth.ErrInvalidCredentials) +} + +func TestLogin_UserNotFound(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + _, err := svc.Login(ctx, auth.LoginRequest{ + Username: "nonexistent", + Password: "password123", + }) + assert.ErrorIs(t, err, auth.ErrInvalidCredentials) +} + +// --- Refresh --- + +func TestRefresh_Success(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + // 注册 + regResp, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "charlie", + Password: "password123", + }) + require.NoError(t, err) + + // 刷新 + refreshResp, err := svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: regResp.RefreshToken, + }) + require.NoError(t, err) + assert.Equal(t, "charlie", refreshResp.User.Username) + assert.NotEmpty(t, refreshResp.AccessToken) + assert.NotEmpty(t, refreshResp.RefreshToken) + // 新旧 refresh token 应不同(rotation) + assert.NotEqual(t, regResp.RefreshToken, refreshResp.RefreshToken) +} + +func TestRefresh_UsedTokenFails(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + regResp, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "charlie", + Password: "password123", + }) + require.NoError(t, err) + + // 第一次刷新 + _, err = svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: regResp.RefreshToken, + }) + require.NoError(t, err) + + // 用旧 token 再次刷新 → 应失败 + _, err = svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: regResp.RefreshToken, + }) + assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed) +} + +func TestRefresh_InvalidToken(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + _, err := svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: "completely-invalid-token", + }) + assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed) +} + +// --- Logout --- + +func TestLogout_Success(t *testing.T) { + svc, _ := newTestService(t) + ctx := context.Background() + + regResp, err := svc.Register(ctx, auth.RegisterRequest{ + Username: "dave", + Password: "password123", + }) + require.NoError(t, err) + + // 登出 + err = svc.Logout(ctx, regResp.User.ID, regResp.RefreshToken) + require.NoError(t, err) + + // 登出后 refresh token 应失效 + _, err = svc.Refresh(ctx, auth.RefreshRequest{ + RefreshToken: regResp.RefreshToken, + }) + assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed) +}