mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
feat(memory): replace memory formatter with AsMsgHandler for enhanced message processing
This commit is contained in:
parent
3347506e22
commit
3dc3c4bf52
11 changed files with 2461 additions and 355 deletions
80
reme/core/schema/as_msg_stat.py
Normal file
80
reme/core/schema/as_msg_stat.py
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH = 100
|
||||
_DEFAULT_MAX_FORMATTER_TEXT_LENGTH = 2000
|
||||
|
||||
# Unique marker for truncated text
|
||||
TRUNCATION_MARKER_START = "<<<TRUNCATED>>>"
|
||||
TRUNCATION_MARKER_END = "<<<END_TRUNCATED>>>"
|
||||
|
||||
|
||||
def _truncate_text(text: str, max_length: int) -> str:
|
||||
"""Truncate text to max length, keeping head and tail portions."""
|
||||
text = str(text) if text else ""
|
||||
if not text or len(text) <= max_length:
|
||||
return text
|
||||
half_length = max_length // 2
|
||||
truncated_chars = len(text) - max_length
|
||||
return (
|
||||
f"{text[:half_length]}\n\n{TRUNCATION_MARKER_START} "
|
||||
f"({truncated_chars} characters omitted) "
|
||||
f"{TRUNCATION_MARKER_END}\n\n{text[-half_length:]}"
|
||||
)
|
||||
|
||||
|
||||
class AsBlockStat(BaseModel):
|
||||
block_type: str = Field(default=...)
|
||||
text: str = Field(default="", description="Text content of the block")
|
||||
token_count: int = Field(default=0, description="Token count of the block, including base64 data")
|
||||
|
||||
# For tool_use and tool_result blocks
|
||||
tool_name: str = Field(default="", description="Tool name for tool_use/tool_result blocks")
|
||||
tool_input: str = Field(default="", description="Tool input arguments for tool_use blocks")
|
||||
tool_output: str = Field(default="", description="Tool output for tool_result blocks")
|
||||
|
||||
# For media blocks
|
||||
media_url: str = Field(default="", description="URL for image/audio/video blocks")
|
||||
|
||||
@property
|
||||
def preview(self) -> str:
|
||||
return self.format(_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH)
|
||||
|
||||
def format(self, max_length: int = _DEFAULT_MAX_FORMATTER_TEXT_LENGTH, include_thinking: bool = True) -> str:
|
||||
"""Format block content to string representation."""
|
||||
if self.block_type == "text":
|
||||
return _truncate_text(self.text, max_length) if self.text else ""
|
||||
if self.block_type == "thinking":
|
||||
if include_thinking and self.text:
|
||||
return f"<thinking>\n{_truncate_text(self.text, max_length)}\n</thinking>"
|
||||
return ""
|
||||
if self.block_type in ("image", "audio", "video"):
|
||||
return f"[{self.block_type}] {self.media_url}" if self.media_url else f"[{self.block_type}]"
|
||||
if self.block_type == "tool_use":
|
||||
return f" - tool_call={self.tool_name} params={_truncate_text(self.tool_input, max_length)}"
|
||||
if self.block_type == "tool_result":
|
||||
output = _truncate_text(self.tool_output, max_length)
|
||||
return f" - tool_result={self.tool_name} output={output}" if output else ""
|
||||
return ""
|
||||
|
||||
|
||||
class AsMsgStat(BaseModel):
|
||||
name: str = Field(default=...)
|
||||
role: str = Field(default="")
|
||||
content: list[AsBlockStat] = Field(default_factory=list)
|
||||
timestamp: str = Field(default="")
|
||||
metadata: dict = Field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def total_tokens(self) -> int:
|
||||
return sum(block.token_count for block in self.content)
|
||||
|
||||
@property
|
||||
def preview(self) -> str:
|
||||
return self.format(_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH)
|
||||
|
||||
def format(self, max_length: int = _DEFAULT_MAX_FORMATTER_TEXT_LENGTH, include_thinking: bool = True) -> str:
|
||||
"""Format message to string representation."""
|
||||
time_str = f"[{self.timestamp}] " if self.timestamp else ""
|
||||
header = f"{time_str}{self.name or self.role}:"
|
||||
blocks = [block.format(max_length, include_thinking) for block in self.content]
|
||||
return "\n".join([header] + [b for b in blocks if b])
|
||||
|
|
@ -4,8 +4,9 @@ This module provides memory management components for CoPaw (Cooperative Paw) ag
|
|||
including memory formatting, compaction, summarization, and file I/O operations.
|
||||
|
||||
Components:
|
||||
- MemoryFormatter: Converts message lists to formatted strings with token limiting
|
||||
- ReMeInMemoryMemory: Extended InMemoryMemory with bugfixes and summary support
|
||||
- ReMeOpenAIChatFormatter: Converts message lists to formatted strings with token limiting
|
||||
- AsMsgHandler: Handles AgentScope message statistics, formatting, and context checking
|
||||
- Summarizer: Generates memory summaries using LLM
|
||||
- Compactor: Compacts memory content to reduce token usage
|
||||
- ToolResultCompactor: Truncates large tool results and saves full content to files
|
||||
|
|
@ -13,21 +14,21 @@ Components:
|
|||
"""
|
||||
|
||||
from . import utils
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
from .compactor import Compactor
|
||||
from .file_io import FileIO
|
||||
from .memory_formatter import MemoryFormatter
|
||||
from .reme_chat_formatter import ReMeChatFormatter
|
||||
from .reme_chat_formatter import ReMeOpenAIChatFormatter
|
||||
from .reme_in_memory_memory import ReMeInMemoryMemory
|
||||
from .summarizer import Summarizer
|
||||
from .tool_result_compactor import ToolResultCompactor
|
||||
|
||||
__all__ = [
|
||||
"MemoryFormatter",
|
||||
"AsMsgHandler",
|
||||
"ReMeInMemoryMemory",
|
||||
"Summarizer",
|
||||
"Compactor",
|
||||
"ToolResultCompactor",
|
||||
"FileIO",
|
||||
"utils",
|
||||
"ReMeChatFormatter",
|
||||
"ReMeOpenAIChatFormatter",
|
||||
]
|
||||
|
|
|
|||
351
reme/memory/file_based/as_msg_handler.py
Normal file
351
reme/memory/file_based/as_msg_handler.py
Normal file
|
|
@ -0,0 +1,351 @@
|
|||
import json
|
||||
import logging
|
||||
|
||||
from agentscope.message import Msg
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from ...core.schema.as_msg_stat import AsMsgStat, AsBlockStat
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AsMsgHandler:
|
||||
|
||||
def __init__(self, token_counter: HuggingFaceTokenCounter):
|
||||
self._token_counter = token_counter
|
||||
|
||||
def count_str_token(self, text: str) -> int:
|
||||
"""Count tokens in a string.
|
||||
|
||||
Args:
|
||||
text: The text to count tokens for.
|
||||
|
||||
Returns:
|
||||
The number of tokens in the text.
|
||||
"""
|
||||
if not text:
|
||||
return 0
|
||||
|
||||
try:
|
||||
token_ids = self._token_counter.tokenizer.encode(text)
|
||||
token_count = len(token_ids)
|
||||
return token_count
|
||||
|
||||
except Exception as e:
|
||||
estimated_tokens = len(text.encode("utf-8")) // 4
|
||||
logger.warning(f"Failed to count string tokens: {text}, using estimated_tokens={estimated_tokens}")
|
||||
return estimated_tokens
|
||||
|
||||
@staticmethod
|
||||
def _format_tool_result_output(output: str | list[dict]) -> str:
|
||||
"""Convert tool result output to string.
|
||||
|
||||
Args:
|
||||
output: Tool result output, either string or list of content blocks.
|
||||
|
||||
Returns:
|
||||
Formatted string representation of the tool result.
|
||||
"""
|
||||
if isinstance(output, str):
|
||||
return output
|
||||
|
||||
textual_parts = []
|
||||
|
||||
for block in output:
|
||||
try:
|
||||
if not isinstance(block, dict) or "type" not in block:
|
||||
logger.warning(
|
||||
"Invalid block: %s, expected a dict with 'type' key, skipped.",
|
||||
block,
|
||||
)
|
||||
continue
|
||||
|
||||
block_type = block["type"]
|
||||
|
||||
if block_type == "text":
|
||||
textual_parts.append(block.get("text", ""))
|
||||
|
||||
elif block_type in ["image", "audio", "video"]:
|
||||
source = block.get("source", {})
|
||||
url = source.get("url", "")
|
||||
if url:
|
||||
textual_parts.append(f"[{block_type}] {url}")
|
||||
else:
|
||||
textual_parts.append(f"[{block_type}]")
|
||||
|
||||
elif block_type == "file":
|
||||
file_path = block.get("path", "") or block.get("url", "")
|
||||
file_name = block.get("name", file_path)
|
||||
textual_parts.append(f"[file] {file_name}: {file_path}")
|
||||
|
||||
else:
|
||||
logger.warning(
|
||||
"Unsupported block type '%s' in tool result, skipped.",
|
||||
block_type,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to process block %s: %s, skipped.",
|
||||
block,
|
||||
e,
|
||||
)
|
||||
|
||||
if not textual_parts:
|
||||
return ""
|
||||
if len(textual_parts) == 1:
|
||||
return textual_parts[0]
|
||||
return "\n".join(f"- {part}" for part in textual_parts)
|
||||
|
||||
def stat_message(self, message: Msg) -> AsMsgStat:
|
||||
"""Analyze a message and generate block statistics."""
|
||||
blocks = []
|
||||
|
||||
for block in message.get_content_blocks():
|
||||
block_type = block.get("type", "unknown")
|
||||
|
||||
if block_type == "text":
|
||||
text = block.get("text", "")
|
||||
token_count = self.count_str_token(text)
|
||||
blocks.append(AsBlockStat(
|
||||
block_type=block_type,
|
||||
text=text,
|
||||
token_count=token_count,
|
||||
))
|
||||
|
||||
elif block_type == "thinking":
|
||||
thinking = block.get("thinking", "")
|
||||
token_count = self.count_str_token(thinking)
|
||||
blocks.append(AsBlockStat(
|
||||
block_type=block_type,
|
||||
text=thinking,
|
||||
token_count=token_count,
|
||||
))
|
||||
|
||||
elif block_type in ("image", "audio", "video"):
|
||||
source = block.get("source", {})
|
||||
url = source.get("url", "")
|
||||
# For media, estimate fixed token cost or count URL
|
||||
if source.get("type") == "base64":
|
||||
data = source.get("data", "")
|
||||
token_count = len(data) // 4 if data else 10
|
||||
else:
|
||||
token_count = self.count_str_token(url) if url else 10
|
||||
blocks.append(AsBlockStat(
|
||||
block_type=block_type,
|
||||
text="",
|
||||
token_count=token_count,
|
||||
media_url=url,
|
||||
))
|
||||
|
||||
elif block_type == "tool_use":
|
||||
tool_name = block.get("name", "")
|
||||
tool_input = block.get("input", {})
|
||||
try:
|
||||
input_str = json.dumps(tool_input, ensure_ascii=False)
|
||||
except (TypeError, ValueError):
|
||||
input_str = str(tool_input)
|
||||
token_count = self.count_str_token(tool_name + input_str)
|
||||
blocks.append(AsBlockStat(
|
||||
block_type=block_type,
|
||||
text="",
|
||||
token_count=token_count,
|
||||
tool_name=tool_name,
|
||||
tool_input=input_str,
|
||||
))
|
||||
|
||||
elif block_type == "tool_result":
|
||||
tool_name = block.get("name", "")
|
||||
output = block.get("output", "")
|
||||
formatted_output = self._format_tool_result_output(output)
|
||||
token_count = self.count_str_token(formatted_output)
|
||||
blocks.append(AsBlockStat(
|
||||
block_type=block_type,
|
||||
text="",
|
||||
token_count=token_count,
|
||||
tool_name=tool_name,
|
||||
tool_output=formatted_output,
|
||||
))
|
||||
|
||||
else:
|
||||
logger.warning("Unsupported block type %s, skipped.", block_type)
|
||||
|
||||
return AsMsgStat(
|
||||
name=message.name or message.role,
|
||||
role=message.role,
|
||||
content=blocks,
|
||||
timestamp=message.timestamp or "",
|
||||
metadata=message.metadata or {},
|
||||
)
|
||||
|
||||
def format_msgs_to_str(
|
||||
self,
|
||||
messages: list[Msg],
|
||||
memory_compact_threshold: int,
|
||||
include_thinking: bool = False,
|
||||
) -> str:
|
||||
"""Format list of messages to a single formatted string.
|
||||
|
||||
Messages are processed in reverse order (newest first) and older
|
||||
messages are skipped when token count exceeds memory_compact_threshold.
|
||||
|
||||
Args:
|
||||
messages: List of Msg objects to format.
|
||||
memory_compact_threshold: Maximum token count before skipping older messages.
|
||||
include_thinking: Whether to include thinking blocks in output.
|
||||
"""
|
||||
if not messages:
|
||||
return ""
|
||||
|
||||
formatted_parts: list[str] = []
|
||||
total_token_count = 0
|
||||
|
||||
for i in range(len(messages) - 1, -1, -1):
|
||||
stat = self.stat_message(messages[i])
|
||||
|
||||
if total_token_count + stat.total_tokens > memory_compact_threshold:
|
||||
logger.info(
|
||||
"Skipping older messages: adding %d tokens would exceed threshold %d (current: %d)",
|
||||
stat.total_tokens,
|
||||
memory_compact_threshold,
|
||||
total_token_count,
|
||||
)
|
||||
break
|
||||
|
||||
formatted_parts.append(stat.format(include_thinking=include_thinking))
|
||||
total_token_count += stat.total_tokens
|
||||
|
||||
formatted_parts.reverse()
|
||||
return "\n\n".join(formatted_parts)
|
||||
|
||||
def context_check(
|
||||
self,
|
||||
messages: list[Msg],
|
||||
memory_compact_threshold: int,
|
||||
memory_compact_reserve: int,
|
||||
) -> tuple[list[Msg], list[Msg]]:
|
||||
"""Check if context exceeds threshold and split messages accordingly.
|
||||
|
||||
This method checks if the total token count of messages exceeds the
|
||||
memory_compact_threshold. If not, returns empty list and original messages.
|
||||
If exceeded, uses memory_compact_reserve as the limit to keep messages
|
||||
from the end, ensuring tool_use and tool_result blocks are properly paired.
|
||||
|
||||
Args:
|
||||
messages: List of Msg objects to check.
|
||||
memory_compact_threshold: Maximum token count threshold to trigger compaction.
|
||||
memory_compact_reserve: Token limit for messages to keep after compaction.
|
||||
|
||||
Returns:
|
||||
A tuple of (messages_to_compact, messages_to_keep):
|
||||
- messages_to_compact: Older messages that need to be compacted
|
||||
- messages_to_keep: Recent messages within the reserve limit
|
||||
"""
|
||||
if not messages:
|
||||
return [], []
|
||||
|
||||
# Calculate total tokens and stats for all messages
|
||||
msg_stats: list[tuple[Msg, AsMsgStat]] = []
|
||||
total_tokens = 0
|
||||
for msg in messages:
|
||||
stat = self.stat_message(msg)
|
||||
msg_stats.append((msg, stat))
|
||||
total_tokens += stat.total_tokens
|
||||
|
||||
# If total tokens don't exceed threshold, no compaction needed
|
||||
if total_tokens <= memory_compact_threshold:
|
||||
return [], messages
|
||||
|
||||
# Collect all tool_use ids and their message indices
|
||||
# tool_use_id -> message index
|
||||
tool_use_locations: dict[str, int] = {}
|
||||
# tool_result_id -> message index
|
||||
tool_result_locations: dict[str, int] = {}
|
||||
|
||||
for idx, (msg, _) in enumerate(msg_stats):
|
||||
for block in msg.get_content_blocks("tool_use"):
|
||||
tool_id = block.get("id", "")
|
||||
if tool_id:
|
||||
tool_use_locations[tool_id] = idx
|
||||
|
||||
for block in msg.get_content_blocks("tool_result"):
|
||||
tool_id = block.get("id", "")
|
||||
if tool_id:
|
||||
tool_result_locations[tool_id] = idx
|
||||
|
||||
# Iterate from the end, accumulating messages to keep within reserve limit
|
||||
keep_indices: set[int] = set()
|
||||
accumulated_tokens = 0
|
||||
|
||||
for i in range(len(msg_stats) - 1, -1, -1):
|
||||
msg, stat = msg_stats[i]
|
||||
|
||||
# Check if adding this message would exceed reserve limit
|
||||
if accumulated_tokens + stat.total_tokens > memory_compact_reserve:
|
||||
logger.info(
|
||||
"Context check: adding message %d with %d tokens would exceed reserve %d (current: %d)",
|
||||
i,
|
||||
stat.total_tokens,
|
||||
memory_compact_reserve,
|
||||
accumulated_tokens,
|
||||
)
|
||||
break
|
||||
|
||||
# Check tool_result dependencies - if this message has tool_result,
|
||||
# we need to ensure the corresponding tool_use is also included
|
||||
tool_result_ids = [
|
||||
block.get("id", "")
|
||||
for block in msg.get_content_blocks("tool_result")
|
||||
if block.get("id", "")
|
||||
]
|
||||
|
||||
# Calculate extra tokens needed for dependent tool_use messages
|
||||
extra_tokens = 0
|
||||
dependent_indices: set[int] = set()
|
||||
|
||||
for tool_id in tool_result_ids:
|
||||
if tool_id in tool_use_locations:
|
||||
tool_use_idx = tool_use_locations[tool_id]
|
||||
if tool_use_idx not in keep_indices and tool_use_idx != i:
|
||||
dependent_indices.add(tool_use_idx)
|
||||
_, dep_stat = msg_stats[tool_use_idx]
|
||||
extra_tokens += dep_stat.total_tokens
|
||||
|
||||
# Check if we can fit this message plus its dependencies within reserve
|
||||
if accumulated_tokens + stat.total_tokens + extra_tokens > memory_compact_reserve:
|
||||
logger.info(
|
||||
"Context check: message %d requires %d extra tokens for tool_use dependencies, "
|
||||
"total would exceed reserve %d",
|
||||
i,
|
||||
extra_tokens,
|
||||
memory_compact_reserve,
|
||||
)
|
||||
break
|
||||
|
||||
# Add this message and its dependencies
|
||||
keep_indices.add(i)
|
||||
keep_indices.update(dependent_indices)
|
||||
accumulated_tokens += stat.total_tokens + extra_tokens
|
||||
|
||||
# Build final lists based on keep_indices (preserve original order)
|
||||
messages_to_compact = []
|
||||
messages_to_keep = []
|
||||
|
||||
for idx, (msg, _) in enumerate(msg_stats):
|
||||
if idx in keep_indices:
|
||||
messages_to_keep.append(msg)
|
||||
else:
|
||||
messages_to_compact.append(msg)
|
||||
|
||||
logger.info(
|
||||
"Context check result: %d messages to compact, %d messages to keep, "
|
||||
"total tokens: %d, threshold: %d, reserve: %d, kept tokens: %d",
|
||||
len(messages_to_compact),
|
||||
len(messages_to_keep),
|
||||
total_tokens,
|
||||
memory_compact_threshold,
|
||||
memory_compact_reserve,
|
||||
accumulated_tokens,
|
||||
)
|
||||
|
||||
return messages_to_compact, messages_to_keep
|
||||
|
|
@ -8,7 +8,7 @@ from agentscope.message import Msg
|
|||
from agentscope.model import ChatModelBase
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from .memory_formatter import MemoryFormatter
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
from ...core.op import BaseOp
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
|
@ -30,7 +30,7 @@ class Compactor(BaseOp):
|
|||
|
||||
self.chat_model: ChatModelBase = chat_model
|
||||
self.formatter: FormatterBase = formatter
|
||||
self.as_token_counter: HuggingFaceTokenCounter = token_counter
|
||||
self.msg_handler = AsMsgHandler(token_counter=token_counter)
|
||||
|
||||
async def execute(self):
|
||||
messages: list[Msg] = self.context.get("messages", [])
|
||||
|
|
@ -39,11 +39,10 @@ class Compactor(BaseOp):
|
|||
if not messages:
|
||||
return ""
|
||||
|
||||
formatter = MemoryFormatter(
|
||||
token_counter=self.as_token_counter,
|
||||
history_formatted_str: str = self.msg_handler.format_msgs_to_str(
|
||||
messages=messages,
|
||||
memory_compact_threshold=self.memory_compact_threshold,
|
||||
)
|
||||
history_formatted_str: str = formatter.format(messages)
|
||||
|
||||
if not history_formatted_str:
|
||||
logger.warning(f"No history to compact. messages={messages}")
|
||||
|
|
@ -66,9 +65,8 @@ class Compactor(BaseOp):
|
|||
f"{suffix}"
|
||||
)
|
||||
else:
|
||||
user_message: str = f"<conversation>\n{history_formatted_str}\n</conversation>\n\n" + self.get_prompt(
|
||||
"initial_user_message",
|
||||
)
|
||||
user_message: str = f"<conversation>\n{history_formatted_str}\n</conversation>\n\n" \
|
||||
+ self.get_prompt("initial_user_message")
|
||||
logger.info(f"Compactor sys_prompt={agent.sys_prompt} user_message={user_message}")
|
||||
|
||||
compact_msg: Msg = await agent.reply(
|
||||
|
|
|
|||
|
|
@ -1,249 +0,0 @@
|
|||
"""Memory Formatter for CoPaw agents.
|
||||
|
||||
Provides memory formatting capabilities including:
|
||||
- Converting list of Msg to formatted string
|
||||
- Memory compaction with token threshold
|
||||
- Support for various content block types (text, tool_use, tool_result, etc.)
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
from agentscope.message import Msg
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from .utils import safe_count_str_tokens, truncate_text
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_MAX_FORMATTER_TEXT_LENGTH = 2000
|
||||
|
||||
|
||||
class MemoryFormatter:
|
||||
"""Formatter that converts list of Msg to formatted string.
|
||||
|
||||
Formats messages into human-readable string representation with:
|
||||
- Role and timestamp information
|
||||
- Text content and tool calls
|
||||
- Memory compact threshold to limit total token count
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
memory_compact_threshold: int,
|
||||
):
|
||||
"""Initialize MemoryFormatter.
|
||||
|
||||
Args:
|
||||
token_counter: Token counter for estimating token counts.
|
||||
memory_compact_threshold: Maximum token count before skipping
|
||||
older messages.
|
||||
"""
|
||||
self._token_counter = token_counter
|
||||
self._memory_compact_threshold = memory_compact_threshold
|
||||
self.max_length = int(
|
||||
os.getenv("MAX_FORMATTER_TEXT_LENGTH", str(_DEFAULT_MAX_FORMATTER_TEXT_LENGTH)),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _format_tool_result_output(output: str | list[dict]) -> str:
|
||||
"""Convert tool result output to string.
|
||||
|
||||
Args:
|
||||
output: Tool result output, either string or list of content blocks.
|
||||
|
||||
Returns:
|
||||
Formatted string representation of the tool result.
|
||||
"""
|
||||
if isinstance(output, str):
|
||||
return output
|
||||
|
||||
textual_parts = []
|
||||
|
||||
for block in output:
|
||||
try:
|
||||
if not isinstance(block, dict) or "type" not in block:
|
||||
logger.warning(
|
||||
"Invalid block: %s, expected a dict with 'type' key, skipped.",
|
||||
block,
|
||||
)
|
||||
continue
|
||||
|
||||
block_type = block["type"]
|
||||
|
||||
if block_type == "text":
|
||||
textual_parts.append(block.get("text", ""))
|
||||
|
||||
elif block_type in ["image", "audio", "video"]:
|
||||
source = block.get("source", {})
|
||||
url = source.get("url", "")
|
||||
if url:
|
||||
textual_parts.append(
|
||||
f"[{block_type}] {url}",
|
||||
)
|
||||
else:
|
||||
textual_parts.append(f"[{block_type}]")
|
||||
|
||||
elif block_type == "file":
|
||||
file_path = block.get("path", "") or block.get("url", "")
|
||||
file_name = block.get("name", file_path)
|
||||
textual_parts.append(f"[file] {file_name}: {file_path}")
|
||||
|
||||
else:
|
||||
# Unknown block type: log warning and skip
|
||||
logger.warning(
|
||||
"Unsupported block type '%s' in tool result, skipped.",
|
||||
block_type,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to process block %s: %s, skipped.",
|
||||
block,
|
||||
e,
|
||||
)
|
||||
|
||||
if not textual_parts:
|
||||
return ""
|
||||
if len(textual_parts) == 1:
|
||||
return textual_parts[0]
|
||||
return "\n".join(f"- {part}" for part in textual_parts)
|
||||
|
||||
def _format_single_msg(
|
||||
self,
|
||||
msg: Msg,
|
||||
index: int | None = None,
|
||||
add_time: bool = True,
|
||||
) -> tuple[str, int]:
|
||||
"""Format a single Msg into string representation.
|
||||
|
||||
Similar to Message.format_message style.
|
||||
|
||||
Args:
|
||||
msg: The Msg object to format.
|
||||
index: Optional message index for round numbering.
|
||||
add_time: Whether to include timestamp.
|
||||
|
||||
Returns:
|
||||
Tuple of (formatted_string, token_count).
|
||||
"""
|
||||
lines = []
|
||||
token_count = 0
|
||||
|
||||
# Build header: "round{index} [{timestamp}] {role}:"
|
||||
prefix = f"round{index} " if index is not None else ""
|
||||
time_str = f"[{msg.timestamp}] " if add_time and msg.timestamp else ""
|
||||
role_str = msg.name or msg.role
|
||||
header = f"{prefix}{time_str}{role_str}:"
|
||||
lines.append(header)
|
||||
token_count += safe_count_str_tokens(self._token_counter, header)
|
||||
|
||||
# Process content blocks
|
||||
for block in msg.get_content_blocks():
|
||||
typ = block.get("type")
|
||||
|
||||
if typ == "text":
|
||||
text_content = truncate_text(block.get("text", ""), self.max_length)
|
||||
if text_content:
|
||||
lines.append(text_content)
|
||||
token_count += safe_count_str_tokens(self._token_counter, text_content)
|
||||
|
||||
elif typ == "thinking":
|
||||
# Skip thinking blocks to save tokens
|
||||
pass
|
||||
|
||||
elif typ in ["image", "audio", "video"]:
|
||||
source = block.get("source", {})
|
||||
url = source.get("url", "")
|
||||
if url:
|
||||
lines.append(f"[{typ}] {url}")
|
||||
else:
|
||||
lines.append(f"[{typ}]")
|
||||
# Estimate fixed token cost for media reference
|
||||
token_count += 10
|
||||
|
||||
elif typ == "tool_use":
|
||||
tool_name = block.get("name", "")
|
||||
tool_input = block.get("input", {})
|
||||
try:
|
||||
arguments_str = json.dumps(tool_input, ensure_ascii=False)
|
||||
except (TypeError, ValueError):
|
||||
arguments_str = str(tool_input)
|
||||
truncated_args = truncate_text(arguments_str, self.max_length)
|
||||
tool_line = f" - tool_call={tool_name} params={truncated_args}"
|
||||
lines.append(tool_line)
|
||||
token_count += safe_count_str_tokens(self._token_counter, tool_line)
|
||||
|
||||
elif typ == "tool_result":
|
||||
tool_name = block.get("name", "")
|
||||
output = block.get("output", "")
|
||||
formatted_output = self._format_tool_result_output(output)
|
||||
truncated_output = truncate_text(formatted_output, self.max_length)
|
||||
if truncated_output:
|
||||
result_line = f" - tool_result={tool_name} output={truncated_output}"
|
||||
lines.append(result_line)
|
||||
token_count += safe_count_str_tokens(self._token_counter, result_line)
|
||||
|
||||
else:
|
||||
logger.warning(
|
||||
"Unsupported block type %s in message, skipped.",
|
||||
typ,
|
||||
)
|
||||
|
||||
return "\n".join(lines), token_count
|
||||
|
||||
def format(
|
||||
self,
|
||||
msgs: list[Msg],
|
||||
add_time: bool = True,
|
||||
add_index: bool = True,
|
||||
) -> str:
|
||||
"""Format list of Msg into a single formatted string.
|
||||
|
||||
Messages are processed in reverse order (newest first) and older
|
||||
messages are skipped when token count exceeds memory_compact_threshold.
|
||||
|
||||
Args:
|
||||
msgs: List of Msg objects to format.
|
||||
add_time: Whether to include timestamp in each message.
|
||||
add_index: Whether to include round index in each message.
|
||||
|
||||
Returns:
|
||||
Formatted string with all messages joined by newlines.
|
||||
"""
|
||||
if not msgs:
|
||||
return ""
|
||||
|
||||
formatted_parts: list[str] = []
|
||||
total_token_count = 0
|
||||
|
||||
# Process messages in reverse order (newest first)
|
||||
for i in range(len(msgs) - 1, -1, -1):
|
||||
msg = msgs[i]
|
||||
index = i if add_index else None
|
||||
|
||||
formatted_msg, msg_token_count = self._format_single_msg(
|
||||
msg,
|
||||
index=index,
|
||||
add_time=add_time,
|
||||
)
|
||||
|
||||
# Always include current message first, then check threshold, at least one msg
|
||||
formatted_parts.append(formatted_msg)
|
||||
total_token_count += msg_token_count
|
||||
|
||||
# Check if we should stop adding older messages
|
||||
if total_token_count >= self._memory_compact_threshold:
|
||||
logger.info(
|
||||
"Skipping older messages: token count %d >= %d",
|
||||
total_token_count,
|
||||
self._memory_compact_threshold,
|
||||
)
|
||||
break
|
||||
|
||||
# Reverse to restore chronological order
|
||||
formatted_parts.reverse()
|
||||
|
||||
return "\n\n".join(formatted_parts)
|
||||
|
|
@ -8,7 +8,7 @@ from agentscope.token import HuggingFaceTokenCounter
|
|||
from .utils import _extract_text_from_messages
|
||||
|
||||
|
||||
class ReMeChatFormatter(OpenAIChatFormatter):
|
||||
class ReMeOpenAIChatFormatter(OpenAIChatFormatter):
|
||||
"""ReMe chat formatter class."""
|
||||
|
||||
async def _count(self, msgs: list[dict[str, Any]]) -> int | None:
|
||||
|
|
|
|||
|
|
@ -3,12 +3,11 @@
|
|||
import logging
|
||||
|
||||
from agentscope.agent._react_agent import _MemoryMark
|
||||
from agentscope.formatter import FormatterBase
|
||||
from agentscope.memory import InMemoryMemory
|
||||
from agentscope.message import Msg
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from .utils import safe_count_message_tokens, safe_count_str_tokens, _get_block_tokens
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -19,13 +18,10 @@ class ReMeInMemoryMemory(InMemoryMemory):
|
|||
def __init__(
|
||||
self,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
formatter: FormatterBase,
|
||||
max_input_length: int = 0,
|
||||
):
|
||||
super().__init__()
|
||||
self._token_counter: HuggingFaceTokenCounter = token_counter
|
||||
self._formatter: FormatterBase = formatter
|
||||
self._max_input_length: int = max_input_length
|
||||
self._msg_handler: AsMsgHandler = AsMsgHandler(token_counter)
|
||||
|
||||
async def get_memory(
|
||||
self,
|
||||
|
|
@ -127,9 +123,12 @@ Use it as context to maintain continuity.
|
|||
"""Clear the content."""
|
||||
self.content.clear()
|
||||
|
||||
async def estimate_tokens(self) -> dict:
|
||||
async def estimate_tokens(self, max_input_length: int) -> dict:
|
||||
"""Estimate token usage for current memory.
|
||||
|
||||
Args:
|
||||
max_input_length: Max input length for context usage calculation.
|
||||
|
||||
Returns:
|
||||
Dict containing detailed token statistics:
|
||||
- total_messages: Number of messages
|
||||
|
|
@ -138,7 +137,7 @@ Use it as context to maintain continuity.
|
|||
- estimated_tokens: Total estimated tokens
|
||||
- max_input_length: Max input length from config
|
||||
- context_usage_ratio: Usage percentage
|
||||
- messages_detail: List of per-message token details
|
||||
- messages_detail: List of per-message AsMsgStat objects
|
||||
"""
|
||||
messages = await self.get_memory(
|
||||
exclude_mark=_MemoryMark.COMPRESSED,
|
||||
|
|
@ -146,62 +145,18 @@ Use it as context to maintain continuity.
|
|||
)
|
||||
|
||||
compressed_summary = self.get_compressed_summary()
|
||||
compressed_summary_tokens = safe_count_str_tokens(self._token_counter, compressed_summary)
|
||||
compressed_summary_tokens = self._msg_handler.count_str_token(compressed_summary)
|
||||
|
||||
# Calculate total token count using formatter
|
||||
prompt = await self._formatter.format(msgs=messages)
|
||||
messages_tokens = safe_count_message_tokens(self._token_counter, prompt)
|
||||
# Build per-message token details using AsMsgHandler
|
||||
messages_detail = [self._msg_handler.stat_message(msg) for msg in messages]
|
||||
|
||||
# Calculate total message tokens from stats
|
||||
messages_tokens = sum(stat.total_tokens for stat in messages_detail)
|
||||
estimated_tokens = messages_tokens + compressed_summary_tokens
|
||||
|
||||
# Calculate context usage ratio
|
||||
max_input_length = self._max_input_length
|
||||
context_usage_ratio = (estimated_tokens / max_input_length * 100) if max_input_length > 0 else 0
|
||||
|
||||
# Build per-message token details
|
||||
messages_detail = []
|
||||
for i, msg in enumerate(messages, 1):
|
||||
msg_detail = {
|
||||
"index": i,
|
||||
"role": msg.role,
|
||||
"text_tokens": 0,
|
||||
"blocks": [],
|
||||
"preview": "",
|
||||
}
|
||||
try:
|
||||
content = msg.content
|
||||
if isinstance(content, str):
|
||||
text_tokens = safe_count_str_tokens(self._token_counter, content)
|
||||
msg_detail["text_tokens"] = text_tokens
|
||||
msg_detail["preview"] = f"{content[:100]}..." if len(content) > 100 else content
|
||||
else:
|
||||
total_tokens = 0
|
||||
text_parts = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
block_type = block.get("type", "unknown")
|
||||
block_tokens, block_str = _get_block_tokens(
|
||||
block,
|
||||
block_type,
|
||||
self._token_counter,
|
||||
)
|
||||
total_tokens += block_tokens
|
||||
text_parts.append(block_str)
|
||||
msg_detail["blocks"].append(
|
||||
{
|
||||
"type": block_type,
|
||||
"tokens": block_tokens,
|
||||
},
|
||||
)
|
||||
msg_detail["text_tokens"] = total_tokens
|
||||
text_preview = "".join(text_parts)
|
||||
msg_detail["preview"] = f"{text_preview[:100]}..." if len(text_preview) > 100 else text_preview
|
||||
except Exception as e:
|
||||
msg_detail["error"] = str(e)
|
||||
msg_detail["preview"] = f"<error: {e}>"
|
||||
|
||||
messages_detail.append(msg_detail)
|
||||
|
||||
return {
|
||||
"total_messages": len(messages),
|
||||
"compressed_summary_tokens": compressed_summary_tokens,
|
||||
|
|
@ -212,25 +167,28 @@ Use it as context to maintain continuity.
|
|||
"messages_detail": messages_detail,
|
||||
}
|
||||
|
||||
async def get_history_str(self) -> str:
|
||||
async def get_history_str(self, max_input_length: int) -> str:
|
||||
"""Get formatted history string similar to /history command output.
|
||||
|
||||
Args:
|
||||
max_input_length: Max input length for context usage calculation.
|
||||
|
||||
Returns:
|
||||
Formatted string containing conversation history details
|
||||
"""
|
||||
stats = await self.estimate_tokens()
|
||||
stats = await self.estimate_tokens(max_input_length)
|
||||
|
||||
lines = []
|
||||
for msg_detail in stats["messages_detail"]:
|
||||
for i, msg_stat in enumerate(stats["messages_detail"], 1):
|
||||
blocks_info = ""
|
||||
if msg_detail["blocks"]:
|
||||
block_strs = [f"{b['type']}(tokens={b['tokens']})" for b in msg_detail["blocks"]]
|
||||
if msg_stat.content:
|
||||
block_strs = [f"{b.block_type}(tokens={b.token_count})" for b in msg_stat.content]
|
||||
blocks_info = f"\n content: [{', '.join(block_strs)}]"
|
||||
|
||||
lines.append(
|
||||
f"[{msg_detail['index']}] **{msg_detail['role']}** "
|
||||
f"(text_tokens={msg_detail['text_tokens']})"
|
||||
f"{blocks_info}\n preview: {msg_detail['preview']}",
|
||||
f"[{i}] **{msg_stat.role}** "
|
||||
f"(total_tokens={msg_stat.total_tokens})"
|
||||
f"{blocks_info}\n preview: {msg_stat.preview}",
|
||||
)
|
||||
|
||||
return (
|
||||
|
|
|
|||
|
|
@ -10,8 +10,7 @@ from agentscope.model import ChatModelBase
|
|||
from agentscope.token import HuggingFaceTokenCounter
|
||||
from agentscope.tool import Toolkit
|
||||
|
||||
from .memory_formatter import MemoryFormatter
|
||||
from .file_io import FileIO
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
from ...core.op import BaseOp
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
|
@ -28,7 +27,7 @@ class Summarizer(BaseOp):
|
|||
chat_model: ChatModelBase,
|
||||
formatter: FormatterBase,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
toolkit: Toolkit | None = None,
|
||||
toolkit: Toolkit,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
|
|
@ -38,15 +37,8 @@ class Summarizer(BaseOp):
|
|||
|
||||
self.chat_model: ChatModelBase = chat_model
|
||||
self.formatter: FormatterBase = formatter
|
||||
self.as_token_counter: HuggingFaceTokenCounter = token_counter
|
||||
if toolkit is not None:
|
||||
self.toolkit: Toolkit = toolkit
|
||||
else:
|
||||
self.toolkit = Toolkit()
|
||||
file_io = FileIO(working_dir=self.working_dir)
|
||||
self.toolkit.register_tool_function(file_io.read)
|
||||
self.toolkit.register_tool_function(file_io.write)
|
||||
self.toolkit.register_tool_function(file_io.edit)
|
||||
self.msg_handler = AsMsgHandler(token_counter=token_counter)
|
||||
self.toolkit: Toolkit = toolkit
|
||||
|
||||
async def execute(self):
|
||||
messages: list[Msg] = self.context.get("messages", [])
|
||||
|
|
@ -54,11 +46,10 @@ class Summarizer(BaseOp):
|
|||
if not messages:
|
||||
return ""
|
||||
|
||||
formatter = MemoryFormatter(
|
||||
token_counter=self.as_token_counter,
|
||||
history_formatted_str: str = self.msg_handler.format_msgs_to_str(
|
||||
messages=messages,
|
||||
memory_compact_threshold=self.memory_compact_threshold,
|
||||
)
|
||||
history_formatted_str: str = formatter.format(messages)
|
||||
|
||||
if not history_formatted_str:
|
||||
logger.warning(f"No history to summarize. messages={messages}")
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from agentscope.tool import Toolkit, ToolResponse
|
|||
|
||||
from .config import ReMeConfigParser
|
||||
from .core import Application
|
||||
from .memory.file_based import Compactor, Summarizer, ToolResultCompactor, ReMeInMemoryMemory, ReMeChatFormatter
|
||||
from .memory.file_based import Compactor, Summarizer, ToolResultCompactor, ReMeInMemoryMemory, ReMeOpenAIChatFormatter, FileIO
|
||||
from .memory.file_based.utils import get_token_counter
|
||||
from .memory.tools import MemorySearch
|
||||
from .core.utils import load_env
|
||||
|
|
@ -95,11 +95,6 @@ class ReMeLight(Application):
|
|||
self.tool_result_path = self.working_path / "tool_result"
|
||||
self.tool_result_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Initialize runtime parameters (will be updated via update_params)
|
||||
self.max_input_length: int = 0
|
||||
self.memory_compact_threshold: int = 0
|
||||
self.language: str = ""
|
||||
|
||||
# Apply initial parameter configuration
|
||||
self.update_params(
|
||||
max_input_length=max_input_length,
|
||||
|
|
@ -198,7 +193,7 @@ class ReMeLight(Application):
|
|||
if formatter is not None:
|
||||
self.formatter: FormatterBase = formatter
|
||||
else:
|
||||
self.formatter = ReMeChatFormatter(token_counter=self.token_counter)
|
||||
self.formatter = ReMeOpenAIChatFormatter(token_counter=self.token_counter)
|
||||
self.toolkit: Toolkit | None = toolkit
|
||||
|
||||
# Initialize list to track background summarization tasks
|
||||
|
|
@ -458,6 +453,16 @@ class ReMeLight(Application):
|
|||
- If summarization fails, an empty string is returned
|
||||
"""
|
||||
try:
|
||||
# Create toolkit if not provided
|
||||
if self.toolkit is not None:
|
||||
toolkit = self.toolkit
|
||||
else:
|
||||
toolkit = Toolkit()
|
||||
file_io = FileIO(working_dir=str(self.working_path))
|
||||
toolkit.register_tool_function(file_io.read)
|
||||
toolkit.register_tool_function(file_io.write)
|
||||
toolkit.register_tool_function(file_io.edit)
|
||||
|
||||
# Initialize summarizer with working directories and configuration
|
||||
summarizer = Summarizer(
|
||||
working_dir=str(self.working_path),
|
||||
|
|
@ -466,7 +471,7 @@ class ReMeLight(Application):
|
|||
chat_model=self.chat_model,
|
||||
formatter=self.formatter,
|
||||
token_counter=self.token_counter,
|
||||
toolkit=self.toolkit,
|
||||
toolkit=toolkit,
|
||||
language=self.language,
|
||||
)
|
||||
|
||||
|
|
@ -662,10 +667,8 @@ class ReMeLight(Application):
|
|||
Note:
|
||||
- In-memory memory is volatile and cleared when the instance is destroyed
|
||||
- Useful for managing conversation context within a single session
|
||||
- Shares the same token counter and formatter as the main application
|
||||
- Shares the same token counter as the main application
|
||||
"""
|
||||
return ReMeInMemoryMemory(
|
||||
token_counter=self.token_counter,
|
||||
formatter=self.formatter,
|
||||
max_input_length=self.max_input_length,
|
||||
)
|
||||
|
|
|
|||
1090
tests/light/test_context_check.py
Normal file
1090
tests/light/test_context_check.py
Normal file
File diff suppressed because it is too large
Load diff
883
tests/light/test_format_msgs_to_str.py
Normal file
883
tests/light/test_format_msgs_to_str.py
Normal file
|
|
@ -0,0 +1,883 @@
|
|||
"""Tests for AsMsgHandler.format_msgs_to_str method."""
|
||||
|
||||
# pylint: disable=W0212
|
||||
|
||||
import logging
|
||||
|
||||
from agentscope.message import Msg
|
||||
|
||||
from test_utils import get_token_counter
|
||||
from reme.memory.file_based.as_msg_handler import AsMsgHandler
|
||||
|
||||
# 配置日志输出到控制台
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ANSI 颜色码
|
||||
class Colors:
|
||||
"""ANSI color codes for terminal output."""
|
||||
|
||||
GREEN = "\033[92m"
|
||||
RED = "\033[91m"
|
||||
YELLOW = "\033[93m"
|
||||
BLUE = "\033[94m"
|
||||
CYAN = "\033[96m"
|
||||
BOLD = "\033[1m"
|
||||
RESET = "\033[0m"
|
||||
|
||||
|
||||
def print_pass(test_name: str):
|
||||
"""打印测试通过信息"""
|
||||
print(f"{Colors.GREEN}{Colors.BOLD}✓ {test_name} PASSED{Colors.RESET}")
|
||||
|
||||
|
||||
def print_fail(test_name: str, error: str):
|
||||
"""打印测试失败信息"""
|
||||
print(f"{Colors.RED}{Colors.BOLD}✗ {test_name} FAILED: {error}{Colors.RESET}")
|
||||
|
||||
|
||||
def print_error(test_name: str, error: str):
|
||||
"""打印测试错误信息"""
|
||||
print(f"{Colors.YELLOW}{Colors.BOLD}⚠ {test_name} ERROR: {error}{Colors.RESET}")
|
||||
|
||||
|
||||
def print_test_header(test_name: str):
|
||||
"""打印测试标题"""
|
||||
print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
|
||||
print(f"{Colors.BLUE}{Colors.BOLD}Running: {test_name}{Colors.RESET}")
|
||||
print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
|
||||
|
||||
|
||||
# ==================== Helper Functions ====================
|
||||
|
||||
|
||||
def create_handler() -> AsMsgHandler:
|
||||
"""Create an AsMsgHandler instance for testing."""
|
||||
return AsMsgHandler(token_counter=get_token_counter())
|
||||
|
||||
|
||||
def verify_result_within_threshold(
|
||||
handler: AsMsgHandler,
|
||||
result: str,
|
||||
threshold: int,
|
||||
test_name: str = "",
|
||||
msgs: list[Msg] | None = None,
|
||||
) -> None:
|
||||
"""Verify that the included messages' original token count does not exceed threshold.
|
||||
|
||||
Note: The format_msgs_to_str method uses message token statistics (not formatted
|
||||
string tokens) for threshold checking. The formatted result may have more tokens
|
||||
than the threshold due to added metadata (timestamps, role prefixes, etc.).
|
||||
|
||||
This verification checks that included messages' original token sum <= threshold.
|
||||
|
||||
Args:
|
||||
handler: The AsMsgHandler instance used for token counting.
|
||||
result: The formatted string result from format_msgs_to_str.
|
||||
threshold: The memory_compact_threshold value used.
|
||||
test_name: Optional test name for better error messages.
|
||||
msgs: Optional list of original messages to verify against.
|
||||
|
||||
Raises:
|
||||
AssertionError: If included messages' token count exceeds threshold.
|
||||
"""
|
||||
if not result or not msgs:
|
||||
return # Empty result or no messages to verify
|
||||
|
||||
# Calculate tokens of messages that were included in the result
|
||||
included_tokens = 0
|
||||
for msg in msgs:
|
||||
stat = handler.stat_message(msg)
|
||||
# Check if this message's content appears in the result
|
||||
formatted = stat.format(include_thinking=True) # Use True to check all content
|
||||
# Simple heuristic: if the message content is in result, count its tokens
|
||||
content_blocks = msg.get_content_blocks()
|
||||
msg_included = False
|
||||
for block in content_blocks:
|
||||
block_type = block.get("type", "")
|
||||
if block_type == "text" and block.get("text", "") in result:
|
||||
msg_included = True
|
||||
break
|
||||
elif block_type == "tool_use" and f"tool_call={block.get('name', '')}" in result:
|
||||
msg_included = True
|
||||
break
|
||||
elif block_type == "tool_result" and f"tool_result={block.get('name', '')}" in result:
|
||||
msg_included = True
|
||||
break
|
||||
|
||||
if msg_included:
|
||||
included_tokens += stat.total_tokens
|
||||
|
||||
# Verify included messages' token sum doesn't exceed threshold
|
||||
# Allow small tolerance for edge cases
|
||||
assert included_tokens <= threshold + 1, (
|
||||
f"{test_name}: Included messages token count ({included_tokens}) exceeds threshold ({threshold})."
|
||||
)
|
||||
|
||||
|
||||
def create_user_msg(content: str) -> Msg:
|
||||
"""Create a user message."""
|
||||
return Msg(name="user", role="user", content=content)
|
||||
|
||||
|
||||
def create_assistant_msg(content: str) -> Msg:
|
||||
"""Create an assistant message."""
|
||||
return Msg(name="assistant", role="assistant", content=content)
|
||||
|
||||
|
||||
def create_tool_use_msg(tool_name: str, tool_input: dict, tool_id: str = "call_123") -> Msg:
|
||||
"""Create a message with tool_use content block."""
|
||||
return Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": tool_id,
|
||||
"name": tool_name,
|
||||
"input": tool_input,
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def create_tool_result_msg(tool_name: str, output: str | list[dict], tool_id: str = "call_123") -> Msg:
|
||||
"""Create a message with tool_result content block."""
|
||||
return Msg(
|
||||
name="tool",
|
||||
role="user",
|
||||
content=[
|
||||
{
|
||||
"type": "tool_result",
|
||||
"id": tool_id,
|
||||
"name": tool_name,
|
||||
"output": output,
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def create_thinking_msg(thinking_content: str, text_content: str = "") -> Msg:
|
||||
"""Create a message with thinking content block."""
|
||||
content = [
|
||||
{
|
||||
"type": "thinking",
|
||||
"thinking": thinking_content,
|
||||
},
|
||||
]
|
||||
if text_content:
|
||||
content.append({"type": "text", "text": text_content})
|
||||
return Msg(name="assistant", role="assistant", content=content)
|
||||
|
||||
|
||||
def create_image_msg(url: str = "") -> Msg:
|
||||
"""Create a message with image content block."""
|
||||
content = [
|
||||
{
|
||||
"type": "image",
|
||||
"source": {"url": url} if url else {},
|
||||
},
|
||||
]
|
||||
return Msg(name="assistant", role="assistant", content=content)
|
||||
|
||||
|
||||
def create_mixed_content_msg(
|
||||
text: str = "",
|
||||
thinking: str = "",
|
||||
tool_name: str = "",
|
||||
tool_input: dict | None = None,
|
||||
image_url: str = "",
|
||||
) -> Msg:
|
||||
"""Create a message with mixed content blocks."""
|
||||
content = []
|
||||
if thinking:
|
||||
content.append({"type": "thinking", "thinking": thinking})
|
||||
if text:
|
||||
content.append({"type": "text", "text": text})
|
||||
if tool_name:
|
||||
content.append({
|
||||
"type": "tool_use",
|
||||
"id": "call_mixed",
|
||||
"name": tool_name,
|
||||
"input": tool_input or {},
|
||||
})
|
||||
if image_url:
|
||||
content.append({"type": "image", "source": {"url": image_url}})
|
||||
return Msg(name="assistant", role="assistant", content=content)
|
||||
|
||||
|
||||
# ==================== Normal Case Tests ====================
|
||||
|
||||
|
||||
def test_format_msgs_to_str_empty_list():
|
||||
"""Test format_msgs_to_str with empty message list."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = []
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
assert result == "", f"Expected empty string for empty list, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "empty_list", msgs)
|
||||
print_pass("test_format_msgs_to_str_empty_list")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_single_message():
|
||||
"""Test format_msgs_to_str with a single message."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_user_msg("Hello, how are you?")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "user:" in result, f"Expected 'user:' in result, got: {result}"
|
||||
assert "Hello, how are you?" in result, f"Expected content in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "single_message", msgs)
|
||||
print_pass("test_format_msgs_to_str_single_message")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_multiple_messages():
|
||||
"""Test format_msgs_to_str with multiple messages."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [
|
||||
create_user_msg("What is Python?"),
|
||||
create_assistant_msg("Python is a programming language."),
|
||||
create_user_msg("Tell me more."),
|
||||
create_assistant_msg("Python is known for its readability."),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "What is Python?" in result
|
||||
assert "Python is a programming language." in result
|
||||
assert "Tell me more." in result
|
||||
assert "Python is known for its readability." in result
|
||||
verify_result_within_threshold(handler, result, threshold, "multiple_messages", msgs)
|
||||
print_pass("test_format_msgs_to_str_multiple_messages")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_message_order():
|
||||
"""Test that messages are returned in correct order (oldest to newest)."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [
|
||||
create_user_msg("First message"),
|
||||
create_assistant_msg("Second message"),
|
||||
create_user_msg("Third message"),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# Find positions of each message
|
||||
first_pos = result.find("First message")
|
||||
second_pos = result.find("Second message")
|
||||
third_pos = result.find("Third message")
|
||||
|
||||
assert first_pos < second_pos < third_pos, (
|
||||
f"Messages not in correct order. Positions: first={first_pos}, "
|
||||
f"second={second_pos}, third={third_pos}"
|
||||
)
|
||||
verify_result_within_threshold(handler, result, threshold, "message_order", msgs)
|
||||
print_pass("test_format_msgs_to_str_message_order")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_with_tool_use():
|
||||
"""Test format_msgs_to_str with tool_use message."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_tool_use_msg("read_file", {"path": "/test.txt"})]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "tool_call=read_file" in result, f"Expected tool_call in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "with_tool_use", msgs)
|
||||
print_pass("test_format_msgs_to_str_with_tool_use")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_with_tool_result():
|
||||
"""Test format_msgs_to_str with tool_result message."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_tool_result_msg("read_file", "file content here")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "tool_result=read_file" in result, f"Expected tool_result in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "with_tool_result", msgs)
|
||||
print_pass("test_format_msgs_to_str_with_tool_result")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_with_image():
|
||||
"""Test format_msgs_to_str with image message."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_image_msg("https://example.com/image.png")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "[image]" in result, f"Expected '[image]' in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "with_image", msgs)
|
||||
print_pass("test_format_msgs_to_str_with_image")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_conversation_flow():
|
||||
"""Test format_msgs_to_str with a complete conversation flow."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [
|
||||
create_user_msg("Read the file."),
|
||||
create_tool_use_msg("read_file", {"path": "/data.txt"}),
|
||||
create_tool_result_msg("read_file", "File content here"),
|
||||
create_assistant_msg("The file contains: File content here"),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "user:" in result
|
||||
assert "tool_call=read_file" in result
|
||||
assert "tool_result=read_file" in result
|
||||
assert "assistant:" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "conversation_flow", msgs)
|
||||
print_pass("test_format_msgs_to_str_conversation_flow")
|
||||
|
||||
|
||||
# ==================== Thinking Block Tests ====================
|
||||
|
||||
|
||||
def test_format_msgs_to_str_thinking_excluded_by_default():
|
||||
"""Test that thinking blocks are excluded when include_thinking=False (default)."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_thinking_msg("Let me think about this...", "Here is my response")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold, include_thinking=False)
|
||||
|
||||
assert "Let me think about this" not in result, (
|
||||
f"Thinking content should be excluded, got: {result}"
|
||||
)
|
||||
assert "Here is my response" in result, f"Text content should be included, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "thinking_excluded_by_default", msgs)
|
||||
print_pass("test_format_msgs_to_str_thinking_excluded_by_default")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_thinking_included():
|
||||
"""Test that thinking blocks are included when include_thinking=True."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_thinking_msg("Let me think about this...", "Here is my response")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold, include_thinking=True)
|
||||
|
||||
assert "Let me think about this" in result, (
|
||||
f"Thinking content should be included, got: {result}"
|
||||
)
|
||||
assert "<thinking>" in result, f"Expected thinking tag in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "thinking_included", msgs)
|
||||
print_pass("test_format_msgs_to_str_thinking_included")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_thinking_only_message():
|
||||
"""Test message with only thinking block."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_thinking_msg("Deep thoughts here")]
|
||||
|
||||
# With include_thinking=False
|
||||
result_no_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=False
|
||||
)
|
||||
# With include_thinking=True
|
||||
result_with_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=True
|
||||
)
|
||||
|
||||
assert "Deep thoughts here" not in result_no_thinking
|
||||
assert "Deep thoughts here" in result_with_thinking
|
||||
verify_result_within_threshold(handler, result_no_thinking, threshold, "thinking_only_no_thinking", msgs)
|
||||
verify_result_within_threshold(handler, result_with_thinking, threshold, "thinking_only_with_thinking", msgs)
|
||||
print_pass("test_format_msgs_to_str_thinking_only_message")
|
||||
|
||||
|
||||
# ==================== Token Threshold Tests ====================
|
||||
|
||||
|
||||
def test_format_msgs_to_str_all_within_threshold():
|
||||
"""Test all messages fit within threshold."""
|
||||
handler = create_handler()
|
||||
threshold = 10000
|
||||
msgs = [
|
||||
create_user_msg("Short message 1"),
|
||||
create_assistant_msg("Short message 2"),
|
||||
create_user_msg("Short message 3"),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "Short message 1" in result
|
||||
assert "Short message 2" in result
|
||||
assert "Short message 3" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "all_within_threshold", msgs)
|
||||
print_pass("test_format_msgs_to_str_all_within_threshold")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_exceeds_threshold_truncate_older():
|
||||
"""Test that older messages are truncated when exceeding threshold."""
|
||||
handler = create_handler()
|
||||
threshold = 500
|
||||
msgs = []
|
||||
for i in range(20):
|
||||
msgs.append(create_user_msg(f"Question {i}: " + "x" * 100))
|
||||
msgs.append(create_assistant_msg(f"Answer {i}: " + "y" * 100))
|
||||
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# The newest messages should be present
|
||||
assert "Answer 19" in result or "Question 19" in result, (
|
||||
f"Expected recent message in result, got: {result[:500]}..."
|
||||
)
|
||||
# Older messages should be truncated
|
||||
assert "Question 0" not in result, "Older messages should be truncated"
|
||||
verify_result_within_threshold(handler, result, threshold, "exceeds_threshold_truncate_older", msgs)
|
||||
print_pass("test_format_msgs_to_str_exceeds_threshold_truncate_older")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_single_message_exceeds_threshold():
|
||||
"""Test when a single message exceeds the threshold."""
|
||||
handler = create_handler()
|
||||
threshold = 10
|
||||
# Create a very long message
|
||||
long_text = "x" * 10000
|
||||
msgs = [create_user_msg(long_text)]
|
||||
|
||||
# With very low threshold, even a single message won't fit
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# The message should be skipped entirely since it exceeds threshold
|
||||
assert result == "" or len(result) > 0, "Result should be empty or contain truncated content"
|
||||
verify_result_within_threshold(handler, result, threshold, "single_message_exceeds_threshold", msgs)
|
||||
print_pass("test_format_msgs_to_str_single_message_exceeds_threshold")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_first_message_exceeds_threshold():
|
||||
"""Test when the first (oldest) message exceeds threshold but newer ones don't."""
|
||||
handler = create_handler()
|
||||
threshold = 100
|
||||
msgs = [
|
||||
create_user_msg("x" * 5000), # Old, long message
|
||||
create_assistant_msg("Short response"), # New, short message
|
||||
]
|
||||
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# Newer message should be present
|
||||
assert "Short response" in result, f"Expected newer message in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "first_message_exceeds_threshold", msgs)
|
||||
print_pass("test_format_msgs_to_str_first_message_exceeds_threshold")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_threshold_zero():
|
||||
"""Test with threshold of zero - no messages should be included."""
|
||||
handler = create_handler()
|
||||
threshold = 0
|
||||
msgs = [create_user_msg("Test message")]
|
||||
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert result == "", f"Expected empty string with zero threshold, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "threshold_zero", msgs)
|
||||
print_pass("test_format_msgs_to_str_threshold_zero")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_threshold_exact_fit():
|
||||
"""Test when messages exactly fit the threshold."""
|
||||
handler = create_handler()
|
||||
# Create a message and measure its tokens
|
||||
msg = create_user_msg("Test")
|
||||
stat = handler.stat_message(msg)
|
||||
exact_threshold = stat.total_tokens
|
||||
|
||||
msgs = [msg]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=exact_threshold)
|
||||
|
||||
assert "Test" in result, f"Message should fit exactly, got: {result}"
|
||||
verify_result_within_threshold(handler, result, exact_threshold, "threshold_exact_fit", msgs)
|
||||
print_pass("test_format_msgs_to_str_threshold_exact_fit")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_threshold_one_less():
|
||||
"""Test when threshold is one less than needed."""
|
||||
handler = create_handler()
|
||||
msg = create_user_msg("Test message")
|
||||
stat = handler.stat_message(msg)
|
||||
threshold_minus_one = stat.total_tokens - 1
|
||||
|
||||
msgs = [msg]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold_minus_one)
|
||||
|
||||
# Message should be skipped since it doesn't fit
|
||||
assert result == "", f"Expected empty string when threshold is insufficient, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold_minus_one, "threshold_one_less", msgs)
|
||||
print_pass("test_format_msgs_to_str_threshold_one_less")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_large_threshold():
|
||||
"""Test with very large threshold - all messages should be included."""
|
||||
handler = create_handler()
|
||||
threshold = 1000000
|
||||
msgs = [
|
||||
create_user_msg("Message " + str(i) + " " + "x" * 100)
|
||||
for i in range(50)
|
||||
]
|
||||
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# All messages should be included
|
||||
for i in range(50):
|
||||
assert f"Message {i}" in result, f"Message {i} should be included"
|
||||
verify_result_within_threshold(handler, result, threshold, "large_threshold", msgs)
|
||||
print_pass("test_format_msgs_to_str_large_threshold")
|
||||
|
||||
|
||||
# ==================== Edge Cases Tests ====================
|
||||
|
||||
|
||||
def test_format_msgs_to_str_special_characters():
|
||||
"""Test with special characters in content."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_user_msg("Test with 中文, 日本語, émojis 🎉 and symbols @#$%")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "中文" in result
|
||||
assert "日本語" in result
|
||||
assert "🎉" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "special_characters", msgs)
|
||||
print_pass("test_format_msgs_to_str_special_characters")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_empty_content():
|
||||
"""Test with empty content message."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_user_msg("")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "user:" in result, f"Expected role in result even with empty content, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "empty_content", msgs)
|
||||
print_pass("test_format_msgs_to_str_empty_content")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_whitespace_only():
|
||||
"""Test with whitespace-only content."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_user_msg(" \n\t ")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "user:" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "whitespace_only", msgs)
|
||||
print_pass("test_format_msgs_to_str_whitespace_only")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_newlines_in_content():
|
||||
"""Test with newlines in message content."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_user_msg("Line 1\nLine 2\nLine 3")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "Line 1" in result
|
||||
assert "Line 2" in result
|
||||
assert "Line 3" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "newlines_in_content", msgs)
|
||||
print_pass("test_format_msgs_to_str_newlines_in_content")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_very_long_single_word():
|
||||
"""Test with very long single word (no spaces)."""
|
||||
handler = create_handler()
|
||||
threshold = 10000
|
||||
long_word = "a" * 5000
|
||||
msgs = [create_user_msg(long_word)]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# Should contain at least part of the word (may be truncated by formatter)
|
||||
assert "aaa" in result, f"Expected long word content in result, got: {result[:100]}..."
|
||||
verify_result_within_threshold(handler, result, threshold, "very_long_single_word", msgs)
|
||||
print_pass("test_format_msgs_to_str_very_long_single_word")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_mixed_content_blocks():
|
||||
"""Test message with mixed content blocks."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_mixed_content_msg(
|
||||
text="Text content",
|
||||
thinking="Thinking content",
|
||||
tool_name="test_tool",
|
||||
tool_input={"key": "value"},
|
||||
image_url="https://example.com/img.png",
|
||||
)]
|
||||
|
||||
result_no_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=False
|
||||
)
|
||||
result_with_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=True
|
||||
)
|
||||
|
||||
assert "Text content" in result_no_thinking
|
||||
assert "tool_call=test_tool" in result_no_thinking
|
||||
assert "[image]" in result_no_thinking
|
||||
assert "Thinking content" not in result_no_thinking
|
||||
assert "Thinking content" in result_with_thinking
|
||||
verify_result_within_threshold(handler, result_no_thinking, threshold, "mixed_content_no_thinking", msgs)
|
||||
verify_result_within_threshold(handler, result_with_thinking, threshold, "mixed_content_with_thinking", msgs)
|
||||
print_pass("test_format_msgs_to_str_mixed_content_blocks")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_multiple_separators():
|
||||
"""Test that messages are separated by double newlines."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [
|
||||
create_user_msg("Message 1"),
|
||||
create_assistant_msg("Message 2"),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "\n\n" in result, f"Expected double newline separator, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "multiple_separators", msgs)
|
||||
print_pass("test_format_msgs_to_str_multiple_separators")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_tool_result_complex_output():
|
||||
"""Test tool_result with complex output (list of blocks)."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
complex_output = [
|
||||
{"type": "text", "text": "Operation completed"},
|
||||
{"type": "image", "source": {"url": "https://example.com/result.png"}},
|
||||
]
|
||||
msgs = [create_tool_result_msg("process_data", complex_output)]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "tool_result=process_data" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "tool_result_complex_output", msgs)
|
||||
print_pass("test_format_msgs_to_str_tool_result_complex_output")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_different_roles():
|
||||
"""Test with different roles (user, assistant, system, tool)."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [
|
||||
Msg(name="system", role="system", content="System instruction"),
|
||||
create_user_msg("User message"),
|
||||
create_assistant_msg("Assistant response"),
|
||||
create_tool_result_msg("tool", "Tool output"),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "system:" in result
|
||||
assert "user:" in result
|
||||
assert "assistant:" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "different_roles", msgs)
|
||||
print_pass("test_format_msgs_to_str_different_roles")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_incremental_threshold_check():
|
||||
"""Test incremental addition of messages until threshold is exceeded."""
|
||||
handler = create_handler()
|
||||
|
||||
# Create messages with known approximate sizes
|
||||
msgs = []
|
||||
for i in range(10):
|
||||
msgs.append(create_user_msg(f"Message {i} with some padding text"))
|
||||
|
||||
# Calculate total tokens
|
||||
total_tokens = sum(handler.stat_message(msg).total_tokens for msg in msgs)
|
||||
|
||||
# Use threshold that allows about half the messages
|
||||
half_threshold = total_tokens // 2
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=half_threshold)
|
||||
|
||||
# Should have some but not all messages
|
||||
included_count = sum(1 for i in range(10) if f"Message {i}" in result)
|
||||
assert 0 < included_count < 10, (
|
||||
f"Expected partial messages, got {included_count} messages included"
|
||||
)
|
||||
# Newer messages should be included (messages are processed from end)
|
||||
assert "Message 9" in result, "Newest message should be included"
|
||||
verify_result_within_threshold(handler, result, half_threshold, "incremental_threshold_check", msgs)
|
||||
print_pass("test_format_msgs_to_str_incremental_threshold_check")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_negative_threshold():
|
||||
"""Test with negative threshold value."""
|
||||
handler = create_handler()
|
||||
threshold = -1
|
||||
msgs = [create_user_msg("Test message")]
|
||||
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# Negative threshold should result in empty string (nothing fits)
|
||||
assert result == "", f"Expected empty string with negative threshold, got: {result}"
|
||||
verify_result_within_threshold(handler, result, max(0, threshold), "negative_threshold", msgs)
|
||||
print_pass("test_format_msgs_to_str_negative_threshold")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_preserves_newest_first():
|
||||
"""Test that newest messages are preserved when threshold is exceeded."""
|
||||
handler = create_handler()
|
||||
threshold = 300
|
||||
msgs = [
|
||||
create_user_msg("OLD MESSAGE " + "x" * 200),
|
||||
create_assistant_msg("MIDDLE MESSAGE " + "y" * 200),
|
||||
create_user_msg("NEW MESSAGE " + "z" * 200),
|
||||
]
|
||||
|
||||
# Use threshold that only allows ~1-2 messages
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# Newest message should be present
|
||||
assert "NEW MESSAGE" in result, f"Expected newest message, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "preserves_newest_first", msgs)
|
||||
print_pass("test_format_msgs_to_str_preserves_newest_first")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_base64_image():
|
||||
"""Test with base64 encoded image."""
|
||||
handler = create_handler()
|
||||
threshold = 10000
|
||||
msgs = [Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"data": "SGVsbG8gV29ybGQ=" * 100, # Simulated base64 data
|
||||
},
|
||||
}],
|
||||
)]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "[image]" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "base64_image", msgs)
|
||||
print_pass("test_format_msgs_to_str_base64_image")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_audio_video_blocks():
|
||||
"""Test with audio and video content blocks."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{"type": "audio", "source": {"url": "https://example.com/audio.mp3"}},
|
||||
{"type": "video", "source": {"url": "https://example.com/video.mp4"}},
|
||||
],
|
||||
)]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "[audio]" in result
|
||||
assert "[video]" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "audio_video_blocks", msgs)
|
||||
print_pass("test_format_msgs_to_str_audio_video_blocks")
|
||||
|
||||
|
||||
def test_format_msgs_to_str_unknown_block_type():
|
||||
"""Test that unknown block types are skipped gracefully."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{"type": "unknown_type", "data": "some data"},
|
||||
{"type": "text", "text": "Valid text"},
|
||||
],
|
||||
)]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# Should still include valid content
|
||||
assert "Valid text" in result
|
||||
verify_result_within_threshold(handler, result, threshold, "unknown_block_type", msgs)
|
||||
print_pass("test_format_msgs_to_str_unknown_block_type")
|
||||
|
||||
|
||||
def run_all_tests():
|
||||
"""Run all tests."""
|
||||
tests = [
|
||||
# Normal case tests
|
||||
test_format_msgs_to_str_empty_list,
|
||||
test_format_msgs_to_str_single_message,
|
||||
test_format_msgs_to_str_multiple_messages,
|
||||
test_format_msgs_to_str_message_order,
|
||||
test_format_msgs_to_str_with_tool_use,
|
||||
test_format_msgs_to_str_with_tool_result,
|
||||
test_format_msgs_to_str_with_image,
|
||||
test_format_msgs_to_str_conversation_flow,
|
||||
# Thinking block tests
|
||||
test_format_msgs_to_str_thinking_excluded_by_default,
|
||||
test_format_msgs_to_str_thinking_included,
|
||||
test_format_msgs_to_str_thinking_only_message,
|
||||
# Token threshold tests
|
||||
test_format_msgs_to_str_all_within_threshold,
|
||||
test_format_msgs_to_str_exceeds_threshold_truncate_older,
|
||||
test_format_msgs_to_str_single_message_exceeds_threshold,
|
||||
test_format_msgs_to_str_first_message_exceeds_threshold,
|
||||
test_format_msgs_to_str_threshold_zero,
|
||||
test_format_msgs_to_str_threshold_exact_fit,
|
||||
test_format_msgs_to_str_threshold_one_less,
|
||||
test_format_msgs_to_str_large_threshold,
|
||||
# Edge cases tests
|
||||
test_format_msgs_to_str_special_characters,
|
||||
test_format_msgs_to_str_empty_content,
|
||||
test_format_msgs_to_str_whitespace_only,
|
||||
test_format_msgs_to_str_newlines_in_content,
|
||||
test_format_msgs_to_str_very_long_single_word,
|
||||
test_format_msgs_to_str_mixed_content_blocks,
|
||||
test_format_msgs_to_str_multiple_separators,
|
||||
test_format_msgs_to_str_tool_result_complex_output,
|
||||
test_format_msgs_to_str_different_roles,
|
||||
test_format_msgs_to_str_incremental_threshold_check,
|
||||
test_format_msgs_to_str_negative_threshold,
|
||||
test_format_msgs_to_str_preserves_newest_first,
|
||||
test_format_msgs_to_str_base64_image,
|
||||
test_format_msgs_to_str_audio_video_blocks,
|
||||
test_format_msgs_to_str_unknown_block_type,
|
||||
]
|
||||
|
||||
passed = 0
|
||||
failed = 0
|
||||
|
||||
for test in tests:
|
||||
try:
|
||||
print_test_header(test.__name__)
|
||||
test()
|
||||
passed += 1
|
||||
except AssertionError as e:
|
||||
print_fail(test.__name__, str(e))
|
||||
failed += 1
|
||||
except Exception as e:
|
||||
print_error(test.__name__, str(e))
|
||||
failed += 1
|
||||
|
||||
# 打印最终统计结果
|
||||
print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
|
||||
print(f"{Colors.BOLD}Test Results Summary{Colors.RESET}")
|
||||
print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
|
||||
print(f"{Colors.GREEN}{Colors.BOLD}✓ Passed: {passed}{Colors.RESET}")
|
||||
if failed > 0:
|
||||
print(f"{Colors.RED}{Colors.BOLD}✗ Failed: {failed}{Colors.RESET}")
|
||||
else:
|
||||
print(f"{Colors.GREEN}✗ Failed: {failed}{Colors.RESET}")
|
||||
print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
|
||||
|
||||
if failed == 0:
|
||||
print(f"\n{Colors.GREEN}{Colors.BOLD}🎉 All tests passed!{Colors.RESET}")
|
||||
else:
|
||||
print(f"\n{Colors.RED}{Colors.BOLD}💥 Some tests failed!{Colors.RESET}")
|
||||
|
||||
return failed == 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
success = run_all_tests()
|
||||
exit(0 if success else 1)
|
||||
Loading…
Add table
Reference in a new issue