[功能] 支持 PostgreSQL 数据库,添加数据库迁移逻辑并更新相关文档

This commit is contained in:
ryan
2026-03-17 10:24:12 +08:00
parent 2c17f3289b
commit cfd7c3d7ca
13 changed files with 465 additions and 163 deletions
+219 -73
View File
@@ -1,17 +1,66 @@
package model
import (
"fmt"
"github.com/glebarez/sqlite"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/schema"
"log/slog"
"openflare/common"
"openflare/utils/security"
"os"
"reflect"
"sync"
)
var DB *gorm.DB
type dbModel struct {
value any
tableName string
hasIDPK bool
}
func registeredModels() []any {
return []any{
&File{},
&User{},
&Option{},
&ProxyRoute{},
&ConfigVersion{},
&Node{},
&NodeSystemProfile{},
&ApplyLog{},
&NodeMetricSnapshot{},
&NodeRequestReport{},
&NodeAccessLog{},
&NodeHealthEvent{},
&TLSCertificate{},
&ManagedDomain{},
}
}
func buildDBModels() ([]dbModel, error) {
models := registeredModels()
result := make([]dbModel, 0, len(models))
namer := schema.NamingStrategy{}
cache := &sync.Map{}
for _, item := range models {
parsed, err := schema.Parse(item, cache, namer)
if err != nil {
return nil, err
}
hasIDPK := len(parsed.PrimaryFields) == 1 && parsed.PrimaryFields[0].DBName == "id"
result = append(result, dbModel{
value: item,
tableName: parsed.Table,
hasIDPK: hasIDPK,
})
}
return result, nil
}
func migrateProxyRouteEnableHTTPSColumn(db *gorm.DB) error {
if !db.Migrator().HasTable(&ProxyRoute{}) {
return nil
@@ -48,88 +97,185 @@ func CountTable(tableName string) (num int64) {
return
}
func InitDB() (err error) {
var db *gorm.DB
if os.Getenv("SQL_DSN") != "" {
// Use MySQL
db, err = gorm.Open(mysql.Open(os.Getenv("SQL_DSN")), &gorm.Config{
PrepareStmt: true, // precompile SQL
func openDatabase() (*gorm.DB, string, error) {
if common.SQLDSN != "" {
db, err := gorm.Open(postgres.Open(common.SQLDSN), &gorm.Config{
PrepareStmt: true,
})
} else {
// Use SQLite
db, err = gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{
PrepareStmt: true, // precompile SQL
})
slog.Info("SQL_DSN not set, using SQLite as database")
if err != nil {
return nil, "", err
}
return db, "postgres", nil
}
if err == nil {
DB = db
if err = migrateProxyRouteEnableHTTPSColumn(db); err != nil {
db, err := gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{
PrepareStmt: true,
})
if err != nil {
return nil, "", err
}
slog.Info("database DSN not set, using SQLite as database", "sqlite_path", common.SQLitePath)
return db, "sqlite", nil
}
func autoMigrateAll(db *gorm.DB) error {
for _, item := range registeredModels() {
if err := db.AutoMigrate(item); err != nil {
return err
}
err := db.AutoMigrate(&File{})
if err != nil {
return err
}
return nil
}
func isDatabaseEmpty(db *gorm.DB) (bool, error) {
for _, item := range registeredModels() {
var count int64
if err := db.Model(item).Limit(1).Count(&count).Error; err != nil {
return false, err
}
err = db.AutoMigrate(&User{})
if err != nil {
return err
if count > 0 {
return false, nil
}
err = db.AutoMigrate(&Option{})
if err != nil {
return err
}
err = db.AutoMigrate(&ProxyRoute{})
if err != nil {
return err
}
err = db.AutoMigrate(&ConfigVersion{})
if err != nil {
return err
}
err = db.AutoMigrate(&Node{})
if err != nil {
return err
}
err = db.AutoMigrate(&NodeSystemProfile{})
if err != nil {
return err
}
err = db.AutoMigrate(&ApplyLog{})
if err != nil {
return err
}
err = db.AutoMigrate(&NodeMetricSnapshot{})
if err != nil {
return err
}
err = db.AutoMigrate(&NodeRequestReport{})
if err != nil {
return err
}
err = db.AutoMigrate(&NodeAccessLog{})
if err != nil {
return err
}
err = db.AutoMigrate(&NodeHealthEvent{})
if err != nil {
return err
}
err = db.AutoMigrate(&TLSCertificate{})
if err != nil {
return err
}
err = db.AutoMigrate(&ManagedDomain{})
if err != nil {
return err
}
err = createRootAccountIfNeed()
}
return true, nil
}
func sqliteSourceExists() bool {
info, err := os.Stat(common.SQLitePath)
if err != nil {
return false
}
return !info.IsDir()
}
func migrateSQLiteDataIfNeeded(target *gorm.DB, backend string) error {
if backend != "postgres" {
return nil
}
empty, err := isDatabaseEmpty(target)
if err != nil {
return err
} else {
}
if !empty {
slog.Info("skip sqlite migration because target database already has data", "backend", backend)
return nil
}
if !sqliteSourceExists() {
slog.Info("skip sqlite migration because sqlite source file was not found", "sqlite_path", common.SQLitePath)
return nil
}
source, err := gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{
PrepareStmt: true,
})
if err != nil {
return fmt.Errorf("open sqlite source database failed: %w", err)
}
sourceSQLDB, err := source.DB()
if err != nil {
return fmt.Errorf("get sqlite source database handle failed: %w", err)
}
defer func() {
_ = sourceSQLDB.Close()
}()
models, err := buildDBModels()
if err != nil {
return err
}
slog.Info("starting sqlite to postgres database migration", "sqlite_path", common.SQLitePath)
err = target.Transaction(func(tx *gorm.DB) error {
for _, item := range models {
if err := migrateTableData(source, tx, item); err != nil {
return err
}
if item.hasIDPK {
if err := resetPostgresSequence(tx, item.tableName); err != nil {
return err
}
}
}
return nil
})
if err != nil {
return err
}
slog.Info("sqlite to postgres database migration completed", "sqlite_path", common.SQLitePath)
return nil
}
func migrateTableData(source *gorm.DB, target *gorm.DB, item dbModel) error {
if !source.Migrator().HasTable(item.value) {
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", 0, "status", "skipped_missing_source_table")
return nil
}
var total int64
if err := source.Model(item.value).Count(&total).Error; err != nil {
return fmt.Errorf("count sqlite table %s failed: %w", item.tableName, err)
}
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", total, "status", "starting")
if total == 0 {
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", total, "status", "completed")
return nil
}
modelType := reflect.TypeOf(item.value).Elem()
sliceType := reflect.SliceOf(modelType)
migrated := int64(0)
offset := 0
const batchSize = 200
for {
batchPtr := reflect.New(sliceType)
query := source.Model(item.value).Limit(batchSize).Offset(offset)
if item.hasIDPK {
query = query.Order("id ASC")
}
if err := query.Find(batchPtr.Interface()).Error; err != nil {
return fmt.Errorf("read sqlite table %s failed: %w", item.tableName, err)
}
batchLen := batchPtr.Elem().Len()
if batchLen == 0 {
break
}
if err := target.Create(batchPtr.Interface()).Error; err != nil {
return fmt.Errorf("write target table %s failed: %w", item.tableName, err)
}
migrated += int64(batchLen)
offset += batchLen
slog.Info("database migration progress", "table", item.tableName, "migrated", migrated, "total", total, "status", "running")
}
slog.Info("database migration progress", "table", item.tableName, "migrated", migrated, "total", total, "status", "completed")
return nil
}
func resetPostgresSequence(db *gorm.DB, tableName string) error {
sql := fmt.Sprintf(
"SELECT setval(pg_get_serial_sequence('%s', 'id'), COALESCE(MAX(id), 1), MAX(id) IS NOT NULL) FROM \"%s\"",
tableName,
tableName,
)
return db.Exec(sql).Error
}
func InitDB() (err error) {
db, backend, err := openDatabase()
if err != nil {
slog.Error("open database failed", "error", err)
os.Exit(1)
}
return err
DB = db
if err = migrateProxyRouteEnableHTTPSColumn(db); err != nil {
return err
}
if err = autoMigrateAll(db); err != nil {
return err
}
if err = migrateSQLiteDataIfNeeded(db, backend); err != nil {
return err
}
return createRootAccountIfNeed()
}
func CloseDB() error {
+123
View File
@@ -0,0 +1,123 @@
package model
import (
"path/filepath"
"testing"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
func openTestSQLiteDB(t *testing.T, name string) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), name)), &gorm.Config{})
if err != nil {
t.Fatalf("open sqlite db: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("get sql db: %v", err)
}
t.Cleanup(func() {
_ = sqlDB.Close()
})
return db
}
func findDBModelByTableName(t *testing.T, tableName string) dbModel {
t.Helper()
models, err := buildDBModels()
if err != nil {
t.Fatalf("build db models: %v", err)
}
for _, item := range models {
if item.tableName == tableName {
return item
}
}
t.Fatalf("db model not found for table %s", tableName)
return dbModel{}
}
func TestIsDatabaseEmpty(t *testing.T) {
db := openTestSQLiteDB(t, "empty.db")
empty, err := isDatabaseEmpty(db)
if err != nil {
t.Fatalf("isDatabaseEmpty returned error: %v", err)
}
if !empty {
t.Fatal("expected database to be empty")
}
if err := db.Create(&User{
Username: "alice",
Password: "secret",
DisplayName: "Alice",
Role: 1,
Status: 1,
}).Error; err != nil {
t.Fatalf("seed user: %v", err)
}
empty, err = isDatabaseEmpty(db)
if err != nil {
t.Fatalf("isDatabaseEmpty after seed returned error: %v", err)
}
if empty {
t.Fatal("expected database to be non-empty")
}
}
func TestMigrateTableDataCopiesRows(t *testing.T) {
source := openTestSQLiteDB(t, "source.db")
target := openTestSQLiteDB(t, "target.db")
user := User{
Id: 1,
Username: "root",
Password: "hashed",
DisplayName: "Root User",
Role: 100,
Status: 1,
}
option := Option{
Key: "AgentHeartbeatInterval",
Value: "10000",
}
if err := source.Create(&user).Error; err != nil {
t.Fatalf("seed source user: %v", err)
}
if err := source.Create(&option).Error; err != nil {
t.Fatalf("seed source option: %v", err)
}
if err := migrateTableData(source, target, findDBModelByTableName(t, "users")); err != nil {
t.Fatalf("migrate users: %v", err)
}
if err := migrateTableData(source, target, findDBModelByTableName(t, "options")); err != nil {
t.Fatalf("migrate options: %v", err)
}
var gotUser User
if err := target.First(&gotUser, 1).Error; err != nil {
t.Fatalf("query migrated user: %v", err)
}
if gotUser.Username != user.Username || gotUser.DisplayName != user.DisplayName {
t.Fatalf("unexpected migrated user: %+v", gotUser)
}
var gotOption Option
if err := target.First(&gotOption, "key = ?", option.Key).Error; err != nil {
t.Fatalf("query migrated option: %v", err)
}
if gotOption.Value != option.Value {
t.Fatalf("unexpected migrated option value: %s", gotOption.Value)
}
}