Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 9 additions & 4 deletions apps/api/internal/httpapi/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -1152,7 +1152,7 @@ func (s *Server) addReaction(w http.ResponseWriter, r *http.Request) {
if err == nil && event.ID != "" {
s.publishEvent(r.Context(), event)
}
writeEventMutationResult(w, http.StatusCreated, event, err)
s.writeReactionMutationResult(w, r, http.StatusCreated, act.user.ID, chi.URLParam(r, "message_id"), event, err)
}

func (s *Server) removeReaction(w http.ResponseWriter, r *http.Request) {
Expand All @@ -1172,7 +1172,7 @@ func (s *Server) removeReaction(w http.ResponseWriter, r *http.Request) {
if err == nil && event.ID != "" {
s.publishEvent(r.Context(), event)
}
writeEventMutationResult(w, http.StatusOK, event, err)
s.writeReactionMutationResult(w, r, http.StatusOK, act.user.ID, chi.URLParam(r, "message_id"), event, err)
}

func (s *Server) listEvents(w http.ResponseWriter, r *http.Request) {
Expand Down Expand Up @@ -1581,7 +1581,12 @@ func writeThreadReplyCreateResult(w http.ResponseWriter, message store.Message,
writeJSON(w, status, map[string]any{"message": message, "thread_state": state, "events": events})
}

func writeEventMutationResult(w http.ResponseWriter, changedStatus int, event store.Event, err error) {
func (s *Server) writeReactionMutationResult(w http.ResponseWriter, r *http.Request, changedStatus int, userID, messageID string, event store.Event, err error) {
if err != nil {
writeStoreError(w, err)
return
}
message, err := s.store.GetMessage(r.Context(), messageID, userID)
if err != nil {
writeStoreError(w, err)
return
Expand All @@ -1590,7 +1595,7 @@ func writeEventMutationResult(w http.ResponseWriter, changedStatus int, event st
if event.ID != "" {
status = changedStatus
}
writeJSON(w, status, map[string]any{"event": event})
writeJSON(w, status, map[string]any{"event": event, "reactions": message.Reactions})
}

func writeJSON(w http.ResponseWriter, status int, body any) {
Expand Down
13 changes: 9 additions & 4 deletions apps/api/internal/httpapi/server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -246,16 +246,21 @@ func TestChatAPIVerticalSlice(t *testing.T) {
expectStatusAsUser(t, second.ID, http.MethodPost, server.URL+"/api/messages/"+created.Message.ID+"/attachments", strings.NewReader(`{"upload_id":"`+privateUpload.ID+`"}`), http.StatusForbidden)

reaction := postJSON[struct {
Event store.Event `json:"event"`
Event store.Event `json:"event"`
Reactions []store.ReactionSummary `json:"reactions"`
}](t, server.URL+"/api/messages/"+created.Message.ID+"/reactions", map[string]string{"emoji": "lobster"})
if reaction.Event.Type != "reaction.added" {
t.Fatalf("unexpected reaction event: %s", reaction.Event.Type)
}
if len(reaction.Reactions) != 1 || reaction.Reactions[0].Emoji != "lobster" || reaction.Reactions[0].Count != 1 || !reaction.Reactions[0].ReactedByMe {
t.Fatalf("unexpected reaction summaries: %#v", reaction.Reactions)
}
duplicateReaction, duplicateStatus := postJSONWithStatus[struct {
Event store.Event `json:"event"`
Event store.Event `json:"event"`
Reactions []store.ReactionSummary `json:"reactions"`
}](t, server.URL+"/api/messages/"+created.Message.ID+"/reactions", map[string]string{"emoji": "lobster"})
if duplicateStatus != http.StatusOK || duplicateReaction.Event.ID != "" {
t.Fatalf("expected duplicate reaction no-op, status=%d event=%#v", duplicateStatus, duplicateReaction.Event)
if duplicateStatus != http.StatusOK || duplicateReaction.Event.ID != "" || len(duplicateReaction.Reactions) != 1 {
t.Fatalf("expected duplicate reaction no-op, status=%d response=%#v", duplicateStatus, duplicateReaction)
}
deleteJSON(t, server.URL+"/api/messages/"+created.Message.ID+"/reactions/lobster")

Expand Down
5 changes: 3 additions & 2 deletions apps/api/internal/store/postgres/dms.go
Original file line number Diff line number Diff line change
Expand Up @@ -208,8 +208,9 @@ func (s *Store) ListDirectMessages(ctx context.Context, conversationID, userID s
return store.MessagePage{}, err
}
return s.listMessagePage(ctx, messagePageScope{
where: "m.direct_conversation_id = $1 AND m.parent_message_id IS NULL",
args: []any{conversationID},
where: "m.direct_conversation_id = $1 AND m.parent_message_id IS NULL",
args: []any{conversationID},
userID: userID,
}, page)
}

Expand Down
44 changes: 42 additions & 2 deletions apps/api/internal/store/postgres/message_pages.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,9 @@ import (
)

type messagePageScope struct {
where string
args []any
where string
args []any
userID string
}

type messagePageMode string
Expand Down Expand Up @@ -84,9 +85,48 @@ func (s *Store) listMessagePage(ctx context.Context, scope messagePageScope, req
if err != nil {
return store.MessagePage{}, err
}
messages, err = s.hydrateReactions(ctx, scope.userID, messages)
if err != nil {
return store.MessagePage{}, err
}
return s.buildMessagePage(ctx, scope, messages)
}

func (s *Store) hydrateReactions(ctx context.Context, userID string, messages []store.Message) ([]store.Message, error) {
ids := make([]string, len(messages))
for i, m := range messages {
ids[i] = m.ID
}
if len(ids) == 0 {
return messages, nil
}

rows, err := s.q.ListReactionsForMessages(ctx, storedb.ListReactionsForMessagesParams{
UserID: userID,
MessageIds: ids,
})
if err != nil {
return nil, fmt.Errorf("load reactions: %w", err)
}

reactionsByMsg := make(map[string][]store.ReactionSummary, len(ids))
for _, row := range rows {
reactionsByMsg[row.MessageID] = append(reactionsByMsg[row.MessageID], store.ReactionSummary{
Emoji: row.Emoji,
Count: row.ReactionCount,
ReactedByMe: row.ReactedByMe,
})
}

for i := range messages {
messages[i].Reactions = reactionsByMsg[messages[i].ID]
if messages[i].Reactions == nil {
messages[i].Reactions = []store.ReactionSummary{}
}
}
return messages, nil
}

func (s *Store) hydrateThreadStates(ctx context.Context, messages []store.Message) ([]store.Message, error) {
rootIDs := make([]string, 0, len(messages))
for _, message := range messages {
Expand Down
36 changes: 33 additions & 3 deletions apps/api/internal/store/postgres/postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -481,8 +481,9 @@ func (s *Store) ListMessages(ctx context.Context, channelID, userID string, page
return store.MessagePage{}, err
}
return s.listMessagePage(ctx, messagePageScope{
where: "m.channel_id = $1 AND m.parent_message_id IS NULL",
args: []any{channelID},
where: "m.channel_id = $1 AND m.parent_message_id IS NULL",
args: []any{channelID},
userID: userID,
}, page)
}

Expand All @@ -498,6 +499,10 @@ func (s *Store) GetMessage(ctx context.Context, messageID, userID string) (store
if err != nil {
return store.Message{}, err
}
messages, err = s.hydrateReactions(ctx, userID, messages)
if err != nil {
return store.Message{}, err
}
return messages[0], nil
}

Expand All @@ -517,6 +522,10 @@ func (s *Store) GetMessageByNonce(ctx context.Context, authorID, nonce string) (
if err != nil {
return store.Message{}, err
}
messages, err = s.hydrateReactions(ctx, authorID, messages)
if err != nil {
return store.Message{}, err
}
return messages[0], nil
}

Expand Down Expand Up @@ -704,6 +713,12 @@ func (s *Store) getThread(ctx context.Context, rootMessageID, userID string, lim
if err != nil {
return store.Message{}, nil, store.ThreadState{}, err
}
threadMessages := append([]store.Message{root}, replies...)
threadMessages, err = s.hydrateReactions(ctx, userID, threadMessages)
if err != nil {
return store.Message{}, nil, store.ThreadState{}, err
}
root, replies = threadMessages[0], threadMessages[1:]
state, err := getThreadState(ctx, s.db, rootMessageID)
return root, replies, state, err
}
Expand Down Expand Up @@ -927,6 +942,9 @@ func (s *Store) reaction(ctx context.Context, input store.CreateReactionInput, a
return store.Event{}, err
}
qtx := s.q.WithTx(tx)
if _, err := qtx.LockMessageForReaction(ctx, input.MessageID); err != nil {
return store.Event{}, err
}
var affected int64
if add {
affected, err = qtx.AddReaction(ctx, storedb.AddReactionParams{MessageID: input.MessageID, UserID: input.UserID, Emoji: input.Emoji, CreatedAt: now()})
Expand All @@ -939,11 +957,23 @@ func (s *Store) reaction(ctx context.Context, input store.CreateReactionInput, a
if affected == 0 {
return store.Event{}, tx.Commit()
}
count, err := qtx.CountMessageReaction(ctx, storedb.CountMessageReactionParams{
MessageID: input.MessageID,
Emoji: input.Emoji,
})
if err != nil {
return store.Event{}, err
}
eventType := "reaction.added"
if !add {
eventType = "reaction.removed"
}
payload := map[string]string{"message_id": input.MessageID, "emoji": input.Emoji}
payload := map[string]any{
"message_id": input.MessageID,
"emoji": input.Emoji,
"user_id": input.UserID,
"count": count,
}
if msg.DirectConversationID != "" {
payload["direct_conversation_id"] = msg.DirectConversationID
}
Expand Down
23 changes: 23 additions & 0 deletions apps/api/internal/store/postgres/sqlc/queries.sql
Original file line number Diff line number Diff line change
Expand Up @@ -1242,6 +1242,29 @@ WHERE message_id = sqlc.arg(message_id)
AND user_id = sqlc.arg(user_id)
AND emoji = sqlc.arg(emoji);

-- name: CountMessageReaction :one
SELECT COUNT(*)
FROM reactions
WHERE message_id = sqlc.arg(message_id)
AND emoji = sqlc.arg(emoji);

-- name: LockMessageForReaction :one
SELECT id
FROM messages
WHERE id = sqlc.arg(message_id)
FOR UPDATE;

-- name: ListReactionsForMessages :many
SELECT
r.message_id,
r.emoji,
COUNT(*)::bigint AS reaction_count,
BOOL_OR(r.user_id = sqlc.arg(user_id)) AS reacted_by_me
FROM reactions r
WHERE r.message_id = ANY(sqlc.arg(message_ids)::text[])
GROUP BY r.message_id, r.emoji
ORDER BY r.message_id, reaction_count DESC, r.emoji;

-- name: ListEventsAfter :many
SELECT e.id, e.cursor, e.workspace_id, COALESCE(e.channel_id, '') AS channel_id, e.type, e.seq, e.payload_json, e.created_at
FROM events e
Expand Down
85 changes: 85 additions & 0 deletions apps/api/internal/store/postgres/storedb/queries.sql.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading