Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 14 additions & 7 deletions amplifier_module_context_simple/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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:
Expand Down
56 changes: 56 additions & 0 deletions amplifier_module_context_simple/_text_estimate.py
Original file line number Diff line number Diff line change
@@ -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
110 changes: 110 additions & 0 deletions tests/test_typed_image_estimation.py
Original file line number Diff line number Diff line change
@@ -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
Loading