Files
TrulyMEM-TrueHumanMEM-local/core/tool_limiter.py
root 3de418b7e8 refactor: 工具调用限制全部移入 config.json,代码零硬编码
- 删除 server.py BackendServer.__init__ 中 _tool_limits 的硬编码字典
- _load_config 从 ~/.trulymem/config.json 读取全部限制值
- 首次启动无配置文件时自动生成带默认值的 config.json
- _create_tool_limiter 直接索引 _tool_limits 字典,无硬编码兜底
- 如需修改限制值,编辑 ~/.trulymem/config.json 即可生效(下次重启)
2026-04-28 20:30:24 +08:00

141 lines
5.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
工具调用限制器 - 限制每轮对话中各类工具的调用次数
"""
from typing import Optional
from dataclasses import dataclass
@dataclass
class ToolLimits:
"""工具调用限制配置
实际值由 server.py 从 config.json 加载后传入,此处默认值仅作安全兜底。
如需修改限制,请编辑 ~/.trulymem/config.json。
"""
persona_update_max: int = 1
task_update_max: int = 20 # 工作记忆链修改create/set_state/delete/link_info
task_query_max: int = 30 # 工作记忆链查询memory_recall 查任务相关)
memory_query_max: int = 30
memory_update_max: int = 15
@dataclass
class ToolCallCount:
"""工具调用计数"""
persona_update: int = 0
task_update: int = 0
task_query: 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:
"""
分类工具调用
返回: (category, operation)
category: 'persona', 'task', 'memory'
operation: 'query', 'update'
"""
if tool_name in ('persona_update', 'persona_remove', 'persona_clear'):
return ('persona', 'update')
if tool_name in ('task_create', 'task_set_state', 'task_delete', 'task_link_info', 'task_archive'):
return ('task', 'update')
if tool_name == 'task_query':
return ('task', 'query')
if tool_name == 'memory_recall':
# 尝试区分工作记忆链查询 vs 一般记忆查询
query = (arguments.get('queryIntent', '') + ' ' + ' '.join(
arguments.get('seedEntities', []))).strip().lower()
task_keywords = ['task', '任务', '工作记忆', '当前轮', '会话', '过程', '流程']
if any(kw in query for kw in task_keywords):
return ('task', 'query')
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 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}次)")
elif 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:
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':
self.counts.persona_update += 1
elif category == 'task':
if operation == 'query':
self.counts.task_query += 1
else:
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_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.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()