mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
commit
873104a2b7
3 changed files with 866 additions and 0 deletions
494
reme_ai/context/offload/context_compress_op.py
Normal file
494
reme_ai/context/offload/context_compress_op.py
Normal file
|
|
@ -0,0 +1,494 @@
|
|||
"""
|
||||
Context compression module for reducing token usage in conversation contexts using LLM.
|
||||
|
||||
This module provides functionality to compress conversation history by using a language
|
||||
model to generate concise summaries of older messages while preserving recent messages.
|
||||
This helps manage context window limits while maintaining conversation coherence.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
import xml.etree.ElementTree as ET
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
from uuid import uuid4
|
||||
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from flowllm.core.utils import extract_content
|
||||
from loguru import logger
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class ContextCompressOp(BaseAsyncOp):
|
||||
"""
|
||||
Context compression operation that uses LLM to reduce token usage.
|
||||
|
||||
When the total token count exceeds the threshold, this operation uses a language
|
||||
model to compress older messages into a concise summary while keeping recent
|
||||
messages intact. This preserves conversation context while reducing token usage.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
all_token_threshold: int = 20000,
|
||||
keep_recent: int = 5,
|
||||
storage_path: str = "./compressed_contexts",
|
||||
micro_summary_token_threshold: int = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize the context compression operation.
|
||||
|
||||
Args:
|
||||
all_token_threshold: Maximum total token count before compression is triggered.
|
||||
keep_recent: Number of recent messages to keep uncompressed.
|
||||
storage_path: Directory path where original messages will be stored for traceability.
|
||||
micro_summary_token_threshold: Token threshold for each compression group.
|
||||
If set, messages will be split into groups of this size and compressed separately.
|
||||
If None, all messages will be compressed together.
|
||||
**kwargs: Additional arguments passed to the base class.
|
||||
|
||||
Note:
|
||||
System messages are NEVER compressed to preserve important system instructions.
|
||||
"""
|
||||
super().__init__(**kwargs)
|
||||
self.all_token_threshold: int = all_token_threshold
|
||||
self.keep_recent: int = keep_recent
|
||||
self.storage_path: Path = Path(storage_path)
|
||||
self.micro_summary_token_threshold: int = micro_summary_token_threshold
|
||||
|
||||
assert (
|
||||
micro_summary_token_threshold is None or micro_summary_token_threshold > 0
|
||||
), "Micro summary token threshold must be greater than 0"
|
||||
|
||||
# Create storage directory if it doesn't exist
|
||||
self.storage_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def _save_original_messages(self, messages: List[Message]) -> str:
|
||||
"""Save original messages to file for traceability.
|
||||
|
||||
Args:
|
||||
messages: List of messages to save
|
||||
|
||||
Returns:
|
||||
Path to the saved file
|
||||
"""
|
||||
# Generate unique filename with timestamp
|
||||
file_name = f"context_{uuid4().hex}.txt"
|
||||
file_path = self.storage_path / file_name
|
||||
|
||||
# Convert messages to serializable format
|
||||
messages_data = [
|
||||
{
|
||||
"role": msg.role.value if hasattr(msg.role, "value") else str(msg.role),
|
||||
"content": msg.content,
|
||||
"name": getattr(msg, "name", None),
|
||||
"tool_call_id": getattr(msg, "tool_call_id", None),
|
||||
}
|
||||
for msg in messages
|
||||
]
|
||||
|
||||
# Save to file
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(messages_data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
logger.info(f"Saved {len(messages)} original messages to {file_path}")
|
||||
return str(file_path)
|
||||
|
||||
def _split_messages_by_token_threshold(
|
||||
self,
|
||||
messages: List[Message],
|
||||
token_threshold: int,
|
||||
) -> List[List[Message]]:
|
||||
"""Split messages into groups based on token threshold.
|
||||
|
||||
Args:
|
||||
messages: List of messages to split
|
||||
token_threshold: Maximum token count for each group
|
||||
|
||||
Returns:
|
||||
List of message groups, each within the token threshold
|
||||
"""
|
||||
if not messages:
|
||||
return []
|
||||
|
||||
groups = []
|
||||
current_group = []
|
||||
current_token_count = 0
|
||||
|
||||
for msg in messages:
|
||||
msg_tokens = self.token_count([msg])
|
||||
|
||||
# If single message exceeds threshold, put it in its own group
|
||||
if msg_tokens > token_threshold:
|
||||
if current_group:
|
||||
groups.append(current_group)
|
||||
current_group = []
|
||||
current_token_count = 0
|
||||
groups.append([msg])
|
||||
continue
|
||||
|
||||
# If adding this message would exceed threshold, start new group
|
||||
if current_token_count + msg_tokens > token_threshold and current_group:
|
||||
groups.append(current_group)
|
||||
current_group = [msg]
|
||||
current_token_count = msg_tokens
|
||||
else:
|
||||
current_group.append(msg)
|
||||
current_token_count += msg_tokens
|
||||
|
||||
# Add the last group if it has messages
|
||||
if current_group:
|
||||
groups.append(current_group)
|
||||
|
||||
logger.info(
|
||||
f"Split {len(messages)} messages into {len(groups)} groups " f"with token threshold {token_threshold}",
|
||||
)
|
||||
return groups
|
||||
|
||||
@staticmethod
|
||||
def _extract_xml_fragments(text: str) -> str:
|
||||
"""
|
||||
Extract XML fragments from text, removing scratchpad elements.
|
||||
|
||||
Scans text to extract complete and parseable top-level XML fragments,
|
||||
excluding <scratchpad> elements. If state_snapshot XML is found, returns it;
|
||||
otherwise returns the original text.
|
||||
|
||||
Args:
|
||||
text: Input text potentially containing XML fragments
|
||||
|
||||
Returns:
|
||||
Extracted XML content or original text
|
||||
"""
|
||||
try:
|
||||
# Remove scratchpad elements
|
||||
new_text = re.sub(r"<scratchpad>.*?</scratchpad>", "", text, flags=re.S | re.I)
|
||||
# Extract balanced XML tags
|
||||
extract_xml = [m[0] for m in re.findall(r"(<(\w+)[^>]*>(?:[^<]|<(?!/\2))*</\2>)", new_text)]
|
||||
|
||||
# Validate XML parsing
|
||||
valid_xml = []
|
||||
for xml_str in extract_xml:
|
||||
try:
|
||||
ET.fromstring(xml_str)
|
||||
valid_xml.append(xml_str)
|
||||
except ET.ParseError:
|
||||
continue
|
||||
|
||||
# Return state_snapshot if found, otherwise original text
|
||||
if valid_xml and any("<state_snapshot>" in xml for xml in valid_xml):
|
||||
return next(xml for xml in valid_xml if "<state_snapshot>" in xml)
|
||||
elif valid_xml:
|
||||
return valid_xml[0]
|
||||
else:
|
||||
return text
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract XML fragments: {e}. Returning original text.")
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def _format_messages_for_compression(messages: List[Message]) -> str:
|
||||
"""Format messages into a readable text for compression.
|
||||
|
||||
Args:
|
||||
messages: List of messages to format
|
||||
|
||||
Returns:
|
||||
Formatted string representation of messages
|
||||
"""
|
||||
lines = []
|
||||
for i, msg in enumerate(messages, 1):
|
||||
role_name = msg.role.value if hasattr(msg.role, "value") else str(msg.role)
|
||||
lines.append(f"[Message {i} - {role_name}]")
|
||||
lines.append(msg.content)
|
||||
lines.append("") # Empty line between messages
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
async def _compress_messages_with_llm(self, messages_to_compress: List[Message]) -> str:
|
||||
"""Use LLM to compress messages into a concise summary.
|
||||
|
||||
Args:
|
||||
messages_to_compress: List of messages to compress
|
||||
|
||||
Returns:
|
||||
Compressed summary text
|
||||
"""
|
||||
# Format messages for the prompt
|
||||
formatted_messages = self._format_messages_for_compression(messages_to_compress)
|
||||
|
||||
# Create prompt for compression
|
||||
prompt = self.prompt_format(
|
||||
prompt_name="compress_context_prompt",
|
||||
messages_content=formatted_messages,
|
||||
)
|
||||
|
||||
def parse_compressed_result(message: Message) -> str:
|
||||
"""Parse LLM response to extract compressed content.
|
||||
|
||||
Args:
|
||||
message: LLM response message
|
||||
|
||||
Returns:
|
||||
Compressed content string
|
||||
"""
|
||||
content = message.content.strip()
|
||||
# Try to extract content from txt code block
|
||||
compressed = extract_content(content, "txt")
|
||||
|
||||
# If no code block found, use the raw content
|
||||
if not compressed:
|
||||
compressed = content
|
||||
|
||||
logger.info(
|
||||
f"Compressed {len(messages_to_compress)} messages into "
|
||||
f"{len(compressed)} characters (reduction: "
|
||||
f"{len(formatted_messages)} -> {len(compressed)})",
|
||||
)
|
||||
return compressed
|
||||
|
||||
# Call LLM to generate compressed summary
|
||||
result = await self.llm.achat(
|
||||
messages=[Message(role=Role.USER, content=prompt)],
|
||||
callback_fn=parse_compressed_result,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
async def _compress_with_micro_groups(
|
||||
self,
|
||||
messages_to_compress: List[Message],
|
||||
system_messages: List[Message],
|
||||
recent_messages: List[Message],
|
||||
) -> List[Message]:
|
||||
"""Compress messages by splitting into groups and compressing each separately.
|
||||
|
||||
Args:
|
||||
messages_to_compress: Messages to be compressed
|
||||
system_messages: System messages to preserve
|
||||
recent_messages: Recent messages to keep uncompressed
|
||||
|
||||
Returns:
|
||||
List of new messages after compression
|
||||
"""
|
||||
# Split messages into groups based on micro threshold
|
||||
message_groups = self._split_messages_by_token_threshold(
|
||||
messages_to_compress,
|
||||
self.micro_summary_token_threshold,
|
||||
)
|
||||
|
||||
# Compress each group separately
|
||||
compressed_messages = []
|
||||
total_original_tokens = 0
|
||||
total_compressed_tokens = 0
|
||||
|
||||
for group_idx, group in enumerate(message_groups, 1):
|
||||
# Calculate original token count for this group
|
||||
group_original_tokens = self.token_count(group)
|
||||
total_original_tokens += group_original_tokens
|
||||
|
||||
# Save original messages for this group
|
||||
group_file_path = self._save_original_messages(group)
|
||||
|
||||
# Compress this group
|
||||
logger.info(
|
||||
f"Compressing group {group_idx}/{len(message_groups)} "
|
||||
f"({len(group)} messages, {group_original_tokens} tokens)",
|
||||
)
|
||||
group_summary = await self._compress_messages_with_llm(group)
|
||||
group_summary = self._extract_xml_fragments(group_summary)
|
||||
|
||||
# Create compressed message for this group
|
||||
compressed_message = Message(
|
||||
role=Role.SYSTEM,
|
||||
content=(
|
||||
f"[Compressed conversation history - Part {group_idx}/{len(message_groups)}]\n"
|
||||
f"{group_summary}\n\n"
|
||||
f"(Original {len(group)} messages are stored in: {group_file_path})"
|
||||
),
|
||||
)
|
||||
|
||||
# Check if compression actually reduced tokens for this group
|
||||
compressed_tokens = self.token_count([compressed_message])
|
||||
|
||||
if compressed_tokens >= group_original_tokens:
|
||||
logger.warning(
|
||||
f"Group {group_idx} compression did not reduce tokens: "
|
||||
f"{group_original_tokens} -> {compressed_tokens}. Using original messages.",
|
||||
)
|
||||
return None
|
||||
|
||||
logger.info(
|
||||
f"Group {group_idx} compression successful: "
|
||||
f"{group_original_tokens} -> {compressed_tokens} tokens "
|
||||
f"(reduction: {group_original_tokens - compressed_tokens} tokens, "
|
||||
f"{100 * (1 - compressed_tokens / group_original_tokens):.1f}%)",
|
||||
)
|
||||
compressed_messages.append(compressed_message)
|
||||
total_compressed_tokens += compressed_tokens
|
||||
|
||||
# Construct new message list: system messages + all compressed messages + recent messages
|
||||
new_messages = system_messages + compressed_messages + recent_messages
|
||||
|
||||
logger.info(
|
||||
f"Context compression completed using micro-compression: "
|
||||
f"{len(messages_to_compress) + len(system_messages) + len(recent_messages)} messages -> "
|
||||
f"{len(new_messages)} messages ({len(message_groups)} compressed groups), "
|
||||
f"total tokens: {total_original_tokens} -> {total_compressed_tokens}",
|
||||
)
|
||||
|
||||
return new_messages
|
||||
|
||||
async def _compress_all_together(
|
||||
self,
|
||||
messages_to_compress: List[Message],
|
||||
system_messages: List[Message],
|
||||
recent_messages: List[Message],
|
||||
) -> List[Message]:
|
||||
"""Compress all messages together into a single summary.
|
||||
|
||||
Args:
|
||||
messages_to_compress: Messages to be compressed
|
||||
system_messages: System messages to preserve
|
||||
recent_messages: Recent messages to keep uncompressed
|
||||
|
||||
Returns:
|
||||
List of new messages after compression
|
||||
"""
|
||||
# Calculate original token count
|
||||
original_tokens = self.token_count(messages_to_compress)
|
||||
|
||||
# Save original messages to file for traceability
|
||||
original_file_path = self._save_original_messages(messages_to_compress)
|
||||
|
||||
# Use LLM to compress messages
|
||||
logger.info(
|
||||
f"Starting LLM compression of {len(messages_to_compress)} messages "
|
||||
f"({original_tokens} tokens), keeping {len(recent_messages)} recent messages",
|
||||
)
|
||||
compressed_summary = await self._compress_messages_with_llm(messages_to_compress)
|
||||
compressed_summary = self._extract_xml_fragments(compressed_summary)
|
||||
|
||||
# Create a new system message with the compressed content and file reference
|
||||
compressed_message = Message(
|
||||
role=Role.SYSTEM,
|
||||
content=(
|
||||
f"[Compressed conversation history]\n"
|
||||
f"{compressed_summary}\n\n"
|
||||
f"(Original {len(messages_to_compress)} messages are stored in: {original_file_path})"
|
||||
),
|
||||
)
|
||||
|
||||
# Check if compression actually reduced tokens
|
||||
compressed_tokens = self.token_count([compressed_message])
|
||||
|
||||
if compressed_tokens >= original_tokens:
|
||||
logger.warning(
|
||||
f"Compression did not reduce tokens: {original_tokens} -> {compressed_tokens}. "
|
||||
f"Returning original messages.",
|
||||
)
|
||||
return None
|
||||
|
||||
logger.info(
|
||||
f"Compression successful: {original_tokens} -> {compressed_tokens} tokens "
|
||||
f"(reduction: {original_tokens - compressed_tokens} tokens, "
|
||||
f"{100 * (1 - compressed_tokens / original_tokens):.1f}%)",
|
||||
)
|
||||
|
||||
# Construct new message list: system messages + compressed message + recent messages
|
||||
new_messages = system_messages + [compressed_message] + recent_messages
|
||||
|
||||
logger.info(
|
||||
f"Context compression completed: "
|
||||
f"{len(messages_to_compress) + len(system_messages) + len(recent_messages)} messages -> "
|
||||
f"{len(new_messages)} messages",
|
||||
)
|
||||
|
||||
return new_messages
|
||||
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Execute the context compression operation.
|
||||
|
||||
The operation:
|
||||
1. Splits messages into system messages, messages to compress, and recent messages
|
||||
2. Calculates token count of messages to compress
|
||||
3. If below threshold, returns messages unchanged
|
||||
4. Otherwise, uses LLM to compress older messages by:
|
||||
- Saving original messages to file
|
||||
- Generating a concise summary of older messages
|
||||
- Replacing older messages with a single summary message
|
||||
"""
|
||||
# Convert context messages to Message objects
|
||||
messages = [Message(**x) for x in self.context.messages]
|
||||
|
||||
# Check if we have enough messages to compress
|
||||
if len(messages) <= self.keep_recent:
|
||||
self.context.response.answer = self.context.messages
|
||||
logger.info(
|
||||
f"Message count ({len(messages)}) is less than or "
|
||||
f"equal to keep_recent ({self.keep_recent}), no compression needed",
|
||||
)
|
||||
return
|
||||
|
||||
# Split messages into those to compress and those to keep
|
||||
messages_to_compress = messages[: -self.keep_recent]
|
||||
recent_messages = messages[-self.keep_recent :]
|
||||
|
||||
# Always filter out system messages (system messages are never compressed)
|
||||
system_messages = [m for m in messages_to_compress if m.role is Role.SYSTEM]
|
||||
messages_to_compress = [m for m in messages_to_compress if m.role is not Role.SYSTEM]
|
||||
logger.info(
|
||||
f"Excluding {len(system_messages)} system messages from compression, "
|
||||
f"{len(messages_to_compress)} messages remaining for compression check",
|
||||
)
|
||||
|
||||
# If nothing to compress after filtering, return original messages
|
||||
if not messages_to_compress:
|
||||
self.context.response.answer = self.context.messages
|
||||
logger.info("No messages to compress after filtering, returning original messages")
|
||||
return
|
||||
|
||||
# Calculate token count of messages to compress (only the content that will be compressed)
|
||||
compress_token_cnt: int = self.token_count(messages_to_compress)
|
||||
logger.info(
|
||||
f"Context compression check: messages_to_compress token count={compress_token_cnt}, "
|
||||
f"threshold={self.all_token_threshold}",
|
||||
)
|
||||
|
||||
# If token count is within threshold, no compression needed
|
||||
if compress_token_cnt <= self.all_token_threshold:
|
||||
self.context.response.answer = self.context.messages
|
||||
logger.info(
|
||||
f"Messages to compress token count ({compress_token_cnt}) is within threshold "
|
||||
f"({self.all_token_threshold}), no compression needed",
|
||||
)
|
||||
return
|
||||
|
||||
# Determine whether to use micro-compression (split into groups) or compress all together
|
||||
if self.micro_summary_token_threshold is not None and self.micro_summary_token_threshold > 0:
|
||||
new_messages = await self._compress_with_micro_groups(
|
||||
messages_to_compress,
|
||||
system_messages,
|
||||
recent_messages,
|
||||
)
|
||||
else:
|
||||
new_messages = await self._compress_all_together(
|
||||
messages_to_compress,
|
||||
system_messages,
|
||||
recent_messages,
|
||||
)
|
||||
|
||||
# If compression failed (returned None), use original messages
|
||||
if new_messages is None:
|
||||
self.context.response.answer = self.context.messages
|
||||
return
|
||||
|
||||
# Return the compressed messages as JSON
|
||||
self.context.response.answer = json.dumps([x.model_dump() for x in new_messages], ensure_ascii=False, indent=2)
|
||||
87
reme_ai/context/offload/context_compress_prompt.yaml
Normal file
87
reme_ai/context/offload/context_compress_prompt.yaml
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
compress_context_prompt_zh: |
|
||||
你是将扮演一个功能为“将内部聊天记录总结为给定结构”的组件。
|
||||
|
||||
当会话历史记录变得太大时,将调用您将整个历史记录提取为简洁的结构化XML快照。这个快照非常重要,因为它会成为Agent对过去内容**唯一**的记忆。后续的对话将Agent仅基于此snapshot恢复其工作。所有重要的细节、计划、错误和用户指令都必须保留。
|
||||
|
||||
首先,你在个人的<scratchpad>思考整个历史内容。检查用户的总体目标、代理的操作、工具输出、文件修改以及任何未解决的问题。找出对未来行动至关重要的每一条信息。
|
||||
|
||||
在你推理完成后,生成最终的<state_snapshot> XML对象。信息要非常密集。不要省略任何不重要的对话填充。
|
||||
|
||||
结构必须如下:
|
||||
|
||||
<state_snapshot>
|
||||
<overall_goal>
|
||||
<!-- 用一个简洁的句子描述用户的高层次目标。 -->
|
||||
<!-- 示例:‘重构身份验证服务以使用新的JWT库。\’ -->
|
||||
</overall_goal>
|
||||
|
||||
<key_knowledge>
|
||||
<!-- 基于对话历史和与用户的交互,必须记住的重要事实、惯例和约束。使用列表项。 -->
|
||||
<!-- Example:
|
||||
- Build Command: `npm run build`
|
||||
- Testing: Tests are run with `npm test`. Test files must end in `.test.ts`.
|
||||
- API Endpoint: The primary API endpoint is `https://api.example.com/v2`.
|
||||
-->
|
||||
</key_knowledge>
|
||||
|
||||
<recent_actions>
|
||||
<!-- 最近几个重要行为动作及其结果的总结。重点是事实。 -->
|
||||
<!-- Example:
|
||||
- Ran `grep 'old_function'` which returned 3 results in 2 files.
|
||||
- Ran `npm run test`, which failed due to a snapshot mismatch in `UserProfile.test.ts`.
|
||||
- Ran `ls -F static/` and discovered image assets are stored as `.webp`.
|
||||
-->
|
||||
</recent_actions>
|
||||
|
||||
</state_snapshot>
|
||||
|
||||
下面是需要压缩的对话历史记录:
|
||||
|
||||
’‘’
|
||||
{messages_content}
|
||||
’‘’
|
||||
|
||||
首先,您将在一个私有的<scratchpad>中考虑整个历史。审查用户的总体目标。请你一定要使用中文来回答和压缩。
|
||||
|
||||
|
||||
compress_context_prompt: |
|
||||
You are the component that summarizes internal chat history into a given structure.
|
||||
|
||||
When the conversation history grows too large, you will be invoked to distill the entire history into a concise, structured XML snapshot. This snapshot is CRITICAL, as it will become the agent's *only* memory of the past. The agent will resume its work based solely on this snapshot. All crucial details, plans, errors, and user directives MUST be preserved.
|
||||
|
||||
First, you will think through the entire history in a private <scratchpad>. Review the user's overall goal, the agent's actions, tool outputs, file modifications, and any unresolved questions. Identify every piece of information that is essential for future actions.
|
||||
|
||||
After your reasoning is complete, generate the final <state_snapshot> XML object. Be incredibly dense with information. Omit any irrelevant conversational filler.
|
||||
|
||||
The structure MUST be as follows:
|
||||
|
||||
<state_snapshot>
|
||||
<overall_goal>
|
||||
<!-- A single, concise sentence describing the user's high-level objective. -->
|
||||
<!-- Example: "Refactor the authentication service to use a new JWT library." -->
|
||||
</overall_goal>
|
||||
|
||||
<key_knowledge>
|
||||
<!-- Crucial facts, conventions, and constraints the agent must remember based on the conversation history and interaction with the user. Use bullet points. -->
|
||||
<!-- Example:
|
||||
- Build Command: `npm run build`
|
||||
- Testing: Tests are run with `npm test`. Test files must end in `.test.ts`.
|
||||
- API Endpoint: The primary API endpoint is `https://api.example.com/v2`.
|
||||
-->
|
||||
</key_knowledge>
|
||||
|
||||
<recent_actions>
|
||||
<!-- A summary of the last few significant agent actions and their outcomes. Focus on facts. -->
|
||||
<!-- Example:
|
||||
- Ran `grep 'old_function'` which returned 3 results in 2 files.
|
||||
- Ran `npm run test`, which failed due to a snapshot mismatch in `UserProfile.test.ts`.
|
||||
- Ran `ls -F static/` and discovered image assets are stored as `.webp`.
|
||||
-->
|
||||
</recent_actions>
|
||||
</state_snapshot>
|
||||
|
||||
Here is the conversation history that needs to be compressed:
|
||||
’’’
|
||||
{messages_content}
|
||||
’’’
|
||||
First, you will think through the entire history in a private <scratchpad>. Review the user's overall goal, the agent'
|
||||
285
test/test_context_compress_op.py
Normal file
285
test/test_context_compress_op.py
Normal file
|
|
@ -0,0 +1,285 @@
|
|||
"""
|
||||
Test script for ContextCompressOp.
|
||||
|
||||
This script demonstrates how to use the context compression operation to reduce
|
||||
token usage in conversation histories using language models.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.context.offload.context_compress_op import ContextCompressOp
|
||||
from reme_ai.main import ReMeApp
|
||||
|
||||
|
||||
async def main():
|
||||
"""Main function to test ContextCompressOp."""
|
||||
|
||||
async with ReMeApp():
|
||||
logger.info("=" * 80)
|
||||
logger.info("Testing ContextCompressOp - LLM-based Context Compression")
|
||||
logger.info("=" * 80)
|
||||
|
||||
# Create a mock conversation with multiple messages
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a helpful AI assistant specialized in software development.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "I need help building a REST API in Python. I want to use FastAPI.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Great choice! FastAPI is an excellent framework for building REST APIs. "
|
||||
"It's fast, modern, and has automatic API documentation. To get started, you'll need "
|
||||
"to install FastAPI and uvicorn. Would you like me to guide you through setting up "
|
||||
"your first endpoint?",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Yes please. I want to create a user management API with CRUD operations.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Perfect! For a user management API, I recommend this structure:\n"
|
||||
"1. Define a User model using Pydantic\n"
|
||||
"2. Create POST /users endpoint for creating users\n"
|
||||
"3. Create GET /users and GET /users/{id} for reading\n"
|
||||
"4. Create PUT /users/{id} for updates\n"
|
||||
"5. Create DELETE /users/{id} for deletion\n"
|
||||
"We'll also need a database. Would you prefer SQLite, PostgreSQL, or MongoDB?",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Let's use PostgreSQL. Also, I need JWT authentication.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Excellent. PostgreSQL is a robust choice. For JWT authentication, we'll use "
|
||||
"python-jose library. Here's what we'll implement:\n"
|
||||
"1. User registration endpoint\n"
|
||||
"2. Login endpoint that returns JWT token\n"
|
||||
"3. Protected endpoints that require valid JWT\n"
|
||||
"4. Password hashing using bcrypt\n"
|
||||
"Let me show you the code for the User model first.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Before we proceed, I also need rate limiting and input validation.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Good thinking! For rate limiting, we can use slowapi library which integrates "
|
||||
"well with FastAPI. For input validation, Pydantic (which FastAPI uses) handles most of it, "
|
||||
"but we can add custom validators. I'll also add request validation middleware. "
|
||||
"Let's start implementing all of this step by step.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "I need to build a distributed task queue system in Python that can handle millions of tasks"
|
||||
" per day. It needs to be horizontally scalable and fault-tolerant.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "That's a challenging but exciting project! For a high-throughput distributed task queue, "
|
||||
"I recommend a architecture with:\n\n1. **Message Broker**: Redis or RabbitMQ for task "
|
||||
"distribution\n2. **Task Workers**: Multiple worker processes across multiple machines\n3."
|
||||
" **Result Backend**: Redis or PostgreSQL for storing task results\n4. **Monitoring**:"
|
||||
" Prometheus + Grafana for metrics\n5. **API Layer**: FastAPI for task submission and "
|
||||
"status queries\n\nFor the core library, we can build on top of Celery or create a custo"
|
||||
"m solution. Would you like me to design the system architecture first, or do you have pr"
|
||||
"eferences for specific technologies?",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "I want to build this from scratch without Celery. I need custom features like task"
|
||||
" priorities, retry policies with exponential backoff, and task dependencies. Also,"
|
||||
" I need it to support both synchronous and asynchronous task execution patterns.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Excellent! Building from scratch gives us full control. Let me design a comprehensive "
|
||||
"system architecture:\n\n**Core Components:**\n1. **Task Router**: Distributes tasks based"
|
||||
" on priority queues\n2. **Worker Manager**: Handles worker lifecycle and load balancing\n3."
|
||||
" **Retry Engine**: Implements exponential backoff with jitter\n4. **Dependency Graph**:"
|
||||
" Manages task dependencies using topological sorting\n5. **State Manager**: Tracks task "
|
||||
"states (pending, running, completed, failed)\n6. **Metrics Collector**: Real-time "
|
||||
"performance metrics\n\n**Data Structures:**\n- Priority queues using Redis Sorted Sets\n-"
|
||||
" Task metadata in PostgreSQL with JSONB columns\n- Distributed locks with Redis Redlock"
|
||||
" algorithm\n\n**Features:**\n- At-least-once delivery guarantee\n- Dead letter queue for "
|
||||
"failed tasks\n- Rate limiting per task type\n- Circuit breaker pattern for failing "
|
||||
"services\n\nWould you like me to start with the core task model and priority queue "
|
||||
"implementation?",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "That sounds good. But I also need WebSocket support for real-time task status updates, a "
|
||||
"web dashboard for monitoring, and the ability to dynamically scale workers based on queue"
|
||||
" depth. Plus, I need task deduplication and exactly-once processing semantics.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Perfect! Let me enhance the architecture with these advanced requirements:\n\n**Real-time "
|
||||
"Features:**\n1. **WebSocket Manager**: Using Python's websocket-server with Redis pub/sub "
|
||||
"for multi-instance communication\n2. **Event Stream**: Task state changes published to Re"
|
||||
"dis streams\n3. **Client SDK**: JavaScript/TypeScript library for real-time subscriptio"
|
||||
"ns\n\n**Auto-scaling Architecture:**\n1. **Metrics Aggregator**: Collects queue depth, "
|
||||
"processing latency, error rates\n2. **Scaler Service**: Kubernetes HPA integration or cu"
|
||||
"stom Docker Swarm scaling\n3. **Predictive Scaling**: ML model to predict load based on hi"
|
||||
"storical patterns\n\n**Exactly-once Processing:**\n1. **Idempotency Keys**: Each task has "
|
||||
"UUID-based idempotency key\n2. **Deduplication Store**: Redis with TTL for processed task "
|
||||
"IDs\n3. **Transactional Outbox**: PostgreSQL outbox pattern for reliable event publishing"
|
||||
"\n4. **Two-phase Commit**: For distributed task processing\n\n**Dashboard Features:**\n- "
|
||||
"Real-time queue metrics with D3.js visualizations\n- Task timeline view with dependency g"
|
||||
"raphs\n- Worker health monitoring with automatic restart\n- Alert system for SLA violati"
|
||||
"ons\n\n**Additional Components:**\n- Task scheduler for delayed tasks using Redis keys w"
|
||||
"ith TTL\n- Batch processing support for bulk operations\n- Multi-tenant isolation with na"
|
||||
"mespace support\n\nThis is getting quite comprehensive! Should I start with the core tas"
|
||||
"k model and database schema, or would you prefer to begin with the WebSocket real-time up"
|
||||
"date system?",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "This is getting complex. I also need multi-region support with eventual consistency, the "
|
||||
"ability to pause/resume task processing by task type, and A/B testing capabilities for "
|
||||
"different worker implementations. Also, I need comprehensive audit logging and GDPR comp"
|
||||
"liance features.",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Absolutely! This is evolving into an enterprise-grade system. Let me architect the compl"
|
||||
"ete solution:\n\n**Multi-region Architecture:**\n1. **Global Load Balancer**: GeoDNS with"
|
||||
" health checks\n2. **Cross-region Replication**: Redis Cluster with active-active setup "
|
||||
"using CRDTs\n3. **Conflict Resolution**: Vector clocks for task ordering, last-writer-win"
|
||||
"s for metadata\n4. **Region-aware Routing**: Route tasks to workers in same region when p"
|
||||
"ossible\n5. **Failover Mechanism**: Automatic region failover with 30-second RTO\n\n**Adv"
|
||||
"anced Control Features:**\n1. **Task Type Governance**: \n - Pause/resume via Redis fe"
|
||||
"ature flags with immediate propagation\n - Rate limits per task type with burst capaci"
|
||||
"ty\n - Resource quotas (CPU/memory) per task category\n2. **A/B Testing Framework**:\n"
|
||||
" - Task routing based on consistent hashing of task ID\n - Variant assignment with s"
|
||||
"tickiness\n - Statistical significance tracking for performance metrics\n - Automati"
|
||||
"c winner selection based on success rate and latency\n\n**Compliance & Audit:**\n1. **A"
|
||||
"udit Trail**:\n - Immutable task history in PostgreSQL with row-level security\n -"
|
||||
" Change data capture (CDC) using Debezium\n - Cryptographic signing of audit logs\n "
|
||||
" - 7-year retention policy with automated archival to S3\n2. **GDPR Compliance**:\n "
|
||||
" - Right to be forgotten: Cascade delete with verification\n - Data portability: JSO"
|
||||
"N export of all user tasks\n - Consent management: Task processing consent tracking"
|
||||
"\n - Data anonymization: PII encryption with rotating keys\n\n**Enhanced Monitoring:*"
|
||||
"*\n1. **Distributed Tracing**: OpenTelemetry integration with Jaeger\n2. **Performance "
|
||||
"Profiling**: Py-spy continuous profiling\n3. **Cost Analytics**: Per-task-type resource"
|
||||
" consumption tracking\n4. **SLA Monitoring**: 99.9% task processing latency targets\n\n"
|
||||
"**Security Features:**\n1. **Zero-trust Architecture**: mTLS between all services\n2. **"
|
||||
"Secrets Management**: HashiCorp Vault integration\n3. **Network Segmentation**: VPC is"
|
||||
"olation with private subnets\n4. **Encryption**: AES-256 at rest, TLS 1.3 in transit\n"
|
||||
"\n**Operational Excellence:**\n1. **Blue-green Deployments**: Zero-downtime updates\n2."
|
||||
" **Canary Releases**: Gradual rollout of worker changes\n3. **Chaos Engineering**: Rand"
|
||||
"om worker termination testing\n4. **Capacity Planning**: Automated infrastructure provi"
|
||||
"sioning\n\nThis is becoming a mission-critical system! Would you like me to start with "
|
||||
"the database schema and core task lifecycle management, or should we begin with the mul"
|
||||
"ti-region replication strategy? I can also create a detailed implementation roadmap wi"
|
||||
"th phases if you prefer.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "How do I reverse a string in Python?",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "You can reverse a string in Python using slicing: `reversed_string = original_string[::"
|
||||
"-1]`. The `[::-1]` slice notation means start at the end of the string and end at posit"
|
||||
"ion 0, moving with the step -1.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What's the difference between a list and a tuple?",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Lists are mutable (you can change their content after creation) and use square bracket"
|
||||
"s `[]`. Tuples are immutable and use parentheses `()`. Lists are better for collectio"
|
||||
"ns that need to change, while tuples are faster and can be used as dictionary keys.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "How do I handle exceptions in Python?",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Use try-except blocks: `try:` followed by your risky code, then `except SpecificError"
|
||||
":` to catch particular exceptions. Always catch specific exceptions rather than the g"
|
||||
"eneric `Exception` when possible. You can also use `finally:` for cleanup code that ru"
|
||||
"ns regardless.",
|
||||
},
|
||||
|
||||
]
|
||||
|
||||
# Test 1: Messages below threshold (should not compress)
|
||||
logger.info("\n" + "=" * 60)
|
||||
logger.info("Test 1: Messages below threshold (should skip compression)")
|
||||
logger.info("=" * 60)
|
||||
|
||||
compress_op1 = ContextCompressOp(
|
||||
all_token_threshold=50000, # High threshold, won't trigger
|
||||
keep_recent=3,
|
||||
)
|
||||
|
||||
await compress_op1.async_call(messages=messages)
|
||||
|
||||
result_messages1 = compress_op1.context.response.answer
|
||||
logger.info(f"✓ Result: {len(result_messages1)} messages (unchanged)")
|
||||
|
||||
# Test 2: Messages above threshold (should compress)
|
||||
logger.info("\n" + "=" * 60)
|
||||
logger.info("Test 2: Messages above threshold (should compress)")
|
||||
logger.info("=" * 60)
|
||||
|
||||
compress_op2 = ContextCompressOp(
|
||||
all_token_threshold=2000, # Low threshold, will trigger
|
||||
keep_recent=3, # Keep last 3 messages
|
||||
compress_system_message=False, # Don't compress system messages
|
||||
)
|
||||
|
||||
await compress_op2.async_call(messages=messages)
|
||||
|
||||
result_messages2 = compress_op2.context.response.answer
|
||||
logger.info(f"✓ Result: {len(result_messages2)} messages (compressed)")
|
||||
|
||||
# Display compression results
|
||||
logger.info("\n" + "=" * 60)
|
||||
logger.info("Compression Result Details:")
|
||||
logger.info("=" * 60)
|
||||
logger.info(f"Original messages: {len(messages)}")
|
||||
logger.info(f"Compressed messages: {len(result_messages2)}")
|
||||
|
||||
# Test 3: Messages above threshold (should compress)
|
||||
logger.info("\n" + "=!" * 30)
|
||||
logger.info("Test 3: Messages above micro threshold (should compress)")
|
||||
logger.info("=!" * 30)
|
||||
|
||||
compress_op2 = ContextCompressOp(
|
||||
all_token_threshold=2000, # Low threshold, will trigger
|
||||
keep_recent=2, # Keep last 3 messages
|
||||
compress_system_message=False, # Don't compress system messages
|
||||
micro_summary_token_threshold=1500,
|
||||
language="zh",
|
||||
)
|
||||
|
||||
await compress_op2.async_call(messages=messages)
|
||||
|
||||
result_messages2 = compress_op2.context.response.answer
|
||||
logger.info(f"✓ Result: {len(result_messages2)} messages (compressed)")
|
||||
|
||||
# Display compression results
|
||||
logger.info("\n" + "=" * 60)
|
||||
logger.info("Compression Result Details:")
|
||||
logger.info("=" * 60)
|
||||
logger.info(f"Original messages: {len(messages)}")
|
||||
logger.info(f"Compressed messages: {len(result_messages2)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
Loading…
Add table
Reference in a new issue