diff --git a/solstone/think/providers/_image.py b/solstone/think/providers/_image.py index d188a76c2..a25e99621 100644 --- a/solstone/think/providers/_image.py +++ b/solstone/think/providers/_image.py @@ -9,27 +9,39 @@ import logging import reprlib from typing import Any -from PIL import Image +from PIL import Image, UnidentifiedImageError LOG = logging.getLogger(__name__) +# llama.cpp mtmd decodes image inputs through stb_image. mtmd documents +# jpeg/png/tga/bmp/gif support, and ggml-org/llama.cpp#15542 tells clients to +# convert unsupported formats to PNG before sending. This encoder only emits +# PNG/JPEG/GIF/WEBP, so WebP is excluded from the local/STB lane. +STB_IMAGE_MEDIA_TYPES = frozenset({"image/png", "image/jpeg", "image/gif"}) +CLOUD_IMAGE_MEDIA_TYPES = STB_IMAGE_MEDIA_TYPES | frozenset({"image/webp"}) + _PIL_FORMATS = { "PNG": ("PNG", "image/png"), "JPEG": ("JPEG", "image/jpeg"), "GIF": ("GIF", "image/gif"), "WEBP": ("WEBP", "image/webp"), } +_MEDIA_TYPE_FORMATS = { + media_type: save_format for save_format, media_type in _PIL_FORMATS.values() +} def is_image_part(part: Any) -> bool: return isinstance(part, Image.Image) or isinstance(part, bytes | bytearray) -def encode_image_part(part: Any) -> tuple[str, str]: +def encode_image_part( + part: Any, *, accepts: frozenset[str] = STB_IMAGE_MEDIA_TYPES +) -> tuple[str, str]: if isinstance(part, Image.Image): - return _encode_pil_image(part) + return _encode_pil_image(part, accepts=accepts) if isinstance(part, bytes | bytearray): - return _encode_image_bytes(part) + return _encode_image_bytes(part, accepts=accepts) raise _image_error("unsupported image part", part) @@ -41,14 +53,56 @@ def _image_error(message: str, part: Any) -> ValueError: return ValueError(f"{message}: {_part_repr(part)}") -def _encode_pil_image(image: Image.Image) -> tuple[str, str]: +def _accepted_types(accepts: frozenset[str]) -> str: + return ", ".join(sorted(accepts)) + + +def _not_sent_message( + action: str, + source_format: str, + accepts: frozenset[str], + detail: str | None = None, +) -> str: + message = ( + f"{action} image source format {source_format or 'UNKNOWN'} " + f"for accepted media types [{_accepted_types(accepts)}]; " + "abandoned without sending image" + ) + if detail: + message = f"{message} ({detail})" + return message + + +def _select_image_encoding( + source_format: str, + accepts: frozenset[str], + part: Any, +) -> tuple[str, str]: + normalized = source_format.upper() + preferred = _PIL_FORMATS.get(normalized) + if preferred is not None: + save_format, media_type = preferred + if media_type in accepts: + return save_format, media_type + if "image/png" in accepts: + return "PNG", "image/png" + raise _image_error( + _not_sent_message("cannot encode", normalized, accepts), + part, + ) + + +def _encode_pil_image( + image: Image.Image, *, accepts: frozenset[str] +) -> tuple[str, str]: if image.width <= 0 or image.height <= 0: raise _image_error("cannot encode zero-size image part", image) source_format = (image.format or "").upper() - save_format, media_type = _PIL_FORMATS.get( + save_format, media_type = _select_image_encoding( source_format, - ("PNG", "image/png"), + accepts, + image, ) prepared = _prepare_pil_image(image, save_format) @@ -89,10 +143,38 @@ def _prepare_pil_image(image: Image.Image, save_format: str) -> Image.Image: raise _image_error(f"unsupported PIL format: {save_format}", image) -def _encode_image_bytes(part: bytes | bytearray) -> tuple[str, str]: +def _encode_image_bytes( + part: bytes | bytearray, *, accepts: frozenset[str] +) -> tuple[str, str]: data = bytes(part) media_type = _sniff_image_media_type(data, part) - return media_type, base64.b64encode(data).decode("ascii") + source_format = _MEDIA_TYPE_FORMATS[media_type] + _, accepted_media_type = _select_image_encoding(source_format, accepts, part) + if accepted_media_type == media_type: + return media_type, base64.b64encode(data).decode("ascii") + image = _decode_image_bytes(data, source_format, accepts, part) + return _encode_pil_image(image, accepts=accepts) + + +def _decode_image_bytes( + data: bytes, + source_format: str, + accepts: frozenset[str], + part: bytes | bytearray, +) -> Image.Image: + try: + image = Image.open(io.BytesIO(data)) + image.load() + except ( + UnidentifiedImageError, + OSError, + Image.DecompressionBombError, + ) as exc: + raise _image_error( + _not_sent_message("cannot decode", source_format, accepts, str(exc)), + part, + ) from exc + return image def _sniff_image_media_type(data: bytes, part: bytes | bytearray) -> str: @@ -107,4 +189,9 @@ def _sniff_image_media_type(data: bytes, part: bytes | bytearray) -> str: raise _image_error("unrecognized image bytes", part) -__all__ = ["is_image_part", "encode_image_part"] +__all__ = [ + "CLOUD_IMAGE_MEDIA_TYPES", + "STB_IMAGE_MEDIA_TYPES", + "is_image_part", + "encode_image_part", +] diff --git a/solstone/think/providers/local.py b/solstone/think/providers/local.py index ace7d5323..22518cf03 100644 --- a/solstone/think/providers/local.py +++ b/solstone/think/providers/local.py @@ -548,7 +548,6 @@ def run_generate( endpoint = resolve_local_endpoint() # Validate the requested logical id; served id comes from the server. normalize_model_id(model) - messages = _build_messages(contents, system_instruction) if endpoint.is_bundled: from solstone.think.providers import local_server from solstone.think.providers.local_admission import ( @@ -645,6 +644,7 @@ def run_generate( ) raise + messages = _build_messages(contents, system_instruction) body = _build_request_body( endpoint.served_model_id, messages, @@ -721,11 +721,11 @@ async def run_agenerate( raise TypeError(f"Unsupported local generate options: {unknown}") endpoint = resolve_local_endpoint() normalize_model_id(model) - messages = _build_messages(contents, system_instruction) import httpx if not endpoint.is_bundled: + messages = _build_messages(contents, system_instruction) from solstone.think.providers.local_admission import ( LocalAdmissionTimeout, acquire_local_slot_async, diff --git a/solstone/think/providers/openhands.py b/solstone/think/providers/openhands.py index c9e8fe463..b8904a783 100644 --- a/solstone/think/providers/openhands.py +++ b/solstone/think/providers/openhands.py @@ -349,9 +349,16 @@ def _build_generate_llm( def _data_url(part: Any) -> str: - from solstone.think.providers._image import encode_image_part + from solstone.think.providers._image import ( + CLOUD_IMAGE_MEDIA_TYPES, + encode_image_part, + ) - media_type, payload = encode_image_part(part) + # Use the permissive cloud image set here even though this module is not cloud-only: + # local.run_cogitate delegates to openhands.run_cogitate. This is safe today only + # because no cogitate tool returns image content; if that changes, accepts must be + # threaded from the calling provider instead of hardcoded here. + media_type, payload = encode_image_part(part, accepts=CLOUD_IMAGE_MEDIA_TYPES) return f"data:{media_type};base64,{payload}" diff --git a/tests/test_importer_images.py b/tests/test_importer_images.py index 7639fe731..9b35bd1cc 100644 --- a/tests/test_importer_images.py +++ b/tests/test_importer_images.py @@ -111,6 +111,42 @@ def test_process_vision_failure_propagates_before_success_entry(tmp_path, monkey assert not list((tmp_path / "chronicle").glob("**/import.image")) +def test_process_real_encoder_failure_leaves_no_image_artifacts(tmp_path, monkeypatch): + mod = __import__("solstone.think.importers.images", fromlist=["importer"]) + _configure_journal(tmp_path, monkeypatch) + image_path = tmp_path / "cmyk.jpg" + Image.new("CMYK", (8, 8)).save(image_path, "JPEG") + import_id = "20260115_120000" + + from solstone.think import models + from solstone.think.models import LOCAL_MODEL + from solstone.think.providers import local as local_provider + from solstone.think.providers.local_endpoint import LocalEndpoint + + monkeypatch.setattr( + models, + "resolve_provider", + lambda _agent_type: ("local", LOCAL_MODEL), + ) + monkeypatch.setattr( + local_provider, + "resolve_local_endpoint", + lambda: LocalEndpoint( + "http://127.0.0.1:9", + "served-model", + None, + is_bundled=False, + ), + ) + + with pytest.raises(ValueError, match="unsupported PIL mode for JPEG: CMYK"): + mod.importer.process(image_path, tmp_path, import_id=import_id) + + assert not list((tmp_path / "chronicle").glob("**/import.image")) + assert not list((tmp_path / "chronicle").glob("**/image_transcript.md")) + assert not (tmp_path / "imports" / import_id / "content_manifest.jsonl").exists() + + def test_registry_entry(): assert FILE_IMPORTER_REGISTRY["image"] == "solstone.think.importers.images" importer = get_file_importer("image") diff --git a/tests/test_local.py b/tests/test_local.py index bb31945fe..e6effae4f 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -919,6 +919,66 @@ def _patch_bundled_server(monkeypatch): ) +def test_run_generate_bundled_encodes_image_once(monkeypatch): + provider = _provider() + monkeypatch.setattr(provider, "resolve_local_endpoint", _bundled_endpoint) + _patch_bundled_server(monkeypatch) + png = b"\x89PNG\r\n\x1a\npayload" + calls = [] + captured = {} + + def count_encode(part): + calls.append(part) + return "image/png", base64.b64encode(b"encoded").decode("ascii") + + class TokenResponse: + text = "" + + def raise_for_status(self): + return None + + def json(self): + return {"tokens": [1]} + + class ChatResponse: + text = "" + + def raise_for_status(self): + return None + + def json(self): + return { + "model": LOCAL_MODEL, + "choices": [ + { + "message": {"content": "ok"}, + "finish_reason": "stop", + } + ], + } + + def fake_post(url, json, timeout): + del timeout + if url.endswith("/tokenize"): + return TokenResponse() + if url.endswith("/v1/chat/completions"): + captured["body"] = json + return ChatResponse() + raise AssertionError(f"unexpected local provider URL: {url}") + + import httpx + + monkeypatch.setattr(provider, "encode_image_part", count_encode) + monkeypatch.setattr(httpx, "post", fake_post) + + result = provider.run_generate(["look", png], model=LOCAL_MODEL) + + assert result["text"] == "ok" + assert "body" in captured + assert len(calls) == 1 + assert calls == [png] + + def test_run_generate_byo_posts_to_normalized_endpoint_and_skips_connect(monkeypatch): provider = _provider() monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) diff --git a/tests/test_provider_image.py b/tests/test_provider_image.py index c5a4357f7..725dd90fe 100644 --- a/tests/test_provider_image.py +++ b/tests/test_provider_image.py @@ -7,7 +7,12 @@ import io import pytest from PIL import Image -from solstone.think.providers._image import encode_image_part, is_image_part +from solstone.think.providers._image import ( + CLOUD_IMAGE_MEDIA_TYPES, + STB_IMAGE_MEDIA_TYPES, + encode_image_part, + is_image_part, +) def _png_bytes(size: tuple[int, int] = (4, 3)) -> bytes: @@ -17,6 +22,19 @@ def _png_bytes(size: tuple[int, int] = (4, 3)) -> bytes: return buf.getvalue() +def _image_bytes(source_format: str, size: tuple[int, int] = (4, 3)) -> bytes: + image = Image.new("RGB", size, color="red") + buf = io.BytesIO() + image.save(buf, format=source_format) + return buf.getvalue() + + +def _pil_image(source_format: str, size: tuple[int, int] = (4, 3)) -> Image.Image: + image = Image.open(io.BytesIO(_image_bytes(source_format, size))) + image.load() + return image + + def _decoded_image(b64: str) -> Image.Image: return Image.open(io.BytesIO(base64.b64decode(b64))) @@ -86,3 +104,145 @@ def test_cmyk_image_raises_with_part_type_and_repr(): message = str(exc_info.value) assert "Image" in message assert "CMYK" in message + + +def test_default_webp_pil_transcodes_to_png(): + image = _pil_image("WEBP", (5, 4)) + + media_type, b64 = encode_image_part(image) + + decoded = _decoded_image(b64) + assert media_type == "image/png" + assert decoded.size == image.size + assert decoded.format == "PNG" + + +def test_stb_webp_pil_transcodes_to_png(): + image = _pil_image("WEBP", (5, 4)) + + media_type, b64 = encode_image_part(image, accepts=STB_IMAGE_MEDIA_TYPES) + + decoded = _decoded_image(b64) + assert media_type == "image/png" + assert decoded.size == image.size + assert decoded.format == "PNG" + + +def test_cloud_webp_pil_stays_webp(): + image = _pil_image("WEBP", (5, 4)) + + media_type, b64 = encode_image_part(image, accepts=CLOUD_IMAGE_MEDIA_TYPES) + + decoded = _decoded_image(b64) + assert media_type == "image/webp" + assert decoded.size == image.size + assert decoded.format == "WEBP" + + +def test_webp_bytes_stb_transcodes_and_cloud_preserves_bytes(): + data = _image_bytes("WEBP", (5, 4)) + + stb_media_type, stb_b64 = encode_image_part(data, accepts=STB_IMAGE_MEDIA_TYPES) + cloud_media_type, cloud_b64 = encode_image_part( + data, + accepts=CLOUD_IMAGE_MEDIA_TYPES, + ) + + stb_decoded = _decoded_image(stb_b64) + assert stb_media_type == "image/png" + assert stb_decoded.size == (5, 4) + assert stb_decoded.format == "PNG" + assert cloud_media_type == "image/webp" + assert base64.b64decode(cloud_b64) == data + + +@pytest.mark.parametrize( + ("source_format", "media_type"), + [("PNG", "image/png"), ("JPEG", "image/jpeg")], +) +@pytest.mark.parametrize("accepts", [STB_IMAGE_MEDIA_TYPES, CLOUD_IMAGE_MEDIA_TYPES]) +def test_png_and_jpeg_bytes_preserved_byte_for_byte_under_accept_sets( + source_format: str, + media_type: str, + accepts: frozenset[str], +): + data = _image_bytes(source_format, (5, 4)) + + encoded_media_type, b64 = encode_image_part(data, accepts=accepts) + + assert encoded_media_type == media_type + assert base64.b64decode(b64) == data + + +def test_tiff_pil_transcodes_to_png_and_tiff_bytes_still_raise(): + image = _pil_image("TIFF", (5, 4)) + + for accepts in [STB_IMAGE_MEDIA_TYPES, CLOUD_IMAGE_MEDIA_TYPES]: + media_type, b64 = encode_image_part(image, accepts=accepts) + decoded = _decoded_image(b64) + assert media_type == "image/png" + assert decoded.size == image.size + assert decoded.format == "PNG" + + with pytest.raises(ValueError, match="unrecognized image bytes"): + encode_image_part(_image_bytes("TIFF", (5, 4))) + + +@pytest.mark.parametrize("mode", ["CMYK", "I;16"]) +def test_unsupported_pil_modes_still_raise(mode: str): + image = Image.new(mode, (2, 2)) + + with pytest.raises(ValueError) as exc_info: + encode_image_part(image) + + message = str(exc_info.value) + assert f"unsupported PIL mode for PNG: {mode}" in message + + +def test_accepts_constants_include_png_and_successes_stay_inside_accepts(): + assert "image/png" in STB_IMAGE_MEDIA_TYPES + assert "image/png" in CLOUD_IMAGE_MEDIA_TYPES + parts = [ + _image_bytes("PNG", (3, 2)), + _image_bytes("JPEG", (3, 2)), + _image_bytes("GIF", (3, 2)), + _image_bytes("WEBP", (3, 2)), + _pil_image("PNG", (3, 2)), + _pil_image("JPEG", (3, 2)), + _pil_image("GIF", (3, 2)), + _pil_image("WEBP", (3, 2)), + _pil_image("TIFF", (3, 2)), + ] + + for accepts in [ + STB_IMAGE_MEDIA_TYPES, + CLOUD_IMAGE_MEDIA_TYPES, + frozenset({"image/png"}), + ]: + for part in parts: + media_type, _ = encode_image_part(part, accepts=accepts) + assert media_type in accepts + + +def test_selection_failure_message_names_format_accepts_and_not_sent(): + image = _pil_image("WEBP", (3, 2)) + + with pytest.raises(ValueError) as exc_info: + encode_image_part(image, accepts=frozenset({"image/jpeg"})) + + message = str(exc_info.value) + assert "cannot encode image source format WEBP" in message + assert "accepted media types [image/jpeg]" in message + assert "abandoned without sending image" in message + + +def test_corrupt_webp_decode_failure_message_names_format_accepts_and_not_sent(): + corrupt_webp = b"RIFF\x00\x00\x00\x00WEBPnot-a-real-webp" + + with pytest.raises(ValueError) as exc_info: + encode_image_part(corrupt_webp) + + message = str(exc_info.value) + assert "cannot decode image source format WEBP" in message + assert "accepted media types [image/gif, image/jpeg, image/png]" in message + assert "abandoned without sending image" in message diff --git a/tests/test_provider_vision_parity.py b/tests/test_provider_vision_parity.py new file mode 100644 index 000000000..e9204306a --- /dev/null +++ b/tests/test_provider_vision_parity.py @@ -0,0 +1,234 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import base64 +import importlib +import inspect +import io +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any + +import pytest +from PIL import Image + +from solstone.think.models import LOCAL_MODEL +from solstone.think.providers import PROVIDER_REGISTRY, local, openhands +from tests.openhands_fakes import install_fake_openhands +from tests.test_local import _bundled_endpoint, _patch_bundled_server, _provider + + +@dataclass(frozen=True) +class _VisionLane: + chokepoint: Callable[..., Any] + driver: str + + +@dataclass(frozen=True) +class _ImageCase: + source_format: str + representation: str + make_part: Callable[[], Any] + size: tuple[int, int] + + +_LANES = { + "anthropic": _VisionLane(openhands._message_content, "openhands"), + "google": _VisionLane(openhands._message_content, "openhands"), + "local": _VisionLane(local.run_generate, "local"), + "openai": _VisionLane(openhands._message_content, "openhands"), +} + +_EXPECTED_MEDIA_TYPES = { + "anthropic": { + "GIF": "image/gif", + "JPEG": "image/jpeg", + "PNG": "image/png", + "TIFF": "image/png", + "WEBP": "image/webp", + }, + "google": { + "GIF": "image/gif", + "JPEG": "image/jpeg", + "PNG": "image/png", + "TIFF": "image/png", + "WEBP": "image/webp", + }, + "local": { + "GIF": "image/gif", + "JPEG": "image/jpeg", + "PNG": "image/png", + "TIFF": "image/png", + "WEBP": "image/png", + }, + "openai": { + "GIF": "image/gif", + "JPEG": "image/jpeg", + "PNG": "image/png", + "TIFF": "image/png", + "WEBP": "image/webp", + }, +} + +_FORMAT_BY_MEDIA_TYPE = { + "image/gif": "GIF", + "image/jpeg": "JPEG", + "image/png": "PNG", + "image/webp": "WEBP", +} + +_IMAGE_SIZE = (3, 2) + + +@pytest.fixture(autouse=True) +def _isolate_local_admission(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: + from solstone.think.providers import local_admission + + monkeypatch.setattr( + local_admission, + "_admission_dir", + lambda: tmp_path / "local-inference-admission", + ) + monkeypatch.setattr(local_admission, "record_local_inference", lambda _record: None) + + +def _source_bytes(source_format: str) -> bytes: + image = Image.new("RGB", _IMAGE_SIZE, "red") + buf = io.BytesIO() + image.save(buf, format=source_format) + return buf.getvalue() + + +def _source_pil(source_format: str) -> Image.Image: + image = Image.open(io.BytesIO(_source_bytes(source_format))) + image.load() + return image + + +_IMAGE_CASES = [ + _ImageCase( + source_format=source_format, + representation=representation, + make_part=make_part, + size=_IMAGE_SIZE, + ) + for source_format in ("PNG", "JPEG", "GIF", "WEBP") + for representation, make_part in ( + ("bytes", lambda source_format=source_format: _source_bytes(source_format)), + ("pil", lambda source_format=source_format: _source_pil(source_format)), + ) +] + [ + _ImageCase( + source_format="TIFF", + representation="pil", + make_part=lambda: _source_pil("TIFF"), + size=_IMAGE_SIZE, + ) +] + + +def _case_id(case: _ImageCase) -> str: + return f"{case.representation}-{case.source_format}" + + +def test_provider_vision_lane_table_matches_registry() -> None: + assert set(_LANES) == set(PROVIDER_REGISTRY) + for name, lane in _LANES.items(): + assert inspect.getmodule(lane.chokepoint) is importlib.import_module( + PROVIDER_REGISTRY[name] + ) + + +def _local_data_url(monkeypatch: pytest.MonkeyPatch, part: Any) -> str: + provider = _provider() + monkeypatch.setattr(provider, "resolve_local_endpoint", _bundled_endpoint) + _patch_bundled_server(monkeypatch) + captured: dict[str, Any] = {} + + class TokenResponse: + text = "" + + def raise_for_status(self) -> None: + return None + + def json(self) -> dict[str, Any]: + return {"tokens": [1]} + + class ChatResponse: + text = "" + + def raise_for_status(self) -> None: + return None + + def json(self) -> dict[str, Any]: + return { + "model": LOCAL_MODEL, + "choices": [ + { + "message": {"content": "ok"}, + "finish_reason": "stop", + } + ], + } + + def fake_post(url: str, json: dict[str, Any], timeout: float) -> Any: + del timeout + if url.endswith("/tokenize"): + return TokenResponse() + if url.endswith("/v1/chat/completions"): + captured["body"] = json + return ChatResponse() + raise AssertionError(f"unexpected local provider URL: {url}") + + import httpx + + monkeypatch.setattr(httpx, "post", fake_post) + provider.run_generate(["look", part], model=LOCAL_MODEL) + + assert "body" in captured + content = captured["body"]["messages"][-1]["content"] + return next( + item["image_url"]["url"] for item in content if item.get("type") == "image_url" + ) + + +def _openhands_data_url(monkeypatch: pytest.MonkeyPatch, part: Any) -> str: + fake_openhands = install_fake_openhands(monkeypatch) + blocks = openhands._message_content(["look", part]) + image_blocks = [ + block for block in blocks if isinstance(block, fake_openhands.ImageContent) + ] + assert len(image_blocks) == 1 + return image_blocks[0].image_urls[0] + + +def _data_url_payload(url: str) -> tuple[str, bytes]: + prefix, payload = url.split(",", 1) + media_type = prefix.removeprefix("data:").removesuffix(";base64") + return media_type, base64.b64decode(payload) + + +@pytest.mark.parametrize("name", sorted(_LANES)) +@pytest.mark.parametrize("image_case", _IMAGE_CASES, ids=_case_id) +def test_provider_vision_media_matrix( + monkeypatch: pytest.MonkeyPatch, + name: str, + image_case: _ImageCase, +) -> None: + part = image_case.make_part() + lane = _LANES[name] + if lane.driver == "local": + url = _local_data_url(monkeypatch, part) + else: + url = _openhands_data_url(monkeypatch, part) + + media_type, payload = _data_url_payload(url) + expected_media_type = _EXPECTED_MEDIA_TYPES[name][image_case.source_format] + + assert media_type == expected_media_type + decoded = Image.open(io.BytesIO(payload)) + decoded.load() + assert decoded.size == image_case.size + assert decoded.format == _FORMAT_BY_MEDIA_TYPE[expected_media_type]