backup_test.go raw
1 package api
2
3 import (
4 "archive/zip"
5 "bytes"
6 "encoding/hex"
7 "io"
8 "os"
9 "path/filepath"
10 "strconv"
11 "strings"
12 "testing"
13 "testing/iotest"
14
15 "github.com/sirupsen/logrus"
16 "github.com/stretchr/testify/require"
17 "gorm.io/datatypes"
18
19 "github.com/getAlby/hub/config"
20 "github.com/getAlby/hub/db"
21 "github.com/getAlby/hub/logger"
22 test_db "github.com/getAlby/hub/tests/db"
23 "github.com/getAlby/hub/tests/mocks"
24 )
25
26 // TestCreateBackup creates a backup from the test database (sqlite by
27 // default, postgres when TEST_DATABASE_URI is set) and verifies that the
28 // archive contains a valid sqlite database with the expected data.
29 func TestCreateBackup(t *testing.T) {
30 logger.Init(strconv.Itoa(int(logrus.DebugLevel)))
31
32 workDir := t.TempDir()
33
34 gormDB, err := test_db.NewDB(t)
35 require.NoError(t, err)
36 defer test_db.CloseDB(gormDB)
37
38 appConfig := &config.AppConfig{
39 Workdir: workDir,
40 DatabaseUri: test_db.GetTestDatabaseURI(),
41 }
42 cfg, err := config.NewConfig(appConfig, gormDB)
43 require.NoError(t, err)
44
45 unlockPassword := ""
46
47 // Represent a fully set-up hub: the unlock-password canary is written during
48 // setup and is required for the password check to pass.
49 require.NoError(t, cfg.SaveUnlockPasswordCheck(unlockPassword))
50
51 app := &db.App{
52 Name: "test",
53 AppPubkey: "2b7dea2866958f17c568cf024e113db7a3baa9c253a9016889196b8d0b11c7ae",
54 Metadata: datatypes.JSON("{}"),
55 }
56 require.NoError(t, gormDB.Create(app).Error)
57
58 lnClient := mocks.NewMockLNClient(t)
59 lnClient.On("GetStorageDir").Return("", nil)
60 lnClient.On("ResetRouter", "ALL").Return(nil)
61
62 svc := mocks.NewMockService(t)
63 svc.On("GetLNClient").Return(lnClient)
64 svc.On("StopApp").Return()
65
66 albyOAuthSvc := mocks.NewMockAlbyOAuthService(t)
67 albyOAuthSvc.On("RemoveOAuthAccessToken").Return(nil)
68
69 theAPI := &api{
70 db: gormDB,
71 cfg: cfg,
72 svc: svc,
73 albyOAuthSvc: albyOAuthSvc,
74 }
75
76 var buf bytes.Buffer
77 err = theAPI.CreateBackup(unlockPassword, &buf)
78 require.NoError(t, err)
79
80 // The temporary database created when converting from postgres must
81 // not be left behind in the working directory.
82 entries, err := os.ReadDir(workDir)
83 require.NoError(t, err)
84 require.Empty(t, entries)
85
86 cr, err := decryptingReader(&buf, unlockPassword)
87 require.NoError(t, err)
88 decrypted, err := io.ReadAll(cr)
89 require.NoError(t, err)
90
91 zr, err := zip.NewReader(bytes.NewReader(decrypted), int64(len(decrypted)))
92 require.NoError(t, err)
93
94 dbFile, err := zr.Open("nwc.db")
95 require.NoError(t, err)
96 dbContents, err := io.ReadAll(dbFile)
97 require.NoError(t, err)
98 require.NoError(t, dbFile.Close())
99
100 restoredPath := filepath.Join(workDir, "restored.db")
101 require.NoError(t, os.WriteFile(restoredPath, dbContents, 0600))
102
103 restoredDB, err := db.NewDB(restoredPath, false)
104 require.NoError(t, err)
105 defer func() {
106 require.NoError(t, db.Stop(restoredDB))
107 }()
108
109 var restoredApp db.App
110 require.NoError(t, restoredDB.First(&restoredApp).Error)
111 require.Equal(t, app.Name, restoredApp.Name)
112 require.Equal(t, app.AppPubkey, restoredApp.AppPubkey)
113 }
114
115 // TestRestoreBackupRejectsPathTraversal verifies that a backup archive
116 // containing an entry whose name points outside the restore directory is
117 // rejected and that no file is written outside it.
118 func TestRestoreBackupRejectsPathTraversal(t *testing.T) {
119 logger.Init(strconv.Itoa(int(logrus.DebugLevel)))
120
121 gormDB, err := test_db.NewDB(t)
122 require.NoError(t, err)
123 defer test_db.CloseDB(gormDB)
124
125 if gormDB.Dialector.Name() != "sqlite" {
126 t.Skip("restore is only supported on sqlite")
127 }
128
129 workDir := t.TempDir()
130
131 appConfig := &config.AppConfig{
132 Workdir: workDir,
133 DatabaseUri: test_db.GetTestDatabaseURI(),
134 }
135 cfg, err := config.NewConfig(appConfig, gormDB)
136 require.NoError(t, err)
137
138 theAPI := &api{
139 db: gormDB,
140 cfg: cfg,
141 }
142
143 unlockPassword := ""
144
145 // The restore directory is <workDir>/restore, so a "../" entry targets a
146 // file directly in the working directory, one level above it.
147 const escapeEntryName = "../pwned.txt"
148 escapeTarget := filepath.Join(workDir, "pwned.txt")
149
150 var buf bytes.Buffer
151 cw, err := encryptingWriter(&buf, unlockPassword)
152 require.NoError(t, err)
153 zw := zip.NewWriter(cw)
154 // A valid entry before the malicious one, to verify that a partially
155 // extracted archive is not left behind when a later entry fails.
156 entryWriter, err := zw.Create("nwc.db")
157 require.NoError(t, err)
158 _, err = entryWriter.Write([]byte("backup contents"))
159 require.NoError(t, err)
160 entryWriter, err = zw.Create(escapeEntryName)
161 require.NoError(t, err)
162 _, err = entryWriter.Write([]byte("pwned"))
163 require.NoError(t, err)
164 require.NoError(t, zw.Close())
165
166 err = theAPI.RestoreBackup(unlockPassword, &buf)
167 require.ErrorContains(t, err, "refusing to extract zip entry outside restore directory")
168
169 _, statErr := os.Stat(escapeTarget)
170 require.True(t, os.IsNotExist(statErr), "traversal entry must not be written outside the restore directory")
171
172 // The failed restore must not leave a restore directory (which would be
173 // applied on the next startup) or any staging leftovers.
174 _, statErr = os.Stat(filepath.Join(workDir, "restore"))
175 require.True(t, os.IsNotExist(statErr), "failed restore must not leave a restore directory")
176
177 entries, err := os.ReadDir(workDir)
178 require.NoError(t, err)
179 for _, entry := range entries {
180 require.False(t, strings.HasPrefix(entry.Name(), "albyhub-restore-"), "failed restore must not leave a staging directory")
181 }
182 }
183
184 // legacyBackupFixture is a backup file created with the encryption scheme
185 // used by older versions (PBKDF2 key derivation), encrypted with the
186 // password "test-unlock-password". Its archive contains a single "nwc.db"
187 // entry with the contents "legacy backup contents".
188 const legacyBackupFixture = "0102030405060708101112131415161718191a1b1c1d1e1f8eca79631915f679a00cdd95d3f20d8d169eb9aa5d52642ca13b93886c3c7d7ba4b759462bc9dd8deccf638edcc9b5b9fda3d23dcd904cf6e99bc57ac59c4df6be5aa676542b7cbc9998029420c0ae5a6986c735150ababde5b382560acaebd5894aa4420924f1ced63fde570adc60c43b32e9e14a0ef60c379da5cac1be0000845992ea072ead036e336c7b859e8d018c4ef61667e3f520fe01"
189
190 // TestDecryptingReaderLegacyBackup verifies that backup files created by
191 // older versions can still be decrypted.
192 func TestDecryptingReaderLegacyBackup(t *testing.T) {
193 encrypted, err := hex.DecodeString(legacyBackupFixture)
194 require.NoError(t, err)
195
196 cr, err := decryptingReader(bytes.NewReader(encrypted), "test-unlock-password")
197 require.NoError(t, err)
198
199 decrypted, err := io.ReadAll(cr)
200 require.NoError(t, err)
201
202 zr, err := zip.NewReader(bytes.NewReader(decrypted), int64(len(decrypted)))
203 require.NoError(t, err)
204
205 dbFile, err := zr.Open("nwc.db")
206 require.NoError(t, err)
207 dbContents, err := io.ReadAll(dbFile)
208 require.NoError(t, err)
209 require.NoError(t, dbFile.Close())
210 require.Equal(t, "legacy backup contents", string(dbContents))
211 }
212
213 // TestDecryptingReaderFragmentedReader verifies that a backup file is
214 // decrypted correctly even when the reader delivers one byte at a time,
215 // which would truncate the header if it were not read in full.
216 func TestDecryptingReaderFragmentedReader(t *testing.T) {
217 var buf bytes.Buffer
218 cw, err := encryptingWriter(&buf, "test-unlock-password")
219 require.NoError(t, err)
220
221 zw := zip.NewWriter(cw)
222 entryWriter, err := zw.Create("nwc.db")
223 require.NoError(t, err)
224 _, err = entryWriter.Write([]byte("backup contents"))
225 require.NoError(t, err)
226 require.NoError(t, zw.Close())
227
228 cr, err := decryptingReader(iotest.OneByteReader(bytes.NewReader(buf.Bytes())), "test-unlock-password")
229 require.NoError(t, err)
230
231 decrypted, err := io.ReadAll(cr)
232 require.NoError(t, err)
233
234 zr, err := zip.NewReader(bytes.NewReader(decrypted), int64(len(decrypted)))
235 require.NoError(t, err)
236
237 dbFile, err := zr.Open("nwc.db")
238 require.NoError(t, err)
239 dbContents, err := io.ReadAll(dbFile)
240 require.NoError(t, err)
241 require.NoError(t, dbFile.Close())
242 require.Equal(t, "backup contents", string(dbContents))
243 }
244
245 // TestDecryptingReaderWrongPassword verifies that decryption fails upfront
246 // when the password does not match the backup file.
247 func TestDecryptingReaderWrongPassword(t *testing.T) {
248 var buf bytes.Buffer
249 cw, err := encryptingWriter(&buf, "test-unlock-password")
250 require.NoError(t, err)
251
252 zw := zip.NewWriter(cw)
253 entryWriter, err := zw.Create("nwc.db")
254 require.NoError(t, err)
255 _, err = entryWriter.Write([]byte("backup contents"))
256 require.NoError(t, err)
257 require.NoError(t, zw.Close())
258
259 _, err = decryptingReader(bytes.NewReader(buf.Bytes()), "wrong-password")
260 require.Error(t, err)
261 }
262