diff --git a/strix/core/execution.py b/strix/core/execution.py index 91d94498..85602b45 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -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, diff --git a/strix/core/runner.py b/strix/core/runner.py index 55784d21..048a59de 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -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: diff --git a/strix/core/sessions.py b/strix/core/sessions.py index 9286b662..b1af3a45 100644 --- a/strix/core/sessions.py +++ b/strix/core/sessions.py @@ -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] diff --git a/strix/skills/__init__.py b/strix/skills/__init__.py index bc2d3b40..1e66fae0 100644 --- a/strix/skills/__init__.py +++ b/strix/skills/__init__.py @@ -17,7 +17,7 @@ _FRONTMATTER_PATTERN = re.compile(r"^---\s*\n(?P.*?)\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, diff --git a/strix/tools/mcp/client.py b/strix/tools/mcp/client.py index 4982ffcc..b1f3f902 100644 --- a/strix/tools/mcp/client.py +++ b/strix/tools/mcp/client.py @@ -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 diff --git a/tests/test_execution.py b/tests/test_execution.py index 83d3829b..c3e1776d 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -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] diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 97d78893..56b96887 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -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 ------------- diff --git a/tests/test_runner_root_prompt.py b/tests/test_runner_root_prompt.py index 631c9271..e25d09ab 100644 --- a/tests/test_runner_root_prompt.py +++ b/tests/test_runner_root_prompt.py @@ -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 "