Merge pull request 'feat: 添加 PostgreSQL 服务并挂载数据库迁移脚本' #101
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
111
backend/internal/auth/jwt.go
Normal file
111
backend/internal/auth/jwt.go
Normal file
@@ -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[:])
|
||||
}
|
||||
127
backend/internal/auth/jwt_test.go
Normal file
127
backend/internal/auth/jwt_test.go
Normal file
@@ -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)))
|
||||
}
|
||||
54
backend/internal/auth/middleware.go
Normal file
54
backend/internal/auth/middleware.go
Normal file
@@ -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 <token> 提取并校验 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()
|
||||
}
|
||||
}
|
||||
19
backend/internal/auth/password.go
Normal file
19
backend/internal/auth/password.go
Normal file
@@ -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))
|
||||
}
|
||||
221
backend/internal/auth/service.go
Normal file
221
backend/internal/auth/service.go
Normal file
@@ -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)
|
||||
}
|
||||
190
backend/internal/auth/service_test.go
Normal file
190
backend/internal/auth/service_test.go
Normal file
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user