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:
jinliyl 2026-03-27 21:14:44 +08:00 • committed by GitHub
parent bf79986f9c
commit ff49a77f18
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 82 additions and 33 deletions

View file

@ -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",

View file

@ -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(

View file

@ -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

View file

@ -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")