diff --git a/cmd/server/main.go b/cmd/server/main.go index 7ee0af6..92c8a50 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -19,6 +19,7 @@ import ( "os" "os/signal" "path/filepath" + "runtime" "strings" "syscall" "time" @@ -73,6 +74,8 @@ func main() { } repos := repository.New(db) + service.ApplyRuntimeSettings(context.Background(), cfg, repos, logger) + applyCPUThreadLimit(cfg, logger) services := service.New(cfg, logger, repos) if err := services.Auth.SeedAdmin(context.Background()); err != nil { @@ -126,6 +129,18 @@ func main() { logger.Info("MediaStationGo stopped") } +func applyCPUThreadLimit(cfg *config.Config, logger *zap.Logger) { + if cfg == nil || cfg.App.MaxCPUThreads < 1 { + return + } + prev := runtime.GOMAXPROCS(cfg.App.MaxCPUThreads) + if logger != nil { + logger.Info("runtime CPU thread limit applied", + zap.Int("max_cpu_threads", cfg.App.MaxCPUThreads), + zap.Int("previous", prev)) + } +} + func buildRouter(cfg *config.Config, logger *zap.Logger, svc *service.Container) *gin.Engine { if !cfg.App.Debug { gin.SetMode(gin.ReleaseMode) diff --git a/config.example.yaml b/config.example.yaml index 3cd0b03..804cc0b 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -20,6 +20,7 @@ app: web_dir: ./web/dist ffmpeg_path: ffmpeg ffprobe_path: ffprobe + max_cpu_threads: 2 # hard cap for Go CPU threads; raise only on powerful hosts vaapi_device: /dev/dri/renderD128 cors_origins: [] # leave empty to allow * in dev; set strict allow-list in prod server_url: "" # public URL of the server (used for DLNA / casting) diff --git a/internal/config/config.go b/internal/config/config.go index 9e3e694..066a968 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -80,6 +80,7 @@ type AppConfig struct { // NAS devices can become unresponsive when a scan starts many probe // processes at once, so the default is deliberately conservative. FFprobeMaxConcurrent int `mapstructure:"ffprobe_max_concurrent"` + MaxCPUThreads int `mapstructure:"max_cpu_threads"` VAAPIDevice string `mapstructure:"vaapi_device"` CORSOrigins []string `mapstructure:"cors_origins"` ServerURL string `mapstructure:"server_url"` @@ -223,6 +224,7 @@ func setDefaults(v *viper.Viper) { v.SetDefault("app.ffmpeg_path", "ffmpeg") v.SetDefault("app.ffprobe_path", "ffprobe") v.SetDefault("app.ffprobe_max_concurrent", 1) + v.SetDefault("app.max_cpu_threads", 2) v.SetDefault("app.vaapi_device", "/dev/dri/renderD128") v.SetDefault("app.cors_origins", []string{}) v.SetDefault("app.server_url", "") @@ -308,6 +310,12 @@ func (c *Config) normalize() error { if c.Database.DBPath == "" { c.Database.DBPath = filepath.Join(c.App.DataDir, "mediastation.db") } + if c.App.MaxCPUThreads < 1 { + c.App.MaxCPUThreads = 1 + } + if c.App.MaxCPUThreads > 8 { + c.App.MaxCPUThreads = 8 + } if c.Database.MaxOpenConns <= 1 { c.Database.MaxOpenConns = defaultDatabaseMaxOpenConns } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 182afd8..2fb6679 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -25,6 +25,9 @@ func TestLoadDefaults(t *testing.T) { if cfg.App.Port != 8080 { t.Fatalf("expected default port 8080, got %d", cfg.App.Port) } + if cfg.App.MaxCPUThreads != 2 { + t.Fatalf("expected default MaxCPUThreads 2, got %d", cfg.App.MaxCPUThreads) + } if cfg.Database.DBPath == "" { t.Fatalf("expected non-empty DBPath") } diff --git a/internal/database/database.go b/internal/database/database.go index 90cda72..96f69e1 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -86,8 +86,6 @@ func AutoMigrate(db *gorm.DB) error { func ensurePerformanceIndexes(db *gorm.DB) error { statements := []string{ - `CREATE INDEX IF NOT EXISTS idx_media_created_active ON media(created_at DESC) WHERE deleted_at IS NULL`, - `CREATE INDEX IF NOT EXISTS idx_media_episode_created_active ON media(created_at DESC) WHERE deleted_at IS NULL AND (season_num > 0 OR episode_num > 0)`, `CREATE INDEX IF NOT EXISTS idx_media_library_created_active ON media(library_id, created_at DESC) WHERE deleted_at IS NULL`, `CREATE INDEX IF NOT EXISTS idx_media_library_episode_active ON media(library_id, season_num, episode_num, created_at DESC) WHERE deleted_at IS NULL`, `CREATE INDEX IF NOT EXISTS idx_media_series_active ON media(series_id, season_num, episode_num) WHERE deleted_at IS NULL`, diff --git a/internal/database/database_test.go b/internal/database/database_test.go index 5c4bdef..96e380b 100644 --- a/internal/database/database_test.go +++ b/internal/database/database_test.go @@ -58,8 +58,6 @@ func TestEnsurePerformanceIndexesCreatesHotPathIndexes(t *testing.T) { t.Fatal(err) } for _, name := range []string{ - "idx_media_created_active", - "idx_media_episode_created_active", "idx_media_library_created_active", "idx_media_library_episode_active", "idx_favorites_user_media_active", diff --git a/internal/repository/repository.go b/internal/repository/repository.go index 6ec7250..d2af001 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -114,10 +114,14 @@ func (r *UserRepository) ReleaseDeletedUsername(ctx context.Context, username st // FindByUsername returns the user matching username, or (nil, nil) when absent. func (r *UserRepository) FindByUsername(ctx context.Context, username string) (*model.User, error) { var u model.User - err := r.db.WithContext(ctx).Where("username = ?", username).First(&u).Error - if errors.Is(err, gorm.ErrRecordNotFound) && username != "" { - err = r.db.WithContext(ctx).Where("LOWER(username) = LOWER(?)", username).First(&u).Error - } + err := withSQLiteBusyRetry(ctx, func() error { + u = model.User{} + err := r.db.WithContext(ctx).Where("username = ?", username).First(&u).Error + if errors.Is(err, gorm.ErrRecordNotFound) && username != "" { + err = r.db.WithContext(ctx).Where("LOWER(username) = LOWER(?)", username).First(&u).Error + } + return err + }) if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } @@ -130,7 +134,10 @@ func (r *UserRepository) FindByUsername(ctx context.Context, username string) (* // FindByID returns the user with the matching primary key, or (nil, nil). func (r *UserRepository) FindByID(ctx context.Context, id string) (*model.User, error) { var u model.User - err := r.db.WithContext(ctx).Where("id = ?", id).First(&u).Error + err := withSQLiteBusyRetry(ctx, func() error { + u = model.User{} + return r.db.WithContext(ctx).Where("id = ?", id).First(&u).Error + }) if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } @@ -190,8 +197,10 @@ func (r *UserRepository) UpdatePassword(ctx context.Context, id, hash string) er // TouchLogin updates the last login timestamp. func (r *UserRepository) TouchLogin(ctx context.Context, id string) error { now := time.Now() - return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id). - Update("last_login_at", &now).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id). + Update("last_login_at", &now).Error + }) } // Delete removes a user (soft-delete via gorm.DeletedAt), releases the unique @@ -896,13 +905,18 @@ type RefreshTokenRepository struct{ db *gorm.DB } // Create inserts a new refresh token record. func (r *RefreshTokenRepository) Create(ctx context.Context, t *model.RefreshToken) error { - return r.db.WithContext(ctx).Create(t).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Create(t).Error + }) } // FindByHash returns the refresh token matching the hash, or (nil, nil). func (r *RefreshTokenRepository) FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error) { var t model.RefreshToken - err := r.db.WithContext(ctx).Where("token_hash = ?", hash).First(&t).Error + err := withSQLiteBusyRetry(ctx, func() error { + t = model.RefreshToken{} + return r.db.WithContext(ctx).Where("token_hash = ?", hash).First(&t).Error + }) if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } @@ -914,8 +928,10 @@ func (r *RefreshTokenRepository) FindByHash(ctx context.Context, hash string) (* // RevokeByUserID revokes all refresh tokens for a user. func (r *RefreshTokenRepository) RevokeByUserID(ctx context.Context, userID string) error { - return r.db.WithContext(ctx).Model(&model.RefreshToken{}). - Where("user_id = ?", userID).Update("revoked", true).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.RefreshToken{}). + Where("user_id = ?", userID).Update("revoked", true).Error + }) } // RevokeOldestActiveByUserID keeps at most limit active refresh tokens for a @@ -924,33 +940,39 @@ func (r *RefreshTokenRepository) RevokeOldestActiveByUserID(ctx context.Context, if limit < 1 { limit = 1 } - var tokens []model.RefreshToken - if err := r.db.WithContext(ctx). - Where("user_id = ? AND revoked = ? AND expires_at > ?", userID, false, time.Now()). - Order("created_at desc, id desc"). - Find(&tokens).Error; err != nil { - return err - } - if len(tokens) <= limit { - return nil - } - ids := make([]string, 0, len(tokens)-limit) - for _, token := range tokens[limit:] { - ids = append(ids, token.ID) - } - return r.db.WithContext(ctx).Model(&model.RefreshToken{}). - Where("id IN ?", ids).Update("revoked", true).Error + return withSQLiteBusyRetry(ctx, func() error { + var tokens []model.RefreshToken + if err := r.db.WithContext(ctx). + Where("user_id = ? AND revoked = ? AND expires_at > ?", userID, false, time.Now()). + Order("created_at desc, id desc"). + Find(&tokens).Error; err != nil { + return err + } + if len(tokens) <= limit { + return nil + } + ids := make([]string, 0, len(tokens)-limit) + for _, token := range tokens[limit:] { + ids = append(ids, token.ID) + } + return r.db.WithContext(ctx).Model(&model.RefreshToken{}). + Where("id IN ?", ids).Update("revoked", true).Error + }) } // DeleteExpired removes all expired refresh tokens. func (r *RefreshTokenRepository) DeleteExpired(ctx context.Context) error { - return r.db.WithContext(ctx).Where("expires_at < ?", time.Now()).Delete(&model.RefreshToken{}).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Where("expires_at < ?", time.Now()).Delete(&model.RefreshToken{}).Error + }) } // Revoke revokes a specific refresh token. func (r *RefreshTokenRepository) Revoke(ctx context.Context, hash string) error { - return r.db.WithContext(ctx).Model(&model.RefreshToken{}). - Where("token_hash = ?", hash).Update("revoked", true).Error + return withSQLiteBusyRetry(ctx, func() error { + return r.db.WithContext(ctx).Model(&model.RefreshToken{}). + Where("token_hash = ?", hash).Update("revoked", true).Error + }) } // HashToken returns the SHA256 hash of a token. diff --git a/internal/repository/sqlite_busy_retry.go b/internal/repository/sqlite_busy_retry.go new file mode 100644 index 0000000..7b3e880 --- /dev/null +++ b/internal/repository/sqlite_busy_retry.go @@ -0,0 +1,47 @@ +package repository + +import ( + "context" + "strings" + "time" +) + +const sqliteBusyRetryMaxElapsed = 6 * time.Second + +func withSQLiteBusyRetry(ctx context.Context, op func() error) error { + delay := 25 * time.Millisecond + deadline := time.Now().Add(sqliteBusyRetryMaxElapsed) + for { + err := op() + if !isSQLiteBusyError(err) { + return err + } + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + if time.Now().Add(delay).After(deadline) { + return err + } + timer := time.NewTimer(delay) + select { + case <-ctx.Done(): + timer.Stop() + return ctx.Err() + case <-timer.C: + } + if delay < 500*time.Millisecond { + delay *= 2 + } + } +} + +func isSQLiteBusyError(err error) bool { + if err == nil { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "sqlite_busy") || + strings.Contains(msg, "sqlite_locked") || + strings.Contains(msg, "database is locked") || + strings.Contains(msg, "database table is locked") +} diff --git a/internal/service/auth_user_limits_test.go b/internal/service/auth_user_limits_test.go index b931450..2d9d823 100644 --- a/internal/service/auth_user_limits_test.go +++ b/internal/service/auth_user_limits_test.go @@ -5,15 +5,18 @@ import ( "encoding/json" "errors" "fmt" + "path/filepath" "testing" "time" "github.com/glebarez/sqlite" "github.com/golang-jwt/jwt/v5" "go.uber.org/zap" + "golang.org/x/crypto/bcrypt" "gorm.io/gorm" "github.com/ShukeBta/MediaStationGo/internal/config" + "github.com/ShukeBta/MediaStationGo/internal/database" "github.com/ShukeBta/MediaStationGo/internal/model" "github.com/ShukeBta/MediaStationGo/internal/repository" ) @@ -227,6 +230,67 @@ func TestLoginKeepsOnlyConfiguredActiveRefreshTokens(t *testing.T) { } } +func TestLoginRetriesTransientSQLiteBusy(t *testing.T) { + ctx := context.Background() + cfg := &config.Config{} + cfg.App.DataDir = t.TempDir() + cfg.Database.DBPath = filepath.Join(cfg.App.DataDir, "busy-login.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) + } + if err := repos.User.Create(ctx, &model.User{ + Username: "viewer", + PasswordHash: string(hash), + Role: "user", + Tier: "free", + IsActive: true, + }); 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) + } + release := make(chan struct{}) + go func() { + timer := time.NewTimer(250 * time.Millisecond) + defer timer.Stop() + select { + case <-release: + case <-timer.C: + } + _ = tx.Rollback().Error + }() + defer close(release) + + if _, err := auth.Login(ctx, "viewer", "password"); err != nil { + t.Fatalf("login should survive a transient sqlite write lock: %v", err) + } +} + func TestDefaultPermissionsAreViewerOnly(t *testing.T) { perms := DefaultPermissions("user-1") if !perms.CanViewDashboard || !perms.CanPlayMedia || !perms.CanExternalPlayer { diff --git a/internal/service/downloads.go b/internal/service/downloads.go index 4f8bdb2..55f4dd5 100644 --- a/internal/service/downloads.go +++ b/internal/service/downloads.go @@ -1049,7 +1049,6 @@ func (d *DownloadService) completedTorrentSource(torrent QBitTorrent) string { for _, candidate := range []string{ torrent.ContentPath, filepath.Join(torrent.SavePath, torrent.Name), - torrent.SavePath, } { clean := strings.TrimSpace(candidate) if clean == "" || clean == "." { diff --git a/internal/service/downloads_test.go b/internal/service/downloads_test.go index a53bf2c..2a6aa1a 100644 --- a/internal/service/downloads_test.go +++ b/internal/service/downloads_test.go @@ -87,6 +87,26 @@ func TestDownloadCompleteAutoOrganizesContentPath(t *testing.T) { } } +func TestCompletedTorrentSourceDoesNotFallbackToSavePath(t *testing.T) { + root := t.TempDir() + savePath := filepath.Join(root, "downloads", "日番") + if err := os.MkdirAll(savePath, 0o755); err != nil { + t.Fatal(err) + } + svc := NewDownloadService(zap.NewNop(), newOrganizerTestRepo(t), NewHub(zap.NewNop()), nil) + + got := svc.completedTorrentSource(QBitTorrent{ + Hash: "done123", + Name: "Missing.Payload.S01", + SavePath: savePath, + ContentPath: filepath.Join(savePath, "Missing.Payload.S01", "Missing.Payload.S01E01.mkv"), + }) + + if got != "" { + t.Fatalf("completedTorrentSource fell back to whole save_path %q; want empty", got) + } +} + func TestDownloadPollBaselinesAlreadyCompletedTorrents(t *testing.T) { repos := newOrganizerTestRepo(t) svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) diff --git a/internal/service/runtime_settings.go b/internal/service/runtime_settings.go index 5746a24..edffca0 100644 --- a/internal/service/runtime_settings.go +++ b/internal/service/runtime_settings.go @@ -49,6 +49,16 @@ func ApplyRuntimeSetting(cfg *config.Config, key, value string) { } cfg.App.FFprobeMaxConcurrent = n } + case "app.max_cpu_threads", "runtime.max_cpu_threads": + if n, err := strconv.Atoi(value); err == nil { + if n < 1 { + n = 1 + } + if n > 8 { + n = 8 + } + cfg.App.MaxCPUThreads = n + } case "transcode.enabled", "transcoder.enabled": cfg.Transcoder.Enabled = parseBoolSetting(value, true) case "transcode.hw_enabled", "transcoder.hardware_accel": diff --git a/internal/service/runtime_settings_test.go b/internal/service/runtime_settings_test.go index 4987bc3..b05ecfc 100644 --- a/internal/service/runtime_settings_test.go +++ b/internal/service/runtime_settings_test.go @@ -30,4 +30,14 @@ func TestApplyRuntimeSettingTranscodeSwitches(t *testing.T) { if cfg.Transcoder.MaxConcurrent != 1 { t.Fatalf("max concurrent = %d, want 1", cfg.Transcoder.MaxConcurrent) } + + ApplyRuntimeSetting(cfg, "app.max_cpu_threads", "99") + if cfg.App.MaxCPUThreads != 8 { + t.Fatalf("max cpu threads = %d, want clamp 8", cfg.App.MaxCPUThreads) + } + + ApplyRuntimeSetting(cfg, "app.max_cpu_threads", "0") + if cfg.App.MaxCPUThreads != 1 { + t.Fatalf("max cpu threads = %d, want clamp 1", cfg.App.MaxCPUThreads) + } }