Files
openteam/server/internal/store/db.go
T

86 lines
2.4 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 = 1
// 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)
}