Refactor codebase for Phase 3
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
package dbutil
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/go-sql-driver/mysql"
|
||||
)
|
||||
|
||||
func OpenMySQLAndCreateDatabaseIfMissing(dsn string) (*sql.DB, error) {
|
||||
db, err := sql.Open("mysql", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := db.Ping(); err == nil {
|
||||
return db, nil
|
||||
} else {
|
||||
_ = db.Close()
|
||||
var mysqlErr *mysql.MySQLError
|
||||
if !errors.As(err, &mysqlErr) || mysqlErr.Number != 1049 {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
cfg, err := mysql.ParseDSN(dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cfg.DBName == "" {
|
||||
return nil, errors.New("DSN does not include a database name")
|
||||
}
|
||||
databaseName := cfg.DBName
|
||||
cfg.DBName = ""
|
||||
|
||||
adminDB, err := sql.Open("mysql", cfg.FormatDSN())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer adminDB.Close()
|
||||
if err := adminDB.Ping(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := adminDB.Exec("CREATE DATABASE IF NOT EXISTS " + QuoteMySQLIdentifier(databaseName) + " CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
fmt.Printf("created database %q\n", databaseName)
|
||||
|
||||
db, err = sql.Open("mysql", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := db.Ping(); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func QuoteMySQLIdentifier(name string) string {
|
||||
return "`" + strings.ReplaceAll(name, "`", "``") + "`"
|
||||
}
|
||||
|
||||
func IsDuplicateColumnError(err error) bool {
|
||||
var mysqlErr *mysql.MySQLError
|
||||
return errors.As(err, &mysqlErr) && mysqlErr.Number == 1060
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
package dbutil
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestQuoteMySQLIdentifier(t *testing.T) {
|
||||
got := QuoteMySQLIdentifier("git`ocean")
|
||||
if got != "`git``ocean`" {
|
||||
t.Fatalf("unexpected quote: %q", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user