From b9109d48d5e09e409915b287e41be25fddbbf665 Mon Sep 17 00:00:00 2001 From: Bohdan Horbeshko Date: Sun, 24 Aug 2025 08:17:08 -0400 Subject: [PATCH] That MAM implementation was utterly broken --- telegram/utils.go | 75 +++++++++++++++++++++++++++++++---------------- xmpp/handlers.go | 61 +++++++++++++++++++++++++------------- 2 files changed, 89 insertions(+), 47 deletions(-) diff --git a/telegram/utils.go b/telegram/utils.go index 2fd1bca..b127ae8 100644 --- a/telegram/utils.go +++ b/telegram/utils.go @@ -2375,7 +2375,7 @@ func (c *Client) getNLastMessages(chatID int64, limit *MessageLimit) ([]*client. } // GetMessagesBetween lazily fetches message history between given ids (from exclusive, last inclusive), also calculating completeness flag; negative limit means messages from the end -func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64, limit int32) (messages []*client.Message, complete bool, err error) { +func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64, limit int32, reverse bool) (messages []*client.Message, complete bool, err error) { log.WithFields(log.Fields{ "chat_id": chatID, "from": fromMessageId, @@ -2399,6 +2399,11 @@ func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64, reqOffset = -limit - 1 reqLimit = limit + 1 } + log.WithFields(log.Fields{ + "from": reqFromMessageId, + "offset": reqOffset, + "limit": reqLimit, + }).Debug("calculated history request") newMessages, err = c.client.GetChatHistory(&client.GetChatHistoryRequest{ ChatId: chatID, @@ -2407,12 +2412,28 @@ func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64, Limit: reqLimit, }) if err == nil { - if limit < 0 { + log.Debugf("pre-fetched %v messages, cutting", len(newMessages.Messages)) + if limit > 0 { + complete = true fromPos := -1 - for i, message := range newMessages.Messages { - if message.Id >= fromMessageId { - fromPos = i+1 - break + if lastMessageId != 0 { + for i, message := range newMessages.Messages { + if message.Id < lastMessageId { + complete = false + break + } else if message.Id == lastMessageId { + fromPos = i + break + } + } + } else { + if len(newMessages.Messages) > 0 { + lastMsg, lastMsgErr := c.GetPreviousMessage(chatID, 0) + if lastMsgErr == nil && lastMsg != nil { + if lastMsg.Id != newMessages.Messages[0].Id { + complete = false + } + } } } if fromPos > -1 { @@ -2420,27 +2441,19 @@ func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64, } else { messages = newMessages.Messages } - complete = true } else { + complete = true for _, message := range newMessages.Messages { - if message.Id == fromMessageId { - continue - } - if lastMessageId != 0 { - if message.Id > lastMessageId { - break - complete = true - } else if message.Id == lastMessageId { - complete = true - } + if message.Id <= fromMessageId { + break } messages = append(messages, message) } - if len(messages) == 0 { - complete = true - } } } + if !reverse { + ReverseMessagesSlice(messages) + } return } @@ -3394,12 +3407,22 @@ func GetErrorCode(err error) (int32, bool) { } // ChronologicallySortMessages is… self-explanatory (Achtung: destructive) -func ChronologicallySortMessages(messages []*client.Message) []*client.Message { - sort.Slice(messages, func(i int, j int) bool { - msg1 := messages[i] - msg2 := messages[j] - return msg1.Date < msg2.Date - }) +func ChronologicallySortMessages(messages []*client.Message, reverse bool) []*client.Message { + var sortFunc func(int, int) bool + if reverse { + sortFunc = func(i int, j int) bool { + msg1 := messages[i] + msg2 := messages[j] + return msg1.Date > msg2.Date + } + } else { + sortFunc = func(i int, j int) bool { + msg1 := messages[i] + msg2 := messages[j] + return msg1.Date < msg2.Date + } + } + sort.Slice(messages, sortFunc) return messages } diff --git a/xmpp/handlers.go b/xmpp/handlers.go index b847cac..90aa9d6 100644 --- a/xmpp/handlers.go +++ b/xmpp/handlers.go @@ -2259,6 +2259,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer before := *query.ResultSet.Before if before == "" { rsmLastPage = true + log.Debugf("MAM RSM last page") } else { rsmBefore, ok = parseMessageId(*query.ResultSet.Before) if !ok { @@ -2303,11 +2304,12 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer ((beforeId != 0 || afterId != 0) && (!startTime.IsZero() || !endTime.IsZero() || ids != nil)) || (ids != nil && (!startTime.IsZero() || !endTime.IsZero() || beforeId != 0 || afterId != 0)) { iqAnswerSetError(answer, 501) + log.Debugf("MAM: incompatible parameters") return } // dummy call to circumvent unexported type - messages, _, err := session.GetMessagesBetween(toID, 0, 0, 0) + messages, _, err := session.GetMessagesBetween(toID, 0, 0, 0, false) quotaTs := time.Now().AddDate(0, 0, -int(gateway.MAMThreshold)) var beyond, complete bool @@ -2315,25 +2317,30 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer fromStart := true var overallyFirstMessageId, overallyLastMessageId int64 + reverse := query.FlipPage != nil + if ids != nil { for _, sId := range ids { id, ok := parseMessageId(sId) if !ok { iqAnswerSetError(answer, 400) + log.Debugf("MAM: bogus id") return } msg, err := session.GetMessage(toID, id) if err != nil { iqAnswerSetError(answer, 404) + log.Debugf("MAM: unknown id") return } messages = append(messages, msg) } - messages = telegram.ChronologicallySortMessages(messages) + messages = telegram.ChronologicallySortMessages(messages, reverse) complete = true } else if beforeId != 0 || afterId != 0 { if (beforeId != 0 && afterId != 0) && beforeId > afterId { iqAnswerSetError(answer, 400) + log.Debugf("MAM: before > after") return } @@ -2347,6 +2354,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer } } else { iqAnswerSetError(answer, 404) + log.Debugf("MAM: unknown after") return } } else { @@ -2366,6 +2374,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer beforeMessage, beforeMessageErr := session.GetMessage(toID, beforeId) if beforeMessageErr != nil || beforeMessage == nil { iqAnswerSetError(answer, 404) + log.Debugf("MAM: unknown before") return } overallyLastMessage, overallyLastMessageErr := session.GetPreviousMessage(toID, beforeId) @@ -2382,6 +2391,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer if rsmAfter != 0 { if rsmAfter < afterId || (beforeId != 0 && rsmAfter > beforeId) { iqAnswerSetError(answer, 400) + log.Debugf("MAM: unknown RSM after") return } if rsmAfter != afterId { @@ -2397,6 +2407,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer if rsmBefore != 0 { if rsmBefore < afterId || (beforeId != 0 && rsmBefore > beforeId) { iqAnswerSetError(answer, 400) + log.Debugf("MAM: unknown RSM before") return } if rsmBefore < beforeId { @@ -2416,7 +2427,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer if !beyond { var newComplete bool - messages, newComplete, err = session.GetMessagesBetween(toID, afterId, lastMessageId, rsmLimit) + messages, newComplete, err = session.GetMessagesBetween(toID, afterId, lastMessageId, rsmLimit, reverse) if canBeComplete { complete = newComplete } @@ -2461,12 +2472,13 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer if rsmAfter != 0 { rsmAfterMessage, rsmAfterMessageErr := session.GetMessage(toID, rsmAfter) if rsmAfterMessageErr == nil && rsmAfterMessage != nil { - if int64(rsmAfterMessage.Date) < startTime.Unix() || (!endTime.IsZero() && int64(rsmAfterMessage.Date) > endTime.Unix()) { + if !endTime.IsZero() && int64(rsmAfterMessage.Date) > endTime.Unix() { iqAnswerSetError(answer, 400) + log.Debugf("MAM: RSM after out of range") return } - if rsmAfterMessage.Id > fromMessageId { + if rsmAfterMessage.Id > fromMessageId && int64(rsmAfterMessage.Date) >= startTime.Unix() { fromStart = false overallyFirstMessage, overallyFirstMessageErr := session.GetNextMessage(toID, fromMessageId) if overallyFirstMessageErr == nil && overallyFirstMessage != nil { @@ -2476,6 +2488,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer fromMessageId = rsmAfterMessage.Id } else { iqAnswerSetError(answer, 404) + log.Debugf("MAM: unknown RSM after") return } } @@ -2483,30 +2496,34 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer if rsmBefore != 0 { rsmBeforeMessage, rsmBeforeMessageErr := session.GetMessage(toID, rsmBefore) if rsmBeforeMessageErr == nil && rsmBeforeMessage != nil { - if int64(rsmBeforeMessage.Date) < startTime.Unix() || (!endTime.IsZero() && int64(rsmBeforeMessage.Date) > endTime.Unix()) { + if int64(rsmBeforeMessage.Date) < startTime.Unix() { iqAnswerSetError(answer, 400) + log.Debugf("MAM: RSM before out of range") return } - newLastMessage, newLastMessageErr := session.GetPreviousMessage(toID, rsmBeforeMessage.Id) - if newLastMessageErr == nil && newLastMessage != nil { - if lastMessageId == 0 || lastMessageId != newLastMessage.Id { - canBeComplete = false - lastMessageId = newLastMessage.Id + if endTime.IsZero() || int64(rsmBeforeMessage.Date) <= endTime.Unix() { + newLastMessage, newLastMessageErr := session.GetPreviousMessage(toID, rsmBeforeMessage.Id) + if newLastMessageErr == nil && newLastMessage != nil { + if lastMessageId == 0 || lastMessageId != newLastMessage.Id { + canBeComplete = false + lastMessageId = newLastMessage.Id + } + } else { + // nothing?.. not complete, just empty + beyond = true } - } else { - // nothing?.. not complete, just empty - beyond = true } } else { iqAnswerSetError(answer, 404) + log.Debugf("MAM: unknown RSM before") return } } if !beyond { // yes🗿, twice var newComplete bool - messages, newComplete, err = session.GetMessagesBetween(toID, fromMessageId, lastMessageId, rsmLimit) + messages, newComplete, err = session.GetMessagesBetween(toID, fromMessageId, lastMessageId, rsmLimit, reverse) if canBeComplete { complete = newComplete } @@ -2514,10 +2531,6 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer } } - if query.FlipPage != nil { - telegram.ReverseMessagesSlice(messages) - } - log.Debugf("obtained %v messages", len(messages)) for _, message := range messages { session.SendDelayedMUCMessage(toID, message, iq.From, query.QueryId) @@ -2555,8 +2568,14 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer if lastMsgPositionErr == nil && lastMsgPosition != nil { lastMsgPositionCount = lastMsgPosition.Count } + log.WithFields(log.Fields{ + "overallyFirstMessageId": overallyFirstMessageId, + "overallyLastMessageId": overallyLastMessageId, + "firstMsgPositionCount": firstMsgPositionCount, + "lastMsgPositionCount": lastMsgPositionCount, + }).Debug("RSM count") if firstMsgPositionCount != 0 && lastMsgPositionCount != 0 { - count := int(lastMsgPositionCount - firstMsgPositionCount + 1) + count := int(firstMsgPositionCount - lastMsgPositionCount + 1) rs.Count = &count } } @@ -2568,7 +2587,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer if firstMsgPositionCount != 0 { rsmFirstMsgPosition, rsmFirstMsgPositionErr := session.GetChatMessagePosition(toID, messages[0].Id) if rsmFirstMsgPositionErr == nil && rsmFirstMsgPosition != nil { - index := int(rsmFirstMsgPosition.Count - firstMsgPositionCount) + index := int(firstMsgPositionCount - rsmFirstMsgPosition.Count) rs.First.Index = &index } }