diff --git a/fastapi_mcp/server.py b/fastapi_mcp/server.py index bb751067..3501f8ab 100644 --- a/fastapi_mcp/server.py +++ b/fastapi_mcp/server.py @@ -1,5 +1,6 @@ import json import httpx +import jsonschema from typing import Dict, Optional, Any, List, Union, Literal, Sequence from typing_extensions import Annotated, Doc @@ -141,22 +142,10 @@ def setup_server(self) -> None: # Filter tools based on operation IDs and tags self.tools = self._filter_tools(all_tools, openapi_schema) - mcp_server: Server = Server(self.name, self.description) - - @mcp_server.list_tools() - async def handle_list_tools() -> List[types.Tool]: - return self.tools - - @mcp_server.call_tool() - async def handle_call_tool( - name: str, arguments: Dict[str, Any] - ) -> List[Union[types.TextContent, types.ImageContent, types.EmbeddedResource]]: - # Extract HTTP request info from MCP context + def _extract_http_request_info(request_context: Any) -> Optional[HTTPRequestInfo]: + """Extract HTTP request info from an MCP request context, if present.""" http_request_info = None try: - # Access the MCP server's request context to get the original HTTP Request - request_context = mcp_server.request_context - if request_context and hasattr(request_context, "request"): http_request = request_context.request @@ -174,13 +163,86 @@ async def handle_call_tool( ) except (LookupError, AttributeError) as e: logger.error(f"Could not extract HTTP request info from context: {e}") + return http_request_info + + if hasattr(Server, "list_tools"): + # mcp 1.x: register tool handlers via decorators + mcp_server: Server = Server( + self.name, + version=self.fastapi.version, + instructions=self.description, + ) - return await self._execute_api_tool( - client=self._http_client, - tool_name=name, - arguments=arguments, - operation_map=self.operation_map, - http_request_info=http_request_info, + @mcp_server.list_tools() + async def handle_list_tools() -> List[types.Tool]: + return self.tools + + @mcp_server.call_tool() + async def handle_call_tool( + name: str, arguments: Dict[str, Any] + ) -> List[Union[types.TextContent, types.ImageContent, types.EmbeddedResource]]: + # Extract HTTP request info from MCP context + http_request_info = _extract_http_request_info(mcp_server.request_context) + + return await self._execute_api_tool( + client=self._http_client, + tool_name=name, + arguments=arguments, + operation_map=self.operation_map, + http_request_info=http_request_info, + ) + + else: + # mcp 2.x: the decorator API was removed — handlers are registered + # via constructor callbacks, receive the request context explicitly, + # and must return a CallToolResult instead of a content list. + async def handle_list_tools(context: Any, params: Any) -> types.ListToolsResult: + return types.ListToolsResult(tools=self.tools) + + async def handle_call_tool(context: Any, params: Any) -> types.CallToolResult: + # Mirror mcp 1.x's call_tool contract: validate the arguments + # against the tool's JSON schema and turn handler exceptions + # into error results, not protocol errors. + tool = next((t for t in self.tools if t.name == params.name), None) + if tool is not None: + schema = getattr(tool, "inputSchema", None) + if schema is None: + schema = getattr(tool, "input_schema", None) + if schema is not None: + try: + jsonschema.validate(instance=params.arguments or {}, schema=schema) + except jsonschema.ValidationError as e: + return types.CallToolResult( + content=[types.TextContent(type="text", text=f"Input validation error: {e.message}")], + is_error=True, + ) + + # Extract HTTP request info from the per-request context + http_request_info = _extract_http_request_info(context) + + try: + result = await self._execute_api_tool( + client=self._http_client, + tool_name=params.name, + arguments=params.arguments, + operation_map=self.operation_map, + http_request_info=http_request_info, + ) + except Exception as exc: + # Mirror mcp 1.x's call_tool contract: handler exceptions + # become error results, not protocol errors. + return types.CallToolResult( + content=[types.TextContent(type="text", text=str(exc))], + is_error=True, + ) + return types.CallToolResult(content=list(result)) + + mcp_server: Server = Server( + self.name, + version=self.fastapi.version, + instructions=self.description, + on_list_tools=handle_list_tools, + on_call_tool=handle_call_tool, ) self.server = mcp_server diff --git a/fastapi_mcp/transport/sse.py b/fastapi_mcp/transport/sse.py index 4b64a430..b3d869c5 100644 --- a/fastapi_mcp/transport/sse.py +++ b/fastapi_mcp/transport/sse.py @@ -6,13 +6,18 @@ from fastapi import Request, Response, BackgroundTasks, HTTPException from fastapi.responses import JSONResponse from mcp.shared.message import SessionMessage, ServerMessageMetadata -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError from mcp.server.sse import SseServerTransport from mcp.types import JSONRPCMessage, JSONRPCError, ErrorData logger = logging.getLogger(__name__) +# mcp 1.x exposes JSONRPCMessage as a RootModel; mcp 2.x models it as a plain +# Union type without classmethods like model_validate_json. A TypeAdapter +# handles both shapes uniformly. +_JSONRPC_MESSAGE_ADAPTER = TypeAdapter(JSONRPCMessage) + class FastApiSseTransport(SseServerTransport): async def handle_fastapi_post_message(self, request: Request) -> Response: @@ -57,7 +62,7 @@ async def handle_fastapi_post_message(self, request: Request) -> Response: logger.debug(f"Received JSON: {body.decode()}") try: - message = JSONRPCMessage.model_validate_json(body) + message = _JSONRPC_MESSAGE_ADAPTER.validate_json(body) logger.debug(f"Validated client message: {message}") except ValidationError as err: @@ -104,7 +109,7 @@ async def _send_message_safely( id="unknown", # We don't know the ID from the invalid request error=error_data, ) - error_message = SessionMessage(JSONRPCMessage(root=json_rpc_error)) + error_message = SessionMessage(_JSONRPC_MESSAGE_ADAPTER.validate_python(json_rpc_error)) await writer.send(error_message) else: await writer.send(message) diff --git a/pyproject.toml b/pyproject.toml index 202e2c12..a1f34fef 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,7 +29,8 @@ dependencies = [ "fastapi>=0.100.0", "typer>=0.9.0", "rich>=13.0.0", - "mcp>=1.12.0", + "mcp>=1.12.0,<3.0.0", + "jsonschema>=4.0.0", "pydantic>=2.0.0", "pydantic-settings>=2.5.2", "uvicorn>=0.20.0", diff --git a/tests/_mcp_compat.py b/tests/_mcp_compat.py new file mode 100644 index 00000000..e73c736b --- /dev/null +++ b/tests/_mcp_compat.py @@ -0,0 +1,81 @@ +""" +Compatibility layer for the mcp SDK's in-memory test helper. + +mcp 1.x ships ``create_connected_server_and_client_session`` in +``mcp.shared.memory``. mcp 2.x removed it when the lowlevel server API was +redesigned, but the public ``create_client_server_memory_streams`` helper is +still there — so on 2.x we rebuild the exact same contract from it. Tests +import the helper from this module instead of from mcp directly, so the suite +runs unchanged under both major versions. +""" + +from contextlib import asynccontextmanager +from typing import Any, AsyncGenerator + +import anyio + +def tool_input_schema(tool: Any) -> Any: + """Tool input schema (mcp 1.x: .inputSchema, mcp 2.x: .input_schema).""" + schema = getattr(tool, "inputSchema", None) + if schema is None: + schema = tool.input_schema + return schema + + +def call_result_is_error(result: Any) -> bool: + """True when a CallToolResult flags an error (mcp 1.x: .isError, mcp 2.x: .is_error).""" + value = getattr(result, "isError", None) + if value is None: + value = getattr(result, "is_error", False) + return bool(value) + + +def jsonrpc_message_root(message: Any) -> Any: + """Unwrapped value of a parsed JSONRPCMessage (mcp 1.x: .root, mcp 2.x: already unwrapped).""" + return getattr(message, "root", message) + + +try: # mcp 1.x + from mcp.shared.memory import ( # type: ignore[attr-defined,no-redef] + create_connected_server_and_client_session, + ) + +except ImportError: # mcp 2.x + + from mcp.client.session import ClientSession + from mcp.shared.memory import create_client_server_memory_streams + + @asynccontextmanager + async def create_connected_server_and_client_session( + server: Any, + read_timeout_seconds: Any = None, + raise_exceptions: bool = False, + ) -> AsyncGenerator[Any, None]: + """Create a ClientSession connected to a running MCP server (mcp 2.x).""" + async with create_client_server_memory_streams() as (client_streams, server_streams): + client_read, client_write = client_streams + server_read, server_write = server_streams + + # Create a cancel scope for the server task + async with anyio.create_task_group() as tg: + + async def run_server() -> None: + await server.run( + server_read, + server_write, + server.create_initialization_options(), + raise_exceptions=raise_exceptions, + ) + + tg.start_soon(run_server) + + try: + async with ClientSession( + read_stream=client_read, + write_stream=client_write, + read_timeout_seconds=read_timeout_seconds, + ) as client_session: + await client_session.initialize() + yield client_session + finally: + tg.cancel_scope.cancel() diff --git a/tests/test_mcp_complex_app.py b/tests/test_mcp_complex_app.py index eb6baaff..3c29d1ba 100644 --- a/tests/test_mcp_complex_app.py +++ b/tests/test_mcp_complex_app.py @@ -3,7 +3,7 @@ import pytest import mcp.types as types from mcp.server.lowlevel import Server -from mcp.shared.memory import create_connected_server_and_client_session +from ._mcp_compat import call_result_is_error, create_connected_server_and_client_session from fastapi import FastAPI from fastapi_mcp import FastApiMCP @@ -45,7 +45,7 @@ async def test_call_tool_list_products_default(lowlevel_server_complex_app: Serv async with create_connected_server_and_client_session(lowlevel_server_complex_app) as client_session: response = await client_session.call_tool("list_products", {}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -65,7 +65,7 @@ async def test_call_tool_list_products_with_filters(lowlevel_server_complex_app: {"category": "electronics", "min_price": 10.0, "page": 1, "size": 10, "in_stock_only": True}, ) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -83,7 +83,7 @@ async def test_call_tool_get_product(lowlevel_server_complex_app: Server, exampl async with create_connected_server_and_client_session(lowlevel_server_complex_app) as client_session: response = await client_session.call_tool("get_product", {"product_id": product_id}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -104,7 +104,7 @@ async def test_call_tool_get_product_with_options(lowlevel_server_complex_app: S "get_product", {"product_id": product_id, "include_unavailable": True} ) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -129,7 +129,7 @@ async def test_call_tool_create_order(lowlevel_server_complex_app: Server, examp async with create_connected_server_and_client_session(lowlevel_server_complex_app) as client_session: response = await client_session.call_tool("create_order", order_request) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -155,7 +155,7 @@ async def test_call_tool_create_order_validation_error(lowlevel_server_complex_a async with create_connected_server_and_client_session(lowlevel_server_complex_app) as client_session: response = await client_session.call_tool("create_order", order_request) - assert response.isError + assert call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -169,7 +169,7 @@ async def test_call_tool_get_customer(lowlevel_server_complex_app: Server, examp async with create_connected_server_and_client_session(lowlevel_server_complex_app) as client_session: response = await client_session.call_tool("get_customer", {"customer_id": customer_id}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -195,7 +195,7 @@ async def test_call_tool_get_customer_with_options(lowlevel_server_complex_app: }, ) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -210,7 +210,7 @@ async def test_error_handling_missing_parameter(lowlevel_server_complex_app: Ser # Missing required product_id parameter response = await client_session.call_tool("get_product", {}) - assert response.isError + assert call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) diff --git a/tests/test_mcp_simple_app.py b/tests/test_mcp_simple_app.py index 1c0548f6..cf7b5df6 100644 --- a/tests/test_mcp_simple_app.py +++ b/tests/test_mcp_simple_app.py @@ -3,7 +3,7 @@ import pytest import mcp.types as types from mcp.server.lowlevel import Server -from mcp.shared.memory import create_connected_server_and_client_session +from ._mcp_compat import call_result_is_error, create_connected_server_and_client_session from fastapi import FastAPI from fastapi_mcp import FastApiMCP @@ -58,7 +58,7 @@ async def test_call_tool_get_item_1(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("get_item", {"item_id": 1}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -76,7 +76,7 @@ async def test_call_tool_get_item_2(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("get_item", {"item_id": 2}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -94,7 +94,7 @@ async def test_call_tool_raise_error(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("raise_error", {}) - assert response.isError + assert call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -107,7 +107,7 @@ async def test_error_handling(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("get_item", {}) - assert response.isError + assert call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -128,7 +128,7 @@ async def test_complex_tool_arguments(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("create_item", test_item) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -145,7 +145,7 @@ async def test_call_tool_list_items_default(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("list_items", {}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -163,7 +163,7 @@ async def test_call_tool_list_items_with_pagination(lowlevel_server_simple_app: async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("list_items", {"skip": 1, "limit": 1}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -181,7 +181,7 @@ async def test_call_tool_get_item_not_found(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("get_item", {"item_id": 999}) - assert response.isError + assert call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -203,7 +203,7 @@ async def test_call_tool_update_item(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("update_item", test_update) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -221,7 +221,7 @@ async def test_call_tool_delete_item(lowlevel_server_simple_app: Server): async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("delete_item", {"item_id": 3}) - assert not response.isError + assert not call_result_is_error(response) # The endpoint returns 204 No Content, so we expect an empty response text_content = next(c for c in response.content if isinstance(c, types.TextContent)) assert ( @@ -234,7 +234,7 @@ async def test_call_tool_get_item_with_details(lowlevel_server_simple_app: Serve async with create_connected_server_and_client_session(lowlevel_server_simple_app) as client_session: response = await client_session.call_tool("get_item", {"item_id": 1, "include_details": True}) - assert not response.isError + assert not call_result_is_error(response) assert len(response.content) > 0 text_content = next(c for c in response.content if isinstance(c, types.TextContent)) @@ -375,7 +375,7 @@ async def test_context_extraction_in_tool_handler(fastapi_mcp: FastApiMCP): """Test that handle_call_tool extracts HTTP request info from MCP context.""" from unittest.mock import patch, MagicMock import mcp.types as types - from mcp.server.lowlevel.server import request_ctx + from mcp.server.lowlevel import Server # Create a fake HTTP request object with headers fake_http_request = MagicMock() @@ -389,55 +389,56 @@ async def test_context_extraction_in_tool_handler(fastapi_mcp: FastApiMCP): fake_request_context = MagicMock() fake_request_context.request = fake_http_request - # Test with authorization header extraction from context - token = request_ctx.set(fake_request_context) - try: + call_request = types.CallToolRequest( + method="tools/call", params=types.CallToolRequestParams(name="get_item", arguments={"item_id": 1}) + ) + + async def invoke_tool_handler(with_context: bool): + """Dispatch a tools/call request straight into the lowlevel handler.""" with patch.object(fastapi_mcp, "_execute_api_tool") as mock_execute: mock_execute.return_value = [types.TextContent(type="text", text="success")] - # Create a CallToolRequest like the MCP protocol would - call_request = types.CallToolRequest( - method="tools/call", params=types.CallToolRequestParams(name="get_item", arguments={"item_id": 1}) - ) - try: - # Call the tool handler directly like the MCP server would - await fastapi_mcp.server.request_handlers[types.CallToolRequest](call_request) + if hasattr(Server, "list_tools"): + # mcp 1.x: handlers are dispatched from request_handlers and read + # the request context from the request_ctx ContextVar. + from mcp.server.lowlevel.server import request_ctx + + token = request_ctx.set(fake_request_context if with_context else None) + try: + await fastapi_mcp.server.request_handlers[types.CallToolRequest](call_request) + finally: + request_ctx.reset(token) + else: + # mcp 2.x: the request context is passed to the handler explicitly. + entry = fastapi_mcp.server._request_handlers["tools/call"] + context = fake_request_context if with_context else MagicMock(spec=[]) + await entry.handler(context, call_request.params) except Exception: pass - assert mock_execute.called, "The _execute_api_tool method was not called" - - if mock_execute.called: - # Verify that HTTPRequestInfo was extracted from context and passed to _execute_api_tool - http_request_info = mock_execute.call_args.kwargs["http_request_info"] - assert http_request_info is not None, "HTTPRequestInfo should be extracted from context" - assert http_request_info.method == "POST" - assert http_request_info.path == "/test" - assert "Authorization" in http_request_info.headers - assert http_request_info.headers["Authorization"] == "Bearer token-123" - assert "X-Custom" in http_request_info.headers - assert http_request_info.headers["X-Custom"] == "custom-value-123" - finally: - # Clean up the context variable - request_ctx.reset(token) + return mock_execute - # Test with missing request context (should still work but with None) - with patch.object(fastapi_mcp, "_execute_api_tool") as mock_execute: - mock_execute.return_value = [types.TextContent(type="text", text="success")] + # Test with authorization header extraction from context + mock_execute = await invoke_tool_handler(with_context=True) - call_request = types.CallToolRequest( - method="tools/call", params=types.CallToolRequestParams(name="get_item", arguments={"item_id": 1}) - ) + assert mock_execute.called, "The _execute_api_tool method was not called" - try: - await fastapi_mcp.server.request_handlers[types.CallToolRequest](call_request) - except Exception: - pass + # Verify that HTTPRequestInfo was extracted from context and passed to _execute_api_tool + http_request_info = mock_execute.call_args.kwargs["http_request_info"] + assert http_request_info is not None, "HTTPRequestInfo should be extracted from context" + assert http_request_info.method == "POST" + assert http_request_info.path == "/test" + assert "Authorization" in http_request_info.headers + assert http_request_info.headers["Authorization"] == "Bearer token-123" + assert "X-Custom" in http_request_info.headers + assert http_request_info.headers["X-Custom"] == "custom-value-123" + + # Test with missing request context (should still work but with None) + mock_execute = await invoke_tool_handler(with_context=False) - assert mock_execute.called, "The _execute_api_tool method was not called" + assert mock_execute.called, "The _execute_api_tool method was not called" - if mock_execute.called: - # Verify that HTTPRequestInfo is None when context is not available - http_request_info = mock_execute.call_args.kwargs["http_request_info"] - assert http_request_info is None, "HTTPRequestInfo should be None when context is not available" + # Verify that HTTPRequestInfo is None when context is not available + http_request_info = mock_execute.call_args.kwargs["http_request_info"] + assert http_request_info is None, "HTTPRequestInfo should be None when context is not available" diff --git a/tests/test_openapi_conversion.py b/tests/test_openapi_conversion.py index aefe6433..5f7ff432 100644 --- a/tests/test_openapi_conversion.py +++ b/tests/test_openapi_conversion.py @@ -3,6 +3,7 @@ import mcp.types as types from fastapi_mcp.openapi.convert import convert_openapi_to_mcp_tools +from ._mcp_compat import tool_input_schema from fastapi_mcp.openapi.utils import ( clean_schema_for_display, generate_example_from_schema, @@ -32,7 +33,7 @@ def test_simple_app_conversion(simple_fastapi_app: FastAPI): assert isinstance(tool, types.Tool) assert tool.name in expected_operations assert tool.description is not None - assert tool.inputSchema is not None + assert tool_input_schema(tool) is not None def test_complex_app_conversion(complex_fastapi_app: FastAPI): @@ -57,7 +58,7 @@ def test_complex_app_conversion(complex_fastapi_app: FastAPI): assert isinstance(tool, types.Tool) assert tool.name in expected_operations assert tool.description is not None - assert tool.inputSchema is not None + assert tool_input_schema(tool) is not None def test_describe_full_response_schema(simple_fastapi_app: FastAPI): @@ -171,7 +172,7 @@ def test_parameter_handling(complex_fastapi_app: FastAPI): list_products_tool = next(tool for tool in tools if tool.name == "list_products") - properties = list_products_tool.inputSchema["properties"] + properties = tool_input_schema(list_products_tool)["properties"] assert "product_id" not in properties # This is from get_product, not list_products @@ -206,7 +207,7 @@ def test_parameter_handling(complex_fastapi_app: FastAPI): assert "tag" in properties assert properties["tag"].get("type") == "array" - required = list_products_tool.inputSchema.get("required", []) + required = tool_input_schema(list_products_tool).get("required", []) assert "page" not in required # Has default value assert "category" not in required # Optional parameter @@ -215,14 +216,14 @@ def test_parameter_handling(complex_fastapi_app: FastAPI): assert operation_map["list_products"]["method"] == "get" get_product_tool = next(tool for tool in tools if tool.name == "get_product") - get_product_props = get_product_tool.inputSchema["properties"] + get_product_props = tool_input_schema(get_product_tool)["properties"] assert "product_id" in get_product_props assert get_product_props["product_id"].get("type") == "string" # UUID converted to string assert "description" in get_product_props["product_id"] get_customer_tool = next(tool for tool in tools if tool.name == "get_customer") - get_customer_props = get_customer_tool.inputSchema["properties"] + get_customer_props = tool_input_schema(get_customer_tool)["properties"] assert "fields" in get_customer_props assert get_customer_props["fields"].get("type") == "array" @@ -247,7 +248,7 @@ def test_request_body_handling(complex_fastapi_app: FastAPI): create_order_tool = next(tool for tool in tools if tool.name == "create_order") - properties = create_order_tool.inputSchema["properties"] + properties = tool_input_schema(create_order_tool)["properties"] assert "customer_id" in properties assert "items" in properties @@ -268,7 +269,7 @@ def test_request_body_handling(complex_fastapi_app: FastAPI): assert "default" in properties[param_name] assert properties[param_name]["default"] == original_properties[param_name]["default"] - required = create_order_tool.inputSchema.get("required", []) + required = tool_input_schema(create_order_tool).get("required", []) assert "customer_id" in required assert "items" in required assert "shipping_address_id" in required @@ -321,7 +322,7 @@ def test_missing_type_handling(complex_fastapi_app: FastAPI): tools, operation_map = convert_openapi_to_mcp_tools(openapi_schema) get_product_tool = next(tool for tool in tools if tool.name == "get_product") - get_product_props = get_product_tool.inputSchema["properties"] + get_product_props = tool_input_schema(get_product_tool)["properties"] assert "product_id" in get_product_props assert get_product_props["product_id"].get("type") == "string" # Default type applied @@ -353,7 +354,7 @@ def test_body_params_descriptions_and_defaults(complex_fastapi_app: FastAPI): tools, _ = convert_openapi_to_mcp_tools(openapi_schema) create_order_tool = next(tool for tool in tools if tool.name == "create_order") - properties = create_order_tool.inputSchema["properties"] + properties = tool_input_schema(create_order_tool)["properties"] assert "description" in properties["customer_id"] assert properties["customer_id"]["description"] == "Test customer ID description" @@ -409,7 +410,7 @@ def test_body_params_edge_cases(complex_fastapi_app: FastAPI): tools, _ = convert_openapi_to_mcp_tools(openapi_schema) create_order_tool = next(tool for tool in tools if tool.name == "create_order") - properties = create_order_tool.inputSchema["properties"] + properties = tool_input_schema(create_order_tool)["properties"] assert "customer_id" in properties assert "title" in properties["customer_id"] diff --git a/tests/test_sse_mock_transport.py b/tests/test_sse_mock_transport.py index e833e5da..1e3542b0 100644 --- a/tests/test_sse_mock_transport.py +++ b/tests/test_sse_mock_transport.py @@ -7,8 +7,9 @@ from pydantic import ValidationError from anyio.streams.memory import MemoryObjectSendStream -from fastapi_mcp.transport.sse import FastApiSseTransport -from mcp.types import JSONRPCMessage, JSONRPCError +from fastapi_mcp.transport.sse import FastApiSseTransport, _JSONRPC_MESSAGE_ADAPTER +from ._mcp_compat import jsonrpc_message_root +from mcp.types import JSONRPCError @pytest.fixture @@ -118,11 +119,12 @@ async def test_handle_post_message_general_exception( # Instead of mocking the body method to raise an exception, # we'll patch the body method to return a normal value and then - # patch JSONRPCMessage.model_validate_json to raise the exception + # make the JSON-RPC message parsing raise the exception mock_request.body = AsyncMock(return_value=b'{"jsonrpc": "2.0", "method": "test", "id": "1"}') - # Mock the model_validate_json method to raise an Exception - with patch("mcp.types.JSONRPCMessage.model_validate_json", side_effect=Exception("Test exception")): + # Mock the JSON-RPC message parsing to raise an Exception. The transport + # always parses via _JSONRPC_MESSAGE_ADAPTER (works on mcp 1.x and 2.x). + with patch("fastapi_mcp.transport.sse._JSONRPC_MESSAGE_ADAPTER.validate_json", side_effect=Exception("Test exception")): # Check that the function raises HTTPException with the correct status code with pytest.raises(HTTPException) as excinfo: await mock_transport.handle_fastapi_post_message(mock_request) @@ -147,9 +149,8 @@ async def test_send_message_safely_with_validation_error( assert mock_writer.send.called sent_message = mock_writer.send.call_args[0][0] assert isinstance(sent_message, SessionMessage) - assert isinstance(sent_message.message, JSONRPCMessage) - assert isinstance(sent_message.message.root, JSONRPCError) - assert sent_message.message.root.error.code == -32700 # Parse error code + assert isinstance(jsonrpc_message_root(sent_message.message), JSONRPCError) + assert jsonrpc_message_root(sent_message.message).error.code == -32700 # Parse error code @pytest.mark.anyio @@ -158,9 +159,7 @@ async def test_send_message_safely_with_jsonrpc_message( ) -> None: """Test sending a JSONRPCMessage safely.""" # Create a JSONRPCMessage - message = SessionMessage( - JSONRPCMessage.model_validate({"jsonrpc": "2.0", "id": "123", "method": "test_method", "params": {}}) - ) + message = SessionMessage(_JSONRPC_MESSAGE_ADAPTER.validate_python({"jsonrpc": "2.0", "id": "123", "method": "test_method", "params": {}})) # Call the function await mock_transport._send_message_safely(mock_writer, message) @@ -181,7 +180,7 @@ async def test_send_message_safely_exception_handling( # Create a message message = SessionMessage( - JSONRPCMessage.model_validate({"jsonrpc": "2.0", "id": "123", "method": "test_method", "params": {}}) + _JSONRPC_MESSAGE_ADAPTER.validate_python({"jsonrpc": "2.0", "id": "123", "method": "test_method", "params": {}}) ) # Call the function - it should not raise an exception