mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-26 13:54:48 -06:00
6b6c220986
The compaction summary ran as a single model call sized by the per-message token estimate, which disagreed with the head+tail-capped formatted text, so a long history could overflow the summary call itself; the old prefix-fit also silently dropped the most-recent messages. Summarize the whole selection via _summarize_blocks: greedily pack the formatted blocks into batches that each fit the summary call's own input budget (_summary_input_budget_chars), summarize each, and recursively merge the partials until they collapse to one. The common case (it all fits) stays a single call. Bail to the existing False path when the input is irreducible rather than fabricate a summary; a mid-chunk failure leaves messages untouched (atomic swap only on full success). Bound the summary output reserve to half the context window (_summary_output_tokens), used by BOTH the input-budget sizing and the actual call. compact_max_tokens defaults to the full window (32768); clamped only by max_output_tokens it reserved the entire context for output, flooring the input budget so compaction overflowed (or bailed as irreducible) at the default/small-window config that needs it most. Large windows are unaffected (compact_max_tokens stays binding). Guard an empty summary (keep history instead of swapping in nothing and reporting success). Fold tool-def tokens into the _last_usage-less estimate AND the post-compaction usage anchor, so the compact-before-truncate budget doesn't over-state free space by the tool-def count. Single-source the shared compactor/merge prompt section (_COMPACT_OUTPUT_FORMAT) and the tool-def sizing (_tool_def_chars/_tool_def_tokens). A just-resumed session (no _last_usage) now counts tool-def tokens so it doesn't undercount and skip proactive compaction until its first reply re-anchors the estimate. Prepush review follow-ups: generation-guard the end-of-turn auto-compaction and its resume turn so a force-cancel during the slow summary call can't compact or persist under a new generation (matching the mid-turn and end-of-loop guards); single-source the soft-threshold predicate (_over_soft) shared by the mid-turn policy, _compaction_owed, and the end-of-turn check; add tests for the pre-attempted-compaction guard and the recursion depth ceiling.
245 lines
9.2 KiB
Python
245 lines
9.2 KiB
Python
"""Tests for capacity-aware tool output truncation and context overflow recovery."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
from turnstone.core.session import ChatSession
|
|
from turnstone.core.trajectory import turns_from_dicts
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@pytest.fixture
|
|
def session(tmp_db, mock_openai_client):
|
|
"""Create a ChatSession with defaults for truncation testing."""
|
|
return ChatSession(
|
|
client=mock_openai_client,
|
|
model="test-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
tool_timeout=10,
|
|
context_window=10_000,
|
|
max_tokens=1_000,
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# _truncate_output
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestTruncateOutput:
|
|
def test_no_truncation_when_under_limit(self, session):
|
|
result = session._truncate_output("short text")
|
|
assert result == "short text"
|
|
|
|
def test_truncates_to_tool_truncation_limit(self, session):
|
|
session.tool_truncation = 100
|
|
big = "x" * 500
|
|
result = session._truncate_output(big)
|
|
assert len(result) <= 200 # head + tail + marker
|
|
assert "chars truncated" in result
|
|
|
|
def test_budget_aware_truncation(self, session):
|
|
session.tool_truncation = 100_000
|
|
session._chars_per_token = 4.0
|
|
# Budget of 50 tokens = 200 chars
|
|
big = "x" * 1000
|
|
result = session._truncate_output(big, remaining_budget_tokens=50)
|
|
assert len(result) <= 400 # head + tail + marker
|
|
assert "chars truncated" in result
|
|
|
|
def test_budget_takes_precedence_when_smaller(self, session):
|
|
session.tool_truncation = 10_000
|
|
session._chars_per_token = 4.0
|
|
# Budget of 25 tokens = 100 chars, smaller than tool_truncation
|
|
big = "x" * 500
|
|
result = session._truncate_output(big, remaining_budget_tokens=25)
|
|
assert "chars truncated" in result
|
|
|
|
def test_zero_budget_returns_placeholder(self, session):
|
|
big = "x" * 1000
|
|
result = session._truncate_output(big, remaining_budget_tokens=0)
|
|
assert "exceeded context budget" in result
|
|
assert len(result) < 100
|
|
|
|
def test_negative_budget_returns_placeholder(self, session):
|
|
big = "x" * 1000
|
|
result = session._truncate_output(big, remaining_budget_tokens=-10)
|
|
assert "exceeded context budget" in result
|
|
|
|
def test_none_budget_uses_fixed_limit(self, session):
|
|
session.tool_truncation = 100
|
|
big = "x" * 500
|
|
result = session._truncate_output(big, remaining_budget_tokens=None)
|
|
assert "100 char limit" in result
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# _remaining_token_budget
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestRemainingTokenBudget:
|
|
def test_empty_session(self, session):
|
|
session._tools = [] # isolate the budget formula from the tool-def estimate
|
|
session._system_tokens = 500
|
|
session._msg_tokens = []
|
|
budget = session._remaining_token_budget()
|
|
# 10000 - 500 - 0 - 1000 - 500 (5%) = 8000
|
|
assert budget == 8000
|
|
|
|
def test_partially_full(self, session):
|
|
session._tools = [] # isolate the budget formula from the tool-def estimate
|
|
session._system_tokens = 500
|
|
session._msg_tokens = [2000, 3000]
|
|
budget = session._remaining_token_budget()
|
|
# 10000 - 500 - 5000 - 1000 - 500 = 3000
|
|
assert budget == 3000
|
|
|
|
def test_overfull_returns_zero(self, session):
|
|
session._system_tokens = 500
|
|
session._msg_tokens = [9000]
|
|
assert session._remaining_token_budget() == 0
|
|
|
|
def test_exactly_full_returns_zero(self, session):
|
|
session._system_tokens = 500
|
|
session._msg_tokens = [8000]
|
|
assert session._remaining_token_budget() == 0
|
|
|
|
def test_max_tokens_equals_context_window(self, tmp_db, mock_openai_client):
|
|
"""Regression: max_tokens >= context_window must not zero the budget."""
|
|
s = ChatSession(
|
|
client=mock_openai_client,
|
|
model="test-model",
|
|
ui=MagicMock(),
|
|
instructions=None,
|
|
temperature=0.5,
|
|
tool_timeout=10,
|
|
context_window=32_768,
|
|
max_tokens=32_768,
|
|
)
|
|
s._tools = [] # isolate the budget formula from the tool-def estimate
|
|
s._system_tokens = 500
|
|
s._msg_tokens = [1000]
|
|
budget = s._remaining_token_budget()
|
|
# response_reserve = min(32768, 32768//4) = 8192
|
|
# safety = 32768 * 0.05 = 1638
|
|
# budget = 32768 - 500 - 1000 - 8192 - 1638 = 21438
|
|
assert budget > 20_000
|
|
# Tool output should NOT be collapsed to a placeholder
|
|
big = "x" * 5000
|
|
result = s._truncate_output(big, remaining_budget_tokens=budget)
|
|
assert result == big # 5000 chars fits easily in 21K+ token budget
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Context overflow recovery
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestContextOverflowRecovery:
|
|
"""Test that context-length errors trigger compact-and-retry."""
|
|
|
|
def test_openai_context_length_error_triggers_compact(self, session):
|
|
session.messages = turns_from_dicts([{"role": "user", "content": "hi"}])
|
|
session._msg_tokens = [1]
|
|
|
|
call_count = 0
|
|
|
|
def mock_create_stream(msgs):
|
|
nonlocal call_count
|
|
call_count += 1
|
|
if call_count == 1:
|
|
raise Exception("maximum context length exceeded")
|
|
return iter([])
|
|
|
|
compact_mock = MagicMock()
|
|
with (
|
|
patch.object(session, "_create_stream_with_retry", side_effect=mock_create_stream),
|
|
patch.object(session, "_compact_messages", compact_mock),
|
|
patch.object(
|
|
session, "_stream_response", return_value={"role": "assistant", "content": "ok"}
|
|
),
|
|
patch.object(session, "_full_messages", return_value=[]),
|
|
patch.object(session, "_update_token_table"),
|
|
patch.object(session, "_print_status_line"),
|
|
patch.object(session, "_emit_state"),
|
|
patch("turnstone.core.session.save_message"),
|
|
):
|
|
session.send("hello")
|
|
|
|
compact_mock.assert_called_once_with(auto=True)
|
|
assert call_count == 2
|
|
|
|
def test_anthropic_prompt_too_long_triggers_compact(self, session):
|
|
session.messages = turns_from_dicts([{"role": "user", "content": "hi"}])
|
|
session._msg_tokens = [1]
|
|
|
|
call_count = 0
|
|
|
|
def mock_create_stream(msgs):
|
|
nonlocal call_count
|
|
call_count += 1
|
|
if call_count == 1:
|
|
raise Exception("prompt is too long: 250000 tokens > 200000 maximum")
|
|
return iter([])
|
|
|
|
compact_mock = MagicMock()
|
|
with (
|
|
patch.object(session, "_create_stream_with_retry", side_effect=mock_create_stream),
|
|
patch.object(session, "_compact_messages", compact_mock),
|
|
patch.object(
|
|
session, "_stream_response", return_value={"role": "assistant", "content": "ok"}
|
|
),
|
|
patch.object(session, "_full_messages", return_value=[]),
|
|
patch.object(session, "_update_token_table"),
|
|
patch.object(session, "_print_status_line"),
|
|
patch.object(session, "_emit_state"),
|
|
patch("turnstone.core.session.save_message"),
|
|
):
|
|
session.send("hello")
|
|
|
|
compact_mock.assert_called_once_with(auto=True)
|
|
|
|
def test_non_context_error_propagates(self, session):
|
|
session.messages = turns_from_dicts([{"role": "user", "content": "hi"}])
|
|
session._msg_tokens = [1]
|
|
|
|
with (
|
|
patch.object(
|
|
session,
|
|
"_create_stream_with_retry",
|
|
side_effect=Exception("authentication failed"),
|
|
),
|
|
patch.object(session, "_full_messages", return_value=[]),
|
|
patch.object(session, "_emit_state"),
|
|
patch("turnstone.core.session.save_message"),
|
|
pytest.raises(Exception, match="authentication failed"),
|
|
):
|
|
session.send("hello")
|
|
|
|
def test_compact_failure_raises_original_error(self, session):
|
|
session.messages = turns_from_dicts([{"role": "user", "content": "hi"}])
|
|
session._msg_tokens = [1]
|
|
|
|
with (
|
|
patch.object(
|
|
session,
|
|
"_create_stream_with_retry",
|
|
side_effect=Exception("maximum context length exceeded"),
|
|
),
|
|
patch.object(session, "_compact_messages", side_effect=RuntimeError("compact failed")),
|
|
patch.object(session, "_full_messages", return_value=[]),
|
|
patch.object(session, "_emit_state"),
|
|
patch("turnstone.core.session.save_message"),
|
|
pytest.raises(Exception, match="maximum context length exceeded"),
|
|
):
|
|
session.send("hello")
|