feat: 实现 Phase 3 JWT + 认证服务 #95
@@ -28,6 +28,7 @@ require (
|
|||||||
github.com/go-playground/validator/v10 v10.20.0 // indirect
|
github.com/go-playground/validator/v10 v10.20.0 // indirect
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
||||||
github.com/goccy/go-json v0.10.2 // 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/pgpassfile v1.0.0 // indirect
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||||
github.com/jackc/puddle/v2 v2.2.2 // 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/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 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
|
||||||
github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I=
|
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 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
|
||||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
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=
|
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