diff --git a/reme/core/__init__.py b/reme/core/__init__.py
index 88e42175..053755cc 100644
--- a/reme/core/__init__.py
+++ b/reme/core/__init__.py
@@ -1,4 +1,5 @@
"""Core"""
+
from . import as_llm
from . import as_llm_formatter
from . import embedding
diff --git a/reme/core/application.py b/reme/core/application.py
index bc59c7f3..46f4a934 100644
--- a/reme/core/application.py
+++ b/reme/core/application.py
@@ -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
diff --git a/reme/core/as_llm/__init__.py b/reme/core/as_llm/__init__.py
index 888048e6..9cf527af 100644
--- a/reme/core/as_llm/__init__.py
+++ b/reme/core/as_llm/__init__.py
@@ -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)
diff --git a/reme/core/as_llm_formatter/__init__.py b/reme/core/as_llm_formatter/__init__.py
index 9c3a52cf..88b326a7 100644
--- a/reme/core/as_llm_formatter/__init__.py
+++ b/reme/core/as_llm_formatter/__init__.py
@@ -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)
diff --git a/reme/core/op/base_op.py b/reme/core/op/base_op.py
index 0f0580a8..86cfa3be 100644
--- a/reme/core/op/base_op.py
+++ b/reme/core/op/base_op.py
@@ -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."""
diff --git a/reme/core/schema/as_msg_stat.py b/reme/core/schema/as_msg_stat.py
index 1acb9e8a..4bb69f99 100644
--- a/reme/core/schema/as_msg_stat.py
+++ b/reme/core/schema/as_msg_stat.py
@@ -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:
diff --git a/reme/core/utils/hf_token_counter_utils.py b/reme/core/utils/hf_token_counter_utils.py
index dfbc3dc6..a8ab348c 100644
--- a/reme/core/utils/hf_token_counter_utils.py
+++ b/reme/core/utils/hf_token_counter_utils.py
@@ -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
diff --git a/reme/core/utils/truncate_text_utils.py b/reme/core/utils/truncate_text_utils.py
index ec85ec61..da0c473a 100644
--- a/reme/core/utils/truncate_text_utils.py
+++ b/reme/core/utils/truncate_text_utils.py
@@ -1,3 +1,5 @@
+"""Utility functions for truncating long text strings."""
+
from .std_logger import get_logger
logger = get_logger()
diff --git a/reme/memory/file_based/as_msg_handler.py b/reme/memory/file_based/as_msg_handler.py
index e967f81e..802157ab 100644
--- a/reme/memory/file_based/as_msg_handler.py
+++ b/reme/memory/file_based/as_msg_handler.py
@@ -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
diff --git a/reme/memory/file_based/reme_in_memory_memory.py b/reme/memory/file_based/reme_in_memory_memory.py
index f08ee98e..16f18726 100644
--- a/reme/memory/file_based/reme_in_memory_memory.py
+++ b/reme/memory/file_based/reme_in_memory_memory.py
@@ -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)
)
diff --git a/reme/memory/file_based/sub_agent/compactor.py b/reme/memory/file_based/sub_agent/compactor.py
index 571db449..3292c874 100644
--- a/reme/memory/file_based/sub_agent/compactor.py
+++ b/reme/memory/file_based/sub_agent/compactor.py
@@ -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"\n{history_formatted_str}\n\n\n" \
- + self.get_prompt("initial_user_message")
+ user_message: str = f"\n{history_formatted_str}\n\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(
diff --git a/reme/memory/file_based/sub_agent/summarizer.py b/reme/memory/file_based/sub_agent/summarizer.py
index e063757e..db3522da 100644
--- a/reme/memory/file_based/sub_agent/summarizer.py
+++ b/reme/memory/file_based/sub_agent/summarizer.py
@@ -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
diff --git a/reme/memory/file_based/sub_agent/tool_result_compactor.py b/reme/memory/file_based/sub_agent/tool_result_compactor.py
index 3b49c504..412df6de 100644
--- a/reme/memory/file_based/sub_agent/tool_result_compactor.py
+++ b/reme/memory/file_based/sub_agent/tool_result_compactor.py
@@ -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)
diff --git a/reme/memory/tools/__init__.py b/reme/memory/tools/__init__.py
index 2ad25ef0..af9b851f 100644
--- a/reme/memory/tools/__init__.py
+++ b/reme/memory/tools/__init__.py
@@ -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
diff --git a/reme/memory/tools/file/__init__.py b/reme/memory/tools/file/__init__.py
index e69de29b..8234e60d 100644
--- a/reme/memory/tools/file/__init__.py
+++ b/reme/memory/tools/file/__init__.py
@@ -0,0 +1,7 @@
+"""File-based memory tool implementations."""
+
+from .file_io import FileIO
+
+__all__ = [
+ "FileIO",
+]
diff --git a/reme/reme_light.py b/reme/reme_light.py
index 5c435410..82a11850 100644
--- a/reme/reme_light.py
+++ b/reme/reme_light.py
@@ -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.
diff --git a/tests/light/test_compactor.py b/tests/light/test_compactor.py
index 32dfd9d3..19891e29 100644
--- a/tests/light/test_compactor.py
+++ b/tests/light/test_compactor.py
@@ -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 颜色码
diff --git a/tests/light/test_context_check.py b/tests/light/test_context_check.py
index 300f0b61..f65f961e 100644
--- a/tests/light/test_context_check.py
+++ b/tests/light/test_context_check.py
@@ -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")
diff --git a/tests/light/test_format_msgs_to_str.py b/tests/light/test_format_msgs_to_str.py
index 29e1b2cb..bd69751a 100644
--- a/tests/light/test_format_msgs_to_str.py
+++ b/tests/light/test_format_msgs_to_str.py
@@ -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 "" 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)
diff --git a/tests/light/test_memory_formatter.py b/tests/light/test_memory_formatter.py
index 00f45bb1..8b31718b 100644
--- a/tests/light/test_memory_formatter.py
+++ b/tests/light/test_memory_formatter.py
@@ -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 颜色码
diff --git a/tests/light/test_reme_light.py b/tests/light/test_reme_light.py
index 9e2ac70a..49102826 100644
--- a/tests/light/test_reme_light.py
+++ b/tests/light/test_reme_light.py
@@ -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']}")
diff --git a/tests/light/test_summarizer.py b/tests/light/test_summarizer.py
index 560efabd..bf8e3a78 100644
--- a/tests/light/test_summarizer.py
+++ b/tests/light/test_summarizer.py
@@ -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 颜色码
diff --git a/tests/light/test_tool_result_compactor.py b/tests/light/test_tool_result_compactor.py
index 059e626e..ef97558a 100644
--- a/tests/light/test_tool_result_compactor.py
+++ b/tests/light/test_tool_result_compactor.py
@@ -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
diff --git a/tests/light/test_utils.py b/tests/light/test_utils.py
index f5cae021..fc93009e 100644
--- a/tests/light/test_utils.py
+++ b/tests/light/test_utils.py
@@ -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()