package db import ( "context" "database/sql" "encoding/json" "errors" "fmt" "log/slog" "net/mail" "time" "postern/internal/model" ) func (db *DB) GetAllMailboxMessages(ctx context.Context, userID, mailboxID int) ([]model.Message, error) { query := ` SELECT mm.uid, mm.modseq, m.size, m.blob_address, mm.internal_date, m.date, m.subject, m.message_id, m.in_reply_to, m.from_addr, m.sender, m.reply_to, m.to_addr, m.cc, m.bcc, coalesce(json_group_array(mf.flag), '[]') AS flags FROM ( SELECT uid, message_id, modseq, internal_date FROM mailbox_messages JOIN mailboxes ON mailboxes.id = mailbox_messages.mailbox_id WHERE mailbox_messages.mailbox_id = ? AND mailboxes.user_id = ? AND mailbox_messages.expunged_modseq IS NULL ) numbered JOIN mailbox_messages mm ON mm.mailbox_id = ? AND mm.uid = numbered.uid JOIN messages m ON m.id = numbered.message_id LEFT JOIN message_flags mf ON mf.mailbox_id = ? AND mf.uid = numbered.uid GROUP BY mm.uid ORDER BY mm.uid ASC ` rows, err := db.read.QueryContext(ctx, query, mailboxID, userID, mailboxID, mailboxID) if err != nil { return nil, err } defer func(rows *sql.Rows) { err := rows.Close() if err != nil { slog.Error("closing rows", "error", err) } }(rows) out := make([]model.Message, 0) for rows.Next() { var m model.Message var flagsJSON []byte err := rows.Scan( &m.UID, &m.ModSeq, &m.RFC822Size, &m.BlobHash, &m.InternalDate, &m.EnvelopeDate, &m.EnvelopeSubject, &m.EnvelopeMessageID, &m.EnvelopeInReplyTo, &m.EnvelopeFrom, &m.EnvelopeSender, &m.EnvelopeReplyTo, &m.EnvelopeTo, &m.EnvelopeCc, &m.EnvelopeBcc, &flagsJSON, ) if err != nil { return nil, fmt.Errorf("scanning message: %w", err) } m.ServerSeq = func(ctx context.Context) (uint32, error) { serverSeq, err := db.UIDToServerSeq(ctx, mailboxID, m.UID) if err != nil { return 0, err } return serverSeq, nil } if len(flagsJSON) > 0 { if err := json.Unmarshal(flagsJSON, &m.Flags); err != nil { return nil, fmt.Errorf("unmarshaling flags: %w", err) } } else { m.Flags = []string{} } out = append(out, m) } return out, rows.Err() } func (db *DB) GetMailboxMessagesByUID(ctx context.Context, userID, mailboxID int, uids []uint32) ([]model.Message, error) { allMessages, err := db.GetAllMailboxMessages(ctx, userID, mailboxID) if err != nil { return nil, err } uidSet := make(map[uint32]bool, len(uids)) for _, uid := range uids { uidSet[uid] = true } out := make([]model.Message, 0, len(uids)) for _, m := range allMessages { if uidSet[m.UID] { out = append(out, m) } } return out, nil } func (db *DB) AppendMessage(ctx context.Context, mailboxID, userID int, blobHash string, size int64, internalDate string, flags []string, parsedMsg *mail.Message) (uint32, error) { tx, err := db.write.BeginTx(ctx, nil) if err != nil { return 0, err } defer txRollback(tx) // 1. Allocate a UID and bump modseq, verifying mailbox ownership in one shot. // RETURNING reflects post-update values, so uidnext-1 is the UID to assign. var assignedUID uint32 var newModSeq int64 err = tx.QueryRowContext(ctx, ` UPDATE mailboxes SET uidnext = uidnext + 1, highest_modseq = highest_modseq + 1 WHERE id = ? AND user_id = ? RETURNING uidnext - 1, highest_modseq `, mailboxID, userID).Scan(&assignedUID, &newModSeq) if errors.Is(err, sql.ErrNoRows) { return 0, fmt.Errorf("mailbox %d not found for user %d", mailboxID, userID) } if err != nil { return 0, fmt.Errorf("allocating uid: %w", err) } // 2. Insert the message, or reuse the existing row for identical content. // The no-op DO UPDATE forces RETURNING to yield the existing id on conflict. var internalMessageID int64 var subject, messageID, inReplyTo, fromAddr, sender, replyTo, toAddr, cc, bcc sql.NullString var datestring sql.NullString if parsedMsg != nil && parsedMsg.Header != nil { date, err := parsedMsg.Header.Date() if err == nil && !date.IsZero() { datestring = sql.NullString{ Valid: true, String: date.UTC().Format(time.RFC3339), } } for headerKey, headerValue := range parsedMsg.Header { switch headerKey { case "Subject": subject = sql.NullString{ Valid: true, String: headerValue[0], } case "Message-Id": messageID = sql.NullString{ Valid: true, String: headerValue[0], } case "In-Reply-To": inReplyTo = sql.NullString{ Valid: true, String: headerValue[0], } case "From": fromAddr = sql.NullString{ Valid: true, String: headerValue[0], } case "Sender": sender = sql.NullString{ Valid: true, String: headerValue[0], } case "Reply-To": replyTo = sql.NullString{ Valid: true, String: headerValue[0], } case "To": toAddr = sql.NullString{ Valid: true, String: headerValue[0], } case "Cc": cc = sql.NullString{ Valid: true, String: headerValue[0], } case "Bcc": bcc = sql.NullString{ Valid: true, String: headerValue[0], } } } } err = tx.QueryRowContext(ctx, ` INSERT INTO messages (blob_address, size, subject, message_id, in_reply_to, date, from_addr, sender, reply_to, to_addr, cc, bcc) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(blob_address) DO UPDATE SET blob_address = blob_address RETURNING id `, blobHash, size, subject, messageID, inReplyTo, datestring, fromAddr, sender, replyTo, toAddr, cc, bcc).Scan(&internalMessageID) if err != nil { return 0, fmt.Errorf("inserting message: %w", err) } // 3. Link the message into the mailbox at the allocated UID. _, err = tx.ExecContext(ctx, ` INSERT INTO mailbox_messages (mailbox_id, uid, message_id, modseq, internal_date) VALUES (?, ?, ?, ?, ?) `, mailboxID, assignedUID, internalMessageID, newModSeq, internalDate) if err != nil { return 0, fmt.Errorf("linking message: %w", err) } // 4. Store flags for this mailbox/uid. for _, flag := range flags { _, err = tx.ExecContext(ctx, ` INSERT INTO message_flags (mailbox_id, uid, flag) VALUES (?, ?, ?) ON CONFLICT DO NOTHING `, mailboxID, assignedUID, flag) if err != nil { return 0, fmt.Errorf("inserting flag %q: %w", flag, err) } } if err := tx.Commit(); err != nil { return 0, fmt.Errorf("committing append: %w", err) } return assignedUID, nil } // SetMessageFlags sets flags for a message and returns those flags. // Returning flags is for consistency with addMessageFlags and deleteMessageFlags. func (db *DB) SetMessageFlags(ctx context.Context, mailboxID int, messageID uint32, flags []string) ([]string, error) { tx, err := db.write.BeginTx(ctx, nil) if err != nil { return nil, err } defer txRollback(tx) _, err = tx.ExecContext(ctx, `DELETE FROM message_flags WHERE mailbox_id = ? and uid = ?`, mailboxID, messageID) if err != nil { return nil, err } for _, flag := range flags { _, err = tx.ExecContext(ctx, `INSERT INTO message_flags (mailbox_id, uid, flag) VALUES (?, ?, ?)`, mailboxID, messageID, flag) if err != nil { return nil, err } } return flags, tx.Commit() } func (db *DB) AddMessageFlags(ctx context.Context, mailboxID int, messageID uint32, flags []string) ([]string, error) { tx, err := db.write.BeginTx(ctx, nil) if err != nil { return nil, err } defer txRollback(tx) for _, flag := range flags { _, err = tx.ExecContext(ctx, `INSERT INTO message_flags (mailbox_id, uid, flag) VALUES (?, ?, ?) ON CONFLICT DO NOTHING`, mailboxID, messageID, flag) if err != nil { return nil, err } } flagSet, err := getMessageFlags(ctx, tx, mailboxID, messageID) if err != nil { return nil, err } return flagSet, tx.Commit() } func (db *DB) DeleteMessageFlags(ctx context.Context, mailboxID int, messageID uint32, flags []string) ([]string, error) { tx, err := db.write.BeginTx(ctx, nil) if err != nil { return nil, err } for _, flag := range flags { _, err = tx.ExecContext(ctx, `DELETE FROM message_flags WHERE mailbox_id = ? AND uid = ? AND flag = ?`, mailboxID, messageID, flag) if err != nil { return nil, errors.Join(err, tx.Rollback()) } } flagSet, err := getMessageFlags(ctx, tx, mailboxID, messageID) if err != nil { return nil, errors.Join(err, tx.Rollback()) } return flagSet, tx.Commit() } // getMessageFlags is a helper for addMessageFlags and deleteMessageFlags func getMessageFlags(ctx context.Context, tx *sql.Tx, mailboxID int, messageID uint32) ([]string, error) { var flags []string rows, err := tx.QueryContext(ctx, `SELECT flag FROM message_flags WHERE mailbox_id = ? AND uid = ?`, mailboxID, messageID) if err != nil { return nil, err } defer func(rows *sql.Rows) { _ = rows.Close() }(rows) for rows.Next() { var flag string if err := rows.Scan(&flag); err != nil { return nil, err } flags = append(flags, flag) } return flags, rows.Err() } func (db *DB) CopyMessagesToMailbox(ctx context.Context, fromMailboxID, toMailboxID int, uids []uint32) ([]uint32, error) { tx, err := db.write.BeginTx(ctx, nil) if err != nil { return nil, err } defer txRollback(tx) destUIDs := make([]uint32, len(uids)) for i, uid := range uids { nextUID, err := getAndIncreaseNextUID(ctx, tx, toMailboxID) if err != nil { return nil, err } destUIDs[i] = nextUID _, err = tx.ExecContext( ctx, ` INSERT INTO mailbox_messages (mailbox_id, uid, message_id, modseq, internal_date, expunged_modseq) SELECT ?, ?, message_id, modseq, internal_date, expunged_modseq FROM mailbox_messages WHERE mailbox_id = ? AND uid = ?`, toMailboxID, nextUID, fromMailboxID, uid) if err != nil { return nil, err } _, err = tx.ExecContext( ctx, ` INSERT INTO message_flags (mailbox_id, uid, flag) SELECT ?, ?, flag FROM message_flags WHERE mailbox_id = ? AND uid = ?`, toMailboxID, nextUID, fromMailboxID, uid) if err != nil { return nil, err } } return destUIDs, tx.Commit() } func (db *DB) MoveMessagesToMailbox(ctx context.Context, fromMailboxID, toMailboxID int, uids []uint32) ([]uint32, error) { tx, err := db.write.BeginTx(ctx, nil) if err != nil { return nil, err } defer txRollback(tx) destUIDs := make([]uint32, len(uids)) for i, uid := range uids { nextUID, err := getAndIncreaseNextUID(ctx, tx, toMailboxID) if err != nil { return nil, err } destUIDs[i] = nextUID _, err = tx.ExecContext( ctx, ` UPDATE mailbox_messages SET mailbox_id = ?, uid = ? WHERE mailbox_id = ? AND uid = ?`, toMailboxID, nextUID, fromMailboxID, uid) if err != nil { return nil, err } } return destUIDs, tx.Commit() } func (db *DB) SetGatekeeperDecision(ctx context.Context, userID int, fromAddress, destination string) error { _, err := db.write.ExecContext(ctx, ` INSERT INTO gatekeepers (user_id, from_address, destination, created_at) VALUES (?, ?, ?, ?) ON CONFLICT DO UPDATE SET destination = ?`, userID, fromAddress, destination, time.Now().Unix(), destination) if err != nil { return err } slog.Info("set gatekeeper decision", "user_id", userID, "destination", destination, "from_address", fromAddress) return nil } // getAndIncreaseNextUID returns the next UID to be used. This function increases the mailbox' UID on each call. // It should only be used as part of another transaction. // There are certainly more efficient ways, but this is by far the most readable. func getAndIncreaseNextUID(ctx context.Context, db *sql.Tx, mailboxID int) (uint32, error) { res := db.QueryRowContext( ctx, `UPDATE mailboxes SET uidnext = uidnext + 1 WHERE id = ? RETURNING uidnext`, mailboxID) if res.Err() != nil { return 0, res.Err() } var uidnext uint32 err := res.Scan(&uidnext) if err != nil { return 0, err } return uidnext - 1, nil // The query returns the upcoming UID. We need to return the to-be-used UID } func (db *DB) GetAllFlaggedDeletedMessages(ctx context.Context, mailboxID int) ([]uint32, error) { rows, err := db.read.QueryContext(ctx, `SELECT uid FROM message_flags WHERE flag = '\Deleted' AND mailbox_id = ?`, mailboxID) if err != nil { return nil, err } defer func(rows *sql.Rows) { _ = rows.Close() }(rows) uids := make([]uint32, 0) for rows.Next() { var uid uint32 err := rows.Scan(&uid) if err != nil { return nil, err } uids = append(uids, uid) } return uids, rows.Err() } func (db *DB) DeleteMessageFromMailbox(ctx context.Context, mailboxID int, uid uint32) error { // message_flags are automatically deleted via CASCADE _, err := db.write.ExecContext(ctx, `DELETE FROM mailbox_messages WHERE mailbox_id = ? AND uid = ?`, mailboxID, uid) return err } // UIDExists reports whether a user has a message with the given uid. // RFC 9051: A non-existent unique identifier is ignored without any error message generated. Thus, it is possible for a // UID FETCH command to return an OK without any data or a UID COPY, UID MOVE, or UID STORE to return an OK without // performing any operations. func (db *DB) UIDExists(ctx context.Context, userID int, uid uint32) (bool, error) { query := ` SELECT EXISTS ( SELECT 1 FROM mailbox_messages mm JOIN mailboxes m ON mm.mailbox_id = m.id WHERE m.user_id = ? AND mm.uid = ? )` var exists bool err := db.read.QueryRowContext(ctx, query, userID, uid).Scan(&exists) if err != nil { return false, fmt.Errorf("error checking uid existence: %w", err) } return exists, nil }