mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-11 03:40:03 +00:00
feat(memory): add tool call deduplication with hash-based detection
This commit is contained in:
parent
d9f332ece2
commit
69bf1e9d5f
2 changed files with 51 additions and 1 deletions
|
|
@ -1,4 +1,6 @@
|
|||
import datetime
|
||||
import hashlib
|
||||
import json
|
||||
from abc import ABC
|
||||
from typing import List
|
||||
from uuid import uuid4
|
||||
|
|
@ -119,9 +121,28 @@ class ToolCallResult(BaseModel):
|
|||
evaluation: str = Field(default="", description="Detailed evaluation for the tool invocation")
|
||||
score: float = Field(default=0, description="Score of the Evaluation (0.0 for failure, 1.0 for complete success)")
|
||||
is_summarized: bool = Field(default=False, description="Whether this tool call has been included in a summary")
|
||||
call_hash: str = Field(default="", description="Hash value of input and output combined for deduplication")
|
||||
|
||||
metadata: dict = Field(default_factory=dict)
|
||||
|
||||
def generate_hash(self) -> str:
|
||||
"""Generate hash value from tool input and output for deduplication"""
|
||||
# Convert input to string if it's a dict
|
||||
input_str = json.dumps(self.input, sort_keys=True) if isinstance(self.input, dict) else str(self.input)
|
||||
|
||||
# Combine input and output
|
||||
combined = f"{input_str}|{self.output}"
|
||||
|
||||
# Generate MD5 hash
|
||||
hash_value = hashlib.md5(combined.encode('utf-8')).hexdigest()
|
||||
|
||||
return hash_value
|
||||
|
||||
def ensure_hash(self):
|
||||
"""Ensure call_hash is set, generate if empty"""
|
||||
if not self.call_hash:
|
||||
self.call_hash = self.generate_hash()
|
||||
|
||||
def from_mcp_tool_result(self, tool_result: CallToolResult, max_char_len: int = None):
|
||||
text_list = []
|
||||
for content in tool_result.content:
|
||||
|
|
|
|||
|
|
@ -107,6 +107,10 @@ class ParseToolCallResultOp(BaseAsyncOp):
|
|||
self.context.response.success = False
|
||||
return
|
||||
|
||||
# 确保所有 tool_call_results 都有 hash 值
|
||||
for tool_call_result in tool_call_results:
|
||||
tool_call_result.ensure_hash()
|
||||
|
||||
# 使用基类的 submit_async_task 提交所有评估任务
|
||||
for index, tool_call_result in enumerate(tool_call_results):
|
||||
self.submit_async_task(self._evaluate_single_tool_call, tool_call_result, index)
|
||||
|
|
@ -122,6 +126,7 @@ class ParseToolCallResultOp(BaseAsyncOp):
|
|||
# 处理每个 tool_name 的结果
|
||||
all_memory_list = []
|
||||
all_deleted_memory_ids = []
|
||||
deduplication_stats = {"total_new": 0, "deduplicated": 0, "added": 0}
|
||||
|
||||
for tool_name, tool_call_results in tool_results_by_name.items():
|
||||
nodes: List[VectorNode] = await self.vector_store.async_search(query=tool_name,
|
||||
|
|
@ -144,7 +149,23 @@ class ParseToolCallResultOp(BaseAsyncOp):
|
|||
if tool_memory is None:
|
||||
tool_memory = ToolMemory(workspace_id=workspace_id, when_to_use=tool_name)
|
||||
|
||||
tool_memory.tool_call_results.extend(tool_call_results)
|
||||
# 获取现有的所有 hash 值用于去重
|
||||
existing_hashes = {result.call_hash for result in tool_memory.tool_call_results if result.call_hash}
|
||||
|
||||
# 过滤掉重复的 tool_call_results
|
||||
new_results = []
|
||||
for result in tool_call_results:
|
||||
deduplication_stats["total_new"] += 1
|
||||
if result.call_hash not in existing_hashes:
|
||||
new_results.append(result)
|
||||
existing_hashes.add(result.call_hash)
|
||||
deduplication_stats["added"] += 1
|
||||
else:
|
||||
deduplication_stats["deduplicated"] += 1
|
||||
logger.info(f"Skipping duplicate tool call for {tool_name} with hash {result.call_hash}")
|
||||
|
||||
# 只添加非重复的结果
|
||||
tool_memory.tool_call_results.extend(new_results)
|
||||
|
||||
# 保留最近的 n 个
|
||||
if len(tool_memory.tool_call_results) > self.max_history_tool_call_cnt:
|
||||
|
|
@ -162,11 +183,19 @@ class ParseToolCallResultOp(BaseAsyncOp):
|
|||
# 格式化结果信息
|
||||
formatted_answer = self._format_tool_memories_summary(all_memory_list, all_deleted_memory_ids)
|
||||
|
||||
# 添加去重统计信息
|
||||
dedup_info = (f"\n\nDeduplication Summary:\n"
|
||||
f" Total new calls: {deduplication_stats['total_new']}\n"
|
||||
f" Added: {deduplication_stats['added']}\n"
|
||||
f" Deduplicated: {deduplication_stats['deduplicated']}")
|
||||
formatted_answer += dedup_info
|
||||
|
||||
# 设置返回结果
|
||||
self.context.response.answer = formatted_answer
|
||||
self.context.response.success = True
|
||||
self.context.response.metadata["deleted_memory_ids"] = all_deleted_memory_ids
|
||||
self.context.response.metadata["memory_list"] = all_memory_list
|
||||
self.context.response.metadata["deduplication_stats"] = deduplication_stats
|
||||
|
||||
|
||||
async def main():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue