Inicial
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
package database_test
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/techxcar/backend/pkg/database"
|
||||
)
|
||||
|
||||
func TestNew_invalidURL(t *testing.T) {
|
||||
_, err := database.New("not-a-valid-url")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestNew_valid(t *testing.T) {
|
||||
url := os.Getenv("TEST_DATABASE_URL")
|
||||
if url == "" {
|
||||
t.Skip("TEST_DATABASE_URL not set, skipping integration test")
|
||||
}
|
||||
|
||||
db, err := database.New(url)
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
|
||||
assert.NoError(t, db.Ping())
|
||||
}
|
||||
Reference in New Issue
Block a user