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 }