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
102 changes: 82 additions & 20 deletions fastapi_mcp/server.py
Original file line number Diff line number Diff line change
@@ -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

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

Expand All @@ -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
Expand Down
11 changes: 8 additions & 3 deletions fastapi_mcp/transport/sse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
81 changes: 81 additions & 0 deletions tests/_mcp_compat.py
Original file line number Diff line number Diff line change
@@ -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()
20 changes: 10 additions & 10 deletions tests/test_mcp_complex_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand Down
Loading