perf(memory): batch-load message files in TokenBufferMemory to remove N+1 queries (#38002)

Co-authored-by: xiaweiwei67-stack <293320877+xiaweiwei67-stack@users.noreply.github.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
mengnanjiugaipeifense 2026-07-03 09:59:48 +08:00 committed by GitHub
parent d120995efc
commit cf0f42a8da
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 64 additions and 9 deletions

View File

@ -1,3 +1,4 @@
from collections import defaultdict
from collections.abc import Sequence
from sqlalchemy import select
@ -153,16 +154,36 @@ class TokenBufferMemory:
messages = list(reversed(thread_messages))
# Batch-load message files for the whole thread to avoid an N+1 query pattern.
# Previously each message issued two MessageFile queries (user + assistant),
# i.e. 2N+1 round-trips for N messages. We now use two batched queries keyed by
# message_id, preserving the exact filter semantics (user files include rows
# whose belongs_to is NULL).
message_ids = [message.id for message in messages]
user_files_by_message: dict[str, list[MessageFile]] = defaultdict(list)
assistant_files_by_message: dict[str, list[MessageFile]] = defaultdict(list)
if message_ids:
for message_file in db.session.scalars(
select(MessageFile).where(
MessageFile.message_id.in_(message_ids),
(MessageFile.belongs_to == "user") | (MessageFile.belongs_to.is_(None)),
)
).all():
user_files_by_message[message_file.message_id].append(message_file)
for message_file in db.session.scalars(
select(MessageFile).where(
MessageFile.message_id.in_(message_ids),
MessageFile.belongs_to == "assistant",
)
).all():
assistant_files_by_message[message_file.message_id].append(message_file)
curr_message_tokens = 0
prompt_messages: list[PromptMessage] = []
for message in messages:
# Process user message with files
user_files = db.session.scalars(
select(MessageFile).where(
MessageFile.message_id == message.id,
(MessageFile.belongs_to == "user") | (MessageFile.belongs_to.is_(None)),
)
).all()
user_files = user_files_by_message.get(message.id, [])
if user_files:
user_prompt_message = self._build_prompt_message_with_files(
@ -177,9 +198,7 @@ class TokenBufferMemory:
prompt_messages.append(UserPromptMessage(content=message.query))
# Process assistant message with files
assistant_files = db.session.scalars(
select(MessageFile).where(MessageFile.message_id == message.id, MessageFile.belongs_to == "assistant")
).all()
assistant_files = assistant_files_by_message.get(message.id, [])
if assistant_files:
assistant_prompt_message = self._build_prompt_message_with_files(

View File

@ -629,6 +629,7 @@ class TestGetHistoryPromptMessages:
msg.parent_message_id = None
mock_user_file = MagicMock()
mock_user_file.message_id = msg.id # must match so batched grouping keys it to this message
mock_user_prompt = UserPromptMessage(content="from build")
mock_assistant_prompt = AssistantPromptMessage(content="answer")
@ -679,6 +680,7 @@ class TestGetHistoryPromptMessages:
msg.parent_message_id = None
mock_assistant_file = MagicMock()
mock_assistant_file.message_id = msg.id # must match so batched grouping keys it to this message
mock_user_prompt = UserPromptMessage(content="query")
mock_assistant_prompt = AssistantPromptMessage(content="built")
@ -714,6 +716,40 @@ class TestGetHistoryPromptMessages:
call_kwargs = mock_build.call_args[1]
assert call_kwargs["is_user_message"] is False
def test_message_files_loaded_with_constant_query_count(self):
"""Regression guard against N+1: message files must be batch-loaded.
Regardless of the number of messages in the thread, file loading must use a
constant number of queries (1 messages query + 2 batched file queries),
never 2 queries per message.
"""
mem = self._make_memory()
messages = [_make_message() for _ in range(5)]
for m in messages:
m.parent_message_id = None
scalars_calls = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
# First call returns the thread messages; the batched file queries return none.
r.all.return_value = messages if scalars_calls["n"] == 0 else []
scalars_calls["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=messages),
patch("core.memory.token_buffer_memory.FileUploadConfigManager.convert", return_value=None),
):
mock_db.session.scalars.side_effect = scalars_side_effect
mem.get_history_prompt_messages()
# 1 (messages) + 2 (batched user/assistant files) = 3, independent of message count.
# Before this fix it would have been 1 + 2 * 5 = 11 (an N+1 pattern).
assert scalars_calls["n"] == 3
def test_token_pruning_removes_oldest_messages(self):
"""If tokens exceed limit, oldest messages are removed until within limit."""
conv = _make_conversation()