From dc2dc876419c77598bac1bb55b44e76c2fad5b5a Mon Sep 17 00:00:00 2001 From: Brian Krabach Date: Fri, 25 Sep 2026 03:45:11 -0700 Subject: [PATCH] fix: keep loop image estimates aligned with text-only context budgets --- amplifier_module_loop_streaming/__init__.py | 36 +++-- .../_text_estimate.py | 56 +++++++ tests/test_large_image_runtime.py | 146 ++++++++++++++++++ tests/test_typed_image_estimation.py | 78 ++++++++++ 4 files changed, 307 insertions(+), 9 deletions(-) create mode 100644 amplifier_module_loop_streaming/_text_estimate.py create mode 100644 tests/test_large_image_runtime.py create mode 100644 tests/test_typed_image_estimation.py diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 54c1a59..e66eaa8 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -42,6 +42,7 @@ from amplifier_core.llm_errors import LLMError from amplifier_core.message_models import ChatRequest, Message, ToolSpec +from ._text_estimate import estimate_messages from .steering import SteeringQueue logger = logging.getLogger(__name__) @@ -3669,7 +3670,7 @@ async def check_request_budget( "reporting a concrete budget" ) return None, None - context_estimate = sum(len(str(message)) // 4 for message in base_messages) + context_estimate = estimate_messages(base_messages).tokens budget_kwargs: dict[str, Any] = {"context_estimate": context_estimate} if accepts_named_request_options(request_budget): budget_kwargs["request_options"] = request_options @@ -3726,6 +3727,11 @@ async def check_request_budget( { "attempt": attempt, "context_estimate": context_estimate, + "context_estimate_scope": ( + "text_only" + if estimate_messages(base_messages).has_unmeasured_images + else "all_text" + ), "estimated_input_tokens": estimated, "input_limit_tokens": allowance, "context_token_budget": target, @@ -3755,7 +3761,7 @@ async def count_measured_request( mode="measured", reason="capability_missing" ) return None - context_estimate = sum(len(str(message)) // 4 for message in base_view) + context_estimate = estimate_messages(base_view).tokens kwargs: dict[str, Any] = {"context_estimate": context_estimate} if accepts_named_request_options(request_budget): kwargs["request_options"] = request_options @@ -4106,7 +4112,7 @@ async def recover_context_overflow( return None failed_call_end = time.monotonic() - context_estimate = sum(len(str(message)) // 4 for message in base_messages) + context_estimate = estimate_messages(base_messages).tokens recovery_kwargs: dict[str, Any] = {"context_estimate": context_estimate} if accepts_named_request_options(recovery): recovery_kwargs["request_options"] = request_options @@ -4825,9 +4831,15 @@ async def count_view( "orchestrator:provider_budget", { "attempt": measured_result.get("count_calls", 1) - 1, - "context_estimate": sum( - len(str(message)) // 4 - for message in measured_result.get("base_view", []) + "context_estimate": estimate_messages( + measured_result.get("base_view", []) + ).tokens, + "context_estimate_scope": ( + "text_only" + if estimate_messages( + measured_result.get("base_view", []) + ).has_unmeasured_images + else "all_text" ), "estimated_input_tokens": estimated, "input_limit_tokens": allowance, @@ -5809,9 +5821,15 @@ async def count_final_view( "orchestrator:provider_budget", { "attempt": measured_result.get("count_calls", 1) - 1, - "context_estimate": sum( - len(str(message)) // 4 - for message in measured_result.get("base_view", []) + "context_estimate": estimate_messages( + measured_result.get("base_view", []) + ).tokens, + "context_estimate_scope": ( + "text_only" + if estimate_messages( + measured_result.get("base_view", []) + ).has_unmeasured_images + else "all_text" ), "estimated_input_tokens": estimated, "input_limit_tokens": allowance, diff --git a/amplifier_module_loop_streaming/_text_estimate.py b/amplifier_module_loop_streaming/_text_estimate.py new file mode 100644 index 0000000..8d3cd7e --- /dev/null +++ b/amplifier_module_loop_streaming/_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_large_image_runtime.py b/tests/test_large_image_runtime.py new file mode 100644 index 0000000..7274741 --- /dev/null +++ b/tests/test_large_image_runtime.py @@ -0,0 +1,146 @@ +"""Real context + loop regression; no credentials or paid provider calls.""" + +import base64 +import random +import struct +import zlib +from types import SimpleNamespace + +import pytest +from amplifier_core import ContextLengthError + +pytest.importorskip("amplifier_module_context_simple") +from amplifier_module_context_simple import SimpleContextManager + +from amplifier_module_loop_streaming import StreamingOrchestrator +from tests.test_ephemeral_cache_persist_mode import ( + MockCoordinator, + MockResponse, + ScriptedHooks, +) + + +@pytest.fixture(scope="module") +def image_data(): + def chunk(kind, payload): + return ( + struct.pack(">I", len(payload)) + + kind + + payload + + struct.pack(">I", zlib.crc32(kind + payload)) + ) + + width = 1450 + pixels = random.Random(17).randbytes(width * width * 3) + rows = b"".join( + b"\0" + pixels[y * width * 3 : (y + 1) * width * 3] for y in range(width) + ) + png = ( + b"\x89PNG\r\n\x1a\n" + + chunk(b"IHDR", struct.pack(">IIBBBBB", width, width, 8, 2, 0, 0, 0)) + + chunk(b"IDAT", zlib.compress(rows)) + + chunk(b"IEND", b"") + ) + assert 6_000_000 < len(png) < 8_000_000 + return base64.b64encode(png).decode() + + +class ImageProvider: + def __init__(self, mode, streaming): + self.mode = mode + self.requests = [] + self.estimates = [] + if mode == "absent": + self.request_budget = None + if streaming: + self.stream = self._stream + + def get_info(self): + return SimpleNamespace( + capabilities=["request_budget:provider_count"], defaults={} + ) + + def request_budget(self, request, *, context_estimate, request_options=None): + self.estimates.append(context_estimate) + if self.mode == "unavailable": + return None + if self.mode == "failed": + raise RuntimeError("counter failed") + count = 150_000 if self.mode == "oversized" else 1200 + return { + "estimated_input_tokens": count, + "input_limit_tokens": 100_000, + "context_token_budget": context_estimate if count < 100_000 else 1, + "measurement": { + "kind": "provider_count", + "source": "fixture", + "input_tokens": count, + }, + } + + async def complete(self, request, **kwargs): + self.requests.append(request) + return MockResponse(text="image received") + + async def _stream(self, request, *, tools): + self.requests.append(request) + yield {"content": "image received"} + + def parse_tool_calls(self, response): + return [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("meter", ["estimate", "actual"]) +@pytest.mark.parametrize( + "mode", ["absent", "unavailable", "exact", "oversized", "failed"] +) +async def test_large_image_preserved_and_native_limits_remain_authoritative( + image_data, mode, streaming, meter +): + context = SimpleContextManager( + max_tokens=100_000, compaction_notice_enabled=False, token_meter=meter + ) + image = { + "type": "image", + "source": {"type": "base64", "media_type": "image/png", "data": image_data}, + } + await context.add_message({"role": "user", "content": [image]}) + coordinator = MockCoordinator() + coordinator.register_capability( + "context.request_retention", context.get_messages_for_request_retaining + ) + if meter == "actual": + coordinator.register_capability( + "context.measured_request_view", context.get_measured_request_view + ) + provider = ImageProvider(mode, streaming) + operation = StreamingOrchestrator({}).execute( + "Describe this image", + context, + {"test": provider}, + {}, + ScriptedHooks({}), + coordinator, + ) + if mode in {"oversized", "failed"}: + with pytest.raises(ContextLengthError if mode == "oversized" else RuntimeError): + await operation + assert not provider.requests + else: + await operation + assert len(provider.requests) == 1 + images = [ + block + for msg in provider.requests[0].messages + if isinstance(msg.content, list) + for block in msg.content + if getattr(block, "type", None) == "image" + ] + assert len(images) == 1 + assert images[0].source["data"] == image_data + assert all(estimate < 100 for estimate in provider.estimates) + assert (await context.get_messages())[0]["content"][0]["source"][ + "data" + ] == image_data diff --git a/tests/test_typed_image_estimation.py b/tests/test_typed_image_estimation.py new file mode 100644 index 0000000..41a49d1 --- /dev/null +++ b/tests/test_typed_image_estimation.py @@ -0,0 +1,78 @@ +"""Typed pixels must not become text tokens or change the request body.""" + +import copy + +import pytest + +from amplifier_module_loop_streaming._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