73 lines
2.4 KiB
Go
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
|
|
}
|