186 lines
5.0 KiB
Go
186 lines
5.0 KiB
Go
package database
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"regexp"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/golang-migrate/migrate/v4"
|
|
migratepg "github.com/golang-migrate/migrate/v4/database/postgres"
|
|
_ "github.com/golang-migrate/migrate/v4/source/file"
|
|
"github.com/jackc/pgx/v5/pgxpool"
|
|
"github.com/jackc/pgx/v5/stdlib"
|
|
)
|
|
|
|
var validTenantID = regexp.MustCompile(`^[a-zA-Z0-9_-]{1,63}$`)
|
|
|
|
type DB struct {
|
|
Pool *pgxpool.Pool
|
|
dsn string
|
|
}
|
|
|
|
func New(dsn string) (*DB, error) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
|
|
pool, err := pgxpool.New(ctx, dsn)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("database: failed to create pool: %w", err)
|
|
}
|
|
|
|
if err := pool.Ping(ctx); err != nil {
|
|
pool.Close()
|
|
return nil, fmt.Errorf("database: failed to ping: %w", err)
|
|
}
|
|
|
|
return &DB{Pool: pool, dsn: dsn}, nil
|
|
}
|
|
|
|
func (db *DB) Ping() error {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
defer cancel()
|
|
return db.Pool.Ping(ctx)
|
|
}
|
|
|
|
func (db *DB) Close() {
|
|
db.Pool.Close()
|
|
}
|
|
|
|
func tenantSchema(tenantID string) string {
|
|
return fmt.Sprintf("tenant_%s", strings.ReplaceAll(tenantID, "-", "_"))
|
|
}
|
|
|
|
func (db *DB) SetTenantSchema(ctx context.Context, tenantID string) error {
|
|
if !validTenantID.MatchString(tenantID) {
|
|
return fmt.Errorf("database: invalid tenantID %q", tenantID)
|
|
}
|
|
schema := tenantSchema(tenantID)
|
|
_, err := db.Pool.Exec(ctx, fmt.Sprintf(`SET search_path = "%s", public`, schema))
|
|
return err
|
|
}
|
|
|
|
func (db *DB) ResetSchema(ctx context.Context) error {
|
|
_, err := db.Pool.Exec(ctx, "SET search_path = public")
|
|
return err
|
|
}
|
|
|
|
func (db *DB) stdDB(dsn string) (*sql.DB, error) {
|
|
cfg, err := pgxpool.ParseConfig(dsn)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return stdlib.OpenDB(*cfg.ConnConfig), nil
|
|
}
|
|
|
|
func (db *DB) MigratePublic(dsn, migrationsPath string) error {
|
|
stdDB, err := db.stdDB(dsn)
|
|
if err != nil {
|
|
return fmt.Errorf("migrate: %w", err)
|
|
}
|
|
defer stdDB.Close()
|
|
|
|
driver, err := migratepg.WithInstance(stdDB, &migratepg.Config{SchemaName: "public"})
|
|
if err != nil {
|
|
return fmt.Errorf("migrate: driver: %w", err)
|
|
}
|
|
|
|
m, err := migrate.NewWithDatabaseInstance(
|
|
"file://"+migrationsPath,
|
|
"postgres",
|
|
driver,
|
|
)
|
|
if err != nil {
|
|
return fmt.Errorf("migrate: %w", err)
|
|
}
|
|
|
|
if err := m.Up(); err != nil && err != migrate.ErrNoChange {
|
|
return fmt.Errorf("migrate: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// MigrateTenantSchema ensures the tenant schema exists and runs all pending migrations.
|
|
// Safe to call on existing tenants — golang-migrate tracks state in schema_migrations.
|
|
// The 000001 migration uses IF NOT EXISTS so re-running it is idempotent.
|
|
func (db *DB) MigrateTenantSchema(ctx context.Context, tenantID, migrationsPath string) error {
|
|
if !validTenantID.MatchString(tenantID) {
|
|
return fmt.Errorf("database: invalid tenantID %q", tenantID)
|
|
}
|
|
schema := tenantSchema(tenantID)
|
|
|
|
// Ensure schema exists before handing off to golang-migrate.
|
|
conn, err := db.Pool.Acquire(ctx)
|
|
if err != nil {
|
|
return fmt.Errorf("migrate tenant: acquire: %w", err)
|
|
}
|
|
_, execErr := conn.Exec(ctx, fmt.Sprintf(`CREATE SCHEMA IF NOT EXISTS "%s"`, schema))
|
|
conn.Release()
|
|
if execErr != nil {
|
|
return fmt.Errorf("migrate tenant: create schema: %w", execErr)
|
|
}
|
|
|
|
// Append search_path to DSN so every connection golang-migrate opens
|
|
// automatically resolves unqualified table names to the tenant schema.
|
|
tenantDSN := db.dsn
|
|
sep := "?"
|
|
if strings.Contains(tenantDSN, "?") {
|
|
sep = "&"
|
|
}
|
|
tenantDSN += sep + "search_path=" + schema
|
|
|
|
stdDB, err := db.stdDB(tenantDSN)
|
|
if err != nil {
|
|
return fmt.Errorf("migrate tenant: open stdDB: %w", err)
|
|
}
|
|
defer stdDB.Close()
|
|
|
|
driver, err := migratepg.WithInstance(stdDB, &migratepg.Config{SchemaName: schema})
|
|
if err != nil {
|
|
return fmt.Errorf("migrate tenant: driver: %w", err)
|
|
}
|
|
|
|
m, err := migrate.NewWithDatabaseInstance("file://"+migrationsPath, "postgres", driver)
|
|
if err != nil {
|
|
return fmt.Errorf("migrate tenant: init: %w", err)
|
|
}
|
|
|
|
if err := m.Up(); err != nil && err != migrate.ErrNoChange {
|
|
return fmt.Errorf("migrate tenant %s: %w", tenantID, err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// ProvisionTenantSchema is an alias kept for call-site compatibility.
|
|
func (db *DB) ProvisionTenantSchema(ctx context.Context, tenantID, migrationsPath string) error {
|
|
return db.MigrateTenantSchema(ctx, tenantID, migrationsPath)
|
|
}
|
|
|
|
// MigrateAllTenantSchemas runs pending migrations against every registered tenant.
|
|
// Called at startup so existing tenants always get new migration files applied.
|
|
func (db *DB) MigrateAllTenantSchemas(ctx context.Context, migrationsPath string) error {
|
|
rows, err := db.Pool.Query(ctx, `SELECT id FROM public.tenants`)
|
|
if err != nil {
|
|
return fmt.Errorf("migrate all tenants: query: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
|
|
var ids []string
|
|
for rows.Next() {
|
|
var id string
|
|
if err := rows.Scan(&id); err != nil {
|
|
return fmt.Errorf("migrate all tenants: scan: %w", err)
|
|
}
|
|
ids = append(ids, id)
|
|
}
|
|
|
|
for _, id := range ids {
|
|
if err := db.MigrateTenantSchema(ctx, id, migrationsPath); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return nil
|
|
}
|