93 lines
2.2 KiB
Go
93 lines
2.2 KiB
Go
package schema
|
||
|
||
import (
|
||
"context"
|
||
"database/sql"
|
||
"fmt"
|
||
"net/url"
|
||
"regexp"
|
||
"strings"
|
||
)
|
||
|
||
var dbNameRe = regexp.MustCompile(`^[a-z][a-z0-9_]{1,47}$`)
|
||
|
||
type Runner interface {
|
||
ExecDDL(ctx context.Context, stmts []string) error
|
||
EnsureDatabase(ctx context.Context, dbName string) error
|
||
}
|
||
|
||
type NoopRunner struct{}
|
||
|
||
func (NoopRunner) ExecDDL(context.Context, []string) error { return nil }
|
||
func (NoopRunner) EnsureDatabase(context.Context, string) error { return nil }
|
||
|
||
type PostgresRunner struct {
|
||
DB *sql.DB
|
||
AdminDSN string // 用于 CREATE DATABASE(连到 postgres 库)
|
||
}
|
||
|
||
func (r *PostgresRunner) ExecDDL(ctx context.Context, stmts []string) error {
|
||
tx, err := r.DB.BeginTx(ctx, nil)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
defer func() { _ = tx.Rollback() }()
|
||
for _, s := range stmts {
|
||
if _, err := tx.ExecContext(ctx, s); err != nil {
|
||
return fmt.Errorf("ddl failed: %w\nsql: %s", err, s)
|
||
}
|
||
}
|
||
return tx.Commit()
|
||
}
|
||
|
||
func (r *PostgresRunner) EnsureDatabase(ctx context.Context, dbName string) error {
|
||
if !dbNameRe.MatchString(dbName) {
|
||
return fmt.Errorf("invalid database name: %s", dbName)
|
||
}
|
||
admin := r.DB
|
||
var err error
|
||
if r.AdminDSN != "" {
|
||
admin, err = sql.Open("postgres", r.AdminDSN)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
defer admin.Close()
|
||
}
|
||
var exists bool
|
||
if err := admin.QueryRowContext(ctx,
|
||
`SELECT EXISTS(SELECT 1 FROM pg_database WHERE datname=$1)`, dbName,
|
||
).Scan(&exists); err != nil {
|
||
return err
|
||
}
|
||
if exists {
|
||
return nil
|
||
}
|
||
// CREATE DATABASE 不能在事务中
|
||
_, err = admin.ExecContext(ctx, fmt.Sprintf(`CREATE DATABASE %s`, quoteIdent(dbName)))
|
||
return err
|
||
}
|
||
|
||
// DSNForDatabase 把原 DSN 的库名替换为目标库。
|
||
func DSNForDatabase(baseDSN, dbName string) (string, error) {
|
||
if !dbNameRe.MatchString(dbName) {
|
||
return "", fmt.Errorf("invalid database name")
|
||
}
|
||
u, err := url.Parse(baseDSN)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
u.Path = "/" + dbName
|
||
return u.String(), nil
|
||
}
|
||
|
||
func QuoteIdentExport(name string) string { return quoteIdent(name) }
|
||
|
||
func SanitizeDBName(tenantID int64, slug string) string {
|
||
name := fmt.Sprintf("appdb_t%d_%s", tenantID, slug)
|
||
name = strings.ToLower(name)
|
||
if len(name) > 48 {
|
||
name = name[:48]
|
||
}
|
||
return name
|
||
}
|