fix(compression): handle CJK token budgeting

This commit is contained in:
miniadmin 2026-06-25 23:07:10 +09:00 committed by Teknium
parent a269a5b38c
commit 3f33a1c5aa
4 changed files with 150 additions and 10 deletions

View file

@ -31,6 +31,7 @@ from agent.model_metadata import (
MINIMUM_CONTEXT_LENGTH,
get_model_context_length,
estimate_messages_tokens_rough,
estimate_tokens_rough,
)
from agent.redact import redact_sensitive_text
from agent.turn_context import drop_stale_api_content
@ -449,11 +450,15 @@ def _estimate_msg_budget_tokens(msg: dict) -> int:
compaction re-fires continuously (#55572). Accounting-only: replay
fields are never mutated or pruned here.
"""
content_len = _content_length_for_budget(msg.get("content") or "")
tokens = content_len // _CHARS_PER_TOKEN + 10 # +10 for role/key overhead
content = msg.get("content") or ""
if isinstance(content, str):
tokens = estimate_tokens_rough(content) + 10 # +10 for role/key overhead
else:
content_len = _content_length_for_budget(content)
tokens = content_len // _CHARS_PER_TOKEN + 10
for tc in msg.get("tool_calls") or []:
if isinstance(tc, dict):
tokens += len(str(tc)) // _CHARS_PER_TOKEN
tokens += estimate_tokens_rough(str(tc))
for key in _REPLAY_BUDGET_KEYS:
tokens += _serialized_length_for_budget(msg.get(key)) // _CHARS_PER_TOKEN
return tokens
@ -3704,6 +3709,19 @@ This compaction should PRIORITISE preserving all information related to the focu
# Phase 4: Assemble compressed message list
compressed = []
for i in range(compress_start):
# If an earlier compaction handoff is in the protected head
# (common after resume / in-place compaction), do not carry it
# forward verbatim. It has already been rehydrated into
# _previous_summary above and _generate_summary() will emit the
# updated replacement below. Keeping both makes repeated
# compactions accumulate old summaries and prevents the live prompt
# from actually shrinking.
if (
summary_idx is not None
and i == summary_idx
and self._is_context_summary_content(messages[i].get("content"))
):
continue
msg = _fresh_compaction_message_copy(messages[i])
if i == 0 and msg.get("role") == "system":
existing = msg.get("content")
@ -3730,7 +3748,7 @@ This compaction should PRIORITISE preserving all information related to the focu
)
_merge_summary_into_tail = False
last_head_role = messages[compress_start - 1].get("role", "user") if compress_start > 0 else "user"
last_head_role = compressed[-1].get("role", "user") if compressed else "user"
first_tail_role = messages[compress_end].get("role", "user") if compress_end < n_messages else "user"
# When the only protected head message is the system prompt, the
# summary becomes the first *visible* message in the API request

View file

@ -2666,16 +2666,39 @@ async def get_model_context_length_async(
)
def _is_cjk_token_dense_char(ch: str) -> bool:
code = ord(ch)
return (
0x1100 <= code <= 0x11FF # Hangul Jamo
or 0x2E80 <= code <= 0x9FFF # CJK radicals/ideographs
or 0xA960 <= code <= 0xA97F # Hangul Jamo Extended-A
or 0xAC00 <= code <= 0xD7AF # Hangul Syllables
or 0xF900 <= code <= 0xFAFF # CJK compatibility ideographs
or 0xFF00 <= code <= 0xFFEF # Fullwidth forms / halfwidth kana
)
def estimate_tokens_rough(text: str) -> int:
"""Rough token estimate (~4 chars/token) for pre-flight checks.
"""Rough token estimate for pre-flight checks.
Uses ceiling division so short texts (1-3 chars) never estimate as
0 tokens, which would cause the compressor and pre-flight checks to
systematically undercount when many short tool results are present.
CJK/Hangul/Kana text is much denser than English under common LLM
tokenizers, so count those codepoints as roughly one token each instead
of applying the English-centric ~4 chars/token rule.
"""
if not text:
return 0
return (len(text) + 3) // 4
text = str(text)
dense = 0
sparse = 0
for ch in text:
if _is_cjk_token_dense_char(ch):
dense += 1
else:
sparse += 1
return dense + ((sparse + 3) // 4)
def estimate_messages_tokens_rough(messages: List[Dict[str, Any]]) -> int:
@ -2687,12 +2710,12 @@ def estimate_messages_tokens_rough(messages: List[Dict[str, Any]]) -> int:
estimated at ~250K tokens and trigger premature context compression.
"""
_IMAGE_TOKEN_COST = 1500
total_chars = 0
text_tokens = 0
image_tokens = 0
for msg in messages:
total_chars += _estimate_message_chars(msg)
text_tokens += _estimate_message_tokens_without_images(msg)
image_tokens += _count_image_tokens(msg, _IMAGE_TOKEN_COST)
return ((total_chars + 3) // 4) + image_tokens
return text_tokens + image_tokens
def _count_image_tokens(msg: Dict[str, Any], cost_per_image: int) -> int:
@ -2754,6 +2777,35 @@ def _estimate_message_chars(msg: Dict[str, Any]) -> int:
return len(str(shadow))
def _estimate_message_tokens_without_images(msg: Dict[str, Any]) -> int:
"""Token estimate for a message shadow with image payloads stripped."""
if not isinstance(msg, dict):
return estimate_tokens_rough(str(msg))
shadow: Dict[str, Any] = {}
for k, v in msg.items():
if k == "_anthropic_content_blocks":
continue
if k == "content":
if isinstance(v, list):
cleaned = []
for part in v:
if isinstance(part, dict):
if part.get("type") in {"image", "image_url", "input_image"}:
cleaned.append({"type": part.get("type"), "image": "[stripped]"})
else:
cleaned.append(part)
else:
cleaned.append(part)
shadow[k] = cleaned
elif isinstance(v, dict) and v.get("_multimodal"):
shadow[k] = v.get("text_summary", "")
else:
shadow[k] = v
else:
shadow[k] = v
return estimate_tokens_rough(str(shadow))
def estimate_request_tokens_rough(
messages: List[Dict[str, Any]],
*,
@ -2770,7 +2822,7 @@ def estimate_request_tokens_rough(
"""
total = 0
if system_prompt:
total += (len(system_prompt) + 3) // 4
total += estimate_tokens_rough(system_prompt)
if messages:
total += estimate_messages_tokens_rough(messages)
if tools:

View file

@ -0,0 +1,50 @@
from unittest.mock import patch
from agent.context_compressor import ContextCompressor, _estimate_msg_budget_tokens
from agent.model_metadata import estimate_messages_tokens_rough, estimate_tokens_rough
def test_cjk_text_is_not_estimated_as_four_chars_per_token():
assert estimate_tokens_rough("a" * 400) == 100
assert estimate_tokens_rough("" * 400) >= 400
def test_message_estimate_counts_korean_content_as_token_dense():
messages = [{"role": "user", "content": "압축 테스트 " + ("" * 1000)}]
assert estimate_messages_tokens_rough(messages) >= 1000
def test_compressor_tail_budget_uses_cjk_aware_message_estimate():
korean_msg = {"role": "assistant", "content": "" * 2000}
english_msg = {"role": "assistant", "content": "a" * 2000}
assert _estimate_msg_budget_tokens(korean_msg) > _estimate_msg_budget_tokens(english_msg)
def test_cjk_tail_does_not_expand_to_english_char_budget():
with patch("agent.context_compressor.get_model_context_length", return_value=65536):
compressor = ContextCompressor(
"test/model",
protect_first_n=3,
protect_last_n=20,
summary_target_ratio=0.2,
quiet_mode=True,
)
messages = [
{"role": "user", "content": "head 1"},
{"role": "assistant", "content": "head 2"},
{"role": "user", "content": "head 3"},
]
for idx in range(40):
role = "assistant" if idx % 2 else "user"
messages.append({"role": role, "content": "" * 1200})
compress_start = compressor._align_boundary_forward(
messages,
compressor._protect_head_size(messages),
)
compress_end = compressor._find_tail_cut_by_tokens(messages, compress_start)
assert len(messages) - compress_end < 31

View file

@ -89,3 +89,23 @@ def test_handoff_in_protected_head_populates_previous_summary_before_update():
assert compressor._previous_summary == old_summary
assert seen_turns
assert all(old_summary not in str(msg.get("content", "")) for msg in seen_turns)
def test_handoff_in_protected_head_is_replaced_not_duplicated():
"""Re-compaction must replace a protected old handoff with the updated one."""
compressor = _compressor()
old_summary = "OLD-PROTECTED-HANDOFF unique old summary body"
with patch("agent.context_compressor.call_llm", return_value=_response("UPDATED summary body")):
compressed = compressor.compress(_messages_with_handoff(old_summary))
summary_messages = [
msg
for msg in compressed
if isinstance(msg, dict)
and str(msg.get("content") or "").startswith(SUMMARY_PREFIX)
]
assert len(summary_messages) == 1
assert "UPDATED summary body" in summary_messages[0]["content"]
assert old_summary not in summary_messages[0]["content"]
assert old_summary not in "\n".join(str(msg.get("content") or "") for msg in compressed)