From 52d75171b2113e88a95d2192b5dfdbb1491e2ae2 Mon Sep 17 00:00:00 2001 From: ShukeBta Date: Thu, 11 Jun 2026 20:06:55 +0800 Subject: [PATCH] fix: allow login during sqlite write pressure --- internal/handler/active_user.go | 24 ++++++- internal/handler/permission_handler.go | 15 +++++ internal/handler/permissions.go | 5 ++ internal/repository/repository.go | 25 +++++-- internal/repository/sqlite_busy_retry.go | 4 +- internal/service/auth.go | 2 +- internal/service/auth_user_limits_test.go | 73 +++++++++++++++++++++ internal/service/db_errors.go | 7 ++ internal/service/permission.go | 7 ++ internal/service/token_svc.go | 79 +++++++++++++++++++++-- 10 files changed, 225 insertions(+), 16 deletions(-) create mode 100644 internal/service/db_errors.go diff --git a/internal/handler/active_user.go b/internal/handler/active_user.go index 19dd9e4..1f70fd0 100644 --- a/internal/handler/active_user.go +++ b/internal/handler/active_user.go @@ -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 } diff --git a/internal/handler/permission_handler.go b/internal/handler/permission_handler.go index 7cd60e7..c261d48 100644 --- a/internal/handler/permission_handler.go +++ b/internal/handler/permission_handler.go @@ -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 diff --git a/internal/handler/permissions.go b/internal/handler/permissions.go index a2e1e55..6753fe4 100644 --- a/internal/handler/permissions.go +++ b/internal/handler/permissions.go @@ -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 } diff --git a/internal/repository/repository.go b/internal/repository/repository.go index d2af001..e15a868 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -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 ─────────────────────────────────────────────────────────── diff --git a/internal/repository/sqlite_busy_retry.go b/internal/repository/sqlite_busy_retry.go index 7b3e880..da1e247 100644 --- a/internal/repository/sqlite_busy_retry.go +++ b/internal/repository/sqlite_busy_retry.go @@ -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 } diff --git a/internal/service/auth.go b/internal/service/auth.go index e07bf30..1293c34 100644 --- a/internal/service/auth.go +++ b/internal/service/auth.go @@ -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 } diff --git a/internal/service/auth_user_limits_test.go b/internal/service/auth_user_limits_test.go index 2d9d823..9690c39 100644 --- a/internal/service/auth_user_limits_test.go +++ b/internal/service/auth_user_limits_test.go @@ -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 { diff --git a/internal/service/db_errors.go b/internal/service/db_errors.go new file mode 100644 index 0000000..8084e34 --- /dev/null +++ b/internal/service/db_errors.go @@ -0,0 +1,7 @@ +package service + +import "github.com/ShukeBta/MediaStationGo/internal/repository" + +func IsTransientDatabaseLock(err error) bool { + return repository.IsSQLiteBusyError(err) +} diff --git a/internal/service/permission.go b/internal/service/permission.go index b913668..87e2382 100644 --- a/internal/service/permission.go +++ b/internal/service/permission.go @@ -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) { diff --git a/internal/service/token_svc.go b/internal/service/token_svc.go index 62fc437..a08a384 100644 --- a/internal/service/token_svc.go +++ b/internal/service/token_svc.go @@ -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 {