From 5b32ffb77f3238e3fd84b42e8ba976cc27824b8a Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 11 Apr 2025 10:51:07 -0400 Subject: [PATCH] Fix bug with tools that return lists --- src/fastmcp/server/server.py | 27 ++++++++++++-- tests/server/test_server.py | 72 +++++++++++------------------------- 2 files changed, 44 insertions(+), 55 deletions(-) diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 80c6b9dd7..e2c26efcc 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -8,7 +8,6 @@ from contextlib import ( AbstractAsyncContextManager, asynccontextmanager, ) -from itertools import chain from typing import TYPE_CHECKING, Any, Generic, Literal import anyio @@ -618,7 +617,8 @@ class FastMCP(Generic[LifespanResultT]): def _convert_to_content( result: Any, -) -> Sequence[TextContent | ImageContent | EmbeddedResource]: + _process_as_single_item: bool = False, +) -> list[TextContent | ImageContent | EmbeddedResource]: """Convert a result to a sequence of content objects.""" if result is None: return [] @@ -629,8 +629,27 @@ def _convert_to_content( if isinstance(result, Image): return [result.to_image_content()] - if isinstance(result, list | tuple): - return list(chain.from_iterable(_convert_to_content(item) for item in result)) # type: ignore[reportUnknownVariableType] + if isinstance(result, list | tuple) and not _process_as_single_item: + # if the result is a list, then it could either be a list of MCP types, + # or a "regular" list that the tool is returning, or a mix of both. + # + # so we extract all the MCP types / images and convert them as individual content elements, + # and aggregate the rest as a single content element + + mcp_types = [] + other_content = [] + + for item in result: + if isinstance(item, (TextContent, ImageContent, EmbeddedResource, Image)): + mcp_types.append(_convert_to_content(item)[0]) + else: + other_content.append(item) + if other_content: + other_content = _convert_to_content( + other_content, _process_as_single_item=True + ) + + return other_content + mcp_types if not isinstance(result, str): try: diff --git a/tests/server/test_server.py b/tests/server/test_server.py index d79ce8f6d..8d2df726c 100644 --- a/tests/server/test_server.py +++ b/tests/server/test_server.py @@ -1,4 +1,5 @@ import base64 +import json from pathlib import Path from typing import TYPE_CHECKING @@ -25,13 +26,11 @@ if TYPE_CHECKING: class TestServer: - @pytest.mark.anyio async def test_create_server(self): mcp = FastMCP(instructions="Server instructions") assert mcp.name == "FastMCP" assert mcp.instructions == "Server instructions" - @pytest.mark.anyio async def test_non_ascii_description(self): """Test that FastMCP handles non-ASCII characters in descriptions correctly""" mcp = FastMCP() @@ -59,7 +58,6 @@ class TestServer: assert isinstance(content, TextContent) assert "¡Hola, 世界! 👋" == content.text - @pytest.mark.anyio async def test_add_tool_decorator(self): mcp = FastMCP() @@ -69,7 +67,6 @@ class TestServer: assert len(mcp._tool_manager.list_tools()) == 1 - @pytest.mark.anyio async def test_add_tool_decorator_incorrect_usage(self): mcp = FastMCP() @@ -79,7 +76,6 @@ class TestServer: def add(x: int, y: int) -> int: return x + y - @pytest.mark.anyio async def test_add_resource_decorator(self): mcp = FastMCP() @@ -89,7 +85,6 @@ class TestServer: assert len(mcp._resource_manager._templates) == 1 - @pytest.mark.anyio async def test_add_resource_decorator_incorrect_usage(self): mcp = FastMCP() @@ -106,6 +101,10 @@ def tool_fn(x: int, y: int) -> int: return x + y +def tool_fn_list() -> list[str | int]: + return ["x", 2] + + def error_tool_fn() -> None: raise ValueError("Test error") @@ -122,14 +121,12 @@ def mixed_content_tool_fn() -> list[TextContent | ImageContent]: class TestServerTools: - @pytest.mark.anyio async def test_add_tool(self): mcp = FastMCP() mcp.add_tool(tool_fn) mcp.add_tool(tool_fn) assert len(mcp._tool_manager.list_tools()) == 1 - @pytest.mark.anyio async def test_list_tools(self): mcp = FastMCP() mcp.add_tool(tool_fn) @@ -137,7 +134,6 @@ class TestServerTools: tools = await client.list_tools() assert len(tools.tools) == 1 - @pytest.mark.anyio async def test_call_tool(self): mcp = FastMCP() mcp.add_tool(tool_fn) @@ -146,7 +142,6 @@ class TestServerTools: assert not hasattr(result, "error") assert len(result.content) > 0 - @pytest.mark.anyio async def test_tool_exception_handling(self): mcp = FastMCP() mcp.add_tool(error_tool_fn) @@ -158,7 +153,6 @@ class TestServerTools: assert "Test error" in content.text assert result.isError is True - @pytest.mark.anyio async def test_tool_error_handling(self): mcp = FastMCP() mcp.add_tool(error_tool_fn) @@ -170,7 +164,6 @@ class TestServerTools: assert "Test error" in content.text assert result.isError is True - @pytest.mark.anyio async def test_tool_error_details(self): """Test that exception details are properly formatted in the response""" mcp = FastMCP() @@ -183,7 +176,6 @@ class TestServerTools: assert "Test error" in content.text assert result.isError is True - @pytest.mark.anyio async def test_tool_return_value_conversion(self): mcp = FastMCP() mcp.add_tool(tool_fn) @@ -194,7 +186,16 @@ class TestServerTools: assert isinstance(content, TextContent) assert content.text == "3" - @pytest.mark.anyio + async def test_tool_returns_list(self): + mcp = FastMCP() + mcp.add_tool(tool_fn_list) + async with client_session(mcp._mcp_server) as client: + result = await client.call_tool("tool_fn_list", {}) + assert len(result.content) == 1 + content = result.content[0] + assert isinstance(content, TextContent) + assert json.loads(content.text) == ["x", 2] + async def test_tool_image_helper(self, tmp_path: Path): # Create a test image image_path = tmp_path / "test.png" @@ -213,12 +214,12 @@ class TestServerTools: decoded = base64.b64decode(content.data) assert decoded == b"fake png data" - @pytest.mark.anyio async def test_tool_mixed_content(self): mcp = FastMCP() mcp.add_tool(mixed_content_tool_fn) async with client_session(mcp._mcp_server) as client: result = await client.call_tool("mixed_content_tool_fn", {}) + assert len(result.content) == 2 content1 = result.content[0] content2 = result.content[1] @@ -228,10 +229,9 @@ class TestServerTools: assert content2.mimeType == "image/png" assert content2.data == "abc" - @pytest.mark.anyio async def test_tool_mixed_list_with_image(self, tmp_path: Path): """Test that lists containing Image objects and other types are handled - correctly""" + correctly. Note that the non-MCP content will be grouped together.""" # Create a test image image_path = tmp_path / "test.png" image_path.write_bytes(b"test image data") @@ -248,24 +248,20 @@ class TestServerTools: mcp.add_tool(mixed_list_fn) async with client_session(mcp._mcp_server) as client: result = await client.call_tool("mixed_list_fn", {}) - assert len(result.content) == 4 + assert len(result.content) == 3 # Check text conversion content1 = result.content[0] assert isinstance(content1, TextContent) - assert content1.text == "text message" + assert json.loads(content1.text) == ["text message", {"key": "value"}] # Check image conversion content2 = result.content[1] assert isinstance(content2, ImageContent) assert content2.mimeType == "image/png" assert base64.b64decode(content2.data) == b"test image data" - # Check dict conversion + # Check direct TextContent content3 = result.content[2] assert isinstance(content3, TextContent) - assert '"key": "value"' in content3.text - # Check direct TextContent - content4 = result.content[3] - assert isinstance(content4, TextContent) - assert content4.text == "direct content" + assert content3.text == "direct content" async def test_parameter_descriptions(self): mcp = FastMCP("Test Server") @@ -291,7 +287,6 @@ class TestServerTools: class TestServerResources: - @pytest.mark.anyio async def test_text_resource(self): mcp = FastMCP() @@ -308,7 +303,6 @@ class TestServerResources: assert isinstance(result.contents[0], TextResourceContents) assert result.contents[0].text == "Hello, world!" - @pytest.mark.anyio async def test_binary_resource(self): mcp = FastMCP() @@ -328,7 +322,6 @@ class TestServerResources: assert isinstance(result.contents[0], BlobResourceContents) assert result.contents[0].blob == base64.b64encode(b"Binary data").decode() - @pytest.mark.anyio async def test_file_resource_text(self, tmp_path: Path): mcp = FastMCP() @@ -346,7 +339,6 @@ class TestServerResources: assert isinstance(result.contents[0], TextResourceContents) assert result.contents[0].text == "Hello from file!" - @pytest.mark.anyio async def test_file_resource_binary(self, tmp_path: Path): mcp = FastMCP() @@ -372,7 +364,6 @@ class TestServerResources: class TestServerResourceTemplates: - @pytest.mark.anyio async def test_resource_with_params(self): """Test that a resource with function parameters raises an error if the URI parameters don't match""" @@ -384,7 +375,6 @@ class TestServerResourceTemplates: def get_data_fn(param: str) -> str: return f"Data: {param}" - @pytest.mark.anyio async def test_resource_with_uri_params(self): """Test that a resource with URI parameters is automatically a template""" mcp = FastMCP() @@ -395,7 +385,6 @@ class TestServerResourceTemplates: def get_data() -> str: return "Data" - @pytest.mark.anyio async def test_resource_with_untyped_params(self): """Test that a resource with untyped parameters raises an error""" mcp = FastMCP() @@ -404,7 +393,6 @@ class TestServerResourceTemplates: def get_data(param) -> str: return "Data" - @pytest.mark.anyio async def test_resource_matching_params(self): """Test that a resource with matching URI and function parameters works""" mcp = FastMCP() @@ -418,7 +406,6 @@ class TestServerResourceTemplates: assert isinstance(result.contents[0], TextResourceContents) assert result.contents[0].text == "Data for test" - @pytest.mark.anyio async def test_resource_mismatched_params(self): """Test that mismatched parameters raise an error""" mcp = FastMCP() @@ -429,7 +416,6 @@ class TestServerResourceTemplates: def get_data(user: str) -> str: return f"Data for {user}" - @pytest.mark.anyio async def test_resource_multiple_params(self): """Test that multiple parameters work correctly""" mcp = FastMCP() @@ -445,7 +431,6 @@ class TestServerResourceTemplates: assert isinstance(result.contents[0], TextResourceContents) assert result.contents[0].text == "Data for cursor/fastmcp" - @pytest.mark.anyio async def test_resource_multiple_mismatched_params(self): """Test that mismatched parameters raise an error""" mcp = FastMCP() @@ -468,7 +453,6 @@ class TestServerResourceTemplates: assert isinstance(result.contents[0], TextResourceContents) assert result.contents[0].text == "Static data" - @pytest.mark.anyio async def test_template_to_resource_conversion(self): """Test that templates are properly converted to resources when accessed""" mcp = FastMCP() @@ -491,7 +475,6 @@ class TestServerResourceTemplates: class TestContextInjection: """Test context injection in tools.""" - @pytest.mark.anyio async def test_context_detection(self): """Test that context parameters are properly detected.""" mcp = FastMCP() @@ -502,7 +485,6 @@ class TestContextInjection: tool = mcp._tool_manager.add_tool(tool_with_context) assert tool.context_kwarg == "ctx" - @pytest.mark.anyio async def test_context_injection(self): """Test that context is properly injected into tool calls.""" mcp = FastMCP() @@ -520,7 +502,6 @@ class TestContextInjection: assert "Request" in content.text assert "42" in content.text - @pytest.mark.anyio async def test_async_context(self): """Test that context works in async functions.""" mcp = FastMCP() @@ -538,7 +519,6 @@ class TestContextInjection: assert "Async request" in content.text assert "42" in content.text - @pytest.mark.anyio async def test_context_logging(self): from unittest.mock import patch @@ -576,7 +556,6 @@ class TestContextInjection: level="error", data="Error message", logger=None ) - @pytest.mark.anyio async def test_optional_context(self): """Test that context is optional.""" mcp = FastMCP() @@ -592,7 +571,6 @@ class TestContextInjection: assert isinstance(content, TextContent) assert content.text == "42" - @pytest.mark.anyio async def test_context_resource_access(self): """Test that context can access resources.""" mcp = FastMCP() @@ -620,7 +598,6 @@ class TestContextInjection: class TestServerPrompts: """Test prompt functionality in FastMCP server.""" - @pytest.mark.anyio async def test_prompt_decorator(self): """Test that the prompt decorator registers prompts correctly.""" mcp = FastMCP() @@ -637,7 +614,6 @@ class TestServerPrompts: assert isinstance(content[0].content, TextContent) assert content[0].content.text == "Hello, world!" - @pytest.mark.anyio async def test_prompt_decorator_with_name(self): """Test prompt decorator with custom name.""" mcp = FastMCP() @@ -653,7 +629,6 @@ class TestServerPrompts: assert isinstance(content[0].content, TextContent) assert content[0].content.text == "Hello, world!" - @pytest.mark.anyio async def test_prompt_decorator_with_description(self): """Test prompt decorator with custom description.""" mcp = FastMCP() @@ -678,7 +653,6 @@ class TestServerPrompts: def fn() -> str: return "Hello, world!" - @pytest.mark.anyio async def test_list_prompts(self): """Test listing prompts through MCP protocol.""" mcp = FastMCP() @@ -700,7 +674,6 @@ class TestServerPrompts: assert prompt.arguments[1].name == "optional" assert prompt.arguments[1].required is False - @pytest.mark.anyio async def test_get_prompt(self): """Test getting a prompt through MCP protocol.""" mcp = FastMCP() @@ -718,7 +691,6 @@ class TestServerPrompts: assert isinstance(content, TextContent) assert content.text == "Hello, World!" - @pytest.mark.anyio async def test_get_prompt_with_resource(self): """Test getting a prompt that returns resource content.""" mcp = FastMCP() @@ -748,7 +720,6 @@ class TestServerPrompts: assert resource.text == "File contents" assert resource.mimeType == "text/plain" - @pytest.mark.anyio async def test_get_unknown_prompt(self): """Test error when getting unknown prompt.""" mcp = FastMCP() @@ -756,7 +727,6 @@ class TestServerPrompts: with pytest.raises(McpError, match="Unknown prompt"): await client.get_prompt("unknown") - @pytest.mark.anyio async def test_get_prompt_missing_args(self): """Test error when required arguments are missing.""" mcp = FastMCP()