This commit is contained in:
truewhile
2026-08-28 16:27:11 +08:00
parent 90064a5480
commit 994f64f753
19 changed files with 1304 additions and 51 deletions
+2 -2
View File
@@ -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="
)
+41
View File
@@ -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
}
+36
View File
@@ -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)
}
}
+223
View File
@@ -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
}
+71
View File
@@ -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"])
}
}
+11 -11
View File
@@ -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)
+1 -1
View File
@@ -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
}
+16 -13
View File
@@ -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) {
+25 -2
View File
@@ -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
}
+145
View File
@@ -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,
})
}
}
+101
View File
@@ -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())
}
}
+8
View File
@@ -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))
}
+90
View File
@@ -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
}
+5 -4
View File
@@ -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
+1
View File
@@ -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(