254 lines
7.0 KiB
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/app/system/internal/conf"
|
|
"kra/pkg/database/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...)
|
|
}
|