From 213abc42448e021679cae8dde460c01f025ab0da Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 23 May 2025 12:38:24 -0400 Subject: [PATCH] Pass client headers through to OpenAPI client --- src/fastmcp/server/openapi.py | 34 +++- tests/client/test_openapi.py | 157 ++++++++++++++++++ tests/deprecated/test_route_type_ignore.py | 4 +- tests/server/http/test_http_dependencies.py | 88 ++++++++-- tests/server/openapi/test_openapi.py | 48 +++--- .../utilities/openapi/test_openapi_fastapi.py | 18 +- 6 files changed, 296 insertions(+), 53 deletions(-) create mode 100644 tests/client/test_openapi.py 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/tests/client/test_openapi.py b/tests/client/test_openapi.py new file mode 100644 index 000000000..b2639e614 --- /dev/null +++ b/tests/client/test_openapi.py @@ -0,0 +1,157 @@ +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) + + @pytest.fixture(autouse=True, 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(autouse=True, 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" + + 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" 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..b41c215b2 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) 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)