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()