Files
ai_site/platform/internal/schema/runner.go
2026-07-31 10:19:22 +08:00

93 lines
2.2 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}