package models import ( "context" "fmt" "os" "strings" "testing" "github.com/google/uuid" "github.com/stretchr/testify/require" "gorm.io/driver/postgres" "gorm.io/gorm" ) func TestMigrateUpgradesLegacyPlansWithoutClosingLockConnection(t *testing.T) { var databaseName string require.NoError(t, DB().Raw("SELECT current_database()").Scan(&databaseName).Error) require.Contains(t, databaseName, "test") original := db schema := "migrate_test_" + strings.ReplaceAll(uuid.NewString(), "-", "") require.NoError(t, original.Exec("CREATE SCHEMA "+schema).Error) var isolatedSQLDB interface{ Close() error } t.Cleanup(func() { SetDB(original) if isolatedSQLDB != nil { _ = isolatedSQLDB.Close() } original.Exec("DROP SCHEMA IF EXISTS " + schema + " CASCADE") }) dsn := fmt.Sprintf("host=%s port=%s user=%s password=%s dbname=%s sslmode=disable search_path=%s", testDatabaseEnv("DATABASE_HOST", "POSTGRES_HOST", "localhost"), testDatabaseEnv("DATABASE_PORT", "POSTGRES_PORT", "5432"), testDatabaseEnv("DATABASE_USER", "POSTGRES_USER", "rsmon"), testDatabaseEnv("DATABASE_PASSWORD", "POSTGRES_PASSWORD", "rsmon"), databaseName, schema, ) isolated, err := gorm.Open(postgres.Open(dsn), &gorm.Config{}) require.NoError(t, err) sqlDB, err := isolated.DB() require.NoError(t, err) isolatedSQLDB = sqlDB RegisterCallbacks(isolated) SetDB(isolated.Set("gorm:association_autoupdate", false)) require.NoError(t, DB().Exec(`CREATE TABLE plans ( id BIGSERIAL PRIMARY KEY, name TEXT NOT NULL, price BIGINT NOT NULL DEFAULT 0, total_monitors BIGINT NOT NULL DEFAULT 0, "default" BOOLEAN NOT NULL DEFAULT FALSE )`).Error) require.NoError(t, DB().Exec(`CREATE TABLE accounts ( id BIGSERIAL PRIMARY KEY, name TEXT NOT NULL, plan_id BIGINT REFERENCES plans(id) )`).Error) require.NoError(t, DB().Exec(`CREATE TABLE subscriptions ( id BIGSERIAL PRIMARY KEY, account_id BIGINT NOT NULL REFERENCES accounts(id), plan_id BIGINT NOT NULL REFERENCES plans(id), provider VARCHAR(16) NOT NULL DEFAULT 'manual', status VARCHAR(24) NOT NULL DEFAULT 'active', billing_cycle VARCHAR(8) NOT NULL DEFAULT 'monthly', currency VARCHAR(3) NOT NULL DEFAULT 'RUB', amount_minor BIGINT NOT NULL DEFAULT 0, metadata_json JSONB NOT NULL DEFAULT '{}' )`).Error) require.NoError(t, DB().Exec(`CREATE TABLE subscription_events ( id BIGSERIAL PRIMARY KEY, subscription_id BIGINT NOT NULL, account_id BIGINT NOT NULL, provider VARCHAR(16) NOT NULL DEFAULT 'manual', kind VARCHAR(32) NOT NULL, from_plan_id BIGINT, to_plan_id BIGINT, payload_json JSONB NOT NULL DEFAULT '{}', created_at TIMESTAMPTZ NOT NULL DEFAULT now() )`).Error) require.NoError(t, DB().Exec(`INSERT INTO plans (id, name, price, total_monitors, "default") VALUES (42, 'Legacy free', 0, 10, TRUE)`).Error) require.NoError(t, DB().Exec("INSERT INTO accounts (id, name, plan_id) VALUES (7, 'Legacy account', 42)").Error) require.NoError(t, DB().Exec("INSERT INTO subscriptions (id, account_id, plan_id) VALUES (9, 7, 42)").Error) require.NoError(t, DB().Exec(`INSERT INTO subscription_events (id, subscription_id, account_id, kind, from_plan_id, to_plan_id) VALUES (11, 9, 7, 'legacy_change', 42, 42)`).Error) Migrate() require.True(t, DB().Migrator().HasColumn("plans", "code")) require.False(t, DB().Migrator().HasTable("plans_legacy")) require.True(t, DB().Migrator().HasTable("plans_legacy_migrated")) var codes []string require.NoError(t, DB().Table("plans").Order("code").Pluck("code", &codes).Error) require.ElementsMatch(t, CanonicalPlanCodes(), codes) var accountPlanCode string require.NoError(t, DB().Table("accounts").Select("plans.code"). Joins("JOIN plans ON plans.id = accounts.plan_id").Where("accounts.id = 7"). Scan(&accountPlanCode).Error) require.Equal(t, "free", accountPlanCode) var subscriptionPlanCode string require.NoError(t, DB().Table("subscriptions").Select("plans.code"). Joins("JOIN plans ON plans.id = subscriptions.plan_id").Where("subscriptions.id = 9"). Scan(&subscriptionPlanCode).Error) require.Equal(t, "free", subscriptionPlanCode) var eventPlanCodes struct { FromCode string ToCode string } require.NoError(t, DB().Table("subscription_events e"). Select("fp.code AS from_code, tp.code AS to_code"). Joins("JOIN plans fp ON fp.id = e.from_plan_id"). Joins("JOIN plans tp ON tp.id = e.to_plan_id"). Where("e.id = 11").Scan(&eventPlanCodes).Error) require.Equal(t, "free", eventPlanCodes.FromCode) require.Equal(t, "free", eventPlanCodes.ToCode) require.NoError(t, DB().Exec("SELECT 1").Error) // The previous implementation retained this name after remapping. Its // canonical FKs must prevent a retry from remapping the same rows again. require.NoError(t, DB().Exec("ALTER TABLE plans_legacy_migrated RENAME TO plans_legacy").Error) require.NotPanics(t, Migrate) accountPlanCode = "" require.NoError(t, DB().Table("accounts").Select("plans.code"). Joins("JOIN plans ON plans.id = accounts.plan_id").Where("accounts.id = 7"). Scan(&accountPlanCode).Error) require.Equal(t, "free", accountPlanCode) require.True(t, DB().Migrator().HasTable("plans_legacy_migrated")) } func TestMigrationAdvisoryLockIsReleasedAfterPanic(t *testing.T) { sqlDB, err := DB().DB() require.NoError(t, err) otherConn, err := sqlDB.Conn(context.Background()) require.NoError(t, err) t.Cleanup(func() { _ = otherConn.Close() }) require.PanicsWithValue(t, "migration failed", func() { withMigrationAdvisoryLock(func() { panic("migration failed") }) }) const migrateAdvisoryLock = int64(1234567890) var acquired bool require.NoError(t, otherConn.QueryRowContext(context.Background(), "SELECT pg_try_advisory_lock($1)", migrateAdvisoryLock).Scan(&acquired)) require.True(t, acquired) _, err = otherConn.ExecContext(context.Background(), "SELECT pg_advisory_unlock($1)", migrateAdvisoryLock) require.NoError(t, err) } func testDatabaseEnv(primary, fallback, defaultValue string) string { if value := os.Getenv(primary); value != "" { return value } if value := os.Getenv(fallback); value != "" { return value } return defaultValue }