86 lines
2.4 KiB
Go
86 lines
2.4 KiB
Go
package store
|
||
|
||
import (
|
||
"fmt"
|
||
"log"
|
||
"os"
|
||
"path/filepath"
|
||
"strings"
|
||
|
||
"github.com/glebarez/sqlite"
|
||
"gorm.io/gorm"
|
||
gormlogger "gorm.io/gorm/logger"
|
||
)
|
||
|
||
// schemaVersion 当前 schema 版本。struct 变更(加列/改列/删列)时递增,
|
||
// 触发一次 AutoMigrate 并把新版本写入库(SQLite 用 PRAGMA user_version)。
|
||
// AutoMigrate 对已有表的列判定不收敛(每次都重建表:CREATE __temp + INSERT SELECT + DROP),
|
||
// 大表上一次重建数十秒且每次重启重演,所以之后版本未变就直接跳过。
|
||
const schemaVersion = 2
|
||
|
||
// Open 打开数据库连接并自动迁移。
|
||
// 开发默认 SQLite(dsn 支持 file:...?_journal_mode=WAL),生产可切 postgres。
|
||
func Open(driver, dsn string) (*gorm.DB, error) {
|
||
var dialector gorm.Dialector
|
||
switch driver {
|
||
case "postgres":
|
||
dialector = postgresDialector(dsn)
|
||
default:
|
||
// 确保 SQLite 文件所在目录存在
|
||
if dir := sqliteDir(dsn); dir != "" {
|
||
_ = os.MkdirAll(dir, 0o755)
|
||
}
|
||
dialector = sqlite.Open(dsn)
|
||
}
|
||
|
||
db, err := gorm.Open(dialector, &gorm.Config{
|
||
Logger: gormlogger.Default.LogMode(gormlogger.Warn),
|
||
})
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
if driver != "postgres" && currentSQLiteVersion(db) >= schemaVersion {
|
||
log.Printf("store: connected driver=%s (schema up-to-date v%d, skip migrate)", driver, schemaVersion)
|
||
return db, nil
|
||
}
|
||
|
||
if err := db.AutoMigrate(AllModels()...); err != nil {
|
||
return nil, err
|
||
}
|
||
if driver != "postgres" {
|
||
setSQLiteVersion(db, schemaVersion)
|
||
}
|
||
log.Printf("store: connected driver=%s (migrated, schema v%d)", driver, schemaVersion)
|
||
// 将已有渠道的超时时间从 120000ms 更新为 300000ms(幂等操作)
|
||
db.Model(&Channel{}).Where("timeout_ms = ?", 120000).Update("timeout_ms", 300000)
|
||
return db, nil
|
||
}
|
||
|
||
// currentSQLiteVersion 读取 PRAGMA user_version。
|
||
func currentSQLiteVersion(db *gorm.DB) int {
|
||
var v int
|
||
db.Raw("PRAGMA user_version").Scan(&v)
|
||
return v
|
||
}
|
||
|
||
// setSQLiteVersion 写入 PRAGMA user_version。
|
||
func setSQLiteVersion(db *gorm.DB, v int) {
|
||
db.Exec(fmt.Sprintf("PRAGMA user_version = %d", v))
|
||
}
|
||
|
||
// sqliteDir 提取 SQLite DSN 中的目录部分(忽略 file: 前缀与查询参数)。
|
||
func sqliteDir(dsn string) string {
|
||
d := dsn
|
||
if i := strings.IndexByte(d, '?'); i >= 0 {
|
||
d = d[:i]
|
||
}
|
||
if strings.HasPrefix(d, "file:") {
|
||
d = d[len("file:"):]
|
||
}
|
||
if d == "" || d == ":memory:" || strings.Contains(d, "::") {
|
||
return ""
|
||
}
|
||
return filepath.Dir(d)
|
||
}
|