mirror of
https://github.com/usestrix/strix.git
synced 2026-10-08 03:08:08 +00:00
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:
parent
7f1d148778
commit
5a30961852
8 changed files with 106 additions and 16 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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 -------------
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue