refactor: deduplicate pin handler and DB scan logic

This commit is contained in:
jevb
2026-03-21 21:42:14 +01:00
parent fe76034b7e
commit ce2b2369bf
2 changed files with 30 additions and 149 deletions
+14 -88
View File
@@ -23,8 +23,8 @@ func MountChannelRoutes(r chi.Router, database *db.DB) {
r.Get("/", handleListChannels(database))
r.Get("/{id}/messages", handleGetMessages(database))
r.Get("/{id}/pins", handleGetPins(database))
r.Post("/{id}/pins/{messageId}", handlePinMessage(database))
r.Delete("/{id}/pins/{messageId}", handleUnpinMessage(database))
r.Post("/{id}/pins/{messageId}", handleSetPinned(database, true))
r.Delete("/{id}/pins/{messageId}", handleSetPinned(database, false))
})
r.With(AuthMiddleware(database)).Get("/api/v1/search", handleSearch(database))
}
@@ -334,87 +334,20 @@ func handleGetPins(database *db.DB) http.HandlerFunc {
}
}
// handlePinMessage pins a message in a channel.
func handlePinMessage(database *db.DB) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
channelID, ok := parseIDParam(w, r, "id")
if !ok {
return
}
raw := chi.URLParam(r, "messageId")
messageID, err := strconv.ParseInt(raw, 10, 64)
if err != nil || messageID <= 0 {
writeJSON(w, http.StatusBadRequest, errorResponse{
Error: "BAD_REQUEST",
Message: "messageId must be a positive integer",
})
return
}
// Permission check: user must have MANAGE_MESSAGES on this channel.
role, _ := r.Context().Value(RoleKey).(*db.Role)
if !hasChannelPermREST(database, role, channelID, permissions.ManageMessages) {
writeJSON(w, http.StatusForbidden, errorResponse{
Error: "FORBIDDEN",
Message: "no permission to manage messages in this channel",
})
return
}
// Verify message exists and belongs to this channel.
msg, err := database.GetMessage(messageID)
if err != nil {
slog.Error("handlePinMessage GetMessage", "err", err, "message_id", messageID)
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL",
Message: "failed to look up message",
})
return
}
if msg == nil {
writeJSON(w, http.StatusNotFound, errorResponse{
Error: "NOT_FOUND",
Message: "message not found",
})
return
}
if msg.ChannelID != channelID {
writeJSON(w, http.StatusNotFound, errorResponse{
Error: "NOT_FOUND",
Message: "message not found in this channel",
})
return
}
if err := database.SetMessagePinned(messageID, true); err != nil {
slog.Error("handlePinMessage SetMessagePinned", "err", err, "message_id", messageID)
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL",
Message: "failed to pin message",
})
return
}
w.WriteHeader(http.StatusNoContent)
// handleSetPinned pins or unpins a message in a channel.
func handleSetPinned(database *db.DB, pinned bool) http.HandlerFunc {
action := "pin"
if !pinned {
action = "unpin"
}
}
// handleUnpinMessage unpins a message in a channel.
func handleUnpinMessage(database *db.DB) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
channelID, ok := parseIDParam(w, r, "id")
if !ok {
return
}
raw := chi.URLParam(r, "messageId")
messageID, err := strconv.ParseInt(raw, 10, 64)
if err != nil || messageID <= 0 {
writeJSON(w, http.StatusBadRequest, errorResponse{
Error: "BAD_REQUEST",
Message: "messageId must be a positive integer",
})
messageID, ok := parseIDParam(w, r, "messageId")
if !ok {
return
}
@@ -431,33 +364,26 @@ func handleUnpinMessage(database *db.DB) http.HandlerFunc {
// Verify message exists and belongs to this channel.
msg, err := database.GetMessage(messageID)
if err != nil {
slog.Error("handleUnpinMessage GetMessage", "err", err, "message_id", messageID)
slog.Error("handleSetPinned GetMessage", "err", err, "action", action, "message_id", messageID)
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL",
Message: "failed to look up message",
})
return
}
if msg == nil {
if msg == nil || msg.ChannelID != channelID {
writeJSON(w, http.StatusNotFound, errorResponse{
Error: "NOT_FOUND",
Message: "message not found",
})
return
}
if msg.ChannelID != channelID {
writeJSON(w, http.StatusNotFound, errorResponse{
Error: "NOT_FOUND",
Message: "message not found in this channel",
})
return
}
if err := database.SetMessagePinned(messageID, false); err != nil {
slog.Error("handleUnpinMessage SetMessagePinned", "err", err, "message_id", messageID)
if err := database.SetMessagePinned(messageID, pinned); err != nil {
slog.Error("handleSetPinned SetMessagePinned", "err", err, "action", action, "message_id", messageID)
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL",
Message: "failed to unpin message",
Message: "failed to " + action + " message",
})
return
}
+16 -61
View File
@@ -277,54 +277,7 @@ func (d *DB) GetMessagesForAPI(channelID, before int64, limit int, requestingUse
}
defer rows.Close() //nolint:errcheck
var msgs []MessageAPIResponse
var msgIDs []int64
for rows.Next() {
var m MessageAPIResponse
var deleted, pinned int
if scanErr := rows.Scan(
&m.ID, &m.ChannelID, &m.User.ID, &m.User.Username, &m.User.Avatar,
&m.Content, &m.ReplyTo, &m.EditedAt, &deleted, &pinned, &m.Timestamp,
); scanErr != nil {
return nil, fmt.Errorf("GetMessagesForAPI scan: %w", scanErr)
}
m.Deleted = deleted != 0
m.Pinned = pinned != 0
m.Attachments = []AttachmentInfo{}
m.Reactions = []ReactionInfo{}
msgs = append(msgs, m)
msgIDs = append(msgIDs, m.ID)
}
if rows.Err() != nil {
return nil, fmt.Errorf("GetMessagesForAPI rows: %w", rows.Err())
}
if msgs == nil {
return []MessageAPIResponse{}, nil
}
// Batch-fetch reactions for all message IDs.
reactMap, err := d.getReactionsBatch(msgIDs, requestingUserID)
if err != nil {
return nil, fmt.Errorf("GetMessagesForAPI reactions: %w", err)
}
for i := range msgs {
if r, ok := reactMap[msgs[i].ID]; ok {
msgs[i].Reactions = r
}
}
// Batch-fetch attachments for all message IDs.
attMap, err := d.GetAttachmentsByMessageIDs(msgIDs)
if err != nil {
return nil, fmt.Errorf("GetMessagesForAPI attachments: %w", err)
}
for i := range msgs {
if a, ok := attMap[msgs[i].ID]; ok {
msgs[i].Attachments = a
}
}
return msgs, nil
return d.scanAndEnrichMessages(rows, requestingUserID)
}
// getReactionsBatch returns aggregated reactions for multiple messages.
@@ -456,6 +409,12 @@ func (d *DB) GetPinnedMessages(channelID int64, requestingUserID int64) ([]Messa
}
defer rows.Close() //nolint:errcheck
return d.scanAndEnrichMessages(rows, requestingUserID)
}
// scanAndEnrichMessages scans rows into MessageAPIResponse slice and
// batch-fetches reactions and attachments. Caller must defer rows.Close().
func (d *DB) scanAndEnrichMessages(rows *sql.Rows, requestingUserID int64) ([]MessageAPIResponse, error) {
var msgs []MessageAPIResponse
var msgIDs []int64
for rows.Next() {
@@ -465,7 +424,7 @@ func (d *DB) GetPinnedMessages(channelID int64, requestingUserID int64) ([]Messa
&m.ID, &m.ChannelID, &m.User.ID, &m.User.Username, &m.User.Avatar,
&m.Content, &m.ReplyTo, &m.EditedAt, &deleted, &pinned, &m.Timestamp,
); scanErr != nil {
return nil, fmt.Errorf("GetPinnedMessages scan: %w", scanErr)
return nil, fmt.Errorf("scanAndEnrichMessages scan: %w", scanErr)
}
m.Deleted = deleted != 0
m.Pinned = pinned != 0
@@ -475,7 +434,7 @@ func (d *DB) GetPinnedMessages(channelID int64, requestingUserID int64) ([]Messa
msgIDs = append(msgIDs, m.ID)
}
if rows.Err() != nil {
return nil, fmt.Errorf("GetPinnedMessages rows: %w", rows.Err())
return nil, fmt.Errorf("scanAndEnrichMessages rows: %w", rows.Err())
}
if msgs == nil {
return []MessageAPIResponse{}, nil
@@ -484,7 +443,7 @@ func (d *DB) GetPinnedMessages(channelID int64, requestingUserID int64) ([]Messa
// Batch-fetch reactions for all message IDs.
reactMap, err := d.getReactionsBatch(msgIDs, requestingUserID)
if err != nil {
return nil, fmt.Errorf("GetPinnedMessages reactions: %w", err)
return nil, fmt.Errorf("scanAndEnrichMessages reactions: %w", err)
}
for i := range msgs {
if r, ok := reactMap[msgs[i].ID]; ok {
@@ -495,7 +454,7 @@ func (d *DB) GetPinnedMessages(channelID int64, requestingUserID int64) ([]Messa
// Batch-fetch attachments for all message IDs.
attMap, err := d.GetAttachmentsByMessageIDs(msgIDs)
if err != nil {
return nil, fmt.Errorf("GetPinnedMessages attachments: %w", err)
return nil, fmt.Errorf("scanAndEnrichMessages attachments: %w", err)
}
for i := range msgs {
if a, ok := attMap[msgs[i].ID]; ok {
@@ -509,22 +468,18 @@ func (d *DB) GetPinnedMessages(channelID int64, requestingUserID int64) ([]Messa
// SetMessagePinned updates the pinned column on a message.
// Returns ErrNotFound if the message does not exist.
func (d *DB) SetMessagePinned(id int64, pinned bool) error {
msg, err := d.GetMessage(id)
if err != nil {
return err
}
if msg == nil {
return fmt.Errorf("SetMessagePinned: message %d: %w", id, ErrNotFound)
}
val := 0
if pinned {
val = 1
}
_, err = d.sqlDB.Exec(`UPDATE messages SET pinned = ? WHERE id = ?`, val, id)
res, err := d.sqlDB.Exec(`UPDATE messages SET pinned = ? WHERE id = ? AND deleted = 0`, val, id)
if err != nil {
return fmt.Errorf("SetMessagePinned: %w", err)
}
n, _ := res.RowsAffected()
if n == 0 {
return fmt.Errorf("SetMessagePinned: message %d: %w", id, ErrNotFound)
}
return nil
}