diff --git a/internal/database/database.go b/internal/database/database.go index 09548c9..00de091 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -131,11 +131,7 @@ func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *za return nil } - srcCfg := *cfg - srcCfg.Database.Type = "sqlite" - src, err := gorm.Open(sqlite.Open(buildSQLiteDSN(&srcCfg)), &gorm.Config{ - Logger: logger.Default.LogMode(logger.Silent), - }) + src, err := openSQLiteMigrationSource(cfg, sqlitePath) if err != nil { return fmt.Errorf("open sqlite migration source: %w", err) } @@ -164,6 +160,15 @@ func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *za return nil } +func openSQLiteMigrationSource(cfg *config.Config, sqlitePath string) (*gorm.DB, error) { + srcCfg := *cfg + srcCfg.Database.Type = "sqlite" + srcCfg.Database.DBPath = sqlitePath + return gorm.Open(sqlite.Open(buildSQLiteDSN(&srcCfg)), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) +} + func sqliteMigrationSourcePath(cfg *config.Config, log *zap.Logger) (string, error) { configured := strings.TrimSpace(cfg.Database.DBPath) if configured != "" { diff --git a/internal/database/database_test.go b/internal/database/database_test.go index 6e9fe87..6a7fdb5 100644 --- a/internal/database/database_test.go +++ b/internal/database/database_test.go @@ -302,3 +302,67 @@ func TestSQLiteMigrationFallsBackToDataDirDefaultPath(t *testing.T) { t.Fatalf("library count = %d, want 1", libCount) } } + +func TestOpenSQLiteMigrationSourceUsesFallbackSourcePath(t *testing.T) { + dir := t.TempDir() + sqlitePath := filepath.Join(dir, "mediastation.db") + src, err := gorm.Open(sqlite.Open(sqlitePath), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := src.AutoMigrate(&model.User{}, &model.Library{}, &model.Setting{}); err != nil { + t.Fatal(err) + } + if err := src.Create(&model.User{Username: "real-admin", PasswordHash: "hash", Role: "admin", IsActive: true}).Error; err != nil { + t.Fatal(err) + } + if err := src.Create(&model.Library{Name: "Movies", Path: "/media/movies", Type: "movie", Enabled: true}).Error; err != nil { + t.Fatal(err) + } + sqlDB, _ := src.DB() + _ = sqlDB.Close() + + dst, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := dst.AutoMigrate(&model.User{}, &model.Library{}, &model.Setting{}); err != nil { + t.Fatal(err) + } + + cfg := &config.Config{} + cfg.App.DataDir = dir + cfg.Database.DBPath = filepath.Join(dir, "disabled-sqlite-migration.db") + sourcePath, err := sqliteMigrationSourcePath(cfg, nil) + if err != nil { + t.Fatal(err) + } + if sourcePath != sqlitePath { + t.Fatalf("source path = %q, want %q", sourcePath, sqlitePath) + } + + src2, err := openSQLiteMigrationSource(cfg, sourcePath) + if err != nil { + t.Fatal(err) + } + sqlDB2, _ := src2.DB() + defer func() { + if sqlDB2 != nil { + _ = sqlDB2.Close() + } + }() + copied, err := copyModelTables(src2, dst, 2) + if err != nil { + t.Fatal(err) + } + if copied != 2 { + t.Fatalf("copied rows = %d, want 2", copied) + } + var userCount int64 + if err := dst.Model(&model.User{}).Where("username = ?", "real-admin").Count(&userCount).Error; err != nil { + t.Fatal(err) + } + if userCount != 1 { + t.Fatalf("migrated user count = %d, want 1", userCount) + } +} diff --git a/internal/service/media_classifier_test.go b/internal/service/media_classifier_test.go index 5748777..0485abd 100644 --- a/internal/service/media_classifier_test.go +++ b/internal/service/media_classifier_test.go @@ -1,7 +1,6 @@ package service import ( - "path/filepath" "testing" "github.com/glebarez/sqlite" @@ -176,7 +175,7 @@ func TestSubscriptionResolveClassifiedSavePath(t *testing.T) { if err := repos.Setting.Set(t.Context(), "organizer.smart_classify", "true"); err != nil { t.Fatal(err) } - if err := repos.Setting.Set(t.Context(), "qbittorrent.savepath", filepath.Join("D:", "Downloads")); err != nil { + if err := repos.Setting.Set(t.Context(), "qbittorrent.savepath", `D:\Downloads`); err != nil { t.Fatal(err) } svc := NewSubscriptionService(&config.Config{}, zap.NewNop(), repos, nil, nil, nil) @@ -187,7 +186,7 @@ func TestSubscriptionResolveClassifiedSavePath(t *testing.T) { t.Fatalf("classification = %q/%q, want tv/综艺", mediaType, category) } got := svc.resolveSubscriptionSavePath(t.Context(), sub, mediaType, category) - want := filepath.Join("D:", "Downloads", "综艺") + want := `D:\Downloads\综艺` if got != want { t.Fatalf("save path = %q, want %q", got, want) }