mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
refactor(core): update registry registration syntax and improve code formatting
This commit is contained in:
parent
57f8a7b42c
commit
fc7b1cdba8
24 changed files with 616 additions and 468 deletions
|
|
@ -1,4 +1,5 @@
|
|||
"""Core"""
|
||||
|
||||
from . import as_llm
|
||||
from . import as_llm_formatter
|
||||
from . import embedding
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
"""Utility functions for truncating long text strings."""
|
||||
|
||||
from .std_logger import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -15,10 +15,10 @@ class Compactor(BaseOp):
|
|||
"""Compactor class for compacting memory messages."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
memory_compact_threshold: int,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
**kwargs,
|
||||
self,
|
||||
memory_compact_threshold: int,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.memory_compact_threshold: int = memory_compact_threshold
|
||||
|
|
@ -58,8 +58,9 @@ class Compactor(BaseOp):
|
|||
f"{suffix}"
|
||||
)
|
||||
else:
|
||||
user_message: str = f"<conversation>\n{history_formatted_str}\n</conversation>\n\n" \
|
||||
+ self.get_prompt("initial_user_message")
|
||||
user_message: str = f"<conversation>\n{history_formatted_str}\n</conversation>\n\n" + self.get_prompt(
|
||||
"initial_user_message",
|
||||
)
|
||||
logger.info(f"Compactor sys_prompt={agent.sys_prompt} user_message={user_message}")
|
||||
|
||||
compact_msg: Msg = await agent.reply(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -0,0 +1,7 @@
|
|||
"""File-based memory tool implementations."""
|
||||
|
||||
from .file_io import FileIO
|
||||
|
||||
__all__ = [
|
||||
"FileIO",
|
||||
]
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 颜色码
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2,19 +2,15 @@
|
|||
|
||||
# pylint: disable=W0212
|
||||
|
||||
import logging
|
||||
import sys
|
||||
|
||||
from agentscope.message import Msg
|
||||
|
||||
from test_utils import get_token_counter
|
||||
from reme.core.utils import get_std_logger
|
||||
from reme.memory.file_based.as_msg_handler import AsMsgHandler
|
||||
|
||||
# 配置日志输出到控制台
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
logger = get_std_logger()
|
||||
|
||||
|
||||
# ANSI 颜色码
|
||||
|
|
@ -72,7 +68,7 @@ def verify_result_within_threshold(
|
|||
Note: The format_msgs_to_str method uses message token statistics (not formatted
|
||||
string tokens) for threshold checking. The formatted result may have more tokens
|
||||
than the threshold due to added metadata (timestamps, role prefixes, etc.).
|
||||
|
||||
|
||||
This verification checks that included messages' original token sum <= threshold.
|
||||
|
||||
Args:
|
||||
|
|
@ -93,7 +89,7 @@ def verify_result_within_threshold(
|
|||
for msg in msgs:
|
||||
stat = handler.stat_message(msg)
|
||||
# Check if this message's content appears in the result
|
||||
formatted = stat.format(include_thinking=True) # Use True to check all content
|
||||
_ = stat.format(include_thinking=True) # Use True to check all content
|
||||
# Simple heuristic: if the message content is in result, count its tokens
|
||||
content_blocks = msg.get_content_blocks()
|
||||
msg_included = False
|
||||
|
|
@ -102,21 +98,21 @@ def verify_result_within_threshold(
|
|||
if block_type == "text" and block.get("text", "") in result:
|
||||
msg_included = True
|
||||
break
|
||||
elif block_type == "tool_use" and f"tool_call={block.get('name', '')}" in result:
|
||||
if block_type == "tool_use" and f"tool_call={block.get('name', '')}" in result:
|
||||
msg_included = True
|
||||
break
|
||||
elif block_type == "tool_result" and f"tool_result={block.get('name', '')}" in result:
|
||||
if block_type == "tool_result" and f"tool_result={block.get('name', '')}" in result:
|
||||
msg_included = True
|
||||
break
|
||||
|
||||
|
||||
if msg_included:
|
||||
included_tokens += stat.total_tokens
|
||||
|
||||
# Verify included messages' token sum doesn't exceed threshold
|
||||
# Allow small tolerance for edge cases
|
||||
assert included_tokens <= threshold + 1, (
|
||||
f"{test_name}: Included messages token count ({included_tokens}) exceeds threshold ({threshold})."
|
||||
)
|
||||
assert (
|
||||
included_tokens <= threshold + 1
|
||||
), f"{test_name}: Included messages token count ({included_tokens}) exceeds threshold ({threshold})."
|
||||
|
||||
|
||||
def create_user_msg(content: str) -> Msg:
|
||||
|
|
@ -199,12 +195,14 @@ def create_mixed_content_msg(
|
|||
if text:
|
||||
content.append({"type": "text", "text": text})
|
||||
if tool_name:
|
||||
content.append({
|
||||
"type": "tool_use",
|
||||
"id": "call_mixed",
|
||||
"name": tool_name,
|
||||
"input": tool_input or {},
|
||||
})
|
||||
content.append(
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "call_mixed",
|
||||
"name": tool_name,
|
||||
"input": tool_input or {},
|
||||
},
|
||||
)
|
||||
if image_url:
|
||||
content.append({"type": "image", "source": {"url": image_url}})
|
||||
return Msg(name="assistant", role="assistant", content=content)
|
||||
|
|
@ -274,8 +272,7 @@ def test_format_msgs_to_str_message_order():
|
|||
third_pos = result.find("Third message")
|
||||
|
||||
assert first_pos < second_pos < third_pos, (
|
||||
f"Messages not in correct order. Positions: first={first_pos}, "
|
||||
f"second={second_pos}, third={third_pos}"
|
||||
f"Messages not in correct order. Positions: first={first_pos}, " f"second={second_pos}, third={third_pos}"
|
||||
)
|
||||
verify_result_within_threshold(handler, result, threshold, "message_order", msgs)
|
||||
print_pass("test_format_msgs_to_str_message_order")
|
||||
|
|
@ -347,9 +344,7 @@ def test_format_msgs_to_str_thinking_excluded_by_default():
|
|||
msgs = [create_thinking_msg("Let me think about this...", "Here is my response")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold, include_thinking=False)
|
||||
|
||||
assert "Let me think about this" not in result, (
|
||||
f"Thinking content should be excluded, got: {result}"
|
||||
)
|
||||
assert "Let me think about this" not in result, f"Thinking content should be excluded, got: {result}"
|
||||
assert "Here is my response" in result, f"Text content should be included, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "thinking_excluded_by_default", msgs)
|
||||
print_pass("test_format_msgs_to_str_thinking_excluded_by_default")
|
||||
|
|
@ -362,9 +357,7 @@ def test_format_msgs_to_str_thinking_included():
|
|||
msgs = [create_thinking_msg("Let me think about this...", "Here is my response")]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold, include_thinking=True)
|
||||
|
||||
assert "Let me think about this" in result, (
|
||||
f"Thinking content should be included, got: {result}"
|
||||
)
|
||||
assert "Let me think about this" in result, f"Thinking content should be included, got: {result}"
|
||||
assert "<thinking>" in result, f"Expected thinking tag in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "thinking_included", msgs)
|
||||
print_pass("test_format_msgs_to_str_thinking_included")
|
||||
|
|
@ -375,14 +368,18 @@ def test_format_msgs_to_str_thinking_only_message():
|
|||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_thinking_msg("Deep thoughts here")]
|
||||
|
||||
|
||||
# With include_thinking=False
|
||||
result_no_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=False
|
||||
msgs,
|
||||
memory_compact_threshold=threshold,
|
||||
include_thinking=False,
|
||||
)
|
||||
# With include_thinking=True
|
||||
result_with_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=True
|
||||
msgs,
|
||||
memory_compact_threshold=threshold,
|
||||
include_thinking=True,
|
||||
)
|
||||
|
||||
assert "Deep thoughts here" not in result_no_thinking
|
||||
|
|
@ -425,9 +422,9 @@ def test_format_msgs_to_str_exceeds_threshold_truncate_older():
|
|||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# The newest messages should be present
|
||||
assert "Answer 19" in result or "Question 19" in result, (
|
||||
f"Expected recent message in result, got: {result[:500]}..."
|
||||
)
|
||||
assert (
|
||||
"Answer 19" in result or "Question 19" in result
|
||||
), f"Expected recent message in result, got: {result[:500]}..."
|
||||
# Older messages should be truncated
|
||||
assert "Question 0" not in result, "Older messages should be truncated"
|
||||
verify_result_within_threshold(handler, result, threshold, "exceeds_threshold_truncate_older", msgs)
|
||||
|
|
@ -517,10 +514,7 @@ def test_format_msgs_to_str_large_threshold():
|
|||
"""Test with very large threshold - all messages should be included."""
|
||||
handler = create_handler()
|
||||
threshold = 1000000
|
||||
msgs = [
|
||||
create_user_msg("Message " + str(i) + " " + "x" * 100)
|
||||
for i in range(50)
|
||||
]
|
||||
msgs = [create_user_msg("Message " + str(i) + " " + "x" * 100) for i in range(50)]
|
||||
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
|
|
@ -604,19 +598,25 @@ def test_format_msgs_to_str_mixed_content_blocks():
|
|||
"""Test message with mixed content blocks."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [create_mixed_content_msg(
|
||||
text="Text content",
|
||||
thinking="Thinking content",
|
||||
tool_name="test_tool",
|
||||
tool_input={"key": "value"},
|
||||
image_url="https://example.com/img.png",
|
||||
)]
|
||||
msgs = [
|
||||
create_mixed_content_msg(
|
||||
text="Text content",
|
||||
thinking="Thinking content",
|
||||
tool_name="test_tool",
|
||||
tool_input={"key": "value"},
|
||||
image_url="https://example.com/img.png",
|
||||
),
|
||||
]
|
||||
|
||||
result_no_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=False
|
||||
msgs,
|
||||
memory_compact_threshold=threshold,
|
||||
include_thinking=False,
|
||||
)
|
||||
result_with_thinking = handler.format_msgs_to_str(
|
||||
msgs, memory_compact_threshold=threshold, include_thinking=True
|
||||
msgs,
|
||||
memory_compact_threshold=threshold,
|
||||
include_thinking=True,
|
||||
)
|
||||
|
||||
assert "Text content" in result_no_thinking
|
||||
|
|
@ -682,7 +682,7 @@ def test_format_msgs_to_str_different_roles():
|
|||
def test_format_msgs_to_str_incremental_threshold_check():
|
||||
"""Test incremental addition of messages until threshold is exceeded."""
|
||||
handler = create_handler()
|
||||
|
||||
|
||||
# Create messages with known approximate sizes
|
||||
msgs = []
|
||||
for i in range(10):
|
||||
|
|
@ -690,16 +690,14 @@ def test_format_msgs_to_str_incremental_threshold_check():
|
|||
|
||||
# Calculate total tokens
|
||||
total_tokens = sum(handler.stat_message(msg).total_tokens for msg in msgs)
|
||||
|
||||
|
||||
# Use threshold that allows about half the messages
|
||||
half_threshold = total_tokens // 2
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=half_threshold)
|
||||
|
||||
# Should have some but not all messages
|
||||
included_count = sum(1 for i in range(10) if f"Message {i}" in result)
|
||||
assert 0 < included_count < 10, (
|
||||
f"Expected partial messages, got {included_count} messages included"
|
||||
)
|
||||
assert 0 < included_count < 10, f"Expected partial messages, got {included_count} messages included"
|
||||
# Newer messages should be included (messages are processed from end)
|
||||
assert "Message 9" in result, "Newest message should be included"
|
||||
verify_result_within_threshold(handler, result, half_threshold, "incremental_threshold_check", msgs)
|
||||
|
|
@ -743,17 +741,21 @@ def test_format_msgs_to_str_base64_image():
|
|||
"""Test with base64 encoded image."""
|
||||
handler = create_handler()
|
||||
threshold = 10000
|
||||
msgs = [Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"data": "SGVsbG8gV29ybGQ=" * 100, # Simulated base64 data
|
||||
},
|
||||
}],
|
||||
)]
|
||||
msgs = [
|
||||
Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"data": "SGVsbG8gV29ybGQ=" * 100, # Simulated base64 data
|
||||
},
|
||||
},
|
||||
],
|
||||
),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "[image]" in result
|
||||
|
|
@ -765,14 +767,16 @@ def test_format_msgs_to_str_audio_video_blocks():
|
|||
"""Test with audio and video content blocks."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{"type": "audio", "source": {"url": "https://example.com/audio.mp3"}},
|
||||
{"type": "video", "source": {"url": "https://example.com/video.mp4"}},
|
||||
],
|
||||
)]
|
||||
msgs = [
|
||||
Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{"type": "audio", "source": {"url": "https://example.com/audio.mp3"}},
|
||||
{"type": "video", "source": {"url": "https://example.com/video.mp4"}},
|
||||
],
|
||||
),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
assert "[audio]" in result
|
||||
|
|
@ -785,14 +789,16 @@ def test_format_msgs_to_str_unknown_block_type():
|
|||
"""Test that unknown block types are skipped gracefully."""
|
||||
handler = create_handler()
|
||||
threshold = 4000
|
||||
msgs = [Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{"type": "unknown_type", "data": "some data"},
|
||||
{"type": "text", "text": "Valid text"},
|
||||
],
|
||||
)]
|
||||
msgs = [
|
||||
Msg(
|
||||
name="assistant",
|
||||
role="assistant",
|
||||
content=[
|
||||
{"type": "unknown_type", "data": "some data"},
|
||||
{"type": "text", "text": "Valid text"},
|
||||
],
|
||||
),
|
||||
]
|
||||
result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
|
||||
|
||||
# Should still include valid content
|
||||
|
|
@ -880,4 +886,4 @@ def run_all_tests():
|
|||
|
||||
if __name__ == "__main__":
|
||||
success = run_all_tests()
|
||||
exit(0 if success else 1)
|
||||
sys.exit(0 if success else 1)
|
||||
|
|
|
|||
|
|
@ -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 颜色码
|
||||
|
|
|
|||
|
|
@ -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']}")
|
||||
|
|
|
|||
|
|
@ -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 颜色码
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue