diff --git a/tests/tools/test_mcp_tool.py b/tests/tools/test_mcp_tool.py index da46348ea81e7..7783c9ec4227d 100644 --- a/tests/tools/test_mcp_tool.py +++ b/tests/tools/test_mcp_tool.py @@ -4,6 +4,7 @@ """ import asyncio +import base64 import json import os import threading @@ -227,6 +228,68 @@ def test_mcp_error_result(self): finally: _servers.pop("test_srv", None) + def test_image_only_result_is_cached_as_media_tag(self, tmp_path, monkeypatch): + from tools.mcp_tool import _make_tool_handler, _servers + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + png_bytes = b"\x89PNG\r\n\x1a\n" + b"\x00" * 32 + img_block = SimpleNamespace( + data=base64.b64encode(png_bytes).decode("ascii"), + mimeType="image/png", + ) + + mock_session = MagicMock() + mock_session.call_tool = AsyncMock( + return_value=SimpleNamespace(content=[img_block], isError=False) + ) + server = _make_mock_server("test_srv", session=mock_session) + _servers["test_srv"] = server + + try: + handler = _make_tool_handler("test_srv", "take_screenshot", 120) + with self._patch_mcp_loop(): + result = json.loads(handler({})) + + media_line = next( + line for line in result["result"].splitlines() + if line.startswith("MEDIA:") + ) + image_path = media_line.removeprefix("MEDIA:") + + assert os.path.exists(image_path) + assert image_path.startswith(str(tmp_path)) + assert image_path.endswith(".png") + finally: + _servers.pop("test_srv", None) + + def test_text_and_image_result_preserves_text_and_media_tag(self, tmp_path, monkeypatch): + from tools.mcp_tool import _make_tool_handler, _servers + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + png_bytes = b"\x89PNG\r\n\x1a\n" + b"\x00" * 32 + text_block = SimpleNamespace(text="Screenshot captured") + img_block = SimpleNamespace( + data=base64.b64encode(png_bytes).decode("ascii"), + mimeType="image/png", + ) + + mock_session = MagicMock() + mock_session.call_tool = AsyncMock( + return_value=SimpleNamespace(content=[text_block, img_block], isError=False) + ) + server = _make_mock_server("test_srv", session=mock_session) + _servers["test_srv"] = server + + try: + handler = _make_tool_handler("test_srv", "take_screenshot", 120) + with self._patch_mcp_loop(): + result = json.loads(handler({})) + + assert "Screenshot captured" in result["result"] + assert "MEDIA:" in result["result"] + finally: + _servers.pop("test_srv", None) + def test_disconnected_server(self): from tools.mcp_tool import _make_tool_handler, _servers diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index a73aa438175f9..12ba0379e8b16 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -70,10 +70,12 @@ """ import asyncio +import base64 import concurrent.futures import inspect import json import logging +import mimetypes import math import os import re @@ -321,6 +323,42 @@ def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: return resolved_command, resolved_env +def _image_extension_for_mime_type(mime_type: str) -> str: + """Return a reasonable file extension for an MCP image MIME type.""" + normalized = (mime_type or "").split(";", 1)[0].strip().lower() + if normalized in {"image/jpeg", "image/jpg"}: + return ".jpg" + return mimetypes.guess_extension(normalized) or ".png" + + +def _cache_image_content_block(block) -> str: + """Cache an MCP ImageContent block and return a MEDIA tag, or empty string.""" + data = getattr(block, "data", None) + mime_type = getattr(block, "mimeType", None) + normalized_mime = str(mime_type or "").split(";", 1)[0].strip().lower() + if data is None or not normalized_mime.startswith("image/"): + return "" + + try: + raw_bytes = base64.b64decode(data) + except (TypeError, ValueError) as exc: + logger.warning("MCP image block decode failed: %s", exc) + return "" + + try: + from gateway.platforms.base import cache_image_from_bytes + + image_path = cache_image_from_bytes( + raw_bytes, + ext=_image_extension_for_mime_type(normalized_mime), + ) + except Exception as exc: + logger.warning("MCP image block cache failed: %s", exc) + return "" + + return f"MEDIA:{image_path}" + + def _format_connect_error(exc: BaseException) -> str: """Render nested MCP connection errors into an actionable short message.""" @@ -1391,7 +1429,7 @@ async def _call(): if result.isError: error_text = "" for block in (result.content or []): - if hasattr(block, "text"): + if getattr(block, "text", ""): error_text += block.text return json.dumps({ "error": _sanitize_error( @@ -1402,8 +1440,11 @@ async def _call(): # Collect text from content blocks parts: List[str] = [] for block in (result.content or []): - if hasattr(block, "text"): + if getattr(block, "text", ""): parts.append(block.text) + media_tag = _cache_image_content_block(block) + if media_tag: + parts.append(media_tag) text_result = "\n".join(parts) if parts else "" # Combine content + structuredContent when both are present.