From cf0f42a8dad2e0fe6087ea55b1ca6ffdae7c5cdc Mon Sep 17 00:00:00 2001 From: mengnanjiugaipeifense Date: Fri, 3 Jul 2026 09:59:48 +0800 Subject: [PATCH] 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 --- api/core/memory/token_buffer_memory.py | 37 ++++++++++++++----- .../core/memory/test_token_buffer_memory.py | 36 ++++++++++++++++++ 2 files changed, 64 insertions(+), 9 deletions(-) diff --git a/api/core/memory/token_buffer_memory.py b/api/core/memory/token_buffer_memory.py index c7db6d986f7..b83f3b84274 100644 --- a/api/core/memory/token_buffer_memory.py +++ b/api/core/memory/token_buffer_memory.py @@ -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( diff --git a/api/tests/unit_tests/core/memory/test_token_buffer_memory.py b/api/tests/unit_tests/core/memory/test_token_buffer_memory.py index 72c24bda96c..007486f3c34 100644 --- a/api/tests/unit_tests/core/memory/test_token_buffer_memory.py +++ b/api/tests/unit_tests/core/memory/test_token_buffer_memory.py @@ -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()