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
36 changes: 27 additions & 9 deletions amplifier_module_loop_streaming/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
56 changes: 56 additions & 0 deletions amplifier_module_loop_streaming/_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
146 changes: 146 additions & 0 deletions tests/test_large_image_runtime.py
Original file line number Diff line number Diff line change
@@ -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
78 changes: 78 additions & 0 deletions tests/test_typed_image_estimation.py
Original file line number Diff line number Diff line change
@@ -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
Loading