diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index f82b238d0..ff2c5ded3 100644 --- a/src/fastmcp/client/transports.py +++ b/src/fastmcp/client/transports.py @@ -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 diff --git a/src/fastmcp/server/openapi.py b/src/fastmcp/server/openapi.py index dfce4e21e..33e6440e6 100644 --- a/src/fastmcp/server/openapi.py +++ b/src/fastmcp/server/openapi.py @@ -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) diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 8cc182e22..43f06742e 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -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 diff --git a/tests/client/test_openapi.py b/tests/client/test_openapi.py new file mode 100644 index 000000000..903582f59 --- /dev/null +++ b/tests/client/test_openapi.py @@ -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" diff --git a/tests/deprecated/test_route_type_ignore.py b/tests/deprecated/test_route_type_ignore.py index 1382d7137..575c0780e 100644 --- a/tests/deprecated/test_route_type_ignore.py +++ b/tests/deprecated/test_route_type_ignore.py @@ -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 diff --git a/tests/server/http/test_http_dependencies.py b/tests/server/http/test_http_dependencies.py index 680717adb..192090792 100644 --- a/tests/server/http/test_http_dependencies.py +++ b/tests/server/http/test_http_dependencies.py @@ -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) diff --git a/tests/server/openapi/test_openapi.py b/tests/server/openapi/test_openapi.py index f4d16ae08..89ff261ba 100644 --- a/tests/server/openapi/test_openapi.py +++ b/tests/server/openapi/test_openapi.py @@ -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 diff --git a/tests/utilities/openapi/test_openapi_fastapi.py b/tests/utilities/openapi/test_openapi_fastapi.py index 3f0ee8b41..9b7951ec0 100644 --- a/tests/utilities/openapi/test_openapi_fastapi.py +++ b/tests/utilities/openapi/test_openapi_fastapi.py @@ -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)