kra-new/internal/data/api_token.go

145 lines
4.6 KiB
Go

package data
import (
"context"
"errors"
"time"
"kra/internal/biz"
"gorm.io/gorm"
)
type apiTokenRepo struct{ data *Data }
func NewAPITokenRepo(data *Data) biz.APITokenRepo { return &apiTokenRepo{data: data} }
type apiTokenPO struct {
ID uint `gorm:"primaryKey"`
CreatedAt time.Time
UpdatedAt time.Time
DeletedAt gorm.DeletedAt `gorm:"index"`
UserID uint
AuthorityID uint
Token string `gorm:"type:text"`
Status bool
ExpiresAt time.Time
Remark string
}
func (apiTokenPO) TableName() string { return "sys_api_tokens" }
type jwtBlacklistPO struct {
ID uint `gorm:"primaryKey"`
CreatedAt time.Time
UpdatedAt time.Time
DeletedAt gorm.DeletedAt `gorm:"index"`
JWT string `gorm:"column:jwt;type:text"`
}
func (jwtBlacklistPO) TableName() string { return "jwt_blacklists" }
func (r *apiTokenRepo) UserHasAuthority(ctx context.Context, userID, authorityID uint) (*biz.User, bool, error) {
var po userPO
if err := r.data.gormDB.WithContext(ctx).First(&po, userID).Error; err != nil {
return nil, false, err
}
var count int64
err := r.data.gormDB.WithContext(ctx).Model(&userAuthorityPO{}).Where("sys_user_id = ? AND sys_authority_authority_id = ?", userID, authorityID).Count(&count).Error
if err != nil {
return nil, false, err
}
user, err := (&userRepo{data: r.data}).loadUser(ctx, &po)
return user, count > 0 || po.AuthorityID == authorityID, err
}
func (r *apiTokenRepo) CreateAPIToken(ctx context.Context, v *biz.APIToken) error {
po := apiTokenPO{UserID: v.UserID, AuthorityID: v.AuthorityID, Token: v.Token, Status: v.Status, ExpiresAt: v.ExpiresAt, Remark: v.Remark}
if err := r.data.gormDB.WithContext(ctx).Create(&po).Error; err != nil {
return err
}
v.ID = po.ID
v.CreatedAt = po.CreatedAt
return nil
}
func (r *apiTokenRepo) ListAPITokens(ctx context.Context, page, size int, userID uint, status *bool) ([]*biz.APIToken, int64, error) {
db := r.data.gormDB.WithContext(ctx).Model(&apiTokenPO{})
if userID != 0 {
db = db.Where("user_id = ?", userID)
}
if status != nil {
db = db.Where("status = ?", *status)
}
var total int64
if err := db.Count(&total).Error; err != nil {
return nil, 0, err
}
var pos []apiTokenPO
if err := applyPagination(db.Order("created_at desc"), page, size, 100).Find(&pos).Error; err != nil {
return nil, 0, err
}
userIDs := make([]uint, 0, len(pos))
for _, po := range pos {
userIDs = append(userIDs, po.UserID)
}
var userPOs []userPO
if len(userIDs) > 0 {
if err := r.data.gormDB.WithContext(ctx).Where("id IN ?", userIDs).Find(&userPOs).Error; err != nil {
return nil, 0, err
}
}
users := make(map[uint]*biz.User, len(userPOs))
for i := range userPOs {
users[userPOs[i].ID] = baseBizUser(&userPOs[i])
}
out := make([]*biz.APIToken, 0, len(pos))
for _, po := range pos {
v := &biz.APIToken{ID: po.ID, CreatedAt: po.CreatedAt, UpdatedAt: po.UpdatedAt, UserID: po.UserID, AuthorityID: po.AuthorityID, Token: po.Token, Status: po.Status, ExpiresAt: po.ExpiresAt, Remark: po.Remark, User: users[po.UserID]}
out = append(out, v)
}
return out, total, nil
}
func (r *apiTokenRepo) DisableAPIToken(ctx context.Context, id uint) (string, error) {
var po apiTokenPO
if err := r.data.gormDB.WithContext(ctx).First(&po, id).Error; err != nil {
return "", err
}
return po.Token, r.data.gormDB.WithContext(ctx).Model(&po).Update("status", false).Error
}
func (r *apiTokenRepo) DisableAndBlacklistAPIToken(ctx context.Context, id uint) error {
return r.data.gormDB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var po apiTokenPO
if err := tx.First(&po, id).Error; err != nil {
return err
}
// Blacklist the token before disabling it while making both writes atomic.
if err := tx.Create(&jwtBlacklistPO{JWT: po.Token}).Error; err != nil {
return err
}
return tx.Model(&po).Update("status", false).Error
})
}
func (r *apiTokenRepo) BlacklistToken(ctx context.Context, token string) error {
if token == "" {
return errors.New("token is empty")
}
return r.data.gormDB.WithContext(ctx).Create(&jwtBlacklistPO{JWT: token}).Error
}
func (r *apiTokenRepo) IsTokenDisabled(ctx context.Context, token string) (bool, error) {
var blacklistCount int64
if err := r.data.gormDB.WithContext(ctx).Model(&jwtBlacklistPO{}).Where("jwt = ?", token).Count(&blacklistCount).Error; err != nil {
return false, err
}
if blacklistCount > 0 {
return true, nil
}
var po apiTokenPO
result := r.data.gormDB.WithContext(ctx).Where("token = ?", token).Limit(1).Find(&po)
if result.Error != nil {
return false, result.Error
}
if result.RowsAffected == 0 {
return false, nil
}
return !po.Status || time.Now().After(po.ExpiresAt), nil
}