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 }