mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
feat(memory): replace memory formatter with AsMsgHandler for enhanced message processing
This commit is contained in:
commit
6b742b6719
50 changed files with 3958 additions and 1921 deletions
75
README.md
75
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<br>(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<br>History compaction]
|
||||
ReMeLight --> SummaryMemory[summary_memory<br>Write memory to files]
|
||||
ReMeLight --> CompactToolResult[compact_tool_result<br>Oversized tool output compaction]
|
||||
ReMeLight --> MemSearch[memory_search<br>Semantic search]
|
||||
ReMeLight --> InMemory[get_in_memory_memory<br>ReMeInMemoryMemory]
|
||||
CompactMemory --> Compactor[Compactor<br>ReActAgent]
|
||||
SummaryMemory --> Summarizer[Summarizer<br>ReActAgent + file tools]
|
||||
CompactToolResult --> ToolResultCompactor[ToolResultCompactor<br>Truncate + save to file]
|
||||
Summarizer --> FileIO[FileIO<br>read / write / edit]
|
||||
FileIO --> MemoryFiles[memory/YYYY-MM-DD.md]
|
||||
ToolResultCompactor --> ToolResultFiles[tool_result/*.txt]
|
||||
MemoryFiles -.->|File change| FileWatcher[Async File Watcher]
|
||||
|
|
|
|||
75
README_ZH.md
75
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<br>(继承 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<br>历史对话压缩]
|
||||
ReMeLight --> SummaryMemory[summary_memory<br>记忆写入文件]
|
||||
ReMeLight --> CompactToolResult[compact_tool_result<br>超长工具输出压缩]
|
||||
ReMeLight --> MemSearch[memory_search<br>语义搜索]
|
||||
ReMeLight --> InMemory[get_in_memory_memory<br>ReMeInMemoryMemory]
|
||||
CompactMemory --> Compactor[Compactor<br>ReActAgent]
|
||||
SummaryMemory --> Summarizer[Summarizer<br>ReActAgent + 文件工具]
|
||||
CompactToolResult --> ToolResultCompactor[ToolResultCompactor<br>截断 + 转存文件]
|
||||
Summarizer --> FileIO[FileIO<br>read / write / edit]
|
||||
FileIO --> MemoryFiles[memory/YYYY-MM-DD.md]
|
||||
ToolResultCompactor --> ToolResultFiles[tool_result/*.txt]
|
||||
MemoryFiles -.->|文件变更| FileWatcher[异步文件监控]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
|
|
|||
9
reme/core/as_llm/__init__.py
Normal file
9
reme/core/as_llm/__init__.py
Normal file
|
|
@ -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)
|
||||
9
reme/core/as_llm_formatter/__init__.py
Normal file
9
reme/core/as_llm_formatter/__init__.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
83
reme/core/schema/as_msg_stat.py
Normal file
83
reme/core/schema/as_msg_stat.py
Normal file
|
|
@ -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"<thinking>\n{truncate_text(self.text, max_length)}\n</thinking>"
|
||||
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])
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"] = {}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
23
reme/core/utils/hf_token_counter_utils.py
Normal file
23
reme/core/utils/hf_token_counter_utils.py
Normal file
|
|
@ -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
|
||||
109
reme/core/utils/std_logger.py
Normal file
109
reme/core/utils/std_logger.py
Normal file
|
|
@ -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
|
||||
55
reme/core/utils/truncate_text_utils.py
Normal file
55
reme/core/utils/truncate_text_utils.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
"""Utility functions for truncating long text strings."""
|
||||
|
||||
from .std_logger import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
TRUNCATION_MARKER_START = "<<<TRUNCATED>>>"
|
||||
TRUNCATION_MARKER_END = "<<<END_TRUNCATED>>>"
|
||||
|
||||
|
||||
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
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
400
reme/memory/file_based/as_msg_handler.py
Normal file
400
reme/memory/file_based/as_msg_handler.py
Normal file
|
|
@ -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
|
||||
0
reme/memory/file_based/component/__init__.py
Normal file
0
reme/memory/file_based/component/__init__.py
Normal file
|
|
@ -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:
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
@ -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):
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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"<error: {e}>"
|
||||
|
||||
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 (
|
||||
|
|
|
|||
|
|
@ -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 = "<<<TRUNCATED>>>"
|
||||
TRUNCATION_MARKER_END = "<<<END_TRUNCATED>>>"
|
||||
|
||||
|
||||
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
|
||||
7
reme/memory/tools/file/__init__.py
Normal file
7
reme/memory/tools/file/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""File-based memory tool implementations."""
|
||||
|
||||
from .file_io import FileIO
|
||||
|
||||
__all__ = [
|
||||
"FileIO",
|
||||
]
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
1285
tests/light/test_context_check.py
Normal file
1285
tests/light/test_context_check.py
Normal file
File diff suppressed because it is too large
Load diff
891
tests/light/test_format_msgs_to_str.py
Normal file
891
tests/light/test_format_msgs_to_str.py
Normal file
|
|
@ -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 "<thinking>" in result, f"Expected thinking tag in result, got: {result}"
|
||||
verify_result_within_threshold(handler, result, threshold, "thinking_included", msgs)
|
||||
print_pass("test_format_msgs_to_str_thinking_included")
|
||||
|
||||
|
||||
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)
|
||||
|
|
@ -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()
|
||||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
)
|
||||
|
||||
# 第一轮对话
|
||||
|
|
|
|||
|
|
@ -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)<<<END_TRUNCATED>>>tail"
|
||||
truncated_content = "head<<<TRUNCATED>>>(100 chars omitted)<<<END_TRUNCATED>>>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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue