refactor(summary): remove system message handling in compression logic

This commit is contained in:
jinli.yl 2025-11-25 12:13:22 +08:00
parent 2143a73819
commit 55eb1f87bc
3 changed files with 8 additions and 26 deletions

View file

@ -58,15 +58,6 @@ class MessageCompactOp(BaseAsyncOp):
if keep_recent_count > 0: if keep_recent_count > 0:
messages_to_compress = messages_to_compress[:-keep_recent_count] messages_to_compress = messages_to_compress[:-keep_recent_count]
# Extract system message (should be exactly one)
system_message = [x for x in messages if x.role is Role.SYSTEM]
assert len(system_message) <= 1, f"Expected at most one system message, got {len(system_message)}"
if len(system_message) == 0:
system_message = Message(role=Role.SYSTEM, content="")
else:
system_message = system_message[0]
# If nothing to compress after filtering, return original messages # If nothing to compress after filtering, return original messages
if not messages_to_compress: if not messages_to_compress:
self.context.response.answer = self.context.messages self.context.response.answer = self.context.messages
@ -125,7 +116,6 @@ class MessageCompactOp(BaseAsyncOp):
if preview_char_length > 0: if preview_char_length > 0:
compact_result += f"\npreview: {original_content[:preview_char_length]}..." compact_result += f"\npreview: {original_content[:preview_char_length]}..."
tool_message.content = compact_result tool_message.content = compact_result
system_message.content += f"\n\n{compact_result}"
# Store write_file_dict in context for potential batch writing # Store write_file_dict in context for potential batch writing
if write_file_dict: if write_file_dict:

View file

@ -162,7 +162,6 @@ class MessageCompressOp(BaseAsyncOp):
async def _compress_with_groups( async def _compress_with_groups(
self, self,
system_message: Message,
message_groups: List[List[Message]], message_groups: List[List[Message]],
) -> Tuple[dict, list]: ) -> Tuple[dict, list]:
"""Compress multiple message groups and prepare them for storage. """Compress multiple message groups and prepare them for storage.
@ -173,7 +172,6 @@ class MessageCompressOp(BaseAsyncOp):
message. Otherwise, original messages are preserved in the return list. message. Otherwise, original messages are preserved in the return list.
Args: Args:
system_message: The system message to append compressed summaries to.
message_groups: List of message groups, where each group is a list of Message message_groups: List of message groups, where each group is a list of Message
objects to be compressed together. objects to be compressed together.
@ -190,7 +188,7 @@ class MessageCompressOp(BaseAsyncOp):
chat_id: str = self.context.get("chat_id", uuid4().hex) chat_id: str = self.context.get("chat_id", uuid4().hex)
# Create a copy of system_message to avoid modifying the original # Create a copy of system_message to avoid modifying the original
system_message_copy = Message(role=system_message.role, content=system_message.content) compress_message = Message(role=Role.ASSISTANT, content="")
for g_idx, messages in enumerate(message_groups): for g_idx, messages in enumerate(message_groups):
group_original_tokens = self.token_count(messages) group_original_tokens = self.token_count(messages)
@ -218,7 +216,7 @@ class MessageCompressOp(BaseAsyncOp):
) )
return_messages.extend(messages) return_messages.extend(messages)
else: else:
system_message_copy.content += compress_content + "\n\n" compress_message.content += compress_content + "\n\n"
write_file_dict[store_path.as_posix()] = messages_str write_file_dict[store_path.as_posix()] = messages_str
logger.info( logger.info(
f"Group {g_idx} compression successful: " f"Group {g_idx} compression successful: "
@ -227,7 +225,7 @@ class MessageCompressOp(BaseAsyncOp):
f"{100 * (1 - compressed_tokens / group_original_tokens):.1f}%)", f"{100 * (1 - compressed_tokens / group_original_tokens):.1f}%)",
) )
return_messages = [system_message_copy] + return_messages return_messages = [compress_message] + return_messages
return write_file_dict, return_messages return write_file_dict, return_messages
async def async_execute(self): async def async_execute(self):
@ -257,13 +255,7 @@ class MessageCompressOp(BaseAsyncOp):
messages = [Message(**x) for x in self.context.messages] messages = [Message(**x) for x in self.context.messages]
# Extract system message (should be exactly one) # Extract system message (should be exactly one)
system_message = [x for x in messages if x.role is Role.SYSTEM] system_messages = [x for x in messages if x.role is Role.SYSTEM]
assert len(system_message) <= 1, f"Expected at most one system message, got {len(system_message)}"
if len(system_message) == 0:
system_message = Message(role=Role.SYSTEM, content="")
else:
system_message = system_message[0]
messages_without_system = [x for x in messages if x.role is not Role.SYSTEM] messages_without_system = [x for x in messages if x.role is not Role.SYSTEM]
if keep_recent_count > 0: if keep_recent_count > 0:
@ -296,13 +288,13 @@ class MessageCompressOp(BaseAsyncOp):
else: else:
message_groups = [messages_to_compress] message_groups = [messages_to_compress]
write_file_dict, return_messages = await self._compress_with_groups(system_message, message_groups) write_file_dict, return_messages = await self._compress_with_groups(message_groups)
# Store write_file_dict in context for potential batch writing # Store write_file_dict in context for potential batch writing
if write_file_dict: if write_file_dict:
self.context.write_file_dict = write_file_dict self.context.write_file_dict = write_file_dict
self.context.response.answer = [x.simple_dump() for x in (return_messages + recent_messages)] self.context.response.answer = [x.simple_dump() for x in (system_messages + return_messages + recent_messages)]
self.context.response.metadata["write_file_dict"] = write_file_dict self.context.response.metadata["write_file_dict"] = write_file_dict
async def async_default_execute(self, e: Exception = None, **_kwargs): async def async_default_execute(self, e: Exception = None, **_kwargs):

View file

@ -70,9 +70,9 @@ async def test_agentic_retrieve_basic():
] ]
# llm = "qwen3_coder_plus" # llm = "qwen3_coder_plus"
# llm = "qwen3_30b_instruct" llm = "qwen3_30b_instruct"
# llm = "qwen3_30b_thinking" # llm = "qwen3_30b_thinking"
llm = "qwen3_coder_30b_instruct" # llm = "qwen3_coder_30b_instruct"
# llm = "qwen3_max_instruct" # llm = "qwen3_max_instruct"
op = AgenticRetrieveOp(llm=llm) op = AgenticRetrieveOp(llm=llm)