refactor(core): update registry registration syntax and improve code formatting

This commit is contained in:
jinli.yl 2026-03-06 01:57:29 +08:00
parent 57f8a7b42c
commit fc7b1cdba8
24 changed files with 616 additions and 468 deletions

View file

@ -1,4 +1,5 @@
"""Core"""
from . import as_llm
from . import as_llm_formatter
from . import embedding

View file

@ -25,26 +25,26 @@ class Application:
"""Application wrapper that wires together service context, flows, and runtimes."""
def __init__(
self,
*args,
llm_api_key: str | None = None,
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
working_dir: str | None = None,
config_path: str | None = None,
enable_logo: bool = True,
log_to_console: bool = True,
parser: type[PydanticConfigParser] | None = None,
default_as_llm_config: dict | None = None,
default_as_llm_formatter_config: dict | None = None,
default_llm_config: dict | None = None,
default_embedding_model_config: dict | None = None,
default_vector_store_config: dict | None = None,
default_file_store_config: dict | None = None,
default_token_counter_config: dict | None = None,
default_file_watcher_config: dict | None = None,
**kwargs,
self,
*args,
llm_api_key: str | None = None,
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
working_dir: str | None = None,
config_path: str | None = None,
enable_logo: bool = True,
log_to_console: bool = True,
parser: type[PydanticConfigParser] | None = None,
default_as_llm_config: dict | None = None,
default_as_llm_formatter_config: dict | None = None,
default_llm_config: dict | None = None,
default_embedding_model_config: dict | None = None,
default_vector_store_config: dict | None = None,
default_file_store_config: dict | None = None,
default_token_counter_config: dict | None = None,
default_file_watcher_config: dict | None = None,
**kwargs,
):
self.service_context = ServiceContext(
*args,
@ -142,8 +142,8 @@ class Application:
ray.init(num_cpus=self.service_config.ray_max_workers)
if (
self.service_context.thread_pool is None
or self.service_context.thread_pool._shutdown # pylint: disable=protected-access
self.service_context.thread_pool is None
or self.service_context.thread_pool._shutdown # pylint: disable=protected-access
):
self.service_context.thread_pool = ThreadPoolExecutor(
max_workers=self.service_config.thread_pool_max_workers,
@ -319,10 +319,10 @@ class Application:
stream_queue = asyncio.Queue()
task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs))
async for chunk in execute_stream_task(
stream_queue=stream_queue,
task=task,
task_name=name,
output_format="str",
stream_queue=stream_queue,
task=task,
task_name=name,
output_format="str",
):
yield chunk

View file

@ -1,7 +1,9 @@
"""Module for registering AgentScope LLM models."""
from agentscope.model import DashScopeChatModel
from agentscope.model import OpenAIChatModel
from ..registry_factory import R
R.as_llms.register(OpenAIChatModel, "openai")
R.as_llms.register(DashScopeChatModel, "dashscope")
R.as_llms.register("openai")(OpenAIChatModel)
R.as_llms.register("dashscope")(DashScopeChatModel)

View file

@ -1,7 +1,9 @@
"""Module for registering AgentScope LLM formatters."""
from agentscope.formatter import DashScopeChatFormatter
from agentscope.formatter import OpenAIChatFormatter
from ..registry_factory import R
R.as_llm_formatters.register(OpenAIChatFormatter, "openai")
R.as_llm_formatters.register(DashScopeChatFormatter, "dashscope")
R.as_llm_formatters.register("openai")(OpenAIChatFormatter)
R.as_llm_formatters.register("dashscope")(DashScopeChatFormatter)

View file

@ -7,6 +7,8 @@ from abc import ABCMeta
from pathlib import Path
from typing import Callable, Optional, Any
from agentscope.formatter import FormatterBase
from agentscope.model import ChatModelBase
from loguru import logger
from tqdm import tqdm
@ -21,8 +23,7 @@ from ..service_context import ServiceContext
from ..token_counter import BaseTokenCounter
from ..utils import camel_to_snake, CacheHandler, timer
from ..vector_store import BaseVectorStore
from agentscope.model import ChatModelBase
from agentscope.formatter import FormatterBase
class BaseOp(metaclass=ABCMeta):
"""Base operator class for LLM workflow execution and composition."""

View file

@ -1,3 +1,5 @@
"""Schema definitions for AgentScope message statistics."""
from pydantic import BaseModel, Field
_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH = 100
@ -5,6 +7,8 @@ _DEFAULT_MAX_FORMATTER_TEXT_LENGTH = 2000
class AsBlockStat(BaseModel):
"""Statistics and metadata for a single content block in an AgentScope message."""
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")
@ -19,10 +23,20 @@ class AsBlockStat(BaseModel):
@property
def preview(self) -> str:
"""Return a short preview of the block content."""
return self.format(_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH)
# pylint: disable=too-many-return-statements
def format(self, max_length: int = _DEFAULT_MAX_FORMATTER_TEXT_LENGTH, include_thinking: bool = True) -> str:
"""Format block content to string representation."""
"""Format block content to string representation.
Args:
max_length: Maximum length of text content in the output.
include_thinking: Whether to include thinking block content.
Returns:
Formatted string representation of the block.
"""
from ..utils import truncate_text
if self.block_type == "text":
@ -33,15 +47,17 @@ class AsBlockStat(BaseModel):
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":
if self.block_type in ("tool_use", "tool_result"):
if self.block_type == "tool_use":
return f" - tool_call={self.tool_name} params={truncate_text(self.tool_input, max_length)}"
output = truncate_text(self.tool_output, max_length)
return f" - tool_result={self.tool_name} output={output}" if output else ""
return ""
class AsMsgStat(BaseModel):
"""Statistics and metadata for a complete AgentScope message."""
name: str = Field(default=...)
role: str = Field(default="")
content: list[AsBlockStat] = Field(default_factory=list)
@ -50,10 +66,12 @@ class AsMsgStat(BaseModel):
@property
def total_tokens(self) -> int:
"""Return the total token count across all content blocks."""
return sum(block.token_count for block in self.content)
@property
def preview(self) -> str:
"""Return a short preview of the message content."""
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:

View file

@ -6,10 +6,10 @@ _token_counter = None
def get_hf_token_counter(
pretrained_model_name_or_path="Qwen/Qwen2.5-7B-Instruct",
use_mirror=True,
use_fast=True,
trust_remote_code=True,
pretrained_model_name_or_path="Qwen/Qwen2.5-7B-Instruct",
use_mirror=True,
use_fast=True,
trust_remote_code=True,
):
"""Get or initialize the global token counter instance."""
global _token_counter

View file

@ -1,3 +1,5 @@
"""Utility functions for truncating long text strings."""
from .std_logger import get_logger
logger = get_logger()

View file

@ -1,3 +1,5 @@
"""Handler for AgentScope message processing, token counting, and context management."""
import json
from agentscope.message import Msg
@ -10,6 +12,7 @@ logger = get_std_logger()
class AsMsgHandler:
"""Handles token counting, formatting, and context compaction for AgentScope messages."""
def __init__(self, token_counter: HuggingFaceTokenCounter):
self._token_counter = token_counter
@ -33,7 +36,7 @@ class AsMsgHandler:
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}")
logger.warning(f"Failed to count string tokens: {text}, e={e}")
return estimated_tokens
@staticmethod
@ -107,20 +110,24 @@ class AsMsgHandler:
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,
))
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,
))
blocks.append(
AsBlockStat(
block_type=block_type,
text=thinking,
token_count=token_count,
),
)
elif block_type in ("image", "audio", "video"):
source = block.get("source", {})
@ -131,12 +138,14 @@ class AsMsgHandler:
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,
))
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", "")
@ -146,26 +155,30 @@ class AsMsgHandler:
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,
))
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,
))
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)
@ -179,10 +192,10 @@ class AsMsgHandler:
)
def format_msgs_to_str(
self,
messages: list[Msg],
memory_compact_threshold: int,
include_thinking: bool = False,
self,
messages: list[Msg],
memory_compact_threshold: int,
include_thinking: bool = False,
) -> str:
"""Format list of messages to a single formatted string.
@ -219,10 +232,10 @@ class AsMsgHandler:
return "\n\n".join(formatted_parts)
def context_check(
self,
messages: list[Msg],
memory_compact_threshold: int,
memory_compact_reserve: int,
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.
@ -294,9 +307,7 @@ class AsMsgHandler:
# 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", "")
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

View file

@ -20,11 +20,11 @@ class ReMeInMemoryMemory(InMemoryMemory):
self._msg_handler: AsMsgHandler = AsMsgHandler(token_counter)
async def get_memory(
self,
mark: str | None = None,
exclude_mark: str | None = _MemoryMark.COMPRESSED,
prepend_summary: bool = True,
**_kwargs,
self,
mark: str | None = None,
exclude_mark: str | None = _MemoryMark.COMPRESSED,
prepend_summary: bool = True,
**_kwargs,
) -> list[Msg]:
"""Get the messages from the memory by mark (if provided).
@ -188,10 +188,10 @@ Use it as context to maintain continuity.
)
return (
f"**Conversation History**\n\n"
f"- Total messages: {stats['total_messages']}\n"
f"- Estimated tokens: {stats['estimated_tokens']}\n"
f"- Max input length: {stats['max_input_length']}\n"
f"- Context usage: {stats['context_usage_ratio']:.1f}%\n"
f"- Compressed summary tokens: {stats['compressed_summary_tokens']}\n\n" + "\n\n".join(lines)
f"**Conversation History**\n\n"
f"- Total messages: {stats['total_messages']}\n"
f"- Estimated tokens: {stats['estimated_tokens']}\n"
f"- Max input length: {stats['max_input_length']}\n"
f"- Context usage: {stats['context_usage_ratio']:.1f}%\n"
f"- Compressed summary tokens: {stats['compressed_summary_tokens']}\n\n" + "\n\n".join(lines)
)

View file

@ -15,10 +15,10 @@ class Compactor(BaseOp):
"""Compactor class for compacting memory messages."""
def __init__(
self,
memory_compact_threshold: int,
token_counter: HuggingFaceTokenCounter,
**kwargs,
self,
memory_compact_threshold: int,
token_counter: HuggingFaceTokenCounter,
**kwargs,
):
super().__init__(**kwargs)
self.memory_compact_threshold: int = memory_compact_threshold
@ -58,8 +58,9 @@ 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(

View file

@ -18,13 +18,13 @@ class Summarizer(BaseOp):
"""Summarizer class for summarizing memory messages."""
def __init__(
self,
working_dir: str,
memory_dir: str,
memory_compact_threshold: int,
token_counter: HuggingFaceTokenCounter,
toolkit: Toolkit,
**kwargs,
self,
working_dir: str,
memory_dir: str,
memory_compact_threshold: int,
token_counter: HuggingFaceTokenCounter,
toolkit: Toolkit,
**kwargs,
):
super().__init__(**kwargs)
self.working_dir: str = working_dir

View file

@ -17,11 +17,11 @@ class ToolResultCompactor(BaseOp):
"""Truncate large tool_result outputs and save full content to files."""
def __init__(
self,
tool_result_dir: str | Path,
tool_result_threshold: int,
retention_days: int = 7,
**kwargs,
self,
tool_result_dir: str | Path,
tool_result_threshold: int,
retention_days: int = 7,
**kwargs,
):
super().__init__(**kwargs)
self.tool_result_dir = Path(tool_result_dir)

View file

@ -1,14 +1,17 @@
"""memory tools"""
from .base_memory_tool import BaseMemoryTool
# chunk tools
from .chunk.memory_get import MemoryGet
from .chunk.memory_search import MemorySearch
from .delegate_task import DelegateTask
# history tools
from .history.add_history import AddHistory
from .history.read_history import ReadHistory
from .history.read_history_v2 import ReadHistoryV2
# profiles tools
from .profiles.add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles
from .profiles.add_profile import AddProfile
@ -16,6 +19,7 @@ from .profiles.delete_profile import DeleteProfile
from .profiles.read_all_profiles import ReadAllProfiles
from .profiles.update_profile import UpdateProfile
from .profiles.update_profiles_v1 import UpdateProfilesV1
# record tools
from .record.add_and_retrieve_similar_memory import AddAndRetrieveSimilarMemory
from .record.add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory

View file

@ -0,0 +1,7 @@
"""File-based memory tool implementations."""
from .file_io import FileIO
__all__ = [
"FileIO",
]

View file

@ -15,7 +15,6 @@ Key Features:
"""
import asyncio
import logging
from pathlib import Path
from agentscope.formatter import FormatterBase
@ -26,32 +25,31 @@ from agentscope.tool import Toolkit, ToolResponse
from .config import ReMeConfigParser
from .core import Application
from .core.utils import get_hf_token_counter
from .memory.file_based import Compactor, Summarizer, ToolResultCompactor, ReMeInMemoryMemory, ReMeOpenAIChatFormatter, \
FileIO
from .memory.file_based.utils import get_token_counter
from .core.utils import get_hf_token_counter, get_std_logger
from .memory.file_based import Compactor, Summarizer, ToolResultCompactor, ReMeInMemoryMemory
from .memory.tools import MemorySearch
from .memory.tools.file import FileIO
logger = logging.getLogger(__name__)
logger = get_std_logger()
class ReMeLight(Application):
"""ReMe Light Application Class"""
def __init__(
self,
working_dir: str = ".reme",
llm_api_key: str | None = None,
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
default_as_llm_config: dict | None = None,
default_embedding_model_config: dict | None = None,
default_file_store_config: dict | None = None,
vector_weight: float = 0.7,
candidate_multiplier: float = 3.0,
tool_result_threshold: int = 1000,
retention_days: int = 7,
self,
working_dir: str = ".reme",
llm_api_key: str | None = None,
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
default_as_llm_config: dict | None = None,
default_embedding_model_config: dict | None = None,
default_file_store_config: dict | None = None,
vector_weight: float = 0.7,
candidate_multiplier: float = 3.0,
tool_result_threshold: int = 1000,
retention_days: int = 7,
):
# Initialize working directory structure
self.working_path = Path(working_dir).absolute()
@ -94,6 +92,15 @@ class ReMeLight(Application):
@staticmethod
def calculate_memory_compact_threshold(max_input_length: float, compact_ratio: float) -> int:
"""Calculate the memory compaction threshold based on input length and ratio.
Args:
max_input_length: Maximum input length in tokens.
compact_ratio: Ratio of the input length to use as the threshold.
Returns:
Computed compaction threshold as an integer.
"""
return int(max_input_length * compact_ratio * 0.9)
def _cleanup_tool_results(self) -> int:
@ -156,15 +163,15 @@ class ReMeLight(Application):
return messages
async def compact_memory(
self,
messages: list[Msg],
as_llm: str | ChatModelBase = "default",
as_llm_formatter: str | FormatterBase = "default",
token_counter: HuggingFaceTokenCounter | None = None,
language: str = "zh",
max_input_length: float = 128 * 1024,
compact_ratio: float = 0.7,
previous_summary: str = "",
self,
messages: list[Msg],
as_llm: str | ChatModelBase = "default",
as_llm_formatter: str | FormatterBase = "default",
token_counter: HuggingFaceTokenCounter | None = None,
language: str = "zh",
max_input_length: float = 128 * 1024,
compact_ratio: float = 0.7,
previous_summary: str = "",
) -> str:
"""Compact a list of messages into a condensed summary."""
try:
@ -173,9 +180,9 @@ class ReMeLight(Application):
compactor = Compactor(
memory_compact_threshold=self.calculate_memory_compact_threshold(max_input_length, compact_ratio),
token_counter=token_counter,
as_llm=as_llm,
as_llm_formatter=as_llm_formatter,
token_counter=token_counter,
language=language if language == "zh" else "",
)
@ -190,58 +197,69 @@ class ReMeLight(Application):
logger.exception(f"Error compacting memory: {e}")
return ""
async def summary_memory(self, messages: list[Msg]) -> str:
async def summary_memory(
self,
messages: list[Msg],
as_llm: str | ChatModelBase = "default",
as_llm_formatter: str | FormatterBase = "default",
token_counter: HuggingFaceTokenCounter | None = None,
toolkit: Toolkit | None = None,
language: str = "zh",
max_input_length: float = 128 * 1024,
compact_ratio: float = 0.7,
) -> str:
"""Generate a comprehensive summary of the given messages."""
try:
# Create toolkit if not provided
if self.toolkit is not None:
toolkit = self.toolkit
else:
if token_counter is None:
token_counter = get_hf_token_counter()
if toolkit is None:
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),
memory_dir=str(self.memory_path),
memory_compact_threshold=self.memory_compact_threshold,
chat_model=self.chat_model,
formatter=self.formatter,
token_counter=self.token_counter,
memory_compact_threshold=self.calculate_memory_compact_threshold(max_input_length, compact_ratio),
token_counter=token_counter,
toolkit=toolkit,
language=self.language,
as_llm=as_llm,
as_llm_formatter=as_llm_formatter,
language=language if language == "zh" else "",
)
# Execute summarization on the provided messages
return await summarizer.call(messages=messages, service_context=self.service_context)
except Exception as e:
# Log error and return empty string to indicate failure
logger.exception(f"Error summarizing memory: {e}")
return ""
def add_async_summary_task(self, messages: list[Msg], **kwargs):
"""Add an asynchronous summary task for the given messages."""
remaining_tasks = []
for task in self.summary_tasks:
if task.done():
if task.cancelled():
logger.warning("Summary task was cancelled.")
continue
exc = task.exception()
if exc is not None:
logger.error(f"Summary task failed: {exc}")
else:
result = task.result()
logger.info(f"Summary task completed: {result}")
else:
remaining_tasks.append(task)
self.summary_tasks = remaining_tasks
task = asyncio.create_task(self.summary_memory(messages=messages, **kwargs))
self.summary_tasks.append(task)
async def await_summary_tasks(self) -> str:
"""
Wait for all background summary tasks to complete and collect results.
This method iterates through all pending summary tasks, waits for their
completion, and collects their results or error information. It's used
to synchronize with background summarization operations before shutdown
or when results are needed.
Returns:
str: A concatenated string containing the status and results of
all summary tasks, with each task on a new line
Note:
- Completed tasks are processed immediately without waiting
- Incomplete tasks are awaited with a timeout
- Cancelled tasks and exceptions are logged and included in results
- The task list is cleared after processing all tasks
"""
"""Wait for all background summary tasks to complete and collect results."""
result = ""
for task in self.summary_tasks:
if task.done():
@ -279,48 +297,6 @@ class ReMeLight(Application):
self.summary_tasks.clear()
return result
def add_async_summary_task(self, messages: list[Msg]):
"""
Add an asynchronous summary task for the given messages.
This method creates a background task to summarize the provided messages
without blocking the main execution flow. Before adding a new task, it
cleans up any completed tasks from the task list to prevent memory leaks.
Args:
messages (list[Msg]): The list of messages to be summarized in the
background task
Note:
- Completed tasks are removed from the tracking list before adding
- Task status (success, failure, cancellation) is logged for monitoring
- The new task is created using asyncio.create_task for true async execution
- Failed or cancelled tasks are logged but do not prevent new tasks
"""
# Clean up completed summary tasks before adding a new one
remaining_tasks = []
for task in self.summary_tasks:
if task.done():
# Process completed task status
if task.cancelled():
logger.warning("Summary task was cancelled.")
continue
exc = task.exception()
if exc is not None:
logger.error(f"Summary task failed: {exc}")
else:
# Log successful completion with result summary
result = task.result()
logger.info(f"Summary task completed: {result}")
else:
# Keep incomplete tasks in the tracking list
remaining_tasks.append(task)
self.summary_tasks = remaining_tasks
# Create and track the new background summarization task
task = asyncio.create_task(self.summary_memory(messages=messages))
self.summary_tasks.append(task)
async def memory_search(self, query: str, max_results: int = 5, min_score: float = 0.1) -> ToolResponse:
"""
Perform semantic memory search using vector and full-text search.

View file

@ -1,7 +1,6 @@
"""Tests for Compactor."""
import asyncio
import logging
from agentscope.message import Msg
@ -10,14 +9,10 @@ from test_utils import (
get_formatter,
get_token_counter,
)
from reme.core.utils import get_std_logger
from reme.memory.file_based import Compactor
# 配置日志输出到控制台
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
logger = get_std_logger()
# ANSI 颜色码

View file

@ -1,18 +1,12 @@
"""Tests for AsMsgHandler.context_check method."""
import logging
from agentscope.message import Msg
from test_utils import get_token_counter
from reme.core.utils import get_std_logger
from reme.memory.file_based.as_msg_handler import AsMsgHandler
# Configure logging
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
logger = get_std_logger()
# ANSI color codes
@ -101,8 +95,7 @@ def verify_context_check_invariants(
# 2. Reserve requirement check
kept_tokens = sum(handler.stat_message(m).total_tokens for m in to_keep)
assert kept_tokens <= memory_compact_reserve or len(to_keep) == 0, (
f"[{test_name}] Reserve violation: kept_tokens ({kept_tokens}) > "
f"reserve ({memory_compact_reserve})"
f"[{test_name}] Reserve violation: kept_tokens ({kept_tokens}) > " f"reserve ({memory_compact_reserve})"
)
# 3. Order requirement check - both lists should preserve original order
@ -143,9 +136,7 @@ def verify_context_check_invariants(
all_returned = set(id(m) for m in to_compact) | set(id(m) for m in to_keep)
all_original = set(id(m) for m in messages)
assert all_returned == all_original, (
f"[{test_name}] Message set mismatch: returned messages differ from original"
)
assert all_returned == all_original, f"[{test_name}] Message set mismatch: returned messages differ from original"
def create_user_msg(content: str) -> Msg:
@ -234,7 +225,7 @@ def test_empty_messages():
memory_compact_threshold=threshold,
memory_compact_reserve=reserve,
)
assert to_compact == [], f"Expected empty compact list, got: {to_compact}"
assert not to_compact, f"Expected empty compact list, got: {to_compact}"
assert to_keep == [], f"Expected empty keep list, got: {to_keep}"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_empty_messages")
print_pass("test_empty_messages")
@ -254,10 +245,18 @@ def test_below_threshold_returns_all():
memory_compact_threshold=threshold, # Very high threshold
memory_compact_reserve=reserve,
)
assert to_compact == [], f"Expected empty compact list, got: {len(to_compact)}"
assert not to_compact, f"Expected empty compact list, got: {len(to_compact)}"
assert len(to_keep) == 3, f"Expected 3 messages to keep, got: {len(to_keep)}"
assert to_keep == messages, "Messages to keep should be the original messages"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_below_threshold_returns_all")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_below_threshold_returns_all",
)
print_pass("test_below_threshold_returns_all")
@ -280,7 +279,15 @@ def test_above_threshold_triggers_compaction():
# Should have some messages compacted and some kept
assert len(to_compact) + len(to_keep) == len(messages), "Total messages should match"
assert len(to_compact) > 0, "Expected some messages to be compacted"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_above_threshold_triggers_compaction")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_above_threshold_triggers_compaction",
)
print_pass("test_above_threshold_triggers_compaction")
@ -304,7 +311,15 @@ def test_message_order_preserved():
all_messages = to_compact + to_keep
for i, msg in enumerate(all_messages):
assert msg in messages, f"Message {i} not found in original messages"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_message_order_preserved")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_message_order_preserved",
)
print_pass("test_message_order_preserved")
@ -323,9 +338,17 @@ def test_single_message_below_threshold():
memory_compact_threshold=threshold,
memory_compact_reserve=reserve,
)
assert to_compact == [], "Should not compact single message below threshold"
assert not to_compact, "Should not compact single message below threshold"
assert len(to_keep) == 1, "Should keep the single message"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_single_message_below_threshold")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_single_message_below_threshold",
)
print_pass("test_single_message_below_threshold")
@ -343,7 +366,15 @@ def test_single_message_above_threshold():
# Message exceeds both threshold and reserve, so it's compacted
assert len(to_compact) == 1, "Single large message should be compacted"
assert len(to_keep) == 0, "Nothing can fit in reserve"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_single_message_above_threshold")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_single_message_above_threshold",
)
print_pass("test_single_message_above_threshold")
@ -388,12 +419,12 @@ def test_exact_threshold_boundary():
"""Test messages exactly at threshold boundary."""
handler = create_handler()
messages = [create_user_msg("Test message")]
# Get exact token count
stat = handler.stat_message(messages[0])
exact_tokens = stat.total_tokens
threshold, reserve = exact_tokens, exact_tokens
# Test at exact boundary
to_compact, to_keep = handler.context_check(
messages=messages,
@ -401,9 +432,17 @@ def test_exact_threshold_boundary():
memory_compact_reserve=reserve,
)
# At exact boundary (<=), should not trigger compaction
assert to_compact == [], "Should not compact at exact boundary"
assert not to_compact, "Should not compact at exact boundary"
assert len(to_keep) == 1, "Should keep message at exact boundary"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_exact_threshold_boundary")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_exact_threshold_boundary",
)
print_pass("test_exact_threshold_boundary")
@ -423,7 +462,15 @@ def test_reserve_larger_than_threshold():
# Compaction triggered but reserve can hold everything
# Total messages should be preserved
assert len(to_compact) + len(to_keep) == 2
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_reserve_larger_than_threshold")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_reserve_larger_than_threshold",
)
print_pass("test_reserve_larger_than_threshold")
@ -447,20 +494,22 @@ def test_tool_use_result_paired():
memory_compact_threshold=threshold, # Trigger compaction
memory_compact_reserve=reserve, # Enough for tool pair
)
# If tool_result is kept, tool_use should also be kept
tool_result_in_keep = any(
any(b.get("type") == "tool_result" for b in m.get_content_blocks())
for m in to_keep
)
tool_use_in_keep = any(
any(b.get("type") == "tool_use" for b in m.get_content_blocks())
for m in to_keep
)
tool_result_in_keep = any(any(b.get("type") == "tool_result" for b in m.get_content_blocks()) for m in to_keep)
tool_use_in_keep = any(any(b.get("type") == "tool_use" for b in m.get_content_blocks()) for m in to_keep)
if tool_result_in_keep:
assert tool_use_in_keep, "tool_use should be kept when tool_result is kept"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_tool_use_result_paired")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_tool_use_result_paired",
)
print_pass("test_tool_use_result_paired")
@ -480,7 +529,15 @@ def test_tool_use_without_result():
)
# Should not crash, just process normally
assert len(to_compact) + len(to_keep) == 3
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_tool_use_without_result")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_tool_use_without_result",
)
print_pass("test_tool_use_without_result")
@ -500,7 +557,15 @@ def test_tool_result_without_use():
)
# Should not crash even with orphan tool_result
assert len(to_compact) + len(to_keep) == 3
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_tool_result_without_use")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_tool_result_without_use",
)
print_pass("test_tool_result_without_use")
@ -523,7 +588,7 @@ def test_multiple_tool_pairs():
memory_compact_threshold=threshold,
memory_compact_reserve=reserve,
)
# Verify tool pairs integrity - for each kept tool_result, its tool_use should be kept
for msg in to_keep:
for block in msg.get_content_blocks("tool_result"):
@ -537,7 +602,15 @@ def test_multiple_tool_pairs():
tool_use_found = True
break
assert tool_use_found, f"tool_use for {tool_id} should be kept with tool_result"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_multiple_tool_pairs")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_multiple_tool_pairs",
)
print_pass("test_multiple_tool_pairs")
@ -552,7 +625,7 @@ def test_tool_dependency_causes_extra_inclusion():
messages = [
create_user_msg("Start " * 100), # Large message
create_tool_use_msg("call_dep", "dep_tool", large_tool_input), # Medium
create_user_msg("Middle " * 100), # Large message
create_user_msg("Middle " * 100), # Large message
create_tool_result_msg("call_dep", "dep_tool", "Result"), # Small
create_assistant_msg("End"), # Small
]
@ -562,22 +635,27 @@ def test_tool_dependency_causes_extra_inclusion():
memory_compact_threshold=threshold, # Trigger compaction
memory_compact_reserve=reserve, # Medium reserve
)
# Check pair integrity
result_kept = any(
any(b.get("id") == "call_dep" and b.get("type") == "tool_result"
for b in m.get_content_blocks())
any(b.get("id") == "call_dep" and b.get("type") == "tool_result" for b in m.get_content_blocks())
for m in to_keep
)
use_kept = any(
any(b.get("id") == "call_dep" and b.get("type") == "tool_use"
for b in m.get_content_blocks())
for m in to_keep
any(b.get("id") == "call_dep" and b.get("type") == "tool_use" for b in m.get_content_blocks()) for m in to_keep
)
if result_kept:
assert use_kept, "Dependent tool_use should be included with tool_result"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_tool_dependency_causes_extra_inclusion")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_tool_dependency_causes_extra_inclusion",
)
print_pass("test_tool_dependency_causes_extra_inclusion")
@ -598,24 +676,30 @@ def test_tool_dependency_exceeds_reserve():
memory_compact_threshold=threshold, # Trigger compaction
memory_compact_reserve=reserve, # Small reserve - can't fit the pair
)
# The tool pair is too large, so it should be excluded or partially handled
# Either both are compacted (pair excluded) or neither is kept
result_kept = any(
any(b.get("id") == "call_big" and b.get("type") == "tool_result"
for b in m.get_content_blocks())
any(b.get("id") == "call_big" and b.get("type") == "tool_result" for b in m.get_content_blocks())
for m in to_keep
)
if result_kept:
# If result is kept, use must also be kept (pair integrity)
use_kept = any(
any(b.get("id") == "call_big" and b.get("type") == "tool_use"
for b in m.get_content_blocks())
any(b.get("id") == "call_big" and b.get("type") == "tool_use" for b in m.get_content_blocks())
for m in to_keep
)
assert use_kept, "Pair integrity violated"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_tool_dependency_exceeds_reserve")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_tool_dependency_exceeds_reserve",
)
print_pass("test_tool_dependency_exceeds_reserve")
@ -636,19 +720,26 @@ def test_interleaved_tool_pairs():
memory_compact_threshold=threshold,
memory_compact_reserve=reserve,
)
# Verify pair integrity for interleaved pairs
for msg in to_keep:
for block in msg.get_content_blocks("tool_result"):
tool_id = block.get("id", "")
if tool_id:
use_found = any(
any(ub.get("id") == tool_id and ub.get("type") == "tool_use"
for ub in km.get_content_blocks())
any(ub.get("id") == tool_id and ub.get("type") == "tool_use" for ub in km.get_content_blocks())
for km in to_keep
)
assert use_found, f"Interleaved tool_use {tool_id} should be kept"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_interleaved_tool_pairs")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_interleaved_tool_pairs",
)
print_pass("test_interleaved_tool_pairs")
@ -671,7 +762,15 @@ def test_message_with_empty_content():
memory_compact_reserve=reserve,
)
assert len(to_compact) + len(to_keep) == 2
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_message_with_empty_content")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_message_with_empty_content",
)
print_pass("test_message_with_empty_content")
@ -689,7 +788,15 @@ def test_message_with_whitespace_only():
memory_compact_reserve=reserve,
)
assert len(to_compact) + len(to_keep) == 2
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_message_with_whitespace_only")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_message_with_whitespace_only",
)
print_pass("test_message_with_whitespace_only")
@ -706,7 +813,15 @@ def test_very_long_single_message():
)
# Single huge message - either kept alone or compacted
assert len(to_compact) + len(to_keep) == 1
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_very_long_single_message")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_very_long_single_message",
)
print_pass("test_very_long_single_message")
@ -723,7 +838,15 @@ def test_many_small_messages():
# Should compact older messages and keep recent ones
assert len(to_compact) + len(to_keep) == 100
assert len(to_keep) > 0, "Should keep some messages"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_many_small_messages")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_many_small_messages",
)
print_pass("test_many_small_messages")
@ -760,7 +883,15 @@ def test_special_characters_content():
memory_compact_reserve=reserve,
)
assert len(to_compact) + len(to_keep) == 2
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_special_characters_content")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_special_characters_content",
)
print_pass("test_special_characters_content")
@ -776,11 +907,11 @@ def test_all_messages_fit_exactly_in_reserve():
create_user_msg("Message 1"),
create_assistant_msg("Message 2"),
]
# Calculate total tokens
total = sum(handler.stat_message(m).total_tokens for m in messages)
threshold, reserve = total - 1, total
to_compact, to_keep = handler.context_check(
messages=messages,
memory_compact_threshold=threshold, # Just below total to trigger
@ -788,7 +919,15 @@ def test_all_messages_fit_exactly_in_reserve():
)
# All should be kept since reserve can hold everything
assert len(to_keep) == 2, f"All messages should fit in reserve, got {len(to_keep)}"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_all_messages_fit_exactly_in_reserve")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_all_messages_fit_exactly_in_reserve",
)
print_pass("test_all_messages_fit_exactly_in_reserve")
@ -800,20 +939,28 @@ def test_first_message_only_compacted():
create_assistant_msg("Small"), # Small
create_user_msg("Tiny"), # Tiny
]
# Calculate tokens to set appropriate reserve
small_msg_tokens = handler.stat_message(messages[1]).total_tokens
tiny_msg_tokens = handler.stat_message(messages[2]).total_tokens
threshold, reserve = 50, small_msg_tokens + tiny_msg_tokens + 10
to_compact, to_keep = handler.context_check(
messages=messages,
memory_compact_threshold=threshold, # Low to trigger
memory_compact_reserve=reserve, # Fits last 2
)
assert len(to_compact) >= 1, "At least first message should be compacted"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_first_message_only_compacted")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_first_message_only_compacted",
)
print_pass("test_first_message_only_compacted")
@ -825,20 +972,28 @@ def test_last_message_only_kept():
create_assistant_msg("Large " * 200),
create_user_msg("Tiny"), # Only this fits
]
tiny_tokens = handler.stat_message(messages[2]).total_tokens
threshold, reserve = 10, tiny_tokens + 5
to_compact, to_keep = handler.context_check(
messages=messages,
memory_compact_threshold=threshold,
memory_compact_reserve=reserve, # Only fits last message
)
if len(to_keep) == 1:
# Last message should be the one kept
assert to_keep[0] == messages[2], "Only last message should be kept"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_last_message_only_kept")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_last_message_only_kept",
)
print_pass("test_last_message_only_kept")
@ -857,7 +1012,15 @@ def test_all_messages_compacted():
)
assert len(to_compact) == 2, "All messages should be compacted"
assert len(to_keep) == 0, "No messages should be kept"
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_all_messages_compacted")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_all_messages_compacted",
)
print_pass("test_all_messages_compacted")
@ -929,7 +1092,15 @@ def test_tool_use_with_empty_id():
)
# Should handle gracefully
assert len(to_compact) + len(to_keep) == 3
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_tool_use_with_empty_id")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_tool_use_with_empty_id",
)
print_pass("test_tool_use_with_empty_id")
@ -949,7 +1120,15 @@ def test_tool_result_with_empty_id():
)
# Should handle gracefully
assert len(to_compact) + len(to_keep) == 3
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_tool_result_with_empty_id")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_tool_result_with_empty_id",
)
print_pass("test_tool_result_with_empty_id")
@ -970,7 +1149,15 @@ def test_duplicate_tool_ids():
)
# Should not crash with duplicate IDs
assert len(to_compact) + len(to_keep) == 4
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_duplicate_tool_ids")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_duplicate_tool_ids",
)
print_pass("test_duplicate_tool_ids")
@ -1000,7 +1187,15 @@ def test_message_with_multiple_tool_blocks():
memory_compact_reserve=reserve,
)
assert len(to_compact) + len(to_keep) == 5
verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_message_with_multiple_tool_blocks")
verify_context_check_invariants(
handler,
messages,
to_compact,
to_keep,
threshold,
reserve,
"test_message_with_multiple_tool_blocks",
)
print_pass("test_message_with_multiple_tool_blocks")

View file

@ -2,19 +2,15 @@
# pylint: disable=W0212
import logging
import sys
from agentscope.message import Msg
from test_utils import get_token_counter
from reme.core.utils import get_std_logger
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__)
logger = get_std_logger()
# ANSI 颜色码
@ -72,7 +68,7 @@ def verify_result_within_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:
@ -93,7 +89,7 @@ def verify_result_within_threshold(
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
_ = 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
@ -102,21 +98,21 @@ def verify_result_within_threshold(
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:
if 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:
if 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})."
)
assert (
included_tokens <= threshold + 1
), f"{test_name}: Included messages token count ({included_tokens}) exceeds threshold ({threshold})."
def create_user_msg(content: str) -> Msg:
@ -199,12 +195,14 @@ def create_mixed_content_msg(
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 {},
})
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)
@ -274,8 +272,7 @@ def test_format_msgs_to_str_message_order():
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}"
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")
@ -347,9 +344,7 @@ def test_format_msgs_to_str_thinking_excluded_by_default():
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 "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")
@ -362,9 +357,7 @@ def test_format_msgs_to_str_thinking_included():
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 "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")
@ -375,14 +368,18 @@ def test_format_msgs_to_str_thinking_only_message():
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
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
msgs,
memory_compact_threshold=threshold,
include_thinking=True,
)
assert "Deep thoughts here" not in result_no_thinking
@ -425,9 +422,9 @@ def test_format_msgs_to_str_exceeds_threshold_truncate_older():
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]}..."
)
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)
@ -517,10 +514,7 @@ 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)
]
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)
@ -604,19 +598,25 @@ 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",
)]
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
msgs,
memory_compact_threshold=threshold,
include_thinking=False,
)
result_with_thinking = handler.format_msgs_to_str(
msgs, memory_compact_threshold=threshold, include_thinking=True
msgs,
memory_compact_threshold=threshold,
include_thinking=True,
)
assert "Text content" in result_no_thinking
@ -682,7 +682,7 @@ def 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):
@ -690,16 +690,14 @@ def test_format_msgs_to_str_incremental_threshold_check():
# 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"
)
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)
@ -743,17 +741,21 @@ 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
},
}],
)]
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
@ -765,14 +767,16 @@ 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"}},
],
)]
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
@ -785,14 +789,16 @@ 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"},
],
)]
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
@ -880,4 +886,4 @@ def run_all_tests():
if __name__ == "__main__":
success = run_all_tests()
exit(0 if success else 1)
sys.exit(0 if success else 1)

View file

@ -2,19 +2,13 @@
# pylint: disable=W0212
import logging
from agentscope.message import Msg
from test_utils import get_token_counter
from reme.core.utils import get_std_logger
from reme.memory.file_based import MemoryFormatter
# 配置日志输出到控制台
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
logger = get_std_logger()
# ANSI 颜色码

View file

@ -3,6 +3,7 @@
import asyncio
from agentscope.message import Msg
from reme.reme_light import ReMeLight
@ -127,9 +128,6 @@ async def main():
# 初始化 ReMeLight
reme = ReMeLight(
working_dir=".reme", # 记忆文件存储目录
max_input_length=128000, # 模型上下文窗口(tokens)
memory_compact_ratio=0.7, # 达到 max_input_length * 0.7 时触发压缩
language="zh", # 摘要语言(zh / "")
tool_result_threshold=1000, # 超过此字符数的工具输出自动转存
retention_days=7, # tool_result/ 文件保留天数
)
@ -176,7 +174,7 @@ async def main():
# 将消息添加到内存中以便估算
for msg in messages:
await memory.add(msg)
token_stats = await memory.estimate_tokens()
token_stats = await memory.estimate_tokens(max_input_length=128000)
print(f"当前上下文使用率: {token_stats['context_usage_ratio']:.1f}%")
print(f"消息 Token 数: {token_stats['messages_tokens']}")
print(f"预估总 Token 数: {token_stats['estimated_tokens']}")

View file

@ -2,7 +2,6 @@
import asyncio
import datetime
import logging
import tempfile
from pathlib import Path
@ -13,14 +12,10 @@ from test_utils import (
get_formatter,
get_token_counter,
)
from reme.core.utils import get_std_logger
from reme.memory.file_based import Summarizer
# 配置日志输出到控制台
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
logger = get_std_logger()
# ANSI 颜色码

View file

@ -6,7 +6,6 @@ from datetime import datetime, timedelta
from pathlib import Path
from agentscope.message import Msg
from reme.memory.file_based.tool_result_compactor import ToolResultCompactor
from reme.memory.file_based.utils import TRUNCATION_MARKER_START

View file

@ -1,50 +1,13 @@
"""Test utilities for copaw tests."""
import os
from pathlib import Path
from typing import Any
from loguru import logger
_token_counter = None
def get_token_counter():
"""Get or initialize the global token counter instance.
"""Get HF token counter instance."""
from reme.core.utils import get_hf_token_counter
Returns:
TokenCounterBase: The token counter instance for Qwen models.
Raises:
RuntimeError: If token counter initialization fails.
"""
global _token_counter
if _token_counter is None:
from agentscope.token import HuggingFaceTokenCounter
# Use Qwen tokenizer for DashScope models
# Qwen3 series uses the same tokenizer as Qwen2.5
# Try local tokenizer first, fall back to online if not found
local_tokenizer_path = Path(__file__).parent.parent.parent / "tokenizer"
if local_tokenizer_path.exists() and (local_tokenizer_path / "tokenizer.json").exists():
tokenizer_path = str(local_tokenizer_path)
logger.info(f"Using local Qwen tokenizer from {tokenizer_path}")
else:
tokenizer_path = "Qwen/Qwen2.5-7B-Instruct"
logger.info(
"Local tokenizer not found, downloading from HuggingFace",
)
_token_counter = HuggingFaceTokenCounter(
pretrained_model_name_or_path=tokenizer_path,
use_mirror=True, # Use HF mirror for users in China
use_fast=True,
trust_remote_code=True,
)
logger.debug("Token counter initialized with Qwen tokenizer")
return _token_counter
return get_hf_token_counter()
def get_dash_chat_model(model_name: str = "qwen3.5-plus"):
@ -54,8 +17,8 @@ def get_dash_chat_model(model_name: str = "qwen3.5-plus"):
load_env()
return OpenAIChatModel(
api_key=os.environ["REME_LLM_API_KEY"],
client_kwargs={"base_url": os.environ["REME_LLM_BASE_URL"]},
api_key=os.environ["LLM_API_KEY"],
client_kwargs={"base_url": os.environ["LLM_BASE_URL"]},
model_name=model_name,
)
@ -63,27 +26,5 @@ def get_dash_chat_model(model_name: str = "qwen3.5-plus"):
def get_formatter():
"""Get formatter instance."""
from agentscope.formatter import OpenAIChatFormatter
from agentscope.token import HuggingFaceTokenCounter
from reme.memory.file_based.utils import _extract_text_from_messages
class ReMeChatFormatter(OpenAIChatFormatter):
"""ReMe chat formatter class."""
async def _count(self, msgs: list[dict[str, Any]]) -> int | None:
"""Count the number of tokens in the input messages. If token counter
is not provided, `None` will be returned.
Args:
msgs (`list[Msg]`):
The input messages to count tokens for.
"""
if self.token_counter is None:
return None
assert isinstance(self.token_counter, HuggingFaceTokenCounter)
text = _extract_text_from_messages(msgs)
token_ids = self.token_counter.tokenizer.encode(text)
token_count = len(token_ids)
return token_count
return ReMeChatFormatter(token_counter=get_token_counter())
return OpenAIChatFormatter()