Merge pull request #575 from jlowin/headers-passthrough

Ensure client headers are passed through to remote servers
This commit is contained in:
Jeremiah Lowin 2025-05-23 15:56:34 -04:00 committed by GitHub
commit e8bde2953e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 379 additions and 63 deletions

View file

@ -25,6 +25,7 @@ from pydantic import AnyUrl
from typing_extensions import Unpack
from fastmcp.server import FastMCP as FastMCPServer
from fastmcp.server.dependencies import get_http_request
from fastmcp.server.server import FastMCP
from fastmcp.utilities.logging import get_logger
from fastmcp.utilities.mcp_config import MCPConfig, infer_transport_type_from_url
@ -34,6 +35,11 @@ if TYPE_CHECKING:
logger = get_logger(__name__)
EXCLUDE_HEADERS = {
"content-type",
"content-length",
}
class SessionKwargs(TypedDict, total=False):
"""Keyword arguments for the MCP ClientSession constructor."""
@ -132,7 +138,21 @@ class SSETransport(ClientTransport):
async def connect_session(
self, **session_kwargs: Unpack[SessionKwargs]
) -> AsyncIterator[ClientSession]:
client_kwargs = {}
client_kwargs: dict[str, Any] = {
"headers": self.headers,
}
# load headers from an active HTTP request, if available. This will only be true
# if the client is used in a FastMCP Proxy, in which case the MCP client headers
# need to be forwarded to the remote server.
try:
active_request = get_http_request()
for name, value in active_request.headers.items():
if name not in self.headers and name not in EXCLUDE_HEADERS:
client_kwargs["headers"][name] = str(value)
except RuntimeError:
client_kwargs["headers"] = self.headers
# sse_read_timeout has a default value set, so we can't pass None without overriding it
# instead we simply leave the kwarg out if it's not provided
if self.sse_read_timeout is not None:
@ -143,9 +163,7 @@ class SSETransport(ClientTransport):
)
client_kwargs["timeout"] = read_timeout_seconds.total_seconds()
async with sse_client(
self.url, headers=self.headers, **client_kwargs
) as transport:
async with sse_client(self.url, **client_kwargs) as transport:
read_stream, write_stream = transport
async with ClientSession(
read_stream, write_stream, **session_kwargs
@ -180,7 +198,23 @@ class StreamableHttpTransport(ClientTransport):
async def connect_session(
self, **session_kwargs: Unpack[SessionKwargs]
) -> AsyncIterator[ClientSession]:
client_kwargs = {}
client_kwargs: dict[str, Any] = {
"headers": self.headers,
}
# load headers from an active HTTP request, if available. This will only be true
# if the client is used in a FastMCP Proxy, in which case the MCP client headers
# need to be forwarded to the remote server.
try:
active_request = get_http_request()
for name, value in active_request.headers.items():
if name not in self.headers and name not in EXCLUDE_HEADERS:
client_kwargs["headers"][name] = str(value)
except RuntimeError:
client_kwargs["headers"] = self.headers
print(client_kwargs)
# sse_read_timeout has a default value set, so we can't pass None without overriding it
# instead we simply leave the kwarg out if it's not provided
if self.sse_read_timeout is not None:
@ -188,9 +222,7 @@ class StreamableHttpTransport(ClientTransport):
if session_kwargs.get("read_timeout_seconds", None) is not None:
client_kwargs["timeout"] = session_kwargs.get("read_timeout_seconds")
async with streamablehttp_client(
self.url, headers=self.headers, **client_kwargs
) as transport:
async with streamablehttp_client(self.url, **client_kwargs) as transport:
read_stream, write_stream, _ = transport
async with ClientSession(
read_stream, write_stream, **session_kwargs

View file

@ -17,6 +17,7 @@ from pydantic.networks import AnyUrl
from fastmcp.exceptions import ToolError
from fastmcp.resources import Resource, ResourceTemplate
from fastmcp.server.dependencies import get_http_request
from fastmcp.server.server import FastMCP
from fastmcp.tools.tool import Tool, _convert_to_content
from fastmcp.utilities import openapi
@ -365,6 +366,20 @@ class OpenAPITool(Tool):
# Prepare headers - fix typing by ensuring all values are strings
headers = {}
# Try to get headers from the current MCP client HTTP request
try:
http_request = get_http_request()
# Add headers from the MCP client request
for name, value in http_request.headers.items():
# Don't override headers that are already set on the client
if name not in self._client.headers:
headers[name] = str(value)
except RuntimeError:
# No active HTTP request (e.g., STDIO transport), continue without client headers
pass
# Add any OpenAPI-defined header parameters (these take precedence over client headers)
for p in self._route.parameters:
if (
p.location == "header"
@ -519,10 +534,24 @@ class OpenAPIResource(Resource):
if value is not None and value != "":
query_params[param.name] = value
# Prepare headers from MCP client request if available
headers = {}
try:
http_request = get_http_request()
# Add headers from the MCP client request
for name, value in http_request.headers.items():
# Don't override headers that are already set on the client
if name not in self._client.headers:
headers[name] = str(value)
except RuntimeError:
# No active HTTP request (e.g., STDIO transport), continue without client headers
pass
response = await self._client.request(
method=self._route.method,
url=path,
params=query_params,
headers=headers,
timeout=self._timeout,
)
@ -733,6 +762,7 @@ class FastMCPOpenAPI(FastMCP):
) -> str:
"""Generate a default name from the route path."""
# First check for OpenAPI operationId which takes precedence
if route.operation_id:
return route.operation_id
@ -848,7 +878,7 @@ class FastMCPOpenAPI(FastMCP):
# Get a unique resource name
resource_name = self._get_unique_name(name, "resources")
resource_uri = f"resource://openapi/{resource_name}"
resource_uri = f"resource://{resource_name}"
base_description = (
route.description or route.summary or f"Represents {route.path}"
)
@ -896,7 +926,7 @@ class FastMCPOpenAPI(FastMCP):
path_params = [p.name for p in route.parameters if p.location == "path"]
path_params.sort() # Sort for consistent URIs
uri_template_str = f"resource://openapi/{template_name}"
uri_template_str = f"resource://{template_name}"
if path_params:
uri_template_str += "/" + "/".join(f"{{{p}}}" for p in path_params)

View file

@ -1211,6 +1211,7 @@ class FastMCP(Generic[LifespanResultT]):
route_map_fn: OpenAPIRouteMapFn | None = None,
mcp_component_fn: OpenAPIComponentFn | None = None,
all_routes_as_tools: bool = False,
httpx_client_kwargs: dict[str, Any] | None = None,
**settings: Any,
) -> FastMCPOpenAPI:
"""
@ -1234,8 +1235,13 @@ class FastMCP(Generic[LifespanResultT]):
elif all_routes_as_tools:
route_maps = [RouteMap(methods="*", pattern=r".*", mcp_type=MCPType.TOOL)]
if httpx_client_kwargs is None:
httpx_client_kwargs = {}
httpx_client_kwargs.setdefault("base_url", "http://fastapi")
client = httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://fastapi"
transport=httpx.ASGITransport(app=app),
**httpx_client_kwargs,
)
name = name or app.title

View file

@ -0,0 +1,192 @@
import json
import sys
from collections.abc import Generator
import pytest
import uvicorn
from fastapi import FastAPI, Request
from mcp.types import TextContent, TextResourceContents
from fastmcp import Client, FastMCP
from fastmcp.client.transports import SSETransport, StreamableHttpTransport
from fastmcp.utilities.tests import run_server_in_process
def fastmcp_server_for_headers() -> FastMCP:
app = FastAPI()
@app.get("/headers")
def get_headers(request: Request):
return request.headers
@app.get("/headers/{header_name}")
def get_header_by_name(header_name: str, request: Request):
return request.headers[header_name]
@app.post("/headers")
def post_headers(request: Request):
return request.headers
mcp = FastMCP.from_fastapi(
app, httpx_client_kwargs={"headers": {"X-SERVER": "test-abc"}}
)
return mcp
class TestClientHeaders:
def run_shttp_server(self, host: str, port: int) -> None:
try:
app = fastmcp_server_for_headers().http_app(transport="streamable-http")
server = uvicorn.Server(
config=uvicorn.Config(
app=app,
host=host,
port=port,
log_level="error",
lifespan="on",
)
)
server.run()
except Exception as e:
print(f"Server error: {e}")
sys.exit(1)
sys.exit(0)
def run_sse_server(self, host: str, port: int) -> None:
try:
app = fastmcp_server_for_headers().http_app(transport="sse")
server = uvicorn.Server(
config=uvicorn.Config(
app=app,
host=host,
port=port,
log_level="error",
lifespan="on",
)
)
server.run()
except Exception as e:
print(f"Server error: {e}")
sys.exit(1)
sys.exit(0)
def run_proxy_server(self, host: str, port: int, remote_url: str) -> None:
try:
client = Client(transport=StreamableHttpTransport(remote_url))
app = FastMCP.as_proxy(client).http_app(transport="streamable-http")
server = uvicorn.Server(
config=uvicorn.Config(
app=app,
host=host,
port=port,
log_level="error",
lifespan="on",
)
)
server.run()
except Exception as e:
print(f"Server error: {e}")
sys.exit(1)
sys.exit(0)
@pytest.fixture(scope="class")
def shttp_server(self) -> Generator[str, None, None]:
with run_server_in_process(self.run_shttp_server) as url:
yield f"{url}/mcp"
@pytest.fixture(scope="class")
def sse_server(self) -> Generator[str, None, None]:
with run_server_in_process(self.run_sse_server) as url:
yield f"{url}/sse"
@pytest.fixture(scope="class")
def proxy_server(self, shttp_server: str) -> Generator[str, None, None]:
with run_server_in_process(self.run_proxy_server, shttp_server + "/mcp") as url:
yield f"{url}/mcp"
async def test_client_headers_sse_resource(self, sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
) as client:
result = await client.read_resource("resource://get_headers_headers_get")
assert isinstance(result[0], TextResourceContents)
headers = json.loads(result[0].text)
assert headers["x-test"] == "test-123"
async def test_client_headers_shttp_resource(self, shttp_server: str):
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-TEST": "test-123"}
)
) as client:
result = await client.read_resource("resource://get_headers_headers_get")
assert isinstance(result[0], TextResourceContents)
headers = json.loads(result[0].text)
assert headers["x-test"] == "test-123"
async def test_client_headers_sse_resource_template(self, sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
) as client:
result = await client.read_resource(
"resource://get_header_by_name_headers__header_name__get/x-test"
)
assert isinstance(result[0], TextResourceContents)
header = json.loads(result[0].text)
assert header == "test-123"
async def test_client_headers_shttp_resource_template(self, shttp_server: str):
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-TEST": "test-123"}
)
) as client:
result = await client.read_resource(
"resource://get_header_by_name_headers__header_name__get/x-test"
)
assert isinstance(result[0], TextResourceContents)
header = json.loads(result[0].text)
assert header == "test-123"
async def test_client_headers_sse_tool(self, sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
) as client:
result = await client.call_tool("post_headers_headers_post")
assert isinstance(result[0], TextContent)
headers = json.loads(result[0].text)
assert headers["x-test"] == "test-123"
async def test_client_headers_shttp_tool(self, shttp_server: str):
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-TEST": "test-123"}
)
) as client:
result = await client.call_tool("post_headers_headers_post")
assert isinstance(result[0], TextContent)
headers = json.loads(result[0].text)
assert headers["x-test"] == "test-123"
async def test_client_doesnt_override_server_headers(self, shttp_server: str):
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-SERVER": "test-client"}
)
) as client:
result = await client.read_resource("resource://get_headers_headers_get")
assert isinstance(result[0], TextResourceContents)
headers = json.loads(result[0].text)
assert headers["x-server"] == "test-abc"
async def test_client_headers_proxy(self, proxy_server: str):
"""
Test that client headers are passed through the proxy to the remove server.
"""
async with Client(transport=StreamableHttpTransport(proxy_server)) as client:
await client.ping()
result = await client.read_resource("resource://get_headers_headers_get")
assert isinstance(result[0], TextResourceContents)
headers = json.loads(result[0].text)
assert headers["x-server"] == "test-abc"

View file

@ -109,5 +109,5 @@ class TestRouteTypeIgnoreDeprecation:
resource_uris = [str(r.uri) for r in resources.values()]
# Analytics should be excluded
assert "resource://openapi/get_items" in resource_uris
assert "resource://openapi/get_analytics" not in resource_uris
assert "resource://get_items" in resource_uris
assert "resource://get_analytics" not in resource_uris

View file

@ -7,7 +7,7 @@ import uvicorn
from mcp.types import TextContent, TextResourceContents
from fastmcp.client import Client
from fastmcp.client.transports import StreamableHttpTransport
from fastmcp.client.transports import SSETransport, StreamableHttpTransport
from fastmcp.server.dependencies import get_http_request
from fastmcp.server.server import FastMCP
from fastmcp.utilities.tests import run_server_in_process
@ -41,9 +41,28 @@ def fastmcp_server():
return server
def run_server(host: str, port: int) -> None:
def run_shttp_server(host: str, port: int) -> None:
try:
app = fastmcp_server().http_app()
app = fastmcp_server().http_app(transport="streamable-http")
server = uvicorn.Server(
config=uvicorn.Config(
app=app,
host=host,
port=port,
log_level="error",
lifespan="on",
)
)
server.run()
except Exception as e:
print(f"Server error: {e}")
sys.exit(1)
sys.exit(0)
def run_sse_server(host: str, port: int) -> None:
try:
app = fastmcp_server().http_app(transport="sse")
server = uvicorn.Server(
config=uvicorn.Config(
app=app,
@ -61,15 +80,23 @@ def run_server(host: str, port: int) -> None:
@pytest.fixture(autouse=True, scope="module")
def sse_server() -> Generator[str, None, None]:
with run_server_in_process(run_server) as url:
def shttp_server() -> Generator[str, None, None]:
with run_server_in_process(run_shttp_server) as url:
yield f"{url}/mcp"
async def test_http_headers_resource(sse_server: str):
@pytest.fixture(autouse=True, scope="module")
def sse_server() -> Generator[str, None, None]:
with run_server_in_process(run_sse_server) as url:
yield f"{url}/sse"
async def test_http_headers_resource_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
raw_result = await client.read_resource("request://headers")
assert isinstance(raw_result[0], TextResourceContents)
@ -78,10 +105,24 @@ async def test_http_headers_resource(sse_server: str):
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_tool(sse_server: str):
async def test_http_headers_resource_sse(sse_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
raw_result = await client.read_resource("request://headers")
assert isinstance(raw_result[0], TextResourceContents)
json_result = json.loads(raw_result[0].text)
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_tool_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
result = await client.call_tool("get_headers_tool")
assert isinstance(result[0], TextContent)
@ -90,10 +131,35 @@ async def test_http_headers_tool(sse_server: str):
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_prompt(sse_server: str):
async def test_http_headers_tool_sse(sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
result = await client.call_tool("get_headers_tool")
assert isinstance(result[0], TextContent)
json_result = json.loads(result[0].text)
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_prompt_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
result = await client.get_prompt("get_headers_prompt")
assert isinstance(result.messages[0].content, TextContent)
json_result = json.loads(result.messages[0].content.text)
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_prompt_sse(sse_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
result = await client.get_prompt("get_headers_prompt")
assert isinstance(result.messages[0].content, TextContent)

View file

@ -249,7 +249,7 @@ class TestTools:
# Check that the user was created via MCP
async with Client(fastmcp_openapi_server) as client:
user_response = await client.read_resource(
"resource://openapi/get_user_users__user_id__get/4"
"resource://get_user_users__user_id__get/4"
)
assert isinstance(user_response[0], TextResourceContents)
response_text = user_response[0].text
@ -283,7 +283,7 @@ class TestTools:
# Check that the user was updated via MCP
async with Client(fastmcp_openapi_server) as client:
user_response = await client.read_resource(
"resource://openapi/get_user_users__user_id__get/1"
"resource://get_user_users__user_id__get/1"
)
assert isinstance(user_response[0], TextResourceContents)
response_text = user_response[0].text
@ -325,7 +325,7 @@ class TestResources:
async with Client(fastmcp_openapi_server) as client:
resources = await client.list_resources()
assert len(resources) == 4
assert resources[0].uri == AnyUrl("resource://openapi/get_users_users_get")
assert resources[0].uri == AnyUrl("resource://get_users_users_get")
assert resources[0].name == "get_users_users_get"
async def test_get_resource(
@ -343,7 +343,7 @@ class TestResources:
)
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
"resource://openapi/get_users_users_get"
"resource://get_users_users_get"
)
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
@ -360,7 +360,7 @@ class TestResources:
"""Test reading a resource that returns bytes."""
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
"resource://openapi/ping_bytes_ping_bytes_get"
"resource://ping_bytes_ping_bytes_get"
)
assert isinstance(resource_response[0], BlobResourceContents)
assert base64.b64decode(resource_response[0].blob) == b"pong"
@ -372,9 +372,7 @@ class TestResources:
):
"""Test reading a resource that returns a string."""
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
"resource://openapi/ping_ping_get"
)
resource_response = await client.read_resource("resource://ping_ping_get")
assert isinstance(resource_response[0], TextResourceContents)
assert resource_response[0].text == "pong"
@ -392,7 +390,7 @@ class TestResourceTemplates:
assert resource_templates[0].name == "get_user_users__user_id__get"
assert (
resource_templates[0].uriTemplate
== r"resource://openapi/get_user_users__user_id__get/{user_id}"
== r"resource://get_user_users__user_id__get/{user_id}"
)
assert (
resource_templates[1].name
@ -400,7 +398,7 @@ class TestResourceTemplates:
)
assert (
resource_templates[1].uriTemplate
== r"resource://openapi/get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
== r"resource://get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
)
async def test_get_resource_template(
@ -415,7 +413,7 @@ class TestResourceTemplates:
user_id = 2
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
f"resource://openapi/get_user_users__user_id__get/{user_id}"
f"resource://get_user_users__user_id__get/{user_id}"
)
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
@ -438,7 +436,7 @@ class TestResourceTemplates:
is_active = True
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
f"resource://openapi/get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
f"resource://get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
)
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
@ -555,7 +553,7 @@ class TestTagTransfer:
# Manually create a resource from template
params = {"user_id": 1}
resource = await get_user_template.create_resource(
"resource://openapi/get_user_users__user_id__get/1", params
"resource://get_user_users__user_id__get/1", params
)
# Verify tags are preserved from template to resource
@ -672,7 +670,7 @@ class TestOpenAPI30Compatibility:
async with Client(openapi_30_server) as client:
resources = await client.list_resources()
assert len(resources) == 1
assert resources[0].uri == AnyUrl("resource://openapi/listProducts")
assert resources[0].uri == AnyUrl("resource://listProducts")
async def test_resource_template_discovery(self, openapi_30_server):
"""Test that resource templates are correctly discovered from an OpenAPI 3.0 spec."""
@ -680,7 +678,7 @@ class TestOpenAPI30Compatibility:
templates = await client.list_resource_templates()
assert len(templates) == 1
assert templates[0].name == "getProduct"
assert templates[0].uriTemplate == r"resource://openapi/getProduct/{product_id}"
assert templates[0].uriTemplate == r"resource://getProduct/{product_id}"
async def test_tool_discovery(self, openapi_30_server):
"""Test that tools are correctly discovered from an OpenAPI 3.0 spec."""
@ -694,9 +692,7 @@ class TestOpenAPI30Compatibility:
async def test_resource_access(self, openapi_30_server):
"""Test reading a resource from an OpenAPI 3.0 server."""
async with Client(openapi_30_server) as client:
resource_response = await client.read_resource(
"resource://openapi/listProducts"
)
resource_response = await client.read_resource("resource://listProducts")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)
@ -707,9 +703,7 @@ class TestOpenAPI30Compatibility:
async def test_resource_template_access(self, openapi_30_server):
"""Test reading a resource from template from an OpenAPI 3.0 server."""
async with Client(openapi_30_server) as client:
resource_response = await client.read_resource(
"resource://openapi/getProduct/p1"
)
resource_response = await client.read_resource("resource://getProduct/p1")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)
@ -852,7 +846,7 @@ class TestOpenAPI31Compatibility:
async with Client(openapi_31_server) as client:
resources = await client.list_resources()
assert len(resources) == 1
assert resources[0].uri == AnyUrl("resource://openapi/listOrders")
assert resources[0].uri == AnyUrl("resource://listOrders")
async def test_resource_template_discovery(self, openapi_31_server):
"""Test that resource templates are correctly discovered from an OpenAPI 3.1 spec."""
@ -860,7 +854,7 @@ class TestOpenAPI31Compatibility:
templates = await client.list_resource_templates()
assert len(templates) == 1
assert templates[0].name == "getOrder"
assert templates[0].uriTemplate == r"resource://openapi/getOrder/{order_id}"
assert templates[0].uriTemplate == r"resource://getOrder/{order_id}"
async def test_tool_discovery(self, openapi_31_server):
"""Test that tools are correctly discovered from an OpenAPI 3.1 spec."""
@ -874,9 +868,7 @@ class TestOpenAPI31Compatibility:
async def test_resource_access(self, openapi_31_server):
"""Test reading a resource from an OpenAPI 3.1 server."""
async with Client(openapi_31_server) as client:
resource_response = await client.read_resource(
"resource://openapi/listOrders"
)
resource_response = await client.read_resource("resource://listOrders")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)
@ -887,9 +879,7 @@ class TestOpenAPI31Compatibility:
async def test_resource_template_access(self, openapi_31_server):
"""Test reading a resource from template from an OpenAPI 3.1 server."""
async with Client(openapi_31_server) as client:
resource_response = await client.read_resource(
"resource://openapi/getOrder/o1"
)
resource_response = await client.read_resource("resource://getOrder/o1")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)
@ -927,7 +917,7 @@ class TestMountFastMCP:
assert len(resources) == 4 # Updated to account for new search endpoint
# We're checking the key used by mcp to store the resource
# The prefixed URI is used as the key, but the resource's original uri is preserved
prefixed_uri = "resource://fastapi/openapi/get_users_users_get"
prefixed_uri = "resource://fastapi/get_users_users_get"
resource = mcp._resource_manager.get_resources().get(prefixed_uri)
assert resource is not None

View file

@ -9,7 +9,7 @@ from fastmcp.utilities.openapi import parse_openapi_to_http_routes
@pytest.fixture
def fastapi_server() -> FastAPI:
def fastapi_app() -> FastAPI:
"""Fixture that returns a FastAPI app for live OpenAPI schema testing."""
from enum import Enum
@ -228,9 +228,9 @@ def fastapi_server() -> FastAPI:
@pytest.fixture
def fastapi_openapi_schema(fastapi_server) -> dict[str, Any]:
def fastapi_openapi_schema(fastapi_app) -> dict[str, Any]:
"""Fixture that returns the OpenAPI schema from a live FastAPI server."""
return fastapi_server.openapi()
return fastapi_app.openapi()
@pytest.fixture
@ -472,11 +472,11 @@ def test_tag_consistency_across_related_endpoints(route_map):
)
def test_tag_order_preservation(fastapi_server):
def test_tag_order_preservation(fastapi_app):
"""Test that tag order is preserved in the parsed routes."""
# Add a new endpoint with specifically ordered tags
@fastapi_server.get(
@fastapi_app.get(
"/test-tag-order",
tags=["first", "second", "third"],
operation_id="test_tag_order",
@ -485,7 +485,7 @@ def test_tag_order_preservation(fastapi_server):
return {"result": "testing tag order"}
# Get the updated schema and parse routes
routes = parse_openapi_to_http_routes(fastapi_server.openapi())
routes = parse_openapi_to_http_routes(fastapi_app.openapi())
# Find our test route
test_route = next((r for r in routes if r.path == "/test-tag-order"), None)
@ -497,11 +497,11 @@ def test_tag_order_preservation(fastapi_server):
)
def test_duplicate_tags_handling(fastapi_server):
def test_duplicate_tags_handling(fastapi_app):
"""Test handling of duplicate tags in the OpenAPI schema."""
# Add an endpoint with duplicate tags
@fastapi_server.get(
@fastapi_app.get(
"/test-duplicate-tags",
tags=["duplicate", "items", "duplicate"],
operation_id="test_duplicate_tags",
@ -510,7 +510,7 @@ def test_duplicate_tags_handling(fastapi_server):
return {"result": "testing duplicate tags"}
# Get the updated schema and parse routes
routes = parse_openapi_to_http_routes(fastapi_server.openapi())
routes = parse_openapi_to_http_routes(fastapi_app.openapi())
# Find our test route
test_route = next((r for r in routes if r.path == "/test-duplicate-tags"), None)