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
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,8 @@
from pathlib import Path
from typing import Any

from ._text_estimate import estimate_messages

logger = logging.getLogger(__name__)

# Format version for transcript.jsonl header
Expand Down Expand Up @@ -526,11 +528,16 @@ async def get_messages_for_request(
{
"usage_fraction": usage_fraction,
"token_count": conversation_tokens,
"estimate_scope": (
"text_only"
if estimate_messages(assembled).has_unmeasured_images
else "all_text"
),
"budget": available,
},
)

# Hard budget enforcement — never return more tokens than the budget allows.
# Text-budget enforcement; unmeasured images require provider validation.
# The async LLM summarization is the preferred path (produces proper summaries),
# but this is the last gate before the provider call. It operates on the
# assembled view so self._messages and self._running_token_estimate are left
Expand Down Expand Up @@ -695,12 +702,12 @@ def _calculate_budget(self, token_budget: int | None, provider: Any | None) -> i
# ── Token Estimation ──────────────────────────────────────────────────────

def _estimate_tokens(self, messages: list[dict[str, Any]]) -> int:
"""Rough token estimation (chars / 4) for a list of messages."""
return sum(len(str(msg)) // 4 for msg in messages)
"""Text estimate; typed image cost remains unknown to this heuristic."""
return estimate_messages(messages).tokens

def _estimate_tokens_single(self, message: dict[str, Any]) -> int:
"""Rough token estimation for a single message."""
return len(str(message)) // 4
return estimate_messages([message]).tokens

# ── Persistence ───────────────────────────────────────────────────────────

Expand Down
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
101 changes: 101 additions & 0 deletions modules/context-managed/tests/test_typed_image_estimation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
"""Typed pixels must not become text tokens or change the request body."""

import copy

import pytest

from amplifier_module_context_managed._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_large_image_does_not_trigger_inline_compaction():
from amplifier_module_context_managed import ManagedContextManager

context = ManagedContextManager(max_tokens=1000)
image = {
"type": "image",
"source": {
"type": "base64",
"media_type": "image/png",
"data": "x" * 8_500_000,
},
}
message = {"role": "user", "content": [{"type": "text", "text": "CURRENT"}, image]}
await context.add_message(message)
view = await context.get_messages_for_request(token_budget=1000)
assert any(m.get("content") == message["content"] for m in view)
assert context._running_token_estimate < 100
assert any(
m.get("content") == message["content"] for m in await context.get_messages()
)