diff --git a/strix/llm/compaction.py b/strix/llm/compaction.py index cfd224be..56f84db5 100644 --- a/strix/llm/compaction.py +++ b/strix/llm/compaction.py @@ -28,6 +28,9 @@ logger = logging.getLogger(__name__) _CHECKPOINT_TAG = "" _TOOL_OUTPUT_MAX_CHARS = 2_000 _MIN_ITEMS_TO_COMPACT = 6 +# Floor for how much of the head we still try to summarise even on a tiny model. +_MIN_SUMMARY_INPUT_TOKENS = 1_000 +_HEAD_TRUNCATED_MARKER = "\n\n[... older conversation omitted to fit the summary request ...]\n\n" # Substrings that identify a context-window-overflow error across providers. # Deliberately excludes rate-limit/throttle wording, which must not trigger @@ -182,6 +185,40 @@ def _previous_summary(head: list[Any]) -> str | None: return None +def _fit_to_tokens(model: str, text: str, max_tokens: int) -> str: + """Head+tail-truncate ``text`` so it fits within ``max_tokens``. + + Keeps the start (objective/setup) and the end (most recent activity), which + matter most for continuity, and drops the middle. Prevents the summary + request itself from overflowing the model window on very large histories. + """ + if count_tokens(model, text) <= max_tokens: + return text + # Convert the token budget to a rough character budget (chars ~= 4x tokens) + # split between head and tail, then confirm and tighten by real token count. + budget_chars = max_tokens * 4 + head_chars = budget_chars // 2 + tail_chars = budget_chars - head_chars + candidate = text[:head_chars] + _HEAD_TRUNCATED_MARKER + text[len(text) - tail_chars :] + while count_tokens(model, candidate) > max_tokens and (head_chars > 0 or tail_chars > 0): + head_chars = int(head_chars * 0.8) + tail_chars = int(tail_chars * 0.8) + candidate = text[:head_chars] + _HEAD_TRUNCATED_MARKER + text[len(text) - tail_chars :] + return candidate + + +def _summary_input_budget(model: str, previous: str | None) -> int: + """Token room left for the head after instructions and the summary output.""" + context = load_settings().context + overhead = count_tokens(model, _SUMMARY_INSTRUCTIONS) + if previous: + overhead += count_tokens(model, previous) + # Leave slack for the prompt's wrapper text ("Conversation to summarise:", + # the update instructions, etc.) that is not part of ``overhead``. + room = context_window(model) - context.summary_max_tokens - overhead - 256 + return max(_MIN_SUMMARY_INPUT_TOKENS, room) + + def _build_summary_prompt(serialized_head: str, previous: str | None) -> str: previous_block = ( f"\n\nA previous checkpoint summary follows. Update it: keep what is " @@ -263,9 +300,13 @@ async def maybe_compact( if not head: return False + previous = _previous_summary(head) + serialized_head = _fit_to_tokens( + model, _serialize_items(head), _summary_input_budget(model, previous) + ) summary = await _summarize( model, - _build_summary_prompt(_serialize_items(head), _previous_summary(head)), + _build_summary_prompt(serialized_head, previous), context.summary_max_tokens, ) if summary is None: diff --git a/tests/test_compaction.py b/tests/test_compaction.py index b3954c1a..ae4a0733 100644 --- a/tests/test_compaction.py +++ b/tests/test_compaction.py @@ -171,6 +171,37 @@ async def test_maybe_compact_updates_previous_summary(monkeypatch: pytest.Monkey assert "OLD SUMMARY TEXT" in captured["prompt"] +def test_fit_to_tokens_truncates_oversized_text(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(compaction, "count_tokens", lambda _m, t: len(t)) + text = "x" * 10_000 + + fitted = compaction._fit_to_tokens("m", text, 500) + + assert len(fitted) <= 500 + assert compaction._HEAD_TRUNCATED_MARKER in fitted + # Small text is returned untouched. + assert compaction._fit_to_tokens("m", "short", 500) == "short" + + +@pytest.mark.asyncio +async def test_maybe_compact_bounds_summary_prompt(monkeypatch: pytest.MonkeyPatch) -> None: + # A tiny window with a huge head must not send an oversized summary request. + _patch_budget(monkeypatch, keep_tokens=30, window=4_000) + captured: dict[str, str] = {} + + async def fake_acompletion(**kwargs: Any) -> Any: + captured["prompt"] = kwargs["messages"][0]["content"] + return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="S"))]) + + monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion) + big_turns = [{"role": "user", "content": "y" * 2_000} for _ in range(50)] + session = FakeSession(big_turns) + + assert await compaction.maybe_compact(session, model="m") is True + # count_tokens==len(chars); prompt must fit the model window. + assert len(captured["prompt"]) <= 4_000 + + @pytest.mark.asyncio async def test_maybe_compact_skips_when_summary_fails(monkeypatch: pytest.MonkeyPatch) -> None: _patch_budget(monkeypatch, keep_tokens=30, window=50)