diff --git a/Server/api/channel_handler.go b/Server/api/channel_handler.go index fed85086..860a86e4 100644 --- a/Server/api/channel_handler.go +++ b/Server/api/channel_handler.go @@ -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 } diff --git a/Server/db/message_queries.go b/Server/db/message_queries.go index d50b4905..2c3252db 100644 --- a/Server/db/message_queries.go +++ b/Server/db/message_queries.go @@ -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 }