Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 63 additions & 0 deletions tests/tools/test_mcp_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
"""

import asyncio
import base64
import json
import os
import threading
Expand Down Expand Up @@ -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

Expand Down
45 changes: 43 additions & 2 deletions tools/mcp_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""

Expand Down Expand Up @@ -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(
Expand All @@ -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.
Expand Down