Files
backend_v2/internal/repository/refresh_token_repository.go
toom1996 f4f6e02e3c update
2026-08-31 19:59:56 +08:00

73 lines
2.4 KiB
Go

package repository
import (
"context"
"errors"
"time"
"fashionapi/internal/model"
"gorm.io/gorm"
)
// RefreshTokenRepository 刷新令牌数据访问接口。
//
// refresh token 落库是该账号体系能做「续期 + 踢下线」的关键:
// 服务端持有每个 login 会话的行,吊销它即可强行让对应用户重新登录。
type RefreshTokenRepository interface {
// Create 写入一条刷新令牌记录。
Create(ctx context.Context, rt *model.RefreshToken) error
// FindByHash 按 SHA256(raw) 查找,不存在时返回 ErrNotFound。
FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error)
// Revoke 将指定 id 的令牌置为已吊销。
Revoke(ctx context.Context, id uint64) error
// RevokeAllByUser 吊销某用户的所有令牌(全设备登出 / 踢下线),返回受影响行数。
RevokeAllByUser(ctx context.Context, userID uint32) (int64, error)
// DeleteExpired 清理已过期或已吊销的令牌,避免表无限增长。
DeleteExpired(ctx context.Context) (int64, error)
}
type refreshTokenRepository struct {
db *gorm.DB
}
// NewRefreshTokenRepository 创建刷新令牌仓储。
func NewRefreshTokenRepository(db *gorm.DB) RefreshTokenRepository {
return &refreshTokenRepository{db: db}
}
func (r *refreshTokenRepository) Create(ctx context.Context, rt *model.RefreshToken) error {
return r.db.WithContext(ctx).Create(rt).Error
}
func (r *refreshTokenRepository) FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error) {
var rt model.RefreshToken
err := r.db.WithContext(ctx).Where("token_hash = ?", hash).First(&rt).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrNotFound
}
return nil, err
}
return &rt, nil
}
func (r *refreshTokenRepository) Revoke(ctx context.Context, id uint64) error {
return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
Where("id = ?", id).Update("revoked", 1).Error
}
func (r *refreshTokenRepository) RevokeAllByUser(ctx context.Context, userID uint32) (int64, error) {
res := r.db.WithContext(ctx).Model(&model.RefreshToken{}).
Where("user_id = ?", userID).Update("revoked", 1)
return res.RowsAffected, res.Error
}
func (r *refreshTokenRepository) DeleteExpired(ctx context.Context) (int64, error) {
now := uint32(time.Now().Unix())
res := r.db.WithContext(ctx).
Where("expires_at < ? OR revoked = ?", now, 1).
Delete(&model.RefreshToken{})
return res.RowsAffected, res.Error
}