That MAM implementation was utterly broken

This commit is contained in:
Bohdan Horbeshko 2025-08-24 08:17:08 -04:00
parent e59ef598a3
commit b9109d48d5
2 changed files with 89 additions and 47 deletions

View file

@ -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 // 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{ log.WithFields(log.Fields{
"chat_id": chatID, "chat_id": chatID,
"from": fromMessageId, "from": fromMessageId,
@ -2399,6 +2399,11 @@ func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64,
reqOffset = -limit - 1 reqOffset = -limit - 1
reqLimit = 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{ newMessages, err = c.client.GetChatHistory(&client.GetChatHistoryRequest{
ChatId: chatID, ChatId: chatID,
@ -2407,12 +2412,28 @@ func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64,
Limit: reqLimit, Limit: reqLimit,
}) })
if err == nil { if err == nil {
if limit < 0 { log.Debugf("pre-fetched %v messages, cutting", len(newMessages.Messages))
if limit > 0 {
complete = true
fromPos := -1 fromPos := -1
if lastMessageId != 0 {
for i, message := range newMessages.Messages { for i, message := range newMessages.Messages {
if message.Id >= fromMessageId { if message.Id < lastMessageId {
fromPos = i+1 complete = false
break 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 { if fromPos > -1 {
@ -2420,26 +2441,18 @@ func (c *Client) GetMessagesBetween(chatID, fromMessageId, lastMessageId int64,
} else { } else {
messages = newMessages.Messages messages = newMessages.Messages
} }
complete = true
} else { } else {
complete = true
for _, message := range newMessages.Messages { for _, message := range newMessages.Messages {
if message.Id == fromMessageId { if message.Id <= fromMessageId {
continue
}
if lastMessageId != 0 {
if message.Id > lastMessageId {
break break
complete = true
} else if message.Id == lastMessageId {
complete = true
}
} }
messages = append(messages, message) messages = append(messages, message)
} }
if len(messages) == 0 {
complete = true
} }
} }
if !reverse {
ReverseMessagesSlice(messages)
} }
return return
} }
@ -3394,12 +3407,22 @@ func GetErrorCode(err error) (int32, bool) {
} }
// ChronologicallySortMessages is… self-explanatory (Achtung: destructive) // ChronologicallySortMessages is… self-explanatory (Achtung: destructive)
func ChronologicallySortMessages(messages []*client.Message) []*client.Message { func ChronologicallySortMessages(messages []*client.Message, reverse bool) []*client.Message {
sort.Slice(messages, func(i int, j int) bool { 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] msg1 := messages[i]
msg2 := messages[j] msg2 := messages[j]
return msg1.Date < msg2.Date return msg1.Date < msg2.Date
}) }
}
sort.Slice(messages, sortFunc)
return messages return messages
} }

View file

@ -2259,6 +2259,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
before := *query.ResultSet.Before before := *query.ResultSet.Before
if before == "" { if before == "" {
rsmLastPage = true rsmLastPage = true
log.Debugf("MAM RSM last page")
} else { } else {
rsmBefore, ok = parseMessageId(*query.ResultSet.Before) rsmBefore, ok = parseMessageId(*query.ResultSet.Before)
if !ok { 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)) || ((beforeId != 0 || afterId != 0) && (!startTime.IsZero() || !endTime.IsZero() || ids != nil)) ||
(ids != nil && (!startTime.IsZero() || !endTime.IsZero() || beforeId != 0 || afterId != 0)) { (ids != nil && (!startTime.IsZero() || !endTime.IsZero() || beforeId != 0 || afterId != 0)) {
iqAnswerSetError(answer, 501) iqAnswerSetError(answer, 501)
log.Debugf("MAM: incompatible parameters")
return return
} }
// dummy call to circumvent unexported type // 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)) quotaTs := time.Now().AddDate(0, 0, -int(gateway.MAMThreshold))
var beyond, complete bool var beyond, complete bool
@ -2315,25 +2317,30 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
fromStart := true fromStart := true
var overallyFirstMessageId, overallyLastMessageId int64 var overallyFirstMessageId, overallyLastMessageId int64
reverse := query.FlipPage != nil
if ids != nil { if ids != nil {
for _, sId := range ids { for _, sId := range ids {
id, ok := parseMessageId(sId) id, ok := parseMessageId(sId)
if !ok { if !ok {
iqAnswerSetError(answer, 400) iqAnswerSetError(answer, 400)
log.Debugf("MAM: bogus id")
return return
} }
msg, err := session.GetMessage(toID, id) msg, err := session.GetMessage(toID, id)
if err != nil { if err != nil {
iqAnswerSetError(answer, 404) iqAnswerSetError(answer, 404)
log.Debugf("MAM: unknown id")
return return
} }
messages = append(messages, msg) messages = append(messages, msg)
} }
messages = telegram.ChronologicallySortMessages(messages) messages = telegram.ChronologicallySortMessages(messages, reverse)
complete = true complete = true
} else if beforeId != 0 || afterId != 0 { } else if beforeId != 0 || afterId != 0 {
if (beforeId != 0 && afterId != 0) && beforeId > afterId { if (beforeId != 0 && afterId != 0) && beforeId > afterId {
iqAnswerSetError(answer, 400) iqAnswerSetError(answer, 400)
log.Debugf("MAM: before > after")
return return
} }
@ -2347,6 +2354,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
} }
} else { } else {
iqAnswerSetError(answer, 404) iqAnswerSetError(answer, 404)
log.Debugf("MAM: unknown after")
return return
} }
} else { } else {
@ -2366,6 +2374,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
beforeMessage, beforeMessageErr := session.GetMessage(toID, beforeId) beforeMessage, beforeMessageErr := session.GetMessage(toID, beforeId)
if beforeMessageErr != nil || beforeMessage == nil { if beforeMessageErr != nil || beforeMessage == nil {
iqAnswerSetError(answer, 404) iqAnswerSetError(answer, 404)
log.Debugf("MAM: unknown before")
return return
} }
overallyLastMessage, overallyLastMessageErr := session.GetPreviousMessage(toID, beforeId) 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 != 0 {
if rsmAfter < afterId || (beforeId != 0 && rsmAfter > beforeId) { if rsmAfter < afterId || (beforeId != 0 && rsmAfter > beforeId) {
iqAnswerSetError(answer, 400) iqAnswerSetError(answer, 400)
log.Debugf("MAM: unknown RSM after")
return return
} }
if rsmAfter != afterId { if rsmAfter != afterId {
@ -2397,6 +2407,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
if rsmBefore != 0 { if rsmBefore != 0 {
if rsmBefore < afterId || (beforeId != 0 && rsmBefore > beforeId) { if rsmBefore < afterId || (beforeId != 0 && rsmBefore > beforeId) {
iqAnswerSetError(answer, 400) iqAnswerSetError(answer, 400)
log.Debugf("MAM: unknown RSM before")
return return
} }
if rsmBefore < beforeId { if rsmBefore < beforeId {
@ -2416,7 +2427,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
if !beyond { if !beyond {
var newComplete bool var newComplete bool
messages, newComplete, err = session.GetMessagesBetween(toID, afterId, lastMessageId, rsmLimit) messages, newComplete, err = session.GetMessagesBetween(toID, afterId, lastMessageId, rsmLimit, reverse)
if canBeComplete { if canBeComplete {
complete = newComplete complete = newComplete
} }
@ -2461,12 +2472,13 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
if rsmAfter != 0 { if rsmAfter != 0 {
rsmAfterMessage, rsmAfterMessageErr := session.GetMessage(toID, rsmAfter) rsmAfterMessage, rsmAfterMessageErr := session.GetMessage(toID, rsmAfter)
if rsmAfterMessageErr == nil && rsmAfterMessage != nil { 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) iqAnswerSetError(answer, 400)
log.Debugf("MAM: RSM after out of range")
return return
} }
if rsmAfterMessage.Id > fromMessageId { if rsmAfterMessage.Id > fromMessageId && int64(rsmAfterMessage.Date) >= startTime.Unix() {
fromStart = false fromStart = false
overallyFirstMessage, overallyFirstMessageErr := session.GetNextMessage(toID, fromMessageId) overallyFirstMessage, overallyFirstMessageErr := session.GetNextMessage(toID, fromMessageId)
if overallyFirstMessageErr == nil && overallyFirstMessage != nil { if overallyFirstMessageErr == nil && overallyFirstMessage != nil {
@ -2476,6 +2488,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
fromMessageId = rsmAfterMessage.Id fromMessageId = rsmAfterMessage.Id
} else { } else {
iqAnswerSetError(answer, 404) iqAnswerSetError(answer, 404)
log.Debugf("MAM: unknown RSM after")
return return
} }
} }
@ -2483,11 +2496,13 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
if rsmBefore != 0 { if rsmBefore != 0 {
rsmBeforeMessage, rsmBeforeMessageErr := session.GetMessage(toID, rsmBefore) rsmBeforeMessage, rsmBeforeMessageErr := session.GetMessage(toID, rsmBefore)
if rsmBeforeMessageErr == nil && rsmBeforeMessage != nil { 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) iqAnswerSetError(answer, 400)
log.Debugf("MAM: RSM before out of range")
return return
} }
if endTime.IsZero() || int64(rsmBeforeMessage.Date) <= endTime.Unix() {
newLastMessage, newLastMessageErr := session.GetPreviousMessage(toID, rsmBeforeMessage.Id) newLastMessage, newLastMessageErr := session.GetPreviousMessage(toID, rsmBeforeMessage.Id)
if newLastMessageErr == nil && newLastMessage != nil { if newLastMessageErr == nil && newLastMessage != nil {
if lastMessageId == 0 || lastMessageId != newLastMessage.Id { if lastMessageId == 0 || lastMessageId != newLastMessage.Id {
@ -2498,15 +2513,17 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
// nothing?.. not complete, just empty // nothing?.. not complete, just empty
beyond = true beyond = true
} }
}
} else { } else {
iqAnswerSetError(answer, 404) iqAnswerSetError(answer, 404)
log.Debugf("MAM: unknown RSM before")
return return
} }
} }
if !beyond { // yes🗿, twice if !beyond { // yes🗿, twice
var newComplete bool var newComplete bool
messages, newComplete, err = session.GetMessagesBetween(toID, fromMessageId, lastMessageId, rsmLimit) messages, newComplete, err = session.GetMessagesBetween(toID, fromMessageId, lastMessageId, rsmLimit, reverse)
if canBeComplete { if canBeComplete {
complete = newComplete 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)) log.Debugf("obtained %v messages", len(messages))
for _, message := range messages { for _, message := range messages {
session.SendDelayedMUCMessage(toID, message, iq.From, query.QueryId) 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 { if lastMsgPositionErr == nil && lastMsgPosition != nil {
lastMsgPositionCount = lastMsgPosition.Count lastMsgPositionCount = lastMsgPosition.Count
} }
log.WithFields(log.Fields{
"overallyFirstMessageId": overallyFirstMessageId,
"overallyLastMessageId": overallyLastMessageId,
"firstMsgPositionCount": firstMsgPositionCount,
"lastMsgPositionCount": lastMsgPositionCount,
}).Debug("RSM count")
if firstMsgPositionCount != 0 && lastMsgPositionCount != 0 { if firstMsgPositionCount != 0 && lastMsgPositionCount != 0 {
count := int(lastMsgPositionCount - firstMsgPositionCount + 1) count := int(firstMsgPositionCount - lastMsgPositionCount + 1)
rs.Count = &count rs.Count = &count
} }
} }
@ -2568,7 +2587,7 @@ func handleSetQueryMAM2(s xmpp.Sender, iq *stanza.IQ, query *extensions.MAM2Quer
if firstMsgPositionCount != 0 { if firstMsgPositionCount != 0 {
rsmFirstMsgPosition, rsmFirstMsgPositionErr := session.GetChatMessagePosition(toID, messages[0].Id) rsmFirstMsgPosition, rsmFirstMsgPositionErr := session.GetChatMessagePosition(toID, messages[0].Id)
if rsmFirstMsgPositionErr == nil && rsmFirstMsgPosition != nil { if rsmFirstMsgPositionErr == nil && rsmFirstMsgPosition != nil {
index := int(rsmFirstMsgPosition.Count - firstMsgPositionCount) index := int(firstMsgPositionCount - rsmFirstMsgPosition.Count)
rs.First.Index = &index rs.First.Index = &index
} }
} }