internal/db/message.go (view raw)
1package db
2
3import (
4 "context"
5 "database/sql"
6 "encoding/json"
7 "errors"
8 "fmt"
9 "log/slog"
10 "net/mail"
11 "time"
12
13 "postern/internal/model"
14)
15
16func (db *DB) GetAllMailboxMessages(ctx context.Context, userID, mailboxID int) ([]model.Message, error) {
17 query := `
18SELECT
19 mm.uid,
20 mm.modseq,
21 m.size,
22 m.blob_address,
23 mm.internal_date,
24 m.date,
25 m.subject,
26 m.message_id,
27 m.in_reply_to,
28 m.from_addr,
29 m.sender,
30 m.reply_to,
31 m.to_addr,
32 m.cc,
33 m.bcc,
34 coalesce(json_group_array(mf.flag), '[]') AS flags
35FROM (
36 SELECT
37 uid,
38 message_id,
39 modseq,
40 internal_date
41 FROM mailbox_messages
42 JOIN mailboxes ON mailboxes.id = mailbox_messages.mailbox_id
43 WHERE mailbox_messages.mailbox_id = ?
44 AND mailboxes.user_id = ?
45 AND mailbox_messages.expunged_modseq IS NULL
46) numbered
47JOIN mailbox_messages mm ON mm.mailbox_id = ? AND mm.uid = numbered.uid
48JOIN messages m ON m.id = numbered.message_id
49LEFT JOIN message_flags mf ON mf.mailbox_id = ? AND mf.uid = numbered.uid
50GROUP BY mm.uid
51ORDER BY mm.uid ASC
52`
53
54 rows, err := db.read.QueryContext(ctx, query, mailboxID, userID, mailboxID, mailboxID)
55 if err != nil {
56 return nil, err
57 }
58 defer func(rows *sql.Rows) {
59 err := rows.Close()
60 if err != nil {
61 slog.Error("closing rows", "error", err)
62 }
63 }(rows)
64
65 out := make([]model.Message, 0)
66 for rows.Next() {
67 var m model.Message
68 var flagsJSON []byte
69 err := rows.Scan(
70 &m.UID,
71 &m.ModSeq,
72 &m.RFC822Size,
73 &m.BlobHash,
74 &m.InternalDate,
75 &m.EnvelopeDate,
76 &m.EnvelopeSubject,
77 &m.EnvelopeMessageID,
78 &m.EnvelopeInReplyTo,
79 &m.EnvelopeFrom,
80 &m.EnvelopeSender,
81 &m.EnvelopeReplyTo,
82 &m.EnvelopeTo,
83 &m.EnvelopeCc,
84 &m.EnvelopeBcc,
85 &flagsJSON,
86 )
87 if err != nil {
88 return nil, fmt.Errorf("scanning message: %w", err)
89 }
90 m.ServerSeq = func(ctx context.Context) (uint32, error) {
91 serverSeq, err := db.UIDToServerSeq(ctx, mailboxID, m.UID)
92 if err != nil {
93 return 0, err
94 }
95 return serverSeq, nil
96 }
97 if len(flagsJSON) > 0 {
98 if err := json.Unmarshal(flagsJSON, &m.Flags); err != nil {
99 return nil, fmt.Errorf("unmarshaling flags: %w", err)
100 }
101 } else {
102 m.Flags = []string{}
103 }
104 out = append(out, m)
105 }
106
107 return out, rows.Err()
108}
109
110func (db *DB) GetMailboxMessagesByUID(ctx context.Context, userID, mailboxID int, uids []uint32) ([]model.Message, error) {
111 allMessages, err := db.GetAllMailboxMessages(ctx, userID, mailboxID)
112 if err != nil {
113 return nil, err
114 }
115
116 uidSet := make(map[uint32]bool, len(uids))
117 for _, uid := range uids {
118 uidSet[uid] = true
119 }
120
121 out := make([]model.Message, 0, len(uids))
122 for _, m := range allMessages {
123 if uidSet[m.UID] {
124 out = append(out, m)
125 }
126 }
127
128 return out, nil
129}
130
131func (db *DB) AppendMessage(ctx context.Context, mailboxID, userID int,
132 blobHash string, size int64, internalDate string, flags []string, parsedMsg *mail.Message) (uint32, error) {
133
134 tx, err := db.write.BeginTx(ctx, nil)
135 if err != nil {
136 return 0, err
137 }
138 defer txRollback(tx)
139
140 // 1. Allocate a UID and bump modseq, verifying mailbox ownership in one shot.
141 // RETURNING reflects post-update values, so uidnext-1 is the UID to assign.
142 var assignedUID uint32
143 var newModSeq int64
144 err = tx.QueryRowContext(ctx, `
145 UPDATE mailboxes
146 SET uidnext = uidnext + 1,
147 highest_modseq = highest_modseq + 1
148 WHERE id = ? AND user_id = ?
149 RETURNING uidnext - 1, highest_modseq
150 `, mailboxID, userID).Scan(&assignedUID, &newModSeq)
151 if errors.Is(err, sql.ErrNoRows) {
152 return 0, fmt.Errorf("mailbox %d not found for user %d", mailboxID, userID)
153 }
154 if err != nil {
155 return 0, fmt.Errorf("allocating uid: %w", err)
156 }
157
158 // 2. Insert the message, or reuse the existing row for identical content.
159 // The no-op DO UPDATE forces RETURNING to yield the existing id on conflict.
160 var internalMessageID int64
161 var subject, messageID, inReplyTo, fromAddr, sender, replyTo, toAddr, cc, bcc sql.NullString
162 var datestring sql.NullString
163 if parsedMsg != nil && parsedMsg.Header != nil {
164 date, err := parsedMsg.Header.Date()
165 if err == nil && !date.IsZero() {
166 datestring = sql.NullString{
167 Valid: true,
168 String: date.UTC().Format(time.RFC3339),
169 }
170 }
171
172 for headerKey, headerValue := range parsedMsg.Header {
173 switch headerKey {
174 case "Subject":
175 subject = sql.NullString{
176 Valid: true,
177 String: headerValue[0],
178 }
179 case "Message-Id":
180 messageID = sql.NullString{
181 Valid: true,
182 String: headerValue[0],
183 }
184 case "In-Reply-To":
185 inReplyTo = sql.NullString{
186 Valid: true,
187 String: headerValue[0],
188 }
189 case "From":
190 fromAddr = sql.NullString{
191 Valid: true,
192 String: headerValue[0],
193 }
194 case "Sender":
195 sender = sql.NullString{
196 Valid: true,
197 String: headerValue[0],
198 }
199 case "Reply-To":
200 replyTo = sql.NullString{
201 Valid: true,
202 String: headerValue[0],
203 }
204 case "To":
205 toAddr = sql.NullString{
206 Valid: true,
207 String: headerValue[0],
208 }
209 case "Cc":
210 cc = sql.NullString{
211 Valid: true,
212 String: headerValue[0],
213 }
214 case "Bcc":
215 bcc = sql.NullString{
216 Valid: true,
217 String: headerValue[0],
218 }
219 }
220 }
221 }
222
223 err = tx.QueryRowContext(ctx, `
224 INSERT INTO messages (blob_address, size, subject, message_id, in_reply_to, date, from_addr, sender, reply_to, to_addr, cc, bcc)
225 VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
226 ON CONFLICT(blob_address) DO UPDATE SET blob_address = blob_address
227 RETURNING id
228 `, blobHash, size, subject, messageID, inReplyTo, datestring, fromAddr, sender, replyTo, toAddr, cc, bcc).Scan(&internalMessageID)
229 if err != nil {
230 return 0, fmt.Errorf("inserting message: %w", err)
231 }
232
233 // 3. Link the message into the mailbox at the allocated UID.
234 _, err = tx.ExecContext(ctx, `
235 INSERT INTO mailbox_messages (mailbox_id, uid, message_id, modseq, internal_date)
236 VALUES (?, ?, ?, ?, ?)
237 `, mailboxID, assignedUID, internalMessageID, newModSeq, internalDate)
238 if err != nil {
239 return 0, fmt.Errorf("linking message: %w", err)
240 }
241
242 // 4. Store flags for this mailbox/uid.
243 for _, flag := range flags {
244 _, err = tx.ExecContext(ctx, `
245 INSERT INTO message_flags (mailbox_id, uid, flag)
246 VALUES (?, ?, ?)
247 ON CONFLICT DO NOTHING
248 `, mailboxID, assignedUID, flag)
249 if err != nil {
250 return 0, fmt.Errorf("inserting flag %q: %w", flag, err)
251 }
252 }
253
254 if err := tx.Commit(); err != nil {
255 return 0, fmt.Errorf("committing append: %w", err)
256 }
257 return assignedUID, nil
258}
259
260// SetMessageFlags sets flags for a message and returns those flags.
261// Returning flags is for consistency with addMessageFlags and deleteMessageFlags.
262func (db *DB) SetMessageFlags(ctx context.Context, mailboxID int, messageID uint32, flags []string) ([]string, error) {
263 tx, err := db.write.BeginTx(ctx, nil)
264 if err != nil {
265 return nil, err
266 }
267 defer txRollback(tx)
268
269 _, err = tx.ExecContext(ctx, `DELETE FROM message_flags WHERE mailbox_id = ? and uid = ?`, mailboxID, messageID)
270 if err != nil {
271 return nil, err
272 }
273
274 for _, flag := range flags {
275 _, err = tx.ExecContext(ctx, `INSERT INTO message_flags (mailbox_id, uid, flag) VALUES (?, ?, ?)`, mailboxID, messageID, flag)
276 if err != nil {
277 return nil, err
278 }
279 }
280 return flags, tx.Commit()
281}
282
283func (db *DB) AddMessageFlags(ctx context.Context, mailboxID int, messageID uint32, flags []string) ([]string, error) {
284 tx, err := db.write.BeginTx(ctx, nil)
285 if err != nil {
286 return nil, err
287 }
288 defer txRollback(tx)
289
290 for _, flag := range flags {
291 _, err = tx.ExecContext(ctx, `INSERT INTO message_flags (mailbox_id, uid, flag) VALUES (?, ?, ?) ON CONFLICT DO NOTHING`, mailboxID, messageID, flag)
292 if err != nil {
293 return nil, err
294 }
295 }
296
297 flagSet, err := getMessageFlags(ctx, tx, mailboxID, messageID)
298 if err != nil {
299 return nil, err
300 }
301
302 return flagSet, tx.Commit()
303}
304
305func (db *DB) DeleteMessageFlags(ctx context.Context, mailboxID int, messageID uint32, flags []string) ([]string, error) {
306 tx, err := db.write.BeginTx(ctx, nil)
307 if err != nil {
308 return nil, err
309 }
310
311 for _, flag := range flags {
312 _, err = tx.ExecContext(ctx, `DELETE FROM message_flags WHERE mailbox_id = ? AND uid = ? AND flag = ?`, mailboxID, messageID, flag)
313 if err != nil {
314 return nil, errors.Join(err, tx.Rollback())
315 }
316 }
317
318 flagSet, err := getMessageFlags(ctx, tx, mailboxID, messageID)
319 if err != nil {
320 return nil, errors.Join(err, tx.Rollback())
321 }
322
323 return flagSet, tx.Commit()
324}
325
326// getMessageFlags is a helper for addMessageFlags and deleteMessageFlags
327func getMessageFlags(ctx context.Context, tx *sql.Tx, mailboxID int, messageID uint32) ([]string, error) {
328 var flags []string
329 rows, err := tx.QueryContext(ctx, `SELECT flag FROM message_flags WHERE mailbox_id = ? AND uid = ?`, mailboxID, messageID)
330 if err != nil {
331 return nil, err
332 }
333 defer func(rows *sql.Rows) {
334 _ = rows.Close()
335 }(rows)
336
337 for rows.Next() {
338 var flag string
339 if err := rows.Scan(&flag); err != nil {
340 return nil, err
341 }
342 flags = append(flags, flag)
343 }
344
345 return flags, rows.Err()
346}
347
348func (db *DB) CopyMessagesToMailbox(ctx context.Context, fromMailboxID, toMailboxID int, uids []uint32) ([]uint32, error) {
349 tx, err := db.write.BeginTx(ctx, nil)
350 if err != nil {
351 return nil, err
352 }
353 defer txRollback(tx)
354
355 destUIDs := make([]uint32, len(uids))
356 for i, uid := range uids {
357 nextUID, err := getAndIncreaseNextUID(ctx, tx, toMailboxID)
358 if err != nil {
359 return nil, err
360 }
361 destUIDs[i] = nextUID
362
363 _, err = tx.ExecContext(
364 ctx,
365 `
366INSERT INTO mailbox_messages (mailbox_id, uid, message_id, modseq, internal_date, expunged_modseq)
367SELECT ?, ?, message_id, modseq, internal_date, expunged_modseq
368FROM mailbox_messages
369WHERE mailbox_id = ? AND uid = ?`,
370 toMailboxID, nextUID, fromMailboxID, uid)
371 if err != nil {
372 return nil, err
373 }
374
375 _, err = tx.ExecContext(
376 ctx,
377 `
378INSERT INTO message_flags (mailbox_id, uid, flag)
379SELECT ?, ?, flag
380FROM message_flags
381WHERE mailbox_id = ? AND uid = ?`,
382 toMailboxID, nextUID, fromMailboxID, uid)
383 if err != nil {
384 return nil, err
385 }
386 }
387
388 return destUIDs, tx.Commit()
389}
390
391func (db *DB) MoveMessagesToMailbox(ctx context.Context, fromMailboxID, toMailboxID int, uids []uint32) ([]uint32, error) {
392 tx, err := db.write.BeginTx(ctx, nil)
393 if err != nil {
394 return nil, err
395 }
396 defer txRollback(tx)
397
398 destUIDs := make([]uint32, len(uids))
399 for i, uid := range uids {
400 nextUID, err := getAndIncreaseNextUID(ctx, tx, toMailboxID)
401 if err != nil {
402 return nil, err
403 }
404 destUIDs[i] = nextUID
405
406 _, err = tx.ExecContext(
407 ctx,
408 `
409UPDATE mailbox_messages
410SET mailbox_id = ?, uid = ?
411WHERE mailbox_id = ? AND uid = ?`,
412 toMailboxID, nextUID, fromMailboxID, uid)
413 if err != nil {
414 return nil, err
415 }
416 }
417
418 return destUIDs, tx.Commit()
419}
420
421func (db *DB) SetGatekeeperDecision(ctx context.Context, userID int, fromAddress, destination string) error {
422 _, err := db.write.ExecContext(ctx,
423 `
424INSERT INTO gatekeepers (user_id, from_address, destination, created_at)
425VALUES (?, ?, ?, ?) ON CONFLICT DO UPDATE SET destination = ?`,
426 userID, fromAddress, destination, time.Now().Unix(), destination)
427 if err != nil {
428 return err
429 }
430
431 slog.Info("set gatekeeper decision", "user_id", userID, "destination", destination, "from_address", fromAddress)
432 return nil
433}
434
435// getAndIncreaseNextUID returns the next UID to be used. This function increases the mailbox' UID on each call.
436// It should only be used as part of another transaction.
437// There are certainly more efficient ways, but this is by far the most readable.
438func getAndIncreaseNextUID(ctx context.Context, db *sql.Tx, mailboxID int) (uint32, error) {
439 res := db.QueryRowContext(
440 ctx,
441 `UPDATE mailboxes SET uidnext = uidnext + 1 WHERE id = ? RETURNING uidnext`, mailboxID)
442 if res.Err() != nil {
443 return 0, res.Err()
444 }
445
446 var uidnext uint32
447 err := res.Scan(&uidnext)
448 if err != nil {
449 return 0, err
450 }
451 return uidnext - 1, nil // The query returns the upcoming UID. We need to return the to-be-used UID
452}
453
454func (db *DB) GetAllFlaggedDeletedMessages(ctx context.Context, mailboxID int) ([]uint32, error) {
455 rows, err := db.read.QueryContext(ctx, `SELECT uid FROM message_flags WHERE flag = '\Deleted' AND mailbox_id = ?`, mailboxID)
456 if err != nil {
457 return nil, err
458 }
459
460 defer func(rows *sql.Rows) {
461 _ = rows.Close()
462 }(rows)
463
464 uids := make([]uint32, 0)
465 for rows.Next() {
466 var uid uint32
467 err := rows.Scan(&uid)
468 if err != nil {
469 return nil, err
470 }
471 uids = append(uids, uid)
472 }
473 return uids, rows.Err()
474}
475
476func (db *DB) DeleteMessageFromMailbox(ctx context.Context, mailboxID int, uid uint32) error {
477 // message_flags are automatically deleted via CASCADE
478 _, err := db.write.ExecContext(ctx, `DELETE FROM mailbox_messages WHERE mailbox_id = ? AND uid = ?`, mailboxID, uid)
479 return err
480}
481
482// UIDExists reports whether a user has a message with the given uid.
483// RFC 9051: A non-existent unique identifier is ignored without any error message generated. Thus, it is possible for a
484// UID FETCH command to return an OK without any data or a UID COPY, UID MOVE, or UID STORE to return an OK without
485// performing any operations.
486func (db *DB) UIDExists(ctx context.Context, userID int, uid uint32) (bool, error) {
487 query := `
488 SELECT EXISTS (
489 SELECT 1 FROM mailbox_messages mm
490 JOIN mailboxes m ON mm.mailbox_id = m.id
491 WHERE m.user_id = ? AND mm.uid = ?
492 )`
493
494 var exists bool
495 err := db.read.QueryRowContext(ctx, query, userID, uid).Scan(&exists)
496 if err != nil {
497 return false, fmt.Errorf("error checking uid existence: %w", err)
498 }
499 return exists, nil
500}