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...) }