Files
rill/internal/database/database.go
T
2026-09-19 16:02:30 +08:00

119 lines
3.1 KiB
Go

// Package database 负责数据库连接的建立、健康检查与关闭。
package database
import (
"context"
"fmt"
"os"
"path/filepath"
"time"
"github.com/glebarez/sqlite"
"gorm.io/driver/mysql"
"gorm.io/gorm"
"rill/internal/config"
)
const (
driverSQLite = "sqlite3"
driverMySQL = "mysql"
)
// Open 依据配置建立数据库连接,并在 connect_timeout 内完成连通性检查。
func Open(cfg *config.Config) (*gorm.DB, error) {
timeout, err := time.ParseDuration(cfg.Database.ConnectTimeout)
if err != nil {
return nil, fmt.Errorf("解析 database.connect_timeout 失败: %w", err)
}
dialector, err := newDialector(cfg)
if err != nil {
return nil, err
}
db, err := gorm.Open(dialector, &gorm.Config{
Logger: newLogger(cfg.LogLevel()),
DisableAutomaticPing: true,
TranslateError: true,
DisableForeignKeyConstraintWhenMigrating: cfg.Database.Driver == driverSQLite,
})
if err != nil {
return nil, fmt.Errorf("打开数据库失败: %w", err)
}
if err := applyPool(db, cfg); err != nil {
Close(db)
return nil, err
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
if err := Ping(ctx, db); err != nil {
Close(db)
return nil, fmt.Errorf("连接数据库失败: %w", err)
}
return db, nil
}
func newDialector(cfg *config.Config) (gorm.Dialector, error) {
switch cfg.Database.Driver {
case driverSQLite:
path := filepath.Clean(cfg.Database.SQLite.Path)
if dir := filepath.Dir(path); dir != "." {
if err := os.MkdirAll(dir, 0o755); err != nil {
return nil, fmt.Errorf("创建数据库目录 %s 失败: %w", dir, err)
}
}
dsn := fmt.Sprintf("file:%s?_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)&_pragma=foreign_keys(1)", path)
return sqlite.Open(dsn), nil
case driverMySQL:
m := cfg.Database.MySQL
dsn := fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=%s&parseTime=True&loc=Local",
m.User, m.Password, m.Host, m.Port, m.Database, m.Charset)
return mysql.Open(dsn), nil
default:
return nil, fmt.Errorf("不支持的数据库驱动: %q", cfg.Database.Driver)
}
}
// applyPool 应用连接池配置:SQLite 使用单连接串行化写入,避免并发写触发 SQLITE_BUSY。
func applyPool(db *gorm.DB, cfg *config.Config) error {
sqlDB, err := db.DB()
if err != nil {
return fmt.Errorf("获取数据库连接池失败: %w", err)
}
if cfg.Database.Driver == driverMySQL {
m := cfg.Database.MySQL
sqlDB.SetMaxOpenConns(m.MaxOpenConns)
sqlDB.SetMaxIdleConns(m.MaxIdleConns)
if lifetime, err := time.ParseDuration(m.ConnMaxLifetime); err == nil {
sqlDB.SetConnMaxLifetime(lifetime)
}
return nil
}
sqlDB.SetMaxOpenConns(1)
sqlDB.SetMaxIdleConns(1)
return nil
}
// Ping 检查数据库连通性。
func Ping(ctx context.Context, db *gorm.DB) error {
sqlDB, err := db.DB()
if err != nil {
return err
}
return sqlDB.PingContext(ctx)
}
// Close 关闭数据库连接池。
func Close(db *gorm.DB) error {
sqlDB, err := db.DB()
if err != nil {
return err
}
return sqlDB.Close()
}