diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index 767a1a5..d23f941 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -48,6 +48,8 @@ from amplifier_core import ModuleCoordinator, TextBlock from amplifier_core.llm_errors import ContextLengthError +from ._text_estimate import estimate_messages + logger = logging.getLogger(__name__) DEFAULT_MAX_TOOL_RESULT_BYTES = 128 * 1024 @@ -1640,6 +1642,11 @@ async def get_messages_for_request( "source": meter_source, "used_tokens": token_count, "estimated_tokens": estimated_tokens, + "estimate_scope": ( + "text_only" + if estimate_messages(sticky_view).has_unmeasured_images + else "all_text" + ), "measured_tokens": self._last_measured_prompt_tokens, "foreground_usage_stale": self._foreground_usage_stale, "budget": effective_budget, @@ -1907,7 +1914,7 @@ def _measure_working_tokens( """Return (token_count, source, estimated_tokens) used to evaluate the compaction trigger this call. - `estimated_tokens` is ALWAYS the len(str)//4 heuristic over + `estimated_tokens` is ALWAYS the text-only chars/4 heuristic over `working_messages` (see _estimate_tokens) -- computed unconditionally so the estimator-vs-real-usage drift this meter exists to close is observable via `_last_token_meter_stats` regardless of mode. @@ -1917,7 +1924,7 @@ def _measure_working_tokens( - token_meter == "estimate" (default): ALWAYS `estimated_tokens`, source "estimate" -- regardless of whether a real measurement is available. This is what keeps the default mode's behavior - byte-identical to before this meter existed. + independent of prior provider measurements. - token_meter == "actual": the last real usage recorded from `llm:response` (input_tokens + cache_write_tokens -- see `_on_llm_response`), source "measured", if one has arrived this @@ -2642,9 +2649,9 @@ def _truncate_tool_wave( # total-vs-total after this first mutation replaces the # caller-supplied (already total) seed value. true_total = self._estimate_tokens(messages) + system_tokens - old_len = len(str(msg)) // 4 + old_len = self._estimate_tokens([msg]) messages[i] = self._truncate_tool_result(msg) - new_len = len(str(messages[i])) // 4 + new_len = self._estimate_tokens([messages[i]]) true_total += new_len - old_len truncated += 1 current_tokens = true_total @@ -2758,7 +2765,7 @@ def _remove_messages_with_protection( # `messages` is constant here, the base total only needs computing # once, and the removed-token total only needs an O(1) delta per # newly-removed index. - token_lens = [len(str(msg)) // 4 for msg in messages] + token_lens = [self._estimate_tokens([msg]) for msg in messages] tool_call_id_to_indices: dict[str, list[int]] = {} for idx, m in enumerate(messages): tcid = m.get("tool_call_id") @@ -3337,8 +3344,8 @@ def _derive_budget(self, token_budget: int | None, provider: Any | None) -> int: return self.max_tokens_fallback def _estimate_tokens(self, messages: list[dict[str, Any]]) -> int: - """Rough token estimation (chars / 4).""" - return sum(len(str(msg)) // 4 for msg in messages) + """Text estimate; image cost is unknown without a provider count.""" + return estimate_messages(messages).tokens class _MeasuredViewTransaction: diff --git a/amplifier_module_context_simple/_text_estimate.py b/amplifier_module_context_simple/_text_estimate.py new file mode 100644 index 0000000..8d3cd7e --- /dev/null +++ b/amplifier_module_context_simple/_text_estimate.py @@ -0,0 +1,56 @@ +"""Text-only size hints for structured messages, never image token counts. + +Keep this small helper local: runtime modules must not depend on a host or on +another context/loop implementation. Image cost requires the selected provider. +Only actual content-block positions are inspected; quoted data and tool arguments +remain ordinary text. The original messages and image payloads are never changed. +""" + +from typing import Any, NamedTuple + + +class TextEstimate(NamedTuple): + tokens: int + has_unmeasured_images: bool + + +def estimate_messages(messages: list[dict[str, Any]]) -> TextEstimate: + tokens = 0 + has_images = False + for message in messages: + content = message.get("content") + if isinstance(content, list): + content, found = _content_without_image_payloads(content) + message = {**message, "content": content} + has_images |= found + tokens += len(str(message)) // 4 + return TextEstimate(tokens, has_images) + + +def _content_without_image_payloads(content: list[Any]) -> tuple[list[Any], bool]: + result = [] + has_images = False + for block in content: + if not isinstance(block, dict): + # Core content blocks may remain typed in an in-memory view. + if callable(getattr(block, "model_dump", None)): + block = block.model_dump(exclude_none=True) + else: + result.append(block) + continue + kind = block.get("type") + if kind in {"image", "image_url", "input_image"}: + # Count the textual envelope, not encoded pixels or URL transport. + # This is explicitly a partial estimate, not a zero-cost image. + block = { + k: v + for k, v in block.items() + if k not in {"source", "image_url", "url", "data", "file_id"} + } + has_images = True + elif kind == "tool_result" and isinstance(block.get("content"), list): + nested, found = _content_without_image_payloads(block["content"]) + block = {**block, "content": nested} + has_images |= found + result.append(block) + return result, has_images diff --git a/tests/test_typed_image_estimation.py b/tests/test_typed_image_estimation.py new file mode 100644 index 0000000..e3d7e4e --- /dev/null +++ b/tests/test_typed_image_estimation.py @@ -0,0 +1,110 @@ +"""Typed pixels must not become text tokens or change the request body.""" + +import copy + +import pytest + +from amplifier_module_context_simple._text_estimate import estimate_messages + + +def test_typed_image_size_is_unknown_and_transport_size_does_not_drive_text_budget(): + small = { + "role": "user", + "content": [ + {"type": "text", "text": "Describe it"}, + { + "type": "image", + "source": {"type": "base64", "media_type": "image/png", "data": "x"}, + }, + ], + } + large = copy.deepcopy(small) + large["content"][1]["source"]["data"] = "x" * 8_500_000 + before = copy.deepcopy(large) + assert estimate_messages([large]) == estimate_messages([small]) + assert estimate_messages([large]).has_unmeasured_images + assert large == before + + +@pytest.mark.parametrize( + "content", + [ + "data:image/png;base64," + "x" * 10000, + [ + { + "type": "text", + "text": '{"type":"image","source":{"data":"' + "x" * 10000 + '"}}', + } + ], + [ + { + "type": "tool_call", + "name": "test", + "arguments": {"type": "image", "data": "x" * 10000}, + } + ], + ], +) +def test_quoted_images_and_tool_arguments_still_count_as_text(content): + message = {"role": "user", "content": content} + estimate = estimate_messages([message]) + assert estimate.tokens == len(str(message)) // 4 + assert not estimate.has_unmeasured_images + + +@pytest.mark.parametrize( + "image", + [ + { + "type": "image", + "source": {"type": "url", "url": "https://example.test/" + "x" * 10000}, + }, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64," + "x" * 10000}, + }, + {"type": "input_image", "image_url": "data:image/png;base64," + "x" * 10000}, + ], +) +def test_multiple_tool_screenshots_have_unknown_pixel_cost(image): + message = { + "role": "tool", + "content": [{"type": "tool_result", "content": [image, image]}], + } + before = copy.deepcopy(message) + estimate = estimate_messages([message]) + assert estimate.tokens < 100 + assert estimate.has_unmeasured_images + assert message == before + + +@pytest.mark.asyncio +async def test_protected_image_survives_hard_fit_while_real_text_overflow_still_fails(): + from amplifier_core import ContextLengthError + + from amplifier_module_context_simple import SimpleContextManager + + context = SimpleContextManager(max_tokens=1000, compaction_notice_enabled=False) + image = { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "x" * 8_500_000, + }, + } + await context.add_message( + {"role": "user", "content": [{"type": "text", "text": "CURRENT"}, image]} + ) + view = await context.get_messages_for_request_retaining( + retain_contents=[], hard_fit=True, token_budget=1000 + ) + assert view[0]["content"][1] == image + assert context._last_token_meter_stats["estimate_scope"] == "text_only" + assert context._last_token_meter_stats["measured_tokens"] is None + await context.add_message({"role": "user", "content": "LARGE-TEXT" * 1000}) + with pytest.raises(ContextLengthError): + await context.get_messages_for_request_retaining( + retain_contents=[], hard_fit=True, token_budget=1000 + ) + assert (await context.get_messages())[0]["content"][1] == image