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