mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
feat(memory): improve skills tool result truncation (#182)
* fix(memory): correct line numbering and improve tool result truncation - Changed default start_line from 0 to 1 in truncate_text_output function - Refactored _truncate method to be a standalone method in ToolResultCompactor - Improved tool result compaction logic to handle text blocks more efficiently - Added detection of skill-related tool calls for special handling - Implemented conditional byte limits based on tool type for better memory management - Updated version number from 0.3.1.4 to 0.3.1.5 * feat(file_utils): add encoding parameter to truncate_text_output function - Added encoding parameter with default value "utf-8" to truncate_text_output function - Updated all encode/decode calls to use the specified encoding parameter - Modified ToolResultCompactor to pass encoding parameter when calling truncate_text_output - Added error handling for skill tool ID detection in message processing loop - Fixed potential AttributeError when accessing raw_input field that might be None * fix(file-store): handle corrupted ChromaDB initialization and improve tool result truncation - Add shutil import for directory removal operations - Extract ChromaDB client creation into separate _create_chroma_client method - Implement retry mechanism with database wipe on ChromaDB initialization failure - Add proper exception handling in tool result compaction to prevent truncation errors - Move file writing logic outside of exception handling scope for better error management - Add warning log when truncation fails and return original content as fallback
This commit is contained in:
parent
bf79986f9c
commit
ff49a77f18
4 changed files with 82 additions and 33 deletions
|
|
@ -6,7 +6,7 @@ from . import extension
|
|||
from . import memory
|
||||
from .reme import ReMe
|
||||
|
||||
__version__ = "0.3.1.4"
|
||||
__version__ = "0.3.1.5"
|
||||
|
||||
__all__ = [
|
||||
"config",
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
import json
|
||||
import random
|
||||
import shutil
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
|
|
@ -108,13 +109,9 @@ class ChromaFileStore(BaseFileStore):
|
|||
except Exception as e:
|
||||
logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}")
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Initialize ChromaDB client and collection."""
|
||||
if self.client is not None:
|
||||
return
|
||||
|
||||
# Initialize persistent ChromaDB client
|
||||
self.client = chromadb.PersistentClient(
|
||||
def _create_chroma_client(self):
|
||||
"""Create a ChromaDB PersistentClient instance."""
|
||||
return chromadb.PersistentClient(
|
||||
path=str(self.db_path),
|
||||
settings=Settings(
|
||||
anonymized_telemetry=False,
|
||||
|
|
@ -122,6 +119,23 @@ class ChromaFileStore(BaseFileStore):
|
|||
),
|
||||
)
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Initialize ChromaDB client and collection."""
|
||||
if self.client is not None:
|
||||
return
|
||||
|
||||
# Initialize persistent ChromaDB client, retry once after wiping db_path on failure
|
||||
try:
|
||||
self.client = self._create_chroma_client()
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"ChromaDB failed to initialize at {self.db_path} ({e}). " f"Deleting corrupted database and retrying.",
|
||||
)
|
||||
if self.db_path.exists():
|
||||
shutil.rmtree(self.db_path)
|
||||
logger.info(f"Deleted ChromaDB directory: {self.db_path}")
|
||||
self.client = self._create_chroma_client()
|
||||
|
||||
# Get or create the chunks collection
|
||||
# ChromaDB uses cosine distance by default for similarity
|
||||
self.chunks_collection = self.client.get_or_create_collection(
|
||||
|
|
|
|||
|
|
@ -37,36 +37,43 @@ class ToolResultCompactor(BaseOp):
|
|||
self.encoding = encoding
|
||||
self.tool_result_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def _compact(self, output: str | list[dict], max_bytes: int) -> str | list[dict]:
|
||||
"""Truncate output to max_bytes, saving full content to file if needed."""
|
||||
|
||||
def _truncate(content: str) -> str:
|
||||
if not content:
|
||||
return content
|
||||
def _truncate(self, content: str, max_bytes: int) -> str:
|
||||
if not content:
|
||||
return content
|
||||
|
||||
try:
|
||||
if TRUNCATION_NOTICE_MARKER in content:
|
||||
return truncate_text_output(content, max_bytes=max_bytes)
|
||||
return truncate_text_output(content, max_bytes=max_bytes, encoding=self.encoding)
|
||||
|
||||
if len(content.encode(self.encoding)) <= max_bytes + 100:
|
||||
return content
|
||||
|
||||
saved_path: str | None = None
|
||||
try:
|
||||
fp = self.tool_result_dir / f"{uuid.uuid4().hex}.txt"
|
||||
fp.write_text(content, encoding=self.encoding)
|
||||
saved_path = str(fp)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to save full tool result to file: %s", e)
|
||||
fp = self.tool_result_dir / f"{uuid.uuid4().hex}.txt"
|
||||
fp.write_text(content, encoding=self.encoding)
|
||||
saved_path = str(fp)
|
||||
|
||||
return truncate_text_output(content, 1, content.count("\n") + 1, max_bytes, file_path=saved_path)
|
||||
return truncate_text_output(
|
||||
content,
|
||||
1,
|
||||
content.count("\n") + 1,
|
||||
max_bytes,
|
||||
file_path=saved_path,
|
||||
encoding=self.encoding,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to truncate content, returning original: %s", e)
|
||||
return content
|
||||
|
||||
def _compact(self, output: str | list[dict], max_bytes: int) -> str | list[dict]:
|
||||
"""Truncate output to max_bytes, saving full content to file if needed."""
|
||||
|
||||
if isinstance(output, str):
|
||||
return _truncate(output)
|
||||
return self._truncate(output, max_bytes)
|
||||
if isinstance(output, list):
|
||||
return [
|
||||
{**b, "text": _truncate(b.get("text", ""))} if isinstance(b, dict) and b.get("type") == "text" else b
|
||||
for b in output
|
||||
]
|
||||
for b in output:
|
||||
if isinstance(b, dict) and b.get("type") == "text":
|
||||
b["text"] = self._truncate(b.get("text", ""), max_bytes)
|
||||
return output
|
||||
|
||||
async def execute(self) -> list[Msg]:
|
||||
|
|
@ -84,6 +91,27 @@ class ToolResultCompactor(BaseOp):
|
|||
recent_n += 1
|
||||
split_index = max(0, len(messages) - max(recent_n, self.recent_n))
|
||||
|
||||
skills_tool_ids = set()
|
||||
try:
|
||||
for msg in messages:
|
||||
if not isinstance(msg.content, list):
|
||||
continue
|
||||
|
||||
for block in msg.content:
|
||||
if isinstance(block, dict) and block.get("type") == "tool_use":
|
||||
tool_id = block.get("id", "")
|
||||
if not tool_id:
|
||||
continue
|
||||
|
||||
if (
|
||||
block.get("name", "").lower() == "read_file"
|
||||
and "skill.md" in (block.get("raw_input") or "").lower()
|
||||
):
|
||||
skills_tool_ids.add(tool_id)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to detect skill tool ids: %s", e)
|
||||
logger.info(f"skills_tool_ids: {skills_tool_ids}")
|
||||
|
||||
for idx, msg in enumerate(messages):
|
||||
if not isinstance(msg.content, list):
|
||||
continue
|
||||
|
|
@ -91,7 +119,12 @@ class ToolResultCompactor(BaseOp):
|
|||
max_bytes = self.recent_max_bytes if is_recent else self.old_max_bytes
|
||||
for block in msg.content:
|
||||
if isinstance(block, dict) and block.get("type") == "tool_result" and block.get("output"):
|
||||
block["output"] = self._compact(block["output"], max_bytes)
|
||||
tool_use_id = block.get("id", "")
|
||||
if tool_use_id in skills_tool_ids:
|
||||
effective_max_bytes = self.recent_max_bytes
|
||||
else:
|
||||
effective_max_bytes = max_bytes
|
||||
block["output"] = self._compact(block["output"], effective_max_bytes)
|
||||
|
||||
return messages
|
||||
|
||||
|
|
|
|||
|
|
@ -29,10 +29,11 @@ TRUNCATION_NOTICE_MARKER = "<<<TRUNCATED>>>"
|
|||
# pylint: disable=too-many-return-statements
|
||||
def truncate_text_output(
|
||||
text: str,
|
||||
start_line: int = 0,
|
||||
start_line: int = 1,
|
||||
total_lines: int = 0,
|
||||
max_bytes: int = DEFAULT_MAX_BYTES,
|
||||
file_path: str | None = None,
|
||||
encoding: str = "utf-8",
|
||||
) -> str:
|
||||
"""Truncate file output by bytes with line integrity.
|
||||
|
||||
|
|
@ -52,6 +53,7 @@ def truncate_text_output(
|
|||
contains a truncation notice (values are parsed from the notice instead).
|
||||
max_bytes: Maximum size in bytes.
|
||||
file_path: Optional file path to include in the truncation notice.
|
||||
encoding: Character encoding used for byte-length calculation and decoding.
|
||||
|
||||
Returns:
|
||||
Truncated text with notice if truncated.
|
||||
|
|
@ -67,7 +69,7 @@ def truncate_text_output(
|
|||
original_content = parts[0]
|
||||
old_notice = parts[1]
|
||||
|
||||
text_bytes = original_content.encode("utf-8")
|
||||
text_bytes = original_content.encode(encoding)
|
||||
|
||||
# Allow a small slack to avoid re-truncating near-limit content
|
||||
if len(text_bytes) <= max_bytes + 100:
|
||||
|
|
@ -82,7 +84,7 @@ def truncate_text_output(
|
|||
total_lines_parsed = int(total_match.group(1))
|
||||
|
||||
truncated_bytes = text_bytes[:max_bytes]
|
||||
result = truncated_bytes.decode("utf-8", errors="ignore")
|
||||
result = truncated_bytes.decode(encoding, errors="ignore")
|
||||
newline_count = result.count("\n")
|
||||
|
||||
next_line = start_line_parsed + max(1, newline_count)
|
||||
|
|
@ -99,13 +101,13 @@ def truncate_text_output(
|
|||
return result + TRUNCATION_NOTICE_MARKER + new_notice
|
||||
|
||||
else:
|
||||
text_bytes = text.encode("utf-8")
|
||||
text_bytes = text.encode(encoding)
|
||||
|
||||
if len(text_bytes) <= max_bytes:
|
||||
return text
|
||||
|
||||
truncated = text_bytes[:max_bytes]
|
||||
result = truncated.decode("utf-8", errors="ignore")
|
||||
result = truncated.decode(encoding, errors="ignore")
|
||||
|
||||
newline_count = result.count("\n")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue