package store import ( "database/sql" "fmt" "log" ) // migrateUserProfileColumns adds avatar_url and state when missing (existing DBs). func migrateUserProfileColumns(db *sql.DB) error { cols := []string{"avatar_url", "state"} for _, col := range cols { stmt := fmt.Sprintf(`ALTER TABLE users ADD COLUMN IF NOT EXISTS %s TEXT NOT NULL DEFAULT ''`, col) if _, err := db.Exec(stmt); err != nil { return fmt.Errorf("add column %s: %w", col, err) } } return nil } const migrateLockKey int64 = 0x706c756d5f6d6967 // "plum_mig" // applyMigrations runs versioned migrations under an advisory lock. // Fresh databases apply schemaSQL as version 001; later versions are incremental. func applyMigrations(db *sql.DB, schemaSQL string) error { tx, err := db.Begin() if err != nil { return err } defer tx.Rollback() if _, err := tx.Exec(`SELECT pg_advisory_xact_lock($1)`, migrateLockKey); err != nil { return fmt.Errorf("migrate lock: %w", err) } if _, err := tx.Exec(` CREATE TABLE IF NOT EXISTS schema_migrations ( version TEXT PRIMARY KEY, applied_at TIMESTAMPTZ NOT NULL DEFAULT now() )`); err != nil { return fmt.Errorf("schema_migrations: %w", err) } if err := tx.Commit(); err != nil { return err } applied, err := appliedVersions(db) if err != nil { return err } migrations := []struct { version string run func(*sql.DB) error }{ {"001_schema", func(db *sql.DB) error { return applySchema(db, schemaSQL) }}, {"002_user_profile_columns", migrateUserProfileColumns}, } for _, m := range migrations { if applied[m.version] { continue } log.Printf("migrate: applying %s", m.version) if err := m.run(db); err != nil { return fmt.Errorf("migrate %s: %w", m.version, err) } if _, err := db.Exec(`INSERT INTO schema_migrations (version) VALUES ($1)`, m.version); err != nil { return fmt.Errorf("record %s: %w", m.version, err) } } return nil } func appliedVersions(db *sql.DB) (map[string]bool, error) { rows, err := db.Query(`SELECT version FROM schema_migrations`) if err != nil { return nil, err } defer rows.Close() out := map[string]bool{} for rows.Next() { var v string if err := rows.Scan(&v); err != nil { return nil, err } out[v] = true } return out, rows.Err() }