owpengram-server/internal/store/postgres/message_polls.go

137 lines
5 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package postgres
import (
"context"
"errors"
"fmt"
"time"
"github.com/jackc/pgx/v5"
"telesrv/internal/domain"
"telesrv/internal/store/postgres/sqlcgen"
)
// 私聊消息 poll 投票/关闭:消息可见性在 message_boxes 校验(与 reaction 同款 SELECT FOR
// UPDATE 定位poll 级语义委托 poll.go 的共享 SQLpolls 行 FOR UPDATE 防 quiz 并发双投)。
func (s *MessageStore) VoteMessagePoll(ctx context.Context, req domain.VotePrivateMessagePollRequest) (domain.PrivateMessagePollResult, error) {
return s.mutateMessagePoll(ctx, req.UserID, req.Peer, req.MessageID, req.Date, func(ctx context.Context, tx pgx.Tx, def domain.PollDefinition, date int) error {
return applyPollVote(ctx, tx, def, req.UserID, req.Options, date)
})
}
func (s *MessageStore) CloseMessagePoll(ctx context.Context, req domain.ClosePrivateMessagePollRequest) (domain.PrivateMessagePollResult, error) {
return s.mutateMessagePoll(ctx, req.UserID, req.Peer, req.MessageID, req.Date, func(ctx context.Context, tx pgx.Tx, def domain.PollDefinition, _ int) error {
return closePollAsCreator(ctx, tx, def, req.UserID)
})
}
func (s *MessageStore) mutateMessagePoll(
ctx context.Context,
userID int64,
peer domain.Peer,
messageID int,
date int,
mutate func(ctx context.Context, tx pgx.Tx, def domain.PollDefinition, date int) error,
) (domain.PrivateMessagePollResult, error) {
if userID == 0 || peer.Type != domain.PeerTypeUser || peer.ID == 0 || messageID <= 0 || messageID > domain.MaxMessageBoxID {
return domain.PrivateMessagePollResult{}, domain.ErrMessageIDInvalid
}
if date == 0 {
date = int(time.Now().Unix())
}
beginner, ok := s.db.(txBeginner)
if !ok {
return domain.PrivateMessagePollResult{}, fmt.Errorf("mutate message poll: db does not support transactions")
}
tx, err := beginner.Begin(ctx)
if err != nil {
return domain.PrivateMessagePollResult{}, fmt.Errorf("begin message poll tx: %w", err)
}
committed := false
defer func() {
if !committed {
_ = tx.Rollback(ctx)
}
}()
var target struct {
privateMessageID int64
messageSenderID int64
mediaJSON string
}
if err := tx.QueryRow(ctx, `
SELECT private_message_id, message_sender_id, COALESCE(media::text, '{}')
FROM message_boxes
WHERE owner_user_id = $1
AND box_id = $2
AND peer_type = $3
AND peer_id = $4
AND NOT deleted
LIMIT 1`, userID, int32(messageID), string(peer.Type), peer.ID).Scan(&target.privateMessageID, &target.messageSenderID, &target.mediaJSON); err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return domain.PrivateMessagePollResult{}, domain.ErrMessageIDInvalid
}
return domain.PrivateMessagePollResult{}, fmt.Errorf("get message for poll: %w", err)
}
media, err := decodeMessageMedia(target.mediaJSON)
if err != nil {
return domain.PrivateMessagePollResult{}, fmt.Errorf("decode poll message media: %w", err)
}
if media == nil || media.Kind != domain.MessageMediaKindPoll || media.Poll == nil || media.Poll.ID == 0 {
return domain.PrivateMessagePollResult{}, domain.ErrMessageIDInvalid
}
pollID := media.Poll.ID
defs, err := loadPollDefinitions(ctx, tx, []int64{pollID}, true)
if err != nil {
return domain.PrivateMessagePollResult{}, err
}
def, ok := defs[pollID]
if !ok {
return domain.PrivateMessagePollResult{}, domain.ErrPollNotFound
}
if err := mutate(ctx, tx, def, date); err != nil {
return domain.PrivateMessagePollResult{}, err
}
boxes, err := sqlcgen.New(tx).ListVisibleMessageBoxesByPrivateMessage(ctx, sqlcgen.ListVisibleMessageBoxesByPrivateMessageParams{
OwnerUserIds: privateMessageOwnerIDs(userID, peer.ID),
MessageSenderID: target.messageSenderID,
PrivateMessageID: target.privateMessageID,
})
if err != nil {
return domain.PrivateMessagePollResult{}, fmt.Errorf("list visible poll boxes: %w", err)
}
res := domain.PrivateMessagePollResult{PollID: pollID, Messages: make([]domain.Message, 0, len(boxes))}
for _, box := range boxes {
msg, err := messageFromVisibleBoxRow(box)
if err != nil {
return domain.PrivateMessagePollResult{}, err
}
res.Messages = append(res.Messages, msg)
}
// poll enrichment按各 box owner 视角)由 enrichPrivateMessageReactions 统一挂载。
if err := s.enrichPrivateMessageReactions(ctx, tx, userID, res.Messages); err != nil {
return domain.PrivateMessagePollResult{}, err
}
if err := tx.Commit(ctx); err != nil {
return domain.PrivateMessagePollResult{}, fmt.Errorf("commit message poll tx: %w", err)
}
committed = true
return res, nil
}
// enrichPrivateMessagePolls 把页内全部 poll media 按各消息 owner 视角 enrich
// 由 enrichPrivateMessageReactions 统一挂载(所有私聊读路径共用一个 choke point
func (s *MessageStore) enrichPrivateMessagePolls(ctx context.Context, db sqlcgen.DBTX, viewerUserID int64, messages []domain.Message) error {
refs := make([]pollMediaRef, 0, 2)
for i := range messages {
viewer := messages[i].OwnerUserID
if viewer == 0 {
viewer = viewerUserID
}
refs = append(refs, pollMediaRef{media: messages[i].Media, viewer: viewer})
}
return enrichPollMediaRefs(ctx, db, refs)
}