feat(llm): scrub images from every request to a text-only model (#1466)

* feat(llm): scrub images from every request to a text-only model

Replace images at the MCP call before any result_transform, give load_skill the
same agent_browser screenshot swap as the system prompt, and add a
call_model_input_filter that turns any image left in the input into a tool error
while chaining any filter already set.

* chore: fix the ruff and bandit failures on main
This commit is contained in:
ian-at-strix 2026-10-07 16:59:19 -04:00 • committed by GitHub
parent 7f1d148778
commit 5a30961852
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 106 additions and 16 deletions

View file

@ -4,9 +4,11 @@ from __future__ import annotations
import asyncio
import contextlib
import inspect
import logging
import uuid
from collections.abc import Callable
from dataclasses import replace
from functools import cache
from typing import TYPE_CHECKING, Any, cast
@ -30,6 +32,7 @@ from strix.core.sessions import (
enforce_image_budget,
open_agent_session,
replace_session_items,
scrub_images_from_items,
seed_initial_input,
strip_all_images_from_session,
)
@ -44,6 +47,7 @@ if TYPE_CHECKING:
from agents.lifecycle import RunHooks
from agents.memory import Session, SQLiteSession
from agents.result import RunResultBase
from agents.run_config import CallModelData, ModelInputData
from strix.core.agents import AgentCoordinator, Status
@ -120,6 +124,28 @@ async def _compact_session(
)
_TEXT_ONLY_IMAGE_TEXT = "[error: this model cannot view images; use `snapshot -i` instead]"
def _with_image_scrub(run_config: RunConfig, context: dict[str, Any]) -> RunConfig:
if context.get("supports_images", True):
return run_config
# Chain any filter already set; it sees the scrubbed input.
inner = run_config.call_model_input_filter
async def _scrub(data: CallModelData[Any]) -> ModelInputData:
model_data = replace(
data.model_data,
input=scrub_images_from_items(data.model_data.input, text=_TEXT_ONLY_IMAGE_TEXT),
)
if inner is None:
return model_data
result = inner(replace(data, model_data=model_data))
return await result if inspect.isawaitable(result) else result
return replace(run_config, call_model_input_filter=_scrub)
_MAX_TRANSIENT_MODEL_RETRIES = 5
_TRANSIENT_MODEL_RETRY_BASE_DELAY_S = 2.0
_TRANSIENT_MODEL_RETRY_MAX_DELAY_S = 90.0
@ -747,7 +773,7 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
stream = Runner.run_streamed(
agent,
input=input_data,
run_config=run_config,
run_config=_with_image_scrub(run_config, context),
context=context,
max_turns=max_turns,
session=session,

View file

@ -602,7 +602,7 @@ async def run_strix_scan(
)
if not interactive and result is not None:
final = getattr(result, "final_output", None)
# Lifecycle tools mark the root completed.
# Lifecycle tools mark the root completed.
async with coordinator._lock:
root_completed = coordinator.statuses.get(root_id) == "completed"
if not root_completed:

View file

@ -199,13 +199,13 @@ async def enforce_image_budget(session: Session, max_images: int) -> bool:
return await _rewrite_session(session, _transform)
def scrub_images_from_items(items: list[Any]) -> list[Any]:
def scrub_images_from_items(items: list[Any], *, text: str = _INHERITED_IMAGE_TEXT) -> list[Any]:
"""Return a copy of ``items`` with every image block replaced by text."""
def _scrub(obj: Any) -> Any:
if isinstance(obj, dict):
if obj.get("type") == "input_image":
return {"type": "input_text", "text": _INHERITED_IMAGE_TEXT}
return {"type": "input_text", "text": text}
return {k: _scrub(v) for k, v in obj.items()}
if isinstance(obj, list):
return [_scrub(v) for v in obj]

View file

@ -17,7 +17,7 @@ _FRONTMATTER_PATTERN = re.compile(r"^---\s*\n(?P<body>.*?)\n---\s*\n", re.DOTALL
# Only skills with `template: jinja` in their frontmatter are rendered; the rest hold
# literal `{{ }}` payloads.
_SKILL_TEMPLATES = Environment(
autoescape=False, # noqa: S701 - prompts, not HTML
autoescape=False, # noqa: S701 # nosec B701 - prompts, not HTML
undefined=StrictUndefined,
trim_blocks=True,
lstrip_blocks=True,

View file

@ -32,6 +32,7 @@ from agents.mcp import (
)
from mcp.client.stdio import stdio_client
from mcp.shared._httpx_utils import create_mcp_http_client
from mcp.types import TextContent
from strix.tools.mcp.failures import HttpStatusRecorder
from strix.tools.mcp.session import McpConnectionUnavailableError, SupervisedMcpSession
@ -183,9 +184,19 @@ def _build_server(config: McpConnectionConfig) -> BuiltMcpServer:
)
def _mcp_result_to_tool_output(
server: MCPServer, result: Any, *, supports_images: bool = True
) -> Any:
def _without_images(result: Any) -> Any:
content = [
TextContent(
type="text", text=f"[{item.mimeType} image omitted: this model cannot view images]"
)
if item.type == "image"
else item
for item in result.content
]
return result.model_copy(update={"content": content})
def _mcp_result_to_tool_output(server: MCPServer, result: Any) -> Any:
"""Serialize a ``CallToolResult`` to a tool output, mirroring the agents SDK.
This reproduces the serialization in ``agents.mcp.util.MCPUtil.invoke_mcp_tool``
@ -201,13 +212,6 @@ def _mcp_result_to_tool_output(
for item in result.content:
if item.type == "text":
outputs.append({"type": "text", "text": item.text})
elif item.type == "image" and not supports_images:
outputs.append(
{
"type": "text",
"text": f"[{item.mimeType} image omitted: this model cannot view images]",
}
)
elif item.type == "image":
outputs.append(
{"type": "image", "image_url": f"data:{item.mimeType};base64,{item.data}"}
@ -243,9 +247,11 @@ async def dispatch_mcp_call(
corrupt the content the agent receives).
"""
result = await server.call_tool(tool_name, arguments)
if not supports_images:
result = _without_images(result)
if result_transform is not None:
return result_transform(label, result.model_dump(mode="json"))
tool_output = _mcp_result_to_tool_output(server, result, supports_images=supports_images)
tool_output = _mcp_result_to_tool_output(server, result)
if getattr(result, "isError", False):
return _errored_tool_output(tool_output)
return tool_output

View file

@ -9,9 +9,11 @@ from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from agents import RunConfig
from agents.exceptions import MaxTurnsExceeded
from agents.items import MessageOutputItem
from agents.memory import SQLiteSession
from agents.run_config import CallModelData, ModelInputData
from agents.tool_context import ToolContext
from openai.types.responses import ResponseOutputMessage, ResponseOutputRefusal, ResponseOutputText
@ -1531,3 +1533,32 @@ async def test_autonomous_nudge_does_not_offer_the_user() -> None:
)
assert "wait_for_user" not in items[0]["content"]
@pytest.mark.asyncio
async def test_text_only_filter_scrubs_images_then_chains_existing_filter() -> None:
run_config = RunConfig(model="m")
assert execution._with_image_scrub(run_config, {"supports_images": True}) is run_config
seen: list[ModelInputData] = []
def _inner(data: CallModelData[Any]) -> ModelInputData:
seen.append(data.model_data)
return data.model_data
run_config = RunConfig(model="m", call_model_input_filter=_inner)
scrub: Any = execution._with_image_scrub(run_config, {"supports_images": False})
image = {"type": "input_image", "image_url": "data:image/png;base64,aGk="}
output = {"type": "function_call_output", "call_id": "c1", "output": [image]}
data = CallModelData(
model_data=ModelInputData(input=[cast("Any", output)], instructions="sys"),
agent=MagicMock(),
context=None,
)
result = await scrub.call_model_input_filter(data)
assert result is seen[0]
assert result.input[0]["output"] == [
{"type": "input_text", "text": execution._TEXT_ONLY_IMAGE_TEXT}
]
assert data.model_data.input[0]["output"] == [image]

View file

@ -917,6 +917,17 @@ async def test_call_mcp_replaces_images_for_a_text_only_model() -> None:
{"type": "text", "text": "[image/png image omitted: this model cannot view images]"},
]
transformed: list[Any] = []
registry.add(
name="pro",
server=_ImageServer("pro", [_mcp_tool("shot")]),
purpose=None,
tool_count=1,
result_transform=lambda _label, result: transformed.append(result),
)
await call_mcp.on_invoke_tool(ctx, json.dumps({"connection": "pro", "tool": "shot"}))
assert [block["type"] for block in transformed[0]["content"]] == ["text", "text"]
# --- generic MCP tools are the only MCP surface every agent gets -------------

View file

@ -13,6 +13,7 @@ from typing import Any
import httpx
import pytest
from agents import ModelSettings
from agents.tool_context import ToolContext
from openai import RateLimitError
import strix.tools.mcp as mcp_pkg
@ -24,6 +25,7 @@ from strix.core import runner
from strix.core.agents import AgentCoordinator
from strix.core.inputs import make_model_settings
from strix.runtime import session_manager
from strix.tools.load_skill.tool import load_skill
from strix.tools.mcp import BearerAuth, McpConnectionConfig, McpConnectionRequest
@ -296,6 +298,20 @@ def test_text_only_prompt_drops_screenshot_guidance() -> None:
assert "<!--" not in prompt
@pytest.mark.asyncio
async def test_text_only_load_skill_drops_screenshot_guidance() -> None:
ctx = ToolContext(
context={"supports_images": False},
tool_name="load_skill",
tool_call_id="call-1",
tool_arguments="{}",
)
out = await load_skill.on_invoke_tool(ctx, '{"skills": ["agent_browser"]}')
assert "view_image" not in out
assert "text-only model" in out
def test_scope_is_sent_as_its_own_system_message_on_cache_point_routes() -> None:
settings = make_model_settings(None, model_name="anthropic/claude-sonnet-5-5")
prompt = render_system_prompt(