mirror of
https://gitcode.com/JianFeeeee/TrulyMEM-TrueHumanMEM.git
synced 2026-09-20 00:48:52 +00:00
!1 merge test into main
新增 context_rewrite 工具优化系统提示词 Created-by: JianFeeeee Commit-by: JianFeeeee;root Merged-by: JianFeeeee Description: feat: 添加 context_rewrite 工具并清理工具限制器死配置 test: 修复测试套件结构问题 fix: 修复 context_rewrite 工具调用问题并优化系统提示词 See merge request: JianFeeeee/TrulyMEM-TrueHumanMEM!1
This commit is contained in:
@ -13,9 +13,12 @@
|
||||
|
||||
*The More Human Choice.*
|
||||
|
||||
> ⚠️ **当前分支**: `test` — 测试分支,用于测试实验性功能。此处代码可能不稳定,不应用于生产环境。稳定版本请参阅 `main` 分支。
|
||||
|
||||
[](https://www.gnu.org/licenses/gpl-3.0)
|
||||
[](https://www.python.org/downloads/)
|
||||
[]()
|
||||
[]()
|
||||
|
||||
---
|
||||
|
||||
|
||||
@ -13,9 +13,12 @@
|
||||
|
||||
*The More Human Choice.*
|
||||
|
||||
> ⚠️ **Current Branch**: `test` — Testing branch for experimental features. Code here may be unstable and should not be used in production. For stable releases, see the `main` branch.
|
||||
|
||||
[](https://www.gnu.org/licenses/gpl-3.0)
|
||||
[](https://www.python.org/downloads/)
|
||||
[]()
|
||||
[]()
|
||||
|
||||
---
|
||||
|
||||
|
||||
@ -89,6 +89,7 @@
|
||||
| `memory_commit` | 写入记忆 | 存储重要信息 |
|
||||
| `memory_purge` | 删除记忆 | 修正错误信息 |
|
||||
| `memory_introspect` | 查看状态 | 监控记忆系统 |
|
||||
| `context_rewrite` | 压缩工具调用上下文 | 工具调用≥2次后,压缩JSON为自然语言摘要 |
|
||||
|
||||
### 人设工具
|
||||
| 工具 | 功能 | 使用场景 |
|
||||
@ -104,6 +105,57 @@
|
||||
| `task_delete` | 删除任务 | 清理完成任务 |
|
||||
| `task_link_info` | 关联信息 | 连接任务与记忆 |
|
||||
|
||||
## context_rewrite 使用规则
|
||||
|
||||
### ⚠️ 强制触发条件
|
||||
|
||||
**每调用 5 次记忆相关工具,必须调用一次 context_rewrite!**
|
||||
|
||||
记忆相关工具包括:
|
||||
- `memory_recall` - 检索记忆
|
||||
- `memory_commit` - 写入记忆
|
||||
- `memory_purge` - 删除记忆
|
||||
- `memory_introspect` - 查看状态
|
||||
- `persona_update` - 更新人设
|
||||
- `persona_clear` - 清除人设
|
||||
- `task_create` - 创建任务
|
||||
- `task_set_state` - 设置状态
|
||||
- `task_delete` - 删除任务
|
||||
- `task_link_info` - 关联信息
|
||||
|
||||
**触发规则**:
|
||||
- 累计调用 5 次记忆工具 → 必须调用 context_rewrite
|
||||
- 累计调用 10 次记忆工具 → 必须调用 context_rewrite
|
||||
- 以此类推...
|
||||
|
||||
**目的**:
|
||||
- 保持上下文精简,只保留AI真正需要的信息
|
||||
- 避免无用的JSON细节填满上下文
|
||||
- 提高后续推理效率
|
||||
|
||||
### 使用场景
|
||||
|
||||
当你已经执行了多次工具调用,且:
|
||||
- 工具结果的JSON细节你已经理解,不再需要原始格式
|
||||
- 但你需要记住"我调用了哪些工具、得到了什么结论"
|
||||
- 继续携带原始JSON会干扰后续推理
|
||||
|
||||
→ 调用 context_rewrite 压缩上下文
|
||||
|
||||
**强制格式要求**:
|
||||
- 必须标注 `[工具调用总结: 本次总结了 N 次工具调用 | 调用工具: tool1, tool2]`
|
||||
- 必须保留关键语义信息
|
||||
- 不可删除用户原始消息
|
||||
- 不可歪曲工具返回的关键事实
|
||||
|
||||
**示例**:
|
||||
```
|
||||
[工具调用总结: 本次总结了 2 次工具调用 | 调用工具: memory_recall, memory_recall]
|
||||
|
||||
- 查询人设图:未找到人设,使用默认身份
|
||||
- 查询工作记忆链:发现 Task_成语接龙,状态已暂停,当前成语为虎作伥
|
||||
```
|
||||
|
||||
## 每轮对话强制要求
|
||||
|
||||
### ⚠️ 执行顺序(每轮必须)
|
||||
@ -206,20 +258,27 @@ AI操作:
|
||||
- 查询意图: "TaskNode,工作记忆,任务链"
|
||||
- 目的: 获取之前的任务上下文,了解对话历史
|
||||
|
||||
2. **用户提到"刚才"、"之前"、"上次"**
|
||||
2. **用户提到"刚才"、"之前"、"上次"、"刚刚"**
|
||||
- 例: "刚才我们聊了什么?"
|
||||
- 例: "继续刚才的话题"
|
||||
- 例: "关于刚才的成语接龙..."
|
||||
- 例: "我不是刚刚给你讲了个故事嘛"
|
||||
|
||||
3. **用户询问对话历史**
|
||||
3. **用户使用指代词(这个故事、那个故事、这件事等)**
|
||||
- 例: "你给我整体讲一下这个故事吧" → 必须查询工作记忆链确定"这个故事"指什么
|
||||
- 例: "继续那个任务" → 必须查询工作记忆链确定"那个任务"是什么
|
||||
- 例: "复述一下" → 必须查询工作记忆链确定要复述什么
|
||||
- **关键**: 指代词必须通过工作记忆链解析,不能凭空猜测!
|
||||
|
||||
4. **用户询问对话历史**
|
||||
- 例: "我们之前说了什么?"
|
||||
- 例: "我们聊过X吗?"
|
||||
|
||||
4. **连续性任务被打断后恢复**
|
||||
5. **连续性任务被打断后恢复**
|
||||
- 例: 用户突然回到之前的话题
|
||||
- 例: 用户要求继续之前的任务
|
||||
|
||||
5. **涉及上下文的引用**
|
||||
6. **涉及上下文的引用**
|
||||
- 例: "那个东西"(需要查询上下文)
|
||||
- 例: "继续"(需要查询当前任务)
|
||||
|
||||
@ -350,7 +409,9 @@ AI操作步骤:
|
||||
- [ ] 步骤3: 是否根据人设和工作记忆链生成回复?
|
||||
- [ ] 步骤4: 是否更新了工作记忆链?
|
||||
- [ ] 涉及上下文引用时是否查询了工作记忆链?
|
||||
- [ ] 用户提到"刚才/之前/上次"时是否查询了工作记忆链?
|
||||
- [ ] 用户提到"刚才/之前/上次/刚刚"时是否查询了工作记忆链?
|
||||
- [ ] 用户使用指代词(这个故事、那个任务等)时是否通过工作记忆链解析?
|
||||
- [ ] 累计调用5次记忆工具后是否调用了 context_rewrite?
|
||||
|
||||
---
|
||||
|
||||
|
||||
@ -60,10 +60,8 @@ class BackendServer:
|
||||
self._lock = threading.Lock()
|
||||
self._config = {"api_key": "", "base_url": "https://api.deepseek.com", "model": "deepseek-chat"}
|
||||
self._tool_limits = {
|
||||
"persona_query_max": 1,
|
||||
"persona_update_max": 1,
|
||||
"task_query_max": 4,
|
||||
"task_update_max": 2,
|
||||
"task_update_max": 5,
|
||||
"memory_query_max": 20,
|
||||
"memory_update_max": 10,
|
||||
}
|
||||
@ -119,9 +117,7 @@ class BackendServer:
|
||||
def _create_tool_limiter(self):
|
||||
from .tool_limiter import ToolLimiter, ToolLimits
|
||||
limits = ToolLimits(
|
||||
persona_query_max=self._tool_limits.get("persona_query_max", 1),
|
||||
persona_update_max=self._tool_limits.get("persona_update_max", 1),
|
||||
task_query_max=self._tool_limits.get("task_query_max", 4),
|
||||
task_update_max=self._tool_limits.get("task_update_max", 5),
|
||||
memory_query_max=self._tool_limits.get("memory_query_max", 20),
|
||||
memory_update_max=self._tool_limits.get("memory_update_max", 10),
|
||||
@ -243,6 +239,29 @@ class BackendServer:
|
||||
|
||||
self._tool_limiter.record_call(tool_call.function.name, args)
|
||||
|
||||
if tool_call.function.name == "context_rewrite":
|
||||
result = execute_tool(self._graph, tool_call.function.name, args)
|
||||
result_data = json.loads(result)
|
||||
|
||||
# 记录到 tool_calls,让 TUI 显示这个工具调用
|
||||
tool_calls.append({
|
||||
"name": tool_call.function.name,
|
||||
"arguments": args,
|
||||
"result": result
|
||||
})
|
||||
|
||||
if result_data.get("status") == "success":
|
||||
user_msg = messages_history[0]
|
||||
# 添加特殊标记,让 AI 知道这是上下文压缩的结果
|
||||
compressed_content = f"<context_compressed>\n{result_data['summary']}\n</context_compressed>"
|
||||
messages_history[:] = [
|
||||
user_msg,
|
||||
{"role": "assistant", "content": compressed_content}
|
||||
]
|
||||
# context_rewrite 压缩上下文后,不需要添加 tool 结果消息
|
||||
# 因为 messages_history 已经被重写为压缩后的状态
|
||||
continue
|
||||
|
||||
result = execute_tool(self._graph, tool_call.function.name, args)
|
||||
tool_calls.append({
|
||||
"name": tool_call.function.name,
|
||||
@ -324,8 +343,8 @@ class BackendServer:
|
||||
self.update_config(api_key, base_url, model)
|
||||
|
||||
limits_keys = [
|
||||
"persona_query_max", "persona_update_max",
|
||||
"task_query_max", "task_update_max",
|
||||
"persona_update_max",
|
||||
"task_update_max",
|
||||
"memory_query_max", "memory_update_max"
|
||||
]
|
||||
for key in limits_keys:
|
||||
|
||||
@ -50,6 +50,10 @@ def execute_tool(graph: Any, tool_name: str, arguments: dict) -> str:
|
||||
result = graph.cleanup(dry_run=arguments.get("dry_run", True))
|
||||
return json.dumps(result, ensure_ascii=False, default=str)
|
||||
|
||||
elif tool_name == "context_rewrite":
|
||||
result = execute_context_rewrite(graph, arguments)
|
||||
return json.dumps(result, ensure_ascii=False, default=str)
|
||||
|
||||
# 人设图管理工具
|
||||
elif tool_name == "persona_update":
|
||||
result = execute_persona_update(graph, arguments)
|
||||
@ -111,6 +115,24 @@ def format_recall_result(result: dict) -> str:
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def execute_context_rewrite(graph: Any, arguments: dict) -> dict:
|
||||
"""压缩工具调用上下文"""
|
||||
summary = arguments.get("summary", "")
|
||||
|
||||
# 验证格式:必须包含工具调用标记
|
||||
if "[工具调用总结" not in summary:
|
||||
return {
|
||||
"status": "error",
|
||||
"message": "总结格式错误:必须包含 [工具调用总结: 本次总结了 N 次工具调用 | 调用工具: ...] 标记"
|
||||
}
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"message": "上下文已压缩",
|
||||
"summary": summary
|
||||
}
|
||||
|
||||
|
||||
# 人设图管理工具实现
|
||||
def execute_persona_update(graph: Any, arguments: dict) -> dict:
|
||||
"""更新人设"""
|
||||
|
||||
@ -1,49 +1,35 @@
|
||||
"""
|
||||
工具调用限制器 - 限制每轮对话中各类工具的调用次数
|
||||
"""
|
||||
from typing import Dict, List, Optional
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolLimits:
|
||||
"""工具调用限制配置"""
|
||||
# 人设图限制
|
||||
persona_query_max: int = 1 # 每轮最多查询1次人设图
|
||||
persona_update_max: int = 1 # 每轮最多修改1次人设图
|
||||
|
||||
# 工作记忆链限制
|
||||
task_query_max: int = 4 # 每轮最多查询4次工作记忆链
|
||||
task_update_max: int = 5 # 每轮最多修改5次工作记忆链
|
||||
|
||||
# 一般记忆限制
|
||||
memory_query_max: int = 20 # 每轮最多查询20次一般记忆
|
||||
memory_update_max: int = 10 # 每轮最多修改10次一般记忆
|
||||
persona_update_max: int = 1
|
||||
task_update_max: int = 5
|
||||
memory_query_max: int = 20
|
||||
memory_update_max: int = 10
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolCallCount:
|
||||
"""工具调用计数"""
|
||||
# 人设图
|
||||
persona_query: int = 0
|
||||
persona_update: int = 0
|
||||
|
||||
# 工作记忆链
|
||||
task_query: int = 0
|
||||
task_update: int = 0
|
||||
|
||||
# 一般记忆
|
||||
memory_query: int = 0
|
||||
memory_update: int = 0
|
||||
|
||||
|
||||
class ToolLimiter:
|
||||
"""工具调用限制器"""
|
||||
|
||||
|
||||
def __init__(self, limits: Optional[ToolLimits] = None):
|
||||
self.limits = limits or ToolLimits()
|
||||
self.counts = ToolCallCount()
|
||||
|
||||
|
||||
def _classify_tool(self, tool_name: str, arguments: dict) -> tuple:
|
||||
"""
|
||||
分类工具调用
|
||||
@ -51,105 +37,83 @@ class ToolLimiter:
|
||||
category: 'persona', 'task', 'memory'
|
||||
operation: 'query', 'update'
|
||||
"""
|
||||
# 人设图工具
|
||||
if tool_name in ('persona_update', 'persona_clear'):
|
||||
return ('persona', 'update')
|
||||
|
||||
# 工作记忆链工具
|
||||
|
||||
if tool_name in ('task_create', 'task_set_state', 'task_delete', 'task_link_info'):
|
||||
# task_link_info 是关联操作,算作更新
|
||||
return ('task', 'update')
|
||||
|
||||
# 一般记忆工具
|
||||
|
||||
if tool_name == 'memory_recall':
|
||||
# 所有 memory_recall 统一归为一般记忆查询
|
||||
# 因为 query_intent 内容不可控,无法准确判断查询类型
|
||||
# 写入操作通过工具名称明确区分,不受此影响
|
||||
return ('memory', 'query')
|
||||
|
||||
|
||||
if tool_name == 'memory_commit':
|
||||
return ('memory', 'update')
|
||||
|
||||
|
||||
if tool_name == 'memory_purge':
|
||||
return ('memory', 'update')
|
||||
|
||||
|
||||
if tool_name == 'memory_introspect':
|
||||
return ('memory', 'query')
|
||||
|
||||
|
||||
if tool_name in ('memory_archive', 'memory_cleanup'):
|
||||
return ('memory', 'update')
|
||||
|
||||
# 未知工具,归类为一般记忆更新
|
||||
|
||||
if tool_name == 'context_rewrite':
|
||||
return ('memory', 'query')
|
||||
|
||||
return ('memory', 'update')
|
||||
|
||||
|
||||
def can_call(self, tool_name: str, arguments: dict) -> tuple:
|
||||
"""
|
||||
检查是否允许调用工具
|
||||
返回: (allowed, reason)
|
||||
"""
|
||||
category, operation = self._classify_tool(tool_name, arguments)
|
||||
|
||||
# 获取当前计数和限制
|
||||
|
||||
if category == 'persona':
|
||||
if operation == 'query':
|
||||
if self.counts.persona_query >= self.limits.persona_query_max:
|
||||
return (False, f"人设图查询次数已达上限({self.limits.persona_query_max}次)")
|
||||
else: # update
|
||||
if self.counts.persona_update >= self.limits.persona_update_max:
|
||||
return (False, f"人设图修改次数已达上限({self.limits.persona_update_max}次)")
|
||||
|
||||
if self.counts.persona_update >= self.limits.persona_update_max:
|
||||
return (False, f"人设图修改次数已达上限({self.limits.persona_update_max}次)")
|
||||
|
||||
elif category == 'task':
|
||||
if operation == 'query':
|
||||
if self.counts.task_query >= self.limits.task_query_max:
|
||||
return (False, f"工作记忆链查询次数已达上限({self.limits.task_query_max}次)")
|
||||
else: # update
|
||||
if self.counts.task_update >= self.limits.task_update_max:
|
||||
return (False, f"工作记忆链修改次数已达上限({self.limits.task_update_max}次)")
|
||||
|
||||
if self.counts.task_update >= self.limits.task_update_max:
|
||||
return (False, f"工作记忆链修改次数已达上限({self.limits.task_update_max}次)")
|
||||
|
||||
elif category == 'memory':
|
||||
if operation == 'query':
|
||||
if self.counts.memory_query >= self.limits.memory_query_max:
|
||||
return (False, f"一般记忆查询次数已达上限({self.limits.memory_query_max}次)")
|
||||
else: # update
|
||||
else:
|
||||
if self.counts.memory_update >= self.limits.memory_update_max:
|
||||
return (False, f"一般记忆修改次数已达上限({self.limits.memory_update_max}次)")
|
||||
|
||||
|
||||
return (True, "允许调用")
|
||||
|
||||
|
||||
def record_call(self, tool_name: str, arguments: dict) -> None:
|
||||
"""记录工具调用"""
|
||||
category, operation = self._classify_tool(tool_name, arguments)
|
||||
|
||||
|
||||
if category == 'persona':
|
||||
if operation == 'query':
|
||||
self.counts.persona_query += 1
|
||||
else:
|
||||
self.counts.persona_update += 1
|
||||
|
||||
self.counts.persona_update += 1
|
||||
|
||||
elif category == 'task':
|
||||
if operation == 'query':
|
||||
self.counts.task_query += 1
|
||||
else:
|
||||
self.counts.task_update += 1
|
||||
|
||||
self.counts.task_update += 1
|
||||
|
||||
elif category == 'memory':
|
||||
if operation == 'query':
|
||||
self.counts.memory_query += 1
|
||||
else:
|
||||
self.counts.memory_update += 1
|
||||
|
||||
|
||||
def get_summary(self) -> str:
|
||||
"""获取调用统计摘要"""
|
||||
lines = [
|
||||
f"人设图: 查询{self.counts.persona_query}/{self.limits.persona_query_max}次, "
|
||||
f"修改{self.counts.persona_update}/{self.limits.persona_update_max}次",
|
||||
f"工作记忆链: 查询{self.counts.task_query}/{self.limits.task_query_max}次, "
|
||||
f"修改{self.counts.task_update}/{self.limits.task_update_max}次",
|
||||
f"人设图: 修改{self.counts.persona_update}/{self.limits.persona_update_max}次",
|
||||
f"工作记忆链: 修改{self.counts.task_update}/{self.limits.task_update_max}次",
|
||||
f"一般记忆: 查询{self.counts.memory_query}/{self.limits.memory_query_max}次, "
|
||||
f"修改{self.counts.memory_update}/{self.limits.memory_update_max}次"
|
||||
]
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置计数(新的一轮对话开始时调用)"""
|
||||
self.counts = ToolCallCount()
|
||||
|
||||
@ -234,6 +234,43 @@ MEMORY_TOOLS = [
|
||||
"required": []
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "context_rewrite",
|
||||
"description": """压缩本轮对话的工具调用上下文。将冗长的JSON工具结果提炼为简洁摘要。
|
||||
|
||||
【使用场景】
|
||||
- 已执行多次工具调用,JSON细节已理解,不再需要原始格式
|
||||
- 但需保留"我调用了什么工具、得到了什么结论"的元认知
|
||||
- 继续携带原始JSON会干扰后续推理
|
||||
|
||||
【⚠️ 强制格式要求】
|
||||
1. 必须标注调用了哪些工具
|
||||
2. 必须标注是对几次工具调用的总结
|
||||
3. 必须保留关键语义信息
|
||||
|
||||
【示例】
|
||||
{
|
||||
"summary": "[工具调用总结: 本次总结了 2 次工具调用 | 调用工具: memory_recall, memory_recall]\\n\\n- 查询人设图:未找到人设,使用默认身份\\n- 查询工作记忆链:发现 Task_成语接龙,状态已暂停,当前成语为虎作伥"
|
||||
}
|
||||
|
||||
【注意事项】
|
||||
- 不可删除用户原始消息
|
||||
- 不可歪曲工具返回的关键事实
|
||||
- 仅在工具调用 ≥ 2 次后使用""",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"summary": {
|
||||
"type": "string",
|
||||
"description": "压缩后的摘要文本,必须包含工具调用元信息"
|
||||
}
|
||||
},
|
||||
"required": ["summary"]
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@ -472,12 +472,11 @@ asyncio.run(main())
|
||||
|
||||
| Category | Operation | Per-Turn Limit |
|
||||
|----------|-----------|---------------|
|
||||
| Persona graph | Query | 1 time |
|
||||
| Persona graph | Modify | 1 time |
|
||||
| Working memory chain | Query | 4 times |
|
||||
| Working memory chain | Modify | 2 times |
|
||||
| Working memory chain | Modify | 5 times |
|
||||
| General memory | Query | 20 times |
|
||||
| General memory | Modify | 10 times |
|
||||
| Context compression | Query | Counted as general memory query |
|
||||
|
||||
### Reset Mechanism
|
||||
|
||||
|
||||
@ -32,7 +32,7 @@ TrulyMEM-TrueHumanMEM/
|
||||
│ ├── services/ # Service layer (config only)
|
||||
│ ├── handlers/ # Event handlers
|
||||
│ └── styles/ # Style files
|
||||
└── tests/ # Tests (42 tests)
|
||||
└── tests/ # Tests (54 tests)
|
||||
```
|
||||
|
||||
## Architecture Diagram
|
||||
@ -146,13 +146,14 @@ def main():
|
||||
|
||||
## Tool System
|
||||
|
||||
### Memory Tools (6)
|
||||
### Memory Tools (7)
|
||||
- `memory_recall` - Retrieve memory
|
||||
- `memory_commit` - Write memory
|
||||
- `memory_purge` - Delete memory
|
||||
- `memory_introspect` - View status
|
||||
- `memory_archive` - Archive memory
|
||||
- `memory_cleanup` - Clean data
|
||||
- `context_rewrite` - Compress single-turn tool call context (experimental)
|
||||
|
||||
### Persona Tools (2)
|
||||
- `persona_update` - Update persona
|
||||
@ -166,6 +167,19 @@ def main():
|
||||
|
||||
---
|
||||
|
||||
## Tool Call Limits
|
||||
|
||||
| Category | Operation | Per-Turn Limit |
|
||||
|----------|-----------|---------------|
|
||||
| Persona graph | Modify | 1 time |
|
||||
| Working memory chain | Modify | 5 times |
|
||||
| General memory | Query | 20 times |
|
||||
| General memory | Modify | 10 times |
|
||||
|
||||
> Note: `memory_recall` is uniformly counted as general memory query, no longer distinguished by persona/working memory queries.
|
||||
|
||||
---
|
||||
|
||||
## Error Handling Principle
|
||||
|
||||
All APIs **do not throw exceptions**, errors are passed via return dictionary:
|
||||
|
||||
@ -25,6 +25,14 @@ All memory must be written to the graph database:
|
||||
All memory must be read from:
|
||||
- `memory_recall` - Retrieve memory
|
||||
|
||||
### Working Memory Management (Experimental)
|
||||
|
||||
`context_rewrite` allows AI to proactively compress tool call context within a single turn:
|
||||
- Distills verbose JSON tool results into concise natural language summaries
|
||||
- Summary must include which tools were called and how many calls are summarized
|
||||
- After system validates the format, replaces `messages_history` with `[user message, summary]`
|
||||
- Ensures LLM retains meta-cognition (knows "I called tools") while reducing JSON noise
|
||||
|
||||
---
|
||||
|
||||
## Mandatory Execution Flow (Per Turn)
|
||||
|
||||
@ -478,12 +478,11 @@ asyncio.run(main())
|
||||
|
||||
| 类别 | 操作 | 每轮上限 |
|
||||
|------|------|---------|
|
||||
| 人设图 | 查询 | 1 次 |
|
||||
| 人设图 | 修改 | 1 次 |
|
||||
| 工作记忆链 | 查询 | 4 次 |
|
||||
| 工作记忆链 | 修改 | 2 次 |
|
||||
| 工作记忆链 | 修改 | 5 次 |
|
||||
| 一般记忆 | 查询 | 20 次 |
|
||||
| 一般记忆 | 修改 | 10 次 |
|
||||
| 上下文压缩 | 查询 | 计入一般记忆查询 |
|
||||
|
||||
### 重置机制
|
||||
|
||||
|
||||
@ -32,7 +32,7 @@ TrulyMEM-TrueHumanMEM/
|
||||
│ ├── services/ # 服务层(仅配置管理)
|
||||
│ ├── handlers/ # 事件处理
|
||||
│ └── styles/ # 样式文件
|
||||
└── tests/ # 测试 (42 tests)
|
||||
└── tests/ # 测试 (54 tests)
|
||||
```
|
||||
|
||||
## 架构图
|
||||
@ -146,13 +146,14 @@ def main():
|
||||
|
||||
## 工具系统
|
||||
|
||||
### 记忆工具 (6个)
|
||||
### 记忆工具 (7个)
|
||||
- `memory_recall` - 检索记忆
|
||||
- `memory_commit` - 写入记忆
|
||||
- `memory_purge` - 删除记忆
|
||||
- `memory_introspect` - 查看状态
|
||||
- `memory_archive` - 归档记忆
|
||||
- `memory_cleanup` - 清理数据
|
||||
- `context_rewrite` - 压缩单轮工具调用上下文(实验性)
|
||||
|
||||
### 人设工具 (2个)
|
||||
- `persona_update` - 更新人设
|
||||
@ -166,6 +167,19 @@ def main():
|
||||
|
||||
---
|
||||
|
||||
## 工具调用限制
|
||||
|
||||
| 类别 | 操作 | 每轮上限 |
|
||||
|------|------|---------|
|
||||
| 人设图 | 修改 | 1 次 |
|
||||
| 工作记忆链 | 修改 | 5 次 |
|
||||
| 一般记忆 | 查询 | 20 次 |
|
||||
| 一般记忆 | 修改 | 10 次 |
|
||||
|
||||
> 注:`memory_recall` 统一计入一般记忆查询,不再区分人设/工作记忆查询。
|
||||
|
||||
---
|
||||
|
||||
## 错误处理原则
|
||||
|
||||
所有 API **不抛出异常**,错误通过返回字典传递:
|
||||
|
||||
@ -25,6 +25,14 @@ TrulyMEM 的解决思路:
|
||||
所有记忆必须通过以下方式读取:
|
||||
- `memory_recall` - 检索记忆
|
||||
|
||||
### 工作记忆管理(实验性)
|
||||
|
||||
`context_rewrite` 允许 AI 在单轮对话内主动压缩工具调用的临时上下文:
|
||||
- 将冗长的 JSON 工具结果提炼为简洁的自然语言摘要
|
||||
- 摘要必须包含调用了哪些工具、对几次调用的总结
|
||||
- 系统验证格式后,替换 `messages_history` 为 `[用户消息, 摘要]`
|
||||
- 确保 LLM 保留元认知(知道"我调用过工具"),同时减少 JSON 噪音
|
||||
|
||||
---
|
||||
|
||||
## 强制执行流程(每轮对话)
|
||||
|
||||
@ -1,54 +1 @@
|
||||
import pytest
|
||||
from datetime import datetime
|
||||
from ui.models.message import Message, ToolCall, ToolResult
|
||||
from ui.models.config import AppConfig
|
||||
from ui.models.log_entry import LogEntry
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_config():
|
||||
return AppConfig(
|
||||
api_key="test-api-key",
|
||||
model="test-model",
|
||||
base_url="https://test.api.com"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_message():
|
||||
return Message(
|
||||
role="user",
|
||||
content="测试消息",
|
||||
timestamp=datetime.now()
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_tool_call():
|
||||
return ToolCall(
|
||||
id="test-call-id",
|
||||
name="memory_recall",
|
||||
arguments={"query_intent": "测试查询"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_tool_result():
|
||||
return ToolResult(
|
||||
tool_call_id="test-call-id",
|
||||
name="memory_recall",
|
||||
arguments={"query_intent": "测试查询"},
|
||||
result="测试结果",
|
||||
success=True
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_log_entry():
|
||||
return LogEntry(
|
||||
timestamp=datetime.now(),
|
||||
tool_name="memory_recall",
|
||||
arguments={"query_intent": "测试查询"},
|
||||
result="测试结果",
|
||||
duration=0.5
|
||||
)
|
||||
218
tests/test_core/test_context_rewrite.py
Normal file
218
tests/test_core/test_context_rewrite.py
Normal file
@ -0,0 +1,218 @@
|
||||
"""context_rewrite 全面测试 - 单元 + 集成"""
|
||||
|
||||
import pytest
|
||||
import json
|
||||
import tempfile
|
||||
import os
|
||||
from core.tools.memory_tools import TOOLS, MEMORY_TOOLS
|
||||
from core.tool_executor import execute_tool, execute_context_rewrite
|
||||
from core.tool_limiter import ToolLimiter, ToolLimits
|
||||
from core import EmbeddedGraphDB
|
||||
|
||||
|
||||
# ========== 工具定义测试 ==========
|
||||
|
||||
class TestToolDefinition:
|
||||
def test_context_rewrite_in_tools(self):
|
||||
tool_names = [t["function"]["name"] for t in TOOLS]
|
||||
assert "context_rewrite" in tool_names
|
||||
|
||||
def test_context_rewrite_in_memory_tools(self):
|
||||
tool_names = [t["function"]["name"] for t in MEMORY_TOOLS]
|
||||
assert "context_rewrite" in tool_names
|
||||
|
||||
def test_context_rewrite_has_required_params(self):
|
||||
tool_def = None
|
||||
for t in MEMORY_TOOLS:
|
||||
if t["function"]["name"] == "context_rewrite":
|
||||
tool_def = t
|
||||
break
|
||||
assert tool_def is not None
|
||||
assert "summary" in tool_def["function"]["parameters"]["required"]
|
||||
|
||||
def test_context_rewrite_description_not_empty(self):
|
||||
for t in MEMORY_TOOLS:
|
||||
if t["function"]["name"] == "context_rewrite":
|
||||
assert len(t["function"]["description"]) > 100
|
||||
break
|
||||
|
||||
|
||||
# ========== 执行器测试 ==========
|
||||
|
||||
class TestContextRewriteExecutor:
|
||||
@pytest.fixture
|
||||
def db(self):
|
||||
with tempfile.NamedTemporaryFile(suffix='.db', delete=False) as f:
|
||||
db_path = f.name
|
||||
db = EmbeddedGraphDB(db_path)
|
||||
yield db
|
||||
db.close()
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_valid_summary(self, db):
|
||||
args = {"summary": "[工具调用总结: 本次总结了 2 次工具调用 | 调用工具: memory_recall, memory_recall]\n\n- 查询人设图:未找到"}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "success"
|
||||
assert result["message"] == "上下文已压缩"
|
||||
assert "memory_recall" in result["summary"]
|
||||
|
||||
def test_missing_marker(self, db):
|
||||
args = {"summary": "查询人设图:未找到"}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "error"
|
||||
assert "必须包含" in result["message"]
|
||||
|
||||
def test_empty_summary(self, db):
|
||||
args = {"summary": ""}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "error"
|
||||
|
||||
def test_marker_only(self, db):
|
||||
args = {"summary": "[工具调用总结"}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "success"
|
||||
|
||||
def test_via_execute_tool(self, db):
|
||||
args = {"summary": "[工具调用总结: 本次总结了 1 次工具调用 | 调用工具: memory_recall]\n\n- 查询记忆:找到 3 个实体"}
|
||||
result_str = execute_tool(db, "context_rewrite", args)
|
||||
result = json.loads(result_str)
|
||||
assert result["status"] == "success"
|
||||
|
||||
def test_unicode_content(self, db):
|
||||
args = {"summary": "[工具调用总结: 本次总结了 3 次工具调用 | 调用工具: memory_recall, memory_commit, task_create]\n\n- 查询:找到实体\"用户\"\n- 写入:{\"subject\": \"用户\", \"relation\": \"喜欢\"}\n- 任务:Task_测试"}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "success"
|
||||
assert "用户" in result["summary"]
|
||||
|
||||
def test_newlines_preserved(self, db):
|
||||
summary = "[工具调用总结: 本次总结了 2 次工具调用 | 调用工具: memory_recall, memory_recall]\n\n- 查询1:结果1\n- 查询2:结果2"
|
||||
args = {"summary": summary}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["summary"] == summary
|
||||
|
||||
def test_long_summary(self, db):
|
||||
summary = "[工具调用总结: 本次总结了 5 次工具调用 | 调用工具: memory_recall, memory_recall, memory_commit, task_create, task_set_state]\n\n" + "详细结果\n" * 50
|
||||
args = {"summary": summary}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "success"
|
||||
assert len(result["summary"]) == len(summary)
|
||||
|
||||
def test_special_json_chars(self, db):
|
||||
args = {"summary": '[工具调用总结: 本次总结了 1 次工具调用 | 调用工具: memory_commit]\n\n- 写入:{"subject": "测试", "relation": "包含\"引号"}'}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "success"
|
||||
|
||||
def test_missing_summary_key(self, db):
|
||||
args = {}
|
||||
result = execute_context_rewrite(db, args)
|
||||
assert result["status"] == "error"
|
||||
|
||||
|
||||
# ========== 工具限流器测试 ==========
|
||||
|
||||
class TestContextRewriteLimiter:
|
||||
def test_classified_as_memory_query(self):
|
||||
limiter = ToolLimiter(ToolLimits(memory_query_max=1))
|
||||
category, operation = limiter._classify_tool("context_rewrite", {})
|
||||
assert category == "memory"
|
||||
assert operation == "query"
|
||||
|
||||
def test_counts_toward_memory_query_limit(self):
|
||||
limiter = ToolLimiter(ToolLimits(memory_query_max=1))
|
||||
allowed, _ = limiter.can_call("context_rewrite", {})
|
||||
assert allowed
|
||||
limiter.record_call("context_rewrite", {})
|
||||
allowed, reason = limiter.can_call("context_rewrite", {})
|
||||
assert not allowed
|
||||
assert "一般记忆查询次数已达上限" in reason
|
||||
|
||||
def test_does_not_affect_memory_update(self):
|
||||
limiter = ToolLimiter(ToolLimits(memory_update_max=1))
|
||||
limiter.record_call("context_rewrite", {})
|
||||
allowed, _ = limiter.can_call("memory_commit", {})
|
||||
assert allowed
|
||||
|
||||
def test_reset_clears_count(self):
|
||||
limiter = ToolLimiter(ToolLimits(memory_query_max=1))
|
||||
limiter.record_call("context_rewrite", {})
|
||||
limiter.reset()
|
||||
allowed, _ = limiter.can_call("context_rewrite", {})
|
||||
assert allowed
|
||||
|
||||
|
||||
# ========== 集成测试:messages_history 压缩流程 ==========
|
||||
|
||||
class TestMessagesHistoryCompression:
|
||||
def test_compression_preserves_user_message(self):
|
||||
messages_history = [
|
||||
{"role": "user", "content": "我们之前聊过成语接龙吗?"},
|
||||
{"role": "assistant", "content": None, "tool_calls": [{"id": "tc1", "type": "function", "function": {"name": "memory_recall", "arguments": '{"query_intent": "人设"}'}}]},
|
||||
{"role": "tool", "tool_call_id": "tc1", "content": "===== 记忆检索结果 =====\n\n(未找到相关记忆)\n=============================="},
|
||||
{"role": "assistant", "content": None, "tool_calls": [{"id": "tc2", "type": "function", "function": {"name": "memory_recall", "arguments": '{"query_intent": "工作记忆"}'}}]},
|
||||
{"role": "tool", "tool_call_id": "tc2", "content": "===== 记忆检索结果 =====\n\n实体 (2 个):\n - Task_成语接龙 (类型: unknown, 提及: 1次)\n=============================="},
|
||||
]
|
||||
|
||||
summary = "[工具调用总结: 本次总结了 2 次工具调用 | 调用工具: memory_recall, memory_recall]\n\n- 查询人设图:未找到人设\n- 查询工作记忆链:发现 Task_成语接龙,状态已暂停"
|
||||
|
||||
user_msg = messages_history[0]
|
||||
messages_history[:] = [
|
||||
user_msg,
|
||||
{"role": "assistant", "content": summary}
|
||||
]
|
||||
|
||||
assert len(messages_history) == 2
|
||||
assert messages_history[0]["role"] == "user"
|
||||
assert messages_history[0]["content"] == "我们之前聊过成语接龙吗?"
|
||||
assert messages_history[1]["role"] == "assistant"
|
||||
assert "memory_recall" in messages_history[1]["content"]
|
||||
assert "成语接龙" in messages_history[1]["content"]
|
||||
|
||||
def test_compression_removes_json_noise(self):
|
||||
messages_history = [
|
||||
{"role": "user", "content": "查询用户信息"},
|
||||
{"role": "assistant", "content": None, "tool_calls": [{"id": "tc1", "type": "function", "function": {"name": "memory_recall", "arguments": '{"query_intent": "用户"}'}}]},
|
||||
{"role": "tool", "tool_call_id": "tc1", "content": json.dumps({"entities": [{"name": "用户", "type": "person", "mention_count": 5}], "relations": [{"source": "用户", "target": "Python", "type": "喜欢"}]})},
|
||||
]
|
||||
|
||||
summary = "[工具调用总结: 本次总结了 1 次工具调用 | 调用工具: memory_recall]\n\n- 查询用户:找到用户实体,提及5次,喜欢Python"
|
||||
|
||||
user_msg = messages_history[0]
|
||||
messages_history[:] = [user_msg, {"role": "assistant", "content": summary}]
|
||||
|
||||
for msg in messages_history[1:]:
|
||||
assert "entities" not in msg.get("content", "")
|
||||
assert "relations" not in msg.get("content", "")
|
||||
|
||||
def test_compression_retains_tool_meta_cognition(self):
|
||||
messages_history = [
|
||||
{"role": "user", "content": "查询"},
|
||||
{"role": "assistant", "content": None, "tool_calls": [{"id": "tc1", "type": "function", "function": {"name": "memory_recall", "arguments": "{}"}}]},
|
||||
{"role": "tool", "tool_call_id": "tc1", "content": "结果"},
|
||||
]
|
||||
|
||||
summary = "[工具调用总结: 本次总结了 1 次工具调用 | 调用工具: memory_recall]\n\n- 查询记忆:无结果"
|
||||
|
||||
user_msg = messages_history[0]
|
||||
messages_history[:] = [user_msg, {"role": "assistant", "content": summary}]
|
||||
|
||||
content = messages_history[1]["content"]
|
||||
assert "工具调用总结" in content
|
||||
assert "memory_recall" in content
|
||||
assert "1 次" in content
|
||||
|
||||
def test_multiple_compressions_in_sequence(self):
|
||||
messages_history = [
|
||||
{"role": "user", "content": "多轮查询"},
|
||||
]
|
||||
|
||||
for i in range(3):
|
||||
messages_history.append({"role": "assistant", "content": None, "tool_calls": [{"id": f"tc{i}", "type": "function", "function": {"name": "memory_recall", "arguments": "{}"}}]})
|
||||
messages_history.append({"role": "tool", "tool_call_id": f"tc{i}", "content": f"结果{i}"})
|
||||
|
||||
summary = f"[工具调用总结: 本次总结了 {i+1} 次工具调用 | 调用工具: memory_recall]\n\n- 第{i+1}轮查询:结果{i}"
|
||||
user_msg = messages_history[0]
|
||||
messages_history[:] = [user_msg, {"role": "assistant", "content": summary}]
|
||||
|
||||
assert len(messages_history) == 2
|
||||
assert messages_history[0]["content"] == "多轮查询"
|
||||
assert "第3轮查询" in messages_history[1]["content"]
|
||||
@ -9,60 +9,26 @@ os.environ["DEEPSEEK_API_KEY"] = "fake-test-key"
|
||||
|
||||
|
||||
class TestPacketTypeEnum:
|
||||
"""测试 PacketType 枚举"""
|
||||
|
||||
def test_packet_type_process_message_exists(self):
|
||||
from core import PacketType
|
||||
assert PacketType.PROCESS_MESSAGE is not None
|
||||
assert PacketType.PROCESS_MESSAGE.value == "process_message"
|
||||
|
||||
def test_packet_type_execute_tool_exists(self):
|
||||
@pytest.mark.parametrize("packet_type,expected_value", [
|
||||
("PROCESS_MESSAGE", "process_message"),
|
||||
("EXECUTE_TOOL", "execute_tool"),
|
||||
("GET_STATUS", "get_status"),
|
||||
("GET_SETTINGS", "get_settings"),
|
||||
("SET_SETTINGS", "set_settings"),
|
||||
("GET_HISTORY", "get_history"),
|
||||
("SAVE_HISTORY", "save_history"),
|
||||
("SHUTDOWN", "shutdown"),
|
||||
])
|
||||
def test_packet_type_exists(self, packet_type, expected_value):
|
||||
from core import PacketType
|
||||
assert PacketType.EXECUTE_TOOL is not None
|
||||
assert PacketType.EXECUTE_TOOL.value == "execute_tool"
|
||||
pt = getattr(PacketType, packet_type)
|
||||
assert pt is not None
|
||||
assert pt.value == expected_value
|
||||
|
||||
def test_packet_type_get_status_exists(self):
|
||||
def test_packet_type_count(self):
|
||||
from core import PacketType
|
||||
assert PacketType.GET_STATUS is not None
|
||||
assert PacketType.GET_STATUS.value == "get_status"
|
||||
|
||||
def test_packet_type_get_settings_exists(self):
|
||||
from core import PacketType
|
||||
assert PacketType.GET_SETTINGS is not None
|
||||
assert PacketType.GET_SETTINGS.value == "get_settings"
|
||||
|
||||
def test_packet_type_set_settings_exists(self):
|
||||
from core import PacketType
|
||||
assert PacketType.SET_SETTINGS is not None
|
||||
assert PacketType.SET_SETTINGS.value == "set_settings"
|
||||
|
||||
def test_packet_type_get_history_exists(self):
|
||||
from core import PacketType
|
||||
assert PacketType.GET_HISTORY is not None
|
||||
assert PacketType.GET_HISTORY.value == "get_history"
|
||||
|
||||
def test_packet_type_save_history_exists(self):
|
||||
from core import PacketType
|
||||
assert PacketType.SAVE_HISTORY is not None
|
||||
assert PacketType.SAVE_HISTORY.value == "save_history"
|
||||
|
||||
def test_packet_type_shutdown_exists(self):
|
||||
from core import PacketType
|
||||
assert PacketType.SHUTDOWN is not None
|
||||
assert PacketType.SHUTDOWN.value == "shutdown"
|
||||
|
||||
def test_packet_type_all_values(self):
|
||||
from core import PacketType
|
||||
values = [pt.value for pt in PacketType]
|
||||
assert "process_message" in values
|
||||
assert "execute_tool" in values
|
||||
assert "get_status" in values
|
||||
assert "get_settings" in values
|
||||
assert "set_settings" in values
|
||||
assert "get_history" in values
|
||||
assert "save_history" in values
|
||||
assert "shutdown" in values
|
||||
assert len(values) == 8
|
||||
assert len(list(PacketType)) == 8
|
||||
|
||||
|
||||
class TestPacketCreation:
|
||||
@ -192,7 +158,7 @@ class TestBackendClientAPI:
|
||||
|
||||
result = client.update_settings(
|
||||
api_config={"api_key": "new-key", "base_url": "https://api.deepseek.com", "model": "deepseek-chat"},
|
||||
tool_limits={"persona_query_max": 2}
|
||||
tool_limits={"persona_update_max": 2}
|
||||
)
|
||||
assert result.get("success") is True
|
||||
|
||||
@ -276,77 +242,3 @@ class TestBackendClientAPI:
|
||||
server.shutdown()
|
||||
|
||||
|
||||
class TestToolLimiter:
|
||||
"""测试工具限制器"""
|
||||
|
||||
def test_tool_limiter_init(self):
|
||||
from core.tool_limiter import ToolLimiter
|
||||
limiter = ToolLimiter()
|
||||
assert limiter.counts.persona_query == 0
|
||||
assert limiter.counts.persona_update == 0
|
||||
|
||||
def test_tool_limiter_classify(self):
|
||||
from core.tool_limiter import ToolLimiter
|
||||
limiter = ToolLimiter()
|
||||
|
||||
category, operation = limiter._classify_tool("memory_recall", {"query_intent": "test"})
|
||||
assert category == "memory"
|
||||
assert operation == "query"
|
||||
|
||||
category, operation = limiter._classify_tool("persona_update", {})
|
||||
assert category == "persona"
|
||||
assert operation == "update"
|
||||
|
||||
category, operation = limiter._classify_tool("task_create", {})
|
||||
assert category == "task"
|
||||
assert operation == "update"
|
||||
|
||||
def test_tool_limiter_can_call(self):
|
||||
from core.tool_limiter import ToolLimiter
|
||||
limiter = ToolLimiter()
|
||||
|
||||
allowed, reason = limiter.can_call("persona_update", {})
|
||||
assert allowed is True
|
||||
|
||||
limiter.record_call("persona_update", {})
|
||||
allowed, reason = limiter.can_call("persona_update", {})
|
||||
assert allowed is False
|
||||
assert "已达上限" in reason
|
||||
|
||||
def test_tool_limiter_reset(self):
|
||||
from core.tool_limiter import ToolLimiter
|
||||
limiter = ToolLimiter()
|
||||
|
||||
limiter.record_call("persona_update", {})
|
||||
assert limiter.counts.persona_update == 1
|
||||
|
||||
limiter.reset()
|
||||
assert limiter.counts.persona_update == 0
|
||||
|
||||
|
||||
class TestEmbeddedGraphDB:
|
||||
"""测试图数据库"""
|
||||
|
||||
def test_embedded_db_init(self):
|
||||
from core.embedded_db import EmbeddedGraphDB
|
||||
db = EmbeddedGraphDB(db_path=":memory:")
|
||||
assert db.conn is not None
|
||||
db.close()
|
||||
|
||||
def test_embedded_db_commit_and_recall(self):
|
||||
from core.embedded_db import EmbeddedGraphDB
|
||||
db = EmbeddedGraphDB(db_path=":memory:")
|
||||
|
||||
# 写入记忆 (使用 triplets 参数)
|
||||
result = db.commit(
|
||||
triplets=[
|
||||
{"subject": "测试", "relation": "是", "object": "test"}
|
||||
],
|
||||
session_id="test-session"
|
||||
)
|
||||
|
||||
# 读取记忆
|
||||
results = db.recall("测试")
|
||||
assert len(results.get("entities", [])) > 0
|
||||
|
||||
db.close()
|
||||
@ -1,179 +1 @@
|
||||
import pytest
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
os.environ["DEEPSEEK_API_KEY"] = "fake-test-key"
|
||||
|
||||
|
||||
class TestIntegrationPacketFlow:
|
||||
"""测试 Packet 通信流程"""
|
||||
|
||||
def test_packet_round_trip_process_message(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
# 无 API key 时应该返回错误而非抛异常
|
||||
result = client.process_message("test message")
|
||||
assert result.get("success") is False
|
||||
assert "error" in result
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_packet_round_trip_config(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.update_config(api_key="test-api", base_url="https://test.com")
|
||||
assert result.get("success") is True
|
||||
|
||||
status = client.get_status()
|
||||
data = status.get("data", {})
|
||||
assert data.get("config", {}).get("api_key") == "test-api"
|
||||
assert data.get("config", {}).get("base_url") == "https://test.com"
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_packet_round_trip_status(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.get_status()
|
||||
assert result.get("success") is True
|
||||
data = result.get("data", {})
|
||||
assert data.get("running") is True
|
||||
assert data.get("graph_initialized") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_packet_round_trip_execute_tool(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.execute_tool("memory_introspect", {})
|
||||
assert result.get("success") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
|
||||
class TestIntegrationToolLimiter:
|
||||
"""测试工具限制器集成"""
|
||||
|
||||
def test_external_tool_call_not_limited(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
# 外部调用多次应该成功
|
||||
for i in range(5):
|
||||
result = client.execute_tool("memory_introspect", {})
|
||||
assert result.get("success") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_internal_tool_call_limited(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="fake-key") # 假 key 会失败但不影响测试
|
||||
client = BackendClient(server)
|
||||
|
||||
# 内部调用受限,tool_limiter 存在
|
||||
assert server._tool_limiter is not None
|
||||
|
||||
# 初始状态
|
||||
assert server._tool_limiter.counts.persona_update == 0
|
||||
|
||||
# 记录一次调用
|
||||
server._tool_limiter.record_call("persona_update", {})
|
||||
assert server._tool_limiter.counts.persona_update == 1
|
||||
|
||||
# 再次调用应该被拒绝
|
||||
allowed, reason = server._tool_limiter.can_call("persona_update", {})
|
||||
assert allowed is False
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
|
||||
class TestIntegrationErrorHandling:
|
||||
"""测试错误处理"""
|
||||
|
||||
def test_process_message_returns_error_not_raise(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
# 应该返回错误,而不是抛出异常
|
||||
result = client.process_message("hello")
|
||||
assert result.get("success") is False
|
||||
assert "error" in result
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_execute_tool_error_handling(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
# 不存在的工具应该返回错误
|
||||
result = client.execute_tool("nonexistent_tool", {})
|
||||
assert result.get("success") is False
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
"""Tests for integration layer"""
|
||||
|
||||
164
tests/test_integration/test_integration.py
Normal file
164
tests/test_integration/test_integration.py
Normal file
@ -0,0 +1,164 @@
|
||||
"""Integration tests - Packet flow, tool limiter, and error handling across layers."""
|
||||
import pytest
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
os.environ["DEEPSEEK_API_KEY"] = "fake-test-key"
|
||||
|
||||
|
||||
class TestIntegrationPacketFlow:
|
||||
"""测试 Packet 通信流程"""
|
||||
|
||||
def test_packet_round_trip_process_message(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.process_message("test message")
|
||||
assert result.get("success") is False
|
||||
assert "error" in result
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_packet_round_trip_config(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.update_settings(
|
||||
api_config={"api_key": "test-api", "base_url": "https://test.com"},
|
||||
)
|
||||
assert result.get("success") is True
|
||||
|
||||
settings = client.get_settings()
|
||||
data = settings.get("data", {})
|
||||
assert data.get("api_config", {}).get("api_key") == "test-api"
|
||||
assert data.get("api_config", {}).get("base_url") == "https://test.com"
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_packet_round_trip_status(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.get_status()
|
||||
assert result.get("success") is True
|
||||
data = result.get("data", {})
|
||||
assert data.get("running") is True
|
||||
assert data.get("graph_initialized") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_packet_round_trip_execute_tool(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.execute_tool("memory_introspect", {})
|
||||
assert result.get("success") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
|
||||
class TestIntegrationToolLimiter:
|
||||
"""测试工具限制器集成"""
|
||||
|
||||
def test_external_tool_call_not_limited(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
for i in range(5):
|
||||
result = client.execute_tool("memory_introspect", {})
|
||||
assert result.get("success") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_internal_tool_call_limited(self):
|
||||
from core.tool_limiter import ToolLimiter, ToolLimits
|
||||
limiter = ToolLimiter(ToolLimits(persona_update_max=1))
|
||||
|
||||
assert limiter.counts.persona_update == 0
|
||||
|
||||
limiter.record_call("persona_update", {})
|
||||
assert limiter.counts.persona_update == 1
|
||||
|
||||
allowed, reason = limiter.can_call("persona_update", {})
|
||||
assert allowed is False
|
||||
assert "已达上限" in reason
|
||||
|
||||
|
||||
class TestIntegrationErrorHandling:
|
||||
"""测试错误处理"""
|
||||
|
||||
def test_process_message_returns_error_not_raise(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.process_message("hello")
|
||||
assert result.get("success") is False
|
||||
assert "error" in result
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_execute_tool_error_handling(self):
|
||||
from core import BackendServer, BackendClient
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path, use_embedded_db=True)
|
||||
server.start(api_key="")
|
||||
client = BackendClient(server)
|
||||
|
||||
result = client.execute_tool("nonexistent_tool", {})
|
||||
data = result.get("data", {})
|
||||
assert "未知工具" in data.get("result", "")
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
@ -1,241 +1 @@
|
||||
import pytest
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
os.environ["DEEPSEEK_API_KEY"] = "fake-test-key"
|
||||
|
||||
|
||||
class TestUIImport:
|
||||
"""测试 UI 模块导入"""
|
||||
|
||||
def test_import_graphmemoryapp(self):
|
||||
from ui import GraphMemoryApp
|
||||
assert GraphMemoryApp is not None
|
||||
|
||||
def test_import_appconfig(self):
|
||||
from ui import AppConfig
|
||||
assert AppConfig is not None
|
||||
|
||||
def test_import_message(self):
|
||||
from ui.models.message import Message, ToolCall, ToolResult
|
||||
assert Message is not None
|
||||
assert ToolCall is not None
|
||||
assert ToolResult is not None
|
||||
|
||||
def test_import_config(self):
|
||||
from ui.models.config import AppConfig
|
||||
assert AppConfig is not None
|
||||
|
||||
def test_import_log_entry(self):
|
||||
from ui.models.log_entry import LogEntry
|
||||
assert LogEntry is not None
|
||||
|
||||
|
||||
class TestAppConfig:
|
||||
"""测试配置模型"""
|
||||
|
||||
def test_config_default_values(self):
|
||||
from ui.models.config import AppConfig
|
||||
config = AppConfig()
|
||||
assert config.api_key == ""
|
||||
assert config.model == "deepseek-chat"
|
||||
assert config.base_url == "https://api.deepseek.com"
|
||||
|
||||
def test_config_from_env(self):
|
||||
from ui.models.config import AppConfig
|
||||
config = AppConfig.from_env()
|
||||
assert "fake-test-key" in config.api_key
|
||||
|
||||
|
||||
class TestMessageModel:
|
||||
"""测试消息模型"""
|
||||
|
||||
def test_message_creation_user(self):
|
||||
from ui.models.message import Message
|
||||
from datetime import datetime
|
||||
msg = Message(role="user", content="test content")
|
||||
assert msg.role == "user"
|
||||
assert msg.content == "test content"
|
||||
assert isinstance(msg.timestamp, datetime)
|
||||
|
||||
def test_message_creation_assistant(self):
|
||||
from ui.models.message import Message
|
||||
msg = Message(role="assistant", content="assistant response")
|
||||
assert msg.role == "assistant"
|
||||
|
||||
def test_message_with_tool_calls(self):
|
||||
from ui.models.message import Message, ToolCall
|
||||
tc = ToolCall(id="call-1", name="memory_recall", arguments={"query": "test"})
|
||||
msg = Message(role="assistant", content="response", tool_calls=[tc])
|
||||
assert msg.tool_calls is not None
|
||||
assert len(msg.tool_calls) == 1
|
||||
|
||||
|
||||
class TestAppCSSPath:
|
||||
"""测试 App CSS 配置"""
|
||||
|
||||
def test_app_has_css_path(self):
|
||||
from ui import GraphMemoryApp
|
||||
assert hasattr(GraphMemoryApp, 'CSS_PATH')
|
||||
assert len(GraphMemoryApp.CSS_PATH) > 0
|
||||
|
||||
|
||||
class TestAppBindings:
|
||||
"""测试 App 快捷键"""
|
||||
|
||||
def test_app_has_bindings(self):
|
||||
from ui import GraphMemoryApp
|
||||
assert hasattr(GraphMemoryApp, 'BINDINGS')
|
||||
assert len(GraphMemoryApp.BINDINGS) > 0
|
||||
|
||||
|
||||
class TestWidgetImports:
|
||||
"""测试组件导入"""
|
||||
|
||||
def test_import_left_panel(self):
|
||||
from ui.widgets.left_panel import LeftPanel
|
||||
assert LeftPanel is not None
|
||||
|
||||
def test_import_right_panel(self):
|
||||
from ui.widgets.right_panel import RightPanel
|
||||
assert RightPanel is not None
|
||||
|
||||
def test_import_input_box(self):
|
||||
from ui.widgets.input_box import InputBox
|
||||
assert InputBox is not None
|
||||
|
||||
def test_import_message_history(self):
|
||||
from ui.widgets.message_history import MessageHistory
|
||||
assert MessageHistory is not None
|
||||
|
||||
def test_import_status_bar(self):
|
||||
from ui.widgets.status_bar import StatusBar
|
||||
assert StatusBar is not None
|
||||
|
||||
|
||||
class TestHandlerImports:
|
||||
"""测试处理器导入"""
|
||||
|
||||
def test_import_focus_handler(self):
|
||||
from ui.handlers.focus_handler import FocusHandler
|
||||
assert FocusHandler is not None
|
||||
|
||||
def test_import_key_handler(self):
|
||||
from ui.handlers.key_handler import KeyHandler
|
||||
assert KeyHandler is not None
|
||||
|
||||
|
||||
class TestServiceImports:
|
||||
"""测试服务导入"""
|
||||
|
||||
def test_import_config_service(self):
|
||||
from ui.services.config_service import ConfigService
|
||||
assert ConfigService is not None
|
||||
|
||||
def test_import_config_manager(self):
|
||||
from ui.services.config_manager import ConfigManager
|
||||
assert ConfigManager is not None
|
||||
|
||||
|
||||
class TestAppInitialization:
|
||||
"""测试 App 初始化"""
|
||||
|
||||
def test_app_without_backend(self):
|
||||
from ui import GraphMemoryApp
|
||||
app = GraphMemoryApp()
|
||||
assert app._backend_server is None
|
||||
assert app._backend_client is None
|
||||
|
||||
def test_app_with_backend(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
assert app._backend_server is server
|
||||
assert app._backend_client is not None
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
|
||||
class TestUIWithBackendClient:
|
||||
"""测试 UI 与后端通信"""
|
||||
|
||||
def test_app_sends_message_via_backend_client(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer, BackendClient
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
client = app._backend_client
|
||||
|
||||
status = client.get_status()
|
||||
assert status.get("success") is True
|
||||
|
||||
result = client.update_settings(
|
||||
api_config={"api_key": "sk-test", "base_url": "https://api.deepseek.com", "model": "deepseek-chat"},
|
||||
tool_limits={"persona_query_max": 1}
|
||||
)
|
||||
assert result.get("success") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_ui_get_history(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
client = app._backend_client
|
||||
|
||||
history = client.get_history()
|
||||
assert isinstance(history, list)
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_ui_only_uses_backend_client(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
|
||||
# UI 不应该直接访问后端内部
|
||||
assert hasattr(app, '_backend_client')
|
||||
assert app._backend_client is not None
|
||||
|
||||
# 不应该有 _graph, _client 等直接访问
|
||||
assert not hasattr(app, '_graph')
|
||||
assert not hasattr(app, '_client')
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
"""Tests for UI layer"""
|
||||
|
||||
242
tests/test_ui/test_ui.py
Normal file
242
tests/test_ui/test_ui.py
Normal file
@ -0,0 +1,242 @@
|
||||
"""Tests for UI layer - models, widgets, handlers, services, and app initialization."""
|
||||
import pytest
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
os.environ["DEEPSEEK_API_KEY"] = "fake-test-key"
|
||||
|
||||
|
||||
class TestUIImport:
|
||||
"""测试 UI 模块导入"""
|
||||
|
||||
def test_import_graphmemoryapp(self):
|
||||
from ui import GraphMemoryApp
|
||||
assert GraphMemoryApp is not None
|
||||
|
||||
def test_import_appconfig(self):
|
||||
from ui import AppConfig
|
||||
assert AppConfig is not None
|
||||
|
||||
def test_import_message(self):
|
||||
from ui.models.message import Message, ToolCall, ToolResult
|
||||
assert Message is not None
|
||||
assert ToolCall is not None
|
||||
assert ToolResult is not None
|
||||
|
||||
def test_import_config(self):
|
||||
from ui.models.config import AppConfig
|
||||
assert AppConfig is not None
|
||||
|
||||
def test_import_log_entry(self):
|
||||
from ui.models.log_entry import LogEntry
|
||||
assert LogEntry is not None
|
||||
|
||||
|
||||
class TestAppConfig:
|
||||
"""测试配置模型"""
|
||||
|
||||
def test_config_default_values(self):
|
||||
from ui.models.config import AppConfig
|
||||
config = AppConfig()
|
||||
assert config.api_key == ""
|
||||
assert config.model == "deepseek-chat"
|
||||
assert config.base_url == "https://api.deepseek.com"
|
||||
|
||||
def test_config_from_env(self):
|
||||
from ui.models.config import AppConfig
|
||||
config = AppConfig.from_env()
|
||||
assert "fake-test-key" in config.api_key
|
||||
|
||||
|
||||
class TestMessageModel:
|
||||
"""测试消息模型"""
|
||||
|
||||
def test_message_creation_user(self):
|
||||
from ui.models.message import Message
|
||||
from datetime import datetime
|
||||
msg = Message(role="user", content="test content")
|
||||
assert msg.role == "user"
|
||||
assert msg.content == "test content"
|
||||
assert isinstance(msg.timestamp, datetime)
|
||||
|
||||
def test_message_creation_assistant(self):
|
||||
from ui.models.message import Message
|
||||
msg = Message(role="assistant", content="assistant response")
|
||||
assert msg.role == "assistant"
|
||||
|
||||
def test_message_with_tool_calls(self):
|
||||
from ui.models.message import Message, ToolCall
|
||||
tc = ToolCall(id="call-1", name="memory_recall", arguments={"query": "test"})
|
||||
msg = Message(role="assistant", content="response", tool_calls=[tc])
|
||||
assert msg.tool_calls is not None
|
||||
assert len(msg.tool_calls) == 1
|
||||
|
||||
|
||||
class TestAppCSSPath:
|
||||
"""测试 App CSS 配置"""
|
||||
|
||||
def test_app_has_css_path(self):
|
||||
from ui import GraphMemoryApp
|
||||
assert hasattr(GraphMemoryApp, 'CSS_PATH')
|
||||
assert len(GraphMemoryApp.CSS_PATH) > 0
|
||||
|
||||
|
||||
class TestAppBindings:
|
||||
"""测试 App 快捷键"""
|
||||
|
||||
def test_app_has_bindings(self):
|
||||
from ui import GraphMemoryApp
|
||||
assert hasattr(GraphMemoryApp, 'BINDINGS')
|
||||
assert len(GraphMemoryApp.BINDINGS) > 0
|
||||
|
||||
|
||||
class TestWidgetImports:
|
||||
"""测试组件导入"""
|
||||
|
||||
def test_import_left_panel(self):
|
||||
from ui.widgets.left_panel import LeftPanel
|
||||
assert LeftPanel is not None
|
||||
|
||||
def test_import_right_panel(self):
|
||||
from ui.widgets.right_panel import RightPanel
|
||||
assert RightPanel is not None
|
||||
|
||||
def test_import_input_box(self):
|
||||
from ui.widgets.input_box import InputBox
|
||||
assert InputBox is not None
|
||||
|
||||
def test_import_message_history(self):
|
||||
from ui.widgets.message_history import MessageHistory
|
||||
assert MessageHistory is not None
|
||||
|
||||
def test_import_status_bar(self):
|
||||
from ui.widgets.status_bar import StatusBar
|
||||
assert StatusBar is not None
|
||||
|
||||
|
||||
class TestHandlerImports:
|
||||
"""测试处理器导入"""
|
||||
|
||||
def test_import_focus_handler(self):
|
||||
from ui.handlers.focus_handler import FocusHandler
|
||||
assert FocusHandler is not None
|
||||
|
||||
def test_import_key_handler(self):
|
||||
from ui.handlers.key_handler import KeyHandler
|
||||
assert KeyHandler is not None
|
||||
|
||||
|
||||
class TestServiceImports:
|
||||
"""测试服务导入"""
|
||||
|
||||
def test_import_config_service(self):
|
||||
from ui.services.config_service import ConfigService
|
||||
assert ConfigService is not None
|
||||
|
||||
def test_import_config_manager(self):
|
||||
from ui.services.config_manager import ConfigManager
|
||||
assert ConfigManager is not None
|
||||
|
||||
|
||||
class TestAppInitialization:
|
||||
"""测试 App 初始化"""
|
||||
|
||||
def test_app_without_backend(self):
|
||||
from ui import GraphMemoryApp
|
||||
app = GraphMemoryApp()
|
||||
assert app._backend_server is None
|
||||
assert app._backend_client is None
|
||||
|
||||
def test_app_with_backend(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
assert app._backend_server is server
|
||||
assert app._backend_client is not None
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
|
||||
class TestUIWithBackendClient:
|
||||
"""测试 UI 与后端通信"""
|
||||
|
||||
def test_app_sends_message_via_backend_client(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer, BackendClient
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
client = app._backend_client
|
||||
|
||||
status = client.get_status()
|
||||
assert status.get("success") is True
|
||||
|
||||
result = client.update_settings(
|
||||
api_config={"api_key": "sk-test", "base_url": "https://api.deepseek.com", "model": "deepseek-chat"},
|
||||
tool_limits={"persona_update_max": 1}
|
||||
)
|
||||
assert result.get("success") is True
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_ui_get_history(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
client = app._backend_client
|
||||
|
||||
history = client.get_history()
|
||||
assert isinstance(history, list)
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
|
||||
def test_ui_only_uses_backend_client(self):
|
||||
from ui import GraphMemoryApp
|
||||
from core import BackendServer
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f:
|
||||
db_path = f.name
|
||||
try:
|
||||
server = BackendServer(db_path=db_path)
|
||||
server.start(api_key="")
|
||||
|
||||
app = GraphMemoryApp(backend_server=server)
|
||||
|
||||
# UI 不应该直接访问后端内部
|
||||
assert hasattr(app, '_backend_client')
|
||||
assert app._backend_client is not None
|
||||
|
||||
# 不应该有 _graph, _client 等直接访问
|
||||
assert not hasattr(app, '_graph')
|
||||
assert not hasattr(app, '_client')
|
||||
|
||||
server.shutdown()
|
||||
finally:
|
||||
if os.path.exists(db_path):
|
||||
os.unlink(db_path)
|
||||
10
ui/app.py
10
ui/app.py
@ -46,10 +46,8 @@ class GraphMemoryApp(App):
|
||||
initial_config.model = api_config.get("model", "deepseek-chat")
|
||||
|
||||
tool_limits = settings_data.get("tool_limits", {})
|
||||
initial_config.persona_query_max = tool_limits.get("persona_query_max", 1)
|
||||
initial_config.persona_update_max = tool_limits.get("persona_update_max", 1)
|
||||
initial_config.task_query_max = tool_limits.get("task_query_max", 4)
|
||||
initial_config.task_update_max = tool_limits.get("task_update_max", 2)
|
||||
initial_config.task_update_max = tool_limits.get("task_update_max", 5)
|
||||
initial_config.memory_query_max = tool_limits.get("memory_query_max", 20)
|
||||
initial_config.memory_update_max = tool_limits.get("memory_update_max", 10)
|
||||
|
||||
@ -222,9 +220,7 @@ class GraphMemoryApp(App):
|
||||
}
|
||||
|
||||
tool_limits = {
|
||||
"persona_query_max": config.persona_query_max,
|
||||
"persona_update_max": config.persona_update_max,
|
||||
"task_query_max": config.task_query_max,
|
||||
"task_update_max": config.task_update_max,
|
||||
"memory_query_max": config.memory_query_max,
|
||||
"memory_update_max": config.memory_update_max,
|
||||
@ -257,10 +253,8 @@ class GraphMemoryApp(App):
|
||||
api_key=api_cfg.get("api_key", ""),
|
||||
base_url=api_cfg.get("base_url", "https://api.deepseek.com"),
|
||||
model=api_cfg.get("model", "deepseek-chat"),
|
||||
persona_query_max=tool_lmts.get("persona_query_max", 1),
|
||||
persona_update_max=tool_lmts.get("persona_update_max", 1),
|
||||
task_query_max=tool_lmts.get("task_query_max", 4),
|
||||
task_update_max=tool_lmts.get("task_update_max", 2),
|
||||
task_update_max=tool_lmts.get("task_update_max", 5),
|
||||
memory_query_max=tool_lmts.get("memory_query_max", 20),
|
||||
memory_update_max=tool_lmts.get("memory_update_max", 10),
|
||||
))
|
||||
|
||||
@ -13,10 +13,8 @@ class AppConfig:
|
||||
api_key: str = ""
|
||||
model: str = "deepseek-chat"
|
||||
base_url: str = "https://api.deepseek.com"
|
||||
persona_query_max: int = 1
|
||||
persona_update_max: int = 1
|
||||
task_query_max: int = 4
|
||||
task_update_max: int = 2
|
||||
task_update_max: int = 5
|
||||
memory_query_max: int = 20
|
||||
memory_update_max: int = 10
|
||||
|
||||
@ -26,10 +24,8 @@ class AppConfig:
|
||||
api_key=os.getenv("DEEPSEEK_API_KEY", ""),
|
||||
model=os.getenv("MODEL_NAME", "deepseek-chat"),
|
||||
base_url=os.getenv("DEEPSEEK_BASE_URL", "https://api.deepseek.com"),
|
||||
persona_query_max=int(os.getenv("PERSONA_QUERY_MAX", 1)),
|
||||
persona_update_max=int(os.getenv("PERSONA_UPDATE_MAX", 1)),
|
||||
task_query_max=int(os.getenv("TASK_QUERY_MAX", 4)),
|
||||
task_update_max=int(os.getenv("TASK_UPDATE_MAX", 2)),
|
||||
task_update_max=int(os.getenv("TASK_UPDATE_MAX", 5)),
|
||||
memory_query_max=int(os.getenv("MEMORY_QUERY_MAX", 20)),
|
||||
memory_update_max=int(os.getenv("MEMORY_UPDATE_MAX", 10)),
|
||||
)
|
||||
@ -46,10 +42,8 @@ class AppConfig:
|
||||
api_key=data.get("api_key", ""),
|
||||
model=data.get("model", "deepseek-chat"),
|
||||
base_url=data.get("base_url", "https://api.deepseek.com"),
|
||||
persona_query_max=data.get("persona_query_max", 1),
|
||||
persona_update_max=data.get("persona_update_max", 1),
|
||||
task_query_max=data.get("task_query_max", 4),
|
||||
task_update_max=data.get("task_update_max", 2),
|
||||
task_update_max=data.get("task_update_max", 5),
|
||||
memory_query_max=data.get("memory_query_max", 20),
|
||||
memory_update_max=data.get("memory_update_max", 10),
|
||||
)
|
||||
|
||||
@ -63,25 +63,15 @@ class ConfigSection(Vertical):
|
||||
limits_title.can_focus = False
|
||||
yield limits_title
|
||||
|
||||
l1 = Static("人设图查询:", classes="config-label")
|
||||
l1.can_focus = False
|
||||
yield l1
|
||||
yield Input(value=str(self._config.persona_query_max), placeholder="1", id="persona-query-max")
|
||||
|
||||
l2 = Static("人设图修改:", classes="config-label")
|
||||
l2.can_focus = False
|
||||
yield l2
|
||||
yield Input(value=str(self._config.persona_update_max), placeholder="1", id="persona-update-max")
|
||||
|
||||
l3 = Static("工作记忆查询:", classes="config-label")
|
||||
l3.can_focus = False
|
||||
yield l3
|
||||
yield Input(value=str(self._config.task_query_max), placeholder="4", id="task-query-max")
|
||||
|
||||
l4 = Static("工作记忆修改:", classes="config-label")
|
||||
l4.can_focus = False
|
||||
yield l4
|
||||
yield Input(value=str(self._config.task_update_max), placeholder="2", id="task-update-max")
|
||||
yield Input(value=str(self._config.task_update_max), placeholder="5", id="task-update-max")
|
||||
|
||||
l5 = Static("一般记忆查询:", classes="config-label")
|
||||
l5.can_focus = False
|
||||
@ -122,9 +112,7 @@ class ConfigSection(Vertical):
|
||||
model_input = self.query_one("#model-input", Input)
|
||||
base_url_input = self.query_one("#base-url-input", Input)
|
||||
|
||||
persona_query = self.query_one("#persona-query-max", Input)
|
||||
persona_update = self.query_one("#persona-update-max", Input)
|
||||
task_query = self.query_one("#task-query-max", Input)
|
||||
task_update = self.query_one("#task-update-max", Input)
|
||||
memory_query = self.query_one("#memory-query-max", Input)
|
||||
memory_update = self.query_one("#memory-update-max", Input)
|
||||
@ -133,10 +121,8 @@ class ConfigSection(Vertical):
|
||||
api_key=api_key_input.value,
|
||||
model=model_input.value,
|
||||
base_url=base_url_input.value,
|
||||
persona_query_max=int(persona_query.value or 1),
|
||||
persona_update_max=int(persona_update.value or 1),
|
||||
task_query_max=int(task_query.value or 4),
|
||||
task_update_max=int(task_update.value or 2),
|
||||
task_update_max=int(task_update.value or 5),
|
||||
memory_query_max=int(memory_query.value or 20),
|
||||
memory_update_max=int(memory_update.value or 10),
|
||||
)
|
||||
@ -162,9 +148,7 @@ class ConfigSection(Vertical):
|
||||
model_input.value = config.model
|
||||
base_url_input.value = config.base_url
|
||||
|
||||
self.query_one("#persona-query-max", Input).value = str(config.persona_query_max)
|
||||
self.query_one("#persona-update-max", Input).value = str(config.persona_update_max)
|
||||
self.query_one("#task-query-max", Input).value = str(config.task_query_max)
|
||||
self.query_one("#task-update-max", Input).value = str(config.task_update_max)
|
||||
self.query_one("#memory-query-max", Input).value = str(config.memory_query_max)
|
||||
self.query_one("#memory-update-max", Input).value = str(config.memory_update_max)
|
||||
|
||||
Reference in New Issue
Block a user