db_migrate.go raw

   1  package db
   2  
   3  import (
   4  	"fmt"
   5  	"slices"
   6  
   7  	"gorm.io/gorm"
   8  
   9  	"github.com/getAlby/hub/logger"
  10  )
  11  
  12  var expectedTables = []string{
  13  	"apps",
  14  	"app_permissions",
  15  	"request_events",
  16  	"response_events",
  17  	"transactions",
  18  	"swaps",
  19  	"user_configs",
  20  	"migrations",
  21  	"forwards",
  22  }
  23  
  24  // MigrateDB copies all rows from one database to another. Both databases
  25  // must have an up-to-date schema (they are checked against expectedTables).
  26  // Orphaned request and response events are deleted from the source database
  27  // before copying, as they would violate foreign key constraints in the
  28  // destination database.
  29  func MigrateDB(from, to *gorm.DB) error {
  30  	if err := checkSchema(from); err != nil {
  31  		return fmt.Errorf("source database schema check failed: %w", err)
  32  	}
  33  
  34  	if err := checkSchema(to); err != nil {
  35  		return fmt.Errorf("destination database schema check failed: %w", err)
  36  	}
  37  
  38  	// NOTE: we assume that excess request events have already been cleaned up due to the background task
  39  	// and only a maximum of ~1000 remain.
  40  	logger.Logger.Info("Deleting orphaned request events.")
  41  	err := from.Exec("DELETE FROM request_events WHERE app_id NOT IN (SELECT id FROM apps);").Error
  42  	if err != nil {
  43  		return fmt.Errorf("failed to delete orphaned request events: %w", err)
  44  	}
  45  
  46  	// NOTE: we assume that excess response events have already been cleaned up due to the background task
  47  	// and only a maximum of ~1000 remain.
  48  	logger.Logger.Info("Deleting orphaned response events.")
  49  	err = from.Exec("DELETE FROM response_events WHERE request_id NOT IN (SELECT id FROM request_events);").Error
  50  	if err != nil {
  51  		return fmt.Errorf("failed to delete orphaned response events: %w", err)
  52  	}
  53  
  54  	tx := to.Begin()
  55  	defer tx.Rollback()
  56  
  57  	if err := tx.Error; err != nil {
  58  		return fmt.Errorf("failed to start transaction: %w", err)
  59  	}
  60  
  61  	// Table migration order matters: referenced tables must be migrated
  62  	// before referencing tables.
  63  
  64  	logger.Logger.Info("migrating apps...")
  65  	if err := migrateTable[App](from, tx); err != nil {
  66  		return fmt.Errorf("failed to migrate apps: %w", err)
  67  	}
  68  
  69  	logger.Logger.Info("migrating app_permissions...")
  70  	if err := migrateTable[AppPermission](from, tx); err != nil {
  71  		return fmt.Errorf("failed to migrate app_permissions: %w", err)
  72  	}
  73  
  74  	logger.Logger.Info("migrating request_events...")
  75  	if err := migrateTable[RequestEvent](from, tx); err != nil {
  76  		return fmt.Errorf("failed to migrate request_events: %w", err)
  77  	}
  78  
  79  	logger.Logger.Info("migrating response_events...")
  80  	if err := migrateTable[ResponseEvent](from, tx); err != nil {
  81  		return fmt.Errorf("failed to migrate response_events: %w", err)
  82  	}
  83  
  84  	logger.Logger.Info("migrating transactions...")
  85  	if err := migrateTable[Transaction](from, tx); err != nil {
  86  		return fmt.Errorf("failed to migrate transactions: %w", err)
  87  	}
  88  
  89  	logger.Logger.Info("migrating swaps...")
  90  	if err := migrateTable[Swap](from, tx); err != nil {
  91  		return fmt.Errorf("failed to migrate swaps: %w", err)
  92  	}
  93  
  94  	logger.Logger.Info("migrating forwards...")
  95  	if err := migrateTable[Forward](from, tx); err != nil {
  96  		return fmt.Errorf("failed to migrate forwards: %w", err)
  97  	}
  98  
  99  	logger.Logger.Info("migrating user_configs...")
 100  	if err := migrateTable[UserConfig](from, tx); err != nil {
 101  		return fmt.Errorf("failed to migrate user_configs: %w", err)
 102  	}
 103  
 104  	if to.Dialector.Name() == "postgres" {
 105  		logger.Logger.Info("resetting sequences...")
 106  		if err := resetSequences(tx); err != nil {
 107  			return fmt.Errorf("failed to reset sequences: %w", err)
 108  		}
 109  	}
 110  
 111  	tx.Commit()
 112  	if err := tx.Error; err != nil {
 113  		return fmt.Errorf("failed to commit transaction: %w", err)
 114  	}
 115  
 116  	return nil
 117  }
 118  
 119  func migrateTable[T any](from, to *gorm.DB) error {
 120  	var data []T
 121  	if err := from.Find(&data).Error; err != nil {
 122  		return fmt.Errorf("failed to fetch data: %w", err)
 123  	}
 124  
 125  	if len(data) == 0 {
 126  		return nil
 127  	}
 128  
 129  	// to avoid "failed to migrate transactions: failed to insert data: extended protocol limited to 65535 parameters"
 130  	// see https://stackoverflow.com/questions/77372430/extended-protocol-limited-to-65535-parameters-golang-gorm
 131  	// max statements is 65535
 132  	// but it's the number of records * columns
 133  	// to be safe, using a lower value of 1000.
 134  	// this will fail if any table has more than 65 columns, which I doubt we will have
 135  	max := 1000
 136  	for i := 0; i < len(data); i += max {
 137  		j := min(i+max, len(data))
 138  
 139  		if err := to.Create(data[i:j]).Error; err != nil {
 140  			return fmt.Errorf("failed to insert data: %w", err)
 141  		}
 142  	}
 143  
 144  	return nil
 145  }
 146  
 147  func checkSchema(db *gorm.DB) error {
 148  	tables, err := listTables(db)
 149  	if err != nil {
 150  		return fmt.Errorf("failed to list database tables: %w", err)
 151  	}
 152  
 153  	for _, table := range expectedTables {
 154  		if !slices.Contains(tables, table) {
 155  			return fmt.Errorf("table missing from the database: %q", table)
 156  		}
 157  	}
 158  
 159  	for _, table := range tables {
 160  		if !slices.Contains(expectedTables, table) {
 161  			return fmt.Errorf("unexpected table found in the database: %q", table)
 162  		}
 163  	}
 164  
 165  	return nil
 166  }
 167  
 168  func listTables(db *gorm.DB) ([]string, error) {
 169  	var query string
 170  
 171  	switch db.Dialector.Name() {
 172  	case "sqlite":
 173  		query = "SELECT name FROM sqlite_master WHERE type='table'  AND name NOT LIKE 'sqlite_%';"
 174  	case "postgres":
 175  		query = "SELECT tablename FROM pg_tables WHERE schemaname = 'public';"
 176  	default:
 177  		return nil, fmt.Errorf("unsupported database: %q", db.Dialector.Name())
 178  	}
 179  
 180  	rows, err := db.Raw(query).Rows()
 181  	if err != nil {
 182  		return nil, fmt.Errorf("failed to query table names: %w", err)
 183  	}
 184  	defer func() {
 185  		if err := rows.Close(); err != nil {
 186  			logger.Logger.WithError(err).Error("failed to close rows")
 187  		}
 188  	}()
 189  
 190  	var tables []string
 191  	for rows.Next() {
 192  		var table string
 193  		if err := rows.Scan(&table); err != nil {
 194  			return nil, fmt.Errorf("failed to scan table name: %w", err)
 195  		}
 196  		tables = append(tables, table)
 197  	}
 198  
 199  	return tables, nil
 200  }
 201  
 202  func resetSequences(db *gorm.DB) error {
 203  	type resetReq struct {
 204  		table string
 205  		seq   string
 206  	}
 207  
 208  	resetReqs := []resetReq{
 209  		{"apps", "apps_2_id_seq"},
 210  		{"app_permissions", "app_permissions_2_id_seq"},
 211  		{"request_events", "request_events_id_seq"},
 212  		{"response_events", "response_events_id_seq"},
 213  		{"transactions", "transactions_id_seq"},
 214  		{"swaps", "swaps_id_seq"},
 215  		{"forwards", "forwards_id_seq"},
 216  		{"user_configs", "user_configs_id_seq"},
 217  	}
 218  
 219  	for _, req := range resetReqs {
 220  		if err := resetPostgresSequence(db, req.table, req.seq); err != nil {
 221  			return fmt.Errorf("failed to reset sequence %q for %q: %w", req.seq, req.table, err)
 222  		}
 223  	}
 224  
 225  	return nil
 226  }
 227  
 228  func resetPostgresSequence(db *gorm.DB, table string, seq string) error {
 229  	query := fmt.Sprintf("SELECT setval('%s', (SELECT MAX(id) FROM %s));", seq, table)
 230  	if err := db.Exec(query).Error; err != nil {
 231  		return fmt.Errorf("failed to execute setval(): %w", err)
 232  	}
 233  
 234  	return nil
 235  }
 236