kra-new/internal/data/database.go

254 lines
7.0 KiB
Go

package data
import (
"fmt"
"log/slog"
"net"
"net/url"
"os"
"path/filepath"
"regexp"
"strings"
"time"
oracle "github.com/dzwvip/gorm-oracle"
"github.com/glebarez/sqlite"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/driver/sqlserver"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"gorm.io/gorm/schema"
"kra/internal/conf"
"kra/internal/data/gormkit"
)
var databaseNamePattern = regexp.MustCompile(`^[A-Za-z][A-Za-z0-9_$-]*$`)
// databaseConnectionConfigured mirrors the initialization contract of the
// administration backend: an explicitly selected database (or a standalone
// DSN) means a real database connection is expected to exist. The internal
// bootstrap database is deliberately excluded from this state.
func databaseConnectionConfigured(c *conf.Data_Database) bool {
if c == nil {
return false
}
if strings.TrimSpace(c.Name) != "" {
return true
}
return strings.TrimSpace(c.Source) != "" && strings.TrimSpace(c.Host) == "" && strings.TrimSpace(c.Path) == ""
}
func normalizedDriver(driver string) string {
switch strings.ToLower(strings.TrimSpace(driver)) {
case "postgres", "postgresql", "pgsql":
return "pgsql"
case "sqlserver", "mssql":
return "mssql"
case "sqlite", "sqlite3":
return "sqlite"
case "oracle":
return "oracle"
case "", "mysql":
return "mysql"
default:
return ""
}
}
func databaseDSN(c *conf.Data_Database, name string) (string, error) {
if c.Source != "" && c.Host == "" && c.Path == "" {
return c.Source, nil
}
driver := normalizedDriver(c.Driver)
if name == "" {
name = c.Name
}
host := c.Host
if host == "" {
host = "127.0.0.1"
}
switch driver {
case "mysql":
port := c.Port
if port == "" {
port = "3306"
}
query := c.Config
if query == "" {
query = "timeout=5s&parseTime=True&loc=Local&charset=utf8mb4"
}
return fmt.Sprintf("%s:%s@tcp(%s)/%s?%s", c.User, c.Password, net.JoinHostPort(host, port), name, query), nil
case "pgsql":
port := c.Port
if port == "" {
port = "5432"
}
extra := c.Config
if extra == "" {
extra = "sslmode=disable TimeZone=Local"
}
return fmt.Sprintf("host=%s port=%s user=%s password=%s dbname=%s %s", host, port, c.User, c.Password, name, extra), nil
case "mssql":
port := c.Port
if port == "" {
port = "1433"
}
q := url.Values{"database": {name}, "encrypt": {"disable"}}
return (&url.URL{Scheme: "sqlserver", User: url.UserPassword(c.User, c.Password), Host: net.JoinHostPort(host, port), RawQuery: q.Encode()}).String(), nil
case "oracle":
port := c.Port
if port == "" {
port = "1521"
}
return fmt.Sprintf("oracle://%s:%s@%s/%s?%s", url.PathEscape(c.User), url.PathEscape(c.Password), net.JoinHostPort(host, port), url.PathEscape(name), c.Config), nil
case "sqlite":
path := c.Path
if path == "" {
path = "."
}
if name == "" {
name = "kra"
}
if filepath.Ext(name) == "" {
name += ".db"
}
return filepath.Join(path, name), nil
}
return "", fmt.Errorf("unsupported database driver %q", c.Driver)
}
func gormConfig(config *conf.Data_Database, appLogger ...*slog.Logger) *gorm.Config {
level := logger.Info
switch strings.ToLower(config.LogMode) {
case "silent":
level = logger.Silent
case "error":
level = logger.Error
case "warn":
level = logger.Warn
}
var log *slog.Logger
if len(appLogger) > 0 {
log = appLogger[0]
}
return &gorm.Config{Logger: gormkit.NewLogger(log, level), NamingStrategy: schema.NamingStrategy{TablePrefix: config.Prefix, SingularTable: config.Singular}}
}
func openWithDriver(driver, dsn string, appLogger ...*slog.Logger) (*gorm.DB, error) {
return openWithDriverConfig(driver, dsn, &conf.Data_Database{Driver: driver}, appLogger...)
}
func openWithDriverConfig(driver, dsn string, config *conf.Data_Database, appLogger ...*slog.Logger) (*gorm.DB, error) {
gormConfig := gormConfig(config, appLogger...)
var db *gorm.DB
var err error
switch normalizedDriver(driver) {
case "mysql":
db, err = gorm.Open(mysql.Open(dsn), gormConfig)
case "pgsql":
db, err = gorm.Open(postgres.Open(dsn), gormConfig)
case "mssql":
db, err = gorm.Open(sqlserver.Open(dsn), gormConfig)
case "oracle":
db, err = gorm.Open(oracle.Open(dsn), gormConfig)
case "sqlite":
db, err = gorm.Open(sqlite.Open(dsn), gormConfig)
default:
return nil, fmt.Errorf("unsupported database driver %q", driver)
}
if err != nil {
return nil, err
}
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
if config.MaxIdleConns > 0 {
sqlDB.SetMaxIdleConns(int(config.MaxIdleConns))
}
if config.MaxOpenConns > 0 {
sqlDB.SetMaxOpenConns(int(config.MaxOpenConns))
}
if config.ConnMaxLifetime > 0 {
sqlDB.SetConnMaxLifetime(time.Duration(config.ConnMaxLifetime) * time.Second)
}
if config.Engine != "" && (normalizedDriver(driver) == "mysql" || normalizedDriver(driver) == "mssql") {
db = db.Set("gorm:table_options", "ENGINE="+config.Engine)
}
return db, nil
}
func openDatabase(c *conf.Data_Database, create bool, template string, appLogger ...*slog.Logger) (*gorm.DB, error) {
driver := normalizedDriver(c.Driver)
if driver == "" {
return nil, fmt.Errorf("unsupported database driver %q", c.Driver)
}
if driver == "sqlite" {
dsn, err := databaseDSN(c, "")
if err != nil {
return nil, err
}
if err = os.MkdirAll(filepath.Dir(dsn), 0o755); err != nil {
return nil, err
}
return openWithDriverConfig(driver, dsn, c, appLogger...)
}
if create && driver != "oracle" {
if !databaseNamePattern.MatchString(c.Name) {
return nil, fmt.Errorf("invalid database name %q", c.Name)
}
bootstrap := ""
switch driver {
case "pgsql":
bootstrap = "postgres"
case "mssql":
bootstrap = "master"
}
dsn, err := databaseDSN(c, bootstrap)
if err != nil {
return nil, err
}
adminDB, err := openWithDriverConfig(driver, dsn, c, appLogger...)
if err != nil {
return nil, fmt.Errorf("connect database server: %w", err)
}
var statement string
switch driver {
case "mysql":
statement = "CREATE DATABASE IF NOT EXISTS `" + c.Name + "` CHARACTER SET utf8mb4 COLLATE utf8mb4_general_ci"
case "pgsql":
if template == "" {
template = "template0"
}
if !databaseNamePattern.MatchString(template) {
return nil, fmt.Errorf("invalid PostgreSQL template %q", template)
}
var count int64
if err = adminDB.Raw("SELECT count(*) FROM pg_database WHERE datname = ?", c.Name).Scan(&count).Error; err == nil && count == 0 {
statement = `CREATE DATABASE "` + c.Name + `" TEMPLATE "` + template + `"`
}
case "mssql":
statement = "IF DB_ID(N'" + c.Name + "') IS NULL CREATE DATABASE [" + c.Name + "]"
}
if statement != "" {
err = adminDB.Exec(statement).Error
}
if sqlDB, e := adminDB.DB(); e == nil {
_ = sqlDB.Close()
}
if err != nil {
return nil, fmt.Errorf("create database: %w", err)
}
}
dsn, err := databaseDSN(c, "")
if err != nil {
return nil, err
}
return openWithDriverConfig(driver, dsn, c, appLogger...)
}
func openFallbackDatabase(appLogger ...*slog.Logger) (*gorm.DB, error) {
return openWithDriver("sqlite", "file:kra-bootstrap?mode=memory&cache=shared", appLogger...)
}