mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 14:46:36 +08:00
[功能] 支持 PostgreSQL 数据库,添加数据库迁移逻辑并更新相关文档
This commit is contained in:
+219
-73
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user