This commit is contained in:
Luciano Milani
2026-07-02 12:47:55 +01:00
commit 5de37bb512
132 changed files with 28495 additions and 0 deletions
+185
View File
@@ -0,0 +1,185 @@
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
}