all repos — postern @ main

Modern mail management

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}