diff --git a/reme_ai/schema/memory.py b/reme_ai/schema/memory.py index 821f74e4..27745f7b 100644 --- a/reme_ai/schema/memory.py +++ b/reme_ai/schema/memory.py @@ -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: diff --git a/reme_ai/summary/tool/parse_tool_call_result_op.py b/reme_ai/summary/tool/parse_tool_call_result_op.py index 835e4e06..9a049c25 100644 --- a/reme_ai/summary/tool/parse_tool_call_result_op.py +++ b/reme_ai/summary/tool/parse_tool_call_result_op.py @@ -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():