diff --git a/Server/api/channel_handler.go b/Server/api/channel_handler.go index e8c2e5d4..d0bf7118 100644 --- a/Server/api/channel_handler.go +++ b/Server/api/channel_handler.go @@ -135,8 +135,14 @@ func handleGetMessages(database *db.DB) http.HandlerFunc { limit = v } + // Extract requesting user ID for reaction "me" flag. + var userID int64 + if user, ok := r.Context().Value(UserKey).(*db.User); ok && user != nil { + userID = user.ID + } + // Fetch one extra to determine has_more. - msgs, err := database.GetMessages(channelID, before, limit+1) + msgs, err := database.GetMessagesForAPI(channelID, before, limit+1, userID) if err != nil { writeJSON(w, http.StatusInternalServerError, errorResponse{ Error: "INTERNAL", @@ -152,8 +158,8 @@ func handleGetMessages(database *db.DB) http.HandlerFunc { } type response struct { - Messages []db.MessageWithUser `json:"messages"` - HasMore bool `json:"has_more"` + Messages []db.MessageAPIResponse `json:"messages"` + HasMore bool `json:"has_more"` } writeJSON(w, http.StatusOK, response{Messages: msgs, HasMore: hasMore}) } diff --git a/Server/db/message_queries.go b/Server/db/message_queries.go index 39e606a3..104faf35 100644 --- a/Server/db/message_queries.go +++ b/Server/db/message_queries.go @@ -193,7 +193,7 @@ func (d *DB) SearchMessages(query string, channelID *int64, limit int) ([]Messag if channelID != nil { rows, err = d.sqlDB.Query( - `SELECT m.id, m.channel_id, c.name, u.username, m.content, m.timestamp + `SELECT m.id, m.channel_id, c.name, u.id, u.username, u.avatar, m.content, m.timestamp FROM messages_fts f JOIN messages m ON f.rowid = m.id JOIN channels c ON m.channel_id = c.id @@ -204,7 +204,7 @@ func (d *DB) SearchMessages(query string, channelID *int64, limit int) ([]Messag ) } else { rows, err = d.sqlDB.Query( - `SELECT m.id, m.channel_id, c.name, u.username, m.content, m.timestamp + `SELECT m.id, m.channel_id, c.name, u.id, u.username, u.avatar, m.content, m.timestamp FROM messages_fts f JOIN messages m ON f.rowid = m.id JOIN channels c ON m.channel_id = c.id @@ -222,7 +222,9 @@ func (d *DB) SearchMessages(query string, channelID *int64, limit int) ([]Messag var results []MessageSearchResult for rows.Next() { var r MessageSearchResult - if scanErr := rows.Scan(&r.MessageID, &r.ChannelID, &r.ChannelName, &r.Username, &r.Content, &r.Timestamp); scanErr != nil { + if scanErr := rows.Scan(&r.MessageID, &r.ChannelID, &r.ChannelName, + &r.User.ID, &r.User.Username, &r.User.Avatar, + &r.Content, &r.Timestamp); scanErr != nil { return nil, fmt.Errorf("SearchMessages scan: %w", scanErr) } results = append(results, r) @@ -236,6 +238,124 @@ func (d *DB) SearchMessages(query string, channelID *int64, limit int) ([]Messag return results, nil } +// GetMessagesForAPI returns messages in the API.md response shape, including +// user object, reactions (with me flag), and attachments. +func (d *DB) GetMessagesForAPI(channelID, before int64, limit int, requestingUserID int64) ([]MessageAPIResponse, error) { + var ( + rows *sql.Rows + err error + ) + if before > 0 { + rows, err = d.sqlDB.Query( + `SELECT m.id, m.channel_id, m.user_id, u.username, u.avatar, + m.content, m.reply_to, m.edited_at, m.deleted, m.pinned, m.timestamp + FROM messages m JOIN users u ON m.user_id = u.id + WHERE m.channel_id = ? AND m.id < ? AND m.deleted = 0 + ORDER BY m.id DESC LIMIT ?`, + channelID, before, limit, + ) + } else { + rows, err = d.sqlDB.Query( + `SELECT m.id, m.channel_id, m.user_id, u.username, u.avatar, + m.content, m.reply_to, m.edited_at, m.deleted, m.pinned, m.timestamp + FROM messages m JOIN users u ON m.user_id = u.id + WHERE m.channel_id = ? AND m.deleted = 0 + ORDER BY m.id DESC LIMIT ?`, + channelID, limit, + ) + } + if err != nil { + return nil, fmt.Errorf("GetMessagesForAPI: %w", err) + } + defer rows.Close() + + 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 + } + } + + return msgs, nil +} + +// getReactionsBatch returns aggregated reactions for multiple messages. +func (d *DB) getReactionsBatch(msgIDs []int64, requestingUserID int64) (map[int64][]ReactionInfo, error) { + if len(msgIDs) == 0 { + return map[int64][]ReactionInfo{}, nil + } + + // Build placeholders for IN clause. + placeholders := "" + args := make([]any, 0, len(msgIDs)+len(msgIDs)) + for i, id := range msgIDs { + if i > 0 { + placeholders += "," + } + placeholders += "?" + args = append(args, id) + } + + // Query: aggregate count + check if requesting user reacted. + query := fmt.Sprintf( + `SELECT r.message_id, r.emoji, COUNT(*) as cnt, + MAX(CASE WHEN r.user_id = ? THEN 1 ELSE 0 END) as me + FROM reactions r + WHERE r.message_id IN (%s) + GROUP BY r.message_id, r.emoji`, + placeholders, + ) + args = append([]any{requestingUserID}, args...) + + rows, err := d.sqlDB.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("getReactionsBatch: %w", err) + } + defer rows.Close() + + result := make(map[int64][]ReactionInfo) + for rows.Next() { + var msgID int64 + var ri ReactionInfo + var me int + if scanErr := rows.Scan(&msgID, &ri.Emoji, &ri.Count, &me); scanErr != nil { + return nil, fmt.Errorf("getReactionsBatch scan: %w", scanErr) + } + ri.Me = me != 0 + result[msgID] = append(result[msgID], ri) + } + return result, nil +} + // UpdateReadState upserts the read state for a user in a channel. func (d *DB) UpdateReadState(userID, channelID, lastReadMessageID int64) error { _, err := d.sqlDB.Exec( diff --git a/Server/db/models.go b/Server/db/models.go index 7fa3fd71..20b5c7e7 100644 --- a/Server/db/models.go +++ b/Server/db/models.go @@ -98,12 +98,50 @@ type ReactionCount struct { // MessageSearchResult is a row returned by the FTS5 message search. type MessageSearchResult struct { - MessageID int64 - ChannelID int64 - ChannelName string - Username string - Content string - Timestamp string + MessageID int64 `json:"message_id"` + ChannelID int64 `json:"channel_id"` + ChannelName string `json:"channel_name"` + User UserPublic `json:"user"` + Content string `json:"content"` + Timestamp string `json:"timestamp"` +} + +// UserPublic is the public-facing user shape for API responses. +type UserPublic struct { + ID int64 `json:"id"` + Username string `json:"username"` + Avatar *string `json:"avatar,omitempty"` +} + +// MessageAPIResponse matches the API.md shape for GET /channels/{id}/messages. +type MessageAPIResponse struct { + ID int64 `json:"id"` + ChannelID int64 `json:"channel_id"` + User UserPublic `json:"user"` + Content string `json:"content"` + ReplyTo *int64 `json:"reply_to"` + Attachments []AttachmentInfo `json:"attachments"` + Reactions []ReactionInfo `json:"reactions"` + Pinned bool `json:"pinned"` + EditedAt *string `json:"edited_at"` + Deleted bool `json:"deleted"` + Timestamp string `json:"timestamp"` +} + +// AttachmentInfo is the attachment shape in API responses. +type AttachmentInfo struct { + ID string `json:"id"` + Filename string `json:"filename"` + Size int64 `json:"size"` + Mime string `json:"mime"` + URL string `json:"url"` +} + +// ReactionInfo is the reaction shape in API responses. +type ReactionInfo struct { + Emoji string `json:"emoji"` + Count int `json:"count"` + Me bool `json:"me"` } // VoiceState represents a row in the voice_states table.