diff --git a/internal/config/defaults.go b/internal/config/defaults.go index 7f273aa..f64d462 100644 --- a/internal/config/defaults.go +++ b/internal/config/defaults.go @@ -3,8 +3,8 @@ package config import "github.com/spf13/viper" const ( - defaultDatabaseMaxOpenConns = 4 - defaultDatabaseMaxIdleConns = 2 + defaultDatabaseMaxOpenConns = 16 + defaultDatabaseMaxIdleConns = 4 defaultLicenseServerURL = "https://mgosever.3jzs.com" defaultLicensePublicKey = "MCowBQYDK2VwAyEABRXnXy+urjrbKit6Yu/HiezWgP0NdsZW3tsegJWRrtI=" ) diff --git a/internal/config/save.go b/internal/config/save.go new file mode 100644 index 0000000..cbedb10 --- /dev/null +++ b/internal/config/save.go @@ -0,0 +1,41 @@ +package config + +import ( + "fmt" + "os" + + "gopkg.in/yaml.v3" +) + +// SaveDatabaseConfig updates or creates config.yaml with the specified database configuration. +func SaveDatabaseConfig(dbType, dsn string) error { + configPath := "config.yaml" + data := make(map[string]any) + + content, err := os.ReadFile(configPath) + if err == nil { + if err := yaml.Unmarshal(content, &data); err != nil { + data = make(map[string]any) + } + } else if !os.IsNotExist(err) { + return fmt.Errorf("read config.yaml: %w", err) + } + + dbSection, ok := data["database"].(map[string]any) + if !ok { + dbSection = make(map[string]any) + } + dbSection["type"] = dbType + dbSection["dsn"] = dsn + data["database"] = dbSection + + out, err := yaml.Marshal(data) + if err != nil { + return fmt.Errorf("marshal config.yaml: %w", err) + } + + if err := os.WriteFile(configPath, out, 0644); err != nil { + return fmt.Errorf("write config.yaml: %w", err) + } + return nil +} diff --git a/internal/config/save_test.go b/internal/config/save_test.go new file mode 100644 index 0000000..261d7e2 --- /dev/null +++ b/internal/config/save_test.go @@ -0,0 +1,36 @@ +package config + +import ( + "os" + "path/filepath" + "testing" +) + +func TestSaveDatabaseConfig(t *testing.T) { + dir := t.TempDir() + wd, _ := os.Getwd() + defer func() { _ = os.Chdir(wd) }() + if err := os.Chdir(dir); err != nil { + t.Fatalf("chdir: %v", err) + } + + dsn := "postgres://admin:pass@127.0.0.1:5432/mmtl?sslmode=disable" + if err := SaveDatabaseConfig("postgres", dsn); err != nil { + t.Fatalf("SaveDatabaseConfig error: %v", err) + } + + if _, err := os.Stat(filepath.Join(dir, "config.yaml")); err != nil { + t.Fatalf("expected config.yaml to exist: %v", err) + } + + loaded, err := Load() + if err != nil { + t.Fatalf("Load error: %v", err) + } + if loaded.Database.Type != "postgres" { + t.Fatalf("expected database.type=postgres, got %s", loaded.Database.Type) + } + if loaded.Database.DSN != dsn { + t.Fatalf("expected dsn=%s, got %s", dsn, loaded.Database.DSN) + } +} diff --git a/internal/database/database_admin.go b/internal/database/database_admin.go new file mode 100644 index 0000000..6b97032 --- /dev/null +++ b/internal/database/database_admin.go @@ -0,0 +1,223 @@ +package database + +import ( + "context" + "fmt" + "net/url" + "strings" + "time" + + "go.uber.org/zap" + "gorm.io/driver/postgres" + "gorm.io/gorm" + "gorm.io/gorm/logger" + + "github.com/ShukeBta/MMTL/internal/config" + "github.com/ShukeBta/MMTL/internal/model" +) + +// DatabaseStatus describes the currently active database engine and runtime metrics. +type DatabaseStatus struct { + Type string `json:"type"` + DSN string `json:"dsn,omitempty"` + DBPath string `json:"db_path,omitempty"` + OpenConns int `json:"open_conns"` + InUse int `json:"in_use"` + Idle int `json:"idle"` + MaxOpenConns int `json:"max_open_conns"` + TableCounts map[string]int64 `json:"table_counts"` +} + +// PostgresTestResult returns latency and version info after testing connection. +type PostgresTestResult struct { + Success bool `json:"success"` + LatencyMS int64 `json:"latency_ms"` + Version string `json:"version,omitempty"` + Message string `json:"message,omitempty"` + Error string `json:"error,omitempty"` +} + +// DatabaseMigrationResult returns row counts and execution duration of migration. +type DatabaseMigrationResult struct { + Success bool `json:"success"` + TotalRows int64 `json:"total_rows"` + TableRows map[string]int64 `json:"table_rows"` + DurationMS int64 `json:"duration_ms"` + Message string `json:"message,omitempty"` + Error string `json:"error,omitempty"` +} + +// InspectDatabaseStatus queries the currently active database for metrics and table rows. +func InspectDatabaseStatus(db *gorm.DB, cfg *config.Config) *DatabaseStatus { + st := &DatabaseStatus{ + Type: "sqlite", + TableCounts: make(map[string]int64), + } + if cfg != nil { + st.DBPath = cfg.Database.DBPath + if cfg.Database.Type == "postgres" || (cfg.Database.Type == "auto" && strings.TrimSpace(cfg.Database.DSN) != "") { + st.Type = "postgres" + st.DSN = MaskDSN(cfg.Database.DSN) + } + } + if isPostgres(db) { + st.Type = "postgres" + } + + if db != nil { + if sqlDB, err := db.DB(); err == nil { + stats := sqlDB.Stats() + st.OpenConns = stats.OpenConnections + st.InUse = stats.InUse + st.Idle = stats.Idle + st.MaxOpenConns = stats.MaxOpenConnections + } + + // Count rows for major model tables + for _, m := range model.AllModels() { + if tbl, err := modelTableName(db, m); err == nil { + if db.Migrator().HasTable(tbl) { + var count int64 + if err := db.Raw("SELECT COUNT(1) FROM " + quoteIdent(tbl)).Scan(&count).Error; err == nil { + st.TableCounts[tbl] = count + } + } + } + } + } + return st +} + +// TestPostgres establishes a temporary connection to verify reachability and permissions. +func TestPostgres(dsn string) (*PostgresTestResult, error) { + dsn = strings.TrimSpace(dsn) + if dsn == "" { + return &PostgresTestResult{ + Success: false, + Error: "PostgreSQL DSN 不能为空", + }, nil + } + + start := time.Now() + testDB, err := gorm.Open(postgres.Open(dsn), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + return &PostgresTestResult{ + Success: false, + Error: fmt.Sprintf("连接失败: %v", err), + }, nil + } + + sqlDB, err := testDB.DB() + if err != nil { + return &PostgresTestResult{ + Success: false, + Error: fmt.Sprintf("获取底层连接失败: %v", err), + }, nil + } + defer sqlDB.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + if err := sqlDB.PingContext(ctx); err != nil { + return &PostgresTestResult{ + Success: false, + Error: fmt.Sprintf("Ping 超时或失败: %v", err), + }, nil + } + + var version string + if err := testDB.WithContext(ctx).Raw("SELECT version()").Scan(&version).Error; err != nil { + version = "PostgreSQL (unknown version)" + } + + latency := time.Since(start).Milliseconds() + return &PostgresTestResult{ + Success: true, + LatencyMS: latency, + Version: version, + Message: "连接成功", + }, nil +} + +// MigrateCurrentToPostgres performs schema initialization and full table data copy into target PostgreSQL. +func MigrateCurrentToPostgres(src *gorm.DB, targetDSN string, batchSize int, log *zap.Logger) (*DatabaseMigrationResult, error) { + targetDSN = strings.TrimSpace(targetDSN) + if targetDSN == "" { + return nil, fmt.Errorf("target PostgreSQL DSN cannot be empty") + } + if src == nil { + return nil, fmt.Errorf("current database is not available") + } + + started := time.Now() + targetDB, err := gorm.Open(postgres.Open(targetDSN), &gorm.Config{ + Logger: newGormLogger(log), + }) + if err != nil { + return nil, fmt.Errorf("open target PostgreSQL: %w", err) + } + targetSQLDB, err := targetDB.DB() + if err == nil { + defer targetSQLDB.Close() + } + + // 1. 初始化目标库 Schema、类型与索引 + if err := AutoMigrate(targetDB); err != nil { + return nil, fmt.Errorf("auto migrate target PostgreSQL: %w", err) + } + + // 2. 安全重置目标数据库的初始默认数据 + if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, targetDB, log); err != nil { + return nil, fmt.Errorf("reset target bootstrap data: %w", err) + } + + // 3. 执行数据批量复制 + tableRows, totalRows, err := copyModelTables(src, targetDB, batchSize) + if err != nil { + return nil, fmt.Errorf("copy tables: %w", err) + } + + // 4. 标记迁移完成 + if err := markSQLiteMigrationComplete(targetDB); err != nil { + return nil, fmt.Errorf("mark migration complete: %w", err) + } + + duration := time.Since(started).Milliseconds() + return &DatabaseMigrationResult{ + Success: true, + TotalRows: totalRows, + TableRows: tableRows, + DurationMS: duration, + Message: fmt.Sprintf("成功迁移 %d 条记录至 PostgreSQL", totalRows), + }, nil +} + +// MaskDSN masks the password in a connection string for safe API responses. +func MaskDSN(rawDSN string) string { + rawDSN = strings.TrimSpace(rawDSN) + if rawDSN == "" { + return "" + } + if u, err := url.Parse(rawDSN); err == nil && u.User != nil { + if pass, hasPassword := u.User.Password(); hasPassword && pass != "" { + rawUserPass := u.User.String() + user := u.User.Username() + maskedUserPass := user + ":******" + return strings.Replace(rawDSN, rawUserPass+"@", maskedUserPass+"@", 1) + } + } + // Fallback for keyword-style DSN (e.g. host=... password=...) + if strings.Contains(rawDSN, "password=") { + parts := strings.Fields(rawDSN) + for i, p := range parts { + if strings.HasPrefix(p, "password=") { + parts[i] = "password=******" + } + } + return strings.Join(parts, " ") + } + return rawDSN +} diff --git a/internal/database/database_admin_test.go b/internal/database/database_admin_test.go new file mode 100644 index 0000000..bb1acb2 --- /dev/null +++ b/internal/database/database_admin_test.go @@ -0,0 +1,71 @@ +package database + +import ( + "testing" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" + + "github.com/ShukeBta/MMTL/internal/config" + "github.com/ShukeBta/MMTL/internal/model" +) + +func TestMaskDSN(t *testing.T) { + cases := []struct { + in string + want string + }{ + { + in: "postgres://admin:secret123@localhost:5432/mmtl?sslmode=disable", + want: "postgres://admin:******@localhost:5432/mmtl?sslmode=disable", + }, + { + in: "host=localhost port=5432 user=admin password=secret dbname=mmtl sslmode=disable", + want: "host=localhost port=5432 user=admin password=****** dbname=mmtl sslmode=disable", + }, + { + in: "sqlite://data/mmtl.db", + want: "sqlite://data/mmtl.db", + }, + { + in: "", + want: "", + }, + } + + for _, c := range cases { + got := MaskDSN(c.in) + if got != c.want { + t.Errorf("MaskDSN(%q) = %q, want %q", c.in, got, c.want) + } + } +} + +func TestInspectDatabaseStatus(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.User{}, &model.Media{}); err != nil { + t.Fatal(err) + } + _ = db.Create(&model.User{Username: "testuser", PasswordHash: "h", Role: "user"}).Error + + cfg := &config.Config{} + cfg.Database.Type = "sqlite" + cfg.Database.DBPath = "./data/mmtl.db" + + st := InspectDatabaseStatus(db, cfg) + if st == nil { + t.Fatal("expected non-nil DatabaseStatus") + } + if st.Type != "sqlite" { + t.Fatalf("expected sqlite, got %s", st.Type) + } + if st.DBPath != "./data/mmtl.db" { + t.Fatalf("expected db_path, got %s", st.DBPath) + } + if st.TableCounts["users"] != 1 { + t.Fatalf("expected 1 user, got %d", st.TableCounts["users"]) + } +} diff --git a/internal/database/database_test.go b/internal/database/database_test.go index 8b9f098..73fb690 100644 --- a/internal/database/database_test.go +++ b/internal/database/database_test.go @@ -159,7 +159,7 @@ func TestCopyModelTablesMigratesExistingSQLiteRows(t *testing.T) { t.Fatal(err) } - copied, err := copyModelTables(src, dst, 2) + _, copied, err := copyModelTables(src, dst, 2) if err != nil { t.Fatal(err) } @@ -222,7 +222,7 @@ func TestCopyModelTablesResumesPartialSQLiteMigration(t *testing.T) { t.Fatal(err) } - copied, err := copyModelTables(src, dst, 2) + _, copied, err := copyModelTables(src, dst, 2) if err != nil { t.Fatal(err) } @@ -240,7 +240,7 @@ func TestCopyModelTablesResumesPartialSQLiteMigration(t *testing.T) { t.Fatalf("genres = %q, want %q", got.Genres, media.Genres) } - copied, err = copyModelTables(src, dst, 2) + _, copied, err = copyModelTables(src, dst, 2) if err != nil { t.Fatal(err) } @@ -332,7 +332,7 @@ func TestSQLiteMigrationFallsBackToDataDirDefaultPath(t *testing.T) { if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src2, dst, nil); err != nil { t.Fatal(err) } - copied, err := copyModelTables(src2, dst, 2) + _, copied, err := copyModelTables(src2, dst, 2) if err != nil { t.Fatal(err) } @@ -408,13 +408,13 @@ func TestOpenSQLiteMigrationSourceUsesFallbackSourcePath(t *testing.T) { _ = 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) - } + _, 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) diff --git a/internal/database/sqlite_migration.go b/internal/database/sqlite_migration.go index 29f92f6..f5a274e 100644 --- a/internal/database/sqlite_migration.go +++ b/internal/database/sqlite_migration.go @@ -48,7 +48,7 @@ func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *za if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target, log); err != nil { return err } - copied, err := copyModelTables(src, target, 500) + _, copied, err := copyModelTables(src, target, 500) if err != nil { return err } diff --git a/internal/database/sqlite_migration_copy.go b/internal/database/sqlite_migration_copy.go index fbf3f8c..4053c7c 100644 --- a/internal/database/sqlite_migration_copy.go +++ b/internal/database/sqlite_migration_copy.go @@ -13,52 +13,53 @@ import ( "github.com/ShukeBta/MMTL/internal/model" ) -func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) { +func copyModelTables(src, target *gorm.DB, batchSize int) (map[string]int64, int64, error) { if batchSize <= 0 { batchSize = 500 } - var copied int64 + tableCounts := make(map[string]int64) + var totalCopied int64 for _, m := range model.AllModels() { table, err := modelTableName(src, m) if err != nil { - return copied, err + return tableCounts, totalCopied, err } primaryColumns, err := modelPrimaryColumns(src, m) if err != nil { - return copied, fmt.Errorf("inspect model %T primary keys: %w", m, err) + return tableCounts, totalCopied, fmt.Errorf("inspect model %T primary keys: %w", m, err) } exists, err := sqliteTableExists(src, table) if err != nil { - return copied, err + return tableCounts, totalCopied, err } if !exists { continue } var sourceCount int64 if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&sourceCount).Error; err != nil { - return copied, fmt.Errorf("count sqlite table %s: %w", table, err) + return tableCounts, totalCopied, fmt.Errorf("count sqlite table %s: %w", table, err) } if sourceCount == 0 { continue } var targetCount int64 if err := target.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&targetCount).Error; err != nil { - return copied, fmt.Errorf("count target table %s: %w", table, err) + return tableCounts, totalCopied, fmt.Errorf("count target table %s: %w", table, err) } modelType := reflect.TypeOf(m) if modelType.Kind() != reflect.Ptr { - return copied, fmt.Errorf("model %T is not a pointer", m) + return tableCounts, totalCopied, fmt.Errorf("model %T is not a pointer", m) } sliceType := reflect.SliceOf(modelType.Elem()) slicePtr := reflect.New(sliceType) if err := src.Unscoped().Find(slicePtr.Interface()).Error; err != nil { - return copied, fmt.Errorf("read sqlite table %s: %w", table, err) + return tableCounts, totalCopied, fmt.Errorf("read sqlite table %s: %w", table, err) } filtered := slicePtr.Elem() if targetCount > 0 { primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns) if err != nil { - return copied, err + return tableCounts, totalCopied, err } filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet) } @@ -68,11 +69,13 @@ func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) { filteredPtr := reflect.New(filtered.Type()) filteredPtr.Elem().Set(filtered) if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil { - return copied, fmt.Errorf("copy sqlite table %s: %w", table, err) + return tableCounts, totalCopied, fmt.Errorf("copy sqlite table %s: %w", table, err) } - copied += int64(filtered.Len()) + copiedForTable := int64(filtered.Len()) + tableCounts[table] = copiedForTable + totalCopied += copiedForTable } - return copied, nil + return tableCounts, totalCopied, nil } func modelPrimaryColumns(db *gorm.DB, m any) ([]string, error) { diff --git a/internal/database/sqlite_runtime.go b/internal/database/sqlite_runtime.go index 57b1f7c..a0ce354 100644 --- a/internal/database/sqlite_runtime.go +++ b/internal/database/sqlite_runtime.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "path/filepath" + "strings" "gorm.io/gorm" @@ -32,16 +33,37 @@ func installSQLiteWriteGate(db *gorm.DB) { gate.Unlock() } } + rawLock := func(tx *gorm.DB) { + if tx.Statement != nil && isReadOnlySQL(tx.Statement.SQL.String()) { + return + } + lock(tx) + } _ = db.Callback().Create().Before("gorm:create").Register("mmtl:sqlite_write_lock", lock) _ = db.Callback().Create().After("gorm:create").Register("mmtl:sqlite_write_unlock", unlock) _ = db.Callback().Update().Before("gorm:update").Register("mmtl:sqlite_write_lock", lock) _ = db.Callback().Update().After("gorm:update").Register("mmtl:sqlite_write_unlock", unlock) _ = db.Callback().Delete().Before("gorm:delete").Register("mmtl:sqlite_write_lock", lock) _ = db.Callback().Delete().After("gorm:delete").Register("mmtl:sqlite_write_unlock", unlock) - _ = db.Callback().Raw().Before("gorm:raw").Register("mmtl:sqlite_write_lock", lock) + _ = db.Callback().Raw().Before("gorm:raw").Register("mmtl:sqlite_write_lock", rawLock) _ = db.Callback().Raw().After("gorm:raw").Register("mmtl:sqlite_write_unlock", unlock) } +func isReadOnlySQL(sql string) bool { + trimmed := strings.TrimSpace(sql) + if len(trimmed) == 0 { + return false + } + upper := strings.ToUpper(trimmed) + if strings.HasPrefix(upper, "SELECT") || strings.HasPrefix(upper, "EXPLAIN") { + return true + } + if strings.HasPrefix(upper, "WITH") && !strings.Contains(upper, "INSERT") && !strings.Contains(upper, "UPDATE") && !strings.Contains(upper, "DELETE") { + return true + } + return false +} + // sqliteWriteGate serializes in-process SQLite writes while respecting the // statement context, so request cancellation can break out of a queued write. type sqliteWriteGate struct { @@ -84,7 +106,7 @@ func buildSQLiteDSN(cfg *config.Config) string { } dsn := dbPath + "?_pragma=foreign_keys(1)" if cfg.Database.WALMode { - dsn += "&_pragma=journal_mode(WAL)" + dsn += "&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)" } if cfg.Database.BusyTimeout > 0 { dsn += fmt.Sprintf("&_pragma=busy_timeout(%d)", cfg.Database.BusyTimeout) @@ -92,6 +114,7 @@ func buildSQLiteDSN(cfg *config.Config) string { if cfg.Database.CacheSize != 0 { dsn += fmt.Sprintf("&_pragma=cache_size(%d)", cfg.Database.CacheSize) } + dsn += "&_pragma=temp_store(MEMORY)&_pragma=mmap_size(268435456)" return dsn } diff --git a/internal/handler/admin_database.go b/internal/handler/admin_database.go new file mode 100644 index 0000000..6b1fbdc --- /dev/null +++ b/internal/handler/admin_database.go @@ -0,0 +1,145 @@ +package handler + +import ( + "fmt" + "net/http" + "net/url" + "strings" + + "github.com/gin-gonic/gin" + + "github.com/ShukeBta/MMTL/internal/service" +) + +type DatabaseConnectionPayload struct { + Type string `json:"type"` + DSN string `json:"dsn"` + Host string `json:"host"` + Port int `json:"port"` + User string `json:"user"` + Password string `json:"password"` + DBName string `json:"dbname"` + SSLMode string `json:"sslmode"` +} + +func (p *DatabaseConnectionPayload) BuildDSN() string { + raw := strings.TrimSpace(p.DSN) + if raw != "" { + return raw + } + host := strings.TrimSpace(p.Host) + if host == "" { + return "" + } + port := p.Port + if port <= 0 { + port = 5432 + } + user := strings.TrimSpace(p.User) + dbname := strings.TrimSpace(p.DBName) + if dbname == "" { + dbname = "mmtl" + } + sslmode := strings.TrimSpace(p.SSLMode) + if sslmode == "" { + sslmode = "disable" + } + + userInfo := url.User(user) + if p.Password != "" { + userInfo = url.UserPassword(user, p.Password) + } + + u := url.URL{ + Scheme: "postgres", + User: userInfo, + Host: fmt.Sprintf("%s:%d", host, port), + Path: "/" + dbname, + RawQuery: "sslmode=" + url.QueryEscape(sslmode), + } + return u.String() +} + +func getDatabaseStatusHandler(svc *service.Container) gin.HandlerFunc { + return func(c *gin.Context) { + if svc.Database == nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "database service unavailable"}) + return + } + status := svc.Database.GetStatus(c.Request.Context()) + c.JSON(http.StatusOK, status) + } +} + +func testDatabaseHandler(svc *service.Container) gin.HandlerFunc { + return func(c *gin.Context) { + var req DatabaseConnectionPayload + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid payload: " + err.Error()}) + return + } + dsn := req.BuildDSN() + if dsn == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": "请提供有效的 PostgreSQL 连接信息或 DSN"}) + return + } + res, err := svc.Database.TestPostgres(c.Request.Context(), dsn) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + c.JSON(http.StatusOK, res) + } +} + +func migrateDatabaseHandler(svc *service.Container) gin.HandlerFunc { + return func(c *gin.Context) { + var req DatabaseConnectionPayload + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid payload: " + err.Error()}) + return + } + dsn := req.BuildDSN() + if dsn == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": "请提供目标 PostgreSQL 连接信息或 DSN"}) + return + } + res, err := svc.Database.MigrateToPostgres(c.Request.Context(), dsn) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "迁移失败: " + err.Error()}) + return + } + c.JSON(http.StatusOK, res) + } +} + +func saveDatabaseConfigHandler(svc *service.Container) gin.HandlerFunc { + return func(c *gin.Context) { + var req DatabaseConnectionPayload + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid payload: " + err.Error()}) + return + } + dbType := strings.ToLower(strings.TrimSpace(req.Type)) + if dbType == "" { + dbType = "postgres" + } + var dsn string + if dbType == "postgres" { + dsn = req.BuildDSN() + if dsn == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": "请提供有效的 PostgreSQL 连接信息或 DSN"}) + return + } + } + + if err := svc.Database.SaveConfig(c.Request.Context(), dbType, dsn); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{ + "message": "数据库配置已成功保存至配置文件,重启服务后将以新数据库运行", + "type": dbType, + }) + } +} diff --git a/internal/handler/admin_database_test.go b/internal/handler/admin_database_test.go new file mode 100644 index 0000000..b948a85 --- /dev/null +++ b/internal/handler/admin_database_test.go @@ -0,0 +1,101 @@ +package handler + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + + "github.com/ShukeBta/MMTL/internal/config" + "github.com/ShukeBta/MMTL/internal/service" +) + +func TestBuildDSN(t *testing.T) { + cases := []struct { + payload DatabaseConnectionPayload + want string + }{ + { + payload: DatabaseConnectionPayload{ + DSN: "postgres://myuser:mypass@10.0.0.1:5432/mydb?sslmode=require", + }, + want: "postgres://myuser:mypass@10.0.0.1:5432/mydb?sslmode=require", + }, + { + payload: DatabaseConnectionPayload{ + Host: "127.0.0.1", + Port: 5432, + User: "postgres", + Password: "secretpassword", + DBName: "mmtl_prod", + SSLMode: "disable", + }, + want: "postgres://postgres:secretpassword@127.0.0.1:5432/mmtl_prod?sslmode=disable", + }, + } + + for _, c := range cases { + got := c.payload.BuildDSN() + if got != c.want { + t.Errorf("BuildDSN() = %q, want %q", got, c.want) + } + } +} + +func TestGetDatabaseStatusHandler(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg := &config.Config{} + cfg.Database.Type = "sqlite" + cfg.Database.DBPath = "./data/mmtl.db" + + svc := &service.Container{ + Database: service.NewDatabaseAdminService(cfg, nil, nil, nil), + } + + r := gin.New() + r.GET("/api/admin/database/status", getDatabaseStatusHandler(svc)) + + req := httptest.NewRequest(http.MethodGet, "/api/admin/database/status", nil) + rec := httptest.NewRecorder() + r.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("expected status 200, got %d: %s", rec.Code, rec.Body.String()) + } + + var resp map[string]any + if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + if resp["type"] != "sqlite" { + t.Fatalf("expected type=sqlite, got %v", resp["type"]) + } +} + +func TestSaveDatabaseConfigHandler(t *testing.T) { + gin.SetMode(gin.TestMode) + dir := t.TempDir() + cfg := &config.Config{} + cfg.App.DataDir = dir + cfg.Database.Type = "sqlite" + + svc := &service.Container{ + Database: service.NewDatabaseAdminService(cfg, nil, nil, nil), + } + + r := gin.New() + r.POST("/api/admin/database/save-config", saveDatabaseConfigHandler(svc)) + + body := bytes.NewBufferString(`{"type":"postgres","host":"localhost","port":5432,"user":"admin","password":"pwd","dbname":"mmtl"}`) + req := httptest.NewRequest(http.MethodPost, "/api/admin/database/save-config", body) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + r.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("expected status 200, got %d: %s", rec.Code, rec.Body.String()) + } +} diff --git a/internal/handler/routes_admin.go b/internal/handler/routes_admin.go index 12b5eae..e2d90b6 100644 --- a/internal/handler/routes_admin.go +++ b/internal/handler/routes_admin.go @@ -21,6 +21,7 @@ func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service. registerAdminRecognitionWordRoutes(admin, svc) registerAdminStrmRoutes(admin, svc) registerAdminScraperRoutes(admin, svc) + registerAdminDatabaseRoutes(admin, svc) } func registerAdminScraperRoutes(admin *gin.RouterGroup, svc *service.Container) { @@ -138,3 +139,10 @@ func registerAdminRecognitionWordRoutes(admin *gin.RouterGroup, svc *service.Con admin.POST("/recognition-words/sync", syncRecognitionWordsHandler(svc)) admin.POST("/recognition-words/test", testRecognitionWordsHandler(svc)) } + +func registerAdminDatabaseRoutes(admin *gin.RouterGroup, svc *service.Container) { + admin.GET("/database/status", getDatabaseStatusHandler(svc)) + admin.POST("/database/test", testDatabaseHandler(svc)) + admin.POST("/database/migrate", migrateDatabaseHandler(svc)) + admin.POST("/database/save-config", saveDatabaseConfigHandler(svc)) +} diff --git a/internal/service/database_admin.go b/internal/service/database_admin.go new file mode 100644 index 0000000..dc0719c --- /dev/null +++ b/internal/service/database_admin.go @@ -0,0 +1,90 @@ +package service + +import ( + "context" + "fmt" + "strings" + + "go.uber.org/zap" + "gorm.io/gorm" + + "github.com/ShukeBta/MMTL/internal/config" + "github.com/ShukeBta/MMTL/internal/database" + "github.com/ShukeBta/MMTL/internal/repository" +) + +// DatabaseAdminService manages database configuration, connectivity testing, and migration. +type DatabaseAdminService struct { + cfg *config.Config + log *zap.Logger + repos *repository.Container + db *gorm.DB +} + +// NewDatabaseAdminService creates a new DatabaseAdminService. +func NewDatabaseAdminService(cfg *config.Config, log *zap.Logger, repos *repository.Container, db *gorm.DB) *DatabaseAdminService { + if log == nil { + log = zap.NewNop() + } + return &DatabaseAdminService{ + cfg: cfg, + log: log, + repos: repos, + db: db, + } +} + +// GetStatus returns the status of the currently active database. +func (s *DatabaseAdminService) GetStatus(ctx context.Context) *database.DatabaseStatus { + return database.InspectDatabaseStatus(s.db, s.cfg) +} + +// TestPostgres verifies connectivity and permissions to the specified PostgreSQL DSN. +func (s *DatabaseAdminService) TestPostgres(ctx context.Context, dsn string) (*database.PostgresTestResult, error) { + return database.TestPostgres(dsn) +} + +// MigrateToPostgres copies all records from the current active database to the target PostgreSQL database. +func (s *DatabaseAdminService) MigrateToPostgres(ctx context.Context, targetDSN string) (*database.DatabaseMigrationResult, error) { + s.log.Info("starting user-initiated database migration to PostgreSQL", zap.String("target", database.MaskDSN(targetDSN))) + res, err := database.MigrateCurrentToPostgres(s.db, targetDSN, 500, s.log) + if err != nil { + s.log.Error("database migration to PostgreSQL failed", zap.Error(err)) + return nil, err + } + s.log.Info("database migration to PostgreSQL completed successfully", + zap.Int64("total_rows", res.TotalRows), + zap.Int64("duration_ms", res.DurationMS), + ) + return res, nil +} + +// SaveConfig persists the database configuration to config.yaml and the database settings table. +func (s *DatabaseAdminService) SaveConfig(ctx context.Context, dbType, dsn string) error { + dbType = strings.TrimSpace(dbType) + dsn = strings.TrimSpace(dsn) + if dbType == "" { + dbType = "postgres" + } + if dbType == "postgres" && dsn == "" { + return fmt.Errorf("PostgreSQL DSN 不能为空") + } + + // 1. 保存到本地 config.yaml + if err := config.SaveDatabaseConfig(dbType, dsn); err != nil { + return fmt.Errorf("保存配置文件失败: %w", err) + } + + // 2. 更新内存配置 + s.cfg.Database.Type = dbType + s.cfg.Database.DSN = dsn + + // 3. 同时更新 settings 存储库作为副本 + if s.repos != nil && s.repos.Setting != nil { + _ = s.repos.Setting.Set(ctx, "database.type", dbType) + _ = s.repos.Setting.Set(ctx, "database.dsn", dsn) + } + + s.log.Info("database configuration saved", zap.String("type", dbType), zap.String("dsn", database.MaskDSN(dsn))) + return nil +} diff --git a/internal/service/service.go b/internal/service/service.go index caf3381..f2f1c93 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -58,11 +58,12 @@ type Container struct { Device *DeviceService Cache *RuntimeCacheService Sessions *SessionTrackerService - RecognitionWords *RecognitionWordsService - Danmaku *DanmakuService - Strm *StrmService + RecognitionWords *RecognitionWordsService + Danmaku *DanmakuService + Strm *StrmService + Database *DatabaseAdminService - stopCtx context.Context + stopCtx context.Context stopCancel context.CancelFunc // ReloadHTTPServer 由 cmd/server 注入。HTTPS 相关设置保存后,handler diff --git a/internal/service/service_builder.go b/internal/service/service_builder.go index fbe7c9d..bd6a0e6 100644 --- a/internal/service/service_builder.go +++ b/internal/service/service_builder.go @@ -118,6 +118,7 @@ func (b *serviceContainerBuilder) initContentServices() { func (b *serviceContainerBuilder) initAccessAndStorageServices() { b.c.PlayProfiles = NewPlayProfileService(b.log, b.repos) b.c.Permissions = NewPermissionService(b.log, b.repos) + b.c.Database = NewDatabaseAdminService(b.cfg, b.log, b.repos, b.repos.DB) b.c.Emby.SetRuntimeCache(b.c.Cache) b.c.Emby.SetSubtitleService(b.c.Subtitle) b.c.Scheduler = NewSchedulerService( diff --git a/web/src/api/admin.ts b/web/src/api/admin.ts index 36fbf16..9d3eda2 100644 --- a/web/src/api/admin.ts +++ b/web/src/api/admin.ts @@ -1,6 +1,45 @@ import { api } from './client' import type { AccessLog, Setting, User } from '../types' +export interface DatabaseStatus { + type: 'sqlite' | 'postgres' + dsn?: string + db_path?: string + open_conns: number + in_use: number + idle: number + max_open_conns: number + table_counts?: Record +} + +export interface PostgresTestResult { + success: boolean + latency_ms?: number + version?: string + message?: string + error?: string +} + +export interface DatabaseMigrationResult { + success: boolean + total_rows: number + table_rows?: Record + duration_ms: number + message?: string + error?: string +} + +export interface DatabaseConnectionPayload { + type?: string + dsn?: string + host?: string + port?: number + user?: string + password?: string + dbname?: string + sslmode?: string +} + export interface SystemUpdateStatus { image: string current_version?: string @@ -55,21 +94,33 @@ export const adminAPI = { systemUpdateApply: () => api.post('/admin/system/update/apply').then((r) => r.data), - testAdultScraper: (payload: { - engine?: string - server_url?: string - token?: string - javdb_url?: string - javbus_url?: string - cookie?: string - }) => - api - .post<{ - success: boolean - latency_ms?: number - providers?: string[] - message?: string - error?: string - }>('/admin/adult/test-scraper', payload) - .then((r) => r.data), -} + testAdultScraper: (payload: { + engine?: string + server_url?: string + token?: string + javdb_url?: string + javbus_url?: string + cookie?: string + }) => + api + .post<{ + success: boolean + latency_ms?: number + providers?: string[] + message?: string + error?: string + }>('/admin/adult/test-scraper', payload) + .then((r) => r.data), + + getDatabaseStatus: () => + api.get('/admin/database/status').then((r) => r.data), + + testDatabaseConnection: (payload: DatabaseConnectionPayload) => + api.post('/admin/database/test', payload).then((r) => r.data), + + migrateDatabase: (payload: DatabaseConnectionPayload) => + api.post('/admin/database/migrate', payload).then((r) => r.data), + + saveDatabaseConfig: (payload: DatabaseConnectionPayload) => + api.post<{ message: string; type: string }>('/admin/database/save-config', payload).then((r) => r.data), + } diff --git a/web/src/pages/DatabaseSettingsPanel.tsx b/web/src/pages/DatabaseSettingsPanel.tsx new file mode 100644 index 0000000..7711a83 --- /dev/null +++ b/web/src/pages/DatabaseSettingsPanel.tsx @@ -0,0 +1,449 @@ +import { useEffect, useState } from 'react' +import toast from 'react-hot-toast' +import { + Activity, + ArrowRightLeft, + CheckCircle2, + Database, + HardDrive, + HelpCircle, + Loader2, + RefreshCw, + Save, + Server, + ShieldCheck, + Zap, +} from 'lucide-react' + +import { + adminAPI, + type DatabaseConnectionPayload, + type DatabaseMigrationResult, + type DatabaseStatus, + type PostgresTestResult, +} from '../api/admin' +import { confirmAction } from '../components/confirmAction' + +export function DatabaseSettingsPanel() { + const [status, setStatus] = useState(null) + const [loading, setLoading] = useState(true) + const [mode, setMode] = useState<'form' | 'dsn'>('form') + + // 表单状态 + const [formData, setFormData] = useState({ + host: '127.0.0.1', + port: 5432, + user: 'postgres', + password: '', + dbname: 'mmtl', + sslmode: 'disable', + dsn: '', + }) + + // 测试与操作状态 + const [testing, setTesting] = useState(false) + const [testResult, setTestResult] = useState(null) + const [migrating, setMigrating] = useState(false) + const [migrationResult, setMigrationResult] = useState(null) + const [saving, setSaving] = useState(false) + + const refreshStatus = () => { + setLoading(true) + return adminAPI + .getDatabaseStatus() + .then(setStatus) + .catch((err) => toast.error('获取数据库状态失败: ' + (err.message || '网络错误'))) + .finally(() => setLoading(false)) + } + + useEffect(() => { + refreshStatus().catch(() => undefined) + }, []) + + const getPayload = (): DatabaseConnectionPayload => { + if (mode === 'dsn') { + return { type: 'postgres', dsn: formData.dsn?.trim() || '' } + } + return { + type: 'postgres', + host: formData.host?.trim() || '', + port: Number(formData.port) || 5432, + user: formData.user?.trim() || '', + password: formData.password || '', + dbname: formData.dbname?.trim() || 'mmtl', + sslmode: formData.sslmode || 'disable', + } + } + + const handleTestConnection = async () => { + const payload = getPayload() + if (mode === 'form' && (!payload.host || !payload.user)) { + toast.error('请填写 PostgreSQL 主机和用户名') + return + } + if (mode === 'dsn' && !payload.dsn) { + toast.error('请填写 PostgreSQL DSN') + return + } + + setTesting(true) + setTestResult(null) + try { + const res = await adminAPI.testDatabaseConnection(payload) + setTestResult(res) + if (res.success) { + toast.success(`连接成功!延迟: ${res.latency_ms}ms`) + } else { + toast.error(res.error || '连接失败') + } + } catch (err: any) { + const errorMsg = err.response?.data?.error || err.message || '测试连接异常' + setTestResult({ success: false, error: errorMsg }) + toast.error(errorMsg) + } finally { + setTesting(false) + } + } + + const handleMigrate = async () => { + const payload = getPayload() + if (status?.type === 'postgres') { + const ok = await confirmAction({ + title: '覆盖/同步确认', + message: '当前已经处于 PostgreSQL 模式,继续迁移将覆盖/合并目标库的数据,确定继续吗?', + }) + if (!ok) return + } else { + const ok = await confirmAction({ + title: '开始数据库迁移', + message: '即将把当前 SQLite 数据库中的所有媒体、用户、播放记录、设置等全量迁移到目标 PostgreSQL 数据库。确定开始吗?', + }) + if (!ok) return + } + + setMigrating(true) + setMigrationResult(null) + try { + const res = await adminAPI.migrateDatabase(payload) + setMigrationResult(res) + if (res.success) { + toast.success(`数据迁移完成!共迁移 ${res.total_rows} 条记录`) + } else { + toast.error(res.error || '数据迁移失败') + } + } catch (err: any) { + const errorMsg = err.response?.data?.error || err.message || '迁移发生错误' + setMigrationResult({ success: false, total_rows: 0, duration_ms: 0, error: errorMsg }) + toast.error(errorMsg) + } finally { + setMigrating(false) + } + } + + const handleSaveAndSwitch = async () => { + const payload = getPayload() + const ok = await confirmAction({ + title: '切换数据库', + message: '保存后系统配置将更新为使用 PostgreSQL。需要重启 MMTL 服务使新数据库生效。确定保存吗?', + }) + if (!ok) return + + setSaving(true) + try { + const res = await adminAPI.saveDatabaseConfig(payload) + toast.success(res.message || '数据库配置已保存,请重启服务生效') + refreshStatus() + } catch (err: any) { + toast.error(err.response?.data?.error || err.message || '保存配置失败') + } finally { + setSaving(false) + } + } + + return ( +
+ {/* 头部标题 */} +
+
+ +
+

数据库设置与迁移

+

+ 管理系统底层数据库,支持在 SQLite(本地嵌入式)与 PostgreSQL(高性能关系库)之间平滑切换与数据迁移 +

+
+
+ +
+ + {/* 当前数据库状态卡片 */} +
+
+
+ + 当前运行引擎 +
+
+ + + {status?.type === 'postgres' ? 'PostgreSQL' : 'SQLite (WAL 优化)'} + +
+
+ +
+
+

存储位置 / 连接

+

+ {status?.type === 'postgres' + ? status.dsn || '配置的 PostgreSQL 实例' + : status?.db_path || './data/mmtl.db'} +

+
+
+

连接池活跃 / 最大

+

+ 活跃: {status?.in_use ?? 0} · 空闲: {status?.idle ?? 0} · 上限:{' '} + {status?.max_open_conns ?? 16} +

+
+
+

核心表记录概览

+

+ 媒体: {status?.table_counts?.media ?? 0} · 用户: {status?.table_counts?.users ?? 0} · + 播放记录: {status?.table_counts?.playback_histories ?? 0} +

+
+
+
+ + {/* 配置 PostgreSQL */} +
+
+
+ + 配置目标 PostgreSQL +
+
+ + +
+
+ + {mode === 'form' ? ( +
+
+ + setFormData({ ...formData, host: e.target.value })} + placeholder="例如 127.0.0.1 或 postgres" + className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none" + /> +
+
+ + setFormData({ ...formData, port: Number(e.target.value) })} + placeholder="5432" + className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none" + /> +
+
+ + setFormData({ ...formData, dbname: e.target.value })} + placeholder="mmtl" + className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none" + /> +
+
+ + setFormData({ ...formData, user: e.target.value })} + placeholder="postgres" + className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none" + /> +
+
+ + setFormData({ ...formData, password: e.target.value })} + placeholder="••••••••" + className="w-full rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none" + /> +
+
+ + +
+
+ ) : ( +
+ + setFormData({ ...formData, dsn: e.target.value })} + placeholder="postgres://user:password@127.0.0.1:5432/mmtl?sslmode=disable" + className="w-full font-mono text-xs rounded-lg border border-gray-200 bg-sand-200/40 px-3 py-2 text-ink-600 placeholder:text-gray-400 focus:border-brand-500 focus:outline-none" + /> +
+ )} + + {/* 测试结果卡片 */} + {testResult && ( +
+ {testResult.success ? ( + + ) : ( + + )} +
+

+ {testResult.success ? `连接测试通过 (${testResult.latency_ms} ms)` : '连接测试未通过'} +

+ {testResult.version &&

{testResult.version}

} + {testResult.error &&

{testResult.error}

} +
+
+ )} + + {/* 迁移结果卡片 */} + {migrationResult && ( +
+
+ + {migrationResult.message || (migrationResult.success ? '迁移完成' : '迁移失败')} + {migrationResult.success && ( + + (耗时: {migrationResult.duration_ms} ms) + + )} +
+ {migrationResult.table_rows && Object.keys(migrationResult.table_rows).length > 0 && ( +
+ {Object.entries(migrationResult.table_rows).map(([tbl, count]) => ( +
+ {tbl}: + {count} 条 +
+ ))} +
+ )} + {migrationResult.error && ( +

{migrationResult.error}

+ )} +
+ )} + + {/* 操作按钮区 */} +
+ + +
+ + + +
+
+
+
+ ) +} diff --git a/web/src/pages/SettingsPage.tsx b/web/src/pages/SettingsPage.tsx index 266a336..8db7465 100644 --- a/web/src/pages/SettingsPage.tsx +++ b/web/src/pages/SettingsPage.tsx @@ -8,6 +8,7 @@ import { libraryAPI } from '../api/library' import type { Library, Setting } from '../types' import { APIConfigsPanel } from '../components/APIConfigsPanel' import { AdultSettingsPanel } from './AdultSettingsPanel' +import { DatabaseSettingsPanel } from './DatabaseSettingsPanel' import { RecognitionWordsPanel } from './RecognitionWordsPanel' import { SettingRow } from './SettingsRow' import { ALL_KEYS, GROUPS } from './settingsGroups' @@ -167,6 +168,7 @@ export function SettingsPage() { {!loading && (
+ {group.key === 'database' && } {group.key === 'api-configs' && } {group.key === 'recognition-words' && } {group.key === 'adult' && } diff --git a/web/src/pages/settingsGroups.ts b/web/src/pages/settingsGroups.ts index 877f12c..fe24508 100644 --- a/web/src/pages/settingsGroups.ts +++ b/web/src/pages/settingsGroups.ts @@ -7,8 +7,16 @@ import type { SettingGroup } from './settingsGroupTypes' export type { SettingGroup } from './settingsGroupTypes' +export const databaseSettingsGroup: SettingGroup = { + key: 'database', + label: '数据库', + description: '配置底层数据库(SQLite / PostgreSQL)及数据平滑迁移', + items: [], +} + export const GROUPS: SettingGroup[] = [ generalSettingsGroup, + databaseSettingsGroup, apiConfigsSettingsGroup, recognitionWordsSettingsGroup, danmakuSettingsGroup,