fix: preserve tool results during context compaction

This commit is contained in:
Haoran
2026-08-23 11:50:52 +08:00
parent 61c3e64595
commit 6d70d2e071
24 changed files with 843 additions and 497 deletions

View File

@@ -18,10 +18,13 @@ s08_context_compact.py - Context Compact
| no | yes
| v
| +--------------------+
| | micro_compact | shorten old tool results
| | micro_compact | save + shorten old results
| +--------------------+
| |
| v
| fit_tool_results persist oversized new results
| |
| v
| still over limit?
| | no | yes
v v v
@@ -298,15 +301,53 @@ class ContextCompactor:
transcript.write(json.dumps(message, default=str, ensure_ascii=False) + "\n")
return path
def persist_large_output(self, tool_use_id: str, output: str) -> str:
if len(output) <= self.LARGE_RESULT_CHAR_LIMIT:
return output
def persisted_output_path(self, output: str) -> str | None:
candidate = None
if output.startswith("<persisted-output>\n"):
candidate = next(
(line.removeprefix("Full output: ")
for line in output.splitlines()
if line.startswith("Full output: ")),
None,
)
prefix = "[Earlier tool result saved at "
if output.startswith(prefix) and output.endswith("]"):
candidate = output.removeprefix(prefix).removesuffix("]")
if not candidate:
return None
path = Path(candidate)
if (not path.resolve().is_relative_to(self.tool_results_dir.resolve())
or not path.is_file()):
return None
return str(path)
def save_output(self, tool_use_id: str, output: str) -> Path:
self.tool_results_dir.mkdir(parents=True, exist_ok=True)
safe_id = re.sub(r"[^A-Za-z0-9._-]", "_", str(tool_use_id))[:120] or "unknown"
path = self.tool_results_dir / f"{safe_id}.txt"
if not path.exists():
path.write_text(output)
return f"<persisted-output>\nFull output: {path}\nPreview:\n{output[:2000]}\n</persisted-output>"
path.write_text(output, encoding="utf-8")
return path
def persisted_preview(self, tool_use_id: str, output: str,
preview_chars: int = 2000) -> str:
saved_path = self.persisted_output_path(output)
if saved_path:
path = Path(saved_path)
try:
with path.open(encoding="utf-8") as saved:
preview = saved.read(preview_chars)
except OSError:
preview = output[:preview_chars]
else:
path = self.save_output(tool_use_id, output)
preview = output[:preview_chars]
return (f"<persisted-output>\nFull output: {path}\n"
f"Preview:\n{preview}\n</persisted-output>")
def persist_large_output(self, tool_use_id: str, output: str) -> str:
if len(output) <= self.LARGE_RESULT_CHAR_LIMIT:
return output
return self.persisted_preview(tool_use_id, output)
def tool_result_budget(self, messages: list, max_chars: int | None = None) -> list:
if not messages:
@@ -328,11 +369,21 @@ class ContextCompactor:
total = sum(len(str(item.get("content", ""))) for item in blocks)
return messages
def is_archive_marker(self, message: dict) -> bool:
content = message.get("content")
match = (re.fullmatch(r"\[\d+ messages archived at (.+)\]", content)
if isinstance(content, str) else None)
if not match:
return False
path = Path(match.group(1))
return (path.resolve().is_relative_to(self.transcript_dir.resolve())
and path.is_file())
def snip_compact(self, messages: list, max_messages: int = 50) -> list:
if len(messages) <= max_messages:
return messages
head_end = 3
tail_start = len(messages) - (max_messages - head_end)
tail_start = len(messages) - (max_messages - head_end - 1)
if self.has_tool_use(messages[head_end - 1]):
while head_end < tail_start and self.is_tool_result(messages[head_end]):
head_end += 1
@@ -341,12 +392,16 @@ class ContextCompactor:
tail_start -= 1
if head_end >= tail_start:
return messages
middle = messages[head_end:tail_start]
if len(middle) == 1 and self.is_archive_marker(middle[0]):
return messages
transcript_path = self.write_transcript(messages)
marker = {"role": "user", "content":
f"[{tail_start - head_end} messages archived at {transcript_path}]"}
return [*messages[:head_end], marker, *messages[tail_start:]]
def micro_compact(self, messages: list) -> list:
def micro_compact(self, messages: list,
target_chars: int | None = None) -> list:
results = [
(message_index, block_index, block)
for message_index, message in enumerate(messages)
@@ -357,18 +412,38 @@ class ContextCompactor:
unseen = self.unseen_tool_result_positions(messages)
consumed = [entry for entry in results if entry[:2] not in unseen]
for _, _, block in consumed[:-self.KEEP_RECENT_RESULTS]:
if (target_chars is not None
and self.estimate_chars(messages) <= target_chars):
break
content = str(block.get("content", ""))
if len(content) <= 120:
continue
saved_path = next(
(line.removeprefix("Full output: ") for line in content.splitlines()
if line.startswith("Full output: ")),
None,
)
block["content"] = (
f"[Earlier tool result saved at {saved_path}]"
if saved_path else "[Earlier tool result omitted.]"
)
saved_path = self.persisted_output_path(content)
if not saved_path:
saved_path = str(self.save_output(
block.get("tool_use_id", "unknown"), content))
block["content"] = f"[Earlier tool result saved at {saved_path}]"
return messages
def fit_tool_results(self, messages: list, target_chars: int) -> list:
results = [
block
for message in messages
if message.get("role") == "user" and isinstance(message.get("content"), list)
for block in message["content"]
if isinstance(block, dict) and block.get("type") == "tool_result"
]
for block in sorted(
results,
key=lambda item: len(str(item.get("content", ""))),
reverse=True):
if self.estimate_chars(messages) <= target_chars:
break
output = str(block.get("content", ""))
replacement = self.persisted_preview(
block.get("tool_use_id", "unknown"), output, preview_chars=1000)
if len(replacement) < len(output):
block["content"] = replacement
return messages
def summary_input(self, messages: list) -> str:
@@ -426,7 +501,10 @@ class ContextCompactor:
messages = self.tool_result_budget(messages)
messages = self.snip_compact(messages)
if self.estimate_chars(messages) > self.CONTEXT_CHAR_LIMIT:
messages = self.micro_compact(messages)
target = int(self.CONTEXT_CHAR_LIMIT * 0.8)
messages = self.micro_compact(messages, target)
if self.estimate_chars(messages) > self.CONTEXT_CHAR_LIMIT:
messages = self.fit_tool_results(messages, target)
if self.estimate_chars(messages) > self.CONTEXT_CHAR_LIMIT:
print("[auto compact]")
messages = self.compact_history(messages, active_request)