diff --git a/solstone/think/providers/openhands.py b/solstone/think/providers/openhands.py index 6898d66ff..d4a67bf1a 100644 --- a/solstone/think/providers/openhands.py +++ b/solstone/think/providers/openhands.py @@ -158,8 +158,20 @@ def _local_output_reserve_tokens(window: int) -> int: return window // 4 +_LOCAL_TOKENIZER_DIVERGENCE_FACTOR = 1.125 + + def _local_condenser_max_tokens(window: int) -> int: - return window * 11 // 16 + """Modeled input boundary with a measured client/server tokenizer margin. + + This factor is not a system/tool overhead term. The condenser token count + already includes the system prompt and tool schemas. It is a safety factor + for one production observation (n=1): 12,437 served tokens vs 11,237 + LiteLLM-estimated tokens, or +10.7%, rounded upward. + """ + + modeled_input_budget = window - _local_output_reserve_tokens(window) + return math.floor(modeled_input_budget / _LOCAL_TOKENIZER_DIVERGENCE_FACTOR) def _resolve_allowed_roots(config: dict[str, Any]) -> list[Path]: diff --git a/tests/test_cogitate_context_headroom.py b/tests/test_cogitate_context_headroom.py new file mode 100644 index 000000000..a4883c2ef --- /dev/null +++ b/tests/test_cogitate_context_headroom.py @@ -0,0 +1,302 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import asyncio +import math +import os +from collections.abc import Sequence +from dataclasses import dataclass +from typing import Any + +import pytest +from pydantic import Field, PrivateAttr + +from tests._logging_isolation import preserve_global_logging + +os.environ.setdefault("OPENHANDS_SUPPRESS_BANNER", "1") + +with preserve_global_logging(): + from openhands.sdk import Agent, Conversation, MessageEvent + from openhands.sdk.context.condenser.utils import get_total_token_count + from openhands.sdk.llm import LLMResponse, Message, TextContent + from openhands.sdk.testing import TestLLM + from openhands.sdk.tool import ToolAnnotations, ToolDefinition, ToolExecutor + from openhands.sdk.tool.registry import register_tool + from openhands.sdk.tool.schema import Action, Observation + from openhands.sdk.tool.spec import Tool + +from solstone.think.providers import openhands + +_TOOL_NAME = "context_headroom_probe" +_DIVERGENCE_FACTOR = 1.125 +_SEED_TEXT = ( + "Earlier segment observation records local context, tool use, finalization, " + "and verification details for this run.\n" * 40 +) +_PAD_UNIT = " measured evidence padding" + + +@dataclass(frozen=True) +class _RequestRecord: + kind: str + estimated_input: int + max_output_tokens: int + + +class _ProbeAction(Action): + value: str = Field(default="", description="Probe value.") + + +class _ProbeObservation(Observation): + pass + + +class _ProbeExecutor(ToolExecutor): + def __call__(self, action: Any, conversation: Any = None) -> _ProbeObservation: + del action, conversation + return _ProbeObservation.from_text("ok") + + +class _ProbeTool(ToolDefinition[_ProbeAction, _ProbeObservation]): + name = _TOOL_NAME + + @classmethod + def create(cls, *args: Any, **kwargs: Any) -> list[Any]: + del args, kwargs + return [] + + +class _RecordingLLM(TestLLM): + _records: list[_RequestRecord] = PrivateAttr(default_factory=list) + + @property + def records(self) -> list[_RequestRecord]: + return self._records + + def completion( + self, + messages: list[Message], + tools: Sequence[ToolDefinition] | None = None, + _return_metrics: bool = False, + add_security_risk_prediction: bool = False, + on_token: Any | None = None, + **kwargs: Any, + ) -> LLMResponse: + tool_list = list(tools or []) + self._records.append( + _RequestRecord( + kind="agent" if tool_list else "condenser", + estimated_input=self.get_token_count( + messages, + tools=tool_list, + add_security_risk_prediction=True, + ), + max_output_tokens=self.effective_max_output_tokens, + ) + ) + return super().completion( + messages=messages, + tools=tools, + _return_metrics=_return_metrics, + add_security_risk_prediction=add_security_risk_prediction, + on_token=on_token, + **kwargs, + ) + + +def _assistant_message(text: str) -> Message: + return Message(role="assistant", content=[TextContent(text=text)]) + + +def _probe_tool() -> _ProbeTool: + return _ProbeTool( + description="Probe tool that keeps the request on the tool-calling path.", + action_type=_ProbeAction, + observation_type=_ProbeObservation, + executor=_ProbeExecutor(), + annotations=ToolAnnotations( + title=_TOOL_NAME, + readOnlyHint=True, + destructiveHint=False, + idempotentHint=True, + openWorldHint=False, + ), + ) + + +def _recording_llm(*, window: int, reserve: int) -> _RecordingLLM: + return _RecordingLLM.from_messages( + [_assistant_message("summary") for _ in range(300)], + model="openai/gpt-4o-mini", + native_tool_calling=False, + max_input_tokens=window, + max_output_tokens=reserve, + input_cost_per_token=0, + output_cost_per_token=0, + ) + + +def _conversation( + *, + llm: _RecordingLLM, + boundary: int, + tmp_path, +) -> Conversation: + register_tool(_TOOL_NAME, _probe_tool()) + agent = Agent( + llm=llm, + tools=[Tool(name=_TOOL_NAME)], + include_default_tools=[], + system_prompt="System prompt for local condenser headroom tests.", + condenser=openhands._build_local_condenser(llm, max_tokens=boundary), + ) + return Conversation( + agent=agent, + workspace=str(tmp_path), + persistence_dir=str(tmp_path / "history"), + visualizer=None, + stuck_detection=False, + ) + + +def _user_event(text: str) -> MessageEvent: + return MessageEvent( + source="user", + llm_message=Message(role="user", content=[TextContent(text=text)]), + ) + + +def _prepared_count( + *, + conversation: Conversation, + llm: _RecordingLLM, + candidate: str | None = None, +) -> int: + events = list(conversation.state.events) + if candidate is not None: + events.append(_user_event(candidate)) + return get_total_token_count(events, llm) + + +def _run_turn(conversation: Conversation, message: str) -> None: + conversation.send_message(message) + asyncio.run(conversation.arun()) + + +def _seed_history( + *, + conversation: Conversation, + llm: _RecordingLLM, + target_low: int, +) -> None: + seed_turns = 0 + while True: + next_count = _prepared_count( + conversation=conversation, + llm=llm, + candidate=_SEED_TEXT, + ) + if seed_turns >= openhands._LOCAL_CONDENSER_KEEP_FIRST and ( + next_count >= target_low - 256 + ): + return + assert next_count < target_low, ( + "harness precondition failed: seed history would cross the target " + f"before final padding; next_count={next_count}, target_low={target_low}" + ) + _run_turn(conversation, _SEED_TEXT) + seed_turns += 1 + + +def _pad_to_band( + *, + conversation: Conversation, + llm: _RecordingLLM, + lower_exclusive: int, + upper_inclusive: int, + label: str, +) -> tuple[str, int]: + text = "Final measured turn." + count = _prepared_count(conversation=conversation, llm=llm, candidate=text) + while count <= lower_exclusive: + text += _PAD_UNIT + count = _prepared_count(conversation=conversation, llm=llm, candidate=text) + + assert count <= upper_inclusive, ( + f"harness precondition failed: {label} prepared count {count} not in " + f"({lower_exclusive}, {upper_inclusive}]" + ) + return text, count + + +@pytest.mark.parametrize(("window", "reserve"), [(16384, 4096), (32768, 8192)]) +@pytest.mark.parametrize("case", ["above-boundary", "below-boundary"]) +def test_local_context_headroom_constrains_agent_turns(window, reserve, case, tmp_path): + expected_boundary = math.floor((window - reserve) / _DIVERGENCE_FACTOR) + source_boundary = openhands._local_condenser_max_tokens(window) + # Historical shipped condenser boundary, used only as the regression band ceiling. + old_boundary = window * 11 // 16 + assert math.ceil(expected_boundary * _DIVERGENCE_FACTOR) + reserve == window + + llm = _recording_llm(window=window, reserve=reserve) + conversation = _conversation( + llm=llm, + boundary=source_boundary, + tmp_path=tmp_path / f"{window}-{case}", + ) + try: + if case == "above-boundary": + lower = expected_boundary + upper = old_boundary + else: + lower = expected_boundary - 512 + upper = expected_boundary - 1 + + _seed_history(conversation=conversation, llm=llm, target_low=lower) + final_message, prepared_count = _pad_to_band( + conversation=conversation, + llm=llm, + lower_exclusive=lower, + upper_inclusive=upper, + label=case, + ) + + start_index = len(llm.records) + conversation.send_message(final_message) + assert _prepared_count(conversation=conversation, llm=llm) == prepared_count + asyncio.run(conversation.arun()) + + observed = llm.records[start_index:] + assert observed, "expected the final turn to reach the LLM spy" + if case == "above-boundary": + assert observed[0].kind == "condenser", ( + "expected condensation before any agent-turn request; observed " + f"{observed}" + ) + assert any(record.kind == "agent" for record in observed), ( + "expected an agent turn after condensation" + ) + else: + assert observed[0].kind == "agent", ( + "expected below-boundary prompt to reach the agent without " + f"condensation; observed {observed}" + ) + assert all(record.kind == "agent" for record in observed), ( + f"unexpected condenser call below boundary; observed {observed}" + ) + + for record in llm.records: + if record.kind != "agent": + continue + assert ( + math.ceil(record.estimated_input * _DIVERGENCE_FACTOR) + + record.max_output_tokens + <= window + ), ( + "agent turn exceeded local tokenizer-divergence headroom: " + f"{record}, window={window}" + ) + finally: + conversation.close() diff --git a/tests/test_cogitate_local_condenser.py b/tests/test_cogitate_local_condenser.py index 8af99689a..adaa796d8 100644 --- a/tests/test_cogitate_local_condenser.py +++ b/tests/test_cogitate_local_condenser.py @@ -3,6 +3,8 @@ from __future__ import annotations +import math + from tests._logging_isolation import preserve_global_logging with preserve_global_logging(): @@ -76,18 +78,17 @@ def test_build_cogitate_agent_skips_condenser_for_small_endpoint_window(): def test_local_condenser_window_invariants(): - condenser_max_tokens = openhands._local_condenser_max_tokens( - LOCAL_MIN_CONTEXT_TOKENS - ) - output_reserve_tokens = openhands._local_output_reserve_tokens( - LOCAL_MIN_CONTEXT_TOKENS - ) - assert condenser_max_tokens // 2 + output_reserve_tokens < LOCAL_MIN_CONTEXT_TOKENS - assert condenser_max_tokens + output_reserve_tokens <= LOCAL_MIN_CONTEXT_TOKENS - assert condenser_max_tokens < LOCAL_MIN_CONTEXT_TOKENS - assert 11000 <= condenser_max_tokens <= 11500 - assert openhands._local_output_reserve_tokens(16384) == 4096 - assert openhands._local_condenser_max_tokens(16384) == 11264 + for window, expected_reserve in ((16384, 4096), (32768, 8192)): + boundary = openhands._local_condenser_max_tokens(window) + reserve = openhands._local_output_reserve_tokens(window) + + assert reserve == expected_reserve + assert ( + math.ceil(boundary * openhands._LOCAL_TOKENIZER_DIVERGENCE_FACTOR) + reserve + <= window + ) + assert boundary < window - reserve + assert boundary > 0 assert openhands._LOCAL_CONDENSER_KEEP_FIRST < 240 // 2 - 1 diff --git a/tests/test_openhands_errors.py b/tests/test_openhands_errors.py index a21c4ee76..6178c70bf 100644 --- a/tests/test_openhands_errors.py +++ b/tests/test_openhands_errors.py @@ -372,7 +372,9 @@ def test_run_cogitate_local_window_threaded_once( assert window_calls == [endpoint] assert llm_calls[0]["endpoint"] is endpoint assert llm_calls[0]["served_window"] == 16384 - assert agent_calls[0]["condenser_max_tokens"] == 11264 + assert agent_calls[0][ + "condenser_max_tokens" + ] == openhands._local_condenser_max_tokens(16384) llm = fake_openhands.LLM.instances[0] assert llm.max_input_tokens == 16384 assert llm.max_output_tokens == 4096