diff --git a/README.md b/README.md
index 9c69d232..27acedf5 100644
--- a/README.md
+++ b/README.md
@@ -67,14 +67,15 @@ working_dir/
capabilities for AI Agents:
| Method | Function | Key Components |
-|------------------------|------------------------------------|-------------------------------------------------------------------------------------------------------------------------------------------------------------|
+|------------------------|------------------------------------|----------------------------------------------------------------------------------------------------------------------------------------------------------|
| `start` | 🚀 Start memory system | Initialize file store, file watcher, Embedding cache; clean up expired tool result files |
| `close` | 📕 Close and clean up | Clean tool result files, stop file watcher, save Embedding cache |
| `compact_memory` | 📦 Compact history to summary | [Compactor](reme/memory/file_based/compactor.py) — ReActAgent generates structured context checkpoint |
| `summary_memory` | 📝 Write important memory to files | [Summarizer](reme/memory/file_based/summarizer.py) — ReActAgent + file tools (read / write / edit) |
| `compact_tool_result` | ✂️ Compact oversized tool output | [ToolResultCompactor](reme/memory/file_based/tool_result_compactor.py) — Truncate and save to `tool_result/`, keep file reference in message |
+| `pre_reasoning_hook` | 🔄 Pre-reasoning hook | Auto compact tool results + generate summary + async trigger memory summarization task |
| `memory_search` | 🔍 Semantic memory search | [MemorySearch](reme/memory/tools/chunk/memory_search.py) — Vector + BM25 hybrid retrieval |
-| `get_in_memory_memory` | 🗂️ Create in-memory instance | [ReMeInMemoryMemory](reme/memory/file_based/reme_in_memory_memory.py) — Token-aware memory management, supports compression summary and state serialization |
+| `get_in_memory_memory` | 🗂️ Create in-memory instance | [ReMeInMemoryMemory](reme/memory/file_based/reme_in_memory_memory.py) — Token-aware memory management, supports compression summary and state serialization (static method) |
---
@@ -96,7 +97,6 @@ pip install -e ".[light]"
| `LLM_BASE_URL` | LLM base URL | `https://dashscope.aliyuncs.com/compatible-mode/v1` |
| `EMBEDDING_API_KEY` | Embedding API key (Optional) | `sk-xxx` |
| `EMBEDDING_BASE_URL` | Embedding base URL (Optional) | `https://dashscope.aliyuncs.com/compatible-mode/v1` |
-| `LLM_MODEL_NAME` | LLM model name | `qwen3.5-plus` |
#### Python Usage
@@ -108,38 +108,56 @@ from reme.reme_light import ReMeLight
async def main():
+ # Initialize ReMeLight
reme = ReMeLight(
- working_dir=".reme", # Memory file storage directory
- max_input_length=128000, # Model context window (tokens)
- memory_compact_ratio=0.7, # Trigger compaction when reaching max_input_length * 0.7
- language="zh", # Summary language (zh / "")
- tool_result_threshold=1000, # Auto-save tool outputs exceeding this character count
- retention_days=7, # tool_result/ file retention days
+ default_as_llm_config={"model_name": "qwen3.5-35b-a3b"},
+ # default_embedding_model_config={"model_name": "text-embedding-v4"},
+ default_file_store_config={"fts_enabled": True, "vector_enabled": False},
)
await reme.start()
- messages = [...]
+ messages = [...] # Conversation message list
# 1. Compact oversized tool outputs (prevent tool results from overflowing context)
messages = await reme.compact_tool_result(messages)
- # 2. Compact history to structured summary (trigger: context approaching limit), can pass previous summary for incremental update
- summary = await reme.compact_memory(messages=messages, previous_summary="")
+ # 2. Compact history to structured summary (can pass previous summary for incremental update)
+ summary = await reme.compact_memory(
+ messages=messages,
+ previous_summary="",
+ max_input_length=128000, # Model context window (tokens)
+ compact_ratio=0.7, # Trigger compaction when reaching max_input_length * 0.7
+ language="zh", # Summary language (zh / "")
+ )
# 3. Submit async summary task in background (non-blocking, writes to memory/YYYY-MM-DD.md)
reme.add_async_summary_task(messages=messages)
- # 4. Semantic memory search (Vector + BM25 hybrid retrieval)
+ # 4. Pre-reasoning hook (auto compact tool results + generate summary)
+ processed_messages, compressed_summary = await reme.pre_reasoning_hook(
+ messages=messages,
+ system_prompt="You are a helpful AI assistant.",
+ compressed_summary="",
+ max_input_length=128000,
+ compact_ratio=0.7,
+ memory_compact_reserve=10000,
+ enable_tool_result_compact=True,
+ tool_result_compact_keep_n=3,
+ )
+
+ # 5. Semantic memory search (Vector + BM25 hybrid retrieval)
result = await reme.memory_search(query="Python version preference", max_results=5)
- # 5. Get in-memory instance (ReMeInMemoryMemory, manages single conversation context) AgentScope InMemoryMemory
- memory = reme.get_in_memory_memory()
- token_stats = await memory.estimate_tokens()
+ # 6. Get in-memory instance (static method, manages single conversation context)
+ memory = ReMeLight.get_in_memory_memory()
+ for msg in messages:
+ await memory.add(msg)
+ token_stats = await memory.estimate_tokens(max_input_length=128000)
print(f"Current context usage: {token_stats['context_usage_ratio']:.1f}%")
print(f"Message tokens: {token_stats['messages_tokens']}")
print(f"Estimated total tokens: {token_stats['estimated_tokens']}")
- # 6. Wait for background tasks before closing
+ # 7. Wait for background tasks before closing
summary_result = await reme.await_summary_tasks()
# Close ReMeLight
@@ -150,6 +168,9 @@ if __name__ == "__main__":
asyncio.run(main())
```
+> 📂 Full example code: [test_reme_light.py](tests/light/test_reme_light.py)
+> 📋 Example output: [test_reme_light.log](tests/light/test_reme_light.log) (223,838 tokens → 1,105 tokens, 99.5% compression ratio)
+
### File-Based ReMeLight Memory System Architecture
[CoPaw MemoryManager](https://github.com/agentscope-ai/CoPaw/blob/main/src/copaw/agents/memory/memory_manager.py)
@@ -157,18 +178,18 @@ inherits `ReMeLight` and integrates memory capabilities into the Agent reasoning
```mermaid
graph TB
- CoPaw["CoPaw MemoryManager\n(inherits ReMeLight)"] -->|pre_reasoning hook| Hook[MemoryCompactionHook]
+ CoPaw["CoPaw MemoryManager
(inherits ReMeLight)"] -->|pre_reasoning hook| Hook[MemoryCompactionHook]
CoPaw --> ReMeLight[ReMeLight]
Hook -->|exceeds threshold| ReMeLight
- ReMeLight --> CompactMemory[compact_memory\nHistory compaction]
- ReMeLight --> SummaryMemory[summary_memory\nWrite memory to files]
- ReMeLight --> CompactToolResult[compact_tool_result\nOversized tool output compaction]
- ReMeLight --> MemSearch[memory_search\nSemantic search]
- ReMeLight --> InMemory[get_in_memory_memory\nReMeInMemoryMemory]
- CompactMemory --> Compactor[Compactor\nReActAgent]
- SummaryMemory --> Summarizer[Summarizer\nReActAgent + file tools]
- CompactToolResult --> ToolResultCompactor[ToolResultCompactor\nTruncate + save to file]
- Summarizer --> FileIO[FileIO\nread / write / edit]
+ ReMeLight --> CompactMemory[compact_memory
History compaction]
+ ReMeLight --> SummaryMemory[summary_memory
Write memory to files]
+ ReMeLight --> CompactToolResult[compact_tool_result
Oversized tool output compaction]
+ ReMeLight --> MemSearch[memory_search
Semantic search]
+ ReMeLight --> InMemory[get_in_memory_memory
ReMeInMemoryMemory]
+ CompactMemory --> Compactor[Compactor
ReActAgent]
+ SummaryMemory --> Summarizer[Summarizer
ReActAgent + file tools]
+ CompactToolResult --> ToolResultCompactor[ToolResultCompactor
Truncate + save to file]
+ Summarizer --> FileIO[FileIO
read / write / edit]
FileIO --> MemoryFiles[memory/YYYY-MM-DD.md]
ToolResultCompactor --> ToolResultFiles[tool_result/*.txt]
MemoryFiles -.->|File change| FileWatcher[Async File Watcher]
diff --git a/README_ZH.md b/README_ZH.md
index b45b261c..ded1d8d1 100644
--- a/README_ZH.md
+++ b/README_ZH.md
@@ -66,9 +66,10 @@ working_dir/
| `close` | 📕 关闭并清理 | 清理工具结果文件、停止文件监控、保存 Embedding 缓存 |
| `compact_memory` | 📦 压缩历史对话为摘要 | [Compactor](reme/memory/file_based/compactor.py) — ReActAgent 生成结构化上下文检查点 |
| `summary_memory` | 📝 将重要记忆写入文件 | [Summarizer](reme/memory/file_based/summarizer.py) — ReActAgent + 文件工具(read / write / edit) |
-| `compact_tool_result` | ✂️ 压缩超长工具输出 | [ToolResultCompactor](reme/memory/file_based/tool_result_compactor.py) — 截断并转存到 `tool_result/`,消息中保留文件引用 | |
+| `compact_tool_result` | ✂️ 压缩超长工具输出 | [ToolResultCompactor](reme/memory/file_based/tool_result_compactor.py) — 截断并转存到 `tool_result/`,消息中保留文件引用 |
+| `pre_reasoning_hook` | 🔄 推理前预处理钩子 | 自动压缩工具结果 + 生成摘要 + 异步触发记忆总结任务 |
| `memory_search` | 🔍 语义搜索记忆 | [MemorySearch](reme/memory/tools/chunk/memory_search.py) — 向量 + BM25 混合检索 |
-| `get_in_memory_memory` | 🗂️ 创建会话内存实例 | [ReMeInMemoryMemory](reme/memory/file_based/reme_in_memory_memory.py) — Token 感知的内存管理,支持压缩摘要和状态序列化 |
+| `get_in_memory_memory` | 🗂️ 创建会话内存实例 | [ReMeInMemoryMemory](reme/memory/file_based/reme_in_memory_memory.py) — Token 感知的内存管理,支持压缩摘要和状态序列化(静态方法) |
---
@@ -90,7 +91,6 @@ pip install -e ".[light]"
| `LLM_BASE_URL` | LLM base URL | `https://dashscope.aliyuncs.com/compatible-mode/v1` |
| `EMBEDDING_API_KEY` | Embedding API key | `sk-xxx` |
| `EMBEDDING_BASE_URL` | Embedding base URL | `https://dashscope.aliyuncs.com/compatible-mode/v1` |
-| `LLM_MODEL_NAME` | LLM model name | `qwen3.5-plus` |
#### Python使用
@@ -102,38 +102,56 @@ from reme.reme_light import ReMeLight
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/ 文件保留天数
+ default_as_llm_config={"model_name": "qwen3.5-35b-a3b"},
+ # default_embedding_model_config={"model_name": "text-embedding-v4"},
+ default_file_store_config={"fts_enabled": True, "vector_enabled": False},
)
await reme.start()
- messages = [...]
+ messages = [...] # 对话消息列表
# 1. 压缩超长工具输出(防止工具结果撑爆上下文)
messages = await reme.compact_tool_result(messages)
- # 2. 将历史对话压缩为结构化摘要(触发时机:上下文接近上限),可传入上轮摘要,实现增量更新
- summary = await reme.compact_memory(messages=messages, previous_summary="")
+ # 2. 将历史对话压缩为结构化摘要(可传入上轮摘要,实现增量更新)
+ summary = await reme.compact_memory(
+ messages=messages,
+ previous_summary="",
+ max_input_length=128000, # 模型上下文窗口(tokens)
+ compact_ratio=0.7, # 达到 max_input_length * 0.7 时触发压缩
+ language="zh", # 摘要语言(zh / "")
+ )
# 3. 后台异步提交摘要任务(不阻塞对话,摘要写入 memory/YYYY-MM-DD.md)
reme.add_async_summary_task(messages=messages)
- # 4. 语义搜索记忆(向量 + BM25 混合检索)
+ # 4. 推理前预处理钩子(自动压缩工具结果 + 生成摘要)
+ processed_messages, compressed_summary = await reme.pre_reasoning_hook(
+ messages=messages,
+ system_prompt="你是一个有帮助的 AI 助手。",
+ compressed_summary="",
+ max_input_length=128000,
+ compact_ratio=0.7,
+ memory_compact_reserve=10000,
+ enable_tool_result_compact=True,
+ tool_result_compact_keep_n=3,
+ )
+
+ # 5. 语义搜索记忆(向量 + BM25 混合检索)
result = await reme.memory_search(query="Python 版本偏好", max_results=5)
- # 5. 获取会话内存实例(ReMeInMemoryMemory,管理单次对话的上下文)AgentScope InMemoryMemory
- memory = reme.get_in_memory_memory()
- token_stats = await memory.estimate_tokens()
+ # 6. 获取会话内存实例(静态方法,管理单次对话的上下文)
+ memory = ReMeLight.get_in_memory_memory()
+ for msg in messages:
+ await memory.add(msg)
+ 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']}")
- # 6. 关闭前等待后台任务完成
+ # 7. 关闭前等待后台任务完成
summary_result = await reme.await_summary_tasks()
# 关闭 ReMeLight
@@ -144,6 +162,9 @@ if __name__ == "__main__":
asyncio.run(main())
```
+> 📂 完整示例代码:[test_reme_light.py](tests/light/test_reme_light.py)
+> 📋 运行结果示例:[test_reme_light.log](tests/light/test_reme_light.log)(223,838 tokens → 1,105 tokens,压缩率 99.5%)
+
### 基于文件的 ReMeLight 记忆系统架构
[CoPaw MemoryManager](https://github.com/agentscope-ai/CoPaw/blob/main/src/copaw/agents/memory/memory_manager.py) 继承
@@ -151,18 +172,18 @@ if __name__ == "__main__":
```mermaid
graph TB
- CoPaw["CoPaw MemoryManager\n(继承 ReMeLight)"] -->|pre_reasoning hook| Hook[MemoryCompactionHook]
+ CoPaw["CoPaw MemoryManager
(继承 ReMeLight)"] -->|pre_reasoning hook| Hook[MemoryCompactionHook]
CoPaw --> ReMeLight[ReMeLight]
Hook -->|超出阈值| ReMeLight
- ReMeLight --> CompactMemory[compact_memory\n历史对话压缩]
- ReMeLight --> SummaryMemory[summary_memory\n记忆写入文件]
- ReMeLight --> CompactToolResult[compact_tool_result\n超长工具输出压缩]
- ReMeLight --> MemSearch[memory_search\n语义搜索]
- ReMeLight --> InMemory[get_in_memory_memory\nReMeInMemoryMemory]
- CompactMemory --> Compactor[Compactor\nReActAgent]
- SummaryMemory --> Summarizer[Summarizer\nReActAgent + 文件工具]
- CompactToolResult --> ToolResultCompactor[ToolResultCompactor\n截断 + 转存文件]
- Summarizer --> FileIO[FileIO\nread / write / edit]
+ ReMeLight --> CompactMemory[compact_memory
历史对话压缩]
+ ReMeLight --> SummaryMemory[summary_memory
记忆写入文件]
+ ReMeLight --> CompactToolResult[compact_tool_result
超长工具输出压缩]
+ ReMeLight --> MemSearch[memory_search
语义搜索]
+ ReMeLight --> InMemory[get_in_memory_memory
ReMeInMemoryMemory]
+ CompactMemory --> Compactor[Compactor
ReActAgent]
+ SummaryMemory --> Summarizer[Summarizer
ReActAgent + 文件工具]
+ CompactToolResult --> ToolResultCompactor[ToolResultCompactor
截断 + 转存文件]
+ Summarizer --> FileIO[FileIO
read / write / edit]
FileIO --> MemoryFiles[memory/YYYY-MM-DD.md]
ToolResultCompactor --> ToolResultFiles[tool_result/*.txt]
MemoryFiles -.->|文件变更| FileWatcher[异步文件监控]
diff --git a/reme/config/light.yaml b/reme/config/light.yaml
index 2a83371a..bc85c10d 100644
--- a/reme/config/light.yaml
+++ b/reme/config/light.yaml
@@ -1,11 +1,27 @@
+as_llms:
+ default:
+ backend: openai
+ model_name: qwen3.5-plus
+
+as_llm_formatters:
+ default:
+ backend: openai
+
embedding_models:
default:
backend: openai
+ dimensions: 1024
+ use_dimensions: false
+ enable_cache: true
+ max_batch_size: 10
+ max_cache_size: 2000
+ max_input_length: 8192
file_stores:
default:
backend: chroma
embedding_model: default
+ store_name: "reme"
file_watchers:
default:
diff --git a/reme/core/__init__.py b/reme/core/__init__.py
index 5872e2ad..053755cc 100644
--- a/reme/core/__init__.py
+++ b/reme/core/__init__.py
@@ -1,5 +1,7 @@
"""Core"""
+from . import as_llm
+from . import as_llm_formatter
from . import embedding
from . import enumeration
from . import file_store
@@ -21,6 +23,8 @@ from .service_context import ServiceContext
__all__ = [
# Submodules
+ "as_llm",
+ "as_llm_formatter",
"embedding",
"enumeration",
"file_watcher",
diff --git a/reme/core/application.py b/reme/core/application.py
index 49537807..46f4a934 100644
--- a/reme/core/application.py
+++ b/reme/core/application.py
@@ -1,6 +1,7 @@
"""High-level entry point for configuring and running ReMe services and flows."""
import asyncio
+import os
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
@@ -35,6 +36,8 @@ class Application:
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,
@@ -55,6 +58,8 @@ class Application:
config_path=config_path,
enable_logo=enable_logo,
log_to_console=log_to_console,
+ default_as_llm_config=default_as_llm_config,
+ default_as_llm_formatter_config=default_as_llm_formatter_config,
default_llm_config=default_llm_config,
default_embedding_model_config=default_embedding_model_config,
default_vector_store_config=default_vector_store_config,
@@ -147,6 +152,26 @@ class Application:
if self.service_context.service_config.enable_logo:
print_logo(service_config=self.service_config)
+ for name, config in self.service_config.as_llms.items():
+ if config.backend not in R.as_llms:
+ logger.warning(f"AS LLM backend {config.backend} is not supported.")
+ else:
+ config_dict = config.model_dump(exclude={"backend"})
+ if not config_dict.get("api_key", ""):
+ config_dict["api_key"] = os.getenv("LLM_API_KEY", "")
+ if "client_kwargs" not in config_dict:
+ config_dict["client_kwargs"] = {}
+ if not config_dict["client_kwargs"].get("base_url", ""):
+ config_dict["client_kwargs"]["base_url"] = os.getenv("LLM_BASE_URL", "")
+ self.service_context.as_llms[name] = R.as_llms[config.backend](**config_dict)
+
+ for name, config in self.service_config.as_llm_formatters.items():
+ if config.backend not in R.as_llm_formatters:
+ logger.warning(f"AS LLM formatter backend {config.backend} is not supported.")
+ else:
+ config_dict = config.model_dump(exclude={"backend"})
+ self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config.backend](**config_dict)
+
for name, config in self.service_config.llms.items():
if config.backend not in R.llms:
logger.warning(f"LLM backend {config.backend} is not supported.")
diff --git a/reme/core/as_llm/__init__.py b/reme/core/as_llm/__init__.py
new file mode 100644
index 00000000..9cf527af
--- /dev/null
+++ b/reme/core/as_llm/__init__.py
@@ -0,0 +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("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
new file mode 100644
index 00000000..88b326a7
--- /dev/null
+++ b/reme/core/as_llm_formatter/__init__.py
@@ -0,0 +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("openai")(OpenAIChatFormatter)
+R.as_llm_formatters.register("dashscope")(DashScopeChatFormatter)
diff --git a/reme/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py
index 36c6def7..78a91b8b 100644
--- a/reme/core/embedding/base_embedding_model.py
+++ b/reme/core/embedding/base_embedding_model.py
@@ -81,12 +81,12 @@ class BaseEmbeddingModel(ABC):
@property
def api_key(self) -> str | None:
"""Get API key from environment variable."""
- return os.getenv("REME_EMBEDDING_API_KEY") or self._api_key
+ return os.getenv("EMBEDDING_API_KEY") or self._api_key
@property
def base_url(self) -> str | None:
"""Get base URL from environment variable."""
- return os.getenv("REME_EMBEDDING_BASE_URL") or self._base_url
+ return os.getenv("EMBEDDING_BASE_URL") or self._base_url
def _truncate_text(self, text: str) -> str:
"""Truncate text to max_input_length if it exceeds the limit."""
diff --git a/reme/core/op/base_op.py b/reme/core/op/base_op.py
index af5158c4..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
@@ -42,6 +44,8 @@ class BaseOp(metaclass=ABCMeta):
language: str = "",
prompt_name: str = "",
prompt_path: str = "",
+ as_llm: str | ChatModelBase = "default",
+ as_llm_formatter: str | FormatterBase = "default",
llm: str | BaseLLM = "default",
embedding_model: str | BaseEmbeddingModel = "default",
vector_store: str | BaseVectorStore = "default",
@@ -64,6 +68,8 @@ class BaseOp(metaclass=ABCMeta):
self.language = language
self.prompt = self._get_prompt_handler(prompt_name, prompt_path)
+ self._as_llm = as_llm
+ self._as_llm_formatter = as_llm_formatter
self._llm = llm
self._embedding_model = embedding_model
self._vector_store = vector_store
@@ -129,6 +135,20 @@ class BaseOp(metaclass=ABCMeta):
"""Access the service configuration."""
return self.service_context.service_config
+ @property
+ def as_llm(self) -> ChatModelBase:
+ """Get the AgentScope LLM instance from ServiceContext."""
+ if isinstance(self._as_llm, str):
+ self._as_llm = self.service_context.as_llms[self._as_llm]
+ return self._as_llm
+
+ @property
+ def as_llm_formatter(self) -> FormatterBase:
+ """Get the AgentScope LLM formatter instance from ServiceContext."""
+ if isinstance(self._as_llm_formatter, str):
+ self._as_llm_formatter = self.service_context.as_llm_formatters[self._as_llm_formatter]
+ return self._as_llm_formatter
+
@property
def llm(self) -> BaseLLM:
"""Get the LLM instance from ServiceContext."""
diff --git a/reme/core/registry_factory.py b/reme/core/registry_factory.py
index b319b049..f54ad3c1 100644
--- a/reme/core/registry_factory.py
+++ b/reme/core/registry_factory.py
@@ -34,6 +34,8 @@ class RegistryFactory:
def __init__(self):
self.llms = Registry()
+ self.as_llms = Registry()
+ self.as_llm_formatters = Registry()
self.embedding_models = Registry()
self.vector_stores = Registry()
self.file_stores = Registry()
diff --git a/reme/core/schema/__init__.py b/reme/core/schema/__init__.py
index de167c69..b6445a28 100644
--- a/reme/core/schema/__init__.py
+++ b/reme/core/schema/__init__.py
@@ -1,5 +1,6 @@
"""schema"""
+from .as_msg_stat import AsBlockStat, AsMsgStat
from .cut_point_result import CutPointResult
from .file_metadata import FileMetadata
from .memory_chunk import MemoryChunk
@@ -27,6 +28,8 @@ from .truncation_result import TruncationResult
from .vector_node import VectorNode
__all__ = [
+ "AsBlockStat",
+ "AsMsgStat",
"CutPointResult",
"CmdConfig",
"ContentBlock",
diff --git a/reme/core/schema/as_msg_stat.py b/reme/core/schema/as_msg_stat.py
new file mode 100644
index 00000000..2e861863
--- /dev/null
+++ b/reme/core/schema/as_msg_stat.py
@@ -0,0 +1,83 @@
+"""Schema definitions for AgentScope message statistics."""
+
+from pydantic import BaseModel, Field
+
+_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH = 100
+_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")
+
+ # For tool_use and tool_result blocks
+ tool_name: str = Field(default="", description="Tool name for tool_use/tool_result blocks")
+ tool_input: str = Field(default="", description="Tool input arguments for tool_use blocks")
+ tool_output: str = Field(default="", description="Tool output for tool_result blocks")
+
+ # For media blocks
+ media_url: str = Field(default="", description="URL for image/audio/video blocks")
+
+ @property
+ def preview(self) -> str:
+ """Return 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.
+
+ 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":
+ return truncate_text(self.text, max_length) if self.text else ""
+ if self.block_type == "thinking":
+ if include_thinking and self.text:
+ return f"\n{truncate_text(self.text, max_length)}\n"
+ 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 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)}"
+ else:
+ 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)
+ timestamp: str = Field(default="")
+ metadata: dict = Field(default_factory=dict)
+
+ @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:
+ """Format message to string representation."""
+ time_str = f"[{self.timestamp}] " if self.timestamp else ""
+ header = f"{time_str}{self.name or self.role}:"
+ blocks = [block.format(max_length, include_thinking) for block in self.content]
+ return "\n".join([header] + [b for b in blocks if b])
diff --git a/reme/core/schema/service_config.py b/reme/core/schema/service_config.py
index 5b89367a..5e4cd212 100644
--- a/reme/core/schema/service_config.py
+++ b/reme/core/schema/service_config.py
@@ -58,69 +58,60 @@ class FlowConfig(ToolCall):
cache_expire_hours: float = Field(default=0.1)
-class LLMConfig(BaseModel):
+class BasicConfig(BaseModel):
+ """Configuration for basic service settings and parameters."""
+
+ model_config = ConfigDict(extra="allow")
+
+ backend: str = Field(default="")
+
+
+class ModelConfig(BasicConfig):
+ """Configuration for model-based services with backend and model name."""
+
+ model_name: str = Field(default="")
+
+
+class LLMConfig(ModelConfig):
"""Configuration for Large Language Model backend and model identification."""
- model_config = ConfigDict(extra="allow")
- backend: str = Field(default="")
- model_name: str = Field(default="")
-
-
-class EmbeddingModelConfig(BaseModel):
+class EmbeddingModelConfig(ModelConfig):
"""Configuration for embedding model backends and identity."""
- model_config = ConfigDict(extra="allow")
- backend: str = Field(default="")
- model_name: str = Field(default="")
-
-
-class VectorStoreConfig(BaseModel):
- """Configuration for vector database storage and associated embeddings."""
-
- model_config = ConfigDict(extra="allow")
-
- backend: str = Field(default="local")
- collection_name: str = Field(default="reme")
- embedding_model: str = Field(default="default")
-
-
-class FileStoreConfig(BaseModel):
- """Configuration for file store database storage and associated embeddings."""
-
- model_config = ConfigDict(extra="allow")
-
- backend: str = Field(default="sqlite")
- store_name: str = Field(default="reme")
- embedding_model: str = Field(default="default")
-
-
-class TokenCounterConfig(BaseModel):
+class TokenCounterConfig(ModelConfig):
"""Configuration for token counting services and model mapping."""
- model_config = ConfigDict(extra="allow")
- backend: str = Field(default="base")
- model_name: str = Field(default="")
+class StoreConfig(BasicConfig):
+ """Configuration for storage services with embedding model support."""
+
+ embedding_model: str = Field(default="default")
-class FileWatcherConfig(BaseModel):
+class VectorStoreConfig(StoreConfig):
+ """Configuration for vector database storage and associated embeddings."""
+
+ collection_name: str = Field(default="reme")
+
+
+class FileStoreConfig(StoreConfig):
+ """Configuration for file store database storage and associated embeddings."""
+
+ store_name: str = Field(default="reme")
+
+
+class FileWatcherConfig(BasicConfig):
"""Configuration for file watcher service."""
- model_config = ConfigDict(extra="allow")
-
- backend: str = Field(default="")
file_store: str = Field(default="")
watch_paths: list[str] = Field(default_factory=list)
-class ServiceConfig(BaseModel):
+class ServiceConfig(BasicConfig):
"""Root configuration schema aggregating all service-level settings and components."""
- model_config = ConfigDict(extra="allow")
-
- backend: str = Field(default="")
app_name: str = Field(default=os.getenv("APP_NAME", "ReMe"))
working_dir: str = Field(default=".reme")
enable_logo: bool = Field(default=True)
@@ -137,6 +128,8 @@ class ServiceConfig(BaseModel):
cmd: CmdConfig = Field(default_factory=CmdConfig)
ops: dict[str, OpConfig] = Field(default_factory=dict)
flows: dict[str, FlowConfig] = Field(default_factory=dict)
+ as_llms: dict[str, BasicConfig] = Field(default_factory=dict)
+ as_llm_formatters: dict[str, BasicConfig] = Field(default_factory=dict)
llms: dict[str, LLMConfig] = Field(default_factory=dict)
embedding_models: dict[str, EmbeddingModelConfig] = Field(default_factory=dict)
vector_stores: dict[str, VectorStoreConfig] = Field(default_factory=dict)
diff --git a/reme/core/service_context.py b/reme/core/service_context.py
index ecb6c3a3..92d0d566 100644
--- a/reme/core/service_context.py
+++ b/reme/core/service_context.py
@@ -11,6 +11,8 @@ from .schema import ServiceConfig
from .utils import load_env, PydanticConfigParser
if TYPE_CHECKING:
+ from agentscope.model import ChatModelBase
+ from agentscope.formatter import FormatterBase
from .llm import BaseLLM
from .embedding import BaseEmbeddingModel
from .vector_store import BaseVectorStore
@@ -36,6 +38,8 @@ class ServiceContext(BaseDict):
config_path: str | None = None,
enable_logo: bool = True,
log_to_console: bool = True,
+ 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,
@@ -64,6 +68,10 @@ class ServiceContext(BaseDict):
if args:
input_args.extend(args)
+ if default_as_llm_config:
+ self._update_section_config(kwargs, "as_llms", **default_as_llm_config)
+ if default_as_llm_formatter_config:
+ self._update_section_config(kwargs, "as_llm_formatters", **default_as_llm_formatter_config)
if default_llm_config:
self._update_section_config(kwargs, "llms", **default_llm_config)
if default_embedding_model_config:
@@ -90,6 +98,8 @@ class ServiceContext(BaseDict):
self.service_config: ServiceConfig = service_config
self.thread_pool: ThreadPoolExecutor | None = None
+ self.as_llms: dict[str, "ChatModelBase"] = {}
+ self.as_llm_formatters: dict[str, "FormatterBase"] = {}
self.llms: dict[str, "BaseLLM"] = {}
self.embedding_models: dict[str, "BaseEmbeddingModel"] = {}
self.token_counters: dict[str, "BaseTokenCounter"] = {}
diff --git a/reme/core/utils/__init__.py b/reme/core/utils/__init__.py
index b43784c8..c1f35adb 100644
--- a/reme/core/utils/__init__.py
+++ b/reme/core/utils/__init__.py
@@ -11,12 +11,15 @@ from .horse import play_horse_easter_egg
from .http_client import HttpClient
from .llm_utils import extract_content, format_messages, deduplicate_memories
from .logger_utils import init_logger
+from .std_logger import get_logger as get_std_logger
from .logo_utils import print_logo
from .mcp_client import MCPClient
from .pydantic_config_parser import PydanticConfigParser
from .pydantic_utils import create_pydantic_model
from .singleton import singleton
from .time import timer, get_now_time
+from .hf_token_counter_utils import get_hf_token_counter
+from .truncate_text_utils import truncate_text, is_truncated
__all__ = [
"convert_dashscope_to_agentscope",
@@ -39,6 +42,7 @@ __all__ = [
"format_messages",
"deduplicate_memories",
"init_logger",
+ "get_std_logger",
"print_logo",
"MCPClient",
"PydanticConfigParser",
@@ -46,4 +50,7 @@ __all__ = [
"singleton",
"timer",
"get_now_time",
+ "get_hf_token_counter",
+ "truncate_text",
+ "is_truncated",
]
diff --git a/reme/core/utils/hf_token_counter_utils.py b/reme/core/utils/hf_token_counter_utils.py
new file mode 100644
index 00000000..a8ab348c
--- /dev/null
+++ b/reme/core/utils/hf_token_counter_utils.py
@@ -0,0 +1,23 @@
+"""Utility functions for working with text."""
+
+from agentscope.token import HuggingFaceTokenCounter
+
+_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,
+):
+ """Get or initialize the global token counter instance."""
+ global _token_counter
+ if _token_counter is None:
+ _token_counter = HuggingFaceTokenCounter(
+ pretrained_model_name_or_path=pretrained_model_name_or_path,
+ use_mirror=use_mirror,
+ use_fast=use_fast,
+ trust_remote_code=trust_remote_code,
+ )
+ return _token_counter
diff --git a/reme/core/utils/std_logger.py b/reme/core/utils/std_logger.py
new file mode 100644
index 00000000..e2de0e9c
--- /dev/null
+++ b/reme/core/utils/std_logger.py
@@ -0,0 +1,109 @@
+"""Standard logging module configuration with loguru-like features."""
+
+import logging
+import os
+import sys
+from datetime import datetime
+from logging.handlers import TimedRotatingFileHandler
+
+# Store created logger instances
+_loggers: dict[str, logging.Logger] = {}
+
+
+class CustomFormatter(logging.Formatter):
+ """Custom formatter with colorized output support."""
+
+ # ANSI color codes
+ COLORS = {
+ logging.DEBUG: "\033[36m", # Cyan
+ logging.INFO: "\033[32m", # Green
+ logging.WARNING: "\033[33m", # Yellow
+ logging.ERROR: "\033[31m", # Red
+ logging.CRITICAL: "\033[35m", # Magenta
+ }
+ RESET = "\033[0m"
+
+ def __init__(self, fmt: str, colorize: bool = False):
+ super().__init__(fmt)
+ self.colorize = colorize
+
+ def format(self, record: logging.LogRecord) -> str:
+ # Add custom attribute: simplified filename and line number
+ record.file_line = f"{record.filename}:{record.lineno}"
+
+ if self.colorize:
+ color = self.COLORS.get(record.levelno, self.RESET)
+ record.levelname = f"{color}{record.levelname}{self.RESET}"
+
+ return super().format(record)
+
+
+def get_logger(
+ name: str = "reme",
+ log_dir: str = "logs",
+ level: str = "INFO",
+ log_to_console: bool = True,
+ log_to_file: bool = True,
+ log_file_prefix: str = "reme",
+ rotation: str = "midnight",
+ retention_days: int = 7,
+) -> logging.Logger:
+ """Get a configured logger instance.
+
+ Args:
+ name: Logger name for distinguishing different loggers.
+ log_dir: Directory path for log files.
+ level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL).
+ log_to_console: Whether to output logs to console.
+ log_to_file: Whether to output logs to file.
+ log_file_prefix: Prefix for log file names (e.g., 'reme' -> 'reme_2024-01-01.log').
+ rotation: Log rotation time, defaults to midnight.
+ retention_days: Number of days to retain log files.
+
+ Returns:
+ Configured Logger instance.
+ """
+ # Return existing logger if already created
+ if name in _loggers:
+ return _loggers[name]
+
+ # Create new logger without using root logger
+ logger = logging.getLogger(name)
+ logger.setLevel(getattr(logging, level.upper(), logging.INFO))
+ logger.propagate = False # Do not propagate to root logger
+
+ # Clear existing handlers
+ logger.handlers.clear()
+
+ # Log format
+ log_format = "%(asctime)s | %(levelname)s | %(file_line)s | %(funcName)s | %(message)s"
+
+ # Configure file logging
+ if log_to_file:
+ os.makedirs(log_dir, exist_ok=True)
+ current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
+ log_filename = f"{log_file_prefix}_{current_ts}.log"
+ log_filepath = os.path.join(log_dir, log_filename)
+
+ file_handler = TimedRotatingFileHandler(
+ log_filepath,
+ when=rotation,
+ interval=1,
+ backupCount=retention_days,
+ encoding="utf-8",
+ )
+ file_handler.setLevel(getattr(logging, level.upper(), logging.INFO))
+ file_handler.setFormatter(CustomFormatter(log_format, colorize=False))
+ file_handler.suffix = "%Y-%m-%d"
+ logger.addHandler(file_handler)
+
+ # Configure console logging
+ if log_to_console:
+ console_handler = logging.StreamHandler(sys.stdout)
+ console_handler.setLevel(getattr(logging, level.upper(), logging.INFO))
+ console_handler.setFormatter(CustomFormatter(log_format, colorize=True))
+ logger.addHandler(console_handler)
+
+ # Cache logger
+ _loggers[name] = logger
+ return logger
diff --git a/reme/core/utils/truncate_text_utils.py b/reme/core/utils/truncate_text_utils.py
new file mode 100644
index 00000000..da0c473a
--- /dev/null
+++ b/reme/core/utils/truncate_text_utils.py
@@ -0,0 +1,55 @@
+"""Utility functions for truncating long text strings."""
+
+from .std_logger import get_logger
+
+logger = get_logger()
+
+TRUNCATION_MARKER_START = "<<>>"
+TRUNCATION_MARKER_END = "<<>>"
+
+
+def truncate_text(text: str, max_length: int) -> str:
+ """Truncate text to max length, keeping head and tail portions.
+
+ Args:
+ text: The text to truncate
+ max_length: Maximum allowed length
+
+ Returns:
+ Truncated text with unique markers indicating truncation
+ """
+ text = str(text) if text else ""
+ if not text:
+ return text
+
+ if len(text) <= max_length:
+ return text
+
+ half_length = max_length // 2
+ truncated_chars = len(text) - max_length
+ logger.debug(
+ "Text truncated: original %d chars, kept head %d + tail %d, removed %d chars.",
+ len(text),
+ half_length,
+ half_length,
+ truncated_chars,
+ )
+ return (
+ f"{text[:half_length]}\n\n{TRUNCATION_MARKER_START} "
+ f"({truncated_chars} characters omitted) "
+ f"{TRUNCATION_MARKER_END}\n\n{text[-half_length:]}"
+ )
+
+
+def is_truncated(text: str) -> bool:
+ """Check if the text has been truncated (contains truncation markers).
+
+ Args:
+ text: The text to check
+
+ Returns:
+ bool: True if text contains truncation markers, False otherwise
+ """
+ if not text:
+ return False
+ return TRUNCATION_MARKER_START in text and TRUNCATION_MARKER_END in text
diff --git a/reme/memory/file_based/__init__.py b/reme/memory/file_based/__init__.py
index 29f33749..2e01cc41 100644
--- a/reme/memory/file_based/__init__.py
+++ b/reme/memory/file_based/__init__.py
@@ -4,30 +4,23 @@ This module provides memory management components for CoPaw (Cooperative Paw) ag
including memory formatting, compaction, summarization, and file I/O operations.
Components:
- - MemoryFormatter: Converts message lists to formatted strings with token limiting
- ReMeInMemoryMemory: Extended InMemoryMemory with bugfixes and summary support
+ - AsMsgHandler: Handles AgentScope message statistics, formatting, and context checking
- Summarizer: Generates memory summaries using LLM
- Compactor: Compacts memory content to reduce token usage
- ToolResultCompactor: Truncates large tool results and saves full content to files
- - FileIO: File I/O operations with configurable working directory
"""
-from . import utils
-from .compactor import Compactor
-from .file_io import FileIO
-from .memory_formatter import MemoryFormatter
-from .reme_chat_formatter import ReMeChatFormatter
+from .as_msg_handler import AsMsgHandler
from .reme_in_memory_memory import ReMeInMemoryMemory
-from .summarizer import Summarizer
-from .tool_result_compactor import ToolResultCompactor
+from .component.compactor import Compactor
+from .component.summarizer import Summarizer
+from .component.tool_result_compactor import ToolResultCompactor
__all__ = [
- "MemoryFormatter",
+ "AsMsgHandler",
"ReMeInMemoryMemory",
"Summarizer",
"Compactor",
"ToolResultCompactor",
- "FileIO",
- "utils",
- "ReMeChatFormatter",
]
diff --git a/reme/memory/file_based/as_msg_handler.py b/reme/memory/file_based/as_msg_handler.py
new file mode 100644
index 00000000..9db6cac2
--- /dev/null
+++ b/reme/memory/file_based/as_msg_handler.py
@@ -0,0 +1,400 @@
+"""Handler for AgentScope message processing, token counting, and context management."""
+
+import json
+
+from agentscope.message import Msg
+from agentscope.token import HuggingFaceTokenCounter
+
+from ...core.schema import AsMsgStat, AsBlockStat
+from ...core.utils import get_std_logger
+
+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
+
+ def count_str_token(self, text: str) -> int:
+ """Count tokens in a string.
+
+ Args:
+ text: The text to count tokens for.
+
+ Returns:
+ The number of tokens in the text.
+ """
+ if not text:
+ return 0
+
+ try:
+ token_ids = self._token_counter.tokenizer.encode(text)
+ token_count = len(token_ids)
+ return token_count
+
+ except Exception as e:
+ estimated_tokens = len(text.encode("utf-8")) // 4
+ logger.warning(f"Failed to count string tokens: {text}, e={e}")
+ return estimated_tokens
+
+ def _format_tool_result_output(self, output: str | list[dict]) -> tuple[str, int]:
+ """Convert tool result output to string."""
+ if isinstance(output, str):
+ return output, self.count_str_token(output)
+
+ textual_parts = []
+ total_token_count = 0
+ for block in output:
+ try:
+ if not isinstance(block, dict) or "type" not in block:
+ logger.warning(
+ "Invalid block: %s, expected a dict with 'type' key, skipped.",
+ block,
+ )
+ continue
+
+ block_type = block["type"]
+
+ if block_type == "text":
+ textual_parts.append(block.get("text", ""))
+ total_token_count += self.count_str_token(textual_parts[-1])
+
+ elif block_type in ["image", "audio", "video"]:
+ source = block.get("source", {})
+ if source.get("type") == "base64":
+ data = source.get("data", "")
+ total_token_count += len(data) // 4 if data else 10
+ else:
+ url = source.get("url", "")
+ total_token_count += self.count_str_token(url) if url else 10
+ textual_parts.append(f"[{block_type}] {url}")
+
+ elif block_type == "file":
+ file_path = block.get("path", "") or block.get("url", "")
+ file_name = block.get("name", file_path)
+ textual_parts.append(f"[file] {file_name}: {file_path}")
+ total_token_count += self.count_str_token(file_path)
+
+ else:
+ logger.warning(
+ "Unsupported block type '%s' in tool result, skipped.",
+ block_type,
+ )
+
+ except Exception as e:
+ logger.warning(
+ "Failed to process block %s: %s, skipped.",
+ block,
+ e,
+ )
+
+ return "\n".join(textual_parts), total_token_count
+
+ def stat_message(self, message: Msg) -> AsMsgStat:
+ """Analyze a message and generate block statistics."""
+ blocks = []
+ if isinstance(message.content, str):
+ blocks.append(
+ AsBlockStat(
+ block_type="text",
+ text=message.content,
+ token_count=self.count_str_token(message.content),
+ ),
+ )
+ return AsMsgStat(
+ name=message.name or message.role,
+ role=message.role,
+ content=blocks,
+ timestamp=message.timestamp or "",
+ metadata=message.metadata or {},
+ )
+
+ for block in message.content:
+ block_type = block.get("type", "unknown")
+
+ if block_type == "text":
+ text = block.get("text", "")
+ token_count = self.count_str_token(text)
+ blocks.append(
+ AsBlockStat(
+ block_type=block_type,
+ text=text,
+ token_count=token_count,
+ ),
+ )
+
+ elif block_type == "thinking":
+ thinking = block.get("thinking", "")
+ token_count = self.count_str_token(thinking)
+ blocks.append(
+ AsBlockStat(
+ block_type=block_type,
+ text=thinking,
+ token_count=token_count,
+ ),
+ )
+
+ elif block_type in ("image", "audio", "video"):
+ source = block.get("source", {})
+ url = source.get("url", "")
+ if source.get("type") == "base64":
+ data = source.get("data", "")
+ token_count = len(data) // 4 if data else 10
+ else:
+ token_count = self.count_str_token(url) if url else 10
+ blocks.append(
+ AsBlockStat(
+ block_type=block_type,
+ text="",
+ token_count=token_count,
+ media_url=url,
+ ),
+ )
+
+ elif block_type == "tool_use":
+ tool_name = block.get("name", "")
+ tool_input = block.get("raw_input", "")
+ try:
+ input_str = json.dumps(tool_input, ensure_ascii=False)
+ except (TypeError, ValueError):
+ input_str = str(tool_input)
+ token_count = self.count_str_token(tool_name + input_str)
+ blocks.append(
+ AsBlockStat(
+ block_type=block_type,
+ text="",
+ token_count=token_count,
+ tool_name=tool_name,
+ tool_input=input_str,
+ ),
+ )
+
+ elif block_type == "tool_result":
+ tool_name = block.get("name", "")
+ output = block.get("output", "")
+ formatted_output, token_count = self._format_tool_result_output(output)
+ blocks.append(
+ AsBlockStat(
+ block_type=block_type,
+ text="",
+ token_count=token_count,
+ tool_name=tool_name,
+ tool_output=formatted_output,
+ ),
+ )
+
+ else:
+ logger.warning("Unsupported block type %s, skipped.", block_type)
+
+ return AsMsgStat(
+ name=message.name or message.role,
+ role=message.role,
+ content=blocks,
+ timestamp=message.timestamp or "",
+ metadata=message.metadata or {},
+ )
+
+ def count_msgs_token(self, messages: list[Msg]) -> int:
+ """Count total token count of a list of messages."""
+ return sum(self.stat_message(msg).total_tokens for msg in messages)
+
+ def format_msgs_to_str(
+ self,
+ messages: list[Msg],
+ memory_compact_threshold: int,
+ include_thinking: bool = False,
+ ) -> str:
+ """Format list of messages to a single formatted string.
+
+ Messages are processed in reverse order (newest first) and older
+ messages are skipped when token count exceeds memory_compact_threshold.
+
+ Args:
+ messages: List of Msg objects to format.
+ memory_compact_threshold: Maximum token count before skipping older messages.
+ include_thinking: Whether to include thinking blocks in output.
+ """
+ if not messages:
+ return ""
+
+ formatted_parts: list[str] = []
+ total_token_count = 0
+
+ for i in range(len(messages) - 1, -1, -1):
+ stat = self.stat_message(messages[i])
+ formatted_content = stat.format(include_thinking=include_thinking)
+ content_token_count = self.count_str_token(formatted_content)
+
+ if total_token_count + content_token_count > memory_compact_threshold:
+ logger.info(
+ "Skipping older messages: adding %d tokens would exceed threshold %d (current: %d)",
+ content_token_count,
+ memory_compact_threshold,
+ total_token_count,
+ )
+ break
+
+ formatted_parts.append(formatted_content)
+ total_token_count += content_token_count
+
+ formatted_parts.reverse()
+ return "\n\n".join(formatted_parts)
+
+ @staticmethod
+ def validate_tool_ids_alignment(messages: list[Msg]) -> bool:
+ """Check if tool_use_ids and tool_result_ids are properly aligned.
+
+ Args:
+ messages: List of Msg objects to validate.
+
+ Returns:
+ True if all tool_use ids have corresponding tool_result ids and vice versa.
+ """
+ tool_use_ids: set[str] = set()
+ tool_result_ids: set[str] = set()
+
+ for msg in messages:
+ for block in msg.get_content_blocks("tool_use"):
+ if tool_id := block.get("id"):
+ tool_use_ids.add(tool_id)
+ for block in msg.get_content_blocks("tool_result"):
+ if tool_id := block.get("id"):
+ tool_result_ids.add(tool_id)
+
+ return tool_use_ids == tool_result_ids
+
+ def context_check(
+ self,
+ messages: list[Msg],
+ memory_compact_threshold: int,
+ memory_compact_reserve: int,
+ ) -> tuple[list[Msg], list[Msg], bool]:
+ """Check if context exceeds threshold and split messages accordingly.
+
+ Only when total tokens exceed memory_compact_threshold, messages are split into
+ messages_to_keep (within reserve limit) and messages_to_compact (older messages).
+
+ Args:
+ messages: List of Msg objects to check.
+ memory_compact_threshold: Maximum token count threshold to trigger compaction.
+ memory_compact_reserve: Token limit for messages to keep.
+
+ Returns:
+ A tuple of (messages_to_compact, messages_to_keep, tools_aligned):
+ - messages_to_compact: Older messages that exceed reserve limit
+ - messages_to_keep: Recent messages within the reserve limit
+ - tools_aligned: Whether tool_use and tool_result ids are aligned in messages_to_keep
+ """
+ if not messages:
+ return [], [], True
+
+ # Calculate total tokens and stats for all messages
+ msg_stats: list[tuple[Msg, AsMsgStat]] = []
+ total_tokens = 0
+ for msg in messages:
+ stat = self.stat_message(msg)
+ msg_stats.append((msg, stat))
+ total_tokens += stat.total_tokens
+
+ # If total tokens don't exceed threshold, no split needed
+ if total_tokens < memory_compact_threshold:
+ return [], messages, True
+
+ # Collect all tool_use ids and their message indices
+ # tool_use_id -> message index
+ tool_use_locations: dict[str, int] = {}
+ # tool_result_id -> message index
+ tool_result_locations: dict[str, int] = {}
+
+ for idx, (msg, _) in enumerate(msg_stats):
+ for block in msg.get_content_blocks("tool_use"):
+ tool_id = block.get("id", "")
+ if tool_id:
+ tool_use_locations[tool_id] = idx
+
+ for block in msg.get_content_blocks("tool_result"):
+ tool_id = block.get("id", "")
+ if tool_id:
+ tool_result_locations[tool_id] = idx
+
+ # Iterate from the end, accumulating messages to keep within reserve limit
+ keep_indices: set[int] = set()
+ accumulated_tokens = 0
+
+ for i in range(len(msg_stats) - 1, -1, -1):
+ msg, stat = msg_stats[i]
+
+ # Check if adding this message would exceed reserve limit
+ if accumulated_tokens + stat.total_tokens > memory_compact_reserve:
+ logger.info(
+ "Context check: adding message %d with %d tokens would exceed reserve %d (current: %d)",
+ i,
+ stat.total_tokens,
+ memory_compact_reserve,
+ accumulated_tokens,
+ )
+ break
+
+ # Check tool_result dependencies - if this message has tool_result,
+ # we need to ensure the corresponding tool_use is also included
+ tool_result_ids = [
+ block.get("id", "") for block in msg.get_content_blocks("tool_result") if block.get("id", "")
+ ]
+
+ # Calculate extra tokens needed for dependent tool_use messages
+ extra_tokens = 0
+ dependent_indices: set[int] = set()
+
+ for tool_id in tool_result_ids:
+ if tool_id in tool_use_locations:
+ tool_use_idx = tool_use_locations[tool_id]
+ if tool_use_idx not in keep_indices and tool_use_idx != i:
+ dependent_indices.add(tool_use_idx)
+ _, dep_stat = msg_stats[tool_use_idx]
+ extra_tokens += dep_stat.total_tokens
+
+ # Check if we can fit this message plus its dependencies within reserve
+ if accumulated_tokens + stat.total_tokens + extra_tokens > memory_compact_reserve:
+ logger.info(
+ "Context check: message %d requires %d extra tokens for tool_use dependencies, "
+ "total would exceed reserve %d",
+ i,
+ extra_tokens,
+ memory_compact_reserve,
+ )
+ break
+
+ # Add this message and its dependencies
+ keep_indices.add(i)
+ keep_indices.update(dependent_indices)
+ accumulated_tokens += stat.total_tokens + extra_tokens
+
+ # Build final lists based on keep_indices (preserve original order)
+ messages_to_compact = []
+ messages_to_keep = []
+
+ for idx, (msg, _) in enumerate(msg_stats):
+ if idx in keep_indices:
+ messages_to_keep.append(msg)
+ else:
+ messages_to_compact.append(msg)
+
+ # Validate tool ids alignment for messages_to_keep
+ tools_aligned = self.validate_tool_ids_alignment(messages_to_keep)
+
+ logger.info(
+ "Context check result: %d messages to compact, %d messages to keep, "
+ "total tokens: %d, threshold: %d, reserve: %d, kept tokens: %d, "
+ "tools_aligned: %s",
+ len(messages_to_compact),
+ len(messages_to_keep),
+ total_tokens,
+ memory_compact_threshold,
+ memory_compact_reserve,
+ accumulated_tokens,
+ tools_aligned,
+ )
+
+ return messages_to_compact, messages_to_keep, tools_aligned
diff --git a/reme/memory/file_based/component/__init__.py b/reme/memory/file_based/component/__init__.py
new file mode 100644
index 00000000..e69de29b
diff --git a/reme/memory/file_based/compactor.py b/reme/memory/file_based/component/compactor.py
similarity index 75%
rename from reme/memory/file_based/compactor.py
rename to reme/memory/file_based/component/compactor.py
index d8186d86..3292c874 100644
--- a/reme/memory/file_based/compactor.py
+++ b/reme/memory/file_based/component/compactor.py
@@ -1,17 +1,14 @@
"""Compactor module for memory compaction operations."""
-import logging
-
from agentscope.agent import ReActAgent
-from agentscope.formatter import FormatterBase
from agentscope.message import Msg
-from agentscope.model import ChatModelBase
from agentscope.token import HuggingFaceTokenCounter
-from .memory_formatter import MemoryFormatter
-from ...core.op import BaseOp
+from ..as_msg_handler import AsMsgHandler
+from ....core.op import BaseOp
+from ....core.utils import get_std_logger
-logger = logging.getLogger(__name__)
+logger = get_std_logger()
class Compactor(BaseOp):
@@ -20,17 +17,13 @@ class Compactor(BaseOp):
def __init__(
self,
memory_compact_threshold: int,
- chat_model: ChatModelBase,
- formatter: FormatterBase,
token_counter: HuggingFaceTokenCounter,
**kwargs,
):
super().__init__(**kwargs)
self.memory_compact_threshold: int = memory_compact_threshold
- self.chat_model: ChatModelBase = chat_model
- self.formatter: FormatterBase = formatter
- self.as_token_counter: HuggingFaceTokenCounter = token_counter
+ self.msg_handler = AsMsgHandler(token_counter=token_counter)
async def execute(self):
messages: list[Msg] = self.context.get("messages", [])
@@ -39,11 +32,10 @@ class Compactor(BaseOp):
if not messages:
return ""
- formatter = MemoryFormatter(
- token_counter=self.as_token_counter,
+ history_formatted_str: str = self.msg_handler.format_msgs_to_str(
+ messages=messages,
memory_compact_threshold=self.memory_compact_threshold,
)
- history_formatted_str: str = formatter.format(messages)
if not history_formatted_str:
logger.warning(f"No history to compact. messages={messages}")
@@ -51,9 +43,9 @@ class Compactor(BaseOp):
agent = ReActAgent(
name="reme_compactor",
- model=self.chat_model,
+ model=self.as_llm,
sys_prompt=self.get_prompt("system_prompt"),
- formatter=self.formatter,
+ formatter=self.as_llm_formatter,
)
if previous_summary:
diff --git a/reme/memory/file_based/compactor.yaml b/reme/memory/file_based/component/compactor.yaml
similarity index 100%
rename from reme/memory/file_based/compactor.yaml
rename to reme/memory/file_based/component/compactor.yaml
diff --git a/reme/memory/file_based/summarizer.py b/reme/memory/file_based/component/summarizer.py
similarity index 62%
rename from reme/memory/file_based/summarizer.py
rename to reme/memory/file_based/component/summarizer.py
index 462ea9c5..db3522da 100644
--- a/reme/memory/file_based/summarizer.py
+++ b/reme/memory/file_based/component/summarizer.py
@@ -1,20 +1,17 @@
"""Summarizer module for memory summarization operations."""
import datetime
-import logging
from agentscope.agent import ReActAgent
-from agentscope.formatter import FormatterBase
from agentscope.message import Msg
-from agentscope.model import ChatModelBase
from agentscope.token import HuggingFaceTokenCounter
from agentscope.tool import Toolkit
-from .memory_formatter import MemoryFormatter
-from .file_io import FileIO
-from ...core.op import BaseOp
+from ..as_msg_handler import AsMsgHandler
+from ....core.op import BaseOp
+from ....core.utils import get_std_logger
-logger = logging.getLogger(__name__)
+logger = get_std_logger()
class Summarizer(BaseOp):
@@ -25,10 +22,8 @@ class Summarizer(BaseOp):
working_dir: str,
memory_dir: str,
memory_compact_threshold: int,
- chat_model: ChatModelBase,
- formatter: FormatterBase,
token_counter: HuggingFaceTokenCounter,
- toolkit: Toolkit | None = None,
+ toolkit: Toolkit,
**kwargs,
):
super().__init__(**kwargs)
@@ -36,17 +31,8 @@ class Summarizer(BaseOp):
self.memory_dir: str = memory_dir
self.memory_compact_threshold: int = memory_compact_threshold
- self.chat_model: ChatModelBase = chat_model
- self.formatter: FormatterBase = formatter
- self.as_token_counter: HuggingFaceTokenCounter = token_counter
- if toolkit is not None:
- self.toolkit: Toolkit = toolkit
- else:
- self.toolkit = Toolkit()
- file_io = FileIO(working_dir=self.working_dir)
- self.toolkit.register_tool_function(file_io.read)
- self.toolkit.register_tool_function(file_io.write)
- self.toolkit.register_tool_function(file_io.edit)
+ self.msg_handler = AsMsgHandler(token_counter=token_counter)
+ self.toolkit: Toolkit = toolkit
async def execute(self):
messages: list[Msg] = self.context.get("messages", [])
@@ -54,11 +40,10 @@ class Summarizer(BaseOp):
if not messages:
return ""
- formatter = MemoryFormatter(
- token_counter=self.as_token_counter,
+ history_formatted_str: str = self.msg_handler.format_msgs_to_str(
+ messages=messages,
memory_compact_threshold=self.memory_compact_threshold,
)
- history_formatted_str: str = formatter.format(messages)
if not history_formatted_str:
logger.warning(f"No history to summarize. messages={messages}")
@@ -66,9 +51,9 @@ class Summarizer(BaseOp):
agent = ReActAgent(
name="reme_summarizer",
- model=self.chat_model,
+ model=self.as_llm,
sys_prompt="You are a helpful assistant.",
- formatter=self.formatter,
+ formatter=self.as_llm_formatter,
toolkit=self.toolkit,
)
diff --git a/reme/memory/file_based/summarizer.yaml b/reme/memory/file_based/component/summarizer.yaml
similarity index 100%
rename from reme/memory/file_based/summarizer.yaml
rename to reme/memory/file_based/component/summarizer.yaml
diff --git a/reme/memory/file_based/tool_result_compactor.py b/reme/memory/file_based/component/tool_result_compactor.py
similarity index 96%
rename from reme/memory/file_based/tool_result_compactor.py
rename to reme/memory/file_based/component/tool_result_compactor.py
index 5b1ef0a4..412df6de 100644
--- a/reme/memory/file_based/tool_result_compactor.py
+++ b/reme/memory/file_based/component/tool_result_compactor.py
@@ -1,16 +1,16 @@
"""Tool Result Compactor: truncate large tool results and save full content to files."""
-import logging
import uuid
from datetime import datetime, timedelta
from pathlib import Path
from agentscope.message import Msg
-from .utils import is_truncated, truncate_text
-from ...core.op import BaseOp
+from ....core.op import BaseOp
+from ....core.utils import get_std_logger
+from ....core.utils import truncate_text, is_truncated
-logger = logging.getLogger(__name__)
+logger = get_std_logger()
class ToolResultCompactor(BaseOp):
diff --git a/reme/memory/file_based/memory_formatter.py b/reme/memory/file_based/memory_formatter.py
deleted file mode 100644
index 6c0e22c3..00000000
--- a/reme/memory/file_based/memory_formatter.py
+++ /dev/null
@@ -1,249 +0,0 @@
-"""Memory Formatter for CoPaw agents.
-
-Provides memory formatting capabilities including:
-- Converting list of Msg to formatted string
-- Memory compaction with token threshold
-- Support for various content block types (text, tool_use, tool_result, etc.)
-"""
-
-import json
-import logging
-import os
-
-from agentscope.message import Msg
-from agentscope.token import HuggingFaceTokenCounter
-
-from .utils import safe_count_str_tokens, truncate_text
-
-logger = logging.getLogger(__name__)
-
-_DEFAULT_MAX_FORMATTER_TEXT_LENGTH = 2000
-
-
-class MemoryFormatter:
- """Formatter that converts list of Msg to formatted string.
-
- Formats messages into human-readable string representation with:
- - Role and timestamp information
- - Text content and tool calls
- - Memory compact threshold to limit total token count
- """
-
- def __init__(
- self,
- token_counter: HuggingFaceTokenCounter,
- memory_compact_threshold: int,
- ):
- """Initialize MemoryFormatter.
-
- Args:
- token_counter: Token counter for estimating token counts.
- memory_compact_threshold: Maximum token count before skipping
- older messages.
- """
- self._token_counter = token_counter
- self._memory_compact_threshold = memory_compact_threshold
- self.max_length = int(
- os.getenv("MAX_FORMATTER_TEXT_LENGTH", str(_DEFAULT_MAX_FORMATTER_TEXT_LENGTH)),
- )
-
- @staticmethod
- def _format_tool_result_output(output: str | list[dict]) -> str:
- """Convert tool result output to string.
-
- Args:
- output: Tool result output, either string or list of content blocks.
-
- Returns:
- Formatted string representation of the tool result.
- """
- if isinstance(output, str):
- return output
-
- textual_parts = []
-
- for block in output:
- try:
- if not isinstance(block, dict) or "type" not in block:
- logger.warning(
- "Invalid block: %s, expected a dict with 'type' key, skipped.",
- block,
- )
- continue
-
- block_type = block["type"]
-
- if block_type == "text":
- textual_parts.append(block.get("text", ""))
-
- elif block_type in ["image", "audio", "video"]:
- source = block.get("source", {})
- url = source.get("url", "")
- if url:
- textual_parts.append(
- f"[{block_type}] {url}",
- )
- else:
- textual_parts.append(f"[{block_type}]")
-
- elif block_type == "file":
- file_path = block.get("path", "") or block.get("url", "")
- file_name = block.get("name", file_path)
- textual_parts.append(f"[file] {file_name}: {file_path}")
-
- else:
- # Unknown block type: log warning and skip
- logger.warning(
- "Unsupported block type '%s' in tool result, skipped.",
- block_type,
- )
-
- except Exception as e:
- logger.warning(
- "Failed to process block %s: %s, skipped.",
- block,
- e,
- )
-
- if not textual_parts:
- return ""
- if len(textual_parts) == 1:
- return textual_parts[0]
- return "\n".join(f"- {part}" for part in textual_parts)
-
- def _format_single_msg(
- self,
- msg: Msg,
- index: int | None = None,
- add_time: bool = True,
- ) -> tuple[str, int]:
- """Format a single Msg into string representation.
-
- Similar to Message.format_message style.
-
- Args:
- msg: The Msg object to format.
- index: Optional message index for round numbering.
- add_time: Whether to include timestamp.
-
- Returns:
- Tuple of (formatted_string, token_count).
- """
- lines = []
- token_count = 0
-
- # Build header: "round{index} [{timestamp}] {role}:"
- prefix = f"round{index} " if index is not None else ""
- time_str = f"[{msg.timestamp}] " if add_time and msg.timestamp else ""
- role_str = msg.name or msg.role
- header = f"{prefix}{time_str}{role_str}:"
- lines.append(header)
- token_count += safe_count_str_tokens(self._token_counter, header)
-
- # Process content blocks
- for block in msg.get_content_blocks():
- typ = block.get("type")
-
- if typ == "text":
- text_content = truncate_text(block.get("text", ""), self.max_length)
- if text_content:
- lines.append(text_content)
- token_count += safe_count_str_tokens(self._token_counter, text_content)
-
- elif typ == "thinking":
- # Skip thinking blocks to save tokens
- pass
-
- elif typ in ["image", "audio", "video"]:
- source = block.get("source", {})
- url = source.get("url", "")
- if url:
- lines.append(f"[{typ}] {url}")
- else:
- lines.append(f"[{typ}]")
- # Estimate fixed token cost for media reference
- token_count += 10
-
- elif typ == "tool_use":
- tool_name = block.get("name", "")
- tool_input = block.get("input", {})
- try:
- arguments_str = json.dumps(tool_input, ensure_ascii=False)
- except (TypeError, ValueError):
- arguments_str = str(tool_input)
- truncated_args = truncate_text(arguments_str, self.max_length)
- tool_line = f" - tool_call={tool_name} params={truncated_args}"
- lines.append(tool_line)
- token_count += safe_count_str_tokens(self._token_counter, tool_line)
-
- elif typ == "tool_result":
- tool_name = block.get("name", "")
- output = block.get("output", "")
- formatted_output = self._format_tool_result_output(output)
- truncated_output = truncate_text(formatted_output, self.max_length)
- if truncated_output:
- result_line = f" - tool_result={tool_name} output={truncated_output}"
- lines.append(result_line)
- token_count += safe_count_str_tokens(self._token_counter, result_line)
-
- else:
- logger.warning(
- "Unsupported block type %s in message, skipped.",
- typ,
- )
-
- return "\n".join(lines), token_count
-
- def format(
- self,
- msgs: list[Msg],
- add_time: bool = True,
- add_index: bool = True,
- ) -> str:
- """Format list of Msg into a single formatted string.
-
- Messages are processed in reverse order (newest first) and older
- messages are skipped when token count exceeds memory_compact_threshold.
-
- Args:
- msgs: List of Msg objects to format.
- add_time: Whether to include timestamp in each message.
- add_index: Whether to include round index in each message.
-
- Returns:
- Formatted string with all messages joined by newlines.
- """
- if not msgs:
- return ""
-
- formatted_parts: list[str] = []
- total_token_count = 0
-
- # Process messages in reverse order (newest first)
- for i in range(len(msgs) - 1, -1, -1):
- msg = msgs[i]
- index = i if add_index else None
-
- formatted_msg, msg_token_count = self._format_single_msg(
- msg,
- index=index,
- add_time=add_time,
- )
-
- # Always include current message first, then check threshold, at least one msg
- formatted_parts.append(formatted_msg)
- total_token_count += msg_token_count
-
- # Check if we should stop adding older messages
- if total_token_count >= self._memory_compact_threshold:
- logger.info(
- "Skipping older messages: token count %d >= %d",
- total_token_count,
- self._memory_compact_threshold,
- )
- break
-
- # Reverse to restore chronological order
- formatted_parts.reverse()
-
- return "\n\n".join(formatted_parts)
diff --git a/reme/memory/file_based/reme_chat_formatter.py b/reme/memory/file_based/reme_chat_formatter.py
deleted file mode 100644
index b70e700c..00000000
--- a/reme/memory/file_based/reme_chat_formatter.py
+++ /dev/null
@@ -1,29 +0,0 @@
-"""ReMe chat formatter."""
-
-from typing import Any
-
-from agentscope.formatter import OpenAIChatFormatter
-from agentscope.token import HuggingFaceTokenCounter
-
-from .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
diff --git a/reme/memory/file_based/reme_in_memory_memory.py b/reme/memory/file_based/reme_in_memory_memory.py
index 0e915382..16f18726 100644
--- a/reme/memory/file_based/reme_in_memory_memory.py
+++ b/reme/memory/file_based/reme_in_memory_memory.py
@@ -1,31 +1,23 @@
"""Custom memory implementation with bugfixes and extensions."""
-import logging
-
-from agentscope.agent._react_agent import _MemoryMark
-from agentscope.formatter import FormatterBase
+from agentscope.agent._react_agent import _MemoryMark # noqa
from agentscope.memory import InMemoryMemory
from agentscope.message import Msg
from agentscope.token import HuggingFaceTokenCounter
-from .utils import safe_count_message_tokens, safe_count_str_tokens, _get_block_tokens
+from .as_msg_handler import AsMsgHandler
+from ...core.utils import get_std_logger
-logger = logging.getLogger(__name__)
+logger = get_std_logger()
class ReMeInMemoryMemory(InMemoryMemory):
"""Extended InMemoryMemory with bugfixes and summary support."""
- def __init__(
- self,
- token_counter: HuggingFaceTokenCounter,
- formatter: FormatterBase,
- max_input_length: int = 0,
- ):
+ def __init__(self, token_counter: HuggingFaceTokenCounter):
super().__init__()
self._token_counter: HuggingFaceTokenCounter = token_counter
- self._formatter: FormatterBase = formatter
- self._max_input_length: int = max_input_length
+ self._msg_handler: AsMsgHandler = AsMsgHandler(token_counter)
async def get_memory(
self,
@@ -127,9 +119,12 @@ Use it as context to maintain continuity.
"""Clear the content."""
self.content.clear()
- async def estimate_tokens(self) -> dict:
+ async def estimate_tokens(self, max_input_length: int) -> dict:
"""Estimate token usage for current memory.
+ Args:
+ max_input_length: Max input length for context usage calculation.
+
Returns:
Dict containing detailed token statistics:
- total_messages: Number of messages
@@ -138,7 +133,7 @@ Use it as context to maintain continuity.
- estimated_tokens: Total estimated tokens
- max_input_length: Max input length from config
- context_usage_ratio: Usage percentage
- - messages_detail: List of per-message token details
+ - messages_detail: List of per-message AsMsgStat objects
"""
messages = await self.get_memory(
exclude_mark=_MemoryMark.COMPRESSED,
@@ -146,62 +141,18 @@ Use it as context to maintain continuity.
)
compressed_summary = self.get_compressed_summary()
- compressed_summary_tokens = safe_count_str_tokens(self._token_counter, compressed_summary)
+ compressed_summary_tokens = self._msg_handler.count_str_token(compressed_summary)
- # Calculate total token count using formatter
- prompt = await self._formatter.format(msgs=messages)
- messages_tokens = safe_count_message_tokens(self._token_counter, prompt)
+ # Build per-message token details using AsMsgHandler
+ messages_detail = [self._msg_handler.stat_message(msg) for msg in messages]
+
+ # Calculate total message tokens from stats
+ messages_tokens = sum(stat.total_tokens for stat in messages_detail)
estimated_tokens = messages_tokens + compressed_summary_tokens
# Calculate context usage ratio
- max_input_length = self._max_input_length
context_usage_ratio = (estimated_tokens / max_input_length * 100) if max_input_length > 0 else 0
- # Build per-message token details
- messages_detail = []
- for i, msg in enumerate(messages, 1):
- msg_detail = {
- "index": i,
- "role": msg.role,
- "text_tokens": 0,
- "blocks": [],
- "preview": "",
- }
- try:
- content = msg.content
- if isinstance(content, str):
- text_tokens = safe_count_str_tokens(self._token_counter, content)
- msg_detail["text_tokens"] = text_tokens
- msg_detail["preview"] = f"{content[:100]}..." if len(content) > 100 else content
- else:
- total_tokens = 0
- text_parts = []
- for block in content:
- if not isinstance(block, dict):
- continue
- block_type = block.get("type", "unknown")
- block_tokens, block_str = _get_block_tokens(
- block,
- block_type,
- self._token_counter,
- )
- total_tokens += block_tokens
- text_parts.append(block_str)
- msg_detail["blocks"].append(
- {
- "type": block_type,
- "tokens": block_tokens,
- },
- )
- msg_detail["text_tokens"] = total_tokens
- text_preview = "".join(text_parts)
- msg_detail["preview"] = f"{text_preview[:100]}..." if len(text_preview) > 100 else text_preview
- except Exception as e:
- msg_detail["error"] = str(e)
- msg_detail["preview"] = f""
-
- messages_detail.append(msg_detail)
-
return {
"total_messages": len(messages),
"compressed_summary_tokens": compressed_summary_tokens,
@@ -212,25 +163,28 @@ Use it as context to maintain continuity.
"messages_detail": messages_detail,
}
- async def get_history_str(self) -> str:
+ async def get_history_str(self, max_input_length: int) -> str:
"""Get formatted history string similar to /history command output.
+ Args:
+ max_input_length: Max input length for context usage calculation.
+
Returns:
Formatted string containing conversation history details
"""
- stats = await self.estimate_tokens()
+ stats = await self.estimate_tokens(max_input_length)
lines = []
- for msg_detail in stats["messages_detail"]:
+ for i, msg_stat in enumerate(stats["messages_detail"], 1):
blocks_info = ""
- if msg_detail["blocks"]:
- block_strs = [f"{b['type']}(tokens={b['tokens']})" for b in msg_detail["blocks"]]
+ if msg_stat.content:
+ block_strs = [f"{b.block_type}(tokens={b.token_count})" for b in msg_stat.content]
blocks_info = f"\n content: [{', '.join(block_strs)}]"
lines.append(
- f"[{msg_detail['index']}] **{msg_detail['role']}** "
- f"(text_tokens={msg_detail['text_tokens']})"
- f"{blocks_info}\n preview: {msg_detail['preview']}",
+ f"[{i}] **{msg_stat.role}** "
+ f"(total_tokens={msg_stat.total_tokens})"
+ f"{blocks_info}\n preview: {msg_stat.preview}",
)
return (
diff --git a/reme/memory/file_based/utils.py b/reme/memory/file_based/utils.py
deleted file mode 100644
index f460f96f..00000000
--- a/reme/memory/file_based/utils.py
+++ /dev/null
@@ -1,271 +0,0 @@
-"""Utility functions for working with text."""
-
-import logging
-from pathlib import Path
-
-from agentscope.token import HuggingFaceTokenCounter
-
-logger = logging.getLogger(__name__)
-
-# Unique marker for truncated text
-TRUNCATION_MARKER_START = "<<>>"
-TRUNCATION_MARKER_END = "<<>>"
-
-
-def truncate_text(text: str, max_length: int) -> str:
- """Truncate text to max length, keeping head and tail portions.
-
- Args:
- text: The text to truncate
- max_length: Maximum allowed length
-
- Returns:
- Truncated text with unique markers indicating truncation
- """
- text = str(text) if text else ""
- if not text:
- return text
-
- if len(text) <= max_length:
- return text
-
- half_length = max_length // 2
- truncated_chars = len(text) - max_length
- logger.debug(
- "Text truncated: original %d chars, kept head %d + tail %d, removed %d chars.",
- len(text),
- half_length,
- half_length,
- truncated_chars,
- )
- return (
- f"{text[:half_length]}\n\n{TRUNCATION_MARKER_START} "
- f"({truncated_chars} characters omitted) "
- f"{TRUNCATION_MARKER_END}\n\n{text[-half_length:]}"
- )
-
-
-def is_truncated(text: str) -> bool:
- """Check if the text has been truncated (contains truncation markers).
-
- Args:
- text: The text to check
-
- Returns:
- bool: True if text contains truncation markers, False otherwise
- """
- if not text:
- return False
- return TRUNCATION_MARKER_START in text and TRUNCATION_MARKER_END in text
-
-
-def _extract_text_from_messages(messages: list[dict]) -> str:
- """Extract text content from messages and concatenate into a string.
-
- Handles various message formats:
- - Simple string content: {"role": "user", "content": "hello"}
- - List content with text blocks:
- {"role": "user", "content": [{"type": "text", "text": "hello"}]}
- - List content with tool_result blocks:
- {"role": "user", "content": [{"type": "tool_result", "output": "..."}]}
-
- Args:
- messages: List of message dictionaries in chat format.
-
- Returns:
- str: Concatenated text content from all messages.
- """
- parts = []
- for msg in messages:
- content = msg.get("content", "")
- if isinstance(content, str):
- parts.append(content)
- elif isinstance(content, list):
- for block in content:
- if isinstance(block, dict):
- block_type = block.get("type", "")
- if block_type == "tool_result":
- output = block.get("output", "")
- if isinstance(output, str) and output:
- parts.append(output)
- elif isinstance(output, list):
- for sub in output:
- if isinstance(sub, dict):
- sub_text = sub.get("text") or sub.get("content", "")
- if sub_text:
- parts.append(str(sub_text))
- else:
- text = block.get("text") or block.get("content", "")
- if text:
- parts.append(str(text))
- elif isinstance(block, str):
- parts.append(block)
- return "\n".join(parts)
-
-
-def safe_count_message_tokens(
- token_counter: HuggingFaceTokenCounter,
- messages: list[dict],
-) -> int:
- """Safely count tokens in messages with fallback estimation.
-
- This is a wrapper around count_message_tokens that catches exceptions
- and falls back to a character-based estimation (len // 4) if the
- tokenizer fails.
-
- Args:
- token_counter: Token counter instance.
- messages: List of message dictionaries in chat format.
-
- Returns:
- int: The estimated number of tokens in the messages.
- """
- try:
- text = _extract_text_from_messages(messages)
- token_ids = token_counter.tokenizer.encode(text)
- token_count = len(token_ids)
- return token_count
-
- except Exception as e:
- # Fallback to character-based estimation
- text = _extract_text_from_messages(messages)
- estimated_tokens = len(text) // 4
- logger.warning(
- "Failed to count tokens: %s, using estimated_tokens=%d",
- e,
- estimated_tokens,
- )
- return estimated_tokens
-
-
-def safe_count_str_tokens(
- token_counter: HuggingFaceTokenCounter,
- text: str,
-) -> int:
- """Safely count tokens in a string with fallback estimation.
-
- Uses the tokenizer to count tokens in the given text. If the tokenizer
- fails, falls back to a character-based estimation (len // 4).
-
- Args:
- token_counter: Token counter instance.
- text: The string to count tokens for.
-
- Returns:
- int: The estimated number of tokens in the string.
- """
- try:
- token_ids = token_counter.tokenizer.encode(text)
- token_count = len(token_ids)
- return token_count
- except Exception as e:
- # Fallback to character-based estimation
- estimated_tokens = len(text) // 4
- logger.warning(
- "Failed to count string tokens: %s, using estimated_tokens=%d",
- e,
- estimated_tokens,
- )
- return estimated_tokens
-
-
-def _get_block_tokens( # pylint: disable=too-many-return-statements
- block: dict,
- block_type: str,
- token_counter: HuggingFaceTokenCounter,
-) -> tuple[int, str]:
- """Get token count and content string for different block types.
-
- Args:
- block: The content block dict
- block_type: The type of the block
-
- Returns:
- Tuple of (token count, content string)
- """
- if block_type == "text":
- text = block.get("text", "")
- return (safe_count_str_tokens(token_counter, text), text) if text else (0, "")
-
- if block_type == "thinking":
- thinking = block.get("thinking", "")
- return (safe_count_str_tokens(token_counter, thinking), thinking) if thinking else (0, "")
-
- if block_type == "tool_use":
- # Count input dict and raw_input string
- input_dict = block.get("input", {})
- raw_input = block.get("raw_input", "")
- input_str = str(input_dict) if input_dict else ""
- total = input_str + raw_input
- return (safe_count_str_tokens(token_counter, total), total) if total else (0, "")
-
- if block_type == "tool_result":
- output = block.get("output")
- if isinstance(output, str):
- return (safe_count_str_tokens(token_counter, output), output) if output else (0, "")
-
- if isinstance(output, list):
- # Recursively count tokens in nested blocks
- total_tokens = 0
- total_str = ""
- for item in output:
- if isinstance(item, dict):
- item_type = item.get("type", "unknown")
- item_tokens, item_str = _get_block_tokens(item, item_type, token_counter)
- total_tokens += item_tokens
- total_str += item_str
- return total_tokens, total_str
- return 0, ""
-
- if block_type in ("image", "audio", "video"):
- # For media blocks, count the URL or indicate base64 size
- source = block.get("source", {})
- if source.get("type") == "url":
- url = source.get("url", "")
- return safe_count_str_tokens(token_counter, url), url
- if source.get("type") == "base64":
- # Base64 data can be large, return approximate token count
- data = source.get("data", "")
- return (len(data) // 4, "[base64]") if data else (0, "")
- return 0, ""
-
- return 0, ""
-
-
-_token_counter = None
-
-
-def get_token_counter():
- """Get or initialize the global token counter instance.
-
- Returns:
- TokenCounterBase: The token counter instance for Qwen models.
-
- Raises:
- RuntimeError: If token counter initialization fails.
- """
- global _token_counter
- if _token_counter is None:
- # 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
diff --git a/reme/memory/tools/file/__init__.py b/reme/memory/tools/file/__init__.py
new file mode 100644
index 00000000..8234e60d
--- /dev/null
+++ 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/memory/file_based/file_io.py b/reme/memory/tools/file/file_io.py
similarity index 100%
rename from reme/memory/file_based/file_io.py
rename to reme/memory/tools/file/file_io.py
diff --git a/reme/reme_light.py b/reme/reme_light.py
index 7c13c72e..5266ae76 100644
--- a/reme/reme_light.py
+++ b/reme/reme_light.py
@@ -15,53 +15,26 @@ Key Features:
"""
import asyncio
-import logging
-import os
-import platform
from pathlib import Path
from agentscope.formatter import FormatterBase
from agentscope.message import Msg, TextBlock
-from agentscope.model import ChatModelBase, OpenAIChatModel
+from agentscope.model import ChatModelBase
from agentscope.token import HuggingFaceTokenCounter
from agentscope.tool import Toolkit, ToolResponse
from .config import ReMeConfigParser
from .core import Application
-from .memory.file_based import Compactor, Summarizer, ToolResultCompactor, ReMeInMemoryMemory, ReMeChatFormatter
-from .memory.file_based.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, AsMsgHandler
from .memory.tools import MemorySearch
-from .core.utils import load_env
+from .memory.tools.file import FileIO
-logger = logging.getLogger(__name__)
+logger = get_std_logger()
class ReMeLight(Application):
- """
- ReMe Light Application Class
-
- A specialized application class that extends ReMe's core Application framework
- with advanced memory management capabilities. This class is designed to handle
- long-running conversations by providing intelligent memory compaction,
- summarization, and semantic search features.
-
- Attributes:
- working_path (Path): Absolute path to the working directory for storing data
- memory_path (Path): Path to the memory storage directory
- tool_result_path (Path): Path to store large tool result files
- chat_model (ChatModelBase): Language model for generating summaries and processing
- formatter (FormatterBase): Formatter for structuring model inputs/outputs
- token_counter (HuggingFaceTokenCounter): Token counting utility for length management
- toolkit (Toolkit): Collection of tools available to the application
- max_input_length (int): Maximum allowed input length in tokens
- memory_compact_threshold (int): Threshold at which memory compaction triggers
- language (str): Language code for localization ("zh" for Chinese, empty for English)
- vector_weight (float): Weight for vector search in hybrid search (0.0-1.0)
- candidate_multiplier (float): Multiplier for candidate retrieval in search
- tool_result_threshold (int): Size threshold for tool result compaction
- retention_days (int): Number of days to retain tool result files
- summary_tasks (list[asyncio.Task]): List of background summarization tasks
- """
+ """ReMe Light Application Class"""
def __init__(
self,
@@ -70,81 +43,27 @@ class ReMeLight(Application):
llm_base_url: str | None = None,
embedding_api_key: str | None = None,
embedding_base_url: str | None = None,
- chat_model: ChatModelBase | None = None,
- formatter: FormatterBase | None = None,
- token_counter: HuggingFaceTokenCounter | None = None,
- toolkit: Toolkit | None = None,
- max_input_length: int = 128000,
- memory_compact_ratio: float = 0.7,
- language: str = "zh",
+ 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
- # All application data will be stored under this path
self.working_path = Path(working_dir).absolute()
self.working_path.mkdir(parents=True, exist_ok=True)
-
- # Create memory storage directory for persistent memory files
self.memory_path = self.working_path / "memory"
self.memory_path.mkdir(parents=True, exist_ok=True)
-
- # Create tool result directory for storing large tool outputs
self.tool_result_path = self.working_path / "tool_result"
self.tool_result_path.mkdir(parents=True, exist_ok=True)
- # Initialize runtime parameters (will be updated via update_params)
- self.max_input_length: int = 0
- self.memory_compact_threshold: int = 0
- self.language: str = ""
-
- # Apply initial parameter configuration
- self.update_params(
- max_input_length=max_input_length,
- memory_compact_ratio=memory_compact_ratio,
- language=language,
- )
-
- # Store configuration parameters
self.vector_weight: float = vector_weight
self.candidate_multiplier: float = candidate_multiplier
self.tool_result_threshold: int = tool_result_threshold
self.retention_days: int = retention_days
- load_env()
-
- llm_model_name = self._safe_str("LLM_MODEL_NAME", "")
- embedding_model_name = self._safe_str("EMBEDDING_MODEL_NAME", "")
- embedding_dimensions = self._safe_int("EMBEDDING_DIMENSIONS", 1024)
- embedding_cache_enabled = self._safe_str("EMBEDDING_CACHE_ENABLED", "true").lower() == "true"
- embedding_max_cache_size = self._safe_int("EMBEDDING_MAX_CACHE_SIZE", 2000)
- embedding_max_input_length = self._safe_int("EMBEDDING_MAX_INPUT_LENGTH", 8192)
- embedding_max_batch_size = self._safe_int("EMBEDDING_MAX_BATCH_SIZE", 10)
-
- # Determine if vector search should be enabled based on configuration
- # Vector search requires either an API key or a local model name
- vector_enabled = bool(embedding_api_key) or bool(embedding_model_name)
- if vector_enabled:
- logger.info("Vector search enabled.")
- else:
- logger.warning(
- "Vector search disabled. Memory search functionality will be restricted. "
- "To enable, configure: EMBEDDING_API_KEY, EMBEDDING_BASE_URL, EMBEDDING_MODEL_NAME.",
- )
-
- # Check if full-text search (FTS) is enabled via environment variable
- fts_enabled = os.environ.get("FTS_ENABLED", "true").lower() == "true"
-
- # Determine the memory store backend to use
- # "auto" selects based on platform (local for Windows, chroma otherwise)
- memory_store_backend = os.environ.get("MEMORY_STORE_BACKEND", "auto")
- if memory_store_backend == "auto":
- memory_backend = "local" if platform.system() == "Windows" else "chroma"
- else:
- memory_backend = memory_store_backend
-
# Initialize the parent Application class with comprehensive configuration
super().__init__(
llm_api_key=llm_api_key,
@@ -156,21 +75,9 @@ class ReMeLight(Application):
enable_logo=False,
log_to_console=False,
parser=ReMeConfigParser,
- default_embedding_model_config={
- "model_name": embedding_model_name,
- "dimensions": embedding_dimensions,
- "enable_cache": embedding_cache_enabled,
- "use_dimensions": False,
- "max_cache_size": embedding_max_cache_size,
- "max_input_length": embedding_max_input_length,
- "max_batch_size": embedding_max_batch_size,
- },
- default_file_store_config={
- "backend": memory_backend,
- "store_name": "copaw",
- "vector_enabled": vector_enabled,
- "fts_enabled": fts_enabled,
- },
+ default_as_llm_config=default_as_llm_config,
+ default_embedding_model_config=default_embedding_model_config,
+ default_file_store_config=default_file_store_config,
default_file_watcher_config={
"watch_paths": [
str(self.working_path / "MEMORY.md"),
@@ -180,107 +87,21 @@ class ReMeLight(Application):
},
)
- if chat_model is not None:
- self.chat_model: ChatModelBase = chat_model
- else:
- # add more params later
- self.chat_model = OpenAIChatModel(
- api_key=os.environ["LLM_API_KEY"],
- client_kwargs={"base_url": os.environ["LLM_BASE_URL"]},
- model_name=llm_model_name,
- )
-
- if token_counter is not None:
- self.token_counter: HuggingFaceTokenCounter = token_counter
- else:
- self.token_counter = get_token_counter()
-
- if formatter is not None:
- self.formatter: FormatterBase = formatter
- else:
- self.formatter = ReMeChatFormatter(token_counter=self.token_counter)
- self.toolkit: Toolkit | None = toolkit
-
# Initialize list to track background summarization tasks
self.summary_tasks: list[asyncio.Task] = []
- def update_params(
- self,
- max_input_length: int,
- memory_compact_ratio: float,
- language: str,
- ):
- """
- Update runtime parameters for memory management.
-
- This method allows dynamic adjustment of memory-related parameters during
- runtime. It recalculates the memory compaction threshold based on the
- new input length and compaction ratio.
-
- Args:
- max_input_length (int): New maximum input length in tokens
- memory_compact_ratio (float): Ratio at which to trigger compaction (0.0-1.0)
- language (str): Language code for localization ("zh" or other)
-
- Note:
- The memory_compact_threshold is calculated as:
- max_input_length * memory_compact_ratio * 0.9
- The 0.9 factor provides a safety margin before reaching the absolute limit
- """
- # Update the maximum allowed input length
- self.max_input_length = max_input_length
-
- # Calculate compaction threshold with safety margin
- # This ensures compaction happens before hitting the hard limit
- self.memory_compact_threshold = int(max_input_length * memory_compact_ratio * 0.9)
-
- # Set language for localization
- if language == "zh":
- self.language = "zh"
- else:
- self.language = ""
-
@staticmethod
- def _safe_str(key: str, default: str) -> str:
- """
- Safely retrieve a string value from an environment variable.
+ 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:
- key (str): The name of the environment variable to retrieve
- default (str): The default value to return if the variable is not set
+ max_input_length: Maximum input length in tokens.
+ compact_ratio: Ratio of the input length to use as the threshold.
Returns:
- str: The value of the environment variable, or the default if not set
+ Computed compaction threshold as an integer.
"""
- return os.environ.get(key, default)
-
- @staticmethod
- def _safe_int(key: str, default: int) -> int:
- """
- Safely retrieve an integer value from an environment variable.
-
- This method handles cases where the environment variable is not set
- or contains a non-integer value by returning the specified default.
-
- Args:
- key (str): The name of the environment variable to retrieve
- default (int): The default value to return on failure or if not set
-
- Returns:
- int: The integer value of the environment variable, or the default
-
- Note:
- Logs a warning if the value exists but cannot be parsed as an integer
- """
- value = os.environ.get(key)
- if value is None:
- return default
-
- try:
- return int(value)
- except ValueError:
- logger.warning(f"Invalid int value '{value}' for key '{key}', using default {default}")
- return default
+ return int(max_input_length * compact_ratio * 0.9)
def _cleanup_tool_results(self) -> int:
"""
@@ -292,10 +113,6 @@ class ReMeLight(Application):
Returns:
int: The number of files that were successfully deleted
-
- Note:
- Exceptions during cleanup are logged but do not raise errors,
- ensuring the application continues to function even if cleanup fails
"""
try:
# Create a compactor instance with current configuration
@@ -312,67 +129,18 @@ class ReMeLight(Application):
return 0
async def start(self):
- """
- Start the application lifecycle.
-
- This method initializes the application by calling the parent class's
- start method and performs initial cleanup of expired tool result files.
-
- Returns:
- The result from the parent class's start method
-
- Note:
- Tool result cleanup runs after successful startup to ensure
- the application is fully initialized before performing maintenance
- """
- # Initialize parent application components
+ """Start the application lifecycle."""
result = await super().start()
- # Perform initial cleanup of old tool result files
self._cleanup_tool_results()
return result
async def close(self) -> bool:
- """
- Close the application and perform cleanup.
-
- This method performs final cleanup of expired tool result files before
- shutting down the application through the parent class's close method.
-
- Returns:
- bool: True if shutdown was successful, False otherwise
-
- Note:
- Cleanup is performed before calling parent close to ensure
- all resources are available during the cleanup process
- """
- # Clean up tool results before shutting down
+ """Close the application and perform cleanup."""
self._cleanup_tool_results()
- # Shutdown parent application components
return await super().close()
- async def compact_tool_result(
- self,
- messages: list[Msg],
- ) -> list[Msg]:
- """
- Compact tool results by truncating large outputs and saving full content to files.
-
- This method processes a list of messages and identifies tool results that exceed
- the configured size threshold. Large tool outputs are truncated in the message
- list while their full content is saved to files for later retrieval.
-
- Args:
- messages (list[Msg]): List of messages to process for tool result compaction
-
- Returns:
- list[Msg]: The processed message list with large tool results compacted
-
- Note:
- - Tool results below the threshold remain unchanged in the messages
- - Large results are replaced with truncated versions and file references
- - Expired files are cleaned up as part of the compaction process
- - If compaction fails, the original messages are returned unchanged
- """
+ async def compact_tool_result(self, messages: list[Msg]) -> list[Msg]:
+ """Compact tool results by truncating large outputs and saving full content to files."""
try:
# Create compactor with instance configuration
compactor = ToolResultCompactor(
@@ -394,38 +162,30 @@ class ReMeLight(Application):
logger.exception(f"Error compacting tool results: {e}")
return messages
- async def compact_memory(self, messages: list[Msg], previous_summary: str = "") -> str:
- """
- Compact a list of messages into a condensed summary.
-
- This method uses the Compactor to reduce the length of message history
- while preserving essential information. It's useful when conversation
- history approaches the maximum input length limit.
-
- Args:
- messages (list[Msg]): The list of messages to compact
- previous_summary (str): Optional previous summary to incorporate
- into the compaction process for continuity
-
- Returns:
- str: A compacted summary of the messages, or empty string on failure
-
- Note:
- - Compaction uses the configured language model to generate summaries
- - The compaction threshold determines when compaction is triggered
- - If compaction fails, an empty string is returned
- """
+ 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 = "",
+ ) -> str:
+ """Compact a list of messages into a condensed summary."""
try:
- # Initialize compactor with current configuration
+ if token_counter is None:
+ token_counter = get_hf_token_counter()
+
compactor = Compactor(
- memory_compact_threshold=self.memory_compact_threshold,
- chat_model=self.chat_model,
- formatter=self.formatter,
- token_counter=self.token_counter,
- language=self.language,
+ 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,
+ language=language if language == "zh" else "",
)
- # Execute compaction with optional previous summary context
return await compactor.call(
messages=messages,
previous_summary=previous_summary,
@@ -437,66 +197,138 @@ class ReMeLight(Application):
logger.exception(f"Error compacting memory: {e}")
return ""
- async def summary_memory(self, messages: list[Msg]) -> str:
- """
- Generate a comprehensive summary of the given messages.
-
- This method uses the Summarizer to create a detailed summary of the
- conversation history, which can be stored as persistent memory. Unlike
- compaction, summarization aims to capture key information in a format
- suitable for long-term storage and retrieval.
-
- Args:
- messages (list[Msg]): The list of messages to summarize
-
- Returns:
- str: A generated summary of the messages, or empty string on failure
-
- Note:
- - Summarization may use tools from the toolkit to enhance the summary
- - The summary is typically stored in the memory directory
- - If summarization fails, an empty string is returned
- """
+ 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:
- # Initialize summarizer with working directories and configuration
+ 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)
+
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,
- toolkit=self.toolkit,
- language=self.language,
+ memory_compact_threshold=self.calculate_memory_compact_threshold(max_input_length, compact_ratio),
+ token_counter=token_counter,
+ toolkit=toolkit,
+ 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 pre_reasoning_hook(
+ self,
+ messages: list[Msg],
+ system_prompt: str = "",
+ compressed_summary: str = "",
+ 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,
+ memory_compact_reserve: int = 10000,
+ enable_tool_result_compact: bool = True,
+ tool_result_compact_keep_n: int = 3,
+ ) -> tuple[list[Msg], str]:
+ """Hook called before reasoning."""
+ if token_counter is None:
+ token_counter = get_hf_token_counter()
+
+ msg_handler = AsMsgHandler(token_counter=token_counter)
+
+ system_token_count = msg_handler.count_str_token(system_prompt)
+ compressed_token_count = msg_handler.count_str_token(compressed_summary)
+ memory_compact_threshold = self.calculate_memory_compact_threshold(max_input_length, compact_ratio)
+ left_compact_threshold = memory_compact_threshold - (system_token_count + compressed_token_count)
+ logger.info(f"Left compact threshold: {left_compact_threshold}")
+
+ if enable_tool_result_compact and tool_result_compact_keep_n > 0:
+ compact_msgs = messages[:-tool_result_compact_keep_n]
+ await self.compact_tool_result(compact_msgs)
+
+ messages_to_compact, messages_to_keep, is_valid = msg_handler.context_check(
+ messages=messages,
+ memory_compact_threshold=left_compact_threshold,
+ memory_compact_reserve=memory_compact_reserve,
+ )
+
+ if not messages_to_compact:
+ return messages, compressed_summary
+
+ if not is_valid:
+ logger.warning("Invalid messages to compact, skipping.")
+ return messages, compressed_summary
+
+ self.add_async_summary_task(
+ messages=messages_to_compact,
+ as_llm=as_llm,
+ as_llm_formatter=as_llm_formatter,
+ token_counter=token_counter,
+ toolkit=toolkit,
+ language=language,
+ max_input_length=max_input_length,
+ compact_ratio=compact_ratio,
+ )
+
+ compressed_summary = await self.compact_memory(
+ messages=messages_to_compact,
+ as_llm=as_llm,
+ as_llm_formatter=as_llm_formatter,
+ token_counter=token_counter,
+ language=language,
+ max_input_length=max_input_length,
+ compact_ratio=compact_ratio,
+ previous_summary=compressed_summary,
+ )
+
+ return messages_to_keep, compressed_summary
+
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():
@@ -526,6 +358,7 @@ class ReMeLight(Application):
except asyncio.CancelledError:
logger.warning("Summary task was cancelled while waiting.")
result += "Summary task was cancelled.\n"
+
except Exception as e:
logger.exception(f"Summary task failed: {e}")
result += f"Summary task failed: {e}\n"
@@ -534,48 +367,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.
@@ -613,12 +404,25 @@ class ReMeLight(Application):
# Validate and clamp max_results to valid range [1, 100]
if isinstance(max_results, int):
max_results = min(max(max_results, 1), 100)
+
+ elif isinstance(max_results, str):
+ try:
+ max_results = min(max(int(max_results), 1), 100)
+ except ValueError:
+ max_results = 5
else:
max_results = 5
# Validate and clamp min_score to valid range [0.001, 0.999]
if isinstance(min_score, (int, float)):
- min_score = min(max(min_score, 0.001), 0.999)
+ min_score = float(min(max(min_score, 0.001), 0.999))
+
+ elif isinstance(min_score, str):
+ try:
+ min_score = float(min(max(float(min_score), 0.001), 0.999))
+ except ValueError:
+ min_score = 0.1
+
else:
min_score = 0.1
@@ -646,26 +450,10 @@ class ReMeLight(Application):
],
)
- def get_in_memory_memory(self):
- """
- Create and return an in-memory memory instance.
+ @staticmethod
+ def get_in_memory_memory(token_counter: HuggingFaceTokenCounter | None = None):
+ """Create and return an in-memory memory instance."""
+ if token_counter is None:
+ token_counter = get_hf_token_counter()
- This method instantiates a ReMeInMemoryMemory object configured with
- the current application's token counter, formatter, and input length limits.
- The in-memory memory provides fast, temporary storage for conversation
- context without persistence.
-
- Returns:
- ReMeInMemoryMemory: A configured in-memory memory instance ready
- for storing and retrieving conversation messages
-
- Note:
- - In-memory memory is volatile and cleared when the instance is destroyed
- - Useful for managing conversation context within a single session
- - Shares the same token counter and formatter as the main application
- """
- return ReMeInMemoryMemory(
- token_counter=self.token_counter,
- formatter=self.formatter,
- max_input_length=self.max_input_length,
- )
+ return ReMeInMemoryMemory(token_counter=token_counter)
diff --git a/reme/extension/cli/__init__.py b/test/cli/__init__.py
similarity index 100%
rename from reme/extension/cli/__init__.py
rename to test/cli/__init__.py
diff --git a/reme/extension/cli/fb_cli.py b/test/cli/fb_cli.py
similarity index 100%
rename from reme/extension/cli/fb_cli.py
rename to test/cli/fb_cli.py
diff --git a/reme/extension/cli/fb_cli.yaml b/test/cli/fb_cli.yaml
similarity index 100%
rename from reme/extension/cli/fb_cli.yaml
rename to test/cli/fb_cli.yaml
diff --git a/reme/extension/cli/fb_compactor.py b/test/cli/fb_compactor.py
similarity index 100%
rename from reme/extension/cli/fb_compactor.py
rename to test/cli/fb_compactor.py
diff --git a/reme/extension/cli/fb_compactor.yaml b/test/cli/fb_compactor.yaml
similarity index 100%
rename from reme/extension/cli/fb_compactor.yaml
rename to test/cli/fb_compactor.yaml
diff --git a/reme/extension/cli/fb_context_checker.py b/test/cli/fb_context_checker.py
similarity index 100%
rename from reme/extension/cli/fb_context_checker.py
rename to test/cli/fb_context_checker.py
diff --git a/reme/extension/cli/fb_summarizer.py b/test/cli/fb_summarizer.py
similarity index 100%
rename from reme/extension/cli/fb_summarizer.py
rename to test/cli/fb_summarizer.py
diff --git a/reme/extension/cli/fb_summarizer.yaml b/test/cli/fb_summarizer.yaml
similarity index 100%
rename from reme/extension/cli/fb_summarizer.yaml
rename to test/cli/fb_summarizer.yaml
diff --git a/reme/extension/reme_cli.py b/test/reme_cli.py
similarity index 100%
rename from reme/extension/reme_cli.py
rename to test/reme_cli.py
diff --git a/tests/light/test_compactor.py b/tests/light/test_compactor.py
index 32dfd9d3..8ae2a051 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,12 @@ 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 颜色码
@@ -101,9 +98,10 @@ def create_compactor():
"""Create a Compactor instance for testing."""
return Compactor(
memory_compact_threshold=4000,
- chat_model=get_dash_chat_model(),
- formatter=get_formatter(),
token_counter=get_token_counter(),
+ as_llm=get_dash_chat_model(),
+ as_llm_formatter=get_formatter(),
+ language="zh",
)
@@ -286,9 +284,9 @@ def test_low_threshold():
"""Test compaction with low memory threshold."""
compactor = Compactor(
memory_compact_threshold=500,
- chat_model=get_dash_chat_model(),
- formatter=get_formatter(),
token_counter=get_token_counter(),
+ as_llm=get_dash_chat_model(),
+ as_llm_formatter=get_formatter(),
)
messages = [
@@ -309,9 +307,9 @@ def test_high_threshold():
"""Test compaction with high memory threshold."""
compactor = Compactor(
memory_compact_threshold=10000,
- chat_model=get_dash_chat_model(),
- formatter=get_formatter(),
token_counter=get_token_counter(),
+ as_llm=get_dash_chat_model(),
+ as_llm_formatter=get_formatter(),
)
messages = [
diff --git a/tests/light/test_context_check.py b/tests/light/test_context_check.py
new file mode 100644
index 00000000..872d6e2d
--- /dev/null
+++ b/tests/light/test_context_check.py
@@ -0,0 +1,1285 @@
+"""Tests for AsMsgHandler.context_check method."""
+
+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
+
+logger = get_std_logger()
+
+
+# ANSI color codes
+class Colors:
+ """ANSI color codes for terminal output."""
+
+ GREEN = "\033[92m"
+ RED = "\033[91m"
+ YELLOW = "\033[93m"
+ BLUE = "\033[94m"
+ CYAN = "\033[96m"
+ BOLD = "\033[1m"
+ RESET = "\033[0m"
+
+
+def print_pass(test_name: str):
+ """Print test passed message."""
+ print(f"{Colors.GREEN}{Colors.BOLD}✓ {test_name} PASSED{Colors.RESET}")
+
+
+def print_fail(test_name: str, error: str):
+ """Print test failed message."""
+ print(f"{Colors.RED}{Colors.BOLD}✗ {test_name} FAILED: {error}{Colors.RESET}")
+
+
+def print_error(test_name: str, error: str):
+ """Print test error message."""
+ print(f"{Colors.YELLOW}{Colors.BOLD}⚠ {test_name} ERROR: {error}{Colors.RESET}")
+
+
+def print_test_header(test_name: str):
+ """Print test header."""
+ print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+ print(f"{Colors.BLUE}{Colors.BOLD}Running: {test_name}{Colors.RESET}")
+ print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+
+
+def create_handler() -> AsMsgHandler:
+ """Create an AsMsgHandler instance for testing."""
+ return AsMsgHandler(token_counter=get_token_counter())
+
+
+def verify_context_check_invariants(
+ handler: AsMsgHandler,
+ messages: list[Msg],
+ to_compact: list[Msg],
+ to_keep: list[Msg],
+ memory_compact_threshold: int,
+ memory_compact_reserve: int,
+ test_name: str,
+):
+ """Verify that context_check results satisfy all invariants.
+
+ This function checks:
+ 1. Threshold requirement: If total tokens <= threshold, no compaction should occur
+ 2. Reserve requirement: Kept messages' total tokens should not exceed reserve
+ 3. Order requirement: Both to_compact and to_keep should preserve original order
+
+ Args:
+ handler: The AsMsgHandler instance
+ messages: Original messages list
+ to_compact: Messages to compact returned by context_check
+ to_keep: Messages to keep returned by context_check
+ memory_compact_threshold: The threshold parameter used
+ memory_compact_reserve: The reserve parameter used
+ test_name: Name of the test for error reporting
+
+ Raises:
+ AssertionError: If any invariant is violated
+ """
+ # Calculate total tokens of original messages
+ total_tokens = sum(handler.stat_message(m).total_tokens for m in messages)
+
+ # 1. Threshold requirement check
+ if total_tokens <= memory_compact_threshold:
+ assert len(to_compact) == 0, (
+ f"[{test_name}] Threshold violation: total_tokens ({total_tokens}) <= "
+ f"threshold ({memory_compact_threshold}), but to_compact is not empty "
+ f"(has {len(to_compact)} messages)"
+ )
+ assert to_keep == messages, (
+ f"[{test_name}] Threshold violation: total_tokens ({total_tokens}) <= "
+ f"threshold ({memory_compact_threshold}), but to_keep differs from original messages"
+ )
+
+ # 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})"
+ )
+
+ # 3. Order requirement check - both lists should preserve original order
+ # Create a mapping of message id to original index
+ msg_to_idx = {id(m): i for i, m in enumerate(messages)}
+
+ # Check to_compact order
+ compact_indices = [msg_to_idx.get(id(m), -1) for m in to_compact]
+ for i in range(len(compact_indices) - 1):
+ assert compact_indices[i] < compact_indices[i + 1], (
+ f"[{test_name}] Order violation in to_compact: message at original index "
+ f"{compact_indices[i]} appears before message at index {compact_indices[i + 1]}"
+ )
+
+ # Check to_keep order
+ keep_indices = [msg_to_idx.get(id(m), -1) for m in to_keep]
+ for i in range(len(keep_indices) - 1):
+ assert keep_indices[i] < keep_indices[i + 1], (
+ f"[{test_name}] Order violation in to_keep: message at original index "
+ f"{keep_indices[i]} appears before message at index {keep_indices[i + 1]}"
+ )
+
+ # 4. Additional check: to_compact indices should all be less than to_keep indices
+ # (compact messages come from the beginning, keep messages come from the end)
+ if to_compact and to_keep:
+ max_compact_idx = max(compact_indices) if compact_indices else -1
+ min_keep_idx = min(keep_indices) if keep_indices else len(messages)
+ assert max_compact_idx < min_keep_idx, (
+ f"[{test_name}] Partition violation: max compact index ({max_compact_idx}) >= "
+ f"min keep index ({min_keep_idx}). Compact and keep should be a clean partition."
+ )
+
+ # 5. Check that all messages are accounted for (no duplicates, no missing)
+ assert len(to_compact) + len(to_keep) == len(messages), (
+ f"[{test_name}] Count mismatch: to_compact ({len(to_compact)}) + "
+ f"to_keep ({len(to_keep)}) != original ({len(messages)})"
+ )
+
+ 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"
+
+
+def create_user_msg(content: str) -> Msg:
+ """Create a user message."""
+ return Msg(name="user", role="user", content=content)
+
+
+def create_assistant_msg(content: str) -> Msg:
+ """Create an assistant message."""
+ return Msg(name="assistant", role="assistant", content=content)
+
+
+def create_tool_use_msg(tool_id: str, tool_name: str, tool_input: dict) -> Msg:
+ """Create a message with tool_use content block."""
+ return Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ {
+ "type": "tool_use",
+ "id": tool_id,
+ "name": tool_name,
+ "input": tool_input,
+ },
+ ],
+ )
+
+
+def create_tool_result_msg(tool_id: str, tool_name: str, output: str) -> Msg:
+ """Create a message with tool_result content block."""
+ return Msg(
+ name="tool",
+ role="user",
+ content=[
+ {
+ "type": "tool_result",
+ "id": tool_id,
+ "name": tool_name,
+ "output": output,
+ },
+ ],
+ )
+
+
+def create_mixed_tool_msg(
+ tool_use_id: str,
+ tool_use_name: str,
+ tool_use_input: dict,
+ tool_result_id: str,
+ tool_result_name: str,
+ tool_result_output: str,
+) -> Msg:
+ """Create a message with both tool_use and tool_result blocks."""
+ return Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ {
+ "type": "tool_use",
+ "id": tool_use_id,
+ "name": tool_use_name,
+ "input": tool_use_input,
+ },
+ {
+ "type": "tool_result",
+ "id": tool_result_id,
+ "name": tool_result_name,
+ "output": tool_result_output,
+ },
+ ],
+ )
+
+
+# =============================================================================
+# Normal Cases
+# =============================================================================
+
+
+def test_empty_messages():
+ """Test context_check with empty messages list."""
+ handler = create_handler()
+ messages = []
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ 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")
+
+
+def test_below_threshold_returns_all():
+ """Test that messages below threshold are all kept."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Hello"),
+ create_assistant_msg("Hi there!"),
+ create_user_msg("How are you?"),
+ ]
+ threshold, reserve = 10000, 5000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Very high threshold
+ memory_compact_reserve=reserve,
+ )
+ 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",
+ )
+ print_pass("test_below_threshold_returns_all")
+
+
+def test_above_threshold_triggers_compaction():
+ """Test that messages above threshold are split correctly."""
+ handler = create_handler()
+ # Create messages that will exceed threshold
+ messages = [
+ create_user_msg("First message " * 100),
+ create_assistant_msg("Second message " * 100),
+ create_user_msg("Third message " * 100),
+ create_assistant_msg("Fourth message " * 100),
+ ]
+ threshold, reserve = 100, 200
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Low threshold to trigger compaction
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_above_threshold_triggers_compaction")
+
+
+def test_message_order_preserved():
+ """Test that message order is preserved in both lists."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("First " * 50),
+ create_assistant_msg("Second " * 50),
+ create_user_msg("Third " * 50),
+ create_assistant_msg("Fourth " * 50),
+ create_user_msg("Fifth " * 10),
+ ]
+ threshold, reserve = 100, 150
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Low threshold
+ memory_compact_reserve=reserve,
+ )
+ # Check order preservation - compact messages should appear first in original
+ 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",
+ )
+ print_pass("test_message_order_preserved")
+
+
+# =============================================================================
+# Edge Cases - Threshold and Reserve Boundaries
+# =============================================================================
+
+
+def test_single_message_below_threshold():
+ """Test single message below threshold."""
+ handler = create_handler()
+ messages = [create_user_msg("Short message")]
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ 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",
+ )
+ print_pass("test_single_message_below_threshold")
+
+
+def test_single_message_above_threshold():
+ """Test single message that exceeds threshold - nothing can be kept in reserve."""
+ handler = create_handler()
+ long_content = "Very long message " * 1000
+ messages = [create_user_msg(long_content)]
+ threshold, reserve = 10, 5
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Very low threshold
+ memory_compact_reserve=reserve, # Even lower reserve
+ )
+ # 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",
+ )
+ print_pass("test_single_message_above_threshold")
+
+
+def test_reserve_zero():
+ """Test with reserve=0, no messages can be kept."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Hello"),
+ create_assistant_msg("Hi there!"),
+ ]
+ threshold, reserve = 1, 0
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Trigger compaction
+ memory_compact_reserve=reserve, # Zero reserve
+ )
+ # All messages should be compacted since reserve is 0
+ assert len(to_compact) == 2, f"All messages should be compacted, got {len(to_compact)}"
+ assert len(to_keep) == 0, f"No messages should be kept, got {len(to_keep)}"
+ verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_reserve_zero")
+ print_pass("test_reserve_zero")
+
+
+def test_threshold_zero():
+ """Test with threshold=0, always triggers compaction."""
+ handler = create_handler()
+ messages = [create_user_msg("A")] # Minimal message
+ threshold, reserve = 0, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Zero threshold - always triggers
+ memory_compact_reserve=reserve,
+ )
+ # Even minimal message triggers compaction with threshold=0
+ # But reserve is high so it should be kept
+ assert len(to_compact) == 0 or len(to_keep) == 1, "Message should fit in reserve"
+ verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_threshold_zero")
+ print_pass("test_threshold_zero")
+
+
+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,
+ memory_compact_threshold=threshold, # Exactly at boundary
+ memory_compact_reserve=reserve,
+ )
+ # At exact boundary (<=), should not trigger compaction
+ 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",
+ )
+ print_pass("test_exact_threshold_boundary")
+
+
+def test_reserve_larger_than_threshold():
+ """Test when reserve is larger than threshold (unusual but valid config)."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Message one " * 20),
+ create_assistant_msg("Message two " * 20),
+ ]
+ threshold, reserve = 50, 10000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Low threshold
+ memory_compact_reserve=reserve, # High reserve
+ )
+ # 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",
+ )
+ print_pass("test_reserve_larger_than_threshold")
+
+
+# =============================================================================
+# Edge Cases - Tool Use/Result Pairing
+# =============================================================================
+
+
+def test_tool_use_result_paired():
+ """Test that tool_use and tool_result pairs are kept together."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Please run the tool " * 50),
+ create_tool_use_msg("call_001", "test_tool", {"arg": "value"}),
+ create_tool_result_msg("call_001", "test_tool", "Tool output"),
+ create_assistant_msg("The tool returned results"),
+ ]
+ threshold, reserve = 50, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ 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)
+
+ 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",
+ )
+ print_pass("test_tool_use_result_paired")
+
+
+def test_tool_use_without_result():
+ """Test tool_use message without corresponding tool_result."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Run the tool"),
+ create_tool_use_msg("call_orphan", "orphan_tool", {"arg": "value"}),
+ create_assistant_msg("Something happened"),
+ ]
+ threshold, reserve = 10, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_tool_use_without_result")
+
+
+def test_tool_result_without_use():
+ """Test tool_result message without corresponding tool_use."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Here's a result"),
+ create_tool_result_msg("call_orphan", "orphan_tool", "Some output"),
+ create_assistant_msg("Got it"),
+ ]
+ threshold, reserve = 10, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_tool_result_without_use")
+
+
+def test_multiple_tool_pairs():
+ """Test multiple tool_use/tool_result pairs."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Start task " * 50),
+ create_tool_use_msg("call_001", "tool_a", {"a": 1}),
+ create_tool_result_msg("call_001", "tool_a", "Result A"),
+ create_tool_use_msg("call_002", "tool_b", {"b": 2}),
+ create_tool_result_msg("call_002", "tool_b", "Result B"),
+ create_tool_use_msg("call_003", "tool_c", {"c": 3}),
+ create_tool_result_msg("call_003", "tool_c", "Result C"),
+ create_assistant_msg("All done"),
+ ]
+ threshold, reserve = 50, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ 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"):
+ tool_id = block.get("id", "")
+ if tool_id:
+ # Find corresponding tool_use
+ tool_use_found = False
+ for keep_msg in to_keep:
+ for use_block in keep_msg.get_content_blocks("tool_use"):
+ if use_block.get("id") == tool_id:
+ 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",
+ )
+ print_pass("test_multiple_tool_pairs")
+
+
+def test_tool_dependency_causes_extra_inclusion():
+ """Test that tool_use is included even if it exceeds simple reserve calculation."""
+ handler = create_handler()
+ # Create a scenario where:
+ # - First message (tool_use) is large
+ # - Later message (tool_result) references it
+ # - Reserve alone wouldn't fit tool_use, but dependency requires it
+ large_tool_input = {"data": "x" * 200}
+ 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_tool_result_msg("call_dep", "dep_tool", "Result"), # Small
+ create_assistant_msg("End"), # Small
+ ]
+ threshold, reserve = 100, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ 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())
+ 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
+ )
+
+ 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",
+ )
+ print_pass("test_tool_dependency_causes_extra_inclusion")
+
+
+def test_tool_dependency_exceeds_reserve():
+ """Test when tool_result + its tool_use dependency would exceed reserve."""
+ handler = create_handler()
+ # tool_use is very large, making the pair not fit in reserve
+ very_large_input = {"data": "x" * 2000}
+ messages = [
+ create_user_msg("First"),
+ create_tool_use_msg("call_big", "big_tool", very_large_input), # Very large
+ create_tool_result_msg("call_big", "big_tool", "Small result"),
+ create_assistant_msg("Last message"),
+ ]
+ threshold, reserve = 10, 100
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ 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())
+ 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())
+ 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",
+ )
+ print_pass("test_tool_dependency_exceeds_reserve")
+
+
+def test_interleaved_tool_pairs():
+ """Test interleaved tool_use/tool_result (not strictly sequential)."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Multi-tool task " * 30),
+ create_tool_use_msg("call_a", "tool_a", {"a": 1}),
+ create_tool_use_msg("call_b", "tool_b", {"b": 2}), # Two uses before results
+ create_tool_result_msg("call_a", "tool_a", "Result A"),
+ create_tool_result_msg("call_b", "tool_b", "Result B"),
+ create_assistant_msg("Both done"),
+ ]
+ threshold, reserve = 50, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ 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())
+ 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",
+ )
+ print_pass("test_interleaved_tool_pairs")
+
+
+# =============================================================================
+# Edge Cases - Message Content Variations
+# =============================================================================
+
+
+def test_message_with_empty_content():
+ """Test message with empty string content."""
+ handler = create_handler()
+ messages = [
+ create_user_msg(""), # Empty content
+ create_assistant_msg("Response"),
+ ]
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ 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",
+ )
+ print_pass("test_message_with_empty_content")
+
+
+def test_message_with_whitespace_only():
+ """Test message with whitespace-only content."""
+ handler = create_handler()
+ messages = [
+ create_user_msg(" \n\t "), # Whitespace only
+ create_assistant_msg("Response"),
+ ]
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ 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",
+ )
+ print_pass("test_message_with_whitespace_only")
+
+
+def test_very_long_single_message():
+ """Test very long single message that exceeds any reasonable reserve."""
+ handler = create_handler()
+ huge_content = "x" * 100000 # Very long
+ messages = [create_user_msg(huge_content)]
+ threshold, reserve = 100, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_very_long_single_message")
+
+
+def test_many_small_messages():
+ """Test many small messages."""
+ handler = create_handler()
+ messages = [create_user_msg(f"Msg {i}") for i in range(100)]
+ threshold, reserve = 100, 200
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Low threshold
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_many_small_messages")
+
+
+def test_unicode_content():
+ """Test messages with unicode characters."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("你好世界!🎉 Emoji and 中文"),
+ create_assistant_msg("مرحبا العالم 🌍 Arabic and more"),
+ create_user_msg("日本語テスト 🇯🇵"),
+ ]
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ assert len(to_compact) + len(to_keep) == 3
+ verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_unicode_content")
+ print_pass("test_unicode_content")
+
+
+def test_special_characters_content():
+ """Test messages with special characters."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Special chars: <>&\"'`~!@#$%^&*()[]{}|\\"),
+ create_assistant_msg("More: \n\r\t\0 nulls and newlines"),
+ ]
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ 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",
+ )
+ print_pass("test_special_characters_content")
+
+
+# =============================================================================
+# Edge Cases - Boundary Conditions
+# =============================================================================
+
+
+def test_all_messages_fit_exactly_in_reserve():
+ """Test when all messages fit exactly in reserve after threshold exceeded."""
+ handler = create_handler()
+ messages = [
+ 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
+ memory_compact_reserve=reserve, # Exactly fits all
+ )
+ # 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",
+ )
+ print_pass("test_all_messages_fit_exactly_in_reserve")
+
+
+def test_first_message_only_compacted():
+ """Test when only the first message is compacted."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Large first message " * 100), # Large
+ 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",
+ )
+ print_pass("test_first_message_only_compacted")
+
+
+def test_last_message_only_kept():
+ """Test when only the last message can be kept."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Large " * 200),
+ 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",
+ )
+ print_pass("test_last_message_only_kept")
+
+
+def test_all_messages_compacted():
+ """Test when all messages need to be compacted (nothing fits in reserve)."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Large message " * 100),
+ create_assistant_msg("Large message " * 100),
+ ]
+ threshold, reserve = 10, 1
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold, # Trigger compaction
+ memory_compact_reserve=reserve, # Too small for anything
+ )
+ 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",
+ )
+ print_pass("test_all_messages_compacted")
+
+
+# =============================================================================
+# Edge Cases - Message Roles
+# =============================================================================
+
+
+def test_system_message():
+ """Test handling of system role messages."""
+ handler = create_handler()
+ system_msg = Msg(name="system", role="system", content="You are a helpful assistant.")
+ messages = [
+ system_msg,
+ create_user_msg("Hello"),
+ create_assistant_msg("Hi there!"),
+ ]
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ assert len(to_compact) + len(to_keep) == 3
+ verify_context_check_invariants(handler, messages, to_compact, to_keep, threshold, reserve, "test_system_message")
+ print_pass("test_system_message")
+
+
+def test_mixed_roles():
+ """Test messages with various roles (user, assistant, system)."""
+ handler = create_handler()
+ # agentscope.message.Msg only supports: user, assistant, system
+ messages = [
+ Msg(name="system", role="system", content="System prompt"),
+ Msg(name="user", role="user", content="User message"),
+ Msg(name="assistant", role="assistant", content="Assistant response"),
+ Msg(name="tool", role="user", content="Tool output as user role"),
+ Msg(name="helper", role="assistant", content="Another assistant message"),
+ ]
+ threshold, reserve = 1000, 500
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ 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_mixed_roles")
+ print_pass("test_mixed_roles")
+
+
+# =============================================================================
+# Edge Cases - Tool Block Variations
+# =============================================================================
+
+
+def test_tool_use_with_empty_id():
+ """Test tool_use block with empty id."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Run tool"),
+ create_tool_use_msg("", "test_tool", {"arg": "value"}), # Empty ID
+ create_assistant_msg("Done"),
+ ]
+ threshold, reserve = 10, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_tool_use_with_empty_id")
+
+
+def test_tool_result_with_empty_id():
+ """Test tool_result block with empty id."""
+ handler = create_handler()
+ messages = [
+ create_user_msg("Got result"),
+ create_tool_result_msg("", "test_tool", "Output"), # Empty ID
+ create_assistant_msg("Noted"),
+ ]
+ threshold, reserve = 10, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_tool_result_with_empty_id")
+
+
+def test_duplicate_tool_ids():
+ """Test messages with duplicate tool IDs (unusual but possible)."""
+ handler = create_handler()
+ messages = [
+ create_tool_use_msg("call_dup", "tool_a", {"a": 1}),
+ create_tool_result_msg("call_dup", "tool_a", "Result A"),
+ create_tool_use_msg("call_dup", "tool_b", {"b": 2}), # Same ID, different tool
+ create_tool_result_msg("call_dup", "tool_b", "Result B"),
+ ]
+ threshold, reserve = 10, 1000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ memory_compact_reserve=reserve,
+ )
+ # 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",
+ )
+ print_pass("test_duplicate_tool_ids")
+
+
+def test_message_with_multiple_tool_blocks():
+ """Test single message containing multiple tool blocks."""
+ handler = create_handler()
+ msg_with_multiple_tools = Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ {"type": "tool_use", "id": "call_1", "name": "tool1", "input": {}},
+ {"type": "tool_use", "id": "call_2", "name": "tool2", "input": {}},
+ {"type": "tool_use", "id": "call_3", "name": "tool3", "input": {}},
+ ],
+ )
+ messages = [
+ create_user_msg("Do multiple things"),
+ msg_with_multiple_tools,
+ create_tool_result_msg("call_1", "tool1", "Result 1"),
+ create_tool_result_msg("call_2", "tool2", "Result 2"),
+ create_tool_result_msg("call_3", "tool3", "Result 3"),
+ ]
+ threshold, reserve = 10, 2000
+ to_compact, to_keep, _ = handler.context_check(
+ messages=messages,
+ memory_compact_threshold=threshold,
+ 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",
+ )
+ print_pass("test_message_with_multiple_tool_blocks")
+
+
+# =============================================================================
+# Run All Tests
+# =============================================================================
+
+
+def run_all_tests():
+ """Run all tests."""
+ tests = [
+ # Normal cases
+ test_empty_messages,
+ test_below_threshold_returns_all,
+ test_above_threshold_triggers_compaction,
+ test_message_order_preserved,
+ # Edge cases - boundaries
+ test_single_message_below_threshold,
+ test_single_message_above_threshold,
+ test_reserve_zero,
+ test_threshold_zero,
+ test_exact_threshold_boundary,
+ test_reserve_larger_than_threshold,
+ # Edge cases - tool pairing
+ test_tool_use_result_paired,
+ test_tool_use_without_result,
+ test_tool_result_without_use,
+ test_multiple_tool_pairs,
+ test_tool_dependency_causes_extra_inclusion,
+ test_tool_dependency_exceeds_reserve,
+ test_interleaved_tool_pairs,
+ # Edge cases - content variations
+ test_message_with_empty_content,
+ test_message_with_whitespace_only,
+ test_very_long_single_message,
+ test_many_small_messages,
+ test_unicode_content,
+ test_special_characters_content,
+ # Edge cases - boundaries
+ test_all_messages_fit_exactly_in_reserve,
+ test_first_message_only_compacted,
+ test_last_message_only_kept,
+ test_all_messages_compacted,
+ # Edge cases - roles
+ test_system_message,
+ test_mixed_roles,
+ # Edge cases - tool blocks
+ test_tool_use_with_empty_id,
+ test_tool_result_with_empty_id,
+ test_duplicate_tool_ids,
+ test_message_with_multiple_tool_blocks,
+ ]
+
+ passed = 0
+ failed = 0
+
+ for test in tests:
+ try:
+ print_test_header(test.__name__)
+ test()
+ passed += 1
+ except AssertionError as e:
+ print_fail(test.__name__, str(e))
+ failed += 1
+ except Exception as e:
+ print_error(test.__name__, str(e))
+ failed += 1
+
+ # Print summary
+ print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+ print(f"{Colors.BOLD}Test Results Summary{Colors.RESET}")
+ print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+ print(f"{Colors.GREEN}{Colors.BOLD}✓ Passed: {passed}{Colors.RESET}")
+ if failed > 0:
+ print(f"{Colors.RED}{Colors.BOLD}✗ Failed: {failed}{Colors.RESET}")
+ else:
+ print(f"{Colors.GREEN}✗ Failed: {failed}{Colors.RESET}")
+ print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+
+ if failed == 0:
+ print(f"\n{Colors.GREEN}{Colors.BOLD}🎉 All tests passed!{Colors.RESET}")
+ else:
+ print(f"\n{Colors.RED}{Colors.BOLD}💥 Some tests failed!{Colors.RESET}")
+
+
+if __name__ == "__main__":
+ run_all_tests()
diff --git a/tests/light/test_format_msgs_to_str.py b/tests/light/test_format_msgs_to_str.py
new file mode 100644
index 00000000..93dd7a2f
--- /dev/null
+++ b/tests/light/test_format_msgs_to_str.py
@@ -0,0 +1,891 @@
+"""Tests for AsMsgHandler.format_msgs_to_str method."""
+
+# pylint: disable=W0212
+
+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
+
+logger = get_std_logger()
+
+
+# ANSI 颜色码
+class Colors:
+ """ANSI color codes for terminal output."""
+
+ GREEN = "\033[92m"
+ RED = "\033[91m"
+ YELLOW = "\033[93m"
+ BLUE = "\033[94m"
+ CYAN = "\033[96m"
+ BOLD = "\033[1m"
+ RESET = "\033[0m"
+
+
+def print_pass(test_name: str):
+ """打印测试通过信息"""
+ print(f"{Colors.GREEN}{Colors.BOLD}✓ {test_name} PASSED{Colors.RESET}")
+
+
+def print_fail(test_name: str, error: str):
+ """打印测试失败信息"""
+ print(f"{Colors.RED}{Colors.BOLD}✗ {test_name} FAILED: {error}{Colors.RESET}")
+
+
+def print_error(test_name: str, error: str):
+ """打印测试错误信息"""
+ print(f"{Colors.YELLOW}{Colors.BOLD}⚠ {test_name} ERROR: {error}{Colors.RESET}")
+
+
+def print_test_header(test_name: str):
+ """打印测试标题"""
+ print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+ print(f"{Colors.BLUE}{Colors.BOLD}Running: {test_name}{Colors.RESET}")
+ print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+
+
+# ==================== Helper Functions ====================
+
+
+def create_handler() -> AsMsgHandler:
+ """Create an AsMsgHandler instance for testing."""
+ return AsMsgHandler(token_counter=get_token_counter())
+
+
+def verify_result_within_threshold(
+ handler: AsMsgHandler,
+ result: str,
+ threshold: int,
+ test_name: str = "",
+ msgs: list[Msg] | None = None,
+) -> None:
+ """Verify that the included messages' original token count does not exceed threshold.
+
+ Note: The format_msgs_to_str method uses message token statistics (not formatted
+ string tokens) for threshold checking. The formatted result may have more tokens
+ than the threshold due to added metadata (timestamps, role prefixes, etc.).
+
+ This verification checks that included messages' original token sum <= threshold.
+
+ Args:
+ handler: The AsMsgHandler instance used for token counting.
+ result: The formatted string result from format_msgs_to_str.
+ threshold: The memory_compact_threshold value used.
+ test_name: Optional test name for better error messages.
+ msgs: Optional list of original messages to verify against.
+
+ Raises:
+ AssertionError: If included messages' token count exceeds threshold.
+ """
+ if not result or not msgs:
+ return # Empty result or no messages to verify
+
+ # Calculate tokens of messages that were included in the result
+ included_tokens = 0
+ for msg in msgs:
+ stat = handler.stat_message(msg)
+ # Check if this message's content appears in the result
+ _ = stat.format(include_thinking=True) # Use True to check all content
+ # Simple heuristic: if the message content is in result, count its tokens
+ content_blocks = msg.get_content_blocks()
+ msg_included = False
+ for block in content_blocks:
+ block_type = block.get("type", "")
+ if block_type == "text" and block.get("text", "") in result:
+ msg_included = True
+ break
+ if block_type == "tool_use" and f"tool_call={block.get('name', '')}" in result:
+ msg_included = True
+ break
+ 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})."
+
+
+def create_user_msg(content: str) -> Msg:
+ """Create a user message."""
+ return Msg(name="user", role="user", content=content)
+
+
+def create_assistant_msg(content: str) -> Msg:
+ """Create an assistant message."""
+ return Msg(name="assistant", role="assistant", content=content)
+
+
+def create_tool_use_msg(tool_name: str, tool_input: dict, tool_id: str = "call_123") -> Msg:
+ """Create a message with tool_use content block."""
+ return Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ {
+ "type": "tool_use",
+ "id": tool_id,
+ "name": tool_name,
+ "input": tool_input,
+ },
+ ],
+ )
+
+
+def create_tool_result_msg(tool_name: str, output: str | list[dict], tool_id: str = "call_123") -> Msg:
+ """Create a message with tool_result content block."""
+ return Msg(
+ name="tool",
+ role="user",
+ content=[
+ {
+ "type": "tool_result",
+ "id": tool_id,
+ "name": tool_name,
+ "output": output,
+ },
+ ],
+ )
+
+
+def create_thinking_msg(thinking_content: str, text_content: str = "") -> Msg:
+ """Create a message with thinking content block."""
+ content = [
+ {
+ "type": "thinking",
+ "thinking": thinking_content,
+ },
+ ]
+ if text_content:
+ content.append({"type": "text", "text": text_content})
+ return Msg(name="assistant", role="assistant", content=content)
+
+
+def create_image_msg(url: str = "") -> Msg:
+ """Create a message with image content block."""
+ content = [
+ {
+ "type": "image",
+ "source": {"url": url} if url else {},
+ },
+ ]
+ return Msg(name="assistant", role="assistant", content=content)
+
+
+def create_mixed_content_msg(
+ text: str = "",
+ thinking: str = "",
+ tool_name: str = "",
+ tool_input: dict | None = None,
+ image_url: str = "",
+) -> Msg:
+ """Create a message with mixed content blocks."""
+ content = []
+ if thinking:
+ content.append({"type": "thinking", "thinking": thinking})
+ if text:
+ content.append({"type": "text", "text": text})
+ if tool_name:
+ content.append(
+ {
+ "type": "tool_use",
+ "id": "call_mixed",
+ "name": tool_name,
+ "input": tool_input or {},
+ },
+ )
+ if image_url:
+ content.append({"type": "image", "source": {"url": image_url}})
+ return Msg(name="assistant", role="assistant", content=content)
+
+
+# ==================== Normal Case Tests ====================
+
+
+def test_format_msgs_to_str_empty_list():
+ """Test format_msgs_to_str with empty message list."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = []
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+ assert result == "", f"Expected empty string for empty list, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "empty_list", msgs)
+ print_pass("test_format_msgs_to_str_empty_list")
+
+
+def test_format_msgs_to_str_single_message():
+ """Test format_msgs_to_str with a single message."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_user_msg("Hello, how are you?")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "user:" in result, f"Expected 'user:' in result, got: {result}"
+ assert "Hello, how are you?" in result, f"Expected content in result, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "single_message", msgs)
+ print_pass("test_format_msgs_to_str_single_message")
+
+
+def test_format_msgs_to_str_multiple_messages():
+ """Test format_msgs_to_str with multiple messages."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ create_user_msg("What is Python?"),
+ create_assistant_msg("Python is a programming language."),
+ create_user_msg("Tell me more."),
+ create_assistant_msg("Python is known for its readability."),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "What is Python?" in result
+ assert "Python is a programming language." in result
+ assert "Tell me more." in result
+ assert "Python is known for its readability." in result
+ verify_result_within_threshold(handler, result, threshold, "multiple_messages", msgs)
+ print_pass("test_format_msgs_to_str_multiple_messages")
+
+
+def test_format_msgs_to_str_message_order():
+ """Test that messages are returned in correct order (oldest to newest)."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ create_user_msg("First message"),
+ create_assistant_msg("Second message"),
+ create_user_msg("Third message"),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # Find positions of each message
+ first_pos = result.find("First message")
+ second_pos = result.find("Second message")
+ third_pos = result.find("Third message")
+
+ assert first_pos < second_pos < third_pos, (
+ f"Messages not in correct order. Positions: first={first_pos}, " f"second={second_pos}, third={third_pos}"
+ )
+ verify_result_within_threshold(handler, result, threshold, "message_order", msgs)
+ print_pass("test_format_msgs_to_str_message_order")
+
+
+def test_format_msgs_to_str_with_tool_use():
+ """Test format_msgs_to_str with tool_use message."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_tool_use_msg("read_file", {"path": "/test.txt"})]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "tool_call=read_file" in result, f"Expected tool_call in result, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "with_tool_use", msgs)
+ print_pass("test_format_msgs_to_str_with_tool_use")
+
+
+def test_format_msgs_to_str_with_tool_result():
+ """Test format_msgs_to_str with tool_result message."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_tool_result_msg("read_file", "file content here")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "tool_result=read_file" in result, f"Expected tool_result in result, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "with_tool_result", msgs)
+ print_pass("test_format_msgs_to_str_with_tool_result")
+
+
+def test_format_msgs_to_str_with_image():
+ """Test format_msgs_to_str with image message."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_image_msg("https://example.com/image.png")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "[image]" in result, f"Expected '[image]' in result, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "with_image", msgs)
+ print_pass("test_format_msgs_to_str_with_image")
+
+
+def test_format_msgs_to_str_conversation_flow():
+ """Test format_msgs_to_str with a complete conversation flow."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ create_user_msg("Read the file."),
+ create_tool_use_msg("read_file", {"path": "/data.txt"}),
+ create_tool_result_msg("read_file", "File content here"),
+ create_assistant_msg("The file contains: File content here"),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "user:" in result
+ assert "tool_call=read_file" in result
+ assert "tool_result=read_file" in result
+ assert "assistant:" in result
+ verify_result_within_threshold(handler, result, threshold, "conversation_flow", msgs)
+ print_pass("test_format_msgs_to_str_conversation_flow")
+
+
+# ==================== Thinking Block Tests ====================
+
+
+def test_format_msgs_to_str_thinking_excluded_by_default():
+ """Test that thinking blocks are excluded when include_thinking=False (default)."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_thinking_msg("Let me think about this...", "Here is my response")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold, include_thinking=False)
+
+ assert "Let me think about this" not in result, f"Thinking content should be excluded, got: {result}"
+ assert "Here is my response" in result, f"Text content should be included, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "thinking_excluded_by_default", msgs)
+ print_pass("test_format_msgs_to_str_thinking_excluded_by_default")
+
+
+def test_format_msgs_to_str_thinking_included():
+ """Test that thinking blocks are included when include_thinking=True."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_thinking_msg("Let me think about this...", "Here is my response")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold, include_thinking=True)
+
+ assert "Let me think about this" in result, f"Thinking content should be included, got: {result}"
+ assert "" in result, f"Expected thinking tag in result, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "thinking_included", msgs)
+ print_pass("test_format_msgs_to_str_thinking_included")
+
+
+def test_format_msgs_to_str_thinking_only_message():
+ """Test message with only thinking block."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_thinking_msg("Deep thoughts here")]
+
+ # With include_thinking=False
+ result_no_thinking = handler.format_msgs_to_str(
+ msgs,
+ memory_compact_threshold=threshold,
+ include_thinking=False,
+ )
+ # With include_thinking=True
+ result_with_thinking = handler.format_msgs_to_str(
+ msgs,
+ memory_compact_threshold=threshold,
+ include_thinking=True,
+ )
+
+ assert "Deep thoughts here" not in result_no_thinking
+ assert "Deep thoughts here" in result_with_thinking
+ verify_result_within_threshold(handler, result_no_thinking, threshold, "thinking_only_no_thinking", msgs)
+ verify_result_within_threshold(handler, result_with_thinking, threshold, "thinking_only_with_thinking", msgs)
+ print_pass("test_format_msgs_to_str_thinking_only_message")
+
+
+# ==================== Token Threshold Tests ====================
+
+
+def test_format_msgs_to_str_all_within_threshold():
+ """Test all messages fit within threshold."""
+ handler = create_handler()
+ threshold = 10000
+ msgs = [
+ create_user_msg("Short message 1"),
+ create_assistant_msg("Short message 2"),
+ create_user_msg("Short message 3"),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "Short message 1" in result
+ assert "Short message 2" in result
+ assert "Short message 3" in result
+ verify_result_within_threshold(handler, result, threshold, "all_within_threshold", msgs)
+ print_pass("test_format_msgs_to_str_all_within_threshold")
+
+
+def test_format_msgs_to_str_exceeds_threshold_truncate_older():
+ """Test that older messages are truncated when exceeding threshold."""
+ handler = create_handler()
+ threshold = 500
+ msgs = []
+ for i in range(20):
+ msgs.append(create_user_msg(f"Question {i}: " + "x" * 100))
+ msgs.append(create_assistant_msg(f"Answer {i}: " + "y" * 100))
+
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # The newest messages should be present
+ assert (
+ "Answer 19" in result or "Question 19" in result
+ ), f"Expected recent message in result, got: {result[:500]}..."
+ # Older messages should be truncated
+ assert "Question 0" not in result, "Older messages should be truncated"
+ verify_result_within_threshold(handler, result, threshold, "exceeds_threshold_truncate_older", msgs)
+ print_pass("test_format_msgs_to_str_exceeds_threshold_truncate_older")
+
+
+def test_format_msgs_to_str_single_message_exceeds_threshold():
+ """Test when a single message exceeds the threshold."""
+ handler = create_handler()
+ threshold = 10
+ # Create a very long message
+ long_text = "x" * 10000
+ msgs = [create_user_msg(long_text)]
+
+ # With very low threshold, even a single message won't fit
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # The message should be skipped entirely since it exceeds threshold
+ assert result == "" or len(result) > 0, "Result should be empty or contain truncated content"
+ verify_result_within_threshold(handler, result, threshold, "single_message_exceeds_threshold", msgs)
+ print_pass("test_format_msgs_to_str_single_message_exceeds_threshold")
+
+
+def test_format_msgs_to_str_first_message_exceeds_threshold():
+ """Test when the first (oldest) message exceeds threshold but newer ones don't."""
+ handler = create_handler()
+ threshold = 100
+ msgs = [
+ create_user_msg("x" * 5000), # Old, long message
+ create_assistant_msg("Short response"), # New, short message
+ ]
+
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # Newer message should be present
+ assert "Short response" in result, f"Expected newer message in result, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "first_message_exceeds_threshold", msgs)
+ print_pass("test_format_msgs_to_str_first_message_exceeds_threshold")
+
+
+def test_format_msgs_to_str_threshold_zero():
+ """Test with threshold of zero - no messages should be included."""
+ handler = create_handler()
+ threshold = 0
+ msgs = [create_user_msg("Test message")]
+
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert result == "", f"Expected empty string with zero threshold, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "threshold_zero", msgs)
+ print_pass("test_format_msgs_to_str_threshold_zero")
+
+
+def test_format_msgs_to_str_threshold_exact_fit():
+ """Test when messages exactly fit the threshold."""
+ handler = create_handler()
+ # Create a message and measure its formatted string tokens
+ msg = create_user_msg("Test")
+ stat = handler.stat_message(msg)
+ formatted_content = stat.format(include_thinking=False)
+ exact_threshold = handler.count_str_token(formatted_content)
+
+ msgs = [msg]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=exact_threshold)
+
+ assert "Test" in result, f"Message should fit exactly, got: {result}"
+ verify_result_within_threshold(handler, result, exact_threshold, "threshold_exact_fit", msgs)
+ print_pass("test_format_msgs_to_str_threshold_exact_fit")
+
+
+def test_format_msgs_to_str_threshold_one_less():
+ """Test when threshold is one less than needed."""
+ handler = create_handler()
+ msg = create_user_msg("Test message")
+ stat = handler.stat_message(msg)
+ formatted_content = stat.format(include_thinking=False)
+ threshold_minus_one = handler.count_str_token(formatted_content) - 1
+
+ msgs = [msg]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold_minus_one)
+
+ # Message should be skipped since it doesn't fit
+ assert result == "", f"Expected empty string when threshold is insufficient, got: {result}"
+ verify_result_within_threshold(handler, result, threshold_minus_one, "threshold_one_less", msgs)
+ print_pass("test_format_msgs_to_str_threshold_one_less")
+
+
+def test_format_msgs_to_str_large_threshold():
+ """Test with very large threshold - all messages should be included."""
+ handler = create_handler()
+ threshold = 1000000
+ msgs = [create_user_msg("Message " + str(i) + " " + "x" * 100) for i in range(50)]
+
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # All messages should be included
+ for i in range(50):
+ assert f"Message {i}" in result, f"Message {i} should be included"
+ verify_result_within_threshold(handler, result, threshold, "large_threshold", msgs)
+ print_pass("test_format_msgs_to_str_large_threshold")
+
+
+# ==================== Edge Cases Tests ====================
+
+
+def test_format_msgs_to_str_special_characters():
+ """Test with special characters in content."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_user_msg("Test with 中文, 日本語, émojis 🎉 and symbols @#$%")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "中文" in result
+ assert "日本語" in result
+ assert "🎉" in result
+ verify_result_within_threshold(handler, result, threshold, "special_characters", msgs)
+ print_pass("test_format_msgs_to_str_special_characters")
+
+
+def test_format_msgs_to_str_empty_content():
+ """Test with empty content message."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_user_msg("")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "user:" in result, f"Expected role in result even with empty content, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "empty_content", msgs)
+ print_pass("test_format_msgs_to_str_empty_content")
+
+
+def test_format_msgs_to_str_whitespace_only():
+ """Test with whitespace-only content."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_user_msg(" \n\t ")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "user:" in result
+ verify_result_within_threshold(handler, result, threshold, "whitespace_only", msgs)
+ print_pass("test_format_msgs_to_str_whitespace_only")
+
+
+def test_format_msgs_to_str_newlines_in_content():
+ """Test with newlines in message content."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [create_user_msg("Line 1\nLine 2\nLine 3")]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "Line 1" in result
+ assert "Line 2" in result
+ assert "Line 3" in result
+ verify_result_within_threshold(handler, result, threshold, "newlines_in_content", msgs)
+ print_pass("test_format_msgs_to_str_newlines_in_content")
+
+
+def test_format_msgs_to_str_very_long_single_word():
+ """Test with very long single word (no spaces)."""
+ handler = create_handler()
+ threshold = 10000
+ long_word = "a" * 5000
+ msgs = [create_user_msg(long_word)]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # Should contain at least part of the word (may be truncated by formatter)
+ assert "aaa" in result, f"Expected long word content in result, got: {result[:100]}..."
+ verify_result_within_threshold(handler, result, threshold, "very_long_single_word", msgs)
+ print_pass("test_format_msgs_to_str_very_long_single_word")
+
+
+def test_format_msgs_to_str_mixed_content_blocks():
+ """Test message with mixed content blocks."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ create_mixed_content_msg(
+ text="Text content",
+ thinking="Thinking content",
+ tool_name="test_tool",
+ tool_input={"key": "value"},
+ image_url="https://example.com/img.png",
+ ),
+ ]
+
+ result_no_thinking = handler.format_msgs_to_str(
+ msgs,
+ memory_compact_threshold=threshold,
+ include_thinking=False,
+ )
+ result_with_thinking = handler.format_msgs_to_str(
+ msgs,
+ memory_compact_threshold=threshold,
+ include_thinking=True,
+ )
+
+ assert "Text content" in result_no_thinking
+ assert "tool_call=test_tool" in result_no_thinking
+ assert "[image]" in result_no_thinking
+ assert "Thinking content" not in result_no_thinking
+ assert "Thinking content" in result_with_thinking
+ verify_result_within_threshold(handler, result_no_thinking, threshold, "mixed_content_no_thinking", msgs)
+ verify_result_within_threshold(handler, result_with_thinking, threshold, "mixed_content_with_thinking", msgs)
+ print_pass("test_format_msgs_to_str_mixed_content_blocks")
+
+
+def test_format_msgs_to_str_multiple_separators():
+ """Test that messages are separated by double newlines."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ create_user_msg("Message 1"),
+ create_assistant_msg("Message 2"),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "\n\n" in result, f"Expected double newline separator, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "multiple_separators", msgs)
+ print_pass("test_format_msgs_to_str_multiple_separators")
+
+
+def test_format_msgs_to_str_tool_result_complex_output():
+ """Test tool_result with complex output (list of blocks)."""
+ handler = create_handler()
+ threshold = 4000
+ complex_output = [
+ {"type": "text", "text": "Operation completed"},
+ {"type": "image", "source": {"url": "https://example.com/result.png"}},
+ ]
+ msgs = [create_tool_result_msg("process_data", complex_output)]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "tool_result=process_data" in result
+ verify_result_within_threshold(handler, result, threshold, "tool_result_complex_output", msgs)
+ print_pass("test_format_msgs_to_str_tool_result_complex_output")
+
+
+def test_format_msgs_to_str_different_roles():
+ """Test with different roles (user, assistant, system, tool)."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ Msg(name="system", role="system", content="System instruction"),
+ create_user_msg("User message"),
+ create_assistant_msg("Assistant response"),
+ create_tool_result_msg("tool", "Tool output"),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "system:" in result
+ assert "user:" in result
+ assert "assistant:" in result
+ verify_result_within_threshold(handler, result, threshold, "different_roles", msgs)
+ print_pass("test_format_msgs_to_str_different_roles")
+
+
+def test_format_msgs_to_str_incremental_threshold_check():
+ """Test incremental addition of messages until threshold is exceeded."""
+ handler = create_handler()
+
+ # Create messages with known approximate sizes
+ msgs = []
+ for i in range(10):
+ msgs.append(create_user_msg(f"Message {i} with some padding text"))
+
+ # Calculate total tokens
+ total_tokens = sum(handler.stat_message(msg).total_tokens for msg in msgs)
+
+ # Use threshold that allows about half the messages
+ half_threshold = total_tokens // 2
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=half_threshold)
+
+ # Should have some but not all messages
+ included_count = sum(1 for i in range(10) if f"Message {i}" in result)
+ assert 0 < included_count < 10, f"Expected partial messages, got {included_count} messages included"
+ # Newer messages should be included (messages are processed from end)
+ assert "Message 9" in result, "Newest message should be included"
+ verify_result_within_threshold(handler, result, half_threshold, "incremental_threshold_check", msgs)
+ print_pass("test_format_msgs_to_str_incremental_threshold_check")
+
+
+def test_format_msgs_to_str_negative_threshold():
+ """Test with negative threshold value."""
+ handler = create_handler()
+ threshold = -1
+ msgs = [create_user_msg("Test message")]
+
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # Negative threshold should result in empty string (nothing fits)
+ assert result == "", f"Expected empty string with negative threshold, got: {result}"
+ verify_result_within_threshold(handler, result, max(0, threshold), "negative_threshold", msgs)
+ print_pass("test_format_msgs_to_str_negative_threshold")
+
+
+def test_format_msgs_to_str_preserves_newest_first():
+ """Test that newest messages are preserved when threshold is exceeded."""
+ handler = create_handler()
+ threshold = 300
+ msgs = [
+ create_user_msg("OLD MESSAGE " + "x" * 200),
+ create_assistant_msg("MIDDLE MESSAGE " + "y" * 200),
+ create_user_msg("NEW MESSAGE " + "z" * 200),
+ ]
+
+ # Use threshold that only allows ~1-2 messages
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # Newest message should be present
+ assert "NEW MESSAGE" in result, f"Expected newest message, got: {result}"
+ verify_result_within_threshold(handler, result, threshold, "preserves_newest_first", msgs)
+ print_pass("test_format_msgs_to_str_preserves_newest_first")
+
+
+def test_format_msgs_to_str_base64_image():
+ """Test with base64 encoded image."""
+ handler = create_handler()
+ threshold = 10000
+ msgs = [
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ {
+ "type": "image",
+ "source": {
+ "type": "base64",
+ "data": "SGVsbG8gV29ybGQ=" * 100, # Simulated base64 data
+ },
+ },
+ ],
+ ),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "[image]" in result
+ verify_result_within_threshold(handler, result, threshold, "base64_image", msgs)
+ print_pass("test_format_msgs_to_str_base64_image")
+
+
+def test_format_msgs_to_str_audio_video_blocks():
+ """Test with audio and video content blocks."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ {"type": "audio", "source": {"url": "https://example.com/audio.mp3"}},
+ {"type": "video", "source": {"url": "https://example.com/video.mp4"}},
+ ],
+ ),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ assert "[audio]" in result
+ assert "[video]" in result
+ verify_result_within_threshold(handler, result, threshold, "audio_video_blocks", msgs)
+ print_pass("test_format_msgs_to_str_audio_video_blocks")
+
+
+def test_format_msgs_to_str_unknown_block_type():
+ """Test that unknown block types are skipped gracefully."""
+ handler = create_handler()
+ threshold = 4000
+ msgs = [
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ {"type": "unknown_type", "data": "some data"},
+ {"type": "text", "text": "Valid text"},
+ ],
+ ),
+ ]
+ result = handler.format_msgs_to_str(msgs, memory_compact_threshold=threshold)
+
+ # Should still include valid content
+ assert "Valid text" in result
+ verify_result_within_threshold(handler, result, threshold, "unknown_block_type", msgs)
+ print_pass("test_format_msgs_to_str_unknown_block_type")
+
+
+def run_all_tests():
+ """Run all tests."""
+ tests = [
+ # Normal case tests
+ test_format_msgs_to_str_empty_list,
+ test_format_msgs_to_str_single_message,
+ test_format_msgs_to_str_multiple_messages,
+ test_format_msgs_to_str_message_order,
+ test_format_msgs_to_str_with_tool_use,
+ test_format_msgs_to_str_with_tool_result,
+ test_format_msgs_to_str_with_image,
+ test_format_msgs_to_str_conversation_flow,
+ # Thinking block tests
+ test_format_msgs_to_str_thinking_excluded_by_default,
+ test_format_msgs_to_str_thinking_included,
+ test_format_msgs_to_str_thinking_only_message,
+ # Token threshold tests
+ test_format_msgs_to_str_all_within_threshold,
+ test_format_msgs_to_str_exceeds_threshold_truncate_older,
+ test_format_msgs_to_str_single_message_exceeds_threshold,
+ test_format_msgs_to_str_first_message_exceeds_threshold,
+ test_format_msgs_to_str_threshold_zero,
+ test_format_msgs_to_str_threshold_exact_fit,
+ test_format_msgs_to_str_threshold_one_less,
+ test_format_msgs_to_str_large_threshold,
+ # Edge cases tests
+ test_format_msgs_to_str_special_characters,
+ test_format_msgs_to_str_empty_content,
+ test_format_msgs_to_str_whitespace_only,
+ test_format_msgs_to_str_newlines_in_content,
+ test_format_msgs_to_str_very_long_single_word,
+ test_format_msgs_to_str_mixed_content_blocks,
+ test_format_msgs_to_str_multiple_separators,
+ test_format_msgs_to_str_tool_result_complex_output,
+ test_format_msgs_to_str_different_roles,
+ test_format_msgs_to_str_incremental_threshold_check,
+ test_format_msgs_to_str_negative_threshold,
+ test_format_msgs_to_str_preserves_newest_first,
+ test_format_msgs_to_str_base64_image,
+ test_format_msgs_to_str_audio_video_blocks,
+ test_format_msgs_to_str_unknown_block_type,
+ ]
+
+ passed = 0
+ failed = 0
+
+ for test in tests:
+ try:
+ print_test_header(test.__name__)
+ test()
+ passed += 1
+ except AssertionError as e:
+ print_fail(test.__name__, str(e))
+ failed += 1
+ except Exception as e:
+ print_error(test.__name__, str(e))
+ failed += 1
+
+ # 打印最终统计结果
+ print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+ print(f"{Colors.BOLD}Test Results Summary{Colors.RESET}")
+ print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+ print(f"{Colors.GREEN}{Colors.BOLD}✓ Passed: {passed}{Colors.RESET}")
+ if failed > 0:
+ print(f"{Colors.RED}{Colors.BOLD}✗ Failed: {failed}{Colors.RESET}")
+ else:
+ print(f"{Colors.GREEN}✗ Failed: {failed}{Colors.RESET}")
+ print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
+
+ if failed == 0:
+ print(f"\n{Colors.GREEN}{Colors.BOLD}🎉 All tests passed!{Colors.RESET}")
+ else:
+ print(f"\n{Colors.RED}{Colors.BOLD}💥 Some tests failed!{Colors.RESET}")
+
+ return failed == 0
+
+
+if __name__ == "__main__":
+ success = run_all_tests()
+ sys.exit(0 if success else 1)
diff --git a/tests/light/test_memory_formatter.py b/tests/light/test_memory_formatter.py
deleted file mode 100644
index 00f45bb1..00000000
--- a/tests/light/test_memory_formatter.py
+++ /dev/null
@@ -1,489 +0,0 @@
-"""Tests for MemoryFormatter."""
-
-# pylint: disable=W0212
-
-import logging
-
-from agentscope.message import Msg
-
-from test_utils import get_token_counter
-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__)
-
-
-# ANSI 颜色码
-class Colors:
- """ANSI color codes for terminal output."""
-
- GREEN = "\033[92m"
- RED = "\033[91m"
- YELLOW = "\033[93m"
- BLUE = "\033[94m"
- CYAN = "\033[96m"
- BOLD = "\033[1m"
- RESET = "\033[0m"
-
-
-def print_pass(test_name: str):
- """打印测试通过信息"""
- print(f"{Colors.GREEN}{Colors.BOLD}✓ {test_name} PASSED{Colors.RESET}")
-
-
-def print_fail(test_name: str, error: str):
- """打印测试失败信息"""
- print(f"{Colors.RED}{Colors.BOLD}✗ {test_name} FAILED: {error}{Colors.RESET}")
-
-
-def print_error(test_name: str, error: str):
- """打印测试错误信息"""
- print(f"{Colors.YELLOW}{Colors.BOLD}⚠ {test_name} ERROR: {error}{Colors.RESET}")
-
-
-def print_test_header(test_name: str):
- """打印测试标题"""
- print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
- print(f"{Colors.BLUE}{Colors.BOLD}Running: {test_name}{Colors.RESET}")
- print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
-
-
-def create_user_msg(content: str) -> Msg:
- """Create a user message."""
- return Msg(name="user", role="user", content=content)
-
-
-def create_assistant_msg(content: str) -> Msg:
- """Create an assistant message."""
- return Msg(name="assistant", role="assistant", content=content)
-
-
-def create_tool_use_msg(tool_name: str, tool_input: dict) -> Msg:
- """Create a message with tool_use content block."""
- return Msg(
- name="assistant",
- role="assistant",
- content=[
- {
- "type": "tool_use",
- "id": "call_123",
- "name": tool_name,
- "input": tool_input,
- },
- ],
- )
-
-
-def create_tool_result_msg(tool_name: str, output: str | list[dict]) -> Msg:
- """Create a message with tool_result content block."""
- return Msg(
- name="tool",
- role="user",
- content=[
- {
- "type": "tool_result",
- "id": "call_123",
- "name": tool_name,
- "output": output,
- },
- ],
- )
-
-
-def create_thinking_msg(thinking_content: str) -> Msg:
- """Create a message with thinking content block."""
- return Msg(
- name="assistant",
- role="assistant",
- content=[
- {
- "type": "thinking",
- "text": thinking_content,
- },
- ],
- )
-
-
-def create_image_msg(url: str = "") -> Msg:
- """Create a message with image content block."""
- content = [
- {
- "type": "image",
- "source": {"url": url} if url else {},
- },
- ]
- return Msg(name="assistant", role="assistant", content=content)
-
-
-def create_formatter(memory_compact_threshold: int = 4000) -> MemoryFormatter:
- """Create a MemoryFormatter instance for testing."""
- return MemoryFormatter(
- token_counter=get_token_counter(),
- memory_compact_threshold=memory_compact_threshold,
- )
-
-
-# ==================== _format_tool_result_output Tests ====================
-
-
-def test_format_tool_result_output_string():
- """Test _format_tool_result_output with string input."""
- result = MemoryFormatter._format_tool_result_output("Hello, world!")
- assert result == "Hello, world!", f"Expected 'Hello, world!', got: {result}"
- print_pass("test_format_tool_result_output_string")
-
-
-def test_format_tool_result_output_text_block():
- """Test _format_tool_result_output with text block."""
- output = [{"type": "text", "text": "This is text content"}]
- result = MemoryFormatter._format_tool_result_output(output)
- assert result == "This is text content", f"Expected 'This is text content', got: {result}"
- print_pass("test_format_tool_result_output_text_block")
-
-
-def test_format_tool_result_output_image_block():
- """Test _format_tool_result_output with image block."""
- output = [{"type": "image", "source": {"url": "https://example.com/image.png"}}]
- result = MemoryFormatter._format_tool_result_output(output)
- assert "[image]" in result, f"Expected '[image]' in result, got: {result}"
- assert "https://example.com/image.png" in result, f"Expected URL in result, got: {result}"
- print_pass("test_format_tool_result_output_image_block")
-
-
-def test_format_tool_result_output_file_block():
- """Test _format_tool_result_output with file block."""
- output = [{"type": "file", "path": "/path/to/file.txt", "name": "file.txt"}]
- result = MemoryFormatter._format_tool_result_output(output)
- assert "[file]" in result, f"Expected '[file]' in result, got: {result}"
- assert "file.txt" in result, f"Expected 'file.txt' in result, got: {result}"
- print_pass("test_format_tool_result_output_file_block")
-
-
-def test_format_tool_result_output_multiple_blocks():
- """Test _format_tool_result_output with multiple blocks."""
- output = [
- {"type": "text", "text": "First part"},
- {"type": "text", "text": "Second part"},
- ]
- result = MemoryFormatter._format_tool_result_output(output)
- assert "First part" in result, f"Expected 'First part' in result, got: {result}"
- assert "Second part" in result, f"Expected 'Second part' in result, got: {result}"
- # Multiple parts should be joined with newlines and bullets
- assert "- " in result, f"Expected bullet format in result, got: {result}"
- print_pass("test_format_tool_result_output_multiple_blocks")
-
-
-def test_format_tool_result_output_empty_list():
- """Test _format_tool_result_output with empty list."""
- result = MemoryFormatter._format_tool_result_output([])
- assert result == "", f"Expected empty string, got: {result}"
- print_pass("test_format_tool_result_output_empty_list")
-
-
-def test_format_tool_result_output_invalid_block():
- """Test _format_tool_result_output with invalid block (missing type)."""
- output = [{"text": "No type key"}]
- result = MemoryFormatter._format_tool_result_output(output)
- assert result == "", f"Expected empty string for invalid block, got: {result}"
- print_pass("test_format_tool_result_output_invalid_block")
-
-
-def test_format_tool_result_output_unknown_type():
- """Test _format_tool_result_output with unknown block type."""
- output = [{"type": "unknown_type", "data": "some data"}]
- result = MemoryFormatter._format_tool_result_output(output)
- assert result == "", f"Expected empty string for unknown type, got: {result}"
- print_pass("test_format_tool_result_output_unknown_type")
-
-
-# ==================== format (single message) Tests ====================
-
-
-def test_format_empty_messages():
- """Test format with empty message list."""
- formatter = create_formatter()
- result = formatter.format([])
- assert result == "", f"Expected empty string, got: {result}"
- print_pass("test_format_empty_messages")
-
-
-def test_format_single_user_message():
- """Test format with a single user message."""
- formatter = create_formatter()
- msgs = [create_user_msg("Hello, how are you?")]
- result = formatter.format(msgs)
-
- assert "user:" in result, f"Expected 'user:' in result, got: {result}"
- assert "Hello, how are you?" in result, f"Expected content in result, got: {result}"
- print_pass("test_format_single_user_message")
-
-
-def test_format_single_assistant_message():
- """Test format with a single assistant message."""
- formatter = create_formatter()
- msgs = [create_assistant_msg("I am fine, thank you!")]
- result = formatter.format(msgs)
-
- assert "assistant:" in result, f"Expected 'assistant:' in result, got: {result}"
- assert "I am fine, thank you!" in result, f"Expected content in result, got: {result}"
- print_pass("test_format_single_assistant_message")
-
-
-def test_format_with_tool_use():
- """Test format with tool_use message."""
- formatter = create_formatter()
- msgs = [create_tool_use_msg("read_file", {"path": "/test.txt"})]
- result = formatter.format(msgs)
-
- assert "tool_call=read_file" in result, f"Expected 'tool_call=read_file' in result, got: {result}"
- assert "params=" in result, f"Expected 'params=' in result, got: {result}"
- print_pass("test_format_with_tool_use")
-
-
-def test_format_with_tool_result():
- """Test format with tool_result message."""
- formatter = create_formatter()
- msgs = [create_tool_result_msg("read_file", "file content here")]
- result = formatter.format(msgs)
-
- assert "tool_result=read_file" in result, f"Expected 'tool_result=read_file' in result, got: {result}"
- assert "output=" in result, f"Expected 'output=' in result, got: {result}"
- print_pass("test_format_with_tool_result")
-
-
-def test_format_with_thinking_block():
- """Test that thinking blocks are skipped."""
- formatter = create_formatter()
- msgs = [create_thinking_msg("Let me think about this...")]
- result = formatter.format(msgs)
-
- # Thinking content should NOT appear in the result
- assert "Let me think about this" not in result, f"Thinking content should be skipped, got: {result}"
- print_pass("test_format_with_thinking_block")
-
-
-def test_format_with_image():
- """Test format with image content block."""
- formatter = create_formatter()
- msgs = [create_image_msg("https://example.com/image.png")]
- result = formatter.format(msgs)
-
- assert "[image]" in result, f"Expected '[image]' in result, got: {result}"
- print_pass("test_format_with_image")
-
-
-# ==================== format (multiple messages) Tests ====================
-
-
-def test_format_conversation():
- """Test format with a conversation."""
- formatter = create_formatter()
- msgs = [
- create_user_msg("What is Python?"),
- create_assistant_msg("Python is a programming language."),
- create_user_msg("Tell me more."),
- create_assistant_msg("Python is known for its readability and simplicity."),
- ]
- result = formatter.format(msgs)
-
- assert "round0" in result, f"Expected 'round0' in result, got: {result}"
- assert "round1" in result, f"Expected 'round1' in result, got: {result}"
- assert "round2" in result, f"Expected 'round2' in result, got: {result}"
- assert "round3" in result, f"Expected 'round3' in result, got: {result}"
- print_pass("test_format_conversation")
-
-
-def test_format_without_index():
- """Test format without round index."""
- formatter = create_formatter()
- msgs = [
- create_user_msg("Hello"),
- create_assistant_msg("Hi there!"),
- ]
- result = formatter.format(msgs, add_index=False)
-
- assert "round" not in result, f"Expected no 'round' prefix, got: {result}"
- print_pass("test_format_without_index")
-
-
-def test_format_without_time():
- """Test format without timestamp."""
- formatter = create_formatter()
- msgs = [create_user_msg("Test message")]
- result = formatter.format(msgs, add_time=False)
-
- # The result should not have timestamp brackets at the beginning
- # Note: this test may need adjustment based on actual timestamp format
- assert "user:" in result, f"Expected 'user:' in result, got: {result}"
- print_pass("test_format_without_time")
-
-
-def test_format_with_tool_conversation():
- """Test format with tool use and result in conversation."""
- formatter = create_formatter()
- msgs = [
- create_user_msg("Read the file."),
- create_tool_use_msg("read_file", {"path": "/data.txt"}),
- create_tool_result_msg("read_file", "File content here"),
- create_assistant_msg("The file contains: File content here"),
- ]
- result = formatter.format(msgs)
-
- assert "user:" in result
- assert "tool_call=read_file" in result
- assert "tool_result=read_file" in result
- assert "assistant:" in result
- print_pass("test_format_with_tool_conversation")
-
-
-# ==================== Token Threshold Tests ====================
-
-
-def test_format_low_threshold():
- """Test that older messages are skipped with low threshold."""
- formatter = create_formatter(memory_compact_threshold=100)
- msgs = []
- for i in range(20):
- msgs.append(create_user_msg(f"Question {i}: " + "x" * 50))
- msgs.append(create_assistant_msg(f"Answer {i}: " + "y" * 50))
-
- result = formatter.format(msgs)
-
- # With low threshold, not all messages should be included
- # The newest messages should be present
- assert "round39" in result or "round38" in result, f"Expected recent round in result, got: {result}"
- # Older messages might be truncated
- logger.info(f"Result length: {len(result)}")
- print_pass("test_format_low_threshold")
-
-
-def test_format_high_threshold():
- """Test that all messages are included with high threshold."""
- formatter = create_formatter(memory_compact_threshold=100000)
- msgs = [
- create_user_msg("Message 1"),
- create_assistant_msg("Response 1"),
- create_user_msg("Message 2"),
- create_assistant_msg("Response 2"),
- ]
- result = formatter.format(msgs)
-
- # All messages should be included
- assert "round0" in result
- assert "round1" in result
- assert "round2" in result
- assert "round3" in result
- print_pass("test_format_high_threshold")
-
-
-# ==================== Edge Cases Tests ====================
-
-
-def test_format_long_text_truncation():
- """Test that long text is truncated."""
- formatter = create_formatter()
- long_text = "x" * 5000 # Much longer than default max length
- msgs = [create_user_msg(long_text)]
- result = formatter.format(msgs)
-
- # The result should be shorter due to truncation
- assert len(result) < len(long_text), f"Expected truncated result, got length: {len(result)}"
- print_pass("test_format_long_text_truncation")
-
-
-def test_format_special_characters():
- """Test format with special characters in content."""
- formatter = create_formatter()
- msgs = [create_user_msg("Test with 中文, 日本語, émojis 🎉")]
- result = formatter.format(msgs)
-
- assert "中文" in result, f"Expected Chinese characters in result, got: {result}"
- print_pass("test_format_special_characters")
-
-
-def test_format_tool_result_with_complex_output():
- """Test format with complex tool result output."""
- formatter = create_formatter()
- complex_output = [
- {"type": "text", "text": "Operation completed"},
- {"type": "image", "source": {"url": "https://example.com/result.png"}},
- ]
- msgs = [create_tool_result_msg("process_data", complex_output)]
- result = formatter.format(msgs)
-
- assert "tool_result=process_data" in result, f"Expected tool result in result, got: {result}"
- print_pass("test_format_tool_result_with_complex_output")
-
-
-def run_all_tests():
- """Run all tests."""
- tests = [
- # _format_tool_result_output tests
- test_format_tool_result_output_string,
- test_format_tool_result_output_text_block,
- test_format_tool_result_output_image_block,
- test_format_tool_result_output_file_block,
- test_format_tool_result_output_multiple_blocks,
- test_format_tool_result_output_empty_list,
- test_format_tool_result_output_invalid_block,
- test_format_tool_result_output_unknown_type,
- # format tests (single message)
- test_format_empty_messages,
- test_format_single_user_message,
- test_format_single_assistant_message,
- test_format_with_tool_use,
- test_format_with_tool_result,
- test_format_with_thinking_block,
- test_format_with_image,
- # format tests (multiple messages)
- test_format_conversation,
- test_format_without_index,
- test_format_without_time,
- test_format_with_tool_conversation,
- # threshold tests
- test_format_low_threshold,
- test_format_high_threshold,
- # edge cases
- test_format_long_text_truncation,
- test_format_special_characters,
- test_format_tool_result_with_complex_output,
- ]
-
- passed = 0
- failed = 0
-
- for test in tests:
- try:
- print_test_header(test.__name__)
- test()
- passed += 1
- except AssertionError as e:
- print_fail(test.__name__, str(e))
- failed += 1
- except Exception as e:
- print_error(test.__name__, str(e))
- failed += 1
-
- # 打印最终统计结果
- print(f"\n{Colors.CYAN}{'=' * 60}{Colors.RESET}")
- print(f"{Colors.BOLD}Test Results Summary{Colors.RESET}")
- print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
- print(f"{Colors.GREEN}{Colors.BOLD}✓ Passed: {passed}{Colors.RESET}")
- if failed > 0:
- print(f"{Colors.RED}{Colors.BOLD}✗ Failed: {failed}{Colors.RESET}")
- else:
- print(f"{Colors.GREEN}✗ Failed: {failed}{Colors.RESET}")
- print(f"{Colors.CYAN}{'=' * 60}{Colors.RESET}")
-
- if failed == 0:
- print(f"\n{Colors.GREEN}{Colors.BOLD}🎉 All tests passed!{Colors.RESET}")
- else:
- print(f"\n{Colors.RED}{Colors.BOLD}💥 Some tests failed!{Colors.RESET}")
-
-
-if __name__ == "__main__":
- run_all_tests()
diff --git a/tests/light/test_reme_light.py b/tests/light/test_reme_light.py
index 9e2ac70a..4ec8f52a 100644
--- a/tests/light/test_reme_light.py
+++ b/tests/light/test_reme_light.py
@@ -1,197 +1,194 @@
-"""测试 ReMeLight"""
+"""测试 ReMeLight
+
+演示 ReMeLight 的完整功能,并使用 AsMsgHandler 跟踪每步 Token 变化:
+1. compact_tool_result - 压缩超长工具输出
+2. compact_memory - 生成压缩摘要
+3. summary_memory - 生成完整摘要并写入文件
+4. pre_reasoning_hook - 推理前预处理钩子
+5. memory_search - 语义搜索记忆
+6. ReMeInMemoryMemory.estimate_tokens - 估算 Token 使用
+7. ReMeInMemoryMemory.get_history_str - 获取格式化历史记录
+"""
import asyncio
-
-from agentscope.message import Msg
+import logging
+from test_utils import build_sample_messages, get_msg_handler
from reme.reme_light import ReMeLight
-# ==================== 消息创建辅助函数 ====================
-def create_user_msg(content: str) -> Msg:
- """创建用户消息"""
- return Msg(name="user", role="user", content=content)
-
-
-def create_assistant_msg(content: str) -> Msg:
- """创建助手消息"""
- return Msg(name="assistant", role="assistant", content=content)
-
-
-def create_tool_use_msg(tool_id: str, tool_name: str, tool_input: dict) -> Msg:
- """创建工具调用消息"""
- return Msg(
- name="assistant",
- role="assistant",
- content=[
- {
- "type": "tool_use",
- "id": tool_id,
- "name": tool_name,
- "input": tool_input,
- },
- ],
- )
-
-
-def create_tool_result_msg(tool_id: str, tool_name: str, output: str) -> Msg:
- """创建工具结果消息"""
- return Msg(
- name="tool",
- role="user",
- content=[
- {
- "type": "tool_result",
- "id": tool_id,
- "name": tool_name,
- "output": output,
- },
- ],
- )
-
-
-def create_thinking_msg(thinking_content: str) -> Msg:
- """创建思考消息"""
- return Msg(
- name="assistant",
- role="assistant",
- content=[
- {
- "type": "thinking",
- "text": thinking_content,
- },
- ],
- )
-
-
-# ==================== 构建模拟对话历史 ====================
-def build_sample_messages() -> list[Msg]:
- """构建一段包含多种消息类型的模拟对话"""
- messages = [
- # 用户询问 Python 版本
- create_user_msg("我想设置一个 Python 开发环境,你有什么建议?"),
- # 助手思考
- create_thinking_msg("用户想要搭建 Python 开发环境,我需要了解他的需求和偏好..."),
- # 助手回复
- create_assistant_msg(
- "好的!我建议使用 Python 3.11 或 3.12 版本,它们性能更好且功能丰富。"
- "你希望用于什么类型的开发?Web、数据科学还是其他?",
- ),
- # 用户提供更多信息
- create_user_msg("主要是做 Web 开发,使用 FastAPI 框架。另外我喜欢用 pyenv 管理版本。"),
- # 助手调用工具查询
- create_tool_use_msg(
- tool_id="call_001",
- tool_name="search_web",
- tool_input={"query": "FastAPI Python version compatibility 2024"},
- ),
- # 工具返回结果(模拟较长的输出)
- create_tool_result_msg(
- tool_id="call_001",
- tool_name="search_web",
- output=(
- "FastAPI 官方推荐使用 Python 3.8+ 版本,但 3.11/3.12 性能最佳。\n"
- "主要依赖:\n"
- "- Starlette: ASGI 框架\n"
- "- Pydantic v2: 数据验证\n"
- "- Uvicorn: ASGI 服务器\n"
- "最新版本 FastAPI 0.109+ 完全支持 Python 3.12。\n"
- "建议搭配 uv 或 pip-tools 进行依赖管理。"
- ),
- ),
- # 助手总结建议
- create_assistant_msg(
- "根据查询结果,我的建议是:\n"
- "1. **Python 版本**: 使用 Python 3.11 或 3.12(通过 pyenv 安装)\n"
- "2. **框架**: FastAPI 0.109+ 完全兼容这些版本\n"
- "3. **依赖管理**: 推荐使用 uv(更快)或 pip-tools\n"
- "4. **ASGI 服务器**: Uvicorn 配合 gunicorn 用于生产环境\n\n"
- "需要我帮你生成一个项目模板吗?",
- ),
- # 用户确认偏好
- create_user_msg("好的,我决定用 Python 3.12 + FastAPI + uv。请记住我的这些偏好。"),
- # 助手确认
- create_assistant_msg(
- "已记录你的开发偏好:\n"
- "- Python 版本: 3.12 (通过 pyenv 管理)\n"
- "- Web 框架: FastAPI\n"
- "- 包管理器: uv\n"
- "以后有相关问题我会参考这些偏好给你建议!",
- ),
- ]
- return messages
+def print_token_change(_step_name: str, before: int, after: int):
+ """打印 Token 变化统计。"""
+ change = after - before
+ change_pct = (change / before * 100) if before > 0 else 0
+ print(f" 📊 Token 统计: {before:,} → {after:,} (变化: {change:+,}, {change_pct:+.1f}%)")
# ==================== 主测试流程 ====================
async def main():
- """ReMeLight 主测试流程,演示完整的记忆管理功能。"""
+ """测试 ReMeLight 的完整功能,并跟踪每步 Token 变化。"""
+ # 初始化 AsMsgHandler 用于 Token 统计
+ msg_handler = get_msg_handler()
+
# 初始化 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/ 文件保留天数
+ default_as_llm_config={"model_name": "qwen3.5-35b-a3b"},
+ # default_embedding_model_config={"model_name": "text-embedding-v4"},
+ default_file_store_config={"fts_enabled": True, "vector_enabled": False},
)
+ logging.getLogger("reme").setLevel(logging.WARNING)
await reme.start()
- print("=" * 60)
+ print("=" * 70)
print("ReMeLight 已启动")
- print("=" * 60)
+ print("=" * 70)
- # 构建模拟对话历史
- messages = build_sample_messages()
- print(f"\n[原始消息数量]: {len(messages)} 条")
+ # 构建模拟对话历史(包含超长 tool_result,确保超过 128K token)
+ original_messages = build_sample_messages(include_large_tool_result=True)
+ initial_tokens = msg_handler.count_msgs_token(original_messages)
- # 1. 压缩超长工具输出(防止工具结果撑爆上下文)
- print("\n" + "-" * 40)
- print("[步骤 1] 压缩超长工具输出...")
- messages = await reme.compact_tool_result(messages)
- print(f"处理后消息数量: {len(messages)} 条")
+ print(f"\n[原始消息]: {len(original_messages)} 条, {initial_tokens:,} tokens")
+ print(f" 目标阈值: 128K = {128 * 1024:,} tokens")
+ print(f" 超出阈值: {initial_tokens > 128 * 1024}")
- # 2. 将历史对话压缩为结构化摘要(触发时机:上下文接近上限)
- print("\n" + "-" * 40)
- print("[步骤 2] 生成结构化压缩摘要...")
- summary = await reme.compact_memory(
+ # ==================== 1. compact_tool_result ====================
+ print("\n" + "=" * 70)
+ print("[步骤 1] compact_tool_result - 压缩超长工具输出")
+ print("=" * 70)
+
+ # 重新获取原始消息
+ messages = build_sample_messages(include_large_tool_result=True)
+ tokens_before = msg_handler.count_msgs_token(messages)
+ messages_after_step1 = await reme.compact_tool_result(messages)
+ tokens_after = msg_handler.count_msgs_token(messages_after_step1)
+
+ print(f" 消息数量: {len(messages)} → {len(messages_after_step1)}")
+ print_token_change("compact_tool_result", tokens_before, tokens_after)
+
+ # ==================== 2. compact_memory ====================
+ print("\n" + "=" * 70)
+ print("[步骤 2] compact_memory - 生成结构化压缩摘要")
+ print("=" * 70)
+
+ # 重新获取原始消息
+ messages = build_sample_messages(include_large_tool_result=True)
+ tokens_before = msg_handler.count_msgs_token(messages)
+ compact_summary = await reme.compact_memory(
messages=messages,
- previous_summary="", # 可传入上轮摘要,实现增量更新
+ previous_summary="",
)
- print(f"压缩摘要:\n{summary[:500]}..." if len(summary) > 500 else f"压缩摘要:\n{summary}")
+ summary_tokens = msg_handler.count_str_token(compact_summary)
- # 3. 后台异步提交摘要任务(不阻塞对话,摘要写入 memory/YYYY-MM-DD.md)
- print("\n" + "-" * 40)
- print("[步骤 3] 提交后台异步摘要任务...")
- reme.add_async_summary_task(messages=messages)
- print("异步任务已提交")
+ print(f" 输入消息 tokens: {tokens_before:,}")
+ print(f" 压缩摘要长度: {len(compact_summary)} 字符, {summary_tokens:,} tokens")
+ print(f" 压缩比: {summary_tokens / tokens_before * 100:.1f}%" if tokens_before > 0 else " 压缩比: N/A")
+ print(f" 摘要预览: {compact_summary[:200]}..." if len(compact_summary) > 200 else f" 摘要: {compact_summary}")
- # 4. 语义搜索记忆(向量 + BM25 混合检索)
- print("\n" + "-" * 40)
- print("[步骤 4] 语义搜索记忆...")
- result = await reme.memory_search(query="Python 版本偏好", max_results=5)
- print(f"搜索结果: {result}")
+ # ==================== 3. summary_memory ====================
+ print("\n" + "=" * 70)
+ print("[步骤 3] summary_memory - 生成完整摘要并写入文件")
+ print("=" * 70)
- # 5. 获取会话内存实例(ReMeInMemoryMemory,管理单次对话的上下文)
- print("\n" + "-" * 40)
- print("[步骤 5] 获取会话内存实例并估算 Token 使用...")
- memory = reme.get_in_memory_memory()
- # 将消息添加到内存中以便估算
+ # 重新获取原始消息
+ messages = build_sample_messages(include_large_tool_result=True)
+ tokens_before = msg_handler.count_msgs_token(messages)
+ summary_result = await reme.summary_memory(messages=messages)
+
+ print(f" 输入消息 tokens: {tokens_before:,}")
+ print(f" 摘要结果长度: {len(summary_result)} 字符")
+ print(f" 摘要预览: {summary_result[:200]}..." if len(summary_result) > 200 else f" 摘要: {summary_result}")
+
+ # ==================== 4. pre_reasoning_hook ====================
+ print("\n" + "=" * 70)
+ print("[步骤 4] pre_reasoning_hook - 推理前预处理")
+ print("=" * 70)
+
+ # 重新获取原始消息
+ messages = build_sample_messages(include_large_tool_result=True)
+ tokens_before = msg_handler.count_msgs_token(messages)
+ processed_messages, compressed_summary = await reme.pre_reasoning_hook(
+ messages=messages,
+ system_prompt="你是一个有帮助的 AI 助手。",
+ compressed_summary="",
+ max_input_length=128000,
+ compact_ratio=0.7,
+ memory_compact_reserve=10000,
+ enable_tool_result_compact=True,
+ tool_result_compact_keep_n=3,
+ )
+ tokens_after = msg_handler.count_msgs_token(processed_messages)
+ compressed_summary_tokens = msg_handler.count_str_token(compressed_summary)
+
+ print(f" 消息数量: {len(messages)} → {len(processed_messages)}")
+ print_token_change("pre_reasoning_hook", tokens_before, tokens_after)
+ print(f" 压缩摘要: {len(compressed_summary)} 字符, {compressed_summary_tokens:,} tokens")
+ print(f" 总上下文: {tokens_after + compressed_summary_tokens:,} tokens")
+
+ # ==================== 5. memory_search ====================
+ print("\n" + "=" * 70)
+ print("[步骤 5] memory_search - 语义搜索记忆")
+ print("=" * 70)
+
+ search_result = await reme.memory_search(query="Python 版本偏好", max_results=5)
+ if search_result.content:
+ print(f" 搜索结果: {search_result.content}")
+ else:
+ print(" 未找到相关记忆")
+
+ # ==================== 6 & 7. ReMeInMemoryMemory ====================
+ print("\n" + "=" * 70)
+ print("[步骤 6] ReMeInMemoryMemory - 会话内存管理")
+ print("=" * 70)
+
+ # 重新获取原始消息
+ messages = build_sample_messages(include_large_tool_result=True)
+ memory = ReMeLight.get_in_memory_memory()
for msg in messages:
await memory.add(msg)
- token_stats = await memory.estimate_tokens()
- print(f"当前上下文使用率: {token_stats['context_usage_ratio']:.1f}%")
- print(f"消息 Token 数: {token_stats['messages_tokens']}")
- print(f"预估总 Token 数: {token_stats['estimated_tokens']}")
+ print(f" 已添加 {len(messages)} 条原始消息到内存")
- # 6. 关闭前等待后台任务完成
- print("\n" + "-" * 40)
- print("[步骤 6] 等待后台任务完成...")
- summary_result = await reme.await_summary_tasks()
- print(f"后台摘要任务完成,结果长度: {len(summary_result)} 字符")
+ # 6.1 estimate_tokens
+ print("\n[6.1] estimate_tokens - 估算 Token 使用:")
+ token_stats = await memory.estimate_tokens(max_input_length=128000)
+ print(f" - 总消息数: {token_stats['total_messages']}")
+ print(f" - 消息 Token 数: {token_stats['messages_tokens']:,}")
+ print(f" - 压缩摘要 Token 数: {token_stats['compressed_summary_tokens']:,}")
+ print(f" - 预估总 Token 数: {token_stats['estimated_tokens']:,}")
+ print(f" - 最大输入长度: {token_stats['max_input_length']:,}")
+ print(f" - 上下文使用率: {token_stats['context_usage_ratio']:.2f}%")
+
+ # 6.2 get_history_str
+ print("\n[6.2] get_history_str - 格式化历史记录:")
+ history_str = await memory.get_history_str(max_input_length=128000)
+ print(history_str[:1000] + "..." if len(history_str) > 1000 else history_str)
+
+ # ==================== 等待后台任务完成 ====================
+ print("\n" + "=" * 70)
+ print("[步骤 7] 等待后台任务完成")
+ print("=" * 70)
+ await_result = await reme.await_summary_tasks()
+ print(f" 后台任务完成,结果长度: {len(await_result)} 字符")
+
+ # ==================== 总结 ====================
+ print("\n" + "=" * 70)
+ print("📊 Token 变化总结")
+ print("=" * 70)
+ print(f" 原始消息: {initial_tokens:,} tokens")
+ print(f" Step 1 compact_tool_result 后: {msg_handler.count_msgs_token(messages_after_step1):,} tokens")
+ print(f" Step 2 compact_memory 摘要: {summary_tokens:,} tokens")
+ print(
+ f" Step 4 pre_reasoning_hook 后: {tokens_after:,} tokens + 摘要 {compressed_summary_tokens:,} "
+ f"tokens = {tokens_after + compressed_summary_tokens:,} tokens",
+ )
+ print(
+ f" 最大节省: {initial_tokens - tokens_after:,} "
+ f"tokens ({(initial_tokens - tokens_after) / initial_tokens * 100:.1f}%)",
+ )
+ print(f" 目标阈值: {128 * 1024:,} tokens")
# 关闭 ReMeLight
await reme.close()
- print("\n" + "=" * 60)
+ print("\n" + "=" * 70)
print("ReMeLight 已关闭")
- print("=" * 60)
+ print("=" * 70)
if __name__ == "__main__":
diff --git a/tests/light/test_summarizer.py b/tests/light/test_summarizer.py
index 560efabd..a2f2d975 100644
--- a/tests/light/test_summarizer.py
+++ b/tests/light/test_summarizer.py
@@ -2,25 +2,22 @@
import asyncio
import datetime
-import logging
import tempfile
from pathlib import Path
from agentscope.message import Msg
+from agentscope.tool import Toolkit
from test_utils import (
get_dash_chat_model,
get_formatter,
get_token_counter,
)
+from reme.core.utils import get_std_logger
from reme.memory.file_based import Summarizer
+from reme.memory.tools.file import FileIO
-# 配置日志输出到控制台
-logging.basicConfig(
- level=logging.INFO,
- format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
-)
-logger = logging.getLogger(__name__)
+logger = get_std_logger()
# ANSI 颜色码
@@ -100,6 +97,16 @@ def create_tool_result_msg(tool_name: str, output: str) -> Msg:
)
+def create_toolkit(working_dir: str) -> Toolkit:
+ """Create a default Toolkit with FileIO tools for testing."""
+ toolkit = Toolkit()
+ file_io = FileIO(working_dir=working_dir)
+ toolkit.register_tool_function(file_io.read)
+ toolkit.register_tool_function(file_io.write)
+ toolkit.register_tool_function(file_io.edit)
+ return toolkit
+
+
def create_summarizer(working_dir: str = None, memory_dir: str = "memory"):
"""Create a Summarizer instance for testing."""
if working_dir is None:
@@ -114,9 +121,10 @@ def create_summarizer(working_dir: str = None, memory_dir: str = "memory"):
working_dir=working_dir,
memory_dir=memory_dir,
memory_compact_threshold=4000,
- chat_model=get_dash_chat_model(),
- formatter=get_formatter(),
token_counter=get_token_counter(),
+ toolkit=create_toolkit(working_dir),
+ as_llm=get_dash_chat_model(),
+ as_llm_formatter=get_formatter(),
),
working_dir,
)
@@ -195,9 +203,10 @@ def test_consecutive_summaries():
working_dir=working_dir,
memory_dir=memory_dir,
memory_compact_threshold=4000,
- chat_model=get_dash_chat_model(),
- formatter=get_formatter(),
token_counter=get_token_counter(),
+ toolkit=create_toolkit(working_dir),
+ as_llm=get_dash_chat_model(),
+ as_llm_formatter=get_formatter(),
)
# 第一轮对话
diff --git a/tests/light/test_tool_result_compactor.py b/tests/light/test_tool_result_compactor.py
index 059e626e..b6cb7c69 100644
--- a/tests/light/test_tool_result_compactor.py
+++ b/tests/light/test_tool_result_compactor.py
@@ -6,9 +6,8 @@ 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
+from reme.memory.file_based import ToolResultCompactor
+from reme.core.utils import is_truncated
def create_tool_result_msg(output: str | list, tool_name: str = "test_tool") -> Msg:
@@ -52,7 +51,7 @@ class TestToolResultCompactor:
_ = asyncio.run(op.call(messages=messages))
output = messages[0].content[0]["output"]
- assert TRUNCATION_MARKER_START in output
+ assert is_truncated(output)
assert "[Full content saved to:" in output
# Verify file was created
@@ -69,7 +68,7 @@ class TestToolResultCompactor:
"""Test that already truncated content is not re-truncated."""
with tempfile.TemporaryDirectory() as tmpdir:
op = ToolResultCompactor(tool_result_dir=tmpdir, tool_result_threshold=100)
- truncated_content = f"head{TRUNCATION_MARKER_START}(100 chars omitted)<<>>tail"
+ truncated_content = "head<<>>(100 chars omitted)<<>>tail"
messages = [create_tool_result_msg(truncated_content)]
asyncio.run(op.call(messages=messages))
@@ -87,7 +86,7 @@ class TestToolResultCompactor:
asyncio.run(op.call(messages=messages))
text_block = messages[0].content[0]["output"][0]
- assert TRUNCATION_MARKER_START in text_block["text"]
+ assert is_truncated(text_block["text"])
assert len(list(Path(tmpdir).glob("*.txt"))) == 1
def test_list_output_no_truncation_when_short(self):
@@ -116,9 +115,9 @@ class TestToolResultCompactor:
asyncio.run(op.call(messages=messages))
output = messages[0].content[0]["output"]
- assert TRUNCATION_MARKER_START in output[0]["text"]
+ assert is_truncated(output[0]["text"])
assert output[1]["text"] == "short" # unchanged
- assert TRUNCATION_MARKER_START in output[2]["text"]
+ assert is_truncated(output[2]["text"])
assert len(list(Path(tmpdir).glob("*.txt"))) == 2
def test_list_output_mixed_block_types(self):
@@ -134,7 +133,7 @@ class TestToolResultCompactor:
asyncio.run(op.call(messages=messages))
output = messages[0].content[0]["output"]
- assert TRUNCATION_MARKER_START in output[0]["text"]
+ assert is_truncated(output[0]["text"])
assert output[1] == {"type": "image", "source": {"type": "url", "url": "http://example.com/img.png"}}
assert len(list(Path(tmpdir).glob("*.txt"))) == 1
diff --git a/tests/light/test_utils.py b/tests/light/test_utils.py
index f5cae021..19740fe6 100644
--- a/tests/light/test_utils.py
+++ b/tests/light/test_utils.py
@@ -1,50 +1,22 @@
"""Test utilities for copaw tests."""
import os
-from pathlib import Path
-from typing import Any
-from loguru import logger
+from agentscope.message import Msg, ThinkingBlock, TextBlock, ToolUseBlock, ToolResultBlock
-_token_counter = None
+from reme.memory.file_based import AsMsgHandler
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.
+ return get_hf_token_counter()
- 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
+def get_msg_handler() -> AsMsgHandler:
+ """Get AsMsgHandler instance."""
+ return AsMsgHandler(token_counter=get_token_counter())
def get_dash_chat_model(model_name: str = "qwen3.5-plus"):
@@ -54,8 +26,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 +35,422 @@ 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."""
+ return OpenAIChatFormatter()
- 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
+def generate_large_code_content(target_tokens: int = 50000) -> str:
+ """生成大量代码内容,用于测试超长 tool_result。
- 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
+ Args:
+ target_tokens: 目标 token 数(约 4 字符/token)
- return ReMeChatFormatter(token_counter=get_token_counter())
+ Returns:
+ 生成的代码内容字符串
+ """
+ code_template = '''
+# === File: src/module_{idx}/handlers.py ===
+"""Handler module {idx} for processing requests."""
+
+import asyncio
+import logging
+from typing import Any, Dict, List, Optional
+from dataclasses import dataclass, field
+from datetime import datetime
+
+logger = logging.getLogger(__name__)
+
+
+@dataclass
+class RequestContext_{idx}:
+ """Context for request processing in module {idx}."""
+ request_id: str
+ user_id: str
+ timestamp: datetime = field(default_factory=datetime.now)
+ metadata: Dict[str, Any] = field(default_factory=dict)
+ headers: Dict[str, str] = field(default_factory=dict)
+ query_params: Dict[str, str] = field(default_factory=dict)
+ body: Optional[bytes] = None
+ processed: bool = False
+ error_message: Optional[str] = None
+
+
+class Handler_{idx}:
+ """Main handler class for module {idx}."""
+
+ def __init__(self, config: Dict[str, Any]):
+ self.config = config
+ self.cache: Dict[str, Any] = {{}}
+ self.metrics: Dict[str, int] = {{
+ "requests_processed": 0,
+ "errors": 0,
+ "cache_hits": 0,
+ "cache_misses": 0,
+ }}
+ self._initialized = False
+ logger.info(f"Handler_{idx} initialized with config: {{config}}")
+
+ async def initialize(self) -> None:
+ """Initialize the handler with async resources."""
+ if self._initialized:
+ logger.warning("Handler_{idx} already initialized")
+ return
+
+ # Simulate async initialization
+ await asyncio.sleep(0.01)
+ self._initialized = True
+ logger.info("Handler_{idx} initialization complete")
+
+ async def process_request(self, context: RequestContext_{idx}) -> Dict[str, Any]:
+ """Process an incoming request.
+
+ Args:
+ context: The request context containing all request data
+
+ Returns:
+ Dict containing the response data
+ """
+ if not self._initialized:
+ raise RuntimeError("Handler not initialized")
+
+ self.metrics["requests_processed"] += 1
+
+ try:
+ # Check cache first
+ cache_key = f"{{context.request_id}}_{{context.user_id}}"
+ if cache_key in self.cache:
+ self.metrics["cache_hits"] += 1
+ return self.cache[cache_key]
+
+ self.metrics["cache_misses"] += 1
+
+ # Process the request
+ result = await self._do_process(context)
+
+ # Cache the result
+ self.cache[cache_key] = result
+ context.processed = True
+
+ return result
+
+ except Exception as e:
+ self.metrics["errors"] += 1
+ context.error_message = str(e)
+ logger.exception(f"Error processing request {{context.request_id}}: {{e}}")
+ raise
+
+ async def _do_process(self, context: RequestContext_{idx}) -> Dict[str, Any]:
+ """Internal processing logic."""
+ # Simulate some processing
+ await asyncio.sleep(0.001)
+
+ return {{
+ "status": "success",
+ "request_id": context.request_id,
+ "user_id": context.user_id,
+ "processed_at": datetime.now().isoformat(),
+ "module": "module_{idx}",
+ "data": {{
+ "result": f"Processed by handler_{idx}",
+ "metadata": context.metadata,
+ }}
+ }}
+
+ def get_metrics(self) -> Dict[str, int]:
+ """Return current metrics."""
+ return self.metrics.copy()
+
+ async def cleanup(self) -> None:
+ """Cleanup resources."""
+ self.cache.clear()
+ self._initialized = False
+ logger.info("Handler_{idx} cleaned up")
+
+'''
+
+ # 每个模块约 2000 字符 ≈ 500 tokens
+ # 目标 target_tokens,需要 target_tokens / 500 个模块
+ num_modules = max(1, target_tokens // 500)
+
+ parts = [f"# 大型项目代码检索结果\n# 共找到 {num_modules} 个相关模块\n"]
+ for i in range(num_modules):
+ parts.append(code_template.format(idx=i))
+
+ return "".join(parts)
+
+
+def build_sample_messages(include_large_tool_result: bool = True) -> list[Msg]:
+ """构建一段包含多种消息类型的模拟对话。
+
+ Args:
+ include_large_tool_result: 是否包含大型 tool_result,确保超过 128K token
+
+ Returns:
+ 消息列表
+ """
+ messages = [
+ Msg(
+ name="user",
+ role="user",
+ content="我想设置一个 Python 开发环境,你有什么建议?",
+ ),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ThinkingBlock(type="thinking", thinking="用户想要搭建 Python 开发环境,我需要了解他的需求和偏好..."),
+ TextBlock(
+ type="text",
+ text="好的!我建议使用 Python 3.11 或 3.12 版本,它们性能更好且功能丰富。"
+ "你希望用于什么类型的开发?Web、数据科学还是其他?",
+ ),
+ ],
+ ),
+ Msg(name="user", role="user", content="主要是做 Web 开发,使用 FastAPI 框架。另外我喜欢用 pyenv 管理版本。"),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ThinkingBlock(type="thinking", thinking="用户要做 FastAPI Web 开发,我需要搜索相关信息..."),
+ ToolUseBlock(
+ type="tool_use",
+ id="call_001",
+ name="search_web",
+ input={"query": "FastAPI Python version compatibility 2024"},
+ ),
+ ],
+ ),
+ Msg(
+ name="system",
+ role="system",
+ content=[
+ ToolResultBlock(
+ type="tool_result",
+ id="call_001",
+ name="search_web",
+ output=[
+ TextBlock(
+ type="text",
+ text="FastAPI 官方推荐使用 Python 3.8+ 版本,但 3.11/3.12 性能最佳。\n"
+ "主要依赖:\n"
+ "- Starlette: ASGI 框架\n"
+ "- Pydantic v2: 数据验证\n"
+ "- Uvicorn: ASGI 服务器\n"
+ "最新版本 FastAPI 0.109+ 完全支持 Python 3.12。\n"
+ "建议搭配 uv 或 pip-tools 进行依赖管理。",
+ ),
+ ],
+ ),
+ ],
+ ),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ThinkingBlock(
+ type="thinking",
+ thinking="根据搜索结果,FastAPI 完全支持 Python 3.12,结合用户使用 pyenv 的习惯,我来给出完整的环境配置建议...",
+ ),
+ TextBlock(
+ type="text",
+ text="根据最新信息,我为你整理了 FastAPI 开发环境配置方案:\n\n"
+ "**1. Python 版本**\n"
+ "推荐 Python 3.12,性能最佳且完全兼容 FastAPI 0.109+。\n\n"
+ "**2. pyenv 安装步骤**\n"
+ "```bash\n"
+ "pyenv install 3.12.2\n"
+ "pyenv local 3.12.2\n"
+ "```\n\n"
+ "**3. 虚拟环境与依赖管理**\n"
+ "建议使用 `uv` 作为包管理器,它比 pip 快 10-100 倍。\n\n"
+ "需要我帮你生成项目初始化脚本吗?",
+ ),
+ ],
+ ),
+ Msg(name="user", role="user", content="好的,请帮我生成一个完整的项目初始化脚本,包括 pyproject.toml 配置。"),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ThinkingBlock(
+ type="thinking",
+ thinking="用户需要完整的项目初始化脚本,我需要创建 pyproject.toml 并设置好 FastAPI 相关依赖...",
+ ),
+ ToolUseBlock(
+ type="tool_use",
+ id="call_002",
+ name="create_file",
+ input={
+ "path": "pyproject.toml",
+ "content": "[project]\nname = 'my-fastapi-app'\nversion = '0.1.0'\nrequires-python = "
+ "'>=3.11'\ndependencies = ['fastapi>=0.109.0', 'uvicorn[standard]>=0.27.0', "
+ "'pydantic>=2.0']",
+ },
+ ),
+ ],
+ ),
+ Msg(
+ name="system",
+ role="system",
+ content=[
+ ToolResultBlock(
+ type="tool_result",
+ id="call_002",
+ name="create_file",
+ output=[
+ TextBlock(type="text", text="文件 pyproject.toml 已创建成功。"),
+ ],
+ ),
+ ],
+ ),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ToolUseBlock(
+ type="tool_use",
+ id="call_003",
+ name="create_file",
+ input={
+ "path": "init_project.sh",
+ "content": "#!/bin/bash\nset -e\n\n# 使用 pyenv 设置 Python 版本\npyenv install 3.12.2 "
+ "--skip-existing\npyenv local 3.12.2\n\n# 创建虚拟环境\npython -m venv "
+ ".venv\nsource .venv/bin/activate\n\n# 安装 uv 并使用它安装依赖\npip install uv"
+ "\nuv pip install -e .\n\necho '环境初始化完成!'",
+ },
+ ),
+ ],
+ ),
+ Msg(
+ name="system",
+ role="system",
+ content=[
+ ToolResultBlock(
+ type="tool_result",
+ id="call_003",
+ name="create_file",
+ output=[
+ TextBlock(type="text", text="文件 init_project.sh 已创建成功。"),
+ ],
+ ),
+ ],
+ ),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ TextBlock(
+ type="text",
+ text="我已为你创建了两个文件:\n\n"
+ "1. **pyproject.toml** - 项目配置文件,包含 FastAPI 核心依赖\n"
+ "2. **init_project.sh** - 一键初始化脚本\n\n"
+ "运行以下命令即可初始化项目:\n"
+ "```bash\n"
+ "chmod +x init_project.sh && ./init_project.sh\n"
+ "```\n\n"
+ "还有什么需要帮助的吗?",
+ ),
+ ],
+ ),
+ Msg(name="user", role="user", content="太棒了!请帮我搜索一下项目中所有的 handler 相关代码。"),
+ ]
+
+ # 添加大型代码搜索结果(确保超过 128K token)
+ if include_large_tool_result:
+ # 生成超大的代码搜索结果,目标 ~140K tokens
+ large_code_content = generate_large_code_content(target_tokens=140000)
+
+ messages.extend(
+ [
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ThinkingBlock(
+ type="thinking",
+ thinking="用户要我搜索项目中的 handler 代码,我需要使用代码搜索工具...",
+ ),
+ ToolUseBlock(
+ type="tool_use",
+ id="call_004",
+ name="search_codebase",
+ input={"query": "handler class implementation"},
+ ),
+ ],
+ ),
+ Msg(
+ name="system",
+ role="system",
+ content=[
+ ToolResultBlock(
+ type="tool_result",
+ id="call_004",
+ name="search_codebase",
+ output=[
+ TextBlock(type="text", text=large_code_content),
+ ],
+ ),
+ ],
+ ),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ThinkingBlock(type="thinking", thinking="搜索返回了大量 handler 代码,我需要为用户整理一下..."),
+ TextBlock(
+ type="text",
+ text="我已经找到了项目中所有的 handler 相关代码。\n\n"
+ "这些 handler 类包含:\n"
+ "- 请求处理逻辑\n"
+ "- 缓存管理\n"
+ "- 指标统计\n"
+ "- 异步初始化\n\n"
+ "你需要我详细解释某个具体的 handler 吗?",
+ ),
+ ],
+ ),
+ ],
+ )
+
+ # 添加更多对话
+ messages.extend(
+ [
+ Msg(
+ name="user",
+ role="user",
+ content="还有一个问题,我应该如何配置 VS Code 来获得最佳的 FastAPI 开发体验?",
+ ),
+ Msg(
+ name="assistant",
+ role="assistant",
+ content=[
+ ThinkingBlock(
+ type="thinking",
+ thinking="用户询问 VS Code 配置,我需要推荐适合 FastAPI 开发的扩展和设置...",
+ ),
+ TextBlock(
+ type="text",
+ text="VS Code 的 FastAPI 开发配置建议:\n\n"
+ "**推荐扩展:**\n"
+ "- Python (Microsoft)\n"
+ "- Pylance - 类型检查和智能补全\n"
+ "- Ruff - 快速 linter 和 formatter\n"
+ "- REST Client - API 测试\n\n"
+ "**settings.json 配置:**\n"
+ "```json\n"
+ "{\n"
+ ' "python.defaultInterpreterPath": ".venv/bin/python",\n'
+ ' "[python]": {\n'
+ ' "editor.defaultFormatter": "charliermarsh.ruff",\n'
+ ' "editor.formatOnSave": true\n'
+ " }\n"
+ "}\n"
+ "```\n\n"
+ "这样配置后,你就能获得完整的类型提示和自动格式化支持了!",
+ ),
+ ],
+ ),
+ ],
+ )
+
+ return messages