mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-29 03:26:37 +08:00
fix: allow login during sqlite write pressure
This commit is contained in:
@@ -19,7 +19,15 @@ func activeUserRequired(svc *service.Container) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
u, err := svc.Repo.User.FindByID(c.Request.Context(), userID)
|
||||
if err != nil || u == nil {
|
||||
if err != nil {
|
||||
if service.IsTransientDatabaseLock(err) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"code": 40101, "message": "user not found"})
|
||||
return
|
||||
}
|
||||
if u == nil {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"code": 40101, "message": "user not found"})
|
||||
return
|
||||
}
|
||||
@@ -40,7 +48,19 @@ func activeEmbyUserRequired(svc *service.Container) gin.HandlerFunc {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
userID, _ := uid.(string)
|
||||
u, err := svc.Repo.User.FindByID(c.Request.Context(), userID)
|
||||
if userID == "" || err != nil || u == nil {
|
||||
if userID == "" {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"Code": 40101, "Message": "User not found"})
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
if service.IsTransientDatabaseLock(err) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"Code": 40101, "Message": "User not found"})
|
||||
return
|
||||
}
|
||||
if u == nil {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"Code": 40101, "Message": "User not found"})
|
||||
return
|
||||
}
|
||||
|
||||
@@ -116,6 +116,21 @@ func (h *PermissionHandler) GetMyPermissions(c *gin.Context) {
|
||||
|
||||
perms, err := h.svc.Permissions.Effective(c.Request.Context(), currentUserID)
|
||||
if err != nil {
|
||||
if service.IsTransientDatabaseLock(err) {
|
||||
role := middleware.GetUserRole(c)
|
||||
tier := middleware.GetUserTier(c)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"code": 0,
|
||||
"message": "ok",
|
||||
"data": gin.H{
|
||||
"permissions": service.FallbackPermissions(currentUserID, role),
|
||||
"role": role,
|
||||
"tier": tier,
|
||||
"is_super": role == "admin" || tier == "plus",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
h.log.Error("get my permissions failed", zap.Error(err))
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
|
||||
return
|
||||
|
||||
@@ -47,6 +47,11 @@ func myPermissionsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
row, err := svc.Permissions.Effective(c.Request.Context(), toString(uid))
|
||||
if err != nil {
|
||||
if service.IsTransientDatabaseLock(err) {
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
c.JSON(http.StatusOK, service.FallbackPermissions(toString(uid), toString(role)))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
@@ -865,13 +865,18 @@ type PermissionRepository struct{ db *gorm.DB }
|
||||
|
||||
// Create inserts a new permission record.
|
||||
func (r *PermissionRepository) Create(ctx context.Context, p *model.UserPermission) error {
|
||||
return r.db.WithContext(ctx).Create(p).Error
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Create(p).Error
|
||||
})
|
||||
}
|
||||
|
||||
// FindByUserID returns the permission record for a user, or (nil, nil) when absent.
|
||||
func (r *PermissionRepository) FindByUserID(ctx context.Context, userID string) (*model.UserPermission, error) {
|
||||
var p model.UserPermission
|
||||
err := r.db.WithContext(ctx).Where("user_id = ?", userID).First(&p).Error
|
||||
err := withSQLiteBusyRetry(ctx, func() error {
|
||||
p = model.UserPermission{}
|
||||
return r.db.WithContext(ctx).Where("user_id = ?", userID).First(&p).Error
|
||||
})
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -883,19 +888,25 @@ func (r *PermissionRepository) FindByUserID(ctx context.Context, userID string)
|
||||
|
||||
// Update updates permission fields for a user.
|
||||
func (r *PermissionRepository) Update(ctx context.Context, userID string, updates map[string]bool) error {
|
||||
return r.db.WithContext(ctx).Model(&model.UserPermission{}).
|
||||
Where("user_id = ?", userID).Updates(updates).Error
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Model(&model.UserPermission{}).
|
||||
Where("user_id = ?", userID).Updates(updates).Error
|
||||
})
|
||||
}
|
||||
|
||||
// Upsert creates or updates a permission record.
|
||||
func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermission) error {
|
||||
return r.db.WithContext(ctx).Where("user_id = ?", p.UserID).
|
||||
Assign(*p).FirstOrCreate(p).Error
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Where("user_id = ?", p.UserID).
|
||||
Assign(*p).FirstOrCreate(p).Error
|
||||
})
|
||||
}
|
||||
|
||||
// Delete removes a permission record.
|
||||
func (r *PermissionRepository) Delete(ctx context.Context, userID string) error {
|
||||
return r.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
|
||||
return withSQLiteBusyRetry(ctx, func() error {
|
||||
return r.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
|
||||
})
|
||||
}
|
||||
|
||||
// ─── Refresh Token ───────────────────────────────────────────────────────────
|
||||
|
||||
@@ -13,7 +13,7 @@ func withSQLiteBusyRetry(ctx context.Context, op func() error) error {
|
||||
deadline := time.Now().Add(sqliteBusyRetryMaxElapsed)
|
||||
for {
|
||||
err := op()
|
||||
if !isSQLiteBusyError(err) {
|
||||
if !IsSQLiteBusyError(err) {
|
||||
return err
|
||||
}
|
||||
if ctxErr := ctx.Err(); ctxErr != nil {
|
||||
@@ -35,7 +35,7 @@ func withSQLiteBusyRetry(ctx context.Context, op func() error) error {
|
||||
}
|
||||
}
|
||||
|
||||
func isSQLiteBusyError(err error) bool {
|
||||
func IsSQLiteBusyError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -159,7 +159,7 @@ func (s *AuthService) Login(ctx context.Context, username, password string) (*Lo
|
||||
return nil, ErrInvalidCredentials
|
||||
}
|
||||
// 签发令牌对
|
||||
tokens, err := s.tokenSvc.IssuePair(ctx, u.ID, u.Role, u.Tier)
|
||||
tokens, err := s.tokenSvc.IssuePairBestEffort(ctx, u.ID, u.Role, u.Tier)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -291,6 +291,79 @@ func TestLoginRetriesTransientSQLiteBusy(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginReturnsTokensWhenSQLiteWriteLockPersists(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
cfg := &config.Config{}
|
||||
cfg.App.DataDir = t.TempDir()
|
||||
cfg.Database.DBPath = filepath.Join(cfg.App.DataDir, "busy-login-degraded.db")
|
||||
cfg.Database.WALMode = true
|
||||
cfg.Database.BusyTimeout = 20
|
||||
cfg.Database.MaxOpenConns = 4
|
||||
cfg.Database.MaxIdleConns = 2
|
||||
cfg.Secrets.JWTSecret = "test-secret"
|
||||
log := zap.NewNop()
|
||||
db, err := database.Open(cfg, log)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = sqlDB.Close() }()
|
||||
if err := database.AutoMigrate(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
permissions := NewPermissionService(log, repos)
|
||||
auth := NewAuthService(cfg, log, repos, NewTokenService(cfg, log, repos), permissions)
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte("password"), bcrypt.MinCost)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
user := &model.User{
|
||||
Username: "viewer",
|
||||
PasswordHash: string(hash),
|
||||
Role: "user",
|
||||
Tier: "free",
|
||||
IsActive: true,
|
||||
}
|
||||
if err := repos.User.Create(ctx, user); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tx := repos.DB.Begin()
|
||||
if err := tx.Exec("UPDATE users SET updated_at = updated_at WHERE username = ?", "viewer").Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp, err := auth.Login(ctx, "viewer", "password")
|
||||
if err != nil {
|
||||
t.Fatalf("login should return tokens while refresh token store is delayed: %v", err)
|
||||
}
|
||||
if resp == nil || resp.Tokens == nil || resp.Tokens.AccessToken == "" || resp.Tokens.RefreshToken == "" {
|
||||
t.Fatalf("login returned incomplete token pair: %#v", resp)
|
||||
}
|
||||
|
||||
if err := tx.Rollback().Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wantHash := repository.HashToken(resp.Tokens.RefreshToken)
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for {
|
||||
var count int64
|
||||
if err := repos.DB.Model(&model.RefreshToken{}).Where("token_hash = ?", wantHash).Count(&count).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if count == 1 {
|
||||
return
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatal("delayed refresh token store did not complete")
|
||||
}
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultPermissionsAreViewerOnly(t *testing.T) {
|
||||
perms := DefaultPermissions("user-1")
|
||||
if !perms.CanViewDashboard || !perms.CanPlayMedia || !perms.CanExternalPlayer {
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
package service
|
||||
|
||||
import "github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
|
||||
func IsTransientDatabaseLock(err error) bool {
|
||||
return repository.IsSQLiteBusyError(err)
|
||||
}
|
||||
@@ -78,6 +78,13 @@ func adminGrant(userID string) *model.UserPermission {
|
||||
}
|
||||
}
|
||||
|
||||
func FallbackPermissions(userID, role string) *model.UserPermission {
|
||||
if role == "admin" {
|
||||
return adminGrant(userID)
|
||||
}
|
||||
return DefaultPermissions(userID)
|
||||
}
|
||||
|
||||
// Effective returns the permission set the React UI should consume.
|
||||
// Admins skip the table entirely and get a synthetic all-grant row.
|
||||
func (s *PermissionService) Effective(ctx context.Context, userID string) (*model.UserPermission, error) {
|
||||
|
||||
@@ -25,6 +25,8 @@ const (
|
||||
RefreshTokenLength = 32
|
||||
)
|
||||
|
||||
const loginRefreshTokenStoreTimeout = 750 * time.Millisecond
|
||||
|
||||
// Claims 是 JWT 载荷(复制自 middleware 以避免循环导入)。
|
||||
type Claims struct {
|
||||
UserID string `json:"uid"`
|
||||
@@ -62,6 +64,17 @@ var (
|
||||
|
||||
// IssuePair 为用户签发新的令牌对。
|
||||
func (s *TokenService) IssuePair(ctx context.Context, userID, role, tier string) (*TokenPair, error) {
|
||||
return s.issuePair(ctx, userID, role, tier, false)
|
||||
}
|
||||
|
||||
// IssuePairBestEffort 为登录签发令牌。SQLite 被后台扫描长期写锁占用时,
|
||||
// 登录不能因为 refresh token 暂时无法落库而失败:先返回可用 access token,
|
||||
// 再在后台把 refresh token 补写进库。
|
||||
func (s *TokenService) IssuePairBestEffort(ctx context.Context, userID, role, tier string) (*TokenPair, error) {
|
||||
return s.issuePair(ctx, userID, role, tier, true)
|
||||
}
|
||||
|
||||
func (s *TokenService) issuePair(ctx context.Context, userID, role, tier string, bestEffort bool) (*TokenPair, error) {
|
||||
// 生成 Access Token
|
||||
accessToken, err := s.issueAccessToken(userID, role, tier)
|
||||
if err != nil {
|
||||
@@ -81,11 +94,23 @@ func (s *TokenService) IssuePair(ctx context.Context, userID, role, tier string)
|
||||
TokenHash: tokenHash,
|
||||
ExpiresAt: time.Now().Add(RefreshTokenDuration),
|
||||
}
|
||||
if err := s.repo.RefreshToken.Create(ctx, rt); err != nil {
|
||||
return nil, err
|
||||
storeCtx := ctx
|
||||
cancel := func() {}
|
||||
if bestEffort {
|
||||
storeCtx, cancel = context.WithTimeout(context.Background(), loginRefreshTokenStoreTimeout)
|
||||
}
|
||||
if err := s.repo.RefreshToken.RevokeOldestActiveByUserID(ctx, userID, s.maxActiveRefreshTokens(ctx)); err != nil {
|
||||
s.log.Warn("failed to enforce refresh token session limit", zap.String("user_id", userID), zap.Error(err))
|
||||
err = s.storeRefreshToken(storeCtx, rt)
|
||||
cancel()
|
||||
if err != nil {
|
||||
if !bestEffort {
|
||||
return nil, err
|
||||
}
|
||||
if s.log != nil {
|
||||
s.log.Warn("refresh token store delayed; login will continue",
|
||||
zap.String("user_id", userID),
|
||||
zap.Error(err))
|
||||
}
|
||||
go s.storeRefreshTokenEventually(userID, tokenHash, rt.ExpiresAt)
|
||||
}
|
||||
|
||||
return &TokenPair{
|
||||
@@ -96,6 +121,52 @@ func (s *TokenService) IssuePair(ctx context.Context, userID, role, tier string)
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *TokenService) storeRefreshToken(ctx context.Context, rt *model.RefreshToken) error {
|
||||
if err := s.repo.RefreshToken.Create(ctx, rt); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.repo.RefreshToken.RevokeOldestActiveByUserID(ctx, rt.UserID, s.maxActiveRefreshTokens(ctx)); err != nil && s.log != nil {
|
||||
s.log.Warn("failed to enforce refresh token session limit", zap.String("user_id", rt.UserID), zap.Error(err))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *TokenService) storeRefreshTokenEventually(userID, tokenHash string, expiresAt time.Time) {
|
||||
delay := 500 * time.Millisecond
|
||||
for attempt := 1; attempt <= 30; attempt++ {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
err := s.storeRefreshToken(ctx, &model.RefreshToken{
|
||||
UserID: userID,
|
||||
TokenHash: tokenHash,
|
||||
ExpiresAt: expiresAt,
|
||||
})
|
||||
cancel()
|
||||
if err == nil {
|
||||
return
|
||||
}
|
||||
if !repository.IsSQLiteBusyError(err) && !errors.Is(err, context.DeadlineExceeded) && !errors.Is(err, context.Canceled) {
|
||||
if s.log != nil {
|
||||
s.log.Warn("refresh token delayed store failed permanently", zap.String("user_id", userID), zap.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
if s.log != nil && (attempt == 1 || attempt%10 == 0) {
|
||||
s.log.Warn("refresh token delayed store still waiting",
|
||||
zap.String("user_id", userID),
|
||||
zap.Int("attempt", attempt),
|
||||
zap.Error(err))
|
||||
}
|
||||
timer := time.NewTimer(delay)
|
||||
<-timer.C
|
||||
if delay < 10*time.Second {
|
||||
delay *= 2
|
||||
}
|
||||
}
|
||||
if s.log != nil {
|
||||
s.log.Warn("refresh token delayed store gave up", zap.String("user_id", userID))
|
||||
}
|
||||
}
|
||||
|
||||
func (s *TokenService) maxActiveRefreshTokens(ctx context.Context) int {
|
||||
cfg := loadBotConfig(ctx, s.repo)
|
||||
if cfg.MaxLoggedClients < 1 {
|
||||
|
||||
Reference in New Issue
Block a user